test/app/services/asr/model_capabilities.py

231 lines
8.3 KiB
Python

# -*- coding: utf-8 -*-
"""Shared capability-to-model asset definitions."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal, Optional
from app.core.config import settings
from app.services.asr.manager import get_model_manager
from app.services.asr.model_plan import (
get_active_qwen_model,
load_supported_model_ids,
)
ModelSource = Literal["modelscope"]
@dataclass(frozen=True)
class ModelAsset:
source: ModelSource
model_id: str
description: str
revision: Optional[str] = None
required_patterns: tuple[str, ...] = ()
alternative_required_patterns: tuple[tuple[str, ...], ...] = ()
min_total_size_bytes: int = 0
_VAD_ASSETS = (
ModelAsset(
source="modelscope",
model_id=settings.VAD_MODEL,
description="VAD",
revision="v2.0.2",
required_patterns=("configuration.json", "config.yaml", "model.pb"),
min_total_size_bytes=1_000_000,
),
)
_DIARIZATION_ASSETS = (
ModelAsset(
source="modelscope",
model_id="iic/speech_campplus_speaker-diarization_common",
description="CAM++ Diarization",
required_patterns=(
"configuration.json",
"config.yaml",
"onnx/asd.onnx",
"onnx/face_recog_ir101.onnx",
"onnx/fqa.onnx",
"onnx/version-RFB-320.onnx",
),
min_total_size_bytes=50_000_000,
),
ModelAsset(
source="modelscope",
model_id=settings.SV_MODEL,
description="Configured Speaker Verification",
required_patterns=("configuration.json", "config.yaml", "campplus_cn_common.bin"),
min_total_size_bytes=10_000_000,
),
ModelAsset(
source="modelscope",
model_id=settings.REALTIME_SV_MODEL,
description="Realtime Speaker Verification",
required_patterns=("configuration.json",),
min_total_size_bytes=10_000_000,
),
ModelAsset(
source="modelscope",
model_id="damo/speech_campplus_sv_zh-cn_16k-common",
description="CAM++ Speaker Verification",
required_patterns=("configuration.json", "config.yaml", "campplus_cn_common.bin"),
min_total_size_bytes=10_000_000,
),
ModelAsset(
source="modelscope",
model_id="damo/speech_campplus-transformer_scl_zh-cn_16k-common",
description="CAM++ Transformer",
required_patterns=("configuration.json", "campplus_cn_encoder.pt", "transformer_backend.pt"),
min_total_size_bytes=10_000_000,
),
)
def get_download_modelscope_assets() -> list[ModelAsset]:
"""Return the full static ModelScope export set used by predownload/export."""
return _dedupe_assets([
*_VAD_ASSETS,
*_DIARIZATION_ASSETS,
])
def get_runtime_required_modelscope_assets(
*,
include_realtime_punc: bool,
) -> list[ModelAsset]:
"""Return ModelScope assets required by the current runtime plan."""
_ = include_realtime_punc
return _dedupe_assets([*_VAD_ASSETS, *_DIARIZATION_ASSETS])
def _dedupe_assets(assets: list[ModelAsset]) -> list[ModelAsset]:
deduped: list[ModelAsset] = []
seen: set[tuple[str, str]] = set()
for asset in assets:
key = (asset.source, asset.model_id)
if key in seen:
continue
seen.add(key)
deduped.append(asset)
return deduped
def get_camplusplus_replacement_paths(cache_dir: str) -> dict[str, str]:
"""Return the CAM++ offline replacement map for local cache paths."""
return {
"damo/speech_campplus_sv_zh-cn_16k-common": f"{cache_dir}/damo/speech_campplus_sv_zh-cn_16k-common",
"iic/speech_campplus_sv_zh-cn_16k-common": f"{cache_dir}/iic/speech_campplus_sv_zh-cn_16k-common",
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": f"{cache_dir}/damo/speech_campplus-transformer_scl_zh-cn_16k-common",
"damo/speech_campplus-transformer_scl_zh-cn-16k-common": f"{cache_dir}/damo/speech_campplus-transformer_scl_zh-cn-16k-common",
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch": f"{cache_dir}/damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
}
_QWEN_MODELSCOPE_MODEL_IDS = {
"Qwen/Qwen3-ASR-0.6B": "Qwen/Qwen3-ASR-0.6B",
"Qwen/Qwen3-ASR-1.7B": "Qwen/Qwen3-ASR-1.7B",
"Qwen/Qwen3-ForcedAligner-0.6B": "Qwen/Qwen3-ForcedAligner-0.6B",
}
def get_qwen_modelscope_model_id(model_id: str) -> Optional[str]:
"""Map runtime Qwen model ID to the corresponding ModelScope model ID.
Returns:
ModelScope model ID if available, None otherwise.
"""
return _QWEN_MODELSCOPE_MODEL_IDS.get(model_id)
def get_enabled_qwen_modelscope_assets(
*,
include_forced_aligner: bool = True,
) -> list[ModelAsset]:
"""Return ModelScope assets required by the runtime Qwen plan."""
manager = get_model_manager()
assets: list[ModelAsset] = []
model_id = get_active_qwen_model()
model_config = manager.get_declared_entry_config(model_id)
offline_model = model_config.offline_model_path
if offline_model:
ms_model_id = get_qwen_modelscope_model_id(offline_model)
if ms_model_id:
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_model_id,
description=f"{model_config.name} Offline (ModelScope)",
required_patterns=("config.json",),
alternative_required_patterns=(
("model.safetensors",),
("model-*.safetensors",),
),
min_total_size_bytes=500_000_000,
)
)
forced_aligner = str(model_config.extra_kwargs.get("forced_aligner_path") or "").strip()
if forced_aligner and include_forced_aligner:
ms_aligner_id = get_qwen_modelscope_model_id(forced_aligner)
if ms_aligner_id:
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_aligner_id,
description=f"{model_config.name} Forced Aligner (ModelScope)",
required_patterns=("config.json", "model.safetensors"),
min_total_size_bytes=500_000_000,
)
)
return assets
def get_all_qwen_modelscope_assets(
*,
include_forced_aligner: bool = True,
) -> list[ModelAsset]:
"""Return all declared Qwen ModelScope assets for offline bundles."""
manager = get_model_manager()
assets: list[ModelAsset] = []
seen_model_ids: set[str] = set()
for model_id in sorted(load_supported_model_ids()):
if not model_id.startswith("qwen3-asr-"):
continue
model_config = manager.get_declared_entry_config(model_id)
offline_model = model_config.offline_model_path
if offline_model:
ms_model_id = get_qwen_modelscope_model_id(offline_model)
if ms_model_id and ms_model_id not in seen_model_ids:
seen_model_ids.add(ms_model_id)
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_model_id,
description=f"{model_config.name} Offline (ModelScope)",
required_patterns=("config.json",),
alternative_required_patterns=(
("model.safetensors",),
("model-*.safetensors",),
),
min_total_size_bytes=500_000_000,
)
)
forced_aligner = str(model_config.extra_kwargs.get("forced_aligner_path") or "").strip()
if forced_aligner and include_forced_aligner:
ms_aligner_id = get_qwen_modelscope_model_id(forced_aligner)
if ms_aligner_id and ms_aligner_id not in seen_model_ids:
seen_model_ids.add(ms_aligner_id)
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_aligner_id,
description="Qwen3 Forced Aligner (ModelScope)",
required_patterns=("config.json", "model.safetensors"),
min_total_size_bytes=500_000_000,
)
)
return assets