94 lines
3.1 KiB
Python
94 lines
3.1 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Offline/realtime model selection helpers."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import List, Optional
|
|
|
|
from ...core.exceptions import InvalidParameterException
|
|
from .manager import get_model_manager
|
|
from .model_plan import (
|
|
get_active_qwen_model,
|
|
get_default_model_id,
|
|
get_runtime_model_ids,
|
|
)
|
|
|
|
|
|
def get_active_qwen_model_id() -> str:
|
|
"""Return the currently active Qwen model id."""
|
|
active_qwen_model = get_active_qwen_model()
|
|
return active_qwen_model or "qwen3-asr-0.6b"
|
|
|
|
|
|
def get_offline_model_ids() -> List[str]:
|
|
"""Return enabled offline-capable models for docs and APIs."""
|
|
manager = get_model_manager()
|
|
runtime_models = get_runtime_model_ids()
|
|
|
|
def sort_key(model_id: str) -> tuple[int, str]:
|
|
if model_id.startswith("qwen"):
|
|
return (0, model_id)
|
|
return (1, model_id)
|
|
|
|
offline_models = [
|
|
model_id
|
|
for model_id in runtime_models
|
|
if manager.get_declared_entry_config(model_id).has_offline_model
|
|
]
|
|
return sorted(offline_models, key=sort_key)
|
|
|
|
|
|
def get_default_offline_model_id() -> str:
|
|
"""Return the default offline-capable model."""
|
|
default_model = get_default_model_id()
|
|
if default_model:
|
|
try:
|
|
if get_model_manager().get_declared_entry_config(default_model).has_offline_model:
|
|
return default_model
|
|
except InvalidParameterException:
|
|
pass
|
|
return get_active_qwen_model_id()
|
|
|
|
|
|
def validate_offline_model_id(model_id: Optional[str]) -> str:
|
|
"""Validate offline-capable model ids for REST transcription requests."""
|
|
available_models = get_offline_model_ids()
|
|
|
|
if not model_id or not model_id.strip():
|
|
return get_default_offline_model_id()
|
|
|
|
requested_model = model_id.strip()
|
|
if requested_model.lower() == "qwen3-asr":
|
|
active_qwen_model = get_active_qwen_model_id()
|
|
if active_qwen_model.startswith("qwen") and active_qwen_model in available_models:
|
|
return active_qwen_model
|
|
raise InvalidParameterException("当前环境未启用 Qwen3-ASR 模型")
|
|
|
|
if requested_model not in available_models:
|
|
raise InvalidParameterException(
|
|
f"不支持的离线模型ID: {requested_model}。可用模型: {', '.join(available_models)}"
|
|
)
|
|
|
|
return requested_model
|
|
|
|
|
|
def validate_realtime_model_id(model_id: Optional[str]) -> str:
|
|
"""Validate realtime-capable model ids for websocket protocols."""
|
|
available_models = get_offline_model_ids()
|
|
|
|
if not model_id:
|
|
return get_default_offline_model_id()
|
|
|
|
if model_id.lower() == "qwen3-asr":
|
|
active_qwen_model = get_active_qwen_model_id()
|
|
if active_qwen_model.startswith("qwen") and active_qwen_model in available_models:
|
|
return active_qwen_model
|
|
raise InvalidParameterException("当前环境未启用 Qwen3-ASR 模型")
|
|
|
|
if model_id not in available_models:
|
|
raise InvalidParameterException(
|
|
f"不支持的模型ID: {model_id}。可用模型: {', '.join(available_models)}"
|
|
)
|
|
|
|
return model_id
|