2024-07-28 16:16:52 +08:00
|
|
|
from src.utils.logging_config import logger
|
2024-07-14 23:59:52 +08:00
|
|
|
|
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
def select_model(config):
|
|
|
|
|
|
|
|
|
|
model_provider = config.model_provider
|
|
|
|
|
model_name = config.model_name
|
|
|
|
|
|
2024-07-22 00:00:54 +08:00
|
|
|
logger.info(f"Selecting model from {model_provider} with {model_name or 'default'}")
|
2024-07-14 23:59:52 +08:00
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
if model_provider == "deepseek":
|
2024-07-28 16:16:52 +08:00
|
|
|
from src.models.chat_model import DeepSeek
|
2024-07-07 01:58:23 +08:00
|
|
|
return DeepSeek(model_name)
|
2024-07-07 17:21:07 +08:00
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
elif model_provider == "zhipu":
|
2024-07-28 16:16:52 +08:00
|
|
|
from src.models.chat_model import Zhipu
|
2024-07-07 01:58:23 +08:00
|
|
|
return Zhipu(model_name)
|
2024-07-07 17:21:07 +08:00
|
|
|
|
2024-07-09 05:04:20 +08:00
|
|
|
elif 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
|
|
|
|
2024-07-20 15:27:33 +08:00
|
|
|
elif model_provider == "vllm":
|
2024-07-28 16:16:52 +08:00
|
|
|
from src.models.chat_model import VLLM
|
2024-07-20 15:27:33 +08:00
|
|
|
return VLLM(model_name)
|
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
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")
|