add user name prompt

This commit is contained in:
jinli.yl 2024-07-29 11:07:00 +08:00
parent d2cd11707b
commit 94a26dfc95
24 changed files with 100 additions and 53 deletions

View file

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

View file

@ -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"
-rank_model="gte-rerank"

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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