Files
vrsub/tests/test_vad_profiler.py
T
cat-shark d0d96a8f89 feat: 每视频自适应 VAD 调参模块(信号分析 + 幻觉词扣分评分)
背景:实测 CJOD-255 全程 BGM 覆盖(56% 低能量、几乎无静音),固定 VAD
参数把音乐当语音 → whisper 全段解码 → 80% 碎片化 + 32% 漏句 + 敏感段丢失。

实现 nodes/vad_profiler.py:
- profile_audio:1s 能量网格分析 → 静音比例/BGM 覆盖/长停顿识别
- suggest_vad_parameters:按信号特征推荐 threshold/min_silence/speech_pad
  (BGM 覆盖 -> 降 threshold 增人声敏感;长停顿 -> 降 speech_pad 防时间漂移;
   静音占比高 -> 升 threshold 剔虚警)
- score_transcript:启发式评分(不用参考字幕,符合部署实际)——
  碎片率 + 幻觉词(感谢观看/晚安/音乐)扣分 + 平均字数适中
- vad_parameters_for_audio:信号分析 + 可选片段网格验证,选最优参数
- _grid_search_vad / _candidate_params:候选网格搜索与异常回退

测试 tests/test_vad_profiler.py:19 个用例覆盖信号分析、四个建议分支、
幻觉词/碎片评分、网格选优、异常回退、空音频/低采样率、候选展开。
2026-09-06 08:14:59 +08:00

253 lines
9.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""每视频自适应 VAD 调参测试(先红后绿)。
验证信号分析 profile_audio、启发式参数建议 suggest_vad_parameters、
转录质量评分 score_transcript(幻觉词扣分)与 vad_parameters_for_audio
(信号分析 + 片段网格验证)。
评分**不依赖参考字幕**(实际部署无参考),用转录质量 + 幻觉词扣分。
"""
from __future__ import annotations
import wave
from pathlib import Path
import pytest
from nodes.vad_profiler import (
AudioProfile,
HALLUCINATION_TOKENS,
profile_audio,
score_transcript,
suggest_vad_parameters,
vad_parameters_for_audio,
_pick_representative_start,
)
class _Seg:
"""模拟 whisper segment:仅含 text。"""
def __init__(self, text):
self.text = text
def _make_wav(path: Path, silence_seconds: int, voice_seconds: int) -> None:
"""生成 [静音 N 秒 + 语音 N 秒] 的 16kHz 单声道 WAV。
silence 用 0 样本(静音),voice 用较大振幅样本(语音)。
"""
import array
rate = 16000
silence = array.array("h", [0] * rate * silence_seconds)
voice = array.array("h", [8000] * rate * voice_seconds)
samples = silence + voice
with wave.open(str(path), "wb") as f:
f.setnchannels(1)
f.setsampwidth(2)
f.setframerate(rate)
f.writeframes(samples.tobytes())
def test_profile_audio_detects_bgm_heavy() -> None:
"""纯静音+语音的视频:silence_ratio 高、非 BGM 覆盖;mediam_rms 合理。"""
wav = Path("/tmp/test_vad_plain.wav")
_make_wav(wav, silence_seconds=40, voice_seconds=10)
p = profile_audio(wav, 16000)
assert p.silence_ratio > 0.5
assert p.bgm_heavy is False
assert p.duration_seconds == pytest.approx(50.0, abs=1)
def test_suggest_vad_parameters_plain_silence() -> None:
"""静音占比高且无长停顿 -> 建议 threshold 偏高(静音权重分支)。"""
wav = Path("/tmp/test_vad_plain2.wav")
# 交替短静音避免触发 long_silence<5s 连续静音)。
_make_wav(wav, silence_seconds=3, voice_seconds=1)
_make_wav(wav, silence_seconds=3, voice_seconds=1)
p = profile_audio(wav, 16000)
assert p.long_silence is False
params = suggest_vad_parameters(p)
assert params["threshold"] >= 0.5
def test_suggest_vad_parameters_bgm() -> None:
"""BGM 覆盖(低能量占比高但静音少)-> threshold 降低、min_silence 减小。"""
p = AudioProfile(
silence_ratio=0.1, # 几乎无静音
lowish_ratio=0.6, # 大量低能量(音乐)
bgm_heavy=True, # 直接标记 BGM 覆盖
long_silence=False,
rms_bins=[500] * 100,
)
params = suggest_vad_parameters(p)
assert params["threshold"] < 0.5 # 降低
assert params["min_silence_duration_ms"] < 1000 # 减小
def test_score_transcript_penalizes_hallucination() -> None:
"""幻觉词(感谢观看/晚安/音乐)多 -> 评分低。"""
good = [_Seg("ありがとうございます本日は"), _Seg("かしこまりました")]
bad = [_Seg("ご視聴ありがとうございました"), _Seg("おやすみなさい"), _Seg("音楽")]
assert score_transcript(good) > score_transcript(bad)
def test_score_transcript_penalizes_fragments() -> None:
"""碎片(纯单字/语气词)多 -> 评分低。"""
clean = [_Seg("今日はとても暑いですね"), _Seg("それでは始めましょう")]
frag = [_Seg("あ"), _Seg("うん"), _Seg("はい"), _Seg("あっ")]
assert score_transcript(clean) > score_transcript(frag)
def test_vad_parameters_for_audio_uses_grid_when_provider() -> None:
"""提供 whisper 回调时:从候选网格选评分最高的参数(threshold 最小者)。"""
wav = Path("/tmp/test_vad_grid.wav")
_make_wav(wav, silence_seconds=10, voice_seconds=5)
def fake_whisper(**kwargs):
# 假设最优 = threshold 最小的候选(suggest 0.5 时候选含 0.4)。
t = kwargs.get("vad_parameters", {}).get("threshold", 0.5)
if t <= 0.4:
# 最优组合:无幻觉的正常长句(分数最高)。
return [_Seg("今日はお客様のために精神整備を務めさせていただきます")]
# 其它组合:幻觉套话(分数低)。
return [_Seg("ご視聴ありがとうございました"), _Seg("おやすみなさい")]
params = vad_parameters_for_audio(
wav, sample_rate=16000, whisper_invoke=fake_whisper, run_dir=Path("/tmp")
)
# 网格应从候选里选出评分最高的 threshold=0.4 的组合。
assert params["threshold"] == pytest.approx(0.4, abs=0.1)
def test_vad_parameters_for_audio_fallback_without_provider() -> None:
"""无 whisper 回调:退化为信号分析建议,不报错。"""
wav = Path("/tmp/test_vad_noprov.wav")
_make_wav(wav, silence_seconds=10, voice_seconds=5)
params = vad_parameters_for_audio(wav, sample_rate=16000, whisper_invoke=None)
assert "threshold" in params
assert "min_silence_duration_ms" in params
def test_hallucination_tokens_included() -> None:
"""幻觉词集合应含常用套话(感谢观看/晚安/音乐)。"""
all_str = " ".join(HALLUCINATION_TOKENS)
assert "ご視聴ありがとうございました" in all_str
assert "おやすみなさい" in all_str
assert "音楽" in all_str
def test_suggest_vad_parameters_long_silence() -> None:
"""长停顿常见 -> 建议 threshold 0.5、speech_pad 200(防时间轴压缩漂移)。"""
p = AudioProfile(
silence_ratio=0.2, lowish_ratio=0.3, bgm_heavy=False, long_silence=True,
rms_bins=[500] * 100,
)
params = suggest_vad_parameters(p)
assert params["threshold"] == 0.5
assert params["speech_pad_ms"] == 200
def test_suggest_vad_parameters_regular() -> None:
"""常规音频(无特殊标记)-> 默认 0.5/1000/400。"""
p = AudioProfile(
silence_ratio=0.2, lowish_ratio=0.3, bgm_heavy=False, long_silence=False,
rms_bins=[1500] * 100,
)
params = suggest_vad_parameters(p)
assert params["threshold"] == 0.5
assert params["min_silence_duration_ms"] == 1000
assert params["speech_pad_ms"] == 400
def test_pick_representative_start_all_silence() -> None:
"""全静音音频:窗口恐怖沉默用 score>0.55 惩罚,仍返回起始 0。"""
p = AudioProfile(
rms_bins=[50] * 100, silence_ratio=0.9, lowish_ratio=0.9,
bgm_heavy=False, long_silence=False,
)
start = _pick_representative_start(p, 90)
assert start >= 0 and start < 10
def test_grid_search_exception_skipped() -> None:
"""网格搜索某组参数调用抛异常:跳过该组,不崩溃。"""
called = {"n": 0}
def failing_whisper(**kwargs):
called["n"] += 1
raise RuntimeError("mock download fail")
from nodes.vad_profiler import _grid_search_vad
params = _grid_search_vad(
failing_whisper, Path("/tmp/x.wav"), Path("/tmp"), 60, 0,
{"threshold": 0.5, "min_silence_duration_ms": 1000, "speech_pad_ms": 400},
)
# 全部失败时回退 suggested。
assert params == {"threshold": 0.5, "min_silence_duration_ms": 1000, "speech_pad_ms": 400}
assert called["n"] > 0
def test_candidate_params_expands_grid() -> None:
"""候选网格应含建议值邻域(threshold±0.1、silence 倍、pad 组合)。"""
from nodes.vad_profiler import _candidate_params
cands = _candidate_params(
{"threshold": 0.5, "min_silence_duration_ms": 1000, "speech_pad_ms": 400}
)
assert len(cands) == 27
assert {"threshold": 0.4, "min_silence_duration_ms": 500, "speech_pad_ms": 0} in cands
assert {"threshold": 0.6, "min_silence_duration_ms": 2000, "speech_pad_ms": 400} in cands
def test_median_rms_empty_returns_zero() -> None:
"""rms_bins 为空时 median_rms 返回 0。"""
p = AudioProfile(rms_bins=[])
assert p.median_rms == 0.0
def test_profile_audio_empty_wav() -> None:
"""空的 WAV(无样本)-> 返回 duration 0 的空 profile,不崩溃。"""
import array
wav = Path("/tmp/test_vad_empty.wav")
with wave.open(str(wav), "wb") as f:
f.setnchannels(1); f.setsampwidth(2); f.setframerate(16000)
f.writeframes(array.array("h", []).tobytes())
p = profile_audio(wav, 16000)
assert p.duration_seconds == 0.0
def test_grid_search_ignores_non_list() -> None:
"""网格搜索回调返回非 list(如 None)应被跳过,回退建议值。"""
from nodes.vad_profiler import _grid_search_vad
params = _grid_search_vad(
lambda **kw: None, Path("/tmp/x.wav"), Path("/tmp"), 60, 0,
{"threshold": 0.5, "min_silence_duration_ms": 1000, "speech_pad_ms": 400},
)
assert params == {"threshold": 0.5, "min_silence_duration_ms": 1000, "speech_pad_ms": 400}
def test_score_transcript_empty_list() -> None:
"""空列表评分返回 0。"""
assert score_transcript([]) == 0.0
def test_score_transcript_ignores_empty_text_seg() -> None:
"""含空文本的 segment 被跳过,不报错。"""
assert score_transcript([_Seg(""), _Seg("今日は暑いです")]) > 0
def test_profile_audio_different_sample_rate() -> None:
"""采样率与请求不一致时用 wave 实际 rate 近似,不崩溃。"""
import array
wav = Path("/tmp/test_vad_rate.wav")
rate = 8000
samples = array.array("h", [5000] * rate * 2) # 2s 语音
with wave.open(str(wav), "wb") as f:
f.setnchannels(1); f.setsampwidth(2); f.setframerate(rate)
f.writeframes(samples.tobytes())
p = profile_audio(wav, 16000)
assert p.duration_seconds == pytest.approx(2.0, abs=0.2)