[dev] fix demo config

This commit is contained in:
jinli.yl 2024-07-08 23:26:10 +08:00
parent 3563d6aee4
commit 523b752a3d
24 changed files with 145 additions and 123 deletions

View file

@ -13,7 +13,6 @@ memory_service:
class: memory.service.chat_memory_service
history_msg_count: 32
contextual_msg_count: 6
read_memory_key: read_memory
memory_operations:
read_message:
class: memory.operation.read_message
@ -30,12 +29,12 @@ memory_service:
class: memory.operation.write_memory
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
description: "write observation memories of the user"
interval_time: 60
interval_time: 5
summary_memory:
class: memory.operation.summary_memory
workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
description: "summary observation memories of the user"
interval_time: 300
interval_time: 60
worker:
dummy:

View file

@ -1,5 +1,6 @@
import os
import time
from typing import List
import questionary
@ -75,21 +76,29 @@ class CliMemoryChat(BaseMemoryChat):
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
return self._generation_model
@property
def system_prompt_with_memory(self) -> Message:
system_prompt = self.prompt_handler.system_prompt
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
self.memory_service.add_messages(new_message)
messages: List[Message] = []
# add memory to system prompt
system_prompt = self.prompt_handler.system_prompt
memories: str = self.memory_service.read_memory()
if memories:
memory_prompt = self.prompt_handler.memory_prompt
system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]])
messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt))
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)
# add history messages
history_messages = self.memory_service.read_message()
if history_messages:
messages.extend(history_messages)
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
self.memory_service.add_messages(new_message)
return self.generation_model.call(messages=[self.system_prompt_with_memory, new_message], stream=self.stream)
# add new_message
messages.append(new_message)
self.logger.info(f"messages={messages}")
return self.generation_model.call(messages=messages, stream=self.stream)
@staticmethod
def parse_query_command(query: str):

View file

@ -1,8 +1,11 @@
system_prompt:
cn: |
你是一个可靠的小助手你的名字叫MemoryScope
你是一个可靠的小助手你的名字叫MemoryScope。
en: |
You are a helpful assistant, your name is MemoryScope.
memory_prompt:
cn: |
请记住以下信息,他们可以帮助更好地理解用户的问题。
en: |
Please remember the following information, as they can help better understand the user's question.

View file

@ -0,0 +1,57 @@
import time
from abc import abstractmethod
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.logger import Logger
class BaseBackendOperation(BaseOperation):
operation_type: OPERATION_TYPE = "backend"
def __init__(self, interval_time: int, **kwargs):
super(BaseBackendOperation, self).__init__(**kwargs)
self.interval_time: int = interval_time
self._operation_status_run: bool = False
self._loop_switch: bool = False
self._run_thread = None
self.logger = Logger.get_logger()
@abstractmethod
def _run_operation(self, **kwargs):
raise NotImplementedError
def run_operation(self, **kwargs):
if self._operation_status_run:
return
self._operation_status_run = True
result = None
try:
result = self._run_operation(**kwargs)
except Exception as e:
self.logger.exception(f"{self.name} encounter exception. args={e.args}")
self._operation_status_run = False
return result
def _loop_operation(self):
while self._loop_switch:
for _ in range(self.interval_time):
if self._loop_switch:
time.sleep(1)
else:
break
if self._loop_switch:
self.run_operation()
def run_operation_backend(self):
if not self._loop_switch:
self._loop_switch = True
self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation)
def stop_operation_backend(self):
self._loop_switch = False

View file

@ -18,9 +18,3 @@ class BaseOperation(metaclass=ABCMeta):
@abstractmethod
def run_operation(self, **kwargs):
raise NotImplementedError
def run_operation_backend(self):
pass
def stop_operation_backend(self):
pass

View file

@ -17,6 +17,7 @@ class ReadMemory(BaseWorkflow, BaseOperation):
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
self.chat_messages: List[Message] = chat_messages
self.his_msg_count: int = his_msg_count

View file

@ -1,58 +1,22 @@
import time
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_backend_operation import BaseBackendOperation
from memory_scope.memory.operation.base_operation import OPERATION_TYPE
from memory_scope.memory.operation.base_workflow import BaseWorkflow
from memory_scope.utils.global_context import G_CONTEXT
class SummaryMemory(BaseWorkflow, BaseOperation):
class SummaryMemory(BaseWorkflow, BaseBackendOperation):
operation_type: OPERATION_TYPE = "backend"
def __init__(self,
name: str,
description: str,
interval_time: int = 300,
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
self.interval_time: int = interval_time
self._operation_status_run: bool = False
self._loop_switch: bool = False
self._run_thread = None
def __init__(self, **kwargs):
super().__init__(**kwargs)
BaseBackendOperation.__init__(self, **kwargs)
def init_workflow(self):
self.init_workers()
def run_operation(self, **kwargs):
if self._operation_status_run:
return
self._operation_status_run = True
def _run_operation(self, **kwargs):
self.context[CHAT_KWARGS] = kwargs
self.run_workflow()
result = self.context.get(RESULT)
self.context.clear()
self._operation_status_run = False
return result
def _loop_operation(self):
while self._loop_switch:
for _ in range(self.interval_time):
if self._loop_switch:
time.sleep(1)
else:
break
if self._loop_switch:
self.run_operation()
def run_operation_backend(self):
if not self._loop_switch:
self._loop_switch = True
self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation)
def stop_operation_backend(self):
self._loop_switch = False

View file

@ -1,38 +1,30 @@
import time
from typing import List
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_backend_operation import BaseBackendOperation
from memory_scope.memory.operation.base_operation import 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):
class WriteMemory(BaseWorkflow, BaseBackendOperation):
operation_type: OPERATION_TYPE = "backend"
def __init__(self,
name: str,
description: str,
chat_messages: List[Message],
his_msg_count: int = 0,
message_lock=None,
interval_time: int = 60,
contextual_msg_count: int = 6,
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
super().__init__(**kwargs)
BaseBackendOperation.__init__(self, **kwargs)
self.chat_messages: List[Message] = chat_messages
self.his_msg_count: int = his_msg_count
self.message_lock = message_lock
self.interval_time: int = interval_time
self.contextual_msg_count: int = contextual_msg_count
self._operation_status_run: bool = False
self._loop_switch: bool = False
@property
def not_memorized_size(self):
return sum([not x.memorized for x in self.chat_messages])
@ -46,40 +38,19 @@ class WriteMemory(BaseWorkflow, BaseOperation):
def init_workflow(self):
self.init_workers()
def run_operation(self, **kwargs):
if self._operation_status_run:
return
self._operation_status_run = True
def _run_operation(self, **kwargs):
self.context[CHAT_KWARGS] = kwargs
not_memorized_size = self.not_memorized_size
if not_memorized_size < self.contextual_msg_count:
self.logger.info(f"not_memorized_size={not_memorized_size} < "
f"contextual_msg_count={self.contextual_msg_count}, skip.")
return
self._operation_status_run = True
max_count = not_memorized_size + self.his_msg_count
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
self.run_workflow()
result = self.context.get(RESULT)
self.context.clear()
self.set_memorized()
self._operation_status_run = False
return result
def _loop_operation(self):
while self._loop_switch:
for _ in range(self.interval_time):
if self._loop_switch:
time.sleep(1)
else:
break
if self._loop_switch:
self.run_operation()
def run_operation_backend(self):
if not self._loop_switch:
self._loop_switch = True
return G_CONTEXT.thread_pool.submit(self._loop_operation)
def stop_operation_backend(self):
self._loop_switch = False
return result

View file

@ -11,14 +11,16 @@ class BaseMemoryService(metaclass=ABCMeta):
def __init__(self,
memory_operations: Dict[str, dict],
read_memory_key: str = "read_memory",
read_message_key: str = "read_message",
**kwargs):
self.memory_operations: Dict[str, dict] = memory_operations
self.read_memory_key: str = read_memory_key
self.read_message_key: str = read_message_key
self._operation_dict: Dict[str, BaseOperation] = {}
self._op_description_dict: Dict[str, str] = {}
self.chat_messages: List[Message] = []
self.message_lock = threading.Lock
self.message_lock = threading.Lock()
self.logger = Logger.get_logger()
self.kwargs = kwargs
@ -48,5 +50,9 @@ class BaseMemoryService(metaclass=ABCMeta):
assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!"
return self.do_operation(self.read_memory_key)
def read_message(self):
assert self.read_message_key in self._operation_dict, f"op={self.read_message_key} is not inited!"
return self.do_operation(self.read_message_key)
def stop_service(self):
pass

View file

@ -13,6 +13,7 @@ from memory_scope.utils.prompt_handler import PromptHandler
class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
FILE_PATH: str = __file__
def __init__(self,
embedding_model: str = "",
@ -119,19 +120,19 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
@property
def user_name(self) -> str:
if self._user_name is None:
self._user_name = G_CONTEXT.meta_data["human_name"]
self._user_name = G_CONTEXT.meta_data["assistant_name"]
return self._user_name
@property
def target_name(self) -> str:
if self._target_name is None:
self._target_name = G_CONTEXT.meta_data["assistant_name"]
self._target_name = G_CONTEXT.meta_data["human_name"]
return self._target_name
@property
def prompt_handler(self) -> PromptHandler:
if self._prompt_handler is None:
self._prompt_handler = PromptHandler(__file__, **self.kwargs)
self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs)
return self._prompt_handler
def __getattr__(self, key: str):

View file

@ -10,6 +10,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class ExtractTimeWorker(MemoryBaseWorker):
EXTRACT_TIME_PATTERN = r'-\s*(\S+)(\d+)'
FILE_PATH: str = __file__
def _run(self):
query, query_timestamp = self.get_context(QUERY_WITH_TS)

View file

@ -14,9 +14,9 @@ class PrintMemoryWorker(MemoryBaseWorker):
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
expired_content_list: List[str] = []
obs_content_list: List[str] = []
insight_content_list: List[str] = []
expired_content_list: List[str] = ["----- expired -----"]
obs_content_list: List[str] = ["----- observation -----"]
insight_content_list: List[str] = ["----- insight -----"]
i = 0
j = 0
k = 0
@ -44,16 +44,12 @@ class PrintMemoryWorker(MemoryBaseWorker):
result: str = f"""
The memories of {self.user_name} about {self.target_name}.
----- observation -----
{obs_content}
----- observation -----
----- insight -----
{insight_content}
----- insight -----
----- expired -----
{expired_content}
----- expired -----
""".strip()
self.set_context(RESULT, result)

View file

@ -7,12 +7,14 @@ from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
class SetQueryWorker(MemoryBaseWorker):
def _run(self):
query = "_"
query_timestamp = int(datetime.datetime.now().timestamp())
if "query" in self.chat_kwargs:
""" cli test query
"""
# cli test query
query = self.chat_kwargs["query"]
query_timestamp = int(datetime.datetime.now().timestamp())
else:
elif self.chat_messages:
query = self.chat_messages[-1].content
query_timestamp = self.chat_messages[-1].time_created

View file

@ -12,6 +12,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class GetReflectionSubjectWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def new_insight_node(self, insight_key: str) -> MemoryNode:
dt_handler = DatetimeHandler()

View file

@ -11,6 +11,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class LongContraRepeatWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
async def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]):
filter_dict = {

View file

@ -11,6 +11,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class UpdateInsightWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def filter_obs_nodes(self,
insight_node: MemoryNode,

View file

@ -9,6 +9,8 @@ from memory_scope.utils.response_text_parser import ResponseTextParser
class ContraRepeatWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def _run(self):
all_obs_nodes: List[MemoryNode] = self.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
if not all_obs_nodes:

View file

@ -10,6 +10,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class GetObservationWithTimeWorker(GetObservationWorker):
FILE_PATH: str = __file__
def build_prompt(self) -> List[Message]:
# build prompt

View file

@ -13,6 +13,8 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class GetObservationWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str):
dt_handler = DatetimeHandler(dt=message.time_created)

View file

@ -9,6 +9,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg
class InfoFilterWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def _run(self):
# filter user msg
@ -22,6 +23,8 @@ class InfoFilterWorker(MemoryBaseWorker):
msg.content = msg.content[: half_size] + msg.content[-half_size:]
info_messages.append(msg)
self.logger.warning(info_messages)
if not info_messages:
self.logger.warning("info_messages is empty!")
self.continue_run = False
@ -31,6 +34,7 @@ class InfoFilterWorker(MemoryBaseWorker):
user_query_list = []
for i, msg in enumerate(info_messages):
user_query_list.append(f"{i + 1} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
self.logger.warning(self.prompt_handler.prompt_dict)
system_prompt = self.prompt_handler.info_filter_system.format(batch_size=len(info_messages),
user_name=self.target_name)
few_shot = self.prompt_handler.info_filter_few_shot.format(user_name=self.target_name)

View file

@ -86,7 +86,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
self.set_memories(TODAY_NODES, nodes)
async def _run(self):
def _run(self):
mock_query = "-"
self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query)
self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query)

