mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
[dev] change submit_async_task to submit_thread_task
This commit is contained in:
parent
a8fcfeeaa2
commit
88cc1d8c32
8 changed files with 87 additions and 85 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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]]))
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue