from __future__ import annotations from types import SimpleNamespace from imagefind.speech_quality import ( aggregate_transcript_quality, normalize_language, pcm16_voiced_regions, select_language, transcript_quality, ) def test_aggregate_quality_rejects_repeated_short_hallucinations() -> None: segments = ([{"text": "好"}] * 11) + ([{"text": "拜拜"}] * 7) + ([{"text": "正常对话内容"}] * 4) result = aggregate_transcript_quality(segments, base_score=1.0) assert result.score < 0.7 assert "whole_repeated_phrase" in result.flags assert "whole_short_segment_dominance" in result.flags assert "whole_ending_hallucination" in result.flags assert ("拜拜", 7) in result.repeated_phrases def test_sports_samples_patterns_cannot_remain_ready() -> None: sports_2 = ([{"text": "好"}] * 11) + ([{"text": "拜拜"}] * 7) + ([{"text": "嗯"}] * 5) sports_2 += [{"text": f"正常对话{i}"} for i in range(19)] sports_3 = ([{"text": "拜拜"}] * 8) + ([{"text": "啊"}] * 7) + ([{"text": "嗯"}] * 7) sports_3 += ([{"text": "谢谢大家收看"}] * 3) + [{"text": f"正常内容{i}"} for i in range(32)] for segments in (sports_2, sports_3): result = aggregate_transcript_quality(segments, base_score=1.0) assert result.score < 0.7 assert "whole_repeated_phrase" in result.flags def test_traditional_ending_phrase_is_flagged() -> None: result = aggregate_transcript_quality([{"text": "謝謝大家收看"}] * 3, base_score=1.0) assert "whole_ending_hallucination" in result.flags def test_language_policy_treats_stream_metadata_as_weak_hint() -> None: assert select_language(["en", "zh"], policy="zh_priority") == "zh" assert select_language(["en", "zh"], policy="auto") == "en" assert select_language(["en", "en"], policy="zh_priority", stream_language="zh-CN") == "zh" assert select_language([], policy="zh_priority") == "zh" assert select_language([], policy="zh_priority", stream_language="eng") == "zh" assert select_language([], policy="auto", stream_language="eng") == "en" assert normalize_language("<|zh|>") == "zh" def test_chinese_priority_rejects_weak_or_implausible_language_detection() -> None: assert select_language(["jw", "nn", "jw"], policy="zh_priority") == "zh" assert select_language(["en"], policy="zh_priority") == "zh" assert select_language(["en", "en", "zh"], policy="zh_priority") == "zh" def test_chinese_priority_does_not_trust_short_korean_or_japanese_detection() -> None: # Whisper's one-token detector commonly confuses short/noisy Mandarin # clips with ko/ja. Choosing that language makes the whole decode wrong. assert select_language(["ko", "ko", "ko"], policy="zh_priority") == "zh" assert select_language(["ja", "ja"], policy="zh_priority") == "zh" assert select_language(["jw", "jw"], policy="auto") == "jw" def test_quality_accepts_normal_chinese_english_and_mixed_transcripts() -> None: samples = ( ("这是一个正常的中文语音识别结果,包含完整的句子和清晰的信息。", "zh"), ("This is a clear English transcript with enough useful information.", "en"), ("今天我们 discuss ImageFind 的 audio search 功能。", "zh"), ) for text, language in samples: quality = transcript_quality(text, 12_000, speech_ratio=0.7, expected_language=language) assert quality.accepted is True assert quality.score >= 0.9 assert quality.flags == () def test_quality_rejects_repeated_wrong_script_and_replacement_characters() -> None: repeated = transcript_quality( "කකකකකකකකකකකකකකකකකකකක", 12_000, speech_ratio=0.7, expected_language="zh", ) damaged = transcript_quality( "这是一段损坏的字幕�内容", 8_000, speech_ratio=0.6, expected_language="zh", ) assert repeated.accepted is False assert {"repeated_characters", "unexpected_script"}.issubset(repeated.flags) assert damaged.accepted is False assert "invalid_characters" in damaged.flags def test_quality_rejects_full_english_hallucination_when_chinese_is_forced() -> None: result = transcript_quality( "This entire segment was decoded in the wrong language.", 8_000, speech_ratio=0.8, expected_language="zh", ) assert result.accepted is False assert "language_script_conflict" in result.flags def test_quality_rejects_short_phrase_hallucination_but_accepts_silence() -> None: hallucination = transcript_quality( "谢谢观看", 30_000, speech_ratio=0.01, expected_language="zh", ) silence = transcript_quality("", 30_000, speech_ratio=0.0, expected_language="zh") assert hallucination.accepted is False assert "short_hallucination" in hallucination.flags assert silence.accepted is True assert silence.units == 0 def test_quality_does_not_reject_sparse_legitimate_speech_by_chunk_duration() -> None: quality = transcript_quality("今天天气不错", 30_000, speech_ratio=0.08, expected_language="zh") assert quality.accepted is True assert "too_little_text" not in quality.flags def test_vad_regions_bridge_short_gaps_and_skip_short_noise(monkeypatch) -> None: decisions = [False] * 4 + [True] * 10 + [False] * 8 + [True] * 10 + [False] * 20 + [True] * 3 class Detector: def __init__(self, _mode): self.index = 0 def is_speech(self, _frame, _sample_rate): value = decisions[self.index] self.index += 1 return value monkeypatch.setitem(__import__("sys").modules, "webrtcvad", SimpleNamespace(Vad=Detector)) raw = b"\0\0" * 480 * len(decisions) regions = pcm16_voiced_regions(raw, 16_000) assert len(regions) == 1 start, end, ratio = regions[0] assert start < 4 * 480 assert end > 32 * 480 assert 0 < ratio < 1