[dev] add worker prompt yaml/json

This commit is contained in:
jinli.yl 2024-06-29 00:09:25 +08:00
parent 7f84987762
commit d8aceabbf0
29 changed files with 225 additions and 42 deletions

View file

@ -0,0 +1,8 @@
system_prompt:
cn: |
你是一个可靠的小助手你的名字叫MemoryScope
memory_prompt:
cn: |
请记住以下信息,他们可以帮助更好地理解用户的问题。

View file

@ -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}秒。

View file

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

View file

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

View file

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

View file

@ -0,0 +1,31 @@
from memory_scope.enumeration.language_enum import LanguageEnum
DATATIME_WORD_LIST = {
LanguageEnum.CN:
[
"",
"",
"",
"",
"星期",
"",
"分钟",
"小时",
"",
"上午",
"下午",
"早上",
"早晨",
"晚上",
"中午",
"",
"",
"清晨",
"傍晚",
"凌晨",
"",
],
LanguageEnum.EN: [
]
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -1,15 +0,0 @@
from ..enumeration.language_enum import LanguageEnum
SYSTEM_PROMPT = {
LanguageEnum.CN: """
""",
LanguageEnum.EN: """
"""
}
MEMORY_PROMPT = {
LanguageEnum.CN: """
""",
LanguageEnum.EN: """
"""
}

View file

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

View file

@ -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'(?<!^)(?=[A-Z])', '_', name).lower()
def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs):
@ -36,7 +43,7 @@ def init_instance_by_config(config: dict, default_class_path: str = "memory_scop
class_paths.extend(class_name_split)
module = import_module(".".join(class_paths))
cls_name = under_line_to_hump(class_name)
cls_name = underscore_to_camelcase(class_name)
config_copy.update(kwargs)
return getattr(module, cls_name)(**config_copy)