From 88cc1d8c32d6f9b965e6bf142d65b8286a0f3f6e Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 9 Jul 2024 16:32:29 +0800 Subject: [PATCH] [dev] change submit_async_task to submit_thread_task --- .../memory/operation/base_workflow.py | 1 + memory_scope/memory/worker/base_worker.py | 19 +++++++-- .../summary/long_contra_repeat_worker.py | 12 +++--- .../worker/summary/update_insight_worker.py | 18 ++++---- .../write/get_observation_with_time_worker.py | 26 +++++------- .../worker/write/get_observation_worker.py | 37 ++++++++-------- .../memory/worker/write/load_memory_worker.py | 42 +++++++++---------- memory_scope/utils/tool_functions.py | 17 +++----- 8 files changed, 87 insertions(+), 85 deletions(-) diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index 7144fda2..f7ef319a 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -94,6 +94,7 @@ class BaseWorkflow(object): is_multi_thread=is_backend or self.worker_dict[name], context=self.context, context_lock=self.context_lock, + thread_pool=G_CONTEXT.thread_pool, **kwargs) def _run_sub_workflow(self, worker_list: List[str]) -> bool: diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index aea78d37..5b858ac0 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -1,5 +1,6 @@ import asyncio from abc import ABCMeta, abstractmethod +from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any, Dict from memory_scope.utils.logger import Logger @@ -14,6 +15,7 @@ class BaseWorker(metaclass=ABCMeta): context_lock=None, raise_exception: bool = True, is_multi_thread: bool = False, + thread_pool: ThreadPoolExecutor = None, **kwargs): self.name: str = name @@ -21,25 +23,34 @@ class BaseWorker(metaclass=ABCMeta): self.context_lock = context_lock self.raise_exception: bool = raise_exception self.is_multi_thread: bool = is_multi_thread + self.thread_pool: ThreadPoolExecutor = thread_pool self.kwargs: dict = kwargs self.continue_run: bool = True - self.task_list: list = [] + self.async_task_list: list = [] + self.thread_task_list: list = [] 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") - self.task_list.append((fn, args, kwargs)) + 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.task_list]) + return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.async_task_list]) def gather_async_result(self): results = asyncio.run(self._async_gather()) - self.task_list.clear() + self.async_task_list.clear() return results + def submit_thread_task(self, fn, *args, **kwargs): + self.thread_task_list.append(self.thread_pool.submit(fn, *args, **kwargs)) + + def gather_thread_result(self): + for future in as_completed(self.thread_task_list): + yield future.result() + @abstractmethod def _run(self): raise NotImplementedError diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py index e4b46ad8..ece8d65c 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -13,25 +13,25 @@ from memory_scope.utils.tool_functions import prompt_to_msg class LongContraRepeatWorker(MemoryBaseWorker): FILE_PATH: str = __file__ - async def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]): + def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]): filter_dict = { "user_name": self.user_name, "target_name": self.target_name, "status": MemoryNodeStatus.ACTIVE.value, "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value] } - retrieve_nodes = await self.memory_store.a_retrieve_memories(query=node.content, - top_k=self.long_contra_repeat_top_k, - filter_dict=filter_dict) + retrieve_nodes = self.memory_store.retrieve_memories(query=node.content, + top_k=self.long_contra_repeat_top_k, + filter_dict=filter_dict) return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold] def _run(self): not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES) for node in not_updated_nodes: - self.submit_async_task(fn=self.retrieve_similar_content, node=node) + self.submit_thread_task(fn=self.retrieve_similar_content, node=node) obs_node_dict: Dict[str, MemoryNode] = {} - for origin_node, retrieve_nodes in self.gather_async_result(): + for origin_node, retrieve_nodes in self.gather_thread_result(): if not retrieve_nodes: continue obs_node_dict[origin_node.memory_id] = origin_node diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index 6d4a5052..a3a367d5 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -111,17 +111,17 @@ class UpdateInsightWorker(MemoryBaseWorker): for node in insight_nodes: if node.status == MemoryNodeStatus.ACTIVE.value: - self.submit_async_task(fn=self.filter_obs_nodes, - insight_node=node, - not_updated_nodes=not_updated_nodes) + self.submit_thread_task(fn=self.filter_obs_nodes, + insight_node=node, + not_updated_nodes=not_updated_nodes) else: - self.submit_async_task(fn=self.filter_obs_nodes, - insight_node=node, - not_updated_nodes=not_reflected_nodes) + self.submit_thread_task(fn=self.filter_obs_nodes, + insight_node=node, + not_updated_nodes=not_reflected_nodes) # select top n result_list = [] - for result in self.gather_async_result(): + for result in self.gather_thread_result(): insight_node, filtered_nodes, max_score = result if not filtered_nodes: continue @@ -130,10 +130,10 @@ class UpdateInsightWorker(MemoryBaseWorker): # submit llm update task for insight_node, filtered_nodes, _ in result_sorted: - self.submit_async_task(fn=self.update_insight, insight_node=insight_node, filtered_nodes=filtered_nodes) + self.submit_thread_task(fn=self.update_insight, insight_node=insight_node, filtered_nodes=filtered_nodes) # get result - self.gather_async_result() + self.gather_thread_result() for node in not_updated_nodes: node.obs_updated = True diff --git a/memory_scope/memory/worker/write/get_observation_with_time_worker.py b/memory_scope/memory/worker/write/get_observation_with_time_worker.py index 27e3af94..98dfd0c4 100644 --- a/memory_scope/memory/worker/write/get_observation_with_time_worker.py +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.py @@ -3,7 +3,6 @@ from typing import List from memory_scope.constants.common_constants import NEW_OBS_WITH_TIME_NODES from memory_scope.constants.language_constants import COLON_WORD from memory_scope.memory.worker.write.get_observation_worker import GetObservationWorker -from memory_scope.scheme.memory_node import MemoryNode from memory_scope.scheme.message import Message from memory_scope.utils.datetime_handler import DatetimeHandler from memory_scope.utils.tool_functions import prompt_to_msg @@ -11,21 +10,21 @@ from memory_scope.utils.tool_functions import prompt_to_msg class GetObservationWithTimeWorker(GetObservationWorker): FILE_PATH: str = __file__ + OBS_STORE_KEY: str = NEW_OBS_WITH_TIME_NODES - def build_prompt(self) -> List[Message]: - # build prompt - user_query_list = [] - i = 1 + def filter_messages(self) -> List[Message]: + filter_messages = [] for msg in self.chat_messages: if DatetimeHandler.has_time_word(query=msg.content): - dt_handler = DatetimeHandler(dt=msg.time_created) - dt = dt_handler.string_format(self.prompt_handler.time_string_format) - user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") - i += 1 + filter_messages.append(msg) + return filter_messages - if not user_query_list: - self.logger.warning(f"get obs_with_time user_query_list={user_query_list} is empty") - return [] + def build_message(self, filter_messages: List[Message]) -> List[Message]: + user_query_list = [] + for i, msg in enumerate(filter_messages): + dt_handler = DatetimeHandler(dt=msg.time_created) + dt = dt_handler.string_format(self.prompt_handler.time_string_format) + user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") system_prompt = self.prompt_handler.get_observation_with_time_system.format(num_obs=len(user_query_list), user_name=self.target_name) @@ -37,6 +36,3 @@ class GetObservationWithTimeWorker(GetObservationWorker): obtain_obs_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"obtain_obs_message={obtain_obs_message}") return obtain_obs_message - - def save(self, new_obs_nodes: List[MemoryNode]): - self.set_memories(NEW_OBS_WITH_TIME_NODES, new_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 31f61618..44454448 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -14,6 +14,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class GetObservationWorker(MemoryBaseWorker): FILE_PATH: str = __file__ + OBS_STORE_KEY: str = NEW_OBS_NODES def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str): dt_handler = DatetimeHandler(dt=message.time_created) @@ -40,18 +41,17 @@ class GetObservationWorker(MemoryBaseWorker): obs_reflected=False, obs_updated=False) - def build_prompt(self) -> List[Message]: - # build prompt - user_query_list = [] - i = 1 + def filter_messages(self) -> List[Message]: + filter_messages = [] for msg in self.chat_messages: if not DatetimeHandler.has_time_word(query=msg.content): - user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") - i += 1 + filter_messages.append(msg) + return filter_messages - if not user_query_list: - self.logger.warning(f"get obs user_query_list={user_query_list} is empty") - return [] + def build_message(self, filter_messages: List[Message]) -> List[Message]: + user_query_list = [] + for i, msg in enumerate(filter_messages): + user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list), user_name=self.target_name) @@ -63,15 +63,14 @@ class GetObservationWorker(MemoryBaseWorker): self.logger.info(f"obtain_obs_message={obtain_obs_message}") return obtain_obs_message - def save(self, new_obs_nodes: List[MemoryNode]): - self.set_memories(NEW_OBS_NODES, new_obs_nodes) - def _run(self): - obtain_obs_message = self.build_prompt() - if not obtain_obs_message: - self.logger.warning("get obs message is empty!") + filter_messages = self.filter_messages() + if not filter_messages: + self.logger.warning("get obs filter_messages is empty!") return + obtain_obs_message = self.build_message(filter_messages) + # call LLM response = self.generation_model.call(messages=obtain_obs_message, top_k=self.generation_model_top_k) @@ -111,13 +110,13 @@ class GetObservationWorker(MemoryBaseWorker): # index number needs to be corrected to -1 idx = int(idx) - 1 - if idx >= len(self.messages): - self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}") + if idx >= len(filter_messages): + self.logger.warning(f"idx={idx} is invalid! filter_messages.size={len(filter_messages)}") continue - new_obs_nodes.append(self.add_observation(message=self.messages[idx], + new_obs_nodes.append(self.add_observation(message=filter_messages[idx], time_infer=time_infer, obs_content=obs_content, keywords=keywords)) - self.save(new_obs_nodes) + self.set_memories(self.OBS_STORE_KEY, new_obs_nodes) diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index 6717f7a5..cf209118 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 - async def retrieve_not_reflected_memory(self, query: str): + 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] = await self.memory_store.a_retrieve_memories(query=query, - top_k=self.retrieve_not_reflected_top_k, - filter_dict=filter_dict) + nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, + top_k=self.retrieve_not_reflected_top_k, + filter_dict=filter_dict) self.set_memories(NOT_REFLECTED_NODES, nodes) @timer - async def retrieve_not_updated_memory(self, query: str): + 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] = await self.memory_store.a_retrieve_memories(query=query, - top_k=self.retrieve_not_updated_top_k, - filter_dict=filter_dict) + nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, + top_k=self.retrieve_not_updated_top_k, + filter_dict=filter_dict) self.set_memories(NOT_UPDATED_NODES, nodes) @timer - async def retrieve_insight_memory(self, query: str): + 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] = await self.memory_store.a_retrieve_memories(query=query, - top_k=self.retrieve_insight_top_k, - filter_dict=filter_dict) + nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, + top_k=self.retrieve_insight_top_k, + filter_dict=filter_dict) self.set_memories(INSIGHT_NODES, nodes) @timer - async def retrieve_today_memory(self): + def retrieve_today_memory(self): if not self.today_obs_top_k: return @@ -80,16 +80,16 @@ class LoadMemoryWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "dt": dt_handler.datetime_format(), } - nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content, - top_k=self.today_obs_top_k, - filter_dict=filter_dict) + nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content, + top_k=self.today_obs_top_k, + filter_dict=filter_dict) self.set_memories(TODAY_NODES, nodes) def _run(self): mock_query = "-" - 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() + 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() diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index b06e3f15..98f27f0e 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -4,11 +4,13 @@ import re import time from copy import deepcopy from importlib import import_module +from typing import List import pyfiglet from termcolor import colored from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.scheme.message import Message ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey", "light_red", "light_green", "light_yellow", "light_blue", "light_magenta", "light_cyan", "white"] @@ -76,18 +78,11 @@ def init_instance_by_config(config: dict, return getattr(module, cls_name)(**config_copy) -def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str): +def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str) -> List[Message]: return [ - { - "role": MessageRoleEnum.SYSTEM.value, - "content": system_prompt.strip(), - }, - { - "role": MessageRoleEnum.USER.value, - "content": "\n".join( - [x.strip() for x in [few_shot, system_prompt, user_query]] - ), - }, + Message(role=MessageRoleEnum.SYSTEM.value, content=system_prompt.strip()), + Message(role=MessageRoleEnum.USER.value, + content="\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])) ]