ASR-demo/tests/test_auxiliary_server.py

170 lines
8.0 KiB
Python
Raw Permalink 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.

"""辅助服务输出格式测试,确保不同 ModelScope 版本都能被统一解析。"""
from __future__ import annotations
import unittest
import numpy as np
from types import SimpleNamespace
from unittest.mock import patch
from scripts.auxiliary_server import (
AuxiliaryRuntime,
_coerce_finite_float,
_normalize_diarization_segments,
_parse_form_float,
)
class AuxiliaryServerTests(unittest.TestCase):
"""验证 CAM++ 常见的秒、毫秒和对象返回格式。"""
def test_multipart_numeric_values_are_narrowed_before_model_calls(self):
"""表单值只接受有限数字FileField 或非法文本回退/报错而不传入 float。"""
self.assertEqual(_coerce_finite_float(" 12.5 "), 12.5)
self.assertEqual(_coerce_finite_float(b"12.5"), 12.5)
self.assertIsNone(_coerce_finite_float(float("nan")))
self.assertEqual(_parse_form_float(None, "start_time_ms", default=0.0), 0.0)
with self.assertRaises(ValueError):
_parse_form_float("not-a-number", "start_time_ms")
def test_rejects_invalid_embedding_vectors(self):
"""模型成功返回也不代表向量有效,异常特征不能污染聚类池。"""
for value in ([0, 0], [], [float("nan"), 1], [float("inf"), 0], [[1, 0], [0, 1]]):
with self.assertRaises(RuntimeError):
AuxiliaryRuntime._normalize_embedding(value)
def test_optional_diarization_failure_does_not_block_core_startup(self):
"""完整 CAM++ diarization 不是实时核心加载路径,启动失败应只记录可选状态。"""
runtime = AuxiliaryRuntime()
runtime.assets = {
"vad": {"kind": "vad"},
"diarization": {"kind": "diarization"},
"aligner": {"kind": "forced_aligner"},
}
loaded_kinds = []
with patch("scripts.auxiliary_server._asset_ready", return_value=True), \
patch("scripts.auxiliary_server.model_directory", return_value=runtime.manifest and SimpleNamespace()), \
patch.object(runtime, "_load_asset", side_effect=lambda model_id, config, path: loaded_kinds.append(config["kind"]) or object()):
runtime.preload()
self.assertEqual(loaded_kinds, ["vad"])
self.assertEqual(runtime.status["diarization"]["state"], "optional_not_preloaded")
self.assertEqual(runtime.status["aligner"]["state"], "optional_not_preloaded")
def test_campplus_speaker_is_a_required_core_model(self):
"""实时在线聚类必须使用 CAM++ speaker_verification不能只启动 VAD。"""
runtime = AuxiliaryRuntime()
runtime.assets = {
"vad": {"kind": "vad"},
"campplus": {"kind": "speaker_verification"},
}
with patch("scripts.auxiliary_server._asset_ready", side_effect=[True, False]), \
patch("scripts.auxiliary_server.model_directory", return_value=SimpleNamespace()):
with self.assertRaisesRegex(RuntimeError, "campplus"):
runtime.preload()
def test_campplus_is_preferred_over_other_embedding_models(self):
"""即使同时存在 ERes2Net实时聚类仍优先使用 CAM++。"""
runtime = AuxiliaryRuntime()
runtime.assets = {
"campplus": {"kind": "speaker_verification"},
"eres2net": {"kind": "realtime_speaker_verification"},
}
runtime.models = {"campplus": object(), "eres2net": object()}
self.assertEqual(runtime._speaker_embedding_model_id(), "campplus")
def test_missing_vad_reports_download_path(self):
"""VAD 是核心依赖,缺失时错误必须给出可执行的修复方向。"""
runtime = AuxiliaryRuntime()
runtime.assets = {"vad": {"kind": "vad"}}
with patch("scripts.auxiliary_server._asset_ready", return_value=False), \
patch("scripts.auxiliary_server.model_directory", return_value=SimpleNamespace(__str__=lambda self: "/models/vad")):
with self.assertRaisesRegex(RuntimeError, "download_models.py --auxiliary-only"):
runtime.preload()
def test_uses_public_pipeline_embedding_output(self) -> None:
"""声纹推理必须走 pipeline 的预处理和 output_emb 接口。"""
class FakePipeline:
def __init__(self) -> None:
self.calls: list[tuple[list[str], bool]] = []
def __call__(self, audio_paths: list[str], output_emb: bool = False) -> dict[str, object]:
self.calls.append((audio_paths, output_emb))
return {"outputs": {"text": "No similarity score output"}, "embs": [[1.0, 2.0, 3.0]]}
pipeline = FakePipeline()
result = AuxiliaryRuntime._run_embedding_pipeline(pipeline, "turn.wav")
self.assertEqual(result, [1.0, 2.0, 3.0])
self.assertEqual(pipeline.calls, [(["turn.wav"], True)])
def test_supports_legacy_pipeline_without_output_emb_argument(self) -> None:
"""旧版 pipeline 不支持 output_emb 时,仍应使用其 preprocess/forward。"""
class LegacyPipeline:
def __init__(self) -> None:
self.prepared: list[str] = []
def __call__(self, *_args: object, **_kwargs: object) -> None:
raise TypeError("output_emb is not supported")
def preprocess(self, audio_paths: list[str]) -> list[str]:
return [f"prepared:{audio_paths[0]}"]
def forward(self, prepared: list[str]) -> list[list[float]]:
self.prepared = prepared
return [[0.1, 0.2, 0.3]]
pipeline = LegacyPipeline()
result = AuxiliaryRuntime._run_embedding_pipeline(pipeline, "turn.wav")
self.assertEqual(result, [0.1, 0.2, 0.3])
self.assertEqual(pipeline.prepared, ["prepared:turn.wav"])
def test_normalizes_modelscope_text_seconds(self) -> None:
result = _normalize_diarization_segments({"text": [[0.2, 1.4, 0], [1.4, 2.0, 1]]})
self.assertEqual(
result,
[
{"start_time": 200.0, "end_time": 1400.0, "speaker": "0"},
{"start_time": 1400.0, "end_time": 2000.0, "speaker": "1"},
],
)
def test_normalizes_named_millisecond_fields(self) -> None:
result = _normalize_diarization_segments(
{"segments": [{"start_time": 100, "end_time": 900, "speaker_id": "cluster-a"}]}
)
self.assertEqual(result[0]["start_time"], 100.0)
self.assertEqual(result[0]["end_time"], 900.0)
def test_reads_object_text_attribute(self) -> None:
result = _normalize_diarization_segments(SimpleNamespace(text=[[1, 2, "spk"]]))
self.assertEqual(result[0]["speaker"], "spk")
class OnlineSpeakerTests(unittest.IsolatedAsyncioTestCase):
"""使用独立的新鲜向量验证 A→B→A无需显卡和模型权重。"""
async def test_fresh_embeddings_preserve_a_b_a_and_reset(self):
runtime = AuxiliaryRuntime()
vectors = iter(([1, 0], [0, 1], [0.99, 0.01]))
runtime._extract_embedding_sync = lambda _: np.array(next(vectors), dtype=np.float32)
results = [await runtime.resolve_speaker("turn.wav", "test", i * 2000, i * 2000 + 1000) for i in range(3)]
self.assertEqual([r["speaker_id"] for r in results], [0, 1, 0])
self.assertEqual([c["count"] for c in runtime.speaker_clusters["test"]], [2, 1])
runtime.reset_speaker_session("test")
self.assertNotIn("test", runtime.speaker_clusters)
async def test_missing_embedding_does_not_copy_previous_cluster(self):
runtime = AuxiliaryRuntime()
runtime._extract_embedding_sync = lambda _: np.array([1, 0], dtype=np.float32)
await runtime.resolve_speaker("turn.wav", "test", 0, 1000)
runtime._extract_embedding_sync = lambda _: None
result = await runtime.resolve_speaker("short.wav", "test", 1000, 1200)
self.assertEqual(result["speaker_id"], -1)
self.assertEqual(runtime.speaker_clusters["test"][0]["count"], 1)
if __name__ == "__main__":
unittest.main()