ASR-demo/tests/test_model_manifest.py

74 lines
3.2 KiB
Python
Raw Permalink Normal View History

2026-09-10 05:47:09 +00:00
"""独立模型清单测试,确保不会导入原项目应用。"""
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from scripts.download_models import fix_camplusplus_config
from scripts.model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
class ModelManifestTests(unittest.TestCase):
def setUp(self) -> None:
self.manifest = load_manifest()
def test_default_is_zero_point_six_b_model(self) -> None:
self.assertEqual(resolve_model_id("default", self.manifest), "Qwen/Qwen3-ASR-0.6B")
def test_aliases_resolve_to_individual_models(self) -> None:
self.assertEqual(resolve_model_id("1.7b", self.manifest), "Qwen/Qwen3-ASR-1.7B")
self.assertEqual(resolve_model_id("0.6b", self.manifest), "Qwen/Qwen3-ASR-0.6B")
def test_manifest_has_only_asr_models(self) -> None:
self.assertEqual(
set(self.manifest["models"]),
{"Qwen/Qwen3-ASR-0.6B", "Qwen/Qwen3-ASR-1.7B"},
)
def test_manifest_has_auxiliary_runtime_assets(self) -> None:
assets = auxiliary_models(self.manifest)
self.assertIn("damo/speech_fsmn_vad_zh-cn-16k-common-pytorch", assets)
self.assertIn("iic/speech_campplus_speaker-diarization_common", assets)
self.assertIn("iic/speech_campplus_sv_zh-cn_16k-common", assets)
self.assertIn("Qwen/Qwen3-ForcedAligner-0.6B", assets)
def test_model_directory_is_under_demo_models(self) -> None:
models_dir = Path(__file__).resolve().parents[1] / "models"
for model_id in [*self.manifest["models"], *auxiliary_models(self.manifest)]:
self.assertTrue(model_directory(model_id, self.manifest, models_dir).is_relative_to(models_dir))
def test_camplusplus_config_is_rewritten_to_local_assets(self) -> None:
"""离线模型包不能继续从 ModelScope 解析 CAM++ 依赖。"""
with tempfile.TemporaryDirectory() as temp_dir:
models_dir = Path(temp_dir)
config_dir = models_dir / "iic/speech_campplus_speaker-diarization_common"
config_dir.mkdir(parents=True)
for relative_path in (
"damo/speech_campplus_sv_zh-cn_16k-common",
"iic/speech_campplus_sv_zh-cn_16k-common",
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
):
(models_dir / relative_path).mkdir(parents=True)
config = {
"model": {
"speaker_model": "iic/speech_campplus_sv_zh-cn_16k-common",
"change_locator": "damo/speech_campplus_sv_zh-cn_16k-common",
"vad_model": "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
}
}
(config_dir / "configuration.json").write_text(json.dumps(config), encoding="utf-8")
self.assertTrue(fix_camplusplus_config(models_dir))
updated = json.loads((config_dir / "configuration.json").read_text(encoding="utf-8"))
self.assertEqual(
updated["model"]["speaker_model"],
str(models_dir / "iic/speech_campplus_sv_zh-cn_16k-common"),
)
if __name__ == "__main__":
unittest.main()