override es-store mt with ray

This commit is contained in:
青轩 2024-07-10 16:48:03 +08:00
parent 08c90983ed
commit 2d6713376d
4 changed files with 55 additions and 34 deletions

View file

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

View file

@ -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]):
"""

View file

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

View file

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