288 lines
10 KiB
Python
288 lines
10 KiB
Python
#!/usr/bin/env python3
|
|
# -*- coding: utf-8 -*-
|
|
"""Standalone incremental model downloader.
|
|
|
|
This script intentionally does not import the application package. It is meant
|
|
for offline delivery hosts where only ModelScope and a Python runtime are
|
|
needed to fill /opt/dep/asr/models.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Optional
|
|
|
|
from modelscope.hub.snapshot_download import snapshot_download
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ModelAsset:
|
|
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
|
|
|
|
|
|
MODEL_ASSETS: tuple[ModelAsset, ...] = (
|
|
ModelAsset(
|
|
model_id="damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
|
description="VAD",
|
|
revision="v2.0.2",
|
|
required_patterns=("configuration.json", "config.yaml", "model.pb"),
|
|
min_total_size_bytes=1_000_000,
|
|
),
|
|
ModelAsset(
|
|
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(
|
|
model_id="iic/speech_campplus_sv_zh-cn_16k-common",
|
|
description="Configured Speaker Verification",
|
|
required_patterns=("configuration.json", "config.yaml", "campplus_cn_common.bin"),
|
|
min_total_size_bytes=10_000_000,
|
|
),
|
|
ModelAsset(
|
|
model_id="iic/speech_eres2netv2_sv_zh-cn_16k-common",
|
|
description="Realtime Speaker Verification",
|
|
required_patterns=("configuration.json",),
|
|
min_total_size_bytes=10_000_000,
|
|
),
|
|
ModelAsset(
|
|
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(
|
|
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,
|
|
),
|
|
ModelAsset(
|
|
model_id="Qwen/Qwen3-ASR-0.6B",
|
|
description="Qwen3-ASR-0.6B Offline",
|
|
required_patterns=("config.json",),
|
|
alternative_required_patterns=(("model.safetensors",), ("model-*.safetensors",)),
|
|
min_total_size_bytes=500_000_000,
|
|
),
|
|
ModelAsset(
|
|
model_id="Qwen/Qwen3-ForcedAligner-0.6B",
|
|
description="Qwen3 Forced Aligner",
|
|
required_patterns=("config.json", "model.safetensors"),
|
|
min_total_size_bytes=500_000_000,
|
|
),
|
|
ModelAsset(
|
|
model_id="Qwen/Qwen3-ASR-1.7B",
|
|
description="Qwen3-ASR-1.7B Offline",
|
|
required_patterns=("config.json",),
|
|
alternative_required_patterns=(("model.safetensors",), ("model-*.safetensors",)),
|
|
min_total_size_bytes=500_000_000,
|
|
),
|
|
)
|
|
|
|
|
|
def model_path(models_dir: Path, model_id: str) -> Path:
|
|
return models_dir / model_id
|
|
|
|
|
|
def find_missing_patterns(root: Path, patterns: tuple[str, ...]) -> list[str]:
|
|
return [pattern for pattern in patterns if not any(root.glob(pattern))]
|
|
|
|
|
|
def has_required_alternative(root: Path, pattern_groups: tuple[tuple[str, ...], ...]) -> bool:
|
|
if not pattern_groups:
|
|
return True
|
|
return any(not find_missing_patterns(root, group) for group in pattern_groups)
|
|
|
|
|
|
def directory_size(path: Path) -> int:
|
|
return sum(item.stat().st_size for item in path.rglob("*") if item.is_file())
|
|
|
|
|
|
def check_asset(models_dir: Path, asset: ModelAsset) -> tuple[bool, str]:
|
|
path = model_path(models_dir, asset.model_id)
|
|
if not path.exists() or not path.is_dir():
|
|
return False, "directory_missing"
|
|
if not any(path.iterdir()):
|
|
return False, "directory_empty"
|
|
missing = find_missing_patterns(path, asset.required_patterns)
|
|
if missing:
|
|
return False, "missing=" + ",".join(missing)
|
|
if not has_required_alternative(path, asset.alternative_required_patterns):
|
|
alternatives = " OR ".join(" + ".join(group) for group in asset.alternative_required_patterns)
|
|
return False, "missing=" + alternatives
|
|
size = directory_size(path)
|
|
if size < asset.min_total_size_bytes:
|
|
return False, f"too_small={size}"
|
|
return True, "ok"
|
|
|
|
|
|
def fix_camplusplus_config(models_dir: Path, runtime_models_dir: Optional[Path] = None) -> bool:
|
|
config_file = models_dir / "iic/speech_campplus_speaker-diarization_common/configuration.json"
|
|
if not config_file.exists():
|
|
return False
|
|
|
|
runtime_dir = runtime_models_dir or models_dir
|
|
replacements = {
|
|
"damo/speech_campplus_sv_zh-cn_16k-common": "damo/speech_campplus_sv_zh-cn_16k-common",
|
|
"iic/speech_campplus_sv_zh-cn_16k-common": "iic/speech_campplus_sv_zh-cn_16k-common",
|
|
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": "damo/speech_campplus-transformer_scl_zh-cn_16k-common",
|
|
"damo/speech_campplus-transformer_scl_zh-cn-16k-common": "damo/speech_campplus-transformer_scl_zh-cn-16k-common",
|
|
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch": "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
|
}
|
|
|
|
try:
|
|
config = json.loads(config_file.read_text(encoding="utf-8"))
|
|
except Exception as exc:
|
|
print(f"⚠️ CAM++ 配置读取失败: {exc}")
|
|
return False
|
|
|
|
modified = False
|
|
model_config = config.get("model")
|
|
if isinstance(model_config, dict):
|
|
for key in ("speaker_model", "change_locator", "vad_model"):
|
|
old_value = model_config.get(key)
|
|
relative_path = replacements.get(str(old_value))
|
|
if not relative_path and isinstance(old_value, str):
|
|
for candidate in replacements.values():
|
|
if old_value.endswith(candidate):
|
|
relative_path = candidate
|
|
break
|
|
if relative_path and (models_dir / relative_path).exists():
|
|
model_config[key] = str(runtime_dir / relative_path)
|
|
modified = True
|
|
|
|
if not modified:
|
|
return False
|
|
|
|
config_file.write_text(
|
|
json.dumps(config, indent=4, ensure_ascii=False),
|
|
encoding="utf-8",
|
|
)
|
|
return True
|
|
|
|
|
|
def download_missing(
|
|
models_dir: Path,
|
|
cache_dir: Path,
|
|
auto_mode: bool = False,
|
|
runtime_models_dir: Optional[Path] = None,
|
|
) -> bool:
|
|
models_dir.mkdir(parents=True, exist_ok=True)
|
|
cache_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
missing: list[tuple[ModelAsset, str]] = []
|
|
for asset in MODEL_ASSETS:
|
|
ok, reason = check_asset(models_dir, asset)
|
|
if not ok:
|
|
missing.append((asset, reason))
|
|
|
|
if not missing:
|
|
if not auto_mode:
|
|
print("✅ 所有模型已存在且完整,无需下载")
|
|
fix_camplusplus_config(models_dir, runtime_models_dir=runtime_models_dir)
|
|
return True
|
|
|
|
print(f"📦 检测到 {len(missing)} 个模型需要下载/补齐")
|
|
if not auto_mode:
|
|
for asset, reason in missing:
|
|
print(f" - {asset.model_id} ({reason})")
|
|
|
|
failed: list[tuple[str, str]] = []
|
|
for index, (asset, _reason) in enumerate(missing, start=1):
|
|
if not auto_mode:
|
|
print(f"\n[{index}/{len(missing)}] {asset.description}")
|
|
print(f" 模型ID: {asset.model_id}")
|
|
try:
|
|
kwargs = {
|
|
"cache_dir": str(cache_dir),
|
|
"local_dir": str(model_path(models_dir, asset.model_id)),
|
|
}
|
|
if asset.revision:
|
|
kwargs["revision"] = asset.revision
|
|
snapshot_download(asset.model_id, **kwargs)
|
|
except Exception as exc:
|
|
print(f"❌ 下载失败: {asset.model_id}: {exc}")
|
|
failed.append((asset.model_id, str(exc)))
|
|
|
|
if fix_camplusplus_config(models_dir, runtime_models_dir=runtime_models_dir) and not auto_mode:
|
|
print("✅ CAM++ 配置已修复为本地模型路径")
|
|
|
|
if failed:
|
|
print("\n失败模型:")
|
|
for model_id, error in failed:
|
|
print(f" - {model_id}: {error}")
|
|
return False
|
|
|
|
still_missing = [
|
|
(asset, reason)
|
|
for asset in MODEL_ASSETS
|
|
for ok, reason in [check_asset(models_dir, asset)]
|
|
if not ok
|
|
]
|
|
if still_missing:
|
|
print("\n仍不完整的模型:")
|
|
for asset, reason in still_missing:
|
|
print(f" - {asset.model_id}: {reason}")
|
|
return False
|
|
|
|
print("✅ 所有模型准备就绪")
|
|
return True
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser(description="Download Qwen3-ASR models without importing app code")
|
|
parser.add_argument(
|
|
"--models-dir",
|
|
default=None,
|
|
help="Model directory to fill; default: MODELSCOPE_PATH, MODELS_DIR, or ./models",
|
|
)
|
|
parser.add_argument(
|
|
"--cache-dir",
|
|
default=None,
|
|
help="ModelScope cache directory; default: parent of models-dir",
|
|
)
|
|
parser.add_argument(
|
|
"--runtime-models-dir",
|
|
default=None,
|
|
help="Runtime model directory to write into CAM++ config; default: models-dir",
|
|
)
|
|
parser.add_argument("--auto-mode", action="store_true", help="Reduce output")
|
|
args = parser.parse_args()
|
|
|
|
import os
|
|
|
|
models_dir = Path(
|
|
args.models_dir
|
|
or os.getenv("MODELSCOPE_PATH")
|
|
or os.getenv("MODELS_DIR")
|
|
or "./models"
|
|
).resolve()
|
|
cache_dir = Path(args.cache_dir or os.getenv("MODELSCOPE_CACHE") or models_dir.parent).resolve()
|
|
runtime_models_dir = Path(args.runtime_models_dir).resolve() if args.runtime_models_dir else None
|
|
return 0 if download_missing(
|
|
models_dir,
|
|
cache_dir,
|
|
auto_mode=args.auto_mode,
|
|
runtime_models_dir=runtime_models_dir,
|
|
) else 1
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|