fix topk limit of es store

This commit is contained in:
xianzhe.xxz 2024-07-08 13:14:38 +08:00
parent 7993136470
commit 2e7053d2b2
2 changed files with 122 additions and 14 deletions

View file

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

View file

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