mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
override es-store mt with ray
This commit is contained in:
parent
08c90983ed
commit
2d6713376d
4 changed files with 55 additions and 34 deletions
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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]):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue