From 6d30e4b18a0d42c1470cb2524d21d5ecc1e62e31 Mon Sep 17 00:00:00 2001 From: hs Date: Thu, 27 Jun 2024 12:22:00 +0800 Subject: [PATCH] [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")