2025-02-25 21:26:37 +08:00
|
|
|
|
import os
|
2025-04-10 11:42:20 +08:00
|
|
|
|
import traceback
|
2025-03-29 17:33:09 +08:00
|
|
|
|
from src import config
|
2024-07-28 16:16:52 +08:00
|
|
|
|
from src.utils.logging_config import logger
|
2025-02-25 21:26:37 +08:00
|
|
|
|
from src.models.chat_model import OpenAIBase
|
2024-07-14 23:59:52 +08:00
|
|
|
|
|
|
|
|
|
|
|
2025-03-29 17:33:09 +08:00
|
|
|
|
def select_model(model_provider=None, model_name=None):
|
|
|
|
|
|
"""根据模型提供者选择模型"""
|
2025-02-20 01:26:12 +08:00
|
|
|
|
model_provider = model_provider or config.model_provider
|
2025-02-28 14:39:03 +08:00
|
|
|
|
model_info = config.model_names.get(model_provider, {})
|
|
|
|
|
|
model_name = model_name or config.model_name or model_info.get("default", "")
|
2025-02-25 21:26:37 +08:00
|
|
|
|
|
2025-03-29 17:33:09 +08:00
|
|
|
|
|
2025-02-25 21:26:37 +08:00
|
|
|
|
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
|
|
|
|
|
|
|
2024-07-07 17:21:07 +08:00
|
|
|
|
|
2025-04-10 11:42:20 +08:00
|
|
|
|
if model_provider is None:
|
|
|
|
|
|
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
|
|
|
|
|
|
|
|
|
|
|
|
if model_provider == "qianfan":
|
2024-07-28 16:16:52 +08:00
|
|
|
|
from src.models.chat_model import Qianfan
|
2024-07-09 05:04:20 +08:00
|
|
|
|
return Qianfan(model_name)
|
2024-07-07 17:21:07 +08:00
|
|
|
|
|
2025-04-10 11:42:20 +08:00
|
|
|
|
if model_provider == "dashscope":
|
2024-07-31 20:22:05 +08:00
|
|
|
|
from src.models.chat_model import DashScope
|
|
|
|
|
|
return DashScope(model_name)
|
|
|
|
|
|
|
2025-04-10 11:42:20 +08:00
|
|
|
|
if model_provider == "openai":
|
2024-09-09 17:07:03 +08:00
|
|
|
|
from src.models.chat_model import OpenModel
|
|
|
|
|
|
return OpenModel(model_name)
|
|
|
|
|
|
|
2025-05-20 20:49:50 +08:00
|
|
|
|
if model_provider == "deepseek":
|
|
|
|
|
|
from langchain_deepseek import ChatDeepSeek
|
|
|
|
|
|
return OpenAIBase(
|
|
|
|
|
|
api_key=os.getenv(model_info["env"][0]),
|
|
|
|
|
|
base_url=model_info["base_url"],
|
|
|
|
|
|
model_name=model_name,
|
|
|
|
|
|
chat_open_ai=ChatDeepSeek(
|
|
|
|
|
|
model=model_name,
|
|
|
|
|
|
api_key=os.getenv(model_info["env"][0]),
|
|
|
|
|
|
base_url=model_info["base_url"],
|
|
|
|
|
|
)
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if model_provider == "together":
|
|
|
|
|
|
from langchain_together import ChatTogether
|
|
|
|
|
|
return OpenAIBase(
|
|
|
|
|
|
api_key=os.getenv(model_info["env"][0]),
|
|
|
|
|
|
base_url=model_info["base_url"],
|
|
|
|
|
|
model_name=model_name,
|
|
|
|
|
|
chat_open_ai=ChatTogether(
|
|
|
|
|
|
model=model_name,
|
|
|
|
|
|
api_key=os.getenv(model_info["env"][0]),
|
|
|
|
|
|
base_url=model_info["base_url"],
|
|
|
|
|
|
)
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2025-04-10 11:42:20 +08:00
|
|
|
|
if model_provider == "custom":
|
2025-03-30 11:36:03 +08:00
|
|
|
|
model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None)
|
2024-10-08 22:16:17 +08:00
|
|
|
|
if model_info is None:
|
|
|
|
|
|
raise ValueError(f"Model {model_name} not found in custom models")
|
|
|
|
|
|
|
|
|
|
|
|
from src.models.chat_model import CustomModel
|
|
|
|
|
|
return CustomModel(model_info)
|
|
|
|
|
|
|
2025-04-10 11:42:20 +08:00
|
|
|
|
# 其他模型,默认使用OpenAIBase
|
|
|
|
|
|
try:
|
|
|
|
|
|
model = OpenAIBase(
|
|
|
|
|
|
api_key=os.getenv(model_info["env"][0]),
|
|
|
|
|
|
base_url=model_info["base_url"],
|
|
|
|
|
|
model_name=model_name,
|
|
|
|
|
|
)
|
|
|
|
|
|
return model
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}")
|