diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 3b2bc4c0..b43b11f9 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -30,11 +30,11 @@ memory_service: workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory description: "write observation memories of the user" interval_time: 5 - summary_memory: - class: memory.operation.summary_memory - workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory - description: "summary observation memories of the user" - interval_time: 60 +# summary_memory: +# class: memory.operation.summary_memory +# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory +# description: "summary observation memories of the user" +# interval_time: 60 worker: dummy: diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index b8ddddbc..8afef7f4 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -15,8 +15,8 @@ class SummaryMemory(BaseWorkflow, BaseBackendOperation): self.init_workers(is_backend=True, **kwargs) def _run_operation(self, **kwargs): + self.context.clear() self.context[CHAT_KWARGS] = kwargs self.run_workflow() result = self.context.get(RESULT) - self.context.clear() return result diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 120ba918..e42482ac 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -39,6 +39,7 @@ class WriteMemory(BaseWorkflow, BaseBackendOperation): self.init_workers(is_backend=True, **kwargs) def _run_operation(self, **kwargs): + self.context.clear() self.context[CHAT_KWARGS] = kwargs not_memorized_size = self.not_memorized_size if not_memorized_size < self.contextual_msg_count: @@ -51,6 +52,5 @@ class WriteMemory(BaseWorkflow, BaseBackendOperation): self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]] self.run_workflow() result = self.context.get(RESULT) - self.context.clear() self.set_memorized() - return result \ No newline at end of file + return result diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 633f8c87..f3e30d4a 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -32,16 +32,14 @@ class BaseWorker(metaclass=ABCMeta): self.logger: Logger = Logger.get_logger() def submit_async_task(self, fn, *args, **kwargs): - if self.is_multi_thread: - raise RuntimeError(f"async_task is not allowed in multi_thread condition") + # if self.is_multi_thread: + # raise RuntimeError(f"async_task is not allowed in multi_thread condition") self.async_task_list.append((fn, args, kwargs)) async def _async_gather(self): return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.async_task_list]) def gather_async_result(self): - if self.is_multi_thread: - raise RuntimeError(f"async_task is not allowed in multi_thread condition") results = asyncio.run(self._async_gather()) self.async_task_list.clear() return results diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index b00d3db8..efc506f6 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -20,6 +20,7 @@ class ContraRepeatWorker(MemoryBaseWorker): return today_obs_nodes: List[MemoryNode] = self.get_memories(TODAY_NODES) + if today_obs_nodes: all_obs_nodes.extend(today_obs_nodes) all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.timestamp, reverse=True)[:self.contra_repeat_max_count] diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index cf209118..3de5b740 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -13,7 +13,7 @@ from memory_scope.utils.timer import timer class LoadMemoryWorker(MemoryBaseWorker): @timer - def retrieve_not_reflected_memory(self, query: str): + async def retrieve_not_reflected_memory(self, query: str): if not self.retrieve_not_reflected_top_k: return @@ -24,13 +24,13 @@ class LoadMemoryWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "obs_reflected": False, } - nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, + nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query, top_k=self.retrieve_not_reflected_top_k, filter_dict=filter_dict) self.set_memories(NOT_REFLECTED_NODES, nodes) @timer - def retrieve_not_updated_memory(self, query: str): + async def retrieve_not_updated_memory(self, query: str): if not self.retrieve_not_updated_top_k: return @@ -41,13 +41,13 @@ class LoadMemoryWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "obs_updated": False, } - nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, + nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query, top_k=self.retrieve_not_updated_top_k, filter_dict=filter_dict) self.set_memories(NOT_UPDATED_NODES, nodes) @timer - def retrieve_insight_memory(self, query: str): + async def retrieve_insight_memory(self, query: str): if not self.retrieve_insight_top_k: return @@ -57,13 +57,13 @@ class LoadMemoryWorker(MemoryBaseWorker): "status": MemoryNodeStatus.ACTIVE.value, "memory_type": MemoryTypeEnum.INSIGHT.value, } - nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, + nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query, top_k=self.retrieve_insight_top_k, filter_dict=filter_dict) self.set_memories(INSIGHT_NODES, nodes) @timer - def retrieve_today_memory(self): + async def retrieve_today_memory(self): if not self.today_obs_top_k: return @@ -80,7 +80,7 @@ class LoadMemoryWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "dt": dt_handler.datetime_format(), } - nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content, + nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict) @@ -88,8 +88,14 @@ class LoadMemoryWorker(MemoryBaseWorker): def _run(self): mock_query = "-" - self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query) - self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query) - self.submit_thread_task(self.retrieve_insight_memory, query=mock_query) - self.submit_thread_task(self.retrieve_today_memory) - self.gather_thread_result() + # self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query) + # self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query) + # self.submit_thread_task(self.retrieve_insight_memory, query=mock_query) + # self.submit_thread_task(self.retrieve_today_memory) + # self.gather_thread_result() + + self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query) + self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query) + self.submit_async_task(self.retrieve_insight_memory, query=mock_query) + self.submit_async_task(self.retrieve_today_memory) + self.gather_async_result() diff --git a/memory_scope/utils/timer.py b/memory_scope/utils/timer.py index be7df83f..9e035302 100644 --- a/memory_scope/utils/timer.py +++ b/memory_scope/utils/timer.py @@ -5,7 +5,7 @@ from memory_scope.utils.logger import Logger class Timer(object): - def __init__(self, name: str, log_time: bool = True, use_ms: bool = False, **kwargs): + def __init__(self, name: str, log_time: bool = True, use_ms: bool = True, **kwargs): self.name: str = name self.log_time: bool = log_time self.use_ms: bool = use_ms diff --git a/tests/thread_test.py b/tests/thread_test.py new file mode 100644 index 00000000..5999f39a --- /dev/null +++ b/tests/thread_test.py @@ -0,0 +1,76 @@ +import sys + +sys.path.append(".") + +import asyncio +from concurrent.futures import ThreadPoolExecutor, as_completed + +from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore +from memory_scope.utils.logger import Logger + +import questionary + + +class ThreadTest(object): + def __init__(self): + self.task_list = [] + config = { + "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, + "use_hybrid": True + + } + self.es_store = LlamaIndexEsMemoryStore(**config) + self.logger = Logger.get_logger("default") + + async def async_func(self, i: int): + self.logger.info(f"i: {i}") + await asyncio.sleep(i) + self.es_store.insert(MemoryNode( + content="xxxxxx", + memory_type="profile", + user_id="6", + status="valid", + memory_id="ggg567", + meta_data={"5": "5"} + )) + + result = self.es_store.retrieve_memories("_", top_k=10, filter_dict={"memory_id": "ggg567"}) + self.logger.info(f"result: {result}") + + 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 run(self): + 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) + + self.gather_async_result() + self.logger.close() + + +executor = ThreadPoolExecutor(max_workers=5) +t = executor.submit(ThreadTest().run) +questionary.text(message="user:", multiline=False, qmark=">").unsafe_ask() + +print(as_completed([t])) +executor.shutdown()