From 3c423118a526c52059a733939931fa04c308fe00 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 26 Jul 2024 01:10:50 +0800 Subject: [PATCH 01/15] [dev] add new config format --- memoryscope/argument/__init__.py | 0 .../argument/cli_chat_demo.yaml | 83 ++--- memoryscope/argument/default_arguments.py | 182 ++++++++++ memoryscope/argument/init_handler.py | 28 ++ memoryscope/argument/memoryscope_arguments.py | 61 ++++ memoryscope/chat/api_memory_chat.py | 322 ++++++++++++++++++ memoryscope/chat/base_memory_chat.py | 14 +- memoryscope/chat/cli_memory_chat.py | 108 ++---- ...mory_chat.yaml => memory_chat_prompt.yaml} | 0 memoryscope/cli.py | 101 +----- .../memory/service/base_memory_service.py | 22 +- .../memory/service/memory_scope_service.py | 3 + .../memory/worker/memory_base_worker.py | 10 +- memoryscope/memoryscope.py | 198 +++++++++++ memoryscope/memoryscope_context.py | 27 ++ memoryscope/utils/global_context.py | 33 -- memoryscope/utils/memory_handler.py | 2 +- memoryscope/utils/prompt_handler.py | 44 ++- memoryscope/utils/tool_functions.py | 35 +- tests/other/read_prompt.yaml | 3 + tests/other/read_yaml.py | 11 + 21 files changed, 978 insertions(+), 309 deletions(-) create mode 100644 memoryscope/argument/__init__.py rename config/cli_chat_dash_cn.yaml => memoryscope/argument/cli_chat_demo.yaml (76%) create mode 100644 memoryscope/argument/default_arguments.py create mode 100644 memoryscope/argument/init_handler.py create mode 100644 memoryscope/argument/memoryscope_arguments.py create mode 100644 memoryscope/chat/api_memory_chat.py rename memoryscope/chat/{cli_memory_chat.yaml => memory_chat_prompt.yaml} (100%) create mode 100644 memoryscope/memoryscope.py create mode 100644 memoryscope/memoryscope_context.py delete mode 100644 memoryscope/utils/global_context.py create mode 100644 tests/other/read_prompt.yaml create mode 100644 tests/other/read_yaml.py diff --git a/memoryscope/argument/__init__.py b/memoryscope/argument/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/config/cli_chat_dash_cn.yaml b/memoryscope/argument/cli_chat_demo.yaml similarity index 76% rename from config/cli_chat_dash_cn.yaml rename to memoryscope/argument/cli_chat_demo.yaml index dd04463e..b8fdee74 100644 --- a/config/cli_chat_dash_cn.yaml +++ b/memoryscope/argument/cli_chat_demo.yaml @@ -1,19 +1,17 @@ global_config: - language: cn - max_workers: 5 - -logger_config: + language: en + thread_pool_max_workers: 5 logger_name: memoryscope - logger_suffix: time + logger_name_time_suffix: %Y%m%d_%H%M%S memory_chat: cli_memory_chat: class: chat.cli_memory_chat - memory_service: memory_scope_service - generation_model: dashscope_generation + memory_service: memoryscope_service + generation_model: generation_model memory_service: - memory_scope_service: + memoryscope_service: class: memory.service.memory_scope_service memory_operations: read_message: @@ -46,13 +44,13 @@ memory_service: workflow: add_memory description: "add a single observation" - summary_observation_memory: - class: memory.operation.summary_observation_op + consolidate_memory: + class: memory.operation.consolidate_memory_op workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory description: "summary user's observation memory" interval_time: 1 - summary_insight_memory: + reflect_and_reconsolidate: class: memory.operation.backend_operation workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory description: "summary user's insight memory" @@ -61,9 +59,9 @@ memory_service: worker: dummy: class: memory.worker.dummy_worker - generation_model: dashscope_generation - embedding_model: dashscope_embedding - rank_model: dashscope_rank + generation_model: generation_model + embedding_model: embedding_model + rank_model: rank_model read_message: class: memory.worker.frontend.read_message_worker set_query: @@ -74,12 +72,10 @@ worker: retrieve_ins_top_k: 100 extract_time: class: memory.worker.frontend.extract_time_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 + generation_model: generation_model semantic_rank: class: memory.worker.frontend.semantic_rank_worker - rank_model: dashscope_rank + rank_model: rank_model fuse_rerank: class: memory.worker.frontend.fuse_rerank_worker fuse_score_threshold: 0.01 @@ -99,9 +95,9 @@ worker: class: memory.worker.frontend.print_memory_worker retrieve_all_memory: class: memory.worker.frontend.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_top_k: 100 - retrieve_expired_top_k: 100 + retrieve_obs_top_k: 1000 + retrieve_ins_top_k: 1000 + retrieve_expired_top_k: 1000 delete_memory: class: memory.worker.backend.update_memory_worker method: delete_memory @@ -113,27 +109,19 @@ worker: method: from_query info_filter: class: memory.worker.backend.info_filter_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 + generation_model: generation_model load_today_memory: class: memory.worker.backend.load_memory_worker retrieve_today_top_k: 100 get_observation: class: memory.worker.backend.get_observation_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 + generation_model: generation_model get_observation_with_time: class: memory.worker.backend.get_observation_with_time_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 + generation_model: generation_model contra_repeat: class: memory.worker.backend.contra_repeat_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 + generation_model: generation_model store_memory: class: memory.worker.backend.update_memory_worker method: from_memory_key @@ -145,37 +133,31 @@ worker: retrieve_insight_top_k: 100 get_reflection_subject: class: memory.worker.backend.get_reflection_subject_worker - generation_model: dashscope_generation + generation_model: generation_model reflect_obs_cnt_threshold: 10 - generation_model_kwargs: - top_k: 1 update_insight: class: memory.worker.backend.update_insight_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - rank_model: dashscope_rank + generation_model: generation_model + rank_model: rank_model long_contra_repeat: class: memory.worker.backend.long_contra_repeat_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 + generation_model: generation_model -models: - dashscope_generation: +model: + generation_model: class: models.llama_index_generation_model module_name: dashscope_generation model_name: qwen-max max_tokens: 2000 - dashscope_embedding: + embedding_model: class: models.llama_index_embedding_model module_name: dashscope_embedding model_name: text-embedding-v2 - dashscope_rank: + rank_model: class: models.llama_index_rank_model module_name: dashscope_rank model_name: gte-rerank - top_n: 10 + top_n: 500 dummy_generation: class: models.dummy_generation_model module_name: dummy_generation @@ -183,10 +165,11 @@ models: memory_store: class: storage.llama_index_es_memory_store - embedding_model: dashscope_embedding + embedding_model: embedding_model index_name: memory_index es_url: http://localhost:9200 - use_hybrid: true + retrieve_type: dense + hybrid_alpha: 1.0 monitor: class: storage.dummy_monitor \ No newline at end of file diff --git a/memoryscope/argument/default_arguments.py b/memoryscope/argument/default_arguments.py new file mode 100644 index 00000000..275285fd --- /dev/null +++ b/memoryscope/argument/default_arguments.py @@ -0,0 +1,182 @@ +DEFAULT_GLOBAL_ARGUMENTS = { + "language": "en", + "thread_pool_max_workers": 5, + "logger_name": "memoryscope", + "logger_name_time_suffix": "%Y%m%d_%H%M%S" +} + +DEFAULT_MEMORY_CHAT_ARGUMENTS = { + "cli_memory_chat": { + "class": "chat.cli_memory_chat", + "memory_service": "memoryscope_service", + "generation_model": "generation_model" + } +} + +DEFAULT_MEMORY_SERVICE_ARGUMENTS = { + "memoryscope_service": { + "class": "memory.service.memory_scope_service", + "memory_operations": { + "read_message": { + "class": "memory.operation.frontend_operation", + "workflow": "read_message", + "description": "read short memory" + }, + "retrieve_memory": { + "class": "memory.operation.frontend_operation", + "workflow": "set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank", + "description": "retrieve long-term memory" + }, + "list_memory": { + "class": "memory.operation.frontend_operation", + "workflow": "set_query,retrieve_top_memory,print_memory", + "description": "read all long-term memory of the user" + }, + "delete_memory": { + "class": "memory.operation.frontend_operation", + "workflow": "set_query,retrieve_all_memory,delete_memory", + "description": "delete a single long-term memory" + }, + "delete_all": { + "class": "memory.operation.frontend_operation", + "workflow": "set_query,retrieve_all_memory,delete_all", + "description": "delete all long-term memory" + }, + "add_memory": { + "class": "memory.operation.frontend_operation", + "workflow": "add_memory", + "description": "add a single observation" + }, + "consolidate_memory": { + "class": "memory.operation.consolidate_memory_op", + "workflow": "info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat," + "store_memory", + "description": "summary user's observation memory", + "interval_time": 1 + }, + "reflect_and_reconsolidate": { + "class": "memory.operation.backend_operation", + "workflow": "load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat," + "store_memory", + "description": "summary user's insight memory", + "interval_time": 15 + } + } + } +} + +DEFAULT_WORKER_ARGUMENTS = { + "dummy": { + "class": "memory.worker.dummy_worker", + "generation_model": "generation_model", + "embedding_model": "embedding_model", + "rank_model": "rank_model" + }, + "read_message": { + "class": "memory.worker.frontend.read_message_worker" + }, + "set_query": { + "class": "memory.worker.frontend.set_query_worker" + }, + "retrieve_obs_ins": { + "class": "memory.worker.frontend.retrieve_memory_worker", + "retrieve_obs_top_k": 100, + "retrieve_ins_top_k": 100 + }, + "extract_time": { + "class": "memory.worker.frontend.extract_time_worker", + "generation_model": "generation_model" + }, + "semantic_rank": { + "class": "memory.worker.frontend.semantic_rank_worker", + "rank_model": "rank_model" + }, + "fuse_rerank": { + "class": "memory.worker.frontend.fuse_rerank_worker", + "fuse_score_threshold": 0.01, + "fuse_ratio_dict": { + "conversation": 0.5, + "observation": 1, + "obs_customized": 1.2, + "insight": 2 + }, + "fuse_time_ratio": 2, + "fuse_rerank_top_k": 10 + }, + "retrieve_top_memory": { + "class": "memory.worker.frontend.retrieve_memory_worker", + "retrieve_obs_top_k": 100, + "retrieve_ins_top_k": 100, + "retrieve_expired_top_k": 100 + }, + "print_memory": { + "class": "memory.worker.frontend.print_memory_worker" + }, + "retrieve_all_memory": { + "class": "memory.worker.frontend.retrieve_memory_worker", + "retrieve_obs_top_k": 1000, + "retrieve_ins_top_k": 1000, + "retrieve_expired_top_k": 1000 + }, + "delete_memory": { + "class": "memory.worker.backend.update_memory_worker", + "method": "delete_memory" + }, + "delete_all": { + "class": "memory.worker.backend.update_memory_worker", + "method": "delete_all" + }, + "add_memory": { + "class": "memory.worker.backend.update_memory_worker", + "method": "from_query" + }, + "info_filter": { + "class": "memory.worker.backend.info_filter_worker", + "generation_model": "generation_model" + }, + "load_today_memory": { + "class": "memory.worker.backend.load_memory_worker", + "retrieve_today_top_k": 100 + }, + "get_observation": { + "class": "memory.worker.backend.get_observation_worker", + "generation_model": "generation_model" + }, + "get_observation_with_time": { + "class": "memory.worker.backend.get_observation_with_time_worker", + "generation_model": "generation_model" + }, + "contra_repeat": { + "class": "memory.worker.backend.contra_repeat_worker", + "generation_model": "generation_model" + }, + "store_memory": { + "class": "memory.worker.backend.update_memory_worker", + "method": "from_memory_key", + "memory_key": "all" + }, + "load_obs_and_insight": { + "class": "memory.worker.backend.load_memory_worker", + "retrieve_not_reflected_top_k": 100, + "retrieve_not_updated_top_k": 100, + "retrieve_insight_top_k": 100 + }, + "get_reflection_subject": { + "class": "memory.worker.backend.get_reflection_subject_worker", + "generation_model": "generation_model", + "reflect_obs_cnt_threshold": 10 + }, + "update_insight": { + "class": "memory.worker.backend.update_insight_worker", + "generation_model": "generation_model", + "rank_model": "rank_model" + }, + "long_contra_repeat": { + "class": "memory.worker.backend.long_contra_repeat_worker", + "generation_model": "generation_model" + } +} + +DEFAULT_MONITOR_ARGUMENTS = { + "class": "storage.dummy_monitor" +} diff --git a/memoryscope/argument/init_handler.py b/memoryscope/argument/init_handler.py new file mode 100644 index 00000000..86a2d840 --- /dev/null +++ b/memoryscope/argument/init_handler.py @@ -0,0 +1,28 @@ +class InitializationHandler(object): + + def __init__(self): + self.file_path: str = __file__ + + self.global_config_dict: dict = {} + + self.memory_chat_dict: dict = {} + + self.memory_service_dict: dict = {} + + self.worker_dict: dict = {} + + self.model_dict: dict = {} + + self.memory_store: dict = {} + + self.monitor: dict = {} + + def update_by_arguments(self): + pass + + def load_from_config(self): + pass + + + def load_from_file(self): + pass diff --git a/memoryscope/argument/memoryscope_arguments.py b/memoryscope/argument/memoryscope_arguments.py new file mode 100644 index 00000000..fb10cd23 --- /dev/null +++ b/memoryscope/argument/memoryscope_arguments.py @@ -0,0 +1,61 @@ +from dataclasses import dataclass, field +from typing import Literal, Dict + + +@dataclass +class MemoryscopeArguments(object): + language: Literal["cn", "en"] = field(default="en", metadata={"help": "support en & cn now"}) + + thread_pool_max_workers: int = field(default=5, metadata={"help": "thread pool max workers"}) + + logger_name: str = field(default="memoryscope") + + logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S") + + memory_chat_class: str = field(default="chat.api_memory_chat", metadata={ + "help": "The memory chat class for dynamic import: chat.cli_memory_chat, chat.api_memory_chat"}) + + human_name: str = field(default="user", metadata={"help": "en: user, cn: 用户"}) + + assistant_name: str = field(default="AI") + + consolidate_memory_interval_time: int = field(default=1, metadata={ + "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) + + reflect_and_reconsolidate_interval_time: int = field(default=15, metadata={ + "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) + + worker_params: Dict[str, dict] = field(default_factory=lambda: {}, metadata={ + "help": "dict format: worker_name -> param_key -> param_value"}) + + generation_backend: str = field(default="openai_generation", metadata={ + "help": "global generation backend: openai_generation, dashscope_generation, etc."}) + + generation_model: str = field(default="gpt-4o", metadata={ + "help": "global generation model: gpt-4o, gpt-4, qwen-max, etc."}) + + generation_params: dict = field(default_factory=lambda: {}, metadata={ + "help": "global generation params: max_tokens, top_p, temperature, etc."}) + + embedding_backend: str = field(default="openai_embedding", metadata={ + "help": "global embedding backend: openai_embedding, dashscope_embedding, etc."}) + + embedding_model: str = field(default="gpt-4o", metadata={ + "help": "global embedding model: text-embedding-ada-002, text-embedding-v2, etc."}) + + embedding_params: dict = field(default_factory=lambda: {}) + + rank_backend: str = field(default="dashscope_rank", metadata={"help": "global rank backend: dashscope_rank, etc."}) + + rank_model: str = field(default="gte-rerank", metadata={"help": "global rank model: gte-rerank, etc."}) + + rank_params: dict = field(default_factory=lambda: {}) + + es_index_name: str = field(default="memory_index") + + es_url: str = field(default="http://localhost:9200") + + # TODO at xianzhe + retrieve_type: str = field(default="dense", metadata={"help": "es_retrieve_type: dense, sparse, hybrid"}) + + hybrid_alpha: float | None = field(default=1.0, metadata={"help": ""}) diff --git a/memoryscope/chat/api_memory_chat.py b/memoryscope/chat/api_memory_chat.py new file mode 100644 index 00000000..2007ef8d --- /dev/null +++ b/memoryscope/chat/api_memory_chat.py @@ -0,0 +1,322 @@ +import os +import time +from typing import List + +import questionary + +from memoryscope.chat.base_memory_chat import BaseMemoryChat +from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.enumeration.message_role_enum import MessageRoleEnum +from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.models.base_model import BaseModel +from memoryscope.scheme.message import Message +from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen +from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.utils.logger import Logger +from memoryscope.utils.prompt_handler import PromptHandler +from memoryscope.utils.tool_functions import char_logo + + +class ApiMemoryChat(BaseMemoryChat): + + def __init__(self, + memory_service: str, + generation_model: str, + stream: bool = True, + human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language], + assistant_name: str = "AI", + **kwargs): + + self._memory_service: BaseMemoryService | str = memory_service + self._generation_model: BaseModel | str = generation_model + self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {}) + + self.stream: bool = stream + self.human_name: str = human_name + self.assistant_name: str = assistant_name + self.kwargs: dict = kwargs + + self._logo = char_logo("MemoryScope") + self._prompt_handler: PromptHandler | None = None + G_CONTEXT.meta_data.update({ + "human_name": human_name, + "assistant_name": assistant_name, + }) + + self.logger = Logger.get_logger() + + @property + def prompt_handler(self) -> PromptHandler: + """ + Lazy initialization property for the prompt handler. + + This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed. + It uses the current file's path and additional keyword arguments for configuration. + + Returns: + PromptHandler: An instance of the PromptHandler configured for this CLI session. + """ + if self._prompt_handler is None: + self._prompt_handler = PromptHandler(__file__, **self.kwargs) + return self._prompt_handler + + def print_logo(self): + """ + Prints the logo of the CLI application to the console. + + The logo is composed of multiple lines, which are iterated through + and printed one by one to provide a visual identity for the chat interface. + """ + for line in self._logo: + print(line) + + @property + def memory_service(self) -> BaseMemoryService: + """ + Property to access the memory service. If the service is initially set as a string, + it will be looked up in the memory service dictionary of global context, initialized, + and then returned as an instance of `BaseMemoryService`. Ensures the memory service + is properly started before use. + + Returns: + BaseMemoryService: An active memory service instance. + + Raises: + ValueError: If the declaration of memory service is not found in the memory service dictionary of global context. + """ + if isinstance(self._memory_service, str): + if self._memory_service not in G_CONTEXT.memory_service_conf_dict: + raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service) + self._memory_service = G_CONTEXT.memory_service_conf_dict[self._memory_service] + self._memory_service.init_service() + self._memory_service.start_backend_service() + return self._memory_service + + @property + def generation_model(self) -> BaseModel: + """ + Property to get the generation model. If the model is set as a string, it will be resolved from the global + context's model dictionary. + + Raises: + ValueError: If the declaration of generation model is not found in the model dictionary of global context . + + Returns: + BaseModel: An actual generation model instance. + """ + if isinstance(self._generation_model, str): + if self._generation_model not in G_CONTEXT.model_conf_dict: + raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}") + self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model] + return self._generation_model + + def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen: + """ + Engages in a conversation with the AI model, utilizing conversation memory. + The function sends the user's query, incorporates conversation history and memory, + and optionally remembers the AI's response based on the user's preference. + + Args: + query (str): The user's input or query for the AI. + remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. + Defaults to False. + + Returns: + - ModelResponse: In non-streaming mode, returns a complete AI response. + - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. + + Side Effects: + - Updates the conversation memory with the query of user and (optionally) the response of AI. + - Retrieves and includes historical messages and memory content in the context of conversation. + """ + new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) + self.memory_service.add_messages(new_message) + + messages: List[Message] = [] + + # Incorporate memory into the system prompt if available + system_prompt = self.prompt_handler.system_prompt + memories: str = self.memory_service.retrieve_memory() + if memories: + memory_prompt = self.prompt_handler.memory_prompt + system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) + messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)) + + # Include past conversation history in the message list + history_messages = self.memory_service.read_message() + if history_messages: + messages.extend(history_messages) + + # Append the current user's message to the conversation context + messages.append(new_message) + self.logger.info(f"messages={messages}") + + # Invoke the Language Model with the constructed message context, respecting streaming setting + generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) + + # In non-streaming interactions, explicitly save the AI's reply to memory if instructed + if remember_response: + assert not self.stream # Ensure we're not in streaming mode when remembering responses + generated.message.role_name = self.assistant_name + self.memory_service.add_messages(generated.message) + + # Return the AI's response directly or as a generator based on the streaming mode + return generated + + @staticmethod + def parse_query_command(query: str): + """ + Parses the user's input query command, separating it into the command and its associated keyword arguments. + + Args: + query (str): The raw input string from the user which includes the command and its arguments. + + Returns: + tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword + arguments as the second element. + """ + query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command + command = query_split[0] # Extract the command + args = query_split[1:] # Extract the arguments following the command + kwargs = {} # Initialize dictionary to hold keyword arguments + + for arg in args: + # Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original) + if not args: + continue + arg_split = arg.split("=") # Split argument into key-value pair + if len(arg_split) >= 2: # Ensure there's both a key and value + k = arg_split[0] # Extract key + v = arg_split[1] # Extract value + if k and v: # Only add to kwargs if both key and value are non-empty + kwargs[k] = v + + return command, kwargs # Return the parsed command and keyword arguments + + def process_commands(self, query: str) -> bool: + """ + Parses and executes commands from user input in the CLI chat interface. + Supports operations like exiting, clearing screen, showing help, toggling stream mode, + executing predefined memory operations, and handling unknown commands. + + Args: + query (str): The user's input command string. + + Returns: + bool: Indicates whether to continue running the CLI after processing the command. + """ + continue_run = True + command, kwargs = self.parse_query_command(query) + + # Print prompt for AI's response + questionary.print("> ", end="", style="fg:yellow") + questionary.print(f"{self.assistant_name}: ", end="", style="bold") + + if command == "exit": + self.memory_service.stop_backend_service() + continue_run = False + + elif command == "clear": + os.system("clear") + + elif command == "help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(text=f" /{cmd}:", style="bold") + questionary.print(text=f" {desc}") + + elif command == "stream": + self.stream = not self.stream + questionary.print(f"set stream: {self.stream}") + + elif command in self.memory_service.op_description_dict: + refresh_time = kwargs.pop("refresh_time", "") + if refresh_time and refresh_time.isdigit(): + refresh_time = int(refresh_time) + self.memory_service.stop_backend_service() + while True: + result = self.memory_service.do_operation(op_name=command, **kwargs) + os.system("clear") + self.print_logo() + if result: + if isinstance(result, list): + result = "\n".join([str(x) for x in result]) + questionary.print(result) + else: + questionary.print(f"command={command} result is empty! kwargs={kwargs}") + time.sleep(refresh_time) + + else: + result = self.memory_service.do_operation(op_name=command, **kwargs) + if result: + if isinstance(result, list): + result = "\n".join([str(x) for x in result]) + questionary.print(result) + else: + questionary.print(f"command={command} result is empty! kwargs={kwargs}") + + else: + questionary.print(f"Unknown command={command} received.") + + return continue_run + + def run(self): + """ + Runs the CLI chat loop, which handles user input, processes commands, + communicates with the AI model, manages conversation memory, and controls + the chat session including streaming responses, command execution, and error handling. + + The loop continues until the user explicitly chooses to exit. + """ + self.print_logo() + self.USER_COMMANDS.update(self.memory_service.op_description_dict) + + while True: + try: + query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask() + if not query: + continue + + query: str = query.strip() + + # Handle special commands prefixed with '/' + if query.startswith("/"): + if self.process_commands(query=query): + continue + else: + break + + # Print prompt for AI's response + questionary.print("> ", end="", style="fg:yellow") + questionary.print(f"{self.assistant_name}: ", end="", style="bold") + + # Fetch and display AI's response, with support for streaming + self.memory_service.start_backend_service() + if self.stream: + model_response = None + for model_response in self.chat_with_memory(query=query): + questionary.print(model_response.delta, end="") + questionary.print("") + + else: + model_response = self.chat_with_memory(query=query) + questionary.print(model_response.message.content) + + # Append AI's response to the conversation memory + model_response.message.role_name = self.assistant_name + self.memory_service.add_messages(model_response.message) + + except KeyboardInterrupt: + # Handle user interruption and confirm exit + questionary.print("User interrupt occurred.") + is_exit = questionary.confirm("Continue exit?").unsafe_ask() + if is_exit: + self.memory_service.stop_backend_service() + break + + except Exception as e: + # Log and handle any unanticipated exceptions + import traceback + traceback.print_exc() + self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.") + continue diff --git a/memoryscope/chat/base_memory_chat.py b/memoryscope/chat/base_memory_chat.py index 2be6ad6b..72cb8a07 100644 --- a/memoryscope/chat/base_memory_chat.py +++ b/memoryscope/chat/base_memory_chat.py @@ -1,6 +1,9 @@ from abc import ABCMeta, abstractmethod +from typing import List from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.scheme.message import Message +from memoryscope.utils.logger import Logger class BaseMemoryChat(metaclass=ABCMeta): @@ -9,13 +12,19 @@ class BaseMemoryChat(metaclass=ABCMeta): It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement. """ + def __init__(self, generation_stream: bool = True, **kwargs): + self.generation_stream: bool = generation_stream + self.kwargs: dict = kwargs + self.logger = Logger.get_logger() + @abstractmethod - def chat_with_memory(self, query: str): + def chat_with_memory(self, query: str, role_name: str = ""): """ Initiates a chat interaction using the memory service, with the provided query as input. Args: query (str): The user's query or message to start the chat. + role_name (str): The role's name. Returns: This method should return the chat response generated after processing the query @@ -23,6 +32,9 @@ class BaseMemoryChat(metaclass=ABCMeta): subclass. """ + def add_message(self, messages: List[Message] | Message): + self.memory_service.add_messages(messages) + @property def memory_service(self) -> BaseMemoryService: """ diff --git a/memoryscope/chat/cli_memory_chat.py b/memoryscope/chat/cli_memory_chat.py index e06fa0fc..895ab7b9 100644 --- a/memoryscope/chat/cli_memory_chat.py +++ b/memoryscope/chat/cli_memory_chat.py @@ -8,11 +8,10 @@ from memoryscope.chat.base_memory_chat import BaseMemoryChat from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger from memoryscope.utils.prompt_handler import PromptHandler from memoryscope.utils.tool_functions import char_logo @@ -26,48 +25,33 @@ class CliMemoryChat(BaseMemoryChat): "exit": "Exit the CLI.", "clear": "Clear the command history.", "help": "Display available CLI commands and their descriptions.", - "stream": "Toggle between getting streamed responses from the model." } def __init__(self, memory_service: str, generation_model: str, - stream: bool = True, - human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language], - assistant_name: str = "AI", + context: MemoryscopeContext, + human_name: str = None, + assistant_name: str = None, **kwargs): - """ - Initializes the CLI chat instance with specified services, models, and personalized settings. - Args: - memory_service (str | BaseMemoryService): The memory service to be used for storing conversation history. - generation_model (str | BaseModel): The model responsible for generating AI responses. - stream (bool, optional): Flag indicating whether responses should be streamed. Defaults to True. - human_name (str, optional): The name assigned to the human user. Defaults to a language-specific user. - assistant_name (str, optional): The name of the AI assistant. Defaults to "AI". - **kwargs: Additional keyword arguments for flexibility or future extensions. + super().__init__(**kwargs) - Side Effects: - - Updates global context with human and AI names. - - Initializes logging for the instance. - """ self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model - self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) + self.context: MemoryscopeContext = context + self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {}) - self.stream: bool = stream self.human_name: str = human_name + if not self.human_name: + self.human_name = DEFAULT_HUMAN_NAME[self.context.language] + self.assistant_name: str = assistant_name - self.kwargs: dict = kwargs + if not self.assistant_name: + self.assistant_name = "AI" self._logo = char_logo("MemoryScope") self._prompt_handler: PromptHandler | None = None - G_CONTEXT.meta_data.update({ - "human_name": human_name, - "assistant_name": assistant_name, - }) - - self.logger = Logger.get_logger() @property def prompt_handler(self) -> PromptHandler: @@ -81,7 +65,7 @@ class CliMemoryChat(BaseMemoryChat): PromptHandler: An instance of the PromptHandler configured for this CLI session. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, **self.kwargs) + self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs) return self._prompt_handler def print_logo(self): @@ -98,7 +82,7 @@ class CliMemoryChat(BaseMemoryChat): def memory_service(self) -> BaseMemoryService: """ Property to access the memory service. If the service is initially set as a string, - it will be looked up in the memory service dictionary of global context, initialized, + it will be looked up in the memory service dictionary of context, initialized, and then returned as an instance of `BaseMemoryService`. Ensures the memory service is properly started before use. @@ -106,13 +90,15 @@ class CliMemoryChat(BaseMemoryChat): BaseMemoryService: An active memory service instance. Raises: - ValueError: If the declaration of memory service is not found in the memory service dictionary of global context. + ValueError: If the declaration of memory service is not found in the memory service dictionary of context. """ if isinstance(self._memory_service, str): - if self._memory_service not in G_CONTEXT.memory_service_dict: - raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service) - self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] - self._memory_service.init_service() + if self._memory_service not in self.context.memory_service_dict: + raise ValueError(f"Missing declaration of memory_service in context: {self._memory_service}") + + self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] + # init service & update kwargs + self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name) self._memory_service.start_backend_service() return self._memory_service @@ -123,37 +109,21 @@ class CliMemoryChat(BaseMemoryChat): context's model dictionary. Raises: - ValueError: If the declaration of generation model is not found in the model dictionary of global context . + ValueError: If the declaration of generation model is not found in the model dictionary of context . Returns: BaseModel: An actual generation model instance. """ if isinstance(self._generation_model, str): - if self._generation_model not in G_CONTEXT.model_dict: + if self._generation_model not in self.context.model_dict: raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}") - self._generation_model = G_CONTEXT.model_dict[self._generation_model] + self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model - def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen: - """ - Engages in a conversation with the AI model, utilizing conversation memory. - The function sends the user's query, incorporates conversation history and memory, - and optionally remembers the AI's response based on the user's preference. - - Args: - query (str): The user's input or query for the AI. - remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. - Defaults to False. - - Returns: - - ModelResponse: In non-streaming mode, returns a complete AI response. - - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. - - Side Effects: - - Updates the conversation memory with the query of user and (optionally) the response of AI. - - Retrieves and includes historical messages and memory content in the context of conversation. - """ - new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) + def chat_with_memory(self, query: str, role_name: str = "") -> ModelResponse | ModelResponseGen: + if not role_name: + role_name = self.human_name + new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) self.memory_service.add_messages(new_message) messages: List[Message] = [] @@ -176,16 +146,9 @@ class CliMemoryChat(BaseMemoryChat): self.logger.info(f"messages={messages}") # Invoke the Language Model with the constructed message context, respecting streaming setting - generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) - - # In non-streaming interactions, explicitly save the AI's reply to memory if instructed - if remember_response: - assert not self.stream # Ensure we're not in streaming mode when remembering responses - generated.message.role_name = self.assistant_name - self.memory_service.add_messages(generated.message) - - # Return the AI's response directly or as a generator based on the streaming mode - return generated + return self.generation_model.call(messages=messages, + stream=self.generation_stream, + **self.generation_model_kwargs) @staticmethod def parse_query_command(query: str): @@ -249,10 +212,6 @@ class CliMemoryChat(BaseMemoryChat): questionary.print(text=f" /{cmd}:", style="bold") questionary.print(text=f" {desc}") - elif command == "stream": - self.stream = not self.stream - questionary.print(f"set stream: {self.stream}") - elif command in self.memory_service.op_description_dict: refresh_time = kwargs.pop("refresh_time", "") if refresh_time and refresh_time.isdigit(): @@ -314,14 +273,13 @@ class CliMemoryChat(BaseMemoryChat): questionary.print("> ", end="", style="fg:yellow") questionary.print(f"{self.assistant_name}: ", end="", style="bold") - # Fetch and display AI's response, with support for streaming + # Fetch and display AI's response self.memory_service.start_backend_service() - if self.stream: + if self.generation_stream: model_response = None for model_response in self.chat_with_memory(query=query): questionary.print(model_response.delta, end="") questionary.print("") - else: model_response = self.chat_with_memory(query=query) questionary.print(model_response.message.content) diff --git a/memoryscope/chat/cli_memory_chat.yaml b/memoryscope/chat/memory_chat_prompt.yaml similarity index 100% rename from memoryscope/chat/cli_memory_chat.yaml rename to memoryscope/chat/memory_chat_prompt.yaml diff --git a/memoryscope/cli.py b/memoryscope/cli.py index 14b82099..55b05f5b 100644 --- a/memoryscope/cli.py +++ b/memoryscope/cli.py @@ -1,105 +1,18 @@ -import datetime import sys -import questionary +from memoryscope.chat.base_memory_chat import BaseMemoryChat +from memoryscope.memoryscope import MemoryScope sys.path.append(".") # noqa: E402 -import json -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - import fire -import yaml -import atexit - -from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.timer import timer -from memoryscope.utils.tool_functions import init_instance_by_config, camelcase_to_underscore -class MemoryScope(object): - - def __init__(self): - self.config: Dict[str, Any] = {} - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - class_name = camelcase_to_underscore(self.__class__.__name__) - self.logger: Logger = Logger.get_logger(f"{class_name}_{datetime_suffix}", to_stream=False) - - def load_config(self, path: str): - with open(path) as f: - if path.endswith("yaml"): - self.config = yaml.load(f, yaml.FullLoader) - elif path.endswith("json"): - self.config = json.load(f) - else: - raise RuntimeError("not supported config file type!") - self.init_global_content_by_config() - atexit.register(self.shutdown) # register clean up function - return self - - @staticmethod - def shutdown(): - questionary.print('Gracefully executing the shutdown function...') - G_CONTEXT.memory_store.close() - G_CONTEXT.monitor.close() - G_CONTEXT.thread_pool.shutdown() - - def set_global_config(self): - G_CONTEXT.global_config = global_config = self.config["global_config"] - G_CONTEXT.language = LanguageEnum(global_config["language"]) - G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"])) - - @timer - def init_global_content_by_config(self): - # set global config - self.set_global_config() - - # init memory_chat - for name, conf in self.config["memory_chat"].items(): - G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name) - - # 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 - if "memory_store" not in self.config: - raise RuntimeError("memory_store config is required!") - memory_store_config = self.config["memory_store"] - embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]] - G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model) - - # init monitor - G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) - - # set worker config - G_CONTEXT.worker_config = self.config["worker"] - - @property - def default_chat_handle(self): - return list(G_CONTEXT.memory_chat_dict.values())[0] - - @property - def default_service(self): - return self.default_chat_handle.memory_service - - -class CliJob(MemoryScope): - - def run(self, config: str): - self.load_config(config) - self.init_global_content_by_config() - self.default_chat_handle.run() +def cli_job(config_path: str): + ms = MemoryScope(config_path=config_path) + memory_chat: BaseMemoryChat = ms.default_memory_chat + memory_chat.run() if __name__ == "__main__": - cli_job = CliJob() - fire.Fire(cli_job.run) + fire.Fire(cli_job) diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/memory/service/base_memory_service.py index 31b2df5e..3102dc0e 100644 --- a/memoryscope/memory/service/base_memory_service.py +++ b/memoryscope/memory/service/base_memory_service.py @@ -1,46 +1,40 @@ -import threading from abc import ABCMeta, abstractmethod from typing import List, Dict from memoryscope.memory.operation.base_operation import BaseOperation +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.scheme.message import Message from memoryscope.utils.logger import Logger class BaseMemoryService(metaclass=ABCMeta): """ - An abstract base class for managing memory operations within a multi-threaded context. + An abstract base class for managing memory operations within a multithreaded context. It sets up the infrastructure for operation handling, message storage, and synchronization, along with logging capabilities and customizable configurations. """ - def __init__(self, - memory_operations: Dict[str, dict], - retrieve_memory_key: str = "retrieve_memory", - read_message_key: str = "read_message", - **kwargs): + def __init__(self, memory_operations: Dict[str, dict], context: MemoryscopeContext, **kwargs): """ Initializes the BaseMemoryService with operation definitions, keys for memory access, and additional keyword arguments for flexibility. Args: memory_operations (Dict[str, dict]): A dictionary defining available memory operations. - retrieve_memory_key (str): The key indicating a retrieve memory operation. Defaults to "retrieve_memory". - read_message_key (str): The key for reading messages. Defaults to "read_message". **kwargs: Additional parameters to customize service behavior. """ - self.memory_operations: Dict[str, dict] = memory_operations - self.retrieve_memory_key: str = retrieve_memory_key - self.read_message_key: str = read_message_key + self.memory_operations_conf: Dict[str, dict] = memory_operations + self.context: MemoryscopeContext = context self._operation_dict: Dict[str, BaseOperation] = {} self._op_description_dict: Dict[str, str] = {} - self.chat_messages: List[Message] = [] - self.message_lock = threading.Lock() self.logger = Logger.get_logger() self.kwargs = kwargs + def update_kwargs(self, **kwargs): + pass + @abstractmethod def add_messages(self, messages: List[Message] | Message): raise NotImplementedError diff --git a/memoryscope/memory/service/memory_scope_service.py b/memoryscope/memory/service/memory_scope_service.py index da4f4b76..fdcfbc30 100644 --- a/memoryscope/memory/service/memory_scope_service.py +++ b/memoryscope/memory/service/memory_scope_service.py @@ -28,6 +28,9 @@ class MemoryScopeService(BaseMemoryService): self.contextual_msg_min_count: int = contextual_msg_min_count assert history_msg_count >= contextual_msg_max_count >= contextual_msg_min_count + self.chat_messages: List[Message] = [] + self.message_lock = threading.Lock() + def add_messages(self, messages: List[Message] | Message): """ Adds a single message or a list of messages to the chat history, ensuring the message list diff --git a/memoryscope/memory/worker/memory_base_worker.py b/memoryscope/memory/worker/memory_base_worker.py index 5d0ee99a..00fb72be 100644 --- a/memoryscope/memory/worker/memory_base_worker.py +++ b/memoryscope/memory/worker/memory_base_worker.py @@ -87,7 +87,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The embedding model used for converting text into vector representations. """ if isinstance(self._embedding_model, str): - self._embedding_model = G_CONTEXT.model_dict[self._embedding_model] + self._embedding_model = G_CONTEXT.model_conf_dict[self._embedding_model] # ⭐ Retrieve the actual model instance when the attribute is a string reference return self._embedding_model @@ -101,7 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The model used for text generation. """ if isinstance(self._generation_model, str): - self._generation_model = G_CONTEXT.model_dict[self._generation_model] + self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model] # ⭐ Retrieve the model instance if currently a string reference return self._generation_model @@ -115,7 +115,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The rank model instance used for ranking tasks. """ if isinstance(self._rank_model, str): - self._rank_model = G_CONTEXT.model_dict[self._rank_model] # Fetch model instance if string reference + self._rank_model = G_CONTEXT.model_conf_dict[self._rank_model] # Fetch model instance if string reference return self._rank_model @property @@ -128,7 +128,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store + self._memory_store = G_CONTEXT.memory_store_conf return self._memory_store @property @@ -141,7 +141,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMonitor: The monitoring component instance. """ if self._monitor is None: - self._monitor = G_CONTEXT.monitor + self._monitor = G_CONTEXT.monitor_conf return self._monitor @property diff --git a/memoryscope/memoryscope.py b/memoryscope/memoryscope.py new file mode 100644 index 00000000..520bd127 --- /dev/null +++ b/memoryscope/memoryscope.py @@ -0,0 +1,198 @@ +import datetime +import json +from concurrent.futures import ThreadPoolExecutor + +import yaml + +from memoryscope.argument import default_arguments +from memoryscope.argument.memoryscope_arguments import MemoryscopeArguments +from memoryscope.chat.base_memory_chat import BaseMemoryChat +from memoryscope.enumeration.language_enum import LanguageEnum +from memoryscope.enumeration.model_enum import ModelEnum +from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.memoryscope_context import MemoryscopeContext +from memoryscope.utils.logger import Logger +from memoryscope.utils.tool_functions import init_instance_by_config + + +class MemoryScope(object): + + def __init__(self, + arguments: MemoryscopeArguments | None = None, + config: dict | None = None, + config_path: str = ""): + + self.global_conf: dict = {} + self.memory_chat_conf_dict: dict = {} + self.memory_service_conf_dict: dict = {} + self.worker_conf_dict: dict = {} + self.model_conf_dict: dict = {} + self.memory_store_conf: dict = {} + self.monitor_conf: dict = {} + + self.context: MemoryscopeContext = MemoryscopeContext() + + if arguments: + self._init_by_arguments(arguments=arguments) + elif config: + self._init_by_config(config=config) + elif config_path: + self._init_by_config_path(config_path=config_path) + else: + raise RuntimeError("At least one of arguments, config, or file_path must not be empty!") + + self.logger = self._init_logger() + + self._init_context_by_config() + + def _init_by_arguments(self, arguments: MemoryscopeArguments): + # prepare global + self.global_conf = { + "language": arguments.language, + "thread_pool_max_workers": arguments.thread_pool_max_workers, + "logger_name": arguments.logger_name, + "logger_name_time_suffix": arguments.logger_name_time_suffix, + } + + # prepare memory chat + self.memory_chat_conf_dict = default_arguments.DEFAULT_MEMORY_CHAT_ARGUMENTS.copy() + memory_chat_config = list(self.memory_chat_conf_dict.values())[0] + memory_chat_config.update({ + "class": arguments.memory_chat_class, + "human_name": arguments.human_name, + "assistant_name": arguments.assistant_name, + }) + + # prepare memory service + self.memory_service_conf_dict = default_arguments.DEFAULT_MEMORY_SERVICE_ARGUMENTS.copy() + memory_service_config = list(self.memory_service_conf_dict.values())[0] + memory_service_config.update({ + "human_name": arguments.human_name, + "assistant_name": arguments.assistant_name, + }) + memory_service_config["memory_operations"]["consolidate_memory"]["interval_time"] = \ + arguments.consolidate_memory_interval_time + memory_service_config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \ + arguments.reflect_and_reconsolidate_interval_time + + # prepare memory service + self.worker_conf_dict = default_arguments.DEFAULT_WORKER_ARGUMENTS.copy() + if arguments.worker_params: + for worker_name, kv_dict in arguments.worker_params.items(): + if worker_name not in self.worker_conf_dict: + continue + self.worker_conf_dict[worker_name].update(kv_dict) + + # prepare models + self.model_conf_dict = { + "generation_model": { + "class": "models.llama_index_generation_model", + "module_name": arguments.generation_backend, + "model_name": arguments.generation_model, + **arguments.generation_params, + }, + "embedding_model": { + "class": "models.llama_index_embedding_model", + "module_name": arguments.embedding_backend, + "model_name": arguments.embedding_model, + **arguments.embedding_params, + }, + "rank_model": { + "class": "models.llama_index_rank_model", + "module_name": arguments.rank_backend, + "model_name": arguments.rank_model, + **arguments.rank_params, + }, + } + + # prepare memory store + self.memory_store_conf = { + "class": "storage.llama_index_es_memory_store", + "embedding_model": "embedding_model", + "index_name": arguments.es_index_name, + "es_url": arguments.es_url, + "retrieve_type": arguments.retrieve_type, + "hybrid_alpha": arguments.hybrid_alpha, + } + + self.monitor_conf = default_arguments.DEFAULT_MONITOR_ARGUMENTS.copy() + + def _init_by_config(self, config: dict): + self.global_conf = config["global_config"] + self.memory_service_conf_dict = config["memory_service"] + self.worker_conf_dict = config["worker"] + self.model_conf_dict = config["model"] + self.memory_store_conf = config["memory_store"] + + # not necessary + self.memory_chat_conf_dict = config.get("memory_chat") + self.monitor_conf = config.get("monitor") + + def _init_by_config_path(self, config_path: str): + with open(config_path) as f: + if config_path.endswith("yaml"): + config = yaml.load(f, yaml.FullLoader) + elif config_path.endswith("json"): + config = json.load(f) + else: + raise RuntimeError("not supported config file type!") + return self._init_by_config(config) + + def _init_logger(self) -> Logger: + logger_name = self.global_conf.get("logger_name") + assert logger_name, "logger_name is empty!" + logger_name_time_suffix = self.global_conf.get("logger_name_time_suffix") + if logger_name_time_suffix: + suffix = datetime.datetime.now().strftime(logger_name_time_suffix) + logger_name = f"{logger_name}_{suffix}" + return Logger.get_logger(logger_name, to_stream=False) + + def _init_context_by_config(self): + # set global config + self.context.language = LanguageEnum(self.global_conf["language"]) + self.context.thread_pool = ThreadPoolExecutor(max_workers=self.global_conf["max_workers"]) + + # init memory_chat + if self.memory_chat_conf_dict: + for name, conf in self.memory_chat_conf_dict.items(): + self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context) + + # set memory_service + assert self.memory_service_conf_dict + for name, conf in self.memory_service_conf_dict.items(): + self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context) + + # init models + assert self.model_conf_dict + for name, conf in self.model_conf_dict.items(): + self.context.model_dict[name] = init_instance_by_config(conf, name=name) + + # init vector_store + assert self.memory_store_conf + emb_model_name: str = self.memory_store_conf[ModelEnum.EMBEDDING_MODEL.value] + embedding_model = self.context.model_dict[emb_model_name] + self.context.memory_store = init_instance_by_config(self.memory_store_conf, embedding_model=embedding_model) + + # init monitor + if self.monitor_conf: + self.context.monitor = init_instance_by_config(self.monitor_conf) + + # set worker config + self.context.worker_config = self.worker_conf_dict + + def close(self): + for _, service in self.context.memory_service_dict.items(): + service.stop_backend_service() + self.context.memory_store.close() + self.context.thread_pool.shutdown() + + if self.context.monitor: + self.context.monitor.close() + + @property + def default_memory_chat(self) -> BaseMemoryChat: + return list(self.context.memory_chat_dict.values())[0] + + @property + def default_service(self) -> BaseMemoryService: + return list(self.context.memory_service_dict.values())[0] diff --git a/memoryscope/memoryscope_context.py b/memoryscope/memoryscope_context.py new file mode 100644 index 00000000..967bb77d --- /dev/null +++ b/memoryscope/memoryscope_context.py @@ -0,0 +1,27 @@ +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field + +from memoryscope.enumeration.language_enum import LanguageEnum + + +@dataclass +class MemoryscopeContext(object): + """ + The context class archives all configs utilized by store, monitor, services and workers. + """ + + language: LanguageEnum = LanguageEnum.EN + + thread_pool: ThreadPoolExecutor | None = None + + memory_store = None + + monitor = None + + memory_chat_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_chat"}) + + memory_service_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_service"}) + + model_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> model"}) + + worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"}) diff --git a/memoryscope/utils/global_context.py b/memoryscope/utils/global_context.py deleted file mode 100644 index 81ac9e07..00000000 --- a/memoryscope/utils/global_context.py +++ /dev/null @@ -1,33 +0,0 @@ -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - -from memoryscope.chat.base_memory_chat import BaseMemoryChat -from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.models.base_model import BaseModel -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.base_monitor import BaseMonitor - - -class GlobalContext(object): - """ - The GlobalContext class archives all configs utilized by store, monitor, services and workers. - """ - - def __init__(self): - self.global_config: Dict[str, Any] = {} - self.worker_config: Dict[str, Dict[str, Any]] = {} - - self.memory_service_dict: Dict[str, BaseMemoryService] = {} - self.model_dict: Dict[str, BaseModel] = {} - self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} - - self.memory_store: BaseMemoryStore | None = None - self.monitor: BaseMonitor | None = None - self.thread_pool: ThreadPoolExecutor | None = None - self.language: LanguageEnum = LanguageEnum.EN - - self.meta_data: Dict[str, Any] = {} - - -G_CONTEXT = GlobalContext() diff --git a/memoryscope/utils/memory_handler.py b/memoryscope/utils/memory_handler.py index 36ec9b9d..a3cbf889 100644 --- a/memoryscope/utils/memory_handler.py +++ b/memoryscope/utils/memory_handler.py @@ -36,7 +36,7 @@ class MemoryHandler(object): BaseMemoryStore: The memory store instance associated with this worker. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store + self._memory_store = G_CONTEXT.memory_store_conf return self._memory_store def clear(self): diff --git a/memoryscope/utils/prompt_handler.py b/memoryscope/utils/prompt_handler.py index f2e9b380..d44559f3 100644 --- a/memoryscope/utils/prompt_handler.py +++ b/memoryscope/utils/prompt_handler.py @@ -1,19 +1,25 @@ import json import os.path +from pathlib import Path from typing import Dict import yaml -from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.enumeration.language_enum import LanguageEnum class PromptHandler(object): """ The `PromptHandler` class manages prompt messages by loading them from YAML or JSON files and dictionaries, - supporting language selection based on a global context, and providing dictionary-like access to the prompt messages. + supporting language selection based on a context, and providing dictionary-like access to the prompt messages. """ - def __init__(self, class_path: str, prompt_file: str = "", prompt_dict: dict = None, **kwargs): + def __init__(self, + class_path: str, + prompt_file: str = "", + prompt_dict: dict = None, + language_enum: LanguageEnum = LanguageEnum.EN, + **kwargs): """ Initializes the PromptHandler with paths to prompt sources and additional keyword arguments. @@ -21,30 +27,32 @@ class PromptHandler(object): class_path (str): The path to the class where prompts are utilized. prompt_file (str, optional): The path to an external file containing prompts. Defaults to "". prompt_dict (dict, optional): A dictionary directly containing prompt definitions. Defaults to None. + language_enum (LanguageEnum): context language. **kwargs: Additional keyword arguments that might be used in prompt handling. """ - self._class_path: str = class_path - self._prompt_dict: Dict[str, str] = {} + class_path: Path = Path(class_path) + self._class_dir: Path = class_path.parent + self._class_name: str = class_path.stem + self._language_enum: LanguageEnum = language_enum self.kwargs = kwargs - file_path = self._class_path.strip(".py") - - self.add_prompt_file(file_path) + self._prompt_dict: Dict[str, str] = {} + self.add_prompt_file((self._class_dir / self._class_name).__str__(), raise_exception=False) if prompt_file: - self.add_prompt_file(prompt_file) - + self.add_prompt_file((self._class_dir / prompt_file).__str__()) if prompt_dict: self.add_prompt_dict(prompt_dict) @staticmethod - def file_path_completion(file_path: str) -> str: + def file_path_completion(file_path: str, raise_exception: bool = True) -> str: """ Attempts to complete the given file path by appending either a `.yaml` or `.json` extension based on the existence of the respective file. If neither exists, an exception is raised. Args: file_path (str): The base path of the file to be completed. + raise_exception (bool): If the file cannot be found, report an error. Returns: str: The completed file path with the appropriate extension. @@ -61,9 +69,10 @@ class PromptHandler(object): if os.path.exists(f"{file_path}.json"): return f"{file_path}.json" - raise RuntimeError(f"{file_path}/yaml/json is not exists!") + if raise_exception: + raise RuntimeError(f"{file_path}/yaml/json is not exists!") - def add_prompt_file(self, file_path: str): + def add_prompt_file(self, file_path: str, raise_exception: bool = True): """ Adds prompt messages from a YAML or JSON file to the internal dictionary. @@ -72,8 +81,11 @@ class PromptHandler(object): Args: file_path (str): The path to the YAML or JSON file containing the prompts. + raise_exception (bool): If the file cannot be found, report an error. """ - file_path = self.file_path_completion(file_path) + file_path = self.file_path_completion(file_path, raise_exception=raise_exception) + if not file_path: + return prompt_dict = {} @@ -102,9 +114,9 @@ class PromptHandler(object): RuntimeError: If a prompt message for the current language is not found. """ for key, language_dict in prompt_dict.items(): - prompts = language_dict.get(G_CONTEXT.language) + prompts = language_dict.get(self._language_enum.value) if not prompts: - raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!") + raise RuntimeError(f"{key}.prompt.{self._language_enum.value} is empty!") self._prompt_dict[key] = prompts.strip() @property diff --git a/memoryscope/utils/tool_functions.py b/memoryscope/utils/tool_functions.py index 2fd432c1..9bb73896 100644 --- a/memoryscope/utils/tool_functions.py +++ b/memoryscope/utils/tool_functions.py @@ -47,10 +47,7 @@ def camelcase_to_underscore(name: str) -> str: return re.sub(r'(? Date: Fri, 26 Jul 2024 10:54:19 +0800 Subject: [PATCH 02/15] finish memory chat --- memoryscope/chat/api_memory_chat.py | 252 +++--------------- memoryscope/chat/base_memory_chat.py | 13 +- memoryscope/chat/cli_memory_chat.py | 30 ++- .../memory/operation/backend_operation.py | 10 +- .../memory/service/base_memory_service.py | 6 +- 5 files changed, 84 insertions(+), 227 deletions(-) diff --git a/memoryscope/chat/api_memory_chat.py b/memoryscope/chat/api_memory_chat.py index 2007ef8d..7c0ab78a 100644 --- a/memoryscope/chat/api_memory_chat.py +++ b/memoryscope/chat/api_memory_chat.py @@ -1,20 +1,14 @@ -import os -import time from typing import List -import questionary - from memoryscope.chat.base_memory_chat import BaseMemoryChat from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger from memoryscope.utils.prompt_handler import PromptHandler -from memoryscope.utils.tool_functions import char_logo class ApiMemoryChat(BaseMemoryChat): @@ -22,28 +16,27 @@ class ApiMemoryChat(BaseMemoryChat): def __init__(self, memory_service: str, generation_model: str, - stream: bool = True, - human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language], - assistant_name: str = "AI", + context: MemoryscopeContext, + human_name: str = None, + assistant_name: str = None, **kwargs): + super().__init__(**kwargs) + self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model + self.context: MemoryscopeContext = context self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {}) - self.stream: bool = stream self.human_name: str = human_name + if not self.human_name: + self.human_name = DEFAULT_HUMAN_NAME[self.context.language] + self.assistant_name: str = assistant_name - self.kwargs: dict = kwargs + if not self.assistant_name: + self.assistant_name = "AI" - self._logo = char_logo("MemoryScope") self._prompt_handler: PromptHandler | None = None - G_CONTEXT.meta_data.update({ - "human_name": human_name, - "assistant_name": assistant_name, - }) - - self.logger = Logger.get_logger() @property def prompt_handler(self) -> PromptHandler: @@ -57,24 +50,14 @@ class ApiMemoryChat(BaseMemoryChat): PromptHandler: An instance of the PromptHandler configured for this CLI session. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, **self.kwargs) + self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs) return self._prompt_handler - def print_logo(self): - """ - Prints the logo of the CLI application to the console. - - The logo is composed of multiple lines, which are iterated through - and printed one by one to provide a visual identity for the chat interface. - """ - for line in self._logo: - print(line) - @property def memory_service(self) -> BaseMemoryService: """ Property to access the memory service. If the service is initially set as a string, - it will be looked up in the memory service dictionary of global context, initialized, + it will be looked up in the memory service dictionary of context, initialized, and then returned as an instance of `BaseMemoryService`. Ensures the memory service is properly started before use. @@ -82,13 +65,15 @@ class ApiMemoryChat(BaseMemoryChat): BaseMemoryService: An active memory service instance. Raises: - ValueError: If the declaration of memory service is not found in the memory service dictionary of global context. + ValueError: If the declaration of memory service is not found in the memory service dictionary of context. """ if isinstance(self._memory_service, str): - if self._memory_service not in G_CONTEXT.memory_service_conf_dict: - raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service) - self._memory_service = G_CONTEXT.memory_service_conf_dict[self._memory_service] - self._memory_service.init_service() + if self._memory_service not in self.context.memory_service_dict: + raise ValueError(f"Missing declaration of memory_service in context: {self._memory_service}") + + self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] + # init service & update kwargs + self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name) self._memory_service.start_backend_service() return self._memory_service @@ -99,18 +84,18 @@ class ApiMemoryChat(BaseMemoryChat): context's model dictionary. Raises: - ValueError: If the declaration of generation model is not found in the model dictionary of global context . + ValueError: If the declaration of generation model is not found in the model dictionary of context . Returns: BaseModel: An actual generation model instance. """ if isinstance(self._generation_model, str): - if self._generation_model not in G_CONTEXT.model_conf_dict: + if self._generation_model not in self.context.model_dict: raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}") - self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model] + self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model - def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen: + def chat_with_memory(self, query: str, role_name: str = "") -> ModelResponse | ModelResponseGen: """ Engages in a conversation with the AI model, utilizing conversation memory. The function sends the user's query, incorporates conversation history and memory, @@ -118,8 +103,7 @@ class ApiMemoryChat(BaseMemoryChat): Args: query (str): The user's input or query for the AI. - remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. - Defaults to False. + role_name (str, optional): The user's name, default value is human_name. Returns: - ModelResponse: In non-streaming mode, returns a complete AI response. @@ -129,8 +113,10 @@ class ApiMemoryChat(BaseMemoryChat): - Updates the conversation memory with the query of user and (optionally) the response of AI. - Retrieves and includes historical messages and memory content in the context of conversation. """ - new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) - self.memory_service.add_messages(new_message) + if not role_name: + role_name = self.human_name + new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) + self.add_messages(new_message) messages: List[Message] = [] @@ -151,172 +137,20 @@ class ApiMemoryChat(BaseMemoryChat): messages.append(new_message) self.logger.info(f"messages={messages}") - # Invoke the Language Model with the constructed message context, respecting streaming setting - generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) + result = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) - # In non-streaming interactions, explicitly save the AI's reply to memory if instructed - if remember_response: - assert not self.stream # Ensure we're not in streaming mode when remembering responses - generated.message.role_name = self.assistant_name - self.memory_service.add_messages(generated.message) + if self.stream: + assert isinstance(result, ModelResponseGen) + model_response: ModelResponse | None = None + for model_response in result: + yield model_response - # Return the AI's response directly or as a generator based on the streaming mode - return generated - - @staticmethod - def parse_query_command(query: str): - """ - Parses the user's input query command, separating it into the command and its associated keyword arguments. - - Args: - query (str): The raw input string from the user which includes the command and its arguments. - - Returns: - tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword - arguments as the second element. - """ - query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command - command = query_split[0] # Extract the command - args = query_split[1:] # Extract the arguments following the command - kwargs = {} # Initialize dictionary to hold keyword arguments - - for arg in args: - # Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original) - if not args: - continue - arg_split = arg.split("=") # Split argument into key-value pair - if len(arg_split) >= 2: # Ensure there's both a key and value - k = arg_split[0] # Extract key - v = arg_split[1] # Extract value - if k and v: # Only add to kwargs if both key and value are non-empty - kwargs[k] = v - - return command, kwargs # Return the parsed command and keyword arguments - - def process_commands(self, query: str) -> bool: - """ - Parses and executes commands from user input in the CLI chat interface. - Supports operations like exiting, clearing screen, showing help, toggling stream mode, - executing predefined memory operations, and handling unknown commands. - - Args: - query (str): The user's input command string. - - Returns: - bool: Indicates whether to continue running the CLI after processing the command. - """ - continue_run = True - command, kwargs = self.parse_query_command(query) - - # Print prompt for AI's response - questionary.print("> ", end="", style="fg:yellow") - questionary.print(f"{self.assistant_name}: ", end="", style="bold") - - if command == "exit": - self.memory_service.stop_backend_service() - continue_run = False - - elif command == "clear": - os.system("clear") - - elif command == "help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(text=f" /{cmd}:", style="bold") - questionary.print(text=f" {desc}") - - elif command == "stream": - self.stream = not self.stream - questionary.print(f"set stream: {self.stream}") - - elif command in self.memory_service.op_description_dict: - refresh_time = kwargs.pop("refresh_time", "") - if refresh_time and refresh_time.isdigit(): - refresh_time = int(refresh_time) - self.memory_service.stop_backend_service() - while True: - result = self.memory_service.do_operation(op_name=command, **kwargs) - os.system("clear") - self.print_logo() - if result: - if isinstance(result, list): - result = "\n".join([str(x) for x in result]) - questionary.print(result) - else: - questionary.print(f"command={command} result is empty! kwargs={kwargs}") - time.sleep(refresh_time) - - else: - result = self.memory_service.do_operation(op_name=command, **kwargs) - if result: - if isinstance(result, list): - result = "\n".join([str(x) for x in result]) - questionary.print(result) - else: - questionary.print(f"command={command} result is empty! kwargs={kwargs}") + if model_response and model_response.message: + self.add_messages(model_response.message) else: - questionary.print(f"Unknown command={command} received.") - - return continue_run - - def run(self): - """ - Runs the CLI chat loop, which handles user input, processes commands, - communicates with the AI model, manages conversation memory, and controls - the chat session including streaming responses, command execution, and error handling. - - The loop continues until the user explicitly chooses to exit. - """ - self.print_logo() - self.USER_COMMANDS.update(self.memory_service.op_description_dict) - - while True: - try: - query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask() - if not query: - continue - - query: str = query.strip() - - # Handle special commands prefixed with '/' - if query.startswith("/"): - if self.process_commands(query=query): - continue - else: - break - - # Print prompt for AI's response - questionary.print("> ", end="", style="fg:yellow") - questionary.print(f"{self.assistant_name}: ", end="", style="bold") - - # Fetch and display AI's response, with support for streaming - self.memory_service.start_backend_service() - if self.stream: - model_response = None - for model_response in self.chat_with_memory(query=query): - questionary.print(model_response.delta, end="") - questionary.print("") - - else: - model_response = self.chat_with_memory(query=query) - questionary.print(model_response.message.content) - - # Append AI's response to the conversation memory - model_response.message.role_name = self.assistant_name - self.memory_service.add_messages(model_response.message) - - except KeyboardInterrupt: - # Handle user interruption and confirm exit - questionary.print("User interrupt occurred.") - is_exit = questionary.confirm("Continue exit?").unsafe_ask() - if is_exit: - self.memory_service.stop_backend_service() - break - - except Exception as e: - # Log and handle any unanticipated exceptions - import traceback - traceback.print_exc() - self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.") - continue + assert isinstance(result, ModelResponse) + model_response: ModelResponse = result + if model_response and model_response.message: + self.add_messages(model_response.message) + return model_response diff --git a/memoryscope/chat/base_memory_chat.py b/memoryscope/chat/base_memory_chat.py index 72cb8a07..d67c66be 100644 --- a/memoryscope/chat/base_memory_chat.py +++ b/memoryscope/chat/base_memory_chat.py @@ -12,8 +12,8 @@ class BaseMemoryChat(metaclass=ABCMeta): It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement. """ - def __init__(self, generation_stream: bool = True, **kwargs): - self.generation_stream: bool = generation_stream + def __init__(self, stream: bool = True, **kwargs): + self.stream: bool = stream self.kwargs: dict = kwargs self.logger = Logger.get_logger() @@ -32,9 +32,6 @@ class BaseMemoryChat(metaclass=ABCMeta): subclass. """ - def add_message(self, messages: List[Message] | Message): - self.memory_service.add_messages(messages) - @property def memory_service(self) -> BaseMemoryService: """ @@ -45,6 +42,12 @@ class BaseMemoryChat(metaclass=ABCMeta): """ raise NotImplementedError + def add_messages(self, messages: List[Message] | Message): + self.memory_service.add_messages(messages) + + def do_memory_operation(self, op_name: str, **kwargs): + return self.memory_service.do_operation(op_name=op_name, **kwargs) + def run(self): """ Abstract method to run the chat system. diff --git a/memoryscope/chat/cli_memory_chat.py b/memoryscope/chat/cli_memory_chat.py index 895ab7b9..a4ec880a 100644 --- a/memoryscope/chat/cli_memory_chat.py +++ b/memoryscope/chat/cli_memory_chat.py @@ -25,6 +25,7 @@ class CliMemoryChat(BaseMemoryChat): "exit": "Exit the CLI.", "clear": "Clear the command history.", "help": "Display available CLI commands and their descriptions.", + "stream": "Toggle between getting streamed responses from the model." } def __init__(self, @@ -121,10 +122,27 @@ class CliMemoryChat(BaseMemoryChat): return self._generation_model def chat_with_memory(self, query: str, role_name: str = "") -> ModelResponse | ModelResponseGen: + """ + Engages in a conversation with the AI model, utilizing conversation memory. + The function sends the user's query, incorporates conversation history and memory, + and optionally remembers the AI's response based on the user's preference. + + Args: + query (str): The user's input or query for the AI. + role_name (str, optional): The user's name, default value is human_name. + + Returns: + - ModelResponse: In non-streaming mode, returns a complete AI response. + - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. + + Side Effects: + - Updates the conversation memory with the query of user and (optionally) the response of AI. + - Retrieves and includes historical messages and memory content in the context of conversation. + """ if not role_name: role_name = self.human_name new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) - self.memory_service.add_messages(new_message) + self.add_messages(new_message) messages: List[Message] = [] @@ -147,7 +165,7 @@ class CliMemoryChat(BaseMemoryChat): # Invoke the Language Model with the constructed message context, respecting streaming setting return self.generation_model.call(messages=messages, - stream=self.generation_stream, + stream=self.stream, **self.generation_model_kwargs) @staticmethod @@ -212,6 +230,10 @@ class CliMemoryChat(BaseMemoryChat): questionary.print(text=f" /{cmd}:", style="bold") questionary.print(text=f" {desc}") + elif command == "stream": + self.stream = not self.stream + questionary.print(f"set stream: {self.stream}") + elif command in self.memory_service.op_description_dict: refresh_time = kwargs.pop("refresh_time", "") if refresh_time and refresh_time.isdigit(): @@ -275,7 +297,7 @@ class CliMemoryChat(BaseMemoryChat): # Fetch and display AI's response self.memory_service.start_backend_service() - if self.generation_stream: + if self.stream: model_response = None for model_response in self.chat_with_memory(query=query): questionary.print(model_response.delta, end="") @@ -286,7 +308,7 @@ class CliMemoryChat(BaseMemoryChat): # Append AI's response to the conversation memory model_response.message.role_name = self.assistant_name - self.memory_service.add_messages(model_response.message) + self.add_messages(model_response.message) except KeyboardInterrupt: # Handle user interruption and confirm exit diff --git a/memoryscope/memory/operation/backend_operation.py b/memoryscope/memory/operation/backend_operation.py index 2fcbe398..c52cf936 100644 --- a/memoryscope/memory/operation/backend_operation.py +++ b/memoryscope/memory/operation/backend_operation.py @@ -1,11 +1,11 @@ import time from typing import List + from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memoryscope.memory.operation.base_workflow import BaseWorkflow from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT from memoryscope.utils.logger import Logger @@ -30,7 +30,7 @@ class BackendOperation(BaseWorkflow, BaseOperation): self._operation_status_run: bool = False self._loop_switch: bool = False - self._run_thread = None + self._backend_task = None self.logger = Logger.get_logger() @@ -114,10 +114,12 @@ class BackendOperation(BaseWorkflow, BaseOperation): """ if not self._loop_switch: self._loop_switch = True - self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation) + self._backend_task = G_CONTEXT.thread_pool.submit(self._loop_operation) - def stop_operation_backend(self): + def stop_operation_backend(self, wait_task_end: bool = False): """ Stops the background operation loop by setting the _loop_switch to False. """ self._loop_switch = False + if wait_task_end and self._backend_task: + self._backend_task.result() diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/memory/service/base_memory_service.py index 3102dc0e..d72fc4c1 100644 --- a/memoryscope/memory/service/base_memory_service.py +++ b/memoryscope/memory/service/base_memory_service.py @@ -25,15 +25,11 @@ class BaseMemoryService(metaclass=ABCMeta): """ self.memory_operations_conf: Dict[str, dict] = memory_operations self.context: MemoryscopeContext = context + self.kwargs = kwargs self._operation_dict: Dict[str, BaseOperation] = {} self._op_description_dict: Dict[str, str] = {} - self.logger = Logger.get_logger() - self.kwargs = kwargs - - def update_kwargs(self, **kwargs): - pass @abstractmethod def add_messages(self, messages: List[Message] | Message): From 149df5ffa6cc636069b48ab9ec41b9852379c068 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 26 Jul 2024 16:23:28 +0800 Subject: [PATCH 03/15] redefine memory chat interface --- memoryscope/chat/api_memory_chat.py | 79 ++++++++-------- memoryscope/chat/base_memory_chat.py | 23 +++-- memoryscope/chat/cli_memory_chat.py | 90 ++++++++++--------- .../memory/service/base_memory_service.py | 67 ++++++-------- .../worker/frontend/set_query_worker.py | 25 +++++- tests/other/test_attr.py | 15 ++++ 6 files changed, 170 insertions(+), 129 deletions(-) create mode 100644 tests/other/test_attr.py diff --git a/memoryscope/chat/api_memory_chat.py b/memoryscope/chat/api_memory_chat.py index 7c0ab78a..da1d73de 100644 --- a/memoryscope/chat/api_memory_chat.py +++ b/memoryscope/chat/api_memory_chat.py @@ -74,7 +74,6 @@ class ApiMemoryChat(BaseMemoryChat): self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] # init service & update kwargs self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name) - self._memory_service.start_backend_service() return self._memory_service @property @@ -95,62 +94,68 @@ class ApiMemoryChat(BaseMemoryChat): self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model - def chat_with_memory(self, query: str, role_name: str = "") -> ModelResponse | ModelResponseGen: - """ - Engages in a conversation with the AI model, utilizing conversation memory. - The function sends the user's query, incorporates conversation history and memory, - and optionally remembers the AI's response based on the user's preference. - - Args: - query (str): The user's input or query for the AI. - role_name (str, optional): The user's name, default value is human_name. - - Returns: - - ModelResponse: In non-streaming mode, returns a complete AI response. - - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. - - Side Effects: - - Updates the conversation memory with the query of user and (optionally) the response of AI. - - Retrieves and includes historical messages and memory content in the context of conversation. - """ + def get_new_message(self, query: str, role_name: str = "") -> Message: if not role_name: role_name = self.human_name - new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) - self.add_messages(new_message) - - messages: List[Message] = [] + return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) + def get_system_message_with_memory(self, memories: str) -> Message: # Incorporate memory into the system prompt if available system_prompt = self.prompt_handler.system_prompt - memories: str = self.memory_service.retrieve_memory() if memories: memory_prompt = self.prompt_handler.memory_prompt system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) - messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)) + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) + + def chat_with_memory(self, + query: str, + role_name: str = "", + remember_response: bool = True): + + chat_messages: List[Message] = [] + + new_message: Message = self.get_new_message(query=query, role_name=role_name) + + # To retrieve memory, prepare the query timestamp and role name by adding new_message. + memories: str = self.memory_service.retrieve_memory(query=new_message.content, + role_name=new_message.role_name, + timestamp=new_message.time_created) + + # format system_message with memories + system_message: Message = self.get_system_message_with_memory(memories=memories) + chat_messages.append(system_message) # Include past conversation history in the message list history_messages = self.memory_service.read_message() if history_messages: - messages.extend(history_messages) + chat_messages.extend(history_messages) # Append the current user's message to the conversation context - messages.append(new_message) - self.logger.info(f"messages={messages}") + chat_messages.append(new_message) + self.logger.info(f"chat_messages={chat_messages}") - result = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) + resp = self.generation_model.call(messages=chat_messages, stream=self.stream, **self.generation_model_kwargs) if self.stream: - assert isinstance(result, ModelResponseGen) + assert isinstance(resp, ModelResponseGen) model_response: ModelResponse | None = None - for model_response in result: + for model_response in resp: yield model_response - if model_response and model_response.message: - self.add_messages(model_response.message) + if remember_response: + if model_response and model_response.message: + model_response.message.role_name = self.assistant_name + self.memory_service.add_messages([new_message, model_response.message]) + else: + self.logger.info("model_response or model_response.message is empty!") else: - assert isinstance(result, ModelResponse) - model_response: ModelResponse = result - if model_response and model_response.message: - self.add_messages(model_response.message) + assert isinstance(resp, ModelResponse) + model_response: ModelResponse = resp + if remember_response: + if model_response and model_response.message: + model_response.message.role_name = self.assistant_name + self.memory_service.add_messages([new_message, model_response.message]) + else: + self.logger.info("model_response or model_response.message is empty!") return model_response diff --git a/memoryscope/chat/base_memory_chat.py b/memoryscope/chat/base_memory_chat.py index d67c66be..dab56a0a 100644 --- a/memoryscope/chat/base_memory_chat.py +++ b/memoryscope/chat/base_memory_chat.py @@ -18,19 +18,27 @@ class BaseMemoryChat(metaclass=ABCMeta): self.logger = Logger.get_logger() @abstractmethod - def chat_with_memory(self, query: str, role_name: str = ""): + def get_new_message(self, query: str, role_name: str = "") -> Message: + raise NotImplementedError + + @abstractmethod + def get_system_message_with_memory(self, memories: str) -> Message: + raise NotImplementedError + + @abstractmethod + def chat_with_memory(self, + query: str, + role_name: str = "", + remember_response: bool = True): """ Initiates a chat interaction using the memory service, with the provided query as input. Args: query (str): The user's query or message to start the chat. role_name (str): The role's name. - - Returns: - This method should return the chat response generated after processing the query - with the associated memory context. The actual return type and content are defined by the implementing - subclass. + remember_response (bool): whether update memory service. """ + raise NotImplementedError @property def memory_service(self) -> BaseMemoryService: @@ -45,6 +53,9 @@ class BaseMemoryChat(metaclass=ABCMeta): def add_messages(self, messages: List[Message] | Message): self.memory_service.add_messages(messages) + def start_backend_service(self): + self.memory_service.start_backend_service() + def do_memory_operation(self, op_name: str, **kwargs): return self.memory_service.do_operation(op_name=op_name, **kwargs) diff --git a/memoryscope/chat/cli_memory_chat.py b/memoryscope/chat/cli_memory_chat.py index a4ec880a..52c01d12 100644 --- a/memoryscope/chat/cli_memory_chat.py +++ b/memoryscope/chat/cli_memory_chat.py @@ -100,7 +100,6 @@ class CliMemoryChat(BaseMemoryChat): self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] # init service & update kwargs self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name) - self._memory_service.start_backend_service() return self._memory_service @property @@ -121,53 +120,71 @@ class CliMemoryChat(BaseMemoryChat): self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model - def chat_with_memory(self, query: str, role_name: str = "") -> ModelResponse | ModelResponseGen: - """ - Engages in a conversation with the AI model, utilizing conversation memory. - The function sends the user's query, incorporates conversation history and memory, - and optionally remembers the AI's response based on the user's preference. - - Args: - query (str): The user's input or query for the AI. - role_name (str, optional): The user's name, default value is human_name. - - Returns: - - ModelResponse: In non-streaming mode, returns a complete AI response. - - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. - - Side Effects: - - Updates the conversation memory with the query of user and (optionally) the response of AI. - - Retrieves and includes historical messages and memory content in the context of conversation. - """ + def get_new_message(self, query: str, role_name: str = "") -> Message: if not role_name: role_name = self.human_name - new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) - self.add_messages(new_message) - - messages: List[Message] = [] + return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) + def get_system_message_with_memory(self, memories: str) -> Message: # Incorporate memory into the system prompt if available system_prompt = self.prompt_handler.system_prompt - memories: str = self.memory_service.retrieve_memory() if memories: memory_prompt = self.prompt_handler.memory_prompt system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) - messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)) + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) + + def chat_with_memory(self, + query: str, + role_name: str = "", + remember_response: bool = True): + + chat_messages: List[Message] = [] + + new_message: Message = self.get_new_message(query=query, role_name=role_name) + + # To retrieve memory, prepare the query timestamp and role name by adding new_message. + memories: str = self.memory_service.retrieve_memory(query=new_message.content, + role_name=new_message.role_name, + timestamp=new_message.time_created) + + # format system_message with memories + system_message: Message = self.get_system_message_with_memory(memories=memories) + chat_messages.append(system_message) # Include past conversation history in the message list history_messages = self.memory_service.read_message() if history_messages: - messages.extend(history_messages) + chat_messages.extend(history_messages) # Append the current user's message to the conversation context - messages.append(new_message) - self.logger.info(f"messages={messages}") + chat_messages.append(new_message) + self.logger.info(f"chat_messages={chat_messages}") # Invoke the Language Model with the constructed message context, respecting streaming setting - return self.generation_model.call(messages=messages, + resp = self.generation_model.call(messages=chat_messages, stream=self.stream, **self.generation_model_kwargs) + if self.stream: + assert isinstance(resp, ModelResponseGen) + model_response: ModelResponse | None = None + for model_response in resp: + questionary.print(model_response.delta, end="") + questionary.print("") + + if remember_response and model_response and model_response.message: + model_response.message.role_name = self.assistant_name + self.memory_service.add_messages([new_message, model_response.message]) + + else: + assert isinstance(resp, ModelResponse) + model_response: ModelResponse = resp + questionary.print(model_response.message.content) + + if remember_response and model_response and model_response.message: + model_response.message.role_name = self.assistant_name + self.memory_service.add_messages([new_message, model_response.message]) + @staticmethod def parse_query_command(query: str): """ @@ -296,19 +313,8 @@ class CliMemoryChat(BaseMemoryChat): questionary.print(f"{self.assistant_name}: ", end="", style="bold") # Fetch and display AI's response - self.memory_service.start_backend_service() - if self.stream: - model_response = None - for model_response in self.chat_with_memory(query=query): - questionary.print(model_response.delta, end="") - questionary.print("") - else: - model_response = self.chat_with_memory(query=query) - questionary.print(model_response.message.content) - - # Append AI's response to the conversation memory - model_response.message.role_name = self.assistant_name - self.add_messages(model_response.message) + self.start_backend_service() + self.chat_with_memory(query=query) except KeyboardInterrupt: # Handle user interruption and confirm exit diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/memory/service/base_memory_service.py index d72fc4c1..c8d964c0 100644 --- a/memoryscope/memory/service/base_memory_service.py +++ b/memoryscope/memory/service/base_memory_service.py @@ -31,25 +31,6 @@ class BaseMemoryService(metaclass=ABCMeta): self._op_description_dict: Dict[str, str] = {} self.logger = Logger.get_logger() - @abstractmethod - def add_messages(self, messages: List[Message] | Message): - raise NotImplementedError - - @abstractmethod - def do_operation(self, op_name: str, **kwargs): - """ - Abstract method defining the interface for executing a specific operation by its name. - This method must be implemented by subclasses to provide the actual operation logic. - - Args: - op_name (str): The name identifying the operation to be performed. - **kwargs: Additional keyword arguments required for the operation execution. - - Raises: - NotImplementedError: This exception is raised when the method is not overridden in a subclass. - """ - raise NotImplementedError - @property def op_description_dict(self) -> Dict[str, str]: """ @@ -63,27 +44,9 @@ class BaseMemoryService(metaclass=ABCMeta): self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()} return self._op_description_dict - def retrieve_memory(self): - """ - Executes the operation associated with retrieved memory. - Asserts that the operation for retrieved memory has been initialized. - - Returns: - Any: The result of the retrieved memory operation. - """ - assert self.retrieve_memory_key in self._operation_dict, f"op={self.retrieve_memory_key} is not inited!" - return self.do_operation(self.retrieve_memory_key) - - def read_message(self): - """ - Executes the operation associated with reading messages. - Asserts that the operation for reading messages has been initialized. - - Returns: - Any: The result of the read message operation. - """ - assert self.read_message_key in self._operation_dict, f"op={self.read_message_key} is not inited!" - return self.do_operation(self.read_message_key) + @abstractmethod + def add_messages(self, messages: List[Message] | Message): + raise NotImplementedError @abstractmethod def init_service(self, **kwargs): @@ -94,3 +57,27 @@ class BaseMemoryService(metaclass=ABCMeta): def stop_backend_service(self): pass + + def do_operation(self, op_name: str, **kwargs): + """ + Executes a specific operation by its name with provided keyword arguments. + + Args: + op_name (str): The name of the operation to execute. + **kwargs: Keyword arguments for the operation's execution. + + Returns: + The result of the operation execution, if any. Otherwise, None. + + Raises: + Warning: If the operation name is not initialized in `_operation_dict`. + """ + 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(**kwargs) + + def __getattr__(self, name: str): + return lambda **kwargs: self.do_operation(name, **kwargs) + + diff --git a/memoryscope/memory/worker/frontend/set_query_worker.py b/memoryscope/memory/worker/frontend/set_query_worker.py index 552baf5e..1f079288 100644 --- a/memoryscope/memory/worker/frontend/set_query_worker.py +++ b/memoryscope/memory/worker/frontend/set_query_worker.py @@ -22,22 +22,39 @@ class SetQueryWorker(MemoryBaseWorker): along with its creation timestamp. """ query = "" # Default query value - query_timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default + timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default if "query" in self.chat_kwargs: - # Check if a specific 'query' has been provided via chat kwargs + # set query if exists query = self.chat_kwargs["query"] if not query: query = "" query = query.strip() + # set ts if exists + _timestamp = self.chat_kwargs.get("timestamp") + if _timestamp and isinstance(_timestamp, int): + timestamp = _timestamp + + # check role_name + role_name = self.chat_kwargs.get("role_name") + if role_name: + assert role_name == self.target_name, \ + f"role_name={role_name} is not supported in human/assistant memory workflow!" + elif self.chat_messages: # If no explicit query is given, use the content of the latest chat message chat_messages = [msg for msg in self.chat_messages if msg.role == MessageRoleEnum.USER.value] if chat_messages: message = chat_messages[-1] query = message.content - query_timestamp = message.time_created + timestamp = message.time_created + + # check role_name + role_name = message.role_name + if role_name: + assert role_name == self.target_name, \ + f"role_name={role_name} is not supported in human/assistant memory workflow!" # Store the determined query and its timestamp in the context - self.set_context(QUERY_WITH_TS, (query, query_timestamp)) + self.set_context(QUERY_WITH_TS, (query, timestamp)) diff --git a/tests/other/test_attr.py b/tests/other/test_attr.py new file mode 100644 index 00000000..eaf2fbf8 --- /dev/null +++ b/tests/other/test_attr.py @@ -0,0 +1,15 @@ +class MyClass: + def __init__(self): + self.existing_attribute = "I exist" + + def do(self, name: str, **kwargs): + print("do %s %s" % (name, kwargs)) + + def __getattr__(self, name): + return lambda **kwargs: self.do(name, **kwargs) + + +# 创建类的实例 +obj = MyClass() + +obj.haha(a=1, b=2) From b1d11a314a7b09c818002170d41d21305b2bd948 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 27 Jul 2024 00:43:47 +0800 Subject: [PATCH 04/15] [dev] remove useless file InitializationHandler --- memoryscope/argument/init_handler.py | 28 -------- memoryscope/argument/memoryscope_arguments.py | 4 +- memoryscope/chat/api_memory_chat.py | 4 +- memoryscope/chat/base_memory_chat.py | 4 +- memoryscope/chat/cli_memory_chat.py | 8 ++- .../memory/operation/backend_operation.py | 11 ++- .../memory/operation/base_operation.py | 2 +- memoryscope/memory/operation/base_workflow.py | 14 ++-- .../memory/service/base_memory_service.py | 26 ++++--- .../memory/service/memory_scope_service.py | 72 ++++++++----------- memoryscope/memoryscope.py | 2 +- memoryscope/memoryscope_context.py | 2 + memoryscope/storage/dummy_memory_store.py | 22 +++--- 13 files changed, 86 insertions(+), 113 deletions(-) delete mode 100644 memoryscope/argument/init_handler.py diff --git a/memoryscope/argument/init_handler.py b/memoryscope/argument/init_handler.py deleted file mode 100644 index 86a2d840..00000000 --- a/memoryscope/argument/init_handler.py +++ /dev/null @@ -1,28 +0,0 @@ -class InitializationHandler(object): - - def __init__(self): - self.file_path: str = __file__ - - self.global_config_dict: dict = {} - - self.memory_chat_dict: dict = {} - - self.memory_service_dict: dict = {} - - self.worker_dict: dict = {} - - self.model_dict: dict = {} - - self.memory_store: dict = {} - - self.monitor: dict = {} - - def update_by_arguments(self): - pass - - def load_from_config(self): - pass - - - def load_from_file(self): - pass diff --git a/memoryscope/argument/memoryscope_arguments.py b/memoryscope/argument/memoryscope_arguments.py index fb10cd23..92572f25 100644 --- a/memoryscope/argument/memoryscope_arguments.py +++ b/memoryscope/argument/memoryscope_arguments.py @@ -55,7 +55,7 @@ class MemoryscopeArguments(object): es_url: str = field(default="http://localhost:9200") - # TODO at xianzhe - retrieve_type: str = field(default="dense", metadata={"help": "es_retrieve_type: dense, sparse, hybrid"}) + retrieve_mode: str = field(default="dense", metadata={ + "help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"}) hybrid_alpha: float | None = field(default=1.0, metadata={"help": ""}) diff --git a/memoryscope/chat/api_memory_chat.py b/memoryscope/chat/api_memory_chat.py index da1d73de..10bc741d 100644 --- a/memoryscope/chat/api_memory_chat.py +++ b/memoryscope/chat/api_memory_chat.py @@ -31,10 +31,12 @@ class ApiMemoryChat(BaseMemoryChat): self.human_name: str = human_name if not self.human_name: self.human_name = DEFAULT_HUMAN_NAME[self.context.language] + self.context.meta_data["human_name"] = self.human_name self.assistant_name: str = assistant_name if not self.assistant_name: self.assistant_name = "AI" + self.context.meta_data["assistant_name"] = self.assistant_name self._prompt_handler: PromptHandler | None = None @@ -73,7 +75,7 @@ class ApiMemoryChat(BaseMemoryChat): self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] # init service & update kwargs - self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name) + self._memory_service.init_service() return self._memory_service @property diff --git a/memoryscope/chat/base_memory_chat.py b/memoryscope/chat/base_memory_chat.py index dab56a0a..ea81c4fb 100644 --- a/memoryscope/chat/base_memory_chat.py +++ b/memoryscope/chat/base_memory_chat.py @@ -56,8 +56,8 @@ class BaseMemoryChat(metaclass=ABCMeta): def start_backend_service(self): self.memory_service.start_backend_service() - def do_memory_operation(self, op_name: str, **kwargs): - return self.memory_service.do_operation(op_name=op_name, **kwargs) + def do_memory_operation(self, operation_name: str, **kwargs): + return self.memory_service.do_operation(name=operation_name, **kwargs) def run(self): """ diff --git a/memoryscope/chat/cli_memory_chat.py b/memoryscope/chat/cli_memory_chat.py index 52c01d12..27a81485 100644 --- a/memoryscope/chat/cli_memory_chat.py +++ b/memoryscope/chat/cli_memory_chat.py @@ -46,10 +46,12 @@ class CliMemoryChat(BaseMemoryChat): self.human_name: str = human_name if not self.human_name: self.human_name = DEFAULT_HUMAN_NAME[self.context.language] + self.context.meta_data["human_name"] = self.human_name self.assistant_name: str = assistant_name if not self.assistant_name: self.assistant_name = "AI" + self.context.meta_data["assistant_name"] = self.assistant_name self._logo = char_logo("MemoryScope") self._prompt_handler: PromptHandler | None = None @@ -99,7 +101,7 @@ class CliMemoryChat(BaseMemoryChat): self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] # init service & update kwargs - self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name) + self._memory_service.init_service() return self._memory_service @property @@ -257,7 +259,7 @@ class CliMemoryChat(BaseMemoryChat): refresh_time = int(refresh_time) self.memory_service.stop_backend_service() while True: - result = self.memory_service.do_operation(op_name=command, **kwargs) + result = self.memory_service.do_operation(name=command, **kwargs) os.system("clear") self.print_logo() if result: @@ -269,7 +271,7 @@ class CliMemoryChat(BaseMemoryChat): time.sleep(refresh_time) else: - result = self.memory_service.do_operation(op_name=command, **kwargs) + result = self.memory_service.do_operation(name=command, **kwargs) if result: if isinstance(result, list): result = "\n".join([str(x) for x in result]) diff --git a/memoryscope/memory/operation/backend_operation.py b/memoryscope/memory/operation/backend_operation.py index c52cf936..bd5d3ac9 100644 --- a/memoryscope/memory/operation/backend_operation.py +++ b/memoryscope/memory/operation/backend_operation.py @@ -115,11 +115,18 @@ class BackendOperation(BaseWorkflow, BaseOperation): if not self._loop_switch: self._loop_switch = True self._backend_task = G_CONTEXT.thread_pool.submit(self._loop_operation) + self.logger.info(f"start operation={operation.name}...") def stop_operation_backend(self, wait_task_end: bool = False): """ Stops the background operation loop by setting the _loop_switch to False. """ self._loop_switch = False - if wait_task_end and self._backend_task: - self._backend_task.result() + if self._backend_task: + if wait_task_end: + self._backend_task.result() + self.logger.info(f"stop operation={self.name}...") + else: + self.logger.info(f"send stop signal to operation={self.name}...") + + diff --git a/memoryscope/memory/operation/base_operation.py b/memoryscope/memory/operation/base_operation.py index 8904530c..df27ea2f 100644 --- a/memoryscope/memory/operation/base_operation.py +++ b/memoryscope/memory/operation/base_operation.py @@ -58,7 +58,7 @@ class BaseOperation(metaclass=ABCMeta): """ pass - def stop_operation_backend(self): + def stop_operation_backend(self, wait_task_end: bool = False): """ Placeholder method to stop any ongoing backend operations. Should be implemented in subclasses where backend operations are managed. diff --git a/memoryscope/memory/operation/base_workflow.py b/memoryscope/memory/operation/base_workflow.py index 926b59a1..1ec0a5f5 100644 --- a/memoryscope/memory/operation/base_workflow.py +++ b/memoryscope/memory/operation/base_workflow.py @@ -6,7 +6,7 @@ from typing import Dict, Any, List from memoryscope.constants.common_constants import WORKFLOW_NAME from memoryscope.memory.worker.base_worker import BaseWorker -from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.utils.logger import Logger from memoryscope.utils.timer import Timer from memoryscope.utils.tool_functions import init_instance_by_config @@ -16,13 +16,13 @@ class BaseWorkflow(object): def __init__(self, name: str, + memoryscope_context: MemoryscopeContext, workflow: str = "", - thread_pool: ThreadPoolExecutor = G_CONTEXT.thread_pool, **kwargs): self.name: str = name + self.memoryscope_context: MemoryscopeContext = memoryscope_context self.workflow: str = workflow - self.thread_pool: ThreadPoolExecutor = thread_pool self.kwargs = kwargs self.workflow_worker_list: List[List[List[str]]] = [] @@ -128,17 +128,17 @@ class BaseWorkflow(object): This method modifies `self.worker_dict` in-place, replacing the keys with actual worker instances. """ for name in list(self.worker_dict.keys()): - if name not in G_CONTEXT.worker_config: - raise RuntimeError(f"worker={name} is not exists in worker_config!") + if name not in self.memoryscope_context.worker_conf_dict: + raise RuntimeError(f"worker={name} is not exists in worker config!") self.worker_dict[name] = init_instance_by_config( - config=G_CONTEXT.worker_config[name], + config=self.memoryscope_context.worker_conf_dict[name], suffix_name="worker", name=name, is_multi_thread=is_backend or self.worker_dict[name], context=self.context, context_lock=self.context_lock, - thread_pool=G_CONTEXT.thread_pool, + thread_pool=self.memoryscope_context.thread_pool, **kwargs) def _run_sub_workflow(self, worker_list: List[str]) -> bool: diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/memory/service/base_memory_service.py index c8d964c0..7b9fdee5 100644 --- a/memoryscope/memory/service/base_memory_service.py +++ b/memoryscope/memory/service/base_memory_service.py @@ -28,26 +28,25 @@ class BaseMemoryService(metaclass=ABCMeta): self.kwargs = kwargs self._operation_dict: Dict[str, BaseOperation] = {} - self._op_description_dict: Dict[str, str] = {} self.logger = Logger.get_logger() @property def op_description_dict(self) -> Dict[str, str]: """ Property to retrieve a dictionary mapping operation keys to their descriptions. - Lazily initializes the dictionary on first access. - Returns: Dict[str, str]: A dictionary where keys are operation identifiers and values are their descriptions. """ - 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 + return {k: v.description for k, v in self._operation_dict.items()} @abstractmethod def add_messages(self, messages: List[Message] | Message): raise NotImplementedError + @abstractmethod + def register_operation(self, name: str, operation_config: dict, **kwargs): + raise NotImplementedError + @abstractmethod def init_service(self, **kwargs): raise NotImplementedError @@ -58,12 +57,12 @@ class BaseMemoryService(metaclass=ABCMeta): def stop_backend_service(self): pass - def do_operation(self, op_name: str, **kwargs): + def do_operation(self, name: str, **kwargs): """ Executes a specific operation by its name with provided keyword arguments. Args: - op_name (str): The name of the operation to execute. + name (str): The name of the operation to execute. **kwargs: Keyword arguments for the operation's execution. Returns: @@ -72,12 +71,11 @@ class BaseMemoryService(metaclass=ABCMeta): Raises: Warning: If the operation name is not initialized in `_operation_dict`. """ - if op_name not in self._operation_dict: - self.logger.warning(f"op_name={op_name} is not inited!") + if name not in self._operation_dict: + self.logger.warning(f"operation={name} is not registered!") return - return self._operation_dict[op_name].run_operation(**kwargs) + return self._operation_dict[name].run_operation(**kwargs) def __getattr__(self, name: str): - return lambda **kwargs: self.do_operation(name, **kwargs) - - + assert name in self._operation_dict, f"operation={name} is not registered!" + return lambda **kwargs: self.do_operation(name=name, **kwargs) diff --git a/memoryscope/memory/service/memory_scope_service.py b/memoryscope/memory/service/memory_scope_service.py index fdcfbc30..eb9d14e8 100644 --- a/memoryscope/memory/service/memory_scope_service.py +++ b/memoryscope/memory/service/memory_scope_service.py @@ -1,3 +1,4 @@ +import threading from typing import List from memoryscope.memory.operation.base_operation import BaseOperation @@ -11,6 +12,8 @@ class MemoryScopeService(BaseMemoryService): history_msg_count: int = 100, contextual_msg_max_count: int = 20, contextual_msg_min_count: int = 0, + human_name: str = None, + assistant_name: str = None, **kwargs): """ init function. @@ -20,13 +23,19 @@ class MemoryScopeService(BaseMemoryService): it will not be included in the context to prevent token overflow. contextual_msg_min_count (int): The minimum context length in a conversation. If it is shorter than this length, no conversation summary will be made and no long-term memory will be generated. - kwargs (dict): other kwargs + human_name (str): human name. + assistant_name (str): assistant name. + kwargs (dict): other kwargs. """ super().__init__(**kwargs) self.history_msg_count: int = history_msg_count self.contextual_msg_max_count: int = contextual_msg_max_count self.contextual_msg_min_count: int = contextual_msg_min_count assert history_msg_count >= contextual_msg_max_count >= contextual_msg_min_count + if human_name: + self.context.meta_data["human_name"] = human_name + if assistant_name: + self.context.meta_data["assistant_name"] = assistant_name self.chat_messages: List[Message] = [] self.message_lock = threading.Lock() @@ -57,43 +66,28 @@ class MemoryScopeService(BaseMemoryService): for _ in range(gap_size): self.chat_messages.pop(0) - def do_operation(self, op_name: str, **kwargs): - """ - Executes a specific operation by its name with provided keyword arguments. - - Args: - op_name (str): The name of the operation to execute. - **kwargs: Keyword arguments for the operation's execution. - - Returns: - The result of the operation execution, if any. Otherwise, None. - - Raises: - Warning: If the operation name is not initialized in `_operation_dict`. - """ - if op_name not in self._operation_dict: - self.logger.warning(f"op_name={op_name} is not inited!") # Warn if operation not initialized + def register_operation(self, name: str, operation_config: dict, **kwargs): + if name in self._operation_dict: + self.logger.warning(f"op_name={name} is registered before!") return - return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation + + operation: BaseOperation = init_instance_by_config( + config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + context=self.context, + contextual_msg_max_count=self.contextual_msg_max_count, + contextual_msg_min_count=self.contextual_msg_min_count) + + # Initialize workflow for each operation + operation.init_workflow(**kwargs) + self._operation_dict[name] = operation + self.logger.info(f"service={self.__class__.__name__} init operation={name}") def init_service(self, **kwargs): - for name, operation_config in self.memory_operations.items(): - if name in self._operation_dict: - self.logger.warning(f"memory operation={name} is repeated!") - continue - - # ⭐ Initialize operation instance by its config - operation: BaseOperation = init_instance_by_config( - config=operation_config, - name=name, - chat_messages=self.chat_messages, - message_lock=self.message_lock, - contextual_msg_max_count=self.contextual_msg_max_count, - contextual_msg_min_count=self.contextual_msg_min_count) - operation.init_workflow(**kwargs) # Initialize workflow for each operation - - self._operation_dict[name] = operation - self.logger.info(f"service={self.__class__.__name__} init operation={name}") + for name, operation_config in self.memory_operations_conf.items(): + self.register_operation(name, operation_config, **kwargs) def start_backend_service(self): """ @@ -101,16 +95,12 @@ class MemoryScopeService(BaseMemoryService): """ for _, operation in self._operation_dict.items(): if operation.operation_type == "backend": - # Run backend operations operation.run_operation_backend() - self.logger.info(f"start operation={operation.name}...") - def stop_backend_service(self): + def stop_backend_service(self, wait_service_end: bool = False): """ Stops all backend operations that are currently running. """ for _, operation in self._operation_dict.items(): if operation.operation_type == "backend": - # Stop backend operations - operation.stop_operation_backend() - self.logger.info(f"stop operation={operation.name}...") + operation.stop_operation_backend(wait_task_end=wait_service_end) diff --git a/memoryscope/memoryscope.py b/memoryscope/memoryscope.py index 520bd127..0648871c 100644 --- a/memoryscope/memoryscope.py +++ b/memoryscope/memoryscope.py @@ -111,7 +111,7 @@ class MemoryScope(object): "embedding_model": "embedding_model", "index_name": arguments.es_index_name, "es_url": arguments.es_url, - "retrieve_type": arguments.retrieve_type, + "retrieve_mode": arguments.retrieve_mode, "hybrid_alpha": arguments.hybrid_alpha, } diff --git a/memoryscope/memoryscope_context.py b/memoryscope/memoryscope_context.py index 967bb77d..ac8d89fb 100644 --- a/memoryscope/memoryscope_context.py +++ b/memoryscope/memoryscope_context.py @@ -25,3 +25,5 @@ class MemoryscopeContext(object): model_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> model"}) worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"}) + + meta_data: dict = field(default_factory=lambda: {}) diff --git a/memoryscope/storage/dummy_memory_store.py b/memoryscope/storage/dummy_memory_store.py index 304818bb..7c77f670 100644 --- a/memoryscope/storage/dummy_memory_store.py +++ b/memoryscope/storage/dummy_memory_store.py @@ -12,6 +12,17 @@ class DummyMemoryStore(BaseMemoryStore): semantic retrieval. Actual storage operations are not implemented. """ + def __init__(self, embedding_model: BaseModel, **kwargs): + """ + Initializes the DummyMemoryStore with an embedding model and additional keyword arguments. + + Args: + embedding_model (BaseModel): The model used to embed data for potential similarity-based retrieval. + **kwargs: Additional keyword arguments for configuration or future expansion. + """ + self.embedding_model: BaseModel = embedding_model + self.kwargs = kwargs + def retrieve_memories(self, query: str = "", top_k: int = 3, @@ -24,17 +35,6 @@ class DummyMemoryStore(BaseMemoryStore): filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]: pass - def __init__(self, embedding_model: BaseModel, **kwargs): - """ - Initializes the DummyMemoryStore with an embedding model and additional keyword arguments. - - Args: - embedding_model (BaseModel): The model used to embed data for potential similarity-based retrieval. - **kwargs: Additional keyword arguments for configuration or future expansion. - """ - self.embedding_model: BaseModel = embedding_model - self.kwargs = kwargs - def batch_insert(self, nodes: List[MemoryNode]): pass From 9248c6708594dd5d52a6f90be0ee5907937921fb Mon Sep 17 00:00:00 2001 From: "xianzhe.xxz" Date: Fri, 26 Jul 2024 17:16:49 +0800 Subject: [PATCH 05/15] fix dashscope rerank: add alpha support, when alpha=1.0, only use vector similarity, when alpha=0.0, only use bm25, when alpha=None, use rrf to fuse the results --- memoryscope/models/llama_index_rank_model.py | 7 +++-- .../storage/llama_index_es_memory_store.py | 5 +++- .../storage/llama_index_sync_elasticsearch.py | 26 ++++++++++++++++--- 3 files changed, 31 insertions(+), 7 deletions(-) diff --git a/memoryscope/models/llama_index_rank_model.py b/memoryscope/models/llama_index_rank_model.py index 9e74cdbd..38debfd7 100644 --- a/memoryscope/models/llama_index_rank_model.py +++ b/memoryscope/models/llama_index_rank_model.py @@ -35,12 +35,13 @@ class LlamaIndexRankModel(BaseModel): documents = [documents] assert query and documents and all(documents), \ f"query or documents is empty! query={query}, documents={len(documents)}" - + assert len(documents) < 500, \ + f"The input documents of Dashscope rerank model should not larger than 500!" # Using -1.0 as dummy scores nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents] model_response.meta_data.update({ - "data": {"nodes": nodes, "query_str": query}, + "data": {"nodes": nodes, "query_str": query, "top_n": len(documents)}, "documents_map": {doc: idx for idx, doc in enumerate(documents)}, }) @@ -76,6 +77,8 @@ class LlamaIndexRankModel(BaseModel): Returns: ModelResponse: A response object encapsulating the ranked nodes. """ + self.model.top_n = model_response.meta_data["data"]["top_n"] + model_response.meta_data["data"].pop("top_n") model_response.raw = self.model.postprocess_nodes(**model_response.meta_data["data"]) async def _async_call(self, **kwargs) -> ModelResponse: diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index b7ce3b21..4a9d7f4f 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -26,7 +26,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.embedding_model: BaseModel = embedding_model self.es_store = SyncElasticsearchStore(index_name=index_name, es_url=es_url, - retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid), + retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid, + alpha=0.5), # weights of vector similarity, + # while the weights of BM25 is 1-alpha. + # when alpha=None, then rrf fusion is uesd. **kwargs) # TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 2c525920..507044a3 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -132,6 +132,21 @@ def _mode_must_match_retrieval_strategy( class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy): + + + def __init__( + self, + *, + distance: DistanceMetric = DistanceMetric.COSINE, + model_id: Optional[str] = None, + hybrid: bool = False, + rrf: Union[bool, Dict[str, Any]] = True, + text_field: Optional[str] = "text_field", + alpha: Optional[float] = None, + ): + super().__init__(distance=distance, model_id=model_id, hybrid=hybrid, rrf=rrf, text_field=text_field) + self.alpha = alpha + def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]: # Add a query to the knn query. # RRF is used to even the score from the knn query and text query @@ -155,18 +170,19 @@ class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy): "match": { self.text_field: { "query": query, + "boost": (1 - self.alpha) if self.alpha is not None else 1.0, } - } + }, } ], "filter": filter, - } + }, }, } - if isinstance(self.rrf, Dict): + if self.alpha is None and isinstance(self.rrf, Dict): query_body["rank"] = {"rrf": self.rrf} - elif isinstance(self.rrf, bool) and self.rrf is True: + elif self.alpha is None and isinstance(self.rrf, bool) and self.rrf is True: query_body["rank"] = {"rrf": {"window_size": top_k}} return query_body @@ -190,6 +206,7 @@ class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy): "field": vector_field, "k": k, "num_candidates": num_candidates, + "boost": self.alpha if self.alpha is not None else 1.0, } if query_vector is not None: @@ -676,6 +693,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): ): total_rank = sum(top_k_scores) top_k_scores = [rank for rank in top_k_scores] + print("top_k_scores:", top_k_scores) # top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores] # top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores] From 8295b019610090fbf1adb37da89da2d5d8c6a2d0 Mon Sep 17 00:00:00 2001 From: "xianzhe.xxz" Date: Fri, 26 Jul 2024 19:58:30 +0800 Subject: [PATCH 06/15] adjust the params of EsStore, set dense retrieve as default, --- .../storage/llama_index_es_memory_store.py | 25 ++++++---- .../storage/llama_index_sync_elasticsearch.py | 21 ++++++--- tests/storages/test_storages_lli_synces.py | 46 ++++++------------- 3 files changed, 45 insertions(+), 47 deletions(-) diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index 4a9d7f4f..af88f91c 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -7,7 +7,7 @@ from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle from memoryscope.models.base_model import BaseModel from memoryscope.scheme.memory_node import MemoryNode from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, _AsyncDenseVectorStrategy, \ +from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, ESCombinedRetrieveStrategy, \ _to_elasticsearch_filter from memoryscope.utils.logger import Logger @@ -18,18 +18,16 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): embedding_model: BaseModel, index_name: str, es_url: str, - use_hybrid: bool = True, - emb_dims: int = 1536, + retrieve_mode: str = "dense", + hybrid_alpha: float = None, **kwargs): + self.emb_dims = None self.index_name = index_name - self.emb_dims = emb_dims self.embedding_model: BaseModel = embedding_model self.es_store = SyncElasticsearchStore(index_name=index_name, es_url=es_url, - retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid, - alpha=0.5), # weights of vector similarity, - # while the weights of BM25 is 1-alpha. - # when alpha=None, then rrf fusion is uesd. + retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, + hybrid_alpha=hybrid_alpha), **kwargs) # TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By @@ -40,7 +38,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.logger = Logger.get_logger() def retrieve_memories(self, - query: str = "", + query: str = "**--**", top_k: int = 3, filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: # if index is not created, return [] @@ -55,10 +53,13 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter, "fields": ['embedding']}, similarity_top_k=top_k, sparse_top_k=top_k) - if not query: + if not query and self.emb_dims: query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) text_nodes = retriever.retrieve(query) + if text_nodes and text_nodes[0].embedding: + self.emb_dims = len(text_nodes[0].embedding) + return [self._text_node_2_memory_node(n) for n in text_nodes] async def a_retrieve_memories(self, @@ -82,6 +83,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) text_nodes: List[NodeWithScore] = await retriever.aretrieve(query) + + if text_nodes and text_nodes[0].embedding: + self.emb_dims = len(text_nodes[0].embedding) + return [self._text_node_2_memory_node(n) for n in text_nodes] def batch_insert(self, nodes: List[MemoryNode]): diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 507044a3..5c74b0ba 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -131,7 +131,7 @@ def _mode_must_match_retrieval_strategy( raise ValueError(f"to enable hybrid mode, it must be set in retrieval strategy") -class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy): +class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): def __init__( @@ -139,13 +139,21 @@ class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy): *, distance: DistanceMetric = DistanceMetric.COSINE, model_id: Optional[str] = None, - hybrid: bool = False, + retrieve_mode: str = "dense", rrf: Union[bool, Dict[str, Any]] = True, text_field: Optional[str] = "text_field", - alpha: Optional[float] = None, - ): - super().__init__(distance=distance, model_id=model_id, hybrid=hybrid, rrf=rrf, text_field=text_field) - self.alpha = alpha + hybrid_alpha: Optional[float] = None, + ): + if retrieve_mode == "dense": + self.alpha = 1.0 + elif retrieve_mode == "sparse": + # self.alpha = 0.0 + raise NotImplementedError + elif retrieve_mode == "hybrid": + # self.alpha = hybrid_alpha + raise NotImplementedError + + super().__init__(distance=distance, model_id=model_id, hybrid=True, rrf=rrf, text_field=text_field) def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]: # Add a query to the knn query. @@ -637,6 +645,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): else: filter = es_filter or [] num_candidates = query.similarity_top_k * 10 if query.similarity_top_k <= 1000 else query.similarity_top_k + hits = self._store.search( query=query.query_str, query_vector=query.query_embedding, diff --git a/tests/storages/test_storages_lli_synces.py b/tests/storages/test_storages_lli_synces.py index 8186179a..0bb8b554 100644 --- a/tests/storages/test_storages_lli_synces.py +++ b/tests/storages/test_storages_lli_synces.py @@ -20,7 +20,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): "index_name": "0708_8", "es_url": "http://localhost:9200", "embedding_model": emb, - "use_hybrid": True + "retrieve_mode": "dense", } self.es_store = LlamaIndexEsMemoryStore(**config) @@ -152,13 +152,6 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): ), ] - def test_retrieve(self): - filter_dict = { - "timestamp": 12, - # "memory_id": "bbb456", - # "score_rank": 0, - } - for node in self.data: self.es_store.insert(node) @@ -171,36 +164,27 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): meta_data={"5": "5"}, timestamp=13 )) + + def test_retrieve(self): + filter_dict = { + "timestamp": 12, + # "memory_id": "bbb456", + # "score_rank": 0, + } + + res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15) print(len(res)) print(res) - self.es_store.update(MemoryNode( - content="test update", - memory_type="profile", - user_id="6", - status="invalid", - memory_id="ggg567", - timestamp=13, - - )) - res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15) + def test_retrieve_wo_query(self,): + filter_dict = { + "memory_id": "bbb456", + } + res = self.es_store.retrieve_memories(filter_dict=filter_dict, top_k=15) print(len(res)) print(res) - self.es_store.delete(MemoryNode( - content="test update", - memory_type="profile", - user_id="6", - status="invalid", - memory_id="ggg567", - timestamp=13, - )) - import asyncio - res = asyncio.run(self.es_store.a_retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)) - # res = self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10) - print(len(res)) - print(res) def tearDown(self): self.es_store.close() From 6aaa26742ed9e4309f517e8f3e04eb61f70d7f39 Mon Sep 17 00:00:00 2001 From: "xianzhe.xxz" Date: Fri, 26 Jul 2024 20:07:01 +0800 Subject: [PATCH 07/15] add score_recall manager --- memoryscope/storage/llama_index_es_memory_store.py | 3 ++- memoryscope/storage/llama_index_sync_elasticsearch.py | 1 - 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index af88f91c..96b1542f 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -144,7 +144,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): return TextNode(id_=memory_node.memory_id, text=memory_node.content, embedding=embedding, - metadata=memory_node.model_dump(exclude={"content", "vector"})) + metadata=memory_node.model_dump(exclude={"content", "vector", "score_recall", "score_rank", "score_rerank"})) @staticmethod def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode: @@ -158,4 +158,5 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): MemoryNode: The converted MemoryNode with text and metadata from the NodeWithScore. """ text_node.metadata["vector"] = text_node.embedding if text_node.embedding else [] + text_node.metadata["score_recall"] = text_node.score return MemoryNode(content=text_node.text, **text_node.metadata) diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 5c74b0ba..0794b284 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -702,7 +702,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore): ): total_rank = sum(top_k_scores) top_k_scores = [rank for rank in top_k_scores] - print("top_k_scores:", top_k_scores) # top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores] # top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores] From 590cbf65c4fc4f9c100e1345e21f655bf361489f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 27 Jul 2024 01:49:20 +0800 Subject: [PATCH 08/15] [dev] rename MEMORY_HANDLER to MEMORY_MANAGER --- memoryscope/constants/common_constants.py | 4 +- .../memory/operation/backend_operation.py | 22 +++----- memoryscope/memory/operation/base_workflow.py | 21 +++++--- ...rvation_op.py => consolidate_operation.py} | 19 +++---- .../memory/operation/frontend_operation.py | 14 +++-- .../memory/service/memory_scope_service.py | 2 +- .../worker/backend/contra_repeat_worker.py | 6 +-- .../worker/backend/get_observation_worker.py | 8 +-- .../backend/get_reflection_subject_worker.py | 8 +-- .../worker/backend/load_memory_worker.py | 8 +-- .../backend/long_contra_repeat_worker.py | 4 +- .../worker/backend/update_insight_worker.py | 10 ++-- .../worker/backend/update_memory_worker.py | 10 ++-- .../worker/frontend/fuse_rerank_worker.py | 2 +- .../worker/frontend/print_memory_worker.py | 2 +- .../worker/frontend/retrieve_memory_worker.py | 2 +- .../worker/frontend/semantic_rank_worker.py | 4 +- .../memory/worker/memory_base_worker.py | 51 +++++++++++-------- .../worker/memory_manager.py} | 13 +++-- memoryscope/models/dummy_generation_model.py | 7 +-- memoryscope/utils/datetime_handler.py | 49 ++++++++++-------- tests/worker/test_workers_cn.py | 36 ++++++------- tests/worker/test_workers_en.py | 36 ++++++------- 23 files changed, 177 insertions(+), 161 deletions(-) rename memoryscope/memory/operation/{summary_observation_op.py => consolidate_operation.py} (83%) rename memoryscope/{utils/memory_handler.py => memory/worker/memory_manager.py} (96%) diff --git a/memoryscope/constants/common_constants.py b/memoryscope/constants/common_constants.py index adf9e548..4eba1d1d 100644 --- a/memoryscope/constants/common_constants.py +++ b/memoryscope/constants/common_constants.py @@ -5,11 +5,13 @@ WORKFLOW_NAME = "workflow_name" +MEMORYSCOPE_CONTEXT = "memoryscope_context" + RESULT = "result" CHAT_MESSAGES = "chat_messages" -MEMORY_HANDLER = "memory_handler" +MEMORY_MANAGER = "memory_manager" CHAT_KWARGS = "chat_kwargs" diff --git a/memoryscope/memory/operation/backend_operation.py b/memoryscope/memory/operation/backend_operation.py index bd5d3ac9..4a53a6b4 100644 --- a/memoryscope/memory/operation/backend_operation.py +++ b/memoryscope/memory/operation/backend_operation.py @@ -1,7 +1,6 @@ import time from typing import List - from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memoryscope.memory.operation.base_workflow import BaseWorkflow @@ -54,17 +53,14 @@ class BackendOperation(BaseWorkflow, BaseOperation): Returns: Any: The result obtained after executing the workflow. """ - self.context.clear() - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs - - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = self.chat_messages + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: self.chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution return self.context.get(RESULT) @@ -114,8 +110,8 @@ class BackendOperation(BaseWorkflow, BaseOperation): """ if not self._loop_switch: self._loop_switch = True - self._backend_task = G_CONTEXT.thread_pool.submit(self._loop_operation) - self.logger.info(f"start operation={operation.name}...") + self._backend_task = self.thread_pool.submit(self._loop_operation) + self.logger.info(f"start operation={self.name}...") def stop_operation_backend(self, wait_task_end: bool = False): """ @@ -128,5 +124,3 @@ class BackendOperation(BaseWorkflow, BaseOperation): self.logger.info(f"stop operation={self.name}...") else: self.logger.info(f"send stop signal to operation={self.name}...") - - diff --git a/memoryscope/memory/operation/base_workflow.py b/memoryscope/memory/operation/base_workflow.py index 1ec0a5f5..7cb75daf 100644 --- a/memoryscope/memory/operation/base_workflow.py +++ b/memoryscope/memory/operation/base_workflow.py @@ -4,7 +4,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from itertools import zip_longest from typing import Dict, Any, List -from memoryscope.constants.common_constants import WORKFLOW_NAME +from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT from memoryscope.memory.worker.base_worker import BaseWorker from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.utils.logger import Logger @@ -22,6 +22,7 @@ class BaseWorkflow(object): self.name: str = name self.memoryscope_context: MemoryscopeContext = memoryscope_context + self.thread_pool: ThreadPoolExecutor = self.memoryscope_context.thread_pool self.workflow: str = workflow self.kwargs = kwargs @@ -133,12 +134,11 @@ class BaseWorkflow(object): self.worker_dict[name] = init_instance_by_config( config=self.memoryscope_context.worker_conf_dict[name], - suffix_name="worker", name=name, is_multi_thread=is_backend or self.worker_dict[name], context=self.context, context_lock=self.context_lock, - thread_pool=self.memoryscope_context.thread_pool, + thread_pool=self.thread_pool, **kwargs) def _run_sub_workflow(self, worker_list: List[str]) -> bool: @@ -150,7 +150,7 @@ class BaseWorkflow(object): return False return True - def run_workflow(self): + def run_workflow(self, **kwargs): """ Executes the workflow by orchestrating the steps defined in `self.workflow_worker_list`. This method supports both sequential and parallel execution of sub-workflows based on the structure @@ -159,9 +159,18 @@ class BaseWorkflow(object): If a workflow part consists of a single item, it is executed sequentially. For parts with multiple items, they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow returns False. + + Args: + **kwargs: Additional keyword arguments to be passed to context. """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): - self.context[WORKFLOW_NAME] = self.name + self.context.clear() + + self.context.update({ + WORKFLOW_NAME: self.name, + MEMORYSCOPE_CONTEXT: self.memoryscope_context, + **kwargs, + }) # Iterate over each part of the workflow for workflow_part in self.workflow_worker_list: @@ -174,7 +183,7 @@ class BaseWorkflow(object): t_list = [] # Submit tasks to the thread pool for sub_workflow in workflow_part: - t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow)) + t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow)) # Check results; if any task returns False, stop the workflow flag = True diff --git a/memoryscope/memory/operation/summary_observation_op.py b/memoryscope/memory/operation/consolidate_operation.py similarity index 83% rename from memoryscope/memory/operation/summary_observation_op.py rename to memoryscope/memory/operation/consolidate_operation.py index 1da2e65c..967c3639 100644 --- a/memoryscope/memory/operation/summary_observation_op.py +++ b/memoryscope/memory/operation/consolidate_operation.py @@ -3,10 +3,10 @@ from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.memory.operation.backend_operation import BackendOperation -class SummaryObservationOp(BackendOperation): +class ConsolidateOperation(BackendOperation): def __init__(self, **kwargs): - super(SummaryObservationOp, self).__init__(**kwargs) + super(ConsolidateOperation, self).__init__(**kwargs) self.message_lock = kwargs.get("message_lock", None) self.contextual_msg_min_count: int = kwargs.get("contextual_msg_min_count", 0) @@ -43,17 +43,14 @@ class SummaryObservationOp(BackendOperation): f"contextual_msg_min_count({self.contextual_msg_min_count}), skip.") return - self.context.clear() - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs - - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = chat_messages + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution result = self.context.get(RESULT) diff --git a/memoryscope/memory/operation/frontend_operation.py b/memoryscope/memory/operation/frontend_operation.py index 8cde7977..abaca44a 100644 --- a/memoryscope/memory/operation/frontend_operation.py +++ b/memoryscope/memory/operation/frontend_operation.py @@ -39,17 +39,15 @@ class FrontendOperation(BaseWorkflow, BaseOperation): Returns: Any: The result obtained from executing the workflow. """ - self.context.clear() - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = self.chat_messages - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: self.chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution return self.context.get(RESULT) diff --git a/memoryscope/memory/service/memory_scope_service.py b/memoryscope/memory/service/memory_scope_service.py index eb9d14e8..bf1c9b94 100644 --- a/memoryscope/memory/service/memory_scope_service.py +++ b/memoryscope/memory/service/memory_scope_service.py @@ -76,7 +76,7 @@ class MemoryScopeService(BaseMemoryService): name=name, chat_messages=self.chat_messages, message_lock=self.message_lock, - context=self.context, + memoryscope_context=self.context, contextual_msg_max_count=self.contextual_msg_max_count, contextual_msg_min_count=self.contextual_msg_min_count) diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.py b/memoryscope/memory/worker/backend/contra_repeat_worker.py index 5146887f..85a18884 100644 --- a/memoryscope/memory/worker/backend/contra_repeat_worker.py +++ b/memoryscope/memory/worker/backend/contra_repeat_worker.py @@ -43,13 +43,13 @@ class ContraRepeatWorker(MemoryBaseWorker): 6. Updates the status of nodes accordingly. 7. Persists the changes back to memory storage. """ - all_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES]) + all_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES]) if not all_obs_nodes: self.logger.info("all_obs_nodes is empty!") # self.continue_run = False return - today_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories(TODAY_NODES) + today_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories(TODAY_NODES) if today_obs_nodes: all_obs_nodes.extend(today_obs_nodes) @@ -121,4 +121,4 @@ class ContraRepeatWorker(MemoryBaseWorker): merge_obs_nodes.append(node) # save context - self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False) + self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False) diff --git a/memoryscope/memory/worker/backend/get_observation_worker.py b/memoryscope/memory/worker/backend/get_observation_worker.py index b36d2b6f..385251d1 100644 --- a/memoryscope/memory/worker/backend/get_observation_worker.py +++ b/memoryscope/memory/worker/backend/get_observation_worker.py @@ -42,11 +42,11 @@ class GetObservationWorker(MemoryBaseWorker): MemoryTypeEnum.CONVERSATION.value: message.content, TIME_INFER: time_infer, "keywords": keywords, - **{k: str(v) for k, v in dt_handler.dt_info_dict.items()}, + **{k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}, } if time_infer: - dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer) + dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer, language=self.language) meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()}) obs_content = (f"{obs_content} ({self.get_language_value(TIME_INFER_WORD)}" f"{self.get_language_value(COLON_WORD)} {time_infer})") @@ -68,7 +68,7 @@ class GetObservationWorker(MemoryBaseWorker): """ filter_messages = [] for msg in self.chat_messages: - if not DatetimeHandler.has_time_word(query=msg.content): + if not DatetimeHandler.has_time_word(query=msg.content, language=self.language): filter_messages.append(msg) self.logger.info(f"after filter_messages.size from {len(self.chat_messages)} to {len(filter_messages)}") @@ -184,4 +184,4 @@ class GetObservationWorker(MemoryBaseWorker): keywords=keywords)) # Stores the extracted and structured observations in the conversation memory - self.memory_handler.set_memories(self.OBS_STORE_KEY, new_obs_nodes) + self.memory_manager.set_memories(self.OBS_STORE_KEY, new_obs_nodes) diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py b/memoryscope/memory/worker/backend/get_reflection_subject_worker.py index 4c9d7b27..711206ae 100644 --- a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py +++ b/memoryscope/memory/worker/backend/get_reflection_subject_worker.py @@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): """ dt_handler = DatetimeHandler() # Prepare metadata with current datetime info - meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()} + meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict.items()} return MemoryNode(user_name=self.user_name, target_name=self.target_name, @@ -58,8 +58,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): - Parsing the model's responses for new insight keys. - Creating new insight nodes and updating the memory status accordingly. """ - not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_REFLECTED_NODES) - insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES) + not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_REFLECTED_NODES) + insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES) # Count unaudited nodes not_reflected_count = len(not_reflected_nodes) @@ -104,7 +104,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2() if new_insight_keys: for insight_key in new_insight_keys: - self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) + self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) # Mark unaudited nodes as reflected for node in not_reflected_nodes: diff --git a/memoryscope/memory/worker/backend/load_memory_worker.py b/memoryscope/memory/worker/backend/load_memory_worker.py index 1f5e4233..22c77792 100644 --- a/memoryscope/memory/worker/backend/load_memory_worker.py +++ b/memoryscope/memory/worker/backend/load_memory_worker.py @@ -33,7 +33,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_reflected_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) + self.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) @timer def retrieve_not_updated_memory(self): @@ -52,7 +52,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_updated_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + self.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) @timer def retrieve_insight_memory(self): @@ -70,7 +70,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_insight_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(INSIGHT_NODES, nodes) + self.memory_manager.set_memories(INSIGHT_NODES, nodes) @timer def retrieve_today_memory(self, dt: str): @@ -93,7 +93,7 @@ class LoadMemoryWorker(MemoryBaseWorker): nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_today_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(TODAY_NODES, nodes) + self.memory_manager.set_memories(TODAY_NODES, nodes) def _run(self): """ diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py b/memoryscope/memory/worker/backend/long_contra_repeat_worker.py index 94b306b4..056c63d7 100644 --- a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/memory/worker/backend/long_contra_repeat_worker.py @@ -63,7 +63,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): The process helps in maintaining conversation coherence by resolving contradictions and redundancies. """ - not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES) + not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES) for node in not_updated_nodes: self.submit_thread_task(fn=self.retrieve_similar_content, node=node) @@ -157,4 +157,4 @@ class LongContraRepeatWorker(MemoryBaseWorker): f"action_status={node.action_status}") # save context - self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes) + self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memoryscope/memory/worker/backend/update_insight_worker.py b/memoryscope/memory/worker/backend/update_insight_worker.py index b56ed6f7..339d7871 100644 --- a/memoryscope/memory/worker/backend/update_insight_worker.py +++ b/memoryscope/memory/worker/backend/update_insight_worker.py @@ -95,7 +95,7 @@ class UpdateInsightWorker(MemoryBaseWorker): content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}" insight_node.content = content insight_node.value = insight_value - insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()}) + insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict.items()}) insight_node.timestamp = dt_handler.timestamp insight_node.dt = dt_handler.datetime_format() if insight_node.action_status == ActionStatusEnum.NONE.value: @@ -175,9 +175,9 @@ class UpdateInsightWorker(MemoryBaseWorker): 6. Gather the results of all update tasks. 7. Mark processed nodes as updated in memory. """ - insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES) - not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES) - not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(keys=[NOT_REFLECTED_NODES, + insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES) + not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES) + not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(keys=[NOT_REFLECTED_NODES, NOT_UPDATED_NODES]) if not insight_nodes: @@ -216,7 +216,7 @@ class UpdateInsightWorker(MemoryBaseWorker): # delete empty nodes empty_nodes = [n for n in insight_nodes if not n.content.strip()] - self.memory_handler.delete_memories(empty_nodes) + self.memory_manager.delete_memories(empty_nodes) for node in not_updated_nodes: node.obs_updated = 1 diff --git a/memoryscope/memory/worker/backend/update_memory_worker.py b/memoryscope/memory/worker/backend/update_memory_worker.py index 87b182a7..9133de68 100644 --- a/memoryscope/memory/worker/backend/update_memory_worker.py +++ b/memoryscope/memory/worker/backend/update_memory_worker.py @@ -46,7 +46,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): if not self.memory_key: return - return self.memory_handler.get_memories(keys=self.memory_key) + return self.memory_manager.get_memories(keys=self.memory_key) def delete_all(self): """ @@ -55,7 +55,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): Returns: List[MemoryNode]: A list of all MemoryNode objects marked for deletion. """ - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: node.action_status = ActionStatusEnum.DELETE.value self.logger.info(f"delete_all.size={len(nodes)}") @@ -74,7 +74,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): return i = 0 - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: if node.content == query: i += 1 @@ -88,7 +88,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): return i = 0 - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: if node.memory_id == memory_id: i += 1 @@ -109,4 +109,4 @@ class UpdateMemoryWorker(MemoryBaseWorker): if not hasattr(self, method): self.logger.info(f"method={method} is missing!") return - self.memory_handler.update_memories(nodes=getattr(self, method)()) + self.memory_manager.update_memories(nodes=getattr(self, method)()) diff --git a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py b/memoryscope/memory/worker/frontend/fuse_rerank_worker.py index 373d82b0..3d394790 100644 --- a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py +++ b/memoryscope/memory/worker/frontend/fuse_rerank_worker.py @@ -62,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker): """ # Parse input parameters from the worker's context extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT) - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RANKED_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RANKED_MEMORY_NODES) # Check if memory nodes are available; warn and return if not if not memory_node_list: diff --git a/memoryscope/memory/worker/frontend/print_memory_worker.py b/memoryscope/memory/worker/frontend/print_memory_worker.py index 6617c49c..e50adad9 100644 --- a/memoryscope/memory/worker/frontend/print_memory_worker.py +++ b/memoryscope/memory/worker/frontend/print_memory_worker.py @@ -22,7 +22,7 @@ class PrintMemoryWorker(MemoryBaseWorker): 3. Set the formatted string back into the worker's context """ # get long-term memory - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES) memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True) observation_memory_list: List[str] = [] diff --git a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py b/memoryscope/memory/worker/frontend/retrieve_memory_worker.py index fd928f7d..e51bbd2d 100644 --- a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memoryscope/memory/worker/frontend/retrieve_memory_worker.py @@ -139,4 +139,4 @@ class RetrieveMemoryWorker(MemoryBaseWorker): self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} type={node.memory_type} " f"store_status={node.store_status} action_status={node.action_status}") - self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) + self.memory_manager.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memoryscope/memory/worker/frontend/semantic_rank_worker.py b/memoryscope/memory/worker/frontend/semantic_rank_worker.py index be5d67dc..9a78b96c 100644 --- a/memoryscope/memory/worker/frontend/semantic_rank_worker.py +++ b/memoryscope/memory/worker/frontend/semantic_rank_worker.py @@ -29,7 +29,7 @@ class SemanticRankWorker(MemoryBaseWorker): """ # query query, _ = self.get_context(QUERY_WITH_TS) - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES) if not memory_node_list: self.logger.warning("Retrieve memory nodes is empty!") return @@ -58,4 +58,4 @@ class SemanticRankWorker(MemoryBaseWorker): self.logger.info(f"Rank stage: Content={node.content}, Score={node.score_rank}") # save ranked nodes back to memory - self.memory_handler.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False) + self.memory_manager.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False) diff --git a/memoryscope/memory/worker/memory_base_worker.py b/memoryscope/memory/worker/memory_base_worker.py index 00fb72be..e800f8b3 100644 --- a/memoryscope/memory/worker/memory_base_worker.py +++ b/memoryscope/memory/worker/memory_base_worker.py @@ -1,14 +1,16 @@ from abc import ABCMeta from typing import List, Dict, Any -from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORY_HANDLER +from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \ + WORKFLOW_NAME, MEMORY_MANAGER +from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.memory.worker.base_worker import BaseWorker +from memoryscope.memory.worker.memory_manager import MemoryManager +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message from memoryscope.storage.base_memory_store import BaseMemoryStore from memoryscope.storage.base_monitor import BaseMonitor -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.memory_handler import MemoryHandler from memoryscope.utils.prompt_handler import PromptHandler @@ -77,6 +79,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): """ return self.get_context(CHAT_KWARGS) + @property + def workflow_name(self) -> str: + return self.get_context(WORKFLOW_NAME) + + @property + def memoryscope_context(self) -> MemoryscopeContext: + return self.get_context(MEMORYSCOPE_CONTEXT) + + @property + def language(self) -> LanguageEnum: + return self.memoryscope_context.language + @property def embedding_model(self) -> BaseModel: """ @@ -87,8 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The embedding model used for converting text into vector representations. """ if isinstance(self._embedding_model, str): - self._embedding_model = G_CONTEXT.model_conf_dict[self._embedding_model] - # ⭐ Retrieve the actual model instance when the attribute is a string reference + self._embedding_model = self.memoryscope_context.model_dict[self._embedding_model] return self._embedding_model @property @@ -101,8 +114,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The model used for text generation. """ if isinstance(self._generation_model, str): - self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model] - # ⭐ Retrieve the model instance if currently a string reference + self._generation_model = self.memoryscope_context.model_dict[self._generation_model] return self._generation_model @property @@ -115,7 +127,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The rank model instance used for ranking tasks. """ if isinstance(self._rank_model, str): - self._rank_model = G_CONTEXT.model_conf_dict[self._rank_model] # Fetch model instance if string reference + self._rank_model = self.memoryscope_context.model_dict[self._rank_model] return self._rank_model @property @@ -128,7 +140,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store_conf + self._memory_store = self.memoryscope_context.memory_store return self._memory_store @property @@ -141,7 +153,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMonitor: The monitoring component instance. """ if self._monitor is None: - self._monitor = G_CONTEXT.monitor_conf + self._monitor = self.memoryscope_context.monitor return self._monitor @property @@ -154,7 +166,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): str: The name of the assistant. """ if self._user_name is None: - self._user_name = G_CONTEXT.meta_data["assistant_name"] + self._user_name = self.memoryscope_context.meta_data["assistant_name"] return self._user_name @property @@ -166,7 +178,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): str: The readable name of the human. """ if self._target_name is None: - self._target_name = G_CONTEXT.meta_data["human_name"] + self._target_name = self.memoryscope_context.meta_data["human_name"] return self._target_name @property @@ -182,19 +194,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._prompt_handler @property - def memory_handler(self) -> MemoryHandler: + def memory_manager(self) -> MemoryManager: """ Lazily initializes and returns the MemoryHandler instance. Returns: MemoryHandler: An instance of MemoryHandler. """ - if not self.has_content(MEMORY_HANDLER): - self.set_context(MEMORY_HANDLER, MemoryHandler()) # Initialize the memory handler if not present - return self.get_context(MEMORY_HANDLER) + if not self.has_content(MEMORY_MANAGER): + self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context)) + return self.get_context(MEMORY_MANAGER) - @staticmethod - def get_language_value(languages: dict | List[dict]) -> Any | List[Any]: + def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]: """ Retrieves the value(s) corresponding to the current language context. @@ -205,5 +216,5 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Any | list[Any]: The value or list of values matching the current language setting. """ if isinstance(languages, list): - return [x[G_CONTEXT.language] for x in languages] - return languages[G_CONTEXT.language] + return [x[self.language] for x in languages] + return languages[self.language] diff --git a/memoryscope/utils/memory_handler.py b/memoryscope/memory/worker/memory_manager.py similarity index 96% rename from memoryscope/utils/memory_handler.py rename to memoryscope/memory/worker/memory_manager.py index a3cbf889..f705d832 100644 --- a/memoryscope/utils/memory_handler.py +++ b/memoryscope/memory/worker/memory_manager.py @@ -2,21 +2,20 @@ from typing import Dict, List from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.scheme.memory_node import MemoryNode from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.utils.global_context import G_CONTEXT from memoryscope.utils.logger import Logger -class MemoryHandler(object): +class MemoryManager(object): """ The `MemoryHandler` class manages memory nodes with memory store. """ - def __init__(self): - """ - Initializes the MemoryHandler. - """ + def __init__(self, memoryscope_context: MemoryscopeContext): + self.memoryscope_context: MemoryscopeContext = memoryscope_context + self._memory_store: BaseMemoryStore | None = None # dict: memory_id -> MemoryNode @@ -36,7 +35,7 @@ class MemoryHandler(object): BaseMemoryStore: The memory store instance associated with this worker. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store_conf + self._memory_store = self.memoryscope_context.memory_store return self._memory_store def clear(self): diff --git a/memoryscope/models/dummy_generation_model.py b/memoryscope/models/dummy_generation_model.py index 5949bf8e..ee6eead6 100644 --- a/memoryscope/models/dummy_generation_model.py +++ b/memoryscope/models/dummy_generation_model.py @@ -20,9 +20,6 @@ class DummyGenerationModel(BaseModel): m_type: ModelEnum = ModelEnum.GENERATION_MODEL class DummyModel: - """ - An inner class representing the dummy model placeholder. - """ pass MODEL_REGISTRY.register("dummy_generation", DummyModel) @@ -79,12 +76,12 @@ class DummyGenerationModel(BaseModel): for delta in call_result: model_response.message.content += delta model_response.delta = delta - time.sleep(0.1) # ⭐ Introduce a delay to simulate streaming + time.sleep(0.1) yield model_response return gen() else: - model_response.message.content = "".join(call_result) # ⭐ Concatenate results for non-streaming + model_response.message.content = "".join(call_result) return model_response def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: diff --git a/memoryscope/utils/datetime_handler.py b/memoryscope/utils/datetime_handler.py index 6f2d6d2f..44a0b3ad 100644 --- a/memoryscope/utils/datetime_handler.py +++ b/memoryscope/utils/datetime_handler.py @@ -1,8 +1,9 @@ import datetime import re +from typing import List from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT -from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.utils.logger import Logger @@ -40,7 +41,7 @@ class DatetimeHandler(object): self._dt_info_dict: dict | None = None - def _parse_dt_info(self): + def _parse_dt_info(self, language: LanguageEnum): """ Parses the datetime object (_dt) into a dictionary containing detailed date and time components, including language-specific weekday representation. @@ -52,17 +53,16 @@ class DatetimeHandler(object): """ return { "year": self._dt.year, - "month": MONTH_DICT[G_CONTEXT.language][self._dt.month - 1], + "month": MONTH_DICT[language][self._dt.month - 1], "day": self._dt.day, "hour": self._dt.hour, "minute": self._dt.minute, "second": self._dt.second, "week": self._dt.isocalendar().week, - "weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1], + "weekday": WEEKDAYS[language][self._dt.isocalendar().weekday - 1], } - @property - def dt_info_dict(self): + def get_dt_info_dict(self, language: LanguageEnum): """ Property method to get the dictionary containing parsed datetime information. If None, initialize using `_parse_dt_info`. @@ -71,7 +71,7 @@ class DatetimeHandler(object): dict: A dictionary with parsed datetime information. """ if self._dt_info_dict is None: - self._dt_info_dict = self._parse_dt_info() + self._dt_info_dict = self._parse_dt_info(language=language) return self._dt_info_dict @classmethod @@ -207,7 +207,7 @@ class DatetimeHandler(object): return date_info @classmethod - def extract_date_parts(cls, input_string: str) -> dict: + def extract_date_parts(cls, input_string: str, language: LanguageEnum) -> dict: """ Extracts various date components from the input string based on the current language context. @@ -217,48 +217,51 @@ class DatetimeHandler(object): Args: input_string (str): The string containing date information to be parsed. + language (str): current language. Returns: dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails. """ - func_name = f"extract_date_parts_{G_CONTEXT.language.value}" + func_name = f"extract_date_parts_{language}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete extract_date_parts func!") + cls.logger.warning(f"language={language} needs to complete extract_date_parts func!") return {} return getattr(cls, func_name)(input_string=input_string) @classmethod - def has_time_word_cn(cls, query: str) -> bool: + def has_time_word_cn(cls, query: str, datetime_word_list: List[str]) -> bool: """ Check if the input query contains any datetime-related words based on the cn language context. Args: query (str): The input string to check for datetime-related words. + datetime_word_list (list[str]): datetime keywords Returns: bool: True if the query contains at least one datetime-related word, False otherwise. """ contain_datetime = False # TODO use re - for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: + for datetime_word in datetime_word_list: if datetime_word in query: contain_datetime = True break return contain_datetime @classmethod - def has_time_word_en(cls, query: str) -> bool: + def has_time_word_en(cls, query: str, datetime_word_list: List[str]) -> bool: """ Check if the input query contains any datetime-related words based on the en language context. Args: query (str): The input string to check for datetime-related words. + datetime_word_list (list[str]): datetime keywords Returns: bool: True if the query contains at least one datetime-related word, False otherwise. """ contain_datetime = False - for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: + for datetime_word in datetime_word_list: datetime_word = datetime_word.lower() # TODO fix strip if datetime_word in [x.strip().lower().strip(",").strip(".").strip("?").strip(":") @@ -268,12 +271,18 @@ class DatetimeHandler(object): return contain_datetime @classmethod - def has_time_word(cls, query: str) -> bool: - func_name = f"has_time_word_{G_CONTEXT.language.value}" + def has_time_word(cls, query: str, language: LanguageEnum) -> bool: + func_name = f"has_time_word_{language}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete has_time_word func!") + cls.logger.warning(f"language={language} needs to complete has_time_word function!") return False - return getattr(cls, func_name)(query=query) + + if language not in DATATIME_WORD_LIST: + cls.logger.warning(f"language={language} is missing in DATATIME_WORD_LIST!") + return False + + datetime_word_list = DATATIME_WORD_LIST[language] + return getattr(cls, func_name)(query=query, datetime_word_list=datetime_word_list) def datetime_format(self, dt_format: str = "%Y%m%d") -> str: """ @@ -287,7 +296,7 @@ class DatetimeHandler(object): """ return self._dt.strftime(dt_format) - def string_format(self, string_format: str) -> str: + def string_format(self, string_format: str, language: LanguageEnum) -> str: """ Format the datetime information stored in the instance using a custom string format. @@ -297,7 +306,7 @@ class DatetimeHandler(object): Returns: str: A formatted datetime string. """ - return string_format.format(**self.dt_info_dict) + return string_format.format(**self.get_dt_info_dict(language=language)) @property def timestamp(self) -> int: diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index a8114b35..c841d73d 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -140,7 +140,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -172,7 +172,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -202,7 +202,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -224,30 +224,30 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"), MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作或有工作经验。"), MemoryNode(user_name="AI", target_name="用户", content="用户跳槽至openai工作。"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"), @@ -257,10 +257,10 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="我爱吃苹果和香蕉"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) worker.logger.info(f"result1={result1}") worker.logger.info(f"result2={result2}") @@ -293,11 +293,11 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(content="用户想知道维持广泛社交关系的方法。"), ] - worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) - worker.memory_handler.set_memories(INSIGHT_NODES, []) + worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_manager.set_memories(INSIGHT_NODES, []) worker.run() - result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.get_reflection={result}") return worker @@ -320,10 +320,10 @@ class TestWorkersCn(unittest.TestCase): nodes = [ MemoryNode(content="用户喜欢打王者荣耀"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") @@ -345,10 +345,10 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"), MemoryNode(content="用户在上海工作。"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.unit_test_flag = True worker.run() - result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}") diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 6b1c012a..356c3f21 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -150,7 +150,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -193,7 +193,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -226,7 +226,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -248,30 +248,30 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="User works at Alibaba"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"), MemoryNode(user_name="AI", target_name="用户", content="Users working in Meituan"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"), MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="I like to eat watermelon"), @@ -279,10 +279,10 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="I don't like watermelon"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) worker.logger.info(f"result1={result1}") worker.logger.info(f"result2={result2}") @@ -316,11 +316,11 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(content="Users want to know how to maintain extensive social relationships."), ] - worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) - worker.memory_handler.set_memories(INSIGHT_NODES, []) + worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_manager.set_memories(INSIGHT_NODES, []) worker.run() - result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.get_reflection={result}") return worker @@ -343,10 +343,10 @@ class TestWorkersEn(unittest.TestCase): nodes = [ MemoryNode(content="Users like to play King of Glory"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") @@ -368,10 +368,10 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(content="The user works in Beijing, feels stressed, and is looking for ways to relax."), MemoryNode(content="User works in Shanghai."), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.unit_test_flag = True worker.run() - result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}") From a5b45afa8c2de4153d392e2ce899f629b70d85f0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 27 Jul 2024 21:54:19 +0800 Subject: [PATCH 09/15] fix grammer bugs --- config/demo_config_no_stream.yaml | 148 ------------------ config/docker_config.yaml | 83 ---------- examples/docker/__init__.py | 0 examples/docker/docker_config.yaml | 0 memoryscope/argument/cli_chat_demo.yaml | 1 + memoryscope/argument/memoryscope_arguments.py | 7 +- memoryscope/enumeration/model_enum.py | 5 +- .../memory/operation/base_operation.py | 1 - .../memory/service/base_memory_service.py | 2 +- .../worker/backend/contra_repeat_worker.py | 2 +- .../get_observation_with_time_worker.py | 4 +- .../worker/backend/get_observation_worker.py | 2 +- .../backend/get_reflection_subject_worker.py | 5 +- .../worker/backend/info_filter_worker.py | 2 +- .../backend/long_contra_repeat_worker.py | 3 +- .../worker/backend/update_insight_worker.py | 79 ++++++---- .../worker/frontend/extract_time_worker.py | 5 +- .../worker/frontend/semantic_rank_worker.py | 7 + memoryscope/memoryscope.py | 5 +- memoryscope/models/dummy_generation_model.py | 67 +++----- memoryscope/scheme/memory_node.py | 4 +- .../storage/llama_index_es_memory_store.py | 12 +- .../storage/llama_index_sync_elasticsearch.py | 5 +- memoryscope/utils/datetime_handler.py | 1 + memoryscope/utils/logger.py | 2 +- memoryscope/utils/registry.py | 8 +- memoryscope/utils/response_text_parser.py | 25 +-- memoryscope/utils/timer.py | 2 +- memoryscope/utils/tool_functions.py | 21 ++- 29 files changed, 157 insertions(+), 351 deletions(-) delete mode 100644 config/docker_config.yaml create mode 100644 examples/docker/__init__.py create mode 100644 examples/docker/docker_config.yaml diff --git a/config/demo_config_no_stream.yaml b/config/demo_config_no_stream.yaml index b9154131..e69de29b 100644 --- a/config/demo_config_no_stream.yaml +++ b/config/demo_config_no_stream.yaml @@ -1,148 +0,0 @@ -global_config: - language: cn - max_workers: 5 - -memory_chat: - cli_memory_chat: - class: chat.cli_memory_chat - stream: false - memory_service: memory_chat_service - generation_model: dashscope_generation - -memory_service: - memory_chat_service: - class: memory.service.chat_memory_service - history_msg_count: 32 - contextual_msg_count: 6 - memory_operations: - read_message: - class: memory.operation.read_message - description: "read session messages of the user" - read_memory: - class: memory.operation.read_memory - workflow: set_query,retrieve_memory1,[extract_time|semantic_rank],fuse_rerank - description: "read related memories of the user" - list_memory: - class: memory.operation.read_memory - workflow: set_query,retrieve_memory2,print_memory - description: "read all memories of the user" - write_memory: - class: memory.operation.write_memory - workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory - description: "write observation memories of the user" - interval_time: 5 -# summary_memory: -# class: memory.operation.summary_memory -# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory -# description: "summary observation memories of the user" -# interval_time: 60 - -worker: - dummy: - class: memory.worker.dummy_worker - generation_model: dashscope_generation - embedding_model: dashscope_embedding - rank_model: dashscope_rank - set_query: - class: memory.worker.read.set_query_worker - retrieve_memory1: - class: memory.worker.read.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - retrieve_expired_top_k: 0 - extract_time: - class: memory.worker.read.extract_time_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - semantic_rank: - class: memory.worker.read.semantic_rank_worker - fuse_rerank: - class: memory.worker.read.fuse_rerank_worker - fuse_score_threshold: 0.1 - fuse_ratio_dict: - conversation: 0.5 - observation: 1 - obs_customized: 1.2 - insight: 2.0 - fuse_time_ratio: 2.0 - fuse_rerank_top_k: 10 - retrieve_memory2: - class: memory.worker.read.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - retrieve_expired_top_k: 100 - print_memory: - class: memory.worker.read.print_memory_worker - info_filter: - class: memory.worker.write.info_filter_worker - generation_model: dashscope_generation - info_filter_msg_max_size: 200 - generation_model_top_k: 1 - load_memory1: - class: memory.worker.write.load_memory_worker - retrieve_not_reflected_top_k: 0 - retrieve_not_updated_top_k: 0 - retrieve_insight_top_k: 0 - today_obs_top_k: 100 - get_observation: - class: memory.worker.write.get_observation_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - get_observation_with_time: - class: memory.worker.write.get_observation_with_time_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - contra_repeat: - class: memory.worker.write.contra_repeat_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - retrieve_top_k: 30 - contra_repeat_max_count: 50 - store_memory: - class: memory.worker.write.store_memory_worker - store_key: all - load_memory2: - class: memory.worker.write.load_memory_worker - retrieve_not_reflected_top_k: 100 - retrieve_not_updated_top_k: 100 - retrieve_insight_top_k: 100 - today_obs_top_k: 0 - get_reflection_subject: - class: memory.worker.summary.get_reflection_subject_worker - retrieve_top_k: 100 - reflect_obs_cnt_threshold: 32 - generation_model_top_k: 1 - update_insight: - class: memory.worker.summary.update_insight_worker - update_insight_threshold: 0.1 - generation_model_top_k: 1 - update_insight_max_thread: 10 - long_contra_repeat: - class: memory.worker.summary.long_contra_repeat_worker - long_contra_repeat_top_k: 2 - long_contra_repeat_threshold: 0.1 - generation_model_top_k: 1 - -models: - dashscope_generation: - class: models.llama_index_generation_model - module_name: dashscope_generation - model_name: qwen-max - dashscope_embedding: - class: models.llama_index_embedding_model - module_name: dashscope_embedding - model_name: text-embedding-v2 - dashscope_rank: - class: models.llama_index_rank_model - module_name: dashscope_rank - model_name: gte-rerank - -memory_store: - class: storage.llama_index_es_memory_store - embedding_model: dashscope_embedding - index_name: memory_index - es_url: http://localhost:9200 - use_hybrid: false - -monitor: - class: storage.dummy_monitor \ No newline at end of file diff --git a/config/docker_config.yaml b/config/docker_config.yaml deleted file mode 100644 index 2668a5e7..00000000 --- a/config/docker_config.yaml +++ /dev/null @@ -1,83 +0,0 @@ -global_config: - language: cn - max_workers: 5 - dash_scope_apikey: - open_ai_apikey: - -memory_chat: - cli_memory_chat: - class: chat.cli_memory_chat # select class - memory_service: memory_chat_service - generation_model: dashscope_generation - human_name: human - assistant_name: assistant - -memory_service: - memory_chat_service: - class: memory.service.chat_memory_service # select class - history_msg_count: 32 - contextual_msg_count: 6 - read_memory_key: read_memory - memory_operations: - read_message: # define operation - class: memory.operation.read_memory - workflow: dummy_workflow # select workflow - description: "read session messages of the user" - read_memory: - class: memory.operation.read_memory - workflow: dummy_workflow - description: "read related memories of the user" - list_memory: - class: memory.operation.read_memory - workflow: dummy_workflow - description: "read all memories of the user" - write_memory: - class: memory.operation.write_memory - workflow: dummy_workflow - description: "write observation memories of the user" - interval_time: 60 - summary_memory: - class: memory.operation.summary_memory - workflow: dummy_workflow - description: "summary observation memories of the user" - interval_time: 300 - -models: - dashscope_generation: - class: models.llama_index_generation_model # select class - module_name: dashscope_generation - model_name: qwen-max - dashscope_embedding: - class: models.llama_index_embedding_model # select class - module_name: dashscope_embedding - model_name: text-embedding-v2 - dashscope_rank: - class: models.llama_index_rank_model # select class - module_name: dashscope_rank - model_name: gte-rerank - -vector_store: - class: storage.dummy_vector_store # select class - embedding_model: dashscope_embedding - -monitor: - class: storage.dummy_monitor # select class - -worker: - dummy_workflow: - class: memory.worker.dummy_worker - generation_model: dashscope_generation - embedding_model: dashscope_embedding - rank_model: dashscope_rank - retrieve_store_worker: - class: memory.worker.read.retrieve_store_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - fuse_rerank_worker: - class: memory.worker.read.fuse_rerank_worker - fuse_score_threshold: 0.1 - fuse_ratio_dict: - observation: 1 - fuse_time_ratio: 2.0 - fuse_rerank_top_k: 10 - diff --git a/examples/docker/__init__.py b/examples/docker/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/examples/docker/docker_config.yaml b/examples/docker/docker_config.yaml new file mode 100644 index 00000000..e69de29b diff --git a/memoryscope/argument/cli_chat_demo.yaml b/memoryscope/argument/cli_chat_demo.yaml index b8fdee74..16d7f92d 100644 --- a/memoryscope/argument/cli_chat_demo.yaml +++ b/memoryscope/argument/cli_chat_demo.yaml @@ -3,6 +3,7 @@ global_config: thread_pool_max_workers: 5 logger_name: memoryscope logger_name_time_suffix: %Y%m%d_%H%M%S + use_dummy_ranker: true memory_chat: cli_memory_chat: diff --git a/memoryscope/argument/memoryscope_arguments.py b/memoryscope/argument/memoryscope_arguments.py index 92572f25..7f48d281 100644 --- a/memoryscope/argument/memoryscope_arguments.py +++ b/memoryscope/argument/memoryscope_arguments.py @@ -45,6 +45,10 @@ class MemoryscopeArguments(object): embedding_params: dict = field(default_factory=lambda: {}) + use_dummy_ranker: bool = field(default=True, metadata={ + "help": "If a semantic ranking model is not available, MemoryScope will use cosine similarity scoring as a " + "substitute. However, the ranking effectiveness will be somewhat compromised."}) + rank_backend: str = field(default="dashscope_rank", metadata={"help": "global rank backend: dashscope_rank, etc."}) rank_model: str = field(default="gte-rerank", metadata={"help": "global rank model: gte-rerank, etc."}) @@ -58,4 +62,5 @@ class MemoryscopeArguments(object): retrieve_mode: str = field(default="dense", metadata={ "help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"}) - hybrid_alpha: float | None = field(default=1.0, metadata={"help": ""}) + hybrid_alpha: float | None = field(default=1.0, metadata={ + "help": "fuse alpha params used in hybrid mode(not implemented)"}) diff --git a/memoryscope/enumeration/model_enum.py b/memoryscope/enumeration/model_enum.py index 4f7cfb44..8cc76d2f 100644 --- a/memoryscope/enumeration/model_enum.py +++ b/memoryscope/enumeration/model_enum.py @@ -7,8 +7,9 @@ class ModelEnum(str, Enum): Members: GENERATION_MODEL: Represents a model responsible for generating content. - EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a numerical form suitable for machine learning tasks. - RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance or importance. + EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a + numerical form suitable for machine learning tasks. + RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance. """ GENERATION_MODEL = "generation_model" diff --git a/memoryscope/memory/operation/base_operation.py b/memoryscope/memory/operation/base_operation.py index df27ea2f..600e3dd4 100644 --- a/memoryscope/memory/operation/base_operation.py +++ b/memoryscope/memory/operation/base_operation.py @@ -12,7 +12,6 @@ class BaseOperation(metaclass=ABCMeta): operation_type (OPERATION_TYPE): Specifies the type of operation, defaulting to "frontend". name (str): The name of the operation. description (str): A description of the operation. - kwargs (dict): Additional keyword arguments for operation configuration. """ operation_type: OPERATION_TYPE = "frontend" diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/memory/service/base_memory_service.py index 7b9fdee5..a0a476a5 100644 --- a/memoryscope/memory/service/base_memory_service.py +++ b/memoryscope/memory/service/base_memory_service.py @@ -54,7 +54,7 @@ class BaseMemoryService(metaclass=ABCMeta): def start_backend_service(self): pass - def stop_backend_service(self): + def stop_backend_service(self, wait_service_end: bool = False): pass def do_operation(self, name: str, **kwargs): diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.py b/memoryscope/memory/worker/backend/contra_repeat_worker.py index 85a18884..d245ba85 100644 --- a/memoryscope/memory/worker/backend/contra_repeat_worker.py +++ b/memoryscope/memory/worker/backend/contra_repeat_worker.py @@ -80,7 +80,7 @@ class ContraRepeatWorker(MemoryBaseWorker): response_text = response.message.content # parse text - idx_merge_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() + idx_merge_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1() if len(idx_merge_obs_list) <= 0: self.logger.warning("idx_merge_obs_list is empty!") return diff --git a/memoryscope/memory/worker/backend/get_observation_with_time_worker.py b/memoryscope/memory/worker/backend/get_observation_with_time_worker.py index 9655fce0..5c7ca66b 100644 --- a/memoryscope/memory/worker/backend/get_observation_with_time_worker.py +++ b/memoryscope/memory/worker/backend/get_observation_with_time_worker.py @@ -26,7 +26,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): filter_messages = [] for msg in self.chat_messages: # Checks if the message content has any time reference words - if DatetimeHandler.has_time_word(query=msg.content): + if DatetimeHandler.has_time_word(query=msg.content, language=self.language): filter_messages.append(msg) return filter_messages @@ -49,7 +49,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): for i, msg in enumerate(filter_messages): # Create a DatetimeHandler instance for each message's timestamp and format it dt_handler = DatetimeHandler(dt=msg.time_created) - dt = dt_handler.string_format(self.prompt_handler.time_string_format) + dt = dt_handler.string_format(string_format=self.prompt_handler.time_string_format, language=self.language) # Append formatted timestamp-query pairs to the user_query_list user_query_list.append(f"{i + 1} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") diff --git a/memoryscope/memory/worker/backend/get_observation_worker.py b/memoryscope/memory/worker/backend/get_observation_worker.py index 385251d1..0a90803b 100644 --- a/memoryscope/memory/worker/backend/get_observation_worker.py +++ b/memoryscope/memory/worker/backend/get_observation_worker.py @@ -139,7 +139,7 @@ class GetObservationWorker(MemoryBaseWorker): response_text = response.message.content # Parses the generated text to extract observation indices, times, contents, and keywords - idx_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() + idx_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1() if len(idx_obs_list) <= 0: self.logger.warning("idx_obs_list is empty!") return diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py b/memoryscope/memory/worker/backend/get_reflection_subject_worker.py index 711206ae..f9e5511d 100644 --- a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py +++ b/memoryscope/memory/worker/backend/get_reflection_subject_worker.py @@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): """ dt_handler = DatetimeHandler() # Prepare metadata with current datetime info - meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict.items()} + meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()} return MemoryNode(user_name=self.user_name, target_name=self.target_name, @@ -101,7 +101,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): return # Parse LLM response for new insight keys and update memory - new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2() + new_insight_keys = ResponseTextParser(response.message.content, self.language, + self.__class__.__name__).parse_v2() if new_insight_keys: for insight_key in new_insight_keys: self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) diff --git a/memoryscope/memory/worker/backend/info_filter_worker.py b/memoryscope/memory/worker/backend/info_filter_worker.py index ae28c13d..1474429c 100644 --- a/memoryscope/memory/worker/backend/info_filter_worker.py +++ b/memoryscope/memory/worker/backend/info_filter_worker.py @@ -76,7 +76,7 @@ class InfoFilterWorker(MemoryBaseWorker): response_text = response.message.content # parse text - info_score_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() + info_score_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1() if len(info_score_list) != len(info_messages): self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}") diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py b/memoryscope/memory/worker/backend/long_contra_repeat_worker.py index 056c63d7..c12c2cc3 100644 --- a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/memory/worker/backend/long_contra_repeat_worker.py @@ -111,7 +111,8 @@ class LongContraRepeatWorker(MemoryBaseWorker): return # Parses the model's response text to identify updates for memory nodes - idx_obs_info_list = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v1() + idx_obs_info_list = ResponseTextParser(response.message.content, self.language, + self.__class__.__name__).parse_v1() if len(idx_obs_info_list) <= 0: self.logger.warning("idx_obs_info_list is empty!") return diff --git a/memoryscope/memory/worker/backend/update_insight_worker.py b/memoryscope/memory/worker/backend/update_insight_worker.py index 339d7871..a4b66e2e 100644 --- a/memoryscope/memory/worker/backend/update_insight_worker.py +++ b/memoryscope/memory/worker/backend/update_insight_worker.py @@ -8,7 +8,7 @@ from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.utils.datetime_handler import DatetimeHandler from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg +from memoryscope.utils.tool_functions import prompt_to_msg, cosine_similarity class UpdateInsightWorker(MemoryBaseWorker): @@ -27,13 +27,15 @@ class UpdateInsightWorker(MemoryBaseWorker): def filter_obs_nodes(self, insight_node: MemoryNode, - obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float): + obs_nodes: List[MemoryNode], + use_dummy_ranker: bool) -> (MemoryNode, List[MemoryNode], float): """ Filters observed nodes based on their relevance to a given insight node using a ranking model. Args: insight_node (MemoryNode): The insight node used as the basis for filtering. obs_nodes (List[MemoryNode]): A list of observed nodes to be filtered. + use_dummy_ranker (bool): Global parameters, whether to use rank model or not. Returns: tuple: A tuple containing: @@ -53,24 +55,47 @@ class UpdateInsightWorker(MemoryBaseWorker): self.logger.warning("obs_nodes is empty!") return insight_node, filtered_nodes, max_score - # Call the ranking model to get scores for each observed node's content against the insight key - documents = [x.content for x in obs_nodes] - self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}") - response = self.rank_model.call(query=insight_node.key, documents=documents) - if not response.status: - return insight_node, filtered_nodes, max_score + if use_dummy_ranker: + key_vector: List[float] = self.embedding_model.call(text=insight_node.key).embedding_results + if not key_vector: + self.logger.warning(f"embedding call {insight_node.key} failed!") + return insight_node, filtered_nodes, max_score - # Iterate over the ranked scores to filter nodes - for index, score in response.rank_scores.items(): - node = obs_nodes[index] - # Determine if the node should be kept based on the threshold - keep_flag = score >= self.update_insight_threshold - if keep_flag: - filtered_nodes.append(node) - max_score = max(max_score, score) - # Log information about each node's processing - self.logger.info(f"insight_key={insight_node.key} content={node.content} " - f"score={score} keep_flag={keep_flag}") + insight_node.key_vector = key_vector + documents_vector = [x.vector for x in obs_nodes] + score_recall_list = cosine_similarity(key_vector, documents_vector) + assert len(score_recall_list) == len(obs_nodes), \ + f"size is not as excepted. {len(score_recall_list)} v.s. {len(obs_nodes)}" + + for score, node in zip(score_recall_list, obs_nodes): + keep_flag = score >= self.update_insight_threshold + if keep_flag: + filtered_nodes.append(node) + max_score = max(max_score, score) + + # Log information about each node's processing + self.logger.info(f"insight_key={insight_node.key} content={node.content} " + f"score={score} keep_flag={keep_flag}") + + else: + # Call the ranking model to get scores for each observed node's content against the insight key + documents = [x.content for x in obs_nodes] + self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}") + response = self.rank_model.call(query=insight_node.key, documents=documents) + if not response.status: + return insight_node, filtered_nodes, max_score + + # Iterate over the ranked scores to filter nodes + for index, score in response.rank_scores.items(): + node = obs_nodes[index] + # Determine if the node should be kept based on the threshold + keep_flag = score >= self.update_insight_threshold + if keep_flag: + filtered_nodes.append(node) + max_score = max(max_score, score) + # Log information about each node's processing + self.logger.info(f"insight_key={insight_node.key} content={node.content} " + f"score={score} keep_flag={keep_flag}") # Warn if no nodes were filtered if not filtered_nodes: @@ -95,7 +120,7 @@ class UpdateInsightWorker(MemoryBaseWorker): content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}" insight_node.content = content insight_node.value = insight_value - insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict.items()}) + insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}) insight_node.timestamp = dt_handler.timestamp insight_node.dt = dt_handler.datetime_format() if insight_node.action_status == ActionStatusEnum.NONE.value: @@ -136,7 +161,7 @@ class UpdateInsightWorker(MemoryBaseWorker): if not response.status or not response.message.content: return insight_node - insight_value_list = ResponseTextParser(response.message.content, + insight_value_list = ResponseTextParser(response.message.content, self.language, f"update_{insight_node.key}").parse_v1() if not insight_value_list: self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!") @@ -184,17 +209,21 @@ class UpdateInsightWorker(MemoryBaseWorker): self.logger.warning("insight_nodes is empty, stopping processing.") return + use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"] + # Process active insight nodes with corresponding not updated nodes for node in insight_nodes: time.sleep(1) if node.action_status == ActionStatusEnum.NEW.value: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, - obs_nodes=not_reflected_nodes) + obs_nodes=not_reflected_nodes, + use_dummy_ranker=use_dummy_ranker) else: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, - obs_nodes=not_updated_nodes) + obs_nodes=not_updated_nodes, + use_dummy_ranker=use_dummy_ranker) # select top n result_list = [] @@ -221,7 +250,3 @@ class UpdateInsightWorker(MemoryBaseWorker): for node in not_updated_nodes: node.obs_updated = 1 node.action_status = ActionStatusEnum.MODIFIED - - # for node in not_reflected_nodes: - # node.obs_updated = 1 - # node.action_status = ActionStatusEnum.MODIFIED diff --git a/memoryscope/memory/worker/frontend/extract_time_worker.py b/memoryscope/memory/worker/frontend/extract_time_worker.py index 087431c5..6f92c3c8 100644 --- a/memoryscope/memory/worker/frontend/extract_time_worker.py +++ b/memoryscope/memory/worker/frontend/extract_time_worker.py @@ -33,13 +33,14 @@ class ExtractTimeWorker(MemoryBaseWorker): query, query_timestamp = self.get_context(QUERY_WITH_TS) # Identify if the query contains datetime keywords - contain_datetime = DatetimeHandler.has_time_word(query) + contain_datetime = DatetimeHandler.has_time_word(query, self.language) if not contain_datetime: self.logger.info(f"contain_datetime={contain_datetime}") return # Prepare the prompt with necessary contextual details - query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format) + query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format, + self.language) system_prompt = self.prompt_handler.extract_time_system few_shot = self.prompt_handler.extract_time_few_shot user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str) diff --git a/memoryscope/memory/worker/frontend/semantic_rank_worker.py b/memoryscope/memory/worker/frontend/semantic_rank_worker.py index 9a78b96c..4dc7b303 100644 --- a/memoryscope/memory/worker/frontend/semantic_rank_worker.py +++ b/memoryscope/memory/worker/frontend/semantic_rank_worker.py @@ -34,6 +34,13 @@ class SemanticRankWorker(MemoryBaseWorker): self.logger.warning("Retrieve memory nodes is empty!") return + use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"] + if use_dummy_ranker: + for node in memory_node_list: + node.score_rank = node.score_recall + self.logger.warning("use score_recall instead of score_rank!") + return + # drop repeated memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if n.content.strip()} memory_node_list = list(memory_node_dict.values()) diff --git a/memoryscope/memoryscope.py b/memoryscope/memoryscope.py index 0648871c..647b8418 100644 --- a/memoryscope/memoryscope.py +++ b/memoryscope/memoryscope.py @@ -52,6 +52,7 @@ class MemoryScope(object): "thread_pool_max_workers": arguments.thread_pool_max_workers, "logger_name": arguments.logger_name, "logger_name_time_suffix": arguments.logger_name_time_suffix, + "use_dummy_ranker": arguments.use_dummy_ranker, } # prepare memory chat @@ -151,6 +152,7 @@ class MemoryScope(object): # set global config self.context.language = LanguageEnum(self.global_conf["language"]) self.context.thread_pool = ThreadPoolExecutor(max_workers=self.global_conf["max_workers"]) + self.context.meta_data["use_dummy_ranker"] = self.global_conf["use_dummy_ranker"] # init memory_chat if self.memory_chat_conf_dict: @@ -181,8 +183,9 @@ class MemoryScope(object): self.context.worker_config = self.worker_conf_dict def close(self): + # wait service to stop for _, service in self.context.memory_service_dict.items(): - service.stop_backend_service() + service.stop_backend_service(wait_service_end=True) self.context.memory_store.close() self.context.thread_pool.shutdown() diff --git a/memoryscope/models/dummy_generation_model.py b/memoryscope/models/dummy_generation_model.py index ee6eead6..d1ad0a54 100644 --- a/memoryscope/models/dummy_generation_model.py +++ b/memoryscope/models/dummy_generation_model.py @@ -19,35 +19,37 @@ class DummyGenerationModel(BaseModel): """ m_type: ModelEnum = ModelEnum.GENERATION_MODEL - class DummyModel: - pass + MODEL_REGISTRY.register("dummy_generation", object) - MODEL_REGISTRY.register("dummy_generation", DummyModel) - - def before_call(self, **kwargs): + def before_call(self, model_response: ModelResponse, **kwargs): """ - Prepares the input data before making a call to the model's generate function. - Accepts either a 'prompt' or a list of 'messages'. If both are provided or missing, - a RuntimeError is raised. Transforms the input into a standardized format for processing. + Prepares the input data before making a call to the language model. + It accepts either a 'prompt' directly or a list of 'messages'. + If 'prompt' is provided, it sets the data accordingly. + If 'messages' are provided, it constructs a list of ChatMessage objects from the list. + Raises an error if neither 'prompt' nor 'messages' are supplied. Args: - **kwargs: Arbitrary keyword arguments including 'prompt' or 'messages'. - + model_response: model_response + **kwargs: Arbitrary keyword arguments including 'prompt' and 'messages'. + Raises: - RuntimeError: If neither 'prompt' nor 'messages' is provided, or both are provided. + RuntimeError: When both 'prompt' and 'messages' inputs are not provided. """ prompt: str = kwargs.pop("prompt", "") messages: List[Message] | List[dict] = kwargs.pop("messages", []) if prompt: - self.data = {"prompt": prompt} + data = {"prompt": prompt} elif messages: if isinstance(messages[0], dict): - self.data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]} + data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]} else: - self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]} + data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]} else: - raise RuntimeError("Both 'prompt' and 'messages' are empty!") + raise RuntimeError("prompt and messages are both empty!") + data.update(**kwargs) + model_response.meta_data["data"] = data def after_call(self, model_response: ModelResponse, @@ -84,35 +86,8 @@ class DummyGenerationModel(BaseModel): model_response.message.content = "".join(call_result) return model_response - def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: - """ - Generates a dummy response based on the input data, supporting both immediate - and streamed response types. + def _call(self, model_response: ModelResponse, stream: bool = False, **kwargs): + return model_response - Args: - stream (bool, optional): If True, indicates the response should be generated - in a streaming manner. Defaults to False. - **kwargs: Additional keyword arguments not used in this dummy implementation. - - Returns: - Union[ModelResponse, ModelResponseGen]: A dummy response object or a generator - object capable of streaming responses. - """ - assert "prompt" in self.data or "messages" in self.data - results = ModelResponse(m_type=self.m_type) - return results - - async def _async_call(self, **kwargs) -> ModelResponse: - """ - Asynchronous version of `_call`, providing the same functionality but designed - to be used in asynchronous contexts. - - Args: - **kwargs: Additional keyword arguments not used in this dummy implementation. - - Returns: - ModelResponse: A dummy response object suitable for asynchronous use. - """ - assert "prompt" in self.data or "messages" in self.data - results = ModelResponse(m_type=self.m_type) - return results + async def _async_call(self, model_response: ModelResponse, **kwargs): + return model_response diff --git a/memoryscope/scheme/memory_node.py b/memoryscope/scheme/memory_node.py index de884ae7..32205a54 100644 --- a/memoryscope/scheme/memory_node.py +++ b/memoryscope/scheme/memory_node.py @@ -23,6 +23,8 @@ class MemoryNode(BaseModel): key: str = Field("", description="memory key") + key_vector: List[float] = Field([], description="memory key embedding result") + value: str = Field("", description="memory value") score_recall: float = Field(0, description="embedding similarity score used in recall stage") @@ -37,7 +39,7 @@ class MemoryNode(BaseModel): store_status: str = Field("valid", description="store_status: valid / expired") - vector: List[float] = Field([], description="content embedding result, return empty") + vector: List[float] = Field([], description="content embedding result") timestamp: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()), description="timestamp of the memory node") diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index 96b1542f..22d43dd7 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -1,5 +1,5 @@ import random -from typing import Dict, List, Optional +from typing import Dict, List from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle @@ -24,10 +24,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.emb_dims = None self.index_name = index_name self.embedding_model: BaseModel = embedding_model + retrieval_strategy = ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, hybrid_alpha=hybrid_alpha) self.es_store = SyncElasticsearchStore(index_name=index_name, es_url=es_url, - retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, - hybrid_alpha=hybrid_alpha), + retrieval_strategy=retrieval_strategy, **kwargs) # TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By @@ -144,7 +144,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): return TextNode(id_=memory_node.memory_id, text=memory_node.content, embedding=embedding, - metadata=memory_node.model_dump(exclude={"content", "vector", "score_recall", "score_rank", "score_rerank"})) + metadata=memory_node.model_dump(exclude={"content", + "vector", + "score_recall", + "score_rank", + "score_rerank"})) @staticmethod def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode: diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 0794b284..4bfa2de4 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -133,7 +133,6 @@ def _mode_must_match_retrieval_strategy( class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): - def __init__( self, *, @@ -613,7 +612,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore): ] = None, es_filter: Optional[List[Dict]] = None, fields: List[str] = [], - **kwargs: Any, ) -> VectorStoreQueryResult: """ Asynchronously queries the Elasticsearch index for the top k most similar nodes @@ -626,6 +624,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): A custom function to modify the Elasticsearch query body. Defaults to None. es_filter (List[Dict], optional): Additional filters to apply during the query. If filters are present in the query, these filters will not be used. Defaults to None. + fields (List[str], optional): . Returns: VectorStoreQueryResult: The result of the query, including nodes, their IDs, @@ -700,7 +699,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy) and self.retrieval_strategy.hybrid ): - total_rank = sum(top_k_scores) + # total_rank = sum(top_k_scores) top_k_scores = [rank for rank in top_k_scores] # top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores] # top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores] diff --git a/memoryscope/utils/datetime_handler.py b/memoryscope/utils/datetime_handler.py index 44a0b3ad..f41cfcfa 100644 --- a/memoryscope/utils/datetime_handler.py +++ b/memoryscope/utils/datetime_handler.py @@ -302,6 +302,7 @@ class DatetimeHandler(object): Args: string_format (str): A format string where placeholders are keys from `dt_info_dict`. + language (str): current language. Returns: str: A formatted datetime string. diff --git a/memoryscope/utils/logger.py b/memoryscope/utils/logger.py index 9558b3b0..9f400c6d 100644 --- a/memoryscope/utils/logger.py +++ b/memoryscope/utils/logger.py @@ -26,7 +26,7 @@ class Logger(logging.Logger): max_bytes: int = 1024 * 1024 * 1024, backup_count: int = 10): """ - Initializes the Logger instance, setting up handlers for console and/or file logging based on provided parameters. + Initializes the Logger instance, setting up handlers for console and file logging based on provided parameters. Args: name (str): Identifier for the logger. diff --git a/memoryscope/utils/registry.py b/memoryscope/utils/registry.py index df396306..935a6c04 100644 --- a/memoryscope/utils/registry.py +++ b/memoryscope/utils/registry.py @@ -12,7 +12,8 @@ class Registry(object): Attributes: name (str): The name of the registry. - module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are the modules themselves. + module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are + the modules themselves. """ def __init__(self, name: str): @@ -31,7 +32,7 @@ class Registry(object): Args: module_name (str): The name of module to be registered. - modules (List[Any] | Dict[str, Any]): The module to be registered. + module (List[Any] | Dict[str, Any]): The module to be registered. Raises: NotImplementedError: If the input is already registered. @@ -46,7 +47,8 @@ class Registry(object): def batch_register(self, modules: List[Any] | Dict[str, Any]): """ - Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary mapping names to modules. + Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary + mapping names to modules. Args: modules (List[Any] | Dict[str, Any]): A list of modules or a dictionary mapping module names to the modules. diff --git a/memoryscope/utils/response_text_parser.py b/memoryscope/utils/response_text_parser.py index b7de8982..452b5014 100644 --- a/memoryscope/utils/response_text_parser.py +++ b/memoryscope/utils/response_text_parser.py @@ -2,28 +2,22 @@ import re from typing import List from memoryscope.constants.language_constants import NONE_WORD -from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.utils.logger import Logger class ResponseTextParser(object): """ - The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract specific + The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract patterns from the text and filter out unnecessary information, while also logging the processing steps and outcomes. """ - pattern_v1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets - - def __init__(self, response_text: str, logger_prefix: str = ""): - """ - Initializes the `ResponseTextParser` instance with the provided response text and sets up a logger. - - Args: - response_text (str): The raw response text that needs to be parsed and processed. - """ + PATTERN_V1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets + def __init__(self, response_text: str, language: LanguageEnum, logger_prefix: str = ""): # Strips leading and trailing whitespace from the response text self.response_text: str = response_text.strip() + self.language: LanguageEnum = language # The prefix of log. Defaults to "". self.logger_prefix: str = logger_prefix @@ -43,7 +37,7 @@ class ResponseTextParser(object): line = line.strip() if not line: continue - matches = [match.group(1) for match in self.pattern_v1.finditer(line)] + matches = [match.group(1) for match in self.PATTERN_V1.finditer(line)] if matches: result.append(matches) self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2) @@ -51,18 +45,15 @@ class ResponseTextParser(object): def parse_v2(self) -> List[str]: """ - Extract lines which contain NONE_WORD in Chinese or English. + Extract lines which contain NONE_WORD. - Args: - prefix (str): The prefix of log. Defaults to "". - Returns: Contents match the specific patterns. """ result = [] for line in self.response_text.split("\n"): line = line.strip() - if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language): + if not line or line.lower() == NONE_WORD.get(self.language): continue result.append(line) self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2) diff --git a/memoryscope/utils/timer.py b/memoryscope/utils/timer.py index f63ff880..ac7ca1f0 100644 --- a/memoryscope/utils/timer.py +++ b/memoryscope/utils/timer.py @@ -26,7 +26,7 @@ class Timer(object): Args: name (str): The log name. time_log_type (str): The log type. Defaults to 'End'. - use_ms (bool): Use 'ms' as the time scale or not. Defaults to True. + use_ms (bool): Use 'ms' as the timescale or not. Defaults to True. stack_level (int): The stack level of log. Defaults to 2. float_precision (int): The precision of cost time. Defaults to 4. diff --git a/memoryscope/utils/tool_functions.py b/memoryscope/utils/tool_functions.py index 9bb73896..6d5a6834 100644 --- a/memoryscope/utils/tool_functions.py +++ b/memoryscope/utils/tool_functions.py @@ -6,6 +6,7 @@ from copy import deepcopy from importlib import import_module from typing import List +import numpy as np import pyfiglet from termcolor import colored @@ -18,7 +19,7 @@ ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey", def underscore_to_camelcase(name: str, is_first_title: bool = True) -> str: """ - Converts a underscore_notation string to CamelCase. + Converts an underscore_notation string to CamelCase. Args: name (str): The underscore_notation string to be converted. @@ -188,3 +189,21 @@ def contains_keyword(text, keywords) -> bool: escaped_keywords = map(re.escape, keywords) pattern = re.compile('|'.join(escaped_keywords), re.IGNORECASE) return pattern.search(text) is not None + + +def cosine_similarity(query: List[float], documents: List[List[float]]): + query = np.array(query) + documents = np.array(documents) + + query_norm = np.linalg.norm(query) + if query_norm == 0: + raise ValueError("Query vector norm is zero, which will result in a division by zero") + + documents_norm = np.linalg.norm(documents, axis=1) + if np.any(documents_norm == 0): + raise ValueError("One of the document vectors has zero norm, which will result in a division by zero") + + dot_product = np.dot(documents, query) + + cosine_similarities = dot_product / (query_norm * documents_norm) + return cosine_similarities.tolist() From 6174f84b23e1a55badfbfce03ad0a170ee771a79 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 27 Jul 2024 21:59:17 +0800 Subject: [PATCH 10/15] [dev] modify commit --- memoryscope/storage/llama_index_es_memory_store.py | 3 +-- memoryscope/storage/llama_index_sync_elasticsearch.py | 2 +- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index 63fb0cd3..22d43dd7 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -27,8 +27,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): retrieval_strategy = ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, hybrid_alpha=hybrid_alpha) self.es_store = SyncElasticsearchStore(index_name=index_name, es_url=es_url, - retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, - hybrid_alpha=hybrid_alpha), + retrieval_strategy=retrieval_strategy, **kwargs) # TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 4786dad9..50eeef3d 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -267,7 +267,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str } } ) - result['bool'].update({"should": operands}) # тнР Add 'should' clause for OR logic + result['bool'].update({"should": operands}) # Add 'should' clause for OR logic result['bool'].update({"minimum_should_match": 1}) # Ensure at least one 'should' match else: key_str = f"metadata.{key}.keyword" if isinstance(value, str) else f"metadata.{key}" From 70ecdb1d94679472e3ce0665b112c07ed4c209bf Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 28 Jul 2024 19:23:54 +0800 Subject: [PATCH 11/15] add config manager for global configs --- .../config}/demo_config_no_stream.yaml | 0 memoryscope/__init__.py | 3 +- memoryscope/argument/default_arguments.py | 182 ---------------- memoryscope/cli.py | 12 +- memoryscope/constants/common_constants.py | 2 + memoryscope/{argument => core}/__init__.py | 0 memoryscope/{ => core}/chat/__init__.py | 0 .../{ => core}/chat/api_memory_chat.py | 109 ++++++---- .../{ => core}/chat/base_memory_chat.py | 37 +--- .../{ => core}/chat/cli_memory_chat.py | 27 +-- .../{ => core}/chat/memory_chat_prompt.yaml | 0 .../{memory => core/config}/__init__.py | 0 .../config/arguments.py} | 15 +- memoryscope/core/config/config_manager.py | 170 +++++++++++++++ .../config/demo_config.yaml} | 9 +- memoryscope/core/memoryscope.py | 105 +++++++++ memoryscope/{ => core}/memoryscope_context.py | 0 .../operation => core/models}/__init__.py | 0 memoryscope/{ => core}/models/base_model.py | 6 +- .../models/dummy_generation_model.py | 2 +- .../models/llama_index_embedding_model.py | 2 +- .../models/llama_index_generation_model.py | 2 +- .../models/llama_index_rank_model.py | 2 +- .../service => core/operation}/__init__.py | 0 .../operation/backend_operation.py | 6 +- .../operation/base_operation.py | 0 .../operation/base_workflow.py | 10 +- .../operation/consolidate_memory_op.py} | 6 +- .../operation/frontend_operation.py | 4 +- .../worker => core/service}/__init__.py | 0 .../service/base_memory_service.py | 6 +- .../service/memory_scope_service.py | 6 +- .../backend => core/storage}/__init__.py | 0 .../{ => core}/storage/base_memory_store.py | 0 .../{ => core}/storage/base_monitor.py | 0 .../{ => core}/storage/dummy_memory_store.py | 4 +- .../{ => core}/storage/dummy_monitor.py | 2 +- .../storage/llama_index_es_memory_store.py | 25 ++- .../storage/llama_index_sync_elasticsearch.py | 26 +-- .../frontend => core/utils}/__init__.py | 0 .../{ => core}/utils/datetime_handler.py | 12 +- memoryscope/{ => core}/utils/logger.py | 0 .../{ => core}/utils/prompt_handler.py | 0 memoryscope/{ => core}/utils/registry.py | 0 .../{ => core}/utils/response_text_parser.py | 2 +- memoryscope/{ => core}/utils/timer.py | 4 +- .../{ => core}/utils/tool_functions.py | 0 .../{models => core/worker}/__init__.py | 0 .../worker/backend}/__init__.py | 0 .../worker/backend/contra_repeat_worker.py | 6 +- .../worker/backend/contra_repeat_worker.yaml | 0 .../get_observation_with_time_worker.py | 6 +- .../get_observation_with_time_worker.yaml | 0 .../worker/backend/get_observation_worker.py | 8 +- .../backend/get_observation_worker.yaml | 0 .../backend/get_reflection_subject_worker.py | 8 +- .../get_reflection_subject_worker.yaml | 0 .../worker/backend/info_filter_worker.py | 6 +- .../worker/backend/info_filter_worker.yaml | 0 .../worker/backend/load_memory_worker.py | 6 +- .../backend/long_contra_repeat_worker.py | 6 +- .../backend/long_contra_repeat_worker.yaml | 0 .../worker/backend/update_insight_worker.py | 8 +- .../worker/backend/update_insight_worker.yaml | 0 .../worker/backend/update_memory_worker.py | 4 +- .../{memory => core}/worker/base_worker.py | 4 +- .../{memory => core}/worker/dummy_worker.py | 2 +- .../worker/frontend}/__init__.py | 0 .../worker/frontend/extract_time_worker.py | 6 +- .../worker/frontend/extract_time_worker.yaml | 0 .../worker/frontend/fuse_rerank_worker.py | 4 +- .../worker/frontend/print_memory_worker.py | 4 +- .../worker/frontend/print_memory_worker.yaml | 0 .../worker/frontend/read_message_worker.py | 2 +- .../worker/frontend/retrieve_memory_worker.py | 7 +- .../worker/frontend/semantic_rank_worker.py | 29 +-- .../worker/frontend/set_query_worker.py | 2 +- .../worker/memory_base_worker.py | 14 +- .../{memory => core}/worker/memory_manager.py | 6 +- memoryscope/memoryscope.py | 201 ------------------ tests/operations/test_interface.py | 2 +- tests/other/test_cli.py | 14 ++ tests/storages/test_storages_lli_synces.py | 6 +- tests/worker/test_workers_cn.py | 2 +- tests/worker/test_workers_en.py | 2 +- 85 files changed, 538 insertions(+), 625 deletions(-) rename {config => examples/config}/demo_config_no_stream.yaml (100%) delete mode 100644 memoryscope/argument/default_arguments.py rename memoryscope/{argument => core}/__init__.py (100%) rename memoryscope/{ => core}/chat/__init__.py (100%) rename memoryscope/{ => core}/chat/api_memory_chat.py (57%) rename memoryscope/{ => core}/chat/base_memory_chat.py (66%) rename memoryscope/{ => core}/chat/cli_memory_chat.py (93%) rename memoryscope/{ => core}/chat/memory_chat_prompt.yaml (100%) rename memoryscope/{memory => core/config}/__init__.py (100%) rename memoryscope/{argument/memoryscope_arguments.py => core/config/arguments.py} (82%) create mode 100644 memoryscope/core/config/config_manager.py rename memoryscope/{argument/cli_chat_demo.yaml => core/config/demo_config.yaml} (97%) create mode 100644 memoryscope/core/memoryscope.py rename memoryscope/{ => core}/memoryscope_context.py (100%) rename memoryscope/{memory/operation => core/models}/__init__.py (100%) rename memoryscope/{ => core}/models/base_model.py (96%) rename memoryscope/{ => core}/models/dummy_generation_model.py (98%) rename memoryscope/{ => core}/models/llama_index_embedding_model.py (97%) rename memoryscope/{ => core}/models/llama_index_generation_model.py (98%) rename memoryscope/{ => core}/models/llama_index_rank_model.py (98%) rename memoryscope/{memory/service => core/operation}/__init__.py (100%) rename memoryscope/{memory => core}/operation/backend_operation.py (95%) rename memoryscope/{memory => core}/operation/base_operation.py (100%) rename memoryscope/{memory => core}/operation/base_workflow.py (96%) rename memoryscope/{memory/operation/consolidate_operation.py => core/operation/consolidate_memory_op.py} (92%) rename memoryscope/{memory => core}/operation/frontend_operation.py (91%) rename memoryscope/{memory/worker => core/service}/__init__.py (100%) rename memoryscope/{memory => core}/service/base_memory_service.py (94%) rename memoryscope/{memory => core}/service/memory_scope_service.py (95%) rename memoryscope/{memory/worker/backend => core/storage}/__init__.py (100%) rename memoryscope/{ => core}/storage/base_memory_store.py (100%) rename memoryscope/{ => core}/storage/base_monitor.py (100%) rename memoryscope/{ => core}/storage/dummy_memory_store.py (92%) rename memoryscope/{ => core}/storage/dummy_monitor.py (92%) rename memoryscope/{ => core}/storage/llama_index_es_memory_store.py (87%) rename memoryscope/{ => core}/storage/llama_index_sync_elasticsearch.py (98%) rename memoryscope/{memory/worker/frontend => core/utils}/__init__.py (100%) rename memoryscope/{ => core}/utils/datetime_handler.py (96%) rename memoryscope/{ => core}/utils/logger.py (100%) rename memoryscope/{ => core}/utils/prompt_handler.py (100%) rename memoryscope/{ => core}/utils/registry.py (100%) rename memoryscope/{ => core}/utils/response_text_parser.py (97%) rename memoryscope/{ => core}/utils/timer.py (97%) rename memoryscope/{ => core}/utils/tool_functions.py (100%) rename memoryscope/{models => core/worker}/__init__.py (100%) rename memoryscope/{storage => core/worker/backend}/__init__.py (100%) rename memoryscope/{memory => core}/worker/backend/contra_repeat_worker.py (96%) rename memoryscope/{memory => core}/worker/backend/contra_repeat_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/get_observation_with_time_worker.py (94%) rename memoryscope/{memory => core}/worker/backend/get_observation_with_time_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/get_observation_worker.py (96%) rename memoryscope/{memory => core}/worker/backend/get_observation_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/get_reflection_subject_worker.py (94%) rename memoryscope/{memory => core}/worker/backend/get_reflection_subject_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/info_filter_worker.py (95%) rename memoryscope/{memory => core}/worker/backend/info_filter_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/load_memory_worker.py (96%) rename memoryscope/{memory => core}/worker/backend/long_contra_repeat_worker.py (97%) rename memoryscope/{memory => core}/worker/backend/long_contra_repeat_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/update_insight_worker.py (97%) rename memoryscope/{memory => core}/worker/backend/update_insight_worker.yaml (100%) rename memoryscope/{memory => core}/worker/backend/update_memory_worker.py (96%) rename memoryscope/{memory => core}/worker/base_worker.py (98%) rename memoryscope/{memory => core}/worker/dummy_worker.py (92%) rename memoryscope/{utils => core/worker/frontend}/__init__.py (100%) rename memoryscope/{memory => core}/worker/frontend/extract_time_worker.py (93%) rename memoryscope/{memory => core}/worker/frontend/extract_time_worker.yaml (100%) rename memoryscope/{memory => core}/worker/frontend/fuse_rerank_worker.py (97%) rename memoryscope/{memory => core}/worker/frontend/print_memory_worker.py (94%) rename memoryscope/{memory => core}/worker/frontend/print_memory_worker.yaml (100%) rename memoryscope/{memory => core}/worker/frontend/read_message_worker.py (91%) rename memoryscope/{memory => core}/worker/frontend/retrieve_memory_worker.py (96%) rename memoryscope/{memory => core}/worker/frontend/semantic_rank_worker.py (70%) rename memoryscope/{memory => core}/worker/frontend/set_query_worker.py (97%) rename memoryscope/{memory => core}/worker/memory_base_worker.py (94%) rename memoryscope/{memory => core}/worker/memory_manager.py (97%) delete mode 100644 memoryscope/memoryscope.py create mode 100644 tests/other/test_cli.py diff --git a/config/demo_config_no_stream.yaml b/examples/config/demo_config_no_stream.yaml similarity index 100% rename from config/demo_config_no_stream.yaml rename to examples/config/demo_config_no_stream.yaml diff --git a/memoryscope/__init__.py b/memoryscope/__init__.py index d8b7815a..2be4eb6a 100644 --- a/memoryscope/__init__.py +++ b/memoryscope/__init__.py @@ -1,3 +1,2 @@ """ Version of MemoryScope.""" - -__version__ = "0.1.0-alpha.1" +__version__ = "0.1.0" diff --git a/memoryscope/argument/default_arguments.py b/memoryscope/argument/default_arguments.py deleted file mode 100644 index 275285fd..00000000 --- a/memoryscope/argument/default_arguments.py +++ /dev/null @@ -1,182 +0,0 @@ -DEFAULT_GLOBAL_ARGUMENTS = { - "language": "en", - "thread_pool_max_workers": 5, - "logger_name": "memoryscope", - "logger_name_time_suffix": "%Y%m%d_%H%M%S" -} - -DEFAULT_MEMORY_CHAT_ARGUMENTS = { - "cli_memory_chat": { - "class": "chat.cli_memory_chat", - "memory_service": "memoryscope_service", - "generation_model": "generation_model" - } -} - -DEFAULT_MEMORY_SERVICE_ARGUMENTS = { - "memoryscope_service": { - "class": "memory.service.memory_scope_service", - "memory_operations": { - "read_message": { - "class": "memory.operation.frontend_operation", - "workflow": "read_message", - "description": "read short memory" - }, - "retrieve_memory": { - "class": "memory.operation.frontend_operation", - "workflow": "set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank", - "description": "retrieve long-term memory" - }, - "list_memory": { - "class": "memory.operation.frontend_operation", - "workflow": "set_query,retrieve_top_memory,print_memory", - "description": "read all long-term memory of the user" - }, - "delete_memory": { - "class": "memory.operation.frontend_operation", - "workflow": "set_query,retrieve_all_memory,delete_memory", - "description": "delete a single long-term memory" - }, - "delete_all": { - "class": "memory.operation.frontend_operation", - "workflow": "set_query,retrieve_all_memory,delete_all", - "description": "delete all long-term memory" - }, - "add_memory": { - "class": "memory.operation.frontend_operation", - "workflow": "add_memory", - "description": "add a single observation" - }, - "consolidate_memory": { - "class": "memory.operation.consolidate_memory_op", - "workflow": "info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat," - "store_memory", - "description": "summary user's observation memory", - "interval_time": 1 - }, - "reflect_and_reconsolidate": { - "class": "memory.operation.backend_operation", - "workflow": "load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat," - "store_memory", - "description": "summary user's insight memory", - "interval_time": 15 - } - } - } -} - -DEFAULT_WORKER_ARGUMENTS = { - "dummy": { - "class": "memory.worker.dummy_worker", - "generation_model": "generation_model", - "embedding_model": "embedding_model", - "rank_model": "rank_model" - }, - "read_message": { - "class": "memory.worker.frontend.read_message_worker" - }, - "set_query": { - "class": "memory.worker.frontend.set_query_worker" - }, - "retrieve_obs_ins": { - "class": "memory.worker.frontend.retrieve_memory_worker", - "retrieve_obs_top_k": 100, - "retrieve_ins_top_k": 100 - }, - "extract_time": { - "class": "memory.worker.frontend.extract_time_worker", - "generation_model": "generation_model" - }, - "semantic_rank": { - "class": "memory.worker.frontend.semantic_rank_worker", - "rank_model": "rank_model" - }, - "fuse_rerank": { - "class": "memory.worker.frontend.fuse_rerank_worker", - "fuse_score_threshold": 0.01, - "fuse_ratio_dict": { - "conversation": 0.5, - "observation": 1, - "obs_customized": 1.2, - "insight": 2 - }, - "fuse_time_ratio": 2, - "fuse_rerank_top_k": 10 - }, - "retrieve_top_memory": { - "class": "memory.worker.frontend.retrieve_memory_worker", - "retrieve_obs_top_k": 100, - "retrieve_ins_top_k": 100, - "retrieve_expired_top_k": 100 - }, - "print_memory": { - "class": "memory.worker.frontend.print_memory_worker" - }, - "retrieve_all_memory": { - "class": "memory.worker.frontend.retrieve_memory_worker", - "retrieve_obs_top_k": 1000, - "retrieve_ins_top_k": 1000, - "retrieve_expired_top_k": 1000 - }, - "delete_memory": { - "class": "memory.worker.backend.update_memory_worker", - "method": "delete_memory" - }, - "delete_all": { - "class": "memory.worker.backend.update_memory_worker", - "method": "delete_all" - }, - "add_memory": { - "class": "memory.worker.backend.update_memory_worker", - "method": "from_query" - }, - "info_filter": { - "class": "memory.worker.backend.info_filter_worker", - "generation_model": "generation_model" - }, - "load_today_memory": { - "class": "memory.worker.backend.load_memory_worker", - "retrieve_today_top_k": 100 - }, - "get_observation": { - "class": "memory.worker.backend.get_observation_worker", - "generation_model": "generation_model" - }, - "get_observation_with_time": { - "class": "memory.worker.backend.get_observation_with_time_worker", - "generation_model": "generation_model" - }, - "contra_repeat": { - "class": "memory.worker.backend.contra_repeat_worker", - "generation_model": "generation_model" - }, - "store_memory": { - "class": "memory.worker.backend.update_memory_worker", - "method": "from_memory_key", - "memory_key": "all" - }, - "load_obs_and_insight": { - "class": "memory.worker.backend.load_memory_worker", - "retrieve_not_reflected_top_k": 100, - "retrieve_not_updated_top_k": 100, - "retrieve_insight_top_k": 100 - }, - "get_reflection_subject": { - "class": "memory.worker.backend.get_reflection_subject_worker", - "generation_model": "generation_model", - "reflect_obs_cnt_threshold": 10 - }, - "update_insight": { - "class": "memory.worker.backend.update_insight_worker", - "generation_model": "generation_model", - "rank_model": "rank_model" - }, - "long_contra_repeat": { - "class": "memory.worker.backend.long_contra_repeat_worker", - "generation_model": "generation_model" - } -} - -DEFAULT_MONITOR_ARGUMENTS = { - "class": "storage.dummy_monitor" -} diff --git a/memoryscope/cli.py b/memoryscope/cli.py index 55b05f5b..1d890358 100644 --- a/memoryscope/cli.py +++ b/memoryscope/cli.py @@ -1,17 +1,15 @@ import sys -from memoryscope.chat.base_memory_chat import BaseMemoryChat -from memoryscope.memoryscope import MemoryScope - sys.path.append(".") # noqa: E402 import fire +from memoryscope.core.memoryscope import MemoryScope -def cli_job(config_path: str): - ms = MemoryScope(config_path=config_path) - memory_chat: BaseMemoryChat = ms.default_memory_chat - memory_chat.run() + +def cli_job(**kwargs): + kwargs["memory_chat_type"] = "cli_chat" + MemoryScope(**kwargs).default_memory_chat.run() if __name__ == "__main__": diff --git a/memoryscope/constants/common_constants.py b/memoryscope/constants/common_constants.py index 4eba1d1d..df217846 100644 --- a/memoryscope/constants/common_constants.py +++ b/memoryscope/constants/common_constants.py @@ -9,6 +9,8 @@ MEMORYSCOPE_CONTEXT = "memoryscope_context" RESULT = "result" +MEMORIES = "memories" + CHAT_MESSAGES = "chat_messages" MEMORY_MANAGER = "memory_manager" diff --git a/memoryscope/argument/__init__.py b/memoryscope/core/__init__.py similarity index 100% rename from memoryscope/argument/__init__.py rename to memoryscope/core/__init__.py diff --git a/memoryscope/chat/__init__.py b/memoryscope/core/chat/__init__.py similarity index 100% rename from memoryscope/chat/__init__.py rename to memoryscope/core/chat/__init__.py diff --git a/memoryscope/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py similarity index 57% rename from memoryscope/chat/api_memory_chat.py rename to memoryscope/core/chat/api_memory_chat.py index 10bc741d..d1992d73 100644 --- a/memoryscope/chat/api_memory_chat.py +++ b/memoryscope/core/chat/api_memory_chat.py @@ -1,14 +1,15 @@ -from typing import List +from typing import List, Optional -from memoryscope.chat.base_memory_chat import BaseMemoryChat +from memoryscope.constants.common_constants import MEMORIES from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.core.chat.base_memory_chat import BaseMemoryChat +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.prompt_handler import PromptHandler from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.memoryscope_context import MemoryscopeContext -from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message -from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.prompt_handler import PromptHandler +from memoryscope.scheme.model_response import ModelResponse class ApiMemoryChat(BaseMemoryChat): @@ -96,50 +97,79 @@ class ApiMemoryChat(BaseMemoryChat): self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model - def get_new_message(self, query: str, role_name: str = "") -> Message: - if not role_name: - role_name = self.human_name - return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) - - def get_system_message_with_memory(self, memories: str) -> Message: - # Incorporate memory into the system prompt if available - system_prompt = self.prompt_handler.system_prompt - if memories: - memory_prompt = self.prompt_handler.memory_prompt - system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) - return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) - def chat_with_memory(self, query: str, - role_name: str = "", - remember_response: bool = True): - + role_name: Optional[str] = None, + system_prompt: Optional[str] = None, + memory_prompt: Optional[str] = None, + extra_memories: Optional[str] = None, + add_not_memorized_messages: bool = True, + remember_response: bool = True, + **kwargs): + """ + The core function that carries out conversation with memory accepts user queries through query and returns the + conversation results through model_response. The retrieved memories are stored in the memories within meta_data. + Args: + query (str, optional): User's query, includes the user's question. + role_name (str, optional): User's role name. + system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml". + memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml". + extra_memories (str, optional): Manually added user memory in this function. + add_not_memorized_messages (bool, optional): whether add not memorized messages to LLM. + remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. + Defaults to False. + Returns: + - ModelResponse: In non-streaming mode, returns a complete AI response. + - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. + - Memories: To obtain the memory by invoking the method of model_response.meta_data[MEMORIES] + """ chat_messages: List[Message] = [] - new_message: Message = self.get_new_message(query=query, role_name=role_name) + # prepare query message + if not role_name: + role_name = self.human_name + query_message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) - # To retrieve memory, prepare the query timestamp and role name by adding new_message. - memories: str = self.memory_service.retrieve_memory(query=new_message.content, - role_name=new_message.role_name, - timestamp=new_message.time_created) + # To retrieve memory, prepare the query timestamp and role name by adding query_message. + memories: str = self.memory_service.retrieve_memory(query=query_message.content, + role_name=query_message.role_name, + timestamp=query_message.time_created) # format system_message with memories - system_message: Message = self.get_system_message_with_memory(memories=memories) + system_prompt_list = [] + if system_prompt: + system_prompt_list.append(system_prompt) + else: + system_prompt_list.append(self.prompt_handler.system_prompt) + + if memories: + # add memory prompt + if memory_prompt: + system_prompt_list.append(memory_prompt) + else: + system_prompt_list.append(self.prompt_handler.memory_prompt) + system_prompt_list.append(memories) + + if extra_memories: + system_prompt_list.extend(extra_memories) + + system_prompt_join = "\n".join([x.strip() for x in system_prompt_list]) + system_message = Message(role=MessageRoleEnum.SYSTEM, content=system_prompt_join) chat_messages.append(system_message) # Include past conversation history in the message list - history_messages = self.memory_service.read_message() - if history_messages: - chat_messages.extend(history_messages) + if add_not_memorized_messages: + history_messages = self.memory_service.read_message() + if history_messages: + chat_messages.extend(history_messages) # Append the current user's message to the conversation context - chat_messages.append(new_message) + chat_messages.append(query_message) self.logger.info(f"chat_messages={chat_messages}") resp = self.generation_model.call(messages=chat_messages, stream=self.stream, **self.generation_model_kwargs) if self.stream: - assert isinstance(resp, ModelResponseGen) model_response: ModelResponse | None = None for model_response in resp: yield model_response @@ -147,17 +177,18 @@ class ApiMemoryChat(BaseMemoryChat): if remember_response: if model_response and model_response.message: model_response.message.role_name = self.assistant_name - self.memory_service.add_messages([new_message, model_response.message]) + model_response.meta_data[MEMORIES] = memories + self.memory_service.add_messages([query_message, model_response.message]) else: - self.logger.info("model_response or model_response.message is empty!") + self.logger.warning("model_response or model_response.message is empty!") else: - assert isinstance(resp, ModelResponse) model_response: ModelResponse = resp if remember_response: if model_response and model_response.message: model_response.message.role_name = self.assistant_name - self.memory_service.add_messages([new_message, model_response.message]) + model_response.meta_data[MEMORIES] = memories + self.memory_service.add_messages([query_message, model_response.message]) else: - self.logger.info("model_response or model_response.message is empty!") + self.logger.warning("model_response or model_response.message is empty!") return model_response diff --git a/memoryscope/chat/base_memory_chat.py b/memoryscope/core/chat/base_memory_chat.py similarity index 66% rename from memoryscope/chat/base_memory_chat.py rename to memoryscope/core/chat/base_memory_chat.py index ea81c4fb..995d1d46 100644 --- a/memoryscope/chat/base_memory_chat.py +++ b/memoryscope/core/chat/base_memory_chat.py @@ -1,9 +1,7 @@ from abc import ABCMeta, abstractmethod -from typing import List -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.scheme.message import Message -from memoryscope.utils.logger import Logger +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.logger import Logger class BaseMemoryChat(metaclass=ABCMeta): @@ -17,12 +15,14 @@ class BaseMemoryChat(metaclass=ABCMeta): self.kwargs: dict = kwargs self.logger = Logger.get_logger() - @abstractmethod - def get_new_message(self, query: str, role_name: str = "") -> Message: - raise NotImplementedError + @property + def memory_service(self) -> BaseMemoryService: + """ + Abstract property to access the memory service. - @abstractmethod - def get_system_message_with_memory(self, memories: str) -> Message: + Raises: + NotImplementedError: This method should be implemented in a subclass. + """ raise NotImplementedError @abstractmethod @@ -40,25 +40,6 @@ class BaseMemoryChat(metaclass=ABCMeta): """ raise NotImplementedError - @property - def memory_service(self) -> BaseMemoryService: - """ - Abstract property to access the memory service. - - Raises: - NotImplementedError: This method should be implemented in a subclass. - """ - raise NotImplementedError - - def add_messages(self, messages: List[Message] | Message): - self.memory_service.add_messages(messages) - - def start_backend_service(self): - self.memory_service.start_backend_service() - - def do_memory_operation(self, operation_name: str, **kwargs): - return self.memory_service.do_operation(name=operation_name, **kwargs) - def run(self): """ Abstract method to run the chat system. diff --git a/memoryscope/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py similarity index 93% rename from memoryscope/chat/cli_memory_chat.py rename to memoryscope/core/chat/cli_memory_chat.py index 27a81485..f2fe8a0a 100644 --- a/memoryscope/chat/cli_memory_chat.py +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -4,16 +4,16 @@ from typing import List import questionary -from memoryscope.chat.base_memory_chat import BaseMemoryChat from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.core.chat.base_memory_chat import BaseMemoryChat +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.prompt_handler import PromptHandler +from memoryscope.core.utils.tool_functions import char_logo from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.memoryscope_context import MemoryscopeContext -from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message -from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.prompt_handler import PromptHandler -from memoryscope.utils.tool_functions import char_logo +from memoryscope.scheme.model_response import ModelResponse class CliMemoryChat(BaseMemoryChat): @@ -122,7 +122,7 @@ class CliMemoryChat(BaseMemoryChat): self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model - def get_new_message(self, query: str, role_name: str = "") -> Message: + def get_user_message(self, query: str, role_name: str = "") -> Message: if not role_name: role_name = self.human_name return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) @@ -142,7 +142,7 @@ class CliMemoryChat(BaseMemoryChat): chat_messages: List[Message] = [] - new_message: Message = self.get_new_message(query=query, role_name=role_name) + new_message: Message = self.get_user_message(query=query, role_name=role_name) # To retrieve memory, prepare the query timestamp and role name by adding new_message. memories: str = self.memory_service.retrieve_memory(query=new_message.content, @@ -168,18 +168,11 @@ class CliMemoryChat(BaseMemoryChat): **self.generation_model_kwargs) if self.stream: - assert isinstance(resp, ModelResponseGen) model_response: ModelResponse | None = None for model_response in resp: questionary.print(model_response.delta, end="") questionary.print("") - - if remember_response and model_response and model_response.message: - model_response.message.role_name = self.assistant_name - self.memory_service.add_messages([new_message, model_response.message]) - else: - assert isinstance(resp, ModelResponse) model_response: ModelResponse = resp questionary.print(model_response.message.content) @@ -315,7 +308,7 @@ class CliMemoryChat(BaseMemoryChat): questionary.print(f"{self.assistant_name}: ", end="", style="bold") # Fetch and display AI's response - self.start_backend_service() + self.memory_service.start_backend_service() self.chat_with_memory(query=query) except KeyboardInterrupt: diff --git a/memoryscope/chat/memory_chat_prompt.yaml b/memoryscope/core/chat/memory_chat_prompt.yaml similarity index 100% rename from memoryscope/chat/memory_chat_prompt.yaml rename to memoryscope/core/chat/memory_chat_prompt.yaml diff --git a/memoryscope/memory/__init__.py b/memoryscope/core/config/__init__.py similarity index 100% rename from memoryscope/memory/__init__.py rename to memoryscope/core/config/__init__.py diff --git a/memoryscope/argument/memoryscope_arguments.py b/memoryscope/core/config/arguments.py similarity index 82% rename from memoryscope/argument/memoryscope_arguments.py rename to memoryscope/core/config/arguments.py index 7f48d281..1925e31e 100644 --- a/memoryscope/argument/memoryscope_arguments.py +++ b/memoryscope/core/config/arguments.py @@ -3,7 +3,7 @@ from typing import Literal, Dict @dataclass -class MemoryscopeArguments(object): +class Arguments(object): language: Literal["cn", "en"] = field(default="en", metadata={"help": "support en & cn now"}) thread_pool_max_workers: int = field(default=5, metadata={"help": "thread pool max workers"}) @@ -12,12 +12,8 @@ class MemoryscopeArguments(object): logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S") - memory_chat_class: str = field(default="chat.api_memory_chat", metadata={ - "help": "The memory chat class for dynamic import: chat.cli_memory_chat, chat.api_memory_chat"}) - - human_name: str = field(default="user", metadata={"help": "en: user, cn: 用户"}) - - assistant_name: str = field(default="AI") + memory_chat_type: str = field(default="cli_chat", metadata={ + "help": "cli_chat(Command-line interaction), api_chat(API interface interaction), etc."}) consolidate_memory_interval_time: int = field(default=1, metadata={ "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) @@ -40,7 +36,7 @@ class MemoryscopeArguments(object): embedding_backend: str = field(default="openai_embedding", metadata={ "help": "global embedding backend: openai_embedding, dashscope_embedding, etc."}) - embedding_model: str = field(default="gpt-4o", metadata={ + embedding_model: str = field(default="text-embedding-ada-002", metadata={ "help": "global embedding model: text-embedding-ada-002, text-embedding-v2, etc."}) embedding_params: dict = field(default_factory=lambda: {}) @@ -61,6 +57,3 @@ class MemoryscopeArguments(object): retrieve_mode: str = field(default="dense", metadata={ "help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"}) - - hybrid_alpha: float | None = field(default=1.0, metadata={ - "help": "fuse alpha params used in hybrid mode(not implemented)"}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py new file mode 100644 index 00000000..7ac8095c --- /dev/null +++ b/memoryscope/core/config/config_manager.py @@ -0,0 +1,170 @@ +import json +from dataclasses import fields +from pathlib import Path +from typing import Optional, Literal + +import yaml + +from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.core.config.arguments import Arguments +from memoryscope.enumeration.language_enum import LanguageEnum + + +class ConfigManager(object): + + def __init__(self, + config: dict = None, + config_path: Optional[str] = None, + arguments: Optional[Arguments] = None, + demo_config_name: str = "demo_config.yaml", + **kwargs): + self.config: dict = {} + self.kwargs = kwargs + + if config: + self.config = config + + elif config_path: + self.read_config(config_path) + + else: + self.read_demo_config(demo_config_name) + + if arguments: + self.update_config_by_arguments(arguments) + + elif kwargs: + key_list = [x.name for x in fields(Arguments)] + arguments = Arguments(**{k: v for k, v in kwargs.items() if k in key_list}) + self.update_config_by_arguments(arguments) + + def read_config(self, config_path: str): + if config_path.endswith(".yaml"): + with open(config_path) as f: + self.config = yaml.load(f, yaml.FullLoader) + + elif config_path.endswith(".json"): + with open(config_path) as f: + self.config = json.load(f) + + def read_demo_config(self, demo_config_name: str): + file_path = Path(__file__) + demo_config_path = (file_path.parent / demo_config_name).__str__() + with open(demo_config_path) as f: + self.config = yaml.load(f, yaml.FullLoader) + + @staticmethod + def update_global_by_arguments(config: dict, arguments: Arguments): + config.update({ + "language": arguments.language, + "thread_pool_max_workers": arguments.thread_pool_max_workers, + "logger_name": arguments.logger_name, + "logger_name_time_suffix": arguments.logger_name_time_suffix, + "use_dummy_ranker": arguments.use_dummy_ranker, + }) + + @staticmethod + def update_memory_chat_by_arguments(config: dict, arguments: Arguments): + if arguments.memory_chat_type == "cli_chat": + memory_chat_class = "chat.cli_memory_chat" + elif arguments.memory_chat_type == "api_chat": + memory_chat_class = "chat.api_memory_chat" + else: + raise NotImplementedError(f"known memory_chat_type={arguments.memory_chat_type}") + config.update({ + "class": memory_chat_class, + "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], + "assistant_name": "AI", + }) + + @staticmethod + def update_memory_service_by_arguments(config: dict, arguments: Arguments): + config.update({ + "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], + "assistant_name": "AI", + }) + config["memory_operations"]["consolidate_memory"]["interval_time"] = \ + arguments.consolidate_memory_interval_time + config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \ + arguments.reflect_and_reconsolidate_interval_time + + @staticmethod + def update_worker_by_arguments(config: dict, arguments: Arguments): + for worker_name, kv_dict in arguments.worker_params.items(): + if worker_name not in config: + continue + config[worker_name].update(kv_dict) + + @staticmethod + def update_model_by_arguments(config: dict, arguments: Arguments): + config["generation_model"].update({ + "module_name": arguments.generation_backend, + "model_name": arguments.generation_model, + **arguments.generation_params, + }) + + config["embedding_model"].update({ + "module_name": arguments.embedding_backend, + "model_name": arguments.embedding_model, + **arguments.embedding_params, + }) + + config["rank_model"].update({ + "module_name": arguments.rank_backend, + "model_name": arguments.rank_model, + **arguments.rank_params, + }) + + @staticmethod + def update_memory_store_by_arguments(config: dict, arguments: Arguments): + config.update({ + "index_name": arguments.es_index_name, + "es_url": arguments.es_url, + "retrieve_mode": arguments.retrieve_mode}) + + def update_config_by_arguments(self, arguments: Arguments): + # prepare global + self.update_global_by_arguments(self.config["global"], arguments) + + # prepare memory chat + memory_chat_conf_dict = self.config["memory_chat"] + memory_chat_config = list(memory_chat_conf_dict.values())[0] + self.update_memory_chat_by_arguments(memory_chat_config, arguments) + + # prepare memory service + memory_service_conf_dict = self.config["memory_service"] + memory_service_config = list(memory_service_conf_dict.values())[0] + self.update_memory_service_by_arguments(memory_service_config, arguments) + + # prepare worker + self.update_worker_by_arguments(self.config["worker"], arguments) + + # prepare model + self.update_model_by_arguments(self.config["model"], arguments) + + # prepare memory store + self.update_memory_store_by_arguments(self.config["memory_store"], arguments) + + def add_node_object(self, node: str, name: str, config: dict): + self.config[node][name] = config + + def pop_node_object(self, node: str, name: str): + return self.config[node].pop(name, None) + + def clear_node_all(self, node: str): + self.config[node].clear() + + def dump_config(self, file_type: Literal["json", "yaml"], to_stream: bool = True, file_path: Optional[str] = None): + if file_type == "json": + content = json.dumps(self.config, indent=2, ensure_ascii=False) + elif file_type == "yaml": + content = yaml.dump(self.config, indent=2, allow_unicode=True) + else: + raise NotImplementedError + + if to_stream: + print(content) + + if file_type: + with open(file_path, "w") as f: + f.write(content) diff --git a/memoryscope/argument/cli_chat_demo.yaml b/memoryscope/core/config/demo_config.yaml similarity index 97% rename from memoryscope/argument/cli_chat_demo.yaml rename to memoryscope/core/config/demo_config.yaml index 16d7f92d..e81ba428 100644 --- a/memoryscope/argument/cli_chat_demo.yaml +++ b/memoryscope/core/config/demo_config.yaml @@ -1,9 +1,9 @@ -global_config: +global: language: en thread_pool_max_workers: 5 logger_name: memoryscope - logger_name_time_suffix: %Y%m%d_%H%M%S - use_dummy_ranker: true + logger_name_time_suffix: "%Y%m%d_%H%M%S" + use_dummy_ranker: false memory_chat: cli_memory_chat: @@ -169,8 +169,7 @@ memory_store: embedding_model: embedding_model index_name: memory_index es_url: http://localhost:9200 - retrieve_type: dense - hybrid_alpha: 1.0 + retrieve_mode: dense monitor: class: storage.dummy_monitor \ No newline at end of file diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py new file mode 100644 index 00000000..3e0c7978 --- /dev/null +++ b/memoryscope/core/memoryscope.py @@ -0,0 +1,105 @@ +import datetime +from concurrent.futures import ThreadPoolExecutor + +from memoryscope.core.chat.base_memory_chat import BaseMemoryChat +from memoryscope.core.config.config_manager import ConfigManager +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.enumeration.language_enum import LanguageEnum +from memoryscope.enumeration.model_enum import ModelEnum + + +class MemoryScope(ConfigManager): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + + self.logger = self._init_logger() + + self.context: MemoryscopeContext = MemoryscopeContext() + self.init_context_by_config() + + def _init_logger(self) -> Logger: + global_config = self.config["global"] + logger_name = global_config["logger_name"] + logger_name_time_suffix = global_config["logger_name_time_suffix"] + if logger_name_time_suffix: + suffix = datetime.datetime.now().strftime(logger_name_time_suffix) + logger_name = f"{logger_name}_{suffix}" + return Logger.get_logger(logger_name, to_stream=False) + + def init_context_by_config(self): + # set global config + global_conf = self.config["global"] + self.context.language = LanguageEnum(global_conf["language"]) + self.context.thread_pool = ThreadPoolExecutor(max_workers=global_conf["thread_pool_max_workers"]) + self.context.meta_data["use_dummy_ranker"] = global_conf["use_dummy_ranker"] + + # init memory_chat + memory_chat_conf_dict = self.config["memory_chat"] + if memory_chat_conf_dict: + for name, conf in memory_chat_conf_dict.items(): + self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context) + + # set memory_service + memory_service_conf_dict = self.config["memory_service"] + assert memory_service_conf_dict + for name, conf in memory_service_conf_dict.items(): + self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context) + + # init model + model_conf_dict = self.config["model"] + assert model_conf_dict + for name, conf in model_conf_dict.items(): + self.context.model_dict[name] = init_instance_by_config(conf, name=name) + + # init memory_store + memory_store_conf = self.config["memory_store"] + assert memory_store_conf + emb_model_name: str = memory_store_conf[ModelEnum.EMBEDDING_MODEL.value] + embedding_model = self.context.model_dict[emb_model_name] + self.context.memory_store = init_instance_by_config(memory_store_conf, embedding_model=embedding_model) + + # init monitor + monitor_conf = self.config["monitor"] + if monitor_conf: + self.context.monitor = init_instance_by_config(monitor_conf) + + # set worker config + self.context.worker_conf_dict = self.config["worker"] + + def close(self): + # wait service to stop + for _, service in self.context.memory_service_dict.items(): + service.stop_backend_service(wait_service_end=True) + + self.context.thread_pool.shutdown() + + self.context.memory_store.close() + + if self.context.monitor: + self.context.monitor.close() + + def __enter__(self): + self.init_context_by_config() + + def __exit__(self, exc_type, exc_val, exc_tb): + self.close() + + @property + def memory_chat_dict(self): + return self.context.memory_chat_dict + + @property + def memory_service_dict(self): + return self.context.memory_service_dict + + @property + def default_memory_chat(self) -> BaseMemoryChat: + return list(self.memory_chat_dict.values())[0] + + @property + def default_service(self) -> BaseMemoryService: + return list(self.memory_service_dict.values())[0] diff --git a/memoryscope/memoryscope_context.py b/memoryscope/core/memoryscope_context.py similarity index 100% rename from memoryscope/memoryscope_context.py rename to memoryscope/core/memoryscope_context.py diff --git a/memoryscope/memory/operation/__init__.py b/memoryscope/core/models/__init__.py similarity index 100% rename from memoryscope/memory/operation/__init__.py rename to memoryscope/core/models/__init__.py diff --git a/memoryscope/models/base_model.py b/memoryscope/core/models/base_model.py similarity index 96% rename from memoryscope/models/base_model.py rename to memoryscope/core/models/base_model.py index 83b5aa50..07cbae69 100644 --- a/memoryscope/models/base_model.py +++ b/memoryscope/core/models/base_model.py @@ -3,11 +3,11 @@ import time from abc import abstractmethod, ABCMeta from typing import Any +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.registry import Registry +from memoryscope.core.utils.timer import Timer from memoryscope.enumeration.model_enum import ModelEnum from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.logger import Logger -from memoryscope.utils.registry import Registry -from memoryscope.utils.timer import Timer MODEL_REGISTRY = Registry("models") diff --git a/memoryscope/models/dummy_generation_model.py b/memoryscope/core/models/dummy_generation_model.py similarity index 98% rename from memoryscope/models/dummy_generation_model.py rename to memoryscope/core/models/dummy_generation_model.py index d1ad0a54..e1d81438 100644 --- a/memoryscope/models/dummy_generation_model.py +++ b/memoryscope/core/models/dummy_generation_model.py @@ -3,9 +3,9 @@ from typing import List from llama_index.core.base.llms.types import ChatMessage +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen diff --git a/memoryscope/models/llama_index_embedding_model.py b/memoryscope/core/models/llama_index_embedding_model.py similarity index 97% rename from memoryscope/models/llama_index_embedding_model.py rename to memoryscope/core/models/llama_index_embedding_model.py index 822f619d..efde959b 100644 --- a/memoryscope/models/llama_index_embedding_model.py +++ b/memoryscope/core/models/llama_index_embedding_model.py @@ -2,8 +2,8 @@ from typing import List from llama_index.embeddings.dashscope import DashScopeEmbedding +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.model_response import ModelResponse diff --git a/memoryscope/models/llama_index_generation_model.py b/memoryscope/core/models/llama_index_generation_model.py similarity index 98% rename from memoryscope/models/llama_index_generation_model.py rename to memoryscope/core/models/llama_index_generation_model.py index 63a249d9..230bde67 100644 --- a/memoryscope/models/llama_index_generation_model.py +++ b/memoryscope/core/models/llama_index_generation_model.py @@ -3,9 +3,9 @@ from typing import List from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse from llama_index.llms.dashscope import DashScope +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen diff --git a/memoryscope/models/llama_index_rank_model.py b/memoryscope/core/models/llama_index_rank_model.py similarity index 98% rename from memoryscope/models/llama_index_rank_model.py rename to memoryscope/core/models/llama_index_rank_model.py index 38debfd7..9e54c545 100644 --- a/memoryscope/models/llama_index_rank_model.py +++ b/memoryscope/core/models/llama_index_rank_model.py @@ -4,8 +4,8 @@ from llama_index.core.data_structs import Node from llama_index.core.schema import NodeWithScore from llama_index.postprocessor.dashscope_rerank import DashScopeRerank +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.model_response import ModelResponse diff --git a/memoryscope/memory/service/__init__.py b/memoryscope/core/operation/__init__.py similarity index 100% rename from memoryscope/memory/service/__init__.py rename to memoryscope/core/operation/__init__.py diff --git a/memoryscope/memory/operation/backend_operation.py b/memoryscope/core/operation/backend_operation.py similarity index 95% rename from memoryscope/memory/operation/backend_operation.py rename to memoryscope/core/operation/backend_operation.py index 4a53a6b4..fd6a55a4 100644 --- a/memoryscope/memory/operation/backend_operation.py +++ b/memoryscope/core/operation/backend_operation.py @@ -2,10 +2,10 @@ import time from typing import List from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES -from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE -from memoryscope.memory.operation.base_workflow import BaseWorkflow +from memoryscope.core.operation.base_operation import BaseOperation, OPERATION_TYPE +from memoryscope.core.operation.base_workflow import BaseWorkflow +from memoryscope.core.utils.logger import Logger from memoryscope.scheme.message import Message -from memoryscope.utils.logger import Logger class BackendOperation(BaseWorkflow, BaseOperation): diff --git a/memoryscope/memory/operation/base_operation.py b/memoryscope/core/operation/base_operation.py similarity index 100% rename from memoryscope/memory/operation/base_operation.py rename to memoryscope/core/operation/base_operation.py diff --git a/memoryscope/memory/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py similarity index 96% rename from memoryscope/memory/operation/base_workflow.py rename to memoryscope/core/operation/base_workflow.py index 7cb75daf..4c6ec8fa 100644 --- a/memoryscope/memory/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -5,11 +5,11 @@ from itertools import zip_longest from typing import Dict, Any, List from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT -from memoryscope.memory.worker.base_worker import BaseWorker -from memoryscope.memoryscope_context import MemoryscopeContext -from memoryscope.utils.logger import Logger -from memoryscope.utils.timer import Timer -from memoryscope.utils.tool_functions import init_instance_by_config +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.timer import Timer +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.base_worker import BaseWorker class BaseWorkflow(object): diff --git a/memoryscope/memory/operation/consolidate_operation.py b/memoryscope/core/operation/consolidate_memory_op.py similarity index 92% rename from memoryscope/memory/operation/consolidate_operation.py rename to memoryscope/core/operation/consolidate_memory_op.py index 967c3639..fec60e12 100644 --- a/memoryscope/memory/operation/consolidate_operation.py +++ b/memoryscope/core/operation/consolidate_memory_op.py @@ -1,12 +1,12 @@ from memoryscope.constants.common_constants import CHAT_KWARGS, CHAT_MESSAGES, RESULT +from memoryscope.core.operation.backend_operation import BackendOperation from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.operation.backend_operation import BackendOperation -class ConsolidateOperation(BackendOperation): +class ConsolidateMemoryOp(BackendOperation): def __init__(self, **kwargs): - super(ConsolidateOperation, self).__init__(**kwargs) + super(ConsolidateMemoryOp, self).__init__(**kwargs) self.message_lock = kwargs.get("message_lock", None) self.contextual_msg_min_count: int = kwargs.get("contextual_msg_min_count", 0) diff --git a/memoryscope/memory/operation/frontend_operation.py b/memoryscope/core/operation/frontend_operation.py similarity index 91% rename from memoryscope/memory/operation/frontend_operation.py rename to memoryscope/core/operation/frontend_operation.py index abaca44a..bed28875 100644 --- a/memoryscope/memory/operation/frontend_operation.py +++ b/memoryscope/core/operation/frontend_operation.py @@ -1,8 +1,8 @@ from typing import List from memoryscope.constants.common_constants import RESULT, CHAT_MESSAGES, CHAT_KWARGS -from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE -from memoryscope.memory.operation.base_workflow import BaseWorkflow +from memoryscope.core.operation.base_operation import BaseOperation, OPERATION_TYPE +from memoryscope.core.operation.base_workflow import BaseWorkflow from memoryscope.scheme.message import Message diff --git a/memoryscope/memory/worker/__init__.py b/memoryscope/core/service/__init__.py similarity index 100% rename from memoryscope/memory/worker/__init__.py rename to memoryscope/core/service/__init__.py diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/core/service/base_memory_service.py similarity index 94% rename from memoryscope/memory/service/base_memory_service.py rename to memoryscope/core/service/base_memory_service.py index a0a476a5..d997fa98 100644 --- a/memoryscope/memory/service/base_memory_service.py +++ b/memoryscope/core/service/base_memory_service.py @@ -1,10 +1,10 @@ from abc import ABCMeta, abstractmethod from typing import List, Dict -from memoryscope.memory.operation.base_operation import BaseOperation -from memoryscope.memoryscope_context import MemoryscopeContext +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.operation.base_operation import BaseOperation +from memoryscope.core.utils.logger import Logger from memoryscope.scheme.message import Message -from memoryscope.utils.logger import Logger class BaseMemoryService(metaclass=ABCMeta): diff --git a/memoryscope/memory/service/memory_scope_service.py b/memoryscope/core/service/memory_scope_service.py similarity index 95% rename from memoryscope/memory/service/memory_scope_service.py rename to memoryscope/core/service/memory_scope_service.py index bf1c9b94..88c644d4 100644 --- a/memoryscope/memory/service/memory_scope_service.py +++ b/memoryscope/core/service/memory_scope_service.py @@ -1,10 +1,10 @@ import threading from typing import List -from memoryscope.memory.operation.base_operation import BaseOperation -from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.core.operation.base_operation import BaseOperation +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.tool_functions import init_instance_by_config from memoryscope.scheme.message import Message -from memoryscope.utils.tool_functions import init_instance_by_config class MemoryScopeService(BaseMemoryService): diff --git a/memoryscope/memory/worker/backend/__init__.py b/memoryscope/core/storage/__init__.py similarity index 100% rename from memoryscope/memory/worker/backend/__init__.py rename to memoryscope/core/storage/__init__.py diff --git a/memoryscope/storage/base_memory_store.py b/memoryscope/core/storage/base_memory_store.py similarity index 100% rename from memoryscope/storage/base_memory_store.py rename to memoryscope/core/storage/base_memory_store.py diff --git a/memoryscope/storage/base_monitor.py b/memoryscope/core/storage/base_monitor.py similarity index 100% rename from memoryscope/storage/base_monitor.py rename to memoryscope/core/storage/base_monitor.py diff --git a/memoryscope/storage/dummy_memory_store.py b/memoryscope/core/storage/dummy_memory_store.py similarity index 92% rename from memoryscope/storage/dummy_memory_store.py rename to memoryscope/core/storage/dummy_memory_store.py index 7c77f670..2d5eeede 100644 --- a/memoryscope/storage/dummy_memory_store.py +++ b/memoryscope/core/storage/dummy_memory_store.py @@ -1,8 +1,8 @@ from typing import Dict, List -from memoryscope.models.base_model import BaseModel +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.storage.base_memory_store import BaseMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.base_memory_store import BaseMemoryStore class DummyMemoryStore(BaseMemoryStore): diff --git a/memoryscope/storage/dummy_monitor.py b/memoryscope/core/storage/dummy_monitor.py similarity index 92% rename from memoryscope/storage/dummy_monitor.py rename to memoryscope/core/storage/dummy_monitor.py index 4818e0fb..816202e5 100644 --- a/memoryscope/storage/dummy_monitor.py +++ b/memoryscope/core/storage/dummy_monitor.py @@ -1,4 +1,4 @@ -from memoryscope.storage.base_monitor import BaseMonitor +from memoryscope.core.storage.base_monitor import BaseMonitor class DummyMonitor(BaseMonitor): diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/core/storage/llama_index_es_memory_store.py similarity index 87% rename from memoryscope/storage/llama_index_es_memory_store.py rename to memoryscope/core/storage/llama_index_es_memory_store.py index 22d43dd7..f3967039 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/core/storage/llama_index_es_memory_store.py @@ -4,12 +4,14 @@ from typing import Dict, List from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle -from memoryscope.models.base_model import BaseModel +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.storage.base_memory_store import BaseMemoryStore +from memoryscope.core.storage.llama_index_sync_elasticsearch import (SyncElasticsearchStore, + ESCombinedRetrieveStrategy, + _to_elasticsearch_filter, + SPECIAL_QUERY) +from memoryscope.core.utils.logger import Logger from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, ESCombinedRetrieveStrategy, \ - _to_elasticsearch_filter -from memoryscope.utils.logger import Logger class LlamaIndexEsMemoryStore(BaseMemoryStore): @@ -38,7 +40,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.logger = Logger.get_logger() def retrieve_memories(self, - query: str = "**--**", + query: str = "", top_k: int = 3, filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: # if index is not created, return [] @@ -53,8 +55,12 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter, "fields": ['embedding']}, similarity_top_k=top_k, sparse_top_k=top_k) + + if not query: + query = SPECIAL_QUERY + if not query and self.emb_dims: - query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) + query = QueryBundle(query_str=SPECIAL_QUERY, embedding=self.dummy_query_vector()) text_nodes = retriever.retrieve(query) if text_nodes and text_nodes[0].embedding: @@ -80,7 +86,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): sparse_top_k=top_k) if not query: - query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) + query = SPECIAL_QUERY + + if not query: + query = QueryBundle(query_str=SPECIAL_QUERY, embedding=self.dummy_query_vector()) text_nodes: List[NodeWithScore] = await retriever.aretrieve(query) diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/core/storage/llama_index_sync_elasticsearch.py similarity index 98% rename from memoryscope/storage/llama_index_sync_elasticsearch.py rename to memoryscope/core/storage/llama_index_sync_elasticsearch.py index 50eeef3d..a94c9168 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/core/storage/llama_index_sync_elasticsearch.py @@ -38,6 +38,8 @@ DISTANCE_STRATEGIES = Literal[ "EUCLIDEAN_DISTANCE", ] +SPECIAL_QUERY: str = "**--**" + def get_elasticsearch_client( url: Optional[str] = None, @@ -134,15 +136,15 @@ def _mode_must_match_retrieval_strategy( class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): def __init__( - self, - *, - distance: DistanceMetric = DistanceMetric.COSINE, - model_id: Optional[str] = None, - retrieve_mode: str = "dense", - rrf: Union[bool, Dict[str, Any]] = True, - text_field: Optional[str] = "text_field", - hybrid_alpha: Optional[float] = None, - ): + self, + *, + distance: DistanceMetric = DistanceMetric.COSINE, + model_id: Optional[str] = None, + retrieve_mode: str = "dense", + rrf: Union[bool, Dict[str, Any]] = True, + text_field: Optional[str] = "text_field", + hybrid_alpha: Optional[float] = None, + ): if retrieve_mode == "dense": self.alpha = 1.0 elif retrieve_mode == "sparse": @@ -151,7 +153,7 @@ class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): elif retrieve_mode == "hybrid": # self.alpha = hybrid_alpha raise NotImplementedError - + super().__init__(distance=distance, model_id=model_id, hybrid=True, rrf=rrf, text_field=text_field) def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]: @@ -159,7 +161,7 @@ class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): # RRF is used to even the score from the knn query and text query # RRF has two optional parameters: {'rank_constant':int, 'window_size':int} # https://www.elastic.co/guide/en/elasticsearch/reference/current/rrf.html - if query == "**--**": + if query == SPECIAL_QUERY: query_body = { "query": { "bool": { @@ -699,7 +701,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy) and self.retrieval_strategy.hybrid ): - # total_rank = sum(top_k_scores) + total_rank = sum(top_k_scores) top_k_scores = [rank for rank in top_k_scores] # top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores] # top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores] diff --git a/memoryscope/memory/worker/frontend/__init__.py b/memoryscope/core/utils/__init__.py similarity index 100% rename from memoryscope/memory/worker/frontend/__init__.py rename to memoryscope/core/utils/__init__.py diff --git a/memoryscope/utils/datetime_handler.py b/memoryscope/core/utils/datetime_handler.py similarity index 96% rename from memoryscope/utils/datetime_handler.py rename to memoryscope/core/utils/datetime_handler.py index f41cfcfa..21ab3f31 100644 --- a/memoryscope/utils/datetime_handler.py +++ b/memoryscope/core/utils/datetime_handler.py @@ -3,8 +3,8 @@ import re from typing import List from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT +from memoryscope.core.utils.logger import Logger from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.utils.logger import Logger class DatetimeHandler(object): @@ -222,9 +222,9 @@ class DatetimeHandler(object): Returns: dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails. """ - func_name = f"extract_date_parts_{language}" + func_name = f"extract_date_parts_{language.value}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={language} needs to complete extract_date_parts func!") + cls.logger.warning(f"language={language.value} needs to complete extract_date_parts func!") return {} return getattr(cls, func_name)(input_string=input_string) @@ -272,13 +272,13 @@ class DatetimeHandler(object): @classmethod def has_time_word(cls, query: str, language: LanguageEnum) -> bool: - func_name = f"has_time_word_{language}" + func_name = f"has_time_word_{language.value}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={language} needs to complete has_time_word function!") + cls.logger.warning(f"language={language.value} needs to complete has_time_word function!") return False if language not in DATATIME_WORD_LIST: - cls.logger.warning(f"language={language} is missing in DATATIME_WORD_LIST!") + cls.logger.warning(f"language={language.value} is missing in DATATIME_WORD_LIST!") return False datetime_word_list = DATATIME_WORD_LIST[language] diff --git a/memoryscope/utils/logger.py b/memoryscope/core/utils/logger.py similarity index 100% rename from memoryscope/utils/logger.py rename to memoryscope/core/utils/logger.py diff --git a/memoryscope/utils/prompt_handler.py b/memoryscope/core/utils/prompt_handler.py similarity index 100% rename from memoryscope/utils/prompt_handler.py rename to memoryscope/core/utils/prompt_handler.py diff --git a/memoryscope/utils/registry.py b/memoryscope/core/utils/registry.py similarity index 100% rename from memoryscope/utils/registry.py rename to memoryscope/core/utils/registry.py diff --git a/memoryscope/utils/response_text_parser.py b/memoryscope/core/utils/response_text_parser.py similarity index 97% rename from memoryscope/utils/response_text_parser.py rename to memoryscope/core/utils/response_text_parser.py index 452b5014..d74b3141 100644 --- a/memoryscope/utils/response_text_parser.py +++ b/memoryscope/core/utils/response_text_parser.py @@ -2,8 +2,8 @@ import re from typing import List from memoryscope.constants.language_constants import NONE_WORD +from memoryscope.core.utils.logger import Logger from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.utils.logger import Logger class ResponseTextParser(object): diff --git a/memoryscope/utils/timer.py b/memoryscope/core/utils/timer.py similarity index 97% rename from memoryscope/utils/timer.py rename to memoryscope/core/utils/timer.py index ac7ca1f0..4ba40903 100644 --- a/memoryscope/utils/timer.py +++ b/memoryscope/core/utils/timer.py @@ -1,7 +1,7 @@ import time from typing import Literal -from memoryscope.utils.logger import Logger +from memoryscope.core.utils.logger import Logger TIME_LOG_TYPE = Literal["end", "wrap", "none"] @@ -75,7 +75,7 @@ class Timer(object): self.logger.info(f"----- {self.name}.begin -----") return self - def __exit__(self, *args, **kwargs): + def __exit__(self, exc_type, exc_value, exc_tb): """ End timing and print the formatted log. """ diff --git a/memoryscope/utils/tool_functions.py b/memoryscope/core/utils/tool_functions.py similarity index 100% rename from memoryscope/utils/tool_functions.py rename to memoryscope/core/utils/tool_functions.py diff --git a/memoryscope/models/__init__.py b/memoryscope/core/worker/__init__.py similarity index 100% rename from memoryscope/models/__init__.py rename to memoryscope/core/worker/__init__.py diff --git a/memoryscope/storage/__init__.py b/memoryscope/core/worker/backend/__init__.py similarity index 100% rename from memoryscope/storage/__init__.py rename to memoryscope/core/worker/backend/__init__.py diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.py b/memoryscope/core/worker/backend/contra_repeat_worker.py similarity index 96% rename from memoryscope/memory/worker/backend/contra_repeat_worker.py rename to memoryscope/core/worker/backend/contra_repeat_worker.py index d245ba85..177523eb 100644 --- a/memoryscope/memory/worker/backend/contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/contra_repeat_worker.py @@ -2,11 +2,11 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_NODES from memoryscope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, CONTAINED_WORD +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class ContraRepeatWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.yaml b/memoryscope/core/worker/backend/contra_repeat_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/contra_repeat_worker.yaml rename to memoryscope/core/worker/backend/contra_repeat_worker.yaml diff --git a/memoryscope/memory/worker/backend/get_observation_with_time_worker.py b/memoryscope/core/worker/backend/get_observation_with_time_worker.py similarity index 94% rename from memoryscope/memory/worker/backend/get_observation_with_time_worker.py rename to memoryscope/core/worker/backend/get_observation_with_time_worker.py index 5c7ca66b..b1fa4c79 100644 --- a/memoryscope/memory/worker/backend/get_observation_with_time_worker.py +++ b/memoryscope/core/worker/backend/get_observation_with_time_worker.py @@ -2,10 +2,10 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_WITH_TIME_NODES from memoryscope.constants.language_constants import COLON_WORD -from memoryscope.memory.worker.backend.get_observation_worker import GetObservationWorker +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.backend.get_observation_worker import GetObservationWorker from memoryscope.scheme.message import Message -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.tool_functions import prompt_to_msg class GetObservationWithTimeWorker(GetObservationWorker): diff --git a/memoryscope/memory/worker/backend/get_observation_with_time_worker.yaml b/memoryscope/core/worker/backend/get_observation_with_time_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/get_observation_with_time_worker.yaml rename to memoryscope/core/worker/backend/get_observation_with_time_worker.yaml diff --git a/memoryscope/memory/worker/backend/get_observation_worker.py b/memoryscope/core/worker/backend/get_observation_worker.py similarity index 96% rename from memoryscope/memory/worker/backend/get_observation_worker.py rename to memoryscope/core/worker/backend/get_observation_worker.py index 0a90803b..78c7eeda 100644 --- a/memoryscope/memory/worker/backend/get_observation_worker.py +++ b/memoryscope/core/worker/backend/get_observation_worker.py @@ -2,14 +2,14 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_NODES, TIME_INFER from memoryscope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_WORD +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class GetObservationWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/get_observation_worker.yaml b/memoryscope/core/worker/backend/get_observation_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/get_observation_worker.yaml rename to memoryscope/core/worker/backend/get_observation_worker.yaml diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py b/memoryscope/core/worker/backend/get_reflection_subject_worker.py similarity index 94% rename from memoryscope/memory/worker/backend/get_reflection_subject_worker.py rename to memoryscope/core/worker/backend/get_reflection_subject_worker.py index f9e5511d..f9323372 100644 --- a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py +++ b/memoryscope/core/worker/backend/get_reflection_subject_worker.py @@ -2,13 +2,13 @@ from typing import List from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES from memoryscope.constants.language_constants import COMMA_WORD +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class GetReflectionSubjectWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.yaml b/memoryscope/core/worker/backend/get_reflection_subject_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/get_reflection_subject_worker.yaml rename to memoryscope/core/worker/backend/get_reflection_subject_worker.yaml diff --git a/memoryscope/memory/worker/backend/info_filter_worker.py b/memoryscope/core/worker/backend/info_filter_worker.py similarity index 95% rename from memoryscope/memory/worker/backend/info_filter_worker.py rename to memoryscope/core/worker/backend/info_filter_worker.py index 1474429c..78307c68 100644 --- a/memoryscope/memory/worker/backend/info_filter_worker.py +++ b/memoryscope/core/worker/backend/info_filter_worker.py @@ -1,11 +1,11 @@ from typing import List from memoryscope.constants.language_constants import COLON_WORD +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.message import Message -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class InfoFilterWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/info_filter_worker.yaml b/memoryscope/core/worker/backend/info_filter_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/info_filter_worker.yaml rename to memoryscope/core/worker/backend/info_filter_worker.yaml diff --git a/memoryscope/memory/worker/backend/load_memory_worker.py b/memoryscope/core/worker/backend/load_memory_worker.py similarity index 96% rename from memoryscope/memory/worker/backend/load_memory_worker.py rename to memoryscope/core/worker/backend/load_memory_worker.py index 22c77792..00c24ab3 100644 --- a/memoryscope/memory/worker/backend/load_memory_worker.py +++ b/memoryscope/core/worker/backend/load_memory_worker.py @@ -1,12 +1,12 @@ from typing import List from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES, TODAY_NODES +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.timer import timer +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.timer import timer class LoadMemoryWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py b/memoryscope/core/worker/backend/long_contra_repeat_worker.py similarity index 97% rename from memoryscope/memory/worker/backend/long_contra_repeat_worker.py rename to memoryscope/core/worker/backend/long_contra_repeat_worker.py index c12c2cc3..527e1201 100644 --- a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/long_contra_repeat_worker.py @@ -2,13 +2,13 @@ from typing import List, Dict from memoryscope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_NODES from memoryscope.constants.language_constants import NONE_WORD, CONTAINED_WORD, CONTRADICTORY_WORD +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class LongContraRepeatWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.yaml b/memoryscope/core/worker/backend/long_contra_repeat_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/long_contra_repeat_worker.yaml rename to memoryscope/core/worker/backend/long_contra_repeat_worker.yaml diff --git a/memoryscope/memory/worker/backend/update_insight_worker.py b/memoryscope/core/worker/backend/update_insight_worker.py similarity index 97% rename from memoryscope/memory/worker/backend/update_insight_worker.py rename to memoryscope/core/worker/backend/update_insight_worker.py index a4b66e2e..fbc7698b 100644 --- a/memoryscope/memory/worker/backend/update_insight_worker.py +++ b/memoryscope/core/worker/backend/update_insight_worker.py @@ -3,12 +3,12 @@ from typing import List from memoryscope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES from memoryscope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg, cosine_similarity +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg, cosine_similarity class UpdateInsightWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/backend/update_insight_worker.yaml b/memoryscope/core/worker/backend/update_insight_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/update_insight_worker.yaml rename to memoryscope/core/worker/backend/update_insight_worker.yaml diff --git a/memoryscope/memory/worker/backend/update_memory_worker.py b/memoryscope/core/worker/backend/update_memory_worker.py similarity index 96% rename from memoryscope/memory/worker/backend/update_memory_worker.py rename to memoryscope/core/worker/backend/update_memory_worker.py index 9133de68..b6803776 100644 --- a/memoryscope/memory/worker/backend/update_memory_worker.py +++ b/memoryscope/core/worker/backend/update_memory_worker.py @@ -1,10 +1,10 @@ from typing import List +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler class UpdateMemoryWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/base_worker.py b/memoryscope/core/worker/base_worker.py similarity index 98% rename from memoryscope/memory/worker/base_worker.py rename to memoryscope/core/worker/base_worker.py index e5bf691c..fc92ce71 100644 --- a/memoryscope/memory/worker/base_worker.py +++ b/memoryscope/core/worker/base_worker.py @@ -3,8 +3,8 @@ from abc import ABCMeta, abstractmethod from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any, Dict -from memoryscope.utils.logger import Logger -from memoryscope.utils.timer import Timer +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.timer import Timer class BaseWorker(metaclass=ABCMeta): diff --git a/memoryscope/memory/worker/dummy_worker.py b/memoryscope/core/worker/dummy_worker.py similarity index 92% rename from memoryscope/memory/worker/dummy_worker.py rename to memoryscope/core/worker/dummy_worker.py index b8eecb7b..bd2c1760 100644 --- a/memoryscope/memory/worker/dummy_worker.py +++ b/memoryscope/core/worker/dummy_worker.py @@ -1,7 +1,7 @@ import datetime from memoryscope.constants.common_constants import RESULT, WORKFLOW_NAME, CHAT_KWARGS -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker class DummyWorker(MemoryBaseWorker): diff --git a/memoryscope/utils/__init__.py b/memoryscope/core/worker/frontend/__init__.py similarity index 100% rename from memoryscope/utils/__init__.py rename to memoryscope/core/worker/frontend/__init__.py diff --git a/memoryscope/memory/worker/frontend/extract_time_worker.py b/memoryscope/core/worker/frontend/extract_time_worker.py similarity index 93% rename from memoryscope/memory/worker/frontend/extract_time_worker.py rename to memoryscope/core/worker/frontend/extract_time_worker.py index 6f92c3c8..b2d864e2 100644 --- a/memoryscope/memory/worker/frontend/extract_time_worker.py +++ b/memoryscope/core/worker/frontend/extract_time_worker.py @@ -3,9 +3,9 @@ from typing import Dict from memoryscope.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT from memoryscope.constants.language_constants import DATATIME_KEY_MAP -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.tool_functions import prompt_to_msg +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker class ExtractTimeWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/frontend/extract_time_worker.yaml b/memoryscope/core/worker/frontend/extract_time_worker.yaml similarity index 100% rename from memoryscope/memory/worker/frontend/extract_time_worker.yaml rename to memoryscope/core/worker/frontend/extract_time_worker.yaml diff --git a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py b/memoryscope/core/worker/frontend/fuse_rerank_worker.py similarity index 97% rename from memoryscope/memory/worker/frontend/fuse_rerank_worker.py rename to memoryscope/core/worker/frontend/fuse_rerank_worker.py index 3d394790..cd749211 100644 --- a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py +++ b/memoryscope/core/worker/frontend/fuse_rerank_worker.py @@ -1,9 +1,9 @@ from typing import Dict, List from memoryscope.constants.common_constants import EXTRACT_TIME_DICT, RANKED_MEMORY_NODES, RESULT -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler class FuseRerankWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/frontend/print_memory_worker.py b/memoryscope/core/worker/frontend/print_memory_worker.py similarity index 94% rename from memoryscope/memory/worker/frontend/print_memory_worker.py rename to memoryscope/core/worker/frontend/print_memory_worker.py index e50adad9..5154371f 100644 --- a/memoryscope/memory/worker/frontend/print_memory_worker.py +++ b/memoryscope/core/worker/frontend/print_memory_worker.py @@ -1,11 +1,11 @@ from typing import List from memoryscope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler class PrintMemoryWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/frontend/print_memory_worker.yaml b/memoryscope/core/worker/frontend/print_memory_worker.yaml similarity index 100% rename from memoryscope/memory/worker/frontend/print_memory_worker.yaml rename to memoryscope/core/worker/frontend/print_memory_worker.yaml diff --git a/memoryscope/memory/worker/frontend/read_message_worker.py b/memoryscope/core/worker/frontend/read_message_worker.py similarity index 91% rename from memoryscope/memory/worker/frontend/read_message_worker.py rename to memoryscope/core/worker/frontend/read_message_worker.py index dfbe4a12..32f3380b 100644 --- a/memoryscope/memory/worker/frontend/read_message_worker.py +++ b/memoryscope/core/worker/frontend/read_message_worker.py @@ -1,6 +1,6 @@ from memoryscope.constants.common_constants import RESULT +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker class ReadMessageWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py b/memoryscope/core/worker/frontend/retrieve_memory_worker.py similarity index 96% rename from memoryscope/memory/worker/frontend/retrieve_memory_worker.py rename to memoryscope/core/worker/frontend/retrieve_memory_worker.py index e51bbd2d..6b5ddea4 100644 --- a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memoryscope/core/worker/frontend/retrieve_memory_worker.py @@ -1,12 +1,12 @@ from typing import List from memoryscope.constants.common_constants import QUERY_WITH_TS, RETRIEVE_MEMORY_NODES +from memoryscope.core.utils.timer import timer +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.timer import timer class RetrieveMemoryWorker(MemoryBaseWorker): @@ -120,6 +120,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): 7. Stores the processed memory nodes for further use. """ query, _ = self.get_context(QUERY_WITH_TS) + self.logger.info(f"retrieve memory with query={query}.") self.submit_thread_task(self.retrieve_from_observation, query=query) self.submit_thread_task(self.retrieve_from_insight, query=query) self.submit_thread_task(self.retrieve_expired_memory, query=query) @@ -136,7 +137,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): memory_node_list = sorted(memory_node_list, key=lambda x: x.score_recall, reverse=True) for node in memory_node_list: node.action_status = ActionStatusEnum.NONE.value - self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} type={node.memory_type} " + self.logger.info(f"recall_stage: content={node.content} score={node.score_recall} type={node.memory_type} " f"store_status={node.store_status} action_status={node.action_status}") self.memory_manager.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memoryscope/memory/worker/frontend/semantic_rank_worker.py b/memoryscope/core/worker/frontend/semantic_rank_worker.py similarity index 70% rename from memoryscope/memory/worker/frontend/semantic_rank_worker.py rename to memoryscope/core/worker/frontend/semantic_rank_worker.py index 4dc7b303..e1ca0a04 100644 --- a/memoryscope/memory/worker/frontend/semantic_rank_worker.py +++ b/memoryscope/core/worker/frontend/semantic_rank_worker.py @@ -1,7 +1,7 @@ from typing import List, Dict from memoryscope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode @@ -39,22 +39,23 @@ class SemanticRankWorker(MemoryBaseWorker): for node in memory_node_list: node.score_rank = node.score_recall self.logger.warning("use score_recall instead of score_rank!") - return - # drop repeated - memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if n.content.strip()} - memory_node_list = list(memory_node_dict.values()) + else: + # drop repeated + memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if + n.content.strip()} + memory_node_list = list(memory_node_dict.values()) - response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list]) - if not response.status or not response.rank_scores: - return + response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list]) + if not response.status or not response.rank_scores: + return - # set score - for idx, score in response.rank_scores.items(): - if idx >= len(memory_node_list): - self.logger.warning(f"Idx={idx} exceeds the maximum length of rank_scores!") - continue - memory_node_list[idx].score_rank = score + # set score + for idx, score in response.rank_scores.items(): + if idx >= len(memory_node_list): + self.logger.warning(f"Idx={idx} exceeds the maximum length of rank_scores!") + continue + memory_node_list[idx].score_rank = score # sort by score memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True) diff --git a/memoryscope/memory/worker/frontend/set_query_worker.py b/memoryscope/core/worker/frontend/set_query_worker.py similarity index 97% rename from memoryscope/memory/worker/frontend/set_query_worker.py rename to memoryscope/core/worker/frontend/set_query_worker.py index 1f079288..85961bc9 100644 --- a/memoryscope/memory/worker/frontend/set_query_worker.py +++ b/memoryscope/core/worker/frontend/set_query_worker.py @@ -1,8 +1,8 @@ import datetime from memoryscope.constants.common_constants import QUERY_WITH_TS +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker class SetQueryWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py similarity index 94% rename from memoryscope/memory/worker/memory_base_worker.py rename to memoryscope/core/worker/memory_base_worker.py index e800f8b3..f4de9044 100644 --- a/memoryscope/memory/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -3,15 +3,15 @@ from typing import List, Dict, Any from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \ WORKFLOW_NAME, MEMORY_MANAGER +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.storage.base_memory_store import BaseMemoryStore +from memoryscope.core.storage.base_monitor import BaseMonitor +from memoryscope.core.utils.prompt_handler import PromptHandler +from memoryscope.core.worker.base_worker import BaseWorker +from memoryscope.core.worker.memory_manager import MemoryManager from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.memory.worker.base_worker import BaseWorker -from memoryscope.memory.worker.memory_manager import MemoryManager -from memoryscope.memoryscope_context import MemoryscopeContext -from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.base_monitor import BaseMonitor -from memoryscope.utils.prompt_handler import PromptHandler class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): diff --git a/memoryscope/memory/worker/memory_manager.py b/memoryscope/core/worker/memory_manager.py similarity index 97% rename from memoryscope/memory/worker/memory_manager.py rename to memoryscope/core/worker/memory_manager.py index f705d832..af7b3016 100644 --- a/memoryscope/memory/worker/memory_manager.py +++ b/memoryscope/core/worker/memory_manager.py @@ -1,11 +1,11 @@ from typing import Dict, List +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.storage.base_memory_store import BaseMemoryStore +from memoryscope.core.utils.logger import Logger from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.utils.logger import Logger class MemoryManager(object): diff --git a/memoryscope/memoryscope.py b/memoryscope/memoryscope.py deleted file mode 100644 index 647b8418..00000000 --- a/memoryscope/memoryscope.py +++ /dev/null @@ -1,201 +0,0 @@ -import datetime -import json -from concurrent.futures import ThreadPoolExecutor - -import yaml - -from memoryscope.argument import default_arguments -from memoryscope.argument.memoryscope_arguments import MemoryscopeArguments -from memoryscope.chat.base_memory_chat import BaseMemoryChat -from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.memoryscope_context import MemoryscopeContext -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config - - -class MemoryScope(object): - - def __init__(self, - arguments: MemoryscopeArguments | None = None, - config: dict | None = None, - config_path: str = ""): - - self.global_conf: dict = {} - self.memory_chat_conf_dict: dict = {} - self.memory_service_conf_dict: dict = {} - self.worker_conf_dict: dict = {} - self.model_conf_dict: dict = {} - self.memory_store_conf: dict = {} - self.monitor_conf: dict = {} - - self.context: MemoryscopeContext = MemoryscopeContext() - - if arguments: - self._init_by_arguments(arguments=arguments) - elif config: - self._init_by_config(config=config) - elif config_path: - self._init_by_config_path(config_path=config_path) - else: - raise RuntimeError("At least one of arguments, config, or file_path must not be empty!") - - self.logger = self._init_logger() - - self._init_context_by_config() - - def _init_by_arguments(self, arguments: MemoryscopeArguments): - # prepare global - self.global_conf = { - "language": arguments.language, - "thread_pool_max_workers": arguments.thread_pool_max_workers, - "logger_name": arguments.logger_name, - "logger_name_time_suffix": arguments.logger_name_time_suffix, - "use_dummy_ranker": arguments.use_dummy_ranker, - } - - # prepare memory chat - self.memory_chat_conf_dict = default_arguments.DEFAULT_MEMORY_CHAT_ARGUMENTS.copy() - memory_chat_config = list(self.memory_chat_conf_dict.values())[0] - memory_chat_config.update({ - "class": arguments.memory_chat_class, - "human_name": arguments.human_name, - "assistant_name": arguments.assistant_name, - }) - - # prepare memory service - self.memory_service_conf_dict = default_arguments.DEFAULT_MEMORY_SERVICE_ARGUMENTS.copy() - memory_service_config = list(self.memory_service_conf_dict.values())[0] - memory_service_config.update({ - "human_name": arguments.human_name, - "assistant_name": arguments.assistant_name, - }) - memory_service_config["memory_operations"]["consolidate_memory"]["interval_time"] = \ - arguments.consolidate_memory_interval_time - memory_service_config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \ - arguments.reflect_and_reconsolidate_interval_time - - # prepare memory service - self.worker_conf_dict = default_arguments.DEFAULT_WORKER_ARGUMENTS.copy() - if arguments.worker_params: - for worker_name, kv_dict in arguments.worker_params.items(): - if worker_name not in self.worker_conf_dict: - continue - self.worker_conf_dict[worker_name].update(kv_dict) - - # prepare models - self.model_conf_dict = { - "generation_model": { - "class": "models.llama_index_generation_model", - "module_name": arguments.generation_backend, - "model_name": arguments.generation_model, - **arguments.generation_params, - }, - "embedding_model": { - "class": "models.llama_index_embedding_model", - "module_name": arguments.embedding_backend, - "model_name": arguments.embedding_model, - **arguments.embedding_params, - }, - "rank_model": { - "class": "models.llama_index_rank_model", - "module_name": arguments.rank_backend, - "model_name": arguments.rank_model, - **arguments.rank_params, - }, - } - - # prepare memory store - self.memory_store_conf = { - "class": "storage.llama_index_es_memory_store", - "embedding_model": "embedding_model", - "index_name": arguments.es_index_name, - "es_url": arguments.es_url, - "retrieve_mode": arguments.retrieve_mode, - "hybrid_alpha": arguments.hybrid_alpha, - } - - self.monitor_conf = default_arguments.DEFAULT_MONITOR_ARGUMENTS.copy() - - def _init_by_config(self, config: dict): - self.global_conf = config["global_config"] - self.memory_service_conf_dict = config["memory_service"] - self.worker_conf_dict = config["worker"] - self.model_conf_dict = config["model"] - self.memory_store_conf = config["memory_store"] - - # not necessary - self.memory_chat_conf_dict = config.get("memory_chat") - self.monitor_conf = config.get("monitor") - - def _init_by_config_path(self, config_path: str): - with open(config_path) as f: - if config_path.endswith("yaml"): - config = yaml.load(f, yaml.FullLoader) - elif config_path.endswith("json"): - config = json.load(f) - else: - raise RuntimeError("not supported config file type!") - return self._init_by_config(config) - - def _init_logger(self) -> Logger: - logger_name = self.global_conf.get("logger_name") - assert logger_name, "logger_name is empty!" - logger_name_time_suffix = self.global_conf.get("logger_name_time_suffix") - if logger_name_time_suffix: - suffix = datetime.datetime.now().strftime(logger_name_time_suffix) - logger_name = f"{logger_name}_{suffix}" - return Logger.get_logger(logger_name, to_stream=False) - - def _init_context_by_config(self): - # set global config - self.context.language = LanguageEnum(self.global_conf["language"]) - self.context.thread_pool = ThreadPoolExecutor(max_workers=self.global_conf["max_workers"]) - self.context.meta_data["use_dummy_ranker"] = self.global_conf["use_dummy_ranker"] - - # init memory_chat - if self.memory_chat_conf_dict: - for name, conf in self.memory_chat_conf_dict.items(): - self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context) - - # set memory_service - assert self.memory_service_conf_dict - for name, conf in self.memory_service_conf_dict.items(): - self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context) - - # init models - assert self.model_conf_dict - for name, conf in self.model_conf_dict.items(): - self.context.model_dict[name] = init_instance_by_config(conf, name=name) - - # init vector_store - assert self.memory_store_conf - emb_model_name: str = self.memory_store_conf[ModelEnum.EMBEDDING_MODEL.value] - embedding_model = self.context.model_dict[emb_model_name] - self.context.memory_store = init_instance_by_config(self.memory_store_conf, embedding_model=embedding_model) - - # init monitor - if self.monitor_conf: - self.context.monitor = init_instance_by_config(self.monitor_conf) - - # set worker config - self.context.worker_config = self.worker_conf_dict - - def close(self): - # wait service to stop - for _, service in self.context.memory_service_dict.items(): - service.stop_backend_service(wait_service_end=True) - self.context.memory_store.close() - self.context.thread_pool.shutdown() - - if self.context.monitor: - self.context.monitor.close() - - @property - def default_memory_chat(self) -> BaseMemoryChat: - return list(self.context.memory_chat_dict.values())[0] - - @property - def default_service(self) -> BaseMemoryService: - return list(self.context.memory_service_dict.values())[0] diff --git a/tests/operations/test_interface.py b/tests/operations/test_interface.py index d89efb93..05bdaef5 100644 --- a/tests/operations/test_interface.py +++ b/tests/operations/test_interface.py @@ -1,7 +1,7 @@ from memoryscope.cli import MemoryScope from memoryscope.scheme.message import Message -ms = MemoryScope().load_config("config/demo_config_no_stream.yaml") +ms = MemoryScope().read_config("config/demo_config_no_stream.yaml") memory_service = ms.default_service memory_chat = ms.default_chat_handle diff --git a/tests/other/test_cli.py b/tests/other/test_cli.py new file mode 100644 index 00000000..89a99f17 --- /dev/null +++ b/tests/other/test_cli.py @@ -0,0 +1,14 @@ +import fire + + +class CLI: + def run(self, **kwargs): + """ + 打印传入的 kwargs + """ + for key, value in kwargs.items(): + print(f"{key}: {value}") + + +if __name__ == '__main__': + fire.Fire(CLI().run) diff --git a/tests/storages/test_storages_lli_synces.py b/tests/storages/test_storages_lli_synces.py index 0bb8b554..bf32bc73 100644 --- a/tests/storages/test_storages_lli_synces.py +++ b/tests/storages/test_storages_lli_synces.py @@ -164,7 +164,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): meta_data={"5": "5"}, timestamp=13 )) - + def test_retrieve(self): filter_dict = { "timestamp": 12, @@ -172,12 +172,11 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): # "score_rank": 0, } - res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15) print(len(res)) print(res) - def test_retrieve_wo_query(self,): + def test_retrieve_wo_query(self, ): filter_dict = { "memory_id": "bbb456", } @@ -185,6 +184,5 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): print(len(res)) print(res) - def tearDown(self): self.es_store.close() diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index c841d73d..e7a65fae 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -21,7 +21,7 @@ class TestWorkersCn(unittest.TestCase): self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) ms = MemoryScope() - ms.load_config("config/demo_config_cn.yaml") + ms.read_config("config/demo_config_cn.yaml") ms.init_global_content_by_config() def tearDown(self): diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 356c3f21..0123bc3a 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -21,7 +21,7 @@ class TestWorkersEn(unittest.TestCase): self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) ms = MemoryScope() - ms.load_config("config/demo_config_en.yaml") + ms.read_config("config/demo_config_en.yaml") ms.init_global_content_by_config() def tearDown(self): From 2237a6abafc4d36c1379d42bc367c183d285d92a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 28 Jul 2024 20:12:02 +0800 Subject: [PATCH 12/15] fix test file import problem --- examples/api/__init__.py | 0 .../api}/test_interface.py | 0 memoryscope/cli.py | 1 - memoryscope/core/chat/api_memory_chat.py | 5 +- memoryscope/core/chat/cli_memory_chat.py | 5 +- memoryscope/core/config/arguments.py | 4 +- memoryscope/core/config/config_manager.py | 18 +-- memoryscope/core/config/demo_config.yaml | 78 ++++++------- memoryscope/core/memoryscope.py | 2 + memoryscope/core/utils/prompt_handler.py | 6 +- memoryscope/core/worker/memory_base_worker.py | 2 +- tests/models/test_models_lli_embedding.py | 4 +- tests/models/test_models_lli_generation.py | 4 +- tests/models/test_models_lli_rank.py | 2 +- tests/{operations => other}/init_test.py | 0 tests/other/read_yaml.py | 4 +- tests/storages/test_storages_lli_es.py | 4 +- tests/storages/test_storages_lli_synces.py | 4 +- tests/worker/test_workers_cn.py | 106 +++++++++--------- tests/worker/test_workers_en.py | 104 +++++++++-------- 20 files changed, 175 insertions(+), 178 deletions(-) create mode 100644 examples/api/__init__.py rename {tests/operations => examples/api}/test_interface.py (100%) rename tests/{operations => other}/init_test.py (100%) diff --git a/examples/api/__init__.py b/examples/api/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/operations/test_interface.py b/examples/api/test_interface.py similarity index 100% rename from tests/operations/test_interface.py rename to examples/api/test_interface.py diff --git a/memoryscope/cli.py b/memoryscope/cli.py index 1d890358..85d5bec5 100644 --- a/memoryscope/cli.py +++ b/memoryscope/cli.py @@ -8,7 +8,6 @@ from memoryscope.core.memoryscope import MemoryScope def cli_job(**kwargs): - kwargs["memory_chat_type"] = "cli_chat" MemoryScope(**kwargs).default_memory_chat.run() diff --git a/memoryscope/core/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py index d1992d73..6aee7c75 100644 --- a/memoryscope/core/chat/api_memory_chat.py +++ b/memoryscope/core/chat/api_memory_chat.py @@ -53,7 +53,10 @@ class ApiMemoryChat(BaseMemoryChat): PromptHandler: An instance of the PromptHandler configured for this CLI session. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs) + self._prompt_handler = PromptHandler(__file__, + language=self.context.language, + prompt_file="memory_chat_prompt", + **self.kwargs) return self._prompt_handler @property diff --git a/memoryscope/core/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py index f2fe8a0a..653b92ec 100644 --- a/memoryscope/core/chat/cli_memory_chat.py +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -68,7 +68,10 @@ class CliMemoryChat(BaseMemoryChat): PromptHandler: An instance of the PromptHandler configured for this CLI session. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs) + self._prompt_handler = PromptHandler(__file__, + language=self.context.language, + prompt_file="memory_chat_prompt", + **self.kwargs) return self._prompt_handler def print_logo(self): diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index 1925e31e..f6542d3e 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -12,8 +12,8 @@ class Arguments(object): logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S") - memory_chat_type: str = field(default="cli_chat", metadata={ - "help": "cli_chat(Command-line interaction), api_chat(API interface interaction), etc."}) + memory_chat_class: str = field(default="cli_memory_chat", metadata={ + "help": "cli_memory_chat(Command-line interaction), api_memory_chat(API interface interaction), etc."}) consolidate_memory_interval_time: int = field(default=1, metadata={ "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py index 7ac8095c..b23eed5c 100644 --- a/memoryscope/core/config/config_manager.py +++ b/memoryscope/core/config/config_manager.py @@ -65,14 +65,9 @@ class ConfigManager(object): @staticmethod def update_memory_chat_by_arguments(config: dict, arguments: Arguments): - if arguments.memory_chat_type == "cli_chat": - memory_chat_class = "chat.cli_memory_chat" - elif arguments.memory_chat_type == "api_chat": - memory_chat_class = "chat.api_memory_chat" - else: - raise NotImplementedError(f"known memory_chat_type={arguments.memory_chat_type}") + memory_chat_class_split = config["class"].split(".") config.update({ - "class": memory_chat_class, + "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], "assistant_name": "AI", }) @@ -154,7 +149,7 @@ class ConfigManager(object): def clear_node_all(self, node: str): self.config[node].clear() - def dump_config(self, file_type: Literal["json", "yaml"], to_stream: bool = True, file_path: Optional[str] = None): + def dump_config(self, file_type: Literal["json", "yaml"] = "yaml", file_path: Optional[str] = None) -> str: if file_type == "json": content = json.dumps(self.config, indent=2, ensure_ascii=False) elif file_type == "yaml": @@ -162,9 +157,8 @@ class ConfigManager(object): else: raise NotImplementedError - if to_stream: - print(content) - - if file_type: + if file_path: with open(file_path, "w") as f: f.write(content) + + return content diff --git a/memoryscope/core/config/demo_config.yaml b/memoryscope/core/config/demo_config.yaml index e81ba428..88680e7a 100644 --- a/memoryscope/core/config/demo_config.yaml +++ b/memoryscope/core/config/demo_config.yaml @@ -7,78 +7,78 @@ global: memory_chat: cli_memory_chat: - class: chat.cli_memory_chat + class: core.chat.cli_memory_chat memory_service: memoryscope_service generation_model: generation_model memory_service: memoryscope_service: - class: memory.service.memory_scope_service + class: core.service.memory_scope_service memory_operations: read_message: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: read_message description: "read short memory" retrieve_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank description: "retrieve long-term memory" list_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,retrieve_top_memory,print_memory description: "read all long-term memory of the user" delete_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,retrieve_all_memory,delete_memory description: "delete a single long-term memory" delete_all: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,retrieve_all_memory,delete_all description: "delete all long-term memory" add_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: add_memory description: "add a single observation" consolidate_memory: - class: memory.operation.consolidate_memory_op + class: core.operation.consolidate_memory_op workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory description: "summary user's observation memory" interval_time: 1 reflect_and_reconsolidate: - class: memory.operation.backend_operation + class: core.operation.backend_operation workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory description: "summary user's insight memory" interval_time: 15 worker: dummy: - class: memory.worker.dummy_worker + class: core.worker.dummy_worker generation_model: generation_model embedding_model: embedding_model rank_model: rank_model read_message: - class: memory.worker.frontend.read_message_worker + class: core.worker.frontend.read_message_worker set_query: - class: memory.worker.frontend.set_query_worker + class: core.worker.frontend.set_query_worker retrieve_obs_ins: - class: memory.worker.frontend.retrieve_memory_worker + class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 retrieve_ins_top_k: 100 extract_time: - class: memory.worker.frontend.extract_time_worker + class: core.worker.frontend.extract_time_worker generation_model: generation_model semantic_rank: - class: memory.worker.frontend.semantic_rank_worker + class: core.worker.frontend.semantic_rank_worker rank_model: rank_model fuse_rerank: - class: memory.worker.frontend.fuse_rerank_worker + class: core.worker.frontend.fuse_rerank_worker fuse_score_threshold: 0.01 fuse_ratio_dict: conversation: 0.5 @@ -88,88 +88,88 @@ worker: fuse_time_ratio: 2.0 fuse_rerank_top_k: 10 retrieve_top_memory: - class: memory.worker.frontend.retrieve_memory_worker + class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 retrieve_ins_top_k: 100 retrieve_expired_top_k: 100 print_memory: - class: memory.worker.frontend.print_memory_worker + class: core.worker.frontend.print_memory_worker retrieve_all_memory: - class: memory.worker.frontend.retrieve_memory_worker + class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 1000 retrieve_ins_top_k: 1000 retrieve_expired_top_k: 1000 delete_memory: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: delete_memory delete_all: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: delete_all add_memory: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: from_query info_filter: - class: memory.worker.backend.info_filter_worker + class: core.worker.backend.info_filter_worker generation_model: generation_model load_today_memory: - class: memory.worker.backend.load_memory_worker + class: core.worker.backend.load_memory_worker retrieve_today_top_k: 100 get_observation: - class: memory.worker.backend.get_observation_worker + class: core.worker.backend.get_observation_worker generation_model: generation_model get_observation_with_time: - class: memory.worker.backend.get_observation_with_time_worker + class: core.worker.backend.get_observation_with_time_worker generation_model: generation_model contra_repeat: - class: memory.worker.backend.contra_repeat_worker + class: core.worker.backend.contra_repeat_worker generation_model: generation_model store_memory: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: from_memory_key memory_key: all load_obs_and_insight: - class: memory.worker.backend.load_memory_worker + class: core.worker.backend.load_memory_worker retrieve_not_reflected_top_k: 100 retrieve_not_updated_top_k: 100 retrieve_insight_top_k: 100 get_reflection_subject: - class: memory.worker.backend.get_reflection_subject_worker + class: core.worker.backend.get_reflection_subject_worker generation_model: generation_model reflect_obs_cnt_threshold: 10 update_insight: - class: memory.worker.backend.update_insight_worker + class: core.worker.backend.update_insight_worker generation_model: generation_model rank_model: rank_model long_contra_repeat: - class: memory.worker.backend.long_contra_repeat_worker + class: core.worker.backend.long_contra_repeat_worker generation_model: generation_model model: generation_model: - class: models.llama_index_generation_model + class: core.models.llama_index_generation_model module_name: dashscope_generation model_name: qwen-max max_tokens: 2000 embedding_model: - class: models.llama_index_embedding_model + class: core.models.llama_index_embedding_model module_name: dashscope_embedding model_name: text-embedding-v2 rank_model: - class: models.llama_index_rank_model + class: core.models.llama_index_rank_model module_name: dashscope_rank model_name: gte-rerank top_n: 500 dummy_generation: - class: models.dummy_generation_model + class: core.models.dummy_generation_model module_name: dummy_generation model_name: dummy_generation_model memory_store: - class: storage.llama_index_es_memory_store + class: core.storage.llama_index_es_memory_store embedding_model: embedding_model index_name: memory_index es_url: http://localhost:9200 retrieve_mode: dense monitor: - class: storage.dummy_monitor \ No newline at end of file + class: core.storage.dummy_monitor \ No newline at end of file diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py index 3e0c7978..54b429b9 100644 --- a/memoryscope/core/memoryscope.py +++ b/memoryscope/core/memoryscope.py @@ -82,6 +82,8 @@ class MemoryScope(ConfigManager): if self.context.monitor: self.context.monitor.close() + self.logger.close() + def __enter__(self): self.init_context_by_config() diff --git a/memoryscope/core/utils/prompt_handler.py b/memoryscope/core/utils/prompt_handler.py index d44559f3..a60497e4 100644 --- a/memoryscope/core/utils/prompt_handler.py +++ b/memoryscope/core/utils/prompt_handler.py @@ -16,9 +16,9 @@ class PromptHandler(object): def __init__(self, class_path: str, + language: LanguageEnum | str, prompt_file: str = "", prompt_dict: dict = None, - language_enum: LanguageEnum = LanguageEnum.EN, **kwargs): """ Initializes the PromptHandler with paths to prompt sources and additional keyword arguments. @@ -27,13 +27,13 @@ class PromptHandler(object): class_path (str): The path to the class where prompts are utilized. prompt_file (str, optional): The path to an external file containing prompts. Defaults to "". prompt_dict (dict, optional): A dictionary directly containing prompt definitions. Defaults to None. - language_enum (LanguageEnum): context language. + language (LanguageEnum, str): context language. **kwargs: Additional keyword arguments that might be used in prompt handling. """ class_path: Path = Path(class_path) self._class_dir: Path = class_path.parent self._class_name: str = class_path.stem - self._language_enum: LanguageEnum = language_enum + self._language_enum: LanguageEnum = LanguageEnum(language) self.kwargs = kwargs self._prompt_dict: Dict[str, str] = {} diff --git a/memoryscope/core/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py index f4de9044..85b1ad41 100644 --- a/memoryscope/core/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -190,7 +190,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): PromptHandler: An instance of PromptHandler initialized with specific file path and keyword arguments. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs) + self._prompt_handler = PromptHandler(self.FILE_PATH, language=self.language, **self.kwargs) return self._prompt_handler @property diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index f21d0bc7..67d3bb13 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -5,8 +5,8 @@ sys.path.append(".") # noqa: E402 import asyncio import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel -from memoryscope.utils.logger import Logger +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.utils.logger import Logger class TestLLIEmbedding(unittest.TestCase): diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index 474851b4..1fbe41b5 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -6,8 +6,8 @@ import unittest import time import asyncio from memoryscope.scheme.message import Message -from memoryscope.models.llama_index_generation_model import LlamaIndexGenerationModel -from memoryscope.utils.logger import Logger +from memoryscope.core.models.llama_index_generation_model import LlamaIndexGenerationModel +from memoryscope.core.utils.logger import Logger class TestLLILLM(unittest.TestCase): diff --git a/tests/models/test_models_lli_rank.py b/tests/models/test_models_lli_rank.py index 25c59fde..e0e90c7c 100644 --- a/tests/models/test_models_lli_rank.py +++ b/tests/models/test_models_lli_rank.py @@ -1,7 +1,7 @@ import asyncio import unittest -from memoryscope.models.llama_index_rank_model import LlamaIndexRankModel +from memoryscope.core.models.llama_index_rank_model import LlamaIndexRankModel class TestLLIReRank(unittest.TestCase): diff --git a/tests/operations/init_test.py b/tests/other/init_test.py similarity index 100% rename from tests/operations/init_test.py rename to tests/other/init_test.py diff --git a/tests/other/read_yaml.py b/tests/other/read_yaml.py index c54bd6ad..ac0513f4 100644 --- a/tests/other/read_yaml.py +++ b/tests/other/read_yaml.py @@ -2,10 +2,10 @@ import sys sys.path.append(".") # noqa: E402 -from memoryscope.utils.prompt_handler import PromptHandler +from memoryscope.core.utils.prompt_handler import PromptHandler if __name__ == "__main__": file_path: str = __file__ print(file_path) - handler = PromptHandler(__file__, "read_prompt") + handler = PromptHandler(__file__, language="cn", prompt_file="read_prompt", ) print(handler.prompt_dict) diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 791286aa..a72d1aab 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -1,8 +1,8 @@ import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore class TestLlamaIndexElasticSearchStore(unittest.TestCase): diff --git a/tests/storages/test_storages_lli_synces.py b/tests/storages/test_storages_lli_synces.py index bf32bc73..f407a389 100644 --- a/tests/storages/test_storages_lli_synces.py +++ b/tests/storages/test_storages_lli_synces.py @@ -1,8 +1,8 @@ import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore class TestLlamaIndexElasticSearchStore(unittest.TestCase): diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index e7a65fae..77cad3ff 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -1,44 +1,51 @@ import datetime import unittest -from memoryscope.cli import MemoryScope from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \ + MEMORYSCOPE_CONTEXT +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config class TestWorkersCn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) - - ms = MemoryScope() - ms.read_config("config/demo_config_cn.yaml") - ms.init_global_content_by_config() + arguments = Arguments( + language="cn", + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", + ) + self.ms = MemoryScope(arguments=arguments) + config = self.ms.dump_config() + self.ms.logger.info(f"config=\n{config}") def tearDown(self): - self.logger.close() + self.ms.close() @unittest.skip def test_extract_time(self): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) query = "明天我去上海出差" query_timestamp = int(datetime.datetime.now().timestamp()) @@ -53,13 +60,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"), @@ -79,13 +85,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"), @@ -115,13 +120,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"), @@ -149,13 +153,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"), @@ -181,13 +184,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术"), @@ -211,13 +213,12 @@ class TestWorkersCn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), @@ -272,13 +273,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), @@ -304,18 +304,17 @@ class TestWorkersCn(unittest.TestCase): @unittest.skip def test_update_insight_worker(self): - reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self) name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户喜欢打王者荣耀"), @@ -327,18 +326,17 @@ class TestWorkersCn(unittest.TestCase): result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") - @unittest.skip + # @unittest.skip def test_long_contra_repeat_worker(self): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 0123bc3a..ff421253 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -1,44 +1,51 @@ import datetime import unittest -from memoryscope.cli import MemoryScope from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \ + MEMORYSCOPE_CONTEXT +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config class TestWorkersEn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) - - ms = MemoryScope() - ms.read_config("config/demo_config_en.yaml") - ms.init_global_content_by_config() + arguments = Arguments( + language="en", + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", + ) + self.ms = MemoryScope(arguments=arguments) + config = self.ms.dump_config() + self.ms.logger.info(f"config=\n{config}") def tearDown(self): - self.logger.close() + self.ms.close() @unittest.skip def test_extract_time(self): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) query = "I will be on a business trip to Shanghai tomorrow." query_timestamp = int(datetime.datetime.now().timestamp()) @@ -53,13 +60,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."), @@ -80,13 +86,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="Do you know where the freshest seafood is in Beijing?"), @@ -129,13 +134,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) # FIXME Does the appearance of 'am' indicate the presence of a time keyword? chat_messages = [ @@ -159,13 +163,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -202,13 +205,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -235,13 +237,12 @@ class TestWorkersEn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), @@ -294,13 +295,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), @@ -327,18 +327,17 @@ class TestWorkersEn(unittest.TestCase): @unittest.skip def test_update_insight_worker(self): - reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self) name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users like to play King of Glory"), @@ -355,13 +354,12 @@ class TestWorkersEn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), From db1e44a826712d40c4987d58dfc5658e8bc98807 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 28 Jul 2024 22:04:45 +0800 Subject: [PATCH 13/15] [dev] add chat examples --- examples/__init__.py | 0 examples/api/__init__.py | 0 examples/api/chat_example.py | 66 ++++++++ examples/api/test_interface.py | 18 --- examples/cli/dash_cli_cn1.sh | 1 + examples/cli/dash_cli_cn2.sh | 10 ++ examples/config/demo_config_no_stream.yaml | 0 examples/docker/__init__.py | 0 memoryscope/__init__.py | 3 + memoryscope/cli.py | 4 +- memoryscope/core/chat/api_memory_chat.py | 41 +++-- memoryscope/core/chat/base_memory_chat.py | 20 ++- memoryscope/core/chat/cli_memory_chat.py | 168 +++------------------ memoryscope/core/config/arguments.py | 2 + memoryscope/core/config/config_manager.py | 37 ++++- memoryscope/core/config/demo_config.yaml | 3 +- memoryscope/core/memoryscope.py | 15 +- memoryscope/core/utils/logger.py | 3 +- 18 files changed, 177 insertions(+), 214 deletions(-) delete mode 100644 examples/__init__.py delete mode 100644 examples/api/__init__.py create mode 100644 examples/api/chat_example.py delete mode 100644 examples/api/test_interface.py create mode 100644 examples/cli/dash_cli_cn1.sh create mode 100644 examples/cli/dash_cli_cn2.sh delete mode 100644 examples/config/demo_config_no_stream.yaml delete mode 100644 examples/docker/__init__.py diff --git a/examples/__init__.py b/examples/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/examples/api/__init__.py b/examples/api/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/examples/api/chat_example.py b/examples/api/chat_example.py new file mode 100644 index 00000000..b7781ef5 --- /dev/null +++ b/examples/api/chat_example.py @@ -0,0 +1,66 @@ +import sys + +sys.path.append(".") + +from memoryscope import MemoryScope, Arguments + +arguments = Arguments( + language="cn", + logger_to_screen=False, + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", +) + + +def chat_example1(): + with MemoryScope(config_path="memoryscope/core/config/demo_config.yaml") as ms: + memory_chat = ms.default_memory_chat + response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") + print("回答:\n" + response.message.content) + print("记忆:\n" + response.meta_data["memories"]) + + +def chat_example2(): + with MemoryScope(arguments=arguments) as ms: + memory_chat = ms.default_memory_chat + + response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") + print("回答:\n" + response.message.content) + print("记忆:\n" + response.meta_data["memories"]) + + +def chat_example3(): + with MemoryScope(arguments=arguments) as ms: + memory_chat = ms.default_memory_chat + + response = memory_chat.chat_with_memory( + query="你知道我消灭了几个星系吗?", + extra_memories="用户是一只哥斯拉,这是真是的,不是虚构的。\n前天用户消灭了星系0x4be?\n大前天用户消灭了星系0x5be?") + print("回答:\n" + response.message.content) + + +def chat_example4(): + with MemoryScope(arguments=arguments) as ms: + memory_chat = ms.default_memory_chat + + response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") + print("回答1:\n" + response.message.content) + memory_chat.memory_service.consolidate_memory() + + response = memory_chat.chat_with_memory(query="你知道我的乐器爱好是什么?", + add_messages=False) + print("回答2:\n" + response.message.content) + print("记忆2:\n" + response.meta_data["memories"]) + + +if __name__ == "__main__": + chat_example1() + # chat_example2() + # chat_example3() + # chat_example4() diff --git a/examples/api/test_interface.py b/examples/api/test_interface.py deleted file mode 100644 index 05bdaef5..00000000 --- a/examples/api/test_interface.py +++ /dev/null @@ -1,18 +0,0 @@ -from memoryscope.cli import MemoryScope -from memoryscope.scheme.message import Message - -ms = MemoryScope().read_config("config/demo_config_no_stream.yaml") -memory_service = ms.default_service -memory_chat = ms.default_chat_handle - -# new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name="我", content="我的爱好是弹琴并且喜欢看电影。") -# memory_service.add_messages(new_message) - -res: Message = memory_chat.chat_with_memory(query="我的爱好是弹琴。", remember_response=True) -print(res.message.content) - -res: Message = memory_chat.chat_with_memory(query="昨天弹出一个光粒,消灭了星系0x4be。", remember_response=True) -print(res.message.content) - -res: Message = memory_chat.chat_with_memory(query="今天弹出一个二向箔,消灭了星系0xa2e。", remember_response=True) -print(res.message.content) diff --git a/examples/cli/dash_cli_cn1.sh b/examples/cli/dash_cli_cn1.sh new file mode 100644 index 00000000..2998bd29 --- /dev/null +++ b/examples/cli/dash_cli_cn1.sh @@ -0,0 +1 @@ +python memoryscope/cli.py -config_path=memoryscope/core/config/demo_config.yaml \ No newline at end of file diff --git a/examples/cli/dash_cli_cn2.sh b/examples/cli/dash_cli_cn2.sh new file mode 100644 index 00000000..14c92cee --- /dev/null +++ b/examples/cli/dash_cli_cn2.sh @@ -0,0 +1,10 @@ +python memoryscope/cli.py \ + -language="cn" \ + -memory_chat_class="cli_memory_chat" \ + -generation_backend="dashscope_generation" \ + -generation_model="qwen-max" \ + -embedding_backend="dashscope_embedding" \ + -embedding_model="text-embedding-v2" \ + -use_dummy_ranker=False \ + -rank_backend="dashscope_rank" \ + -rank_model="gte-rerank" \ No newline at end of file diff --git a/examples/config/demo_config_no_stream.yaml b/examples/config/demo_config_no_stream.yaml deleted file mode 100644 index e69de29b..00000000 diff --git a/examples/docker/__init__.py b/examples/docker/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/memoryscope/__init__.py b/memoryscope/__init__.py index 2be4eb6a..0dc03c13 100644 --- a/memoryscope/__init__.py +++ b/memoryscope/__init__.py @@ -1,2 +1,5 @@ +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope + """ Version of MemoryScope.""" __version__ = "0.1.0" diff --git a/memoryscope/cli.py b/memoryscope/cli.py index 85d5bec5..a429fda0 100644 --- a/memoryscope/cli.py +++ b/memoryscope/cli.py @@ -8,7 +8,9 @@ from memoryscope.core.memoryscope import MemoryScope def cli_job(**kwargs): - MemoryScope(**kwargs).default_memory_chat.run() + with MemoryScope(**kwargs) as ms: + memory_chat = ms.default_memory_chat + memory_chat.run() if __name__ == "__main__": diff --git a/memoryscope/core/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py index 6aee7c75..d1e8d491 100644 --- a/memoryscope/core/chat/api_memory_chat.py +++ b/memoryscope/core/chat/api_memory_chat.py @@ -9,7 +9,7 @@ from memoryscope.core.service.base_memory_service import BaseMemoryService from memoryscope.core.utils.prompt_handler import PromptHandler from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.scheme.message import Message -from memoryscope.scheme.model_response import ModelResponse +from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen class ApiMemoryChat(BaseMemoryChat): @@ -18,6 +18,7 @@ class ApiMemoryChat(BaseMemoryChat): memory_service: str, generation_model: str, context: MemoryscopeContext, + stream: bool = False, human_name: str = None, assistant_name: str = None, **kwargs): @@ -27,6 +28,7 @@ class ApiMemoryChat(BaseMemoryChat): self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model self.context: MemoryscopeContext = context + self.stream: bool = stream self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {}) self.human_name: str = human_name @@ -100,13 +102,31 @@ class ApiMemoryChat(BaseMemoryChat): self._generation_model = self.context.model_dict[self._generation_model] return self._generation_model + def iter_response(self, + remember_response: bool, + resp: ModelResponseGen, + memories: str, + query_message: Message) -> ModelResponseGen: + + model_response: ModelResponse | None = None + for model_response in resp: + yield model_response + + if remember_response: + if model_response and model_response.message: + model_response.message.role_name = self.assistant_name + model_response.meta_data[MEMORIES] = memories + self.memory_service.add_messages([query_message, model_response.message]) + else: + self.logger.warning("model_response or model_response.message is empty!") + def chat_with_memory(self, query: str, role_name: Optional[str] = None, system_prompt: Optional[str] = None, memory_prompt: Optional[str] = None, extra_memories: Optional[str] = None, - add_not_memorized_messages: bool = True, + add_messages: bool = True, remember_response: bool = True, **kwargs): """ @@ -118,7 +138,7 @@ class ApiMemoryChat(BaseMemoryChat): system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml". memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml". extra_memories (str, optional): Manually added user memory in this function. - add_not_memorized_messages (bool, optional): whether add not memorized messages to LLM. + add_messages (bool, optional): whether add not memorized messages to LLM. remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. Defaults to False. Returns: @@ -161,7 +181,7 @@ class ApiMemoryChat(BaseMemoryChat): chat_messages.append(system_message) # Include past conversation history in the message list - if add_not_memorized_messages: + if add_messages: history_messages = self.memory_service.read_message() if history_messages: chat_messages.extend(history_messages) @@ -171,19 +191,8 @@ class ApiMemoryChat(BaseMemoryChat): self.logger.info(f"chat_messages={chat_messages}") resp = self.generation_model.call(messages=chat_messages, stream=self.stream, **self.generation_model_kwargs) - if self.stream: - model_response: ModelResponse | None = None - for model_response in resp: - yield model_response - - if remember_response: - if model_response and model_response.message: - model_response.message.role_name = self.assistant_name - model_response.meta_data[MEMORIES] = memories - self.memory_service.add_messages([query_message, model_response.message]) - else: - self.logger.warning("model_response or model_response.message is empty!") + return self.iter_response(remember_response, resp, memories, query_message) else: model_response: ModelResponse = resp diff --git a/memoryscope/core/chat/base_memory_chat.py b/memoryscope/core/chat/base_memory_chat.py index 995d1d46..455c93e4 100644 --- a/memoryscope/core/chat/base_memory_chat.py +++ b/memoryscope/core/chat/base_memory_chat.py @@ -1,4 +1,5 @@ from abc import ABCMeta, abstractmethod +from typing import Optional from memoryscope.core.service.base_memory_service import BaseMemoryService from memoryscope.core.utils.logger import Logger @@ -10,8 +11,7 @@ class BaseMemoryChat(metaclass=ABCMeta): It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement. """ - def __init__(self, stream: bool = True, **kwargs): - self.stream: bool = stream + def __init__(self, **kwargs): self.kwargs: dict = kwargs self.logger = Logger.get_logger() @@ -28,16 +28,14 @@ class BaseMemoryChat(metaclass=ABCMeta): @abstractmethod def chat_with_memory(self, query: str, - role_name: str = "", - remember_response: bool = True): - """ - Initiates a chat interaction using the memory service, with the provided query as input. + role_name: Optional[str] = None, + system_prompt: Optional[str] = None, + memory_prompt: Optional[str] = None, + extra_memories: Optional[str] = None, + add_messages: bool = True, + remember_response: bool = True, + **kwargs): - Args: - query (str): The user's query or message to start the chat. - role_name (str): The role's name. - remember_response (bool): whether update memory service. - """ raise NotImplementedError def run(self): diff --git a/memoryscope/core/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py index 653b92ec..3d610468 100644 --- a/memoryscope/core/chat/cli_memory_chat.py +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -1,22 +1,14 @@ import os import time -from typing import List +from typing import Optional import questionary -from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME -from memoryscope.core.chat.base_memory_chat import BaseMemoryChat -from memoryscope.core.memoryscope_context import MemoryscopeContext -from memoryscope.core.models.base_model import BaseModel -from memoryscope.core.service.base_memory_service import BaseMemoryService -from memoryscope.core.utils.prompt_handler import PromptHandler +from memoryscope.core.chat.api_memory_chat import ApiMemoryChat from memoryscope.core.utils.tool_functions import char_logo -from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.scheme.message import Message -from memoryscope.scheme.model_response import ModelResponse -class CliMemoryChat(BaseMemoryChat): +class CliMemoryChat(ApiMemoryChat): """ Command-line interface for chatting with an AI that integrates memory functionality. Allows users to interact, manage chat history, adjust streaming settings, and view commands' help. @@ -28,51 +20,9 @@ class CliMemoryChat(BaseMemoryChat): "stream": "Toggle between getting streamed responses from the model." } - def __init__(self, - memory_service: str, - generation_model: str, - context: MemoryscopeContext, - human_name: str = None, - assistant_name: str = None, - **kwargs): - + def __init__(self, **kwargs): super().__init__(**kwargs) - - self._memory_service: BaseMemoryService | str = memory_service - self._generation_model: BaseModel | str = generation_model - self.context: MemoryscopeContext = context - self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {}) - - self.human_name: str = human_name - if not self.human_name: - self.human_name = DEFAULT_HUMAN_NAME[self.context.language] - self.context.meta_data["human_name"] = self.human_name - - self.assistant_name: str = assistant_name - if not self.assistant_name: - self.assistant_name = "AI" - self.context.meta_data["assistant_name"] = self.assistant_name - self._logo = char_logo("MemoryScope") - self._prompt_handler: PromptHandler | None = None - - @property - def prompt_handler(self) -> PromptHandler: - """ - Lazy initialization property for the prompt handler. - - This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed. - It uses the current file's path and additional keyword arguments for configuration. - - Returns: - PromptHandler: An instance of the PromptHandler configured for this CLI session. - """ - if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, - language=self.context.language, - prompt_file="memory_chat_prompt", - **self.kwargs) - return self._prompt_handler def print_logo(self): """ @@ -84,104 +34,30 @@ class CliMemoryChat(BaseMemoryChat): for line in self._logo: print(line) - @property - def memory_service(self) -> BaseMemoryService: - """ - Property to access the memory service. If the service is initially set as a string, - it will be looked up in the memory service dictionary of context, initialized, - and then returned as an instance of `BaseMemoryService`. Ensures the memory service - is properly started before use. - - Returns: - BaseMemoryService: An active memory service instance. - - Raises: - ValueError: If the declaration of memory service is not found in the memory service dictionary of context. - """ - if isinstance(self._memory_service, str): - if self._memory_service not in self.context.memory_service_dict: - raise ValueError(f"Missing declaration of memory_service in context: {self._memory_service}") - - self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] - # init service & update kwargs - self._memory_service.init_service() - return self._memory_service - - @property - def generation_model(self) -> BaseModel: - """ - Property to get the generation model. If the model is set as a string, it will be resolved from the global - context's model dictionary. - - Raises: - ValueError: If the declaration of generation model is not found in the model dictionary of context . - - Returns: - BaseModel: An actual generation model instance. - """ - if isinstance(self._generation_model, str): - if self._generation_model not in self.context.model_dict: - raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}") - self._generation_model = self.context.model_dict[self._generation_model] - return self._generation_model - - def get_user_message(self, query: str, role_name: str = "") -> Message: - if not role_name: - role_name = self.human_name - return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) - - def get_system_message_with_memory(self, memories: str) -> Message: - # Incorporate memory into the system prompt if available - system_prompt = self.prompt_handler.system_prompt - if memories: - memory_prompt = self.prompt_handler.memory_prompt - system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) - return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) - def chat_with_memory(self, query: str, - role_name: str = "", - remember_response: bool = True): - - chat_messages: List[Message] = [] - - new_message: Message = self.get_user_message(query=query, role_name=role_name) - - # To retrieve memory, prepare the query timestamp and role name by adding new_message. - memories: str = self.memory_service.retrieve_memory(query=new_message.content, - role_name=new_message.role_name, - timestamp=new_message.time_created) - - # format system_message with memories - system_message: Message = self.get_system_message_with_memory(memories=memories) - chat_messages.append(system_message) - - # Include past conversation history in the message list - history_messages = self.memory_service.read_message() - if history_messages: - chat_messages.extend(history_messages) - - # Append the current user's message to the conversation context - chat_messages.append(new_message) - self.logger.info(f"chat_messages={chat_messages}") - - # Invoke the Language Model with the constructed message context, respecting streaming setting - resp = self.generation_model.call(messages=chat_messages, - stream=self.stream, - **self.generation_model_kwargs) + role_name: Optional[str] = None, + system_prompt: Optional[str] = None, + memory_prompt: Optional[str] = None, + extra_memories: Optional[str] = None, + add_messages: bool = True, + remember_response: bool = True, + **kwargs): + resp = super().chat_with_memory(query=query, + role_name=role_name, + system_prompt=system_prompt, + memory_prompt=memory_prompt, + extra_memories=extra_memories, + add_messages=add_messages, + remember_response=remember_response, + **kwargs) if self.stream: - model_response: ModelResponse | None = None - for model_response in resp: - questionary.print(model_response.delta, end="") + for _resp in resp: + questionary.print(_resp.delta, end="") questionary.print("") else: - model_response: ModelResponse = resp - questionary.print(model_response.message.content) - - if remember_response and model_response and model_response.message: - model_response.message.role_name = self.assistant_name - self.memory_service.add_messages([new_message, model_response.message]) + questionary.print(resp.message.content) @staticmethod def parse_query_command(query: str): diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index f6542d3e..44ed574f 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -12,6 +12,8 @@ class Arguments(object): logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S") + logger_to_screen: bool = field(default=False, metadata={"help": "If false, it does not print to the screen."}) + memory_chat_class: str = field(default="cli_memory_chat", metadata={ "help": "cli_memory_chat(Command-line interaction), api_memory_chat(API interface interaction), etc."}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py index b23eed5c..e30d51e3 100644 --- a/memoryscope/core/config/config_manager.py +++ b/memoryscope/core/config/config_manager.py @@ -1,5 +1,6 @@ import json from dataclasses import fields +from datetime import datetime from pathlib import Path from typing import Optional, Literal @@ -7,6 +8,7 @@ import yaml from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME from memoryscope.core.config.arguments import Arguments +from memoryscope.core.utils.logger import Logger from memoryscope.enumeration.language_enum import LanguageEnum @@ -23,20 +25,40 @@ class ConfigManager(object): if config: self.config = config + self.logger = self._init_logger() + self.logger.info("init by config mode:") elif config_path: self.read_config(config_path) + self.logger = self._init_logger() + self.logger.info("init by config_path mode:") else: self.read_demo_config(demo_config_name) + if arguments: + self.update_config_by_arguments(arguments) + self.logger = self._init_logger() + self.logger.info(f"init by arguments mode: {arguments.__dict__}") - if arguments: - self.update_config_by_arguments(arguments) + elif kwargs: + kwargs = {k: v for k, v in kwargs.items() if k in [x.name for x in fields(Arguments)]} + arguments = Arguments(**kwargs) + self.update_config_by_arguments(arguments) + self.logger = self._init_logger() + self.logger.info(f"init by kwargs mode: {kwargs}") - elif kwargs: - key_list = [x.name for x in fields(Arguments)] - arguments = Arguments(**{k: v for k, v in kwargs.items() if k in key_list}) - self.update_config_by_arguments(arguments) + else: + raise RuntimeError("can not init config manager without kwargs!") + self.logger.info(self.dump_config()) + + def _init_logger(self) -> Logger: + global_config = self.config["global"] + logger_name = global_config["logger_name"] + logger_name_time_suffix = global_config["logger_name_time_suffix"] + if logger_name_time_suffix: + suffix = datetime.now().strftime(logger_name_time_suffix) + logger_name = f"{logger_name}_{suffix}" + return Logger.get_logger(logger_name, to_stream=global_config["logger_to_screen"]) def read_config(self, config_path: str): if config_path.endswith(".yaml"): @@ -60,16 +82,19 @@ class ConfigManager(object): "thread_pool_max_workers": arguments.thread_pool_max_workers, "logger_name": arguments.logger_name, "logger_name_time_suffix": arguments.logger_name_time_suffix, + "logger_to_screen": arguments.logger_to_screen, "use_dummy_ranker": arguments.use_dummy_ranker, }) @staticmethod def update_memory_chat_by_arguments(config: dict, arguments: Arguments): memory_chat_class_split = config["class"].split(".") + stream = arguments.memory_chat_class in ["cli_memory_chat", ] config.update({ "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], "assistant_name": "AI", + "stream": stream, }) @staticmethod diff --git a/memoryscope/core/config/demo_config.yaml b/memoryscope/core/config/demo_config.yaml index 88680e7a..60db4b34 100644 --- a/memoryscope/core/config/demo_config.yaml +++ b/memoryscope/core/config/demo_config.yaml @@ -3,6 +3,7 @@ global: thread_pool_max_workers: 5 logger_name: memoryscope logger_name_time_suffix: "%Y%m%d_%H%M%S" + logger_to_screen: false use_dummy_ranker: false memory_chat: @@ -86,7 +87,7 @@ worker: obs_customized: 1.2 insight: 2.0 fuse_time_ratio: 2.0 - fuse_rerank_top_k: 10 + fuse_rerank_top_k: 20 retrieve_top_memory: class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py index 54b429b9..1f1ece7b 100644 --- a/memoryscope/core/memoryscope.py +++ b/memoryscope/core/memoryscope.py @@ -1,11 +1,9 @@ -import datetime from concurrent.futures import ThreadPoolExecutor from memoryscope.core.chat.base_memory_chat import BaseMemoryChat from memoryscope.core.config.config_manager import ConfigManager from memoryscope.core.memoryscope_context import MemoryscopeContext from memoryscope.core.service.base_memory_service import BaseMemoryService -from memoryscope.core.utils.logger import Logger from memoryscope.core.utils.tool_functions import init_instance_by_config from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.enumeration.model_enum import ModelEnum @@ -15,21 +13,9 @@ class MemoryScope(ConfigManager): def __init__(self, **kwargs): super().__init__(**kwargs) - - self.logger = self._init_logger() - self.context: MemoryscopeContext = MemoryscopeContext() self.init_context_by_config() - def _init_logger(self) -> Logger: - global_config = self.config["global"] - logger_name = global_config["logger_name"] - logger_name_time_suffix = global_config["logger_name_time_suffix"] - if logger_name_time_suffix: - suffix = datetime.datetime.now().strftime(logger_name_time_suffix) - logger_name = f"{logger_name}_{suffix}" - return Logger.get_logger(logger_name, to_stream=False) - def init_context_by_config(self): # set global config global_conf = self.config["global"] @@ -86,6 +72,7 @@ class MemoryScope(ConfigManager): def __enter__(self): self.init_context_by_config() + return self def __exit__(self, exc_type, exc_val, exc_tb): self.close() diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 9f400c6d..52d286f8 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -105,7 +105,8 @@ class Logger(logging.Logger): by the handlers are freed properly. """ for handler in self.handlers: - handler.close() # ⭐ Close each handler to release resources + # Close each handler to release resources + handler.close() def clear(self): """ From ec5d76956492038a87b51664e8579c5a659deb75 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 28 Jul 2024 22:05:36 +0800 Subject: [PATCH 14/15] [dev] format code --- memoryscope/core/chat/base_memory_chat.py | 1 - memoryscope/core/storage/llama_index_es_memory_store.py | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/memoryscope/core/chat/base_memory_chat.py b/memoryscope/core/chat/base_memory_chat.py index 455c93e4..5707c398 100644 --- a/memoryscope/core/chat/base_memory_chat.py +++ b/memoryscope/core/chat/base_memory_chat.py @@ -35,7 +35,6 @@ class BaseMemoryChat(metaclass=ABCMeta): add_messages: bool = True, remember_response: bool = True, **kwargs): - raise NotImplementedError def run(self): diff --git a/memoryscope/core/storage/llama_index_es_memory_store.py b/memoryscope/core/storage/llama_index_es_memory_store.py index f3967039..dabdbc60 100644 --- a/memoryscope/core/storage/llama_index_es_memory_store.py +++ b/memoryscope/core/storage/llama_index_es_memory_store.py @@ -21,7 +21,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): index_name: str, es_url: str, retrieve_mode: str = "dense", - hybrid_alpha: float = None, + hybrid_alpha: float = None, **kwargs): self.emb_dims = None self.index_name = index_name From d2cd11707b6b00aa860c3fe7b53dd258b5a73bcb Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 29 Jul 2024 00:57:23 +0800 Subject: [PATCH 15/15] [dev] add history_message_strategy to chat with memory func --- examples/api/chat_example.py | 2 +- memoryscope/core/chat/api_memory_chat.py | 21 ++++++++++++---- memoryscope/core/chat/base_memory_chat.py | 25 +++++++++++++++++-- memoryscope/core/chat/cli_memory_chat.py | 6 ++--- .../core/operation/backend_operation.py | 2 +- memoryscope/core/operation/base_operation.py | 2 +- .../core/service/base_memory_service.py | 3 ++- .../core/service/memory_scope_service.py | 13 ++++++---- .../worker/backend/update_insight_worker.py | 15 +++++------ 9 files changed, 63 insertions(+), 26 deletions(-) diff --git a/examples/api/chat_example.py b/examples/api/chat_example.py index b7781ef5..393b004e 100644 --- a/examples/api/chat_example.py +++ b/examples/api/chat_example.py @@ -54,7 +54,7 @@ def chat_example4(): memory_chat.memory_service.consolidate_memory() response = memory_chat.chat_with_memory(query="你知道我的乐器爱好是什么?", - add_messages=False) + history_message_strategy=None) print("回答2:\n" + response.message.content) print("记忆2:\n" + response.meta_data["memories"]) diff --git a/memoryscope/core/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py index d1e8d491..6a8078f7 100644 --- a/memoryscope/core/chat/api_memory_chat.py +++ b/memoryscope/core/chat/api_memory_chat.py @@ -1,4 +1,4 @@ -from typing import List, Optional +from typing import List, Optional, Literal from memoryscope.constants.common_constants import MEMORIES from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME @@ -126,7 +126,7 @@ class ApiMemoryChat(BaseMemoryChat): system_prompt: Optional[str] = None, memory_prompt: Optional[str] = None, extra_memories: Optional[str] = None, - add_messages: bool = True, + history_message_strategy: Literal["auto", None] | int = "auto", remember_response: bool = True, **kwargs): """ @@ -138,7 +138,11 @@ class ApiMemoryChat(BaseMemoryChat): system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml". memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml". extra_memories (str, optional): Manually added user memory in this function. - add_messages (bool, optional): whether add not memorized messages to LLM. + history_message_strategy ("auto", None, int): + - If it is set to "auto", the history messages in the conversation will retain those that have not + yet been summarized. Default to "auto". + - If it is set to None, no conversation history will be saved. + - If it is set to an integer value "n", the most recent "n" messages will be retained. remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. Defaults to False. Returns: @@ -181,8 +185,15 @@ class ApiMemoryChat(BaseMemoryChat): chat_messages.append(system_message) # Include past conversation history in the message list - if add_messages: - history_messages = self.memory_service.read_message() + if history_message_strategy: + history_messages = [] + + if history_message_strategy == "auto": + history_messages = self.memory_service.read_message() + + elif isinstance(history_message_strategy, int): + history_messages = self.memory_service.chat_messages[-history_message_strategy:] + if history_messages: chat_messages.extend(history_messages) diff --git a/memoryscope/core/chat/base_memory_chat.py b/memoryscope/core/chat/base_memory_chat.py index 5707c398..e311f072 100644 --- a/memoryscope/core/chat/base_memory_chat.py +++ b/memoryscope/core/chat/base_memory_chat.py @@ -1,5 +1,5 @@ from abc import ABCMeta, abstractmethod -from typing import Optional +from typing import Optional, Literal from memoryscope.core.service.base_memory_service import BaseMemoryService from memoryscope.core.utils.logger import Logger @@ -32,9 +32,30 @@ class BaseMemoryChat(metaclass=ABCMeta): system_prompt: Optional[str] = None, memory_prompt: Optional[str] = None, extra_memories: Optional[str] = None, - add_messages: bool = True, + history_message_strategy: Literal["auto", None] | int = "auto", remember_response: bool = True, **kwargs): + """ + The core function that carries out conversation with memory accepts user queries through query and returns the + conversation results through model_response. The retrieved memories are stored in the memories within meta_data. + Args: + query (str, optional): User's query, includes the user's question. + role_name (str, optional): User's role name. + system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml". + memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml". + extra_memories (str, optional): Manually added user memory in this function. + history_message_strategy ("auto", None, int): + - If it is set to "auto", the history messages in the conversation will retain those that have not + yet been summarized. Default to "auto". + - If it is set to None, no conversation history will be saved. + - If it is set to an integer value "n", the most recent "n" messages will be retained. + remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. + Defaults to False. + Returns: + - ModelResponse: In non-streaming mode, returns a complete AI response. + - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. + - Memories: To obtain the memory by invoking the method of model_response.meta_data[MEMORIES] + """ raise NotImplementedError def run(self): diff --git a/memoryscope/core/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py index 3d610468..150c79b3 100644 --- a/memoryscope/core/chat/cli_memory_chat.py +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -1,6 +1,6 @@ import os import time -from typing import Optional +from typing import Optional, Literal import questionary @@ -40,7 +40,7 @@ class CliMemoryChat(ApiMemoryChat): system_prompt: Optional[str] = None, memory_prompt: Optional[str] = None, extra_memories: Optional[str] = None, - add_messages: bool = True, + history_message_strategy: Literal["auto", None] | int = "auto", remember_response: bool = True, **kwargs): resp = super().chat_with_memory(query=query, @@ -48,7 +48,7 @@ class CliMemoryChat(ApiMemoryChat): system_prompt=system_prompt, memory_prompt=memory_prompt, extra_memories=extra_memories, - add_messages=add_messages, + history_message_strategy=history_message_strategy, remember_response=remember_response, **kwargs) diff --git a/memoryscope/core/operation/backend_operation.py b/memoryscope/core/operation/backend_operation.py index fd6a55a4..e4b4a41c 100644 --- a/memoryscope/core/operation/backend_operation.py +++ b/memoryscope/core/operation/backend_operation.py @@ -103,7 +103,7 @@ class BackendOperation(BaseWorkflow, BaseOperation): if self._loop_switch: self.run_operation() - def run_operation_backend(self): + def start_operation_backend(self): """ Initiates the background operation loop if it's not already running. Sets the _loop_switch to True and submits the _loop_operation to a thread from the global thread pool. diff --git a/memoryscope/core/operation/base_operation.py b/memoryscope/core/operation/base_operation.py index 600e3dd4..b5276a7c 100644 --- a/memoryscope/core/operation/base_operation.py +++ b/memoryscope/core/operation/base_operation.py @@ -50,7 +50,7 @@ class BaseOperation(metaclass=ABCMeta): """ raise NotImplementedError - def run_operation_backend(self): + def start_operation_backend(self): """ Placeholder method for running an operation specific to the backend. Intended to be overridden by subclasses if backend operations are required. diff --git a/memoryscope/core/service/base_memory_service.py b/memoryscope/core/service/base_memory_service.py index d997fa98..68acc566 100644 --- a/memoryscope/core/service/base_memory_service.py +++ b/memoryscope/core/service/base_memory_service.py @@ -28,6 +28,7 @@ class BaseMemoryService(metaclass=ABCMeta): self.kwargs = kwargs self._operation_dict: Dict[str, BaseOperation] = {} + self.chat_messages: List[Message] = [] self.logger = Logger.get_logger() @property @@ -51,7 +52,7 @@ class BaseMemoryService(metaclass=ABCMeta): def init_service(self, **kwargs): raise NotImplementedError - def start_backend_service(self): + def start_backend_service(self, name: str = None): pass def stop_backend_service(self, wait_service_end: bool = False): diff --git a/memoryscope/core/service/memory_scope_service.py b/memoryscope/core/service/memory_scope_service.py index 88c644d4..e60dfd2f 100644 --- a/memoryscope/core/service/memory_scope_service.py +++ b/memoryscope/core/service/memory_scope_service.py @@ -37,7 +37,6 @@ class MemoryScopeService(BaseMemoryService): if assistant_name: self.context.meta_data["assistant_name"] = assistant_name - self.chat_messages: List[Message] = [] self.message_lock = threading.Lock() def add_messages(self, messages: List[Message] | Message): @@ -89,13 +88,17 @@ class MemoryScopeService(BaseMemoryService): for name, operation_config in self.memory_operations_conf.items(): self.register_operation(name, operation_config, **kwargs) - def start_backend_service(self): + def start_backend_service(self, name: str = None): """ Start all backend operations. """ - for _, operation in self._operation_dict.items(): - if operation.operation_type == "backend": - operation.run_operation_backend() + for op_name, operation in self._operation_dict.items(): + if name: + if op_name == name: + operation.start_operation_backend() + else: + if operation.operation_type == "backend": + operation.start_operation_backend() def stop_backend_service(self, wait_service_end: bool = False): """ diff --git a/memoryscope/core/worker/backend/update_insight_worker.py b/memoryscope/core/worker/backend/update_insight_worker.py index fbc7698b..51f8d7e9 100644 --- a/memoryscope/core/worker/backend/update_insight_worker.py +++ b/memoryscope/core/worker/backend/update_insight_worker.py @@ -56,14 +56,15 @@ class UpdateInsightWorker(MemoryBaseWorker): return insight_node, filtered_nodes, max_score if use_dummy_ranker: - key_vector: List[float] = self.embedding_model.call(text=insight_node.key).embedding_results - if not key_vector: - self.logger.warning(f"embedding call {insight_node.key} failed!") - return insight_node, filtered_nodes, max_score + if not insight_node.key_vector: + key_vector: List[float] = self.embedding_model.call(text=insight_node.key).embedding_results + if not key_vector: + self.logger.warning(f"embedding call {insight_node.key} failed!") + return insight_node, filtered_nodes, max_score - insight_node.key_vector = key_vector - documents_vector = [x.vector for x in obs_nodes] - score_recall_list = cosine_similarity(key_vector, documents_vector) + insight_node.key_vector = key_vector + + score_recall_list = cosine_similarity(insight_node.key_vector, [x.vector for x in obs_nodes]) assert len(score_recall_list) == len(obs_nodes), \ f"size is not as excepted. {len(score_recall_list)} v.s. {len(obs_nodes)}"