- 在 pyproject.toml 中添加了 colorlog、langchain-deepseek 和 langchain-together 依赖。 - 修改了 chat_router.py 中的日志记录方式。 - 重命名 tools_factory.py 中的函数以更好地描述其功能。 - 更新了配置文件以支持新的模型选择。 - 优化了前端组件的样式和功能,包括侧边栏和消息输入框的交互体验。
77 lines
2.6 KiB
Python
77 lines
2.6 KiB
Python
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 == "dashscope":
|
||
from src.models.chat_model import DashScope
|
||
return DashScope(model_name)
|
||
|
||
if model_provider == "openai":
|
||
from src.models.chat_model import OpenModel
|
||
return OpenModel(model_name)
|
||
|
||
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"],
|
||
)
|
||
)
|
||
|
||
if model_provider == "custom":
|
||
model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None)
|
||
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)
|
||
|
||
# 其他模型,默认使用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()}")
|