From f7ce6e5e24080fc428601a4808769f9b345a4a5d Mon Sep 17 00:00:00 2001 From: "xianzhe.xxz" Date: Fri, 26 Jul 2024 19:58:30 +0800 Subject: [PATCH 1/2] adjust the params of EsStore, set dense retrieve as default, --- .../storage/llama_index_es_memory_store.py | 25 ++++++---- .../storage/llama_index_sync_elasticsearch.py | 21 ++++++--- tests/storages/test_storages_lli_synces.py | 46 ++++++------------- 3 files changed, 45 insertions(+), 47 deletions(-) diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index 4a9d7f4f..af88f91c 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -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]): diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 507044a3..5c74b0ba 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -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, diff --git a/tests/storages/test_storages_lli_synces.py b/tests/storages/test_storages_lli_synces.py index 8186179a..0bb8b554 100644 --- a/tests/storages/test_storages_lli_synces.py +++ b/tests/storages/test_storages_lli_synces.py @@ -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() From 95d0047e5cd787706841774e0836150576292ab7 Mon Sep 17 00:00:00 2001 From: "xianzhe.xxz" Date: Fri, 26 Jul 2024 20:07:01 +0800 Subject: [PATCH 2/2] add score_recall manager --- memoryscope/storage/llama_index_es_memory_store.py | 3 ++- memoryscope/storage/llama_index_sync_elasticsearch.py | 1 - 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index af88f91c..96b1542f 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -144,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: @@ -158,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) diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 5c74b0ba..0794b284 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -702,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]