From 94a26dfc957b72b73609e942a043e2ad3be60b2f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 29 Jul 2024 11:07:00 +0800 Subject: [PATCH 1/2] add user name prompt --- examples/api/chat_example.py | 5 ++- examples/cli/dash_cli_cn2.sh | 4 ++- memoryscope/constants/language_constants.py | 5 +++ memoryscope/core/chat/api_memory_chat.py | 5 ++- memoryscope/core/chat/base_memory_chat.py | 6 ++++ memoryscope/core/chat/cli_memory_chat.py | 6 ++-- memoryscope/core/config/arguments.py | 4 +++ memoryscope/core/config/config_manager.py | 10 +++--- memoryscope/core/config/demo_config.yaml | 9 +++-- .../core/service/base_memory_service.py | 4 +-- memoryscope/core/utils/tool_functions.py | 24 ++++--------- .../worker/backend/contra_repeat_worker.py | 4 +-- .../get_observation_with_time_worker.py | 3 +- .../worker/backend/get_observation_worker.py | 3 +- .../backend/get_reflection_subject_worker.py | 3 +- .../core/worker/backend/info_filter_worker.py | 3 +- .../backend/long_contra_repeat_worker.py | 3 +- .../worker/backend/update_insight_worker.py | 5 +-- .../worker/frontend/extract_time_worker.py | 3 +- .../worker/frontend/print_memory_worker.py | 2 +- .../core/worker/frontend/set_query_worker.py | 4 +-- memoryscope/core/worker/memory_base_worker.py | 34 +++++++++++++++++++ tests/worker/test_workers_cn.py | 2 ++ tests/worker/test_workers_en.py | 2 ++ 24 files changed, 100 insertions(+), 53 deletions(-) diff --git a/examples/api/chat_example.py b/examples/api/chat_example.py index 393b004e..11acfff5 100644 --- a/examples/api/chat_example.py +++ b/examples/api/chat_example.py @@ -6,6 +6,8 @@ from memoryscope import MemoryScope, Arguments arguments = Arguments( language="cn", + human_name="用户", + assistant_name="AI", logger_to_screen=False, memory_chat_class="api_memory_chat", generation_backend="dashscope_generation", @@ -48,11 +50,12 @@ def chat_example3(): def chat_example4(): with MemoryScope(arguments=arguments) as ms: memory_chat = ms.default_memory_chat + memory_chat.start_backend_service() response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") print("回答1:\n" + response.message.content) - memory_chat.memory_service.consolidate_memory() + memory_chat.run_service_operation("consolidate_memory") response = memory_chat.chat_with_memory(query="你知道我的乐器爱好是什么?", history_message_strategy=None) print("回答2:\n" + response.message.content) diff --git a/examples/cli/dash_cli_cn2.sh b/examples/cli/dash_cli_cn2.sh index 14c92cee..343ff729 100644 --- a/examples/cli/dash_cli_cn2.sh +++ b/examples/cli/dash_cli_cn2.sh @@ -1,10 +1,12 @@ python memoryscope/cli.py \ -language="cn" \ -memory_chat_class="cli_memory_chat" \ + -human_name="锦鲤" \ + -assistant_name="AI" \ -generation_backend="dashscope_generation" \ -generation_model="qwen-max" \ -embedding_backend="dashscope_embedding" \ -embedding_model="text-embedding-v2" \ -use_dummy_ranker=False \ -rank_backend="dashscope_rank" \ - -rank_model="gte-rerank" \ No newline at end of file + -rank_model="gte-rerank" diff --git a/memoryscope/constants/language_constants.py b/memoryscope/constants/language_constants.py index 94fbcd47..b02895a7 100644 --- a/memoryscope/constants/language_constants.py +++ b/memoryscope/constants/language_constants.py @@ -208,3 +208,8 @@ TIME_INFER_WORD = { LanguageEnum.CN: "推断时间", LanguageEnum.EN: "Inference time" } + +USER_NAME_EXPRESSION = { + LanguageEnum.CN: "用户姓名是{name}。", + LanguageEnum.EN: "User's name is {name}." +} diff --git a/memoryscope/core/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py index 6a8078f7..6028b970 100644 --- a/memoryscope/core/chat/api_memory_chat.py +++ b/memoryscope/core/chat/api_memory_chat.py @@ -1,7 +1,7 @@ from typing import List, Optional, Literal from memoryscope.constants.common_constants import MEMORIES -from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME, USER_NAME_EXPRESSION from memoryscope.core.chat.base_memory_chat import BaseMemoryChat from memoryscope.core.memoryscope_context import MemoryscopeContext from memoryscope.core.models.base_model import BaseModel @@ -175,6 +175,9 @@ class ApiMemoryChat(BaseMemoryChat): system_prompt_list.append(memory_prompt) else: system_prompt_list.append(self.prompt_handler.memory_prompt) + + if self.human_name != DEFAULT_HUMAN_NAME[self.context.language]: + system_prompt_list.append(USER_NAME_EXPRESSION[self.context.language].format(name=self.human_name)) system_prompt_list.append(memories) if extra_memories: diff --git a/memoryscope/core/chat/base_memory_chat.py b/memoryscope/core/chat/base_memory_chat.py index e311f072..b9e48293 100644 --- a/memoryscope/core/chat/base_memory_chat.py +++ b/memoryscope/core/chat/base_memory_chat.py @@ -58,6 +58,12 @@ class BaseMemoryChat(metaclass=ABCMeta): """ raise NotImplementedError + def start_backend_service(self): + self.memory_service.start_backend_service() + + def run_service_operation(self, name: str, **kwargs): + return self.memory_service.run_operation(name, **kwargs) + def run(self): """ Abstract method to run the chat system. diff --git a/memoryscope/core/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py index 150c79b3..e7e2d6f3 100644 --- a/memoryscope/core/chat/cli_memory_chat.py +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -131,7 +131,7 @@ class CliMemoryChat(ApiMemoryChat): refresh_time = int(refresh_time) self.memory_service.stop_backend_service() while True: - result = self.memory_service.do_operation(name=command, **kwargs) + result = self.memory_service.run_operation(name=command, **kwargs) os.system("clear") self.print_logo() if result: @@ -143,7 +143,7 @@ class CliMemoryChat(ApiMemoryChat): time.sleep(refresh_time) else: - result = self.memory_service.do_operation(name=command, **kwargs) + result = self.memory_service.run_operation(name=command, **kwargs) if result: if isinstance(result, list): result = "\n".join([str(x) for x in result]) @@ -187,7 +187,7 @@ class CliMemoryChat(ApiMemoryChat): questionary.print(f"{self.assistant_name}: ", end="", style="bold") # Fetch and display AI's response - self.memory_service.start_backend_service() + self.start_backend_service() self.chat_with_memory(query=query) except KeyboardInterrupt: diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index 44ed574f..dec565d4 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -17,6 +17,10 @@ class Arguments(object): memory_chat_class: str = field(default="cli_memory_chat", metadata={ "help": "cli_memory_chat(Command-line interaction), api_memory_chat(API interface interaction), etc."}) + human_name: str = field(default="user") + + assistant_name: str = field(default="AI") + consolidate_memory_interval_time: int = field(default=1, metadata={ "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py index e30d51e3..81261ec1 100644 --- a/memoryscope/core/config/config_manager.py +++ b/memoryscope/core/config/config_manager.py @@ -6,10 +6,8 @@ from typing import Optional, Literal import yaml -from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME from memoryscope.core.config.arguments import Arguments from memoryscope.core.utils.logger import Logger -from memoryscope.enumeration.language_enum import LanguageEnum class ConfigManager(object): @@ -92,16 +90,16 @@ class ConfigManager(object): stream = arguments.memory_chat_class in ["cli_memory_chat", ] config.update({ "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), - "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], - "assistant_name": "AI", + "human_name": arguments.human_name if arguments.human_name else "", + "assistant_name": arguments.assistant_name if arguments.assistant_name else "", "stream": stream, }) @staticmethod def update_memory_service_by_arguments(config: dict, arguments: Arguments): config.update({ - "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], - "assistant_name": "AI", + "human_name": arguments.human_name if arguments.human_name else "", + "assistant_name": arguments.assistant_name if arguments.assistant_name else "", }) config["memory_operations"]["consolidate_memory"]["interval_time"] = \ arguments.consolidate_memory_interval_time diff --git a/memoryscope/core/config/demo_config.yaml b/memoryscope/core/config/demo_config.yaml index 60db4b34..ad6c1334 100644 --- a/memoryscope/core/config/demo_config.yaml +++ b/memoryscope/core/config/demo_config.yaml @@ -11,6 +11,9 @@ memory_chat: class: core.chat.cli_memory_chat memory_service: memoryscope_service generation_model: generation_model + stream: true + human_name: user + assistant_name: AI memory_service: memoryscope_service: @@ -29,7 +32,7 @@ memory_service: list_memory: class: core.operation.frontend_operation workflow: set_query,retrieve_top_memory,print_memory - description: "read all long-term memory of the user" + description: "read all long-term memory of the user, use `refresh_time=5` to refresh screen." delete_memory: class: core.operation.frontend_operation @@ -49,13 +52,13 @@ memory_service: consolidate_memory: class: core.operation.consolidate_memory_op workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory - description: "summary user's observation memory" + description: "summary user's observation memory, run backend." interval_time: 1 reflect_and_reconsolidate: class: core.operation.backend_operation workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory - description: "summary user's insight memory" + description: "summary user's insight memory, run backend." interval_time: 15 worker: diff --git a/memoryscope/core/service/base_memory_service.py b/memoryscope/core/service/base_memory_service.py index 68acc566..b89f9d08 100644 --- a/memoryscope/core/service/base_memory_service.py +++ b/memoryscope/core/service/base_memory_service.py @@ -58,7 +58,7 @@ class BaseMemoryService(metaclass=ABCMeta): def stop_backend_service(self, wait_service_end: bool = False): pass - def do_operation(self, name: str, **kwargs): + def run_operation(self, name: str, **kwargs): """ Executes a specific operation by its name with provided keyword arguments. @@ -79,4 +79,4 @@ class BaseMemoryService(metaclass=ABCMeta): def __getattr__(self, name: str): assert name in self._operation_dict, f"operation={name} is not registered!" - return lambda **kwargs: self.do_operation(name=name, **kwargs) + return lambda **kwargs: self.run_operation(name=name, **kwargs) diff --git a/memoryscope/core/utils/tool_functions.py b/memoryscope/core/utils/tool_functions.py index 6d5a6834..e7ce7332 100644 --- a/memoryscope/core/utils/tool_functions.py +++ b/memoryscope/core/utils/tool_functions.py @@ -110,24 +110,14 @@ def prompt_to_msg(system_prompt: str, Returns: List[Message]: A list of Message objects, each representing a part of the conversation setup. """ - if concat_system_prompt: - user_message = Message( - role=MessageRoleEnum.USER.value, - content="\n".join( - [x.strip() for x in [few_shot, system_prompt, user_query]] - ), - ) - else: - user_message = Message( - role=MessageRoleEnum.USER.value, - content="\n".join([x.strip() for x in [few_shot, user_query]]), - ) - return [ - Message(role=MessageRoleEnum.SYSTEM.value, content=system_prompt.strip()), # System message - user_message - # User message combining few shot, system prompt, and user query - ] + system_message = Message(role=MessageRoleEnum.SYSTEM.value, content=system_prompt.strip()) + if concat_system_prompt: + user_content_list = [system_prompt, few_shot, user_query] + else: + user_content_list = [few_shot, user_query] + user_message = Message(role=MessageRoleEnum.USER.value, content="\n".join([x.strip() for x in user_content_list])) + return [system_message, user_message] def char_logo(words: str, seed: int = time.time_ns(), color=None): diff --git a/memoryscope/core/worker/backend/contra_repeat_worker.py b/memoryscope/core/worker/backend/contra_repeat_worker.py index 177523eb..87c99b76 100644 --- a/memoryscope/core/worker/backend/contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/contra_repeat_worker.py @@ -3,7 +3,6 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_NODES from memoryscope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, CONTAINED_WORD from memoryscope.core.utils.response_text_parser import ResponseTextParser -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.store_status_enum import StoreStatusEnum from memoryscope.scheme.memory_node import MemoryNode @@ -68,7 +67,8 @@ class ContraRepeatWorker(MemoryBaseWorker): 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)) - contra_repeat_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + contra_repeat_message = self.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/memoryscope/core/worker/backend/get_observation_with_time_worker.py b/memoryscope/core/worker/backend/get_observation_with_time_worker.py index b1fa4c79..f0346806 100644 --- a/memoryscope/core/worker/backend/get_observation_with_time_worker.py +++ b/memoryscope/core/worker/backend/get_observation_with_time_worker.py @@ -3,7 +3,6 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_WITH_TIME_NODES from memoryscope.constants.language_constants import COLON_WORD from memoryscope.core.utils.datetime_handler import DatetimeHandler -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.backend.get_observation_worker import GetObservationWorker from memoryscope.scheme.message import Message @@ -66,7 +65,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): user_name=self.target_name) # Assemble the final message for observation retrieval - obtain_obs_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + obtain_obs_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) # Log the constructed message for debugging purposes self.logger.info(f"obtain_obs_message={obtain_obs_message}") diff --git a/memoryscope/core/worker/backend/get_observation_worker.py b/memoryscope/core/worker/backend/get_observation_worker.py index 78c7eeda..dcaccc4c 100644 --- a/memoryscope/core/worker/backend/get_observation_worker.py +++ b/memoryscope/core/worker/backend/get_observation_worker.py @@ -4,7 +4,6 @@ from memoryscope.constants.common_constants import NEW_OBS_NODES, TIME_INFER from memoryscope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_WORD from memoryscope.core.utils.datetime_handler import DatetimeHandler from memoryscope.core.utils.response_text_parser import ResponseTextParser -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum @@ -102,7 +101,7 @@ class GetObservationWorker(MemoryBaseWorker): user_name=self.target_name) # Combine system prompt, few-shot, and user query into a single message for obtaining observations - obtain_obs_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + obtain_obs_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) # Log the constructed observation message self.logger.info(f"obtain_obs_message={obtain_obs_message}") diff --git a/memoryscope/core/worker/backend/get_reflection_subject_worker.py b/memoryscope/core/worker/backend/get_reflection_subject_worker.py index f9323372..06c6658f 100644 --- a/memoryscope/core/worker/backend/get_reflection_subject_worker.py +++ b/memoryscope/core/worker/backend/get_reflection_subject_worker.py @@ -4,7 +4,6 @@ from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_ from memoryscope.constants.language_constants import COMMA_WORD from memoryscope.core.utils.datetime_handler import DatetimeHandler from memoryscope.core.utils.response_text_parser import ResponseTextParser -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum @@ -90,7 +89,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): user_query="\n".join(user_query_list)) # Construct and log reflection message - reflect_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + reflect_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"reflect_message={reflect_message}") # Invoke Language Model for new insights diff --git a/memoryscope/core/worker/backend/info_filter_worker.py b/memoryscope/core/worker/backend/info_filter_worker.py index 78307c68..d73ed56b 100644 --- a/memoryscope/core/worker/backend/info_filter_worker.py +++ b/memoryscope/core/worker/backend/info_filter_worker.py @@ -2,7 +2,6 @@ from typing import List from memoryscope.constants.language_constants import COLON_WORD from memoryscope.core.utils.response_text_parser import ResponseTextParser -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.scheme.message import Message @@ -63,7 +62,7 @@ class InfoFilterWorker(MemoryBaseWorker): user_name=self.target_name) few_shot = self.prompt_handler.info_filter_few_shot.format(user_name=self.target_name) user_query = self.prompt_handler.info_filter_user_query.format(user_query="\n".join(user_query_list)) - info_filter_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + info_filter_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"info_filter_message={info_filter_message}") # call llm diff --git a/memoryscope/core/worker/backend/long_contra_repeat_worker.py b/memoryscope/core/worker/backend/long_contra_repeat_worker.py index 527e1201..a150bd3f 100644 --- a/memoryscope/core/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/long_contra_repeat_worker.py @@ -3,7 +3,6 @@ from typing import List, Dict from memoryscope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_NODES from memoryscope.constants.language_constants import NONE_WORD, CONTAINED_WORD, CONTRADICTORY_WORD from memoryscope.core.utils.response_text_parser import ResponseTextParser -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum @@ -98,7 +97,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): few_shot = self.prompt_handler.long_contra_repeat_few_shot.format(user_name=self.target_name) user_query = self.prompt_handler.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list)) - long_contra_repeat_message = prompt_to_msg(system_prompt=system_prompt, + long_contra_repeat_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"long_contra_repeat_message={long_contra_repeat_message}") diff --git a/memoryscope/core/worker/backend/update_insight_worker.py b/memoryscope/core/worker/backend/update_insight_worker.py index 51f8d7e9..eed827c9 100644 --- a/memoryscope/core/worker/backend/update_insight_worker.py +++ b/memoryscope/core/worker/backend/update_insight_worker.py @@ -5,7 +5,7 @@ from memoryscope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NO from memoryscope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD from memoryscope.core.utils.datetime_handler import DatetimeHandler from memoryscope.core.utils.response_text_parser import ResponseTextParser -from memoryscope.core.utils.tool_functions import prompt_to_msg, cosine_similarity +from memoryscope.core.utils.tool_functions import cosine_similarity from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.scheme.memory_node import MemoryNode @@ -152,7 +152,8 @@ class UpdateInsightWorker(MemoryBaseWorker): insight_key=insight_node.key, insight_key_value=insight_node.key + self.get_language_value(COLON_WORD) + insight_node.value) # Construct the message for LLM interaction - update_insight_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + update_insight_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, + user_query=user_query) self.logger.info(f"Generated insight update message: {update_insight_message}") # Call the Language Model for insight update diff --git a/memoryscope/core/worker/frontend/extract_time_worker.py b/memoryscope/core/worker/frontend/extract_time_worker.py index b2d864e2..cff18073 100644 --- a/memoryscope/core/worker/frontend/extract_time_worker.py +++ b/memoryscope/core/worker/frontend/extract_time_worker.py @@ -4,7 +4,6 @@ from typing import Dict from memoryscope.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT from memoryscope.constants.language_constants import DATATIME_KEY_MAP from memoryscope.core.utils.datetime_handler import DatetimeHandler -from memoryscope.core.utils.tool_functions import prompt_to_msg from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker @@ -44,7 +43,7 @@ class ExtractTimeWorker(MemoryBaseWorker): system_prompt = self.prompt_handler.extract_time_system few_shot = self.prompt_handler.extract_time_few_shot user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str) - extract_time_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) + extract_time_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"extract_time_message={extract_time_message}") # Invoke the LLM to generate a response diff --git a/memoryscope/core/worker/frontend/print_memory_worker.py b/memoryscope/core/worker/frontend/print_memory_worker.py index 5154371f..00da5855 100644 --- a/memoryscope/core/worker/frontend/print_memory_worker.py +++ b/memoryscope/core/worker/frontend/print_memory_worker.py @@ -51,7 +51,7 @@ class PrintMemoryWorker(MemoryBaseWorker): elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]: j += 1 observation_memory_list.append(f"{dt}] {j}. {node.content} " - f"status({node.obs_reflected},{node.obs_updated})") + f"[status({node.obs_reflected},{node.obs_updated})") elif MemoryTypeEnum(node.memory_type) is MemoryTypeEnum.INSIGHT: k += 1 diff --git a/memoryscope/core/worker/frontend/set_query_worker.py b/memoryscope/core/worker/frontend/set_query_worker.py index 85961bc9..9ab08509 100644 --- a/memoryscope/core/worker/frontend/set_query_worker.py +++ b/memoryscope/core/worker/frontend/set_query_worker.py @@ -39,8 +39,8 @@ class SetQueryWorker(MemoryBaseWorker): # check role_name role_name = self.chat_kwargs.get("role_name") if role_name: - assert role_name == self.target_name, \ - f"role_name={role_name} is not supported in human/assistant memory workflow!" + assert role_name == self.target_name, (f"role_name={role_name} <> target_name={self.target_name} " + f"is not supported in human/assistant memory workflow!") elif self.chat_messages: # If no explicit query is given, use the content of the latest chat message diff --git a/memoryscope/core/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py index 85b1ad41..a7345308 100644 --- a/memoryscope/core/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -3,6 +3,7 @@ from typing import List, Dict, Any from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \ WORKFLOW_NAME, MEMORY_MANAGER +from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME, USER_NAME_EXPRESSION from memoryscope.core.memoryscope_context import MemoryscopeContext from memoryscope.core.models.base_model import BaseModel from memoryscope.core.storage.base_memory_store import BaseMemoryStore @@ -11,6 +12,7 @@ from memoryscope.core.utils.prompt_handler import PromptHandler from memoryscope.core.worker.base_worker import BaseWorker from memoryscope.core.worker.memory_manager import MemoryManager from memoryscope.enumeration.language_enum import LanguageEnum +from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.scheme.message import Message @@ -218,3 +220,35 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): if isinstance(languages, list): return [x[self.language] for x in languages] return languages[self.language] + + def prompt_to_msg(self, + system_prompt: str, + few_shot: str, + user_query: str, + concat_system_prompt: bool = True) -> List[Message]: + """ + Converts input strings into a structured list of message objects suitable for AI interactions. + + Args: + system_prompt (str): The system-level instruction or context. + few_shot (str): An example or demonstration input, often used for illustrating expected behavior. + user_query (str): The actual user query or prompt to be processed. + concat_system_prompt(bool): Concat system prompt again or not in the user message. + A simple method to improve the effectiveness for some LLMs. Defaults to True. + + Returns: + List[Message]: A list of Message objects, each representing a part of the conversation setup. + """ + system_content = "" + if self.target_name != DEFAULT_HUMAN_NAME[self.language]: + system_content += USER_NAME_EXPRESSION[self.language].format(name=self.target_name) + system_content += system_prompt.strip() + system_message = Message(role=MessageRoleEnum.SYSTEM.value, content=system_content) + + if concat_system_prompt: + user_content_list = [system_content, few_shot, user_query] + else: + user_content_list = [few_shot, user_query] + user_message = Message(role=MessageRoleEnum.USER.value, + content="\n".join([x.strip() for x in user_content_list])) + return [system_message, user_message] diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index 77cad3ff..42c4e68c 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -19,6 +19,8 @@ class TestWorkersCn(unittest.TestCase): def setUp(self): arguments = Arguments( language="cn", + human_name="用户", + assistant_name="AI", memory_chat_class="api_memory_chat", generation_backend="dashscope_generation", generation_model="qwen-max", diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index ff421253..eff1e4f0 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -19,6 +19,8 @@ class TestWorkersEn(unittest.TestCase): def setUp(self): arguments = Arguments( language="en", + human_name="user", + assistant_name="AI", memory_chat_class="api_memory_chat", generation_backend="dashscope_generation", generation_model="qwen-max", From 84abc4a325fd58c666f34ac5a9a4a4240b66c5d1 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 29 Jul 2024 11:08:19 +0800 Subject: [PATCH 2/2] [dev] code format --- memoryscope/core/worker/backend/long_contra_repeat_worker.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/memoryscope/core/worker/backend/long_contra_repeat_worker.py b/memoryscope/core/worker/backend/long_contra_repeat_worker.py index a150bd3f..22629742 100644 --- a/memoryscope/core/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/long_contra_repeat_worker.py @@ -98,8 +98,8 @@ class LongContraRepeatWorker(MemoryBaseWorker): user_query = self.prompt_handler.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list)) long_contra_repeat_message = self.prompt_to_msg(system_prompt=system_prompt, - few_shot=few_shot, - user_query=user_query) + few_shot=few_shot, + user_query=user_query) self.logger.info(f"long_contra_repeat_message={long_contra_repeat_message}") # Invokes the language model for processing the constructed prompt