From 39fc87c3783f6f25182ab4ae36cad374835cf761 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 9 Jul 2024 18:00:22 +0800 Subject: [PATCH] [dev] rename timer print lines --- memory_scope/constants/common_constants.py | 2 ++ .../memory/operation/base_workflow.py | 8 ++++--- .../memory/service/chat_memory_service.py | 2 ++ memory_scope/memory/worker/base_worker.py | 6 +++-- .../memory/worker/memory_base_worker.py | 23 +++++++++++-------- .../worker/write/contra_repeat_worker.py | 11 ++++----- .../worker/write/get_observation_worker.py | 2 +- memory_scope/scheme/memory_node.py | 2 +- 8 files changed, 33 insertions(+), 23 deletions(-) diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index f2e5dedd..c27aca87 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -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" diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index f7ef319a..6353ab6d 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -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}-----") diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index 1da819e3..143bdf8b 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -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): diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 5b858ac0..633f8c87 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -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) diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 876e67a7..ec2ba660 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -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 diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index 8b96fc90..b00d3db8 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -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 diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 44454448..df0ebef6 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -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: diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index e1dcb71f..bf2dd5ee 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -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")