mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
fix threading bug
This commit is contained in:
parent
3b9afc3e81
commit
adc7b75d6c
3 changed files with 27 additions and 27 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -143,3 +143,4 @@ docs/sphinx_doc/build/
|
|||
*runs/
|
||||
agentscope.db
|
||||
tmp*.json
|
||||
cradle*
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue