diff --git a/.gitignore b/.gitignore index fca0ec18..41cce170 100644 --- a/.gitignore +++ b/.gitignore @@ -143,3 +143,4 @@ docs/sphinx_doc/build/ *runs/ agentscope.db tmp*.json +cradle* \ No newline at end of file diff --git a/memory_scope/storage/llama_index_es_memory_store.py b/memory_scope/storage/llama_index_es_memory_store.py index a174da2a..a3e6fc0d 100644 --- a/memory_scope/storage/llama_index_es_memory_store.py +++ b/memory_scope/storage/llama_index_es_memory_store.py @@ -141,12 +141,14 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str class LlamaIndexEsMemoryStore(BaseMemoryStore): def __init__(self, - embedding_model: BaseModel, + embedding_model_conf: dict, index_name: str, es_url: str, use_hybrid: bool = True, **kwargs): + from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel + embedding_model = LlamaIndexEmbeddingModel(**embedding_model_conf) self.embedding_model: BaseModel = embedding_model self.es_store = _ElasticsearchStore(index_name=index_name, es_url=es_url, diff --git a/tests/thread_test2.py b/tests/thread_test2.py index e1c5dcf3..9b1a635a 100644 --- a/tests/thread_test2.py +++ b/tests/thread_test2.py @@ -8,58 +8,55 @@ 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): def __init__(self): self.task_list = [] - config = { + embedding_model_conf = { "module_name": "dashscope_embedding", "model_name": "text-embedding-v2", "clazz": "models.llama_index_embedding_model", } - emb = LlamaIndexEmbeddingModel(**config) config = { "index_name": "0708_2", "es_url": "http://localhost:9200", - "embedding_model": emb, + "embedding_model_conf": embedding_model_conf, "use_hybrid": True } - self.es_store = LlamaIndexEsMemoryStore(**config) + self.es_store = LlamaIndexEsMemoryStoreProxy.remote(**config) # 不能在async中初始化 self.logger = logger - async def async_func(self, i: int): - try: - self.logger.info(f"i: {i}") - await asyncio.sleep(i) - # result = self.es_store.retrieve_memories("_", top_k=10, filter_dict={"memory_id": "ggg567"}) - except Exception as e: - self.logger.exception(f"encounter error. e={e.args}") - - def submit_async_task(self, fn, *args, **kwargs): - self.task_list.append((fn, args, kwargs)) - - def gather_async_result(self): - async def async_gather(): - return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.task_list]) - - results = asyncio.run(async_gather()) - self.task_list.clear() - return results + def major_func(self, i: int): + result = ray.get(self.es_store.retrieve_memories.remote("_", top_k=10, filter_dict={"memory_id": "ggg567"})) + return result def run(self): while True: - self.submit_async_task(self.async_func, i=1) - # self.submit_async_task(self.async_func, i=2) - # self.submit_async_task(self.async_func, i=3) + executor_internal = ThreadPoolExecutor(max_workers=5) + f1 = executor_internal.submit(self.major_func, i=1) + f2 = executor_internal.submit(self.major_func, i=2) + executor_internal.shutdown() + f1.result() + f2.result() - self.gather_async_result() executor = ThreadPoolExecutor(max_workers=5) t1 = executor.submit(ThreadTest().run) executor.shutdown() + + +# Shutdown Ray +ray.shutdown() \ No newline at end of file