diff --git a/config/demo_config_cn.yaml b/config/demo_config_cn.yaml index d22194d6..9b639ade 100644 --- a/config/demo_config_cn.yaml +++ b/config/demo_config_cn.yaml @@ -111,6 +111,8 @@ worker: info_filter: class: memory.worker.write.info_filter_worker generation_model: dashscope_generation + generation_model_kwargs: + top_k: 1 load_today_memory: class: memory.worker.write.load_memory_worker retrieve_today_top_k: 100 diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 76dfa7f5..e0442665 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -53,6 +53,8 @@ class CliMemoryChat(BaseMemoryChat): """ self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model + self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) + self.stream: bool = stream self.human_name: str = human_name self.assistant_name: str = assistant_name @@ -174,7 +176,7 @@ class CliMemoryChat(BaseMemoryChat): self.logger.info(f"messages={messages}") # Invoke the Language Model with the constructed message context, respecting streaming setting - generated = self.generation_model.call(messages=messages, stream=self.stream) + generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) # In non-streaming interactions, explicitly save the AI's reply to memory if instructed if remember_response: 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 1bc0d6e2..dd5a61ea 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -24,7 +24,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): 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) + self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]): """ @@ -100,7 +100,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): self.logger.info(f"long_contra_repeat_message={long_contra_repeat_message}") # Invokes the language model for processing the constructed prompt - response = self.generation_model.call(messages=long_contra_repeat_message, top_k=self.generation_model_top_k) + response = self.generation_model.call(messages=long_contra_repeat_message, **self.generation_model_kwargs) # Handles the case where the model's response is empty if not response or not response.message.content: diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index f9df204e..2e6f7777 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -21,7 +21,7 @@ class UpdateInsightWorker(MemoryBaseWorker): def _parse_params(self, **kwargs): self.update_insight_threshold: float = kwargs.get("update_insight_threshold", 0.1) - self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) self.update_insight_max_count: int = kwargs.get("update_insight_max_count", 10) def filter_obs_nodes(self, @@ -113,7 +113,7 @@ class UpdateInsightWorker(MemoryBaseWorker): self.logger.info(f"Generated insight update message: {update_insight_message}") # Call the Language Model for insight update - response = self.generation_model.call(messages=update_insight_message, top_k=self.generation_model_top_k) + response = self.generation_model.call(messages=update_insight_message, **self.generation_model_kwargs) # Handle empty or invalid responses if not response.status or not response.message.content: diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py new file mode 100644 index 00000000..e035fae8 --- /dev/null +++ b/tests/worker/test_workers_en.py @@ -0,0 +1,352 @@ +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, 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 + + +class TestWorkersCn(unittest.TestCase): + """Tests for LLIEmbedding""" + + def setUp(self): + 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" + + 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( + 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="明天我要去高考"), + ] + + 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_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( + 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="我在一家叫京东的公司干活"), + ] + + # chat_messages = [ + # 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_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( + 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="明天是我生日"), + ] + + worker.set_context(CHAT_MESSAGES, chat_messages) + worker.run() + + result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = "\n".join(result) + worker.logger.info(f"result={result}") + + @unittest.skip + def test_contra_repeat(self): + name = "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(user_name="AI", target_name="用户", content="用户在美团干活"), + MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"), + ] + + worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.run() + result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) + for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + + nodes = [ + MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"), + MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), + ] + + worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.run() + result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) + for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + + nodes = [ + MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"), + MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), + ] + + worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.run() + result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) + for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + + nodes = [ + MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"), + MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴干活"), + MemoryNode(user_name="AI", target_name="用户", content="我不爱吃西瓜"), + ] + + worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.run() + result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) + for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + + worker.logger.info(f"result1={result1}") + 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.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.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}")