diff --git a/config/config.json b/config/config.json index 899a84bd..ea868fe3 100644 --- a/config/config.json +++ b/config/config.json @@ -9,5 +9,6 @@ "model_embedding": "config/model/dash_embedding.json", "model_rerank": "config/model/dash_rerank.json", "model_generate": "config/model/dash_generate.json", - "db": "" + "db": "", + "dash_api_key": "" } \ No newline at end of file diff --git a/config/model/dash_embedding.json b/config/model/dash_embedding.json index 8ceaf9b5..70408e30 100644 --- a/config/model/dash_embedding.json +++ b/config/model/dash_embedding.json @@ -1,4 +1,5 @@ { - "method": "DashScopeEmbedding", - "model_name": "text-embedding-v2" + "model_name": "text-embedding-v2", + "method_type": "DashScopeEmbedding", + "clazz": "models.base_embedding_model" } \ No newline at end of file diff --git a/config/worker.json b/config/worker.json index 87724ef3..32c31749 100644 --- a/config/worker.json +++ b/config/worker.json @@ -51,11 +51,8 @@ }, "ExtractTimeWorker": { "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 + "clazz": "worker.summary_long.get_insight", + "parse_time_model": "qwen_1_8_parse_time_service" }, "InfoFilterWorker": { "name": "InfoFilterWorker", diff --git a/memory_scope/chat/memory_service.py b/memory_scope/chat/memory_service.py new file mode 100644 index 00000000..c750621b --- /dev/null +++ b/memory_scope/chat/memory_service.py @@ -0,0 +1,7 @@ + +class MemoryService(object): + + def __init__(self): + pass + + def memory \ No newline at end of file diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 208339c3..3de177eb 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,14 +1,31 @@ -# 使用argparse库的示例 -import argparse import fire -from config import C, init -from chat.memory_chat import MemoryChat -def main(config_path:str): - init(config_path) - +from memory_scope.chat.memory_chat import MemoryChat +from memory_scope.handler.config_handler import ConfigHandler + +""" +1. fire read config +2. init config + 1. global configs: global + db + llm + monitor + 2. worker config list: worker + llm +3. init db +4. init workers, global+worker +5. init llms +6. init monitor +7. new Agent,add db workers llms monitor + 1. new memory service + 1. new pipeline + 2. chat + 3. +""" + + +def main(config_path: str): + config_handler = ConfigHandler(config_path) + agent = MemoryChat() agent.run() + if __name__ == "__main__": fire.Fire(main) diff --git a/memory_scope/db/base_db.py b/memory_scope/db/base_db.py new file mode 100644 index 00000000..59272d0a --- /dev/null +++ b/memory_scope/db/base_db.py @@ -0,0 +1,38 @@ +from abc import ABCMeta, abstractmethod + +from memory_scope.models.base_model import BaseModel + + +class BaseDBClient(metaclass=ABCMeta): + + def __init__(self, index_name: str, embedding_model: BaseModel, content_key: str = "text", **kwargs): + self.index_name: str = index_name + self.embedding_model: BaseModel = embedding_model + self.content_key: str = content_key + self.kwargs: dict = kwargs + + @abstractmethod + def retrieve(self, text: str, limit_size: int): + """ + :param text: + :param limit_size: + :return: + """ + + @abstractmethod + def insert(self, text: str): + """ TODO 是否overwrite + :return: + """ + + @abstractmethod + def insert_batch(self): + """ + :return: + """ + + @abstractmethod + def delete(self): + """ + :return: + """ \ No newline at end of file diff --git a/memory_scope/handler/__init__.py b/memory_scope/handler/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memory_scope/handler/config_handler.py b/memory_scope/handler/config_handler.py new file mode 100644 index 00000000..96251033 --- /dev/null +++ b/memory_scope/handler/config_handler.py @@ -0,0 +1,73 @@ +import json +import os.path +from typing import Dict + +from memory_scope.models.base_model import BaseModel +from memory_scope.utils.tool_functions import init_instance_by_config_v2 +from memory_scope.worker.base_worker import BaseWorker + + +class ConfigHandler(object): + + def __init__(self, path: str): + self.config_name: str = os.path.basename(path) + self.config_base_dir: str = os.path.dirname(path) + + with open(path) as f: + config = json.load(f) + + self.global_configs: Dict[str, str] = {} + self.worker_dict: Dict[str, BaseWorker] = {} + self.model_dict: Dict[str, BaseModel] = {} + + self._init_global_config(config["global"]) + self.worker_base_dir = self.global_configs.get("worker_base_dir", "config") + self.model_base_dir = self.global_configs.get("model_base_dir", "config/model") + self._init_workers(config["workers"]) + self._init_db(config["db"]) + self._init_chat_model(config["chat_model"]) + self._init_monitor(config["monitor"]) + + def _init_global_config(self, global_configs: Dict[str, str]): + """set global_configs & set apikey into env + """ + + def _init_workers(self, worker_config_name: str): + """ load worker config & init workers + """ + with open(os.path.join(self.config_base_dir, worker_config_name)) as f: + worker_config_dict = json.load(f) + + for worker_name, worker_config in worker_config_dict.items(): + if worker_name in self.worker_dict: + continue + self.worker_dict[worker_name] = init_instance_by_config_v2(worker_config, + default_clazz_path=self.worker_base_dir, + suffix_name="worker", + **self.global_configs, + **worker_config) + + self._init_model(worker_config.get("embedding_model")) + self._init_model(worker_config.get("generation_model")) + self._init_model(worker_config.get("rank_model")) + + def _init_model(self, model_name: str): + if not model_name or model_name in self.model_dict: + return + + with open(os.path.join(self.model_base_dir, model_name)) as f: + model_config = json.load(f) + self.model_dict[model_name] = init_instance_by_config_v2(model_config, + default_clazz_path=self.model_base_dir, + suffix_name="", + **model_config) + + def _init_db(self, db_config: dict): + pass + + def _init_chat_model(self, chat_model_config: dict): + chat_model_name = chat_model_config["name"] + self._init_model(chat_model_name) + + def _init_monitor(self, monitor_config: dict): + pass diff --git a/memory_scope/handler/pipeline.py b/memory_scope/handler/pipeline.py new file mode 100644 index 00000000..0bd2dc55 --- /dev/null +++ b/memory_scope/handler/pipeline.py @@ -0,0 +1,151 @@ +import json +import re +import time +from concurrent.futures import ThreadPoolExecutor, as_completed +from importlib import import_module +from itertools import zip_longest +from typing import Dict, Any + +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 pipeline.memory import MemoryServiceRequestModel +from cli.cli_config import C +from utils.tool_functions import init_instance_by_config + + +class Pipeline(object): + def __init__(self, method: MemoryMethodEnum): + self.method = method + self.context_handler = ContextHandler() + + # 线程池 + self.thread_pool = ThreadPoolExecutor(max_workers=C.thread_pool_max_count) + + # 全部初始化的worker + self.worker_dict: Dict[str, BaseWorker] = {} + + # 日志 + 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: + return init_instance_by_config( + config = C.worker.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: + worker = self.worker_dict[worker_name] + # 执行子类实现的_run函数 + worker.run() + # 保存worker的运行信息 + self.run_infos.append(worker.run_info_dict) + # 结束pipeline + if not worker.continue_run: + return False + return True + + @timer + def print_and_init_worker(self, pipeline_list: list[list]): + self.logger.info("----- Pipeline Begin -----") + i: int = 0 + for pipeline_part in pipeline_list: + if len(pipeline_part) == 1: + for w in pipeline_part[0]: + self.logger.info(f"stage{i}: {w}") + self.worker_dict[w] = self.get_worker(w) + i += 1 + else: + for w_zip in zip_longest(*pipeline_part, fillvalue="-"): + self.logger.info(f"stage{i}: {' | '.join(w_zip)}") + i += 1 + for w in w_zip: + if w == "-": + continue + self.worker_dict[w] = self.get_worker(w, is_multi_thread=True) + self.logger.info("----- Pipeline End -----") + + 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 = 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 + pattern = r'(\[[^\]]*\]|[^,]+)' + pipeline_split = re.findall(pattern, pipeline_str) + + pipeline_list = [] + for pipeline_part in pipeline_split: + # e.g., [d,e,f|g,h] + pipeline_part = pipeline_part.strip() + if '[' in pipeline_part or ']' in pipeline_part: + pipeline_part = pipeline_part.replace('[', '').replace(']', '') + + # e.g., ["d,e,f", "g,h"] + line_split = [x.strip() for x in pipeline_part.split("|") if x] + if len(line_split) <= 0: + continue + + # e.g., ["d","e","f"] + pipeline_list.append([x.split(",") for x in line_split]) + + return pipeline_list + + def run(self): + # run workers in multi threads + with self.thread_pool, Timer("ALL_PIPELINE"): + for pipeline_part in self.pipeline_list: + if len(pipeline_part) == 1: + if not self.worker_run(pipeline_part[0]): + break + elif self.max_workers == 1: + for worker_list in pipeline_part: + self.worker_run(worker_list) + else: + t_list = [] + for worker_list in pipeline_part: + time.sleep(0.001) + t_list.append(self.thread_pool.submit(self.worker_run, worker_list)) + + flag = True + for future in as_completed(t_list): + if not future.result(): + flag = False + break + if not flag: + break + + # 获取ext_info + ext_info = self.get_context(RESPONSE_EXT_INFO) + if ext_info is None: + ext_info = {} + self.context_handler.set_context(RESPONSE_EXT_INFO, ext_info) + + # 保存 run_info_list + ext_info["run_infos"] = json.dumps(self.run_infos, ensure_ascii=False) diff --git a/memory_scope/models/__init__.py b/memory_scope/models/__init__.py index 73d7895f..35395203 100644 --- a/memory_scope/models/__init__.py +++ b/memory_scope/models/__init__.py @@ -1,18 +1,5 @@ -from utils.registry import Registry +from llama_index.embeddings.dashscope import DashScopeEmbedding -from llama_index.embeddings.dashscope import ( - DashScopeEmbedding, -) -# from llama_index.postprocessor.dashscope_rerank import DashScopeRerank +from memory_scope.utils.registry import Registry -from llama_index.llms.dashscope import DashScope # type: ignore - - -EMB = Registry('embedding') -EMB.register_module(DashScopeEmbedding) - -# RERANKER = Registry('reranker') -# RERANKER.register_module(DashScopeRerank) - -LLM = Registry('llm') -LLM.register_module(DashScope) +MODEL_REGISTRY = Registry("models") diff --git a/memory_scope/models/base_embedding_model.py b/memory_scope/models/base_embedding_model.py new file mode 100644 index 00000000..3e6d30cf --- /dev/null +++ b/memory_scope/models/base_embedding_model.py @@ -0,0 +1,24 @@ +from llama_index.embeddings.dashscope import DashScopeEmbedding + +from memory_scope.models import MODEL_REGISTRY +from memory_scope.models.base_model import BaseModel +from memory_scope.models.response import ModelResponse, ModelResponseGen + + +class BaseEmbeddingModel(BaseModel): + MODEL_REGISTRY.batch_register([ + DashScopeEmbedding, + ]) + + def before_call(self, **kwargs) -> None: + pass + + def after_call(self, model_response: ModelResponse | ModelResponseGen, + **kwargs) -> ModelResponse | ModelResponseGen: + pass + + def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: + pass + + async def _async_call(self, **kwargs) -> ModelResponse: + pass diff --git a/memory_scope/models/base_generate_model.py b/memory_scope/models/base_generate_model.py new file mode 100644 index 00000000..2f0c3dbd --- /dev/null +++ b/memory_scope/models/base_generate_model.py @@ -0,0 +1,24 @@ +from llama_index.llms.dashscope import DashScope as DashScopeLLM + +from memory_scope.models import MODEL_REGISTRY +from memory_scope.models.base_model import BaseModel +from memory_scope.models.response import ModelResponse, ModelResponseGen + + +class BaseGenerateModel(BaseModel): + MODEL_REGISTRY.batch_register([ + DashScopeLLM + ]) + + def before_call(self, **kwargs) -> None: + pass + + def after_call(self, model_response: ModelResponse | ModelResponseGen, + **kwargs) -> ModelResponse | ModelResponseGen: + pass + + def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: + pass + + async def _async_call(self, **kwargs) -> ModelResponse: + pass diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py new file mode 100644 index 00000000..af415976 --- /dev/null +++ b/memory_scope/models/base_model.py @@ -0,0 +1,104 @@ +import asyncio +import inspect +import time +from abc import abstractmethod, ABCMeta + +from memory_scope.models import MODEL_REGISTRY +from memory_scope.models.response import ModelResponse, ModelResponseGen +from memory_scope.utils.logger import Logger +from memory_scope.utils.timer import Timer + + +class BaseModel(metaclass=ABCMeta): + + def __init__(self, + model_name: str, + method_type: str, + timeout: int = None, + max_retries: int = 3, + retry_interval: float = 1.0, + kwargs_filter: bool = True, + **kwargs): + + self.model_name: str = model_name + self.method_type: str = method_type + self.timeout: int = timeout + self.max_retries: int = max_retries + self.retry_interval: float = retry_interval + self.kwargs: dict = kwargs + + self.data = {} + self.logger = Logger.get_logger() + + obj_cls = MODEL_REGISTRY.get(self.method_type) + if not obj_cls: + raise RuntimeError(f"method_type={self.method_type} is not supported!") + + if kwargs_filter: + allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys()) + kwargs = {key: value for key, value in kwargs.items() if key in allowed_kwargs} + + self.model = obj_cls(**kwargs) + + @abstractmethod + def before_call(self, **kwargs) -> None: + """prepare data before call + :param kwargs: + :return: + """ + + @abstractmethod + def after_call(self, model_response: ModelResponse | ModelResponseGen, + **kwargs) -> ModelResponse | ModelResponseGen: + """ + :param model_response: + :param kwargs: + :return: + """ + + @abstractmethod + def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: + """ + :param kwargs: + :return: + """ + + def call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: + """ + :param stream: only llm needs stream + :param kwargs: + :return: + """ + self.before_call(stream=stream, **kwargs) + with Timer(self.__class__.__name__, log_time=False) as t: + for i in range(self.max_retries): + model_response = self._call(stream=stream, **kwargs) + if not model_response.status and not stream: + self.logger.warning(f"call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} " + f"details={model_response.details}", stacklevel=2) + time.sleep(i * self.retry_interval) + else: + return self.after_call(stream=stream, model_response=model_response, **kwargs) + + @abstractmethod + async def _async_call(self, **kwargs) -> ModelResponse: + """ + :param kwargs: + :return: + """ + + async def async_call(self, **kwargs) -> ModelResponse: + """ 异步不需要stream + :param kwargs: + :return: + """ + self.before_call(**kwargs) + with Timer(self.__class__.__name__, log_time=False) as t: + for i in range(self.max_retries): + model_response = await self._async_call(**kwargs) + if not model_response.status: + self.logger.warning(f"async_call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} " + f"details={model_response.details}", stacklevel=2) + await asyncio.sleep(i * self.retry_interval) + else: + return self.after_call(model_response, **kwargs) diff --git a/memory_scope/models/base_rank_model.py b/memory_scope/models/base_rank_model.py new file mode 100644 index 00000000..ce6326cd --- /dev/null +++ b/memory_scope/models/base_rank_model.py @@ -0,0 +1,22 @@ +from memory_scope.models import MODEL_REGISTRY +from memory_scope.models.base_model import BaseModel +from memory_scope.models.response import ModelResponse, ModelResponseGen + + +class BaseRankModel(BaseModel): + MODEL_REGISTRY.batch_register([ + + ]) + + def before_call(self, **kwargs) -> None: + pass + + def after_call(self, model_response: ModelResponse | ModelResponseGen, + **kwargs) -> ModelResponse | ModelResponseGen: + pass + + def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: + pass + + async def _async_call(self, **kwargs) -> ModelResponse: + pass diff --git a/memory_scope/models/dash_generate_client.py b/memory_scope/models/dash_generate_client.py index 6aefbe0f..aefbbe5f 100644 --- a/memory_scope/models/dash_generate_client.py +++ b/memory_scope/models/dash_generate_client.py @@ -88,7 +88,7 @@ class LLILLM(LLIClient): input_type: llama_input, } - def after_call(self, response_obj, **kwargs): + def after_call(self, response_obj: ChatResponse | CompletionResponse, **kwargs) -> str: self.logger.debug(f"response_obj={response_obj}") if isinstance(response_obj, CompletionResponse): return response_obj.text diff --git a/memory_scope/models/dash_rerank_client.py b/memory_scope/models/dash_rerank_client.py index 480d484f..4736190e 100644 --- a/memory_scope/models/dash_rerank_client.py +++ b/memory_scope/models/dash_rerank_client.py @@ -85,7 +85,7 @@ class LLIReRank(LLIClient): } - def after_call(self, nodes, **kwargs): + def after_call(self, nodes: List[NodeWithScore], **kwargs) -> List[dict]: results = [] for node in nodes: results.append(dict(relevance_score=node.score, diff --git a/memory_scope/models/response.py b/memory_scope/models/response.py new file mode 100644 index 00000000..11890183 --- /dev/null +++ b/memory_scope/models/response.py @@ -0,0 +1,20 @@ +from typing import Generator, List, Dict + +from pydantic import BaseModel, Field + + +class ModelResponse(BaseModel): + text: str = Field("", description="") + + embedding_results: Dict[int, List[float]] | List[float] = Field([], description="") + + rank_scores: Dict[int, float] = Field({}, description="") + + model_type: str = Field("", description="") + + status: bool = Field(True, description="") + + details: str = Field("", description="") + + +ModelResponseGen = Generator[ModelResponse, None, None] diff --git a/memory_scope/pipeline/operator.py b/memory_scope/pipeline/operator.py new file mode 100644 index 00000000..ac1b317b --- /dev/null +++ b/memory_scope/pipeline/operator.py @@ -0,0 +1,17 @@ +# -*- coding: utf-8 -*- +"""A common base class for Pipeline""" +from abc import ABC +from abc import abstractmethod + + +class Operator(ABC): + """ + Abstract base class `Operator` defines a protocol for classes that + implement callable behavior. + The class is designed to be subclassed with an overridden `__call__` + method that specifies the execution logic for the operator. + """ + + @abstractmethod + def __call__(self) -> None: + """Calling function""" diff --git a/memory_scope/utils/logger.py b/memory_scope/utils/logger.py index 684ba6b7..671c521b 100644 --- a/memory_scope/utils/logger.py +++ b/memory_scope/utils/logger.py @@ -2,8 +2,6 @@ import logging from logging.handlers import RotatingFileHandler from pathlib import Path -from constants.common_constants import MEMORY - # remove %(thread)s .%(funcName)s LOG_FORMAT = "%(asctime)s %(levelname)s %(trace_id)s %(module)s:%(lineno)d] %(message)s" DATE_FORMAT = "%Y-%m-%d %H:%M:%S" @@ -90,13 +88,9 @@ class Logger(logging.Logger): if LOGGER_DICT: name = list(LOGGER_DICT.keys())[0] else: - name = MEMORY + name = "default" if name not in LOGGER_DICT: LOGGER_DICT[name] = Logger(name=name, **kwargs) - return LOGGER_DICT[name] - - @classmethod - def get_memory_logger(cls, **kwargs): - return cls.get_logger(MEMORY, **kwargs) + return LOGGER_DICT[name] \ No newline at end of file diff --git a/memory_scope/utils/registry.py b/memory_scope/utils/registry.py index a662180f..0797f9a5 100644 --- a/memory_scope/utils/registry.py +++ b/memory_scope/utils/registry.py @@ -3,52 +3,30 @@ Registry for different modules. Init class according to the class name and verify the input parameters. """ import inspect +from typing import Dict, Any, List -class Registry: - def __init__(self, name): - self.name = name - self.module_dict = dict() +class Registry(object): + def __init__(self, name: str): + self.name: str = name + self.module_dict: Dict[str, Any] = {} - def register_module(self, module, module_name=None): + def register(self, module: Any, module_name: str = None): if module_name is None: module_name = module.__name__ + if module_name in self.module_dict: raise KeyError(f'{module_name} is already registered in {self.name}') self.module_dict[module_name] = module - def get_module(self, module_name): + def batch_register(self, modules: List[Any]): + module_name_dict = {m.__name__: m for m in modules} + self.module_dict.update(module_name_dict) + + def __getitem__(self, module_name: str): assert module_name in self.module_dict, f'{module_name} not found in {self.name}' return self.module_dict[module_name] -def build_from_cfg(config, registry, default_args: dict = None, skip_param_check=False): - - if default_args is None: - default_args = {} - - args = config.copy() - method_type = args.pop('method') - #params = args.get("parameters", {}) or default_args - params = args - if isinstance(method_type, str): - obj_cls = registry.get_module(method_type) - else: - raise TypeError( - f'type must be a str or valid type, but got {type(method_type)}') - - allowed_params = list(inspect.signature(obj_cls.__init__).parameters.keys()) - print(allowed_params, params) - if not skip_param_check: - filter_params = {key: value for key, value in params.items() if key in allowed_params} - else: - filter_params = params - # print( - # f"Registry {registry.name}, " - # f"allowed parameters {allowed_params}, filter parameters {filter_params}", - # flush=True - # ) - print(filter_params) - return obj_cls(**filter_params) diff --git a/memory_scope/utils/timer.py b/memory_scope/utils/timer.py index 0c576363..9e1e3e80 100644 --- a/memory_scope/utils/timer.py +++ b/memory_scope/utils/timer.py @@ -6,12 +6,12 @@ date: 20221106 import time -from utils.logger import Logger +from memory_scope.utils.logger import Logger class Timer(object): - def __init__(self, name: str, log_time: bool = True, use_ms: bool = True, **kwargs): + def __init__(self, name: str, log_time: bool = True, use_ms: bool = False, **kwargs): self.name: str = name self.log_time: bool = log_time self.use_ms: bool = use_ms @@ -61,11 +61,12 @@ class Timer(object): self.logger.info(line, stacklevel=3) - def get_cost_info(self): + @property + def cost_str(self): if self.use_ms: - return f"cost={self.cost:.1f}ms" + return f"{self.cost:.1f}ms" else: - return f"cost={self.cost:.4f}s" + return f"{self.cost:.4f}s" def timer(func): diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index fdf0d8ef..6d8aeb41 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -1,10 +1,10 @@ -import os import re from datetime import datetime +from importlib import import_module from typing import Dict, List from constants.common_constants import WEEKDAYS -from importlib import import_module + def under_line_to_hump(underline_str): sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str) @@ -153,4 +153,25 @@ def init_instance_by_config(config: dict|object, default_module_path: str = None try: return clazz(**config, **try_kwargs) except: - return clazz(**config) \ No newline at end of file + return clazz(**config) + + +def init_instance_by_config_v2(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs): + clazz_path = config.pop("clazz") + if not clazz_path: + raise RuntimeError("empty clazz_path!") + clazz_name_split = clazz_path.split(".") + clazz_name: str = clazz_name_split[-1] + if suffix_name and not clazz_name.endswith(suffix_name): + clazz_name = f"{clazz_name}_{suffix_name}" + + # 构造path + clazz_paths = [] + if default_clazz_path: + clazz_paths.append(default_clazz_path) + clazz_paths.extend(clazz_name_split[:-1]) + clazz_paths.append(clazz_name) + module = import_module(".".join(clazz_paths)) + + cls_name = under_line_to_hump(clazz_name) + return getattr(module, cls_name)(**kwargs)