From ee620a01bb787bfd99d44f995c94e0bedd4c1a21 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 25 Jun 2024 22:55:57 +0800 Subject: [PATCH 01/41] [dev] modify g content --- config/config.yaml | 6 +- memory_scope/chat_v2/__init__.py | 0 memory_scope/chat_v2/base_memory_chat.py | 17 ++++ memory_scope/chat_v2/base_memory_service.py | 4 + memory_scope/chat_v2/cli_memory_chat.py | 83 ++++++++++++++++ memory_scope/chat_v2/global_context.py | 23 +++++ memory_scope/chat_v2/memory_chat.py | 67 +++++++++++++ memory_scope/chat_v2/memory_service.py | 70 ++++++++++++++ memory_scope/cli_job.py | 101 ++++++++++++++++++++ memory_scope/utils/tool_functions.py | 38 ++++---- 10 files changed, 387 insertions(+), 22 deletions(-) create mode 100644 memory_scope/chat_v2/__init__.py create mode 100644 memory_scope/chat_v2/base_memory_chat.py create mode 100644 memory_scope/chat_v2/base_memory_service.py create mode 100644 memory_scope/chat_v2/cli_memory_chat.py create mode 100644 memory_scope/chat_v2/global_context.py create mode 100644 memory_scope/chat_v2/memory_chat.py create mode 100644 memory_scope/chat_v2/memory_service.py create mode 100644 memory_scope/cli_job.py diff --git a/config/config.yaml b/config/config.yaml index d9a0312d..d753b212 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -1,14 +1,16 @@ global_config: - thread_pool_max_count: 5 + language: en + max_workers: 5 dash_scope_apikey: open_ai_apikey: - language: en chat_list: - memory_chat memory_chat: memory_service: memory_chat_service + generation_model: dashscope_generation memory_chat_service: class: memory.base_memory_service + history_msg_count: 5 memory_operations: - name: read_memory class: memory.workflow.base_workflow diff --git a/memory_scope/chat_v2/__init__.py b/memory_scope/chat_v2/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/chat_v2/base_memory_chat.py b/memory_scope/chat_v2/base_memory_chat.py new file mode 100644 index 00000000..0d2566d2 --- /dev/null +++ b/memory_scope/chat_v2/base_memory_chat.py @@ -0,0 +1,17 @@ +from abc import ABCMeta, abstractmethod + + +class BaseMemoryChat(metaclass=ABCMeta): + def __init__(self, **kwargs): + self.kwargs = kwargs + + + @abstractmethod + def chat_with_memory(self, query: str): + """ + :param query: + :return: + """ + + def run(self): + pass diff --git a/memory_scope/chat_v2/base_memory_service.py b/memory_scope/chat_v2/base_memory_service.py new file mode 100644 index 00000000..9cf3fd76 --- /dev/null +++ b/memory_scope/chat_v2/base_memory_service.py @@ -0,0 +1,4 @@ +class BaseMemoryService(object): + def __init__(self, **kwargs): + + self.kwargs = kwargs diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py new file mode 100644 index 00000000..1444edcd --- /dev/null +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -0,0 +1,83 @@ +import datetime + +import questionary +from rich.console import Console + +from .memory_chat import MemoryChat +from enumeration.message_role_enum import MessageRoleEnum +from scheme.message import Message + + +class CliMemoryChat(MemoryChat): + + USER_COMMANDS = { + "/exit": "exit the CLI", + "/memory": "print the current contents of agent memory", + "/retrieve": "retrieve related memory", + "/log": "log chat progress", + # TODO add more commands + } + + def chat_with_memory(self, query): # for testing + query = query.strip() + if not query: + return + + time_created = int(datetime.datetime.now().timestamp()) + message = Message( + role=MessageRoleEnum.USER, content=query, time_created=time_created + ) + messages = [message] + return self.generation_model.call(messages=messages, stream=True) + + def retrieve_all(self): # for testing + return "memory 1. 2. 3." + + def run(self): + console = Console() + while True: + query = questionary.text( + "Enter your message or command:", + multiline=False, + qmark=">", + ).ask() + + query = query.rstrip() + + if query == "": + console.print("Empty input received. Try again!") + continue + + # Handle CLI commands + if query.startswith("/"): + if query.lower() == "/exit": + break + elif query.lower() == "/memory": + console.print(self.memory_service.retrieve_all()) + elif query.lower() == "/help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(cmd, "bold") + questionary.print(f" {desc}") + + continue + + while True: + try: + # with console.status("[bold cyan]Thinking..."): + for msg in self.chat_with_memory(query=query): + console.print(msg.delta, end="") + console.print() + break + except KeyboardInterrupt: + console.print("User interrupt occurred.") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: + break + except Exception as e: + console.print( + f"An exception occurred when running chat_with_memory(): {e}" + ) + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: + break diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py new file mode 100644 index 00000000..71484e68 --- /dev/null +++ b/memory_scope/chat_v2/global_context.py @@ -0,0 +1,23 @@ +from concurrent.futures import ThreadPoolExecutor +from typing import Dict, Any + +import pydantic + +from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat +from memory_scope.enumeration.language_enum import LanguageEnum +from memory_scope.models.base_model import BaseModel +from memory_scope.storage.base_monitor import BaseMonitor +from memory_scope.storage.base_vector_store import BaseVectorStore + + +class GlobalContext(pydantic.BaseModel): + global_config: Dict[str, Any] = pydantic.Field({}, description="global configs") + model_dict: Dict[str, BaseModel] = pydantic.Field({}, description="global model_dict") + memory_chat_dict: Dict[str, BaseMemoryChat] = pydantic.Field({}, description="global memory_chat_dict") + vector_store: BaseVectorStore | None = pydantic.Field(None, description="global vector_store") + monitor: BaseMonitor | None = pydantic.Field(None, description="global monitor") + thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool") + language: LanguageEnum = pydantic.Field(LanguageEnum.CN, description="language: cn / en") + + +G_CONTEXT = GlobalContext() diff --git a/memory_scope/chat_v2/memory_chat.py b/memory_scope/chat_v2/memory_chat.py new file mode 100644 index 00000000..859758de --- /dev/null +++ b/memory_scope/chat_v2/memory_chat.py @@ -0,0 +1,67 @@ +import datetime +from typing import List + +from .base_memory_chat import BaseMemoryChat +from .global_context import GLOBAL_CONTEXT +from enumeration.message_role_enum import MessageRoleEnum +from models.base_model import BaseModel +from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT +from scheme.message import Message +from .memory_service import MemoryService + + +class MemoryChat(BaseMemoryChat): + + def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs): + super().__init__(**kwargs) + self.memory_service = MemoryService(chat_name=chat_name, **kwargs) + self.generation_model_name: str = generation_model + self.history_msg_count: int = history_msg_count + + self._generation_model: BaseModel | None = None + self.history_message_list: List[Message] = [] + + @property + def generation_model(self): + if self._generation_model is None: + self._generation_model = GLOBAL_CONTEXT.model_dict[ + self.generation_model_name + ] + return self._generation_model + + @staticmethod + def get_system_prompt(related_memories: List[str], time_created: int) -> Message: + system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language] + if related_memories: + memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language] + system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories) + return Message( + role=MessageRoleEnum.SYSTEM, + content=system_prompt.strip(), + time_created=time_created, + ) + + def chat_with_memory(self, query: str): + query = query.strip() + if not query: + return + + time_created = int(datetime.datetime.now().timestamp()) + new_message: Message = Message( + role=MessageRoleEnum.USER, content=query, time_created=time_created + ) + related_memories: List[str] = self.memory_service.retrieve(message=new_message) + system_message = self.get_system_prompt(related_memories, time_created) + self.history_message_list.append(new_message) + self.history_message_list = self.history_message_list[-self.history_msg_count :] + all_messages = [system_message] + self.history_message_list + # TODO at xian zhe + return self.generation_model.call(messages=all_messages, stream=True) + + def run(self): + self.memory_service.start_memory_backend() + while True: + query = input("wait for input:") + if query in ["stop", "停止"]: + break + self.chat_with_memory(query=query) diff --git a/memory_scope/chat_v2/memory_service.py b/memory_scope/chat_v2/memory_service.py new file mode 100644 index 00000000..e9fca97a --- /dev/null +++ b/memory_scope/chat_v2/memory_service.py @@ -0,0 +1,70 @@ +from constants.common_constants import RELATED_MEMORIES +from enumeration.memory_method_enum import MemoryMethodEnum +from scheme.message import Message +from utils.pipeline import Pipeline +from .base_memory_service import BaseMemoryService + + +class MemoryService(BaseMemoryService): + def __init__( + self, + chat_name: str, + retrieve_pipeline: str, + retrieve_all_pipeline: str, + summary_short_pipeline: str, + summary_long_pipeline: str, + summary_short_interval_time: int = 60, + summary_short_minimum_count: int = 5, + summary_long_interval_time: int = 60 * 5, + summary_long_minimum_count: int = 5 * 5, + **kwargs + ): + super().__init__(**kwargs) + self.retrieve_pipeline = Pipeline( + chat_name=chat_name, + memory_method_type=MemoryMethodEnum.RETRIEVE, + pipeline_str=retrieve_pipeline, + ) + + self.retrieve_all_pipeline = Pipeline( + chat_name=chat_name, + memory_method_type=MemoryMethodEnum.RETRIEVE_ALL, + pipeline_str=retrieve_all_pipeline, + ) + + self.summary_short_pipeline = Pipeline( + chat_name=chat_name, + memory_method_type=MemoryMethodEnum.SUMMARY_SHORT, + pipeline_str=summary_short_pipeline, + loop_interval_time=summary_short_interval_time, + loop_minimum_count=summary_short_minimum_count, + ) + + self.summary_long_pipeline = Pipeline( + chat_name=chat_name, + memory_method_type=MemoryMethodEnum.SUMMARY_LONG, + pipeline_str=summary_long_pipeline, + loop_interval_time=summary_long_interval_time, + loop_minimum_count=summary_long_minimum_count, + ) + + def retrieve(self, message: Message): + self.retrieve_pipeline.submit_message(message, with_lock=False) + self.summary_short_pipeline.submit_message(message) + self.summary_long_pipeline.submit_message(message) + return self.retrieve_pipeline.run(RELATED_MEMORIES) + + def retrieve_all(self): + return self.retrieve_all_pipeline.run(RELATED_MEMORIES) + + def start_memory_backend(self): + self.summary_short_pipeline.start_loop_run() + self.summary_long_pipeline.start_loop_run() + + def get_worker_list(self) -> list: + worker_set = set() + worker_set.update(self.retrieve_pipeline.worker_set) + worker_set.update(self.retrieve_all_pipeline.worker_set) + worker_set.update(self.summary_short_pipeline.worker_set) + worker_set.update(self.summary_long_pipeline.worker_set) + return sorted(worker_set) diff --git a/memory_scope/cli_job.py b/memory_scope/cli_job.py new file mode 100644 index 00000000..3e02e8e2 --- /dev/null +++ b/memory_scope/cli_job.py @@ -0,0 +1,101 @@ +import json +import os +from concurrent.futures import ThreadPoolExecutor +from typing import Dict, Any + +import yaml + +from chat_v2.global_context import G_CONTEXT +from enumeration.language_enum import LanguageEnum +from enumeration.model_enum import ModelEnum +from utils.logger import Logger +from utils.tool_functions import ( + complete_config_name, + init_instance_by_config, +) + + +class CliJob(object): + + def __init__(self, config_path: str, config_suffix: str = ".yaml"): + self.config_path: str = config_path + self.config_suffix: str = config_suffix + + self.config: Dict[str, Any] = {} + self.global_config: Dict[str, Any] = {} + + self.logger: Logger = Logger.get_logger("memory_chat") + + def init_model(self, model_name: str): + if not model_name or model_name in G_CONTEXT.model_dict: + return + + with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f: + model_config = json.load(f) + GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config) + + def init_workers(self): + """load worker config & init workers""" + worker_config_name: str = self.config["workers"] + with open( + os.path.join(self.config_base_dir, complete_config_name(worker_config_name)) + ) as f: + worker_config_dict = json.load(f) + + for worker_name, worker_config in worker_config_dict.items(): + if worker_name not in self.worker_chat_dict: + continue + + chat_name_list = self.worker_chat_dict[worker_name] + for chat_name in chat_name_list: + if chat_name not in GLOBAL_CONTEXT.worker_dict: + GLOBAL_CONTEXT.worker_dict[chat_name] = {} + GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = ( + init_instance_by_config( + worker_config, + suffix_name="worker", + **GLOBAL_CONTEXT.global_configs, + ) + ) + + self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value)) + self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value)) + self.init_model(worker_config.get(ModelEnum.RANK_MODEL.value)) + + @staticmethod + def set_global_config(): + # TODO at sen, set global_configs & set apikey into env + G_CONTEXT.language = LanguageEnum(G_CONTEXT.global_configs["language"]) + G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(G_CONTEXT.global_configs["max_workers"])) + + def init_global_content_by_config(self): + config_path = self.config_path + if not self.config_path.endswith(self.config_suffix): + config_path += self.config_suffix + + with open(config_path) as f: + self.config = yaml.load(f, yaml.FullLoader) + + G_CONTEXT.global_configs = self.global_config = self.config["global_configs"] + self.set_global_config() + + # init memory_chat + for chat_name in self.global_config["chat_list"]: + memory_chat_config = self.config[chat_name] + G_CONTEXT.memory_chat_dict[chat_name] = init_instance_by_config(memory_chat_config, chat_name=chat_name) + + for model_config in + + GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config) + + # TODO no db and monitor now + GLOBAL_CONTEXT.vector_store = init_instance_by_config( + self.config["vector_store"] + ) + GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) + + @staticmethod + def run(): + with GLOBAL_CONTEXT.thread_pool: + memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0] + memory_chat.run() diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 37ed8c8a..5de34148 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -1,8 +1,8 @@ import re -from importlib import import_module from datetime import datetime +from importlib import import_module -from enumeration.message_role_enum import MessageRoleEnum +from memory_scope.enumeration.message_role_enum import MessageRoleEnum def under_line_to_hump(underline_str): @@ -10,26 +10,24 @@ def under_line_to_hump(underline_str): return sub[0:1].upper() + sub[1:] -def init_instance_by_config( - config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs -): - clazz_path = config.pop("clazz") - if not clazz_path: - raise RuntimeError("empty clazz_path!") - clazz_name_split = clazz_path.split(".") - clazz_name: str = clazz_name_split[-1] - if suffix_name and not clazz_name.endswith(suffix_name): - clazz_name = f"{clazz_name}_{suffix_name}" +def init_instance_by_config(config: dict, default_class_path: str = "", suffix_name: str = "", **kwargs): + class_name = config.pop("class") + if not class_name: + raise RuntimeError("empty class_name!") - # 构造path - clazz_paths = [] - if default_clazz_path: - clazz_paths.append(default_clazz_path) - clazz_paths.extend(clazz_name_split[:-1]) - clazz_paths.append(clazz_name) - module = import_module(".".join(clazz_paths)) + class_name_split = class_name.split(".") + class_name: str = class_name_split[-1] + if suffix_name and not class_name.lower().endswith(suffix_name.lower()): + class_name = f"{class_name}_{suffix_name}" + class_name_split[-1] = class_name - cls_name = under_line_to_hump(clazz_name) + class_paths = [] + if default_class_path: + class_paths.append(default_class_path) + class_paths.extend(class_name_split) + module = import_module(".".join(class_paths)) + + cls_name = under_line_to_hump(class_name) return getattr(module, cls_name)(**config, **kwargs) From 0505b04781a8cb5d73e89cc41a9fd4e9c00277a0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 10:07:45 +0800 Subject: [PATCH 02/41] [dev] add dummy worker --- config/config.yaml | 85 ++++++------ memory_scope/chat_v2/base_memory_chat.py | 2 +- memory_scope/chat_v2/global_context.py | 11 +- memory_scope/cli_job.py | 96 +++++-------- memory_scope/constants/common_constants.py | 4 + memory_scope/memory/base_memory_service.py | 6 + memory_scope/memory/worker/__init__.py | 0 memory_scope/memory/worker/base_worker.py | 73 ++++++++++ memory_scope/memory/worker/dummy_worker.py | 6 + .../memory/worker/memory_base_worker.py | 70 ++++++++++ memory_scope/memory/workflow/__init__.py | 0 .../memory/workflow/backend_v1_workflow.py | 43 ++++++ memory_scope/memory/workflow/base_workflow.py | 126 ++++++++++++++++++ .../memory/workflow/frontend_workflow.py | 12 ++ memory_scope/models/__init__.py | 3 +- memory_scope/models/base_model.py | 2 +- memory_scope/scheme/message.py | 4 +- 17 files changed, 428 insertions(+), 115 deletions(-) create mode 100644 memory_scope/memory/base_memory_service.py create mode 100644 memory_scope/memory/worker/__init__.py create mode 100644 memory_scope/memory/worker/base_worker.py create mode 100644 memory_scope/memory/worker/dummy_worker.py create mode 100644 memory_scope/memory/worker/memory_base_worker.py create mode 100644 memory_scope/memory/workflow/__init__.py create mode 100644 memory_scope/memory/workflow/backend_v1_workflow.py create mode 100644 memory_scope/memory/workflow/base_workflow.py create mode 100644 memory_scope/memory/workflow/frontend_workflow.py diff --git a/config/config.yaml b/config/config.yaml index d753b212..c0279fc7 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -3,34 +3,48 @@ global_config: max_workers: 5 dash_scope_apikey: open_ai_apikey: - chat_list: - - memory_chat memory_chat: - memory_service: memory_chat_service - generation_model: dashscope_generation -memory_chat_service: - class: memory.base_memory_service - history_msg_count: 5 - memory_operations: - - name: read_memory - class: memory.workflow.base_workflow - workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank - work_type: frontend - - name: list_memory - class: memory.workflow.base_workflow - workflow: dummy - work_type: frontend - - name: extract_memory - class: memory.workflow.base_workflow - workflow: dummy - work_type: backend - interval_time: 60 - min_count: 5 - - name: reflect_memory - class: memory.workflow.base_workflow - workflow: dummy - work_type: backend - interval_time: 300 + cli_memory_chat: + class: chat.cli_memory_chat + memory_service: memory_chat_service + generation_model: dashscope_generation +memory_service: + memory_chat_service: + class: memory.base_memory_service + history_msg_count: 5 + memory_operations: + read_memory: + class: memory.workflow.base_workflow + workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank + work_type: frontend + list_memory: + class: memory.workflow.base_workflow + workflow: dummy + work_type: frontend + extract_memory: + class: memory.workflow.base_workflow + workflow: dummy + work_type: backend + interval_time: 60 + min_count: 5 + reflect_memory: + class: memory.workflow.base_workflow + workflow: dummy + work_type: backend + interval_time: 300 +models: + dashscope_generation: + clazz: models.llama_index_generation_model + module_name: DashScope + model_name: qwen-max + dashscope_embedding: + clazz: models.base_embedding_model + module_name: DashScopeEmbedding + model_name: text-embedding-v2 + dashscope_rank: + clazz: models.base_rank_model + module_name: DashScopeRerank + model_name: gte-rerank vector_store: clazz: storage.base_vector_store index_name: memory_test @@ -39,21 +53,8 @@ monitor: clazz: storage.base_monitor index_name: memory_test workers: - - name: update_insight + update_insight: clazz: worker.summary_long.update_insight generation_model: dashscope_generation embedding_model: dashscope_embedding - rank_model: dashscope_rank -models: - - name: dashscope_generation - clazz: models.llama_index_generation_model - module_name: DashScope - model_name: qwen-max - - name: dashscope_embedding - clazz: models.base_embedding_model - module_name: DashScopeEmbedding - model_name: text-embedding-v2 - - name: dashscope_rank - clazz: models.base_rank_model - module_name: DashScopeRerank - model_name: gte-rerank + rank_model: dashscope_rank \ No newline at end of file diff --git a/memory_scope/chat_v2/base_memory_chat.py b/memory_scope/chat_v2/base_memory_chat.py index 0d2566d2..b5bda713 100644 --- a/memory_scope/chat_v2/base_memory_chat.py +++ b/memory_scope/chat_v2/base_memory_chat.py @@ -2,7 +2,7 @@ from abc import ABCMeta, abstractmethod class BaseMemoryChat(metaclass=ABCMeta): - def __init__(self, **kwargs): + def __init__(self, memory_service: str, **kwargs): self.kwargs = kwargs diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py index 71484e68..969c902c 100644 --- a/memory_scope/chat_v2/global_context.py +++ b/memory_scope/chat_v2/global_context.py @@ -5,15 +5,20 @@ import pydantic from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat from memory_scope.enumeration.language_enum import LanguageEnum +from memory_scope.memory.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_vector_store import BaseVectorStore class GlobalContext(pydantic.BaseModel): - global_config: Dict[str, Any] = pydantic.Field({}, description="global configs") - model_dict: Dict[str, BaseModel] = pydantic.Field({}, description="global model_dict") - memory_chat_dict: Dict[str, BaseMemoryChat] = pydantic.Field({}, description="global memory_chat_dict") + global_config: Dict[str, Any] = pydantic.Field({}, description="global config") + worker_config: Dict[str, Any] = pydantic.Field({}, description="worker config") + + memory_service_dict: Dict[str, BaseMemoryService] = pydantic.Field({}, description="memory_service dict") + model_dict: Dict[str, BaseModel] = pydantic.Field({}, description="model dict") + memory_chat_dict: Dict[str, BaseMemoryChat] = pydantic.Field({}, description="memory_chat dict") + vector_store: BaseVectorStore | None = pydantic.Field(None, description="global vector_store") monitor: BaseMonitor | None = pydantic.Field(None, description="global monitor") thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool") diff --git a/memory_scope/cli_job.py b/memory_scope/cli_job.py index 3e02e8e2..48844129 100644 --- a/memory_scope/cli_job.py +++ b/memory_scope/cli_job.py @@ -1,5 +1,3 @@ -import json -import os from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any @@ -7,12 +5,8 @@ import yaml from chat_v2.global_context import G_CONTEXT from enumeration.language_enum import LanguageEnum -from enumeration.model_enum import ModelEnum from utils.logger import Logger -from utils.tool_functions import ( - complete_config_name, - init_instance_by_config, -) +from utils.tool_functions import init_instance_by_config class CliJob(object): @@ -20,82 +14,54 @@ class CliJob(object): def __init__(self, config_path: str, config_suffix: str = ".yaml"): self.config_path: str = config_path self.config_suffix: str = config_suffix - self.config: Dict[str, Any] = {} - self.global_config: Dict[str, Any] = {} - self.logger: Logger = Logger.get_logger("memory_chat") - - def init_model(self, model_name: str): - if not model_name or model_name in G_CONTEXT.model_dict: - return - - with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f: - model_config = json.load(f) - GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config) - - def init_workers(self): - """load worker config & init workers""" - worker_config_name: str = self.config["workers"] - with open( - os.path.join(self.config_base_dir, complete_config_name(worker_config_name)) - ) as f: - worker_config_dict = json.load(f) - - for worker_name, worker_config in worker_config_dict.items(): - if worker_name not in self.worker_chat_dict: - continue - - chat_name_list = self.worker_chat_dict[worker_name] - for chat_name in chat_name_list: - if chat_name not in GLOBAL_CONTEXT.worker_dict: - GLOBAL_CONTEXT.worker_dict[chat_name] = {} - GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = ( - init_instance_by_config( - worker_config, - suffix_name="worker", - **GLOBAL_CONTEXT.global_configs, - ) - ) - - self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value)) - self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value)) - self.init_model(worker_config.get(ModelEnum.RANK_MODEL.value)) + self.logger: Logger = Logger.get_logger("cli_job") @staticmethod - def set_global_config(): - # TODO at sen, set global_configs & set apikey into env - G_CONTEXT.language = LanguageEnum(G_CONTEXT.global_configs["language"]) - G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(G_CONTEXT.global_configs["max_workers"])) + def set_global_config(global_config: Dict[str, Any]): + """ set global_configs & set apikey into env + :return: + TODO at sen + """ + G_CONTEXT.global_config = global_config + G_CONTEXT.language = LanguageEnum(global_config["language"]) + G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"])) def init_global_content_by_config(self): + # load config config_path = self.config_path if not self.config_path.endswith(self.config_suffix): config_path += self.config_suffix - with open(config_path) as f: self.config = yaml.load(f, yaml.FullLoader) - G_CONTEXT.global_configs = self.global_config = self.config["global_configs"] - self.set_global_config() + # set global_config + self.set_global_config(self.config["global_config"]) # init memory_chat - for chat_name in self.global_config["chat_list"]: - memory_chat_config = self.config[chat_name] - G_CONTEXT.memory_chat_dict[chat_name] = init_instance_by_config(memory_chat_config, chat_name=chat_name) + for name, conf in self.config["memory_chat"].items(): + G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name) - for model_config in + # set memory_service + for name, conf in self.config["memory_service"].items(): + G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name) - GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config) + # init models + for name, conf in self.config["models"].items(): + G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) - # TODO no db and monitor now - GLOBAL_CONTEXT.vector_store = init_instance_by_config( - self.config["vector_store"] - ) - GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) + # init vector_store + G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"]) + + # init monitor + G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) + + # set worker config + G_CONTEXT.worker_config = self.config["workers"] @staticmethod def run(): - with GLOBAL_CONTEXT.thread_pool: - memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0] + with G_CONTEXT.thread_pool: + memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index 5f5f9d63..d7379f91 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -1,3 +1,7 @@ +RESULT = "result" + +CHAT_MESSAGES = "chat_messages" + RELATED_MEMORIES = "related_memories" MESSAGES = "messages" diff --git a/memory_scope/memory/base_memory_service.py b/memory_scope/memory/base_memory_service.py new file mode 100644 index 00000000..d05652a3 --- /dev/null +++ b/memory_scope/memory/base_memory_service.py @@ -0,0 +1,6 @@ +from abc import ABCMeta + + +class BaseMemoryService(metaclass=ABCMeta): + def __init__(self, **kwargs): + self.kwargs = kwargs diff --git a/memory_scope/memory/worker/__init__.py b/memory_scope/memory/worker/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py new file mode 100644 index 00000000..a47fd1ed --- /dev/null +++ b/memory_scope/memory/worker/base_worker.py @@ -0,0 +1,73 @@ +from typing import Any, Dict + +from utils.logger import Logger +from utils.timer import Timer + + +class BaseWorker(object): + + def __init__(self, raise_exception: bool = True, **kwargs): + super(BaseWorker, self).__init__(**kwargs) + # 异常是否继续执行 + self.raise_exception: bool = raise_exception + + # True 为正常运行,False会结束整个pipeline + self.continue_run: bool = True + + # 短name + self._name_simple: str = "" + + # 是否多线程环境 + self.is_multi_thread: bool = False + + # pipeline 上下文 + self.context_dict: Dict[str, Any] | None = None + self.context_lock = None + + # 日志 + self.logger: Logger = Logger.get_logger() + + # worker 参数 + self.kwargs: dict = kwargs + + def _run(self): + raise NotImplementedError + + def run(self): + self.logger.info(f"----- Begin {self.name_simple} -----") + with Timer(self.name_simple, log_time=False) as t: + if self.raise_exception: + self._run() + else: + try: + self._run() + except Exception as e: + self.logger.exception(f"run {self.name_simple} failed! args={e.args}") + + self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----") + + def set_context_dict(self, context_dict: dict, context_lock=None): + self.context_dict = context_dict + if context_lock is not None: + self.context_lock = context_lock + self.is_multi_thread = True + + def get_context(self, key: str, default=None): + return self.context_dict.get(key, default) + + def set_context(self, key: str, value: Any): + if self.is_multi_thread: + # add lock to multi thread + with self.context_lock: + self.context_dict[key] = value + else: + self.context_dict[key] = value + + def __getattr__(self, key): + return self.kwargs[key] + + @property + def name_simple(self) -> str: + if not self._name_simple: + self._name_simple = self.__class__.__name__.replace("Worker", "") + return self._name_simple diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py new file mode 100644 index 00000000..87ba5dbe --- /dev/null +++ b/memory_scope/memory/worker/dummy_worker.py @@ -0,0 +1,6 @@ +from memory_base_worker import MemoryBaseWorker + + +class DummyWorker(MemoryBaseWorker): + def _run(self): + pass \ No newline at end of file diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py new file mode 100644 index 00000000..e8d5cc0c --- /dev/null +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -0,0 +1,70 @@ +from typing import List + +from chat.global_context import GLOBAL_CONTEXT +from constants.common_constants import MESSAGES, CHAT_NAME +from models.base_model import BaseModel +from scheme.message import Message +from storage.base_monitor import BaseMonitor +from storage.base_vector_store import BaseVectorStore +from worker.base_worker import BaseWorker + + +class MemoryBaseWorker(BaseWorker): + def __init__(self, + embedding_model: str, + generation_model: str, + rank_model: str, + **kwargs): + super(MemoryBaseWorker, self).__init__(**kwargs) + self.embedding_model_name: str = embedding_model + self.generation_model_name: str = generation_model + self.rank_model_name: str = rank_model + + self._embedding_model: BaseModel | None = None + self._generation_model: BaseModel | None = None + self._rank_model: BaseModel | None = None + + self._vector_store: BaseVectorStore | None = None + self._monitor: BaseMonitor | None = None + + @property + def messages(self) -> List[Message]: + return self.get_context(MESSAGES) + + @messages.setter + def messages(self, value): + self.set_context(MESSAGES, value) + + @property + def chat_name(self): + return self.get_context(CHAT_NAME) + + @property + def embedding_model(self): + if self._embedding_model is None: + self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name) + return self._embedding_model + + @property + def generation_model(self): + if self._generation_model is None: + self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name) + return self._generation_model + + @property + def rank_model(self): + if self._rank_model is None: + self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name) + return self._rank_model + + @property + def vector_store(self): + if self._vector_store is None: + self._vector_store = GLOBAL_CONTEXT.vector_store + return self._vector_store + + @property + def monitor(self): + if self._monitor is None: + self._monitor = GLOBAL_CONTEXT.monitor + return self._monitor diff --git a/memory_scope/memory/workflow/__init__.py b/memory_scope/memory/workflow/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/memory/workflow/backend_v1_workflow.py b/memory_scope/memory/workflow/backend_v1_workflow.py new file mode 100644 index 00000000..8e04bdcd --- /dev/null +++ b/memory_scope/memory/workflow/backend_v1_workflow.py @@ -0,0 +1,43 @@ +import time + +from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES +from memory_scope.memory.workflow.base_workflow import BaseWorkflow + + +class BackendV1Workflow(BaseWorkflow): + + def __init__(self, interval_time: int, min_count: int, **kwargs): + super().__init__(**kwargs) + self.interval_time: int = interval_time + self.min_count: int = min_count + + @property + def not_memorized_size(self): + return sum([not x.memorized for x in self.chat_messages]) + + def set_memorized(self): + for msg in self.chat_messages: + msg.memorized = True + + def _loop(self): + while self.loop_switch: + time.sleep(self.interval_time) + if self.not_memorized_size < self.min_count: + continue + + self.context[CHAT_MESSAGES] = self.chat_messages + self.__call__() + self.context.clear() + self.set_memorized() + + def start_loop_run(self): + if not self.loop_switch: + self.loop_switch = True + return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop) + + def run_workflow(self): + self.context[CHAT_MESSAGES] = self.chat_messages + self.__call__() + result = self.context.get(RESULT) + self.context.clear() + return result diff --git a/memory_scope/memory/workflow/base_workflow.py b/memory_scope/memory/workflow/base_workflow.py new file mode 100644 index 00000000..de79c569 --- /dev/null +++ b/memory_scope/memory/workflow/base_workflow.py @@ -0,0 +1,126 @@ +import re +import threading +from concurrent.futures import ThreadPoolExecutor, as_completed +from itertools import zip_longest +from typing import Dict, Any, List + +from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.memory.worker.base_worker import BaseWorker +from memory_scope.scheme.message import Message +from memory_scope.utils.logger import Logger +from memory_scope.utils.timer import Timer +from memory_scope.utils.tool_functions import init_instance_by_config + + +class BaseWorkflow(object): + + def __init__(self, + name: str, + workflow: str, + thread_pool: ThreadPoolExecutor, + chat_messages: List[Message], + max_history_message_count: int, + **kwargs): + + self.name: str = name + self.workflow: str = workflow + self.thread_pool: ThreadPoolExecutor = thread_pool + self.chat_messages: List[Message] = chat_messages + self.max_history_message_count: int = max_history_message_count + self.kwargs = kwargs + + self.workflow_worker_list: List[List[List[str]]] = [] + self.worker_dict: Dict[str, BaseWorker | bool] = {} + self.context: Dict[str, Any] = {} + self.context_lock = threading.Lock() + + self.logger: Logger = Logger.get_logger() + + if self.workflow: + self._parse_workflow() + self._print_workflow() + + def _parse_workflow(self): + # re-match e.g., [a|b],c,[d,e,f|g,h],j + pattern = r'(\[[^\]]*\]|[^,]+)' + workflow_split = re.findall(pattern, self.workflow) + for workflow_part in workflow_split: + # e.g., [d,e,f|g,h] + workflow_part = workflow_part.strip() + if '[' in workflow_part or ']' in workflow_part: + workflow_part = workflow_part.replace('[', '').replace(']', '') + + # e.g., ["d,e,f", "g,h"] + line_split = [x.strip() for x in workflow_part.split("|") if x] + if len(line_split) <= 0: + continue + + # is under multi thread cond + is_multi_thread: bool = len(line_split) > 1 + + # e.g., ["d","e","f"] + line_split_split: List[List[str]] = [] + for sub_line_split in line_split: + sub_split = [x.strip() for x in sub_line_split.split(",")] + line_split_split.append(sub_split) + # add workers + for sub_item in sub_split: + self.worker_dict[sub_item] = is_multi_thread + self.workflow_worker_list.append(line_split_split) + + def _print_workflow(self): + self.logger.info(f"----- print_workflow_{self.name}_begin -----") + i: int = 0 + for workflow_part in self.workflow_worker_list: + if len(workflow_part) == 1: + for w in workflow_part[0]: + self.logger.info(f"stage{i}: {w}") + i += 1 + else: + for w_zip in zip_longest(*workflow_part, fillvalue="-"): + self.logger.info(f"stage{i}: {' | '.join(w_zip)}") + i += 1 + for w in w_zip: + if w == "-": + continue + self.logger.info(f"----- print_workflow_{self.name}_end -----") + + def init_workers(self): + for name in list(self.worker_dict.keys()): + if name not in G_CONTEXT.worker_config: + raise RuntimeError(f"worker={name} is not exists in worker_config!") + + self.worker_dict[name] = init_instance_by_config( + config=G_CONTEXT.worker_config[name], + suffix_name="worker", + name=name, + is_multi_thread=self.worker_dict[name], + context=self.context, + context_lock=self.context_lock) + + def _run_sub_workflow(self, worker_list: List[str]) -> bool: + for name in worker_list: + worker = self.worker_dict[name] + worker.run() + if not worker.continue_run: + return False + return True + + def run_workflow(self): + with Timer(f"run_workflow_{self.name}"): + for workflow_part in self.workflow_worker_list: + if len(workflow_part) == 1: + if not self._run_sub_workflow(workflow_part[0]): + break + else: + t_list = [] + for sub_workflow in workflow_part: + t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow)) + + flag = True + for future in as_completed(t_list): + if not future.result(): + flag = False + break + if not flag: + break diff --git a/memory_scope/memory/workflow/frontend_workflow.py b/memory_scope/memory/workflow/frontend_workflow.py new file mode 100644 index 00000000..f394f6f9 --- /dev/null +++ b/memory_scope/memory/workflow/frontend_workflow.py @@ -0,0 +1,12 @@ +from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES +from memory_scope.memory.workflow.base_workflow import BaseWorkflow + + +class FrontendWorkflow(BaseWorkflow): + + def run_workflow(self): + self.context[CHAT_MESSAGES] = self.chat_messages[:1 + self.max_history_message_count] + self.__call__() + result = self.context.get(RESULT) + self.context.clear() + return result diff --git a/memory_scope/models/__init__.py b/memory_scope/models/__init__.py index a12698b1..5a9d5579 100644 --- a/memory_scope/models/__init__.py +++ b/memory_scope/models/__init__.py @@ -1,4 +1,3 @@ -from utils.registry import Registry +from memory_scope.utils.registry import Registry -# __all__ = ["LlamaIndexEmbeddingModel", "LlamaIndexGenerationModel", "LlamaIndexRerankModel"] MODEL_REGISTRY = Registry("models") diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 2bef8c6b..6bd2e6c4 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -3,7 +3,7 @@ import time from abc import abstractmethod, ABCMeta from enumeration.model_enum import ModelEnum -from . import MODEL_REGISTRY +from memory_scope.models import MODEL_REGISTRY from .response import ModelResponse, ModelResponseGen from utils.logger import Logger from utils.timer import Timer diff --git a/memory_scope/scheme/message.py b/memory_scope/scheme/message.py index cc44268c..ad180772 100644 --- a/memory_scope/scheme/message.py +++ b/memory_scope/scheme/message.py @@ -6,4 +6,6 @@ class Message(BaseModel): content: str = Field(..., description="The body of the message") - time_created: int = Field("", description="Timestamp when the message was created") + time_created: int = Field(..., description="Timestamp when the message was created") + + memorized: bool = Field(False, description="indicate whether message is memorized") From 83cdd63a844dec8d0c69cc2d2833c5dc95718286 Mon Sep 17 00:00:00 2001 From: "xianzhe.xxz" Date: Wed, 26 Jun 2024 13:34:15 +0800 Subject: [PATCH 03/41] rebase master --- memory_scope/models/base_model.py | 8 +- .../models/llama_index_embedding_model.py | 8 +- memory_scope/models/response.py | 2 +- memory_scope/storage/base_vector_store.py | 3 +- .../llama_index_elastic_search_store.py | 47 +++++++++-- tests/storages/test_storages_lli_es.py | 80 ++++++++++++++----- 6 files changed, 110 insertions(+), 38 deletions(-) diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 6bd2e6c4..28b12fe5 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -2,11 +2,11 @@ import inspect import time from abc import abstractmethod, ABCMeta -from enumeration.model_enum import ModelEnum +from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.models import MODEL_REGISTRY -from .response import ModelResponse, ModelResponseGen -from utils.logger import Logger -from utils.timer import Timer +from memory_scope.models.response import ModelResponse, ModelResponseGen +from memory_scope.utils.logger import Logger +from memory_scope.utils.timer import Timer class BaseModel(metaclass=ABCMeta): diff --git a/memory_scope/models/llama_index_embedding_model.py b/memory_scope/models/llama_index_embedding_model.py index 2e2689c9..a9397116 100644 --- a/memory_scope/models/llama_index_embedding_model.py +++ b/memory_scope/models/llama_index_embedding_model.py @@ -2,10 +2,10 @@ from typing import List from llama_index.embeddings.dashscope import DashScopeEmbedding -from models import MODEL_REGISTRY -from models.base_model import BaseModel -from models.response import ModelResponse, ModelResponseGen -from enumeration.model_enum import ModelEnum +from memory_scope.models import MODEL_REGISTRY +from memory_scope.models.base_model import BaseModel +from memory_scope.models.response import ModelResponse, ModelResponseGen +from memory_scope.enumeration.model_enum import ModelEnum class LlamaIndexEmbeddingModel(BaseModel): diff --git a/memory_scope/models/response.py b/memory_scope/models/response.py index a44e1fab..841024fd 100644 --- a/memory_scope/models/response.py +++ b/memory_scope/models/response.py @@ -3,7 +3,7 @@ from typing import Generator, List, Dict, Any from pydantic import BaseModel, Field -from enumeration.model_enum import ModelEnum +from memory_scope.enumeration.model_enum import ModelEnum class ModelResponse(BaseModel): diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index 9695084f..36bba692 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -1,7 +1,8 @@ from abc import ABCMeta, abstractmethod from typing import Dict, List -from models.base_model import BaseModel +from memory_scope.models.base_model import BaseModel +from memory_scope.scheme.memory_node import MemoryNode class BaseVectorStore(metaclass=ABCMeta): diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index 3e1f2cac..f1e6ecd9 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -10,6 +10,22 @@ from memory_scope.storage.base_vector_store import BaseVectorStore from memory_scope.scheme.memory_node import MemoryNode +class _ElasticsearchStore(ElasticsearchStore): + async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None: + """ + Async delete node from Elasticsearch index. + + Args: + ref_doc_id: ID of the node to delete. + delete_kwargs: Optional. Additional arguments to + pass to AsyncElasticsearch delete_by_query. + + Raises: + Exception: If AsyncElasticsearch delete_by_query fails. + """ + return await self._store.delete( + query={"term": {"_id": ref_doc_id}}, **delete_kwargs + ) def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]: @@ -65,7 +81,7 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): self.index_name: str = index_name self.embedding_model: BaseModel = embedding_model - self.es_store = ElasticsearchStore(index_name=self.index_name, + self.es_store = _ElasticsearchStore(index_name=self.index_name, retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True), **kwargs) @@ -89,8 +105,17 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): return results async def async_retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode: - raise NotImplementedError - ## return await super().async_retrieve(text, limit_size, filter_dict) + filter = _to_elasticsearch_filter(filter_dict) + retriever = self.index.as_retriever( + vector_store_kwargs={ + "es_filter": filter + }, + similarity_top_k=top_k + ) + textnodes = await retriever.aretrieve(query) + results = self._textnodes2memorynodes(textnodes) + + return results def insert(self, node: MemoryNode): node = self._memorynode2textnode(node) @@ -99,20 +124,26 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): def insert_batch(self, node: MemoryNode) -> None: raise NotImplementedError - def delete(self): - raise NotImplementedError - + def delete(self, node: MemoryNode) -> None: + memory_id = node.memory_id + self.es_store.delete(memory_id) + + def update(self, node: MemoryNode) -> None: + self.delete(node) + self.insert(node) + def flush(self): raise NotImplementedError def _memorynode2textnode(self, memory_node: MemoryNode) -> TextNode: content = memory_node.content + memory_id = memory_node.memory_id meta = memory_node.model_dump(exclude={"content"}) - return TextNode(text=content, metadata=meta) + return TextNode(id_=memory_id, text=content, metadata=meta) def _textnode2memorynode(self, text_node: TextNode) -> MemoryNode: content = text_node.text - meta = text_node.metadata + meta = text_node.metadata mem_node = MemoryNode(content=content, **meta) return mem_node diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 4ed04075..e73cd1ca 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -18,7 +18,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): emb = LlamaIndexEmbeddingModel(**config).model config = { - "index_name" : "0625_3", + "index_name" : "0626_1", "es_url" : "http://localhost:9200", "embedding_model" : emb, @@ -28,63 +28,103 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): MemoryNode( content="The lives of two mob hitmen, a boxer, a gangster and his wife, and a pair of diner bandits intertwine in four tales of violence and redemption.", memory_type="observation", - id="0" + user_id="0", + status="valid", + memory_id="aaa123", ), MemoryNode( content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, Batman must accept one of the greatest psychological and physical tests of his ability to fight injustice.", memory_type="observation", - id="1" + user_id="1", + status="valid", + memory_id="bbb456", ), MemoryNode( content="An insomniac office worker and a devil-may-care soapmaker form an underground fight club that evolves into something much, much more.", memory_type="insights", - id="2" + user_id="2", + status="valid", + memory_id="ccc789", + ), MemoryNode( content="A thief who steals corporate secrets through the use of dream-sharing technology is given the inverse task of planting an idea into thed of a C.E.O.", memory_type="insights", - id="3" - + user_id="3", + status="valid", + memory_id="ddd012", ), MemoryNode( content="A computer hacker learns from mysterious rebels about the true nature of his reality and his role in the war against its controllers.", memory_type="profile", - id="4" + user_id="4", + status="valid", + memory_id="eee345", ), MemoryNode( content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven deadly sins as his motives.", memory_type="profile", - id="5" + user_id="5", + status="valid", + memory_id="fff678" ), MemoryNode( content="An organized crime dynasty's aging patriarch transfers control of his clandestine empire to his reluctant son.", memory_type="insights", - id="6"), + user_id="6", + status="valid", + memory_id="ggg901", + ), MemoryNode( content="ggggggggg", memory_type="profile", - id="6"), + user_id="6", + status="valid", + memory_id="ggg234", + ), ] - # @unittest.skip("tmp") - def test_insert(self, ): - for node in self.data: - self.es_store.insert(node) - - # @unittest.skip("tmp") def test_retrieve(self, ): - filter = { - "id": ["1", "2", "3"], - "memory_type": "insights", + "user_id": "6", } + for node in self.data: + self.es_store.insert(node) + self.es_store.insert(MemoryNode( + content="xxxxxx", + memory_type="profile", + user_id="6", + status="valid", + memory_id="ggg567" + )) res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) print(len(res)) print(res) - \ No newline at end of file + self.es_store.update(MemoryNode( + content="test update", + memory_type="profile", + user_id="6", + status="invalid", + memory_id="ggg567" + )) + res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + print(len(res)) + print(res) + + + self.es_store.delete(MemoryNode( + content="test update", + memory_type="profile", + user_id="6", + status="invalid", + memory_id="ggg567" + )) + res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + print(len(res)) + print(res) From d815ba988ab0132b16c4d9c2efd39e3a3394980b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 14:01:48 +0800 Subject: [PATCH 04/41] [dev] add operation --- .../{workflow => operation}/__init__.py | 0 .../memory/operation/base_operation.py | 15 ++++ .../{workflow => operation}/base_workflow.py | 8 +-- .../memory/operation/read_operation.py | 23 ++++++ .../memory/operation/summary_operation.py | 37 ++++++++++ .../memory/operation/write_operation.py | 66 +++++++++++++++++ memory_scope/memory/worker/base_worker.py | 66 +++++++---------- memory_scope/memory/worker/dummy_worker.py | 6 +- .../memory/worker/memory_base_worker.py | 70 ------------------- .../memory/workflow/backend_v1_workflow.py | 43 ------------ .../memory/workflow/frontend_workflow.py | 12 ---- 11 files changed, 171 insertions(+), 175 deletions(-) rename memory_scope/memory/{workflow => operation}/__init__.py (100%) create mode 100644 memory_scope/memory/operation/base_operation.py rename memory_scope/memory/{workflow => operation}/base_workflow.py (93%) create mode 100644 memory_scope/memory/operation/read_operation.py create mode 100644 memory_scope/memory/operation/summary_operation.py create mode 100644 memory_scope/memory/operation/write_operation.py delete mode 100644 memory_scope/memory/worker/memory_base_worker.py delete mode 100644 memory_scope/memory/workflow/backend_v1_workflow.py delete mode 100644 memory_scope/memory/workflow/frontend_workflow.py diff --git a/memory_scope/memory/workflow/__init__.py b/memory_scope/memory/operation/__init__.py similarity index 100% rename from memory_scope/memory/workflow/__init__.py rename to memory_scope/memory/operation/__init__.py diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py new file mode 100644 index 00000000..26a9b9ca --- /dev/null +++ b/memory_scope/memory/operation/base_operation.py @@ -0,0 +1,15 @@ +from abc import ABCMeta, abstractmethod +from typing import Literal + +OPERATION_TYPE = Literal["frontend", "backend"] + + +class BaseOperation(metaclass=ABCMeta): + operation_type: OPERATION_TYPE = "frontend" + + @abstractmethod + def run_operation(self): + raise NotImplementedError + + def run_operation_backend(self): + pass diff --git a/memory_scope/memory/workflow/base_workflow.py b/memory_scope/memory/operation/base_workflow.py similarity index 93% rename from memory_scope/memory/workflow/base_workflow.py rename to memory_scope/memory/operation/base_workflow.py index de79c569..8b5ac67c 100644 --- a/memory_scope/memory/workflow/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -6,7 +6,6 @@ from typing import Dict, Any, List from memory_scope.chat_v2.global_context import G_CONTEXT from memory_scope.memory.worker.base_worker import BaseWorker -from memory_scope.scheme.message import Message from memory_scope.utils.logger import Logger from memory_scope.utils.timer import Timer from memory_scope.utils.tool_functions import init_instance_by_config @@ -18,15 +17,11 @@ class BaseWorkflow(object): name: str, workflow: str, thread_pool: ThreadPoolExecutor, - chat_messages: List[Message], - max_history_message_count: int, **kwargs): self.name: str = name self.workflow: str = workflow self.thread_pool: ThreadPoolExecutor = thread_pool - self.chat_messages: List[Message] = chat_messages - self.max_history_message_count: int = max_history_message_count self.kwargs = kwargs self.workflow_worker_list: List[List[List[str]]] = [] @@ -115,7 +110,8 @@ class BaseWorkflow(object): else: t_list = [] for sub_workflow in workflow_part: - t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow)) + t_list.append(G_CONTEXT.thread_pool.submit( + self._run_sub_workflow, sub_workflow)) flag = True for future in as_completed(t_list): diff --git a/memory_scope/memory/operation/read_operation.py b/memory_scope/memory/operation/read_operation.py new file mode 100644 index 00000000..9f8f294f --- /dev/null +++ b/memory_scope/memory/operation/read_operation.py @@ -0,0 +1,23 @@ +from typing import List + +from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES +from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE +from memory_scope.memory.operation.base_workflow import BaseWorkflow +from memory_scope.scheme.message import Message + + +class ReadOperation(BaseWorkflow, BaseOperation): + operation_type: OPERATION_TYPE = "frontend" + + def __init__(self, chat_messages: List[Message], max_his_msg_count: int = 0, **kwargs): + super().__init__(**kwargs) + self.chat_messages: List[Message] = chat_messages + self.max_his_msg_count: int = max_his_msg_count + + def run_operation(self): + max_count = 1 + self.max_his_msg_count + self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] + self.run_workflow() + result = self.context.get(RESULT) + self.context.clear() + return result diff --git a/memory_scope/memory/operation/summary_operation.py b/memory_scope/memory/operation/summary_operation.py new file mode 100644 index 00000000..68d98c41 --- /dev/null +++ b/memory_scope/memory/operation/summary_operation.py @@ -0,0 +1,37 @@ +import time + +from memory_scope.memory.base_workflow import BaseWorkflow + +from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE + + +class SummaryOperation(BaseWorkflow, BaseOperation): + operation_type: OPERATION_TYPE = "backend" + + def __init__(self, interval_time: int = 300, **kwargs): + super().__init__(**kwargs) + + self.interval_time: int = interval_time + + self._operation_status_run: bool = False + self._loop_switch: bool = False + + def run_operation(self): + if self._operation_status_run: + return + + self._operation_status_run = True + self.run_workflow() + self.context.clear() + self._operation_status_run = False + + def _loop_operation(self): + while self._loop_switch: + time.sleep(self.interval_time) + self.run_operation() + + def run_operation_backend(self): + if not self._loop_switch: + self._loop_switch = True + return G_CONTEXT.thread_pool.submit(self._loop_operation) diff --git a/memory_scope/memory/operation/write_operation.py b/memory_scope/memory/operation/write_operation.py new file mode 100644 index 00000000..efd5d4d1 --- /dev/null +++ b/memory_scope/memory/operation/write_operation.py @@ -0,0 +1,66 @@ +import time +from typing import List + +from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.constants.common_constants import CHAT_MESSAGES +from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE +from memory_scope.memory.operation.base_workflow import BaseWorkflow +from memory_scope.scheme.message import Message + + +class WriteOperation(BaseOperation, BaseWorkflow): + operation_type: OPERATION_TYPE = "backend" + + def __init__(self, + chat_messages: List[Message], + max_his_msg_count: int = 0, + message_lock=None, + interval_time: int = 60, + min_count: int = 5, + **kwargs): + super().__init__(**kwargs) + + self.chat_messages: List[Message] = chat_messages + self.max_his_msg_count: int = max_his_msg_count + self.message_lock = message_lock + self.interval_time: int = interval_time + self.min_count: int = min_count + + self._operation_status_run: bool = False + self._loop_switch: bool = False + + @property + def not_memorized_size(self): + return sum([not x.memorized for x in self.chat_messages]) + + def set_memorized(self): + if self.message_lock: + with self.message_lock: + for msg in self.chat_messages: + msg.memorized = True + + def run_operation(self): + if self._operation_status_run: + return + + self._operation_status_run = True + not_memorized_size = self.not_memorized_size + if not_memorized_size < self.min_count: + return + + max_count = not_memorized_size + self.max_his_msg_count + self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] + self.run_workflow() + self.context.clear() + self.set_memorized() + self._operation_status_run = False + + def _loop_operation(self): + while self._loop_switch: + time.sleep(self.interval_time) + self.run_operation() + + def run_operation_backend(self): + if not self._loop_switch: + self._loop_switch = True + return G_CONTEXT.thread_pool.submit(self._loop_operation) diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index a47fd1ed..78bfae5e 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -1,73 +1,57 @@ +from abc import ABCMeta, abstractmethod from typing import Any, Dict -from utils.logger import Logger -from utils.timer import Timer +from memory_scope.utils.logger import Logger +from memory_scope.utils.timer import Timer -class BaseWorker(object): +class BaseWorker(metaclass=ABCMeta): - def __init__(self, raise_exception: bool = True, **kwargs): - super(BaseWorker, self).__init__(**kwargs) - # 异常是否继续执行 + def __init__(self, + name: str, + context: Dict[str, Any], + context_lock=None, + raise_exception: bool = True, + is_multi_thread: bool = False, + **kwargs): + + self.name: str = name + self.context: Dict[str, Any] = context + self.context_lock = context_lock self.raise_exception: bool = raise_exception - - # True 为正常运行,False会结束整个pipeline - self.continue_run: bool = True - - # 短name - self._name_simple: str = "" - - # 是否多线程环境 - self.is_multi_thread: bool = False - - # pipeline 上下文 - self.context_dict: Dict[str, Any] | None = None - self.context_lock = None - - # 日志 - self.logger: Logger = Logger.get_logger() - - # worker 参数 + self.is_multi_thread: bool = is_multi_thread self.kwargs: dict = kwargs + self.continue_run: bool = True + self.logger: Logger = Logger.get_logger() + + @abstractmethod def _run(self): raise NotImplementedError def run(self): - self.logger.info(f"----- Begin {self.name_simple} -----") - with Timer(self.name_simple, log_time=False) as t: + self.logger.info(f"----- worker_{self.name}_begin -----") + with Timer(self.name, log_time=False) as t: if self.raise_exception: self._run() else: try: self._run() except Exception as e: - self.logger.exception(f"run {self.name_simple} failed! args={e.args}") + self.logger.exception(f"run {self.name} failed! args={e.args}") - self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----") - - def set_context_dict(self, context_dict: dict, context_lock=None): - self.context_dict = context_dict - if context_lock is not None: - self.context_lock = context_lock - self.is_multi_thread = True + self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----") def get_context(self, key: str, default=None): return self.context_dict.get(key, default) def set_context(self, key: str, value: Any): if self.is_multi_thread: - # add lock to multi thread with self.context_lock: self.context_dict[key] = value else: self.context_dict[key] = value def __getattr__(self, key): + # raise exception if not exists return self.kwargs[key] - - @property - def name_simple(self) -> str: - if not self._name_simple: - self._name_simple = self.__class__.__name__.replace("Worker", "") - return self._name_simple diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 87ba5dbe..6b446b0c 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -1,6 +1,6 @@ -from memory_base_worker import MemoryBaseWorker +from memory_scope.memory.worker.base_worker import BaseWorker -class DummyWorker(MemoryBaseWorker): +class DummyWorker(BaseWorker): def _run(self): - pass \ No newline at end of file + self.logger.info("enter dummy worker!") diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py deleted file mode 100644 index e8d5cc0c..00000000 --- a/memory_scope/memory/worker/memory_base_worker.py +++ /dev/null @@ -1,70 +0,0 @@ -from typing import List - -from chat.global_context import GLOBAL_CONTEXT -from constants.common_constants import MESSAGES, CHAT_NAME -from models.base_model import BaseModel -from scheme.message import Message -from storage.base_monitor import BaseMonitor -from storage.base_vector_store import BaseVectorStore -from worker.base_worker import BaseWorker - - -class MemoryBaseWorker(BaseWorker): - def __init__(self, - embedding_model: str, - generation_model: str, - rank_model: str, - **kwargs): - super(MemoryBaseWorker, self).__init__(**kwargs) - self.embedding_model_name: str = embedding_model - self.generation_model_name: str = generation_model - self.rank_model_name: str = rank_model - - self._embedding_model: BaseModel | None = None - self._generation_model: BaseModel | None = None - self._rank_model: BaseModel | None = None - - self._vector_store: BaseVectorStore | None = None - self._monitor: BaseMonitor | None = None - - @property - def messages(self) -> List[Message]: - return self.get_context(MESSAGES) - - @messages.setter - def messages(self, value): - self.set_context(MESSAGES, value) - - @property - def chat_name(self): - return self.get_context(CHAT_NAME) - - @property - def embedding_model(self): - if self._embedding_model is None: - self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name) - return self._embedding_model - - @property - def generation_model(self): - if self._generation_model is None: - self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name) - return self._generation_model - - @property - def rank_model(self): - if self._rank_model is None: - self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name) - return self._rank_model - - @property - def vector_store(self): - if self._vector_store is None: - self._vector_store = GLOBAL_CONTEXT.vector_store - return self._vector_store - - @property - def monitor(self): - if self._monitor is None: - self._monitor = GLOBAL_CONTEXT.monitor - return self._monitor diff --git a/memory_scope/memory/workflow/backend_v1_workflow.py b/memory_scope/memory/workflow/backend_v1_workflow.py deleted file mode 100644 index 8e04bdcd..00000000 --- a/memory_scope/memory/workflow/backend_v1_workflow.py +++ /dev/null @@ -1,43 +0,0 @@ -import time - -from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES -from memory_scope.memory.workflow.base_workflow import BaseWorkflow - - -class BackendV1Workflow(BaseWorkflow): - - def __init__(self, interval_time: int, min_count: int, **kwargs): - super().__init__(**kwargs) - self.interval_time: int = interval_time - self.min_count: int = min_count - - @property - def not_memorized_size(self): - return sum([not x.memorized for x in self.chat_messages]) - - def set_memorized(self): - for msg in self.chat_messages: - msg.memorized = True - - def _loop(self): - while self.loop_switch: - time.sleep(self.interval_time) - if self.not_memorized_size < self.min_count: - continue - - self.context[CHAT_MESSAGES] = self.chat_messages - self.__call__() - self.context.clear() - self.set_memorized() - - def start_loop_run(self): - if not self.loop_switch: - self.loop_switch = True - return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop) - - def run_workflow(self): - self.context[CHAT_MESSAGES] = self.chat_messages - self.__call__() - result = self.context.get(RESULT) - self.context.clear() - return result diff --git a/memory_scope/memory/workflow/frontend_workflow.py b/memory_scope/memory/workflow/frontend_workflow.py deleted file mode 100644 index f394f6f9..00000000 --- a/memory_scope/memory/workflow/frontend_workflow.py +++ /dev/null @@ -1,12 +0,0 @@ -from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES -from memory_scope.memory.workflow.base_workflow import BaseWorkflow - - -class FrontendWorkflow(BaseWorkflow): - - def run_workflow(self): - self.context[CHAT_MESSAGES] = self.chat_messages[:1 + self.max_history_message_count] - self.__call__() - result = self.context.get(RESULT) - self.context.clear() - return result From 08e15f284890cc04a5fed5fd4744b306ad7d361d Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 15:47:35 +0800 Subject: [PATCH 05/41] [dev] add memory service --- config/config.yaml | 22 +++---- memory_scope/memory/base_memory_service.py | 6 -- .../memory/operation/base_operation.py | 3 + .../memory/operation/read_operation.py | 9 ++- .../memory/operation/summary_operation.py | 3 + .../memory/operation/write_operation.py | 15 +++-- memory_scope/memory/service/__init__.py | 0 .../memory/service/base_memory_service.py | 13 ++++ .../memory/service/chat_memory_service.py | 59 +++++++++++++++++++ memory_scope/memory/worker/dummy_worker.py | 4 +- 10 files changed, 105 insertions(+), 29 deletions(-) delete mode 100644 memory_scope/memory/base_memory_service.py create mode 100644 memory_scope/memory/service/__init__.py create mode 100644 memory_scope/memory/service/base_memory_service.py create mode 100644 memory_scope/memory/service/chat_memory_service.py diff --git a/config/config.yaml b/config/config.yaml index c0279fc7..7b99b1a3 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -11,26 +11,22 @@ memory_chat: memory_service: memory_chat_service: class: memory.base_memory_service - history_msg_count: 5 + history_msg_count: 10 memory_operations: read_memory: - class: memory.workflow.base_workflow - workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank - work_type: frontend + class: memory.operation.read_operation + workflow: dummy list_memory: - class: memory.workflow.base_workflow + class: memory.operation.read_operation workflow: dummy - work_type: frontend - extract_memory: - class: memory.workflow.base_workflow + write_memory: + class: memory.operation.write_operation workflow: dummy - work_type: backend interval_time: 60 - min_count: 5 - reflect_memory: - class: memory.workflow.base_workflow + contextual_msg_count: 6 + summary_memory: + class: memory.operation.summary_operation workflow: dummy - work_type: backend interval_time: 300 models: dashscope_generation: diff --git a/memory_scope/memory/base_memory_service.py b/memory_scope/memory/base_memory_service.py deleted file mode 100644 index d05652a3..00000000 --- a/memory_scope/memory/base_memory_service.py +++ /dev/null @@ -1,6 +0,0 @@ -from abc import ABCMeta - - -class BaseMemoryService(metaclass=ABCMeta): - def __init__(self, **kwargs): - self.kwargs = kwargs diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 26a9b9ca..70d1d103 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -7,6 +7,9 @@ OPERATION_TYPE = Literal["frontend", "backend"] class BaseOperation(metaclass=ABCMeta): operation_type: OPERATION_TYPE = "frontend" + def init_workflow(self): + pass + @abstractmethod def run_operation(self): raise NotImplementedError diff --git a/memory_scope/memory/operation/read_operation.py b/memory_scope/memory/operation/read_operation.py index 9f8f294f..485a8f8c 100644 --- a/memory_scope/memory/operation/read_operation.py +++ b/memory_scope/memory/operation/read_operation.py @@ -9,13 +9,16 @@ from memory_scope.scheme.message import Message class ReadOperation(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "frontend" - def __init__(self, chat_messages: List[Message], max_his_msg_count: int = 0, **kwargs): + def __init__(self, chat_messages: List[Message], his_msg_count: int = 0, **kwargs): super().__init__(**kwargs) self.chat_messages: List[Message] = chat_messages - self.max_his_msg_count: int = max_his_msg_count + self.his_msg_count: int = his_msg_count + + def init_workflow(self): + self.init_workers() def run_operation(self): - max_count = 1 + self.max_his_msg_count + max_count = 1 + self.his_msg_count self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] self.run_workflow() result = self.context.get(RESULT) diff --git a/memory_scope/memory/operation/summary_operation.py b/memory_scope/memory/operation/summary_operation.py index 68d98c41..1f36b679 100644 --- a/memory_scope/memory/operation/summary_operation.py +++ b/memory_scope/memory/operation/summary_operation.py @@ -17,6 +17,9 @@ class SummaryOperation(BaseWorkflow, BaseOperation): self._operation_status_run: bool = False self._loop_switch: bool = False + def init_workflow(self): + self.init_workers() + def run_operation(self): if self._operation_status_run: return diff --git a/memory_scope/memory/operation/write_operation.py b/memory_scope/memory/operation/write_operation.py index efd5d4d1..7523c969 100644 --- a/memory_scope/memory/operation/write_operation.py +++ b/memory_scope/memory/operation/write_operation.py @@ -13,18 +13,18 @@ class WriteOperation(BaseOperation, BaseWorkflow): def __init__(self, chat_messages: List[Message], - max_his_msg_count: int = 0, + his_msg_count: int = 0, message_lock=None, interval_time: int = 60, - min_count: int = 5, + contextual_msg_count: int = 6, **kwargs): super().__init__(**kwargs) self.chat_messages: List[Message] = chat_messages - self.max_his_msg_count: int = max_his_msg_count + self.his_msg_count: int = his_msg_count self.message_lock = message_lock self.interval_time: int = interval_time - self.min_count: int = min_count + self.contextual_msg_count: int = contextual_msg_count self._operation_status_run: bool = False self._loop_switch: bool = False @@ -39,16 +39,19 @@ class WriteOperation(BaseOperation, BaseWorkflow): for msg in self.chat_messages: msg.memorized = True + def init_workflow(self): + self.init_workers() + def run_operation(self): if self._operation_status_run: return self._operation_status_run = True not_memorized_size = self.not_memorized_size - if not_memorized_size < self.min_count: + if not_memorized_size < self.contextual_msg_count: return - max_count = not_memorized_size + self.max_his_msg_count + max_count = not_memorized_size + self.his_msg_count self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] self.run_workflow() self.context.clear() diff --git a/memory_scope/memory/service/__init__.py b/memory_scope/memory/service/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py new file mode 100644 index 00000000..b74fa4e6 --- /dev/null +++ b/memory_scope/memory/service/base_memory_service.py @@ -0,0 +1,13 @@ +from abc import ABCMeta, abstractmethod + +from memory_scope.utils.logger import Logger + + +class BaseMemoryService(metaclass=ABCMeta): + def __init__(self, **kwargs): + self.logger = Logger.get_logger() + self.kwargs = kwargs + + @abstractmethod + def get_short_memory(self): + pass diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py new file mode 100644 index 00000000..5fde8502 --- /dev/null +++ b/memory_scope/memory/service/chat_memory_service.py @@ -0,0 +1,59 @@ +import threading +from typing import List, Dict + +from memory_scope.memory.operation.base_operation import BaseOperation +from memory_scope.memory.service.base_memory_service import BaseMemoryService +from memory_scope.scheme.message import Message +from memory_scope.utils.tool_functions import init_instance_by_config + + +class ChatMemoryService(BaseMemoryService): + + def __init__(self, + memory_operations: Dict[str, dict], + history_msg_count: int = 32, + contextual_msg_count: int = 6, + **kwargs): + super().__init__(**kwargs) + + self.op_dict: Dict[str, BaseOperation] = self._init_operation(memory_operations) + self.history_msg_count: int = history_msg_count + self.contextual_msg_count: int = contextual_msg_count + + self.chat_messages: List[Message] = [] + self.message_lock = threading.Lock + + def _init_operation(self, memory_operations: Dict[str, dict]): + op_dict: Dict[str, BaseOperation] = {} + for name, operation_config in memory_operations.items(): + if name in self.op_dict: + self.logger.warning(f"memory operation={name} is repeated!") + continue + self.op_dict[name] = init_instance_by_config(config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + contextual_msg_count=self.contextual_msg_count) + return op_dict + + def submit_message(self, messages: List[Message]): + messages = sorted(messages, key=lambda x: x.time_created) + self.chat_messages.extend(messages) + if len(self.chat_messages) > self.history_msg_count: + gap_size = len(self.chat_messages) - self.history_msg_count + for _ in range(gap_size): + self.chat_messages.pop(0) + + def prepare_service(self): + for _, operation in self.op_dict.items(): + operation.init_workflow() + if operation.operation_type == "backend": + operation.run_operation_backend() + + def do_operation(self, op_name: str): + if op_name not in self.op_dict: + self.logger.warning(f"op_name={op_name} is not inited!") + return + + operation = self.op_dict[op_name] + return operation.run_operation() diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 6b446b0c..9513bc79 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -1,6 +1,8 @@ +from memory_scope.constants.common_constants import RESULT from memory_scope.memory.worker.base_worker import BaseWorker class DummyWorker(BaseWorker): def _run(self): - self.logger.info("enter dummy worker!") + self.set_context(RESULT, ["test 123"]) + self.logger.info("enter dummy worker!") \ No newline at end of file From 375f7857c65323bdac73bc3dc0ad2c7cc1bc58d6 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 16:25:33 +0800 Subject: [PATCH 06/41] [dev] rename operation --- config/config.yaml | 14 +++++++++----- .../{read_operation.py => read_memory.py} | 7 ++++--- .../{summary_operation.py => summary_memory.py} | 5 ++--- .../{write_operation.py => write_memory.py} | 2 +- memory_scope/memory/service/base_memory_service.py | 2 +- memory_scope/memory/service/chat_memory_service.py | 10 ++++++---- 6 files changed, 23 insertions(+), 17 deletions(-) rename memory_scope/memory/operation/{read_operation.py => read_memory.py} (78%) rename memory_scope/memory/operation/{summary_operation.py => summary_memory.py} (89%) rename memory_scope/memory/operation/{write_operation.py => write_memory.py} (97%) diff --git a/config/config.yaml b/config/config.yaml index 7b99b1a3..81a51df6 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -11,21 +11,25 @@ memory_chat: memory_service: memory_chat_service: class: memory.base_memory_service - history_msg_count: 10 + history_msg_count: 32 + contextual_msg_count: 6 memory_operations: + read_user_message: + class: memory.operation.read_memory + workflow: dummy read_memory: - class: memory.operation.read_operation + class: memory.operation.read_memory workflow: dummy list_memory: - class: memory.operation.read_operation + class: memory.operation.read_memory workflow: dummy write_memory: - class: memory.operation.write_operation + class: memory.operation.write_memory workflow: dummy interval_time: 60 contextual_msg_count: 6 summary_memory: - class: memory.operation.summary_operation + class: memory.operation.summary_memory workflow: dummy interval_time: 300 models: diff --git a/memory_scope/memory/operation/read_operation.py b/memory_scope/memory/operation/read_memory.py similarity index 78% rename from memory_scope/memory/operation/read_operation.py rename to memory_scope/memory/operation/read_memory.py index 485a8f8c..8ffe24d2 100644 --- a/memory_scope/memory/operation/read_operation.py +++ b/memory_scope/memory/operation/read_memory.py @@ -6,19 +6,20 @@ from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message -class ReadOperation(BaseWorkflow, BaseOperation): +class ReadMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "frontend" - def __init__(self, chat_messages: List[Message], his_msg_count: int = 0, **kwargs): + def __init__(self, chat_messages: List[Message], his_msg_count: int = 0, contextual_msg_count: int = 0, **kwargs): super().__init__(**kwargs) self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count + self.contextual_msg_count: int = contextual_msg_count def init_workflow(self): self.init_workers() def run_operation(self): - max_count = 1 + self.his_msg_count + max_count = 1 + max(self.his_msg_count, self.contextual_msg_count) self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] self.run_workflow() result = self.context.get(RESULT) diff --git a/memory_scope/memory/operation/summary_operation.py b/memory_scope/memory/operation/summary_memory.py similarity index 89% rename from memory_scope/memory/operation/summary_operation.py rename to memory_scope/memory/operation/summary_memory.py index 1f36b679..3deba4ba 100644 --- a/memory_scope/memory/operation/summary_operation.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -1,12 +1,11 @@ import time -from memory_scope.memory.base_workflow import BaseWorkflow - from memory_scope.chat_v2.global_context import G_CONTEXT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE +from memory_scope.memory.operation.base_workflow import BaseWorkflow -class SummaryOperation(BaseWorkflow, BaseOperation): +class SummaryMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "backend" def __init__(self, interval_time: int = 300, **kwargs): diff --git a/memory_scope/memory/operation/write_operation.py b/memory_scope/memory/operation/write_memory.py similarity index 97% rename from memory_scope/memory/operation/write_operation.py rename to memory_scope/memory/operation/write_memory.py index 7523c969..9d6dce17 100644 --- a/memory_scope/memory/operation/write_operation.py +++ b/memory_scope/memory/operation/write_memory.py @@ -8,7 +8,7 @@ from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message -class WriteOperation(BaseOperation, BaseWorkflow): +class WriteMemory(BaseOperation, BaseWorkflow): operation_type: OPERATION_TYPE = "backend" def __init__(self, diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index b74fa4e6..79299265 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -9,5 +9,5 @@ class BaseMemoryService(metaclass=ABCMeta): self.kwargs = kwargs @abstractmethod - def get_short_memory(self): + def do_operation(self, op_name: str): pass diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index 5fde8502..f822c0dd 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -19,6 +19,7 @@ class ChatMemoryService(BaseMemoryService): self.op_dict: Dict[str, BaseOperation] = self._init_operation(memory_operations) self.history_msg_count: int = history_msg_count self.contextual_msg_count: int = contextual_msg_count + assert self.history_msg_count >= self.contextual_msg_count self.chat_messages: List[Message] = [] self.message_lock = threading.Lock @@ -36,7 +37,10 @@ class ChatMemoryService(BaseMemoryService): contextual_msg_count=self.contextual_msg_count) return op_dict - def submit_message(self, messages: List[Message]): + def submit_messages(self, messages: List[Message] | Message): + if isinstance(messages, Message): + messages = [messages] + messages = sorted(messages, key=lambda x: x.time_created) self.chat_messages.extend(messages) if len(self.chat_messages) > self.history_msg_count: @@ -54,6 +58,4 @@ class ChatMemoryService(BaseMemoryService): if op_name not in self.op_dict: self.logger.warning(f"op_name={op_name} is not inited!") return - - operation = self.op_dict[op_name] - return operation.run_operation() + return self.op_dict[op_name].run_operation() From 3959360fa03e5f09345aa80f1bb75d643bffc21e Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 16:29:15 +0800 Subject: [PATCH 07/41] [dev] rename worker dynamic path --- config/config.yaml | 2 +- memory_scope/chat_v2/global_context.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/config/config.yaml b/config/config.yaml index 81a51df6..bad2cac2 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -10,7 +10,7 @@ memory_chat: generation_model: dashscope_generation memory_service: memory_chat_service: - class: memory.base_memory_service + class: memory.service.base_memory_service history_msg_count: 32 contextual_msg_count: 6 memory_operations: diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py index 969c902c..abfbf2fa 100644 --- a/memory_scope/chat_v2/global_context.py +++ b/memory_scope/chat_v2/global_context.py @@ -5,7 +5,7 @@ import pydantic from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat from memory_scope.enumeration.language_enum import LanguageEnum -from memory_scope.memory.base_memory_service import BaseMemoryService +from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_vector_store import BaseVectorStore From 900ba7b9eb256cba3c3e90b4eb45e60a7f7a000b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 17:53:38 +0800 Subject: [PATCH 08/41] [dev] add operation description --- config/config.yaml | 7 +- memory_scope/chat/memory_chat.py | 12 +- memory_scope/chat_v2/base_memory_chat.py | 3 - memory_scope/chat_v2/base_memory_service.py | 4 - memory_scope/chat_v2/cli_memory_chat.py | 167 +++++++++++------- memory_scope/chat_v2/memory_chat.py | 67 ------- memory_scope/chat_v2/memory_service.py | 70 -------- .../memory/operation/base_operation.py | 4 + memory_scope/memory/operation/read_memory.py | 11 +- .../memory/operation/summary_memory.py | 12 +- memory_scope/memory/operation/write_memory.py | 10 +- .../memory/service/base_memory_service.py | 12 ++ .../memory/service/chat_memory_service.py | 3 + 13 files changed, 157 insertions(+), 225 deletions(-) delete mode 100644 memory_scope/chat_v2/base_memory_service.py delete mode 100644 memory_scope/chat_v2/memory_chat.py delete mode 100644 memory_scope/chat_v2/memory_service.py diff --git a/config/config.yaml b/config/config.yaml index bad2cac2..9c3c9906 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -5,7 +5,7 @@ global_config: open_ai_apikey: memory_chat: cli_memory_chat: - class: chat.cli_memory_chat + class: chat_v2.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation memory_service: @@ -17,20 +17,25 @@ memory_service: read_user_message: class: memory.operation.read_memory workflow: dummy + description: "read session messages of the user" read_memory: class: memory.operation.read_memory workflow: dummy + description: "read related memories of the user" list_memory: class: memory.operation.read_memory workflow: dummy + description: "read all memories of the user" write_memory: class: memory.operation.write_memory workflow: dummy + description: "write observation memories of the user" interval_time: 60 contextual_msg_count: 6 summary_memory: class: memory.operation.summary_memory workflow: dummy + description: "summary observation memories of the user" interval_time: 300 models: dashscope_generation: diff --git a/memory_scope/chat/memory_chat.py b/memory_scope/chat/memory_chat.py index 859758de..73f63fb6 100644 --- a/memory_scope/chat/memory_chat.py +++ b/memory_scope/chat/memory_chat.py @@ -29,17 +29,7 @@ class MemoryChat(BaseMemoryChat): ] return self._generation_model - @staticmethod - def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language] - if related_memories: - memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language] - system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories) - return Message( - role=MessageRoleEnum.SYSTEM, - content=system_prompt.strip(), - time_created=time_created, - ) + def chat_with_memory(self, query: str): query = query.strip() diff --git a/memory_scope/chat_v2/base_memory_chat.py b/memory_scope/chat_v2/base_memory_chat.py index b5bda713..f647cb98 100644 --- a/memory_scope/chat_v2/base_memory_chat.py +++ b/memory_scope/chat_v2/base_memory_chat.py @@ -2,9 +2,6 @@ from abc import ABCMeta, abstractmethod class BaseMemoryChat(metaclass=ABCMeta): - def __init__(self, memory_service: str, **kwargs): - self.kwargs = kwargs - @abstractmethod def chat_with_memory(self, query: str): diff --git a/memory_scope/chat_v2/base_memory_service.py b/memory_scope/chat_v2/base_memory_service.py deleted file mode 100644 index 9cf3fd76..00000000 --- a/memory_scope/chat_v2/base_memory_service.py +++ /dev/null @@ -1,4 +0,0 @@ -class BaseMemoryService(object): - def __init__(self, **kwargs): - - self.kwargs = kwargs diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index 1444edcd..f9e60d8a 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -1,83 +1,124 @@ import datetime +import time +from typing import Dict, List import questionary from rich.console import Console -from .memory_chat import MemoryChat -from enumeration.message_role_enum import MessageRoleEnum -from scheme.message import Message +from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat +from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.memory.service.base_memory_service import BaseMemoryService +from memory_scope.models.base_model import BaseModel +from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT +from memory_scope.scheme.message import Message -class CliMemoryChat(MemoryChat): - +class CliMemoryChat(BaseMemoryChat): USER_COMMANDS = { - "/exit": "exit the CLI", - "/memory": "print the current contents of agent memory", - "/retrieve": "retrieve related memory", - "/log": "log chat progress", - # TODO add more commands + "exit": "exit the CLI", + "help": "get cli commands help", } - def chat_with_memory(self, query): # for testing + def __init__(self, memory_service: str, generation_model: str, **kwargs): + self._memory_service: BaseMemoryService | str = memory_service + self._generation_model: BaseModel | str = generation_model + self.kwargs: dict = kwargs + + @property + def memory_service(self) -> BaseMemoryService: + if isinstance(self._memory_service, str): + self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] + self._memory_service.prepare_service() + return self._memory_service + + @property + def generation_model(self) -> BaseModel: + if isinstance(self._generation_model, str): + self._generation_model = G_CONTEXT.model_dict[self._generation_model] + return self._generation_model + + @staticmethod + def get_system_prompt(related_memories: List[str], time_created: int) -> Message: + system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] + if related_memories: + memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] + system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt] + related_memories]) + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) + + def chat_with_memory(self, query: str): query = query.strip() if not query: return time_created = int(datetime.datetime.now().timestamp()) - message = Message( - role=MessageRoleEnum.USER, content=query, time_created=time_created - ) - messages = [message] - return self.generation_model.call(messages=messages, stream=True) + new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) + related_memories: List[str] = self.memory_service.do_operation("read_memory") + system_message: Message = self.get_system_prompt(related_memories, time_created) + return self.generation_model.call(messages=[system_message, new_message], stream=True) - def retrieve_all(self): # for testing - return "memory 1. 2. 3." - def run(self): - console = Console() +def run(self): + op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() + self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) + + console = Console() + while True: + query = questionary.text( + "Please enter your message or command:", + multiline=False, + qmark=">", + ).ask() + + query: str = query.rstrip() + + if query == "": + console.print("Empty input received. Please try again!") + continue + + # handle cli / commands with memory ops + if query.startswith("/"): + query_split = query.lstrip("/").lower().split(" ") + query = query_split[0] + args = query_split[1:] + if query == "exit": + break + elif query == "help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(cmd, "bold") + questionary.print(f" {desc}") + elif query in op_description_dict: + if not args: + result = self.memory_service.do_operation(op_name=query) + questionary.print(result) + + elif args[0].isdigit(): + refresh_time = int(args[0]) + while True: + time.sleep(refresh_time) + result = self.memory_service.do_operation(op_name=query) + questionary.print(result) + else: + console.print("unknown command received. Please try again!") + else: + console.print("unknown command received. Please try again!") + continue + while True: - query = questionary.text( - "Enter your message or command:", - multiline=False, - qmark=">", - ).ask() - - query = query.rstrip() - - if query == "": - console.print("Empty input received. Try again!") - continue - - # Handle CLI commands - if query.startswith("/"): - if query.lower() == "/exit": + try: + # with console.status("[bold cyan]Thinking..."): + for msg in self.chat_with_memory(query=query): + console.print(msg.delta, end="") + console.print() + break + except KeyboardInterrupt: + console.print("User interrupt occurred.") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: break - elif query.lower() == "/memory": - console.print(self.memory_service.retrieve_all()) - elif query.lower() == "/help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(cmd, "bold") - questionary.print(f" {desc}") - - continue - - while True: - try: - # with console.status("[bold cyan]Thinking..."): - for msg in self.chat_with_memory(query=query): - console.print(msg.delta, end="") - console.print() + except Exception as e: + console.print(f"An exception occurred when running chat_with_memory(): {e}") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: break - except KeyboardInterrupt: - console.print("User interrupt occurred.") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break - except Exception as e: - console.print( - f"An exception occurred when running chat_with_memory(): {e}" - ) - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break diff --git a/memory_scope/chat_v2/memory_chat.py b/memory_scope/chat_v2/memory_chat.py deleted file mode 100644 index 859758de..00000000 --- a/memory_scope/chat_v2/memory_chat.py +++ /dev/null @@ -1,67 +0,0 @@ -import datetime -from typing import List - -from .base_memory_chat import BaseMemoryChat -from .global_context import GLOBAL_CONTEXT -from enumeration.message_role_enum import MessageRoleEnum -from models.base_model import BaseModel -from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT -from scheme.message import Message -from .memory_service import MemoryService - - -class MemoryChat(BaseMemoryChat): - - def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs): - super().__init__(**kwargs) - self.memory_service = MemoryService(chat_name=chat_name, **kwargs) - self.generation_model_name: str = generation_model - self.history_msg_count: int = history_msg_count - - self._generation_model: BaseModel | None = None - self.history_message_list: List[Message] = [] - - @property - def generation_model(self): - if self._generation_model is None: - self._generation_model = GLOBAL_CONTEXT.model_dict[ - self.generation_model_name - ] - return self._generation_model - - @staticmethod - def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language] - if related_memories: - memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language] - system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories) - return Message( - role=MessageRoleEnum.SYSTEM, - content=system_prompt.strip(), - time_created=time_created, - ) - - def chat_with_memory(self, query: str): - query = query.strip() - if not query: - return - - time_created = int(datetime.datetime.now().timestamp()) - new_message: Message = Message( - role=MessageRoleEnum.USER, content=query, time_created=time_created - ) - related_memories: List[str] = self.memory_service.retrieve(message=new_message) - system_message = self.get_system_prompt(related_memories, time_created) - self.history_message_list.append(new_message) - self.history_message_list = self.history_message_list[-self.history_msg_count :] - all_messages = [system_message] + self.history_message_list - # TODO at xian zhe - return self.generation_model.call(messages=all_messages, stream=True) - - def run(self): - self.memory_service.start_memory_backend() - while True: - query = input("wait for input:") - if query in ["stop", "停止"]: - break - self.chat_with_memory(query=query) diff --git a/memory_scope/chat_v2/memory_service.py b/memory_scope/chat_v2/memory_service.py deleted file mode 100644 index e9fca97a..00000000 --- a/memory_scope/chat_v2/memory_service.py +++ /dev/null @@ -1,70 +0,0 @@ -from constants.common_constants import RELATED_MEMORIES -from enumeration.memory_method_enum import MemoryMethodEnum -from scheme.message import Message -from utils.pipeline import Pipeline -from .base_memory_service import BaseMemoryService - - -class MemoryService(BaseMemoryService): - def __init__( - self, - chat_name: str, - retrieve_pipeline: str, - retrieve_all_pipeline: str, - summary_short_pipeline: str, - summary_long_pipeline: str, - summary_short_interval_time: int = 60, - summary_short_minimum_count: int = 5, - summary_long_interval_time: int = 60 * 5, - summary_long_minimum_count: int = 5 * 5, - **kwargs - ): - super().__init__(**kwargs) - self.retrieve_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.RETRIEVE, - pipeline_str=retrieve_pipeline, - ) - - self.retrieve_all_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.RETRIEVE_ALL, - pipeline_str=retrieve_all_pipeline, - ) - - self.summary_short_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.SUMMARY_SHORT, - pipeline_str=summary_short_pipeline, - loop_interval_time=summary_short_interval_time, - loop_minimum_count=summary_short_minimum_count, - ) - - self.summary_long_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.SUMMARY_LONG, - pipeline_str=summary_long_pipeline, - loop_interval_time=summary_long_interval_time, - loop_minimum_count=summary_long_minimum_count, - ) - - def retrieve(self, message: Message): - self.retrieve_pipeline.submit_message(message, with_lock=False) - self.summary_short_pipeline.submit_message(message) - self.summary_long_pipeline.submit_message(message) - return self.retrieve_pipeline.run(RELATED_MEMORIES) - - def retrieve_all(self): - return self.retrieve_all_pipeline.run(RELATED_MEMORIES) - - def start_memory_backend(self): - self.summary_short_pipeline.start_loop_run() - self.summary_long_pipeline.start_loop_run() - - def get_worker_list(self) -> list: - worker_set = set() - worker_set.update(self.retrieve_pipeline.worker_set) - worker_set.update(self.retrieve_all_pipeline.worker_set) - worker_set.update(self.summary_short_pipeline.worker_set) - worker_set.update(self.summary_long_pipeline.worker_set) - return sorted(worker_set) diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 70d1d103..3f531b4e 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -7,6 +7,10 @@ OPERATION_TYPE = Literal["frontend", "backend"] class BaseOperation(metaclass=ABCMeta): operation_type: OPERATION_TYPE = "frontend" + def __init__(self, name: str, description: str = "", **kwargs): + self.name: str = name + self.description: str = description + def init_workflow(self): pass diff --git a/memory_scope/memory/operation/read_memory.py b/memory_scope/memory/operation/read_memory.py index 8ffe24d2..89956137 100644 --- a/memory_scope/memory/operation/read_memory.py +++ b/memory_scope/memory/operation/read_memory.py @@ -9,8 +9,15 @@ from memory_scope.scheme.message import Message class ReadMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "frontend" - def __init__(self, chat_messages: List[Message], his_msg_count: int = 0, contextual_msg_count: int = 0, **kwargs): - super().__init__(**kwargs) + def __init__(self, + name: str, + description: str, + chat_messages: List[Message], + his_msg_count: int = 0, + contextual_msg_count: int = 0, + **kwargs): + super().__init__(name=name, **kwargs) + BaseOperation.__init__(self, name=name, description=description) self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count self.contextual_msg_count: int = contextual_msg_count diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 3deba4ba..30dc96c8 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -1,6 +1,7 @@ import time from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.constants.common_constants import RESULT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow @@ -8,8 +9,13 @@ from memory_scope.memory.operation.base_workflow import BaseWorkflow class SummaryMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "backend" - def __init__(self, interval_time: int = 300, **kwargs): - super().__init__(**kwargs) + def __init__(self, + name: str, + description: str, + interval_time: int = 300, + **kwargs): + super().__init__(name=name, **kwargs) + BaseOperation.__init__(self, name=name, description=description) self.interval_time: int = interval_time @@ -25,8 +31,10 @@ class SummaryMemory(BaseWorkflow, BaseOperation): self._operation_status_run = True self.run_workflow() + result = self.context.get(RESULT) self.context.clear() self._operation_status_run = False + return result def _loop_operation(self): while self._loop_switch: diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 9d6dce17..1e9070b7 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -2,7 +2,7 @@ import time from typing import List from memory_scope.chat_v2.global_context import G_CONTEXT -from memory_scope.constants.common_constants import CHAT_MESSAGES +from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message @@ -12,13 +12,17 @@ class WriteMemory(BaseOperation, BaseWorkflow): operation_type: OPERATION_TYPE = "backend" def __init__(self, + name: str, + description: str, chat_messages: List[Message], his_msg_count: int = 0, message_lock=None, interval_time: int = 60, contextual_msg_count: int = 6, **kwargs): - super().__init__(**kwargs) + + super().__init__(name=name, **kwargs) + BaseOperation.__init__(self, name=name, description=description) self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count @@ -54,9 +58,11 @@ class WriteMemory(BaseOperation, BaseWorkflow): max_count = not_memorized_size + self.his_msg_count self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] self.run_workflow() + result = self.context.get(RESULT) self.context.clear() self.set_memorized() self._operation_status_run = False + return result def _loop_operation(self): while self._loop_switch: diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 79299265..e2d2a452 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -1,5 +1,7 @@ from abc import ABCMeta, abstractmethod +from typing import List, Dict +from memory_scope.scheme.message import Message from memory_scope.utils.logger import Logger @@ -8,6 +10,16 @@ class BaseMemoryService(metaclass=ABCMeta): self.logger = Logger.get_logger() self.kwargs = kwargs + def submit_messages(self, messages: List[Message] | Message): + pass + + def prepare_service(self): + pass + @abstractmethod def do_operation(self, op_name: str): pass + + @abstractmethod + def get_op_description_dict(self) -> Dict[str, str]: + pass diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index f822c0dd..5bb9895c 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -59,3 +59,6 @@ class ChatMemoryService(BaseMemoryService): self.logger.warning(f"op_name={op_name} is not inited!") return return self.op_dict[op_name].run_operation() + + def get_op_description_dict(self) -> Dict[str, str]: + return {k: v.description for k, v in self.op_dict.items()} From 7c30c104fdd2ac6a7bac34c90c698fd8ca8a8c86 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 17:57:10 +0800 Subject: [PATCH 09/41] [dev] add noinspection --- memory_scope/models/{response.py => model_response.py} | 1 + 1 file changed, 1 insertion(+) rename memory_scope/models/{response.py => model_response.py} (97%) diff --git a/memory_scope/models/response.py b/memory_scope/models/model_response.py similarity index 97% rename from memory_scope/models/response.py rename to memory_scope/models/model_response.py index 841024fd..958356db 100644 --- a/memory_scope/models/response.py +++ b/memory_scope/models/model_response.py @@ -27,6 +27,7 @@ class ModelResponse(BaseModel): def __str__(self, max_size=100, **kwargs): result = {} + # noinspection PyBroadException try: all_dict = self.model_dump() except Exception: From 02bb9718be804e969399e035259521543fe4f2ad Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 18:27:39 +0800 Subject: [PATCH 10/41] [dev] format model response --- memory_scope/models/__init__.py | 2 -- memory_scope/models/base_model.py | 10 ++++++---- .../models/llama_index_embedding_model.py | 20 +++++-------------- .../models/llama_index_generation_model.py | 19 ++++++------------ .../models/llama_index_rerank_model.py | 12 ++++------- memory_scope/utils/registry.py | 3 ++- 6 files changed, 23 insertions(+), 43 deletions(-) diff --git a/memory_scope/models/__init__.py b/memory_scope/models/__init__.py index 5a9d5579..8b137891 100644 --- a/memory_scope/models/__init__.py +++ b/memory_scope/models/__init__.py @@ -1,3 +1 @@ -from memory_scope.utils.registry import Registry -MODEL_REGISTRY = Registry("models") diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 28b12fe5..038cde81 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -3,11 +3,13 @@ import time from abc import abstractmethod, ABCMeta from memory_scope.enumeration.model_enum import ModelEnum -from memory_scope.models import MODEL_REGISTRY -from memory_scope.models.response import ModelResponse, ModelResponseGen +from memory_scope.models.model_response import ModelResponse, ModelResponseGen from memory_scope.utils.logger import Logger +from memory_scope.utils.registry import Registry from memory_scope.utils.timer import Timer +MODEL_REGISTRY = Registry("models") + class BaseModel(metaclass=ABCMeta): m_type: ModelEnum | None = None @@ -70,8 +72,8 @@ class BaseModel(metaclass=ABCMeta): :param kwargs: :return: """ - self.before_call(stream=stream, **kwargs) with Timer(self.__class__.__name__, log_time=False) as t: + self.before_call(stream=stream, **kwargs) for i in range(self.max_retries): try: model_response = self._call(stream=stream, **kwargs) @@ -97,8 +99,8 @@ class BaseModel(metaclass=ABCMeta): :param kwargs: :return: """ - self.before_call(**kwargs) with Timer(self.__class__.__name__, log_time=False) as t: + self.before_call(**kwargs) for i in range(self.max_retries): try: model_response = await self._async_call(**kwargs) diff --git a/memory_scope/models/llama_index_embedding_model.py b/memory_scope/models/llama_index_embedding_model.py index a9397116..bb4d35b4 100644 --- a/memory_scope/models/llama_index_embedding_model.py +++ b/memory_scope/models/llama_index_embedding_model.py @@ -2,20 +2,15 @@ from typing import List from llama_index.embeddings.dashscope import DashScopeEmbedding -from memory_scope.models import MODEL_REGISTRY -from memory_scope.models.base_model import BaseModel -from memory_scope.models.response import ModelResponse, ModelResponseGen from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY +from memory_scope.models.model_response import ModelResponse class LlamaIndexEmbeddingModel(BaseModel): m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL - MODEL_REGISTRY.batch_register( - [ - DashScopeEmbedding, - ] - ) + MODEL_REGISTRY.register("dashscope_embedding", DashScopeEmbedding) def before_call(self, **kwargs): text: str | List[str] = kwargs.pop("text", "") @@ -42,16 +37,11 @@ class LlamaIndexEmbeddingModel(BaseModel): :param kwargs: :return: """ - return ModelResponse( - m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data) - ) + return ModelResponse(m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data)) async def _async_call(self, **kwargs) -> ModelResponse: """ :param kwargs: :return: """ - return ModelResponse( - m_type=self.m_type, - raw=await self.model.aget_text_embedding_batch(**self.data), - ) + return ModelResponse(m_type=self.m_type, raw=await self.model.aget_text_embedding_batch(**self.data)) diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index e2662c12..56f0b88b 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -1,24 +1,17 @@ from typing import List, Dict -from llama_index.core.base.llms.types import ( - ChatMessage, - ChatResponse, - CompletionResponse, -) + +from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse from llama_index.llms.dashscope import DashScope -from enumeration.model_enum import ModelEnum -from . import MODEL_REGISTRY -from .base_model import BaseModel -from .response import ModelResponse, ModelResponseGen +from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY +from memory_scope.models.model_response import ModelResponse, ModelResponseGen class LlamaIndexGenerationModel(BaseModel): m_type: ModelEnum = ModelEnum.GENERATION_MODEL - # TODO rename module name at xianzhe - MODEL_REGISTRY.batch_register([ - DashScope, - ]) + MODEL_REGISTRY.register("dashscope_generation", DashScope) def before_call(self, **kwargs) -> None: prompt: str = kwargs.pop("prompt", "") diff --git a/memory_scope/models/llama_index_rerank_model.py b/memory_scope/models/llama_index_rerank_model.py index 144a3b69..1a19e6aa 100644 --- a/memory_scope/models/llama_index_rerank_model.py +++ b/memory_scope/models/llama_index_rerank_model.py @@ -4,19 +4,15 @@ from llama_index.core.data_structs import Node from llama_index.core.schema import NodeWithScore from llama_index.postprocessor.dashscope_rerank import DashScopeRerank -from models import MODEL_REGISTRY -from models.base_model import BaseModel -from models.response import ModelResponse, ModelResponseGen -from enumeration.model_enum import ModelEnum - +from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY +from memory_scope.models.model_response import ModelResponse class LlamaIndexRerankModel(BaseModel): m_type: ModelEnum = ModelEnum.RANK_MODEL - MODEL_REGISTRY.batch_register([ - DashScopeRerank - ]) + MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank) def before_call(self, **kwargs) -> None: assert "query" in kwargs or "documents" in kwargs diff --git a/memory_scope/utils/registry.py b/memory_scope/utils/registry.py index 9939d3ab..88807387 100644 --- a/memory_scope/utils/registry.py +++ b/memory_scope/utils/registry.py @@ -10,7 +10,8 @@ class Registry(object): self.name: str = name self.module_dict: Dict[str, Any] = {} - def register(self, module: Any, module_name: str = None): + def register(self, module_name: str = None, module: Any = None): + assert module is not None if module_name is None: module_name = module.__name__ From 6b3fb6bbc9fd04ac7e3e31058188364ccc726e82 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 18:32:55 +0800 Subject: [PATCH 11/41] [dev] modify base model registry module --- memory_scope/models/base_model.py | 2 +- memory_scope/utils/registry.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 038cde81..e2e8767e 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -33,7 +33,7 @@ class BaseModel(metaclass=ABCMeta): self.data = {} self.logger = Logger.get_logger() - obj_cls = MODEL_REGISTRY.get(self.method_type) + obj_cls = MODEL_REGISTRY[self.method_type] if not obj_cls: raise RuntimeError(f"method_type={self.method_type} is not supported!") diff --git a/memory_scope/utils/registry.py b/memory_scope/utils/registry.py index 88807387..9d7996ea 100644 --- a/memory_scope/utils/registry.py +++ b/memory_scope/utils/registry.py @@ -28,6 +28,6 @@ class Registry(object): raise NotImplementedError self.module_dict.update(module_name_dict) - def get(self, module_name: str): - assert module_name in self.module_dict, f'{module_name} not found in {self.name}' + def __getitem__(self, module_name: str): + assert module_name in self.module_dict, f"{module_name} not found in {self.name}" return self.module_dict[module_name] From aeeb6012da870519f1159612becaf2b039bb8e8f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:14:16 +0800 Subject: [PATCH 12/41] [dev] update cli memory chat --- memory_scope/chat_v2/cli_memory_chat.py | 8 +++++--- memory_scope/memory/operation/base_operation.py | 3 +++ memory_scope/memory/operation/summary_memory.py | 6 +++++- memory_scope/memory/operation/write_memory.py | 3 +++ memory_scope/memory/service/base_memory_service.py | 6 +++++- memory_scope/memory/worker/base_worker.py | 5 ++--- memory_scope/memory/worker/dummy_worker.py | 2 +- 7 files changed, 24 insertions(+), 9 deletions(-) diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index f9e60d8a..1cf6db9e 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -21,9 +21,9 @@ class CliMemoryChat(BaseMemoryChat): } def __init__(self, memory_service: str, generation_model: str, **kwargs): + super().__init__(**kwargs) self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model - self.kwargs: dict = kwargs @property def memory_service(self) -> BaseMemoryService: @@ -43,7 +43,9 @@ class CliMemoryChat(BaseMemoryChat): system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] if related_memories: memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] - system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt] + related_memories]) + all_prompt_list = [system_prompt, memory_prompt] + all_prompt_list.extend(related_memories) + system_prompt = "\n".join([x.strip() for x in all_prompt_list]) return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) def chat_with_memory(self, query: str): @@ -53,7 +55,7 @@ class CliMemoryChat(BaseMemoryChat): time_created = int(datetime.datetime.now().timestamp()) new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) - related_memories: List[str] = self.memory_service.do_operation("read_memory") + related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) return self.generation_model.call(messages=[system_message, new_message], stream=True) diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 3f531b4e..2287f9e3 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -20,3 +20,6 @@ class BaseOperation(metaclass=ABCMeta): def run_operation_backend(self): pass + + def stop_operation_backend(self): + pass diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 30dc96c8..9eac30aa 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -21,6 +21,7 @@ class SummaryMemory(BaseWorkflow, BaseOperation): self._operation_status_run: bool = False self._loop_switch: bool = False + self._run_thread = None def init_workflow(self): self.init_workers() @@ -44,4 +45,7 @@ class SummaryMemory(BaseWorkflow, BaseOperation): def run_operation_backend(self): if not self._loop_switch: self._loop_switch = True - return G_CONTEXT.thread_pool.submit(self._loop_operation) + self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation) + + def stop_operation_backend(self): + self._loop_switch = False diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 1e9070b7..92bc8b47 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -73,3 +73,6 @@ class WriteMemory(BaseOperation, BaseWorkflow): if not self._loop_switch: self._loop_switch = True return G_CONTEXT.thread_pool.submit(self._loop_operation) + + def stop_operation_backend(self): + self._loop_switch = False diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index e2d2a452..05ee4b54 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -6,7 +6,8 @@ from memory_scope.utils.logger import Logger class BaseMemoryService(metaclass=ABCMeta): - def __init__(self, **kwargs): + def __init__(self, read_memory_key: str = "read_memory", **kwargs): + self.read_memory_key: str = read_memory_key self.logger = Logger.get_logger() self.kwargs = kwargs @@ -23,3 +24,6 @@ class BaseMemoryService(metaclass=ABCMeta): @abstractmethod def get_op_description_dict(self) -> Dict[str, str]: pass + + def read_memory(self): + return self.do_operation(self.read_memory_key) diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 78bfae5e..9f9ed979 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -43,15 +43,14 @@ class BaseWorker(metaclass=ABCMeta): self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----") def get_context(self, key: str, default=None): - return self.context_dict.get(key, default) + return self.context.get(key, default) def set_context(self, key: str, value: Any): if self.is_multi_thread: with self.context_lock: self.context_dict[key] = value else: - self.context_dict[key] = value + self.context[key] = value def __getattr__(self, key): - # raise exception if not exists return self.kwargs[key] diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 9513bc79..d0cb7d93 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -5,4 +5,4 @@ from memory_scope.memory.worker.base_worker import BaseWorker class DummyWorker(BaseWorker): def _run(self): self.set_context(RESULT, ["test 123"]) - self.logger.info("enter dummy worker!") \ No newline at end of file + self.logger.info("enter dummy worker!") From a082bd67ae1ce0089fd762b5e12521f70d5c6020 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:34:01 +0800 Subject: [PATCH 13/41] [dev] update memory service operation func --- memory_scope/chat_v2/cli_memory_chat.py | 6 ++-- .../memory/service/base_memory_service.py | 26 ++++++++++---- .../memory/service/chat_memory_service.py | 34 +++++++------------ 3 files changed, 34 insertions(+), 32 deletions(-) diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index 1cf6db9e..0b39c914 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -61,7 +61,7 @@ class CliMemoryChat(BaseMemoryChat): def run(self): - op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() + op_description_dict: Dict[str, str] = self.memory_service.op_description_dict self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) console = Console() @@ -92,14 +92,14 @@ def run(self): questionary.print(f" {desc}") elif query in op_description_dict: if not args: - result = self.memory_service.do_operation(op_name=query) + result = self.memory_service.operate(op_name=query) questionary.print(result) elif args[0].isdigit(): refresh_time = int(args[0]) while True: time.sleep(refresh_time) - result = self.memory_service.do_operation(op_name=query) + result = self.memory_service.operate(op_name=query) questionary.print(result) else: console.print("unknown command received. Please try again!") diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 05ee4b54..84288b01 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -1,6 +1,8 @@ +import threading from abc import ABCMeta, abstractmethod from typing import List, Dict +from memory_scope.memory.operation.base_operation import BaseOperation from memory_scope.scheme.message import Message from memory_scope.utils.logger import Logger @@ -8,22 +10,32 @@ from memory_scope.utils.logger import Logger class BaseMemoryService(metaclass=ABCMeta): def __init__(self, read_memory_key: str = "read_memory", **kwargs): self.read_memory_key: str = read_memory_key + + self._operation_dict: Dict[str, BaseOperation] = {} + self._op_description_dict: Dict[str, str] = {} + + self.chat_messages: List[Message] = [] + self.message_lock = threading.Lock + self.logger = Logger.get_logger() self.kwargs = kwargs - def submit_messages(self, messages: List[Message] | Message): + def add_messages(self, messages: List[Message] | Message): pass def prepare_service(self): pass @abstractmethod - def do_operation(self, op_name: str): - pass + def operate(self, op_name: str): + raise NotImplementedError - @abstractmethod - def get_op_description_dict(self) -> Dict[str, str]: - pass + @property + def op_description_dict(self) -> Dict[str, str]: + if not self._op_description_dict: + self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()} + return self._op_description_dict def read_memory(self): - return self.do_operation(self.read_memory_key) + assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" + return self.operate(self.read_memory_key) diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index 5bb9895c..820d1420 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -1,7 +1,5 @@ -import threading from typing import List, Dict -from memory_scope.memory.operation.base_operation import BaseOperation from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.scheme.message import Message from memory_scope.utils.tool_functions import init_instance_by_config @@ -15,29 +13,24 @@ class ChatMemoryService(BaseMemoryService): contextual_msg_count: int = 6, **kwargs): super().__init__(**kwargs) - - self.op_dict: Dict[str, BaseOperation] = self._init_operation(memory_operations) self.history_msg_count: int = history_msg_count self.contextual_msg_count: int = contextual_msg_count assert self.history_msg_count >= self.contextual_msg_count - self.chat_messages: List[Message] = [] - self.message_lock = threading.Lock + self._init_operation(memory_operations) def _init_operation(self, memory_operations: Dict[str, dict]): - op_dict: Dict[str, BaseOperation] = {} for name, operation_config in memory_operations.items(): - if name in self.op_dict: + if name in self._operation_dict: self.logger.warning(f"memory operation={name} is repeated!") continue - self.op_dict[name] = init_instance_by_config(config=operation_config, - name=name, - chat_messages=self.chat_messages, - message_lock=self.message_lock, - contextual_msg_count=self.contextual_msg_count) - return op_dict + self._operation_dict[name] = init_instance_by_config(config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + contextual_msg_count=self.contextual_msg_count) - def submit_messages(self, messages: List[Message] | Message): + def add_messages(self, messages: List[Message] | Message): if isinstance(messages, Message): messages = [messages] @@ -49,16 +42,13 @@ class ChatMemoryService(BaseMemoryService): self.chat_messages.pop(0) def prepare_service(self): - for _, operation in self.op_dict.items(): + for _, operation in self._operation_dict.items(): operation.init_workflow() if operation.operation_type == "backend": operation.run_operation_backend() - def do_operation(self, op_name: str): - if op_name not in self.op_dict: + def operate(self, op_name: str): + if op_name not in self._operation_dict: self.logger.warning(f"op_name={op_name} is not inited!") return - return self.op_dict[op_name].run_operation() - - def get_op_description_dict(self) -> Dict[str, str]: - return {k: v.description for k, v in self.op_dict.items()} + return self._operation_dict[op_name].run_operation() From d095e87f3d3338db3cbfe5a7be537acf1bc4b8cf Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:39:16 +0800 Subject: [PATCH 14/41] [dev] rename base memory service params --- config/config.yaml | 2 +- .../memory/service/base_memory_service.py | 18 ++++++++++++++---- .../memory/service/chat_memory_service.py | 5 +---- 3 files changed, 16 insertions(+), 9 deletions(-) diff --git a/config/config.yaml b/config/config.yaml index 9c3c9906..4dcae9d8 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -10,7 +10,7 @@ memory_chat: generation_model: dashscope_generation memory_service: memory_chat_service: - class: memory.service.base_memory_service + class: memory.service.chat_memory_service history_msg_count: 32 contextual_msg_count: 6 memory_operations: diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 84288b01..1d488392 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -8,26 +8,36 @@ from memory_scope.utils.logger import Logger class BaseMemoryService(metaclass=ABCMeta): - def __init__(self, read_memory_key: str = "read_memory", **kwargs): + def __init__(self, + memory_operations: Dict[str, dict], + read_memory_key: str = "read_memory", + **kwargs): + self.memory_operations: Dict[str, dict] = memory_operations self.read_memory_key: str = read_memory_key self._operation_dict: Dict[str, BaseOperation] = {} self._op_description_dict: Dict[str, str] = {} - self.chat_messages: List[Message] = [] self.message_lock = threading.Lock self.logger = Logger.get_logger() self.kwargs = kwargs + self._init_operation(memory_operations) + + @abstractmethod + def _init_operation(self, memory_operations: Dict[str, dict]): + raise NotImplementedError + + @abstractmethod def add_messages(self, messages: List[Message] | Message): - pass + raise NotImplementedError def prepare_service(self): pass @abstractmethod - def operate(self, op_name: str): + def do_operation(self, op_name: str): raise NotImplementedError @property diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index 820d1420..be3297f1 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -8,7 +8,6 @@ from memory_scope.utils.tool_functions import init_instance_by_config class ChatMemoryService(BaseMemoryService): def __init__(self, - memory_operations: Dict[str, dict], history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs): @@ -17,8 +16,6 @@ class ChatMemoryService(BaseMemoryService): self.contextual_msg_count: int = contextual_msg_count assert self.history_msg_count >= self.contextual_msg_count - self._init_operation(memory_operations) - def _init_operation(self, memory_operations: Dict[str, dict]): for name, operation_config in memory_operations.items(): if name in self._operation_dict: @@ -47,7 +44,7 @@ class ChatMemoryService(BaseMemoryService): if operation.operation_type == "backend": operation.run_operation_backend() - def operate(self, op_name: str): + def do_operation(self, op_name: str): if op_name not in self._operation_dict: self.logger.warning(f"op_name={op_name} is not inited!") return From 5364f821a2e4e0a59b7795e846665b5d5f8a86e7 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:46:51 +0800 Subject: [PATCH 15/41] [dev] add default path to tool functions --- config/config.yaml | 1 + memory_scope/memory/operation/read_memory.py | 4 ++-- memory_scope/utils/tool_functions.py | 2 +- 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/config/config.yaml b/config/config.yaml index 4dcae9d8..696c39cb 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -18,6 +18,7 @@ memory_service: class: memory.operation.read_memory workflow: dummy description: "read session messages of the user" + contextual_msg_count: 0 read_memory: class: memory.operation.read_memory workflow: dummy diff --git a/memory_scope/memory/operation/read_memory.py b/memory_scope/memory/operation/read_memory.py index 89956137..d647acd3 100644 --- a/memory_scope/memory/operation/read_memory.py +++ b/memory_scope/memory/operation/read_memory.py @@ -13,8 +13,8 @@ class ReadMemory(BaseWorkflow, BaseOperation): name: str, description: str, chat_messages: List[Message], - his_msg_count: int = 0, - contextual_msg_count: int = 0, + his_msg_count: int = 0, # supplement to the current query + contextual_msg_count: int = 0, # for the current context dialogue **kwargs): super().__init__(name=name, **kwargs) BaseOperation.__init__(self, name=name, description=description) diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 5de34148..a0e847c2 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -10,7 +10,7 @@ def under_line_to_hump(underline_str): return sub[0:1].upper() + sub[1:] -def init_instance_by_config(config: dict, default_class_path: str = "", suffix_name: str = "", **kwargs): +def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs): class_name = config.pop("class") if not class_name: raise RuntimeError("empty class_name!") From 632fa86d48e1b90f888ec3be1ac478de6523695a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:47:19 +0800 Subject: [PATCH 16/41] [dev] rename function name --- memory_scope/memory/service/base_memory_service.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 1d488392..b66bd2a7 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -48,4 +48,4 @@ class BaseMemoryService(metaclass=ABCMeta): def read_memory(self): assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" - return self.operate(self.read_memory_key) + return self.do_operation(self.read_memory_key) From 68b0af8ee36970329abba7c30d59d353df1f8ebf Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:52:34 +0800 Subject: [PATCH 17/41] [dev] rename operation name --- config/config.yaml | 3 ++- memory_scope/memory/worker/base_worker.py | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/config/config.yaml b/config/config.yaml index 696c39cb..7d98f49d 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -13,8 +13,9 @@ memory_service: class: memory.service.chat_memory_service history_msg_count: 32 contextual_msg_count: 6 + read_memory_key: read_memory memory_operations: - read_user_message: + contextual_message: class: memory.operation.read_memory workflow: dummy description: "read session messages of the user" diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 9f9ed979..85426cc3 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -52,5 +52,5 @@ class BaseWorker(metaclass=ABCMeta): else: self.context[key] = value - def __getattr__(self, key): + def __getattr__(self, key: str): return self.kwargs[key] From 3e24a00e8dbe532fd52ba6c3e26e179e37c527ce Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:57:35 +0800 Subject: [PATCH 18/41] [dev] add workflow name to context --- config/config.yaml | 1 - memory_scope/constants/common_constants.py | 2 ++ memory_scope/memory/operation/base_workflow.py | 2 ++ memory_scope/memory/worker/dummy_worker.py | 7 ++++--- 4 files changed, 8 insertions(+), 4 deletions(-) diff --git a/config/config.yaml b/config/config.yaml index 7d98f49d..da00b719 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -33,7 +33,6 @@ memory_service: workflow: dummy description: "write observation memories of the user" interval_time: 60 - contextual_msg_count: 6 summary_memory: class: memory.operation.summary_memory workflow: dummy diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index d7379f91..4d35f18b 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -1,3 +1,5 @@ +WORKFLOW_NAME = "workflow_name" + RESULT = "result" CHAT_MESSAGES = "chat_messages" diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index 8b5ac67c..540f6cb8 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -5,6 +5,7 @@ from itertools import zip_longest from typing import Dict, Any, List from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.constants.common_constants import WORKFLOW_NAME from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.utils.logger import Logger from memory_scope.utils.timer import Timer @@ -103,6 +104,7 @@ class BaseWorkflow(object): def run_workflow(self): with Timer(f"run_workflow_{self.name}"): + self.context[WORKFLOW_NAME] = self.name for workflow_part in self.workflow_worker_list: if len(workflow_part) == 1: if not self._run_sub_workflow(workflow_part[0]): diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index d0cb7d93..5ae68efc 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -1,8 +1,9 @@ -from memory_scope.constants.common_constants import RESULT +from memory_scope.constants.common_constants import RESULT, WORKFLOW_NAME from memory_scope.memory.worker.base_worker import BaseWorker class DummyWorker(BaseWorker): def _run(self): - self.set_context(RESULT, ["test 123"]) - self.logger.info("enter dummy worker!") + workflow_name = self.get_context(WORKFLOW_NAME) + self.logger.info(f"enter workflow={workflow_name}.dummy_worker!") + self.set_context(RESULT, f"test {workflow_name}") From 1b088fcca25ab7e1f0d9372ef24023e5f60bc311 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 12:16:49 +0800 Subject: [PATCH 19/41] [dev] add worker_name to workflow --- config/config.yaml | 32 ++++++++++------------ memory_scope/models/base_model.py | 8 +++--- memory_scope/storage/dummy_monitor.py | 12 ++++++++ memory_scope/storage/dummy_vector_store.py | 24 ++++++++++++++++ 4 files changed, 55 insertions(+), 21 deletions(-) create mode 100644 memory_scope/storage/dummy_monitor.py create mode 100644 memory_scope/storage/dummy_vector_store.py diff --git a/config/config.yaml b/config/config.yaml index da00b719..c6a90ca2 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -15,52 +15,50 @@ memory_service: contextual_msg_count: 6 read_memory_key: read_memory memory_operations: - contextual_message: + read_message: class: memory.operation.read_memory - workflow: dummy + workflow: dummy_worker description: "read session messages of the user" contextual_msg_count: 0 read_memory: class: memory.operation.read_memory - workflow: dummy + workflow: dummy_worker description: "read related memories of the user" list_memory: class: memory.operation.read_memory - workflow: dummy + workflow: dummy_worker description: "read all memories of the user" write_memory: class: memory.operation.write_memory - workflow: dummy + workflow: dummy_worker description: "write observation memories of the user" interval_time: 60 summary_memory: class: memory.operation.summary_memory - workflow: dummy + workflow: dummy_worker description: "summary observation memories of the user" interval_time: 300 models: dashscope_generation: clazz: models.llama_index_generation_model - module_name: DashScope + module_name: dashscope_generation model_name: qwen-max dashscope_embedding: - clazz: models.base_embedding_model - module_name: DashScopeEmbedding + clazz: models.llama_index_embedding_model + module_name: dashscope_embedding model_name: text-embedding-v2 dashscope_rank: clazz: models.base_rank_model - module_name: DashScopeRerank + module_name: dashscope_rank model_name: gte-rerank vector_store: - clazz: storage.base_vector_store - index_name: memory_test - password: '' + clazz: storage.dummy_vector_store + embedding_model: dashscope_embedding monitor: - clazz: storage.base_monitor - index_name: memory_test + clazz: storage.dummy_monitor workers: - update_insight: - clazz: worker.summary_long.update_insight + dummy_worker: + clazz: memory.worker.dummy_worker generation_model: dashscope_generation embedding_model: dashscope_embedding rank_model: dashscope_rank \ No newline at end of file diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index e2e8767e..cd63c57a 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -16,7 +16,7 @@ class BaseModel(metaclass=ABCMeta): def __init__(self, model_name: str, - method_type: str, + module_name: str, timeout: int = None, max_retries: int = 3, retry_interval: float = 1.0, @@ -24,7 +24,7 @@ class BaseModel(metaclass=ABCMeta): **kwargs): self.model_name: str = model_name - self.method_type: str = method_type + self.module_name: str = module_name self.timeout: int = timeout self.max_retries: int = max_retries self.retry_interval: float = retry_interval @@ -33,9 +33,9 @@ class BaseModel(metaclass=ABCMeta): self.data = {} self.logger = Logger.get_logger() - obj_cls = MODEL_REGISTRY[self.method_type] + obj_cls = MODEL_REGISTRY[self.module_name] if not obj_cls: - raise RuntimeError(f"method_type={self.method_type} is not supported!") + raise RuntimeError(f"method_type={self.module_name} is not supported!") if kwargs_filter: allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys()) diff --git a/memory_scope/storage/dummy_monitor.py b/memory_scope/storage/dummy_monitor.py new file mode 100644 index 00000000..f39a917b --- /dev/null +++ b/memory_scope/storage/dummy_monitor.py @@ -0,0 +1,12 @@ +from memory_scope.storage.base_monitor import BaseMonitor + + +class DummyMonitor(BaseMonitor): + def add(self): + pass + + def add_token(self): + pass + + def flush(self): + pass diff --git a/memory_scope/storage/dummy_vector_store.py b/memory_scope/storage/dummy_vector_store.py new file mode 100644 index 00000000..f4c7fad8 --- /dev/null +++ b/memory_scope/storage/dummy_vector_store.py @@ -0,0 +1,24 @@ +from typing import Dict, List + +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.storage.base_vector_store import BaseVectorStore + + +class DummyVectorStore(BaseVectorStore): + def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): + pass + + async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): + pass + + def insert(self, node: MemoryNode): + pass + + def insert_batch(self): + pass + + def delete(self): + pass + + def flush(self): + pass From 6d30e4b18a0d42c1470cb2524d21d5ecc1e62e31 Mon Sep 17 00:00:00 2001 From: hs Date: Thu, 27 Jun 2024 12:22:00 +0800 Subject: [PATCH 20/41] [dev] add memoryscope to class path --- config/config.yaml | 13 +- memory_scope/__init__.py | 2 +- memory_scope/chat/base_memory_chat.py | 1 - memory_scope/chat/base_memory_service.py | 54 +++- memory_scope/chat/chat_memory_service.py | 50 ++++ memory_scope/chat/cli_memory_chat.py | 132 +++++++--- memory_scope/chat/global_context.py | 14 +- memory_scope/chat/memory_chat.py | 2 +- memory_scope/chat/memory_service.py | 70 ----- memory_scope/chat_v2/cli_memory_chat.py | 141 +++++----- memory_scope/chat_v2/global_context.py | 2 - memory_scope/cli.py | 167 ++++-------- memory_scope/cli_job.py | 67 ----- memory_scope/constants/common_constants.py | 6 +- .../memory/service/base_memory_service.py | 2 - .../memory/service/chat_memory_service.py | 22 +- .../models/llama_index_rerank_model.py | 8 +- memory_scope/prompts/get_insight_prompt.py | 8 + memory_scope/prompts/get_reflection_prompt.py | 8 + .../prompts/long_contra_repeat_prompt.py | 8 + memory_scope/prompts/memory_chat_prompt.py | 2 +- memory_scope/prompts/update_insight_prompt.py | 7 + memory_scope/utils/response_text_parser.py | 14 +- memory_scope/worker/base_worker.py | 4 +- memory_scope/worker/es/es_insight_worker.py | 4 +- memory_scope/worker/es/es_new_obs_worker.py | 6 +- .../worker/es/es_not_reflected_worker.py | 6 +- memory_scope/worker/es/es_similar_worker.py | 6 +- memory_scope/worker/es/es_today_obs_worker.py | 6 +- memory_scope/worker/es/load_profile_worker.py | 4 +- memory_scope/worker/memory_base_worker.py | 61 +++-- .../worker/retrieve/extract_time_worker.py | 32 ++- .../worker/retrieve/fuse_rerank_worker.py | 65 +++-- .../worker/retrieve/memory_store_worker.py | 22 +- .../worker/retrieve/parse_params_worker.py | 36 --- .../worker/retrieve/semantic_rank_worker.py | 29 +-- memory_scope/worker/summary_long/__init__.py | 0 .../worker/summary_long/get_insight_worker.py | 166 ++++++++++++ .../summary_long/get_reflection_worker.py | 99 +++++++ .../summary_long/long_contra_repeat_worker.py | 129 ++++++++++ .../summary_long/summary_collect_worker.py | 47 ++++ .../summary_long/update_insight_worker.py | 177 +++++++++++++ .../summary_long/update_profile_worker.py | 241 ++++++++++++++++++ memory_scope/worker/summary_short/__init__.py | 0 .../summary_short/contra_repeat_worker.py | 117 +++++++++ .../get_observation_with_time_worker.py | 167 ++++++++++++ .../summary_short/get_observation_worker.py | 144 +++++++++++ .../summary_short/info_filter_worker.py | 70 +++++ test.py | 13 + 49 files changed, 1911 insertions(+), 540 deletions(-) create mode 100644 memory_scope/chat/chat_memory_service.py delete mode 100644 memory_scope/chat/memory_service.py delete mode 100644 memory_scope/cli_job.py create mode 100644 memory_scope/prompts/get_insight_prompt.py create mode 100644 memory_scope/prompts/get_reflection_prompt.py create mode 100644 memory_scope/prompts/long_contra_repeat_prompt.py create mode 100644 memory_scope/prompts/update_insight_prompt.py delete mode 100644 memory_scope/worker/retrieve/parse_params_worker.py create mode 100644 memory_scope/worker/summary_long/__init__.py create mode 100644 memory_scope/worker/summary_long/get_insight_worker.py create mode 100644 memory_scope/worker/summary_long/get_reflection_worker.py create mode 100644 memory_scope/worker/summary_long/long_contra_repeat_worker.py create mode 100644 memory_scope/worker/summary_long/summary_collect_worker.py create mode 100644 memory_scope/worker/summary_long/update_insight_worker.py create mode 100644 memory_scope/worker/summary_long/update_profile_worker.py create mode 100644 memory_scope/worker/summary_short/__init__.py create mode 100644 memory_scope/worker/summary_short/contra_repeat_worker.py create mode 100644 memory_scope/worker/summary_short/get_observation_with_time_worker.py create mode 100644 memory_scope/worker/summary_short/get_observation_worker.py create mode 100644 memory_scope/worker/summary_short/info_filter_worker.py create mode 100644 test.py diff --git a/config/config.yaml b/config/config.yaml index c6a90ca2..bb4fad52 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -5,23 +5,23 @@ global_config: open_ai_apikey: memory_chat: cli_memory_chat: - class: chat_v2.cli_memory_chat + class: memory_scope.chat_v2.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation memory_service: memory_chat_service: - class: memory.service.chat_memory_service + class: memory_scope.memory.service.chat_memory_service history_msg_count: 32 contextual_msg_count: 6 read_memory_key: read_memory memory_operations: read_message: - class: memory.operation.read_memory + class: memory_scope.memory.operation.read_memory workflow: dummy_worker description: "read session messages of the user" contextual_msg_count: 0 read_memory: - class: memory.operation.read_memory + class: memory_scope.memory.operation.read_memory workflow: dummy_worker description: "read related memories of the user" list_memory: @@ -34,7 +34,7 @@ memory_service: description: "write observation memories of the user" interval_time: 60 summary_memory: - class: memory.operation.summary_memory + class: memory_scope.memory.operation.summary_memory workflow: dummy_worker description: "summary observation memories of the user" interval_time: 300 @@ -61,4 +61,5 @@ workers: clazz: memory.worker.dummy_worker generation_model: dashscope_generation embedding_model: dashscope_embedding - rank_model: dashscope_rank \ No newline at end of file + rank_model: dashscope_rank + diff --git a/memory_scope/__init__.py b/memory_scope/__init__.py index 28bac024..d8b7815a 100644 --- a/memory_scope/__init__.py +++ b/memory_scope/__init__.py @@ -1,3 +1,3 @@ """ Version of MemoryScope.""" -__version__ = "0.1.0-alpha.1" \ No newline at end of file +__version__ = "0.1.0-alpha.1" diff --git a/memory_scope/chat/base_memory_chat.py b/memory_scope/chat/base_memory_chat.py index 0d2566d2..63d351dc 100644 --- a/memory_scope/chat/base_memory_chat.py +++ b/memory_scope/chat/base_memory_chat.py @@ -5,7 +5,6 @@ class BaseMemoryChat(metaclass=ABCMeta): def __init__(self, **kwargs): self.kwargs = kwargs - @abstractmethod def chat_with_memory(self, query: str): """ diff --git a/memory_scope/chat/base_memory_service.py b/memory_scope/chat/base_memory_service.py index 9cf3fd76..1f610b01 100644 --- a/memory_scope/chat/base_memory_service.py +++ b/memory_scope/chat/base_memory_service.py @@ -1,4 +1,54 @@ -class BaseMemoryService(object): - def __init__(self, **kwargs): +import threading +from abc import ABCMeta, abstractmethod +from typing import List, Dict +from memory_scope.memory.operation.base_operation import BaseOperation +from memory_scope.scheme.message import Message +from memory_scope.utils.logger import Logger + + +class BaseMemoryService(metaclass=ABCMeta): + def __init__(self, + memory_operations: Dict[str, dict], + read_memory_key: str = "read_memory", + **kwargs): + self.memory_operations: Dict[str, dict] = memory_operations + self.read_memory_key: str = read_memory_key + + self._operation_dict: Dict[str, BaseOperation] = {} + self._op_description_dict: Dict[str, str] = {} + self.chat_messages: List[Message] = [] + self.message_lock = threading.Lock + + self.logger = Logger.get_logger() self.kwargs = kwargs + + self._init_operation(memory_operations) + + @abstractmethod + def _init_operation(self, memory_operations: Dict[str, dict]): + raise NotImplementedError + + @abstractmethod + def add_messages(self, messages: List[Message] | Message): + raise NotImplementedError + + def prepare_service(self): + pass + + @abstractmethod + def do_operation(self, op_name: str): + raise NotImplementedError + + @property + def op_description_dict(self) -> Dict[str, str]: + if not self._op_description_dict: + self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()} + return self._op_description_dict + + def read_memory(self): + assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" + return self.operate(self.read_memory_key) + + # def __getattr__(self, key): + # return self.kwargs[key] diff --git a/memory_scope/chat/chat_memory_service.py b/memory_scope/chat/chat_memory_service.py new file mode 100644 index 00000000..def938e2 --- /dev/null +++ b/memory_scope/chat/chat_memory_service.py @@ -0,0 +1,50 @@ +from typing import List, Dict + +from memory_scope.memory.service.base_memory_service import BaseMemoryService +from memory_scope.scheme.message import Message +from memory_scope.utils.tool_functions import init_instance_by_config + + +class ChatMemoryService(BaseMemoryService): + def __init__(self, + history_msg_count: int = 32, + contextual_msg_count: int = 6, + **kwargs): + super().__init__(**kwargs) + self.history_msg_count: int = history_msg_count + self.contextual_msg_count: int = contextual_msg_count + assert self.history_msg_count >= self.contextual_msg_count + + def _init_operation(self, memory_operations: Dict[str, dict]): + for name, operation_config in memory_operations.items(): + if name in self._operation_dict: + self.logger.warning(f"memory operation={name} is repeated!") + continue + self._operation_dict[name] = init_instance_by_config(config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + contextual_msg_count=self.contextual_msg_count) + + def add_messages(self, messages: List[Message] | Message): + if isinstance(messages, Message): + messages = [messages] + + messages = sorted(messages, key=lambda x: x.time_created) + self.chat_messages.extend(messages) + if len(self.chat_messages) > self.history_msg_count: + gap_size = len(self.chat_messages) - self.history_msg_count + for _ in range(gap_size): + self.chat_messages.pop(0) + + def prepare_service(self): + for _, operation in self._operation_dict.items(): + operation.init_workflow() + if operation.operation_type == "backend": + operation.run_operation_backend() + + def do_operation(self, op_name: str): + if op_name not in self._operation_dict: + self.logger.warning(f"op_name={op_name} is not inited!") + return + return self._operation_dict[op_name].run_operation() diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 1444edcd..1920aba6 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -1,83 +1,139 @@ import datetime +import time +from typing import Dict, List import questionary -from rich.console import Console - -from .memory_chat import MemoryChat -from enumeration.message_role_enum import MessageRoleEnum -from scheme.message import Message +from memory_scope.chat.base_memory_chat import BaseMemoryChat +from memory_scope.chat.global_context import GlobalContext +from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.memory.service.base_memory_service import BaseMemoryService +from memory_scope.models.base_model import BaseModel +from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT +from memory_scope.scheme.message import Message +from ..models.model_response import ModelResponse, ModelResponseGen -class CliMemoryChat(MemoryChat): - +class CliMemoryChat(BaseMemoryChat): USER_COMMANDS = { - "/exit": "exit the CLI", - "/memory": "print the current contents of agent memory", - "/retrieve": "retrieve related memory", - "/log": "log chat progress", - # TODO add more commands + "exit": "exit the CLI", + "help": "get cli commands help", + "stream": "get stream response" } - def chat_with_memory(self, query): # for testing + def __init__(self, memory_service: str, generation_model: str, **kwargs): + super().__init__(**kwargs) + self._memory_service: BaseMemoryService | str = memory_service + self._generation_model: BaseModel | str = generation_model + self.stream: bool = True + + @property + def memory_service(self) -> BaseMemoryService: + if isinstance(self._memory_service, str): + self._memory_service = GlobalContext.memory_service_dict[self._memory_service] + self._memory_service.prepare_service() + return self._memory_service + + @property + def generation_model(self) -> BaseModel: + if isinstance(self._generation_model, str): + self._generation_model = GlobalContext.model_dict[self._generation_model] + return self._generation_model + + @staticmethod + def get_system_prompt(related_memories: List[str], time_created: int) -> Message: + system_prompt = SYSTEM_PROMPT[GlobalContext.language] + if related_memories: + memory_prompt = MEMORY_PROMPT[GlobalContext.language] + all_prompt_list = [system_prompt, memory_prompt] + all_prompt_list.extend(related_memories) + system_prompt = "\n".join([x.strip() for x in all_prompt_list]) + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) + + def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: query = query.strip() if not query: return time_created = int(datetime.datetime.now().timestamp()) - message = Message( - role=MessageRoleEnum.USER, content=query, time_created=time_created - ) - messages = [message] - return self.generation_model.call(messages=messages, stream=True) + new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) + self.submit_messages(new_message) + related_memories: List[str] = self.memory_service.read_memory() + system_message: Message = self.get_system_prompt(related_memories, time_created) + if self.stream: + for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): + yield result - def retrieve_all(self): # for testing - return "memory 1. 2. 3." + self.submit_messages(result.text) def run(self): - console = Console() + op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() + self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) + while True: query = questionary.text( - "Enter your message or command:", + "Please enter your message or command:", multiline=False, qmark=">", ).ask() - query = query.rstrip() + query: str = query.rstrip() if query == "": - console.print("Empty input received. Try again!") + print("Empty input received. Please try again!") continue - # Handle CLI commands + # handle cli / commands with memory ops if query.startswith("/"): - if query.lower() == "/exit": + query_split = query.lstrip("/").lower().split(" ") + query = query_split[0] + args = query_split[1:] + if query == "exit": break - elif query.lower() == "/memory": - console.print(self.memory_service.retrieve_all()) - elif query.lower() == "/help": + elif query == "help": questionary.print("CLI commands", "bold") for cmd, desc in self.USER_COMMANDS.items(): questionary.print(cmd, "bold") - questionary.print(f" {desc}") + print(f" {desc}") + elif query == "stream": + questionary.print(f"stream: {self.stream}") + self.stream = ~self.stream + elif query in op_description_dict: + if not args: + result = self.memory_service.do_operation(op_name=query) + print(result) + elif args[0].isdigit(): + refresh_time = int(args[0]) + try: + while True: + time.sleep(refresh_time) + result = self.memory_service.do_operation(op_name=query) + print(result, flush=True) + except KeyboardInterrupt: + print("stop refresh!") + else: + print("unknown command received. Please try again!") + else: + print("unknown command received. Please try again!") continue while True: try: - # with console.status("[bold cyan]Thinking..."): - for msg in self.chat_with_memory(query=query): - console.print(msg.delta, end="") - console.print() + if self.stream: + for msg in self.chat_with_memory(query=query): + print(msg.text, flush=True) + print() + else: + msg = self.chat_with_memory(query=query) + print(msg.text) break except KeyboardInterrupt: - console.print("User interrupt occurred.") + questionary.print("User interrupt occurred.") retry = questionary.confirm("Retry chat_with_memory()?").ask() if not retry: break except Exception as e: - console.print( - f"An exception occurred when running chat_with_memory(): {e}" - ) + questionary.print(f"An exception occurred when running chat_with_memory(): {e}") retry = questionary.confirm("Retry chat_with_memory()?").ask() if not retry: break diff --git a/memory_scope/chat/global_context.py b/memory_scope/chat/global_context.py index 75a63d62..901f641f 100644 --- a/memory_scope/chat/global_context.py +++ b/memory_scope/chat/global_context.py @@ -1,19 +1,19 @@ from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any -from chat.base_memory_chat import BaseMemoryChat -from enumeration.language_enum import LanguageEnum -from models.base_model import BaseModel -from storage.base_monitor import BaseMonitor -from storage.base_vector_store import BaseVectorStore -from worker.base_worker import BaseWorker +from .base_memory_chat import BaseMemoryChat +from ..enumeration.language_enum import LanguageEnum +from ..models.base_model import BaseModel +from ..storage.base_monitor import BaseMonitor +from ..storage.base_vector_store import BaseVectorStore +from ..worker.base_worker import BaseWorker class GlobalContext(object): def __init__(self): self.global_configs: Dict[str, Any] = {} - self.worker_dict: Dict[str, Dict[str, BaseWorker]] = {} + self.worker_config: Dict[str, Dict[str, BaseWorker]] = {} self.model_dict: Dict[str, BaseModel] = {} diff --git a/memory_scope/chat/memory_chat.py b/memory_scope/chat/memory_chat.py index 73f63fb6..f954af29 100644 --- a/memory_scope/chat/memory_chat.py +++ b/memory_scope/chat/memory_chat.py @@ -43,7 +43,7 @@ class MemoryChat(BaseMemoryChat): related_memories: List[str] = self.memory_service.retrieve(message=new_message) system_message = self.get_system_prompt(related_memories, time_created) self.history_message_list.append(new_message) - self.history_message_list = self.history_message_list[-self.history_msg_count :] + self.history_message_list = self.history_message_list[-self.history_msg_count:] all_messages = [system_message] + self.history_message_list # TODO at xian zhe return self.generation_model.call(messages=all_messages, stream=True) diff --git a/memory_scope/chat/memory_service.py b/memory_scope/chat/memory_service.py deleted file mode 100644 index e9fca97a..00000000 --- a/memory_scope/chat/memory_service.py +++ /dev/null @@ -1,70 +0,0 @@ -from constants.common_constants import RELATED_MEMORIES -from enumeration.memory_method_enum import MemoryMethodEnum -from scheme.message import Message -from utils.pipeline import Pipeline -from .base_memory_service import BaseMemoryService - - -class MemoryService(BaseMemoryService): - def __init__( - self, - chat_name: str, - retrieve_pipeline: str, - retrieve_all_pipeline: str, - summary_short_pipeline: str, - summary_long_pipeline: str, - summary_short_interval_time: int = 60, - summary_short_minimum_count: int = 5, - summary_long_interval_time: int = 60 * 5, - summary_long_minimum_count: int = 5 * 5, - **kwargs - ): - super().__init__(**kwargs) - self.retrieve_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.RETRIEVE, - pipeline_str=retrieve_pipeline, - ) - - self.retrieve_all_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.RETRIEVE_ALL, - pipeline_str=retrieve_all_pipeline, - ) - - self.summary_short_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.SUMMARY_SHORT, - pipeline_str=summary_short_pipeline, - loop_interval_time=summary_short_interval_time, - loop_minimum_count=summary_short_minimum_count, - ) - - self.summary_long_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.SUMMARY_LONG, - pipeline_str=summary_long_pipeline, - loop_interval_time=summary_long_interval_time, - loop_minimum_count=summary_long_minimum_count, - ) - - def retrieve(self, message: Message): - self.retrieve_pipeline.submit_message(message, with_lock=False) - self.summary_short_pipeline.submit_message(message) - self.summary_long_pipeline.submit_message(message) - return self.retrieve_pipeline.run(RELATED_MEMORIES) - - def retrieve_all(self): - return self.retrieve_all_pipeline.run(RELATED_MEMORIES) - - def start_memory_backend(self): - self.summary_short_pipeline.start_loop_run() - self.summary_long_pipeline.start_loop_run() - - def get_worker_list(self) -> list: - worker_set = set() - worker_set.update(self.retrieve_pipeline.worker_set) - worker_set.update(self.retrieve_all_pipeline.worker_set) - worker_set.update(self.summary_short_pipeline.worker_set) - worker_set.update(self.summary_long_pipeline.worker_set) - return sorted(worker_set) diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index 0b39c914..fa0bead1 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -3,8 +3,6 @@ import time from typing import Dict, List import questionary -from rich.console import Console - from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat from memory_scope.chat_v2.global_context import G_CONTEXT from memory_scope.enumeration.message_role_enum import MessageRoleEnum @@ -12,18 +10,21 @@ from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT from memory_scope.scheme.message import Message +from ..models.model_response import ModelResponse, ModelResponseGen class CliMemoryChat(BaseMemoryChat): USER_COMMANDS = { "exit": "exit the CLI", "help": "get cli commands help", + "stream": "get stream response" } def __init__(self, memory_service: str, generation_model: str, **kwargs): super().__init__(**kwargs) self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model + self.stream: bool = True @property def memory_service(self) -> BaseMemoryService: @@ -48,79 +49,91 @@ class CliMemoryChat(BaseMemoryChat): system_prompt = "\n".join([x.strip() for x in all_prompt_list]) return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) - def chat_with_memory(self, query: str): + def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: query = query.strip() if not query: return time_created = int(datetime.datetime.now().timestamp()) new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) + self.submit_messages(new_message) related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) - return self.generation_model.call(messages=[system_message, new_message], stream=True) + if self.stream: + for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): + yield result + self.submit_messages(result.text) -def run(self): - op_description_dict: Dict[str, str] = self.memory_service.op_description_dict - self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) - - console = Console() - while True: - query = questionary.text( - "Please enter your message or command:", - multiline=False, - qmark=">", - ).ask() - - query: str = query.rstrip() - - if query == "": - console.print("Empty input received. Please try again!") - continue - - # handle cli / commands with memory ops - if query.startswith("/"): - query_split = query.lstrip("/").lower().split(" ") - query = query_split[0] - args = query_split[1:] - if query == "exit": - break - elif query == "help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(cmd, "bold") - questionary.print(f" {desc}") - elif query in op_description_dict: - if not args: - result = self.memory_service.operate(op_name=query) - questionary.print(result) - - elif args[0].isdigit(): - refresh_time = int(args[0]) - while True: - time.sleep(refresh_time) - result = self.memory_service.operate(op_name=query) - questionary.print(result) - else: - console.print("unknown command received. Please try again!") - else: - console.print("unknown command received. Please try again!") - continue + def run(self): + op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() + self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) while True: - try: - # with console.status("[bold cyan]Thinking..."): - for msg in self.chat_with_memory(query=query): - console.print(msg.delta, end="") - console.print() - break - except KeyboardInterrupt: - console.print("User interrupt occurred.") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: + query = questionary.text( + "Please enter your message or command:", + multiline=False, + qmark=">", + ).ask() + + query: str = query.rstrip() + + if query == "": + print("Empty input received. Please try again!") + continue + + # handle cli / commands with memory ops + if query.startswith("/"): + query_split = query.lstrip("/").lower().split(" ") + query = query_split[0] + args = query_split[1:] + if query == "exit": break - except Exception as e: - console.print(f"An exception occurred when running chat_with_memory(): {e}") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: + elif query == "help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(cmd, "bold") + print(f" {desc}") + elif query == "stream": + questionary.print(f"stream: {self.stream}") + self.stream = ~self.stream + elif query in op_description_dict: + if not args: + result = self.memory_service.do_operation(op_name=query) + print(result) + + elif args[0].isdigit(): + refresh_time = int(args[0]) + try: + while True: + time.sleep(refresh_time) + result = self.memory_service.do_operation(op_name=query) + print(result, flush=True) + except KeyboardInterrupt: + print("stop refresh!") + else: + print("unknown command received. Please try again!") + else: + print("unknown command received. Please try again!") + continue + + while True: + try: + if self.stream: + for msg in self.chat_with_memory(query=query): + print(msg.text, flush=True) + print() + else: + msg = self.chat_with_memory(query=query) + print(msg.text) break + except KeyboardInterrupt: + questionary.print("User interrupt occurred.") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: + break + except Exception as e: + questionary.print(f"An exception occurred when running chat_with_memory(): {e}") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: + break diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py index abfbf2fa..7b08a2bb 100644 --- a/memory_scope/chat_v2/global_context.py +++ b/memory_scope/chat_v2/global_context.py @@ -24,5 +24,3 @@ class GlobalContext(pydantic.BaseModel): thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool") language: LanguageEnum = pydantic.Field(LanguageEnum.CN, description="language: cn / en") - -G_CONTEXT = GlobalContext() diff --git a/memory_scope/cli.py b/memory_scope/cli.py index a656fdfe..cb040d8d 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,135 +1,76 @@ -import json -import os from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any, List -import sys -import time -import fire -from datetime import datetime +from typing import Dict, Any -from chat.global_context import GLOBAL_CONTEXT -from enumeration.language_enum import LanguageEnum -from enumeration.model_enum import ModelEnum -from utils.logger import Logger -from utils.tool_functions import ( - complete_config_name, - init_instance_by_config, - under_line_to_hump, -) -from chat.memory_chat import MemoryChat -from enumeration.message_role_enum import MessageRoleEnum -from scheme.message import Message -from chat.base_memory_chat import BaseMemoryChat -from models.llama_index_generation_model import LlamaIndexGenerationModel -from models.llama_index_embedding_model import LlamaIndexEmbeddingModel -from models.llama_index_rerank_model import LlamaIndexRerankModel +import yaml +import fire + +from .chat_v2.global_context import G_CONTEXT +from .enumeration.language_enum import LanguageEnum +from .utils.logger import Logger +from .utils.tool_functions import init_instance_by_config class CliJob(object): - def __init__(self, config_path: str): + def __init__(self, config_path: str, config_suffix: str = ".yaml"): self.config_path: str = config_path - self.config_base_dir: str = os.path.dirname(config_path) + self.config_suffix: str = config_suffix self.config: Dict[str, Any] = {} - self.worker_chat_dict: Dict[str, List[str]] = {} - self.logger: Logger = Logger.get_logger("memory_chat") - - def init_memory_chat(self): - for chat_name in GLOBAL_CONTEXT.global_configs["chat_list"]: - memory_chat_config = self.config[chat_name] - memory_chat: BaseMemoryChat = init_instance_by_config( - memory_chat_config, chat_name=chat_name - ) - GLOBAL_CONTEXT.memory_chat_dict[chat_name] = memory_chat - - for worker_name in memory_chat.memory_service.get_worker_list(): - if worker_name not in self.worker_chat_dict: - self.worker_chat_dict[worker_name] = [] - self.worker_chat_dict[worker_name].append(chat_name) - - generation_model = memory_chat_config[ModelEnum.GENERATION_MODEL.value] - self.init_model(generation_model) - - def init_model(self, model_name: str): - if not model_name or model_name in GLOBAL_CONTEXT.model_dict: - return - - with open( - os.path.join( - self.config_base_dir, "model", complete_config_name(model_name) - ) - ) as f: - model_config = json.load(f) - GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config) - - def init_workers(self): - """load worker config & init workers""" - worker_config_name: str = self.config["workers"] - with open( - os.path.join(self.config_base_dir, complete_config_name(worker_config_name)) - ) as f: - worker_config_dict = json.load(f) - - for worker_name, worker_config in worker_config_dict.items(): - if worker_name not in self.worker_chat_dict: - continue - - chat_name_list = self.worker_chat_dict[worker_name] - for chat_name in chat_name_list: - if chat_name not in GLOBAL_CONTEXT.worker_dict: - GLOBAL_CONTEXT.worker_dict[chat_name] = {} - GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = ( - init_instance_by_config( - worker_config, - suffix_name="worker", - **GLOBAL_CONTEXT.global_configs, - ) - ) - - self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value)) - self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value)) - self.init_model(worker_config.get(ModelEnum.RANK_MODEL.value)) + self.logger: Logger = Logger.get_logger("cli_job") @staticmethod - def set_global_config(): - """TODO set global_configs & set apikey into env""" - GLOBAL_CONTEXT.language = LanguageEnum( - GLOBAL_CONTEXT.global_configs["language"] - ) - GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor( - max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"]) + def set_global_config(global_config: Dict[str, Any]): + """set global_configs & set apikey into env + :return: + TODO at sen + """ + G_CONTEXT.global_config = global_config + G_CONTEXT.language = LanguageEnum(global_config["language"]) + G_CONTEXT.thread_pool = ThreadPoolExecutor( + max_workers=int(global_config["max_workers"]) ) def init_global_content_by_config(self): - with open(complete_config_name(self.config_path)) as f: - self.config = json.load(f) + # load config + config_path = self.config_path + if not self.config_path.endswith(self.config_suffix): + config_path += self.config_suffix + with open(config_path) as f: + self.config = yaml.load(f, yaml.FullLoader) - GLOBAL_CONTEXT.global_configs = self.config["global_configs"] - self.set_global_config() + # set global_config + self.set_global_config(self.config["global_config"]) - self.init_memory_chat() + # init memory_chat + for name, conf in self.config["memory_chat"].items(): + G_CONTEXT.memory_chat_dict[name] = init_instance_by_config( + conf, name=name + ) - self.init_workers() + # set memory_service + for name, conf in self.config["memory_service"].items(): + G_CONTEXT.memory_service_dict[name] = init_instance_by_config( + conf, name=name + ) - ## TODO no db and monitor now - # GLOBAL_CONTEXT.vector_store = init_instance_by_config( - # self.config["vector_store"] - # ) - # GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) + # init models + for name, conf in self.config["models"].items(): + G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) + + # init vector_store + G_CONTEXT.vector_store = init_instance_by_config( + self.config["vector_store"] + ) + + # init monitor + G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) + + # set worker config + G_CONTEXT.worker_config = self.config["workers"] @staticmethod def run(): - with GLOBAL_CONTEXT.thread_pool: - memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0] + with G_CONTEXT.thread_pool: + memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() - - -def main(config_path: str): - job = CliJob(config_path=config_path) - job.init_global_content_by_config() - job.run() - - -if __name__ == "__main__": - fire.Fire(main) \ No newline at end of file diff --git a/memory_scope/cli_job.py b/memory_scope/cli_job.py deleted file mode 100644 index 48844129..00000000 --- a/memory_scope/cli_job.py +++ /dev/null @@ -1,67 +0,0 @@ -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - -import yaml - -from chat_v2.global_context import G_CONTEXT -from enumeration.language_enum import LanguageEnum -from utils.logger import Logger -from utils.tool_functions import init_instance_by_config - - -class CliJob(object): - - def __init__(self, config_path: str, config_suffix: str = ".yaml"): - self.config_path: str = config_path - self.config_suffix: str = config_suffix - self.config: Dict[str, Any] = {} - - self.logger: Logger = Logger.get_logger("cli_job") - - @staticmethod - def set_global_config(global_config: Dict[str, Any]): - """ set global_configs & set apikey into env - :return: - TODO at sen - """ - G_CONTEXT.global_config = global_config - G_CONTEXT.language = LanguageEnum(global_config["language"]) - G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"])) - - def init_global_content_by_config(self): - # load config - config_path = self.config_path - if not self.config_path.endswith(self.config_suffix): - config_path += self.config_suffix - with open(config_path) as f: - self.config = yaml.load(f, yaml.FullLoader) - - # set global_config - self.set_global_config(self.config["global_config"]) - - # init memory_chat - for name, conf in self.config["memory_chat"].items(): - G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name) - - # set memory_service - for name, conf in self.config["memory_service"].items(): - G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name) - - # init models - for name, conf in self.config["models"].items(): - G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) - - # init vector_store - G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"]) - - # init monitor - G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) - - # set worker config - G_CONTEXT.worker_config = self.config["workers"] - - @staticmethod - def run(): - with G_CONTEXT.thread_pool: - memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] - memory_chat.run() diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index 4d35f18b..81a65558 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -22,10 +22,6 @@ RELATED_MEMORIES = "related_memories" MODIFIED_MEMORIES = "modified_memories" -RESPONSE_EXT_INFO = "response_ext_info" - -PROMPT_CONFIG = "prompt_config" - MESSAGES = "messages" EXTRACT_TIME_DICT = "extract_time_dict" @@ -116,3 +112,5 @@ DATATIME_KEY_MAP = { "周": "week", "星期几": "weekday", } + +CONTENT_MODIFIED = "content_modified" \ No newline at end of file diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index b66bd2a7..50fefca8 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -23,8 +23,6 @@ class BaseMemoryService(metaclass=ABCMeta): self.logger = Logger.get_logger() self.kwargs = kwargs - self._init_operation(memory_operations) - @abstractmethod def _init_operation(self, memory_operations: Dict[str, dict]): raise NotImplementedError diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index be3297f1..badd26c8 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -6,26 +6,28 @@ from memory_scope.utils.tool_functions import init_instance_by_config class ChatMemoryService(BaseMemoryService): - - def __init__(self, - history_msg_count: int = 32, - contextual_msg_count: int = 6, - **kwargs): + def __init__( + self, history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs + ): super().__init__(**kwargs) self.history_msg_count: int = history_msg_count self.contextual_msg_count: int = contextual_msg_count assert self.history_msg_count >= self.contextual_msg_count + self._init_operation(self.memory_operations) + def _init_operation(self, memory_operations: Dict[str, dict]): for name, operation_config in memory_operations.items(): if name in self._operation_dict: self.logger.warning(f"memory operation={name} is repeated!") continue - self._operation_dict[name] = init_instance_by_config(config=operation_config, - name=name, - chat_messages=self.chat_messages, - message_lock=self.message_lock, - contextual_msg_count=self.contextual_msg_count) + self._operation_dict[name] = init_instance_by_config( + config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + contextual_msg_count=self.contextual_msg_count, + ) def add_messages(self, messages: List[Message] | Message): if isinstance(messages, Message): diff --git a/memory_scope/models/llama_index_rerank_model.py b/memory_scope/models/llama_index_rerank_model.py index 1a19e6aa..c76cbd10 100644 --- a/memory_scope/models/llama_index_rerank_model.py +++ b/memory_scope/models/llama_index_rerank_model.py @@ -19,7 +19,9 @@ class LlamaIndexRerankModel(BaseModel): query: str = kwargs.pop("query", "") documents: List[str] = kwargs.pop("documents", []) - assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}" + assert ( + query and documents + ), f"query or documents is empty! query={query}, documents={len(documents)}" # using -1.0 as dummy scores nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents] @@ -41,7 +43,9 @@ class LlamaIndexRerankModel(BaseModel): return model_response def _call(self, **kwargs) -> ModelResponse: - return ModelResponse(m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data)) + return ModelResponse( + m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data) + ) async def _async_call(self, **kwargs) -> ModelResponse: raise NotImplementedError diff --git a/memory_scope/prompts/get_insight_prompt.py b/memory_scope/prompts/get_insight_prompt.py new file mode 100644 index 00000000..592c9fe2 --- /dev/null +++ b/memory_scope/prompts/get_insight_prompt.py @@ -0,0 +1,8 @@ +from ..enumeration.language_enum import LanguageEnum + + +GET_INSIGHT_SYSTEM_PROMPT = {} + +GET_INSIGHT_FEW_SHOT_PROMPT = {} + +GET_INSIGHT_USER_QUERY_PROMPT = {} diff --git a/memory_scope/prompts/get_reflection_prompt.py b/memory_scope/prompts/get_reflection_prompt.py new file mode 100644 index 00000000..c254ade6 --- /dev/null +++ b/memory_scope/prompts/get_reflection_prompt.py @@ -0,0 +1,8 @@ +from ..enumeration.language_enum import LanguageEnum + + +GET_REFLECTION_SYSTEM_PROMPT = {} + +GET_REFLECTION_FEW_SHOT_PROMPT = {} + +GET_REFLECTION_USER_QUERY_PROMPT = {} diff --git a/memory_scope/prompts/long_contra_repeat_prompt.py b/memory_scope/prompts/long_contra_repeat_prompt.py new file mode 100644 index 00000000..e4593d14 --- /dev/null +++ b/memory_scope/prompts/long_contra_repeat_prompt.py @@ -0,0 +1,8 @@ +from ..enumeration.language_enum import LanguageEnum + + +LONG_CONTRA_REPEAT_SYSTEM_PROMPT = {} + +LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT = {} + +LONG_CONTRA_REPEAT_USER_QUERY_PROMPT = {} diff --git a/memory_scope/prompts/memory_chat_prompt.py b/memory_scope/prompts/memory_chat_prompt.py index cd38dfb3..0c762ec9 100644 --- a/memory_scope/prompts/memory_chat_prompt.py +++ b/memory_scope/prompts/memory_chat_prompt.py @@ -1,4 +1,4 @@ -from enumeration.language_enum import LanguageEnum +from ..enumeration.language_enum import LanguageEnum SYSTEM_PROMPT = { LanguageEnum.CN: """ diff --git a/memory_scope/prompts/update_insight_prompt.py b/memory_scope/prompts/update_insight_prompt.py new file mode 100644 index 00000000..e0325044 --- /dev/null +++ b/memory_scope/prompts/update_insight_prompt.py @@ -0,0 +1,7 @@ +from ..enumeration.language_enum import LanguageEnum + +UPDATE_INSIGHT_SYSTEM_PROMPT = {} + +UPDATE_INSIGHT_FEW_SHOT_PROMPT = {} + +UPDATE_INSIGHT_USER_QUERY_PROMPT = {} diff --git a/memory_scope/utils/response_text_parser.py b/memory_scope/utils/response_text_parser.py index 48dbbd9c..6fcc6f5a 100644 --- a/memory_scope/utils/response_text_parser.py +++ b/memory_scope/utils/response_text_parser.py @@ -4,7 +4,7 @@ from utils.logger import Logger class ResponseTextParser(object): - pattern_v1 = re.compile(r'<(.*?)>') + pattern_v1 = re.compile(r"<(.*?)>") def __init__(self, response_text: str): self.response_text: str = response_text.strip() @@ -12,22 +12,26 @@ class ResponseTextParser(object): def parse_v1(self, prefix: str = ""): result = [] - for line in self.response_text.split('\n'): + for line in self.response_text.split("\n"): line = line.strip() if not line: continue matches = [match.group(1) for match in self.pattern_v1.finditer(line)] if matches: result.append(matches) - self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2) + self.logger.info( + f"{prefix} response_text={self.response_text} result={result}", stacklevel=2 + ) return result def parse_v2(self, prefix: str = ""): result = [] - for line in self.response_text.split('\n'): + for line in self.response_text.split("\n"): line = line.strip() if not line or line == "无": continue result.append(line) - self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2) + self.logger.info( + f"{prefix} response_text={self.response_text} result={result}", stacklevel=2 + ) return result diff --git a/memory_scope/worker/base_worker.py b/memory_scope/worker/base_worker.py index a47fd1ed..7f0dd62f 100644 --- a/memory_scope/worker/base_worker.py +++ b/memory_scope/worker/base_worker.py @@ -1,7 +1,7 @@ from typing import Any, Dict -from utils.logger import Logger -from utils.timer import Timer +from ..utils.logger import Logger +from ..utils.timer import Timer class BaseWorker(object): diff --git a/memory_scope/worker/es/es_insight_worker.py b/memory_scope/worker/es/es_insight_worker.py index 426d126e..985aec13 100644 --- a/memory_scope/worker/es/es_insight_worker.py +++ b/memory_scope/worker/es/es_insight_worker.py @@ -13,9 +13,9 @@ class EsInsightWorker(MemoryBaseWorker): insight_nodes = self.vector_store.retrieve( size=self.kwargs.es_insight_top_k, filter_dict={ - "memoryId": self.memory_id, + "memory_id": self.memory_id, "status": MemoryNodeStatus.ACTIVE.value, - "memoryType": MemoryTypeEnum.INSIGHT.value, + "memory_type": MemoryTypeEnum.INSIGHT.value, }, ) self.logger.info(f"insight_nodes.size={len(insight_nodes)}") diff --git a/memory_scope/worker/es/es_new_obs_worker.py b/memory_scope/worker/es/es_new_obs_worker.py index b706c998..35823e96 100644 --- a/memory_scope/worker/es/es_new_obs_worker.py +++ b/memory_scope/worker/es/es_new_obs_worker.py @@ -12,10 +12,10 @@ class EsNewObsWorker(MemoryBaseWorker): new_obs_nodes = self.vector_store.retrieve( size=self.kwargs.es_new_obs_top_k, filter_dict={ - "memoryId": self.memory_id, + "memory_id": self.memory_id, "status": MemoryNodeStatus.ACTIVE.value, - "memoryType": MemoryTypeEnum.OBSERVATION.value, - f"metaData.{NEW}": "1", + "memory_type": MemoryTypeEnum.OBSERVATION.value, + f"meta_data.{NEW}": "1", }, ) self.logger.info(f"es new obs, size={len(new_obs_nodes)}") diff --git a/memory_scope/worker/es/es_not_reflected_worker.py b/memory_scope/worker/es/es_not_reflected_worker.py index 5a128824..1a9ab0cd 100644 --- a/memory_scope/worker/es/es_not_reflected_worker.py +++ b/memory_scope/worker/es/es_not_reflected_worker.py @@ -14,13 +14,13 @@ class EsNotReflectedWorker(MemoryBaseWorker): not_reflected_obs_nodes = self.vector_store.retrieve( size=self.kwargs.es_new_obs_top_k, filter_dict={ - "memoryId": self.memory_id, + "memory_id": self.memory_id, "status": MemoryNodeStatus.ACTIVE.value, - "memoryType": [ + "memory_type": [ MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value, ], - f"metaData.{REFLECTED}": "0", + f"meta_data.{REFLECTED}": "0", }, ) self.logger.info( diff --git a/memory_scope/worker/es/es_similar_worker.py b/memory_scope/worker/es/es_similar_worker.py index ec2cbe27..ddca69b9 100644 --- a/memory_scope/worker/es/es_similar_worker.py +++ b/memory_scope/worker/es/es_similar_worker.py @@ -19,9 +19,9 @@ class EsSimilarWorker(MemoryBaseWorker): text=query, size=self.es_similar_top_k, filter_dict={ - "memoryId": self.memory_id, + "memory_id": self.memory_id, "status": MemoryNodeStatus.ACTIVE.value, - "memoryType": [ + "memory_type": [ MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.OBS_CUSTOMIZED.value, @@ -30,7 +30,7 @@ class EsSimilarWorker(MemoryBaseWorker): ) for node in similar_obs_nodes: - node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value + node.meta_data[RECALL_TYPE] = MemoryRecallType.SIMILAR.value self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}") for node in similar_obs_nodes: self.logger.info(f"node={node.content} score_similar={node.score_similar}") diff --git a/memory_scope/worker/es/es_today_obs_worker.py b/memory_scope/worker/es/es_today_obs_worker.py index d7d2abfb..8cd1ce54 100644 --- a/memory_scope/worker/es/es_today_obs_worker.py +++ b/memory_scope/worker/es/es_today_obs_worker.py @@ -21,10 +21,10 @@ class EsTodayObsWorker(MemoryBaseWorker): today_obs_nodes = self.vector_store.retrieve( size=self.es_today_obs_top_k, filter_dict={ - "memoryId": self.memory_id, + "memory_id": self.memory_id, "status": MemoryNodeStatus.ACTIVE.value, - "memoryType": MemoryTypeEnum.OBSERVATION.value, - f"metaData.{DT}": time_to_formatted_str(msg_time_created), + "memory_type": MemoryTypeEnum.OBSERVATION.value, + f"meta_Data.{DT}": time_to_formatted_str(msg_time_created), }, ) diff --git a/memory_scope/worker/es/load_profile_worker.py b/memory_scope/worker/es/load_profile_worker.py index f78eada9..90ad2dc9 100644 --- a/memory_scope/worker/es/load_profile_worker.py +++ b/memory_scope/worker/es/load_profile_worker.py @@ -13,9 +13,9 @@ class LoadProfileWorker(MemoryBaseWorker): user_profile_node = self.vector_store( size=10000, filter_dict={ - "memoryId": self.memory_id, + "memory_id": self.memory_id, "status": MemoryNodeStatus.ACTIVE.value, - "memoryType": [ + "memory_type": [ MemoryTypeEnum.PROFILE.value, MemoryTypeEnum.PROFILE_CUSTOMIZED.value, ], diff --git a/memory_scope/worker/memory_base_worker.py b/memory_scope/worker/memory_base_worker.py index e8d5cc0c..7e4fe846 100644 --- a/memory_scope/worker/memory_base_worker.py +++ b/memory_scope/worker/memory_base_worker.py @@ -1,20 +1,20 @@ -from typing import List +from typing import List, Dict -from chat.global_context import GLOBAL_CONTEXT -from constants.common_constants import MESSAGES, CHAT_NAME -from models.base_model import BaseModel -from scheme.message import Message -from storage.base_monitor import BaseMonitor -from storage.base_vector_store import BaseVectorStore -from worker.base_worker import BaseWorker +from ..chat.global_context import GLOBAL_CONTEXT +from ..constants.common_constants import MESSAGES, CHAT_NAME +from ..models.base_model import BaseModel +from ..scheme.message import Message +from ..storage.base_monitor import BaseMonitor +from ..storage.base_vector_store import BaseVectorStore +from ..worker.base_worker import BaseWorker +from ..scheme.memory_node import MemoryNode +from ..constants import common_constants class MemoryBaseWorker(BaseWorker): - def __init__(self, - embedding_model: str, - generation_model: str, - rank_model: str, - **kwargs): + def __init__( + self, embedding_model: str, generation_model: str, rank_model: str, **kwargs + ): super(MemoryBaseWorker, self).__init__(**kwargs) self.embedding_model_name: str = embedding_model self.generation_model_name: str = generation_model @@ -40,25 +40,29 @@ class MemoryBaseWorker(BaseWorker): return self.get_context(CHAT_NAME) @property - def embedding_model(self): + def embedding_model(self) -> BaseModel: if self._embedding_model is None: - self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name) + self._embedding_model = GLOBAL_CONTEXT.model_dict.get( + self.embedding_model_name + ) return self._embedding_model @property - def generation_model(self): + def generation_model(self) -> BaseModel: if self._generation_model is None: - self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name) + self._generation_model = GLOBAL_CONTEXT.model_dict.get( + self.generation_model_name + ) return self._generation_model @property - def rank_model(self): + def rank_model(self) -> BaseModel: if self._rank_model is None: self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name) return self._rank_model @property - def vector_store(self): + def vector_store(self) -> BaseVectorStore: if self._vector_store is None: self._vector_store = GLOBAL_CONTEXT.vector_store return self._vector_store @@ -68,3 +72,22 @@ class MemoryBaseWorker(BaseWorker): if self._monitor is None: self._monitor = GLOBAL_CONTEXT.monitor return self._monitor + + @property + def user_profile_dict(self) -> Dict[str, MemoryNode]: + if not self._user_profile_dict: + self._user_profile_dict = { + user_attr.meta_data.get("memory_key", ""): user_attr + for user_attr in self.get_context(common_constants.USER_PROFILE) + } + return self._user_profile_dict + + @property + def memory_id(self) -> str: + pass + + def __getattr__(self, key): + return self.kwargs[key] + + def get_prompt(self, x): + return x[GLOBAL_CONTEXT.global_configs["language"]] \ No newline at end of file diff --git a/memory_scope/worker/retrieve/extract_time_worker.py b/memory_scope/worker/retrieve/extract_time_worker.py index beb7a4fb..4770360f 100644 --- a/memory_scope/worker/retrieve/extract_time_worker.py +++ b/memory_scope/worker/retrieve/extract_time_worker.py @@ -1,11 +1,16 @@ import re from utils.tool_functions import time_to_formatted_str -from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP, EXTRACT_TIME_DICT +from constants.common_constants import ( + DATATIME_WORD_LIST, + DATATIME_KEY_MAP, + EXTRACT_TIME_DICT, +) from worker.memory_base_worker import MemoryBaseWorker class ExtractTimeWorker(MemoryBaseWorker): + # TODO add en version @staticmethod def get_parse_time_prompt(query: str, query_time_str: str): return f""" @@ -35,27 +40,32 @@ class ExtractTimeWorker(MemoryBaseWorker): return # prepare prompt + # TODO add en version time_format = "{year}年{month}月{day}日,{year}年第{week}周,{weekday},{hour}时{minute}分{second}秒。" - query_time_str = time_to_formatted_str(time=time_created, - date_format="", - string_format=time_format) - extract_time_prompt = self.get_parse_time_prompt(query=query, query_time_str=query_time_str) + query_time_str = time_to_formatted_str( + time=time_created, date_format="", string_format=time_format + ) + extract_time_prompt = self.get_parse_time_prompt( + query=query, query_time_str=query_time_str + ) self.logger.info(f"extract_time_prompt={extract_time_prompt}") # call sft model - self.generation_model.call(prompt=extract_time_prompt, - model_name=self.parse_time_model, - max_token=self.parse_time_max_token, - temperature=self.parse_time_temperature, - top_k=self.parse_time_top_k) + response_text = self.generation_model.call( + prompt=extract_time_prompt, + model_name=self.parse_time_model, + max_token=self.parse_time_max_token, + temperature=self.parse_time_temperature, + top_k=self.parse_time_top_k, + ) # if empty, return if not response_text: return # re-match time info to dict - pattern = r'-\s*(\S+):(\d+)' + pattern = r"-\s*(\S+):(\d+)" matches = re.findall(pattern, response_text) for key, value in matches: if key in DATATIME_KEY_MAP.keys(): diff --git a/memory_scope/worker/retrieve/fuse_rerank_worker.py b/memory_scope/worker/retrieve/fuse_rerank_worker.py index c3ca01c8..9c3f9854 100644 --- a/memory_scope/worker/retrieve/fuse_rerank_worker.py +++ b/memory_scope/worker/retrieve/fuse_rerank_worker.py @@ -1,25 +1,20 @@ from typing import Dict, List -from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \ - TIME_MATCHED +from constants.common_constants import ( + RELATED_MEMORIES, + EXTRACT_TIME_DICT, + ALL_ONLINE_NODES, + TIME_MATCHED, +) from scheme.memory_node import MemoryNode from worker.memory_base_worker import MemoryBaseWorker class FuseRerankWorker(MemoryBaseWorker): - def __init__(self, fuse_time_ratio, fuse_score_threshold, fuse_ratio_dict, *args, **kwargs): - super(FuseRerankWorker, self).__init__(*args, **kwargs) - self.fuse_score_threshold = fuse_score_threshold - self.fuse_ratio_dict = fuse_ratio_dict - # self.default_system_prompt = default_system_prompt - self.fuse_time_ratio = fuse_time_ratio - - @property - def output_max_count(self): - return self.request.user.output_max_count - @staticmethod - def format_time_infer(time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]): + def format_time_infer( + time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str] + ): if time_infer: return time_infer @@ -29,21 +24,21 @@ class FuseRerankWorker(MemoryBaseWorker): if value: time_infer += f"{value}年" elif value == "-1": - time_infer += f"每年" + time_infer += "每年" if "month" in extract_time_dict: value = meta_data.get("msg_month") if value: time_infer += f"{value}月" elif value == "-1": - time_infer += f"每月" + time_infer += "每月" if "day" in extract_time_dict: value = meta_data.get("msg_day") if value: time_infer += f"{value}日" elif value == "-1": - time_infer += f"每日" + time_infer += "每日" if "weekday" in extract_time_dict: value = meta_data.get("msg_weekday") @@ -67,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker): continue # 根据类型给ratio - type_ratio: float = self.fuse_ratio_dict.get(scheme.memory_node.memoryType, 0.1) + type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1) # 时间系数,完全匹配才行 fuse_time_ratio: float = 1.0 @@ -76,7 +71,7 @@ class FuseRerankWorker(MemoryBaseWorker): if extract_time_dict: match_event_flag = True for k, v in extract_time_dict.items(): - event_value = scheme.memory_node.metaData.get(f"event_{k}", "") + event_value = node.meta_data.get(f"event_{k}", "") if event_value in ["-1", v]: continue else: @@ -85,7 +80,7 @@ class FuseRerankWorker(MemoryBaseWorker): match_msg_flag = True for k, v in extract_time_dict.items(): - msg_value = scheme.memory_node.metaData.get(f"msg_{k}", "") + msg_value = node.meta_data.get(f"msg_{k}", "") if msg_value == v: continue else: @@ -94,32 +89,32 @@ class FuseRerankWorker(MemoryBaseWorker): if match_event_flag or match_msg_flag: fuse_time_ratio = self.fuse_time_ratio - scheme.memory_node.metaData[TIME_MATCHED] = "1" + node.meta_data[TIME_MATCHED] = "1" node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio - self.logger.info(f"content={scheme.memory_node.content} f_event={int(match_event_flag)} " - f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}") + self.logger.info( + f"content={node.content} f_event={int(match_event_flag)} " + f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}" + ) filtered_nodes.append(node) # get output & save context - filtered_nodes = sorted(filtered_nodes, key=lambda x: x.score_rerank, reverse=True) + filtered_nodes = sorted( + filtered_nodes, key=lambda x: x.score_rerank, reverse=True + ) filtered_nodes = filtered_nodes[: self.output_max_count] related_memories: List[str] = [] for node in filtered_nodes: - content = scheme.memory_node.content + content = node.content # 如果命中时间逻辑 - if scheme.memory_node.metaData.get(TIME_MATCHED, "") == "1": - # time_infer = scheme.memory_node.metaData.get(TIME_INFER) - # if not time_infer: - # time_infer = self.format_time_infer(time_infer=time_infer, - # extract_time_dict=extract_time_dict, - # meta_data=scheme.memory_node.metaData) - time_infer = self.format_time_infer(time_infer="", - extract_time_dict=extract_time_dict, - meta_data=scheme.memory_node.metaData) + if node.meta_data.get(TIME_MATCHED, "") == "1": + time_infer = self.format_time_infer( + time_infer="", + extract_time_dict=extract_time_dict, + meta_data=node.meta_data, + ) content = f"{time_infer}: {content}" related_memories.append(content) self.set_context(RELATED_MEMORIES, related_memories) - # self.set_context(DEFAULT_SYSTEM_PROMPT, self.default_system_prompt) diff --git a/memory_scope/worker/retrieve/memory_store_worker.py b/memory_scope/worker/retrieve/memory_store_worker.py index 66efef3b..4efea52d 100644 --- a/memory_scope/worker/retrieve/memory_store_worker.py +++ b/memory_scope/worker/retrieve/memory_store_worker.py @@ -1,33 +1,35 @@ from typing import List from utils.user_profile_handler import UserProfileHandler -from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE +from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED from scheme.memory_node import MemoryNode -from scheme.memory_node import MemoryNode -from node.user_attribute import UserAttribute from worker.memory_base_worker import MemoryBaseWorker + class MemoryStoreWorker(MemoryBaseWorker): def _run(self): - modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(MODIFIED_MEMORIES) + modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context( + MODIFIED_MEMORIES + ) if modified_memories: if isinstance(modified_memories[0], MemoryNode): modified_memories = [n.memory_node for n in modified_memories] for n in modified_memories: if not n.id: - n.id = f"{n.memoryId}_{n.scene}_content_{n.content}" + n.id = f"{n.memory_id}_content_{n.content}" n.code = n.id # TODO add batch insert - self.es_client.insert(n.id, body=n.model_dump(exclude=set("content_modified", ))) + n.meta_data.pop(CONTENT_MODIFIED) + self.vector_store.insert(n) else: self.logger.warning("modified_memories is empty!") - new_user_profile: List[UserAttribute] = self.get_context(NEW_USER_PROFILE) + new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE) if new_user_profile: - new_user_nodes: List[MemoryNode] = [n.memory_node for n in UserProfileHandler.to_nodes(new_user_profile)] - for n in new_user_nodes: - self.es_client.insert(n.id, body=n.model_dump(exclude=set("content_modified", ))) + for n in new_user_profile: + n.meta_data.pop(CONTENT_MODIFIED) + self.vector_store.insert(n) else: self.logger.warning("new_user_profile is empty!") diff --git a/memory_scope/worker/retrieve/parse_params_worker.py b/memory_scope/worker/retrieve/parse_params_worker.py deleted file mode 100644 index e26620a7..00000000 --- a/memory_scope/worker/retrieve/parse_params_worker.py +++ /dev/null @@ -1,36 +0,0 @@ -import json - -from config.bailian_memory_config import BailianMemoryConfig -from constants.common_constants import REQUEST, CONFIG -from pipeline.memory import MemoryServiceRequestModel -from worker.base_worker import BaseWorker - - -class ParseParamsWorker(BaseWorker): - - def _run(self): - # 参数合并 - memory_config = {} - - # 更新环境变量 - memory_config.update(self.context_handler.env_configs) - - # 更新请求参数 - request: MemoryServiceRequestModel = self.context_handler.get_context(REQUEST) - memory_config.update(request.model_dump(exclude=set("ext_info", ))) - - # 更新ext_info - if request.ext_info: - memory_config.update(request.ext_info) - - # 存入上下文 - memory_config_model: BailianMemoryConfig = BailianMemoryConfig(**memory_config) - self.context_handler.set_context(CONFIG, memory_config_model) - - # 打印 - self.logger.info(f"memory_config_model={json.dumps(memory_config_model.model_dump(), ensure_ascii=False)}") - - # 上游可能没有传这个参数,可能隐藏在memory_id做区分 - if request.user_profile: - for user_attr in request.user_profile: - user_attr.scene = request.scene diff --git a/memory_scope/worker/retrieve/semantic_rank_worker.py b/memory_scope/worker/retrieve/semantic_rank_worker.py index 02ebda4b..0fc48bc9 100644 --- a/memory_scope/worker/retrieve/semantic_rank_worker.py +++ b/memory_scope/worker/retrieve/semantic_rank_worker.py @@ -1,9 +1,8 @@ from typing import List, Dict -from utils.user_profile_handler import UserProfileHandler from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \ QUERY_KEYWORDS -from enumeration.memory_recall_type import MemoryRecallType +from enumeration.memory_recall_enum import MemoryRecallType from scheme.memory_node import MemoryNode from worker.memory_base_worker import MemoryBaseWorker @@ -11,29 +10,29 @@ from worker.memory_base_worker import MemoryBaseWorker class SemanticRankWorker(MemoryBaseWorker): def user_profile_to_nodes(self) -> List[MemoryNode]: - user_profile_nodes: List[MemoryNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True) + user_profile_nodes: List[MemoryNode] = self.user_profile_dict for node in user_profile_nodes: # 从画像侧召回 - scheme.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE - self.logger.info(f"user profile node={scheme.memory_node.content}") + node.meta_data[RECALL_TYPE] = MemoryRecallType.PROFILE + self.logger.info(f"user profile node={node.content}") return user_profile_nodes def _run(self): all_node_dict: Dict[str, MemoryNode] = {} - # 优先级: similar_obs_nodes < < profile_nodes + # 优先级: similar_obs_nodes < profile_nodes similar_obs_nodes: List[MemoryNode] = self.get_context(SIMILAR_OBS_NODES) if similar_obs_nodes: for node in similar_obs_nodes: - all_node_dict[scheme.memory_node.content] = node + all_node_dict[node.content] = node profile_nodes: List[MemoryNode] = self.user_profile_to_nodes() if profile_nodes: for node in profile_nodes: - all_node_dict[scheme.memory_node.content] = node + all_node_dict[node.content] = node if not all_node_dict: - self.add_run_info(f"all_node_dict is empty!", continue_run=False) + self.add_run_info("all_node_dict is empty!", continue_run=False) return # call recall model @@ -45,18 +44,18 @@ class SemanticRankWorker(MemoryBaseWorker): query_keyword_join = ",".join(query_keywords) query = f"{query} 用户的{query_keyword_join}。" documents = list(all_node_dict.keys()) - result = self.rerank_client.call(query=query, documents=documents) + result = self.rank_model.call(query=query, documents=documents) if not result: - self.add_run_info(f"semantic call recall model failed!") + self.add_run_info("semantic call recall model failed!") return # set score - for rank_node in result: - content = documents[rank_node["index"]] + for index, score in result.rank_scores.items(): + content = documents[index] node = all_node_dict[content] - node.score_rank = rank_node["relevance_score"] - self.logger.info(f"query={query} content={scheme.memory_node.content} score_rank={node.score_rank}") + node.score_rank = score + self.logger.info(f"query={query} content={node.content} score_rank={node.score_rank}") # save context all_online_nodes: List[MemoryNode] = list(all_node_dict.values()) diff --git a/memory_scope/worker/summary_long/__init__.py b/memory_scope/worker/summary_long/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/worker/summary_long/get_insight_worker.py b/memory_scope/worker/summary_long/get_insight_worker.py new file mode 100644 index 00000000..c7be74e8 --- /dev/null +++ b/memory_scope/worker/summary_long/get_insight_worker.py @@ -0,0 +1,166 @@ +from datetime import datetime +from typing import List + +from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict +from ...constants.common_constants import ( + NEW_INSIGHT_NODES, + DT, + NOT_REFLECTED_MERGE_NODES, + NEW_INSIGHT_KEYS, + INSIGHT_KEY, + INSIGHT_VALUE, + REFLECTED, + CONTENT_MODIFIED +) +from ...enumeration.memory_status_enum import MemoryNodeStatus +from ...enumeration.memory_type_enum import MemoryTypeEnum +from ...scheme.memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.get_insight_prompt import ( + GET_INSIGHT_FEW_SHOT_PROMPT, + GET_INSIGHT_SYSTEM_PROMPT, + GET_INSIGHT_USER_QUERY_PROMPT +) + + +class GetInsightWorker(MemoryBaseWorker): + def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode: + created_dt = datetime.now() + dt = time_to_formatted_str(time=created_dt) + + # 组合meta_data + meta_data = { + DT: dt, + INSIGHT_KEY: insight_key, + INSIGHT_VALUE: insight_value, + CONTENT_MODIFIED: True, # 新增的insight需要置为true + } + meta_data.update( + {k: str(v) for k, v in get_datetime_info_dict(created_dt).items()} + ) + + content = f"用户的{insight_key}:{insight_value}" + return MemoryNode( + content=content, + memory_id=self.memory_id, + memory_type=MemoryTypeEnum.INSIGHT.value, + meta_data=meta_data, + status=MemoryNodeStatus.ACTIVE.value, + ) + + def reflect_new_insight_key( + self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode] + ) -> MemoryNode | None: + + # 检索历史memory + hits = self.vector_store.similar_search( + text=insight_key, + size=self.es_insight_similar_top_k, + exact_filters={ + "memory_id": self.memory_id, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [ + MemoryTypeEnum.OBSERVATION.value, + MemoryTypeEnum.OBS_CUSTOMIZED.value, + ], + }, + ) + + # 转化成 MemoryNodeWrap 合并新增nodes + related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits] + related_nodes.extend(not_reflected_merge_nodes) + + # content去重 + related_node_dict = {n.memory_node.content: n for n in related_nodes} + related_nodes = sorted( + list(related_node_dict.values()), key=lambda x: x.memory_node.id + ) + documents = [n.memory_node.content for n in related_nodes] + + # 重排所有记忆 + result = self.rank_model.call(query=insight_key, documents=documents) + if not result: + self.add_run_info( + f"reflect insight_key={insight_key} call rerank client failed!" + ) + return + + # 根据打分过滤 + for rank_node in result: + index = rank_node["index"] + score = rank_node["relevance_score"] + related_nodes[index].score_rank = score + related_nodes_sorted = sorted( + related_nodes, key=lambda x: x.score_rank, reverse=True + )[: self.insight_obs_max_cnt] + + # 生成prompt + user_query_list = [x.memory_node.content for x in related_nodes_sorted] + get_insight_message = self.prompt_to_msg( + system_prompt=self.get_prompt(GET_INSIGHT_SYSTEM_PROMPT), + few_shot=self.get_prompt(GET_INSIGHT_FEW_SHOT_PROMPT), + user_query=self.get_prompt(GET_INSIGHT_USER_QUERY_PROMPT).format( + insight_key=insight_key, user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"get_insight_message={get_insight_message}") + + # call LLM, 提取insight + response_text = self.generation_model.call( + messages=get_insight_message, + model_name=self.get_insight_model, + max_token=self.get_insight_max_token, + temperature=self.get_insight_temperature, + top_k=self.get_insight_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("reflect_upon_user_attr call llm failed!") + return + response_text = response_text.strip() + if response_text in ["无"]: + return + return self.new_insight_node( + insight_key=insight_key, insight_value=response_text + ) + + def _run(self): + new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS) + if not new_insight_keys: + self.add_run_info("new_insight_keys is empty! stop insight.") + return + + not_reflected_merge_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_MERGE_NODES + ) + if not not_reflected_merge_nodes: + self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.") + return + + # submit insight task + for insight_key in new_insight_keys: + self.submit_thread( + self.reflect_new_insight_key, + sleep_time=1, + insight_key=insight_key, + not_reflected_merge_nodes=not_reflected_merge_nodes, + ) + + # save output + new_insight_nodes: List[MemoryNode] = [] + for result in self.join_threads(): + if result: + new_insight_nodes.append(result) + assert isinstance(result, MemoryNode) + insight_key = result.meta_data.get(INSIGHT_KEY, "") + insight_value = result.meta_data.get(INSIGHT_VALUE, "") + self.logger.info( + f"after_get_insight insight_key={insight_key} insight_value={insight_value}" + ) + + self.set_context(NEW_INSIGHT_NODES, new_insight_nodes) + + # set REFLECTED + for node in not_reflected_merge_nodes: + node.meta_data[REFLECTED] = "1" diff --git a/memory_scope/worker/summary_long/get_reflection_worker.py b/memory_scope/worker/summary_long/get_reflection_worker.py new file mode 100644 index 00000000..1647c9b6 --- /dev/null +++ b/memory_scope/worker/summary_long/get_reflection_worker.py @@ -0,0 +1,99 @@ +from typing import List + +from ...utilsresponse_text_parser import ResponseTextParser +from ...constants.common_constants import ( + NEW_OBS_NODES, + NOT_REFLECTED_OBS_NODES, + REFLECTED, + INSIGHT_NODES, + INSIGHT_KEY, + NEW_INSIGHT_KEYS, + NOT_REFLECTED_MERGE_NODES, +) +from ...scheme.memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.get_reflection_prompt import ( + GET_REFLECTION_FEW_SHOT_PROMPT, + GET_REFLECTION_SYSTEM_PROMPT, + GET_REFLECTION_USER_QUERY_PROMPT +) + + +class GetReflectionWorker(MemoryBaseWorker): + def _run(self): + # 过滤得到 not_reflected_merge_nodes + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + not_reflected_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_OBS_NODES + ) + not_reflected_merge_nodes: List[MemoryNode] = [] + if new_obs_nodes: + not_reflected_merge_nodes.extend(new_obs_nodes) + if not_reflected_nodes: + not_reflected_merge_nodes.extend(not_reflected_nodes) + not_reflected_merge_nodes = [ + node + for node in not_reflected_merge_nodes + if node.meta_data.get(REFLECTED, "") == "0" + ] + + # count + not_reflected_count = len(not_reflected_merge_nodes) + if not_reflected_count <= self.reflect_obs_cnt_threshold: + self.logger.info( + f"not_reflected_count={not_reflected_count} is not enough, stop reflect." + ) + return + + # save context + self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes) + + # get profile_keys + exist_keys: List[str] = [] + profile_keys: List[str] = list(self.user_profile_dict.keys()) + exist_keys.extend(profile_keys) + self.logger.info(f"profile_keys={profile_keys}") + + # get insight_keys + insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + if insight_nodes: + insight_keys = [ + n.meta_data.get(INSIGHT_KEY) for n in insight_nodes + ] + insight_keys = [x.strip() for x in insight_keys if x] + exist_keys.extend(insight_keys) + self.logger.info(f"insight_keys={insight_keys}") + + # gen reflect prompt + user_query_list = [n.content for n in not_reflected_merge_nodes] + reflect_message = self.prompt_to_msg( + system_prompt=self.get_prompt(GET_REFLECTION_SYSTEM_PROMPT).format( + num_questions=self.reflect_num_questions + ), + few_shot=self.get_prompt(GET_REFLECTION_FEW_SHOT_PROMPT), + user_query=self.get_prompt(GET_REFLECTION_USER_QUERY_PROMPT).format( + exist_keys=",".join(exist_keys), user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"reflect_message={reflect_message}") + + # # call LLM + response_text = self.generation_model.call( + messages=reflect_message, + model_name=self.reflect_obs_model, + max_token=self.reflect_obs_max_token, + temperature=self.reflect_obs_temperature, + top_k=self.reflect_obs_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("reflect_obs_questions call llm failed!") + return + + # parse text & save + new_insight_keys = ResponseTextParser(response_text).parse_v2( + "get_insight_keys" + ) + if new_insight_keys: + self.set_context(NEW_INSIGHT_KEYS, new_insight_keys) diff --git a/memory_scope/worker/summary_long/long_contra_repeat_worker.py b/memory_scope/worker/summary_long/long_contra_repeat_worker.py new file mode 100644 index 00000000..42f2705e --- /dev/null +++ b/memory_scope/worker/summary_long/long_contra_repeat_worker.py @@ -0,0 +1,129 @@ +from typing import List + +from ...utils.response_text_parser import ResponseTextParser +from ...constants.common_constants import ( + NEW_OBS_NODES, + MSG_TIME, + MODIFIED_MEMORIES, +) +from ...enumeration.memory_status_enum import MemoryNodeStatus +from ...enumeration.memory_type_enum import MemoryTypeEnum +from ...scheme.memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.long_contra_repeat_prompt import ( + LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT, + LONG_CONTRA_REPEAT_SYSTEM_PROMPT, + LONG_CONTRA_REPEAT_USER_QUERY_PROMPT, +) + + +class LongContraRepeatWorker(MemoryBaseWorker): + + def _run(self): + # 合并当前的obs和今日的obs + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + all_obs_nodes: List[MemoryNode] = [] + for new_obs_node in new_obs_nodes: + text = new_obs_node.content + related_nodes = self.vector_store.similar_search( + text=text, + size=self.es_contra_repeat_similar_top_k, + exact_filters={ + "memory_id": self.memory_id, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [ + MemoryTypeEnum.OBSERVATION.value, + MemoryTypeEnum.OBS_CUSTOMIZED.value, + ], + }, + ) + + has_match = False + for related_node in related_nodes: + if related_node.score_similar < self.long_contra_repeat_threshold: + continue + else: + has_match = True + all_obs_nodes.append(related_node) + if has_match: + all_obs_nodes.append(new_obs_node) + + if not all_obs_nodes: + self.add_run_info("all_obs_nodes is empty!") + return + + # gene prompt + user_query_list = [] + all_obs_nodes = sorted( + all_obs_nodes, + key=lambda x: x.meta_data.get(MSG_TIME, ""), + reverse=True, + ) + for i, n in enumerate(all_obs_nodes): + user_query_list.append(f"{i + 1} {n.content}") + merge_obs_message = self.prompt_to_msg( + system_prompt=self.get_prompt(LONG_CONTRA_REPEAT_SYSTEM_PROMPT).format( + num_obs=len(user_query_list) + ), + few_shot=self.get_prompt(LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT), + user_query=self.get_prompt(LONG_CONTRA_REPEAT_USER_QUERY_PROMPT).format( + user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"merge_obs_message={merge_obs_message}") + + # call LLM + response_text = self.generation_model.call( + messages=merge_obs_message, + model_name=self.merge_obs_model, + max_token=self.merge_obs_max_token, + temperature=self.merge_obs_temperature, + top_k=self.merge_obs_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("contra repeat call llm failed!") + return + + # parse text + idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat") + if len(idx_merge_obs_list) <= 0: + self.add_run_info("idx_merge_obs_list is empty!") + return + + # add merged obs + merge_obs_nodes: List[MemoryNode] = [] + for obs_content_list in idx_merge_obs_list: + if not obs_content_list: + continue + + # [6, 逃课] + if len(obs_content_list) != 2: + self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") + continue + + idx, keep_flag = obs_content_list + + if not idx.isdigit(): + self.logger.warning(f"idx={idx} is invalid!") + continue + + # 序号需要修正-1 + idx = int(idx) - 1 + if idx >= len(all_obs_nodes): + self.logger.warning(f"idx={idx} is invalid!") + continue + + if keep_flag not in ["矛盾", "被包含", "无"]: + self.logger.warning(f"keep_flag={keep_flag} is invalid!") + continue + + node: MemoryNode = all_obs_nodes[idx] + if keep_flag != "无": + node.status = MemoryNodeStatus.EXPIRED.value + merge_obs_nodes.append(node) + self.logger.info(f"after contra repeat: {node.content} {node.status}") + + # save context + self.set_context(MODIFIED_MEMORIES, merge_obs_nodes) diff --git a/memory_scope/worker/summary_long/summary_collect_worker.py b/memory_scope/worker/summary_long/summary_collect_worker.py new file mode 100644 index 00000000..62e0699b --- /dev/null +++ b/memory_scope/worker/summary_long/summary_collect_worker.py @@ -0,0 +1,47 @@ +from typing import List, Dict + +from ...constants.common_constants import ( + NEW_INSIGHT_NODES, + MODIFIED_MEMORIES, + INSIGHT_NODES, + NEW_OBS_NODES, + NOT_REFLECTED_OBS_NODES, + NEW, + NOT_REFLECTED_MERGE_NODES, + CONTENT_MODIFIED, +) +from ...scheme.memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker + + +class SummaryCollectWorker(MemoryBaseWorker): + + def _run(self): + insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES) + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + not_reflected_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_OBS_NODES + ) + not_reflected_merge_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_MERGE_NODES + ) + + # 合并逻辑,复杂,务必check + all_node_dict: Dict[str, MemoryNode] = {} + if insight_nodes: + all_node_dict.update( + {n.id: n for n in insight_nodes if n.meta_data.get(CONTENT_MODIFIED, False)} + ) + if new_insight_nodes: + all_node_dict.update({n.content: n for n in new_insight_nodes}) + if new_obs_nodes: + # 设置为非新 + for n in new_obs_nodes: + n.meta_data[NEW] = "0" + all_node_dict.update({n.content: n for n in new_obs_nodes}) + if not_reflected_merge_nodes and not_reflected_nodes: + # 进入reflect阶段 + all_node_dict.update({n.id: n for n in not_reflected_nodes}) + + self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values())) diff --git a/memory_scope/worker/summary_long/update_insight_worker.py b/memory_scope/worker/summary_long/update_insight_worker.py new file mode 100644 index 00000000..93e246a0 --- /dev/null +++ b/memory_scope/worker/summary_long/update_insight_worker.py @@ -0,0 +1,177 @@ +from typing import List + +from ...utils.response_text_parser import ResponseTextParser +from ...constants.common_constants import ( + INSIGHT_NODES, + NEW_OBS_NODES, + INSIGHT_KEY, + INSIGHT_VALUE, + CONTENT_MODIFIED, +) +from ...scheme.memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.update_insight_prompt import ( + UPDATE_INSIGHT_FEW_SHOT_PROMPT, + UPDATE_INSIGHT_SYSTEM_PROMPT, + UPDATE_INSIGHT_USER_QUERY_PROMPT, +) + + +class UpdateInsightWorker(MemoryBaseWorker): + + def filter_obs_nodes( + self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode] + ) -> (MemoryNode, List[MemoryNode], float): + max_score: float = 0 + filtered_nodes: List[MemoryNode] = [] + + insight_key = insight_node.meta_data.get(INSIGHT_KEY, "") + insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "") + if not insight_key or not insight_value: + self.logger.warning( + f"insight_key={insight_key} insight_value={insight_value} is empty!" + ) + return insight_node, filtered_nodes, max_score + + result = self.rank_model.call( + query=insight_key, documents=[x.content for x in new_obs_nodes] + ) + + if not result: + self.add_run_info(f"update_insight={insight_key} call rerank failed!") + return insight_node, filtered_nodes, max_score + + # 找到大于阈值的obs node + + for index, score in result.rank_scores.items(): + node = new_obs_nodes[index] + keep_flag = "filtered" + if score >= self.update_insight_threshold: + filtered_nodes.append(node) + keep_flag = "keep" + max_score = max(max_score, score) + self.logger.info( + f"insight_key={insight_key} insight_value={insight_value} " + f"score={score} keep_flag={keep_flag}" + ) + + if not filtered_nodes: + self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!") + + return insight_node, filtered_nodes, max_score + + def update_insight( + self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode] + ) -> MemoryNode: + + insight_key = insight_node.meta_data.get(INSIGHT_KEY, "") + insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "") + self.logger.info( + f"update_insight insight_key={insight_key} insight_value={insight_value} " + f"doc.size={len(filtered_nodes)}" + ) + + # gen prompt + user_query_list = [] + for node in filtered_nodes: + user_query_list.append(f"句子:{node.content}") + update_insight_message = self.prompt_to_msg( + system_prompt=self.get_prompt(UPDATE_INSIGHT_SYSTEM_PROMPT), + few_shot=self.get_prompt(UPDATE_INSIGHT_FEW_SHOT_PROMPT), + user_query=self.get_prompt(UPDATE_INSIGHT_USER_QUERY_PROMPT).format( + user_query="\n".join(user_query_list), + insight_key=insight_key, + insight_key_value=insight_key + ":" + insight_value, + ), + ) + self.logger.info(f"update_insight_message={update_insight_message}") + + # call LLM + response_text: str = self.generation_model.call( + messages=update_insight_message, + model_name=self.update_insight_model, + max_token=self.update_insight_max_token, + temperature=self.update_insight_temperature, + top_k=self.update_insight_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info( + f"update_insight insight_key={insight_key} call llm failed!" + ) + return insight_node + + profile_list = ResponseTextParser(response_text).parse_v1( + f"update_profile {insight_key}" + ) + if not profile_list: + self.add_run_info( + f"update_insight insight_key={insight_key} profile_list empty 1!" + ) + return insight_node + profile_list = profile_list[0] + if not profile_list: + self.add_run_info( + f"update_insight insight_key={insight_key} profile_list empty 2" + ) + return insight_node + insight_value = profile_list[0] + + if not insight_value or insight_value in ["无", "重复"]: + self.logger.info(f"insight_value={insight_value}, skip.") + return insight_node + + insight_node.meta_data[INSIGHT_VALUE] = insight_value + insight_node.meta_data[CONTENT_MODIFIED] = True + return insight_node + + def _run(self): + # 获取新的obs和insight + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + if not new_obs_nodes: + self.logger.info("new_obs_nodes is empty, stop update sights!") + return + if not insight_nodes: + self.logger.info("insight_nodes is empty, stop update sights!") + return + + # 提交打分任务 + for node in insight_nodes: + self.submit_thread( + self.filter_obs_nodes, + sleep_time=0.1, + insight_node=node, + new_obs_nodes=new_obs_nodes, + ) + + # 选择topN + result_list = [] + for result in self.join_threads(): + insight_node, filtered_nodes, max_score = result + if not filtered_nodes: + continue + result_list.append(result) + result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True) + if len(result_sorted) > self.update_insight_max_thread: + result_sorted = result_sorted[: self.update_insight_max_thread] + + # 提交LLM update任务 + for insight_node, filtered_nodes, _ in result_sorted: + self.submit_thread( + self.update_insight, + sleep_time=1, + insight_node=insight_node, + filtered_nodes=filtered_nodes, + ) + + # 等待结果 + for result in self.join_threads(): + if result: + insight_node: MemoryNode = result + insight_key = insight_node.meta_data.get(INSIGHT_KEY, "") + insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "") + self.logger.info( + f"after_update_insight insight_key={insight_key} insight_value={insight_value}" + ) diff --git a/memory_scope/worker/summary_long/update_profile_worker.py b/memory_scope/worker/summary_long/update_profile_worker.py new file mode 100644 index 00000000..de2b9061 --- /dev/null +++ b/memory_scope/worker/summary_long/update_profile_worker.py @@ -0,0 +1,241 @@ +from typing import List + +from ...utils.response_text_parser import ResponseTextParser +from ...constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE +from ...enumeration.memory_type_enum import MemoryTypeEnum +from ....memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.update_profile_prompt import ( + UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT, + UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT, + UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT, + UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT, + UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT, + UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT +) +from ...chat.global_context import GlobalContext + + +class UpdateProfileWorker(MemoryBaseWorker): + @property + def extra_user_attrs(self): + return GlobalContext.global_configs.get("extra_user_attrs", []) + + def filter_obs_nodes( + self, user_attr: MemoryNode, new_obs_nodes: List[MemoryNode] + ) -> (MemoryNode, List[MemoryNode], float): + max_score: float = 0 + filtered_nodes: List[MemoryNode] = [] + result = self.rank_model.call( + query=user_attr.meta_data.get("description", ""), + documents=[x.content for x in new_obs_nodes], + ) + + if not result: + self.add_run_info( + f"update_user_attr={user_attr.meta_data.get("memory_key", "")} call rerank failed!" + ) + return user_attr, filtered_nodes, max_score + + # 找到大于阈值的obs node + filtered_nodes: List[MemoryNode] = [] + for index, score in result.rank_scores.items(): + node = new_obs_nodes[index] + keep_flag = "filtered" + if score >= self.update_profile_threshold: + filtered_nodes.append(node) + keep_flag = "keep" + max_score = max(max_score, score) + self.logger.info( + f"key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} " + f"content={node.content} score={score} keep_flag={keep_flag}" + ) + + if not filtered_nodes: + self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!") + return user_attr, filtered_nodes, max_score + + def update_user_attr( + self, user_attr: MemoryNode, filtered_nodes: List[MemoryNode] + ) -> MemoryNode: + self.logger.info( + f"update_user_attr memory_key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} " + f"value={user_attr.meta_data.get("value", "")} doc.size={len(filtered_nodes)}" + ) + + # 根据不同的参数类型是否多值,分别给出prompt + user_query_list = [] + for node in filtered_nodes: + user_query_list.append(f"句子:{node.content}") + update_profile = f"{user_attr.meta_data.get("memory_key", "")}({user_attr.meta_data.get("description", "")})" + update_profile_value = update_profile + ":" + ",".join(user_attr.meta_data.get("value", "")) + + if user_attr.meta_data.get("is_unique", 0) == 1: + update_profile_message = self.prompt_to_msg( + system_prompt=self.get_prompt(UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT), + few_shot=self.get_prompt(UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT), + user_query=self.get_prompt(UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT).format( + user_query="\n".join(user_query_list), + update_profile=update_profile, + update_profile_value=update_profile_value, + ), + ) + else: + update_profile_message = self.prompt_to_msg( + system_prompt=self.get_prompt(UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT), + few_shot=self.get_prompt(UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT), + user_query=self.get_prompt(UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT).format( + user_query="\n".join(user_query_list), + update_profile=update_profile, + update_profile_value=update_profile_value, + ), + ) + self.logger.info(f"update_profile_message={update_profile_message}") + + # call LLM + response_text: str = self.generation_model.call( + messages=update_profile_message, + model_name=self.update_profile_model, + max_token=self.update_profile_max_token, + temperature=self.update_profile_temperature, + top_k=self.update_profile_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info( + f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} call llm failed!" + ) + return user_attr + + profile_list = ResponseTextParser(response_text).parse_v1( + f"update_attr {user_attr.meta_data.get("memory_key", "")}" + ) + if not profile_list: + self.add_run_info( + f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 1!" + ) + return user_attr + profile_list = profile_list[0] + if not profile_list: + self.add_run_info( + f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 2" + ) + return user_attr + profile = profile_list[0] + + if not profile or profile in ["无", "重复"]: + self.logger.info(f"profile={profile}, skip.") + return user_attr + + # check 英文中午逗号 + if user_attr.meta_data.get("is_unique", 0) == 1: + user_attr.meta_data["value"] = [profile.strip()] + else: + attr_value_list = profile.replace(",", ",").split(",") + user_attr.meta_data["value"] = [ + x.strip() for x in sorted(list(set(user_attr.meta_data.get("value", "") + attr_value_list))) + ] + return user_attr + + def add_extra_user_attrs(self): + # 解析为空返回 + extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()] + if not extra_user_attr_list: + return + + for user_attr_info in extra_user_attr_list: + user_attr_split = user_attr_info.split(":") + + # 格式不对返回 + if len(user_attr_split) < 1: + continue + user_attr_key = user_attr_split[0] + + user_attr_desc = "" + if len(user_attr_split) >= 2: + user_attr_desc = user_attr_split[1] + + user_attr_unique = 0 + if len(user_attr_split) >= 3: + user_attr_unique = int(user_attr_split[2]) + + # 已经包含返回 + if user_attr_key in self.user_profile_dict: + user_attr = self.user_profile_dict[user_attr_key] + # description为空,补充description + if not user_attr.meta_data.get("description", ""): + user_attr.meta_data["description"] = user_attr_desc + continue + + # 增加新属性 + new_attr = MemoryNode( + memory_id=self.memory_id, + meta_data={ + "memory_key": user_attr_key, + "is_unique": int(user_attr_unique), + "is_mutable": 1, + "description": user_attr_desc + }, + memory_type=MemoryTypeEnum.PROFILE, + status=1, + ) + self.user_profile_dict[user_attr_key] = new_attr + + def _run(self): + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + if not new_obs_nodes: + self.logger.info("new_obs_nodes is empty, stop user profile!") + self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values())) + return + + # 增加环境变量配置的属性 + if self.extra_user_attrs: + self.add_extra_user_attrs() + + new_user_profile: List[MemoryNode] = [] + self.set_context(NEW_USER_PROFILE, new_user_profile) + + for user_attr_key, user_attr in self.user_profile_dict.items(): + # 不可修改直接跳过 + if user_attr.meta_data.get("is_mutable", 0) != 1: + new_user_profile.append(user_attr) + self.logger.info(f"{user_attr_key} is not mutable! continue") + continue + + self.submit_thread( + self.filter_obs_nodes, + sleep_time=0.1, + user_attr=user_attr, + new_obs_nodes=new_obs_nodes, + ) + + # 选择topN + result_list = [] + for result in self.join_threads(): + user_attr, filtered_nodes, max_score = result + if not filtered_nodes: + continue + result_list.append(result) + result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True) + if len(result_sorted) > self.update_profile_max_thread: + result_sorted = result_sorted[: self.update_profile_max_thread] + + # 提交LLM update任务 + for user_attr, filtered_nodes, _ in result_sorted: + self.submit_thread( + self.update_user_attr, + sleep_time=1, + user_attr=user_attr, + filtered_nodes=filtered_nodes, + ) + + # collect result & save + for result in self.join_threads(): + if result: + user_attribute: MemoryNode = result + self.logger.info( + f"after_update_profile memory_key={user_attribute.meta_data.get("memory_key", "")} " + f"desc={user_attribute.meta_data.get("description", "")} value={user_attribute.meta_data.get("value", "")}" + ) + new_user_profile.append(user_attribute) diff --git a/memory_scope/worker/summary_short/__init__.py b/memory_scope/worker/summary_short/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/worker/summary_short/contra_repeat_worker.py b/memory_scope/worker/summary_short/contra_repeat_worker.py new file mode 100644 index 00000000..c1168d58 --- /dev/null +++ b/memory_scope/worker/summary_short/contra_repeat_worker.py @@ -0,0 +1,117 @@ +from typing import List + +from ...utils.response_text_parser import ResponseTextParser +from ...constants.common_constants import ( + NEW_OBS_NODES, + TODAY_OBS_NODES, + MSG_TIME, + NEW_OBS_WITH_TIME_NODES, + MODIFIED_MEMORIES, +) +from ...enumeration.memory_status_enum import MemoryNodeStatus +from ...scheme.memory_node import MemoryNode +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.contra_repeat_prompt import ( + CONTRA_REPEAT_FEW_SHOT_PROMPT, + CONTRA_REPEAT_SYSTEM_PROMPT, + CONTRA_REPEAT_USER_QUERY_PROMPT, +) + + +class ContraRepeatWorker(MemoryBaseWorker): + + def _run(self): + # 合并当前的obs和今日的obs + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + new_obs_with_time_nodes: List[MemoryNode] = self.get_context( + NEW_OBS_WITH_TIME_NODES + ) + today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES) + all_obs_nodes: List[MemoryNode] = [] + if new_obs_nodes: + all_obs_nodes.extend(new_obs_nodes) + if new_obs_with_time_nodes: + all_obs_nodes.extend(new_obs_with_time_nodes) + if today_obs_nodes: + all_obs_nodes.extend(today_obs_nodes) + if not all_obs_nodes: + self.add_run_info("all_obs_nodes is empty!") + return + + # gene prompt + user_query_list = [] + all_obs_nodes = sorted( + all_obs_nodes, + key=lambda x: x.meta_data.get(MSG_TIME, ""), + reverse=True, + ) + for i, n in enumerate(all_obs_nodes): + user_query_list.append(f"{i + 1} {n.content}") + merge_obs_message = self.prompt_to_msg( + system_prompt=self.get_prompt(CONTRA_REPEAT_SYSTEM_PROMPT).format( + num_obs=len(user_query_list) + ), + few_shot=self.get_prompt(CONTRA_REPEAT_FEW_SHOT_PROMPT), + user_query=self.get_prompt(CONTRA_REPEAT_USER_QUERY_PROMPT).format( + user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"merge_obs_message={merge_obs_message}") + + # call LLM + response_text = self.generation_model.call( + messages=merge_obs_message, + model_name=self.merge_obs_model, + max_token=self.merge_obs_max_token, + temperature=self.merge_obs_temperature, + top_k=self.merge_obs_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("contra repeat call llm failed!") + return + + # parse text + idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat") + if len(idx_merge_obs_list) <= 0: + self.add_run_info("idx_merge_obs_list is empty!") + return + + # add merged obs + merge_obs_nodes: List[MemoryNode] = [] + for obs_content_list in idx_merge_obs_list: + if not obs_content_list: + continue + + # [6, 逃课] + if len(obs_content_list) != 2: + self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") + continue + + idx, keep_flag = obs_content_list + + if not idx.isdigit(): + self.logger.warning(f"idx={idx} is invalid!") + continue + + # 序号需要修正-1 + idx = int(idx) - 1 + if idx >= len(all_obs_nodes): + self.logger.warning(f"idx={idx} is invalid!") + continue + + if keep_flag not in ["矛盾", "被包含", "无"]: + self.logger.warning(f"keep_flag={keep_flag} is invalid!") + continue + + node: MemoryNode = all_obs_nodes[idx] + if keep_flag != "无": + node.status = MemoryNodeStatus.EXPIRED.value + merge_obs_nodes.append(node) + self.logger.info( + f"after contra repeat: {node.content} {node.status}" + ) + + # save context + self.set_context(MODIFIED_MEMORIES, merge_obs_nodes) diff --git a/memory_scope/worker/summary_short/get_observation_with_time_worker.py b/memory_scope/worker/summary_short/get_observation_with_time_worker.py new file mode 100644 index 00000000..0ff9aeb4 --- /dev/null +++ b/memory_scope/worker/summary_short/get_observation_with_time_worker.py @@ -0,0 +1,167 @@ +from datetime import datetime +from typing import List + +from ...utils.response_text_parser import ResponseTextParser +from ...utils.tool_functions import ( + time_to_formatted_str, + get_datetime_info_dict, + extract_date_parts, +) +from ...constants.common_constants import ( + REFLECTED, + DT, + TIME_INFER, + NEW, + MSG_TIME, + KEY_WORD, + DATATIME_WORD_LIST, + NEW_OBS_WITH_TIME_NODES, + CONTENT_MODIFIED, +) +from ...enumeration.memory_status_enum import MemoryNodeStatus +from ...enumeration.memory_type_enum import MemoryTypeEnum +from ...scheme.memory_node import MemoryNode +from ...scheme.message import Message +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.get_observation_with_time_prompt import ( + GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT, + GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT, + GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT, +) + + +class GetObservationWithTimeWorker(MemoryBaseWorker): + + def add_observation( + self, message: Message, obs_content: str, time_infer: str, keywords: str + ): + created_dt: datetime = datetime.fromtimestamp(float(message.time_created)) + dt = time_to_formatted_str(time=created_dt) + + # 组合meta_data + meta_data = { + MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话 + REFLECTED: "0", # reflect标记 + DT: dt, # 当天标记 + NEW: "1", # summary-long标记 + MSG_TIME: message.time_created, # 对话时间 + TIME_INFER: time_infer, # 推断的时间 + KEY_WORD: keywords, # 关键词 + CONTENT_MODIFIED: True, # 新增的obs需要置为true + } + + # 事件时间 + meta_data.update( + {f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()} + ) + # 对话时间 + meta_data.update( + {f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()} + ) + + return MemoryNode.init_from_attrs( + content=obs_content, + memory_id=self.memory_id, + memory_type=MemoryTypeEnum.OBSERVATION.value, + meta_data=meta_data, + status=MemoryNodeStatus.ACTIVE.value, + ) + + def _run(self): + # gene prompt + user_query_list = [] + i = 1 + for msg in self.messages: + match = False + for time_keyword in DATATIME_WORD_LIST: + if time_keyword in msg.content: + match = True + break + if match: + dt = time_to_formatted_str( + time=msg.time_created, + date_format="", + string_format="{year}年{month}月{day}日{weekday}{hour}点", + ) + user_query_list.append(f"{i} {dt} 用户:{msg.content}") + i += 1 + + if not user_query_list: + self.add_run_info( + f"get obs with time user_query_list={user_query_list} is empty" + ) + return + + obtain_obs_message = self.prompt_to_msg( + system_prompt=self.get_prompt( + GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT + ).format(num_obs=len(user_query_list)), + few_shot=self.get_prompt(GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT), + user_query=self.get_prompt( + GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT + ).format(user_query="\n".join(user_query_list)), + ) + self.logger.info(f"obtain_obs_message={obtain_obs_message}") + + # call LLM + response_text: str = self.generation_model.call( + messages=obtain_obs_message, + model_name=self.summary_messages_model, + max_token=self.summary_messages_max_token, + temperature=self.summary_messages_temperature, + top_k=self.summary_messages_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("summary call llm failed!", continue_run=False) + return + + # parse text + idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time") + if len(idx_obs_list) <= 0: + self.add_run_info("idx_obs_list is empty!", continue_run=False) + return + + # gene new obs nodes + new_obs_nodes: List[MemoryNode] = [] + for obs_content_list in idx_obs_list: + if not obs_content_list: + continue + + # [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游] + if len(obs_content_list) != 4: + self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") + continue + + idx, time_infer, obs_content, keywords = obs_content_list + + if obs_content in ["无", "重复"]: + continue + + if not idx.isdigit(): + self.logger.warning(f"idx={idx} is invalid!") + continue + + if time_infer == "无": + time_infer = "" + + # 序号需要修正-1 + idx = int(idx) - 1 + if idx >= len(self.messages): + self.logger.warning( + f"idx={idx} is invalid! messages.size={len(self.messages)}" + ) + continue + + new_obs_nodes.append( + self.add_observation( + message=self.messages[idx], + obs_content=obs_content, + time_infer=time_infer, + keywords=keywords, + ) + ) + + # save context + self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes) diff --git a/memory_scope/worker/summary_short/get_observation_worker.py b/memory_scope/worker/summary_short/get_observation_worker.py new file mode 100644 index 00000000..63c74339 --- /dev/null +++ b/memory_scope/worker/summary_short/get_observation_worker.py @@ -0,0 +1,144 @@ +from datetime import datetime +from typing import List + +from ...utils.response_text_parser import ResponseTextParser +from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict +from ...constants.common_constants import ( + REFLECTED, + DT, + NEW_OBS_NODES, + TIME_INFER, + NEW, + MSG_TIME, + KEY_WORD, + DATATIME_WORD_LIST, + CONTENT_MODIFIED, +) +from ...enumeration.memory_status_enum import MemoryNodeStatus +from ...enumeration.memory_type_enum import MemoryTypeEnum +from ...scheme.memory_node import MemoryNode +from ...scheme.message import Message +from ..memory_base_worker import MemoryBaseWorker +from ...prompts.get_observation_prompt import ( + GET_OBSERVATION_FEW_SHOT_PROMPT, + GET_OBSERVATION_SYSTEM_PROMPT, + GET_OBSERVATION_USER_QUERY_PROMPT, +) + + +class GetObservationWorker(MemoryBaseWorker): + + def add_observation(self, message: Message, obs_content: str, keywords: str): + created_dt: datetime = datetime.fromtimestamp(float(message.time_created)) + dt = time_to_formatted_str(time=created_dt) + + # 组合meta_data + meta_data = { + MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话 + REFLECTED: "0", # reflect标记 + DT: dt, # 当天标记 + NEW: "1", # summary-long标记 + MSG_TIME: message.time_created, # 对话时间 + TIME_INFER: "", # 推断的时间 + KEY_WORD: keywords, # 关键词 + CONTENT_MODIFIED: True, # 新增的obs需要置为true + } + meta_data.update( + {k: str(v) for k, v in get_datetime_info_dict(created_dt).items()} + ) + + return MemoryNode( + content=obs_content, + memory_id=self.memory_id, + memory_type=MemoryTypeEnum.OBSERVATION.value, + meta_data=meta_data, + status=MemoryNodeStatus.ACTIVE.value, + ) + + def _run(self): + # gene prompt + user_query_list = [] + i = 1 + for msg in self.messages: + match = False + for time_keyword in DATATIME_WORD_LIST: + if time_keyword in msg.content: + match = True + break + if not match: + user_query_list.append(f"{i} 用户:{msg.content}") + i += 1 + + if not user_query_list: + self.add_run_info(f"get obs user_query_list={user_query_list} is empty") + return + + obtain_obs_message = self.prompt_to_msg( + system_prompt=self.get_prompt(GET_OBSERVATION_SYSTEM_PROMPT).format( + num_obs=len(user_query_list) + ), + few_shot=self.get_prompt(GET_OBSERVATION_FEW_SHOT_PROMPT), + user_query=self.get_prompt(GET_OBSERVATION_USER_QUERY_PROMPT).format( + user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"obtain_obs_message={obtain_obs_message}") + + # call LLM + response_text: str = self.generation_model.call( + messages=obtain_obs_message, + model_name=self.summary_messages_model, + max_token=self.summary_messages_max_token, + temperature=self.summary_messages_temperature, + top_k=self.summary_messages_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("summary call llm failed!", continue_run=False) + return + + # parse text + idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs") + if len(idx_obs_list) <= 0: + self.add_run_info("idx_obs_list is empty!", continue_run=False) + return + + # gene new obs nodes + new_obs_nodes: List[MemoryNode] = [] + for obs_content_list in idx_obs_list: + if not obs_content_list: + continue + + # [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游] + if len(obs_content_list) != 4: + self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") + continue + + idx, time_infer, obs_content, keywords = obs_content_list + + if obs_content in ["无", "重复"]: + continue + + if not idx.isdigit(): + self.logger.warning(f"idx={idx} is invalid!") + continue + + # 序号需要修正-1 + idx = int(idx) - 1 + if idx >= len(self.messages): + self.logger.warning( + f"idx={idx} is invalid! messages.size={len(self.messages)}" + ) + continue + + new_obs_nodes.append( + self.add_observation( + message=self.messages[idx], + obs_content=obs_content, + keywords=keywords, + ) + ) + + # save context + self.set_context(NEW_OBS_NODES, new_obs_nodes) diff --git a/memory_scope/worker/summary_short/info_filter_worker.py b/memory_scope/worker/summary_short/info_filter_worker.py new file mode 100644 index 00000000..f030642e --- /dev/null +++ b/memory_scope/worker/summary_short/info_filter_worker.py @@ -0,0 +1,70 @@ +from ...utils.response_text_parser import ResponseTextParser +from enumeration.message_role_enum import MessageRoleEnum +from worker.memory_base_worker import MemoryBaseWorker +from ...chat.global_context import GlobalContext +from ...prompts.info_filter_prompt import INFO_FILTER_FEW_SHOT_PROMPT, INFO_FILTER_SYSTEM_PROMPT, INFO_FILTER_USER_QUERY_PROMPT + + +class InfoFilterWorker(MemoryBaseWorker): + def _run(self): + # filter user msg + info_messages = [] + for msg in self.messages: + if msg.role != MessageRoleEnum.USER.value: + continue + if len(msg.content) >= self.info_filter_msg_max_size: + continue + info_messages.append(msg) + + # gene prompt + user_query = "\n".join( + [f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)] + ) + info_filter_message = self.prompt_to_msg( + system_prompt=self.get_prompt(INFO_FILTER_SYSTEM_PROMPT).format( + batch_size=len(info_messages) + ), + few_shot=self.get_prompt(INFO_FILTER_FEW_SHOT_PROMPT), + user_query=self.get_prompt(INFO_FILTER_USER_QUERY_PROMPT).format( + user_query=user_query + ), + ) + self.logger.info(f"info_filter_message={info_filter_message}") + + # call llm + response_text = self.generation_model.call( + messages=info_filter_message, + model_name=self.info_filter_model, + max_token=self.info_filter_max_token, + temperature=self.info_filter_temperature, + top_k=self.info_filter_top_k, + ) + + # return if empty + if not response_text: + self.add_run_info("info score call llm failed!", continue_run=False) + return + + # parse text + info_score_list = ResponseTextParser(response_text).parse_v1("info_filter") + if len(info_score_list) != len(info_messages): + self.add_run_info( + f"info_score_size != info_messages_size, " + f"{len(info_score_list)} vs {len(info_messages)}", + continue_run=False, + ) + return + + # 过滤value=0的messages + filtered_messages = [] + for msg, info_score in zip(info_messages, info_score_list): + if not info_score: + continue + score = info_score[0] + # if score in ("1", "2",): + if score in ("2",): + msg.info_score = score + filtered_messages.append(msg) + + # 后续不会关注为0的msg,直接丢弃 + self.messages = filtered_messages diff --git a/test.py b/test.py new file mode 100644 index 00000000..800c0dc0 --- /dev/null +++ b/test.py @@ -0,0 +1,13 @@ +from memory_scope.cli import CliJob +import fire + + +def main(config_path: str): + job = CliJob(config_path=config_path) + job.init_global_content_by_config() + job.run() + + +if __name__ == "__main__": + # fire.Fire(main) + main("config/config.yaml") From 5bcb6301ac3a2cfcd557b427888f8225196d19af Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 12:27:09 +0800 Subject: [PATCH 21/41] [dev] add default class path memory_scope --- config/config.yaml | 12 ++++++------ memory_scope/utils/tool_functions.py | 10 +++++----- 2 files changed, 11 insertions(+), 11 deletions(-) diff --git a/config/config.yaml b/config/config.yaml index bb4fad52..f6b0f85a 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -5,23 +5,23 @@ global_config: open_ai_apikey: memory_chat: cli_memory_chat: - class: memory_scope.chat_v2.cli_memory_chat + class: chat_v2.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation memory_service: memory_chat_service: - class: memory_scope.memory.service.chat_memory_service + class: memory.service.chat_memory_service history_msg_count: 32 contextual_msg_count: 6 read_memory_key: read_memory memory_operations: read_message: - class: memory_scope.memory.operation.read_memory + class: memory.operation.read_memory workflow: dummy_worker description: "read session messages of the user" contextual_msg_count: 0 read_memory: - class: memory_scope.memory.operation.read_memory + class: memory.operation.read_memory workflow: dummy_worker description: "read related memories of the user" list_memory: @@ -34,7 +34,7 @@ memory_service: description: "write observation memories of the user" interval_time: 60 summary_memory: - class: memory_scope.memory.operation.summary_memory + class: memory.operation.summary_memory workflow: dummy_worker description: "summary observation memories of the user" interval_time: 300 @@ -48,7 +48,7 @@ models: module_name: dashscope_embedding model_name: text-embedding-v2 dashscope_rank: - clazz: models.base_rank_model + clazz: models.llama_index_rank_model module_name: dashscope_rank model_name: gte-rerank vector_store: diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index a0e847c2..7ac7e565 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -11,18 +11,18 @@ def under_line_to_hump(underline_str): def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs): - class_name = config.pop("class") - if not class_name: - raise RuntimeError("empty class_name!") + origin_class_path: str = config.pop("class") + if not origin_class_path: + raise RuntimeError("empty class path!") - class_name_split = class_name.split(".") + class_name_split = origin_class_path.split(".") class_name: str = class_name_split[-1] if suffix_name and not class_name.lower().endswith(suffix_name.lower()): class_name = f"{class_name}_{suffix_name}" class_name_split[-1] = class_name class_paths = [] - if default_class_path: + if default_class_path and not origin_class_path.startswith(default_class_path): class_paths.append(default_class_path) class_paths.extend(class_name_split) module = import_module(".".join(class_paths)) From 6dbf44904a89ead556fbf8613567418af4f75ab3 Mon Sep 17 00:00:00 2001 From: hs Date: Thu, 27 Jun 2024 12:29:35 +0800 Subject: [PATCH 22/41] [dev] modify g context attr --- memory_scope/chat_v2/global_context.py | 28 +++++++++++++++++--------- 1 file changed, 18 insertions(+), 10 deletions(-) diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py index 7b08a2bb..88d9e5b2 100644 --- a/memory_scope/chat_v2/global_context.py +++ b/memory_scope/chat_v2/global_context.py @@ -9,18 +9,26 @@ from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_vector_store import BaseVectorStore +from memory_scope.memory.worker.base_worker import BaseWorker -class GlobalContext(pydantic.BaseModel): - global_config: Dict[str, Any] = pydantic.Field({}, description="global config") - worker_config: Dict[str, Any] = pydantic.Field({}, description="worker config") +class GlobalContext(object): + def __init__(self): + self.global_configs: Dict[str, Any] = {} - memory_service_dict: Dict[str, BaseMemoryService] = pydantic.Field({}, description="memory_service dict") - model_dict: Dict[str, BaseModel] = pydantic.Field({}, description="model dict") - memory_chat_dict: Dict[str, BaseMemoryChat] = pydantic.Field({}, description="memory_chat dict") + self.worker_config: Dict[str, Dict[str, BaseWorker]] = {} - vector_store: BaseVectorStore | None = pydantic.Field(None, description="global vector_store") - monitor: BaseMonitor | None = pydantic.Field(None, description="global monitor") - thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool") - language: LanguageEnum = pydantic.Field(LanguageEnum.CN, description="language: cn / en") + self.model_dict: Dict[str, BaseModel] = {} + self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} + + self.vector_store: BaseVectorStore | None = None + + self.monitor: BaseMonitor | None = None + + self.thread_pool: ThreadPoolExecutor | None = None + + self.language: LanguageEnum = LanguageEnum.EN + + +G_CONTEXT = GlobalContext() From a644db8862704d323820db25a1592f0f865ace1a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 13:11:21 +0800 Subject: [PATCH 23/41] [dev] modify g context attr --- memory_scope/chat_v2/global_context.py | 13 +++---------- 1 file changed, 3 insertions(+), 10 deletions(-) diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py index 88d9e5b2..70b193e4 100644 --- a/memory_scope/chat_v2/global_context.py +++ b/memory_scope/chat_v2/global_context.py @@ -1,33 +1,26 @@ from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any -import pydantic - from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_vector_store import BaseVectorStore -from memory_scope.memory.worker.base_worker import BaseWorker class GlobalContext(object): def __init__(self): - self.global_configs: Dict[str, Any] = {} - - self.worker_config: Dict[str, Dict[str, BaseWorker]] = {} + self.global_config: Dict[str, Any] = {} + self.worker_config: Dict[str, Any] = {} + self.memory_service_dict: Dict[str, BaseMemoryService] = {} self.model_dict: Dict[str, BaseModel] = {} - self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} self.vector_store: BaseVectorStore | None = None - self.monitor: BaseMonitor | None = None - self.thread_pool: ThreadPoolExecutor | None = None - self.language: LanguageEnum = LanguageEnum.EN From ea8109c4a1fa4bdf5b5d6c35dc322b8d847b778f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:02:18 +0800 Subject: [PATCH 24/41] [dev] modify default json config --- config/config.json | 104 +++++++++++++++++++------ config/model/dashscope_embedding.json | 5 -- config/model/dashscope_generation.json | 5 -- config/model/dashscope_rank.json | 5 -- config/workers.json | 8 -- memory_scope/cli.py | 71 ++++++++--------- 6 files changed, 113 insertions(+), 85 deletions(-) delete mode 100644 config/model/dashscope_embedding.json delete mode 100644 config/model/dashscope_generation.json delete mode 100644 config/model/dashscope_rank.json delete mode 100644 config/workers.json diff --git a/config/config.json b/config/config.json index 7adce564..394cf9e8 100644 --- a/config/config.json +++ b/config/config.json @@ -1,27 +1,85 @@ { - "global_configs": { - "thread_pool_max_count": 5, - "dash_scope_apikey": "", - "open_ai_apikey": "", - "language": "en", - "chat_list": [ - "memory_chat" - ] + "global_config": { + "language": "en", + "max_workers": 5, + "dash_scope_apikey": null, + "open_ai_apikey": null + }, + "memory_chat": { + "cli_memory_chat": { + "class": "chat_v2.cli_memory_chat", + "memory_service": "memory_chat_service", + "generation_model": "dashscope_generation" + } + }, + "memory_service": { + "memory_chat_service": { + "class": "memory.service.chat_memory_service", + "history_msg_count": 32, + "contextual_msg_count": 6, + "read_memory_key": "read_memory", + "memory_operations": { + "read_message": { + "class": "memory.operation.read_memory", + "workflow": "dummy_worker", + "description": "read session messages of the user", + "contextual_msg_count": 0 + }, + "read_memory": { + "class": "memory.operation.read_memory", + "workflow": "dummy_worker", + "description": "read related memories of the user" + }, + "list_memory": { + "class": "memory.operation.read_memory", + "workflow": "dummy_worker", + "description": "read all memories of the user" + }, + "write_memory": { + "class": "memory.operation.write_memory", + "workflow": "dummy_worker", + "description": "write observation memories of the user", + "interval_time": 60 + }, + "summary_memory": { + "class": "memory.operation.summary_memory", + "workflow": "dummy_worker", + "description": "summary observation memories of the user", + "interval_time": 300 + } + } + } + }, + "models": { + "dashscope_generation": { + "clazz": "models.llama_index_generation_model", + "module_name": "dashscope_generation", + "model_name": "qwen-max" }, - "memory_chat": { - "clazz": "chat.memory_chat", - "retrieve": "parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank", - "generation_model": "dashscope_generation", - "history_msg_count": 3 + "dashscope_embedding": { + "clazz": "models.llama_index_embedding_model", + "module_name": "dashscope_embedding", + "model_name": "text-embedding-v2" }, - "vector_store": { - "clazz": "storage.base_vector_store", - "index_name": "memory_test", - "password": "" - }, - "monitor": { - "clazz": "storage.base_monitor", - "index_name": "memory_test" - }, - "workers": "workers" + "dashscope_rank": { + "clazz": "models.llama_index_rank_model", + "module_name": "dashscope_rank", + "model_name": "gte-rerank" + } + }, + "vector_store": { + "clazz": "storage.dummy_vector_store", + "embedding_model": "dashscope_embedding" + }, + "monitor": { + "clazz": "storage.dummy_monitor" + }, + "workers": { + "dummy_worker": { + "clazz": "memory.worker.dummy_worker", + "generation_model": "dashscope_generation", + "embedding_model": "dashscope_embedding", + "rank_model": "dashscope_rank" + } + } } \ No newline at end of file diff --git a/config/model/dashscope_embedding.json b/config/model/dashscope_embedding.json deleted file mode 100644 index ed5fb740..00000000 --- a/config/model/dashscope_embedding.json +++ /dev/null @@ -1,5 +0,0 @@ -{ - "clazz": "models.base_embedding_model", - "model_name": "text-embedding-v2", - "method_type": "DashScopeEmbedding" -} \ No newline at end of file diff --git a/config/model/dashscope_generation.json b/config/model/dashscope_generation.json deleted file mode 100644 index ba2e3c28..00000000 --- a/config/model/dashscope_generation.json +++ /dev/null @@ -1,5 +0,0 @@ -{ - "clazz": "models.llama_index_generation_model", - "model_name": "qwen-max", - "method_type": "DashScope" -} \ No newline at end of file diff --git a/config/model/dashscope_rank.json b/config/model/dashscope_rank.json deleted file mode 100644 index e2c9e302..00000000 --- a/config/model/dashscope_rank.json +++ /dev/null @@ -1,5 +0,0 @@ -{ - "clazz": "models.base_rank_model", - "model_name": "gte-rerank", - "method_type": "DashScopeRerank" -} \ No newline at end of file diff --git a/config/workers.json b/config/workers.json deleted file mode 100644 index e1bd9d90..00000000 --- a/config/workers.json +++ /dev/null @@ -1,8 +0,0 @@ -{ - "update_insight": { - "clazz": "worker.summary_long.update_insight", - "generation_model": "dashscope_generation", - "embedding_model": "dashscope_embedding", - "rank_model": "dashscope_rank" - } -} \ No newline at end of file diff --git a/memory_scope/cli.py b/memory_scope/cli.py index cb040d8d..f888e674 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,67 +1,54 @@ +import json from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any -import yaml import fire +import yaml -from .chat_v2.global_context import G_CONTEXT -from .enumeration.language_enum import LanguageEnum -from .utils.logger import Logger -from .utils.tool_functions import init_instance_by_config +from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.enumeration.language_enum import LanguageEnum +from memory_scope.utils.logger import Logger +from memory_scope.utils.tool_functions import init_instance_by_config class CliJob(object): - def __init__(self, config_path: str, config_suffix: str = ".yaml"): - self.config_path: str = config_path - self.config_suffix: str = config_suffix + def __init__(self): self.config: Dict[str, Any] = {} - self.logger: Logger = Logger.get_logger("cli_job") - @staticmethod - def set_global_config(global_config: Dict[str, Any]): - """set global_configs & set apikey into env - :return: - TODO at sen - """ - G_CONTEXT.global_config = global_config + def load_config(self, path: str): + with open(path) as f: + if path.endswith("yaml"): + self.config = yaml.load(f, yaml.FullLoader) + elif path.endswith("json"): + self.config = json.load(f) + else: + raise RuntimeError("not supported config file type!") + + def set_global_config(self): + G_CONTEXT.global_config = global_config = self.config["global_config"] G_CONTEXT.language = LanguageEnum(global_config["language"]) - G_CONTEXT.thread_pool = ThreadPoolExecutor( - max_workers=int(global_config["max_workers"]) - ) + G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"])) def init_global_content_by_config(self): - # load config - config_path = self.config_path - if not self.config_path.endswith(self.config_suffix): - config_path += self.config_suffix - with open(config_path) as f: - self.config = yaml.load(f, yaml.FullLoader) - - # set global_config - self.set_global_config(self.config["global_config"]) + # set global config + self.set_global_config() # init memory_chat for name, conf in self.config["memory_chat"].items(): - G_CONTEXT.memory_chat_dict[name] = init_instance_by_config( - conf, name=name - ) + G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name) # set memory_service for name, conf in self.config["memory_service"].items(): - G_CONTEXT.memory_service_dict[name] = init_instance_by_config( - conf, name=name - ) + G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name) # init models for name, conf in self.config["models"].items(): G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) # init vector_store - G_CONTEXT.vector_store = init_instance_by_config( - self.config["vector_store"] - ) + G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"]) # init monitor G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) @@ -69,8 +56,14 @@ class CliJob(object): # set worker config G_CONTEXT.worker_config = self.config["workers"] - @staticmethod - def run(): + def run(self, config: str): + self.load_config(config) + with G_CONTEXT.thread_pool: memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() + + +if __name__ == "__main__": + cli_job = CliJob() + fire.Fire(cli_job.run) From d7b6e20777953b865b1460d812abe5a097b91e6d Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:10:48 +0800 Subject: [PATCH 25/41] [dev] rename clazz to class --- config/config.json | 9 ++++----- config/config.yaml | 7 +++---- memory_scope/chat_v2/cli_memory_chat.py | 18 +++++++++--------- memory_scope/cli.py | 6 +++++- memory_scope/memory/operation/base_workflow.py | 2 +- .../memory/service/chat_memory_service.py | 7 ++----- 6 files changed, 24 insertions(+), 25 deletions(-) diff --git a/config/config.json b/config/config.json index 394cf9e8..e2a4ad5f 100644 --- a/config/config.json +++ b/config/config.json @@ -22,8 +22,7 @@ "read_message": { "class": "memory.operation.read_memory", "workflow": "dummy_worker", - "description": "read session messages of the user", - "contextual_msg_count": 0 + "description": "read session messages of the user" }, "read_memory": { "class": "memory.operation.read_memory", @@ -52,17 +51,17 @@ }, "models": { "dashscope_generation": { - "clazz": "models.llama_index_generation_model", + "class": "models.llama_index_generation_model", "module_name": "dashscope_generation", "model_name": "qwen-max" }, "dashscope_embedding": { - "clazz": "models.llama_index_embedding_model", + "class": "models.llama_index_embedding_model", "module_name": "dashscope_embedding", "model_name": "text-embedding-v2" }, "dashscope_rank": { - "clazz": "models.llama_index_rank_model", + "class": "models.llama_index_rank_model", "module_name": "dashscope_rank", "model_name": "gte-rerank" } diff --git a/config/config.yaml b/config/config.yaml index f6b0f85a..30b70c50 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -19,7 +19,6 @@ memory_service: class: memory.operation.read_memory workflow: dummy_worker description: "read session messages of the user" - contextual_msg_count: 0 read_memory: class: memory.operation.read_memory workflow: dummy_worker @@ -40,15 +39,15 @@ memory_service: interval_time: 300 models: dashscope_generation: - clazz: models.llama_index_generation_model + class: models.llama_index_generation_model module_name: dashscope_generation model_name: qwen-max dashscope_embedding: - clazz: models.llama_index_embedding_model + class: models.llama_index_embedding_model module_name: dashscope_embedding model_name: text-embedding-v2 dashscope_rank: - clazz: models.llama_index_rank_model + class: models.llama_index_rank_model module_name: dashscope_rank model_name: gte-rerank vector_store: diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index fa0bead1..c2012bd8 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -1,8 +1,9 @@ import datetime import time -from typing import Dict, List +from typing import List import questionary + from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat from memory_scope.chat_v2.global_context import G_CONTEXT from memory_scope.enumeration.message_role_enum import MessageRoleEnum @@ -20,11 +21,11 @@ class CliMemoryChat(BaseMemoryChat): "stream": "get stream response" } - def __init__(self, memory_service: str, generation_model: str, **kwargs): - super().__init__(**kwargs) + def __init__(self, memory_service: str, generation_model: str, stream: bool = True, **kwargs): self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model - self.stream: bool = True + self.stream: bool = stream + self.kwargs: dict = kwargs @property def memory_service(self) -> BaseMemoryService: @@ -56,18 +57,17 @@ class CliMemoryChat(BaseMemoryChat): time_created = int(datetime.datetime.now().timestamp()) new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) - self.submit_messages(new_message) + self.memory_service.add_messages(new_message) related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) if self.stream: for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): yield result - self.submit_messages(result.text) + self.memory_service.add_messages(result.text) def run(self): - op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() - self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) + self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) while True: query = questionary.text( @@ -97,7 +97,7 @@ class CliMemoryChat(BaseMemoryChat): elif query == "stream": questionary.print(f"stream: {self.stream}") self.stream = ~self.stream - elif query in op_description_dict: + elif query in self.memory_service.op_description_dict: if not args: result = self.memory_service.do_operation(op_name=query) print(result) diff --git a/memory_scope/cli.py b/memory_scope/cli.py index f888e674..27da1f5d 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,3 +1,7 @@ +import sys + +sys.path.append(".") + import json from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any @@ -58,7 +62,7 @@ class CliJob(object): def run(self, config: str): self.load_config(config) - + self.init_global_content_by_config() with G_CONTEXT.thread_pool: memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index 540f6cb8..74ce7943 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -17,7 +17,7 @@ class BaseWorkflow(object): def __init__(self, name: str, workflow: str, - thread_pool: ThreadPoolExecutor, + thread_pool: ThreadPoolExecutor = G_CONTEXT.thread_pool, **kwargs): self.name: str = name diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index badd26c8..a2197e91 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -6,9 +6,7 @@ from memory_scope.utils.tool_functions import init_instance_by_config class ChatMemoryService(BaseMemoryService): - def __init__( - self, history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs - ): + def __init__(self, history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs): super().__init__(**kwargs) self.history_msg_count: int = history_msg_count self.contextual_msg_count: int = contextual_msg_count @@ -26,8 +24,7 @@ class ChatMemoryService(BaseMemoryService): name=name, chat_messages=self.chat_messages, message_lock=self.message_lock, - contextual_msg_count=self.contextual_msg_count, - ) + contextual_msg_count=self.contextual_msg_count) def add_messages(self, messages: List[Message] | Message): if isinstance(messages, Message): From 3e5ad8f652a82795d17bd244a6b676902516fe9d Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:12:46 +0800 Subject: [PATCH 26/41] [dev] rename Reranker to ranker --- ...{llama_index_rerank_model.py => llama_index_rank_model.py} | 2 +- tests/models/test_models_lli_rerank.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) rename memory_scope/models/{llama_index_rerank_model.py => llama_index_rank_model.py} (97%) diff --git a/memory_scope/models/llama_index_rerank_model.py b/memory_scope/models/llama_index_rank_model.py similarity index 97% rename from memory_scope/models/llama_index_rerank_model.py rename to memory_scope/models/llama_index_rank_model.py index c76cbd10..095fda5e 100644 --- a/memory_scope/models/llama_index_rerank_model.py +++ b/memory_scope/models/llama_index_rank_model.py @@ -9,7 +9,7 @@ from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY from memory_scope.models.model_response import ModelResponse -class LlamaIndexRerankModel(BaseModel): +class LlamaIndexRankModel(BaseModel): m_type: ModelEnum = ModelEnum.RANK_MODEL MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank) diff --git a/tests/models/test_models_lli_rerank.py b/tests/models/test_models_lli_rerank.py index c7bdf2f6..1238f7f8 100644 --- a/tests/models/test_models_lli_rerank.py +++ b/tests/models/test_models_lli_rerank.py @@ -1,7 +1,7 @@ import json import unittest -from memory_scope.models.llama_index_rerank_model import LlamaIndexRerankModel +from memory_scope.models.llama_index_rerank_model import LlamaIndexRankModel class TestLLIReRank(unittest.TestCase): @@ -13,7 +13,7 @@ class TestLLIReRank(unittest.TestCase): "model_name": "gte-rerank", "clazz": "models.llama_index_rerank_model" } - self.reranker = LlamaIndexRerankModel(**config) + self.reranker = LlamaIndexRankModel(**config) def test_rerank(self): query = "吃啥?" From ad206fda2bdb9e0f7ecab6ed476a113c4c96664c Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:14:00 +0800 Subject: [PATCH 27/41] [dev] rename clazz to class in json/yaml --- config/config.json | 6 +++--- config/config.yaml | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/config/config.json b/config/config.json index e2a4ad5f..dc1bc36f 100644 --- a/config/config.json +++ b/config/config.json @@ -67,15 +67,15 @@ } }, "vector_store": { - "clazz": "storage.dummy_vector_store", + "class": "storage.dummy_vector_store", "embedding_model": "dashscope_embedding" }, "monitor": { - "clazz": "storage.dummy_monitor" + "class": "storage.dummy_monitor" }, "workers": { "dummy_worker": { - "clazz": "memory.worker.dummy_worker", + "class": "memory.worker.dummy_worker", "generation_model": "dashscope_generation", "embedding_model": "dashscope_embedding", "rank_model": "dashscope_rank" diff --git a/config/config.yaml b/config/config.yaml index 30b70c50..43f65451 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -51,13 +51,13 @@ models: module_name: dashscope_rank model_name: gte-rerank vector_store: - clazz: storage.dummy_vector_store + class: storage.dummy_vector_store embedding_model: dashscope_embedding monitor: - clazz: storage.dummy_monitor + class: storage.dummy_monitor workers: dummy_worker: - clazz: memory.worker.dummy_worker + class: memory.worker.dummy_worker generation_model: dashscope_generation embedding_model: dashscope_embedding rank_model: dashscope_rank From 879dcf866ce64ea14f7ab5f95d2a89bde268a43b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:22:30 +0800 Subject: [PATCH 28/41] [dev] rename workers to worker --- config/config.json | 2 +- config/config.yaml | 2 +- memory_scope/chat_v2/global_context.py | 2 +- memory_scope/cli.py | 2 +- memory_scope/memory/operation/write_memory.py | 2 +- memory_scope/storage/base_vector_store.py | 6 +++++- memory_scope/utils/tool_functions.py | 7 +++++-- 7 files changed, 15 insertions(+), 8 deletions(-) diff --git a/config/config.json b/config/config.json index dc1bc36f..3c77eb30 100644 --- a/config/config.json +++ b/config/config.json @@ -73,7 +73,7 @@ "monitor": { "class": "storage.dummy_monitor" }, - "workers": { + "worker": { "dummy_worker": { "class": "memory.worker.dummy_worker", "generation_model": "dashscope_generation", diff --git a/config/config.yaml b/config/config.yaml index 43f65451..7424fbea 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -55,7 +55,7 @@ vector_store: embedding_model: dashscope_embedding monitor: class: storage.dummy_monitor -workers: +worker: dummy_worker: class: memory.worker.dummy_worker generation_model: dashscope_generation diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py index 70b193e4..c38ae6c3 100644 --- a/memory_scope/chat_v2/global_context.py +++ b/memory_scope/chat_v2/global_context.py @@ -12,7 +12,7 @@ from memory_scope.storage.base_vector_store import BaseVectorStore class GlobalContext(object): def __init__(self): self.global_config: Dict[str, Any] = {} - self.worker_config: Dict[str, Any] = {} + self.worker_config: Dict[str, Dict[str, Any]] = {} self.memory_service_dict: Dict[str, BaseMemoryService] = {} self.model_dict: Dict[str, BaseModel] = {} diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 27da1f5d..8545315f 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -58,7 +58,7 @@ class CliJob(object): G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) # set worker config - G_CONTEXT.worker_config = self.config["workers"] + G_CONTEXT.worker_config = self.config["worker"] def run(self, config: str): self.load_config(config) diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 92bc8b47..a47f61ad 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -8,7 +8,7 @@ from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message -class WriteMemory(BaseOperation, BaseWorkflow): +class WriteMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "backend" def __init__(self, diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index 36bba692..c5281fbd 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -7,7 +7,11 @@ from memory_scope.scheme.memory_node import MemoryNode class BaseVectorStore(metaclass=ABCMeta): - def __init__(self, index_name: str, embedding_model: BaseModel, content_key: str = "text", **kwargs): + def __init__(self, + index_name: str = "", + embedding_model: BaseModel | None = None, + content_key: str = "text", + **kwargs): self.index_name: str = index_name self.embedding_model: BaseModel = embedding_model self.content_key: str = content_key diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 7ac7e565..aeb8dfbc 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -1,7 +1,9 @@ import re +from copy import deepcopy from datetime import datetime from importlib import import_module +from memory_scope.constants.common_constants import WEEKDAYS from memory_scope.enumeration.message_role_enum import MessageRoleEnum @@ -11,7 +13,8 @@ def under_line_to_hump(underline_str): def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs): - origin_class_path: str = config.pop("class") + config_copy = deepcopy(config) + origin_class_path: str = config_copy.pop("class") if not origin_class_path: raise RuntimeError("empty class path!") @@ -28,7 +31,7 @@ def init_instance_by_config(config: dict, default_class_path: str = "memory_scop module = import_module(".".join(class_paths)) cls_name = under_line_to_hump(class_name) - return getattr(module, cls_name)(**config, **kwargs) + return getattr(module, cls_name)(**config_copy, **kwargs) def complete_config_name(config_name: str, suffix: str = ".json"): From 28624de8785c4d3492d59c5c18ffaef9e0b49157 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:28:19 +0800 Subject: [PATCH 29/41] [dev] modify path to absolute class path --- config/config.json | 2 +- config/config.yaml | 2 +- memory_scope/chat/base_memory_chat.py | 2 - memory_scope/chat/base_memory_service.py | 54 ------- memory_scope/chat/chat_memory_service.py | 50 ------- memory_scope/chat/cli_memory_chat.py | 28 ++-- memory_scope/chat/global_context.py | 24 ++- memory_scope/chat/memory_chat.py | 57 ------- memory_scope/chat_v2/base_memory_chat.py | 14 -- memory_scope/chat_v2/cli_memory_chat.py | 139 ------------------ memory_scope/chat_v2/global_context.py | 27 ---- memory_scope/cli.py | 2 +- .../memory/operation/base_workflow.py | 2 +- .../memory/operation/summary_memory.py | 2 +- memory_scope/memory/operation/write_memory.py | 2 +- memory_scope/storage/base_vector_store.py | 12 -- memory_scope/worker/summary_short/__init__.py | 0 .../chat_v2 => old/worker}/__init__.py | 0 {memory_scope => old}/worker/base_worker.py | 0 {memory_scope => old}/worker/dummy_worker.py | 0 .../worker => old/worker/es}/__init__.py | 0 .../worker/es/es_insight_worker.py | 0 .../worker/es/es_new_obs_worker.py | 0 .../worker/es/es_not_reflected_worker.py | 0 .../worker/es/es_similar_worker.py | 0 .../worker/es/es_today_obs_worker.py | 0 .../worker/es/load_profile_worker.py | 0 .../worker/memory_base_worker.py | 0 .../es => old/worker/retrieve}/__init__.py | 0 .../worker/retrieve/extract_time_worker.py | 0 .../worker/retrieve/fuse_rerank_worker.py | 0 .../worker/retrieve/memory_store_worker.py | 0 .../worker/retrieve/semantic_rank_worker.py | 0 .../worker/summary_long}/__init__.py | 0 .../worker/summary_long/get_insight_worker.py | 0 .../summary_long/get_reflection_worker.py | 0 .../summary_long/long_contra_repeat_worker.py | 0 .../summary_long/summary_collect_worker.py | 0 .../summary_long/update_insight_worker.py | 0 .../summary_long/update_profile_worker.py | 0 .../worker/summary_short}/__init__.py | 0 .../summary_short/contra_repeat_worker.py | 0 .../get_observation_with_time_worker.py | 0 .../summary_short/get_observation_worker.py | 0 .../summary_short/info_filter_worker.py | 0 45 files changed, 30 insertions(+), 389 deletions(-) delete mode 100644 memory_scope/chat/base_memory_service.py delete mode 100644 memory_scope/chat/chat_memory_service.py delete mode 100644 memory_scope/chat/memory_chat.py delete mode 100644 memory_scope/chat_v2/base_memory_chat.py delete mode 100644 memory_scope/chat_v2/cli_memory_chat.py delete mode 100644 memory_scope/chat_v2/global_context.py delete mode 100644 memory_scope/worker/summary_short/__init__.py rename {memory_scope/chat_v2 => old/worker}/__init__.py (100%) rename {memory_scope => old}/worker/base_worker.py (100%) rename {memory_scope => old}/worker/dummy_worker.py (100%) rename {memory_scope/worker => old/worker/es}/__init__.py (100%) rename {memory_scope => old}/worker/es/es_insight_worker.py (100%) rename {memory_scope => old}/worker/es/es_new_obs_worker.py (100%) rename {memory_scope => old}/worker/es/es_not_reflected_worker.py (100%) rename {memory_scope => old}/worker/es/es_similar_worker.py (100%) rename {memory_scope => old}/worker/es/es_today_obs_worker.py (100%) rename {memory_scope => old}/worker/es/load_profile_worker.py (100%) rename {memory_scope => old}/worker/memory_base_worker.py (100%) rename {memory_scope/worker/es => old/worker/retrieve}/__init__.py (100%) rename {memory_scope => old}/worker/retrieve/extract_time_worker.py (100%) rename {memory_scope => old}/worker/retrieve/fuse_rerank_worker.py (100%) rename {memory_scope => old}/worker/retrieve/memory_store_worker.py (100%) rename {memory_scope => old}/worker/retrieve/semantic_rank_worker.py (100%) rename {memory_scope/worker/retrieve => old/worker/summary_long}/__init__.py (100%) rename {memory_scope => old}/worker/summary_long/get_insight_worker.py (100%) rename {memory_scope => old}/worker/summary_long/get_reflection_worker.py (100%) rename {memory_scope => old}/worker/summary_long/long_contra_repeat_worker.py (100%) rename {memory_scope => old}/worker/summary_long/summary_collect_worker.py (100%) rename {memory_scope => old}/worker/summary_long/update_insight_worker.py (100%) rename {memory_scope => old}/worker/summary_long/update_profile_worker.py (100%) rename {memory_scope/worker/summary_long => old/worker/summary_short}/__init__.py (100%) rename {memory_scope => old}/worker/summary_short/contra_repeat_worker.py (100%) rename {memory_scope => old}/worker/summary_short/get_observation_with_time_worker.py (100%) rename {memory_scope => old}/worker/summary_short/get_observation_worker.py (100%) rename {memory_scope => old}/worker/summary_short/info_filter_worker.py (100%) diff --git a/config/config.json b/config/config.json index 3c77eb30..3f895d4b 100644 --- a/config/config.json +++ b/config/config.json @@ -7,7 +7,7 @@ }, "memory_chat": { "cli_memory_chat": { - "class": "chat_v2.cli_memory_chat", + "class": "chat.cli_memory_chat", "memory_service": "memory_chat_service", "generation_model": "dashscope_generation" } diff --git a/config/config.yaml b/config/config.yaml index 7424fbea..e6d78cb1 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -5,7 +5,7 @@ global_config: open_ai_apikey: memory_chat: cli_memory_chat: - class: chat_v2.cli_memory_chat + class: chat.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation memory_service: diff --git a/memory_scope/chat/base_memory_chat.py b/memory_scope/chat/base_memory_chat.py index 63d351dc..f647cb98 100644 --- a/memory_scope/chat/base_memory_chat.py +++ b/memory_scope/chat/base_memory_chat.py @@ -2,8 +2,6 @@ from abc import ABCMeta, abstractmethod class BaseMemoryChat(metaclass=ABCMeta): - def __init__(self, **kwargs): - self.kwargs = kwargs @abstractmethod def chat_with_memory(self, query: str): diff --git a/memory_scope/chat/base_memory_service.py b/memory_scope/chat/base_memory_service.py deleted file mode 100644 index 1f610b01..00000000 --- a/memory_scope/chat/base_memory_service.py +++ /dev/null @@ -1,54 +0,0 @@ -import threading -from abc import ABCMeta, abstractmethod -from typing import List, Dict - -from memory_scope.memory.operation.base_operation import BaseOperation -from memory_scope.scheme.message import Message -from memory_scope.utils.logger import Logger - - -class BaseMemoryService(metaclass=ABCMeta): - def __init__(self, - memory_operations: Dict[str, dict], - read_memory_key: str = "read_memory", - **kwargs): - self.memory_operations: Dict[str, dict] = memory_operations - self.read_memory_key: str = read_memory_key - - self._operation_dict: Dict[str, BaseOperation] = {} - self._op_description_dict: Dict[str, str] = {} - self.chat_messages: List[Message] = [] - self.message_lock = threading.Lock - - self.logger = Logger.get_logger() - self.kwargs = kwargs - - self._init_operation(memory_operations) - - @abstractmethod - def _init_operation(self, memory_operations: Dict[str, dict]): - raise NotImplementedError - - @abstractmethod - def add_messages(self, messages: List[Message] | Message): - raise NotImplementedError - - def prepare_service(self): - pass - - @abstractmethod - def do_operation(self, op_name: str): - raise NotImplementedError - - @property - def op_description_dict(self) -> Dict[str, str]: - if not self._op_description_dict: - self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()} - return self._op_description_dict - - def read_memory(self): - assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" - return self.operate(self.read_memory_key) - - # def __getattr__(self, key): - # return self.kwargs[key] diff --git a/memory_scope/chat/chat_memory_service.py b/memory_scope/chat/chat_memory_service.py deleted file mode 100644 index def938e2..00000000 --- a/memory_scope/chat/chat_memory_service.py +++ /dev/null @@ -1,50 +0,0 @@ -from typing import List, Dict - -from memory_scope.memory.service.base_memory_service import BaseMemoryService -from memory_scope.scheme.message import Message -from memory_scope.utils.tool_functions import init_instance_by_config - - -class ChatMemoryService(BaseMemoryService): - def __init__(self, - history_msg_count: int = 32, - contextual_msg_count: int = 6, - **kwargs): - super().__init__(**kwargs) - self.history_msg_count: int = history_msg_count - self.contextual_msg_count: int = contextual_msg_count - assert self.history_msg_count >= self.contextual_msg_count - - def _init_operation(self, memory_operations: Dict[str, dict]): - for name, operation_config in memory_operations.items(): - if name in self._operation_dict: - self.logger.warning(f"memory operation={name} is repeated!") - continue - self._operation_dict[name] = init_instance_by_config(config=operation_config, - name=name, - chat_messages=self.chat_messages, - message_lock=self.message_lock, - contextual_msg_count=self.contextual_msg_count) - - def add_messages(self, messages: List[Message] | Message): - if isinstance(messages, Message): - messages = [messages] - - messages = sorted(messages, key=lambda x: x.time_created) - self.chat_messages.extend(messages) - if len(self.chat_messages) > self.history_msg_count: - gap_size = len(self.chat_messages) - self.history_msg_count - for _ in range(gap_size): - self.chat_messages.pop(0) - - def prepare_service(self): - for _, operation in self._operation_dict.items(): - operation.init_workflow() - if operation.operation_type == "backend": - operation.run_operation_backend() - - def do_operation(self, op_name: str): - if op_name not in self._operation_dict: - self.logger.warning(f"op_name={op_name} is not inited!") - return - return self._operation_dict[op_name].run_operation() diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 1920aba6..0f5ff092 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -1,10 +1,11 @@ import datetime import time -from typing import Dict, List +from typing import List import questionary + from memory_scope.chat.base_memory_chat import BaseMemoryChat -from memory_scope.chat.global_context import GlobalContext +from memory_scope.chat.global_context import G_CONTEXT from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel @@ -20,30 +21,30 @@ class CliMemoryChat(BaseMemoryChat): "stream": "get stream response" } - def __init__(self, memory_service: str, generation_model: str, **kwargs): - super().__init__(**kwargs) + def __init__(self, memory_service: str, generation_model: str, stream: bool = True, **kwargs): self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model - self.stream: bool = True + self.stream: bool = stream + self.kwargs: dict = kwargs @property def memory_service(self) -> BaseMemoryService: if isinstance(self._memory_service, str): - self._memory_service = GlobalContext.memory_service_dict[self._memory_service] + self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] self._memory_service.prepare_service() return self._memory_service @property def generation_model(self) -> BaseModel: if isinstance(self._generation_model, str): - self._generation_model = GlobalContext.model_dict[self._generation_model] + self._generation_model = G_CONTEXT.model_dict[self._generation_model] return self._generation_model @staticmethod def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[GlobalContext.language] + system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] if related_memories: - memory_prompt = MEMORY_PROMPT[GlobalContext.language] + memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] all_prompt_list = [system_prompt, memory_prompt] all_prompt_list.extend(related_memories) system_prompt = "\n".join([x.strip() for x in all_prompt_list]) @@ -56,18 +57,17 @@ class CliMemoryChat(BaseMemoryChat): time_created = int(datetime.datetime.now().timestamp()) new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) - self.submit_messages(new_message) + self.memory_service.add_messages(new_message) related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) if self.stream: for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): yield result - self.submit_messages(result.text) + self.memory_service.add_messages(result.text) def run(self): - op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() - self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) + self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) while True: query = questionary.text( @@ -97,7 +97,7 @@ class CliMemoryChat(BaseMemoryChat): elif query == "stream": questionary.print(f"stream: {self.stream}") self.stream = ~self.stream - elif query in op_description_dict: + elif query in self.memory_service.op_description_dict: if not args: result = self.memory_service.do_operation(op_name=query) print(result) diff --git a/memory_scope/chat/global_context.py b/memory_scope/chat/global_context.py index 901f641f..f9ec7a57 100644 --- a/memory_scope/chat/global_context.py +++ b/memory_scope/chat/global_context.py @@ -1,31 +1,27 @@ from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any -from .base_memory_chat import BaseMemoryChat -from ..enumeration.language_enum import LanguageEnum -from ..models.base_model import BaseModel -from ..storage.base_monitor import BaseMonitor -from ..storage.base_vector_store import BaseVectorStore -from ..worker.base_worker import BaseWorker +from memory_scope.chat.base_memory_chat import BaseMemoryChat +from memory_scope.enumeration.language_enum import LanguageEnum +from memory_scope.memory.service.base_memory_service import BaseMemoryService +from memory_scope.models.base_model import BaseModel +from memory_scope.storage.base_monitor import BaseMonitor +from memory_scope.storage.base_vector_store import BaseVectorStore class GlobalContext(object): def __init__(self): - self.global_configs: Dict[str, Any] = {} - - self.worker_config: Dict[str, Dict[str, BaseWorker]] = {} + self.global_config: Dict[str, Any] = {} + self.worker_config: Dict[str, Dict[str, Any]] = {} + self.memory_service_dict: Dict[str, BaseMemoryService] = {} self.model_dict: Dict[str, BaseModel] = {} - self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} self.vector_store: BaseVectorStore | None = None - self.monitor: BaseMonitor | None = None - self.thread_pool: ThreadPoolExecutor | None = None - self.language: LanguageEnum = LanguageEnum.EN -GLOBAL_CONTEXT = GlobalContext() +G_CONTEXT = GlobalContext() diff --git a/memory_scope/chat/memory_chat.py b/memory_scope/chat/memory_chat.py deleted file mode 100644 index f954af29..00000000 --- a/memory_scope/chat/memory_chat.py +++ /dev/null @@ -1,57 +0,0 @@ -import datetime -from typing import List - -from .base_memory_chat import BaseMemoryChat -from .global_context import GLOBAL_CONTEXT -from enumeration.message_role_enum import MessageRoleEnum -from models.base_model import BaseModel -from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT -from scheme.message import Message -from .memory_service import MemoryService - - -class MemoryChat(BaseMemoryChat): - - def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs): - super().__init__(**kwargs) - self.memory_service = MemoryService(chat_name=chat_name, **kwargs) - self.generation_model_name: str = generation_model - self.history_msg_count: int = history_msg_count - - self._generation_model: BaseModel | None = None - self.history_message_list: List[Message] = [] - - @property - def generation_model(self): - if self._generation_model is None: - self._generation_model = GLOBAL_CONTEXT.model_dict[ - self.generation_model_name - ] - return self._generation_model - - - - def chat_with_memory(self, query: str): - query = query.strip() - if not query: - return - - time_created = int(datetime.datetime.now().timestamp()) - new_message: Message = Message( - role=MessageRoleEnum.USER, content=query, time_created=time_created - ) - related_memories: List[str] = self.memory_service.retrieve(message=new_message) - system_message = self.get_system_prompt(related_memories, time_created) - self.history_message_list.append(new_message) - self.history_message_list = self.history_message_list[-self.history_msg_count:] - all_messages = [system_message] + self.history_message_list - # TODO at xian zhe - return self.generation_model.call(messages=all_messages, stream=True) - - def run(self): - self.memory_service.start_memory_backend() - while True: - query = input("wait for input:") - if query in ["stop", "停止"]: - break - self.chat_with_memory(query=query) diff --git a/memory_scope/chat_v2/base_memory_chat.py b/memory_scope/chat_v2/base_memory_chat.py deleted file mode 100644 index f647cb98..00000000 --- a/memory_scope/chat_v2/base_memory_chat.py +++ /dev/null @@ -1,14 +0,0 @@ -from abc import ABCMeta, abstractmethod - - -class BaseMemoryChat(metaclass=ABCMeta): - - @abstractmethod - def chat_with_memory(self, query: str): - """ - :param query: - :return: - """ - - def run(self): - pass diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py deleted file mode 100644 index c2012bd8..00000000 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ /dev/null @@ -1,139 +0,0 @@ -import datetime -import time -from typing import List - -import questionary - -from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat -from memory_scope.chat_v2.global_context import G_CONTEXT -from memory_scope.enumeration.message_role_enum import MessageRoleEnum -from memory_scope.memory.service.base_memory_service import BaseMemoryService -from memory_scope.models.base_model import BaseModel -from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT -from memory_scope.scheme.message import Message -from ..models.model_response import ModelResponse, ModelResponseGen - - -class CliMemoryChat(BaseMemoryChat): - USER_COMMANDS = { - "exit": "exit the CLI", - "help": "get cli commands help", - "stream": "get stream response" - } - - def __init__(self, memory_service: str, generation_model: str, stream: bool = True, **kwargs): - self._memory_service: BaseMemoryService | str = memory_service - self._generation_model: BaseModel | str = generation_model - self.stream: bool = stream - self.kwargs: dict = kwargs - - @property - def memory_service(self) -> BaseMemoryService: - if isinstance(self._memory_service, str): - self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] - self._memory_service.prepare_service() - return self._memory_service - - @property - def generation_model(self) -> BaseModel: - if isinstance(self._generation_model, str): - self._generation_model = G_CONTEXT.model_dict[self._generation_model] - return self._generation_model - - @staticmethod - def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] - if related_memories: - memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] - all_prompt_list = [system_prompt, memory_prompt] - all_prompt_list.extend(related_memories) - system_prompt = "\n".join([x.strip() for x in all_prompt_list]) - return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) - - def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: - query = query.strip() - if not query: - return - - time_created = int(datetime.datetime.now().timestamp()) - new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) - self.memory_service.add_messages(new_message) - related_memories: List[str] = self.memory_service.read_memory() - system_message: Message = self.get_system_prompt(related_memories, time_created) - if self.stream: - for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): - yield result - - self.memory_service.add_messages(result.text) - - def run(self): - self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) - - while True: - query = questionary.text( - "Please enter your message or command:", - multiline=False, - qmark=">", - ).ask() - - query: str = query.rstrip() - - if query == "": - print("Empty input received. Please try again!") - continue - - # handle cli / commands with memory ops - if query.startswith("/"): - query_split = query.lstrip("/").lower().split(" ") - query = query_split[0] - args = query_split[1:] - if query == "exit": - break - elif query == "help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(cmd, "bold") - print(f" {desc}") - elif query == "stream": - questionary.print(f"stream: {self.stream}") - self.stream = ~self.stream - elif query in self.memory_service.op_description_dict: - if not args: - result = self.memory_service.do_operation(op_name=query) - print(result) - - elif args[0].isdigit(): - refresh_time = int(args[0]) - try: - while True: - time.sleep(refresh_time) - result = self.memory_service.do_operation(op_name=query) - print(result, flush=True) - except KeyboardInterrupt: - print("stop refresh!") - else: - print("unknown command received. Please try again!") - else: - print("unknown command received. Please try again!") - continue - - while True: - try: - if self.stream: - for msg in self.chat_with_memory(query=query): - print(msg.text, flush=True) - print() - else: - msg = self.chat_with_memory(query=query) - print(msg.text) - break - except KeyboardInterrupt: - questionary.print("User interrupt occurred.") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break - except Exception as e: - questionary.print(f"An exception occurred when running chat_with_memory(): {e}") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break diff --git a/memory_scope/chat_v2/global_context.py b/memory_scope/chat_v2/global_context.py deleted file mode 100644 index c38ae6c3..00000000 --- a/memory_scope/chat_v2/global_context.py +++ /dev/null @@ -1,27 +0,0 @@ -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - -from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat -from memory_scope.enumeration.language_enum import LanguageEnum -from memory_scope.memory.service.base_memory_service import BaseMemoryService -from memory_scope.models.base_model import BaseModel -from memory_scope.storage.base_monitor import BaseMonitor -from memory_scope.storage.base_vector_store import BaseVectorStore - - -class GlobalContext(object): - def __init__(self): - self.global_config: Dict[str, Any] = {} - self.worker_config: Dict[str, Dict[str, Any]] = {} - - self.memory_service_dict: Dict[str, BaseMemoryService] = {} - self.model_dict: Dict[str, BaseModel] = {} - self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} - - self.vector_store: BaseVectorStore | None = None - self.monitor: BaseMonitor | None = None - self.thread_pool: ThreadPoolExecutor | None = None - self.language: LanguageEnum = LanguageEnum.EN - - -G_CONTEXT = GlobalContext() diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 8545315f..a4a04316 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -9,7 +9,7 @@ from typing import Dict, Any import fire import yaml -from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.chat.global_context import G_CONTEXT from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.utils.logger import Logger from memory_scope.utils.tool_functions import init_instance_by_config diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index 74ce7943..8b61b67f 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -4,7 +4,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from itertools import zip_longest from typing import Dict, Any, List -from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.chat.global_context import G_CONTEXT from memory_scope.constants.common_constants import WORKFLOW_NAME from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.utils.logger import Logger diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 9eac30aa..78cc84fb 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -1,6 +1,6 @@ import time -from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.chat.global_context import G_CONTEXT from memory_scope.constants.common_constants import RESULT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index a47f61ad..eecc45c5 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -1,7 +1,7 @@ import time from typing import List -from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.chat.global_context import G_CONTEXT from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index c5281fbd..d5c64a4f 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -19,22 +19,10 @@ class BaseVectorStore(metaclass=ABCMeta): @abstractmethod def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): - """ - :param text: - :param limit_size: - :param filter_dict: - :return: - """ pass @abstractmethod async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): - """ - :param text: - :param limit_size: - :param filter_dict: - :return: - """ pass @abstractmethod diff --git a/memory_scope/worker/summary_short/__init__.py b/memory_scope/worker/summary_short/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/memory_scope/chat_v2/__init__.py b/old/worker/__init__.py similarity index 100% rename from memory_scope/chat_v2/__init__.py rename to old/worker/__init__.py diff --git a/memory_scope/worker/base_worker.py b/old/worker/base_worker.py similarity index 100% rename from memory_scope/worker/base_worker.py rename to old/worker/base_worker.py diff --git a/memory_scope/worker/dummy_worker.py b/old/worker/dummy_worker.py similarity index 100% rename from memory_scope/worker/dummy_worker.py rename to old/worker/dummy_worker.py diff --git a/memory_scope/worker/__init__.py b/old/worker/es/__init__.py similarity index 100% rename from memory_scope/worker/__init__.py rename to old/worker/es/__init__.py diff --git a/memory_scope/worker/es/es_insight_worker.py b/old/worker/es/es_insight_worker.py similarity index 100% rename from memory_scope/worker/es/es_insight_worker.py rename to old/worker/es/es_insight_worker.py diff --git a/memory_scope/worker/es/es_new_obs_worker.py b/old/worker/es/es_new_obs_worker.py similarity index 100% rename from memory_scope/worker/es/es_new_obs_worker.py rename to old/worker/es/es_new_obs_worker.py diff --git a/memory_scope/worker/es/es_not_reflected_worker.py b/old/worker/es/es_not_reflected_worker.py similarity index 100% rename from memory_scope/worker/es/es_not_reflected_worker.py rename to old/worker/es/es_not_reflected_worker.py diff --git a/memory_scope/worker/es/es_similar_worker.py b/old/worker/es/es_similar_worker.py similarity index 100% rename from memory_scope/worker/es/es_similar_worker.py rename to old/worker/es/es_similar_worker.py diff --git a/memory_scope/worker/es/es_today_obs_worker.py b/old/worker/es/es_today_obs_worker.py similarity index 100% rename from memory_scope/worker/es/es_today_obs_worker.py rename to old/worker/es/es_today_obs_worker.py diff --git a/memory_scope/worker/es/load_profile_worker.py b/old/worker/es/load_profile_worker.py similarity index 100% rename from memory_scope/worker/es/load_profile_worker.py rename to old/worker/es/load_profile_worker.py diff --git a/memory_scope/worker/memory_base_worker.py b/old/worker/memory_base_worker.py similarity index 100% rename from memory_scope/worker/memory_base_worker.py rename to old/worker/memory_base_worker.py diff --git a/memory_scope/worker/es/__init__.py b/old/worker/retrieve/__init__.py similarity index 100% rename from memory_scope/worker/es/__init__.py rename to old/worker/retrieve/__init__.py diff --git a/memory_scope/worker/retrieve/extract_time_worker.py b/old/worker/retrieve/extract_time_worker.py similarity index 100% rename from memory_scope/worker/retrieve/extract_time_worker.py rename to old/worker/retrieve/extract_time_worker.py diff --git a/memory_scope/worker/retrieve/fuse_rerank_worker.py b/old/worker/retrieve/fuse_rerank_worker.py similarity index 100% rename from memory_scope/worker/retrieve/fuse_rerank_worker.py rename to old/worker/retrieve/fuse_rerank_worker.py diff --git a/memory_scope/worker/retrieve/memory_store_worker.py b/old/worker/retrieve/memory_store_worker.py similarity index 100% rename from memory_scope/worker/retrieve/memory_store_worker.py rename to old/worker/retrieve/memory_store_worker.py diff --git a/memory_scope/worker/retrieve/semantic_rank_worker.py b/old/worker/retrieve/semantic_rank_worker.py similarity index 100% rename from memory_scope/worker/retrieve/semantic_rank_worker.py rename to old/worker/retrieve/semantic_rank_worker.py diff --git a/memory_scope/worker/retrieve/__init__.py b/old/worker/summary_long/__init__.py similarity index 100% rename from memory_scope/worker/retrieve/__init__.py rename to old/worker/summary_long/__init__.py diff --git a/memory_scope/worker/summary_long/get_insight_worker.py b/old/worker/summary_long/get_insight_worker.py similarity index 100% rename from memory_scope/worker/summary_long/get_insight_worker.py rename to old/worker/summary_long/get_insight_worker.py diff --git a/memory_scope/worker/summary_long/get_reflection_worker.py b/old/worker/summary_long/get_reflection_worker.py similarity index 100% rename from memory_scope/worker/summary_long/get_reflection_worker.py rename to old/worker/summary_long/get_reflection_worker.py diff --git a/memory_scope/worker/summary_long/long_contra_repeat_worker.py b/old/worker/summary_long/long_contra_repeat_worker.py similarity index 100% rename from memory_scope/worker/summary_long/long_contra_repeat_worker.py rename to old/worker/summary_long/long_contra_repeat_worker.py diff --git a/memory_scope/worker/summary_long/summary_collect_worker.py b/old/worker/summary_long/summary_collect_worker.py similarity index 100% rename from memory_scope/worker/summary_long/summary_collect_worker.py rename to old/worker/summary_long/summary_collect_worker.py diff --git a/memory_scope/worker/summary_long/update_insight_worker.py b/old/worker/summary_long/update_insight_worker.py similarity index 100% rename from memory_scope/worker/summary_long/update_insight_worker.py rename to old/worker/summary_long/update_insight_worker.py diff --git a/memory_scope/worker/summary_long/update_profile_worker.py b/old/worker/summary_long/update_profile_worker.py similarity index 100% rename from memory_scope/worker/summary_long/update_profile_worker.py rename to old/worker/summary_long/update_profile_worker.py diff --git a/memory_scope/worker/summary_long/__init__.py b/old/worker/summary_short/__init__.py similarity index 100% rename from memory_scope/worker/summary_long/__init__.py rename to old/worker/summary_short/__init__.py diff --git a/memory_scope/worker/summary_short/contra_repeat_worker.py b/old/worker/summary_short/contra_repeat_worker.py similarity index 100% rename from memory_scope/worker/summary_short/contra_repeat_worker.py rename to old/worker/summary_short/contra_repeat_worker.py diff --git a/memory_scope/worker/summary_short/get_observation_with_time_worker.py b/old/worker/summary_short/get_observation_with_time_worker.py similarity index 100% rename from memory_scope/worker/summary_short/get_observation_with_time_worker.py rename to old/worker/summary_short/get_observation_with_time_worker.py diff --git a/memory_scope/worker/summary_short/get_observation_worker.py b/old/worker/summary_short/get_observation_worker.py similarity index 100% rename from memory_scope/worker/summary_short/get_observation_worker.py rename to old/worker/summary_short/get_observation_worker.py diff --git a/memory_scope/worker/summary_short/info_filter_worker.py b/old/worker/summary_short/info_filter_worker.py similarity index 100% rename from memory_scope/worker/summary_short/info_filter_worker.py rename to old/worker/summary_short/info_filter_worker.py From 093946e999d7be1ec40266b2d6a447980e0fa91c Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 14:46:41 +0800 Subject: [PATCH 30/41] [dev] modify test for llm --- .../models/llama_index_generation_model.py | 29 +++++++++++-------- memory_scope/models/model_response.py | 3 +- tests/models/test_models_lli_generation.py | 15 ++++++---- 3 files changed, 28 insertions(+), 19 deletions(-) diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index 56f0b88b..af1cb59a 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -1,11 +1,14 @@ +import datetime from typing import List, Dict from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse from llama_index.llms.dashscope import DashScope +from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY from memory_scope.models.model_response import ModelResponse, ModelResponseGen +from memory_scope.scheme.message import Message class LlamaIndexGenerationModel(BaseModel): @@ -15,16 +18,16 @@ class LlamaIndexGenerationModel(BaseModel): def before_call(self, **kwargs) -> None: prompt: str = kwargs.pop("prompt", "") - messages: List[Dict[str, str]] = kwargs.pop("messages", []) + messages: List[Message] | List[Dict[str, str]] = kwargs.pop("messages", []) if prompt: input_text = prompt - input_type = 'prompt' + input_type = "prompt" llama_input = input_text elif messages: input_text = messages - input_type = 'messages' - llama_input = [ChatMessage(role=x.role, content=x.content) for x in input_text] + input_type = "messages" + llama_input = [ChatMessage(role=x["role"], content=x["content"]) for x in input_text] else: raise RuntimeError("prompt and messages is both empty!") @@ -34,25 +37,27 @@ class LlamaIndexGenerationModel(BaseModel): model_response: ModelResponse, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: + now_ts = datetime.datetime.now() + model_response.message = Message(role=MessageRoleEnum.ASSISTANT, + content="", + time_created=int(now_ts.timestamp())) + call_result = model_response.raw if stream: def gen() -> ModelResponseGen: - text = "" for response in call_result: - delta = response.delta - text += delta - model_response.text = text - model_response.delta = delta + model_response.message.content += response.delta + model_response.delta = response.delta yield model_response return gen() else: if isinstance(call_result, CompletionResponse): - content = call_result.text + model_response.message.content = call_result.text elif isinstance(call_result, ChatResponse): - content = call_result.message.content + model_response.message.content = call_result.message.content else: raise NotImplementedError - model_response.text = content + return model_response def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: diff --git a/memory_scope/models/model_response.py b/memory_scope/models/model_response.py index 958356db..33047529 100644 --- a/memory_scope/models/model_response.py +++ b/memory_scope/models/model_response.py @@ -4,10 +4,11 @@ from typing import Generator, List, Dict, Any from pydantic import BaseModel, Field from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.scheme.message import Message class ModelResponse(BaseModel): - text: str = Field("", description="generation model result") + message: Message | None = Field(None, description="generation model result") delta: str = Field("", description="New text that just streamed in (only used when streaming)") diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index 5b3329b2..cd6cec73 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -8,19 +8,21 @@ class TestLLILLM(unittest.TestCase): def setUp(self): config = { - "method_type": "DashScope", + "module_name": "dashscope_generation", "model_name": "qwen-max", "clazz": "models.llama_index_generation_model" } self.llm = LlamaIndexGenerationModel(**config) + @unittest.skip("tmp") def test_llm_prompt(self): prompt = "你是谁?" ans = self.llm.call( stream=False, prompt=prompt ) - print(ans.text) + print(ans.message.content) + @unittest.skip("tmp") def test_llm_messages(self): messages = [{"role": "system", "content": "you are a helpful assistant."}, @@ -29,7 +31,8 @@ class TestLLILLM(unittest.TestCase): stream=False, messages=messages ) - print(ans.text) + print(ans.message.content) + @unittest.skip("tmp") def test_llm_prompt_stream(self): prompt = "你如何看待黄金上涨?" @@ -43,8 +46,8 @@ class TestLLILLM(unittest.TestCase): sys.stdout.write(a.delta) sys.stdout.flush() time.sleep(0.1) - @unittest.skip("tmp") - def test_llm_messages(self): + + def test_llm_messages_stream(self): messages = [{"role": "system", "content": "you are a helpful assistant."}, {"role": "user", "content": "你如何看待黄金上涨?"}] ans = self.llm.call( @@ -56,4 +59,4 @@ class TestLLILLM(unittest.TestCase): for a in ans: sys.stdout.write(a.delta) sys.stdout.flush() - time.sleep(0.1) \ No newline at end of file + time.sleep(0.1) From f16034d6a1bd79902c24a3683061f2c52240fc9f Mon Sep 17 00:00:00 2001 From: hs Date: Thu, 27 Jun 2024 14:53:02 +0800 Subject: [PATCH 31/41] [dev] update cli, delete try catch in while loop --- memory_scope/chat/cli_memory_chat.py | 41 ++++++++++++++-------------- 1 file changed, 21 insertions(+), 20 deletions(-) diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 0f5ff092..69db108b 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -64,7 +64,7 @@ class CliMemoryChat(BaseMemoryChat): for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): yield result - self.memory_service.add_messages(result.text) + self.memory_service.add_messages(result.message) def run(self): self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) @@ -118,22 +118,23 @@ class CliMemoryChat(BaseMemoryChat): continue while True: - try: - if self.stream: - for msg in self.chat_with_memory(query=query): - print(msg.text, flush=True) - print() - else: - msg = self.chat_with_memory(query=query) - print(msg.text) - break - except KeyboardInterrupt: - questionary.print("User interrupt occurred.") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break - except Exception as e: - questionary.print(f"An exception occurred when running chat_with_memory(): {e}") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break + # try: + if self.stream: + for msg in self.chat_with_memory(query=query): + print(msg.delta, end="") + print() + else: + msg = self.chat_with_memory(query=query) + print(msg.message.content) + break + # except KeyboardInterrupt: + # questionary.print("User interrupt occurred.") + # retry = questionary.confirm("Retry chat_with_memory()?").ask() + # if not retry: + # break + # except Exception as e: + # questionary.print(f"An exception occurred when running chat_with_memory(): {e}") + # # retry = questionary.confirm("Retry chat_with_memory()?").ask() + # # if not retry: + # # break + # raise e From 862b55f2b087e353f7a5ff021097b0b5007be931 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 16:43:59 +0800 Subject: [PATCH 32/41] [dev] rename prepare service function & reformat llm model data struct --- memory_scope/chat/cli_memory_chat.py | 2 +- .../memory/service/base_memory_service.py | 2 +- .../memory/service/chat_memory_service.py | 7 ++- memory_scope/models/base_model.py | 2 +- .../models/llama_index_embedding_model.py | 2 +- .../models/llama_index_generation_model.py | 25 ++++----- memory_scope/models/llama_index_rank_model.py | 18 +++---- memory_scope/scheme/message.py | 5 +- .../{models => scheme}/model_response.py | 0 tests/models/test_models_lli_embedding.py | 6 ++- tests/models/test_models_lli_generation.py | 53 ++++++++----------- ..._lli_rerank.py => test_models_lli_rank.py} | 5 +- 12 files changed, 59 insertions(+), 68 deletions(-) rename memory_scope/{models => scheme}/model_response.py (100%) rename tests/models/{test_models_lli_rerank.py => test_models_lli_rank.py} (81%) diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 69db108b..b5b7c964 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -31,7 +31,7 @@ class CliMemoryChat(BaseMemoryChat): def memory_service(self) -> BaseMemoryService: if isinstance(self._memory_service, str): self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] - self._memory_service.prepare_service() + self._memory_service.start_service() return self._memory_service @property diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 50fefca8..2a8e78bc 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -31,7 +31,7 @@ class BaseMemoryService(metaclass=ABCMeta): def add_messages(self, messages: List[Message] | Message): raise NotImplementedError - def prepare_service(self): + def start_service(self): pass @abstractmethod diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index a2197e91..f3e9294d 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -37,7 +37,7 @@ class ChatMemoryService(BaseMemoryService): for _ in range(gap_size): self.chat_messages.pop(0) - def prepare_service(self): + def start_service(self): for _, operation in self._operation_dict.items(): operation.init_workflow() if operation.operation_type == "backend": @@ -48,3 +48,8 @@ class ChatMemoryService(BaseMemoryService): self.logger.warning(f"op_name={op_name} is not inited!") return return self._operation_dict[op_name].run_operation() + + def stop_service(self): + for _, operation in self._operation_dict.items(): + if operation.operation_type == "backend": + operation.stop_operation_backend() diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index cd63c57a..1f5703b1 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -3,7 +3,7 @@ import time from abc import abstractmethod, ABCMeta from memory_scope.enumeration.model_enum import ModelEnum -from memory_scope.models.model_response import ModelResponse, ModelResponseGen +from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen from memory_scope.utils.logger import Logger from memory_scope.utils.registry import Registry from memory_scope.utils.timer import Timer diff --git a/memory_scope/models/llama_index_embedding_model.py b/memory_scope/models/llama_index_embedding_model.py index bb4d35b4..bbecbce9 100644 --- a/memory_scope/models/llama_index_embedding_model.py +++ b/memory_scope/models/llama_index_embedding_model.py @@ -4,7 +4,7 @@ from llama_index.embeddings.dashscope import DashScopeEmbedding from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY -from memory_scope.models.model_response import ModelResponse +from memory_scope.scheme.model_response import ModelResponse class LlamaIndexEmbeddingModel(BaseModel): diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index af1cb59a..a3c6f45d 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -1,5 +1,5 @@ import datetime -from typing import List, Dict +from typing import List from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse from llama_index.llms.dashscope import DashScope @@ -7,7 +7,7 @@ from llama_index.llms.dashscope import DashScope from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY -from memory_scope.models.model_response import ModelResponse, ModelResponseGen +from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen from memory_scope.scheme.message import Message @@ -16,31 +16,25 @@ class LlamaIndexGenerationModel(BaseModel): MODEL_REGISTRY.register("dashscope_generation", DashScope) - def before_call(self, **kwargs) -> None: + def before_call(self, **kwargs): prompt: str = kwargs.pop("prompt", "") - messages: List[Message] | List[Dict[str, str]] = kwargs.pop("messages", []) + messages: List[Message] = kwargs.pop("messages", []) if prompt: - input_text = prompt - input_type = "prompt" - llama_input = input_text + self.data = {"prompt": prompt} elif messages: - input_text = messages - input_type = "messages" - llama_input = [ChatMessage(role=x["role"], content=x["content"]) for x in input_text] + self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]} else: raise RuntimeError("prompt and messages is both empty!") - self.data = {input_type: llama_input} - def after_call(self, model_response: ModelResponse, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: - now_ts = datetime.datetime.now() + model_response.message = Message(role=MessageRoleEnum.ASSISTANT, content="", - time_created=int(now_ts.timestamp())) + time_created=int(datetime.datetime.now().timestamp())) call_result = model_response.raw if stream: @@ -61,11 +55,10 @@ class LlamaIndexGenerationModel(BaseModel): return model_response def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: - assert "prompt" in self.data or "messages" in self.data results = ModelResponse(m_type=self.m_type) - if 'prompt' in self.data: + if "prompt" in self.data: if stream: response = self.model.stream_complete(**self.data) else: diff --git a/memory_scope/models/llama_index_rank_model.py b/memory_scope/models/llama_index_rank_model.py index 095fda5e..71e4acef 100644 --- a/memory_scope/models/llama_index_rank_model.py +++ b/memory_scope/models/llama_index_rank_model.py @@ -6,7 +6,7 @@ from llama_index.postprocessor.dashscope_rerank import DashScopeRerank from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY -from memory_scope.models.model_response import ModelResponse +from memory_scope.scheme.model_response import ModelResponse class LlamaIndexRankModel(BaseModel): @@ -15,22 +15,18 @@ class LlamaIndexRankModel(BaseModel): MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank) def before_call(self, **kwargs) -> None: - assert "query" in kwargs or "documents" in kwargs query: str = kwargs.pop("query", "") documents: List[str] = kwargs.pop("documents", []) + if isinstance(documents, str): + documents = [documents] - assert ( - query and documents - ), f"query or documents is empty! query={query}, documents={len(documents)}" + assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}" # using -1.0 as dummy scores nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents] self._get_documents_mapping(documents) - self.data = { - "nodes": nodes, - "query_str": query, - } + self.data = {"nodes": nodes, "query_str": query} def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse: if not model_response.rank_scores: @@ -43,9 +39,7 @@ class LlamaIndexRankModel(BaseModel): return model_response def _call(self, **kwargs) -> ModelResponse: - return ModelResponse( - m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data) - ) + return ModelResponse(m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data)) async def _async_call(self, **kwargs) -> ModelResponse: raise NotImplementedError diff --git a/memory_scope/scheme/message.py b/memory_scope/scheme/message.py index ad180772..f08b3290 100644 --- a/memory_scope/scheme/message.py +++ b/memory_scope/scheme/message.py @@ -1,3 +1,5 @@ +import datetime + from pydantic import Field, BaseModel @@ -6,6 +8,7 @@ class Message(BaseModel): content: str = Field(..., description="The body of the message") - time_created: int = Field(..., description="Timestamp when the message was created") + time_created: int = Field(int(datetime.datetime.now().timestamp()), + description="Timestamp when the message was created") memorized: bool = Field(False, description="indicate whether message is memorized") diff --git a/memory_scope/models/model_response.py b/memory_scope/scheme/model_response.py similarity index 100% rename from memory_scope/models/model_response.py rename to memory_scope/scheme/model_response.py diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index a8078cfd..1bbfae5e 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -1,3 +1,7 @@ +import sys + +sys.path.append(".") # noqa: E402 + import asyncio import unittest @@ -9,7 +13,7 @@ class TestLLIEmbedding(unittest.TestCase): def setUp(self): config = { - "method_type": "DashScopeEmbedding", + "module_name": "dashscope_embedding", "model_name": "text-embedding-v2", "clazz": "models.base_embedding_model" } diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index cd6cec73..6f7b43cd 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -1,6 +1,13 @@ -import unittest +import sys +sys.path.append(".") # noqa: E402 + +import unittest +import time + +from memory_scope.scheme.message import Message from memory_scope.models.llama_index_generation_model import LlamaIndexGenerationModel +from memory_scope.utils.logger import Logger class TestLLILLM(unittest.TestCase): @@ -13,50 +20,36 @@ class TestLLILLM(unittest.TestCase): "clazz": "models.llama_index_generation_model" } self.llm = LlamaIndexGenerationModel(**config) + self.logger = Logger.get_logger() - @unittest.skip("tmp") def test_llm_prompt(self): prompt = "你是谁?" - ans = self.llm.call( - stream=False, - prompt=prompt - ) - print(ans.message.content) + ans = self.llm.call(stream=False, prompt=prompt) + self.logger.info(ans.message.content) - @unittest.skip("tmp") def test_llm_messages(self): - messages = [{"role": "system", "content": "you are a helpful assistant."}, - {"role": "user", "content": "你是谁?"}] - ans = self.llm.call( - stream=False, - messages=messages - ) - print(ans.message.content) + messages = [Message(role="system", content="you are a helpful assistant."), + Message(role="user", content="你是谁?")] + ans = self.llm.call(stream=False, messages=messages) + self.logger.info(ans.message.content) - @unittest.skip("tmp") def test_llm_prompt_stream(self): prompt = "你如何看待黄金上涨?" - ans = self.llm.call( - stream=True, - prompt=prompt - ) - import sys - import time + ans = self.llm.call(stream=True, prompt=prompt) + self.logger.info("-----start-----") for a in ans: sys.stdout.write(a.delta) sys.stdout.flush() time.sleep(0.1) + self.logger.info("-----end-----") def test_llm_messages_stream(self): - messages = [{"role": "system", "content": "you are a helpful assistant."}, - {"role": "user", "content": "你如何看待黄金上涨?"}] - ans = self.llm.call( - stream=True, - messages=messages - ) - import sys - import time + messages = [Message(role="system", content="you are a helpful assistant."), + Message(role="user", content="你如何看待黄金上涨?")] + ans = self.llm.call(stream=True, messages=messages) + self.logger.info("-----start-----") for a in ans: sys.stdout.write(a.delta) sys.stdout.flush() time.sleep(0.1) + self.logger.info("-----end-----") diff --git a/tests/models/test_models_lli_rerank.py b/tests/models/test_models_lli_rank.py similarity index 81% rename from tests/models/test_models_lli_rerank.py rename to tests/models/test_models_lli_rank.py index 1238f7f8..1c3b76a5 100644 --- a/tests/models/test_models_lli_rerank.py +++ b/tests/models/test_models_lli_rank.py @@ -1,7 +1,6 @@ -import json import unittest -from memory_scope.models.llama_index_rerank_model import LlamaIndexRankModel +from memory_scope.models.llama_index_rank_model import LlamaIndexRankModel class TestLLIReRank(unittest.TestCase): @@ -9,7 +8,7 @@ class TestLLIReRank(unittest.TestCase): def setUp(self): config = { - "method_type": "DashScopeRerank", + "module_name": "dashscope_rank", "model_name": "gte-rerank", "clazz": "models.llama_index_rerank_model" } From ee036dc0a8857cbef5c288d9c98b2eb5d8649d34 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 16:46:10 +0800 Subject: [PATCH 33/41] [dev] add logger to test --- memory_scope/cli.py | 2 +- tests/models/test_models_lli_embedding.py | 9 +++++---- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/memory_scope/cli.py b/memory_scope/cli.py index a4a04316..a76c61db 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -19,7 +19,7 @@ class CliJob(object): def __init__(self): self.config: Dict[str, Any] = {} - self.logger: Logger = Logger.get_logger("cli_job") + self.logger: Logger = Logger.get_logger("cli_job", to_stream=False) def load_config(self, path: str): with open(path) as f: diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index 1bbfae5e..5bd3f28b 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -6,7 +6,7 @@ import asyncio import unittest from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel - +from memory_scope.utils.logger import Logger class TestLLIEmbedding(unittest.TestCase): """Tests for LlamaIndexEmbeddingModel""" @@ -18,21 +18,22 @@ class TestLLIEmbedding(unittest.TestCase): "clazz": "models.base_embedding_model" } self.emb = LlamaIndexEmbeddingModel(**config) + self.logger = Logger.get_logger() def test_single_embedding(self): text = "您吃了吗?" result = self.emb.call(text=text) - print(result) + self.logger.info(result) def test_batch_embedding(self): texts = ["您吃了吗?", "吃了吗您?"] result = self.emb.call(text=texts) - print(result) + self.logger.info(result) def test_async_embedding(self): texts = ["您吃了吗?", "吃了吗您?"] # 调用异步函数并等待其结果 result = asyncio.run(self.emb.async_call(text=texts)) - print(result) + self.logger.info(result) From 264a2de946acf4d89b0a869c99fde16ab6fb5262 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 17:39:58 +0800 Subject: [PATCH 34/41] [dev] add process commands to cli --- memory_scope/chat/cli_memory_chat.py | 134 ++++++++++-------- .../memory/service/base_memory_service.py | 3 + tests/models/test_models_lli_embedding.py | 7 +- 3 files changed, 82 insertions(+), 62 deletions(-) diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index b5b7c964..53cfa650 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -11,7 +11,8 @@ from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT from memory_scope.scheme.message import Message -from ..models.model_response import ModelResponse, ModelResponseGen +from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen +from memory_scope.utils.logger import Logger class CliMemoryChat(BaseMemoryChat): @@ -27,6 +28,8 @@ class CliMemoryChat(BaseMemoryChat): self.stream: bool = stream self.kwargs: dict = kwargs + self.logger = Logger.get_logger() + @property def memory_service(self) -> BaseMemoryService: if isinstance(self._memory_service, str): @@ -66,75 +69,84 @@ class CliMemoryChat(BaseMemoryChat): self.memory_service.add_messages(result.message) + def process_commands(self, query: str) -> bool: + continue_run = True + query_split = query.lstrip("/").lower().split(" ") + query = query_split[0] + args = query_split[1:] + if query == "exit": + self.memory_service.stop_service() + continue_run = False + + elif query == "help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(cmd, "bold") + questionary.print(f" {desc}") + + elif query == "stream": + self.stream = bool(args[0]) + questionary.print(f"stream: {self.stream}") + + elif query in self.memory_service.op_description_dict: + if not args: + result = self.memory_service.do_operation(op_name=query) + questionary.print(result) + + elif args[0].isdigit(): + refresh_time = int(args[0]) + while True: + time.sleep(refresh_time) + result = self.memory_service.do_operation(op_name=query) + questionary.print(result, flush=True) + + else: + questionary.print("unknown command received. Please try again!") + + else: + questionary.print("unknown command received. Please try again!") + + return continue_run + def run(self): self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) while True: - query = questionary.text( - "Please enter your message or command:", - multiline=False, - qmark=">", - ).ask() + try: + query = questionary.text( + message="Please enter your message or command:", + multiline=False, + qmark=">", + ).ask() + query: str = query.strip() - query: str = query.rstrip() + if query == "": + questionary.print("Empty input received. Please try again!") + continue - if query == "": - print("Empty input received. Please try again!") - continue - - # handle cli / commands with memory ops - if query.startswith("/"): - query_split = query.lstrip("/").lower().split(" ") - query = query_split[0] - args = query_split[1:] - if query == "exit": - break - elif query == "help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(cmd, "bold") - print(f" {desc}") - elif query == "stream": - questionary.print(f"stream: {self.stream}") - self.stream = ~self.stream - elif query in self.memory_service.op_description_dict: - if not args: - result = self.memory_service.do_operation(op_name=query) - print(result) - - elif args[0].isdigit(): - refresh_time = int(args[0]) - try: - while True: - time.sleep(refresh_time) - result = self.memory_service.do_operation(op_name=query) - print(result, flush=True) - except KeyboardInterrupt: - print("stop refresh!") + # handle cli / commands with memory ops + if query.startswith("/"): + if self.process_commands(query=query): + continue else: - print("unknown command received. Please try again!") - else: - print("unknown command received. Please try again!") - continue + break - while True: - # try: if self.stream: for msg in self.chat_with_memory(query=query): - print(msg.delta, end="") - print() + questionary.print(msg.delta, end="") + questionary.print("") else: msg = self.chat_with_memory(query=query) - print(msg.message.content) - break - # except KeyboardInterrupt: - # questionary.print("User interrupt occurred.") - # retry = questionary.confirm("Retry chat_with_memory()?").ask() - # if not retry: - # break - # except Exception as e: - # questionary.print(f"An exception occurred when running chat_with_memory(): {e}") - # # retry = questionary.confirm("Retry chat_with_memory()?").ask() - # # if not retry: - # # break - # raise e + questionary.print(msg.message.content) + + except KeyboardInterrupt: + questionary.print("User interrupt occurred.") + is_exit = questionary.confirm("continue exit?").ask() + if is_exit: + self.memory_service.stop_service() + break + + except Exception as e: + questionary.print(f"An exception occurred when running cli memory chat. args={e.args}") + self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}") + continue diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 2a8e78bc..41a8c141 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -47,3 +47,6 @@ class BaseMemoryService(metaclass=ABCMeta): def read_memory(self): assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" return self.do_operation(self.read_memory_key) + + def stop_service(self): + pass diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index 5bd3f28b..c71236f5 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -8,6 +8,7 @@ import unittest from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel from memory_scope.utils.logger import Logger + class TestLLIEmbedding(unittest.TestCase): """Tests for LlamaIndexEmbeddingModel""" @@ -18,17 +19,20 @@ class TestLLIEmbedding(unittest.TestCase): "clazz": "models.base_embedding_model" } self.emb = LlamaIndexEmbeddingModel(**config) + print() self.logger = Logger.get_logger() def test_single_embedding(self): text = "您吃了吗?" result = self.emb.call(text=text) - self.logger.info(result) + self.logger.info(result.m_type) + self.logger.info(len(result.embedding_results)) def test_batch_embedding(self): texts = ["您吃了吗?", "吃了吗您?"] result = self.emb.call(text=texts) + print() self.logger.info(result) def test_async_embedding(self): @@ -36,4 +40,5 @@ class TestLLIEmbedding(unittest.TestCase): "吃了吗您?"] # 调用异步函数并等待其结果 result = asyncio.run(self.emb.async_call(text=texts)) + print() self.logger.info(result) From cffeac8565e55d336a5596e44c9b7559048c14c0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 23:21:44 +0800 Subject: [PATCH 35/41] [dev] print char logo in shell env --- memory_scope/chat/cli_memory_chat.py | 45 ++++++++++++------- memory_scope/cli.py | 2 +- .../memory/operation/summary_memory.py | 9 +++- memory_scope/memory/operation/write_memory.py | 9 +++- memory_scope/memory/worker/dummy_worker.py | 5 ++- memory_scope/utils/tool_functions.py | 36 ++++++++++++--- 6 files changed, 78 insertions(+), 28 deletions(-) diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 53cfa650..5e0ca20f 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -13,6 +13,7 @@ from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT from memory_scope.scheme.message import Message from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen from memory_scope.utils.logger import Logger +from memory_scope.utils.tool_functions import char_logo class CliMemoryChat(BaseMemoryChat): @@ -28,8 +29,13 @@ class CliMemoryChat(BaseMemoryChat): self.stream: bool = stream self.kwargs: dict = kwargs + self._logo = char_logo("MemoryScope") self.logger = Logger.get_logger() + def print_logo(self): + for line in self._logo: + print(line) + @property def memory_service(self) -> BaseMemoryService: if isinstance(self._memory_service, str): @@ -63,11 +69,12 @@ class CliMemoryChat(BaseMemoryChat): self.memory_service.add_messages(new_message) related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) + result = self.generation_model.call(messages=[system_message, new_message], stream=self.stream) if self.stream: - for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream): - yield result - - self.memory_service.add_messages(result.message) + for _ in result: + yield _ + else: + return result def process_commands(self, query: str) -> bool: continue_run = True @@ -81,8 +88,8 @@ class CliMemoryChat(BaseMemoryChat): elif query == "help": questionary.print("CLI commands", "bold") for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(cmd, "bold") - questionary.print(f" {desc}") + questionary.print(text=f" /{cmd}:", style="bold") + questionary.print(text=f" {desc}") elif query == "stream": self.stream = bool(args[0]) @@ -109,21 +116,18 @@ class CliMemoryChat(BaseMemoryChat): return continue_run def run(self): - self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) + self.print_logo() + self.USER_COMMANDS.update(self.memory_service.op_description_dict) while True: try: - query = questionary.text( - message="Please enter your message or command:", - multiline=False, - qmark=">", - ).ask() - query: str = query.strip() + query = questionary.text(message="user:", multiline=False, qmark=">").ask() - if query == "": - questionary.print("Empty input received. Please try again!") + if not query: continue + query: str = query.strip() + # handle cli / commands with memory ops if query.startswith("/"): if self.process_commands(query=query): @@ -131,6 +135,9 @@ class CliMemoryChat(BaseMemoryChat): else: break + msg = None + questionary.print("> ", end="", style="fg:yellow") + questionary.print("assistant: ", end="", style="bold") if self.stream: for msg in self.chat_with_memory(query=query): questionary.print(msg.delta, end="") @@ -138,6 +145,7 @@ class CliMemoryChat(BaseMemoryChat): else: msg = self.chat_with_memory(query=query) questionary.print(msg.message.content) + self.memory_service.add_messages(msg.message) except KeyboardInterrupt: questionary.print("User interrupt occurred.") @@ -147,6 +155,9 @@ class CliMemoryChat(BaseMemoryChat): break except Exception as e: - questionary.print(f"An exception occurred when running cli memory chat. args={e.args}") - self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}") + line = f"An exception occurred when running cli memory chat. args={e.args}" + questionary.print(line) + self.logger.exception(line) continue + + questionary.print(f"A memory writing thread is still running, please be patient and wait!") diff --git a/memory_scope/cli.py b/memory_scope/cli.py index a76c61db..a0cdda2b 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,6 +1,6 @@ import sys -sys.path.append(".") +sys.path.append(".") # noqa: E402 import json from concurrent.futures import ThreadPoolExecutor diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 78cc84fb..bef933eb 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -39,8 +39,13 @@ class SummaryMemory(BaseWorkflow, BaseOperation): def _loop_operation(self): while self._loop_switch: - time.sleep(self.interval_time) - self.run_operation() + for _ in range(self.interval_time): + if self._loop_switch: + time.sleep(1) + else: + break + if self._loop_switch: + self.run_operation() def run_operation_backend(self): if not self._loop_switch: diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index eecc45c5..362ec10f 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -66,8 +66,13 @@ class WriteMemory(BaseWorkflow, BaseOperation): def _loop_operation(self): while self._loop_switch: - time.sleep(self.interval_time) - self.run_operation() + for _ in range(self.interval_time): + if self._loop_switch: + time.sleep(1) + else: + break + if self._loop_switch: + self.run_operation() def run_operation_backend(self): if not self._loop_switch: diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 5ae68efc..2564b4eb 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -1,3 +1,5 @@ +import datetime + from memory_scope.constants.common_constants import RESULT, WORKFLOW_NAME from memory_scope.memory.worker.base_worker import BaseWorker @@ -6,4 +8,5 @@ class DummyWorker(BaseWorker): def _run(self): workflow_name = self.get_context(WORKFLOW_NAME) self.logger.info(f"enter workflow={workflow_name}.dummy_worker!") - self.set_context(RESULT, f"test {workflow_name}") + ts = int(datetime.datetime.now().timestamp()) + self.set_context(RESULT, f"test {workflow_name} ts={ts}") diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index aeb8dfbc..a9d0a236 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -1,11 +1,19 @@ +import random import re +import time from copy import deepcopy from datetime import datetime from importlib import import_module +from typing import get_args + +import pyfiglet +from termcolor import colored, COLORS +from termcolor._types import Color from memory_scope.constants.common_constants import WEEKDAYS from memory_scope.enumeration.message_role_enum import MessageRoleEnum +ALL_COLORS = get_args(COLORS) def under_line_to_hump(underline_str): sub = re.sub(r"(_\w)", lambda x: x.group(1)[1].upper(), underline_str) @@ -68,11 +76,10 @@ def get_datetime_info_dict(parse_dt: datetime): } -def time_to_formatted_str( - time: datetime | str | int | float = None, - date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S - string_format: str = "", -) -> str: +def time_to_formatted_str(time: datetime | str | int | float = None, + date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S + string_format: str = "") -> str: + if isinstance(time, str | int | float): if isinstance(time, str): time = float(time) @@ -89,3 +96,22 @@ def time_to_formatted_str( return_str = string_format.format(**get_datetime_info_dict(current_dt)) return return_str + + +def char_logo(words: str, seed: int = time.time_ns(), color: Color = None): + font = pyfiglet.Figlet() + rendered_text = font.renderText(words) + colored_lines = [] + all_colors = list(COLORS.keys()) + random.seed = seed + for line in rendered_text.splitlines(): + line_color = color + if line_color is None: + random.shuffle(all_colors) + line_color = all_colors[0] + colored_line = "" + for char in line: + colored_char = colored(char, line_color, attrs=['bold']) + colored_line += colored_char + colored_lines.append(colored_line) + return colored_lines From 89498c6c3b1722fb35d4f83515d108827707fbe6 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 23:55:24 +0800 Subject: [PATCH 36/41] [dev] refresh print --- memory_scope/chat/cli_memory_chat.py | 10 ++++++---- memory_scope/cli.py | 3 +++ memory_scope/memory/worker/dummy_worker.py | 2 +- memory_scope/utils/tool_functions.py | 22 ++++++++++------------ 4 files changed, 20 insertions(+), 17 deletions(-) diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 5e0ca20f..8a8f3ea3 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -1,4 +1,5 @@ import datetime +import os import time from typing import List @@ -105,7 +106,9 @@ class CliMemoryChat(BaseMemoryChat): while True: time.sleep(refresh_time) result = self.memory_service.do_operation(op_name=query) - questionary.print(result, flush=True) + os.system('clear') + self.print_logo() + questionary.print(result) else: questionary.print("unknown command received. Please try again!") @@ -121,8 +124,7 @@ class CliMemoryChat(BaseMemoryChat): while True: try: - query = questionary.text(message="user:", multiline=False, qmark=">").ask() - + query = questionary.text(message="user:", multiline=False, qmark=">").unsafe_ask() if not query: continue @@ -149,7 +151,7 @@ class CliMemoryChat(BaseMemoryChat): except KeyboardInterrupt: questionary.print("User interrupt occurred.") - is_exit = questionary.confirm("continue exit?").ask() + is_exit = questionary.confirm("continue exit").unsafe_ask() if is_exit: self.memory_service.stop_service() break diff --git a/memory_scope/cli.py b/memory_scope/cli.py index a0cdda2b..427ad145 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -13,6 +13,7 @@ from memory_scope.chat.global_context import G_CONTEXT from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.utils.logger import Logger from memory_scope.utils.tool_functions import init_instance_by_config +from memory_scope.utils.timer import timer class CliJob(object): @@ -35,6 +36,7 @@ class CliJob(object): G_CONTEXT.language = LanguageEnum(global_config["language"]) G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"])) + @timer def init_global_content_by_config(self): # set global config self.set_global_config() @@ -63,6 +65,7 @@ class CliJob(object): def run(self, config: str): self.load_config(config) self.init_global_content_by_config() + with G_CONTEXT.thread_pool: memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 2564b4eb..0485ca89 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -9,4 +9,4 @@ class DummyWorker(BaseWorker): workflow_name = self.get_context(WORKFLOW_NAME) self.logger.info(f"enter workflow={workflow_name}.dummy_worker!") ts = int(datetime.datetime.now().timestamp()) - self.set_context(RESULT, f"test {workflow_name} ts={ts}") + self.set_context(RESULT, f"test {workflow_name} \nts={ts}") diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index a9d0a236..aed4aff1 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -4,16 +4,13 @@ import time from copy import deepcopy from datetime import datetime from importlib import import_module -from typing import get_args import pyfiglet from termcolor import colored, COLORS -from termcolor._types import Color from memory_scope.constants.common_constants import WEEKDAYS from memory_scope.enumeration.message_role_enum import MessageRoleEnum -ALL_COLORS = get_args(COLORS) def under_line_to_hump(underline_str): sub = re.sub(r"(_\w)", lambda x: x.group(1)[1].upper(), underline_str) @@ -39,7 +36,8 @@ def init_instance_by_config(config: dict, default_class_path: str = "memory_scop module = import_module(".".join(class_paths)) cls_name = under_line_to_hump(class_name) - return getattr(module, cls_name)(**config_copy, **kwargs) + config_copy.update(kwargs) + return getattr(module, cls_name)(**config_copy) def complete_config_name(config_name: str, suffix: str = ".json"): @@ -76,16 +74,16 @@ def get_datetime_info_dict(parse_dt: datetime): } -def time_to_formatted_str(time: datetime | str | int | float = None, +def time_to_formatted_str(dt: datetime | str | int | float = None, date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S string_format: str = "") -> str: - if isinstance(time, str | int | float): - if isinstance(time, str): - time = float(time) - current_dt = datetime.fromtimestamp(time) - elif isinstance(time, datetime): - current_dt = time + if isinstance(dt, str | int | float): + if isinstance(dt, str): + dt = float(dt) + current_dt = datetime.fromtimestamp(dt) + elif isinstance(dt, datetime): + current_dt = dt else: current_dt = datetime.now() @@ -98,7 +96,7 @@ def time_to_formatted_str(time: datetime | str | int | float = None, return return_str -def char_logo(words: str, seed: int = time.time_ns(), color: Color = None): +def char_logo(words: str, seed: int = time.time_ns(), color=None): font = pyfiglet.Figlet() rendered_text = font.renderText(words) colored_lines = [] From 5dbef26bae35afe704f3602317f90de67e87df07 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 10:26:02 +0800 Subject: [PATCH 37/41] [dev] rename llm result var name & assistant name --- memory_scope/chat/cli_memory_chat.py | 22 +++++++++++++++------- 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 8a8f3ea3..2eb2e21a 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -24,10 +24,18 @@ class CliMemoryChat(BaseMemoryChat): "stream": "get stream response" } - def __init__(self, memory_service: str, generation_model: str, stream: bool = True, **kwargs): + def __init__(self, + memory_service: str, + generation_model: str, + stream: bool = True, + human_name: str = "human", + assistant_name: str = "assistant", + **kwargs): self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model self.stream: bool = stream + self.human_name: str = human_name + self.assistant_name: str = assistant_name self.kwargs: dict = kwargs self._logo = char_logo("MemoryScope") @@ -66,16 +74,16 @@ class CliMemoryChat(BaseMemoryChat): return time_created = int(datetime.datetime.now().timestamp()) - new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) + new_message: Message = Message(role=MessageRoleEnum.USER.value, content=query, time_created=time_created) self.memory_service.add_messages(new_message) related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) - result = self.generation_model.call(messages=[system_message, new_message], stream=self.stream) + model_response = self.generation_model.call(messages=[system_message, new_message], stream=self.stream) if self.stream: - for _ in result: + for _ in model_response: yield _ else: - return result + return model_response def process_commands(self, query: str) -> bool: continue_run = True @@ -124,7 +132,7 @@ class CliMemoryChat(BaseMemoryChat): while True: try: - query = questionary.text(message="user:", multiline=False, qmark=">").unsafe_ask() + query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask() if not query: continue @@ -139,7 +147,7 @@ class CliMemoryChat(BaseMemoryChat): msg = None questionary.print("> ", end="", style="fg:yellow") - questionary.print("assistant: ", end="", style="bold") + questionary.print(f"{self.assistant_name}: ", end="", style="bold") if self.stream: for msg in self.chat_with_memory(query=query): questionary.print(msg.delta, end="") From 5a533b2721b9949a45e9899b9c04a191b6a5142b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 15:06:27 +0800 Subject: [PATCH 38/41] [dev] modify es search logic & add es test --- config/config.json | 84 ----------- config/config.yaml | 1 + memory_scope/models/base_model.py | 25 +++- memory_scope/scheme/memory_node.py | 6 + memory_scope/scheme/model_response.py | 8 +- memory_scope/storage/base_vector_store.py | 24 +-- .../llama_index_elastic_search_store.py | 140 ++++++++---------- tests/storages/test_storages_lli_es.py | 78 +++++----- 8 files changed, 136 insertions(+), 230 deletions(-) delete mode 100644 config/config.json diff --git a/config/config.json b/config/config.json deleted file mode 100644 index 3f895d4b..00000000 --- a/config/config.json +++ /dev/null @@ -1,84 +0,0 @@ -{ - "global_config": { - "language": "en", - "max_workers": 5, - "dash_scope_apikey": null, - "open_ai_apikey": null - }, - "memory_chat": { - "cli_memory_chat": { - "class": "chat.cli_memory_chat", - "memory_service": "memory_chat_service", - "generation_model": "dashscope_generation" - } - }, - "memory_service": { - "memory_chat_service": { - "class": "memory.service.chat_memory_service", - "history_msg_count": 32, - "contextual_msg_count": 6, - "read_memory_key": "read_memory", - "memory_operations": { - "read_message": { - "class": "memory.operation.read_memory", - "workflow": "dummy_worker", - "description": "read session messages of the user" - }, - "read_memory": { - "class": "memory.operation.read_memory", - "workflow": "dummy_worker", - "description": "read related memories of the user" - }, - "list_memory": { - "class": "memory.operation.read_memory", - "workflow": "dummy_worker", - "description": "read all memories of the user" - }, - "write_memory": { - "class": "memory.operation.write_memory", - "workflow": "dummy_worker", - "description": "write observation memories of the user", - "interval_time": 60 - }, - "summary_memory": { - "class": "memory.operation.summary_memory", - "workflow": "dummy_worker", - "description": "summary observation memories of the user", - "interval_time": 300 - } - } - } - }, - "models": { - "dashscope_generation": { - "class": "models.llama_index_generation_model", - "module_name": "dashscope_generation", - "model_name": "qwen-max" - }, - "dashscope_embedding": { - "class": "models.llama_index_embedding_model", - "module_name": "dashscope_embedding", - "model_name": "text-embedding-v2" - }, - "dashscope_rank": { - "class": "models.llama_index_rank_model", - "module_name": "dashscope_rank", - "model_name": "gte-rerank" - } - }, - "vector_store": { - "class": "storage.dummy_vector_store", - "embedding_model": "dashscope_embedding" - }, - "monitor": { - "class": "storage.dummy_monitor" - }, - "worker": { - "dummy_worker": { - "class": "memory.worker.dummy_worker", - "generation_model": "dashscope_generation", - "embedding_model": "dashscope_embedding", - "rank_model": "dashscope_rank" - } - } -} \ No newline at end of file diff --git a/config/config.yaml b/config/config.yaml index e6d78cb1..053d0132 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -8,6 +8,7 @@ memory_chat: class: chat.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation + memory_service: memory_chat_service: class: memory.service.chat_memory_service diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 1f5703b1..254fb265 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -1,6 +1,7 @@ import inspect import time from abc import abstractmethod, ABCMeta +from typing import Any from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen @@ -28,20 +29,28 @@ class BaseModel(metaclass=ABCMeta): self.timeout: int = timeout self.max_retries: int = max_retries self.retry_interval: float = retry_interval + self.kwargs_filter: bool = kwargs_filter self.kwargs: dict = kwargs self.data = {} + self._model: Any = None + self.logger = Logger.get_logger() - obj_cls = MODEL_REGISTRY[self.module_name] - if not obj_cls: - raise RuntimeError(f"method_type={self.module_name} is not supported!") + @property + def model(self): + if self._model is None: + if self.module_name not in MODEL_REGISTRY.module_dict: + raise RuntimeError(f"method_type={self.module_name} is not supported!") + obj_cls = MODEL_REGISTRY[self.module_name] - if kwargs_filter: - allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys()) - kwargs = {key: value for key, value in kwargs.items() if key in allowed_kwargs} - - self.model = obj_cls(**kwargs) + if self.kwargs_filter: + allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys()) + kwargs = {key: value for key, value in self.kwargs.items() if key in allowed_kwargs} + else: + kwargs = self.kwargs + self._model = obj_cls(**kwargs) + return self._model @abstractmethod def before_call(self, **kwargs) -> None: diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index f15d3064..0c042236 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -24,3 +24,9 @@ class MemoryNode(BaseModel): vector: List[float] = Field([], description="content embedding result, return empty") + @property + def node_keys(self): + return list(self.model_json_schema()["properties"].keys()) + + def __getitem__(self, key: str): + return self.model_dump().get(key) diff --git a/memory_scope/scheme/model_response.py b/memory_scope/scheme/model_response.py index 33047529..8e961332 100644 --- a/memory_scope/scheme/model_response.py +++ b/memory_scope/scheme/model_response.py @@ -28,13 +28,7 @@ class ModelResponse(BaseModel): def __str__(self, max_size=100, **kwargs): result = {} - # noinspection PyBroadException - try: - all_dict = self.model_dump() - except Exception: - all_dict = self.dict() - - for key, value in all_dict.items(): + for key, value in self.model_dump().items(): if key == "raw" or not value: continue diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index d5c64a4f..604866fe 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -10,11 +10,9 @@ class BaseVectorStore(metaclass=ABCMeta): def __init__(self, index_name: str = "", embedding_model: BaseModel | None = None, - content_key: str = "text", **kwargs): self.index_name: str = index_name self.embedding_model: BaseModel = embedding_model - self.content_key: str = content_key self.kwargs: dict = kwargs @abstractmethod @@ -32,23 +30,17 @@ class BaseVectorStore(metaclass=ABCMeta): """ pass - @abstractmethod - def insert_batch(self): - """ - :return: - """ + def insert_batch(self, nodes: List[MemoryNode]): pass - @abstractmethod - def delete(self): - """ - :return: - """ + def delete(self, node: MemoryNode): + pass + + def update(self, node: MemoryNode): pass - @abstractmethod def flush(self): - """ - :return: - """ + pass + + def close(self): pass diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index f1e6ecd9..85f26e7c 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -1,13 +1,12 @@ from typing import Dict, List, Any -from llama_index.core.schema import TextNode -from llama_index.core.vector_stores import VectorStoreQuery -from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext +from llama_index.core import VectorStoreIndex +from llama_index.core.schema import TextNode, NodeWithScore from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy -from llama_index.core.vector_stores.types import MetadataFilters, ExactMatchFilter, VectorStoreQueryMode + from memory_scope.models.base_model import BaseModel -from memory_scope.storage.base_vector_store import BaseVectorStore from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.storage.base_vector_store import BaseVectorStore class _ElasticsearchStore(ElasticsearchStore): @@ -23,9 +22,7 @@ class _ElasticsearchStore(ElasticsearchStore): Raises: Exception: If AsyncElasticsearch delete_by_query fails. """ - return await self._store.delete( - query={"term": {"_id": ref_doc_id}}, **delete_kwargs - ) + return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs) def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]: @@ -40,7 +37,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str """ result = { - "bool" : {} + "bool": {} } for key, value in standard_filters.items(): if isinstance(value, list): @@ -48,10 +45,10 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str for v in value: operands.append( { - "term": - { - f"metadata.{key}.keyword": {"value": v} - } + "term": + { + f"metadata.{key}.keyword": {"value": v} + } } ) result['bool'].update({"should": operands}) @@ -72,82 +69,67 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str class LlamaIndexElasticSearchStore(BaseVectorStore): - def __init__(self, - index_name: str, - embedding_model: BaseModel, - content_key: str = "text", + def __init__(self, + index_name: str, + embedding_model: BaseModel, **kwargs): - - self.index_name: str = index_name - self.embedding_model: BaseModel = embedding_model - + super().__init__(index_name=index_name, embedding_model=embedding_model, **kwargs) self.es_store = _ElasticsearchStore(index_name=self.index_name, - retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True), - **kwargs) - - self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_model, llm=None) + retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True), + **kwargs) self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, - service_context=self.service_context) - - - def retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode: - - filter = _to_elasticsearch_filter(filter_dict) - retriever = self.index.as_retriever( - vector_store_kwargs={ - "es_filter": filter - }, - similarity_top_k=top_k - ) - textnodes = retriever.retrieve(query) - results = self._textnodes2memorynodes(textnodes) + embed_model=self.embedding_model.model) - return results - - async def async_retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode: - filter = _to_elasticsearch_filter(filter_dict) - retriever = self.index.as_retriever( - vector_store_kwargs={ - "es_filter": filter - }, - similarity_top_k=top_k - ) - textnodes = await retriever.aretrieve(query) - results = self._textnodes2memorynodes(textnodes) + self.memory_node_keys = [x for x in MemoryNode().node_keys if x not in ["meta_data", "content"]] + + def retrieve(self, + query: str, + top_k: int, + filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: + if filter_dict is None: + filter_dict = {} + + es_filter = _to_elasticsearch_filter(filter_dict) + retriever = self.index.as_retriever( + vector_store_kwargs={"es_filter": es_filter}, + similarity_top_k=top_k) + text_nodes = retriever.retrieve(query) + return [self._text_node_2_memory_node(n) for n in text_nodes] + + async def async_retrieve(self, + query: str, + top_k: int, + filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: + if filter_dict is None: + filter_dict = {} + + es_filter = _to_elasticsearch_filter(filter_dict) + retriever = self.index.as_retriever( + vector_store_kwargs={"es_filter": es_filter}, + similarity_top_k=top_k) + text_nodes: List[NodeWithScore] = await retriever.aretrieve(query) + return [self._text_node_2_memory_node(n) for n in text_nodes] - return results - def insert(self, node: MemoryNode): - node = self._memorynode2textnode(node) - self.index.insert_nodes([node]) + self.index.insert_nodes([self._memory_node_2_text_node(node)]) - def insert_batch(self, node: MemoryNode) -> None: - raise NotImplementedError - - def delete(self, node: MemoryNode) -> None: + def delete(self, node: MemoryNode): memory_id = node.memory_id - self.es_store.delete(memory_id) - - def update(self, node: MemoryNode) -> None: + return self.es_store.delete(memory_id) + + def update(self, node: MemoryNode): self.delete(node) self.insert(node) - def flush(self): - raise NotImplementedError - - def _memorynode2textnode(self, memory_node: MemoryNode) -> TextNode: - content = memory_node.content - memory_id = memory_node.memory_id - meta = memory_node.model_dump(exclude={"content"}) - return TextNode(id_=memory_id, text=content, metadata=meta) + def close(self): + self.es_store.close() - def _textnode2memorynode(self, text_node: TextNode) -> MemoryNode: - content = text_node.text - meta = text_node.metadata - mem_node = MemoryNode(content=content, **meta) - return mem_node + @staticmethod + def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode: + return TextNode(id_=memory_node.memory_id, + text=memory_node.content, + metadata=memory_node.model_dump(exclude={"content"})) - def _textnodes2memorynodes(self, text_nodes: TextNode) -> MemoryNode: - mem_nodes = [self._textnode2memorynode(node) for node in text_nodes] - return mem_nodes - + @staticmethod + def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode: + return MemoryNode(content=text_node.text, **text_node.metadata) diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index e73cd1ca..5608ce7a 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -1,32 +1,31 @@ import unittest -from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters, FilterCondition, FilterOperator -from llama_index.core.schema import TextNode +from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel from memory_scope.scheme.memory_node import MemoryNode from memory_scope.storage.llama_index_elastic_search_store import LlamaIndexElasticSearchStore -from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel + class TestLlamaIndexElasticSearchStore(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): config = { - "method_type": "DashScopeEmbedding", + "module_name": "dashscope_embedding", "model_name": "text-embedding-v2", "clazz": "models.llama_index_embedding_model" } - emb = LlamaIndexEmbeddingModel(**config).model + emb = LlamaIndexEmbeddingModel(**config) config = { - "index_name" : "0626_1", - "es_url" : "http://localhost:9200", - "embedding_model" : emb, - + "index_name": "0626_1", + "es_url": "http://localhost:9200", + "embedding_model": emb, } self.es_store = LlamaIndexElasticSearchStore(**config) self.data = [ MemoryNode( - content="The lives of two mob hitmen, a boxer, a gangster and his wife, and a pair of diner bandits intertwine in four tales of violence and redemption.", + content="The lives of two mob hitmen, a boxer, a gangster and his wife, " + "and a pair of diner bandits intertwine in four tales of violence and redemption.", memory_type="observation", user_id="0", status="valid", @@ -34,7 +33,9 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): ), MemoryNode( - content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, Batman must accept one of the greatest psychological and physical tests of his ability to fight injustice.", + content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, " + "Batman must accept one of the greatest psychological and physical tests of his " + "ability to fight injustice.", memory_type="observation", user_id="1", status="valid", @@ -42,37 +43,41 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): ), MemoryNode( - content="An insomniac office worker and a devil-may-care soapmaker form an underground fight club that evolves into something much, much more.", + content="An insomniac office worker and a devil-may-care soapmaker form an underground fight " + "club that evolves into something much, much more.", memory_type="insights", user_id="2", status="valid", memory_id="ccc789", - ), MemoryNode( - content="A thief who steals corporate secrets through the use of dream-sharing technology is given the inverse task of planting an idea into thed of a C.E.O.", + content="A thief who steals corporate secrets through the use of dream-sharing technology " + "is given the inverse task of planting an idea into thed of a C.E.O.", memory_type="insights", user_id="3", status="valid", memory_id="ddd012", ), MemoryNode( - content="A computer hacker learns from mysterious rebels about the true nature of his reality and his role in the war against its controllers.", + content="A computer hacker learns from mysterious rebels about the true nature of his reality " + "and his role in the war against its controllers.", memory_type="profile", user_id="4", status="valid", memory_id="eee345", ), MemoryNode( - content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven deadly sins as his motives.", + content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven " + "deadly sins as his motives.", memory_type="profile", user_id="5", status="valid", memory_id="fff678" ), MemoryNode( - content="An organized crime dynasty's aging patriarch transfers control of his clandestine empire to his reluctant son.", + content="An organized crime dynasty's aging patriarch transfers control of his clandestine " + "empire to his reluctant son.", memory_type="insights", user_id="6", status="valid", @@ -85,46 +90,47 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): status="valid", memory_id="ggg234", ), - ] - - def test_retrieve(self, ): - filter = { + + def test_retrieve(self): + filter_dict = { "user_id": "6", } for node in self.data: - self.es_store.insert(node) + self.es_store.insert(node) self.es_store.insert(MemoryNode( content="xxxxxx", - memory_type="profile", - user_id="6", - status="valid", - memory_id="ggg567" + memory_type="profile", + user_id="6", + status="valid", + memory_id="ggg567" )) - res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res) self.es_store.update(MemoryNode( content="test update", - memory_type="profile", - user_id="6", - status="invalid", - memory_id="ggg567" + memory_type="profile", + user_id="6", + status="invalid", + memory_id="ggg567" )) res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) print(len(res)) print(res) - self.es_store.delete(MemoryNode( content="test update", - memory_type="profile", - user_id="6", - status="invalid", - memory_id="ggg567" + memory_type="profile", + user_id="6", + status="invalid", + memory_id="ggg567" )) res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) print(len(res)) print(res) + + def tearDown(self): + self.es_store.close() From b6a3d23ffe1b5c8eb1689bc0ae838a9936d9f455 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 15:30:00 +0800 Subject: [PATCH 39/41] [dev] change dummy vector store to llama es store --- config/config.yaml | 7 ++++-- memory_scope/cli.py | 8 ++++++- memory_scope/storage/base_monitor.py | 8 +++---- memory_scope/storage/base_vector_store.py | 12 ---------- memory_scope/storage/dummy_monitor.py | 2 +- memory_scope/storage/dummy_vector_store.py | 24 ------------------- .../llama_index_elastic_search_store.py | 14 ++++++----- tests/storages/test_storages_lli_es.py | 23 +++++++++++++----- 8 files changed, 42 insertions(+), 56 deletions(-) delete mode 100644 memory_scope/storage/dummy_vector_store.py diff --git a/config/config.yaml b/config/config.yaml index 053d0132..33ed7954 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -8,7 +8,8 @@ memory_chat: class: chat.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation - + human_name: human + assistant_name: assistant memory_service: memory_chat_service: class: memory.service.chat_memory_service @@ -52,8 +53,10 @@ models: module_name: dashscope_rank model_name: gte-rerank vector_store: - class: storage.dummy_vector_store + class: storage.llama_index_elastic_search_store embedding_model: dashscope_embedding + index_name: memory_index + es_url: http://localhost:9200 monitor: class: storage.dummy_monitor worker: diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 427ad145..e76575aa 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -14,6 +14,7 @@ from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.utils.logger import Logger from memory_scope.utils.tool_functions import init_instance_by_config from memory_scope.utils.timer import timer +from memory_scope.enumeration.model_enum import ModelEnum class CliJob(object): @@ -54,7 +55,9 @@ class CliJob(object): G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) # init vector_store - G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"]) + vector_store_config = self.config["vector_store"] + embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]] + G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model) # init monitor G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) @@ -70,6 +73,9 @@ class CliJob(object): memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() + G_CONTEXT.vector_store.close() + G_CONTEXT.monitor.close() + if __name__ == "__main__": cli_job = CliJob() diff --git a/memory_scope/storage/base_monitor.py b/memory_scope/storage/base_monitor.py index 1d84621e..05465cd3 100644 --- a/memory_scope/storage/base_monitor.py +++ b/memory_scope/storage/base_monitor.py @@ -18,8 +18,8 @@ class BaseMonitor(metaclass=ABCMeta): :return: """ - @abstractmethod def flush(self): - """ - :return: - """ + pass + + def close(self): + pass diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index 604866fe..9c5589b8 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -1,20 +1,11 @@ from abc import ABCMeta, abstractmethod from typing import Dict, List -from memory_scope.models.base_model import BaseModel from memory_scope.scheme.memory_node import MemoryNode class BaseVectorStore(metaclass=ABCMeta): - def __init__(self, - index_name: str = "", - embedding_model: BaseModel | None = None, - **kwargs): - self.index_name: str = index_name - self.embedding_model: BaseModel = embedding_model - self.kwargs: dict = kwargs - @abstractmethod def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): pass @@ -25,9 +16,6 @@ class BaseVectorStore(metaclass=ABCMeta): @abstractmethod def insert(self, node: MemoryNode): - """ TODO 是否overwrite - :return: - """ pass def insert_batch(self, nodes: List[MemoryNode]): diff --git a/memory_scope/storage/dummy_monitor.py b/memory_scope/storage/dummy_monitor.py index f39a917b..6c754ae5 100644 --- a/memory_scope/storage/dummy_monitor.py +++ b/memory_scope/storage/dummy_monitor.py @@ -8,5 +8,5 @@ class DummyMonitor(BaseMonitor): def add_token(self): pass - def flush(self): + def close(self): pass diff --git a/memory_scope/storage/dummy_vector_store.py b/memory_scope/storage/dummy_vector_store.py deleted file mode 100644 index f4c7fad8..00000000 --- a/memory_scope/storage/dummy_vector_store.py +++ /dev/null @@ -1,24 +0,0 @@ -from typing import Dict, List - -from memory_scope.scheme.memory_node import MemoryNode -from memory_scope.storage.base_vector_store import BaseVectorStore - - -class DummyVectorStore(BaseVectorStore): - def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): - pass - - async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): - pass - - def insert(self, node: MemoryNode): - pass - - def insert_batch(self): - pass - - def delete(self): - pass - - def flush(self): - pass diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index 85f26e7c..a6f5acb8 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -70,18 +70,20 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str class LlamaIndexElasticSearchStore(BaseVectorStore): def __init__(self, - index_name: str, embedding_model: BaseModel, + index_name: str, + es_url: str, + use_hybrid: bool = True, **kwargs): - super().__init__(index_name=index_name, embedding_model=embedding_model, **kwargs) - self.es_store = _ElasticsearchStore(index_name=self.index_name, - retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True), + + self.embedding_model: BaseModel = embedding_model + self.es_store = _ElasticsearchStore(index_name=index_name, + es_url=es_url, + retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid), **kwargs) self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, embed_model=self.embedding_model.model) - self.memory_node_keys = [x for x in MemoryNode().node_keys if x not in ["meta_data", "content"]] - def retrieve(self, query: str, top_k: int, diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 5608ce7a..7f88bea2 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -40,7 +40,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="1", status="valid", memory_id="bbb456", - + meta_data={"1": "1"} ), MemoryNode( content="An insomniac office worker and a devil-may-care soapmaker form an underground fight " @@ -49,7 +49,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="2", status="valid", memory_id="ccc789", - + meta_data={"2": "2"} ), MemoryNode( content="A thief who steals corporate secrets through the use of dream-sharing technology " @@ -58,6 +58,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="3", status="valid", memory_id="ddd012", + meta_data={"3": "3"} + ), MemoryNode( content="A computer hacker learns from mysterious rebels about the true nature of his reality " @@ -66,6 +68,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="4", status="valid", memory_id="eee345", + meta_data={"4": "4"} + ), MemoryNode( content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven " @@ -73,7 +77,9 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): memory_type="profile", user_id="5", status="valid", - memory_id="fff678" + memory_id="fff678", + meta_data={"5": "5"}, + ), MemoryNode( content="An organized crime dynasty's aging patriarch transfers control of his clandestine " @@ -82,6 +88,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="6", status="valid", memory_id="ggg901", + meta_data={"5": "5"} + ), MemoryNode( content="ggggggggg", @@ -89,6 +97,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="6", status="valid", memory_id="ggg234", + meta_data={"5": "5"} + ), ] @@ -104,7 +114,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): memory_type="profile", user_id="6", status="valid", - memory_id="ggg567" + memory_id="ggg567", + meta_data={"5": "5"} )) res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) @@ -117,7 +128,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): status="invalid", memory_id="ggg567" )) - res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res) @@ -128,7 +139,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): status="invalid", memory_id="ggg567" )) - res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res) From 5f1bcf50d8d1a1948709ff8661a771a03f645d96 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 16:08:42 +0800 Subject: [PATCH 40/41] [dev] add memory base worker & add role name to messages --- memory_scope/chat/cli_memory_chat.py | 30 +-- .../memory/worker/memory_base_worker.py | 77 ++++++++ .../models/llama_index_generation_model.py | 4 +- memory_scope/scheme/message.py | 5 + memory_scope/utils/pipeline.py | 185 ------------------ memory_scope/utils/response_text_parser.py | 2 +- memory_scope/utils/timer.py | 2 +- old/worker/memory_base_worker.py | 93 --------- 8 files changed, 102 insertions(+), 296 deletions(-) create mode 100644 memory_scope/memory/worker/memory_base_worker.py delete mode 100644 memory_scope/utils/pipeline.py delete mode 100644 old/worker/memory_base_worker.py diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 2eb2e21a..21c07b19 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -1,4 +1,3 @@ -import datetime import os import time from typing import List @@ -31,6 +30,7 @@ class CliMemoryChat(BaseMemoryChat): human_name: str = "human", assistant_name: str = "assistant", **kwargs): + self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model self.stream: bool = stream @@ -58,31 +58,35 @@ class CliMemoryChat(BaseMemoryChat): self._generation_model = G_CONTEXT.model_dict[self._generation_model] return self._generation_model - @staticmethod - def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] - if related_memories: + def get_system_prompt(self) -> Message: + system_prompt = SYSTEM_PROMPT[G_CONTEXT.language].strip() + + memories: str = self.memory_service.read_memory() + if memories: memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] - all_prompt_list = [system_prompt, memory_prompt] - all_prompt_list.extend(related_memories) - system_prompt = "\n".join([x.strip() for x in all_prompt_list]) - return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) + system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) + + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: query = query.strip() if not query: return - time_created = int(datetime.datetime.now().timestamp()) - new_message: Message = Message(role=MessageRoleEnum.USER.value, content=query, time_created=time_created) + new_message: Message = Message(role=MessageRoleEnum.USER.value, + role_name=self.human_name, + content=query) + self.memory_service.add_messages(new_message) - related_memories: List[str] = self.memory_service.read_memory() - system_message: Message = self.get_system_prompt(related_memories, time_created) + system_message: Message = self.get_system_prompt() + model_response = self.generation_model.call(messages=[system_message, new_message], stream=self.stream) if self.stream: for _ in model_response: + _.message.role_name = self.assistant_name yield _ else: + model_response.message.role_name = self.assistant_name return model_response def process_commands(self, query: str) -> bool: diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py new file mode 100644 index 00000000..f61619f9 --- /dev/null +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -0,0 +1,77 @@ +from abc import ABCMeta +from typing import List + +from memory_scope.chat.global_context import G_CONTEXT +from memory_scope.memory.worker.base_worker import BaseWorker +from memory_scope.models.base_model import BaseModel +from memory_scope.storage.base_monitor import BaseMonitor +from memory_scope.storage.base_vector_store import BaseVectorStore + + +class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): + + def __init__(self, + embedding_model: str = "", + generation_model: str = "", + rank_model: str = "", + **kwargs): + super(MemoryBaseWorker, self).__init__(**kwargs) + + self._embedding_model: BaseModel | str = embedding_model + self._generation_model: BaseModel | str = generation_model + self._rank_model: BaseModel | str = rank_model + + self._vector_store: BaseVectorStore | None = None + self._monitor: BaseMonitor | None = None + + @property + def messages(self) -> List[Message]: + return self.get_context(MESSAGES) + + @messages.setter + def messages(self, value): + self.set_context(MESSAGES, value) + + @property + def chat_name(self): + return self.get_context(CHAT_NAME) + + @property + def embedding_model(self) -> BaseModel: + if isinstance(self._embedding_model, str): + self._embedding_model = G_CONTEXT.model_dict[self._embedding_model] + return self._embedding_model + + @property + def generation_model(self) -> BaseModel: + if isinstance(self._generation_model, str): + self._generation_model = G_CONTEXT.model_dict[self._generation_model] + return self._generation_model + + @property + def rank_model(self) -> BaseModel: + if isinstance(self._rank_model, str): + self._rank_model = G_CONTEXT.model_dict[self._rank_model] + return self._rank_model + + @property + def vector_store(self) -> BaseVectorStore: + if self._vector_store is None: + self._vector_store = G_CONTEXT.vector_store + return self._vector_store + + @property + def monitor(self): + if self._monitor is None: + self._monitor = G_CONTEXT.monitor + return self._monitor + + @property + def memory_id(self) -> str: + pass + + def __getattr__(self, key: str): + return self.kwargs[key] + + def get_prompt(self, x): + return x[GLOBAL_CONTEXT.global_configs["language"]] \ No newline at end of file diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index a3c6f45d..0d9780e7 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -32,9 +32,7 @@ class LlamaIndexGenerationModel(BaseModel): stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: - model_response.message = Message(role=MessageRoleEnum.ASSISTANT, - content="", - time_created=int(datetime.datetime.now().timestamp())) + model_response.message = Message(role=MessageRoleEnum.ASSISTANT, content="") call_result = model_response.raw if stream: diff --git a/memory_scope/scheme/message.py b/memory_scope/scheme/message.py index f08b3290..43308850 100644 --- a/memory_scope/scheme/message.py +++ b/memory_scope/scheme/message.py @@ -1,4 +1,5 @@ import datetime +from typing import Dict from pydantic import Field, BaseModel @@ -6,9 +7,13 @@ from pydantic import Field, BaseModel class Message(BaseModel): role: str = Field(..., description="The role of the message sender (user, assistant, system)") + role_name: str = Field("", description="role name") + content: str = Field(..., description="The body of the message") time_created: int = Field(int(datetime.datetime.now().timestamp()), description="Timestamp when the message was created") memorized: bool = Field(False, description="indicate whether message is memorized") + + meta_data: Dict[str, str] = Field({}, description="meta data for msg") diff --git a/memory_scope/utils/pipeline.py b/memory_scope/utils/pipeline.py deleted file mode 100644 index 4eb84b0b..00000000 --- a/memory_scope/utils/pipeline.py +++ /dev/null @@ -1,185 +0,0 @@ -import re -import threading -import time -from concurrent.futures import as_completed -from itertools import zip_longest -from typing import Dict, Any, List - -from chat.global_context import GLOBAL_CONTEXT -from constants.common_constants import MESSAGES, CHAT_NAME -from enumeration.memory_method_enum import MemoryMethodEnum -from scheme.message import Message -from utils.logger import Logger -from utils.timer import Timer -from worker.base_worker import BaseWorker - - -class Pipeline(object): - def __init__(self, - chat_name: str, - memory_method_type: MemoryMethodEnum, - pipeline_str: str, - history_msg_count: int = 3, - loop_interval_time: int = 300, - loop_minimum_count: int = 20): - - self.chat_name: str = chat_name - self.memory_method_type: MemoryMethodEnum = memory_method_type - self.pipeline_str: str = pipeline_str - self.history_msg_count: int = history_msg_count - self.loop_interval_time: int = loop_interval_time - self.loop_minimum_count: int = loop_minimum_count - - # pipeline上下文和锁 - self.context: Dict[str, Any] = {} - self.context_lock = threading.Lock() - - # pipeline run config - self.loop_switch: bool = False - self.pipeline_list: list[list] = [] - self.worker_set: set[str] = set() - self.worker_dict: Dict[str, BaseWorker] = {} - self.injected: bool = False - - # message list - self.history_message_list: List[Message] = [] - self.current_message_list: List[Message] = [] - self.message_lock = threading.Lock() - - # 日志 - self.logger: Logger = Logger.get_logger() - - self._parse_pipeline() - - def _parse_pipeline(self): - if not self.pipeline_str: - return - - # re-match e.g., [a|b],c,[d,e,f|g,h],j - pattern = r'(\[[^\]]*\]|[^,]+)' - pipeline_split = re.findall(pattern, self.pipeline_str) - - self.pipeline_list = [] - for pipeline_part in pipeline_split: - # e.g., [d,e,f|g,h] - pipeline_part = pipeline_part.strip() - if '[' in pipeline_part or ']' in pipeline_part: - pipeline_part = pipeline_part.replace('[', '').replace(']', '') - - # e.g., ["d,e,f", "g,h"] - line_split = [x.strip() for x in pipeline_part.split("|") if x] - if len(line_split) <= 0: - continue - - # e.g., ["d","e","f"] - line_split_split = [] - for sub_line_split in line_split: - sub_split = [x.strip() for x in sub_line_split.split(",")] - line_split_split.append(sub_split) - # add to workers - self.worker_set.update(sub_split) - self.pipeline_list.append(line_split_split) - - def _visit_and_inject_workers(self): - if self.injected: - return - - self.worker_dict = GLOBAL_CONTEXT.worker_dict[self.chat_name] - - self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline Begin -----") - i: int = 0 - for pipeline_part in self.pipeline_list: - if len(pipeline_part) == 1: - for w in pipeline_part[0]: - self.logger.info(f"stage{i}: {w}") - i += 1 - if w not in self.worker_dict: - raise RuntimeError(f"worker={w} is not inited.") - # 注入context - self.worker_dict[w].set_context_dict(self.context) - else: - for w_zip in zip_longest(*pipeline_part, fillvalue="-"): - self.logger.info(f"stage{i}: {' | '.join(w_zip)}") - i += 1 - for w in w_zip: - if w == "-": - continue - if w not in self.worker_dict: - raise RuntimeError(f"worker={w} is not inited.") - - # 注入context & lock - self.worker_dict[w].set_context_dict(self.context, self.context_lock) - - self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline End -----") - self.injected = True - - def _worker_run(self, worker_list: list[str]) -> bool: - for worker_name in worker_list: - worker = self.worker_dict[worker_name] - worker.run() - if not worker.continue_run: - return False - return True - - def _run(self): - self._visit_and_inject_workers() - - with Timer(f"pipeline_{self.chat_name}_{self.memory_method_type.value}"): - self.context[MESSAGES] = self.history_message_list + self.current_message_list - self.context[CHAT_NAME] = self.chat_name - - for pipeline_part in self.pipeline_list: - if len(pipeline_part) == 1: - if not self._worker_run(pipeline_part[0]): - break - else: - t_list = [] - for worker_list in pipeline_part: - t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self._worker_run, worker_list)) - - flag = True - for future in as_completed(t_list): - if not future.result(): - flag = False - break - if not flag: - break - - def _thread_loop(self): - while self.loop_switch: - time.sleep(self.loop_interval_time) - if len(self.current_message_list) < self.loop_minimum_count: - continue - self._run() - self.context.clear() - self.history_message_list = self.history_message_list.extend(self.current_message_list)[ - -self.history_msg_count:] - with self.message_lock: - self.current_message_list.clear() - - def start_loop_run(self): - if not self.loop_switch: - self.loop_switch = True - return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop) - - def run(self, result_key: str = None): - self._run() - - # 获取result - result = None - if result_key: - result = self.context.get(result_key) - self.context.clear() - - # 清理 msg - self.history_message_list = self.history_message_list.extend(self.current_message_list)[ - -self.history_msg_count:] - self.current_message_list.clear() - return result - - def submit_message(self, message: Message, with_lock=True): - if with_lock: - with self.message_lock: - self.current_message_list.append(message) - else: - self.current_message_list.append(message) diff --git a/memory_scope/utils/response_text_parser.py b/memory_scope/utils/response_text_parser.py index 6fcc6f5a..d665b458 100644 --- a/memory_scope/utils/response_text_parser.py +++ b/memory_scope/utils/response_text_parser.py @@ -1,6 +1,6 @@ import re -from utils.logger import Logger +from memory_scope.utils.logger import Logger class ResponseTextParser(object): diff --git a/memory_scope/utils/timer.py b/memory_scope/utils/timer.py index a6667e3d..be7df83f 100644 --- a/memory_scope/utils/timer.py +++ b/memory_scope/utils/timer.py @@ -1,6 +1,6 @@ import time -from .logger import Logger +from memory_scope.utils.logger import Logger class Timer(object): diff --git a/old/worker/memory_base_worker.py b/old/worker/memory_base_worker.py deleted file mode 100644 index 7e4fe846..00000000 --- a/old/worker/memory_base_worker.py +++ /dev/null @@ -1,93 +0,0 @@ -from typing import List, Dict - -from ..chat.global_context import GLOBAL_CONTEXT -from ..constants.common_constants import MESSAGES, CHAT_NAME -from ..models.base_model import BaseModel -from ..scheme.message import Message -from ..storage.base_monitor import BaseMonitor -from ..storage.base_vector_store import BaseVectorStore -from ..worker.base_worker import BaseWorker -from ..scheme.memory_node import MemoryNode -from ..constants import common_constants - - -class MemoryBaseWorker(BaseWorker): - def __init__( - self, embedding_model: str, generation_model: str, rank_model: str, **kwargs - ): - super(MemoryBaseWorker, self).__init__(**kwargs) - self.embedding_model_name: str = embedding_model - self.generation_model_name: str = generation_model - self.rank_model_name: str = rank_model - - self._embedding_model: BaseModel | None = None - self._generation_model: BaseModel | None = None - self._rank_model: BaseModel | None = None - - self._vector_store: BaseVectorStore | None = None - self._monitor: BaseMonitor | None = None - - @property - def messages(self) -> List[Message]: - return self.get_context(MESSAGES) - - @messages.setter - def messages(self, value): - self.set_context(MESSAGES, value) - - @property - def chat_name(self): - return self.get_context(CHAT_NAME) - - @property - def embedding_model(self) -> BaseModel: - if self._embedding_model is None: - self._embedding_model = GLOBAL_CONTEXT.model_dict.get( - self.embedding_model_name - ) - return self._embedding_model - - @property - def generation_model(self) -> BaseModel: - if self._generation_model is None: - self._generation_model = GLOBAL_CONTEXT.model_dict.get( - self.generation_model_name - ) - return self._generation_model - - @property - def rank_model(self) -> BaseModel: - if self._rank_model is None: - self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name) - return self._rank_model - - @property - def vector_store(self) -> BaseVectorStore: - if self._vector_store is None: - self._vector_store = GLOBAL_CONTEXT.vector_store - return self._vector_store - - @property - def monitor(self): - if self._monitor is None: - self._monitor = GLOBAL_CONTEXT.monitor - return self._monitor - - @property - def user_profile_dict(self) -> Dict[str, MemoryNode]: - if not self._user_profile_dict: - self._user_profile_dict = { - user_attr.meta_data.get("memory_key", ""): user_attr - for user_attr in self.get_context(common_constants.USER_PROFILE) - } - return self._user_profile_dict - - @property - def memory_id(self) -> str: - pass - - def __getattr__(self, key): - return self.kwargs[key] - - def get_prompt(self, x): - return x[GLOBAL_CONTEXT.global_configs["language"]] \ No newline at end of file From c5386ceaafd36f49c91659d6824a5b9451c8d809 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 16:32:25 +0800 Subject: [PATCH 41/41] [bugfix] deep copy in messages --- memory_scope/memory/operation/read_memory.py | 2 +- memory_scope/memory/operation/write_memory.py | 2 +- .../memory/worker/memory_base_worker.py | 25 +++++++++++-------- memory_scope/scheme/memory_node.py | 13 ++++++++-- memory_scope/utils/tool_functions.py | 7 ++++++ 5 files changed, 35 insertions(+), 14 deletions(-) diff --git a/memory_scope/memory/operation/read_memory.py b/memory_scope/memory/operation/read_memory.py index d647acd3..cc66b6ea 100644 --- a/memory_scope/memory/operation/read_memory.py +++ b/memory_scope/memory/operation/read_memory.py @@ -27,7 +27,7 @@ class ReadMemory(BaseWorkflow, BaseOperation): def run_operation(self): max_count = 1 + max(self.his_msg_count, self.contextual_msg_count) - self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] + self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]] self.run_workflow() result = self.context.get(RESULT) self.context.clear() diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 362ec10f..0df92f0c 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -56,7 +56,7 @@ class WriteMemory(BaseWorkflow, BaseOperation): return max_count = not_memorized_size + self.his_msg_count - self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] + self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]] self.run_workflow() result = self.context.get(RESULT) self.context.clear() diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index f61619f9..4c0d7878 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -2,8 +2,11 @@ from abc import ABCMeta from typing import List from memory_scope.chat.global_context import G_CONTEXT +from memory_scope.constants.common_constants import CHAT_MESSAGES +from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.models.base_model import BaseModel +from memory_scope.scheme.message import Message from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_vector_store import BaseVectorStore @@ -24,17 +27,15 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._vector_store: BaseVectorStore | None = None self._monitor: BaseMonitor | None = None + self._user_id: str | None = None + @property def messages(self) -> List[Message]: - return self.get_context(MESSAGES) + return self.get_context(CHAT_MESSAGES) @messages.setter def messages(self, value): - self.set_context(MESSAGES, value) - - @property - def chat_name(self): - return self.get_context(CHAT_NAME) + self.set_context(CHAT_MESSAGES, value) @property def embedding_model(self) -> BaseModel: @@ -67,11 +68,15 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._monitor @property - def memory_id(self) -> str: - pass + def user_id(self) -> str: + if self._user_id is None: + message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1] + self._user_id = message.role_name + return self._user_id def __getattr__(self, key: str): return self.kwargs[key] - def get_prompt(self, x): - return x[GLOBAL_CONTEXT.global_configs["language"]] \ No newline at end of file + @staticmethod + def get_prompt(prompt: dict) -> str: + return prompt[G_CONTEXT.global_configs["language"]] diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index 0c042236..e801d121 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -1,13 +1,17 @@ +import datetime from typing import Dict, List +import from pydantic import Field, BaseModel +from memory_scope.utils.tool_functions import md5_hash + class MemoryNode(BaseModel): - user_id: str = Field("", description="unique memory id for user") - memory_id: str = Field("", description="unique id for memory item") + user_id: str = Field("", description="unique memory id for user") + content: str = Field("", description="memory content") score_similar: float = Field(0, description="es similar score") @@ -24,9 +28,14 @@ class MemoryNode(BaseModel): vector: List[float] = Field([], description="content embedding result, return empty") + timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node") + @property def node_keys(self): return list(self.model_json_schema()["properties"].keys()) def __getitem__(self, key: str): return self.model_dump().get(key) + + def gen_memory_id(self): + self.memory_id = f"{self.user_id}_{self.timestamp}_{md5_hash(self.content)[:8]}" diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index aed4aff1..11e85e7b 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -1,3 +1,4 @@ +import hashlib import random import re import time @@ -113,3 +114,9 @@ def char_logo(words: str, seed: int = time.time_ns(), color=None): colored_line += colored_char colored_lines.append(colored_line) return colored_lines + + +def md5_hash(input_string: str): + m = hashlib.md5() + m.update(input_string.encode('utf-8')) + return m.hexdigest()