mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-10 22:41:06 +00:00
[dev] rename timer print lines
This commit is contained in:
parent
88cc1d8c32
commit
39fc87c378
8 changed files with 33 additions and 23 deletions
|
|
@ -4,6 +4,8 @@ RESULT = "result"
|
|||
|
||||
CHAT_MESSAGES = "chat_messages"
|
||||
|
||||
CONTEXT_MEMORY_DICT = "context_memory_dict"
|
||||
|
||||
CHAT_KWARGS = "chat_kwargs"
|
||||
|
||||
QUERY_WITH_TS = "query_with_ts"
|
||||
|
|
|
|||
|
|
@ -66,7 +66,7 @@ class BaseWorkflow(object):
|
|||
return self.workflow_worker_list
|
||||
|
||||
def _print_workflow(self):
|
||||
self.logger.info(f"----- print_workflow_{self.name}_begin -----")
|
||||
self.logger.info(f"----- workflow.{self.name}.print.begin -----")
|
||||
i: int = 0
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
if len(workflow_part) == 1:
|
||||
|
|
@ -80,7 +80,7 @@ class BaseWorkflow(object):
|
|||
for w in w_zip:
|
||||
if w == "-":
|
||||
continue
|
||||
self.logger.info(f"----- print_workflow_{self.name}_end -----")
|
||||
self.logger.info(f"----- workflow.{self.name}.print.end -----")
|
||||
|
||||
def init_workers(self, is_backend: bool = False, **kwargs):
|
||||
for name in list(self.worker_dict.keys()):
|
||||
|
|
@ -106,7 +106,8 @@ class BaseWorkflow(object):
|
|||
return True
|
||||
|
||||
def run_workflow(self):
|
||||
with Timer(f"run_workflow_{self.name}"):
|
||||
self.logger.info(f"----- workflow.{self.name}.begin -----")
|
||||
with Timer(self.name, log_time=False) as t:
|
||||
self.context[WORKFLOW_NAME] = self.name
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
if len(workflow_part) == 1:
|
||||
|
|
@ -124,3 +125,4 @@ class BaseWorkflow(object):
|
|||
break
|
||||
if not flag:
|
||||
break
|
||||
self.logger.info(f"----- workflow.{self.name}.end cost={t.cost_str}-----")
|
||||
|
|
|
|||
|
|
@ -19,12 +19,14 @@ class ChatMemoryService(BaseMemoryService):
|
|||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
continue
|
||||
|
||||
self._operation_dict[name] = init_instance_by_config(
|
||||
config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_count=self.contextual_msg_count)
|
||||
self.logger.info(f"service={self.__class__.__name__} init operation={name}")
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
if isinstance(messages, Message):
|
||||
|
|
|
|||
|
|
@ -40,6 +40,8 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
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
|
||||
|
|
@ -56,7 +58,7 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
self.logger.info(f"----- worker_{self.name}_begin -----")
|
||||
self.logger.info(f"----- worker.{self.name}.begin -----")
|
||||
with Timer(self.name, log_time=False) as t:
|
||||
if self.raise_exception:
|
||||
self._run()
|
||||
|
|
@ -66,7 +68,7 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
except Exception as e:
|
||||
self.logger.exception(f"run {self.name} failed! args={e.args}")
|
||||
|
||||
self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----")
|
||||
self.logger.info(f"----- worker.{self.name}.end cost={t.cost_str}-----")
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
return self.context.get(key, default)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from abc import ABCMeta
|
||||
from typing import List, Dict, Set, Any
|
||||
|
||||
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
|
||||
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, CONTEXT_MEMORY_DICT
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
|
@ -33,8 +33,6 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._target_name: str | None = None
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
|
||||
self._contex_memory_dict: Dict[str, MemoryNode] = {}
|
||||
|
||||
@property
|
||||
def chat_messages(self) -> List[Message]:
|
||||
return self.get_context(CHAT_MESSAGES)
|
||||
|
|
@ -71,6 +69,12 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._memory_store = G_CONTEXT.memory_store
|
||||
return self._memory_store
|
||||
|
||||
@property
|
||||
def contex_memory_dict(self) -> Dict[str, MemoryNode]:
|
||||
if not self.has_content(CONTEXT_MEMORY_DICT):
|
||||
self.set_context(CONTEXT_MEMORY_DICT, {})
|
||||
return self.get_context(CONTEXT_MEMORY_DICT)
|
||||
|
||||
def get_memories(self, keys: str | List[str]) -> List[MemoryNode]:
|
||||
memories: List[MemoryNode] = []
|
||||
if isinstance(keys, str):
|
||||
|
|
@ -79,7 +83,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
for key in keys:
|
||||
memory_ids: List[str] = self.get_context(key)
|
||||
if memory_ids:
|
||||
memories.extend([self._contex_memory_dict[x] for x in memory_ids])
|
||||
memories.extend([self.contex_memory_dict[x] for x in memory_ids])
|
||||
return memories
|
||||
|
||||
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode]):
|
||||
|
|
@ -87,17 +91,16 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
nodes = []
|
||||
elif isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
for node in nodes:
|
||||
if node.memory_id in self._contex_memory_dict:
|
||||
if node.memory_id in self.contex_memory_dict:
|
||||
continue
|
||||
self._contex_memory_dict[node.memory_id] = node
|
||||
self.contex_memory_dict[node.memory_id] = node
|
||||
self.set_context(key, [n.memory_id for n in nodes])
|
||||
|
||||
def save_memories(self, keys: str | List[str] = None):
|
||||
if keys is None:
|
||||
self.memory_store.update_memories(list(self._contex_memory_dict.values()))
|
||||
self._contex_memory_dict.clear()
|
||||
self.memory_store.update_memories(list(self.contex_memory_dict.values()))
|
||||
self.contex_memory_dict.clear()
|
||||
return
|
||||
|
||||
if isinstance(keys, str):
|
||||
|
|
@ -108,7 +111,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
t_ids: List[str] = self.get_context(key)
|
||||
if t_ids:
|
||||
ids.update(t_ids)
|
||||
nodes = [self._contex_memory_dict.pop(_) for _ in ids]
|
||||
nodes = [self.contex_memory_dict.pop(_) for _ in ids]
|
||||
self.memory_store.update_memories(nodes)
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
|||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
|
|
@ -29,13 +30,11 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
|
||||
system_prompt = self.prompt_handler.contra_repeat_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.user_id)
|
||||
few_shot = self.prompt_handler.contra_repeat_few_shot.format(user_name=self.user_id)
|
||||
user_name=self.target_name)
|
||||
few_shot = self.prompt_handler.contra_repeat_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.contra_repeat_user_query.format(user_query="\n".join(user_query_list),
|
||||
user_name=self.user_id)
|
||||
contra_repeat_message = self.prompt_to_msg(system_prompt=system_prompt,
|
||||
few_shot=few_shot,
|
||||
user_query=user_query)
|
||||
user_name=self.target_name)
|
||||
contra_repeat_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"contra_repeat_message={contra_repeat_message}")
|
||||
|
||||
# call LLM
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
MemoryTypeEnum.CONVERSATION.value: message.content,
|
||||
TIME_INFER: time_infer,
|
||||
"keywords": keywords,
|
||||
**dt_handler.dt_info_dict,
|
||||
**{k: str(v) for k, v in dt_handler.dt_info_dict.items()},
|
||||
}
|
||||
|
||||
if time_infer:
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from pydantic import Field, BaseModel
|
|||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
memory_id: str = Field(uuid4(), description="unique id for memory")
|
||||
memory_id: str = Field(str(uuid4()), description="unique id for memory")
|
||||
|
||||
user_name: str = Field("", description="the user who owns the memory")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue