添加更多模型支持

This commit is contained in:
Wenjie Zhang 2025-02-25 21:26:37 +08:00
parent d8f2f8490d
commit becbe5864d
8 changed files with 79 additions and 54 deletions

View File

@ -106,11 +106,30 @@ docker compose -f docker/docker-compose.yml --env-file src/.env up --build
| `siliconflow` (默认) | `Qwen/Qwen2.5-7B-Instruct` (免费) | `SILICONFLOW_API_KEY` | | `siliconflow` (默认) | `Qwen/Qwen2.5-7B-Instruct` (免费) | `SILICONFLOW_API_KEY` |
| `openai` | `gpt-4o` | `OPENAI_API_KEY` | | `openai` | `gpt-4o` | `OPENAI_API_KEY` |
| `deepseek` | `deepseek-chat` | `DEEPSEEK_API_KEY` | | `deepseek` | `deepseek-chat` | `DEEPSEEK_API_KEY` |
| `arc`(豆包方舟) | `doubao-1-5-pro-32k-250115` | `ARK_API_KEY` |
| `zhipu`(智谱清言) | `glm-4-flash` | `ZHIPUAI_API_KEY` | | `zhipu`(智谱清言) | `glm-4-flash` | `ZHIPUAI_API_KEY` |
| `dashscope`(阿里) | `qwen-max-latest` | `DASHSCOPE_API_KEY` | | `dashscope`(阿里) | `qwen-max-latest` | `DASHSCOPE_API_KEY` |
| `qianfan`(百度) | `ernie_speed` | `QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY` | | `qianfan`(百度) | `ernie_speed` | `QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY` |
同样支持以 OpenAI 的兼容模型运行模型,可以直接在 Web 设置里面添加。比如使用 vllm 和 Ollama 运行本地模型时。 此外,如果想要添加供应商的模型,确认知识 OpenAI 调用的方法之后,只需要在 [src/static/models.yaml](src/static/models.yaml) 中添加对应的模型配置即可。配置示例如下:
```yaml
ark:
name: 豆包Ark
url: https://console.volcengine.com/ark/region:ark+cn-beijing/model # 模型列表
default: doubao-1-5-pro-32k-250115 # 默认模型
base_url: https://ark.cn-beijing.volces.com/api/v3
env: # 需要配置的环境变量仅限API key
- ARK_API_KEY
models:
- doubao-1-5-pro-32k-250115
- doubao-1-5-lite-32k-250115
- deepseek-r1-250120
```
同样支持添加以 OpenAI 的兼容模式运行的单个模型,可以直接在 Web 设置里面添加。比如使用 vllm 和 Ollama 运行本地模型时。
![](./images/custom_models.png)
### 2. 向量模型支持和重排序模型支持 ### 2. 向量模型支持和重排序模型支持

BIN
images/custom_models.png Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 304 KiB

View File

@ -60,7 +60,9 @@ class Config(SimpleConfig):
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR ## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR
self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(MODEL_NAMES.keys())) self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(MODEL_NAMES.keys()))
self.add_item("model_provider_lite", default="siliconflow", des="模型提供商(用于轻量任务)", choices=list(MODEL_NAMES.keys()))
self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称") self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称")
self.add_item("model_name_lite", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称(用于轻量任务)")
self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys()))
self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(RERANKER_LIST.keys())) self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(RERANKER_LIST.keys()))
self.add_item("model_local_paths", default={}, des="本地模型路径") self.add_item("model_local_paths", default={}, des="本地模型路径")

View File

@ -68,7 +68,9 @@ class DataBaseManager:
assert self.config.enable_knowledge_base, "知识库未启用" assert self.config.enable_knowledge_base, "知识库未启用"
knowledge_base_collections = self.knowledge_base.get_collection_names() knowledge_base_collections = self.knowledge_base.get_collection_names()
if len(self.data["databases"]) != len(knowledge_base_collections): if len(self.data["databases"]) != len(knowledge_base_collections):
logger.warning(f"Database number not match, {knowledge_base_collections}") logger.warning(
f"Database number not match, {knowledge_base_collections}, "
f"self.data['databases']: {self.data['databases']}, ")
for db in self.data["databases"]: for db in self.data["databases"]:
db.update(self.knowledge_base.get_collection_info(db.metaname)) db.update(self.knowledge_base.get_collection_info(db.metaname))

View File

@ -19,8 +19,8 @@ class Startup:
self.retriever = Retriever(self.config, self.dbm, self.model) self.retriever = Retriever(self.config, self.dbm, self.model)
self.model_lite = select_model(self.config, self.model_lite = select_model(self.config,
model_provider=self.config.model_provider_lite or "zhipu", model_provider=self.config.model_provider_lite,
model_name=self.config.model_name_lite or "glm-4-flash") model_name=self.config.model_name_lite)
def restart(self): def restart(self):
logger.info("Restarting...") logger.info("Restarting...")

View File

@ -1,20 +1,29 @@
import os
from src.utils.logging_config import logger from src.utils.logging_config import logger
from src.models.chat_model import OpenAIBase
def select_model(config, model_provider=None, model_name=None): def select_model(config, model_provider=None, model_name=None):
model_provider = model_provider or config.model_provider model_provider = model_provider or config.model_provider
model_name = model_name or config.model_name model_info = config.model_names[model_provider]
model_name = model_name or config.model_name or model_info["default"]
logger.info(f"Selecting model from {model_provider} with {model_name}") logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
if model_provider == "deepseek": if model_provider in [
from src.models.chat_model import DeepSeek "deepseek",
return DeepSeek(model_name) "ark",
"siliconflow",
elif model_provider == "zhipu": "zhipu",
from src.models.chat_model import Zhipu "lingyiwanwu",
return Zhipu(model_name) ]:
return OpenAIBase(
api_key=os.getenv(model_info["env"][0]),
base_url=model_info["base_url"],
model_name=model_name,
)
elif model_provider == "qianfan": elif model_provider == "qianfan":
from src.models.chat_model import Qianfan from src.models.chat_model import Qianfan
@ -28,10 +37,6 @@ def select_model(config, model_provider=None, model_name=None):
from src.models.chat_model import OpenModel from src.models.chat_model import OpenModel
return OpenModel(model_name) return OpenModel(model_name)
elif model_provider == "siliconflow":
from src.models.chat_model import SiliconFlow
return SiliconFlow(model_name)
elif model_provider == "custom": elif model_provider == "custom":
model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None) model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None)
if model_info is None: if model_info is None:

View File

@ -47,28 +47,6 @@ class OpenModel(OpenAIBase):
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
class DeepSeek(OpenAIBase):
def __init__(self, model_name=None):
model_name = model_name or "deepseek-chat"
api_key = os.getenv("DEEPSEEK_API_KEY", "your-default-api-key")
base_url = os.getenv("DEEPSEEK_API_BASE", "https://api.deepseek.com/v1")
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
class Zhipu(OpenAIBase):
def __init__(self, model_name=None):
model_name = model_name or "glm-4-flash"
api_key = os.getenv("ZHIPUAI_API_KEY", "270ea71e9560c0ff406acbcdd48bfd97.e3XOMdWKuZb7Q1Sk")
base_url = "https://open.bigmodel.cn/api/paas/v4/"
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
class SiliconFlow(OpenAIBase):
def __init__(self, model_name=None):
model_name = model_name or "meta-llama/Meta-Llama-3.1-8B-Instruct"
api_key = os.getenv("SILICONFLOW_API_KEY")
base_url = "https://api.siliconflow.cn/v1"
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
class CustomModel(OpenAIBase): class CustomModel(OpenAIBase):
def __init__(self, model_info): def __init__(self, model_info):
model_name = model_info["name"] model_name = model_info["name"]

View File

@ -14,6 +14,7 @@ MODEL_NAMES:
name: DeepSeek name: DeepSeek
url: https://platform.deepseek.com/api-docs/zh-cn/pricing url: https://platform.deepseek.com/api-docs/zh-cn/pricing
default: deepseek-chat default: deepseek-chat
base_url: https://api.deepseek.com/v1
env: env:
- DEEPSEEK_API_KEY - DEEPSEEK_API_KEY
models: models:
@ -23,6 +24,7 @@ MODEL_NAMES:
name: 智谱AI (Zhipu) name: 智谱AI (Zhipu)
url: https://open.bigmodel.cn/dev/api url: https://open.bigmodel.cn/dev/api
default: glm-4-flash default: glm-4-flash
base_url: https://open.bigmodel.cn/api/paas/v4/
env: env:
- ZHIPUAI_API_KEY - ZHIPUAI_API_KEY
models: models:
@ -33,6 +35,19 @@ MODEL_NAMES:
- glm-4-long - glm-4-long
- glm-4-flashx - glm-4-flashx
- glm-4-flash - glm-4-flash
siliconflow:
name: SiliconFlow
url: https://cloud.siliconflow.cn/models
default: Qwen/Qwen2.5-7B-Instruct
base_url: https://api.siliconflow.cn/v1
env:
- SILICONFLOW_API_KEY
models:
- meta-llama/Meta-Llama-3.1-8B-Instruct
- Qwen/Qwen2.5-7B-Instruct
- deepseek-ai/DeepSeek-R1
- deepseek-ai/DeepSeek-V3
- deepseek-ai/DeepSeek-R1-Distill-Qwen-7B
qianfan: qianfan:
name: 百度千帆 (QianFan) name: 百度千帆 (QianFan)
url: https://open.bigmodel.cn/dev/api url: https://open.bigmodel.cn/dev/api
@ -55,28 +70,32 @@ MODEL_NAMES:
- DASHSCOPE_API_KEY - DASHSCOPE_API_KEY
models: models:
- qwen-max-latest - qwen-max-latest
- qwen-plus-latest
- qwen-long-latest - qwen-long-latest
- qwen-turbo-latest
- qwen2.5-72b-instruct - qwen2.5-72b-instruct
- qwen2.5-32b-instruct - qwen2.5-32b-instruct
- qwen2.5-14b-instruct
- qwen2.5-7b-instruct - qwen2.5-7b-instruct
- qwen2.5-3b-instruct
- qwen2.5-1.5b-instruct - qwen2.5-1.5b-instruct
- qwen2.5-0.5b-instruct - qwen2.5-0.5b-instruct
siliconflow: ark:
name: SiliconFlow name: 豆包Ark
url: https://cloud.siliconflow.cn/models url: https://console.volcengine.com/ark/region:ark+cn-beijing/model
default: Qwen/Qwen2.5-7B-Instruct default: doubao-1-5-pro-32k-250115
base_url: https://ark.cn-beijing.volces.com/api/v3
env: env:
- SILICONFLOW_API_KEY - ARK_API_KEY
models: models:
- meta-llama/Meta-Llama-3.1-8B-Instruct - doubao-1-5-pro-32k-250115
- Qwen/Qwen2.5-7B-Instruct - doubao-1-5-lite-32k-250115
- deepseek-ai/DeepSeek-R1 - deepseek-r1-250120
- deepseek-ai/DeepSeek-V3 lingyiwanwu:
- deepseek-ai/DeepSeek-R1-Distill-Qwen-7B name: 零一万物
url: https://platform.lingyiwanwu.com/docs#%E6%A8%A1%E5%9E%8B%E4%B8%8E%E8%AE%A1%E8%B4%B9
base_url: https://api.lingyiwanwu.com/v1
default: yi-lightning
env:
- LINGYIWANWU_API_KEY
models:
- yi-lightning
EMBED_MODEL_INFO: EMBED_MODEL_INFO:
local/BAAI/bge-m3: local/BAAI/bge-m3: