import ast from dataclasses import replace from pathlib import Path from types import SimpleNamespace import numpy as np import pytest import torch from quality_runtime import ( CandidateObservation, ChunkCandidateArtifact, local_candidate_has_coverage_eligibility, verify_trajectory, ) ROOT = Path(__file__).resolve().parents[1] def _isolated_app_function(name, namespace): app_path = ROOT / "app.py" tree = ast.parse(app_path.read_text(encoding="utf-8")) function = next( node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == name ) module = ast.Module( body=[ ast.ImportFrom( module="__future__", names=[ast.alias(name="annotations")], level=0, ), function, ], type_ignores=[], ) ast.fix_missing_locations(module) exec(compile(module, app_path, "exec"), namespace) return namespace[name] def test_closed_loop_pace_rerenders_once_from_untouched_model_waveform(): apply_calls = [] active_calls = [] correction_calls = [] def apply_speed(audio, speed, *, network_conditioned=False): source = np.asarray(audio, dtype=np.float32).copy() apply_calls.append((source, speed, network_conditioned)) return source * np.float32(speed) def active_duration(audio, sample_rate): active_calls.append((np.asarray(audio).copy(), sample_rate)) return 2.0 def active_correction(active_seconds, text, **kwargs): correction_calls.append((active_seconds, text, kwargs)) return 0.95 namespace = { "np": np, "torch": torch, "model": SimpleNamespace(generate=lambda **kwargs: torch.ones(16)), "select_generation_cps": lambda *args, **kwargs: 5.0, "count_network_endpoint_duration_units": lambda text: 8, "count_speech_units": lambda text: 8, "endpoint_generation_plan": lambda *args, **kwargs: ("測試內容。", 10, 12), "effective_generation_cfg": lambda text, cfg, **kwargs: cfg, "set_generation_seed": lambda seed: None, "target_pace_speed": lambda *args, **kwargs: 0.90, "active_voiced_duration_seconds": active_duration, "active_pace_correction_speed": active_correction, "_apply_speed": apply_speed, "_NATIVE_STOP_POLICY": False, "_GENERATE_PARAMETERS": set(), "_STOP_CONTROLLER": None, "SR": 48_000, "STEP_SECONDS": 0.16, "MIN_ENDPOINT_CUE_UNITS": 6, "SHORT_TEXT_CFG_UNITS": 6, "SHORT_TEXT_CFG_MIN": 3.0, "TARGET_CPS": 4.0, "ACTIVE_PACE_TARGET_CPS": 4.0, "CLOSED_LOOP_ACTIVE_PACE_TARGET_CPS": 3.95, "MIN_PACE_SPEED": 0.80, "STOP_THRESHOLD": 0.50, "STOP_CONSECUTIVE": 1, } generate = _isolated_app_function("_generate_chunk", namespace) policy = SimpleNamespace( name="test", cjk_cps=5.0, ascii_cps=4.6, hard_stop_margin_steps=1, ) output = generate( "測試內容。", torch.ones(4), cfg=3.0, steps=10, request_seed=123, policy=policy, network_conditioned=True, ) assert len(apply_calls) == 2 assert np.array_equal(apply_calls[0][0], np.ones(16, dtype=np.float32)) assert np.array_equal(apply_calls[1][0], np.ones(16, dtype=np.float32)) assert apply_calls[0][1] == pytest.approx(0.90) assert apply_calls[1][1] == pytest.approx(0.90 * 0.95) assert all(call[2] is True for call in apply_calls) assert np.array_equal(output, np.ones(16, dtype=np.float32) * np.float32(0.855)) assert len(active_calls) == 2 assert np.array_equal(active_calls[0][0], np.ones(16, dtype=np.float32)) assert np.array_equal(active_calls[1][0], np.ones(16, dtype=np.float32) * 0.90) assert correction_calls[1][2] == { "target_cps": 3.95, "prior_speed": 0.90, "min_total_speed": 0.80, } def test_lazy_transition_f0_measures_only_dp_eligible_rows(): observations = ( CandidateObservation("第一段完整", "第一段完整", 2.0), CandidateObservation("第二段完整", "錯誤內容", 2.0), ) artifacts = ( ChunkCandidateArtifact(rms_db=-20.0), ChunkCandidateArtifact(rms_db=-21.0), ) verification = verify_trajectory( observations, chunk_artifacts=artifacts, speaker_gate_enabled=False, ) waveforms = ( np.ones(100, dtype=np.float32), np.ones(120, dtype=np.float32), ) measured = [] namespace = { "np": np, "replace": replace, "ChunkCandidateArtifact": ChunkCandidateArtifact, "local_candidate_has_coverage_eligibility": ( local_candidate_has_coverage_eligibility ), "SEQUENCE_FALLBACK_MAX_LOCAL_BOUNDARY_SPEAKER_DROP": 0.15, "active_audio_median_f0_hz": ( lambda waveform, sample_rate: measured.append( (waveform.copy(), sample_rate) ) or 210.0 ), "SR": 48_000, } attach = _isolated_app_function("_attach_transition_f0", namespace) disabled = attach( verification, waveforms, collect_transition_f0=False, local_candidate_pool=True, ) updated = attach( verification, waveforms, collect_transition_f0=True, local_candidate_pool=True, ) assert disabled is verification assert len(measured) == 1 assert np.array_equal(measured[0][0], waveforms[0]) assert measured[0][1] == 48_000 assert updated.candidate_results is verification.candidate_results assert updated.rejection_reasons == verification.rejection_reasons assert updated.score == verification.score assert updated.chunk_artifacts[0].median_f0_hz == 210.0 assert updated.chunk_artifacts[1].median_f0_hz is None with pytest.raises(ValueError, match="only valid for local"): attach( verification, waveforms, collect_transition_f0=True, local_candidate_pool=False, ) def test_space_wires_transition_f0_only_to_multichunk_dp_candidates(): source = (ROOT / "app.py").read_text(encoding="utf-8") tree = ast.parse(source) functions = { node.name: ast.get_source_segment(source, node) for node in tree.body if isinstance(node, ast.FunctionDef) } assert "collect_transition_f0=len(chunks) > 1" in functions[ "_qualify_candidate_trajectory_audio" ] assert "collect_transition_f0=True" in functions[ "_verify_refill_candidate_trajectory_audio" ] for name in ( "_verify_independent_whole_audio", "_verify_sequence_trajectory_audio", "_synthesize", ): assert "collect_transition_f0=" not in functions[name] generate = functions["_generate_chunk"] assert "corrected_active_duration = active_voiced_duration_seconds(" in generate assert "final_speed = combined_speed * rerender_speed" in generate assert "_apply_speed(\n corrected" not in generate assert "_apply_speed(\n audio,\n final_speed" in generate