test/scripts/download_models_standalone.py

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())