mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-08 22:21:15 +00:00
[dev] add memory base worker & add role name to messages
This commit is contained in:
parent
b6a3d23ffe
commit
5f1bcf50d8
8 changed files with 102 additions and 296 deletions
|
|
@ -1,4 +1,3 @@
|
|||
import datetime
|
||||
import os
|
||||
import time
|
||||
from typing import List
|
||||
|
|
@ -31,6 +30,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
human_name: str = "human",
|
||||
assistant_name: str = "assistant",
|
||||
**kwargs):
|
||||
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.stream: bool = stream
|
||||
|
|
@ -58,31 +58,35 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
@staticmethod
|
||||
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
|
||||
system_prompt = SYSTEM_PROMPT[G_CONTEXT.language]
|
||||
if related_memories:
|
||||
def get_system_prompt(self) -> Message:
|
||||
system_prompt = SYSTEM_PROMPT[G_CONTEXT.language].strip()
|
||||
|
||||
memories: str = self.memory_service.read_memory()
|
||||
if memories:
|
||||
memory_prompt = MEMORY_PROMPT[G_CONTEXT.language]
|
||||
all_prompt_list = [system_prompt, memory_prompt]
|
||||
all_prompt_list.extend(related_memories)
|
||||
system_prompt = "\n".join([x.strip() for x in all_prompt_list])
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created)
|
||||
system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]])
|
||||
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)
|
||||
|
||||
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
time_created = int(datetime.datetime.now().timestamp())
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value, content=query, time_created=time_created)
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value,
|
||||
role_name=self.human_name,
|
||||
content=query)
|
||||
|
||||
self.memory_service.add_messages(new_message)
|
||||
related_memories: List[str] = self.memory_service.read_memory()
|
||||
system_message: Message = self.get_system_prompt(related_memories, time_created)
|
||||
system_message: Message = self.get_system_prompt()
|
||||
|
||||
model_response = self.generation_model.call(messages=[system_message, new_message], stream=self.stream)
|
||||
if self.stream:
|
||||
for _ in model_response:
|
||||
_.message.role_name = self.assistant_name
|
||||
yield _
|
||||
else:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
return model_response
|
||||
|
||||
def process_commands(self, query: str) -> bool:
|
||||
|
|
|
|||
77
memory_scope/memory/worker/memory_base_worker.py
Normal file
77
memory_scope/memory/worker/memory_base_worker.py
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
from abc import ABCMeta
|
||||
from typing import List
|
||||
|
||||
from memory_scope.chat.global_context import G_CONTEXT
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
||||
|
||||
def __init__(self,
|
||||
embedding_model: str = "",
|
||||
generation_model: str = "",
|
||||
rank_model: str = "",
|
||||
**kwargs):
|
||||
super(MemoryBaseWorker, self).__init__(**kwargs)
|
||||
|
||||
self._embedding_model: BaseModel | str = embedding_model
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self._rank_model: BaseModel | str = rank_model
|
||||
|
||||
self._vector_store: BaseVectorStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
@property
|
||||
def messages(self) -> List[Message]:
|
||||
return self.get_context(MESSAGES)
|
||||
|
||||
@messages.setter
|
||||
def messages(self, value):
|
||||
self.set_context(MESSAGES, value)
|
||||
|
||||
@property
|
||||
def chat_name(self):
|
||||
return self.get_context(CHAT_NAME)
|
||||
|
||||
@property
|
||||
def embedding_model(self) -> BaseModel:
|
||||
if isinstance(self._embedding_model, str):
|
||||
self._embedding_model = G_CONTEXT.model_dict[self._embedding_model]
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
if isinstance(self._generation_model, str):
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
def rank_model(self) -> BaseModel:
|
||||
if isinstance(self._rank_model, str):
|
||||
self._rank_model = G_CONTEXT.model_dict[self._rank_model]
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
def vector_store(self) -> BaseVectorStore:
|
||||
if self._vector_store is None:
|
||||
self._vector_store = G_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
|
||||
@property
|
||||
def monitor(self):
|
||||
if self._monitor is None:
|
||||
self._monitor = G_CONTEXT.monitor
|
||||
return self._monitor
|
||||
|
||||
@property
|
||||
def memory_id(self) -> str:
|
||||
pass
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
return self.kwargs[key]
|
||||
|
||||
def get_prompt(self, x):
|
||||
return x[GLOBAL_CONTEXT.global_configs["language"]]
|
||||
|
|
@ -32,9 +32,7 @@ class LlamaIndexGenerationModel(BaseModel):
|
|||
stream: bool = False,
|
||||
**kwargs) -> ModelResponse | ModelResponseGen:
|
||||
|
||||
model_response.message = Message(role=MessageRoleEnum.ASSISTANT,
|
||||
content="",
|
||||
time_created=int(datetime.datetime.now().timestamp()))
|
||||
model_response.message = Message(role=MessageRoleEnum.ASSISTANT, content="")
|
||||
|
||||
call_result = model_response.raw
|
||||
if stream:
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import datetime
|
||||
from typing import Dict
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
|
@ -6,9 +7,13 @@ from pydantic import Field, BaseModel
|
|||
class Message(BaseModel):
|
||||
role: str = Field(..., description="The role of the message sender (user, assistant, system)")
|
||||
|
||||
role_name: str = Field("", description="role name")
|
||||
|
||||
content: str = Field(..., description="The body of the message")
|
||||
|
||||
time_created: int = Field(int(datetime.datetime.now().timestamp()),
|
||||
description="Timestamp when the message was created")
|
||||
|
||||
memorized: bool = Field(False, description="indicate whether message is memorized")
|
||||
|
||||
meta_data: Dict[str, str] = Field({}, description="meta data for msg")
|
||||
|
|
|
|||
|
|
@ -1,185 +0,0 @@
|
|||
import re
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import as_completed
|
||||
from itertools import zip_longest
|
||||
from typing import Dict, Any, List
|
||||
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from scheme.message import Message
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class Pipeline(object):
|
||||
def __init__(self,
|
||||
chat_name: str,
|
||||
memory_method_type: MemoryMethodEnum,
|
||||
pipeline_str: str,
|
||||
history_msg_count: int = 3,
|
||||
loop_interval_time: int = 300,
|
||||
loop_minimum_count: int = 20):
|
||||
|
||||
self.chat_name: str = chat_name
|
||||
self.memory_method_type: MemoryMethodEnum = memory_method_type
|
||||
self.pipeline_str: str = pipeline_str
|
||||
self.history_msg_count: int = history_msg_count
|
||||
self.loop_interval_time: int = loop_interval_time
|
||||
self.loop_minimum_count: int = loop_minimum_count
|
||||
|
||||
# pipeline上下文和锁
|
||||
self.context: Dict[str, Any] = {}
|
||||
self.context_lock = threading.Lock()
|
||||
|
||||
# pipeline run config
|
||||
self.loop_switch: bool = False
|
||||
self.pipeline_list: list[list] = []
|
||||
self.worker_set: set[str] = set()
|
||||
self.worker_dict: Dict[str, BaseWorker] = {}
|
||||
self.injected: bool = False
|
||||
|
||||
# message list
|
||||
self.history_message_list: List[Message] = []
|
||||
self.current_message_list: List[Message] = []
|
||||
self.message_lock = threading.Lock()
|
||||
|
||||
# 日志
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
self._parse_pipeline()
|
||||
|
||||
def _parse_pipeline(self):
|
||||
if not self.pipeline_str:
|
||||
return
|
||||
|
||||
# re-match e.g., [a|b],c,[d,e,f|g,h],j
|
||||
pattern = r'(\[[^\]]*\]|[^,]+)'
|
||||
pipeline_split = re.findall(pattern, self.pipeline_str)
|
||||
|
||||
self.pipeline_list = []
|
||||
for pipeline_part in pipeline_split:
|
||||
# e.g., [d,e,f|g,h]
|
||||
pipeline_part = pipeline_part.strip()
|
||||
if '[' in pipeline_part or ']' in pipeline_part:
|
||||
pipeline_part = pipeline_part.replace('[', '').replace(']', '')
|
||||
|
||||
# e.g., ["d,e,f", "g,h"]
|
||||
line_split = [x.strip() for x in pipeline_part.split("|") if x]
|
||||
if len(line_split) <= 0:
|
||||
continue
|
||||
|
||||
# e.g., ["d","e","f"]
|
||||
line_split_split = []
|
||||
for sub_line_split in line_split:
|
||||
sub_split = [x.strip() for x in sub_line_split.split(",")]
|
||||
line_split_split.append(sub_split)
|
||||
# add to workers
|
||||
self.worker_set.update(sub_split)
|
||||
self.pipeline_list.append(line_split_split)
|
||||
|
||||
def _visit_and_inject_workers(self):
|
||||
if self.injected:
|
||||
return
|
||||
|
||||
self.worker_dict = GLOBAL_CONTEXT.worker_dict[self.chat_name]
|
||||
|
||||
self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline Begin -----")
|
||||
i: int = 0
|
||||
for pipeline_part in self.pipeline_list:
|
||||
if len(pipeline_part) == 1:
|
||||
for w in pipeline_part[0]:
|
||||
self.logger.info(f"stage{i}: {w}")
|
||||
i += 1
|
||||
if w not in self.worker_dict:
|
||||
raise RuntimeError(f"worker={w} is not inited.")
|
||||
# 注入context
|
||||
self.worker_dict[w].set_context_dict(self.context)
|
||||
else:
|
||||
for w_zip in zip_longest(*pipeline_part, fillvalue="-"):
|
||||
self.logger.info(f"stage{i}: {' | '.join(w_zip)}")
|
||||
i += 1
|
||||
for w in w_zip:
|
||||
if w == "-":
|
||||
continue
|
||||
if w not in self.worker_dict:
|
||||
raise RuntimeError(f"worker={w} is not inited.")
|
||||
|
||||
# 注入context & lock
|
||||
self.worker_dict[w].set_context_dict(self.context, self.context_lock)
|
||||
|
||||
self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline End -----")
|
||||
self.injected = True
|
||||
|
||||
def _worker_run(self, worker_list: list[str]) -> bool:
|
||||
for worker_name in worker_list:
|
||||
worker = self.worker_dict[worker_name]
|
||||
worker.run()
|
||||
if not worker.continue_run:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _run(self):
|
||||
self._visit_and_inject_workers()
|
||||
|
||||
with Timer(f"pipeline_{self.chat_name}_{self.memory_method_type.value}"):
|
||||
self.context[MESSAGES] = self.history_message_list + self.current_message_list
|
||||
self.context[CHAT_NAME] = self.chat_name
|
||||
|
||||
for pipeline_part in self.pipeline_list:
|
||||
if len(pipeline_part) == 1:
|
||||
if not self._worker_run(pipeline_part[0]):
|
||||
break
|
||||
else:
|
||||
t_list = []
|
||||
for worker_list in pipeline_part:
|
||||
t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self._worker_run, worker_list))
|
||||
|
||||
flag = True
|
||||
for future in as_completed(t_list):
|
||||
if not future.result():
|
||||
flag = False
|
||||
break
|
||||
if not flag:
|
||||
break
|
||||
|
||||
def _thread_loop(self):
|
||||
while self.loop_switch:
|
||||
time.sleep(self.loop_interval_time)
|
||||
if len(self.current_message_list) < self.loop_minimum_count:
|
||||
continue
|
||||
self._run()
|
||||
self.context.clear()
|
||||
self.history_message_list = self.history_message_list.extend(self.current_message_list)[
|
||||
-self.history_msg_count:]
|
||||
with self.message_lock:
|
||||
self.current_message_list.clear()
|
||||
|
||||
def start_loop_run(self):
|
||||
if not self.loop_switch:
|
||||
self.loop_switch = True
|
||||
return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop)
|
||||
|
||||
def run(self, result_key: str = None):
|
||||
self._run()
|
||||
|
||||
# 获取result
|
||||
result = None
|
||||
if result_key:
|
||||
result = self.context.get(result_key)
|
||||
self.context.clear()
|
||||
|
||||
# 清理 msg
|
||||
self.history_message_list = self.history_message_list.extend(self.current_message_list)[
|
||||
-self.history_msg_count:]
|
||||
self.current_message_list.clear()
|
||||
return result
|
||||
|
||||
def submit_message(self, message: Message, with_lock=True):
|
||||
if with_lock:
|
||||
with self.message_lock:
|
||||
self.current_message_list.append(message)
|
||||
else:
|
||||
self.current_message_list.append(message)
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
import re
|
||||
|
||||
from utils.logger import Logger
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import time
|
||||
|
||||
from .logger import Logger
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class Timer(object):
|
||||
|
|
|
|||
|
|
@ -1,93 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from ..chat.global_context import GLOBAL_CONTEXT
|
||||
from ..constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from ..models.base_model import BaseModel
|
||||
from ..scheme.message import Message
|
||||
from ..storage.base_monitor import BaseMonitor
|
||||
from ..storage.base_vector_store import BaseVectorStore
|
||||
from ..worker.base_worker import BaseWorker
|
||||
from ..scheme.memory_node import MemoryNode
|
||||
from ..constants import common_constants
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
def __init__(
|
||||
self, embedding_model: str, generation_model: str, rank_model: str, **kwargs
|
||||
):
|
||||
super(MemoryBaseWorker, self).__init__(**kwargs)
|
||||
self.embedding_model_name: str = embedding_model
|
||||
self.generation_model_name: str = generation_model
|
||||
self.rank_model_name: str = rank_model
|
||||
|
||||
self._embedding_model: BaseModel | None = None
|
||||
self._generation_model: BaseModel | None = None
|
||||
self._rank_model: BaseModel | None = None
|
||||
|
||||
self._vector_store: BaseVectorStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
@property
|
||||
def messages(self) -> List[Message]:
|
||||
return self.get_context(MESSAGES)
|
||||
|
||||
@messages.setter
|
||||
def messages(self, value):
|
||||
self.set_context(MESSAGES, value)
|
||||
|
||||
@property
|
||||
def chat_name(self):
|
||||
return self.get_context(CHAT_NAME)
|
||||
|
||||
@property
|
||||
def embedding_model(self) -> BaseModel:
|
||||
if self._embedding_model is None:
|
||||
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(
|
||||
self.embedding_model_name
|
||||
)
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict.get(
|
||||
self.generation_model_name
|
||||
)
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
def rank_model(self) -> BaseModel:
|
||||
if self._rank_model is None:
|
||||
self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name)
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
def vector_store(self) -> BaseVectorStore:
|
||||
if self._vector_store is None:
|
||||
self._vector_store = GLOBAL_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
|
||||
@property
|
||||
def monitor(self):
|
||||
if self._monitor is None:
|
||||
self._monitor = GLOBAL_CONTEXT.monitor
|
||||
return self._monitor
|
||||
|
||||
@property
|
||||
def user_profile_dict(self) -> Dict[str, MemoryNode]:
|
||||
if not self._user_profile_dict:
|
||||
self._user_profile_dict = {
|
||||
user_attr.meta_data.get("memory_key", ""): user_attr
|
||||
for user_attr in self.get_context(common_constants.USER_PROFILE)
|
||||
}
|
||||
return self._user_profile_dict
|
||||
|
||||
@property
|
||||
def memory_id(self) -> str:
|
||||
pass
|
||||
|
||||
def __getattr__(self, key):
|
||||
return self.kwargs[key]
|
||||
|
||||
def get_prompt(self, x):
|
||||
return x[GLOBAL_CONTEXT.global_configs["language"]]
|
||||
Loading…
Add table
Reference in a new issue