[dev] change dummy vector store to llama es store

This commit is contained in:
jinli.yl 2024-06-28 15:30:00 +08:00
parent 5a533b2721
commit b6a3d23ffe
8 changed files with 42 additions and 56 deletions

View file

@ -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:

View file

@ -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()

View file

@ -18,8 +18,8 @@ class BaseMonitor(metaclass=ABCMeta):
:return:
"""
@abstractmethod
def flush(self):
"""
:return:
"""
pass
def close(self):
pass

View file

@ -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]):

View file

@ -8,5 +8,5 @@ class DummyMonitor(BaseMonitor):
def add_token(self):
pass
def flush(self):
def close(self):
pass

View file

@ -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

View file

@ -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,

View file

@ -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)