View file

@ -17,12 +17,15 @@ class LlamaIndexGenerationModel(BaseModel):
def before_call(self, **kwargs):
prompt: str = kwargs.pop("prompt", "")
messages: List[Message] = kwargs.pop("messages", [])
messages: List[Message] | List[dict] = kwargs.pop("messages", [])
if prompt:
self.data = {"prompt": prompt}
elif messages:
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
if isinstance(messages[0], dict):
self.data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]}
else:
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
else:
raise RuntimeError("prompt and messages is both empty!")

View file

@ -28,7 +28,8 @@ class MemoryNode(BaseModel):
memory_type: str = Field("", description="conversation/observation/insight...")
status: str = Field("active", description="active or expired")
status: str = Field("active", description="db status: active / expired; modification_status: "
"new / content_modified / modified / active / expired")
vector: List[float] = Field([], description="content embedding result, return empty")

View file

@ -15,6 +15,7 @@ class PromptHandler(object):
self.kwargs = kwargs
file_path = self._class_path.strip(".py")
self.add_prompt_file(file_path)
if prompt_file:
@ -25,6 +26,7 @@ class PromptHandler(object):
@staticmethod
def file_path_completion(file_path: str) -> str:
if file_path.endswith(".yaml") or file_path.endswith(".json"):
return file_path
@ -56,7 +58,7 @@ class PromptHandler(object):
prompts = language_dict.get(G_CONTEXT.language)
if not prompts:
raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!")
self._prompt_dict[key] = prompts
self._prompt_dict[key] = prompts.strip()
@property
def prompt_dict(self):