mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
[dev] ad memory type & add prompt_cn
This commit is contained in:
parent
f4e0e6bbc6
commit
b2a3c47c36
16 changed files with 334 additions and 265 deletions
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
55
memory_scope/chat/memory_service.py
Normal file
55
memory_scope/chat/memory_service.py
Normal 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()
|
||||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
@ -9,6 +9,8 @@ WORKER = "worker"
|
|||
|
||||
MEMORY = "memory"
|
||||
|
||||
USER_NAME = "user_name"
|
||||
|
||||
DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
|
||||
|
||||
RELATED_MEMORIES = "related_memories"
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ class MemoryMethodEnum(str, Enum):
|
|||
|
||||
RETRIEVE = "retrieve"
|
||||
|
||||
RETRIEVE_ALL = "retrieve_all"
|
||||
|
||||
SUMMARY_SHORT = "summary_short"
|
||||
|
||||
SUMMARY_LONG = "summary_long"
|
||||
|
|
|
|||
9
memory_scope/enumeration/model_type.py
Normal file
9
memory_scope/enumeration/model_type.py
Normal 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"
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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")
|
||||
7
memory_scope/prompts/prompt_cn.py
Normal file
7
memory_scope/prompts/prompt_cn.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
SYSTEM_PROMPT = """
|
||||
|
||||
"""
|
||||
|
||||
MEMORY_PROMPT = """
|
||||
|
||||
"""
|
||||
|
|
@ -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]])
|
||||
},
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue