From d8aceabbf09f369ea0f79a4950fde5b5e76ed84f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 29 Jun 2024 00:09:25 +0800 Subject: [PATCH] [dev] add worker prompt yaml/json --- config/{demo.yaml => demo_config.yaml} | 0 {memory_scope => config}/prompts/__init__.py | 0 config/prompts/cli_memory_chat.yaml | 8 +++ .../prompts/contra_repeat_prompt.py | 0 config/prompts/extract_time_worker.yaml | 12 +++++ .../prompts/get_insight_prompt.py | 0 .../prompts/get_observation_prompt.py | 0 .../get_observation_with_time_prompt.py | 0 .../prompts/get_reflect_subject_prompt.py | 0 .../prompts/get_reflection_prompt.py | 0 .../prompts/info_filter_prompt.py | 0 .../prompts/long_contra_repeat_prompt.py | 0 .../prompts/update_insight_prompt.py | 0 .../prompts/update_profile_prompt.py | 0 memory_scope/chat/cli_memory_chat.py | 17 ++++-- memory_scope/cli.py | 8 +-- memory_scope/constants/common_constants.py | 9 ++-- memory_scope/constants/language_constants.py | 31 +++++++++++ .../memory/operation/base_workflow.py | 2 +- .../memory/operation/summary_memory.py | 2 +- memory_scope/memory/operation/write_memory.py | 2 +- .../memory/worker/memory_base_worker.py | 29 +++++++--- memory_scope/memory/worker/read/__init__.py | 0 .../memory/worker/read/extract_time_worker.py | 53 +++++++++++++++++++ .../memory/worker/read/set_query_worker.py | 17 ++++++ memory_scope/prompts/memory_chat_prompt.py | 15 ------ .../{chat => utils}/global_context.py | 0 memory_scope/utils/prompt_handler.py | 47 ++++++++++++++++ memory_scope/utils/tool_functions.py | 15 ++++-- 29 files changed, 225 insertions(+), 42 deletions(-) rename config/{demo.yaml => demo_config.yaml} (100%) rename {memory_scope => config}/prompts/__init__.py (100%) create mode 100644 config/prompts/cli_memory_chat.yaml rename {memory_scope => config}/prompts/contra_repeat_prompt.py (100%) create mode 100644 config/prompts/extract_time_worker.yaml rename {memory_scope => config}/prompts/get_insight_prompt.py (100%) rename {memory_scope => config}/prompts/get_observation_prompt.py (100%) rename {memory_scope => config}/prompts/get_observation_with_time_prompt.py (100%) rename {memory_scope => config}/prompts/get_reflect_subject_prompt.py (100%) rename {memory_scope => config}/prompts/get_reflection_prompt.py (100%) rename {memory_scope => config}/prompts/info_filter_prompt.py (100%) rename {memory_scope => config}/prompts/long_contra_repeat_prompt.py (100%) rename {memory_scope => config}/prompts/update_insight_prompt.py (100%) rename {memory_scope => config}/prompts/update_profile_prompt.py (100%) create mode 100644 memory_scope/constants/language_constants.py create mode 100644 memory_scope/memory/worker/read/__init__.py create mode 100644 memory_scope/memory/worker/read/extract_time_worker.py create mode 100644 memory_scope/memory/worker/read/set_query_worker.py delete mode 100644 memory_scope/prompts/memory_chat_prompt.py rename memory_scope/{chat => utils}/global_context.py (100%) create mode 100644 memory_scope/utils/prompt_handler.py diff --git a/config/demo.yaml b/config/demo_config.yaml similarity index 100% rename from config/demo.yaml rename to config/demo_config.yaml diff --git a/memory_scope/prompts/__init__.py b/config/prompts/__init__.py similarity index 100% rename from memory_scope/prompts/__init__.py rename to config/prompts/__init__.py diff --git a/config/prompts/cli_memory_chat.yaml b/config/prompts/cli_memory_chat.yaml new file mode 100644 index 00000000..37b7b02a --- /dev/null +++ b/config/prompts/cli_memory_chat.yaml @@ -0,0 +1,8 @@ +system_prompt: + cn: | + 你是一个可靠的小助手,你的名字叫MemoryScope + + +memory_prompt: + cn: | + 请记住以下信息,他们可以帮助更好地理解用户的问题。 diff --git a/memory_scope/prompts/contra_repeat_prompt.py b/config/prompts/contra_repeat_prompt.py similarity index 100% rename from memory_scope/prompts/contra_repeat_prompt.py rename to config/prompts/contra_repeat_prompt.py diff --git a/config/prompts/extract_time_worker.yaml b/config/prompts/extract_time_worker.yaml new file mode 100644 index 00000000..23431946 --- /dev/null +++ b/config/prompts/extract_time_worker.yaml @@ -0,0 +1,12 @@ +extract_time_prompt: + cn: | + 任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。 + 语句:{query} + 时间:{query_time_str} + 回答: + + +time_format_prompt: + cn: | + {year}年{month}月{day}日,{year}年第{week}周,{weekday},{hour}时{minute}分{second}秒。 + diff --git a/memory_scope/prompts/get_insight_prompt.py b/config/prompts/get_insight_prompt.py similarity index 100% rename from memory_scope/prompts/get_insight_prompt.py rename to config/prompts/get_insight_prompt.py diff --git a/memory_scope/prompts/get_observation_prompt.py b/config/prompts/get_observation_prompt.py similarity index 100% rename from memory_scope/prompts/get_observation_prompt.py rename to config/prompts/get_observation_prompt.py diff --git a/memory_scope/prompts/get_observation_with_time_prompt.py b/config/prompts/get_observation_with_time_prompt.py similarity index 100% rename from memory_scope/prompts/get_observation_with_time_prompt.py rename to config/prompts/get_observation_with_time_prompt.py diff --git a/memory_scope/prompts/get_reflect_subject_prompt.py b/config/prompts/get_reflect_subject_prompt.py similarity index 100% rename from memory_scope/prompts/get_reflect_subject_prompt.py rename to config/prompts/get_reflect_subject_prompt.py diff --git a/memory_scope/prompts/get_reflection_prompt.py b/config/prompts/get_reflection_prompt.py similarity index 100% rename from memory_scope/prompts/get_reflection_prompt.py rename to config/prompts/get_reflection_prompt.py diff --git a/memory_scope/prompts/info_filter_prompt.py b/config/prompts/info_filter_prompt.py similarity index 100% rename from memory_scope/prompts/info_filter_prompt.py rename to config/prompts/info_filter_prompt.py diff --git a/memory_scope/prompts/long_contra_repeat_prompt.py b/config/prompts/long_contra_repeat_prompt.py similarity index 100% rename from memory_scope/prompts/long_contra_repeat_prompt.py rename to config/prompts/long_contra_repeat_prompt.py diff --git a/memory_scope/prompts/update_insight_prompt.py b/config/prompts/update_insight_prompt.py similarity index 100% rename from memory_scope/prompts/update_insight_prompt.py rename to config/prompts/update_insight_prompt.py diff --git a/memory_scope/prompts/update_profile_prompt.py b/config/prompts/update_profile_prompt.py similarity index 100% rename from memory_scope/prompts/update_profile_prompt.py rename to config/prompts/update_profile_prompt.py diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 3ce61358..d882aec5 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -4,14 +4,14 @@ import time import questionary from memory_scope.chat.base_memory_chat import BaseMemoryChat -from memory_scope.chat.global_context import G_CONTEXT from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel -from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT from memory_scope.scheme.message import Message from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen +from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.logger import Logger +from memory_scope.utils.prompt_handler import PromptHandler from memory_scope.utils.tool_functions import char_logo @@ -39,8 +39,17 @@ class CliMemoryChat(BaseMemoryChat): self.kwargs: dict = kwargs self._logo = char_logo("MemoryScope") + self._prompt_handler: PromptHandler | None = None + self.logger = Logger.get_logger() + @property + def prompt_handler(self) -> PromptHandler: + if self._prompt_handler is None: + self._prompt_handler = PromptHandler() + self._prompt_handler.add_file_prompts(self.__class__.__name__) + return self._prompt_handler + def print_logo(self): for line in self._logo: print(line) @@ -59,11 +68,11 @@ class CliMemoryChat(BaseMemoryChat): return self._generation_model def get_system_prompt(self) -> Message: - system_prompt = SYSTEM_PROMPT[G_CONTEXT.language].strip() + system_prompt = self.prompt_handler.system_prompt memories: str = self.memory_service.read_memory() if memories: - memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] + memory_prompt = self.prompt_handler.memory_prompt system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) diff --git a/memory_scope/cli.py b/memory_scope/cli.py index e76575aa..a2ef4ae7 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -9,12 +9,12 @@ from typing import Dict, Any import fire import yaml -from memory_scope.chat.global_context import G_CONTEXT from memory_scope.enumeration.language_enum import LanguageEnum -from memory_scope.utils.logger import Logger -from memory_scope.utils.tool_functions import init_instance_by_config -from memory_scope.utils.timer import timer from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.utils.global_context import G_CONTEXT +from memory_scope.utils.logger import Logger +from memory_scope.utils.timer import timer +from memory_scope.utils.tool_functions import init_instance_by_config class CliJob(object): diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index 23c05c88..c38ed926 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -6,11 +6,13 @@ CHAT_MESSAGES = "chat_messages" CHAT_KWARGS = "chat_kwargs" -RELATED_MEMORIES = "related_memories" +QUERY_WITH_TS = "query_with_ts" + + + + -MESSAGES = "messages" -CHAT_NAME = "chat_name" PIPELINE = "pipeline" @@ -20,7 +22,6 @@ MEMORY = "memory" DEFAULT_SYSTEM_PROMPT = "default_system_prompt" -RELATED_MEMORIES = "related_memories" MODIFIED_MEMORIES = "modified_memories" diff --git a/memory_scope/constants/language_constants.py b/memory_scope/constants/language_constants.py new file mode 100644 index 00000000..921d6719 --- /dev/null +++ b/memory_scope/constants/language_constants.py @@ -0,0 +1,31 @@ +from memory_scope.enumeration.language_enum import LanguageEnum + +DATATIME_WORD_LIST = { + LanguageEnum.CN: + [ + "天", + "周", + "月", + "年", + "星期", + "点", + "分钟", + "小时", + "秒", + "上午", + "下午", + "早上", + "早晨", + "晚上", + "中午", + "日", + "夜", + "清晨", + "傍晚", + "凌晨", + "岁", + ], + LanguageEnum.EN: [ + + ] +} diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index 8b61b67f..fee81862 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -4,9 +4,9 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from itertools import zip_longest from typing import Dict, Any, List -from memory_scope.chat.global_context import G_CONTEXT from memory_scope.constants.common_constants import WORKFLOW_NAME from memory_scope.memory.worker.base_worker import BaseWorker +from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.logger import Logger from memory_scope.utils.timer import Timer from memory_scope.utils.tool_functions import init_instance_by_config diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 2cc7d00c..ffb1c297 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -1,9 +1,9 @@ import time -from memory_scope.chat.global_context import G_CONTEXT from memory_scope.constants.common_constants import RESULT, CHAT_KWARGS from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow +from memory_scope.utils.global_context import G_CONTEXT class SummaryMemory(BaseWorkflow, BaseOperation): diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index edfd9746..9a90db0b 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -1,11 +1,11 @@ import time from typing import List -from memory_scope.chat.global_context import G_CONTEXT from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT, CHAT_KWARGS from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message +from memory_scope.utils.global_context import G_CONTEXT class WriteMemory(BaseWorkflow, BaseOperation): diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index fbf7ade5..3f973ed0 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -1,14 +1,15 @@ from abc import ABCMeta -from typing import List +from typing import List, Dict -from memory_scope.chat.global_context import G_CONTEXT -from memory_scope.constants.common_constants import CHAT_MESSAGES +from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.models.base_model import BaseModel from memory_scope.scheme.message import Message from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_vector_store import BaseVectorStore +from memory_scope.utils.global_context import G_CONTEXT +from memory_scope.utils.prompt_handler import PromptHandler class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): @@ -28,15 +29,20 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._monitor: BaseMonitor | None = None self._user_id: str | None = None + self._prompt_handler: PromptHandler | None = None @property - def messages(self) -> List[Message]: + def chat_messages(self) -> List[Message]: return self.get_context(CHAT_MESSAGES) - @messages.setter - def messages(self, value): + @chat_messages.setter + def chat_messages(self, value): self.set_context(CHAT_MESSAGES, value) + @property + def chat_kwargs(self) -> Dict[str, str]: + return self.get_context(CHAT_KWARGS) + @property def embedding_model(self) -> BaseModel: if isinstance(self._embedding_model, str): @@ -62,7 +68,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._vector_store @property - def monitor(self): + def monitor(self) -> BaseMonitor: if self._monitor is None: self._monitor = G_CONTEXT.monitor return self._monitor @@ -74,9 +80,16 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._user_id = message.role_name return self._user_id + @property + def prompt_handler(self) -> PromptHandler: + if self._prompt_handler is None: + self._prompt_handler = PromptHandler() + self._prompt_handler.add_file_prompts(self.__class__.__name__) + return self._prompt_handler + def __getattr__(self, key: str): return self.kwargs[key] @staticmethod - def get_prompt(prompt: dict) -> str: + def get_language_prompt(prompt: dict) -> str: return prompt[G_CONTEXT.language] diff --git a/memory_scope/memory/worker/read/__init__.py b/memory_scope/memory/worker/read/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py new file mode 100644 index 00000000..34107113 --- /dev/null +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -0,0 +1,53 @@ +import re + +from memory_scope.constants.common_constants import DATATIME_KEY_MAP, QUERY_WITH_TS, EXTRACT_TIME_DICT +from memory_scope.constants.language_constants import DATATIME_WORD_LIST +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.utils.tool_functions import time_to_formatted_str + + +class ExtractTimeWorker(MemoryBaseWorker): + Extract_Time_PATTERN = r'-\s*(\S+):(\d+)' + + def _run(self): + query, query_timestamp = self.get_context(QUERY_WITH_TS) + + # find datetime keyword + contain_datetime = False + for datetime_word in self.get_language_prompt(DATATIME_WORD_LIST): + if datetime_word in query: + contain_datetime = True + break + if not contain_datetime: + self.logger.info(f"contain_datetime={contain_datetime}") + return + + # prepare prompt + query_time_str = time_to_formatted_str(dt=query_timestamp, + date_format="", + string_format=self.prompt_handler.time_format_prompt) + extract_time_prompt: str = self.prompt_handler.extract_time_prompt + extract_time_prompt: str = extract_time_prompt.format(query=query, query_time_str=query_time_str) + self.logger.info(f"extract_time_prompt={extract_time_prompt}") + + # call sft model + response = self.generation_model.call(prompt=extract_time_prompt, + model_name=self.extra_time_model, + max_token=self.extra_time_max_token, + temperature=self.extra_time_temperature, + top_k=self.extra_time_top_k) + + # if empty, return + if not response.status or not response.message.content: + return + response_text = response.message.content + + # re-match time info to dict + extract_time_dict = {} + matches = re.findall(self.Extract_Time_PATTERN, response_text) + for key, value in matches: + if key in DATATIME_KEY_MAP.keys(): + extract_time_dict[DATATIME_KEY_MAP[key]] = value + self.set_context(EXTRACT_TIME_DICT, extract_time_dict) + + self.logger.info(f"response_text={response_text} filters={extract_time_dict}") diff --git a/memory_scope/memory/worker/read/set_query_worker.py b/memory_scope/memory/worker/read/set_query_worker.py new file mode 100644 index 00000000..e0a6c84c --- /dev/null +++ b/memory_scope/memory/worker/read/set_query_worker.py @@ -0,0 +1,17 @@ +import datetime + +from memory_scope.constants.common_constants import QUERY_WITH_TS +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker + + +class SetQueryWorker(MemoryBaseWorker): + + def _run(self): + if "query" in self.chat_kwargs: + query = self.chat_kwargs["query"] + query_timestamp = int(datetime.datetime.now().timestamp()) + else: + query = self.messages[-1].content + query_timestamp = self.messages[-1].time_created + + self.set_context(QUERY_WITH_TS, (query, query_timestamp)) diff --git a/memory_scope/prompts/memory_chat_prompt.py b/memory_scope/prompts/memory_chat_prompt.py deleted file mode 100644 index 0c762ec9..00000000 --- a/memory_scope/prompts/memory_chat_prompt.py +++ /dev/null @@ -1,15 +0,0 @@ -from ..enumeration.language_enum import LanguageEnum - -SYSTEM_PROMPT = { - LanguageEnum.CN: """ -""", - LanguageEnum.EN: """ -""" -} - -MEMORY_PROMPT = { - LanguageEnum.CN: """ -""", - LanguageEnum.EN: """ -""" -} diff --git a/memory_scope/chat/global_context.py b/memory_scope/utils/global_context.py similarity index 100% rename from memory_scope/chat/global_context.py rename to memory_scope/utils/global_context.py diff --git a/memory_scope/utils/prompt_handler.py b/memory_scope/utils/prompt_handler.py new file mode 100644 index 00000000..8fc76458 --- /dev/null +++ b/memory_scope/utils/prompt_handler.py @@ -0,0 +1,47 @@ +import json +import os.path +from typing import Dict + +import yaml + +from memory_scope.utils.global_context import G_CONTEXT +from memory_scope.utils.tool_functions import camelcase_to_underscore + + +class PromptHandler(object): + + def __init__(self, default_prompt_dir: str = "config/prompts"): + self._default_prompt_dir: str = default_prompt_dir + self._prompt_dict: Dict[str, str] = {} + + def add_file_prompts(self, name: str, to_underscore: bool = True): + if to_underscore: + name: str = camelcase_to_underscore(name) + class_path = os.path.join(self._default_prompt_dir, name) + if os.path.exists(f"{class_path}.yaml"): + with open(class_path) as f: + prompt_language_dict = yaml.load(f, yaml.FullLoader) + elif os.path.exists(f"{class_path}.json"): + with open(class_path) as f: + prompt_language_dict = json.load(f) + else: + raise RuntimeError(f"{class_path}.yaml/json is not exists!") + + for key, language_dict in prompt_language_dict.items(): + prompts = language_dict.get(G_CONTEXT.language) + if not prompts: + raise RuntimeError(f"{key}.prompt is empty!") + self._prompt_dict[key] = prompts + + @property + def prompt_dict(self): + return self._prompt_dict + + def __getitem__(self, key: str): + return self._prompt_dict[key] + + def __setitem__(self, key: str, value: str): + self._prompt_dict[key] = value + + def __getattr__(self, key: str): + return self._prompt_dict[key] diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 11e85e7b..0e7b79ab 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -13,9 +13,16 @@ from memory_scope.constants.common_constants import WEEKDAYS from memory_scope.enumeration.message_role_enum import MessageRoleEnum -def under_line_to_hump(underline_str): - sub = re.sub(r"(_\w)", lambda x: x.group(1)[1].upper(), underline_str) - return sub[0:1].upper() + sub[1:] +def underscore_to_camelcase(name: str, is_first_title: bool = True): + name_split = name.split("_") + if is_first_title: + return "".join(x.title() for x in name_split[1:]) + else: + return name_split[0] + ''.join(x.title() for x in name_split[1:]) + + +def camelcase_to_underscore(name: str): + return re.sub(r'(?