mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
[dev] add special time parse logic format_time_infer
This commit is contained in:
parent
7751586435
commit
21254b05a7
7 changed files with 187 additions and 6 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
116
memory_scope/memory/worker/read/fuse_rerank_worker.py
Normal file
116
memory_scope/memory/worker/read/fuse_rerank_worker.py
Normal 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))
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
49
memory_scope/memory/worker/read/semantic_rank_worker.py
Normal file
49
memory_scope/memory/worker/read/semantic_rank_worker.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue