ASR-demo/scripts/download_models.py

231 lines
9.4 KiB
Python
Raw Normal View History

2026-09-10 05:47:09 +00:00
#!/usr/bin/env python3
"""为独立服务部署下载 ASR 和配套辅助模型。"""
from __future__ import annotations
import argparse
import json
import os
from pathlib import Path
# 同时支持直接执行脚本和 `python -m scripts.download_models` 两种方式,
# 下载器只依赖本目录中的清单模块,不耦合原项目的包路径。
try:
from .model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
except ImportError:
from model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
def has_model_weights(model_path: Path) -> bool:
"""检查 VLLM 加载 ASR 模型前必须存在的最小本地文件集合。"""
if not model_path.is_dir():
return False
if not (model_path / "config.json").is_file():
return False
return any(model_path.rglob("*.safetensors")) or any(model_path.rglob("*.bin"))
def is_model_ready(model_path: Path, config: dict[str, object]) -> bool:
"""根据模型清单中的专属文件规则检查 ASR 或辅助资产是否完整。"""
if not model_path.is_dir():
return False
required_files = config.get("required_files", [])
if isinstance(required_files, list):
for relative_path in required_files:
if not (model_path / str(relative_path)).is_file():
return False
any_files = config.get("any_files", [])
if isinstance(any_files, list) and any_files:
if not any(
file_path.is_file()
for pattern in any_files
for file_path in model_path.rglob(str(pattern))
):
return False
minimum_size_value = config.get("min_total_size_bytes", 0)
# 模型清单使用 object 表示不同类型的资产字段,因此在转换为整数前必须
# 先收窄类型,避免不合法的清单值在运行时触发难以定位的类型异常。
minimum_size = (
int(minimum_size_value)
if isinstance(minimum_size_value, (int, str))
else 0
)
if minimum_size:
total_size = sum(file_path.stat().st_size for file_path in model_path.rglob("*") if file_path.is_file())
if total_size < minimum_size:
return False
if required_files or any_files or minimum_size:
return True
return has_model_weights(model_path)
def download_model(
model_id: str,
model_path: Path,
cache_dir: Path | None,
revision: str | None,
) -> None:
"""通过 ModelScope 下载一个指定资产,且不导入原项目应用代码。"""
# 延迟导入 ModelScope使模型清单检查和单元测试无需安装重量级依赖。
try:
from modelscope.hub.snapshot_download import snapshot_download
except ImportError as exc:
raise RuntimeError(
"ModelScope is required for downloading; install requirements-download.txt first"
) from exc
model_path.parent.mkdir(parents=True, exist_ok=True)
cache_path: str | None = None
if cache_dir is not None:
cache_dir.mkdir(parents=True, exist_ok=True)
cache_path = str(cache_dir)
print(f"Downloading model asset: {model_id}")
print(f"Local directory: {model_path}")
# 使用显式关键字参数而不是 **dict既便于 Pylance 推断 ModelScope 的真实
# 参数类型,也避免动态字典被误判为其它无关参数的类型签名。
snapshot_download(
model_id,
revision=revision,
cache_dir=cache_path,
local_dir=str(model_path),
)
def fix_camplusplus_config(models_dir: Path) -> bool:
"""将 CAM++ 依赖模型 ID 改写为本地路径,确保服务可以离线启动。
聚类流水线会在 ``configuration.json`` 中保存多个 ModelScope 模型 ID
如果不改写这些 ID即使所有文件已经下载完整辅助服务在无网络环境
启动时仍可能再次访问 ModelScope 获取依赖
"""
config_file = models_dir / "iic/speech_campplus_speaker-diarization_common/configuration.json"
if not config_file.is_file():
return False
replacements = {
"damo/speech_campplus_sv_zh-cn_16k-common": models_dir / "damo/speech_campplus_sv_zh-cn_16k-common",
"iic/speech_campplus_sv_zh-cn_16k-common": models_dir / "iic/speech_campplus_sv_zh-cn_16k-common",
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": models_dir / "damo/speech_campplus-transformer_scl_zh-cn_16k-common",
"damo/speech_campplus-transformer_scl_zh-cn-16k-common": models_dir / "damo/speech_campplus-transformer_scl_zh-cn-16k-common",
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch": models_dir / "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
}
try:
config = json.loads(config_file.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
print(f"Unable to read CAM++ configuration: {exc}")
return False
raw_model_config = config.get("model")
if not isinstance(raw_model_config, dict):
return False
model_config: dict[str, object] = {
str(key): value for key, value in raw_model_config.items()
}
modified = False
for key in ("speaker_model", "change_locator", "vad_model"):
old_value = model_config.get(key)
local_path = replacements.get(old_value) if isinstance(old_value, str) else None
if local_path is not None and local_path.exists():
model_config[key] = str(local_path)
modified = True
if not modified:
return False
config["model"] = model_config
config_file.write_text(json.dumps(config, indent=4, ensure_ascii=False) + "\n", encoding="utf-8")
return True
def main() -> int:
"""检查或下载 ASR 模型及辅助运行时所需的全部资产。"""
# 保持当前部署项目与原项目模型规划器完全独立,同时将孤立服务需要的
# 模型统一准备到本地,方便后续在服务器上离线启动多个常驻服务。
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--model",
default=os.getenv("QWEN3_ASR_MODEL", "default"),
help="ASR model alias (1.7b/0.6b), exact model ID, or default",
)
parser.add_argument(
"--models-dir",
type=Path,
default=Path(os.getenv("MODEL_DIR", str(Path(__file__).resolve().parents[1] / "models"))),
help="Root directory for local model files",
)
model_scope_cache = os.getenv("MODELSCOPE_CACHE")
parser.add_argument(
"--cache-dir",
type=Path,
default=Path(model_scope_cache) if model_scope_cache else None,
help="Optional ModelScope cache directory",
)
parser.add_argument(
"--check-only",
action="store_true",
help="Only check selected assets; do not download",
)
auxiliary_group = parser.add_mutually_exclusive_group()
auxiliary_group.add_argument(
"--skip-auxiliary",
action="store_true",
help="Only download/check the selected ASR model",
)
auxiliary_group.add_argument(
"--auxiliary-only",
action="store_true",
help="Only download/check VAD, speaker, diarization, and aligner assets",
)
args = parser.parse_args()
manifest = load_manifest()
models_dir = args.models_dir.resolve()
cache_dir = args.cache_dir.resolve() if args.cache_dir else None
selected_assets: list[tuple[str, dict[str, object]]] = []
if not args.auxiliary_only:
model_id = resolve_model_id(args.model, manifest)
selected_assets.append((model_id, manifest["models"][model_id]))
if not args.skip_auxiliary:
selected_assets.extend(auxiliary_models(manifest).items())
missing: list[tuple[str, Path, dict[str, object]]] = []
for model_id, config in selected_assets:
model_path = model_directory(model_id, manifest, models_dir)
if is_model_ready(model_path, config):
print(f"Model asset is ready: {model_id}")
else:
missing.append((model_id, model_path, config))
if not missing:
# 即使资产已经存在,也要重新执行一次离线配置修正;这样从其它主机
# 复制过来的模型包也能在启动辅助服务前自动完成本地路径修复。
if fix_camplusplus_config(models_dir):
print("CAM++ configuration updated for offline local model paths")
print(f"All selected model assets are ready: {len(selected_assets)}")
return 0
if args.check_only:
for model_id, model_path, _ in missing:
print(f"Model asset is missing or incomplete: {model_id} ({model_path})")
return 1
failed: list[str] = []
for model_id, model_path, config in missing:
try:
revision = str(config.get("revision") or "") or None
download_model(model_id, model_path, cache_dir, revision)
if not is_model_ready(model_path, config):
print(f"Download finished but model asset is incomplete: {model_path}")
failed.append(model_id)
else:
print(f"Model asset is ready: {model_id}")
except Exception as exc:
print(f"Download failed: {model_id}: {exc}")
failed.append(model_id)
if not failed and fix_camplusplus_config(models_dir):
print("CAM++ configuration updated for offline local model paths")
return 1 if failed else 0
if __name__ == "__main__":
raise SystemExit(main())