ASR-demo/scripts/download_models.py

267 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

#!/usr/bin/env python3
"""Download FunASR ASR and its configured runtime assets."""
from __future__ import annotations
import argparse
import json
import os
from pathlib import Path
from dotenv import load_dotenv
PROJECT_ROOT = Path(__file__).resolve().parents[1]
# 使用与后端启动器相同的项目本地配置。
load_dotenv(PROJECT_ROOT / ".env")
# 同时支持直接执行脚本和 `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:
"""检查流式 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.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("FUNASR_ASR_MODEL", "paraformer-zh-streaming"),
help="ASR model alias, 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 configured auxiliary assets",
)
auxiliary_group.add_argument(
"--funasr-runtime",
action="store_true",
help="Download/check the configured ASR model, FSMN-VAD, CAM++, and punctuation",
)
args = parser.parse_args()
manifest = load_manifest()
models_dir = args.models_dir
if not models_dir.is_absolute():
models_dir = Path(__file__).resolve().parents[1] / models_dir
models_dir = models_dir.resolve()
cache_dir = args.cache_dir.resolve() if args.cache_dir else None
selected_assets: list[tuple[str, dict[str, object]]] = []
if args.funasr_runtime:
# 标点模型在运行时可选,但下载后可获得完整的本地输出。
asr_id = resolve_model_id(args.model, manifest)
selected_assets.append((asr_id, manifest["models"][asr_id]))
assets = auxiliary_models(manifest)
vad_id = next(
model_id for model_id, config in assets.items()
if config.get("kind") == "vad"
)
cam_id = next(
model_id for model_id, config in assets.items()
if config.get("kind") == "speaker_verification"
and model_id.startswith("iic/")
)
punctuation_id = next(
model_id for model_id, config in assets.items()
if config.get("kind") == "punctuation"
)
selected_assets.extend(
(model_id, assets[model_id])
for model_id in (vad_id, cam_id, punctuation_id)
)
else:
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())