mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
add score_recall manager and adjust params of EsStore
This commit is contained in:
commit
160cd326d4
3 changed files with 47 additions and 49 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue