[dev] rename timer print lines

This commit is contained in:
jinli.yl 2024-07-09 18:00:22 +08:00
parent 88cc1d8c32
commit 39fc87c378
8 changed files with 33 additions and 23 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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