From d8b3f6c7e649e7b41474628cafb5eb40c6fa786c Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 18 Jul 2024 16:30:00 +0800 Subject: [PATCH] add reflection test cases & fix name bugs --- config/demo_config_cn.yaml | 11 ++ .../worker/frontend/extract_time_worker.py | 4 +- ...aml => get_reflection_subject_worker.yaml} | 0 ...pt.yaml => long_contra_repeat_worker.yaml} | 0 ...prompt.yaml => update_insight_worker.yaml} | 0 .../memory/worker/write/info_filter_worker.py | 1 - .../worker/write/info_filter_worker.yaml | 2 +- tests/worker/test_workers_cn.py | 136 +++++++++++++++++- 8 files changed, 144 insertions(+), 10 deletions(-) rename memory_scope/memory/worker/summary/{get_reflection_subject_prompt.yaml => get_reflection_subject_worker.yaml} (100%) rename memory_scope/memory/worker/summary/{long_contra_repeat_prompt.yaml => long_contra_repeat_worker.yaml} (100%) rename memory_scope/memory/worker/summary/{update_insight_prompt.yaml => update_insight_worker.yaml} (100%) diff --git a/config/demo_config_cn.yaml b/config/demo_config_cn.yaml index 30d18363..aa1093ce 100644 --- a/config/demo_config_cn.yaml +++ b/config/demo_config_cn.yaml @@ -72,6 +72,8 @@ worker: extract_time: class: memory.worker.frontend.extract_time_worker generation_model: dashscope_generation + generation_model_kwargs: + top_k: 1 semantic_rank: class: memory.worker.frontend.semantic_rank_worker rank_model: dashscope_rank @@ -138,10 +140,19 @@ worker: retrieve_insight_top_k: 100 get_reflection_subject: class: memory.worker.summary.get_reflection_subject_worker + generation_model: dashscope_generation + generation_model_kwargs: + top_k: 1 update_insight: class: memory.worker.summary.update_insight_worker + generation_model: dashscope_generation + generation_model_kwargs: + top_k: 1 long_contra_repeat: class: memory.worker.summary.long_contra_repeat_worker + generation_model: dashscope_generation + generation_model_kwargs: + top_k: 1 models: dashscope_generation: diff --git a/memory_scope/memory/worker/frontend/extract_time_worker.py b/memory_scope/memory/worker/frontend/extract_time_worker.py index 2fb0e7fa..ef93e726 100644 --- a/memory_scope/memory/worker/frontend/extract_time_worker.py +++ b/memory_scope/memory/worker/frontend/extract_time_worker.py @@ -19,7 +19,7 @@ class ExtractTimeWorker(MemoryBaseWorker): FILE_PATH: str = __file__ def _parse_params(self, **kwargs): - self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) def _run(self): """ @@ -47,7 +47,7 @@ class ExtractTimeWorker(MemoryBaseWorker): self.logger.info(f"extract_time_message={extract_time_message}") # Invoke the LLM to generate a response - response = self.generation_model.call(messages=extract_time_message, top_k=self.generation_model_top_k) + response = self.generation_model.call(messages=extract_time_message, **self.generation_model_kwargs) # Handle empty or unsuccessful responses if not response.status or not response.message.content: diff --git a/memory_scope/memory/worker/summary/get_reflection_subject_prompt.yaml b/memory_scope/memory/worker/summary/get_reflection_subject_worker.yaml similarity index 100% rename from memory_scope/memory/worker/summary/get_reflection_subject_prompt.yaml rename to memory_scope/memory/worker/summary/get_reflection_subject_worker.yaml diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_prompt.yaml b/memory_scope/memory/worker/summary/long_contra_repeat_worker.yaml similarity index 100% rename from memory_scope/memory/worker/summary/long_contra_repeat_prompt.yaml rename to memory_scope/memory/worker/summary/long_contra_repeat_worker.yaml diff --git a/memory_scope/memory/worker/summary/update_insight_prompt.yaml b/memory_scope/memory/worker/summary/update_insight_worker.yaml similarity index 100% rename from memory_scope/memory/worker/summary/update_insight_prompt.yaml rename to memory_scope/memory/worker/summary/update_insight_worker.yaml diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 83df42d6..7ef811dd 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -81,7 +81,6 @@ class InfoFilterWorker(MemoryBaseWorker): info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) if len(info_score_list) != len(info_messages): self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}") - return # filter messages filtered_messages: List[Message] = [] diff --git a/memory_scope/memory/worker/write/info_filter_worker.yaml b/memory_scope/memory/worker/write/info_filter_worker.yaml index 8d1b68b3..032ec114 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.yaml +++ b/memory_scope/memory/worker/write/info_filter_worker.yaml @@ -3,7 +3,7 @@ info_filter_system: 任务:对所给{batch_size}个句子中所含有的关于{user_name}的信息打分,分数为0,1,2或3。 注意:其中0表示不包含用户信息,1表示句子中只包含用户假设的信息或者用户虚构的内容比如用户创作的小说或剧本,2表示包含用户的一般信息,时效性信息或者需要猜测才能得到的用户信息,3表示明确含有或者可以确定推断出关于用户的重要信息,或者用户要求记录。 {user_name}的重要信息可以包含用户基本信息,用户画像信息,用户兴趣偏好信息,用户性格,用户价值观,用户人际关系,用户重大事件转折点等等重要信息。 - 对每个句子都做一次信息打分,一共输出{batch_size}个分数。 + 对每个句子都做一次信息打分,一共输出{batch_size}个分数,不需要写最终结果。 请一定要按如下格式依次输出,最后的结果一定要加<>: 思考:思考的依据和过程,30字以内。 结果:<句子序号> <分数:0或1或2或3> diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index e701cb26..63ece3ef 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -1,8 +1,9 @@ +import datetime 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 + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_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 @@ -19,7 +20,28 @@ class TestWorkersCn(unittest.TestCase): ms.init_global_content_by_config() @unittest.skip - def test_info_filter_cn(self): + def test_extract_time(self): + name = "extract_time" + + 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) + + query = "明天我去上海出差" + query_timestamp = int(datetime.datetime.now().timestamp()) + worker.set_context(QUERY_WITH_TS, (query, query_timestamp)) + worker.run() + + result = worker.get_context(EXTRACT_TIME_DICT) + worker.logger.info(f"result={result}") + + @unittest.skip + def test_info_filter(self): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( @@ -45,7 +67,43 @@ class TestWorkersCn(unittest.TestCase): worker.logger.info(f"result={result}") @unittest.skip - def test_get_observation_cn(self): + def test_info_filter2(self): + name = "info_filter" + + 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) + + chat_messages = [ + Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"), + Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"), + Message(role=MessageRoleEnum.USER.value, content="听说篮球运动对身体很好,是真的吗?"), + Message(role=MessageRoleEnum.USER.value, content="最近在北京的工作压力太大,有什么放松的建议吗?"), + Message(role=MessageRoleEnum.USER.value, content="说到朋友,我确实有几位很要好的朋友,我们经常一起出去吃饭。"), + Message(role=MessageRoleEnum.USER.value, content="对了,最近想换工作,你觉得北京的哪个区工作机会更多?"), + Message(role=MessageRoleEnum.USER.value, content="听你这么说,我感觉挺有信心的,谢了!"), + Message(role=MessageRoleEnum.USER.value, content="我很喜欢尝试新的美食,有没有推荐的美食应用?"), + Message(role=MessageRoleEnum.USER.value, content="我有时也喜欢自己在家做饭,你有没有好的海鲜菜谱推荐?"), + Message(role=MessageRoleEnum.USER.value, content="听说打篮球可以长高,这是真的吗?"), + Message(role=MessageRoleEnum.USER.value, content="我在北京阿里云园区工作"), + Message(role=MessageRoleEnum.USER.value, content="我是阿里云百炼的工程师"), + Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?"), + ] + + worker.set_context(CHAT_MESSAGES, chat_messages) + worker.run() + + result = [msg.content for msg in worker.chat_messages] + result = "\n".join(result) + worker.logger.info(f"result={result}") + + @unittest.skip + def test_get_observation(self): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( @@ -79,7 +137,39 @@ class TestWorkersCn(unittest.TestCase): worker.logger.info(f"result={result}") @unittest.skip - def test_get_observation_with_time_cn(self): + def test_get_observation2(self): + name = "get_observation" + + 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) + + chat_messages = [ + Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"), + Message(role=MessageRoleEnum.USER.value, content="最近在北京的工作压力太大,有什么放松的建议吗?"), + Message(role=MessageRoleEnum.USER.value, content="说到朋友,我确实有几位很要好的朋友,我们经常一起出去吃饭。"), + Message(role=MessageRoleEnum.USER.value, content="对了,最近想换工作,你觉得北京的哪个区工作机会更多?"), + Message(role=MessageRoleEnum.USER.value, content="我很喜欢尝试新的美食,有没有推荐的美食应用?"), + Message(role=MessageRoleEnum.USER.value, content="我有时也喜欢自己在家做饭,你有没有好的海鲜菜谱推荐?"), + Message(role=MessageRoleEnum.USER.value, content="我在北京阿里云园区工作"), + Message(role=MessageRoleEnum.USER.value, content="我是阿里云百炼的工程师"), + Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?"), + ] + + worker.set_context(CHAT_MESSAGES, chat_messages) + worker.run() + + result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = "\n".join(result) + worker.logger.info(f"result={result}") + + @unittest.skip + def test_get_observation_with_time(self): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( @@ -108,8 +198,8 @@ class TestWorkersCn(unittest.TestCase): result = "\n".join(result) worker.logger.info(f"result={result}") - # @unittest.skip - def test_contra_repeat_cn(self): + @unittest.skip + def test_contra_repeat(self): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( @@ -166,3 +256,37 @@ class TestWorkersCn(unittest.TestCase): worker.logger.info(f"result2={result2}") worker.logger.info(f"result3={result3}") worker.logger.info(f"result4={result4}") + + # @unittest.skip + def test_get_reflection_subject(self): + name = "get_reflection_subject" + + 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="用户有要好朋友,常一起外出就餐。"), + MemoryNode(content="用户打算换工作,关心北京的工作机会分布。"), + MemoryNode(content="用户喜爱尝试新美食,求美食应用推荐。"), + MemoryNode(content="用户喜欢在家做饭,寻求海鲜菜谱。"), + MemoryNode(content="用户在北京阿里云园区工作。"), + MemoryNode(content="用户是阿里云百炼的工程师。"), + MemoryNode(content="用户目前的工作是大语言模型的应用开发"), + MemoryNode(content="用户想知道维持广泛社交关系的方法。"), + ] + + worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_handler.set_memories(INSIGHT_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