From 9c79fd95d81e1dad3fd10d0db366dcc50e20fdc9 Mon Sep 17 00:00:00 2001 From: hs Date: Tue, 18 Jun 2024 16:50:45 +0800 Subject: [PATCH 1/3] add cli_config & fix api --- memory_scope/cli/cli_config.py | 67 ++++++++++++++++++- memory_scope/db/elastic_search_client.py | 2 +- memory_scope/models/dash_client.py | 4 +- memory_scope/node/memory_wrap_node.py | 2 +- memory_scope/parsers/response_text_parser.py | 2 +- memory_scope/pipeline/memory.py | 9 ++- ...y_service_bailian.py => memory_service.py} | 17 +++-- memory_scope/utils/context_handler.py | 2 +- memory_scope/utils/timer.py | 2 +- memory_scope/utils/tool_functions.py | 31 ++------- memory_scope/utils/user_profile_handler.py | 4 +- memory_scope/worker/base_worker.py | 6 +- memory_scope/worker/es/es_insight_worker.py | 4 +- memory_scope/worker/es/es_keyword_worker.py | 4 +- memory_scope/worker/es/es_new_obs_worker.py | 4 +- .../worker/es/es_not_reflected_worker.py | 4 +- .../worker/es/es_retrieve_all_worker.py | 4 +- memory_scope/worker/es/es_similar_worker.py | 4 +- memory_scope/worker/es/es_today_obs_worker.py | 5 +- memory_scope/worker/es/load_profile_worker.py | 10 +-- memory_scope/worker/memory_base_worker.py | 18 ++--- memory_scope/worker/memory_store_worker.py | 11 ++- memory_scope/worker/parse_params_worker.py | 4 +- .../worker/retrieve/extract_time_worker.py | 2 +- .../worker/retrieve/fuse_rerank_worker.py | 4 +- .../worker/retrieve/semantic_rank_worker.py | 6 +- .../worker/summary_long/get_insight_worker.py | 4 +- .../summary_long/get_reflection_worker.py | 4 +- .../summary_long/long_contra_repeat_worker.py | 4 +- .../summary_long/summary_collect_worker.py | 4 +- .../summary_long/update_insight_worker.py | 4 +- .../summary_long/update_profile_worker.py | 6 +- .../summary_short/contra_repeat_worker.py | 4 +- .../get_observation_with_time_worker.py | 6 +- .../summary_short/get_observation_worker.py | 6 +- .../summary_short/info_filter_worker.py | 2 +- tests/test_memory.py | 8 +-- 37 files changed, 162 insertions(+), 122 deletions(-) rename memory_scope/pipeline/{memory_service_bailian.py => memory_service.py} (93%) diff --git a/memory_scope/cli/cli_config.py b/memory_scope/cli/cli_config.py index fdb68c7c..8812cec3 100644 --- a/memory_scope/cli/cli_config.py +++ b/memory_scope/cli/cli_config.py @@ -1 +1,66 @@ -# runtime params apikey or overwrite config params + +import json +from pydantic import BaseModel + +from models.dash_embedding_client import DashEmbeddingClient +from models.dash_generate_client import DashGenerateClient +from models.dash_rerank_client import DashReRankClient +from models.elastic_search_client import ElasticSearchClient +from pipeline.memory_service import MemoryService +from enumeration.memory_type_enum import MemoryTypeEnum +from utils.tool_functions import init_instance_by_config + +class Wrapper: + """Wrapper class for anything that needs to set up during init""" + + def __init__(self): + self._provider = None + + def register(self, provider): + self._provider = provider + + def __getattr__(self, key): + if self.__dict__.get("_provider", None) is None: + raise AttributeError("Please run init() first using qlib") + return getattr(self._provider, key) + +class Config(BaseModel): + thread_pool_max_count: int = 5 + pipeline: dict = { + "retrive": """ +parse_params,es.load_profile,[retrieve.extract_time|es.es_similar|es.es_keyword],retrieve.semantic_rank,retrieve.fuse_rerank +""".strip(), + "summary_long": """ +parse_params,summary_short.info_filter,[es.es_today_obs|summary_short.get_observation|summary_short.get_observation_with_time],summary_short.contra_repeat,memory_store +""".strip(), + "summary_short": """ +parse_params,[es.load_profile|es.es_new_obs|es.es_insight],[summary_long.update_insight|summary_long.get_reflection,summary_long.get_insight|summary_long.update_profile],summary_long.summary_collect,memory_store +""".strip() + } + model_embedding = _default_embedding_client + model_generate = _default_generate_client + model_rerank = _default_rerank_client + db = _default_es_client + +class UserConfig(BaseModel): + memory_id: str = "" + +C = Wrapper() + +def init(config_path): + config = json.loads(config_path) + C.register(Config(**config)) + + # ## register modules + C.model_embedding = init_instance_by_config(C.model_embedding) + C.model_generate = init_instance_by_config(C.model_generate) + C.model_rerank = init_instance_by_config(C.model_rerank) + C.db = init_instance_by_config(C.db) + + ## register workers + C.worker = json.loads(C.worker) + + ## register services + for k,v in C.pipeline.items(): + C.pipeline[k] = MemoryService(k) + diff --git a/memory_scope/db/elastic_search_client.py b/memory_scope/db/elastic_search_client.py index 74138cf3..ec069f5e 100644 --- a/memory_scope/db/elastic_search_client.py +++ b/memory_scope/db/elastic_search_client.py @@ -2,7 +2,7 @@ from elasticsearch import Elasticsearch from elasticsearch.helpers import bulk from common.dash_embedding_client import DashEmbeddingClient -from common.logger import Logger +from utils.logger import Logger from constants.common_constants import ES_ENV_URL_DICT from enumeration.env_type import EnvType diff --git a/memory_scope/models/dash_client.py b/memory_scope/models/dash_client.py index 6d07fed4..1d526279 100644 --- a/memory_scope/models/dash_client.py +++ b/memory_scope/models/dash_client.py @@ -4,8 +4,8 @@ from http import HTTPStatus import requests -from common.logger import Logger -from common.timer import Timer +from utils.logger import Logger +from utils.timer import Timer from enumeration.env_type import EnvType diff --git a/memory_scope/node/memory_wrap_node.py b/memory_scope/node/memory_wrap_node.py index 84029fe6..899dd323 100644 --- a/memory_scope/node/memory_wrap_node.py +++ b/memory_scope/node/memory_wrap_node.py @@ -1,6 +1,6 @@ from pydantic import Field, BaseModel -from model.memory_node import MemoryNode +from node.memory_node import MemoryNode class MemoryWrapNode(BaseModel): diff --git a/memory_scope/parsers/response_text_parser.py b/memory_scope/parsers/response_text_parser.py index cfeee808..48dbbd9c 100644 --- a/memory_scope/parsers/response_text_parser.py +++ b/memory_scope/parsers/response_text_parser.py @@ -1,6 +1,6 @@ import re -from common.logger import Logger +from utils.logger import Logger class ResponseTextParser(object): diff --git a/memory_scope/pipeline/memory.py b/memory_scope/pipeline/memory.py index 37457eec..769dab4b 100644 --- a/memory_scope/pipeline/memory.py +++ b/memory_scope/pipeline/memory.py @@ -1,12 +1,11 @@ from typing import List, Dict -from pydantic import Field +from pydantic import Field, BaseModel -from model.message import Message -from model.user_attribute import UserAttribute -from request.base_model import RequestBaseModel +from node.message import Message +from node.user_attribute import UserAttribute -class MemoryServiceRequestModel(RequestBaseModel): +class MemoryServiceRequestModel(BaseModel): user: UserConfig = None messages: List[Message] = Field(..., diff --git a/memory_scope/pipeline/memory_service_bailian.py b/memory_scope/pipeline/memory_service.py similarity index 93% rename from memory_scope/pipeline/memory_service_bailian.py rename to memory_scope/pipeline/memory_service.py index cf7ca0f9..3a505a02 100644 --- a/memory_scope/pipeline/memory_service_bailian.py +++ b/memory_scope/pipeline/memory_service.py @@ -6,17 +6,16 @@ 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 +from worker.base_worker import BaseWorker +from utils.context_handler import ContextHandler +from utils.logger import Logger +from utils.timer import timer, Timer from common.tool_functions import under_line_to_hump 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 config.env_config import C, Workers +from pipeline.memory import MemoryServiceRequestModel +from cli.cli_config import C from utils.tool_functions import init_instance_by_config @@ -26,7 +25,7 @@ class MemoryService(object): self.context_handler = ContextHandler() # 线程池 - self.thread_pool = ThreadPoolExecutor(max_workers=C.THREAD_POOL_MAX_COUNT) + self.thread_pool = ThreadPoolExecutor(max_workers=C.thread_pool_max_count) # 全部初始化的worker self.worker_dict: Dict[str, BaseWorker] = {} @@ -40,7 +39,7 @@ class MemoryService(object): def get_worker(self, worker_name: str, is_multi_thread: bool = False) -> BaseWorker: return init_instance_by_config( - config = W.get(worker_name), + config = C.worker.get(worker_name), try_kwargs={ "is_multi_thread": is_multi_thread, "thread_pool": self.thread_pool diff --git a/memory_scope/utils/context_handler.py b/memory_scope/utils/context_handler.py index 71ed7bcd..a4ee67d4 100644 --- a/memory_scope/utils/context_handler.py +++ b/memory_scope/utils/context_handler.py @@ -2,7 +2,7 @@ import os import threading from typing import Dict, Any -from common.logger import Logger +from utils.logger import Logger class ContextHandler(object): def __init__(self): diff --git a/memory_scope/utils/timer.py b/memory_scope/utils/timer.py index 4098d120..0c576363 100644 --- a/memory_scope/utils/timer.py +++ b/memory_scope/utils/timer.py @@ -6,7 +6,7 @@ date: 20221106 import time -from common.logger import Logger +from utils.logger import Logger class Timer(object): diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 959f8f41..849c9e0d 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -4,29 +4,8 @@ from datetime import datetime 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 - - -def get_global_env_type(): - global global_env_type - - if global_env_type is None: - env = os.environ.get("APP_ENV", "") - if env is None or not env: - raise EnvironmentError("Environment variable APP_ENV must be set") - env = env.split("-")[-1] - - if env not in EnvType.__members__.values(): - global_env_type = EnvType.DAILY - else: - global_env_type = EnvType(env) - - return global_env_type - - def under_line_to_hump(underline_str): sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str) return sub[0:1].upper() + sub[1:] @@ -165,10 +144,10 @@ def time_to_formatted_str(time: datetime | str | int | float = None, 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")) +def init_instance_by_config(config: dict, default_module_path=None, try_kwargs={}): + import_module(config.pop("path", default_module_path)) + clazz = getattr(module, config.pop("name")) try: - return clazz(**config.get("kwargs"),**try_kwargs) + return clazz(**config,**try_kwargs) except: - return clazz(**config.get("kwargs")) \ No newline at end of file + return clazz(**config) \ No newline at end of file diff --git a/memory_scope/utils/user_profile_handler.py b/memory_scope/utils/user_profile_handler.py index 9a390f95..6c8bcf07 100644 --- a/memory_scope/utils/user_profile_handler.py +++ b/memory_scope/utils/user_profile_handler.py @@ -2,8 +2,8 @@ import json from typing import List, Dict from enumeration.memory_node_status import MemoryNodeStatus -from model.memory_wrap_node import MemoryWrapNode -from model.user_attribute import UserAttribute +from node.memory_wrap_node import MemoryWrapNode +from node.user_attribute import UserAttribute class UserProfileHandler(object): diff --git a/memory_scope/worker/base_worker.py b/memory_scope/worker/base_worker.py index 1509e2df..0cac7505 100644 --- a/memory_scope/worker/base_worker.py +++ b/memory_scope/worker/base_worker.py @@ -2,9 +2,9 @@ import time from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any, List -from common.context_handler import ContextHandler -from common.logger import Logger -from common.timer import Timer +from utils.context_handler import ContextHandler +from utils.logger import Logger +from utils.timer import Timer class BaseWorker(object): diff --git a/memory_scope/worker/es/es_insight_worker.py b/memory_scope/worker/es/es_insight_worker.py index 620b1e26..a42ad613 100644 --- a/memory_scope/worker/es/es_insight_worker.py +++ b/memory_scope/worker/es/es_insight_worker.py @@ -3,8 +3,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsInsightWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/es/es_keyword_worker.py b/memory_scope/worker/es/es_keyword_worker.py index f251b630..5df1cfd7 100644 --- a/memory_scope/worker/es/es_keyword_worker.py +++ b/memory_scope/worker/es/es_keyword_worker.py @@ -4,8 +4,8 @@ from constants.common_constants import KEY_WORD, KEYWORD_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 node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsKeywordWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/es/es_new_obs_worker.py b/memory_scope/worker/es/es_new_obs_worker.py index b8aa02cc..55a5d3e3 100644 --- a/memory_scope/worker/es/es_new_obs_worker.py +++ b/memory_scope/worker/es/es_new_obs_worker.py @@ -3,8 +3,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsNewObsWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/es/es_not_reflected_worker.py b/memory_scope/worker/es/es_not_reflected_worker.py index e4c9957d..690f9269 100644 --- a/memory_scope/worker/es/es_not_reflected_worker.py +++ b/memory_scope/worker/es/es_not_reflected_worker.py @@ -3,8 +3,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsNotReflectedWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/es/es_retrieve_all_worker.py b/memory_scope/worker/es/es_retrieve_all_worker.py index b0d006d6..8f388dd7 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsRetrieveAllWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/es/es_similar_worker.py b/memory_scope/worker/es/es_similar_worker.py index 9a198e09..c0abd514 100644 --- a/memory_scope/worker/es/es_similar_worker.py +++ b/memory_scope/worker/es/es_similar_worker.py @@ -4,8 +4,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsSimilarWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/es/es_today_obs_worker.py b/memory_scope/worker/es/es_today_obs_worker.py index c1d8f673..d362b331 100644 --- a/memory_scope/worker/es/es_today_obs_worker.py +++ b/memory_scope/worker/es/es_today_obs_worker.py @@ -4,9 +4,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker - +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class EsTodayObsWorker(MemoryBaseWorker): def __init__(self, es_today_obs_top_k, *args, **kwargs): diff --git a/memory_scope/worker/es/load_profile_worker.py b/memory_scope/worker/es/load_profile_worker.py index ef7cbab1..c8c5a182 100644 --- a/memory_scope/worker/es/load_profile_worker.py +++ b/memory_scope/worker/es/load_profile_worker.py @@ -1,13 +1,13 @@ from typing import List, Dict -from common.user_profile_handler import UserProfileHandler +from utils.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.memory_wrap_node import MemoryWrapNode -from model.user_attribute import UserAttribute -from request.memory import MemoryServiceRequestModel -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from node.user_attribute import UserAttribute +from pipeline.memory import MemoryServiceRequestModel +from worker.memory_base_worker import MemoryBaseWorker class LoadProfileWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/memory_base_worker.py b/memory_scope/worker/memory_base_worker.py index 606eac9a..850b0d3d 100644 --- a/memory_scope/worker/memory_base_worker.py +++ b/memory_scope/worker/memory_base_worker.py @@ -3,11 +3,11 @@ from typing import List, Dict, Optional 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.memory.base_worker import BaseWorker -from config.env_config import EnvConfig +from node.message import Message +from node.user_attribute import UserAttribute +from pipeline.memory import MemoryServiceRequestModel +from worker.base_worker import BaseWorker +from cli.cli_config import C class MemoryBaseWorker(BaseWorker): def __init__(self, **kwargs): @@ -52,19 +52,19 @@ class MemoryBaseWorker(BaseWorker): @property def emb_client(self): - return C.emb_client + return C.model_embedding @property def gene_client(self): - return C.gene_client + return C.model_generate @property def rerank_client(self): - return C.rerank_client + return C.model_rerank @property def es_client(self): - return C.es_client + return C.db @property def tenant_id(self): diff --git a/memory_scope/worker/memory_store_worker.py b/memory_scope/worker/memory_store_worker.py index 5740dc89..35ef8b1c 100644 --- a/memory_scope/worker/memory_store_worker.py +++ b/memory_scope/worker/memory_store_worker.py @@ -1,12 +1,11 @@ from typing import List -from common.user_profile_handler import UserProfileHandler +from utils.user_profile_handler import UserProfileHandler from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE -from model.memory_node import MemoryNode -from model.memory.memory_wrap_node import MemoryWrapNode -from model.user_attribute import UserAttribute -from worker.memory.memory_base_worker import MemoryBaseWorker - +from node.memory_node import MemoryNode +from node.memory_wrap_node import MemoryWrapNode +from node.user_attribute import UserAttribute +from worker.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 76bff3ee..e26620a7 100644 --- a/memory_scope/worker/parse_params_worker.py +++ b/memory_scope/worker/parse_params_worker.py @@ -2,8 +2,8 @@ import json from config.bailian_memory_config import BailianMemoryConfig from constants.common_constants import REQUEST, CONFIG -from request.memory import MemoryServiceRequestModel -from worker.memory.base_worker import BaseWorker +from pipeline.memory import MemoryServiceRequestModel +from worker.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 0a492431..fd297d14 100644 --- a/memory_scope/worker/retrieve/extract_time_worker.py +++ b/memory_scope/worker/retrieve/extract_time_worker.py @@ -3,7 +3,7 @@ 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.memory.memory_base_worker import MemoryBaseWorker +from worker.memory_base_worker import MemoryBaseWorker class ExtractTimeWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/retrieve/fuse_rerank_worker.py b/memory_scope/worker/retrieve/fuse_rerank_worker.py index e56bb7aa..7b579eca 100644 --- a/memory_scope/worker/retrieve/fuse_rerank_worker.py +++ b/memory_scope/worker/retrieve/fuse_rerank_worker.py @@ -2,8 +2,8 @@ from typing import Dict, List from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \ TIME_MATCHED -from model.memory.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class FuseRerankWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/retrieve/semantic_rank_worker.py b/memory_scope/worker/retrieve/semantic_rank_worker.py index c20e672c..2864ef74 100644 --- a/memory_scope/worker/retrieve/semantic_rank_worker.py +++ b/memory_scope/worker/retrieve/semantic_rank_worker.py @@ -1,11 +1,11 @@ from typing import List, Dict -from common.user_profile_handler import UserProfileHandler +from utils.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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.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 ede0c44e..1c6b4458 100644 --- a/memory_scope/worker/summary_long/get_insight_worker.py +++ b/memory_scope/worker/summary_long/get_insight_worker.py @@ -6,8 +6,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class GetInsightWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/summary_long/get_reflection_worker.py b/memory_scope/worker/summary_long/get_reflection_worker.py index 45ea8060..0e8fe4ca 100644 --- a/memory_scope/worker/summary_long/get_reflection_worker.py +++ b/memory_scope/worker/summary_long/get_reflection_worker.py @@ -3,8 +3,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class GetReflectionWorker(MemoryBaseWorker): 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 87bb77d8..a517d830 100644 --- a/memory_scope/worker/summary_long/long_contra_repeat_worker.py +++ b/memory_scope/worker/summary_long/long_contra_repeat_worker.py @@ -5,8 +5,8 @@ from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, MODIFIED_MEMORIES 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 node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class LongContraRepeatWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/summary_long/summary_collect_worker.py b/memory_scope/worker/summary_long/summary_collect_worker.py index d42520aa..613638b9 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.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 e30b0e3d..96f49c6a 100644 --- a/memory_scope/worker/summary_long/update_insight_worker.py +++ b/memory_scope/worker/summary_long/update_insight_worker.py @@ -2,8 +2,8 @@ from typing import List from common.response_text_parser import ResponseTextParser 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 +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class UpdateInsightWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/summary_long/update_profile_worker.py b/memory_scope/worker/summary_long/update_profile_worker.py index e2f621b7..1bcb7d13 100644 --- a/memory_scope/worker/summary_long/update_profile_worker.py +++ b/memory_scope/worker/summary_long/update_profile_worker.py @@ -3,9 +3,9 @@ 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.memory_wrap_node import MemoryWrapNode -from model.user_attribute import UserAttribute -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from node.user_attribute import UserAttribute +from worker.memory_base_worker import MemoryBaseWorker class UpdateProfileWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/summary_short/contra_repeat_worker.py b/memory_scope/worker/summary_short/contra_repeat_worker.py index 34ca3760..fa8d8f45 100644 --- a/memory_scope/worker/summary_short/contra_repeat_worker.py +++ b/memory_scope/worker/summary_short/contra_repeat_worker.py @@ -4,8 +4,8 @@ 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.memory_wrap_node import MemoryWrapNode -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from worker.memory_base_worker import MemoryBaseWorker class ContraRepeatWorker(MemoryBaseWorker): 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 4d5a772f..26983ea6 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,9 +7,9 @@ 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.memory_wrap_node import MemoryWrapNode -from model.message import Message -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from node.message import Message +from worker.memory_base_worker import MemoryBaseWorker class GetObservationWithTimeWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/summary_short/get_observation_worker.py b/memory_scope/worker/summary_short/get_observation_worker.py index 3d850259..b4b77b47 100644 --- a/memory_scope/worker/summary_short/get_observation_worker.py +++ b/memory_scope/worker/summary_short/get_observation_worker.py @@ -7,9 +7,9 @@ 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.memory_wrap_node import MemoryWrapNode -from model.message import Message -from worker.memory.memory_base_worker import MemoryBaseWorker +from node.memory_wrap_node import MemoryWrapNode +from node.message import Message +from worker.memory_base_worker import MemoryBaseWorker class GetObservationWorker(MemoryBaseWorker): diff --git a/memory_scope/worker/summary_short/info_filter_worker.py b/memory_scope/worker/summary_short/info_filter_worker.py index 311c8d8b..0723ea54 100644 --- a/memory_scope/worker/summary_short/info_filter_worker.py +++ b/memory_scope/worker/summary_short/info_filter_worker.py @@ -1,6 +1,6 @@ from common.response_text_parser import ResponseTextParser from enumeration.message_role_enum import MessageRoleEnum -from worker.memory.memory_base_worker import MemoryBaseWorker +from worker.memory_base_worker import MemoryBaseWorker class InfoFilterWorker(MemoryBaseWorker): diff --git a/tests/test_memory.py b/tests/test_memory.py index 24d263a1..a6ef0235 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -3,12 +3,12 @@ import os import time from typing import List, Dict -from common.logger import Logger +from utils.logger import Logger from constants.common_constants import NEW_USER_PROFILE, MODIFIED_MEMORIES, RELATED_MEMORIES from enumeration.memory_method_enum import MemoryMethodEnum -from model.memory_node import MemoryNode -from model.user_attribute import UserAttribute -from request.memory import MemoryServiceRequestModel +from node.memory_node import MemoryNode +from node.user_attribute import UserAttribute +from pipeline.memory import MemoryServiceRequestModel from service.memory_service_bailian import MemoryServiceBailian """ From 31b2f80d57271d3d780023e8ed3014f194cbdaeb Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 18 Jun 2024 18:00:15 +0800 Subject: [PATCH 2/3] [dev] move config dir --- config/worker.json | 11 +++--- memory_scope/{cli => }/cli.py | 0 memory_scope/cli/__init__.py | 0 memory_scope/cli/cli_config.py | 66 ---------------------------------- memory_scope/config.py | 36 +++++++++++++++++++ 5 files changed, 41 insertions(+), 72 deletions(-) rename memory_scope/{cli => }/cli.py (100%) delete mode 100644 memory_scope/cli/__init__.py delete mode 100644 memory_scope/cli/cli_config.py create mode 100644 memory_scope/config.py diff --git a/config/worker.json b/config/worker.json index a19d5d10..bcce1b54 100644 --- a/config/worker.json +++ b/config/worker.json @@ -59,12 +59,11 @@ "ExtractTimeWorker": { "module_name": "ExtractTimeWorker", "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_model": "", + "parse_time_model": "qwen_1_8_parse_time_service", + "parse_time_max_token": 100, + "parse_time_temperature": 0.6, + "parse_time_top_k": 1, }, "InfoFilterWorker": { "module_name": "InfoFilterWorker", diff --git a/memory_scope/cli/cli.py b/memory_scope/cli.py similarity index 100% rename from memory_scope/cli/cli.py rename to memory_scope/cli.py diff --git a/memory_scope/cli/__init__.py b/memory_scope/cli/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/memory_scope/cli/cli_config.py b/memory_scope/cli/cli_config.py deleted file mode 100644 index 8812cec3..00000000 --- a/memory_scope/cli/cli_config.py +++ /dev/null @@ -1,66 +0,0 @@ - -import json -from pydantic import BaseModel - -from models.dash_embedding_client import DashEmbeddingClient -from models.dash_generate_client import DashGenerateClient -from models.dash_rerank_client import DashReRankClient -from models.elastic_search_client import ElasticSearchClient -from pipeline.memory_service import MemoryService -from enumeration.memory_type_enum import MemoryTypeEnum -from utils.tool_functions import init_instance_by_config - -class Wrapper: - """Wrapper class for anything that needs to set up during init""" - - def __init__(self): - self._provider = None - - def register(self, provider): - self._provider = provider - - def __getattr__(self, key): - if self.__dict__.get("_provider", None) is None: - raise AttributeError("Please run init() first using qlib") - return getattr(self._provider, key) - -class Config(BaseModel): - thread_pool_max_count: int = 5 - pipeline: dict = { - "retrive": """ -parse_params,es.load_profile,[retrieve.extract_time|es.es_similar|es.es_keyword],retrieve.semantic_rank,retrieve.fuse_rerank -""".strip(), - "summary_long": """ -parse_params,summary_short.info_filter,[es.es_today_obs|summary_short.get_observation|summary_short.get_observation_with_time],summary_short.contra_repeat,memory_store -""".strip(), - "summary_short": """ -parse_params,[es.load_profile|es.es_new_obs|es.es_insight],[summary_long.update_insight|summary_long.get_reflection,summary_long.get_insight|summary_long.update_profile],summary_long.summary_collect,memory_store -""".strip() - } - model_embedding = _default_embedding_client - model_generate = _default_generate_client - model_rerank = _default_rerank_client - db = _default_es_client - -class UserConfig(BaseModel): - memory_id: str = "" - -C = Wrapper() - -def init(config_path): - config = json.loads(config_path) - C.register(Config(**config)) - - # ## register modules - C.model_embedding = init_instance_by_config(C.model_embedding) - C.model_generate = init_instance_by_config(C.model_generate) - C.model_rerank = init_instance_by_config(C.model_rerank) - C.db = init_instance_by_config(C.db) - - ## register workers - C.worker = json.loads(C.worker) - - ## register services - for k,v in C.pipeline.items(): - C.pipeline[k] = MemoryService(k) - diff --git a/memory_scope/config.py b/memory_scope/config.py new file mode 100644 index 00000000..dfca733b --- /dev/null +++ b/memory_scope/config.py @@ -0,0 +1,36 @@ + +import json + +from pipeline.memory_service import MemoryService +from utils.tool_functions import init_instance_by_config + + +class Wrapper: + """Wrapper class for anything that needs to set up during init""" + + def __init__(self): + self._provider = None + + def register(self, provider): + self._provider = provider + + def __getattr__(self, key: str): + if self.__dict__.get("_provider", None) is None: + raise AttributeError("Please run __init__ first!") + return getattr(self._provider, key) + + +C = Wrapper() + + +def init(config_path: str): + config = json.loads(config_path) + C.register(config) + + ## register workers + C.worker = json.loads(C.worker) + + ## register services + for k,v in C.pipeline.items(): + C.pipeline[k] = MemoryService(k) + From 75d01299092be1577cb196b8cd03f9085010a162 Mon Sep 17 00:00:00 2001 From: hs Date: Tue, 18 Jun 2024 19:25:00 +0800 Subject: [PATCH 3/3] fix api --- config/config.json | 6 +- config/worker.json | 74 ++++++++++------------- memory_scope/cli.py | 19 +++--- memory_scope/config.py | 2 - memory_scope/pipeline/memory.py | 6 +- memory_scope/utils/tool_functions.py | 7 ++- memory_scope/worker/memory_base_worker.py | 25 ++++++-- tests/test_memory.py | 18 +++--- 8 files changed, 81 insertions(+), 76 deletions(-) diff --git a/config/config.json b/config/config.json index 9f4d9316..899a84bd 100644 --- a/config/config.json +++ b/config/config.json @@ -2,9 +2,9 @@ "thread_pool_max_count": "", "worker": "config/worker.json", "pipeline": { - "summary_short": "", - "summary_long": "", - "retrieve": "" + "summary_short": "parse_params,[load_profile|es_new_obs|es_insight],[update_insight|get_reflection,get_insight|update_profile],summary_collect,memory_store", + "summary_long": "parse_params,info_filter,[es_today_obs|get_observation|get_observation_with_time],contra_repeat,memory_store", + "retrieve": "parse_params,load_profile,[extract_time|es_similar|es_keyword],semantic_rank,fuse_rerank" }, "model_embedding": "config/model/dash_embedding.json", "model_rerank": "config/model/dash_rerank.json", diff --git a/config/worker.json b/config/worker.json index bcce1b54..87724ef3 100644 --- a/config/worker.json +++ b/config/worker.json @@ -1,52 +1,45 @@ { - "EsInsightWorker": { - "module_name": "EsInsightWorker", - "module_path": "memory_scope/worker", - "kwargs": { - "es_insight_top_k": 128 - } + "es_insight": { + "name": "EsInsightWorker", + "path": "memory_scope/worker", + "es_insight_top_k": 128 }, "es_keyword": { "name": "EsKeywordWorker", "path": "memory_scope/worker", "es_keyword_top_k": 10 }, - "es_keyword2": { - "module_name": "EsKeywordWorker", - "module_path": "memory_scope/worker", - "es_keyword_top_k": 10 - }, "EsNewObsWorker": { - "module_name": "EsNewObsWorker", - "module_path": "memory_scope/worker", + "name": "EsNewObsWorker", + "path": "memory_scope/worker", "kwargs": { "es_new_obs_top_k": 256 } }, "EsNotReflectedWorker": { - "module_name": "EsNotReflectedWorker", - "module_path": "memory_scope/worker", + "name": "EsNotReflectedWorker", + "path": "memory_scope/worker", "kwargs": { "es_not_reflected_top_k": 256 } }, "EsSimilarWorker": { - "module_name": "EsSimilarWorker", - "module_path": "memory_scope/worker", + "name": "EsSimilarWorker", + "path": "memory_scope/worker", "kwargs": { "es_similar_top_k": 128 } }, "EsTodayObsWorker": { - "module_name": "EsTodayObsWorker", - "module_path": "memory_scope/worker", + "name": "EsTodayObsWorker", + "path": "memory_scope/worker", "kwargs": { "es_today_obs_top_k": 128 } }, "GetInsightWorker": { - "module_name": "GetInsightWorker", - "module_path": "memory_scope/worker", + "name": "GetInsightWorker", + "path": "memory_scope/worker", "kwargs": { "es_insight_similar_top_k": 128, "insight_obs_max_cnt": 10, @@ -57,17 +50,16 @@ } }, "ExtractTimeWorker": { - "module_name": "ExtractTimeWorker", - "module_path": "memory_scope/worker", - "parse_time_model": "", + "name": "ExtractTimeWorker", + "path": "memory_scope/worker", "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_scope/worker", + "name": "InfoFilterWorker", + "path": "memory_scope/worker", "kwargs": { "info_filter_msg_max_size": 200, "info_filter_model": "qwen_max", @@ -77,8 +69,8 @@ } }, "GetObservationWithTimeWorker": { - "module_name": "GetObservationWithTimeWorker", - "module_path": "memory_scope/worker", + "name": "GetObservationWithTimeWorker", + "path": "memory_scope/worker", "kwargs": { "summary_messages_model": "qwen_max", "summary_messages_max_token": 500, @@ -87,8 +79,8 @@ } }, "GetObservationWorker": { - "module_name": "GetObservationWorker", - "module_path": "memory_scope/worker", + "name": "GetObservationWorker", + "path": "memory_scope/worker", "kwargs": { "summary_messages_model": "qwen_max", "summary_messages_max_token": 500, @@ -97,8 +89,8 @@ } }, "ContraRepeatWorker": { - "module_name": "ContraRepeatWorker", - "module_path": "memory_scope/worker", + "name": "ContraRepeatWorker", + "path": "memory_scope/worker", "kwargs": { "merge_obs_model": "qwen_max", "merge_obs_max_token": 500, @@ -107,8 +99,8 @@ } }, "FuseRerankWorker": { - "module_name": "FuseRerankWorker", - "module_path": "memory_scope/worker", + "name": "FuseRerankWorker", + "path": "memory_scope/worker", "kwargs": { "fuse_score_threshold": 0.1, "fuse_ratio_dict": { @@ -123,8 +115,8 @@ } }, "UpdateProfileWorker": { - "module_name": "UpdateProfileWorker", - "module_path": "memory_scope/worker", + "name": "UpdateProfileWorker", + "path": "memory_scope/worker", "kwargs": { "update_profile_threshold": 0.1, "update_profile_model": "qwen_max", @@ -135,8 +127,8 @@ } }, "GetReflectionWorker": { - "module_name": "GetReflectionWorker", - "module_path": "memory_scope/worker", + "name": "GetReflectionWorker", + "path": "memory_scope/worker", "kwargs": { "reflect_obs_cnt_threshold": 40, "reflect_num_questions": 3, @@ -147,8 +139,8 @@ } }, "UpdateInsightWorker": { - "module_name": "UpdateInsightWorker", - "module_path": "memory_scope/worker", + "name": "UpdateInsightWorker", + "path": "memory_scope/worker", "kwargs": { "update_insight_threshold": 0.1, "update_insight_model": "qwen_max", diff --git a/memory_scope/cli.py b/memory_scope/cli.py index fa024a08..208339c3 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,15 +1,14 @@ # 使用argparse库的示例 import argparse +import fire +from config import C, init +from chat.memory_chat import MemoryChat - -def main(): - parser = argparse.ArgumentParser(description="示例CLI程序") - parser.add_argument('--echo', help="输出传入的消息") - - args = parser.parse_args() - if args.echo: - print(f"收到的消息: {args.echo}") - +def main(config_path:str): + init(config_path) + + agent = MemoryChat() + agent.run() if __name__ == "__main__": - main() + fire.Fire(main) diff --git a/memory_scope/config.py b/memory_scope/config.py index dfca733b..d69c08c5 100644 --- a/memory_scope/config.py +++ b/memory_scope/config.py @@ -19,10 +19,8 @@ class Wrapper: raise AttributeError("Please run __init__ first!") return getattr(self._provider, key) - C = Wrapper() - def init(config_path: str): config = json.loads(config_path) C.register(config) diff --git a/memory_scope/pipeline/memory.py b/memory_scope/pipeline/memory.py index 769dab4b..062a9786 100644 --- a/memory_scope/pipeline/memory.py +++ b/memory_scope/pipeline/memory.py @@ -11,4 +11,8 @@ class MemoryServiceRequestModel(BaseModel): messages: List[Message] = Field(..., description="summary: 多轮对话的list,默认按照时间正序,最后一条是最新的; retrieve: 最后一条是query") - messages_pick_n: int = Field(1, description="summary:传需要总结的msg的个数;retrieve:不传") + user_profile: List[UserAttribute] = Field([], description="user_profile") + + ext_info: Dict[str, str] = Field({}, description="extra information") + + extra_user_attrs: List = [] \ No newline at end of file diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 849c9e0d..fdf0d8ef 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -144,10 +144,13 @@ def time_to_formatted_str(time: datetime | str | int | float = None, return return_str -def init_instance_by_config(config: dict, default_module_path=None, try_kwargs={}): +def init_instance_by_config(config: dict|object, default_module_path: str = None, try_kwargs: dict = {}, accept_types: type = None): + if isinstance(config, accept_types): + return config + import_module(config.pop("path", default_module_path)) clazz = getattr(module, config.pop("name")) try: - return clazz(**config,**try_kwargs) + return clazz(**config, **try_kwargs) except: return clazz(**config) \ No newline at end of file diff --git a/memory_scope/worker/memory_base_worker.py b/memory_scope/worker/memory_base_worker.py index 850b0d3d..0abef156 100644 --- a/memory_scope/worker/memory_base_worker.py +++ b/memory_scope/worker/memory_base_worker.py @@ -8,6 +8,7 @@ from node.user_attribute import UserAttribute from pipeline.memory import MemoryServiceRequestModel from worker.base_worker import BaseWorker from cli.cli_config import C +from utils.tool_functions import init_instance_by_config class MemoryBaseWorker(BaseWorker): def __init__(self, **kwargs): @@ -21,7 +22,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.request.messages_pick_n:] + messages = self.request.messages self.context_handler.set_context(MESSAGES, messages) return messages @@ -51,20 +52,32 @@ class MemoryBaseWorker(BaseWorker): return self.request.user.prompt @property - def emb_client(self): - return C.model_embedding + def client(self, model_type: str, model_name:str): + models = C.get(model_type) + models["model_name"] = init_instance_by_config( + config = models.get(model_name), + try_kwargs={ + "is_multi_thread": is_multi_thread, + "thread_pool": self.thread_pool + } + ) + return models["model_name"] + + @property + def emb_client(self, model_name: str): + self.client("model_embedding", model_name) @property def gene_client(self): - return C.model_generate + self.client("model_generate", model_name) @property def rerank_client(self): - return C.model_rerank + self.client("model_rerank", model_name) @property def es_client(self): - return C.db + self.client("db", model_name) @property def tenant_id(self): diff --git a/tests/test_memory.py b/tests/test_memory.py index a6ef0235..29a4cf53 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -3,13 +3,13 @@ import os import time from typing import List, Dict -from utils.logger import Logger -from constants.common_constants import NEW_USER_PROFILE, MODIFIED_MEMORIES, RELATED_MEMORIES -from enumeration.memory_method_enum import MemoryMethodEnum -from node.memory_node import MemoryNode -from node.user_attribute import UserAttribute -from pipeline.memory import MemoryServiceRequestModel -from service.memory_service_bailian import MemoryServiceBailian +from memory_scope.utils.logger import Logger +from memory_scope.constants.common_constants import NEW_USER_PROFILE, MODIFIED_MEMORIES, RELATED_MEMORIES +from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum +from memory_scope.node.memory_node import MemoryNode +from memory_scope.node.user_attribute import UserAttribute +from memory_scope.pipeline.memory import MemoryServiceRequestModel +from memory_scope.pipeline.memory_service import MemoryService """ 任务:随机生成一个用户的画像,随机种子0,并根据用户的画像虚拟一段用户和AI的对话。 @@ -139,10 +139,8 @@ for i, msg in enumerate(messages3): def summary_short(messages): - messages_pick_n = len(messages) request: MemoryServiceRequestModel = MemoryServiceRequestModel( messages=messages, - messages_pick_n=messages_pick_n, memory_id=memory_id, workspace_id=workspace_id, api_key=api_key, @@ -171,7 +169,6 @@ def summary_short(messages): def summary_long(): request: MemoryServiceRequestModel = MemoryServiceRequestModel( messages=[], - messages_pick_n=0, memory_id=memory_id, workspace_id=workspace_id, api_key=api_key, @@ -204,7 +201,6 @@ def summary_long(): def retrieve(messages): request: MemoryServiceRequestModel = MemoryServiceRequestModel( messages=messages, - messages_pick_n=1, memory_id=memory_id, workspace_id=workspace_id, api_key=api_key,