fix threading bug

This commit is contained in:
青轩 2024-07-10 14:30:04 +08:00
parent 3b9afc3e81
commit adc7b75d6c
3 changed files with 27 additions and 27 deletions

1
.gitignore vendored
View file

@ -143,3 +143,4 @@ docs/sphinx_doc/build/
*runs/
agentscope.db
tmp*.json
cradle*

View file

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

View file

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