231 lines
8.3 KiB
Python
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
|