From a5b45afa8c2de4153d392e2ce899f629b70d85f0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 27 Jul 2024 21:54:19 +0800 Subject: [PATCH] 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()