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")