diff --git a/memory_scope/cli.py b/memory_scope/cli.py index d40221e0..5ccdb21a 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -68,8 +68,9 @@ class MemoryScope(object): if "memory_store" not in self.config: raise RuntimeError("memory_store config is required!") memory_store_config = self.config["memory_store"] - embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]] - G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model) + # embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]] + embedding_model_conf = self.config["models"][memory_store_config[ModelEnum.EMBEDDING_MODEL.value]] + G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model_conf=embedding_model_conf) # init monitor G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) diff --git a/memory_scope/storage/base_memory_store.py b/memory_scope/storage/base_memory_store.py index bc8aff37..9b041ba9 100644 --- a/memory_scope/storage/base_memory_store.py +++ b/memory_scope/storage/base_memory_store.py @@ -10,10 +10,6 @@ class BaseMemoryStore(metaclass=ABCMeta): def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: pass - @abstractmethod - async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: - pass - @abstractmethod def update_memories(self, nodes: MemoryNode | List[MemoryNode]): """ diff --git a/memory_scope/storage/llama_index_es_memory_store.py b/memory_scope/storage/llama_index_es_memory_store.py index a3e6fc0d..01f33121 100644 --- a/memory_scope/storage/llama_index_es_memory_store.py +++ b/memory_scope/storage/llama_index_es_memory_store.py @@ -138,8 +138,10 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str result['bool'].update({"must": operand}) return result - -class LlamaIndexEsMemoryStore(BaseMemoryStore): +import ray +ray.init(ignore_reinit_error=True) +@ray.remote +class _LlamaIndexEsMemoryStore(BaseMemoryStore): def __init__(self, embedding_model_conf: dict, index_name: str, @@ -172,21 +174,6 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): text_nodes = retriever.retrieve(query) return [self._text_node_2_memory_node(n) for n in text_nodes] - async def a_retrieve_memories(self, - query: str, - top_k: int, - filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: - self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}") - - 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) - text_nodes: List[NodeWithScore] = await retriever.aretrieve(query) - return [self._text_node_2_memory_node(n) for n in text_nodes] - def insert(self, node: MemoryNode): self.index.insert_nodes([self._memory_node_2_text_node(node)]) @@ -256,3 +243,49 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): @staticmethod def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode: return MemoryNode(content=text_node.text, **text_node.metadata) + + + +class LlamaIndexEsMemoryStore(): + def __init__(self, + embedding_model_conf: BaseModel, + index_name: str, + es_url: str, + use_hybrid: bool = True, + **kwargs): + if 'embedding_model' in kwargs: kwargs.pop('embedding_model') + self.proxy_obj = _LlamaIndexEsMemoryStore.remote(embedding_model_conf, index_name, es_url, use_hybrid, **kwargs) + + def retrieve_memories(self, + query: str, + top_k: int, + filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: + return ray.get(self.proxy_obj.retrieve_memories.remote(query, top_k, filter_dict)) + + def insert(self, node: MemoryNode): + return ray.get(self.proxy_obj.insert.remote(node)) + + def delete(self, node: MemoryNode): + return ray.get(self.proxy_obj.delete.remote(node)) + + def update(self, node: MemoryNode): + return ray.get(self.proxy_obj.update.remote(node)) + + def update_batch(self, nodes: List[MemoryNode]): + return ray.get(self.proxy_obj.update_batch.remote(nodes)) + + def close(self): + return ray.get(self.proxy_obj.close.remote()) + + def update_memories(self, nodes: MemoryNode | List[MemoryNode]): + return ray.get(self.proxy_obj.update_memories.remote(nodes)) + + @staticmethod + def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode: + return TextNode(id_=memory_node.memory_id, + text=memory_node.content, + metadata=memory_node.model_dump(exclude={"content"})) + + @staticmethod + def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode: + return MemoryNode(content=text_node.text, **text_node.metadata) diff --git a/tests/thread_test2.py b/tests/thread_test2.py index 9b1a635a..063015a8 100644 --- a/tests/thread_test2.py +++ b/tests/thread_test2.py @@ -8,15 +8,9 @@ from concurrent.futures import ThreadPoolExecutor from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memory_scope.utils.logger import Logger -import ray -# Initialize Ray -ray.init() - logger = Logger.get_logger("default") # Define the class as a Ray actor -@ray.remote -class LlamaIndexEsMemoryStoreProxy(LlamaIndexEsMemoryStore): ... class ThreadTest(object): @@ -35,11 +29,11 @@ class ThreadTest(object): "use_hybrid": True } - self.es_store = LlamaIndexEsMemoryStoreProxy.remote(**config) # 不能在async中初始化 + self.es_store = LlamaIndexEsMemoryStore(**config) # 不能在async中初始化 self.logger = logger def major_func(self, i: int): - result = ray.get(self.es_store.retrieve_memories.remote("_", top_k=10, filter_dict={"memory_id": "ggg567"})) + result = self.es_store.retrieve_memories("_", top_k=10, filter_dict={"memory_id": "ggg567"}) return result def run(self): @@ -57,6 +51,3 @@ executor = ThreadPoolExecutor(max_workers=5) t1 = executor.submit(ThreadTest().run) executor.shutdown() - -# Shutdown Ray -ray.shutdown() \ No newline at end of file