[dev] add special time parse logic format_time_infer

This commit is contained in:
jinli.yl 2024-07-01 14:08:10 +08:00
parent 7751586435
commit 21254b05a7
7 changed files with 187 additions and 6 deletions

View file

@ -69,4 +69,11 @@ worker:
class: memory.worker.read.retrieve_store_worker
retrieve_obs_top_k: 100
retrieve_ins_pf_top_k: 100
fuse_rerank_worker:
class: memory.worker.read.fuse_rerank_worker
fuse_score_threshold: 0.1
fuse_ratio_dict:
observation: 1
fuse_time_ratio: 2.0
fuse_rerank_top_k: 10

View file

@ -8,7 +8,9 @@ CHAT_KWARGS = "chat_kwargs"
QUERY_WITH_TS = "query_with_ts"
RETRIEVE_MEMORY_NODES = "RETRIEVE_MEMORY_NODES"
RETRIEVE_MEMORY_NODES = "retrieve_memory_nodes"
RANKED_MEMORY_NODES = "ranked_memory_nodes"

View file

@ -1,4 +1,5 @@
import re
from typing import Dict
from memory_scope.constants.common_constants import DATATIME_KEY_MAP, QUERY_WITH_TS, EXTRACT_TIME_DICT
from memory_scope.constants.language_constants import DATATIME_WORD_LIST
@ -43,7 +44,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
response_text = response.message.content
# re-match time info to dict
extract_time_dict = {}
extract_time_dict: Dict[str, str] = {}
matches = re.findall(self.EXTRACT_TIME_PATTERN, response_text)
for key, value in matches:
if key in DATATIME_KEY_MAP.keys():

View file

@ -0,0 +1,116 @@
from typing import Dict, List
from memory_scope.constants.common_constants import EXTRACT_TIME_DICT, RANKED_MEMORY_NODES, RESULT
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
class FuseRerankWorker(MemoryBaseWorker):
@staticmethod
def format_time_infer(time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]):
if time_infer:
return time_infer
time_infer = ""
if "year" in extract_time_dict:
value = meta_data.get("msg_year")
if value:
time_infer += f"{value}年"
elif value == "-1":
time_infer += f"每年"
if "month" in extract_time_dict:
value = meta_data.get("msg_month")
if value:
time_infer += f"{value}月"
elif value == "-1":
time_infer += f"每月"
if "day" in extract_time_dict:
value = meta_data.get("msg_day")
if value:
time_infer += f"{value}日"
elif value == "-1":
time_infer += f"每日"
if "weekday" in extract_time_dict:
value = meta_data.get("msg_weekday")
if value:
time_infer += value
return time_infer
@staticmethod
def match_node_time(extract_time_dict: Dict[str, str], node: MemoryNode):
if extract_time_dict:
match_event_flag = True
for k, v in extract_time_dict.items():
event_value = node.meta_data.get(f"event_{k}", "")
if event_value in ["-1", v]:
continue
else:
match_event_flag = False
break
match_msg_flag = True
for k, v in extract_time_dict.items():
msg_value = node.meta_data.get(f"msg_{k}", "")
if msg_value == v:
continue
else:
match_msg_flag = False
break
else:
match_event_flag = False
match_msg_flag = False
node.meta_data["match_event_flag"] = str(int(match_event_flag))
node.meta_data["match_msg_flag"] = str(int(match_msg_flag))
return match_event_flag, match_msg_flag
def _run(self):
# parse input
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
memory_node_list: List[MemoryNode] = self.get_context(RANKED_MEMORY_NODES)
if not memory_node_list:
self.logger.warning(f"ranked memory nodes is empty!")
return
# get reranked nodes
reranked_memory_nodes = []
for node in memory_node_list:
if node.score_rank < self.fuse_score_threshold:
continue
# memory type ratio
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
# memory fuse time ratio
match_event_flag, match_msg_flag = self.match_node_time(extract_time_dict=extract_time_dict, node=node)
fuse_time_ratio: float = self.fuse_time_ratio if match_event_flag or match_msg_flag else 1.0
# fuse rerank score
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
reranked_memory_nodes.append(node)
# build result
memories: List[str] = []
reranked_memory_nodes = sorted(reranked_memory_nodes,
key=lambda x: x.score_rerank,
reverse=True)[: self.fuse_rerank_top_k]
for node in reranked_memory_nodes:
f_event = int(node.meta_data["match_event_flag"])
f_msg = int(node.meta_data["match_msg_flag"])
self.logger.info(f"rerank_stage: content={node.content} score={node.score_rerank} "
f"f_event={f_event} f_msg={f_msg}")
content = node.content
if f_event or f_msg:
time_infer = self.format_time_infer(time_infer="",
extract_time_dict=extract_time_dict,
meta_data=node.memory_node.metaData)
content = f"{time_infer}: {content}"
memories.append(content)
self.set_context(RESULT, "\n".join(memories))

View file

@ -36,9 +36,9 @@ class RetrieveStoreWorker(MemoryBaseWorker):
for result in self._async_run(fn_list=fn_list, query=query):
if result:
memory_node_list.extend(result)
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
for i, node in enumerate(memory_node_list):
self.logger.info(f"{i}: node={node.content} score={node.score_similar} type={node.memory_type}")
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
for node in memory_node_list:
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}")
self.set_context(RETRIEVE_MEMORY_NODES, memory_node_list)

View file

@ -0,0 +1,49 @@
from typing import List, Dict
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
class SemanticRankWorker(MemoryBaseWorker):
def _run(self):
# query
query, _ = self.get_context(QUERY_WITH_TS)
memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES)
if not memory_node_list:
self.logger.warning(f"retrieve memory nodes is empty!")
return
# solve content repeat, insight & profile has higher priority
memory_node_dict: Dict[str, MemoryNode] = {}
memory_type_selected = [MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.PROFILE.value,
MemoryTypeEnum.PROFILE_CUSTOMIZED.value]
for node in memory_node_list:
if node.memory_type in memory_type_selected:
return
memory_node_dict[node.content] = node
for node in memory_node_list:
if node.memory_type not in memory_type_selected:
return
memory_node_dict[node.content] = node
memory_node_list = list(memory_node_dict.values())
response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list])
if not response.status or not response.rank_scores:
return
# set score
rank_memory_nodes: List[MemoryNode] = []
rank_scores = sorted(response.rank_scores.items(), key=lambda x: x[1], reverse=True)
for idx, score in rank_scores:
if idx >= len(memory_node_list):
self.logger.warning(f"idx={idx} exceeds the maximum length of the array")
continue
node = memory_node_list[idx]
node.score_rank = score
self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}")
self.get_context(RANKED_MEMORY_NODES, rank_memory_nodes)

View file

@ -127,3 +127,9 @@ def md5_hash(input_string: str):
m = hashlib.md5()
m.update(input_string.encode('utf-8'))
return m.hexdigest()
def contains_keyword(text, keywords):
escaped_keywords = map(re.escape, keywords)
pattern = re.compile('|'.join(escaped_keywords), re.IGNORECASE)
return pattern.search(text) is not None