diff --git a/config/demo_config.yaml b/config/demo_config.yaml index c2aa5890..9b7fc090 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -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 diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index 7863f55c..915c7536 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -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" diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index 1cbc0381..3ade4b05 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -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(): diff --git a/memory_scope/memory/worker/read/fuse_rerank_worker.py b/memory_scope/memory/worker/read/fuse_rerank_worker.py new file mode 100644 index 00000000..b993e593 --- /dev/null +++ b/memory_scope/memory/worker/read/fuse_rerank_worker.py @@ -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)) diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py index 4da85f70..6462a1e2 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -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) diff --git a/memory_scope/memory/worker/read/semantic_rank_worker.py b/memory_scope/memory/worker/read/semantic_rank_worker.py new file mode 100644 index 00000000..161b6e2f --- /dev/null +++ b/memory_scope/memory/worker/read/semantic_rank_worker.py @@ -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) diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 0e7b79ab..7f31d15f 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -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