| |
| """Host-chain unit tests for the Inflect AX TTS Python SDK. |
| |
| Runs on a plain host without NPU: onnxruntime is used as a numeric stand-in |
| for the AXMODEL runtime (the SDK pipeline above the session layer is identical |
| for both backends; see README.md for what is NOT covered by this). |
| |
| Runnable two ways: |
| python tests/test_host_chain.py # plain script, exit code 0 on pass |
| pytest tests/ # if pytest is installed |
| """ |
|
|
| from __future__ import annotations |
|
|
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
|
|
| SDK_ROOT = Path(__file__).resolve().parents[1] |
| PKG_DIR = Path(__file__).resolve().parents[2] |
| sys.path.insert(0, str(SDK_ROOT)) |
|
|
| from inflect_ax_tts import DUMMY_PHONEME_IDS, InflectTTS |
| from inflect_ax_tts import host_chain |
| from inflect_ax_tts.frontend import intersperse |
|
|
| CONVERT_DIR = PKG_DIR / "model_convert" |
| REF_NPZ = CONVERT_DIR / "export" / "reference_io.npz" |
| ENCODER_ONNX = CONVERT_DIR / "export" / "encoder.onnx" |
| DECODER_ONNX = CONVERT_DIR / "export" / "decoder.onnx" |
|
|
|
|
| def _cosine(a: np.ndarray, b: np.ndarray) -> float: |
| a = np.asarray(a, dtype=np.float64).ravel() |
| b = np.asarray(b, dtype=np.float64).ravel() |
| assert a.shape == b.shape, f"shape {a.shape} vs {b.shape}" |
| return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b) + 1e-30)) |
|
|
|
|
| def _torch_reference_expand(logw, m_p, logs_p, x_len, length_scale): |
| """Replicates export/export_onnx.py:host_expand with torch (ground truth).""" |
| import torch |
|
|
| def sequence_mask(lengths, max_length): |
| return torch.arange(max_length).unsqueeze(0) < lengths.unsqueeze(1) |
|
|
| x_lengths = torch.LongTensor([x_len]) |
| x_mask = sequence_mask(x_lengths, logw.shape[-1]).unsqueeze(1).float() |
| w = torch.exp(torch.from_numpy(logw)) * x_mask * length_scale |
| w_ceil = torch.ceil(w) |
| y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long() |
| y_len = int(y_lengths.item()) |
| y_mask = sequence_mask(y_lengths, y_len).unsqueeze(1).float() |
| attn_mask = x_mask.unsqueeze(2) * y_mask.unsqueeze(-1) |
| |
| cum = torch.cumsum(w_ceil, -1).view(-1) |
| path = sequence_mask(cum, y_len).float().view(1, 1, -1, y_len) |
| path = path - torch.nn.functional.pad(path, (0, 0, 1, 0))[:, :, :-1] |
| attn = path.transpose(2, 3) * attn_mask |
| m_p_e = torch.matmul(attn.squeeze(1), torch.from_numpy(m_p).transpose(1, 2)) |
| logs_p_e = torch.matmul(attn.squeeze(1), torch.from_numpy(logs_p).transpose(1, 2)) |
| return m_p_e[0].numpy(), logs_p_e[0].numpy(), y_len |
|
|
|
|
| |
| |
| |
|
|
| def test_expand_priors_matches_torch_reference(): |
| """numpy host chain (durations + generate_path + matmul) vs torch.""" |
| ref = np.load(REF_NPZ) |
| logw, m_p, logs_p = ref["logw_torch"], ref["m_p_torch"], ref["logs_p_torch"] |
| x_len = int(ref["x_lengths"][0]) |
| for length_scale in (1.0, 0.8, 1.6): |
| exp_m, exp_l, y_len = host_chain.expand_priors( |
| logw, m_p, logs_p, x_len, length_scale |
| ) |
| ref_m, ref_l, ref_y = _torch_reference_expand( |
| logw, m_p, logs_p, x_len, length_scale |
| ) |
| assert y_len == ref_y, f"T' {y_len} vs torch {ref_y}" |
| assert np.allclose(exp_m, ref_m, atol=1e-5), ( |
| f"m_p_e max diff {np.abs(exp_m - ref_m).max():.2e} @ls={length_scale}" |
| ) |
| assert np.allclose(exp_l, ref_l, atol=1e-5), ( |
| f"logs_p_e max diff {np.abs(exp_l - ref_l).max():.2e} @ls={length_scale}" |
| ) |
| print("[PASS] expand_priors == torch host_expand (3 length scales)") |
|
|
|
|
| def test_encoder_padding_semantics_ort(): |
| """Padded static-shape encoder call == unpadded reference (first x_len frames).""" |
| import onnxruntime as ort |
|
|
| ref = np.load(REF_NPZ) |
| tokens = np.zeros((1, 256), dtype=np.int64) |
| x_len = int(ref["x_lengths"][0]) |
| tokens[0, :x_len] = ref["tokens"][0] |
| sess = ort.InferenceSession(str(ENCODER_ONNX), providers=["CPUExecutionProvider"]) |
| m_p, logs_p, logw = sess.run( |
| None, {"tokens": tokens, "x_lengths": np.asarray([x_len], dtype=np.int64)} |
| ) |
| for name, got, want in ( |
| ("m_p", m_p, ref["m_p_torch"]), |
| ("logs_p", logs_p, ref["logs_p_torch"]), |
| ("logw", logw, ref["logw_torch"]), |
| ): |
| cos = _cosine(got[..., :x_len], want[..., :x_len]) |
| assert cos >= 0.999999, f"{name} cos {cos}" |
| print("[PASS] encoder T=256 padding: first-x_len frames match torch ref") |
|
|
|
|
| def test_decoder_chunk_compose_ort(): |
| """SDK chunked decode (Tp=512 path) vs EXPORT reference chain wav.""" |
| import onnxruntime as ort |
|
|
| ref = np.load(REF_NPZ) |
| z_p = ref["z_p"][0] |
| sess = ort.InferenceSession(str(DECODER_ONNX), providers=["CPUExecutionProvider"]) |
|
|
| def run_chunk(z_chunk: np.ndarray) -> np.ndarray: |
| return sess.run(None, {"z_p": z_chunk[None].astype(np.float32)})[0].reshape(-1) |
|
|
| wav = host_chain.decode_waveform(z_p, run_chunk) |
| want = ref["wav_onnx_chain"].reshape(-1) |
| assert wav.shape == want.shape, f"{wav.shape} vs {want.shape}" |
| |
| |
| |
| t_prime = z_p.shape[1] |
| crop = min(64, t_prime // 4) * host_chain.HOP_LENGTH |
| cos = _cosine(wav[:-crop], want[:-crop]) |
| assert cos >= 0.9999, f"decoder compose cos {cos}" |
| head = _cosine(wav[: (t_prime // 2) * host_chain.HOP_LENGTH], |
| want[: (t_prime // 2) * host_chain.HOP_LENGTH]) |
| assert head >= 0.999999, f"decoder compose head cos {head}" |
| print(f"[PASS] decoder chunk compose vs EXPORT chain: cos={cos:.9f} (right-edge crop), head cos={head:.9f}") |
|
|
|
|
| def test_decoder_multi_chunk_compose_ort(): |
| """T'>512 multi-chunk path: overlap crossfade vs single dynamic ORT call. |
| |
| Uses synthetic z_p (not a real prior): checks stitching correctness only, |
| comparing against the dynamic-shape ONNX decoder run over the full length. |
| Boundary regions are excluded from the metric (conv receptive-field bleed |
| is expected; EXPORT_NOTES §5.3). |
| """ |
| import onnxruntime as ort |
|
|
| rng = np.random.default_rng(1234) |
| t_prime = 512 + 300 |
| z_p = rng.standard_normal((192, t_prime), dtype=np.float32) * 0.5 |
| sess = ort.InferenceSession(str(DECODER_ONNX), providers=["CPUExecutionProvider"]) |
| want = sess.run(None, {"z_p": z_p[None]})[0].reshape(-1) |
| got = host_chain.decode_waveform( |
| z_p, lambda z: sess.run(None, {"z_p": z[None]})[0].reshape(-1) |
| ) |
| assert got.shape == want.shape |
| |
| hop = host_chain.HOP_LENGTH |
| lo, hi = (512 + 32) * hop, (t_prime - 32) * hop |
| interior = _cosine(got[lo:hi], want[lo:hi]) |
| head = _cosine(got[: 480 * hop], want[: 480 * hop]) |
| assert interior >= 0.999, f"multi-chunk interior cos {interior}" |
| assert head >= 0.999, f"multi-chunk head cos {head}" |
| print(f"[PASS] multi-chunk compose: head cos={head:.6f}, interior cos={interior:.6f}") |
|
|
|
|
| def test_synthesize_tokens_end_to_end_ort(): |
| """Full SDK pipeline on ORT stand-in: length, determinism, finiteness.""" |
| tts = InflectTTS(ENCODER_ONNX, DECODER_ONNX, backend="onnxruntime") |
| sr, wav1 = tts.synthesize_tokens(DUMMY_PHONEME_IDS, seed=0, variation=0.667) |
| _, wav2 = tts.synthesize_tokens(DUMMY_PHONEME_IDS, seed=0, variation=0.667) |
| _, wav3 = tts.synthesize_tokens(DUMMY_PHONEME_IDS, seed=1, variation=0.667) |
| ref = np.load(REF_NPZ) |
| assert sr == 24000 |
| |
| assert wav1.size == int(ref["wav_onnx_chain"].size), ( |
| f"{wav1.size} vs baseline {ref['wav_onnx_chain'].size}" |
| ) |
| assert np.array_equal(wav1, wav2), "same seed must be bit-reproducible" |
| assert not np.array_equal(wav1, wav3), "different seed must differ" |
| assert np.isfinite(wav1).all() |
| assert np.abs(wav1).max() <= 1.0 |
| print(f"[PASS] e2e synthesize_tokens: {wav1.size} samples, deterministic, bounded") |
|
|
|
|
| def test_frontend_symbol_mapping(): |
| """Symbol table mapping: literal IPA baseline string -> EXPORT baseline ids. |
| |
| export/model_meta.json numeric_baseline was built with |
| cleaned_text_to_sequence("ðə kwɪk") — a literal IPA string, NOT eSpeak |
| output (eSpeak with_stress would add stress marks).""" |
| from inflect_ax_tts.symbols import symbols |
|
|
| symbol_to_id = {s: i for i, s in enumerate(symbols)} |
| ids = [symbol_to_id[ch] for ch in "ðə kwɪk"] |
| assert ids == DUMMY_PHONEME_IDS, f"{ids} vs {DUMMY_PHONEME_IDS}" |
| assert intersperse(ids, 0) == [0, 81, 0, 83, 0, 16, 0, 53, 0, 65, 0, 102, 0, 53, 0] |
| print("[PASS] symbol mapping: 'ðə kwɪk' -> baseline phoneme ids") |
|
|
|
|
| def test_frontend_smoke(): |
| """eSpeak frontend structural check (skipped when espeak is unavailable).""" |
| try: |
| from inflect_ax_tts.frontend import text_to_token_ids |
| ids = text_to_token_ids("the quick") |
| except Exception as exc: |
| print(f"[SKIP] frontend (espeak unavailable: {exc})") |
| return |
| assert len(ids) % 2 == 1, "interspersed sequence must be odd length" |
| assert ids[::2] == [0] * (len(ids) // 2 + 1), "blanks must interleave" |
| from inflect_ax_tts.symbols import symbols |
|
|
| assert all(0 <= i < len(symbols) for i in ids) |
| print(f"[PASS] frontend smoke: 'the quick' -> {len(ids)} interspersed ids") |
|
|
|
|
| ALL_TESTS = [ |
| test_expand_priors_matches_torch_reference, |
| test_encoder_padding_semantics_ort, |
| test_decoder_chunk_compose_ort, |
| test_decoder_multi_chunk_compose_ort, |
| test_synthesize_tokens_end_to_end_ort, |
| test_frontend_symbol_mapping, |
| test_frontend_smoke, |
| ] |
|
|
|
|
| def main() -> int: |
| failures = 0 |
| for test in ALL_TESTS: |
| try: |
| test() |
| except Exception as exc: |
| failures += 1 |
| print(f"[FAIL] {test.__name__}: {exc}") |
| print(f"{'ALL PASS' if not failures else f'{failures} FAILURES'} ({len(ALL_TESTS)} tests)") |
| return 1 if failures else 0 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(main()) |
|
|