From 83de295eff5403a89a21815c4f6630b8496942c9 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 24 Jul 2024 20:05:27 +0800 Subject: [PATCH] fix score_similar -> score_recall --- .../worker/frontend/retrieve_memory_worker.py | 4 ++-- .../summary/get_reflection_subject_worker.py | 2 +- .../summary/long_contra_repeat_worker.py | 4 ++-- .../worker/summary/update_insight_worker.py | 3 ++- .../worker/write/contra_repeat_worker.py | 2 +- .../worker/write/get_observation_worker.py | 2 +- .../memory/worker/write/info_filter_worker.py | 2 +- .../storage/llama_index_es_memory_store.py | 7 ++---- memoryscope/utils/response_text_parser.py | 24 +++++++++++-------- 9 files changed, 26 insertions(+), 24 deletions(-) diff --git a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py b/memoryscope/memory/worker/frontend/retrieve_memory_worker.py index 178e870c..459db4ee 100644 --- a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memoryscope/memory/worker/frontend/retrieve_memory_worker.py @@ -133,10 +133,10 @@ class RetrieveMemoryWorker(MemoryBaseWorker): if not memory_node_list: return - memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True) + memory_node_list = sorted(memory_node_list, key=lambda x: x.score_recall, reverse=True) for node in memory_node_list: node.action_status = ActionStatusEnum.NONE.value - self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type} " + self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} type={node.memory_type} " f"store_status={node.store_status} action_status={node.action_status}") self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memoryscope/memory/worker/summary/get_reflection_subject_worker.py b/memoryscope/memory/worker/summary/get_reflection_subject_worker.py index 1dbb9d6f..4c9d7b27 100644 --- a/memoryscope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memoryscope/memory/worker/summary/get_reflection_subject_worker.py @@ -101,7 +101,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): return # Parse LLM response for new insight keys and update memory - new_insight_keys = ResponseTextParser(response.message.content).parse_v2(self.__class__.__name__) + new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2() if new_insight_keys: for insight_key in new_insight_keys: self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) diff --git a/memoryscope/memory/worker/summary/long_contra_repeat_worker.py b/memoryscope/memory/worker/summary/long_contra_repeat_worker.py index 34ce6853..94b306b4 100644 --- a/memoryscope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memoryscope/memory/worker/summary/long_contra_repeat_worker.py @@ -49,7 +49,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): top_k=self.long_contra_repeat_top_k, filter_dict=filter_dict) # Filter retrieved nodes based on the similarity threshold - return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold] + return node, [n for n in retrieve_nodes if n.score_recall >= self.long_contra_repeat_threshold] def _run(self): """ @@ -111,7 +111,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): return # Parses the model's response text to identify updates for memory nodes - idx_obs_info_list = ResponseTextParser(response.message.content).parse_v1(self.__class__.__name__) + idx_obs_info_list = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v1() if len(idx_obs_info_list) <= 0: self.logger.warning("idx_obs_info_list is empty!") return diff --git a/memoryscope/memory/worker/summary/update_insight_worker.py b/memoryscope/memory/worker/summary/update_insight_worker.py index 498c62e1..b56ed6f7 100644 --- a/memoryscope/memory/worker/summary/update_insight_worker.py +++ b/memoryscope/memory/worker/summary/update_insight_worker.py @@ -136,7 +136,8 @@ class UpdateInsightWorker(MemoryBaseWorker): if not response.status or not response.message.content: return insight_node - insight_value_list = ResponseTextParser(response.message.content).parse_v1(f"update_{insight_node.key}") + insight_value_list = ResponseTextParser(response.message.content, + f"update_{insight_node.key}").parse_v1() if not insight_value_list: self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!") return insight_node diff --git a/memoryscope/memory/worker/write/contra_repeat_worker.py b/memoryscope/memory/worker/write/contra_repeat_worker.py index 05008c9c..5146887f 100644 --- a/memoryscope/memory/worker/write/contra_repeat_worker.py +++ b/memoryscope/memory/worker/write/contra_repeat_worker.py @@ -80,7 +80,7 @@ class ContraRepeatWorker(MemoryBaseWorker): response_text = response.message.content # parse text - idx_merge_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) + idx_merge_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() if len(idx_merge_obs_list) <= 0: self.logger.warning("idx_merge_obs_list is empty!") return diff --git a/memoryscope/memory/worker/write/get_observation_worker.py b/memoryscope/memory/worker/write/get_observation_worker.py index 9fe97e34..b36d2b6f 100644 --- a/memoryscope/memory/worker/write/get_observation_worker.py +++ b/memoryscope/memory/worker/write/get_observation_worker.py @@ -139,7 +139,7 @@ class GetObservationWorker(MemoryBaseWorker): response_text = response.message.content # Parses the generated text to extract observation indices, times, contents, and keywords - idx_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) + idx_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() if len(idx_obs_list) <= 0: self.logger.warning("idx_obs_list is empty!") return diff --git a/memoryscope/memory/worker/write/info_filter_worker.py b/memoryscope/memory/worker/write/info_filter_worker.py index 29a84a86..ae28c13d 100644 --- a/memoryscope/memory/worker/write/info_filter_worker.py +++ b/memoryscope/memory/worker/write/info_filter_worker.py @@ -76,7 +76,7 @@ class InfoFilterWorker(MemoryBaseWorker): response_text = response.message.content # parse text - info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) + info_score_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() if len(info_score_list) != len(info_messages): self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}") diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index 7d9dbd1e..02e0c085 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -149,8 +149,5 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): Returns: MemoryNode: The converted MemoryNode with text and metadata from the NodeWithScore. """ - embedding = text_node.embedding - print("textnode embedding", embedding) - if not embedding: - embedding = [] - return MemoryNode(content=text_node.text, vector=embedding, **text_node.metadata) + text_node.metadata["vector"] = text_node.embedding + return MemoryNode(content=text_node.text, **text_node.metadata) diff --git a/memoryscope/utils/response_text_parser.py b/memoryscope/utils/response_text_parser.py index 4081cbbf..b7de8982 100644 --- a/memoryscope/utils/response_text_parser.py +++ b/memoryscope/utils/response_text_parser.py @@ -14,23 +14,27 @@ class ResponseTextParser(object): pattern_v1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets - def __init__(self, response_text: str): + def __init__(self, response_text: str, logger_prefix: str = ""): """ Initializes the `ResponseTextParser` instance with the provided response text and sets up a logger. Args: response_text (str): The raw response text that needs to be parsed and processed. """ - self.response_text: str = response_text.strip() # Strips leading and trailing whitespace from the response text - self.logger: Logger = Logger.get_logger() # Initializes a logger instance for logging parsing activities - def parse_v1(self, prefix: str = "") -> List[str]: + # Strips leading and trailing whitespace from the response text + self.response_text: str = response_text.strip() + + # The prefix of log. Defaults to "". + self.logger_prefix: str = logger_prefix + + # Initializes a logger instance for logging parsing activities + self.logger: Logger = Logger.get_logger() + + def parse_v1(self) -> List[List[str]]: """ Extract specific patterns from the text which match content within angle brackets. - Args: - prefix (str): The prefix of log. Defaults to "". - Returns: Contents match the specific patterns. """ @@ -42,10 +46,10 @@ class ResponseTextParser(object): matches = [match.group(1) for match in self.pattern_v1.finditer(line)] if matches: result.append(matches) - self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2) + self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2) return result - def parse_v2(self, prefix: str = "") -> List[str]: + def parse_v2(self) -> List[str]: """ Extract lines which contain NONE_WORD in Chinese or English. @@ -61,5 +65,5 @@ class ResponseTextParser(object): if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language): continue result.append(line) - self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2) + self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2) return result