2024-07-07 01:58:23 +08:00
|
|
|
|
import os
|
|
|
|
|
|
import json
|
|
|
|
|
|
import yaml
|
2024-07-28 16:16:52 +08:00
|
|
|
|
from src.utils.logging_config import setup_logger
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
2024-07-14 21:44:27 +08:00
|
|
|
|
logger = setup_logger("Config")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class SimpleConfig(dict):
|
|
|
|
|
|
|
|
|
|
|
|
def __key(self, key):
|
|
|
|
|
|
return "" if key is None else key.lower()
|
|
|
|
|
|
|
|
|
|
|
|
def __str__(self):
|
|
|
|
|
|
return json.dumps(self)
|
|
|
|
|
|
|
|
|
|
|
|
def __setattr__(self, key, value):
|
|
|
|
|
|
self[self.__key(key)] = value
|
|
|
|
|
|
|
|
|
|
|
|
def __getattr__(self, key):
|
|
|
|
|
|
return self.get(self.__key(key))
|
|
|
|
|
|
|
|
|
|
|
|
def __getitem__(self, key):
|
|
|
|
|
|
return super().get(self.__key(key))
|
|
|
|
|
|
|
|
|
|
|
|
def __setitem__(self, key, value):
|
|
|
|
|
|
return super().__setitem__(self.__key(key), value)
|
|
|
|
|
|
|
2024-07-25 20:30:28 +08:00
|
|
|
|
def __dict__(self):
|
|
|
|
|
|
return {k: v for k, v in self.items()}
|
|
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
class Config(SimpleConfig):
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(self, filename=None):
|
|
|
|
|
|
super().__init__()
|
2024-07-25 20:30:28 +08:00
|
|
|
|
self._config_items = {}
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
2024-07-17 18:52:20 +08:00
|
|
|
|
### >>> 默认配置
|
|
|
|
|
|
# 可以在 config/base.yaml 中覆盖
|
2024-07-25 20:30:28 +08:00
|
|
|
|
self.add_item("mode", default="cli", des="运行模式", choices=["cli", "api"])
|
|
|
|
|
|
self.add_item("stream", default=True, des="是否开启流式输出")
|
2024-07-28 16:16:52 +08:00
|
|
|
|
self.add_item("save_dir", default="saves", des="保存目录")
|
2024-07-17 18:52:20 +08:00
|
|
|
|
# 功能选项
|
2024-07-29 01:00:02 +08:00
|
|
|
|
self.add_item("enable_reranker", default=False, des="是否开启重排序")
|
2024-07-31 20:22:05 +08:00
|
|
|
|
self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
|
|
|
|
|
|
self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎")
|
2024-07-17 18:52:20 +08:00
|
|
|
|
|
|
|
|
|
|
# 模型配置
|
|
|
|
|
|
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
|
|
|
|
|
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
|
2024-08-07 17:24:46 +08:00
|
|
|
|
self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"])
|
2024-07-31 20:22:05 +08:00
|
|
|
|
self.add_item("model_name", default=None, des="模型名称")
|
2024-07-25 20:30:28 +08:00
|
|
|
|
self.add_item("embed_model", default="bge-large-zh-v1.5", des="Embedding 模型", choices=["bge-large-zh-v1.5", "zhipu"])
|
|
|
|
|
|
self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"])
|
|
|
|
|
|
self.add_item("model_local_paths", default={}, des="本地模型路径")
|
2024-07-17 18:52:20 +08:00
|
|
|
|
### <<< 默认配置结束
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
2024-07-28 16:16:52 +08:00
|
|
|
|
self.filename = filename or os.path.join(self.save_dir, "config", "config.yaml")
|
2024-07-29 01:00:02 +08:00
|
|
|
|
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
|
2024-07-28 16:16:52 +08:00
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
self.load()
|
2024-07-09 05:04:20 +08:00
|
|
|
|
self.handle_self()
|
|
|
|
|
|
|
2024-07-25 20:30:28 +08:00
|
|
|
|
def add_item(self, key, default, des=None, choices=None):
|
|
|
|
|
|
self.__setattr__(key, default)
|
|
|
|
|
|
self._config_items[key] = {
|
|
|
|
|
|
"default": default,
|
|
|
|
|
|
"des": des,
|
|
|
|
|
|
"choices": choices
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2024-07-09 05:04:20 +08:00
|
|
|
|
def handle_self(self):
|
|
|
|
|
|
### handle local model
|
|
|
|
|
|
model_root_dir = os.getenv("MODEL_ROOT_DIR", "pretrained_models")
|
2024-07-25 20:30:28 +08:00
|
|
|
|
if self.model_local_paths is not None:
|
|
|
|
|
|
for model, model_rel_path in self.model_local_paths.items():
|
|
|
|
|
|
# 如果 model_rel_path 不是绝对路径,那么拼接 model_root_dir
|
|
|
|
|
|
if not model_rel_path.startswith("/"):
|
|
|
|
|
|
self.model_local_paths[model] = os.path.join(model_root_dir, model_rel_path)
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
|
self.model_names = MODEL_NAMES
|
|
|
|
|
|
|
|
|
|
|
|
if self.model_name not in self.model_names[self.model_provider]:
|
|
|
|
|
|
logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name")
|
|
|
|
|
|
self.model_name = self.model_names[self.model_provider][0]
|
|
|
|
|
|
|
|
|
|
|
|
default_model_name = self.model_names[self.model_provider][0]
|
|
|
|
|
|
self.model_name = self.get("model_name") or default_model_name
|
|
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
def load(self):
|
2024-07-17 18:52:20 +08:00
|
|
|
|
"""根据传入的文件覆盖掉默认配置"""
|
2024-07-22 00:00:54 +08:00
|
|
|
|
logger.info(f"Loading config from {self.filename}")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
if self.filename is not None and os.path.exists(self.filename):
|
|
|
|
|
|
if self.filename.endswith(".json"):
|
|
|
|
|
|
with open(self.filename, 'r') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
content = f.read()
|
|
|
|
|
|
if content:
|
|
|
|
|
|
self.update(json.loads(content))
|
|
|
|
|
|
else:
|
|
|
|
|
|
print(f"{self.filename} is empty.")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
elif self.filename.endswith(".yaml"):
|
|
|
|
|
|
with open(self.filename, 'r') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
content = f.read()
|
|
|
|
|
|
if content:
|
|
|
|
|
|
self.update(yaml.safe_load(content))
|
|
|
|
|
|
else:
|
|
|
|
|
|
print(f"{self.filename} is empty.")
|
|
|
|
|
|
else:
|
|
|
|
|
|
logger.warning(f"Unknown config file type {self.filename}")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
else:
|
2024-07-28 16:16:52 +08:00
|
|
|
|
logger.warning(f"\n\n{'='*70}\n{'Config file not found':^70}\n{'You can custum your config in `' + self.filename + '`':^70}\n{'='*70}\n\n")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
def save(self):
|
2024-07-22 00:00:54 +08:00
|
|
|
|
logger.info(f"Saving config to {self.filename}")
|
|
|
|
|
|
if self.filename is None:
|
|
|
|
|
|
logger.warning("Config file is not specified, save to default config/base.yaml")
|
2024-07-28 16:16:52 +08:00
|
|
|
|
self.filename = os.path.join(self.save_dir, "config", "config.yaml")
|
2024-07-29 01:00:02 +08:00
|
|
|
|
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
|
|
|
|
|
if self.filename.endswith(".json"):
|
|
|
|
|
|
with open(self.filename, 'w+') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
json.dump(self.__dict__(), f, indent=4, ensure_ascii=False)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
elif self.filename.endswith(".yaml"):
|
|
|
|
|
|
with open(self.filename, 'w+') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
yaml.dump(self.__dict__(), f, indent=2, allow_unicode=True)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
else:
|
|
|
|
|
|
logger.warning(f"Unknown config file type {self.filename}, save as json")
|
|
|
|
|
|
with open(self.filename, 'w+') as f:
|
2024-07-07 01:58:23 +08:00
|
|
|
|
json.dump(self, f, indent=4)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
|
logger.info(f"Config file {self.filename} saved")
|
|
|
|
|
|
|
|
|
|
|
|
MODEL_NAMES = {
|
|
|
|
|
|
# https://platform.deepseek.com/api-docs/zh-cn/pricing
|
|
|
|
|
|
"deepseek": [
|
|
|
|
|
|
"deepseek-chat",
|
|
|
|
|
|
"deepseek-coder"
|
|
|
|
|
|
],
|
|
|
|
|
|
|
|
|
|
|
|
# https://open.bigmodel.cn/dev/api glm-4-0520、glm-4 、glm-4-air、glm-4-airx、 glm-4-flash
|
|
|
|
|
|
"zhipu": [
|
|
|
|
|
|
"glm-4",
|
|
|
|
|
|
"glm-4-0520",
|
|
|
|
|
|
"glm-4-air",
|
|
|
|
|
|
"glm-4-airx",
|
|
|
|
|
|
"glm-4-flash"
|
|
|
|
|
|
],
|
|
|
|
|
|
|
|
|
|
|
|
# {'ERNIE-4.0-8K-0104', 'ERNIE-Lite-8K-0308', 'ERNIE-Speed-128K', 'ERNIE-3.5-128K(预览版)', 'Yi-34B-Chat', 'ERNIE-4.0-8K-Preview-0518', 'ERNIE-Bot-4', 'ERNIE-3.5-128K', 'ChatGLM2-6B-32K', 'ERNIE-3.5-8K', 'EB-turbo-AppBuilder', 'ERNIE-Lite-AppBuilder-8K', 'ERNIE-4.0-8K-0329', 'AquilaChat-7B', 'Gemma-7B-it', 'Qianfan-Chinese-Llama-2-70B', 'Mixtral-8x7B-Instruct', 'Gemma-7B-It', 'ERNIE Speed-AppBuilder', 'ERNIE-Function-8K', 'ERNIE-4.0-8K-preview', 'ERNIE-Bot', 'Qianfan-BLOOMZ-7B-compressed', 'ERNIE-4.0-8K', 'BLOOMZ-7B', 'ERNIE-Character-8K', 'ERNIE-3.5-8K-0205', 'ERNIE-4.0-8K-0613', 'Llama-2-70B-Chat', 'ERNIE-Character-Fiction-8K', 'ERNIE-4.0-8K-Preview', 'ERNIE-3.5-8K-Preview', 'ERNIE-Speed', 'ERNIE-Tiny-8K', 'ERNIE-4.0-Turbo-8K-Preview', 'Meta-Llama-3-8B', 'ERNIE-4.0-8K-Latest', 'ERNIE 3.5', 'XuanYuan-70B-Chat-4bit', 'Llama-2-13B-Chat', 'ERNIE-Bot-turbo', 'ERNIE-3.5-8K-0613', 'ERNIE-Lite-AppBuilder-8K-0614', 'ERNIE-4.0-preview', 'Llama-2-7B-Chat', 'Qianfan-Chinese-Llama-2-13B', 'ERNIE-Bot-turbo-AI', 'Meta-Llama-3-70B', 'ERNIE-Functions-8K', 'ERNIE-Lite-8K-0922(原ERNIE-Bot-turbo-0922)', 'ERNIE Speed', 'ERNIE-3.5-preview', 'Qianfan-Chinese-Llama-2-7B', 'ERNIE-Speed-8K', 'ERNIE-Lite-8K-0922', 'ChatLaw', 'ERNIE-3.5-8K-0329', 'ERNIE-4.0-Turbo-8K', 'ERNIE-3.5-8K-preview', 'ERNIE-Lite-8K'}
|
|
|
|
|
|
"qianfan": [
|
|
|
|
|
|
"ERNIE-Speed",
|
|
|
|
|
|
"ERNIE-Speed-8K",
|
|
|
|
|
|
"ERNIE-Speed-128K",
|
|
|
|
|
|
"ERNIE-Tiny-8K",
|
|
|
|
|
|
"ERNIE-Lite-8K",
|
|
|
|
|
|
"ERNIE-4.0-8K-Latest"
|
|
|
|
|
|
"Yi-34B-Chat",
|
|
|
|
|
|
],
|
|
|
|
|
|
|
|
|
|
|
|
"vllm": [
|
|
|
|
|
|
"vllm",
|
|
|
|
|
|
],
|
|
|
|
|
|
|
|
|
|
|
|
# https://bailian.console.aliyun.com/?switchAgent=10226727&productCode=p_efm#/model-market
|
|
|
|
|
|
"dashscope": [
|
|
|
|
|
|
"qwen-long",
|
|
|
|
|
|
"qwen2-7b-instruct",
|
|
|
|
|
|
"qwen2-1.5b-instruct",
|
|
|
|
|
|
"llama3.1-8b-instruct",
|
|
|
|
|
|
"llama3-8b-instruct",
|
|
|
|
|
|
"llama3.1-405b-instruct",
|
|
|
|
|
|
"baichuan2-7b-chat-v1",
|
|
|
|
|
|
"qwen2-0.5b-instruct"
|
|
|
|
|
|
]
|
|
|
|
|
|
}
|