Spaces:
Running on Zero
Running on Zero
Download tests/test_inference_cleanup.py from voidful/BlueMagpie-TTS-Demo: direct link, hf CLI and curl.
- Browser
- Download file 7.32 kB
-
https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/ca45ff159bb73d2e280bdf2e91a4923ec473348d/tests/test_inference_cleanup.py
- Command line
-
hf download hf://spaces/voidful/BlueMagpie-TTS-Demo@ca45ff159bb73d2e280bdf2e91a4923ec473348d/tests/test_inference_cleanup.py
-
curl -L -o test_inference_cleanup.py https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/ca45ff159bb73d2e280bdf2e91a4923ec473348d/tests/test_inference_cleanup.py
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 | |