ForcePilot/src/models/__init__.py

81 lines
2.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
import traceback
from src import config
from src.utils.logging_config import logger
from src.models.chat_model import OpenAIBase
def select_model(model_provider=None, model_name=None):
"""根据模型提供者选择模型"""
model_provider = model_provider or config.model_provider
model_info = config.model_names.get(model_provider, {})
model_name = model_name or config.model_name or model_info.get("default", "")
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
if model_provider is None:
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
if model_provider == "qianfan":
from src.models.chat_model import Qianfan
return Qianfan(model_name)
if model_provider == "openai":
from src.models.chat_model import OpenModel
return OpenModel(model_name)
if model_provider == "deepseek" or model_provider == "dashscope":
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"],
api_base=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"],
)
)
if model_provider == "custom":
model_info = get_custom_model(model_name)
from src.models.chat_model import CustomModel
return CustomModel(model_info)
# 其他模型默认使用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()}")
def get_custom_model(model_id):
"""return model_info"""
assert config.custom_models is not None, "custom_models is not set"
modle_info = next((x for x in config.custom_models if x["custom_id"] == model_id), None)
if modle_info is None:
raise ValueError(f"Model {model_id} not found in custom models")
return modle_info