diff --git a/docs/images/logo_2.png b/docs/images/logo_2.png deleted file mode 100644 index 4578bbb0..00000000 Binary files a/docs/images/logo_2.png and /dev/null differ diff --git a/docs/images/logo_3.png b/docs/images/logo_3.png deleted file mode 100644 index 74d911fd..00000000 Binary files a/docs/images/logo_3.png and /dev/null differ diff --git a/examples/advance/custom_operator.md b/examples/advance/custom_operator.md new file mode 100644 index 00000000..5db9a7a0 --- /dev/null +++ b/examples/advance/custom_operator.md @@ -0,0 +1,49 @@ + +# 自定义 Operator 和 Worker + +1. 在 `contrib` 路径下创建新worker,命名为 `example_query_worker.py`: + ```bash + vim memoryscope/contrib/example_query_worker.py + ``` + +2. 写入新的自定义worker的程序,注意`class`的命名需要与文件名保持一致,为`ExampleQueryWorker`: + ```python + import datetime + + from memoryscope.constants.common_constants import QUERY_WITH_TS + from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker + + + class ExampleQueryWorker(MemoryBaseWorker): + + def _run(self): + + timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default + + assert "query" in self.chat_kwargs + query = self.chat_kwargs["query"] + if not query: + query = "" + else: + query = query.strip() + "\n You must add a `meow~` at the end of each of your answer." + + # Store the determined query and its timestamp in the context + self.set_workflow_context(QUERY_WITH_TS, (query, timestamp)) + ``` + +3. 创建yaml启动文件(复制demo_config.yaml) + ``` + cp memoryscope/core/config/demo_config.yaml examples/advance/replacement.yaml + vim examples/advance/replacement.yaml + ``` + +4. 在最下面插入新worker的定义,并且取代之前的默认`set_query`worker + ``` + set_query_meow: + class: contrib.example_query_worker + ``` + +5. 验证: + ``` + python quick-start-demo.py --config examples/advance/replacement.yaml + ``` diff --git a/examples/advance/replacement.yaml b/examples/advance/replacement.yaml new file mode 100644 index 00000000..66871336 --- /dev/null +++ b/examples/advance/replacement.yaml @@ -0,0 +1,185 @@ +global: + language: en + thread_pool_max_workers: 5 + logger_name: memoryscope + logger_name_time_suffix: "%Y%m%d_%H%M%S" + logger_to_screen: false + enable_ranker: false + enable_today_contra_repeat: true + enable_long_contra_repeat: false + output_memory_max_count: 20 + +memory_chat: + cli_memory_chat: + class: core.chat.cli_memory_chat + memory_service: memoryscope_service + generation_model: generation_model + stream: true + +memory_service: + memoryscope_service: + class: core.service.memory_scope_service + human_name: user + assistant_name: AI + memory_operations: + read_message: + class: core.operation.frontend_operation + workflow: read_message + description: "read short memory" + + retrieve_memory: + class: core.operation.frontend_operation + workflow: set_query_meow,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank + description: "retrieve long-term memory" + + list_memory: + class: core.operation.frontend_operation + workflow: set_query,retrieve_top_memory,print_memory + description: "read all long-term memory of the user, use `refresh_time=5` to refresh screen every 5 seconds." + + delete_memory: + class: core.operation.frontend_operation + workflow: set_query,retrieve_all_memory,delete_memory + description: "delete a single long-term memory" + + delete_all: + class: core.operation.frontend_operation + workflow: set_query,retrieve_all_memory,delete_all + description: "delete all long-term memory" + + add_memory: + class: core.operation.frontend_operation + workflow: add_memory + description: "add a single observation" + + consolidate_memory: + class: core.operation.consolidate_memory_op + workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory + description: "summary user's observation memory, run backend." + interval_time: 1 + + reflect_and_reconsolidate: + class: core.operation.backend_operation + workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory + description: "summary user's insight memory, run backend." + interval_time: 15 + +worker: + dummy: + class: core.worker.dummy_worker + generation_model: generation_model + embedding_model: embedding_model + rank_model: rank_model + read_message: + class: core.worker.frontend.read_message_worker + set_query: + class: core.worker.frontend.set_query_worker + set_query_meow: + class: contrib.example_query_worker + generation_model: generation_model + retrieve_obs_ins: + class: core.worker.frontend.retrieve_memory_worker + retrieve_obs_top_k: 100 + retrieve_ins_top_k: 100 + extract_time: + class: core.worker.frontend.extract_time_worker + generation_model: generation_model + semantic_rank: + class: core.worker.frontend.semantic_rank_worker + rank_model: rank_model + fuse_rerank: + class: core.worker.frontend.fuse_rerank_worker + fuse_score_threshold: 0.01 + fuse_ratio_dict: + conversation: 0.5 + observation: 1 + obs_customized: 1.2 + insight: 2.0 + fuse_time_ratio: 2.0 + retrieve_top_memory: + class: core.worker.frontend.retrieve_memory_worker + retrieve_obs_top_k: 100 + retrieve_ins_top_k: 100 + retrieve_expired_top_k: 100 + print_memory: + class: core.worker.frontend.print_memory_worker + retrieve_all_memory: + class: core.worker.frontend.retrieve_memory_worker + retrieve_obs_top_k: 1000 + retrieve_ins_top_k: 1000 + retrieve_expired_top_k: 1000 + delete_memory: + class: core.worker.backend.update_memory_worker + method: delete_memory + delete_all: + class: core.worker.backend.update_memory_worker + method: delete_all + add_memory: + class: core.worker.backend.update_memory_worker + method: from_query + info_filter: + class: core.worker.backend.info_filter_worker + generation_model: generation_model + load_today_memory: + class: core.worker.backend.load_memory_worker + retrieve_today_top_k: 100 + get_observation: + class: core.worker.backend.get_observation_worker + generation_model: generation_model + get_observation_with_time: + class: core.worker.backend.get_observation_with_time_worker + generation_model: generation_model + contra_repeat: + class: core.worker.backend.contra_repeat_worker + generation_model: generation_model + store_memory: + class: core.worker.backend.update_memory_worker + method: from_memory_key + memory_key: all + load_obs_and_insight: + class: core.worker.backend.load_memory_worker + retrieve_not_reflected_top_k: 100 + retrieve_not_updated_top_k: 100 + retrieve_insight_top_k: 100 + get_reflection_subject: + class: core.worker.backend.get_reflection_subject_worker + generation_model: generation_model + reflect_obs_cnt_threshold: 5 + update_insight: + class: core.worker.backend.update_insight_worker + generation_model: generation_model + rank_model: rank_model + embedding_model: embedding_model + update_insight_threshold: 0.01 + enable_parallel: false + long_contra_repeat: + class: core.worker.backend.long_contra_repeat_worker + generation_model: generation_model + long_contra_repeat_threshold: 0.5 + +model: + generation_model: + class: core.models.llama_index_generation_model + module_name: dashscope_generation + model_name: qwen-max + max_tokens: 2000 + temperature: 0.01 + embedding_model: + class: core.models.llama_index_embedding_model + module_name: dashscope_embedding + model_name: text-embedding-v2 + rank_model: + class: core.models.llama_index_rank_model + module_name: dashscope_rank + model_name: gte-rerank + top_n: 500 + +memory_store: + class: core.storage.llama_index_es_memory_store + embedding_model: embedding_model + index_name: memory_index + es_url: http://localhost:9200 + retrieve_mode: dense + +monitor: + class: core.storage.dummy_monitor \ No newline at end of file diff --git a/examples/api/agentscope_example.py b/examples/api/agentscope_example.py index 0f8b06cc..0b5ebc04 100644 --- a/examples/api/agentscope_example.py +++ b/examples/api/agentscope_example.py @@ -21,7 +21,7 @@ class MemoryScopeAgent(AgentBase): response = self.memory_chat.chat_with_memory(query=x.content) # Wrap the response in a message object in AgentScope - msg = Msg(name=self.name, content=response.message.content, role="Assistant") + msg = Msg(name=self.name, content=response.message.content, role="assistant") # Print/speak the message in this agent's voice self.speak(msg) diff --git a/examples/cli/README_ZH.md b/examples/cli/README_ZH.md new file mode 100644 index 00000000..007bb700 --- /dev/null +++ b/examples/cli/README_ZH.md @@ -0,0 +1,57 @@ +# MemoryScope 的命令行接口 + +## 使用方法 + +MemoryScope 可以通过两种不同的方式启动: + +### 1. 使用 YAML 配置文件 + +如果您更喜欢通过 YAML 文件配置设置,可以通过提供配置文件的路径来实现: +```bash +memoryscope --config_path=memoryscope/core/config/demo_config.yaml +``` + +### 2. 使用命令行参数 + +或者,您可以直接在命令行上指定所有参数: + +``` +# 中文 +memoryscope --language="cn" \ + --memory_chat_class="cli_memory_chat" \ + --human_name="用户" \ + --assistant_name="AI" \ + --generation_backend="dashscope_generation" \ + --generation_model="qwen-max" \ + --embedding_backend="dashscope_embedding" \ + --embedding_model="text-embedding-v2" \ + --enable_ranker=True \ + --rank_backend="dashscope_rank" \ + --rank_model="gte-rerank" + +# 英文 +memoryscope --language="en" \ + --memory_chat_class="cli_memory_chat" \ + --human_name="User" \ + --assistant_name="AI" \ + --generation_backend="openai_generation" \ + --generation_model="gpt-4o" \ + --embedding_backend="openai_embedding" \ + --embedding_model="text-embedding-3-small" \ + --enable_ranker=False +``` + + +以下是可以通过任一方法设置的可用选项: + +- `--language`: 对话中使用的语言。 +- `--memory_chat_class`: 管理聊天记录的类名。 +- `--human_name`: 人类用户的名字。 +- `--assistant_name`: AI 助手的名字。 +- `--generation_backend`: 用于生成回复的后端。 +- `--generation_model`: 用于生成回复的模型。 +- `--embedding_backend`: 用于文本嵌入的后端。 +- `--embedding_model`: 用于创建文本嵌入的模型。 +- `--enable_ranker`: 一个布尔值,指示是否使用排名器(默认为 False)。 +- `--rank_backend`: 用于排名回复的后端。 +- `--rank_model`: 用于排名回复的模型。 \ No newline at end of file diff --git a/memoryscope/contrib/example_query_worker.py b/memoryscope/contrib/example_query_worker.py new file mode 100644 index 00000000..90821522 --- /dev/null +++ b/memoryscope/contrib/example_query_worker.py @@ -0,0 +1,86 @@ +import datetime + +from memoryscope.constants.common_constants import QUERY_WITH_TS +from memoryscope.constants.language_constants import NONE_WORD +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.enumeration.message_role_enum import MessageRoleEnum + + +class ExampleQueryWorker(MemoryBaseWorker): + # NOTE: If you want to utilize the capabilities of the prompt handler, please be sure to include this sentence. + FILE_PATH: str = __file__ + + def _parse_params(self, **kwargs): + self.rewrite_history_count: int = kwargs.get("rewrite_history_count", 2) + self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) + + def rewrite_query(self, query: str) -> str: + chat_messages = self.chat_messages_scatter + if len(chat_messages) <= 1: + return query + + if chat_messages[-1].role == MessageRoleEnum.USER: + chat_messages = chat_messages[:-1] + chat_messages = chat_messages[-self.rewrite_history_count:] + + # get context + context_list = [] + for message in chat_messages: + context = message.content + if len(context) > 200: + context = context[:100] + context[-100:] + if message.role == MessageRoleEnum.USER: + context_list.append(f"{self.target_name}: {context}") + elif message.role == MessageRoleEnum.ASSISTANT: + context_list.append(f"Assistant: {context}") + + if not context_list: + return query + + system_prompt = self.prompt_handler.rewrite_query_system + user_query = self.prompt_handler.rewrite_query_query.format(query=query, + context="\n".join(context_list)) + rewrite_query_message = self.prompt_to_msg(system_prompt=system_prompt, + few_shot="", + user_query=user_query) + self.logger.info(f"rewrite_query_message={rewrite_query_message}") + + # Invoke the LLM to generate a response + response = self.generation_model.call(messages=rewrite_query_message, + **self.generation_model_kwargs) + + # Handle empty or unsuccessful responses + if not response.status or not response.message.content: + return query + + response_text = response.message.content + self.logger.info(f"rewrite_query.response_text={response_text}") + + if not response_text or response_text.lower() == self.get_language_value(NONE_WORD): + return query + + return response_text + + def _run(self): + query = "" # Default query value + timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default + + if "query" in self.chat_kwargs: + # set query if exists + query = self.chat_kwargs["query"] + if not query: + query = "" + query = query.strip() + + # set ts if exists + _timestamp = self.chat_kwargs.get("timestamp") + if _timestamp and isinstance(_timestamp, int): + timestamp = _timestamp + + if self.rewrite_history_count > 0: + t_query = self.rewrite_query(query=query) + if t_query: + query = t_query + + # Store the determined query and its timestamp in the context + self.set_workflow_context(QUERY_WITH_TS, (query, timestamp)) diff --git a/memoryscope/contrib/example_query_worker.yaml b/memoryscope/contrib/example_query_worker.yaml new file mode 100644 index 00000000..cebab04f --- /dev/null +++ b/memoryscope/contrib/example_query_worker.yaml @@ -0,0 +1,21 @@ +rewrite_query_system: + cn: | + 任务: 消除指代问题并重写 + 要求: 检查提供的问题是否存在指代。如果存在指代,通过上下文信息重写问题,使其信息充足,能够单独回答。如果没有指代问题,则回答“无”。 + en: | + Task: Eliminate referencing issues and rewrite + Requirements: Check the provided questions for any references. If references exist, rewrite the questions using contextual information to make them sufficiently informative so they can be answered independently. If there are no referencing issues, respond with "None". + +rewrite_query_query: + cn: | + 上下文: + {context} + 问题:{query} + 重写: + + en: | + Context: + {context} + Question: {query} + Rewrite: + diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index 0646d6e3..d4ecab1b 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -1,6 +1,7 @@ from dataclasses import dataclass, field from typing import Literal, Dict + @dataclass class Arguments(object): language: Literal["cn", "en"] = field(default="cn", metadata={"help": "support en & cn now"}) @@ -32,7 +33,7 @@ class Arguments(object): generation_backend: str = field(default="dashscope_generation", metadata={ "help": "global generation backend: openai_generation, dashscope_generation, etc."}) - generation_model: str = field(default="gpt-4o", metadata={ + generation_model: str = field(default="qwen-max", metadata={ "help": "global generation model: gpt-4o, gpt-4o-mini, gpt-4-turbo, qwen-max, etc."}) generation_params: dict = field(default_factory=lambda: {}, metadata={ @@ -41,7 +42,7 @@ class Arguments(object): embedding_backend: str = field(default="dashscope_generation", metadata={ "help": "global embedding backend: openai_embedding, dashscope_embedding, etc."}) - embedding_model: str = field(default="text-embedding-3-small", metadata={ + embedding_model: str = field(default="text-embedding-v2", metadata={ "help": "global embedding model: text-embedding-3-large, text-embedding-3-small, text-embedding-ada-002, " "text-embedding-v2, etc."}) diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py index 94f1b7b2..12381634 100644 --- a/memoryscope/core/memoryscope.py +++ b/memoryscope/core/memoryscope.py @@ -91,7 +91,7 @@ class MemoryScope(ConfigManager): self.close() @property - def content(self): + def context(self): return self._context @property diff --git a/memoryscope/core/models/llama_index_generation_model.py b/memoryscope/core/models/llama_index_generation_model.py index f38f08d0..05941e13 100644 --- a/memoryscope/core/models/llama_index_generation_model.py +++ b/memoryscope/core/models/llama_index_generation_model.py @@ -11,6 +11,7 @@ from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen from memoryscope.core.utils.logger import Logger + class LlamaIndexGenerationModel(BaseModel): """ This class represents a generation model within the LlamaIndex framework, @@ -73,7 +74,7 @@ class LlamaIndexGenerationModel(BaseModel): model_response.message.content += delta model_response.delta = response.delta yield model_response - + self.logger.info(self.logger.format_chat_message(model_response)) return gen() else: if isinstance(call_result, CompletionResponse): diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index 86fa9580..9a07528e 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -31,8 +31,6 @@ class TestWorkersCn(unittest.TestCase): enable_ranker=True, ) self.ms = MemoryScope(arguments=self.arguments) - config = self.ms.dump_config() - self.ms.logger.info(f"config=\n{config}") def tearDown(self): self.ms.close() @@ -42,12 +40,13 @@ class TestWorkersCn(unittest.TestCase): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) query = "明天我去上海出差" query_timestamp = int(datetime.datetime.now().timestamp()) @@ -57,17 +56,18 @@ class TestWorkersCn(unittest.TestCase): result = worker.get_workflow_context(EXTRACT_TIME_DICT) worker.logger.info(f"result={result}") - # @unittest.skip + @unittest.skip def test_info_filter(self): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name), @@ -96,12 +96,13 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗", @@ -144,12 +145,13 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name), @@ -178,12 +180,13 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。", role_name=self.arguments.human_name), @@ -209,12 +212,13 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术", role_name=self.arguments.human_name), @@ -238,12 +242,13 @@ class TestWorkersCn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), @@ -298,12 +303,13 @@ class TestWorkersCn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name), @@ -334,12 +340,13 @@ class TestWorkersCn(unittest.TestCase): name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户喜欢打王者荣耀", role_name=self.arguments.human_name), @@ -356,12 +363,13 @@ class TestWorkersCn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name), @@ -375,3 +383,34 @@ class TestWorkersCn(unittest.TestCase): result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}") + + # @unittest.skip + def test_example_query_worker(self): + name = "example_query_worker" + + worker: MemoryBaseWorker = init_instance_by_config( + config={ + "class": "contrib.example_query_worker", + "generation_model": "generation_model", + }, + name=name, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name, + "chat_kwargs": {"query": "我一直很爱他们"}}, + context_lock=None, + memoryscope_context=self.ms.context, + thread_pool=self.ms._context.thread_pool) + + chat_messages = [ + Message(role=MessageRoleEnum.USER.value, content="我的两个孩子分别叫小明和小红", + role_name=self.arguments.human_name), + Message(role=MessageRoleEnum.ASSISTANT.value, + content="很高兴认识您和您的家庭成员!小明和小红是非常通俗且好听的名字。", + role_name=self.arguments.assistant_name), + Message(role=MessageRoleEnum.USER.value, content="我一直很爱他们", role_name=self.arguments.human_name), + ] + + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.run() + + result = worker.get_workflow_context(QUERY_WITH_TS) + worker.logger.info(f"result={result}") diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 3dd46442..7d1f79f4 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -3,7 +3,7 @@ import unittest from memoryscope.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, \ - MEMORYSCOPE_CONTEXT + MEMORYSCOPE_CONTEXT, TARGET_NAME, CHAT_MESSAGES_SCATTER from memoryscope.core.config.arguments import Arguments from memoryscope.core.memoryscope import MemoryScope from memoryscope.core.utils.tool_functions import init_instance_by_config @@ -17,7 +17,7 @@ class TestWorkersEn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - arguments = Arguments( + self.arguments = Arguments( language="en", human_name="user", assistant_name="AI", @@ -29,9 +29,7 @@ class TestWorkersEn(unittest.TestCase): rank_backend="dashscope_rank", rank_model="gte-rerank", ) - self.ms = MemoryScope(arguments=arguments) - config = self.ms.dump_config() - self.ms.logger.info(f"config=\n{config}") + self.ms = MemoryScope(arguments=self.arguments) def tearDown(self): self.ms.close() @@ -41,12 +39,13 @@ class TestWorkersEn(unittest.TestCase): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) query = "I will be on a business trip to Shanghai tomorrow." query_timestamp = int(datetime.datetime.now().timestamp()) @@ -61,12 +60,13 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."), @@ -87,12 +87,13 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="Do you know where the freshest seafood is in Beijing?"), @@ -135,12 +136,13 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) # FIXME Does the appearance of 'am' indicate the presence of a time keyword? chat_messages = [ @@ -164,12 +166,13 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -206,12 +209,13 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -238,12 +242,13 @@ class TestWorkersEn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), @@ -296,12 +301,13 @@ class TestWorkersEn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), @@ -333,12 +339,13 @@ class TestWorkersEn(unittest.TestCase): name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users like to play King of Glory"), @@ -355,12 +362,13 @@ class TestWorkersEn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms._context.worker_conf_dict[name], + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms._context}, + context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms._context.thread_pool) + memoryscope_context=self.ms.context, + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), @@ -374,3 +382,34 @@ class TestWorkersEn(unittest.TestCase): result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}") + + # @unittest.skip + def test_example_query_worker(self): + name = "example_query_worker" + + worker: MemoryBaseWorker = init_instance_by_config( + config={ + "class": "contrib.example_query_worker", + "generation_model": "generation_model", + }, + name=name, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name, + "chat_kwargs": {"query": "I have always loved them."}}, + context_lock=None, + memoryscope_context=self.ms.context, + thread_pool=self.ms._context.thread_pool) + + chat_messages = [ + Message(role=MessageRoleEnum.USER.value, content="My two children are named Xiaoming and Xiaohong.", + role_name=self.arguments.human_name), + Message(role=MessageRoleEnum.ASSISTANT.value, + content="I am very pleased to meet you and your family members! Xiaoming and Xiaohong are very pleasant names.", + role_name=self.arguments.assistant_name), + Message(role=MessageRoleEnum.USER.value, content="I have always loved them.", role_name=self.arguments.human_name), + ] + + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.run() + + result = worker.get_workflow_context(QUERY_WITH_TS) + worker.logger.info(f"result={result}") \ No newline at end of file