From b898fb5dbcdaed82c430903e0298f446583ce6a7 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 18 Jul 2024 17:34:35 +0800 Subject: [PATCH] fix update insight bug & modify contra repeat prompt --- config/demo_config_cn.yaml | 1 + .../summary/get_reflection_subject_worker.py | 6 +- .../summary/long_contra_repeat_worker.py | 22 +++--- .../summary/long_contra_repeat_worker.yaml | 7 +- .../worker/summary/update_insight_worker.py | 26 +++---- .../worker/summary/update_insight_worker.yaml | 6 ++ .../memory/worker/write/load_memory_worker.py | 3 +- memory_scope/utils/memory_handler.py | 17 +++++ tests/worker/test_workers_cn.py | 68 +++++++++++++++++-- 9 files changed, 125 insertions(+), 31 deletions(-) diff --git a/config/demo_config_cn.yaml b/config/demo_config_cn.yaml index aa1093ce..d22194d6 100644 --- a/config/demo_config_cn.yaml +++ b/config/demo_config_cn.yaml @@ -148,6 +148,7 @@ worker: generation_model: dashscope_generation generation_model_kwargs: top_k: 1 + rank_model: dashscope_rank long_contra_repeat: class: memory.worker.summary.long_contra_repeat_worker generation_model: dashscope_generation diff --git a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py index 8bd95673..46172b87 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -21,7 +21,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): def _parse_params(self, **kwargs): self.reflect_obs_cnt_threshold: int = kwargs.get("reflect_obs_cnt_threshold", 10) - self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) self.reflect_num_questions: int = kwargs.get("reflect_num_questions", 5) def new_insight_node(self, insight_key: str) -> MemoryNode: @@ -88,7 +88,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): self.logger.info(f"reflect_message={reflect_message}") # Invoke Language Model for new insights - response = self.generation_model.call(messages=reflect_message, top_k=self.generation_model_top_k) + response = self.generation_model.call(messages=reflect_message, **self.generation_model_kwargs) # Handle empty response if not response.status or not response.message.content: @@ -98,7 +98,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): new_insight_keys = ResponseTextParser(response.message.content).parse_v2(self.__class__.__name__) if new_insight_keys: for insight_key in new_insight_keys: - insight_nodes.append(self.new_insight_node(insight_key)) + self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) # Mark unaudited nodes as reflected for node in not_reflected_nodes: diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py index 919834bd..1bc0d6e2 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -21,6 +21,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): FILE_PATH: str = __file__ def _parse_params(self, **kwargs): + self.unit_test_flag = False self.long_contra_repeat_top_k: int = kwargs.get("long_contra_repeat_top_k", 2) self.long_contra_repeat_threshold: float = kwargs.get("long_contra_repeat_threshold", 0.1) self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) @@ -66,16 +67,19 @@ class LongContraRepeatWorker(MemoryBaseWorker): for node in not_updated_nodes: self.submit_thread_task(fn=self.retrieve_similar_content, node=node) - obs_node_dict: Dict[str, MemoryNode] = {} - for origin_node, retrieve_nodes in self.gather_thread_result(): - if not retrieve_nodes: - continue - obs_node_dict[origin_node.memory_id] = origin_node - for node in retrieve_nodes: - if node.memory_id in obs_node_dict: + if self.unit_test_flag: + all_obs_nodes: List[MemoryNode] = not_updated_nodes + else: + obs_node_dict: Dict[str, MemoryNode] = {} + for origin_node, retrieve_nodes in self.gather_thread_result(): + if not retrieve_nodes: continue - obs_node_dict[node.memory_id] = node - all_obs_nodes: List[MemoryNode] = sorted(obs_node_dict.values(), key=lambda x: x.timestamp, reverse=True) + obs_node_dict[origin_node.memory_id] = origin_node + for node in retrieve_nodes: + if node.memory_id in obs_node_dict: + continue + obs_node_dict[node.memory_id] = node + all_obs_nodes: List[MemoryNode] = sorted(obs_node_dict.values(), key=lambda x: x.timestamp, reverse=True) if not all_obs_nodes: self.logger.warning("all_obs_nodes is empty, stop.") diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_worker.yaml b/memory_scope/memory/worker/summary/long_contra_repeat_worker.yaml index 0fc7e640..71c59e35 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.yaml +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.yaml @@ -1,6 +1,8 @@ long_contra_repeat_system: cn: | - 对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。 + 对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。 + 注意:只判断与“前面序号”的句子的关系,不要判断“后面序号”。 + 其中矛盾的形式可以有很多种,可以是逻辑上的矛盾,可以是属性上的变化导致的矛盾,比如不能同时在两个地方工作,同一个时刻不能在两个地点,同一个时刻不能干两件事情等等。 对每个句子都做一个判断,最后一共输出{num_obs}条判断。如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。 请一步步思考,并按如下格式输出: 思考:思考的依据和过程,30字以内。 @@ -102,3 +104,6 @@ long_contra_repeat_user_query: cn: | 句子: {user_query} + en: | + Sentences: + {user_query} diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index 61505da1..f9df204e 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -44,8 +44,8 @@ class UpdateInsightWorker(MemoryBaseWorker): filtered_nodes: List[MemoryNode] = [] # Check if insight node key or value is empty and log a warning - if not insight_node.key or not insight_node.value: - self.logger.warning(f"insight_key={insight_node.key} insight_value={insight_node.value} is empty!") + if not insight_node.key: + self.logger.warning(f"insight_key={insight_node.key} is empty!") return insight_node, filtered_nodes, max_score # Call the ranking model to get scores for each observed node's content against the insight key @@ -62,7 +62,7 @@ class UpdateInsightWorker(MemoryBaseWorker): filtered_nodes.append(node) max_score = max(max_score, score) # Log information about each node's processing - self.logger.info(f"insight_key={insight_node.key} insight_value={insight_node.value} " + self.logger.info(f"insight_key={insight_node.key} content={node.content} " f"score={score} keep_flag={keep_flag}") # Warn if no nodes were filtered @@ -74,8 +74,8 @@ class UpdateInsightWorker(MemoryBaseWorker): def update_insight_node(self, insight_node: MemoryNode, insight_value: str): dt_handler = DatetimeHandler() - content = (f"{self.user_name}{self.get_language_value(COLON_WORD)}{insight_node.key}" - f"{self.get_language_value(COLON_WORD)}{insight_value}") + key = self.prompt_handler.insight_string_format.format(name=self.target_name, key=insight_node.key) + content = f"{key}{self.get_language_value(COLON_WORD)}{insight_value}" insight_node.content = content insight_node.value = insight_value insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()}) @@ -97,18 +97,17 @@ class UpdateInsightWorker(MemoryBaseWorker): Returns: MemoryNode: The updated MemoryNode with potentially revised insight value. """ - self.logger.info(f"Updating insight for key={insight_node.key}, value={insight_node.value}, " + self.logger.info(f"Updating insight for key={insight_node.key}, old_value={insight_node.value}, " f"with {len(filtered_nodes)} documents considered.") # Generate the prompt for updating insight user_query_list = [n.content for n in filtered_nodes] - system_prompt = self.prompt_handler.update_insight_system.foramt(user_name=self.target_name) - few_shot = self.prompt_handler.update_insight_few_shot.foramt(user_name=self.target_name) - user_query = self.prompt_handler.update_insight_user_query.foramt( + system_prompt = self.prompt_handler.update_insight_system.format(user_name=self.target_name) + few_shot = self.prompt_handler.update_insight_few_shot.format(user_name=self.target_name) + user_query = self.prompt_handler.update_insight_user_query.format( user_query="\n".join(user_query_list), insight_key=insight_node.key, insight_key_value=insight_node.key + self.get_language_value(COLON_WORD) + insight_node.value) - # Construct the message for LLM interaction update_insight_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"Generated insight update message: {update_insight_message}") @@ -170,11 +169,11 @@ class UpdateInsightWorker(MemoryBaseWorker): if node.action_status == ActionStatusEnum.NONE.value: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, - not_updated_nodes=not_updated_nodes) + obs_nodes=not_updated_nodes) else: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, - not_updated_nodes=not_reflected_nodes) + obs_nodes=not_reflected_nodes) # select top n result_list = [] @@ -190,7 +189,8 @@ class UpdateInsightWorker(MemoryBaseWorker): self.submit_thread_task(fn=self.update_insight, insight_node=insight_node, filtered_nodes=filtered_nodes) # Gather the final results from all update tasks - self.gather_thread_result() + for _ in self.gather_thread_result(): + pass for node in not_updated_nodes: node.obs_updated = 1 diff --git a/memory_scope/memory/worker/summary/update_insight_worker.yaml b/memory_scope/memory/worker/summary/update_insight_worker.yaml index c086d5b3..160e28fc 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.yaml +++ b/memory_scope/memory/worker/summary/update_insight_worker.yaml @@ -105,3 +105,9 @@ update_insight_user_query: {user_query} Category: {insight_key} Existing information: {insight_key_value} + +insight_string_format: + cn: | + {name}的{key} + en: | + The {key} of {name} diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index 5acd4191..26ed12d9 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -103,4 +103,5 @@ class LoadMemoryWorker(MemoryBaseWorker): self.submit_thread_task(self.retrieve_today_memory, query=query, dt=dt) # Waits for all submitted tasks to complete - self.gather_thread_result() + for _ in self.gather_thread_result(): + pass diff --git a/memory_scope/utils/memory_handler.py b/memory_scope/utils/memory_handler.py index 7bbe25e8..06c9ec97 100644 --- a/memory_scope/utils/memory_handler.py +++ b/memory_scope/utils/memory_handler.py @@ -37,6 +37,23 @@ class MemoryHandler(object): self._id_memory_dict.clear() self._key_id_dict.clear() + def add_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True): + if key not in self._key_id_dict: + return self.set_memories(key, nodes, log_repeat) + + if isinstance(nodes, MemoryNode): + nodes = [nodes] + + for node in nodes: + _id = node.memory_id + if _id not in self._key_id_dict[key]: + self._key_id_dict[key].append(_id) + + if _id not in self._id_memory_dict: + self._id_memory_dict[_id] = node + self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content or node.key} " + f"store_status={node.store_status} action_status={node.action_status}") + def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True): if nodes is None: nodes = [] diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index 63ece3ef..e035fae8 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -3,11 +3,12 @@ import unittest from memory_scope.cli import MemoryScope from memory_scope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.memory_node import MemoryNode from memory_scope.scheme.message import Message +from memory_scope.utils import Logger from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.tool_functions import init_instance_by_config @@ -16,9 +17,16 @@ class TestWorkersCn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - ms = MemoryScope().load_config("config/demo_config_cn.yaml") + datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') + self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) + + ms = MemoryScope() + ms.load_config("config/demo_config_cn.yaml") ms.init_global_content_by_config() + def tearDown(self): + self.logger.close() + @unittest.skip def test_extract_time(self): name = "extract_time" @@ -257,7 +265,7 @@ class TestWorkersCn(unittest.TestCase): worker.logger.info(f"result3={result3}") worker.logger.info(f"result4={result4}") - # @unittest.skip + @unittest.skip def test_get_reflection_subject(self): name = "get_reflection_subject" @@ -287,6 +295,58 @@ class TestWorkersCn(unittest.TestCase): worker.memory_handler.set_memories(INSIGHT_NODES, []) worker.run() + result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = "\n".join(result) + worker.logger.info(f"result.get_reflection={result}") + return worker + + @unittest.skip + def test_update_insight_worker(self): + reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + + name = "update_insight" + + worker: MemoryBaseWorker = init_instance_by_config( + config=G_CONTEXT.worker_config[name], + suffix_name="worker", + name=name, + is_multi_thread=False, + context=reflection_worker.context, + context_lock=None, + thread_pool=G_CONTEXT.thread_pool) + + nodes = [ + MemoryNode(content="用户喜欢打王者荣耀"), + ] + worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.run() + result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] result = "\n".join(result) - worker.logger.info(f"result={result}") \ No newline at end of file + worker.logger.info(f"result.update_insight={result}") + + # @unittest.skip + def test_long_contra_repeat_worker(self): + name = "long_contra_repeat" + + worker: MemoryBaseWorker = init_instance_by_config( + config=G_CONTEXT.worker_config[name], + suffix_name="worker", + name=name, + is_multi_thread=False, + context={}, + context_lock=None, + thread_pool=G_CONTEXT.thread_pool) + + nodes = [ + MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), + MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"), + MemoryNode(content="用户在上海工作。"), + ] + worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.unit_test_flag = True + worker.run() + + result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = "\n".join(result) + worker.logger.info(f"result.long_contra_repeat={result}")