BlueMagpie-TTS-Demo / tests /test_inference_cleanup.py
Codex
Add proof-bound email candidate fallback
f06f5d5
Raw History Blame
7.32 kB
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