mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] add worker prompt yaml/json
This commit is contained in:
parent
7f84987762
commit
d8aceabbf0
29 changed files with 225 additions and 42 deletions
8
config/prompts/cli_memory_chat.yaml
Normal file
8
config/prompts/cli_memory_chat.yaml
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
system_prompt:
|
||||
cn: |
|
||||
你是一个可靠的小助手,你的名字叫MemoryScope
|
||||
|
||||
|
||||
memory_prompt:
|
||||
cn: |
|
||||
请记住以下信息,他们可以帮助更好地理解用户的问题。
|
||||
12
config/prompts/extract_time_worker.yaml
Normal file
12
config/prompts/extract_time_worker.yaml
Normal 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}秒。
|
||||
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
31
memory_scope/constants/language_constants.py
Normal file
31
memory_scope/constants/language_constants.py
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
DATATIME_WORD_LIST = {
|
||||
LanguageEnum.CN:
|
||||
[
|
||||
"天",
|
||||
"周",
|
||||
"月",
|
||||
"年",
|
||||
"星期",
|
||||
"点",
|
||||
"分钟",
|
||||
"小时",
|
||||
"秒",
|
||||
"上午",
|
||||
"下午",
|
||||
"早上",
|
||||
"早晨",
|
||||
"晚上",
|
||||
"中午",
|
||||
"日",
|
||||
"夜",
|
||||
"清晨",
|
||||
"傍晚",
|
||||
"凌晨",
|
||||
"岁",
|
||||
],
|
||||
LanguageEnum.EN: [
|
||||
|
||||
]
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
0
memory_scope/memory/worker/read/__init__.py
Normal file
0
memory_scope/memory/worker/read/__init__.py
Normal file
53
memory_scope/memory/worker/read/extract_time_worker.py
Normal file
53
memory_scope/memory/worker/read/extract_time_worker.py
Normal 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}")
|
||||
17
memory_scope/memory/worker/read/set_query_worker.py
Normal file
17
memory_scope/memory/worker/read/set_query_worker.py
Normal 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))
|
||||
|
|
@ -1,15 +0,0 @@
|
|||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
SYSTEM_PROMPT = {
|
||||
LanguageEnum.CN: """
|
||||
""",
|
||||
LanguageEnum.EN: """
|
||||
"""
|
||||
}
|
||||
|
||||
MEMORY_PROMPT = {
|
||||
LanguageEnum.CN: """
|
||||
""",
|
||||
LanguageEnum.EN: """
|
||||
"""
|
||||
}
|
||||
47
memory_scope/utils/prompt_handler.py
Normal file
47
memory_scope/utils/prompt_handler.py
Normal 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]
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue