add score_recall manager and adjust params of EsStore

This commit is contained in:
jinli.yl 2024-07-27 00:43:26 +08:00
commit 160cd326d4
3 changed files with 47 additions and 49 deletions

View file

@ -7,7 +7,7 @@ from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle
from memoryscope.models.base_model import BaseModel
from memoryscope.scheme.memory_node import MemoryNode
from memoryscope.storage.base_memory_store import BaseMemoryStore
from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, _AsyncDenseVectorStrategy, \
from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, ESCombinedRetrieveStrategy, \
_to_elasticsearch_filter
from memoryscope.utils.logger import Logger
@ -18,18 +18,16 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
embedding_model: BaseModel,
index_name: str,
es_url: str,
use_hybrid: bool = True,
emb_dims: int = 1536,
retrieve_mode: str = "dense",
hybrid_alpha: float = None,
**kwargs):
self.emb_dims = None
self.index_name = index_name
self.emb_dims = emb_dims
self.embedding_model: BaseModel = embedding_model
self.es_store = SyncElasticsearchStore(index_name=index_name,
es_url=es_url,
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid,
alpha=0.5), # weights of vector similarity,
# while the weights of BM25 is 1-alpha.
# when alpha=None, then rrf fusion is uesd.
retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode,
hybrid_alpha=hybrid_alpha),
**kwargs)
# TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By
@ -40,7 +38,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
self.logger = Logger.get_logger()
def retrieve_memories(self,
query: str = "",
query: str = "**--**",
top_k: int = 3,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
# if index is not created, return []
@ -55,10 +53,13 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter, "fields": ['embedding']},
similarity_top_k=top_k,
sparse_top_k=top_k)
if not query:
if not query and self.emb_dims:
query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector())
text_nodes = retriever.retrieve(query)
if text_nodes and text_nodes[0].embedding:
self.emb_dims = len(text_nodes[0].embedding)
return [self._text_node_2_memory_node(n) for n in text_nodes]
async def a_retrieve_memories(self,
@ -82,6 +83,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector())
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
if text_nodes and text_nodes[0].embedding:
self.emb_dims = len(text_nodes[0].embedding)
return [self._text_node_2_memory_node(n) for n in text_nodes]
def batch_insert(self, nodes: List[MemoryNode]):
@ -139,7 +144,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
return TextNode(id_=memory_node.memory_id,
text=memory_node.content,
embedding=embedding,
metadata=memory_node.model_dump(exclude={"content", "vector"}))
metadata=memory_node.model_dump(exclude={"content", "vector", "score_recall", "score_rank", "score_rerank"}))
@staticmethod
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
@ -153,4 +158,5 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
MemoryNode: The converted MemoryNode with text and metadata from the NodeWithScore.
"""
text_node.metadata["vector"] = text_node.embedding if text_node.embedding else []
text_node.metadata["score_recall"] = text_node.score
return MemoryNode(content=text_node.text, **text_node.metadata)

View file

@ -131,7 +131,7 @@ def _mode_must_match_retrieval_strategy(
raise ValueError(f"to enable hybrid mode, it must be set in retrieval strategy")
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy):
def __init__(
@ -139,13 +139,21 @@ class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
*,
distance: DistanceMetric = DistanceMetric.COSINE,
model_id: Optional[str] = None,
hybrid: bool = False,
retrieve_mode: str = "dense",
rrf: Union[bool, Dict[str, Any]] = True,
text_field: Optional[str] = "text_field",
alpha: Optional[float] = None,
):
super().__init__(distance=distance, model_id=model_id, hybrid=hybrid, rrf=rrf, text_field=text_field)
self.alpha = alpha
hybrid_alpha: Optional[float] = None,
):
if retrieve_mode == "dense":
self.alpha = 1.0
elif retrieve_mode == "sparse":
# self.alpha = 0.0
raise NotImplementedError
elif retrieve_mode == "hybrid":
# self.alpha = hybrid_alpha
raise NotImplementedError
super().__init__(distance=distance, model_id=model_id, hybrid=True, rrf=rrf, text_field=text_field)
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.
@ -637,6 +645,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
else:
filter = es_filter or []
num_candidates = query.similarity_top_k * 10 if query.similarity_top_k <= 1000 else query.similarity_top_k
hits = self._store.search(
query=query.query_str,
query_vector=query.query_embedding,
@ -693,7 +702,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
):
total_rank = sum(top_k_scores)
top_k_scores = [rank for rank in top_k_scores]
print("top_k_scores:", top_k_scores)
# top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores]
# top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores]

View file

@ -20,7 +20,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
"index_name": "0708_8",
"es_url": "http://localhost:9200",
"embedding_model": emb,
"use_hybrid": True
"retrieve_mode": "dense",
}
self.es_store = LlamaIndexEsMemoryStore(**config)
@ -152,13 +152,6 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
),
]
def test_retrieve(self):
filter_dict = {
"timestamp": 12,
# "memory_id": "bbb456",
# "score_rank": 0,
}
for node in self.data:
self.es_store.insert(node)
@ -171,36 +164,27 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
meta_data={"5": "5"},
timestamp=13
))
def test_retrieve(self):
filter_dict = {
"timestamp": 12,
# "memory_id": "bbb456",
# "score_rank": 0,
}
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",
memory_type="profile",
user_id="6",
status="invalid",
memory_id="ggg567",
timestamp=13,
))
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
def test_retrieve_wo_query(self,):
filter_dict = {
"memory_id": "bbb456",
}
res = self.es_store.retrieve_memories(filter_dict=filter_dict, top_k=15)
print(len(res))
print(res)
self.es_store.delete(MemoryNode(
content="test update",
memory_type="profile",
user_id="6",
status="invalid",
memory_id="ggg567",
timestamp=13,
))
import asyncio
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)
def tearDown(self):
self.es_store.close()