From c437553ff69fc7c5a5d9e53fd477ecfa306126b0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 24 Jul 2024 19:31:44 +0800 Subject: [PATCH] format code --- .../worker/frontend/read_message_worker.py | 1 + .../storage/llama_index_es_memory_store.py | 26 +++++++++---------- .../storage/llama_index_sync_elasticsearch.py | 14 +++++----- memoryscope/utils/logger.py | 1 + memoryscope/utils/response_text_parser.py | 2 +- 5 files changed, 22 insertions(+), 22 deletions(-) diff --git a/memoryscope/memory/worker/frontend/read_message_worker.py b/memoryscope/memory/worker/frontend/read_message_worker.py index 798f1697..dfbe4a12 100644 --- a/memoryscope/memory/worker/frontend/read_message_worker.py +++ b/memoryscope/memory/worker/frontend/read_message_worker.py @@ -7,6 +7,7 @@ class ReadMessageWorker(MemoryBaseWorker): """ Fetches unmemorized chat messages. """ + def _run(self): """ Executes the primary function to fetch unmemorized chat messages. diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/storage/llama_index_es_memory_store.py index b8e87fd3..7d9dbd1e 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/storage/llama_index_es_memory_store.py @@ -1,6 +1,5 @@ -import warnings import random -from typing import Dict, List, Any, Optional, cast +from typing import Dict, List, Optional from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle @@ -8,7 +7,8 @@ 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, _to_elasticsearch_filter +from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, _AsyncDenseVectorStrategy, \ + _to_elasticsearch_filter from memoryscope.utils.logger import Logger @@ -30,7 +30,7 @@ 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. - + self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, embed_model=self.embedding_model.model) @@ -44,18 +44,18 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): 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, "fields": ['embedding']}, + 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()) - + text_nodes = retriever.retrieve(query) return [self._text_node_2_memory_node(n) for n in text_nodes] @@ -71,11 +71,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): retriever = self.index.as_retriever( vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k) - + if query is None: query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) - + text_nodes: List[NodeWithScore] = await retriever.aretrieve(query) return [self._text_node_2_memory_node(n) for n in text_nodes] @@ -114,11 +114,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): Closes the Elasticsearch store, releasing any resources associated with it. """ self.es_store.close() - - def dummy_query_vector(self): + + def dummy_query_vector(self): random_floats = [random.uniform(0, 1) for _ in range(self.emb_dims)] return random_floats - + @staticmethod def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode: """ @@ -129,7 +129,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): Returns: TextNode: The converted TextNode with content and metadata from the MemoryNode. - """ + """ embedding = memory_node.vector if not embedding: embedding = None diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/storage/llama_index_sync_elasticsearch.py index 1532d9f5..2c525920 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/storage/llama_index_sync_elasticsearch.py @@ -18,7 +18,6 @@ from llama_index.core.bridge.pydantic import PrivateAttr from llama_index.core.schema import BaseNode, MetadataMode, TextNode from llama_index.core.vector_stores.types import ( BasePydanticVectorStore, - MetadataFilters, VectorStoreQuery, VectorStoreQueryMode, VectorStoreQueryResult, @@ -141,10 +140,10 @@ class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy): if query == "**--**": query_body = { "query": { - "bool": { - "filter": filter, - } - }, + "bool": { + "filter": filter, + } + }, } else: query_body = { @@ -262,7 +261,6 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str return result - class SyncElasticsearchStore(BasePydanticVectorStore): """ Elasticsearch vector store. @@ -676,9 +674,9 @@ class SyncElasticsearchStore(BasePydanticVectorStore): isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy) and self.retrieval_strategy.hybrid ): - total_rank = sum(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] # top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores] return VectorStoreQueryResult( diff --git a/memoryscope/utils/logger.py b/memoryscope/utils/logger.py index a3d5ea89..9558b3b0 100644 --- a/memoryscope/utils/logger.py +++ b/memoryscope/utils/logger.py @@ -12,6 +12,7 @@ class Logger(logging.Logger): """ The `Logger` class handle the stream of information or errors in activities. """ + def __init__(self, name: str, level: int = logging.INFO, diff --git a/memoryscope/utils/response_text_parser.py b/memoryscope/utils/response_text_parser.py index aa064528..4081cbbf 100644 --- a/memoryscope/utils/response_text_parser.py +++ b/memoryscope/utils/response_text_parser.py @@ -1,5 +1,5 @@ -from typing import List import re +from typing import List from memoryscope.constants.language_constants import NONE_WORD from memoryscope.utils.global_context import G_CONTEXT