minor fix

This commit is contained in:
xianzhe.xxz 2024-07-24 17:25:08 +08:00
parent 7d11369f5a
commit 5c28465fc3
3 changed files with 29 additions and 13 deletions

View file

@ -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]

View file

@ -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)

View file

@ -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
)