mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
[dev] add base class to repo
This commit is contained in:
commit
b98d4fa7f8
22 changed files with 560 additions and 83 deletions
|
|
@ -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": ""
|
||||
}
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
{
|
||||
"method": "DashScopeEmbedding",
|
||||
"model_name": "text-embedding-v2"
|
||||
"model_name": "text-embedding-v2",
|
||||
"method_type": "DashScopeEmbedding",
|
||||
"clazz": "models.base_embedding_model"
|
||||
}
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
7
memory_scope/chat/memory_service.py
Normal file
7
memory_scope/chat/memory_service.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
class MemoryService(object):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def memory
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
38
memory_scope/db/base_db.py
Normal file
38
memory_scope/db/base_db.py
Normal file
|
|
@ -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:
|
||||
"""
|
||||
0
memory_scope/handler/__init__.py
Normal file
0
memory_scope/handler/__init__.py
Normal file
73
memory_scope/handler/config_handler.py
Normal file
73
memory_scope/handler/config_handler.py
Normal file
|
|
@ -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
|
||||
151
memory_scope/handler/pipeline.py
Normal file
151
memory_scope/handler/pipeline.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
24
memory_scope/models/base_embedding_model.py
Normal file
24
memory_scope/models/base_embedding_model.py
Normal file
|
|
@ -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
|
||||
24
memory_scope/models/base_generate_model.py
Normal file
24
memory_scope/models/base_generate_model.py
Normal file
|
|
@ -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
|
||||
104
memory_scope/models/base_model.py
Normal file
104
memory_scope/models/base_model.py
Normal file
|
|
@ -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)
|
||||
22
memory_scope/models/base_rank_model.py
Normal file
22
memory_scope/models/base_rank_model.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
20
memory_scope/models/response.py
Normal file
20
memory_scope/models/response.py
Normal file
|
|
@ -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]
|
||||
17
memory_scope/pipeline/operator.py
Normal file
17
memory_scope/pipeline/operator.py
Normal file
|
|
@ -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"""
|
||||
|
|
@ -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]
|
||||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue