pipeline config

This commit is contained in:
hs 2024-06-18 12:32:37 +08:00
parent cf94bbf927
commit 818d4d794c
32 changed files with 404 additions and 496 deletions

View file

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

View file

@ -0,0 +1,6 @@
{
"embedding": "",
"generate": "",
"rerank": "",
"es": ""
}

View file

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

View file

@ -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
}
},
}
}

View file

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

View file

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

View file

@ -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):
# 多线程环境下,如果是指针下修改,不安全

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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