mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-08 22:21:15 +00:00
pipeline config
This commit is contained in:
parent
cf94bbf927
commit
818d4d794c
32 changed files with 404 additions and 496 deletions
|
|
@ -0,0 +1,11 @@
|
|||
{
|
||||
"thread_pool_max_count": "",
|
||||
"worker": "config/worker.json",
|
||||
"prompt": "config/prompt.json",
|
||||
"pipeline": {
|
||||
"summary_short": "",
|
||||
"summary_long": "",
|
||||
"retrieve": ""
|
||||
},
|
||||
"models": "config/model/dash.json"
|
||||
}
|
||||
|
|
@ -0,0 +1,6 @@
|
|||
{
|
||||
"embedding": "",
|
||||
"generate": "",
|
||||
"rerank": "",
|
||||
"es": ""
|
||||
}
|
||||
|
|
@ -0,0 +1,27 @@
|
|||
{
|
||||
"info_filter_system": "",
|
||||
"info_filter_user_query": "",
|
||||
"get_observation_system": "",
|
||||
"get_observation_user_query": "",
|
||||
"get_observation_with_time_system": "",
|
||||
"get_observation_with_time_few_shot": "",
|
||||
"get_observation_with_time_user_query": "",
|
||||
"contra_repeat_system": "",
|
||||
"contra_repeat_few_shot": "",
|
||||
"contra_repeat_user_query": "",
|
||||
"get_reflect_system": "",
|
||||
"get_reflect_few_shot": "",
|
||||
"get_reflect_user_query": "",
|
||||
"get_insight_system": "",
|
||||
"get_insight_few_shot": "",
|
||||
"get_insight_user_query": "",
|
||||
"update_plural_profile_system": "",
|
||||
"update_plural_profile_few_shot": "",
|
||||
"update_plural_profile_user_query": "",
|
||||
"update_unique_profile_system": "",
|
||||
"update_unique_profile_few_shot": "",
|
||||
"update_unique_profile_user_query": "",
|
||||
"update_insight_system": "",
|
||||
"update_insight_few_shot": "",
|
||||
"update_insight_user_query": ""
|
||||
}
|
||||
|
|
@ -1,74 +1,74 @@
|
|||
{
|
||||
"EsInsightWorker": {
|
||||
"module_name": "EsInsightWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_insight_top_k": 128
|
||||
}
|
||||
},
|
||||
"EsKeywordWorker": {
|
||||
"module_name": "EsKeywordWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_keyword_top_k": 10
|
||||
}
|
||||
},
|
||||
"EsNewObsWorker": {
|
||||
"module_name": "EsNewObsWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_new_obs_top_k": 256
|
||||
}
|
||||
},
|
||||
"EsNotReflectedWorker": {
|
||||
"module_name": "EsNotReflectedWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_not_reflected_top_k": 256
|
||||
}
|
||||
},
|
||||
"EsSimilarWorker": {
|
||||
"module_name": "EsSimilarWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_similar_top_k": 128
|
||||
}
|
||||
},
|
||||
"EsTodayObsWorker": {
|
||||
"module_name": "EsTodayObsWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_today_obs_top_k": 128
|
||||
}
|
||||
},
|
||||
"GetInsightWorker": {
|
||||
"module_name": "GetInsightWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_insight_similar_top_k": 128,
|
||||
"insight_obs_max_cnt": 10,
|
||||
"get_insight_model": "qwen_max",
|
||||
"get_insight_max_token": 500,
|
||||
"get_insight_temperature": 0.6,
|
||||
"get_insight_top_k": 1,
|
||||
"get_insight_top_k": 1
|
||||
}
|
||||
},
|
||||
"ExtractTimeWorker": {
|
||||
"module_name": "ExtractTimeWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"parse_time_model": "qwen_1_8_parse_time_service",
|
||||
"parse_time_max_token": 100,
|
||||
"parse_time_temperature": 0.6,
|
||||
"parse_time_top_k": 1,
|
||||
"parse_time_top_k": 1
|
||||
}
|
||||
},
|
||||
"InfoFilterWorker": {
|
||||
"module_name": "InfoFilterWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"info_filter_msg_max_size": 200,
|
||||
"info_filter_model": qwen_max,
|
||||
"info_filter_model": "qwen_max",
|
||||
"info_filter_max_token": 200,
|
||||
"info_filter_temperature": 0.6,
|
||||
"info_filter_top_k": 1
|
||||
|
|
@ -76,69 +76,69 @@
|
|||
},
|
||||
"GetObservationWithTimeWorker": {
|
||||
"module_name": "GetObservationWithTimeWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"summary_messages_model": qwen_max,
|
||||
"summary_messages_model": "qwen_max",
|
||||
"summary_messages_max_token": 500,
|
||||
"summary_messages_temperature": 0.6,
|
||||
"summary_messages_top_k": 1,
|
||||
"summary_messages_top_k": 1
|
||||
}
|
||||
},
|
||||
"GetObservationWorker": {
|
||||
"module_name": "GetObservationWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"summary_messages_model": qwen_max,
|
||||
"summary_messages_model": "qwen_max",
|
||||
"summary_messages_max_token": 500,
|
||||
"summary_messages_temperature": 0.6,
|
||||
"summary_messages_top_k": 1,
|
||||
"summary_messages_top_k": 1
|
||||
}
|
||||
},
|
||||
"ContraRepeatWorker": {
|
||||
"module_name": "ContraRepeatWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"merge_obs_model": qwen_max,
|
||||
"merge_obs_model": "qwen_max",
|
||||
"merge_obs_max_token": 500,
|
||||
"merge_obs_temperature": 0.6,
|
||||
"merge_obs_top_k": 1,
|
||||
"merge_obs_top_k": 1
|
||||
}
|
||||
},
|
||||
"FuseRerankWorker": {
|
||||
"module_name": "FuseRerankWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"fuse_score_threshold": 0.1,
|
||||
"fuse_ratio_dict": {
|
||||
MemoryTypeEnum.CONVERSATION.value: 0.8,
|
||||
MemoryTypeEnum.OBSERVATION.value: 1.0,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value: 1.0,
|
||||
MemoryTypeEnum.INSIGHT.value: 1.5,
|
||||
MemoryTypeEnum.PROFILE.value: 1.5,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value: 1.5,
|
||||
"conversation": 0.8,
|
||||
"observation": 1.0,
|
||||
"obs_customized": 1.0,
|
||||
"insight": 1.5,
|
||||
"profile": 1.5,
|
||||
"profile_customized": 1.5
|
||||
},
|
||||
"fuse_time_ratio": 2.0,
|
||||
"fuse_time_ratio": 2.0
|
||||
}
|
||||
},
|
||||
"UpdateProfileWorker": {
|
||||
"module_name": "UpdateProfileWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"update_profile_threshold": 0.1,
|
||||
"update_profile_model": "qwen_max",
|
||||
"update_profile_max_token": 500,
|
||||
"update_profile_temperature": 0.6,
|
||||
"update_profile_top_k": 1,
|
||||
"update_profile_max_thread": 10,
|
||||
"update_profile_max_thread": 10
|
||||
}
|
||||
},
|
||||
"GetReflectionWorker": {
|
||||
"module_name": "GetReflectionWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"reflect_obs_cnt_threshold": 40,
|
||||
"reflect_num_questions": 3,
|
||||
"reflect_obs_model": qwen_max,
|
||||
"reflect_obs_model": "qwen_max",
|
||||
"reflect_obs_max_token": 300,
|
||||
"reflect_obs_temperature": 0.6,
|
||||
"reflect_obs_top_k": 1
|
||||
|
|
@ -146,14 +146,14 @@
|
|||
},
|
||||
"UpdateInsightWorker": {
|
||||
"module_name": "UpdateInsightWorker",
|
||||
"module_path": MEMORY_WORKER_PATH,
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"update_insight_threshold": 0.1,
|
||||
"update_insight_model": "qwen_max",
|
||||
"update_insight_max_token": 500,
|
||||
"update_insight_temperature": 0.6,
|
||||
"update_insight_top_k": 1,
|
||||
"update_insight_max_thread": 10,
|
||||
"update_insight_max_thread": 10
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
|
@ -6,25 +6,10 @@ from model.message import Message
|
|||
from model.user_attribute import UserAttribute
|
||||
from request.base_model import RequestBaseModel
|
||||
|
||||
|
||||
class MemoryServiceRequestModel(RequestBaseModel):
|
||||
user: UserConfig = None
|
||||
|
||||
messages: List[Message] = Field(...,
|
||||
description="summary: 多轮对话的list,默认按照时间正序,最后一条是最新的; retrieve: 最后一条是query")
|
||||
|
||||
messages_pick_n: int = Field(1, description="summary:传需要总结的msg的个数;retrieve:不传")
|
||||
|
||||
memory_id: str = Field(..., description="memory id")
|
||||
|
||||
workspace_id: str = Field("", description="workspace id")
|
||||
|
||||
api_key: str = Field("", description="api id")
|
||||
|
||||
scene: str = Field("", description="需要枚举来源: TONGYI_MAIN_CHAT, TONGYI_CHAR_CHAT, BAILIAN, ASSISTANT_API")
|
||||
|
||||
algo_version: str = Field("", description="算法版本,只在做AB实验时透传")
|
||||
|
||||
output_max_count: int = Field(3, description="retrieve时最多的条数,约定3-10条")
|
||||
|
||||
user_profile: List[UserAttribute] = Field([], description="user_profile")
|
||||
|
||||
ext_info: Dict[str, str] = Field({}, description="extra information")
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ from importlib import import_module
|
|||
from itertools import zip_longest
|
||||
from typing import Dict, Any
|
||||
|
||||
from worker.memory.base_worker import BaseWorker
|
||||
|
||||
from common.context_handler import ContextHandler
|
||||
from common.logger import Logger
|
||||
from common.timer import timer, Timer
|
||||
|
|
@ -14,31 +16,17 @@ from constants import common_constants
|
|||
from constants.common_constants import RESPONSE_EXT_INFO, MAX_WORKERS, PIPELINE
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.bailian.base_worker import BaseWorker
|
||||
from config.env_config import C, Workers
|
||||
from utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class MemoryServiceBailian(object):
|
||||
THREAD_POOL_MAX_COUNT: int = 5
|
||||
|
||||
MEMORY_WORKER_PATH: str = "worker.bailian"
|
||||
|
||||
def __init__(self, request: MemoryServiceRequestModel, method: MemoryMethodEnum):
|
||||
# 全局上下文,worker之间交换参数和变量
|
||||
self.context_handler = ContextHandler(scene=request.scene,
|
||||
method=method.value,
|
||||
algo_version=request.algo_version)
|
||||
self.context_handler.set_context(common_constants.REQUEST, request)
|
||||
class MemoryService(object):
|
||||
def __init__(self, method: MemoryMethodEnum):
|
||||
self.method = method
|
||||
self.context_handler = ContextHandler()
|
||||
|
||||
# 线程池
|
||||
max_workers = self.context_handler.get_env_config(MAX_WORKERS)
|
||||
if max_workers:
|
||||
self.max_workers = int(max_workers)
|
||||
else:
|
||||
self.max_workers = self.THREAD_POOL_MAX_COUNT
|
||||
self.thread_pool = ThreadPoolExecutor(max_workers=self.max_workers)
|
||||
|
||||
# 运行信息
|
||||
self.run_infos = []
|
||||
self.thread_pool = ThreadPoolExecutor(max_workers=C.THREAD_POOL_MAX_COUNT)
|
||||
|
||||
# 全部初始化的worker
|
||||
self.worker_dict: Dict[str, BaseWorker] = {}
|
||||
|
|
@ -46,23 +34,18 @@ class MemoryServiceBailian(object):
|
|||
# 日志
|
||||
self.logger: Logger = Logger.get_memory_logger()
|
||||
|
||||
# 初始化pipeline
|
||||
self.pipeline_list = self.get_pipeline()
|
||||
self.print_and_init_worker(self.pipeline_list)
|
||||
|
||||
def get_worker(self, worker_name: str, is_multi_thread: bool = False) -> BaseWorker:
|
||||
# 更新worker name
|
||||
worker_name_split = worker_name.split(".")
|
||||
worker_name = worker_name_split[-1]
|
||||
if common_constants.WORKER not in worker_name:
|
||||
worker_name = f"{worker_name}_{common_constants.WORKER}"
|
||||
|
||||
# 构造path
|
||||
worker_paths = [self.MEMORY_WORKER_PATH]
|
||||
worker_paths.extend(worker_name_split[:-1])
|
||||
worker_paths.append(worker_name)
|
||||
module = import_module(".".join(worker_paths))
|
||||
|
||||
worker_clazz_name = under_line_to_hump(worker_name)
|
||||
return getattr(module, worker_clazz_name)(context_handler=self.context_handler,
|
||||
is_multi_thread=is_multi_thread,
|
||||
thread_pool=self.thread_pool)
|
||||
return init_instance_by_config(
|
||||
config = W.get(worker_name),
|
||||
try_kwargs={
|
||||
"is_multi_thread": is_multi_thread,
|
||||
"thread_pool": self.thread_pool
|
||||
}
|
||||
)
|
||||
|
||||
def worker_run(self, worker_list: list[str]) -> bool:
|
||||
for worker_name in worker_list:
|
||||
|
|
@ -99,9 +82,19 @@ class MemoryServiceBailian(object):
|
|||
def get_context(self, key: str, default=None) -> Any:
|
||||
return self.context_handler.get_context(key, default)
|
||||
|
||||
def flush(self, request: MemoryServiceRequestModel):
|
||||
# 全局上下文,worker之间交换参数和变量
|
||||
self.context_handler.flush()
|
||||
|
||||
# 运行信息
|
||||
self.run_infos = []
|
||||
self.context_handler.set_context(common_constants.REQUEST, request)
|
||||
for pipeline_part in self.pipeline_list:
|
||||
pipeline_part.flush(self.context_handler)
|
||||
|
||||
@timer
|
||||
def get_pipeline(self) -> list[list]:
|
||||
pipeline_str = self.context_handler.get_env_config(PIPELINE)
|
||||
pipeline_str = C.pipeline.get(self.method)
|
||||
self.logger.info(f"pipeline={pipeline_str}")
|
||||
|
||||
# re-match e.g., [a|b],c,[d,e,f|g,h],j
|
||||
|
|
@ -126,12 +119,9 @@ class MemoryServiceBailian(object):
|
|||
return pipeline_list
|
||||
|
||||
def run(self):
|
||||
pipeline_list = self.get_pipeline()
|
||||
self.print_and_init_worker(pipeline_list)
|
||||
|
||||
# run workers in multi threads
|
||||
with self.thread_pool, Timer("ALL_PIPELINE"):
|
||||
for pipeline_part in pipeline_list:
|
||||
for pipeline_part in self.pipeline_list:
|
||||
if len(pipeline_part) == 1:
|
||||
if not self.worker_run(pipeline_part[0]):
|
||||
break
|
||||
|
|
|
|||
|
|
@ -3,26 +3,9 @@ import threading
|
|||
from typing import Dict, Any
|
||||
|
||||
from common.logger import Logger
|
||||
from constants.common_constants import APP_ENV
|
||||
from enumeration.env_type import EnvType
|
||||
|
||||
|
||||
class ContextHandler(object):
|
||||
# 环境类型:日常 预发 开发
|
||||
_env_type: EnvType | None = None
|
||||
|
||||
# 所有环境变量
|
||||
_env_params = os.environ
|
||||
|
||||
# 环境变量解析 -> config
|
||||
_env_configs: Dict[str, str] | None = None
|
||||
|
||||
def __init__(self, scene: str, method: str, prefix: str = "memory", algo_version: str = ""):
|
||||
self.scene: str = scene
|
||||
self.method: str = method
|
||||
self.prefix: str = prefix
|
||||
self.algo_version: str = algo_version
|
||||
|
||||
def __init__(self):
|
||||
# 上下文 所有worker共享
|
||||
self.context_dict: Dict[str, Any] = {}
|
||||
|
||||
|
|
@ -32,89 +15,8 @@ class ContextHandler(object):
|
|||
# 全局锁
|
||||
self.context_lock = threading.Lock()
|
||||
|
||||
# algo_version: 上游参数 > 环境变量参数
|
||||
self._update_algo_version()
|
||||
|
||||
def _update_algo_version(self):
|
||||
# 上游传参优先
|
||||
if self.algo_version:
|
||||
return
|
||||
|
||||
p_key = f"{self.prefix}_{self.method}_algo_version"
|
||||
p_key_with_scene = f"{p_key}_{self.scene}"
|
||||
|
||||
if p_key_with_scene in self._env_params:
|
||||
self.algo_version = self._env_params.get(p_key_with_scene)
|
||||
|
||||
if p_key in self._env_params:
|
||||
self.algo_version = self._env_params.get(p_key_with_scene)
|
||||
|
||||
@property
|
||||
def env_type(self):
|
||||
if self._env_type is None:
|
||||
env_type = self._env_params.get(APP_ENV)
|
||||
assert env_type, f"env_type={env_type} is empty!"
|
||||
self._env_type = EnvType(env_type.lower())
|
||||
|
||||
return self._env_type
|
||||
|
||||
@property
|
||||
def env_configs(self):
|
||||
if self._env_configs is not None:
|
||||
return self._env_configs
|
||||
|
||||
prefix: str = f"{self.prefix}_{self.method}_"
|
||||
all_ket_set = set()
|
||||
for k, v in self._env_params.items():
|
||||
if not k.startswith(prefix):
|
||||
continue
|
||||
|
||||
if not v:
|
||||
continue
|
||||
|
||||
# {key}_{scene}_{algo_version}
|
||||
raw_k = k.removeprefix(prefix)
|
||||
if self.scene in raw_k:
|
||||
raw_k_split = [x.strip("_") for x in raw_k.split(self.scene) if x.strip("_")]
|
||||
if len(raw_k_split) == 1:
|
||||
raw_k = raw_k_split[0]
|
||||
elif len(raw_k_split) == 2:
|
||||
algo_version = raw_k_split[1]
|
||||
if algo_version != self.algo_version:
|
||||
continue
|
||||
|
||||
raw_k = raw_k_split[0]
|
||||
else:
|
||||
self.logger.info(f"_update_env_configs encounter error! k={k} v={v}")
|
||||
continue
|
||||
|
||||
all_ket_set.add(raw_k)
|
||||
|
||||
self._env_configs = {}
|
||||
for k in all_ket_set:
|
||||
v = self.get_param(k)
|
||||
if v:
|
||||
self._env_configs[k] = v
|
||||
self.logger.info(f"update env_configs={self._env_configs}")
|
||||
return self._env_configs
|
||||
|
||||
def get_env_config(self, key: str, default=None) -> str:
|
||||
return self.env_configs.get(key, default)
|
||||
|
||||
def get_param(self, key: str, default=None):
|
||||
p_key = f"{self.prefix}_{self.method}_{key}"
|
||||
p_key_with_scene = f"{p_key}_{self.scene}"
|
||||
# memory_summary_{key}_tongyi_algo_v1 > memory_summary_{key}_tongyi > memory_summary_{key}
|
||||
|
||||
if p_key_with_scene in self._env_params:
|
||||
p_key = p_key_with_scene
|
||||
|
||||
if self.algo_version:
|
||||
p_key_with_version = f"{p_key_with_scene}_{self.algo_version}"
|
||||
if p_key_with_version in self._env_params:
|
||||
p_key = p_key_with_version
|
||||
|
||||
return self._env_params.get(p_key, default)
|
||||
def flush(self):
|
||||
self.context_dict: Dict[str, Any] = {}
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
# 多线程环境下,如果是指针下修改,不安全
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from typing import Dict, List
|
|||
|
||||
from constants.common_constants import WEEKDAYS
|
||||
from enumeration.env_type import EnvType
|
||||
from importlib import import_module
|
||||
|
||||
global_env_type = None
|
||||
|
||||
|
|
@ -162,3 +163,12 @@ def time_to_formatted_str(time: datetime | str | int | float = None,
|
|||
return_str = string_format.format(**get_datetime_info_dict(current_dt))
|
||||
|
||||
return return_str
|
||||
|
||||
|
||||
def init_instance_by_config(config, default_module_path=None, try_kwargs={}):
|
||||
import_module(config.get("module_path", default_module_path))
|
||||
clazz = getattr(module, config.get("module_name"))
|
||||
try:
|
||||
return clazz(**config.get("kwargs"),**try_kwargs)
|
||||
except:
|
||||
return clazz(**config.get("kwargs"))
|
||||
|
|
@ -10,18 +10,14 @@ from common.timer import Timer
|
|||
class BaseWorker(object):
|
||||
|
||||
def __init__(self,
|
||||
context_handler: ContextHandler,
|
||||
is_multi_thread: bool = False,
|
||||
thread_pool: ThreadPoolExecutor = None,
|
||||
raise_exception: bool = True,
|
||||
logger: Logger = None,
|
||||
**kwargs):
|
||||
super(BaseWorker, self).__init__(**kwargs)
|
||||
|
||||
# 原始参数
|
||||
self.context_handler: ContextHandler = context_handler
|
||||
self.is_multi_thread: bool = is_multi_thread
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
self.raise_exception: bool = raise_exception
|
||||
self.logger: Logger = logger
|
||||
|
||||
|
|
@ -36,15 +32,20 @@ class BaseWorker(object):
|
|||
# True 为正常运行,False会结束整个pipeline
|
||||
self.continue_run: bool = True
|
||||
|
||||
# 短name
|
||||
self._name_simple: str = ""
|
||||
|
||||
def flush(self, context_handler: ContextHandler, thread_pool: ThreadPoolExecutor):
|
||||
# 原始参数
|
||||
self.context_handler = context_handler
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
|
||||
# 运行信息,保存到ext_info
|
||||
self.run_infos: List[str] = []
|
||||
|
||||
# 运行时间
|
||||
self.run_cost: float = 0
|
||||
|
||||
# 短name
|
||||
self._name_simple: str = ""
|
||||
|
||||
def _run(self):
|
||||
pass
|
||||
|
||||
|
|
@ -82,25 +83,6 @@ class BaseWorker(object):
|
|||
def set_context(self, key: str, value: Any):
|
||||
self.context_handler.set_context(key, value, self.is_multi_thread)
|
||||
|
||||
def get_param(self, key: str, default=None):
|
||||
return self.context_handler.env_configs.get(key, default)
|
||||
|
||||
@property
|
||||
def env_type(self) -> str:
|
||||
return self.context_handler.env_type.value
|
||||
|
||||
@property
|
||||
def scene(self) -> str:
|
||||
return self.context_handler.scene
|
||||
|
||||
@property
|
||||
def method(self) -> str:
|
||||
return self.context_handler.method
|
||||
|
||||
@property
|
||||
def algo_version(self) -> str:
|
||||
return self.context_handler.algo_version
|
||||
|
||||
@property
|
||||
def name_simple(self) -> str:
|
||||
if not self._name_simple:
|
||||
|
|
|
|||
|
|
@ -3,16 +3,19 @@ from typing import List
|
|||
from constants.common_constants import INSIGHT_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsInsightWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_insight_top_k, *args, **kwargs):
|
||||
super(EsInsightWorker, self).__init__(*args, **kwargs)
|
||||
self.es_insight_top_k = es_insight_top_k
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=self.config.es_insight_top_k,
|
||||
hits = self.es_client.exact_search_v2(size=self.es_insight_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.INSIGHT.value,
|
||||
|
|
|
|||
|
|
@ -3,16 +3,19 @@ from typing import List
|
|||
from constants.common_constants import NEW, NEW_OBS_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNewObsWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_new_obs_top_k, *args, **kwargs):
|
||||
super(EsNewObsWorker, self).__init__(*args, **kwargs)
|
||||
self.es_new_obs_top_k = es_new_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=self.config.es_new_obs_top_k,
|
||||
hits = self.es_client.exact_search_v2(size=self.es_new_obs_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
|
|
|
|||
|
|
@ -3,15 +3,19 @@ from typing import List
|
|||
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNotReflectedWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_not_reflected_top_k, *args, **kwargs):
|
||||
super(EsNotReflectedWorker, self).__init__(*args, **kwargs)
|
||||
self.es_new_obs_top_k = es_new_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=self.config.es_not_reflected_top_k,
|
||||
hits = self.es_client.exact_search_v2(size=self.es_not_reflected_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ from typing import List
|
|||
|
||||
from constants.common_constants import ALL_NODES, ALL_MEMORIES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsRetrieveAllWorker(MemoryBaseWorker):
|
||||
|
|
@ -12,7 +12,7 @@ class EsRetrieveAllWorker(MemoryBaseWorker):
|
|||
# msg_time_created = self.messages[-1].time_created
|
||||
hits = self.es_client.exact_search_v2(size=1000,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
# "memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
|
|
|
|||
|
|
@ -4,77 +4,35 @@ from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
|
|||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsSimilarWorker(MemoryBaseWorker):
|
||||
|
||||
def es_similar_obs(self) -> List[MemoryWrapNode]:
|
||||
query = self.messages[-1].content
|
||||
hits = self.es_client.similar_search(text=query,
|
||||
size=self.config.es_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value})
|
||||
|
||||
# 初始化成MemoryWrapNode,并加入召回源的参数
|
||||
similar_obs_nodes: List[MemoryWrapNode] = []
|
||||
for hit in hits:
|
||||
node = MemoryWrapNode.init_from_es(hit)
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
similar_obs_nodes.append(node)
|
||||
return similar_obs_nodes
|
||||
|
||||
def es_similar_insight(self) -> List[MemoryWrapNode]:
|
||||
query = self.messages[-1].content
|
||||
hits = self.es_client.similar_search(text=query,
|
||||
size=self.config.es_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.INSIGHT.value})
|
||||
|
||||
# 初始化成MemoryWrapNode,并加入召回源的参数
|
||||
similar_obs_nodes: List[MemoryWrapNode] = []
|
||||
for hit in hits:
|
||||
node = MemoryWrapNode.init_from_es(hit)
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
similar_obs_nodes.append(node)
|
||||
return similar_obs_nodes
|
||||
|
||||
def es_similar_obs_custom(self) -> List[MemoryWrapNode]:
|
||||
query = self.messages[-1].content
|
||||
hits = self.es_client.similar_search(text=query,
|
||||
size=self.config.es_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBS_CUSTOMIZED.value})
|
||||
|
||||
# 初始化成MemoryWrapNode,并加入召回源的参数
|
||||
similar_obs_nodes: List[MemoryWrapNode] = []
|
||||
for hit in hits:
|
||||
node = MemoryWrapNode.init_from_es(hit)
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
similar_obs_nodes.append(node)
|
||||
return similar_obs_nodes
|
||||
def __init__(self, es_similar_top_k, *args, **kwargs):
|
||||
super(EsSimilarWorker, self).__init__(*args, **kwargs)
|
||||
self.es_similar_top_k = es_similar_top_k
|
||||
|
||||
def _run(self):
|
||||
for func in [self.es_similar_obs, self.es_similar_insight, self.es_similar_obs_custom]:
|
||||
self.submit_thread(func, sleep_time=0.01)
|
||||
query = self.messages[-1].content
|
||||
hits = self.es_client.similar_search(text=query,
|
||||
size=self.es_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
# 初始化成MemoryWrapNode,并加入召回源的参数
|
||||
similar_obs_nodes: List[MemoryWrapNode] = []
|
||||
for result in self.join_threads():
|
||||
similar_obs_nodes.extend(result)
|
||||
|
||||
for hit in hits:
|
||||
node = MemoryWrapNode.init_from_es(hit)
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
similar_obs_nodes.append(node)
|
||||
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
|
||||
for node in similar_obs_nodes:
|
||||
self.logger.info(f"node={node.memory_node.content} "
|
||||
f"score_similar={node.score_similar} "
|
||||
f"type={node.memory_node.memoryType}")
|
||||
self.logger.info(f"node={node.memory_node.content} score_similar={node.score_similar}")
|
||||
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)
|
||||
|
|
|
|||
|
|
@ -4,20 +4,23 @@ from common.tool_functions import time_to_formatted_str
|
|||
from constants.common_constants import TODAY_OBS_NODES, DT
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsTodayObsWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_today_obs_top_k, *args, **kwargs):
|
||||
super(EsTodayObsWorker, self).__init__(*args, **kwargs)
|
||||
self.es_today_obs_top_k = es_today_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
if not self.messages:
|
||||
self.logger.warning("messages is empty!")
|
||||
return
|
||||
msg_time_created = self.messages[-1].time_created
|
||||
hits = self.es_client.exact_search_v2(size=self.config.es_today_obs_top_k,
|
||||
hits = self.es_client.exact_search_v2(size=self.es_today_obs_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
|
|
|
|||
|
|
@ -4,10 +4,10 @@ from common.user_profile_handler import UserProfileHandler
|
|||
from constants import common_constants
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LoadProfileWorker(MemoryBaseWorker):
|
||||
|
|
@ -15,7 +15,7 @@ class LoadProfileWorker(MemoryBaseWorker):
|
|||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=10000,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.PROFILE.value,
|
||||
|
|
|
|||
|
|
@ -1,37 +1,18 @@
|
|||
from typing import List, Dict, Optional
|
||||
|
||||
from common.dash_embedding_client import DashEmbeddingClient
|
||||
from common.dash_generate_client import DashGenerateClient
|
||||
from common.dash_rerank_client import DashReRankClient
|
||||
from common.elastic_search_client import ElasticSearchClient
|
||||
from config.bailian_memory_config import BailianMemoryConfig
|
||||
from config.bailian_prompt_config import BailianPromptConfig
|
||||
from constants import common_constants
|
||||
from constants.common_constants import CONFIG, MESSAGES, PROMPT_CONFIG
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from model.message import Message
|
||||
from model.user_attribute import UserAttribute
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.bailian.base_worker import BaseWorker
|
||||
|
||||
from worker.memory.base_worker import BaseWorker
|
||||
from config.env_config import EnvConfig
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super(MemoryBaseWorker, self).__init__(**kwargs)
|
||||
|
||||
self._user_profile_dict: Dict[str, UserAttribute] = {}
|
||||
self._request_ext_info: Dict[str, str] = {}
|
||||
|
||||
self._config: Optional[BailianMemoryConfig] = None
|
||||
self._prompt_config: Optional[BailianPromptConfig] = None
|
||||
|
||||
self._dash_embedding_client: Optional[DashEmbeddingClient] = None
|
||||
self._dash_generate_client: Optional[DashGenerateClient] = None
|
||||
self._dash_rerank_client: Optional[DashReRankClient] = None
|
||||
|
||||
self._es_client: Optional[ElasticSearchClient] = None
|
||||
|
||||
@property
|
||||
def request(self) -> MemoryServiceRequestModel:
|
||||
return self.get_context(common_constants.REQUEST)
|
||||
|
|
@ -40,7 +21,7 @@ class MemoryBaseWorker(BaseWorker):
|
|||
def messages(self) -> List[Message]:
|
||||
messages: List[Message] = self.context_handler.get_context(MESSAGES)
|
||||
if messages is None:
|
||||
messages = self.request.messages[-self.config.messages_pick_n:]
|
||||
messages = self.request.messages[-self.request.messages_pick_n:]
|
||||
self.context_handler.set_context(MESSAGES, messages)
|
||||
return messages
|
||||
|
||||
|
|
@ -48,75 +29,50 @@ class MemoryBaseWorker(BaseWorker):
|
|||
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] = {}
|
||||
|
||||
@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_profile}
|
||||
self._user_profile_dict = {user_attr.memory_key: user_attr for user_attr in self.request.user.user_profile}
|
||||
return self._user_profile_dict
|
||||
|
||||
@property
|
||||
def request_ext_info(self):
|
||||
if not self._request_ext_info:
|
||||
self._request_ext_info = self.request.ext_info
|
||||
return self._request_ext_info
|
||||
|
||||
@property
|
||||
def config(self) -> BailianMemoryConfig:
|
||||
if self._config is None:
|
||||
self._config = self.get_context(CONFIG)
|
||||
return self._config
|
||||
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
|
||||
|
||||
@property
|
||||
def prompt_config(self) -> BailianPromptConfig:
|
||||
if self._prompt_config is None:
|
||||
self._prompt_config = self.get_context(PROMPT_CONFIG)
|
||||
if not self._prompt_config:
|
||||
self._prompt_config = BailianPromptConfig()
|
||||
self.set_context(PROMPT_CONFIG, self._prompt_config)
|
||||
return self._prompt_config
|
||||
return self.request.user.prompt
|
||||
|
||||
@property
|
||||
def emb_client(self):
|
||||
if self._dash_embedding_client is None:
|
||||
self._dash_embedding_client = DashEmbeddingClient(request_id=self.config.request_id,
|
||||
dash_scope_uid=self.config.uid,
|
||||
authorization=self.config.api_key,
|
||||
workspace=self.config.workspace_id,
|
||||
env_type=self.env_type,
|
||||
max_retry_count=self.config.dash_embedding_retry_cnt)
|
||||
return self._dash_embedding_client
|
||||
return C.emb_client
|
||||
|
||||
@property
|
||||
def gene_client(self):
|
||||
if self._dash_generate_client is None:
|
||||
self._dash_generate_client = DashGenerateClient(request_id=self.config.request_id,
|
||||
dash_scope_uid=self.config.uid,
|
||||
authorization=self.config.api_key,
|
||||
workspace=self.config.workspace_id,
|
||||
env_type=self.env_type,
|
||||
max_retry_count=self.config.dash_generate_retry_cnt)
|
||||
return self._dash_generate_client
|
||||
return C.gene_client
|
||||
|
||||
@property
|
||||
def rerank_client(self):
|
||||
if self._dash_rerank_client is None:
|
||||
self._dash_rerank_client = DashReRankClient(request_id=self.config.request_id,
|
||||
dash_scope_uid=self.config.uid,
|
||||
authorization=self.config.api_key,
|
||||
workspace=self.config.workspace_id,
|
||||
env_type=self.env_type,
|
||||
max_retry_count=self.config.dash_rerank_retry_cnt)
|
||||
return self._dash_rerank_client
|
||||
return C.rerank_client
|
||||
|
||||
@property
|
||||
def es_client(self):
|
||||
if self._es_client is None:
|
||||
self._es_client = ElasticSearchClient(es_user_name=self.config.es_user_name,
|
||||
es_password=self.config.es_password,
|
||||
es_index_name=self.config.es_index_name,
|
||||
embedding_client=self.emb_client,
|
||||
max_retries=self.config.es_retry_cnt)
|
||||
return self._es_client
|
||||
return C.es_client
|
||||
|
||||
@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):
|
||||
|
|
|
|||
|
|
@ -3,9 +3,9 @@ from typing import List
|
|||
from common.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
|
||||
from model.memory_node import MemoryNode
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class MemoryStoreWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import json
|
|||
from config.bailian_memory_config import BailianMemoryConfig
|
||||
from constants.common_constants import REQUEST, CONFIG
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.bailian.base_worker import BaseWorker
|
||||
from worker.memory.base_worker import BaseWorker
|
||||
|
||||
|
||||
class ParseParamsWorker(BaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,10 +3,16 @@ import re
|
|||
from common.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP
|
||||
from constants.common_constants import EXTRACT_TIME_DICT
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
def __init__(self, parse_time_model, parse_time_max_token, parse_time_temperature, parse_time_top_k, *args, **kwargs):
|
||||
super(ExtractTimeWorker, self).__init__(*args, **kwargs)
|
||||
self.parse_time_model = parse_time_model
|
||||
self.parse_time_max_token = parse_time_max_token
|
||||
self.parse_time_temperature = parse_time_temperature
|
||||
self.parse_time_top_k = parse_time_top_k
|
||||
|
||||
@staticmethod
|
||||
def get_parse_time_prompt(query: str, query_time_str: str):
|
||||
|
|
@ -46,10 +52,10 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
|
||||
# call sft model
|
||||
response_text = self.gene_client.call(prompt=extract_time_prompt,
|
||||
model_name=self.config.parse_time_model,
|
||||
max_token=self.config.parse_time_max_token,
|
||||
temperature=self.config.parse_time_temperature,
|
||||
top_k=self.config.parse_time_top_k)
|
||||
model_name=self.parse_time_model,
|
||||
max_token=self.parse_time_max_token,
|
||||
temperature=self.parse_time_temperature,
|
||||
top_k=self.parse_time_top_k)
|
||||
|
||||
# if empty, return
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -1,12 +1,22 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, DEFAULT_SYSTEM_PROMPT, \
|
||||
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
|
||||
TIME_MATCHED
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class FuseRerankWorker(MemoryBaseWorker):
|
||||
def __init__(self, fuse_time_ratio, fuse_score_threshold, fuse_ratio_dict, *args, **kwargs):
|
||||
super(FuseRerankWorker, self).__init__(*args, **kwargs)
|
||||
self.fuse_score_threshold = fuse_score_threshold
|
||||
self.fuse_ratio_dict = fuse_ratio_dict
|
||||
# self.default_system_prompt = default_system_prompt
|
||||
self.fuse_time_ratio = fuse_time_ratio
|
||||
|
||||
@property
|
||||
def output_max_count(self):
|
||||
return self.request.user.output_max_count
|
||||
|
||||
@staticmethod
|
||||
def format_time_infer(time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]):
|
||||
|
|
@ -53,11 +63,11 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
|
||||
filtered_nodes = []
|
||||
for node in all_online_nodes:
|
||||
if node.score_rank < self.config.fuse_score_threshold:
|
||||
if node.score_rank < self.fuse_score_threshold:
|
||||
continue
|
||||
|
||||
# 根据类型给ratio
|
||||
type_ratio: float = self.config.fuse_ratio_dict.get(node.memory_node.memoryType, 0.1)
|
||||
type_ratio: float = self.fuse_ratio_dict.get(node.memory_node.memoryType, 0.1)
|
||||
|
||||
# 时间系数,完全匹配才行
|
||||
fuse_time_ratio: float = 1.0
|
||||
|
|
@ -83,7 +93,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
break
|
||||
|
||||
if match_event_flag or match_msg_flag:
|
||||
fuse_time_ratio = self.config.fuse_time_ratio
|
||||
fuse_time_ratio = self.fuse_time_ratio
|
||||
node.memory_node.metaData[TIME_MATCHED] = "1"
|
||||
|
||||
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
|
||||
|
|
@ -93,7 +103,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
|
||||
# get output & save context
|
||||
filtered_nodes = sorted(filtered_nodes, key=lambda x: x.score_rerank, reverse=True)
|
||||
filtered_nodes = filtered_nodes[: self.config.output_max_count]
|
||||
filtered_nodes = filtered_nodes[: self.output_max_count]
|
||||
related_memories: List[str] = []
|
||||
for node in filtered_nodes:
|
||||
content = node.memory_node.content
|
||||
|
|
@ -112,4 +122,4 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
related_memories.append(content)
|
||||
|
||||
self.set_context(RELATED_MEMORIES, related_memories)
|
||||
self.set_context(DEFAULT_SYSTEM_PROMPT, self.config.default_system_prompt)
|
||||
# self.set_context(DEFAULT_SYSTEM_PROMPT, self.default_system_prompt)
|
||||
|
|
|
|||
|
|
@ -4,8 +4,8 @@ from common.user_profile_handler import UserProfileHandler
|
|||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
|
||||
QUERY_KEYWORDS
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SemanticRankWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -6,11 +6,19 @@ from constants.common_constants import NEW_INSIGHT_NODES, DT, NOT_REFLECTED_MERG
|
|||
INSIGHT_VALUE, REFLECTED
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetInsightWorker(MemoryBaseWorker):
|
||||
def __init__(self, insight_obs_max_cnt, es_insight_similar_top_k, get_insight_model, get_insight_max_token, get_insight_temperature, get_insight_top_k, **kwargs):
|
||||
super(GetInsightWorker,self).__init__(*args,**kwargs)
|
||||
self.insight_obs_max_cnt = insight_obs_max_cnt
|
||||
self.get_insight_model = get_insight_model
|
||||
self.get_insight_max_token = get_insight_max_token
|
||||
self.get_insight_temperature = get_insight_temperature
|
||||
self.get_insight_top_k = get_insight_top_k
|
||||
self.es_insight_similar_top_k = es_insight_similar_top_k
|
||||
|
||||
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryWrapNode:
|
||||
created_dt = datetime.now()
|
||||
|
|
@ -22,17 +30,17 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
INSIGHT_KEY: insight_key,
|
||||
INSIGHT_VALUE: insight_value,
|
||||
}
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
return MemoryWrapNode.init_from_attrs(content=content,
|
||||
memoryId=self.config.memory_id,
|
||||
memoryId=self.memory_id,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.INSIGHT.value,
|
||||
content_modified=True, # 新增的insight需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.config.tenant_id)
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def reflect_new_insight_key(self,
|
||||
insight_key: str,
|
||||
|
|
@ -40,9 +48,9 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# 检索历史memory
|
||||
hits = self.es_client.similar_search(text=insight_key,
|
||||
size=self.config.es_insight_similar_top_k,
|
||||
size=self.es_insight_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
|
|
@ -70,7 +78,7 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
score = rank_node["relevance_score"]
|
||||
related_nodes[index].score_rank = score
|
||||
related_nodes_sorted = sorted(related_nodes, key=lambda x: x.score_rank, reverse=True)[
|
||||
:self.config.insight_obs_max_cnt]
|
||||
:self.insight_obs_max_cnt]
|
||||
|
||||
# 生成prompt
|
||||
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
|
||||
|
|
@ -83,10 +91,10 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM, 提取insight
|
||||
response_text = self.gene_client.call(messages=get_insight_message,
|
||||
model_name=self.config.get_insight_model,
|
||||
max_token=self.config.get_insight_max_token,
|
||||
temperature=self.config.get_insight_temperature,
|
||||
top_k=self.config.get_insight_top_k)
|
||||
model_name=self.get_insight_model,
|
||||
max_token=self.get_insight_max_token,
|
||||
temperature=self.get_insight_temperature,
|
||||
top_k=self.get_insight_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -3,11 +3,19 @@ from typing import List
|
|||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, NOT_REFLECTED_OBS_NODES, REFLECTED, INSIGHT_NODES, INSIGHT_KEY, \
|
||||
NEW_INSIGHT_KEYS, NOT_REFLECTED_MERGE_NODES
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
def __init__(self, reflect_obs_cnt_threshold, reflect_num_questions, reflect_obs_model, reflect_obs_max_token, reflect_obs_temperature, reflect_obs_top_k, *args, **kwargs):
|
||||
super(GetReflectionWorker,self).__init__(*args, **kwargs)
|
||||
self.reflect_obs_cnt_threshold = reflect_obs_cnt_threshold
|
||||
self.reflect_num_questions = reflect_num_questions
|
||||
self.reflect_obs_model = reflect_obs_model
|
||||
self.reflect_obs_max_token = reflect_obs_max_token
|
||||
self.reflect_obs_temperature = reflect_obs_temperature
|
||||
self.reflect_obs_top_k = reflect_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
|
|
@ -23,7 +31,7 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
if not_reflected_count <= self.config.reflect_obs_cnt_threshold:
|
||||
if not_reflected_count <= self.reflect_obs_cnt_threshold:
|
||||
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop reflect.")
|
||||
return
|
||||
|
||||
|
|
@ -48,18 +56,18 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
user_query_list = [n.memory_node.content for n in not_reflected_merge_nodes]
|
||||
reflect_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_reflect_system.format(
|
||||
num_questions=self.config.reflect_num_questions),
|
||||
num_questions=self.reflect_num_questions),
|
||||
few_shot=self.prompt_config.get_reflect_few_shot,
|
||||
user_query=self.prompt_config.get_reflect_user_query.format(exist_keys=",".join(exist_keys),
|
||||
user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"reflect_message={reflect_message}")
|
||||
|
||||
# call LLM
|
||||
# # call LLM
|
||||
response_text = self.gene_client.call(messages=reflect_message,
|
||||
model_name=self.config.reflect_obs_model,
|
||||
max_token=self.config.reflect_obs_max_token,
|
||||
temperature=self.config.reflect_obs_temperature,
|
||||
top_k=self.config.reflect_obs_top_k)
|
||||
model_name=self.reflect_obs_model,
|
||||
max_token=self.reflect_obs_max_token,
|
||||
temperature=self.reflect_obs_temperature,
|
||||
top_k=self.reflect_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -10,6 +10,13 @@ from worker.bailian.memory_base_worker import MemoryBaseWorker
|
|||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
def __init__(es_contra_repeat_similar_top_k, long_contra_repeat_threshold, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
|
||||
super(LongContraRepeatWorker, self).__init__(*args, **kwargs)
|
||||
self.es_contra_repeat_similar_top_k = es_contra_repeat_similar_top_k
|
||||
self.merge_obs_model = merge_obs_model
|
||||
self.merge_obs_max_token = merge_obs_max_token
|
||||
self.merge_obs_temperature = merge_obs_temperature
|
||||
self.merge_obs_top_k = merge_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
|
|
@ -20,9 +27,9 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
for new_obs_node in new_obs_nodes:
|
||||
text = new_obs_node.memory_node.content
|
||||
hits = self.es_client.similar_search(text=text,
|
||||
size=self.config.es_contra_repeat_similar_top_k,
|
||||
size=self.es_contra_repeat_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
|
|
@ -33,7 +40,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
if related_node.score_similar < self.config.long_contra_repeat_threshold:
|
||||
if related_node.score_similar < self.long_contra_repeat_threshold:
|
||||
continue
|
||||
else:
|
||||
has_match = True
|
||||
|
|
@ -58,10 +65,10 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.config.merge_obs_model,
|
||||
max_token=self.config.merge_obs_max_token,
|
||||
temperature=self.config.merge_obs_temperature,
|
||||
top_k=self.config.merge_obs_top_k)
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ from typing import List, Dict
|
|||
|
||||
from constants.common_constants import NEW_INSIGHT_NODES, MODIFIED_MEMORIES, INSIGHT_NODES, NEW_OBS_NODES, \
|
||||
NOT_REFLECTED_OBS_NODES, NEW, NOT_REFLECTED_MERGE_NODES
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -1,14 +1,20 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from constants.common_constants import INSIGHT_NODES, NEW_OBS_NODES, INSIGHT_KEY, INSIGHT_VALUE, DT
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from constants.common_constants import INSIGHT_NODES, NEW_OBS_NODES, INSIGHT_KEY, INSIGHT_VALUE
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
def __init__(update_insight_threshold, update_insight_max_thread, update_insight_model, update_insight_max_token, update_insight_temperature, update_insight_top_k,*args, **kwargs):
|
||||
super(UpdateInsightWorker, self).__init__(*args, **kwargs)
|
||||
self.update_insight_threshold = update_insight_threshold
|
||||
self.update_insight_max_thread = update_insight_max_thread
|
||||
self.update_insight_model = update_insight_model
|
||||
self.update_insight_max_token = update_insight_max_token
|
||||
self.update_insight_temperature = update_insight_temperature
|
||||
self.update_insight_top_k = update_insight_top_k
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryWrapNode,
|
||||
|
|
@ -36,7 +42,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
score = rank_node["relevance_score"]
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.config.update_insight_threshold:
|
||||
if score >= self.update_insight_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
|
|
@ -48,23 +54,6 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight_node(self, insight_node: MemoryWrapNode, insight_key: str, insight_value: str):
|
||||
created_dt = datetime.now()
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
meta_data = {
|
||||
DT: dt,
|
||||
INSIGHT_KEY: insight_key,
|
||||
INSIGHT_VALUE: insight_value,
|
||||
}
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
insight_node.memory_node.content = content
|
||||
insight_node.memory_node.content_modified = True
|
||||
insight_node.memory_node.metaData = meta_data
|
||||
insight_node.memory_node.tenantId = self.config.tenant_id
|
||||
return insight_node
|
||||
|
||||
def update_insight(self,
|
||||
insight_node: MemoryWrapNode,
|
||||
filtered_nodes: List[MemoryWrapNode]) -> MemoryWrapNode:
|
||||
|
|
@ -89,10 +78,10 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=update_insight_message,
|
||||
model_name=self.config.update_insight_model,
|
||||
max_token=self.config.update_insight_max_token,
|
||||
temperature=self.config.update_insight_temperature,
|
||||
top_k=self.config.update_insight_top_k)
|
||||
model_name=self.update_insight_model,
|
||||
max_token=self.update_insight_max_token,
|
||||
temperature=self.update_insight_temperature,
|
||||
top_k=self.update_insight_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
@ -113,7 +102,9 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
return insight_node
|
||||
|
||||
return self.update_insight_node(insight_node, insight_key, insight_value)
|
||||
insight_node.memory_node.metaData[INSIGHT_VALUE] = insight_value
|
||||
insight_node.memory_node.content_modified = True
|
||||
return insight_node
|
||||
|
||||
def _run(self):
|
||||
# 获取新的obs和insight
|
||||
|
|
@ -141,8 +132,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.config.update_insight_max_thread:
|
||||
result_sorted = result_sorted[:self.config.update_insight_max_thread]
|
||||
if len(result_sorted) > self.update_insight_max_thread:
|
||||
result_sorted = result_sorted[:update_insight_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
|
|
|
|||
|
|
@ -3,12 +3,25 @@ from typing import List
|
|||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateProfileWorker(MemoryBaseWorker):
|
||||
def __init__(update_profile_max_thread, update_profile_threshold, extra_user_attrs, update_profile_model, update_profile_max_token, update_profile_temperature, update_profile_top_k, *args, **kwargs):
|
||||
super(UpdateProfileWorker,self).__init__(*args, **kwargs)
|
||||
self.update_profile_max_thread = update_profile_max_thread
|
||||
self.extra_user_attrs = extra_user_attrs
|
||||
self.update_profile_threshold = update_profile_threshold
|
||||
self.update_profile_model = update_profile_model
|
||||
self.update_profile_max_token = update_profile_max_token
|
||||
self.update_profile_temperature = update_profile_temperature
|
||||
self.update_profile_top_k = update_profile_top_k
|
||||
|
||||
# @property
|
||||
# def extra_user_attrs(self):
|
||||
# return self.request.extra_user_attrs
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
user_attr: UserAttribute,
|
||||
|
|
@ -29,7 +42,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
score = rank_node["relevance_score"]
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.config.update_profile_threshold:
|
||||
if score >= self.update_profile_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
|
|
@ -71,10 +84,10 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=update_profile_message,
|
||||
model_name=self.config.update_profile_model,
|
||||
max_token=self.config.update_profile_max_token,
|
||||
temperature=self.config.update_profile_temperature,
|
||||
top_k=self.config.update_profile_top_k)
|
||||
model_name=self.update_profile_model,
|
||||
max_token=self.update_profile_max_token,
|
||||
temperature=self.update_profile_temperature,
|
||||
top_k=self.update_profile_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
@ -105,7 +118,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
|
||||
def add_extra_user_attrs(self):
|
||||
# 解析为空返回
|
||||
extra_user_attr_list = [x.strip() for x in self.config.extra_user_attrs if x.strip()]
|
||||
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
|
||||
if not extra_user_attr_list:
|
||||
return
|
||||
|
||||
|
|
@ -152,7 +165,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# 增加环境变量配置的属性
|
||||
if self.config.extra_user_attrs:
|
||||
if self.extra_user_attrs:
|
||||
self.add_extra_user_attrs()
|
||||
|
||||
new_user_profile: List[UserAttribute] = []
|
||||
|
|
@ -178,8 +191,8 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.config.update_profile_max_thread:
|
||||
result_sorted = result_sorted[:self.config.update_profile_max_thread]
|
||||
if len(result_sorted) > self.update_profile_max_thread:
|
||||
result_sorted = result_sorted[:self.update_profile_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for user_attr, filtered_nodes, _ in result_sorted:
|
||||
|
|
|
|||
|
|
@ -4,11 +4,17 @@ from common.response_text_parser import ResponseTextParser
|
|||
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
|
||||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
def __init__(self, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
|
||||
super(ContraRepeatWorker, self).__init__(*args, **kwargs)
|
||||
self.merge_obs_model = merge_obs_model
|
||||
self.merge_obs_max_token = merge_obs_max_token
|
||||
self.merge_obs_temperature = merge_obs_temperature
|
||||
self.merge_obs_top_k = merge_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
|
|
@ -39,10 +45,10 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.config.merge_obs_model,
|
||||
max_token=self.config.merge_obs_max_token,
|
||||
temperature=self.config.merge_obs_temperature,
|
||||
top_k=self.config.merge_obs_top_k)
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -7,12 +7,18 @@ from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME,
|
|||
NEW_OBS_WITH_TIME_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.message import Message
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
|
||||
super(GetObservationWithTimeWorker, self).__init__(*args, **kwargs)
|
||||
self.summary_messages_model = summary_messages_model
|
||||
self.summary_messages_max_token = summary_messages_max_token
|
||||
self.summary_messages_temperature = summary_messages_temperature
|
||||
self.summary_messages_top_k = summary_messages_top_k
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, time_infer: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
|
|
@ -35,14 +41,14 @@ class GetObservationWithTimeWorker(MemoryBaseWorker):
|
|||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryWrapNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.config.memory_id,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.config.tenant_id)
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
|
|
@ -74,10 +80,10 @@ class GetObservationWithTimeWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.config.summary_messages_model,
|
||||
max_token=self.config.summary_messages_max_token,
|
||||
temperature=self.config.summary_messages_temperature,
|
||||
top_k=self.config.summary_messages_top_k)
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -7,12 +7,18 @@ from constants.common_constants import REFLECTED, DT, NEW_OBS_NODES, TIME_INFER,
|
|||
DATATIME_WORD_LIST
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.message import Message
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
|
||||
super(GetObservationWorker, self).__init__(*args,**kwargs)
|
||||
self.summary_messages_model = summary_messages_model
|
||||
self.summary_messages_max_token = summary_messages_max_token
|
||||
self.summary_messages_temperature = summary_messages_temperature
|
||||
self.summary_messages_top_k = summary_messages_top_k
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
|
|
@ -28,17 +34,17 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
TIME_INFER: "", # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
}
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryWrapNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.config.memory_id,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.config.tenant_id)
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
|
|
@ -66,10 +72,10 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.config.summary_messages_model,
|
||||
max_token=self.config.summary_messages_max_token,
|
||||
temperature=self.config.summary_messages_temperature,
|
||||
top_k=self.config.summary_messages_top_k)
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
|
|||
|
|
@ -1,9 +1,16 @@
|
|||
from common.response_text_parser import ResponseTextParser
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
def __init__(self, info_filter_msg_max_size, info_filter_model, info_filter_max_token, info_filter_temperature, info_filter_top_k, *args, **kwargs):
|
||||
super(InfoFilterWorker,self).__init__(*args,**kwargs)
|
||||
self.info_filter_msg_max_size
|
||||
self.info_filter_model = info_filter_model
|
||||
self.info_filter_max_token = info_filter_max_token
|
||||
self.info_filter_temperature = info_filter_temperature
|
||||
self.info_filter_top_k = info_filter_top_k
|
||||
|
||||
def _run(self):
|
||||
# filter user msg
|
||||
|
|
@ -11,7 +18,7 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
for msg in self.messages:
|
||||
if msg.role != MessageRoleEnum.USER.value:
|
||||
continue
|
||||
if len(msg.content) >= self.config.info_filter_msg_max_size:
|
||||
if len(msg.content) >= self.info_filter_msg_max_size:
|
||||
continue
|
||||
info_messages.append(msg)
|
||||
|
||||
|
|
@ -25,10 +32,10 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
|
||||
# call llm
|
||||
response_text = self.gene_client.call(messages=info_filter_message,
|
||||
model_name=self.config.info_filter_model,
|
||||
max_token=self.config.info_filter_max_token,
|
||||
temperature=self.config.info_filter_temperature,
|
||||
top_k=self.config.info_filter_top_k)
|
||||
model_name=self.info_filter_model,
|
||||
max_token=self.info_filter_max_token,
|
||||
temperature=self.info_filter_temperature,
|
||||
top_k=self.info_filter_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
|
|
@ -49,7 +56,7 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
continue
|
||||
score = info_score[0]
|
||||
# if score in ("1", "2",):
|
||||
if score in ("3",):
|
||||
if score in ("2",):
|
||||
msg.info_score = score
|
||||
filtered_messages.append(msg)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue