[dev] ad memory type & add prompt_cn

This commit is contained in:
jinli.yl 2024-06-20 15:20:11 +08:00
parent f4e0e6bbc6
commit b2a3c47c36
16 changed files with 334 additions and 265 deletions

View file

@ -1,43 +1,16 @@
from abc import ABCMeta, abstractmethod
from memory_scope.constants.common_constants import RELATED_MEMORIES
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.handler.pipeline_handler import PipelineHandler
from memory_scope.chat.memory_service import MemoryService
class BaseMemoryChat(metaclass=ABCMeta):
def __init__(self,
user_name: str,
retrieve_pipeline: str,
summary_short_pipeline: str,
summary_long_pipeline: str,
**kwargs):
self.user_name: str = user_name
self.retrieve_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.RETRIEVE,
pipeline_str=retrieve_pipeline)
self.summary_short_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
pipeline_str=summary_short_pipeline)
self.summary_long_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
pipeline_str=summary_long_pipeline)
def retrieve(self):
self.retrieve_pipeline_handler.run()
return self.retrieve_pipeline_handler.get_context(RELATED_MEMORIES, [])
def summary_short(self):
self.summary_short_pipeline_handler.run()
def summary_long(self):
self.summary_long_pipeline_handler.run()
def __init__(self, **kwargs):
self.memory_service = MemoryService(**kwargs)
@abstractmethod
def chat(self):
def chat_with_memory(self, query: str):
"""
:param query:
:return:
"""

View file

@ -1,29 +1,51 @@
from memory_scope.chat.memory_service import MemoryService
from memory_scope.handler.init_handler import InitHandler
import datetime
from typing import List
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.handler.global_context import GLOBAL_CONTEXT
from memory_scope.models.base_model import BaseModel
from memory_scope.node.message import Message
from memory_scope.prompts.prompt_cn import SYSTEM_PROMPT, MEMORY_PROMPT
class MemoryChat(object):
class MemoryChat(BaseMemoryChat):
"""
TODO add agent
"""
def __init__(self, init_handler: InitHandler):
self.init_handler: InitHandler = init_handler
def __init__(self,
generation_model: str,
history_msg_count: int,
**kwargs):
super().__init__(**kwargs)
self.model: BaseModel = GLOBAL_CONTEXT.model_dict[generation_model]
self.history_msg_count: int = history_msg_count
self.memory_service: MemoryService = MemoryService(
retrieve_pipeline=init_handler.retrieve_pipeline,
summary_short_pipeline=init_handler.retrieve_pipeline,
summary_long_pipeline=init_handler.retrieve_pipeline,
)
self.memory_service.start_summary_short_backend()
self.memory_service.start_summary_long_backend()
def memory_retrieve(self):
pass
self.history_message_list: List[Message] = []
def memory_summary_short(self):
pass
@staticmethod
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
system_prompt = SYSTEM_PROMPT
if related_memories:
system_prompt = "\n".join([SYSTEM_PROMPT, MEMORY_PROMPT] + related_memories)
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt.strip(), time_created=time_created)
def memory_summary_long(self):
pass
def chat_with_memory(self, query: str):
query = query.strip()
if not query:
return
time_created = int(datetime.datetime.now().timestamp())
new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created)
related_memories: List[str] = self.memory_service.retrieve(message=new_message)
system_message = self.get_system_prompt(related_memories, time_created)
self.history_message_list.append(new_message)
self.history_message_list = self.history_message_list[-self.history_msg_count:]
all_messages = [system_message] + self.history_message_list
return self.model.call(messages=all_messages)
def chat(self):
pass
def chat_with_memory(self):
pass
def chat_with_memory_stream(self):
raise NotImplementedError

View file

@ -0,0 +1,55 @@
from memory_scope.constants.common_constants import RELATED_MEMORIES
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.handler.pipeline_handler import PipelineHandler
from memory_scope.node.message import Message
class MemoryService(object):
def __init__(self,
user_name: str,
retrieve_pipeline: str,
retrieve_all_pipeline: str,
summary_short_pipeline: str,
summary_long_pipeline: str,
summary_short_interval_time: int = 60,
summary_short_minimum_count: int = 5,
summary_long_interval_time: int = 60 * 5,
summary_long_minimum_count: int = 5 * 5,
**kwargs):
self.retrieve_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.RETRIEVE,
pipeline_str=retrieve_pipeline)
self.retrieve_all_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.RETRIEVE_ALL,
pipeline_str=retrieve_all_pipeline)
self.summary_short_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
pipeline_str=summary_short_pipeline,
loop_interval_time=summary_short_interval_time,
loop_minimum_count=summary_short_minimum_count)
self.summary_long_pipeline_handler = PipelineHandler(user_name=user_name,
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
pipeline_str=summary_long_pipeline,
loop_interval_time=summary_long_interval_time,
loop_minimum_count=summary_long_minimum_count)
self.kwargs = kwargs
def retrieve(self, message: Message):
self.retrieve_pipeline_handler.submit_message(message, with_lock=False)
self.summary_short_pipeline_handler.submit_message(message)
self.summary_long_pipeline_handler.submit_message(message)
return self.retrieve_pipeline_handler.run(RELATED_MEMORIES)
def retrieve_all(self):
return self.retrieve_all_pipeline_handler.run(RELATED_MEMORIES)
def start_summary_short_backend(self):
self.summary_short_pipeline_handler.start_loop_run()
def start_summary_long_backend(self):
self.summary_long_pipeline_handler.start_loop_run()

View file

@ -1,11 +1,96 @@
import json
import os
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any
import fire
from memory_scope.job import Job
from handler.global_context import GLOBAL_CONTEXT
from memory_scope.enumeration.model_type import ModelType
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
from memory_scope.utils.tool_functions import complete_config_name, init_instance_by_config_v2
class CliJob(object):
def __init__(self, config_path: str):
self.config_path: str = config_path
self.config_base_dir: str = os.path.dirname(config_path)
self.config: Dict[str, Any] = {}
self.logger: Logger = Logger.get_logger("memory_chat")
def init_memory_chat(self):
for chat in self.config["chat_list"]:
memory_chat_config = self.config[chat]
memory_chat = init_instance_by_config_v2(memory_chat_config)
GLOBAL_CONTEXT.memory_chat_dict[chat] = memory_chat
generation_model = memory_chat_config[ModelType.GENERATION_MODEL.value]
self.init_model(generation_model)
def init_model(self, model_name: str):
if not model_name or model_name in GLOBAL_CONTEXT.model_dict:
return
with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f:
model_config = json.load(f)
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config_v2(model_config)
def init_workers(self):
""" load worker config & init workers
"""
worker_config_name: str = self.config["workers"]
with open(os.path.join(self.config_base_dir, complete_config_name(worker_config_name))) as f:
worker_config_dict = json.load(f)
for worker_name, worker_config in worker_config_dict.items():
if worker_name in GLOBAL_CONTEXT.worker_dict:
raise RuntimeError(f"worker_name={worker_name} is repeated!")
GLOBAL_CONTEXT.worker_dict[worker_name] = init_instance_by_config_v2(worker_config,
suffix_name="worker",
**GLOBAL_CONTEXT.global_configs)
self.init_model(worker_config.get(ModelType.EMBEDDING_MODEL.value))
self.init_model(worker_config.get(ModelType.GENERATION_MODEL.value))
self.init_model(worker_config.get(ModelType.RANK_MODEL.value))
def set_global_config(self):
"""set global_configs & set apikey into env
"""
def init_global_content_by_config(self):
with open(complete_config_name(self.config_path)) as f:
self.config = json.load(f)
GLOBAL_CONTEXT.global_configs = self.config["global_configs"]
self.set_global_config()
GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"]))
self.init_workers()
GLOBAL_CONTEXT.db_client = init_instance_by_config_v2(self.config["db"])
GLOBAL_CONTEXT.monitor = init_instance_by_config_v2(self.config["monitor"])
self.init_memory_chat()
def run(self):
with GLOBAL_CONTEXT.thread_pool, Timer("job", log_time=False) as t:
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
while True:
query = input("wait for input:")
if query in ["stop", "停止"]:
break
memory_chat.chat_with_memory(query=query)
self.logger.info(f"chat complete. cost={t.cost_str}")
def main(config_path: str):
job = Job(config_path=config_path)
job.init_instance_by_config()
job = CliJob(config_path=config_path)
job.init_global_content_by_config()
job.run()

View file

@ -1,34 +0,0 @@
import json
from pipeline.memory_service import MemoryService
from utils.tool_functions import init_instance_by_config
class Wrapper:
"""Wrapper class for anything that needs to set up during init"""
def __init__(self):
self._provider = None
def register(self, provider):
self._provider = provider
def __getattr__(self, key: str):
if self.__dict__.get("_provider", None) is None:
raise AttributeError("Please run __init__ first!")
return getattr(self._provider, key)
C = Wrapper()
def init(config_path: str):
config = json.loads(config_path)
C.register(config)
## register workers
C.worker = json.loads(C.worker)
## register services
for k,v in C.pipeline.items():
C.pipeline[k] = MemoryService(k)

View file

@ -9,6 +9,8 @@ WORKER = "worker"
MEMORY = "memory"
USER_NAME = "user_name"
DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
RELATED_MEMORIES = "related_memories"

View file

@ -6,6 +6,8 @@ class MemoryMethodEnum(str, Enum):
RETRIEVE = "retrieve"
RETRIEVE_ALL = "retrieve_all"
SUMMARY_SHORT = "summary_short"
SUMMARY_LONG = "summary_long"

View file

@ -0,0 +1,9 @@
from enum import Enum
class ModelType(str, Enum):
GENERATION_MODEL = "generation_model"
EMBEDDING_MODEL = "embedding_model"
RANK_MODEL = "rank_model"

View file

@ -1,21 +1,33 @@
import re
import threading
import time
from concurrent.futures import as_completed
from itertools import zip_longest
from typing import Dict, Any
from typing import Dict, Any, List
from memory_scope.constants.common_constants import MESSAGES, USER_NAME
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.handler.global_context import GLOBAL_CONTEXT
from memory_scope.node.message import Message
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
class PipelineHandler(object):
def __init__(self, user_name: str, memory_method_type: MemoryMethodEnum, pipeline_str: str):
def __init__(self,
user_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.user_name: str = user_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
# 日志
self.logger: Logger = Logger.get_logger()
@ -24,11 +36,19 @@ class PipelineHandler(object):
self.context: Dict[str, Any] = {}
self.context_lock = threading.Lock()
# pipeline run config
self.loop_switch: bool = False
# 解析和打印 pipeline
self.pipeline_list: list[list] = []
self._parse_pipeline()
self._print_pipeline()
# message list
self.history_message_list: List[Message] = []
self.current_message_list: List[Message] = []
self.message_lock = threading.Lock()
def _parse_pipeline(self):
# re-match e.g., [a|b],c,[d,e,f|g,h],j
pattern = r'(\[[^\]]*\]|[^,]+)'
@ -70,14 +90,8 @@ class PipelineHandler(object):
GLOBAL_CONTEXT.worker_dict[w].context = self.context
self.logger.info(f"----- {self.user_name} {self.memory_method_type.value} Pipeline End -----")
def get_context(self, key: str, default=None):
return self.context.get(key, default)
def clear_context(self):
self.context.clear()
@staticmethod
def worker_run(worker_list: list[str]) -> bool:
def _worker_run(worker_list: list[str]) -> bool:
for worker_name in worker_list:
worker = GLOBAL_CONTEXT.worker_dict[worker_name]
worker.run()
@ -85,16 +99,19 @@ class PipelineHandler(object):
return False
return True
def run(self):
def _run(self, result_key: str = None):
with Timer(f"pipeline_{self.user_name}_{self.memory_method_type.value}"):
self.context[MESSAGES] = self.history_message_list + self.current_message_list
self.context[USER_NAME] = self.user_name
for pipeline_part in self.pipeline_list:
if len(pipeline_part) == 1:
if not self.worker_run(pipeline_part[0]):
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))
t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self._worker_run, worker_list))
flag = True
for future in as_completed(t_list):
@ -103,3 +120,47 @@ class PipelineHandler(object):
break
if not flag:
break
if result_key:
return self.context.get(result_key)
self.context.clear()
return None
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,78 +0,0 @@
import json
import os
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any
from handler.global_context import GLOBAL_CONTEXT
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
from memory_scope.utils.tool_functions import complete_config_name, init_instance_by_config_v2
class Job(object):
def __init__(self, config_path: str):
self.config_path: str = config_path
self.config_base_dir: str = os.path.dirname(config_path)
self.config: Dict[str, Any] = {}
self.logger: Logger = Logger.get_logger("memory_chat")
def init_memory_chat(self):
for chat in self.config["chat_list"]:
memory_chat_config = self.config[chat]
memory_chat = init_instance_by_config_v2(memory_chat_config)
GLOBAL_CONTEXT.memory_chat_dict[chat] = memory_chat
generation_model = memory_chat_config["generation_model"]
self.init_model(generation_model)
def init_model(self, model_name: str):
if not model_name or model_name in GLOBAL_CONTEXT.model_dict:
return
with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f:
model_config = json.load(f)
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config_v2(model_config)
def init_workers(self):
""" load worker config & init workers
"""
worker_config_name: str = self.config["workers"]
with open(os.path.join(self.config_base_dir, complete_config_name(worker_config_name))) as f:
worker_config_dict = json.load(f)
for worker_name, worker_config in worker_config_dict.items():
if worker_name in GLOBAL_CONTEXT.worker_dict:
raise RuntimeError(f"worker_name={worker_name} is repeated!")
GLOBAL_CONTEXT.worker_dict[worker_name] = init_instance_by_config_v2(worker_config,
suffix_name="worker",
**GLOBAL_CONTEXT.global_configs)
self.init_model(worker_config.get("embedding_model"))
self.init_model(worker_config.get("generation_model"))
self.init_model(worker_config.get("rank_model"))
def set_global_config(self):
"""set global_configs & set apikey into env
"""
def init_instance_by_config(self):
with open(complete_config_name(self.config_path)) as f:
self.config = json.load(f)
GLOBAL_CONTEXT.global_configs = self.config["global_configs"]
self.set_global_config()
GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"]))
self.init_workers()
GLOBAL_CONTEXT.db_client = init_instance_by_config_v2(self.config["db"])
GLOBAL_CONTEXT.monitor = init_instance_by_config_v2(self.config["monitor"])
self.init_memory_chat()
def run(self):
with GLOBAL_CONTEXT.thread_pool, Timer("job"):
pass

View file

@ -6,6 +6,4 @@ class Message(BaseModel):
content: str = Field(..., description="The body of the message")
time_created: str = Field("", description="Timestamp when the message was created")
info_score: str = Field("", description="2 > 1 > 0")
time_created: int = Field("", description="Timestamp when the message was created")

View file

@ -0,0 +1,7 @@
SYSTEM_PROMPT = """
"""
MEMORY_PROMPT = """
"""

View file

@ -5,6 +5,8 @@ from typing import Dict, List
from constants.common_constants import WEEKDAYS
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
def under_line_to_hump(underline_str):
sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str)
@ -181,3 +183,16 @@ def complete_config_name(config_name: str, suffix: str = ".json"):
if not config_name.endswith(suffix):
config_name += suffix
return config_name
def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
return [
{
"role": MessageRoleEnum.SYSTEM.value,
"content": system_prompt.strip(),
},
{
"role": MessageRoleEnum.USER.value,
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
},
]

View file

@ -1,101 +1,53 @@
from typing import List, Dict, Optional
from typing import List
from memory_scope.constants.common_constants import MESSAGES, USER_NAME
from memory_scope.handler.global_context import GLOBAL_CONTEXT
from memory_scope.models.base_model import BaseModel
from memory_scope.node.message import Message
from memory_scope.worker.base_worker import BaseWorker
from constants import common_constants
from constants.common_constants import CONFIG, MESSAGES, PROMPT_CONFIG
from enumeration.message_role_enum import MessageRoleEnum
from node.message import Message
from node.user_attribute import UserAttribute
from pipeline.memory import MemoryServiceRequestModel
from worker.base_worker import BaseWorker
from cli.cli_config import C
from utils.tool_functions import init_instance_by_config
class MemoryBaseWorker(BaseWorker):
def __init__(self, **kwargs):
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
@property
def request(self) -> MemoryServiceRequestModel:
return self.get_context(common_constants.REQUEST)
self._embedding_model: BaseModel | None = None
self._generation_model: BaseModel | None = None
self._rank_model: BaseModel | None = None
@property
def messages(self) -> List[Message]:
messages: List[Message] = self.context_handler.get_context(MESSAGES)
if messages is None:
messages = self.request.messages
self.context_handler.set_context(MESSAGES, messages)
return messages
return self.context[MESSAGES]
@messages.setter
def messages(self, value):
self.context_handler.set_context(MESSAGES, value)
def flush(self, context_handler):
super(MemoryBaseWorker, self).flush(context_handler)
self._user_profile_dict: Dict[str, UserAttribute] = {}
self.context[MESSAGES] = value
@property
def user_profile_dict(self) -> Dict[str, UserAttribute]:
if not self._user_profile_dict:
self._user_profile_dict = {user_attr.memory_key: user_attr for user_attr in self.request.user.user_profile}
return self._user_profile_dict
def embedding_model(self):
if self._embedding_model is None:
GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name)
return self._embedding_model
@property
def request_ext_info(self):
return self.request.user.ext_info
# if not self._request_ext_info:
# self._request_ext_info = self.request.ext_info
# return self._request_ext_info
def generation_model(self):
if self._generation_model is None:
GLOBAL_CONTEXT.model_dict.get(self.generation_model_name)
return self._generation_model
@property
def prompt_config(self) -> BailianPromptConfig:
return self.request.user.prompt
def rank_model(self):
if self._rank_model is None:
GLOBAL_CONTEXT.model_dict.get(self.rank_model_name)
return self._rank_model
@property
def client(self, model_type: str, model_name:str):
models = C.get(model_type)
models["model_name"] = init_instance_by_config(
config = models.get(model_name),
try_kwargs={
"is_multi_thread": is_multi_thread,
"thread_pool": self.thread_pool
}
)
return models["model_name"]
@property
def emb_client(self, model_name: str):
self.client("model_embedding", model_name)
@property
def gene_client(self):
self.client("model_generate", model_name)
@property
def rerank_client(self):
self.client("model_rerank", model_name)
@property
def es_client(self):
self.client("db", model_name)
@property
def tenant_id(self):
return self.request.user.tenant_id
@property
def memory_id(self):
return self.request.user.memory_id
@staticmethod
def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
return [
{
"role": MessageRoleEnum.SYSTEM.value,
"content": system_prompt.strip(),
},
{
"role": MessageRoleEnum.USER.value,
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
},
]
def user_name(self):
return self.context[USER_NAME]