170 lines
8.0 KiB
Python
170 lines
8.0 KiB
Python
"""辅助服务输出格式测试,确保不同 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()
|