fix score_similar -> score_recall

This commit is contained in:
jinli.yl 2024-07-24 20:05:27 +08:00
parent c437553ff6
commit 83de295eff
9 changed files with 26 additions and 24 deletions

View file

@ -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)

View file

@ -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))

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)}")

View file

@ -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)

View file

@ -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