[dev] add base class to repo

This commit is contained in:
huangsen.huang 2024-06-19 14:44:42 +08:00
commit b98d4fa7f8
22 changed files with 560 additions and 83 deletions

View file

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

View file

@ -1,4 +1,5 @@
{
"method": "DashScopeEmbedding",
"model_name": "text-embedding-v2"
"model_name": "text-embedding-v2",
"method_type": "DashScopeEmbedding",
"clazz": "models.base_embedding_model"
}

View file

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

View file

@ -0,0 +1,7 @@
class MemoryService(object):
def __init__(self):
pass
def memory

View file

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

View 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:
"""

View file

View 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

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

View file

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

View 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

View 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

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

View 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

View file

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

View file

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

View 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]

View 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"""

View file

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

View file

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

View file

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

View file

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