From 5f1bcf50d8d1a1948709ff8661a771a03f645d96 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 16:08:42 +0800 Subject: [PATCH] [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