# -*- 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