mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
fix topk limit of es store
This commit is contained in:
parent
7993136470
commit
2e7053d2b2
2 changed files with 122 additions and 14 deletions
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Dict, List, Any
|
||||
from typing import Dict, List, Any, Optional, cast
|
||||
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore
|
||||
|
|
@ -10,6 +10,72 @@ from memory_scope.scheme.memory_node import MemoryNode
|
|||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
|
||||
def _hybrid(
|
||||
self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int,
|
||||
) -> Dict[str, Any]:
|
||||
# Add a query to the knn query.
|
||||
# RRF is used to even the score from the knn query and text query
|
||||
# RRF has two optional parameters: {'rank_constant':int, 'window_size':int}
|
||||
# https://www.elastic.co/guide/en/elasticsearch/reference/current/rrf.html
|
||||
query_body = {
|
||||
"knn": knn,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": [
|
||||
{
|
||||
"match": {
|
||||
self.text_field: {
|
||||
"query": query,
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"filter": filter,
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
if isinstance(self.rrf, Dict):
|
||||
query_body["rank"] = {"rrf": self.rrf}
|
||||
elif isinstance(self.rrf, bool) and self.rrf is True:
|
||||
query_body["rank"] = {"rrf": {"window_size": top_k}}
|
||||
return query_body
|
||||
|
||||
def es_query(
|
||||
self,
|
||||
*,
|
||||
query: Optional[str],
|
||||
query_vector: Optional[List[float]],
|
||||
text_field: str,
|
||||
vector_field: str,
|
||||
k: int,
|
||||
num_candidates: int,
|
||||
filter: List[Dict[str, Any]] = [],
|
||||
) -> Dict[str, Any]:
|
||||
knn = {
|
||||
"filter": filter,
|
||||
"field": vector_field,
|
||||
"k": k,
|
||||
"num_candidates": num_candidates,
|
||||
}
|
||||
|
||||
if query_vector is not None:
|
||||
knn["query_vector"] = query_vector
|
||||
else:
|
||||
# Inference in Elasticsearch. When initializing we make sure to always have
|
||||
# a model_id if don't have an embedding_service.
|
||||
knn["query_vector_builder"] = {
|
||||
"text_embedding": {
|
||||
"model_id": self.model_id,
|
||||
"model_text": query,
|
||||
}
|
||||
}
|
||||
|
||||
if self.hybrid:
|
||||
return self._hybrid(query=cast(str, query), knn=knn, filter=filter, top_k=k)
|
||||
|
||||
return {"knn": knn}
|
||||
|
||||
class _ElasticsearchStore(ElasticsearchStore):
|
||||
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
|
||||
|
|
@ -81,11 +147,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
self.embedding_model: BaseModel = embedding_model
|
||||
self.es_store = _ElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid),
|
||||
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.index.build_index_from_nodes([TextNode(text="text")])
|
||||
# self.index.build_index_from_nodes([TextNode(text="text")])
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve_memories(self,
|
||||
|
|
@ -96,7 +162,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
filter_dict = {}
|
||||
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k)
|
||||
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k, sparse_top_k=top_k)
|
||||
text_nodes = retriever.retrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import unittest
|
|||
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.llama_index_elastic_search_store import LlamaIndexElasticSearchStore
|
||||
from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
|
||||
|
||||
|
||||
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
||||
|
|
@ -12,16 +12,18 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
config = {
|
||||
"module_name": "dashscope_embedding",
|
||||
"model_name": "text-embedding-v2",
|
||||
"clazz": "models.llama_index_embedding_model"
|
||||
"clazz": "models.llama_index_embedding_model",
|
||||
}
|
||||
emb = LlamaIndexEmbeddingModel(**config)
|
||||
|
||||
config = {
|
||||
"index_name": "0626_1",
|
||||
"index_name": "0708_2",
|
||||
"es_url": "http://localhost:9200",
|
||||
"embedding_model": emb,
|
||||
"use_hybrid": True
|
||||
|
||||
}
|
||||
self.es_store = LlamaIndexElasticSearchStore(**config)
|
||||
self.es_store = LlamaIndexEsMemoryStore(**config)
|
||||
self.data = [
|
||||
MemoryNode(
|
||||
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "
|
||||
|
|
@ -100,15 +102,53 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="hhh234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="iii234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="jjj234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="kkk234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
]
|
||||
|
||||
def test_retrieve(self):
|
||||
filter_dict = {
|
||||
"user_id": "6",
|
||||
}
|
||||
# filter_dict = {
|
||||
# "user_id": "6",
|
||||
# }
|
||||
filter_dict = {}
|
||||
|
||||
for node in self.data:
|
||||
self.es_store.insert(node)
|
||||
|
||||
self.es_store.insert(MemoryNode(
|
||||
content="xxxxxx",
|
||||
memory_type="profile",
|
||||
|
|
@ -117,9 +157,10 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
memory_id="ggg567",
|
||||
meta_data={"5": "5"}
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
||||
self.es_store.update(MemoryNode(
|
||||
content="test update",
|
||||
|
|
@ -128,7 +169,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
status="invalid",
|
||||
memory_id="ggg567"
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
|
@ -140,7 +181,8 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
memory_id="ggg567"
|
||||
))
|
||||
import asyncio
|
||||
res = asyncio.run(self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10))
|
||||
res = asyncio.run(self.es_store.a_retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15))
|
||||
#res = self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue