BlueMagpie-TTS-Demo / tests /test_quality_runtime.py
voidful's picture
Add fail-closed stable speaker inference
7e7df2a
Raw History Blame
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)
@pytest.mark.parametrize("offset", [-1, 0.5, 1.0, True, None, "one"])
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]
@pytest.mark.parametrize("seed", [-1, 2**31, 1.0, True, "123", None])
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,
)
)
@pytest.mark.parametrize(
("chunk_count", "expected_limit"),
[(1, 10), (2, 10), (3, 6), (4, 5), (5, 4), (6, 3), (10, 2), (20, 1)],
)
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)