From ee620a01bb787bfd99d44f995c94e0bedd4c1a21 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 25 Jun 2024 22:55:57 +0800 Subject: [PATCH] [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)