[dev] change submit_async_task to submit_thread_task

This commit is contained in:
jinli.yl 2024-07-09 16:32:29 +08:00
parent a8fcfeeaa2
commit 88cc1d8c32
8 changed files with 87 additions and 85 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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