mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
support query without embedding
This commit is contained in:
parent
c0ab3b0fb7
commit
2c14b958a9
3 changed files with 39 additions and 33 deletions
|
|
@ -25,11 +25,11 @@ class MemoryNode(BaseModel):
|
|||
|
||||
value: str = Field("", description="memory value")
|
||||
|
||||
score_similar: float = Field(0, description="es similar score")
|
||||
score_recall: float = Field(0, description="embedding similarity score used in recall stage")
|
||||
|
||||
score_rank: float = Field(0, description="rank model score")
|
||||
score_rank: float = Field(0, description="rank model score used in rank stage")
|
||||
|
||||
score_rerank: float = Field(0, description="rerank score")
|
||||
score_rerank: float = Field(0, description="rerank score used in rerank stage")
|
||||
|
||||
memory_type: str = Field("", description="conversation / observation / insight...")
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ from typing import Dict, List, Any, Optional, cast
|
|||
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle
|
||||
from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy
|
||||
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
|
@ -23,6 +22,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
emb_dims: int = 1536,
|
||||
**kwargs):
|
||||
|
||||
self.emb_dims = emb_dims
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.es_store = SyncElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
|
|
@ -30,9 +30,8 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
**kwargs)
|
||||
# TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By
|
||||
# adding the following lines of code, the display of deprecated information is suppressed.
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore")
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
|
||||
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")])
|
||||
|
|
@ -49,7 +48,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k,
|
||||
sparse_top_k=top_k)
|
||||
if query is None:
|
||||
query = QueryBundle(query_str='-',
|
||||
query = QueryBundle(query_str='**--**',
|
||||
embedding=self.dummy_query_vector())
|
||||
|
||||
text_nodes = retriever.retrieve(query)
|
||||
|
|
@ -69,7 +68,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
similarity_top_k=top_k)
|
||||
|
||||
if query is None:
|
||||
query = QueryBundle(query_str='-',
|
||||
query = QueryBundle(query_str='**--**',
|
||||
embedding=self.dummy_query_vector())
|
||||
|
||||
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
|
||||
|
|
|
|||
|
|
@ -40,7 +40,6 @@ DISTANCE_STRATEGIES = Literal[
|
|||
]
|
||||
|
||||
|
||||
|
||||
def get_elasticsearch_client(
|
||||
url: Optional[str] = None,
|
||||
cloud_id: Optional[str] = None,
|
||||
|
|
@ -133,36 +132,43 @@ def _mode_must_match_retrieval_strategy(
|
|||
raise ValueError(f"to enable hybrid mode, it must be set in retrieval strategy")
|
||||
|
||||
|
||||
|
||||
|
||||
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,
|
||||
if query == "**--**":
|
||||
query_body = {
|
||||
"query": {
|
||||
"bool": {
|
||||
"filter": filter,
|
||||
}
|
||||
},
|
||||
}
|
||||
else:
|
||||
query_body = {
|
||||
"knn": knn,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": [
|
||||
{
|
||||
"match": {
|
||||
self.text_field: {
|
||||
"query": query,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"filter": filter,
|
||||
}
|
||||
},
|
||||
}
|
||||
],
|
||||
"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}}
|
||||
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(
|
||||
|
|
@ -668,8 +674,9 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy)
|
||||
and self.retrieval_strategy.hybrid
|
||||
):
|
||||
total_rank = sum(top_k_scores)
|
||||
top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores]
|
||||
total_rank = sum(top_k_scores)
|
||||
top_k_scores = [rank for rank in 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]
|
||||
|
||||
return VectorStoreQueryResult(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue