diff --git a/config/global_config.json b/config/global_config.json index e69de29b..14fe392a 100644 --- a/config/global_config.json +++ b/config/global_config.json @@ -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" +} \ No newline at end of file diff --git a/config/model/dash.json b/config/model/dash.json index e69de29b..73115431 100644 --- a/config/model/dash.json +++ b/config/model/dash.json @@ -0,0 +1,6 @@ +{ + "embedding": "", + "generate": "", + "rerank": "", + "es": "" +} \ No newline at end of file diff --git a/config/prompt.json b/config/prompt.json index e69de29b..1b5909d5 100644 --- a/config/prompt.json +++ b/config/prompt.json @@ -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": "" +} \ No newline at end of file diff --git a/config/worker_config.json b/config/worker.json similarity index 69% rename from config/worker_config.json rename to config/worker.json index cc33434b..90796974 100644 --- a/config/worker_config.json +++ b/config/worker.json @@ -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 } - }, + } } \ No newline at end of file diff --git a/memory_scope/pipeline/memory.py b/memory_scope/pipeline/memory.py index 61258c45..37457eec 100644 --- a/memory_scope/pipeline/memory.py +++ b/memory_scope/pipeline/memory.py @@ -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") diff --git a/memory_scope/pipeline/memory_service_bailian.py b/memory_scope/pipeline/memory_service_bailian.py index 0b122940..cf7ca0f9 100644 --- a/memory_scope/pipeline/memory_service_bailian.py +++ b/memory_scope/pipeline/memory_service_bailian.py @@ -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 diff --git a/memory_scope/utils/context_handler.py b/memory_scope/utils/context_handler.py index a55afa00..71ed7bcd 100644 --- a/memory_scope/utils/context_handler.py +++ b/memory_scope/utils/context_handler.py @@ -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): # 多线程环境下,如果是指针下修改,不安全 diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index a72f804b..959f8f41 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -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")) \ No newline at end of file diff --git a/memory_scope/worker/base_worker.py b/memory_scope/worker/base_worker.py index 81fdbc51..1509e2df 100644 --- a/memory_scope/worker/base_worker.py +++ b/memory_scope/worker/base_worker.py @@ -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: diff --git a/memory_scope/worker/es/es_insight_worker.py b/memory_scope/worker/es/es_insight_worker.py index 8239f6a7..620b1e26 100644 --- a/memory_scope/worker/es/es_insight_worker.py +++ b/memory_scope/worker/es/es_insight_worker.py @@ -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, diff --git a/memory_scope/worker/es/es_new_obs_worker.py b/memory_scope/worker/es/es_new_obs_worker.py index fbebaddf..b8aa02cc 100644 --- a/memory_scope/worker/es/es_new_obs_worker.py +++ b/memory_scope/worker/es/es_new_obs_worker.py @@ -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, diff --git a/memory_scope/worker/es/es_not_reflected_worker.py b/memory_scope/worker/es/es_not_reflected_worker.py index c031c428..e4c9957d 100644 --- a/memory_scope/worker/es/es_not_reflected_worker.py +++ b/memory_scope/worker/es/es_not_reflected_worker.py @@ -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, diff --git a/memory_scope/worker/es/es_retrieve_all_worker.py b/memory_scope/worker/es/es_retrieve_all_worker.py index e0ba99a8..b0d006d6 100644 --- a/memory_scope/worker/es/es_retrieve_all_worker.py +++ b/memory_scope/worker/es/es_retrieve_all_worker.py @@ -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, diff --git a/memory_scope/worker/es/es_similar_worker.py b/memory_scope/worker/es/es_similar_worker.py index 3b0f1f10..9a198e09 100644 --- a/memory_scope/worker/es/es_similar_worker.py +++ b/memory_scope/worker/es/es_similar_worker.py @@ -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) diff --git a/memory_scope/worker/es/es_today_obs_worker.py b/memory_scope/worker/es/es_today_obs_worker.py index 904ed171..c1d8f673 100644 --- a/memory_scope/worker/es/es_today_obs_worker.py +++ b/memory_scope/worker/es/es_today_obs_worker.py @@ -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, diff --git a/memory_scope/worker/es/load_profile_worker.py b/memory_scope/worker/es/load_profile_worker.py index 6065e804..ef7cbab1 100644 --- a/memory_scope/worker/es/load_profile_worker.py +++ b/memory_scope/worker/es/load_profile_worker.py @@ -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, diff --git a/memory_scope/worker/memory_base_worker.py b/memory_scope/worker/memory_base_worker.py index b9bc8729..606eac9a 100644 --- a/memory_scope/worker/memory_base_worker.py +++ b/memory_scope/worker/memory_base_worker.py @@ -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): diff --git a/memory_scope/worker/memory_store_worker.py b/memory_scope/worker/memory_store_worker.py index 437258a0..5740dc89 100644 --- a/memory_scope/worker/memory_store_worker.py +++ b/memory_scope/worker/memory_store_worker.py @@ -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): diff --git a/memory_scope/worker/parse_params_worker.py b/memory_scope/worker/parse_params_worker.py index ed92d8fd..76bff3ee 100644 --- a/memory_scope/worker/parse_params_worker.py +++ b/memory_scope/worker/parse_params_worker.py @@ -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): diff --git a/memory_scope/worker/retrieve/extract_time_worker.py b/memory_scope/worker/retrieve/extract_time_worker.py index 54094089..0a492431 100644 --- a/memory_scope/worker/retrieve/extract_time_worker.py +++ b/memory_scope/worker/retrieve/extract_time_worker.py @@ -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: diff --git a/memory_scope/worker/retrieve/fuse_rerank_worker.py b/memory_scope/worker/retrieve/fuse_rerank_worker.py index 225ef0b0..e56bb7aa 100644 --- a/memory_scope/worker/retrieve/fuse_rerank_worker.py +++ b/memory_scope/worker/retrieve/fuse_rerank_worker.py @@ -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) diff --git a/memory_scope/worker/retrieve/semantic_rank_worker.py b/memory_scope/worker/retrieve/semantic_rank_worker.py index 08baf118..c20e672c 100644 --- a/memory_scope/worker/retrieve/semantic_rank_worker.py +++ b/memory_scope/worker/retrieve/semantic_rank_worker.py @@ -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): diff --git a/memory_scope/worker/summary_long/get_insight_worker.py b/memory_scope/worker/summary_long/get_insight_worker.py index a2494db0..ede0c44e 100644 --- a/memory_scope/worker/summary_long/get_insight_worker.py +++ b/memory_scope/worker/summary_long/get_insight_worker.py @@ -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: diff --git a/memory_scope/worker/summary_long/get_reflection_worker.py b/memory_scope/worker/summary_long/get_reflection_worker.py index 1b8d4419..45ea8060 100644 --- a/memory_scope/worker/summary_long/get_reflection_worker.py +++ b/memory_scope/worker/summary_long/get_reflection_worker.py @@ -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: diff --git a/memory_scope/worker/summary_long/long_contra_repeat_worker.py b/memory_scope/worker/summary_long/long_contra_repeat_worker.py index c16af9d4..87bb77d8 100644 --- a/memory_scope/worker/summary_long/long_contra_repeat_worker.py +++ b/memory_scope/worker/summary_long/long_contra_repeat_worker.py @@ -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: diff --git a/memory_scope/worker/summary_long/summary_collect_worker.py b/memory_scope/worker/summary_long/summary_collect_worker.py index e26ab6fa..d42520aa 100644 --- a/memory_scope/worker/summary_long/summary_collect_worker.py +++ b/memory_scope/worker/summary_long/summary_collect_worker.py @@ -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): diff --git a/memory_scope/worker/summary_long/update_insight_worker.py b/memory_scope/worker/summary_long/update_insight_worker.py index ee2d889d..e30b0e3d 100644 --- a/memory_scope/worker/summary_long/update_insight_worker.py +++ b/memory_scope/worker/summary_long/update_insight_worker.py @@ -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: diff --git a/memory_scope/worker/summary_long/update_profile_worker.py b/memory_scope/worker/summary_long/update_profile_worker.py index 692f2e50..e2f621b7 100644 --- a/memory_scope/worker/summary_long/update_profile_worker.py +++ b/memory_scope/worker/summary_long/update_profile_worker.py @@ -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: diff --git a/memory_scope/worker/summary_short/contra_repeat_worker.py b/memory_scope/worker/summary_short/contra_repeat_worker.py index 36ccf9af..34ca3760 100644 --- a/memory_scope/worker/summary_short/contra_repeat_worker.py +++ b/memory_scope/worker/summary_short/contra_repeat_worker.py @@ -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: diff --git a/memory_scope/worker/summary_short/get_observation_with_time_worker.py b/memory_scope/worker/summary_short/get_observation_with_time_worker.py index 63c13efe..4d5a772f 100644 --- a/memory_scope/worker/summary_short/get_observation_with_time_worker.py +++ b/memory_scope/worker/summary_short/get_observation_with_time_worker.py @@ -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: diff --git a/memory_scope/worker/summary_short/get_observation_worker.py b/memory_scope/worker/summary_short/get_observation_worker.py index 9ed4a75c..3d850259 100644 --- a/memory_scope/worker/summary_short/get_observation_worker.py +++ b/memory_scope/worker/summary_short/get_observation_worker.py @@ -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: diff --git a/memory_scope/worker/summary_short/info_filter_worker.py b/memory_scope/worker/summary_short/info_filter_worker.py index f369ca4e..311c8d8b 100644 --- a/memory_scope/worker/summary_short/info_filter_worker.py +++ b/memory_scope/worker/summary_short/info_filter_worker.py @@ -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)