mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
[dev] change dummy vector store to llama es store
This commit is contained in:
parent
5a533b2721
commit
b6a3d23ffe
8 changed files with 42 additions and 56 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -18,8 +18,8 @@ class BaseMonitor(metaclass=ABCMeta):
|
|||
:return:
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def flush(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -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]):
|
||||
|
|
|
|||
|
|
@ -8,5 +8,5 @@ class DummyMonitor(BaseMonitor):
|
|||
def add_token(self):
|
||||
pass
|
||||
|
||||
def flush(self):
|
||||
def close(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue