ForcePilot/src/models/__init__.py

53 lines
1.7 KiB
Python
Raw Normal View History

2025-02-25 21:26:37 +08:00
import os
from src.utils.logging_config import logger
2025-02-25 21:26:37 +08:00
from src.models.chat_model import OpenAIBase
def select_model(config, model_provider=None, model_name=None):
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
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
if model_provider in [
"deepseek",
"ark",
"siliconflow",
"zhipu",
"lingyiwanwu",
2025-02-27 19:35:25 +08:00
"together.ai",
2025-02-25 21:26:37 +08:00
]:
return OpenAIBase(
api_key=os.getenv(model_info["env"][0]),
base_url=model_info["base_url"],
model_name=model_name,
)
elif model_provider == "qianfan":
from src.models.chat_model import Qianfan
return Qianfan(model_name)
2024-07-31 20:22:05 +08:00
elif model_provider == "dashscope":
from src.models.chat_model import DashScope
return DashScope(model_name)
2024-09-09 17:07:03 +08:00
elif model_provider == "openai":
from src.models.chat_model import OpenModel
return OpenModel(model_name)
2024-10-08 22:16:17 +08:00
elif model_provider == "custom":
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)
elif model_provider is None:
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
else:
raise ValueError(f"Model provider {model_provider} not supported")