From b6a3d23ffe1b5c8eb1689bc0ae838a9936d9f455 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 28 Jun 2024 15:30:00 +0800 Subject: [PATCH] [dev] change dummy vector store to llama es store --- config/config.yaml | 7 ++++-- memory_scope/cli.py | 8 ++++++- memory_scope/storage/base_monitor.py | 8 +++---- memory_scope/storage/base_vector_store.py | 12 ---------- memory_scope/storage/dummy_monitor.py | 2 +- memory_scope/storage/dummy_vector_store.py | 24 ------------------- .../llama_index_elastic_search_store.py | 14 ++++++----- tests/storages/test_storages_lli_es.py | 23 +++++++++++++----- 8 files changed, 42 insertions(+), 56 deletions(-) delete mode 100644 memory_scope/storage/dummy_vector_store.py diff --git a/config/config.yaml b/config/config.yaml index 053d0132..33ed7954 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -8,7 +8,8 @@ memory_chat: 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 @@ -52,8 +53,10 @@ models: module_name: dashscope_rank model_name: gte-rerank vector_store: - class: storage.dummy_vector_store + class: storage.llama_index_elastic_search_store embedding_model: dashscope_embedding + index_name: memory_index + es_url: http://localhost:9200 monitor: class: storage.dummy_monitor worker: diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 427ad145..e76575aa 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -14,6 +14,7 @@ from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.utils.logger import Logger from memory_scope.utils.tool_functions import init_instance_by_config from memory_scope.utils.timer import timer +from memory_scope.enumeration.model_enum import ModelEnum class CliJob(object): @@ -54,7 +55,9 @@ class CliJob(object): G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) # init vector_store - G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"]) + vector_store_config = self.config["vector_store"] + embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]] + G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model) # init monitor G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) @@ -70,6 +73,9 @@ class CliJob(object): memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() + G_CONTEXT.vector_store.close() + G_CONTEXT.monitor.close() + if __name__ == "__main__": cli_job = CliJob() diff --git a/memory_scope/storage/base_monitor.py b/memory_scope/storage/base_monitor.py index 1d84621e..05465cd3 100644 --- a/memory_scope/storage/base_monitor.py +++ b/memory_scope/storage/base_monitor.py @@ -18,8 +18,8 @@ class BaseMonitor(metaclass=ABCMeta): :return: """ - @abstractmethod def flush(self): - """ - :return: - """ + pass + + def close(self): + pass diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index 604866fe..9c5589b8 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -1,20 +1,11 @@ from abc import ABCMeta, abstractmethod from typing import Dict, List -from memory_scope.models.base_model import BaseModel from memory_scope.scheme.memory_node import MemoryNode class BaseVectorStore(metaclass=ABCMeta): - def __init__(self, - index_name: str = "", - embedding_model: BaseModel | None = None, - **kwargs): - self.index_name: str = index_name - self.embedding_model: BaseModel = embedding_model - self.kwargs: dict = kwargs - @abstractmethod def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): pass @@ -25,9 +16,6 @@ class BaseVectorStore(metaclass=ABCMeta): @abstractmethod def insert(self, node: MemoryNode): - """ TODO 是否overwrite - :return: - """ pass def insert_batch(self, nodes: List[MemoryNode]): diff --git a/memory_scope/storage/dummy_monitor.py b/memory_scope/storage/dummy_monitor.py index f39a917b..6c754ae5 100644 --- a/memory_scope/storage/dummy_monitor.py +++ b/memory_scope/storage/dummy_monitor.py @@ -8,5 +8,5 @@ class DummyMonitor(BaseMonitor): def add_token(self): pass - def flush(self): + def close(self): pass diff --git a/memory_scope/storage/dummy_vector_store.py b/memory_scope/storage/dummy_vector_store.py deleted file mode 100644 index f4c7fad8..00000000 --- a/memory_scope/storage/dummy_vector_store.py +++ /dev/null @@ -1,24 +0,0 @@ -from typing import Dict, List - -from memory_scope.scheme.memory_node import MemoryNode -from memory_scope.storage.base_vector_store import BaseVectorStore - - -class DummyVectorStore(BaseVectorStore): - def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): - pass - - async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): - pass - - def insert(self, node: MemoryNode): - pass - - def insert_batch(self): - pass - - def delete(self): - pass - - def flush(self): - pass diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index 85f26e7c..a6f5acb8 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -70,18 +70,20 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str class LlamaIndexElasticSearchStore(BaseVectorStore): def __init__(self, - index_name: str, embedding_model: BaseModel, + index_name: str, + es_url: str, + use_hybrid: bool = True, **kwargs): - super().__init__(index_name=index_name, embedding_model=embedding_model, **kwargs) - self.es_store = _ElasticsearchStore(index_name=self.index_name, - retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True), + + self.embedding_model: BaseModel = embedding_model + self.es_store = _ElasticsearchStore(index_name=index_name, + es_url=es_url, + retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid), **kwargs) self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, embed_model=self.embedding_model.model) - self.memory_node_keys = [x for x in MemoryNode().node_keys if x not in ["meta_data", "content"]] - def retrieve(self, query: str, top_k: int, diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 5608ce7a..7f88bea2 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -40,7 +40,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="1", status="valid", memory_id="bbb456", - + meta_data={"1": "1"} ), MemoryNode( content="An insomniac office worker and a devil-may-care soapmaker form an underground fight " @@ -49,7 +49,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="2", status="valid", memory_id="ccc789", - + meta_data={"2": "2"} ), MemoryNode( content="A thief who steals corporate secrets through the use of dream-sharing technology " @@ -58,6 +58,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="3", status="valid", memory_id="ddd012", + meta_data={"3": "3"} + ), MemoryNode( content="A computer hacker learns from mysterious rebels about the true nature of his reality " @@ -66,6 +68,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="4", status="valid", memory_id="eee345", + meta_data={"4": "4"} + ), MemoryNode( content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven " @@ -73,7 +77,9 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): memory_type="profile", user_id="5", status="valid", - memory_id="fff678" + memory_id="fff678", + meta_data={"5": "5"}, + ), MemoryNode( content="An organized crime dynasty's aging patriarch transfers control of his clandestine " @@ -82,6 +88,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="6", status="valid", memory_id="ggg901", + meta_data={"5": "5"} + ), MemoryNode( content="ggggggggg", @@ -89,6 +97,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): user_id="6", status="valid", memory_id="ggg234", + meta_data={"5": "5"} + ), ] @@ -104,7 +114,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): memory_type="profile", user_id="6", status="valid", - memory_id="ggg567" + memory_id="ggg567", + meta_data={"5": "5"} )) res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) @@ -117,7 +128,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): status="invalid", memory_id="ggg567" )) - res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res) @@ -128,7 +139,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): status="invalid", memory_id="ggg567" )) - res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10) + res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res)