Spaces:
Running on Zero
Running on Zero
Download tests/test_quality_runtime.py from voidful/BlueMagpie-TTS-Demo: direct link, hf CLI and curl.
- Browser
- Download file 48.6 kB
-
https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/7e7df2a5f9ea6839a43a3ed880a32d7ecaea9176/tests/test_quality_runtime.py
- Command line
-
hf download hf://spaces/voidful/BlueMagpie-TTS-Demo@7e7df2a5f9ea6839a43a3ed880a32d7ecaea9176/tests/test_quality_runtime.py
-
curl -L -o test_quality_runtime.py https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/7e7df2a5f9ea6839a43a3ed880a32d7ecaea9176/tests/test_quality_runtime.py
48.6 kB
| import math | |
| from types import SimpleNamespace | |
| import numpy as np | |
| import pytest | |
| import torch | |
| from production import split_leading_clause, split_text_for_tts | |
| from quality_runtime import ( | |
| BASE_GENERATION_POLICY, | |
| SAFE_DURATION_GENERATION_POLICY, | |
| WHISPER_MODEL_ID, | |
| WHISPER_REVISION, | |
| CandidateObservation, | |
| CascadeResult, | |
| ChunkCandidateArtifact, | |
| FinalOutputRejectedError, | |
| LazyWhisperASR, | |
| NoQualifiedCandidateError, | |
| TrajectoryGateResult, | |
| WhisperRuntime, | |
| active_voiced_duration_seconds, | |
| candidate_chunk_transition_score, | |
| candidate_limit_for_chunk_budget, | |
| cosine_similarity, | |
| generation_policy_for_candidate_offset, | |
| load_pinned_whisper_runtime, | |
| prepare_candidate_audio, | |
| qualify_trajectory_with_joined_output, | |
| require_verified_final_output, | |
| resolve_request_seed, | |
| run_adaptive_cascade, | |
| select_k_candidate_sequences, | |
| speaker_embedding_from_audio, | |
| speaker_evidence_from_audio, | |
| transcribe_whisper, | |
| verify_candidate, | |
| verify_trajectory, | |
| ) | |
| class _Factory: | |
| def __init__(self, value): | |
| self.value = value | |
| self.calls = [] | |
| def from_pretrained(self, *args, **kwargs): | |
| self.calls.append((args, kwargs)) | |
| return self.value | |
| def test_candidate_generation_policy_mapping_uses_safe_duration_after_offset_zero(): | |
| base = generation_policy_for_candidate_offset(0) | |
| assert base is BASE_GENERATION_POLICY | |
| assert (base.name, base.cjk_cps, base.ascii_cps, base.hard_stop_margin_steps) == ( | |
| "base", | |
| 5.2, | |
| 4.6, | |
| 1, | |
| ) | |
| for offset in (1, 2, 5, 9): | |
| safe = generation_policy_for_candidate_offset(offset) | |
| assert safe is SAFE_DURATION_GENERATION_POLICY | |
| assert ( | |
| safe.name, | |
| safe.cjk_cps, | |
| safe.ascii_cps, | |
| safe.hard_stop_margin_steps, | |
| ) == ("safe_duration", 4.6, 4.0, 1) | |
| def test_candidate_generation_policy_mapping_rejects_invalid_offsets(offset): | |
| with pytest.raises(ValueError, match="non-negative integer"): | |
| generation_policy_for_candidate_offset(offset) | |
| def test_explicit_request_seed_is_forwarded_without_calling_random_factory(): | |
| calls = [] | |
| def factory(limit): | |
| calls.append(limit) | |
| return 456 | |
| assert resolve_request_seed(123, factory) == 123 | |
| assert calls == [] | |
| assert resolve_request_seed(None, factory) == 456 | |
| assert calls == [2**31] | |
| def test_request_seed_validation_fails_closed(seed): | |
| factory = lambda _limit: 2**31 if seed is None else 0 | |
| with pytest.raises(ValueError, match=r"\[0, 2147483648\)"): | |
| resolve_request_seed(seed, factory) | |
| def test_quality_gate_accepts_only_eval_canonicalized_pronoun_homophones(): | |
| result = verify_candidate( | |
| CandidateObservation( | |
| target_text="她提醒我", | |
| transcript_text="他提醒我", | |
| audio_duration_seconds=1.0, | |
| ) | |
| ) | |
| assert result.passed | |
| assert result.comparison.cer == 0.0 | |
| assert result.comparison.prefix_cer == 0.0 | |
| assert result.comparison.suffix_cer == 0.0 | |
| class _FakeWhisperModel: | |
| def __init__(self): | |
| self.device = None | |
| self.evaluated = False | |
| self.generate_calls = [] | |
| def to(self, device): | |
| self.device = torch.device(device) | |
| return self | |
| def eval(self): | |
| self.evaluated = True | |
| return self | |
| def generate(self, features, **kwargs): | |
| self.generate_calls.append((features, kwargs)) | |
| return torch.tensor([[1, 2, 3]]).repeat(features.shape[0], 1) | |
| class _FakeProcessor: | |
| def __init__(self): | |
| self.calls = [] | |
| def __call__(self, waveform, **kwargs): | |
| copied = ( | |
| [segment.copy() for segment in waveform] | |
| if isinstance(waveform, list) | |
| else waveform.copy() | |
| ) | |
| self.calls.append((copied, kwargs)) | |
| batch_size = len(waveform) if isinstance(waveform, list) else 1 | |
| return SimpleNamespace(input_features=torch.ones(batch_size, 80, 10)) | |
| def batch_decode(self, token_ids, **kwargs): | |
| assert all(row == [1, 2, 3] for row in token_ids.tolist()) | |
| assert kwargs == {"skip_special_tokens": True} | |
| return [" 合成內容完整。 "] * token_ids.shape[0] | |
| class _StatsEncoder: | |
| def __init__(self, *, invalid=False): | |
| self.lengths = [] | |
| self.wav_lens = [] | |
| self.invalid = invalid | |
| def encode_batch(self, tensor, wav_lens=None): | |
| self.lengths.append(tensor.shape[-1]) | |
| if wav_lens is None: | |
| wav_lens = torch.ones(tensor.shape[0], device=tensor.device) | |
| self.wav_lens.append(wav_lens.detach().cpu().numpy()) | |
| if self.invalid: | |
| return torch.full( | |
| (tensor.shape[0], 1, 2), | |
| math.nan, | |
| device=tensor.device, | |
| ) | |
| rows = [] | |
| for row, relative_length in zip(tensor, wav_lens, strict=True): | |
| sample_count = max(1, int(round(float(relative_length) * tensor.shape[-1]))) | |
| active = row[:sample_count] | |
| rows.append( | |
| torch.stack( | |
| ( | |
| active.mean(), | |
| active.std(unbiased=False), | |
| active.abs().amax(), | |
| ) | |
| ) | |
| ) | |
| return torch.stack(rows) | |
| def _tone(sample_rate=16_000, seconds=1.0, amplitude=0.2, frequency=220.0): | |
| timeline = np.arange(round(sample_rate * seconds), dtype=np.float32) / sample_rate | |
| return amplitude * np.sin(2.0 * np.pi * frequency * timeline).astype(np.float32) | |
| def _gate_result(passed, score=0.0): | |
| return TrajectoryGateResult( | |
| passed=passed, | |
| candidate_results=(), | |
| score=score if passed else math.inf, | |
| rejection_reasons=() if passed else ("rejected",), | |
| ) | |
| def _chunk_verification(target, transcript, *, embedding, rms_db): | |
| observation = CandidateObservation( | |
| target_text=target, | |
| transcript_text=transcript, | |
| audio_duration_seconds=2.0, | |
| speaker_similarity=0.8, | |
| begin_speaker_similarity=0.8, | |
| end_speaker_similarity=0.8, | |
| ) | |
| return observation, ChunkCandidateArtifact( | |
| speaker_embedding=( | |
| None if embedding is None else np.asarray(embedding, dtype=np.float32) | |
| ), | |
| rms_db=rms_db, | |
| ) | |
| def _whole_speaker_verification(*, similarity, boundary_drop, passed=True): | |
| target = "完整而且穩定的候選內容" | |
| observation = CandidateObservation( | |
| target_text=target, | |
| transcript_text=target if passed else "錯誤內容", | |
| audio_duration_seconds=2.5, | |
| speaker_similarity=similarity, | |
| begin_speaker_similarity=0.60, | |
| end_speaker_similarity=0.60 - boundary_drop, | |
| ) | |
| artifact = ChunkCandidateArtifact( | |
| speaker_embedding=np.array([1.0, 0.0], dtype=np.float32), | |
| rms_db=-20.0, | |
| ) | |
| return verify_trajectory( | |
| [observation], | |
| chunk_artifacts=[artifact], | |
| min_speaker_similarity=0.10, | |
| max_boundary_speaker_drop=0.10, | |
| ) | |
| def _joined_verification(target, transcript): | |
| return verify_trajectory( | |
| [ | |
| CandidateObservation( | |
| target_text=target, | |
| transcript_text=transcript, | |
| audio_duration_seconds=3.0, | |
| speaker_similarity=0.8, | |
| begin_speaker_similarity=0.8, | |
| end_speaker_similarity=0.8, | |
| ) | |
| ] | |
| ) | |
| def test_pinned_whisper_loader_uses_exact_revision_without_global_download(): | |
| processor = _FakeProcessor() | |
| model = _FakeWhisperModel() | |
| processor_factory = _Factory(processor) | |
| model_factory = _Factory(model) | |
| runtime = load_pinned_whisper_runtime( | |
| device="cpu", | |
| processor_factory=processor_factory, | |
| model_factory=model_factory, | |
| ) | |
| assert runtime.processor is processor | |
| assert runtime.model is model | |
| assert runtime.device == torch.device("cpu") | |
| assert runtime.dtype == torch.float32 | |
| assert model.evaluated | |
| assert processor_factory.calls == [ | |
| ((WHISPER_MODEL_ID,), {"revision": WHISPER_REVISION}) | |
| ] | |
| model_args, model_kwargs = model_factory.calls[0] | |
| assert model_args == (WHISPER_MODEL_ID,) | |
| assert model_kwargs["revision"] == WHISPER_REVISION | |
| assert model_kwargs["torch_dtype"] == torch.float32 | |
| assert model_kwargs["use_safetensors"] is True | |
| def test_lazy_whisper_runtime_loads_once_and_transcription_is_deterministic(): | |
| processor = _FakeProcessor() | |
| model = _FakeWhisperModel() | |
| runtime = WhisperRuntime(processor, model, torch.device("cpu"), torch.float32) | |
| loads = [] | |
| lazy_asr = LazyWhisperASR(lambda: loads.append("load") or runtime) | |
| audio_8khz = _tone(sample_rate=8_000, seconds=0.5) | |
| first = transcribe_whisper(audio_8khz, 8_000, lazy_asr=lazy_asr) | |
| second = transcribe_whisper(audio_8khz, 8_000, lazy_asr=lazy_asr) | |
| assert first == second == "合成內容完整。" | |
| assert loads == ["load"] | |
| assert processor.calls[0][0].shape == (8_000,) | |
| assert processor.calls[0][1] == { | |
| "sampling_rate": 16_000, | |
| "return_tensors": "pt", | |
| } | |
| _, generation_kwargs = model.generate_calls[0] | |
| assert generation_kwargs == { | |
| "language": "zh", | |
| "task": "transcribe", | |
| "do_sample": False, | |
| "num_beams": 1, | |
| "max_new_tokens": 128, | |
| } | |
| def test_long_whisper_audio_is_batched_below_thirty_second_limit(): | |
| processor = _FakeProcessor() | |
| model = _FakeWhisperModel() | |
| runtime = WhisperRuntime(processor, model, torch.device("cpu"), torch.float32) | |
| audio = np.concatenate( | |
| ( | |
| _tone(seconds=21.0), | |
| np.zeros(16_000, dtype=np.float32), | |
| _tone(seconds=21.0, frequency=240.0), | |
| np.zeros(16_000, dtype=np.float32), | |
| _tone(seconds=21.0, frequency=260.0), | |
| ) | |
| ) | |
| transcript = transcribe_whisper(audio, 16_000, runtime=runtime, max_new_tokens=440) | |
| segments = processor.calls[0][0] | |
| assert isinstance(segments, list) | |
| assert len(segments) == 3 | |
| assert all(segment.size <= 30 * 16_000 for segment in segments) | |
| assert transcript == "合成內容完整。 合成內容完整。 合成內容完整。" | |
| assert len(model.generate_calls) == 1 | |
| assert model.generate_calls[0][1]["max_new_tokens"] == 440 | |
| def test_transcription_rejects_nonfinite_audio_and_conflicting_injection(): | |
| processor = _FakeProcessor() | |
| model = _FakeWhisperModel() | |
| runtime = WhisperRuntime(processor, model, torch.device("cpu"), torch.float32) | |
| with pytest.raises(ValueError, match="non-finite"): | |
| transcribe_whisper(np.array([0.0, math.nan]), 16_000, runtime=runtime) | |
| with pytest.raises(ValueError, match="either runtime or lazy_asr"): | |
| transcribe_whisper( | |
| np.ones(20), | |
| 16_000, | |
| runtime=runtime, | |
| lazy_asr=LazyWhisperASR(lambda: runtime), | |
| ) | |
| def test_candidate_audio_data_or_asr_value_error_becomes_rejection_evidence(): | |
| calls = [] | |
| def transcriber(waveform, sample_rate): | |
| calls.append((waveform.copy(), sample_rate)) | |
| return " 內容完整 " | |
| prepared = prepare_candidate_audio(_tone(seconds=0.25), 16_000, transcriber=transcriber) | |
| assert prepared is not None | |
| assert prepared.transcript_text == "內容完整" | |
| assert prepared.duration_seconds == pytest.approx(0.25) | |
| assert len(calls) == 1 | |
| assert prepare_candidate_audio( | |
| [0.0, math.nan], | |
| 16_000, | |
| transcriber=transcriber, | |
| ) is None | |
| assert prepare_candidate_audio( | |
| _tone(seconds=0.25), | |
| 16_000, | |
| transcriber=lambda *_: (_ for _ in ()).throw(ValueError("bad candidate")), | |
| ) is None | |
| assert len(calls) == 1 | |
| def test_candidate_audio_infrastructure_error_propagates_fail_closed(): | |
| with pytest.raises(RuntimeError, match="ASR unavailable"): | |
| prepare_candidate_audio( | |
| _tone(seconds=0.25), | |
| 16_000, | |
| transcriber=lambda *_: (_ for _ in ()).throw(RuntimeError("ASR unavailable")), | |
| ) | |
| def test_ecapa_ndarray_helper_trims_resamples_normalizes_and_handles_channels(): | |
| encoder = _StatsEncoder() | |
| tone = _tone(sample_rate=8_000, seconds=1.0) | |
| padded = np.concatenate((np.zeros(4_000), tone, np.zeros(4_000))).astype(np.float32) | |
| stereo = np.stack((padded, padded), axis=0) | |
| embedding = speaker_embedding_from_audio(stereo, 8_000, encoder) | |
| assert embedding.shape == (3,) | |
| assert np.isfinite(embedding).all() | |
| assert np.linalg.norm(embedding) == pytest.approx(1.0, abs=1e-6) | |
| assert 16_000 <= encoder.lengths[0] < 32_000 | |
| assert len(encoder.lengths) == 1 | |
| def test_ecapa_ndarray_helper_fails_closed_on_silence_nan_and_bad_embedding(): | |
| with pytest.raises(ValueError, match="active speech"): | |
| speaker_embedding_from_audio(np.zeros(16_000), 16_000, _StatsEncoder()) | |
| with pytest.raises(ValueError, match="non-finite samples"): | |
| speaker_embedding_from_audio( | |
| np.array([0.0, math.nan], dtype=np.float32), | |
| 16_000, | |
| _StatsEncoder(), | |
| ) | |
| with pytest.raises(ValueError, match="invalid embedding"): | |
| speaker_embedding_from_audio(_tone(), 16_000, _StatsEncoder(invalid=True)) | |
| def test_cosine_and_speaker_evidence_are_finite_and_directional(): | |
| encoder = _StatsEncoder() | |
| audio = np.concatenate((_tone(seconds=1.0), _tone(seconds=1.0, amplitude=0.1))) | |
| anchor = speaker_embedding_from_audio(audio, 16_000, encoder) | |
| evidence = speaker_evidence_from_audio( | |
| audio, | |
| 16_000, | |
| encoder, | |
| anchor, | |
| edge_seconds=0.75, | |
| ) | |
| # Anchor extraction is one call; whole/begin/end evidence is one batched | |
| # call instead of three repeated ECAPA forwards for the same candidate. | |
| assert len(encoder.lengths) == 2 | |
| assert encoder.wav_lens[-1].shape == (3,) | |
| assert evidence.similarity == pytest.approx(1.0, abs=1e-5) | |
| assert -1.0 <= evidence.begin_similarity <= 1.0 | |
| assert -1.0 <= evidence.end_similarity <= 1.0 | |
| assert evidence.boundary_drop == max( | |
| 0.0, | |
| evidence.begin_similarity - evidence.end_similarity, | |
| ) | |
| assert evidence.active_duration_seconds > 1.5 | |
| assert cosine_similarity(anchor, anchor) == pytest.approx(1.0) | |
| with pytest.raises(ValueError): | |
| cosine_similarity([0.0, 0.0], [0.0, 0.0]) | |
| def test_long_speaker_evidence_uses_one_bounded_window_batch(): | |
| encoder = _StatsEncoder() | |
| audio = _tone(seconds=20.0) | |
| anchor = np.array([0.0, 0.0, 1.0], dtype=np.float32) | |
| evidence = speaker_evidence_from_audio(audio, 16_000, encoder, anchor) | |
| assert len(encoder.lengths) == 1 | |
| assert encoder.lengths[0] <= 3 * 16_000 | |
| assert encoder.wav_lens[0].shape == (6,) | |
| assert evidence.active_duration_seconds == pytest.approx(20.0, abs=0.05) | |
| assert evidence.speaker_embedding.shape == (3,) | |
| assert math.isfinite(evidence.active_rms_db) | |
| def test_active_union_pace_ignores_one_or_three_second_internal_silence(): | |
| def waveform(pause_seconds): | |
| return np.concatenate( | |
| ( | |
| _tone(seconds=0.75), | |
| np.zeros(round(pause_seconds * 16_000), dtype=np.float32), | |
| _tone(seconds=0.75, frequency=260.0), | |
| ) | |
| ) | |
| one_second = active_voiced_duration_seconds(waveform(1.0), 16_000) | |
| three_seconds = active_voiced_duration_seconds(waveform(3.0), 16_000) | |
| assert one_second == pytest.approx(three_seconds, abs=1.0e-9) | |
| assert 8 / one_second == pytest.approx(8 / three_seconds, abs=1.0e-9) | |
| assert one_second < 1.6 | |
| def test_active_union_duration_controls_speaker_eligibility_at_1_5_seconds(): | |
| below = active_voiced_duration_seconds(_tone(seconds=1.49), 16_000) | |
| boundary = active_voiced_duration_seconds(_tone(seconds=1.50), 16_000) | |
| assert below == pytest.approx(1.49, abs=1.0e-9) | |
| assert boundary == pytest.approx(1.50, abs=1.0e-9) | |
| common = { | |
| "target_text": "這是一段完整而且穩定的合成內容", | |
| "transcript_text": "這是一段完整而且穩定的合成內容", | |
| } | |
| short_result = verify_candidate( | |
| CandidateObservation(**common, audio_duration_seconds=below) | |
| ) | |
| boundary_result = verify_candidate( | |
| CandidateObservation(**common, audio_duration_seconds=boundary) | |
| ) | |
| assert short_result.passed | |
| assert not short_result.speaker_gate_applied | |
| assert not boundary_result.passed | |
| assert boundary_result.speaker_gate_applied | |
| assert "missing_speaker_evidence" in boundary_result.rejection_reasons | |
| def test_short_candidate_uses_exact_semantics_but_bypasses_speaker_hard_gate(): | |
| passed = verify_candidate( | |
| CandidateObservation( | |
| target_text="謝謝你", | |
| transcript_text="謝謝你", | |
| audio_duration_seconds=1.49, | |
| ) | |
| ) | |
| failed = verify_candidate( | |
| CandidateObservation( | |
| target_text="謝謝你", | |
| transcript_text="謝謝", | |
| audio_duration_seconds=1.0, | |
| ) | |
| ) | |
| assert passed.passed | |
| assert not passed.speaker_gate_applied | |
| assert passed.speaker_similarity is None | |
| assert passed.score == 0.0 | |
| assert not failed.passed | |
| assert failed.rejection_reasons == ("semantic_gate",) | |
| def test_candidate_pace_gate_is_optional_and_fails_closed_when_enabled(): | |
| base = CandidateObservation( | |
| target_text="內容完整", | |
| transcript_text="內容完整", | |
| audio_duration_seconds=1.0, | |
| ) | |
| assert verify_candidate(base).passed | |
| assert verify_candidate( | |
| CandidateObservation(**{**base.__dict__, "pace_cps": 4.2}), | |
| max_pace_cps=4.3, | |
| ).passed | |
| missing = verify_candidate(base, max_pace_cps=4.3) | |
| fast = verify_candidate( | |
| CandidateObservation(**{**base.__dict__, "pace_cps": 4.31}), | |
| max_pace_cps=4.3, | |
| ) | |
| assert missing.rejection_reasons == ("missing_pace_evidence",) | |
| assert fast.rejection_reasons == ("pace_too_fast",) | |
| def test_long_candidate_requires_speaker_evidence_and_clamps_negative_drop(): | |
| observation = CandidateObservation( | |
| target_text="這是一段完整而且穩定的合成內容", | |
| transcript_text="這是一段完整而且穩定的合成內容", | |
| audio_duration_seconds=2.0, | |
| speaker_similarity=0.70, | |
| begin_speaker_similarity=0.50, | |
| end_speaker_similarity=0.60, | |
| ) | |
| result = verify_candidate(observation) | |
| assert result.passed | |
| assert result.speaker_gate_applied | |
| assert result.boundary_speaker_drop == 0.0 | |
| assert result.score == pytest.approx(0.05 * 0.30) | |
| missing = verify_candidate( | |
| CandidateObservation( | |
| target_text=observation.target_text, | |
| transcript_text=observation.transcript_text, | |
| audio_duration_seconds=2.0, | |
| ) | |
| ) | |
| assert not missing.passed | |
| assert "missing_speaker_evidence" in missing.rejection_reasons | |
| def test_candidate_rejects_boundary_drop_low_similarity_truncation_and_tail(): | |
| base = dict( | |
| target_text="這是一段完整而且穩定的合成內容", | |
| transcript_text="這是一段完整而且穩定的合成內容", | |
| audio_duration_seconds=2.0, | |
| speaker_similarity=0.7, | |
| begin_speaker_similarity=0.5, | |
| end_speaker_similarity=0.5, | |
| ) | |
| boundary = verify_candidate( | |
| CandidateObservation(**{**base, "begin_speaker_similarity": 0.8}) | |
| ) | |
| similarity = verify_candidate( | |
| CandidateObservation(**{**base, "speaker_similarity": 0.09}) | |
| ) | |
| truncation = verify_candidate(CandidateObservation(**{**base, "truncated": True})) | |
| tail = verify_candidate( | |
| CandidateObservation(**{**base, "transcript_text": base["target_text"] + "多講"}) | |
| ) | |
| assert boundary.rejection_reasons == ("boundary_speaker_drop",) | |
| assert similarity.rejection_reasons == ("speaker_similarity",) | |
| assert truncation.rejection_reasons == ("truncated",) | |
| assert "semantic_gate" in tail.rejection_reasons | |
| # ASR comparison canonicalizes both sides to Simplified Chinese. | |
| assert tail.comparison.extra_tail == "多讲" | |
| def test_candidate_and_trajectory_fail_closed_for_bad_values_or_any_bad_chunk(): | |
| good_short = CandidateObservation("內容完整", "內容完整", 1.0) | |
| bad_short = CandidateObservation("句尾完整", "句尾", 1.0) | |
| invalid = verify_candidate( | |
| CandidateObservation("內容完整", "內容完整", math.nan), | |
| max_cer=math.nan, | |
| ) | |
| passed_trajectory = verify_trajectory([good_short, good_short]) | |
| failed_trajectory = verify_trajectory([good_short, bad_short]) | |
| assert not invalid.passed | |
| assert "invalid_audio_duration" in invalid.rejection_reasons | |
| assert "invalid_gate_config" in invalid.rejection_reasons | |
| assert passed_trajectory.passed | |
| assert passed_trajectory.score == 0.0 | |
| assert not failed_trajectory.passed | |
| assert math.isinf(failed_trajectory.score) | |
| assert failed_trajectory.rejection_reasons == ("chunk_1:semantic_gate",) | |
| assert verify_trajectory([]).rejection_reasons == ("empty_trajectory",) | |
| def test_final_whole_output_gate_rejects_post_join_regression_without_fallback(): | |
| passed = verify_trajectory( | |
| [CandidateObservation("整段內容完整", "整段內容完整", 1.0)] | |
| ) | |
| rejected = verify_trajectory( | |
| [CandidateObservation("整段內容完整", "整段內容", 1.0)] | |
| ) | |
| assert require_verified_final_output(passed) is passed | |
| with pytest.raises(FinalOutputRejectedError, match="final output rejected"): | |
| require_verified_final_output(rejected) | |
| with pytest.raises(FinalOutputRejectedError, match="invalid result"): | |
| require_verified_final_output(None) | |
| def test_joined_output_rejection_preserves_local_results_for_sequence_dp(): | |
| observations = [] | |
| artifacts = [] | |
| for target in ("第一段完整", "第二段完整"): | |
| observation, artifact = _chunk_verification( | |
| target, | |
| target, | |
| embedding=[1.0, 0.0], | |
| rms_db=-20.0, | |
| ) | |
| observations.append(observation) | |
| artifacts.append(artifact) | |
| local = verify_trajectory(observations, chunk_artifacts=artifacts) | |
| joined_failed = _joined_verification("第一段完整第二段完整", "第一段完整") | |
| qualified = qualify_trajectory_with_joined_output(local, joined_failed) | |
| assert local.passed | |
| assert not qualified.passed | |
| assert math.isinf(qualified.score) | |
| assert qualified.candidate_results is local.candidate_results | |
| assert qualified.chunk_artifacts is local.chunk_artifacts | |
| assert qualified.rejection_reasons == ("joined_output:chunk_0:semantic_gate",) | |
| assert all(result.passed for result in qualified.candidate_results) | |
| def test_joined_output_qualification_returns_local_result_only_when_joined_passes(): | |
| local = verify_trajectory( | |
| [CandidateObservation("第一段完整", "第一段完整", 1.0)] | |
| ) | |
| joined = verify_trajectory( | |
| [CandidateObservation("第一段完整", "第一段完整", 1.0)] | |
| ) | |
| assert qualify_trajectory_with_joined_output(local, joined) is local | |
| def test_k_best_sequence_paths_are_cost_ranked_stable_distinct_and_bounded(): | |
| ranked = select_k_candidate_sequences( | |
| [[0.0, 0.2], [0.0, 0.1]], | |
| [[[0.4, 0.0], [0.0, 0.5]]], | |
| max_paths=3, | |
| ) | |
| assert [selection.candidate_indices for selection in ranked] == [ | |
| (0, 1), | |
| (1, 0), | |
| (0, 0), | |
| ] | |
| assert [selection.total_score for selection in ranked] == pytest.approx( | |
| [0.1, 0.2, 0.4] | |
| ) | |
| assert len({selection.candidate_indices for selection in ranked}) == 3 | |
| tied = select_k_candidate_sequences( | |
| [[0.0, 0.0], [0.0, 0.0]], | |
| [[[0.0, 0.0], [0.0, 0.0]]], | |
| max_paths=3, | |
| ) | |
| assert [selection.candidate_indices for selection in tied] == [ | |
| (0, 0), | |
| (0, 1), | |
| (1, 0), | |
| ] | |
| def test_k_best_sequence_selection_fails_closed_for_malformed_or_single_chunk_graphs(): | |
| assert select_k_candidate_sequences([[0.0, 0.1]], (), max_paths=3) == () | |
| assert select_k_candidate_sequences([[0.0], [0.0]], (), max_paths=3) == () | |
| assert ( | |
| select_k_candidate_sequences( | |
| [[0.0, 0.1], [0.0, 0.1]], | |
| [[[0.0]]], | |
| max_paths=3, | |
| ) | |
| == () | |
| ) | |
| for invalid in (0, 4, 1.0, True): | |
| with pytest.raises(ValueError, match="between 1 and 3"): | |
| select_k_candidate_sequences( | |
| [[0.0], [0.0]], | |
| [[[0.0]]], | |
| max_paths=invalid, | |
| ) | |
| def _locally_safe_join_rejected_verification(chunks, seed): | |
| # Crossed embeddings make the two mixed paths cheaper than either | |
| # same-candidate path, giving deterministic k-best callback order. | |
| first_candidate = seed % 2 == 0 | |
| similarities = (0.9, 0.2) if first_candidate else (0.2, 0.9) | |
| embeddings = ( | |
| ([1.0, 0.0], [0.0, 1.0]) | |
| if first_candidate | |
| else ([0.0, 1.0], [1.0, 0.0]) | |
| ) | |
| observations = [] | |
| artifacts = [] | |
| for target, similarity, embedding in zip( | |
| chunks, | |
| similarities, | |
| embeddings, | |
| strict=True, | |
| ): | |
| observations.append( | |
| CandidateObservation( | |
| target_text=target, | |
| transcript_text=target, | |
| audio_duration_seconds=2.0, | |
| speaker_similarity=similarity, | |
| begin_speaker_similarity=similarity, | |
| end_speaker_similarity=similarity, | |
| ) | |
| ) | |
| artifacts.append( | |
| ChunkCandidateArtifact( | |
| speaker_embedding=np.asarray(embedding, dtype=np.float32), | |
| rms_db=-20.0, | |
| ) | |
| ) | |
| local = verify_trajectory(observations, chunk_artifacts=artifacts) | |
| joined_failed = _joined_verification("".join(chunks), chunks[0]) | |
| return qualify_trajectory_with_joined_output(local, joined_failed) | |
| def test_k_best_sequence_final_callback_uses_rank_order_and_first_passing_path(): | |
| chunks = ("第一段完整", "第二段完整") | |
| callback_paths = [] | |
| def generator(candidate_chunks, seed): | |
| return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) | |
| def verifier(trajectory, candidate_chunks, seed): | |
| return _locally_safe_join_rejected_verification(candidate_chunks, seed) | |
| def final_verifier(sequence_result, candidate_chunks): | |
| callback_paths.append(sequence_result.chunk_candidate_indices) | |
| transcript = ( | |
| "".join(candidate_chunks) | |
| if sequence_result.chunk_candidate_indices == (1, 0) | |
| else candidate_chunks[0] | |
| ) | |
| return _joined_verification("".join(candidate_chunks), transcript) | |
| result = run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| sequence_final_verifier=final_verifier, | |
| max_sequence_paths=3, | |
| ) | |
| assert callback_paths == [(0, 1), (1, 0)] | |
| assert result.selection_mode == "sequence_dp" | |
| assert result.chunk_candidate_indices == (1, 0) | |
| assert result.sequence_path_rank == 2 | |
| assert result.sequence_paths_checked == 2 | |
| def test_k_best_sequence_all_paths_rejected_is_no_qualified_candidate(): | |
| chunks = ("第一段完整", "第二段完整") | |
| callback_paths = [] | |
| def generator(candidate_chunks, seed): | |
| return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) | |
| def final_verifier(sequence_result, candidate_chunks): | |
| callback_paths.append(sequence_result.chunk_candidate_indices) | |
| return _joined_verification("".join(candidate_chunks), candidate_chunks[0]) | |
| with pytest.raises(NoQualifiedCandidateError, match="after 2 candidates"): | |
| run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| generator, | |
| lambda trajectory, candidate_chunks, seed: ( | |
| _locally_safe_join_rejected_verification(candidate_chunks, seed) | |
| ), | |
| max_candidates=2, | |
| sequence_final_verifier=final_verifier, | |
| max_sequence_paths=3, | |
| ) | |
| assert callback_paths == [(0, 1), (1, 0), (0, 0)] | |
| assert len(set(callback_paths)) == 3 | |
| def test_k_best_sequence_callback_exception_aborts_fail_closed(): | |
| chunks = ("第一段完整", "第二段完整") | |
| callback_paths = [] | |
| def callback(sequence_result, candidate_chunks): | |
| callback_paths.append(sequence_result.chunk_candidate_indices) | |
| raise ValueError("ASR backend disappeared") | |
| with pytest.raises(RuntimeError, match="refusing unverified audio"): | |
| run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| lambda candidate_chunks, seed: tuple(candidate_chunks), | |
| lambda trajectory, candidate_chunks, seed: ( | |
| _locally_safe_join_rejected_verification(candidate_chunks, seed) | |
| ), | |
| max_candidates=2, | |
| sequence_final_verifier=callback, | |
| ) | |
| assert callback_paths == [(0, 1)] | |
| def test_k_best_sequence_invalid_callback_result_aborts_fail_closed(): | |
| chunks = ("第一段完整", "第二段完整") | |
| with pytest.raises(RuntimeError, match="invalid result"): | |
| run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| lambda candidate_chunks, seed: tuple(candidate_chunks), | |
| lambda trajectory, candidate_chunks, seed: ( | |
| _locally_safe_join_rejected_verification(candidate_chunks, seed) | |
| ), | |
| max_candidates=2, | |
| sequence_final_verifier=lambda *args: None, | |
| ) | |
| def test_k_best_sequence_inconsistent_callback_evidence_never_passes(): | |
| chunks = ("第一段完整", "第二段完整") | |
| failed_candidate = verify_candidate( | |
| CandidateObservation("完整內容", "錯誤內容", 1.0) | |
| ) | |
| with pytest.raises(NoQualifiedCandidateError, match="after 2 candidates"): | |
| run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| lambda candidate_chunks, seed: tuple(candidate_chunks), | |
| lambda trajectory, candidate_chunks, seed: ( | |
| _locally_safe_join_rejected_verification(candidate_chunks, seed) | |
| ), | |
| max_candidates=2, | |
| sequence_final_verifier=lambda *args: TrajectoryGateResult( | |
| passed=True, | |
| candidate_results=(failed_candidate,), | |
| score=0.0, | |
| rejection_reasons=(), | |
| ), | |
| ) | |
| def test_sequence_final_callback_is_not_used_for_single_chunk_or_safe_whole(): | |
| callback_calls = [] | |
| with pytest.raises(NoQualifiedCandidateError): | |
| run_adaptive_cascade( | |
| ("單一完整段落",), | |
| 20, | |
| lambda candidate_chunks, seed: tuple(candidate_chunks), | |
| lambda trajectory, candidate_chunks, seed: TrajectoryGateResult( | |
| passed=False, | |
| candidate_results=( | |
| verify_candidate( | |
| CandidateObservation( | |
| candidate_chunks[0], | |
| candidate_chunks[0], | |
| 1.0, | |
| ) | |
| ), | |
| ), | |
| score=math.inf, | |
| rejection_reasons=("joined_output:semantic_gate",), | |
| chunk_artifacts=(ChunkCandidateArtifact(rms_db=-20.0),), | |
| ), | |
| max_candidates=2, | |
| sequence_final_verifier=lambda *args: callback_calls.append(args), | |
| ) | |
| assert callback_calls == [] | |
| result = run_adaptive_cascade( | |
| ("第一段", "第二段"), | |
| 30, | |
| lambda candidate_chunks, seed: tuple(candidate_chunks), | |
| lambda *args: _gate_result(True, 0.1), | |
| max_candidates=2, | |
| sequence_final_verifier=lambda *args: callback_calls.append(args), | |
| ) | |
| assert result.selection_mode == "whole_trajectory" | |
| assert callback_calls == [] | |
| def test_local_pass_join_fail_expands_and_selects_next_joined_safe_whole(): | |
| calls = [] | |
| def generator(chunks, seed): | |
| calls.append(seed) | |
| return tuple(f"seed-{seed}-{chunk}" for chunk in chunks) | |
| def verifier(trajectory, chunks, seed): | |
| observations = [] | |
| artifacts = [] | |
| for target in chunks: | |
| observation, artifact = _chunk_verification( | |
| target, | |
| target, | |
| embedding=[1.0, 0.0], | |
| rms_db=-20.0, | |
| ) | |
| observations.append(observation) | |
| artifacts.append(artifact) | |
| local = verify_trajectory(observations, chunk_artifacts=artifacts) | |
| joined_transcript = "".join(chunks) if seed == 11 else chunks[0] | |
| joined = _joined_verification("".join(chunks), joined_transcript) | |
| return qualify_trajectory_with_joined_output(local, joined) | |
| result = run_adaptive_cascade( | |
| ("第一段完整", "第二段完整"), | |
| 10, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| assert calls == [10, 11] | |
| assert result.selection_mode == "whole_trajectory" | |
| assert result.candidate_index == 1 | |
| assert result.seed == 11 | |
| def test_all_local_pass_join_fail_can_only_return_as_sequence_dp(): | |
| chunks = ("第一段完整", "第二段完整") | |
| def generator(candidate_chunks, seed): | |
| return tuple(f"seed-{seed}-{chunk}" for chunk in candidate_chunks) | |
| def verifier(trajectory, candidate_chunks, seed): | |
| similarities = (0.9, 0.2) if seed == 20 else (0.2, 0.9) | |
| observations = [] | |
| artifacts = [] | |
| for target, similarity in zip(candidate_chunks, similarities, strict=True): | |
| observations.append( | |
| CandidateObservation( | |
| target_text=target, | |
| transcript_text=target, | |
| audio_duration_seconds=2.0, | |
| speaker_similarity=similarity, | |
| begin_speaker_similarity=similarity, | |
| end_speaker_similarity=similarity, | |
| ) | |
| ) | |
| artifacts.append( | |
| ChunkCandidateArtifact( | |
| speaker_embedding=np.array([1.0, 0.0], dtype=np.float32), | |
| rms_db=-20.0, | |
| ) | |
| ) | |
| local = verify_trajectory(observations, chunk_artifacts=artifacts) | |
| joined = _joined_verification("".join(candidate_chunks), candidate_chunks[0]) | |
| return qualify_trajectory_with_joined_output(local, joined) | |
| result = run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| assert result.selection_mode == "sequence_dp" | |
| assert result.candidate_index is None | |
| assert result.seed is None | |
| assert result.chunk_candidate_indices == (0, 1) | |
| assert result.attempted_seeds == (20, 21) | |
| def test_adaptive_cascade_returns_first_candidate_immediately(): | |
| generation_calls = [] | |
| verification_calls = [] | |
| def generator(chunks, seed): | |
| generation_calls.append((chunks, seed)) | |
| return tuple((chunk, seed) for chunk in chunks) | |
| def verifier(trajectory, chunks, seed): | |
| verification_calls.append((trajectory, chunks, seed)) | |
| assert all(chunk_seed == seed for _, chunk_seed in trajectory) | |
| return _gate_result(True, 0.2) | |
| result = run_adaptive_cascade( | |
| ["第一段", "第二段"], | |
| 12345, | |
| generator, | |
| verifier, | |
| ) | |
| assert isinstance(result, CascadeResult) | |
| assert result.seed == 12345 | |
| assert result.candidate_index == 0 | |
| assert result.attempted_seeds == (12345,) | |
| assert generation_calls == [(("第一段", "第二段"), 12345)] | |
| assert len(verification_calls) == 1 | |
| def test_marginal_stage_one_speaker_pass_expands_to_preferred_candidate(): | |
| calls = [] | |
| def generator(chunks, seed): | |
| calls.append(seed) | |
| return (chunks[0], seed) | |
| def verifier(trajectory, chunks, seed): | |
| if seed == 100: | |
| return _whole_speaker_verification( | |
| similarity=0.20, | |
| boundary_drop=0.04, | |
| ) | |
| if seed == 105: | |
| return _whole_speaker_verification( | |
| similarity=0.35, | |
| boundary_drop=0.02, | |
| ) | |
| return _whole_speaker_verification( | |
| similarity=0.30, | |
| boundary_drop=0.02, | |
| passed=False, | |
| ) | |
| result = run_adaptive_cascade( | |
| ["完整而且穩定的候選內容"], | |
| 100, | |
| generator, | |
| verifier, | |
| max_candidates=10, | |
| ) | |
| assert calls == list(range(100, 110)) | |
| assert result.selection_mode == "whole_trajectory" | |
| assert result.candidate_index == 5 | |
| assert result.seed == 105 | |
| def test_final_stage_can_return_marginal_hard_pass_to_preserve_availability(): | |
| calls = [] | |
| def generator(chunks, seed): | |
| calls.append(seed) | |
| return (chunks[0], seed) | |
| def verifier(trajectory, chunks, seed): | |
| return _whole_speaker_verification( | |
| similarity=0.20, | |
| boundary_drop=0.08, | |
| passed=seed == 200, | |
| ) | |
| result = run_adaptive_cascade( | |
| ["完整而且穩定的候選內容"], | |
| 200, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| assert calls == [200, 201] | |
| assert result.seed == 200 | |
| assert result.candidate_index == 0 | |
| assert result.verification.passed | |
| def test_adaptive_cascade_expands_to_five_and_selects_lowest_verified_score(): | |
| calls = [] | |
| scores = { | |
| 100: None, | |
| 101: 0.4, | |
| 102: None, | |
| 103: 0.1, | |
| 104: 0.1, | |
| } | |
| def generator(chunks, seed): | |
| calls.append((chunks, seed)) | |
| return {"chunks": tuple((chunk, seed) for chunk in chunks), "seed": seed} | |
| def verifier(trajectory, chunks, seed): | |
| assert trajectory["seed"] == seed | |
| assert all(chunk_seed == seed for _, chunk_seed in trajectory["chunks"]) | |
| score = scores[seed] | |
| return _gate_result(score is not None, score or 0.0) | |
| result = run_adaptive_cascade( | |
| ["第一段", "第二段"], | |
| 100, | |
| generator, | |
| verifier, | |
| ) | |
| assert [seed for _, seed in calls] == [100, 101, 102, 103, 104] | |
| assert result.seed == 103 | |
| assert result.candidate_index == 3 | |
| assert result.verification.score == 0.1 | |
| assert result.attempted_seeds == (100, 101, 102, 103, 104) | |
| def test_adaptive_cascade_uses_ten_only_when_stage_five_has_no_verified_candidate(): | |
| calls = [] | |
| def generator(chunks, seed): | |
| calls.append(seed) | |
| return (chunks, seed) | |
| def verifier(trajectory, chunks, seed): | |
| return _gate_result(seed in {107, 109}, score=float(110 - seed)) | |
| result = run_adaptive_cascade( | |
| ["完整內容"], | |
| 100, | |
| generator, | |
| verifier, | |
| max_candidates=10, | |
| ) | |
| assert result.seed == 109 | |
| assert result.candidate_index == 9 | |
| assert result.attempted_seeds == tuple(range(100, 110)) | |
| assert calls == list(range(100, 110)) | |
| def test_sequence_fallback_can_take_a_prefix_from_a_and_suffix_from_b(): | |
| chunks = ("第一段完整", "第二段完整") | |
| trajectories = { | |
| 10: ("audio-a0", "audio-a1"), | |
| 11: ("audio-b0", "audio-b1"), | |
| } | |
| def generator(candidate_chunks, seed): | |
| assert candidate_chunks == chunks | |
| return trajectories[seed] | |
| def verifier(trajectory, candidate_chunks, seed): | |
| transcripts = ( | |
| candidate_chunks | |
| if seed == 999 | |
| else ( | |
| (candidate_chunks[0], "錯誤內容") | |
| if seed == 10 | |
| else ("錯誤內容", candidate_chunks[1]) | |
| ) | |
| ) | |
| observations = [] | |
| artifacts = [] | |
| for target, transcript in zip(candidate_chunks, transcripts, strict=True): | |
| observation, artifact = _chunk_verification( | |
| target, | |
| transcript, | |
| embedding=[1.0, 0.0], | |
| rms_db=-20.0, | |
| ) | |
| observations.append(observation) | |
| artifacts.append(artifact) | |
| return verify_trajectory(observations, chunk_artifacts=artifacts) | |
| result = run_adaptive_cascade( | |
| chunks, | |
| 10, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| assert result.selection_mode == "sequence_dp" | |
| assert result.seed is None | |
| assert result.candidate_index is None | |
| assert result.trajectory == ("audio-a0", "audio-b1") | |
| assert result.chunk_candidate_indices == (0, 1) | |
| assert result.chunk_seeds == (10, 11) | |
| assert result.attempted_seeds == (10, 11) | |
| assert result.verification.passed | |
| assert len(result.verification.candidate_results) == 2 | |
| assert len(result.verification.chunk_artifacts) == 2 | |
| def test_same_seed_whole_trajectory_wins_before_sequence_fallback(): | |
| chunks = ("第一段完整", "第二段完整") | |
| def generator(candidate_chunks, seed): | |
| return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) | |
| def verifier(trajectory, candidate_chunks, seed): | |
| transcripts = ( | |
| (candidate_chunks[0], "錯誤內容") | |
| if seed == 20 | |
| else candidate_chunks | |
| ) | |
| observations = [] | |
| artifacts = [] | |
| for target, transcript in zip(candidate_chunks, transcripts, strict=True): | |
| observation, artifact = _chunk_verification( | |
| target, | |
| transcript, | |
| embedding=[1.0, 0.0], | |
| rms_db=-20.0, | |
| ) | |
| observations.append(observation) | |
| artifacts.append(artifact) | |
| return verify_trajectory(observations, chunk_artifacts=artifacts) | |
| result = run_adaptive_cascade( | |
| chunks, | |
| 20, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| assert result.selection_mode == "whole_trajectory" | |
| assert result.seed == 21 | |
| assert result.candidate_index == 1 | |
| assert result.chunk_candidate_indices == (1, 1) | |
| assert result.chunk_seeds == (21, 21) | |
| assert result.trajectory == ("seed-21-chunk-0", "seed-21-chunk-1") | |
| def test_sequence_fallback_rejects_unsafe_transition_edge(): | |
| chunks = ("第一段完整", "第二段完整") | |
| def generator(candidate_chunks, seed): | |
| return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) | |
| def verifier(trajectory, candidate_chunks, seed): | |
| transcripts = ( | |
| (candidate_chunks[0], "錯誤內容") | |
| if seed == 30 | |
| else ("錯誤內容", candidate_chunks[1]) | |
| ) | |
| embeddings = ( | |
| ([1.0, 0.0], [1.0, 0.0]) | |
| if seed == 30 | |
| else ([1.0, 0.0], [1.0, 0.0, 0.0]) | |
| ) | |
| observations = [] | |
| artifacts = [] | |
| for target, transcript, embedding in zip( | |
| candidate_chunks, | |
| transcripts, | |
| embeddings, | |
| strict=True, | |
| ): | |
| observation, artifact = _chunk_verification( | |
| target, | |
| transcript, | |
| embedding=embedding, | |
| rms_db=-20.0, | |
| ) | |
| observations.append(observation) | |
| artifacts.append(artifact) | |
| return verify_trajectory(observations, chunk_artifacts=artifacts) | |
| with pytest.raises(NoQualifiedCandidateError, match="after 2 candidates"): | |
| run_adaptive_cascade( | |
| chunks, | |
| 30, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| def test_transition_score_requires_safe_evidence_and_penalizes_rms_delta(): | |
| first_observation, first_artifact = _chunk_verification( | |
| "第一段完整", | |
| "第一段完整", | |
| embedding=[1.0, 0.0], | |
| rms_db=-20.0, | |
| ) | |
| second_observation, second_artifact = _chunk_verification( | |
| "第二段完整", | |
| "第二段完整", | |
| embedding=[1.0, 0.0], | |
| rms_db=-16.0, | |
| ) | |
| first = verify_trajectory([first_observation], chunk_artifacts=[first_artifact]) | |
| second = verify_trajectory([second_observation], chunk_artifacts=[second_artifact]) | |
| score = candidate_chunk_transition_score( | |
| first.candidate_results[0], | |
| first.chunk_artifacts[0], | |
| second.candidate_results[0], | |
| second.chunk_artifacts[0], | |
| ) | |
| assert score == pytest.approx(0.20) | |
| unsafe = ChunkCandidateArtifact(speaker_embedding=None, rms_db=-16.0) | |
| assert math.isinf( | |
| candidate_chunk_transition_score( | |
| first.candidate_results[0], | |
| first.chunk_artifacts[0], | |
| second.candidate_results[0], | |
| unsafe, | |
| ) | |
| ) | |
| def test_candidate_limit_obeys_twenty_generated_chunk_budget(chunk_count, expected_limit): | |
| limit = candidate_limit_for_chunk_budget(chunk_count) | |
| assert limit == expected_limit | |
| assert limit * chunk_count <= 20 | |
| def test_candidate_limit_rejects_a_trajectory_that_already_exceeds_budget(): | |
| with pytest.raises(ValueError, match="one trajectory exceeds"): | |
| candidate_limit_for_chunk_budget(21) | |
| with pytest.raises(ValueError, match="between 1 and 10"): | |
| candidate_limit_for_chunk_budget(1, max_candidates=11) | |
| def test_360_character_profiles_stay_inside_generated_chunk_budget(): | |
| profiles = [ | |
| "甲" * 360, | |
| ("甲" * 11 + "。") * 30, | |
| ("甲" * 40 + "。") * 8 + "乙" * 32, | |
| ] | |
| for text in profiles: | |
| chunks = split_text_for_tts(text, max_chars=80, min_chunk_chars=12) | |
| chunks = split_leading_clause( | |
| chunks[0], | |
| search_chars=40, | |
| min_chunk_chars=12, | |
| ) + chunks[1:] | |
| limit = candidate_limit_for_chunk_budget(len(chunks)) | |
| assert limit * len(chunks) <= 20 | |
| def test_candidate_asr_value_error_rejects_seed_and_cascade_continues(): | |
| calls = [] | |
| def generator(chunks, seed): | |
| calls.append(seed) | |
| return _tone(seconds=0.25) | |
| def verifier(trajectory, chunks, seed): | |
| prepared = prepare_candidate_audio( | |
| trajectory, | |
| 16_000, | |
| transcriber=( | |
| (lambda *_: (_ for _ in ()).throw(ValueError("candidate ASR failure"))) | |
| if seed == 50 | |
| else (lambda *_: chunks[0]) | |
| ), | |
| ) | |
| if prepared is None: | |
| return _gate_result(False) | |
| return verify_trajectory( | |
| [ | |
| CandidateObservation( | |
| target_text=chunks[0], | |
| transcript_text=prepared.transcript_text, | |
| audio_duration_seconds=prepared.duration_seconds, | |
| ) | |
| ] | |
| ) | |
| result = run_adaptive_cascade( | |
| ["完整內容"], | |
| 50, | |
| generator, | |
| verifier, | |
| max_candidates=2, | |
| ) | |
| assert result.seed == 51 | |
| assert result.candidate_index == 1 | |
| assert calls == [50, 51] | |
| def test_adaptive_cascade_raises_without_qualified_candidate_and_never_falls_back(): | |
| generated = [] | |
| def generator(chunks, seed): | |
| generated.append(seed) | |
| return (chunks, seed) | |
| with pytest.raises(NoQualifiedCandidateError, match="after 5 candidates"): | |
| run_adaptive_cascade( | |
| ["必須完整"], | |
| 7, | |
| generator, | |
| lambda trajectory, chunks, seed: _gate_result(False), | |
| ) | |
| assert generated == [7, 8, 9, 10, 11] | |
| def test_adaptive_cascade_validates_contract_and_rejects_non_gate_verifier(): | |
| generator = lambda chunks, seed: (chunks, seed) | |
| with pytest.raises(ValueError, match="start with exactly one"): | |
| run_adaptive_cascade( | |
| ["內容"], | |
| 1, | |
| generator, | |
| lambda *args: _gate_result(True), | |
| initial_candidates=2, | |
| ) | |
| with pytest.raises(ValueError, match="between 1 and 10"): | |
| run_adaptive_cascade( | |
| ["內容"], | |
| 1, | |
| generator, | |
| lambda *args: _gate_result(True), | |
| max_candidates=11, | |
| ) | |
| with pytest.raises(TypeError, match="TrajectoryGateResult"): | |
| run_adaptive_cascade(["內容"], 1, generator, lambda *args: True) | |