231 lines
9.4 KiB
Python
231 lines
9.4 KiB
Python
#!/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())
|