[dev] add memory base worker & add role name to messages

This commit is contained in:
jinli.yl 2024-06-28 16:08:42 +08:00
parent b6a3d23ffe
commit 5f1bcf50d8
8 changed files with 102 additions and 296 deletions

View file

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

View 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"]]

View file

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

View file

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

View file

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

View file

@ -1,6 +1,6 @@
import re
from utils.logger import Logger
from memory_scope.utils.logger import Logger
class ResponseTextParser(object):

View file

@ -1,6 +1,6 @@
import time
from .logger import Logger
from memory_scope.utils.logger import Logger
class Timer(object):

View file

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