mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
minor fix
This commit is contained in:
parent
7d11369f5a
commit
5c28465fc3
3 changed files with 29 additions and 13 deletions
|
|
@ -39,7 +39,6 @@ class LlamaIndexEmbeddingModel(BaseModel):
|
|||
model_response.details = "empty embeddings"
|
||||
model_response.status = False
|
||||
return model_response
|
||||
|
||||
if len(embeddings) == 1:
|
||||
# return list[float]
|
||||
embeddings = embeddings[0]
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
use_hybrid: bool = True,
|
||||
emb_dims: int = 1536,
|
||||
**kwargs):
|
||||
|
||||
self.index_name = index_name
|
||||
self.emb_dims = emb_dims
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.es_store = SyncElasticsearchStore(index_name=index_name,
|
||||
|
|
@ -32,21 +32,26 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
# adding the following lines of code, the display of deprecated information is suppressed.
|
||||
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
embed_model=self.embedding_model.model)
|
||||
embed_model=self.embedding_model.model)
|
||||
|
||||
self.index.build_index_from_nodes([TextNode(text="text")])
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve_memories(self,
|
||||
query: Optional[str] = None,
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
# if index is not created, return []
|
||||
exists = self.es_store._store.client.indices.exists(index=self.index_name)
|
||||
if not exists:
|
||||
return []
|
||||
|
||||
if filter_dict is None:
|
||||
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,
|
||||
sparse_top_k=top_k)
|
||||
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 query is None:
|
||||
query = QueryBundle(query_str='**--**',
|
||||
embedding=self.dummy_query_vector())
|
||||
|
|
@ -124,11 +129,14 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
|
||||
Returns:
|
||||
TextNode: The converted TextNode with content and metadata from the MemoryNode.
|
||||
"""
|
||||
"""
|
||||
embedding = memory_node.vector
|
||||
if not embedding:
|
||||
embedding = None
|
||||
return TextNode(id_=memory_node.memory_id,
|
||||
text=memory_node.content,
|
||||
embedding=memory_node.vector,
|
||||
metadata=memory_node.model_dump(exclude={"content"}))
|
||||
embedding=embedding,
|
||||
metadata=memory_node.model_dump(exclude={"content", "vector"}))
|
||||
|
||||
@staticmethod
|
||||
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
|
||||
|
|
@ -141,4 +149,8 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
Returns:
|
||||
MemoryNode: The converted MemoryNode with text and metadata from the NodeWithScore.
|
||||
"""
|
||||
return MemoryNode(content=text_node.text, vector=text_node.embedding, **text_node.metadata)
|
||||
embedding = text_node.embedding
|
||||
print("textnode embedding", embedding)
|
||||
if not embedding:
|
||||
embedding = []
|
||||
return MemoryNode(content=text_node.text, vector=embedding, **text_node.metadata)
|
||||
|
|
|
|||
|
|
@ -169,6 +169,7 @@ class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
|
|||
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(
|
||||
|
|
@ -590,6 +591,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
Callable[[Dict, Union[VectorStoreQuery, None]], Dict]
|
||||
] = None,
|
||||
es_filter: Optional[List[Dict]] = None,
|
||||
fields: List[str] = [],
|
||||
**kwargs: Any,
|
||||
) -> VectorStoreQueryResult:
|
||||
"""
|
||||
|
|
@ -629,15 +631,15 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
num_candidates=num_candidates, # query.similarity_top_k * 10,
|
||||
filter=filter,
|
||||
custom_query=custom_query,
|
||||
fields=fields,
|
||||
)
|
||||
|
||||
top_k_nodes = []
|
||||
top_k_ids = []
|
||||
top_k_scores = []
|
||||
print("hits:", len(hits))
|
||||
for hit in hits:
|
||||
source = hit["_source"]
|
||||
metadata = source.get("metadata", None)
|
||||
embedding = source.get("embedding", None)
|
||||
text = source.get(self.text_field, None)
|
||||
node_id = hit["_id"]
|
||||
|
||||
|
|
@ -645,6 +647,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
# Attempt to parse metadata using the standard method
|
||||
node = metadata_dict_to_node(metadata)
|
||||
node.text = text
|
||||
node.embedding = embedding
|
||||
except Exception:
|
||||
# Legacy support for old metadata format
|
||||
logger.warning(
|
||||
|
|
@ -662,6 +665,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
text=text,
|
||||
metadata=metadata,
|
||||
id_=node_id,
|
||||
embedding=embedding,
|
||||
start_char_idx=start_char_idx,
|
||||
end_char_idx=end_char_idx,
|
||||
relationships=relationships,
|
||||
|
|
@ -682,5 +686,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
return VectorStoreQueryResult(
|
||||
nodes=top_k_nodes,
|
||||
ids=top_k_ids,
|
||||
similarities=_to_llama_similarities(top_k_scores),
|
||||
# similarities=_to_llama_similarities(top_k_scores),
|
||||
similarities=top_k_scores
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue