Files
vrsub/tests/test_subtitle_correction.py
T
cat-shark f2895ee103 feat: 字幕领域纠错模块(通用领域词表 + 有效上下文,不过拟合)
背景:成人视频 ASR 常把性器官(チンポ/マンコ)听错成近音词(手先/
チェーンバー 等),翻译逐字直译导致与画面严重不符。

实验结论:
- 硬编码 ASR 误听例子(手先→肉棒)过拟合——换视频的误听词就失效;
- 通用领域词表(只列性器官的常见日文词+中文)+ 有效上下文 + 通用引导
  可泛化——对从没见过的误听(バナナ/マンゴー→性器官)也能按语境推断。

实现 nodes/subtitle_correction.py:
- PROPER_SESSION_WORDS:通用性器官领域词表(不含 ASR 误听噪声词)
- _is_fragment:纯语气词/碎片过滤
- _build_context:目标 ±60s 有效上下文(过滤碎片,保留动作链)
- _system_prompt:通用引导(无具体误听例子)
- _extract_target_line:从 LLM 输出解析目标行译文
- invoke:逐条领域纠错,产出 corrected.srt

测试 tests/test_subtitle_correction.py(12 个):
- 碎片判定/上下文过滤/目标解析/提示词无噪声例子(9 快速单测)
- 真实 LLM 泛化回归(对未出现误听词也能推断性器官)+ 端到端 invoke
2026-09-06 10:44:09 +08:00

203 lines
8.1 KiB
Python

"""字幕领域纠错节点测试(先红后绿)。
验证:
1. _is_fragment 正确过滤纯语气词/碎片(非过拟合)。
2. _build_context 只保留有效上下文(±60s 内实义句),目标保留。
3. _extract_target_line 从 LLM 输出解析目标行译文。
4. 真实 LLM 集成:对"从未在提示词出现的 ASR 误听"(バナナ/マンゴー→性器官)
能按领域词表+上下文推断,证明不过拟合(V5 验证结果固化为回归)。
提示词不含任何 ASR 误听具体例子,只有通用领域词表。
"""
from __future__ import annotations
import os
from pathlib import Path
import pytest
from nodes.subtitle_correction import (
PROPER_SESSION_WORDS,
_build_context,
_extract_target_line,
_is_fragment,
_system_prompt,
)
from wov_sdk.models import InvokeRequest
WORKSPACE = Path(__file__).resolve().parent.parent
TRANSCRIPT = Path("/home/cat/Downloads/192.168.123.70/202609060835/run_af2987b161a3/steps/asr/transcript.srt")
# ---------------------------------------------------------------------------
# 单元测试:碎片判定 / 上下文构建 / 目标行解析
# ---------------------------------------------------------------------------
def test_is_fragment_pure_words() -> None:
"""纯语气词/单音节应判为碎片。"""
for t in ["あ", "ん", "うん", "はい", "あっ", "あー", "あ〜", "ああ", ""]:
assert _is_fragment(t), f"'{t}' 应为碎片"
def test_is_fragment_meaningful() -> None:
"""有实义的句子不应判为碎片(即使含假名)。"""
for t in ["気持ちいい", "難しい", "ごめんなさい", "手先が違う所に当たり合い"]:
assert not _is_fragment(t), f"'{t}' 不应为碎片"
def test_build_context_filters_fragments() -> None:
"""上下文过滤语气词,但目标条目始终保留。"""
entries = [
{"start": 0.0, "end": 1.0, "text": "あ"},
{"start": 2.0, "end": 3.0, "text": "気持ちいい"},
{"start": 4.0, "end": 5.0, "text": "手先が当たる"},
{"start": 6.0, "end": 7.0, "text": "うん"},
]
ctx = _build_context(entries, target_index=2)
assert "気持ちいい" in ctx
assert "手先が当たる" in ctx
assert "あ\n" not in ctx # 语气词被过滤
assert "うん" not in ctx
def test_build_context_keeps_target() -> None:
"""目标条目即使本身是语气词也保留并标记。"""
entries = [
{"start": 0.0, "end": 1.0, "text": "あ"},
{"start": 2.0, "end": 3.0, "text": "うん"},
]
ctx = _build_context(entries, target_index=1)
assert "<-- 目标" in ctx
assert "うん" in ctx
def test_extract_target_line() -> None:
"""从 LLM 输出解析与目标时间最接近的译文行。"""
content = "6656.60 好舒服\n6704.10 啊,肉棒撞到了别的地方\n6708.10 啊,肉棒顶到舒服的地方"
assert "好舒服" in _extract_target_line(content, 6656.60)
# 目标 6704 附近
assert "肉棒撞到了别的地方" in _extract_target_line(content, 6704.10)
def test_system_prompt_has_generic_domain_terms_no_noise_examples() -> None:
"""系统提示词含通用领域词表,但不含任何 ASR 误听噪声词(不过拟合)。"""
prompt = _system_prompt("zh-CN")
assert "チンポ" in prompt
assert "マンコ" in prompt
assert "バナナ" in prompt or "マンゴー" in prompt
# 关键:绝不含具体 ASR 误听形式(本次视频的 チェーンバー/手先)。
assert "チェーンバー" not in prompt
assert "手先" not in prompt
def test_proper_session_words_are_generic() -> None:
"""领域词表只含通用日文性器官词,不含误听噪声词。"""
assert "チンポ" in PROPER_SESSION_WORDS
assert "マンコ" in PROPER_SESSION_WORDS
assert "チェーンバー" not in PROPER_SESSION_WORDS # 非通用词
assert "手先" not in PROPER_SESSION_WORDS # 非通用词
# ---------------------------------------------------------------------------
# 真实 LLM 集成:泛化验证(B 组场景)
# ---------------------------------------------------------------------------
def _llm_ok() -> bool:
try:
from dotenv import load_dotenv
load_dotenv(WORKSPACE / ".env")
except Exception:
pass
return bool(os.getenv("LLM_API_KEY"))
@pytest.mark.integration
def test_generic_correction_generalizes_to_unseen_mishearing() -> None:
"""泛化回归:对'未在提示词出现'的误听(バナナ/マンゴー→性器官)能正确推断。
提示词只有通用领域词表(含バナナ/マンゴー),没有具体误听例子。
若 LLM 能按上下文把バナナ理解为肉棒、マンゴー理解为小穴,证明不过拟合。
"""
if not _llm_ok():
pytest.skip("未配置 LLM_API_KEY,跳过真实 LLM 集成测试")
from nodes.subtitle_correction import correct_entry
# 模拟新视频:ASR 把 チンポ/マンコ 听成 バナナ/マンゴー(提示词中仅有通用词表)。
entries = [
{"start": 1200.0, "end": 1203.0, "text": "相手がバナナをしゃぶってくれて"},
{"start": 1203.0, "end": 1206.0, "text": "そろそろマンゴーが濡れてきました"},
{"start": 1206.0, "end": 1209.0, "text": "気持ちいいところに当たってるね"},
{"start": 1220.0, "end": 1224.0, "text": "もっとマンゴーを舐めてください"},
{"start": 1224.0, "end": 1226.0, "text": "いっぱい出してね"},
]
# 目标改为含误听词'マンゴー'的条目(索引 3):验证 LLM 结合上下文和
# 领域词表把'マンゴー'推断为小穴,而非字面译'芒果'。
target = correct_entry(entries[3], entries, 3, {"target_language": "zh-CN"})
# 泛化判定:输出应含性器官语义(肉棒/阴部/敏感处等),而非字面"香蕉/芒果"。
flagged = [k for k in ("肉棒", "鸡巴", "阴部", "小穴", "敏感") if k in target]
assert flagged, (
f"泛化失败:模型仍字面直译,输出'{target}'(应结合领域表推断性器官)"
)
@pytest.mark.integration
def test_invoke_end_to_end_real_transcript(tmp_path) -> None:
"""真实 invoke 端到端:读真实 transcript,逐条纠错,产出 corrected.srt。
覆盖 invoke 全流程(文件校验、逐条纠错、SRT 序列化、写文件)。
"""
if not _llm_ok():
pytest.skip("未配置 LLM_API_KEY,跳过真实 LLM 集成测试")
if not TRANSCRIPT.is_file():
pytest.skip("缺少真实 transcript.srt,跳过")
from nodes.subtitle_correction import invoke
out = tmp_path / "out"
resp = invoke(InvokeRequest(
run_id="corr_e2e",
node_instance_id="",
inputs={"srt_uri": str(TRANSCRIPT)},
params={"target_language": "zh-CN"},
output_dir=str(out),
))
assert resp.status == "completed", resp.error
assert Path(resp.outputs["srt_uri"]).is_file()
content = Path(resp.outputs["srt_uri"]).read_text(encoding="utf-8")
assert "--> " in content # 合法 SRT
@pytest.mark.integration
def test_invoke_missing_srt_uri(tmp_path) -> None:
"""缺少 srt_uri -> failed。"""
from nodes.subtitle_correction import invoke
resp = invoke(InvokeRequest(
run_id="x", node_instance_id="", inputs={}, output_dir=str(tmp_path)
))
assert resp.status == "failed"
def test_extract_target_line_no_match_returns_empty() -> None:
"""LLM 输出无法匹配目标时间时返回空串。"""
from nodes.subtitle_correction import _extract_target_line
assert _extract_target_line("随便一段话没有数字", 1234.5) == ""
def test_serialize_srt_roundtrip() -> None:
"""SRT 序列化往返:条目 -> 文本 -> 再解析条数一致。"""
from nodes.subtitle_correction import _serialize_srt
from tests.realdata_contract import parse_srt_entries
entries = [
{"start": 0.0, "end": 2.0, "text": "你好"},
{"start": 2.0, "end": 4.0, "text": "世界"},
]
srt = _serialize_srt(entries)
assert "00:00:00,000 --> 00:00:02,000" in srt
assert len(parse_srt_entries(srt)) == 2