ForcePilot/src/config/__init__.py

176 lines
7.8 KiB
Python
Raw Normal View History

import os
import json
import yaml
from src.utils.logging_config import setup_logger
2024-07-14 21:44:27 +08:00
logger = setup_logger("Config")
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)
def __dict__(self):
return {k: v for k, v in self.items()}
class Config(SimpleConfig):
def __init__(self, filename=None):
super().__init__()
self._config_items = {}
2024-07-17 18:52:20 +08:00
### >>> 默认配置
# 可以在 config/base.yaml 中覆盖
self.add_item("mode", default="cli", des="运行模式", choices=["cli", "api"])
self.add_item("stream", default=True, des="是否开启流式输出")
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="模型名称")
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
### <<< 默认配置结束
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)
self.load()
self.handle_self()
def add_item(self, key, default, des=None, choices=None):
self.__setattr__(key, default)
self._config_items[key] = {
"default": default,
"des": des,
"choices": choices
}
def handle_self(self):
### handle local model
model_root_dir = os.getenv("MODEL_ROOT_DIR", "pretrained_models")
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-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
def load(self):
2024-07-17 18:52:20 +08:00
"""根据传入的文件覆盖掉默认配置"""
logger.info(f"Loading config from {self.filename}")
if self.filename is not None and os.path.exists(self.filename):
if self.filename.endswith(".json"):
with open(self.filename, 'r') as f:
content = f.read()
if content:
self.update(json.loads(content))
else:
print(f"{self.filename} is empty.")
elif self.filename.endswith(".yaml"):
with open(self.filename, 'r') as f:
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}")
else:
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")
def save(self):
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")
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)
if self.filename.endswith(".json"):
with open(self.filename, 'w+') as f:
json.dump(self.__dict__(), f, indent=4, ensure_ascii=False)
elif self.filename.endswith(".yaml"):
with open(self.filename, 'w+') as f:
yaml.dump(self.__dict__(), f, indent=2, allow_unicode=True)
else:
logger.warning(f"Unknown config file type {self.filename}, save as json")
with open(self.filename, 'w+') as f:
json.dump(self, f, indent=4)
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"
]
}