From b4bb7f9ef19cc9f9e4b6942b2d6f7137882d995a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Fri, 12 Jul 2024 18:37:46 +0800 Subject: [PATCH] bug fix --- config/test_config.yaml | 153 ++++++++++++------ .../memory/worker/memory_base_worker.py | 20 +-- .../storage/llama_index_es_memory_store.py | 8 +- 3 files changed, 113 insertions(+), 68 deletions(-) diff --git a/config/test_config.yaml b/config/test_config.yaml index a6f71411..d5ee7c77 100644 --- a/config/test_config.yaml +++ b/config/test_config.yaml @@ -1,89 +1,152 @@ 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 + class: chat.cli_memory_chat 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 + class: memory.service.chat_memory_service 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 + read_message: + class: memory.operation.read_message description: "read session messages of the user" read_memory: class: memory.operation.read_memory - workflow: dummy_workflow + workflow: set_query,[extract_time|retrieve_memory1,semantic_rank],fuse_rerank description: "read related memories of the user" list_memory: class: memory.operation.read_memory - workflow: dummy_workflow + workflow: set_query,retrieve_memory2,print_memory description: "read all memories of the user" write_memory: class: memory.operation.write_memory - workflow: dummy_workflow + workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory description: "write observation memories of the user" - interval_time: 60 + interval_time: 5 summary_memory: class: memory.operation.summary_memory - workflow: dummy_workflow + workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory 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 + interval_time: 30 worker: - dummy_workflow: + dummy: 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 + 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 - fuse_rerank_worker: + 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 + rank_model: dashscope_rank + 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 + preserved_scores: 2,3 + 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: 10 + 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 + dummy_generation: + class: models.dummy_generation_model + module_name: dummy_generation + model_name: dummy_generation_model memory_store: - class: storage.llama_index_es_memory_store + class: storage.llama_index_es_memory_store_sync embedding_model: dashscope_embedding index_name: memory_index - es_url: http://11.160.132.46:9200 - use_hybrid: false \ No newline at end of file + es_url: http://localhost:9200 + use_hybrid: false + +monitor: + class: storage.dummy_monitor \ No newline at end of file diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index db885e95..66beed0e 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -47,31 +47,13 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): @property def chat_messages(self) -> List[Message]: - """ - Getter property to retrieve the list of chat messages from the context. - - Returns: - List[Message]: A list of Message objects representing the chat messages. - """ return self.get_context(CHAT_MESSAGES) @chat_messages.setter - def chat_messages(self, messages: List[Message]) -> None: - """ - Setter property to update the list of chat messages in the context. - - Args: - messages (List[Message]): A list of Message objects to set as the new chat messages. - """ def chat_messages(self, value): - """ - Sets the context for chat messages with the provided value. - - Args: - value: The value to be set for the chat messages context. The type of `value` is inferred from the usage context. - """ self.set_context(CHAT_MESSAGES, value) + @property def chat_kwargs(self) -> Dict[str, str]: """ diff --git a/memory_scope/storage/llama_index_es_memory_store.py b/memory_scope/storage/llama_index_es_memory_store.py index 8af77e11..f36f22ec 100644 --- a/memory_scope/storage/llama_index_es_memory_store.py +++ b/memory_scope/storage/llama_index_es_memory_store.py @@ -168,14 +168,14 @@ ray.init(ignore_reinit_error=True) @ray.remote class _LlamaIndexEsMemoryStore(BaseMemoryStore): def __init__(self, - embedding_model_conf: dict, + embedding_model: dict, index_name: str, es_url: str, use_hybrid: bool = True, **kwargs): from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel - embedding_model = LlamaIndexEmbeddingModel(**embedding_model_conf) + embedding_model = LlamaIndexEmbeddingModel(**embedding_model) self.embedding_model: BaseModel = embedding_model self.es_store = _ElasticsearchStore(index_name=index_name, es_url=es_url, @@ -312,13 +312,13 @@ class _LlamaIndexEsMemoryStore(BaseMemoryStore): class LlamaIndexEsMemoryStore(): def __init__(self, - embedding_model_conf: BaseModel, + embedding_model: BaseModel, index_name: str, es_url: str, use_hybrid: bool = True, **kwargs): if 'embedding_model' in kwargs: kwargs.pop('embedding_model') - self.proxy_obj = _LlamaIndexEsMemoryStore.remote(embedding_model_conf, index_name, es_url, use_hybrid, **kwargs) + self.proxy_obj = _LlamaIndexEsMemoryStore.remote(embedding_model.kwargs, index_name, es_url, use_hybrid, **kwargs) def retrieve_memories(self, query: str,