mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-10 22:41:06 +00:00
fix score_similar -> score_recall
This commit is contained in:
parent
c437553ff6
commit
83de295eff
9 changed files with 26 additions and 24 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue