# --------------------------------------------------------------------- # Copyright (c) 2026 Qualcomm Technologies, Inc. and/or its subsidiaries. # SPDX-License-Identifier: BSD-3-Clause # --------------------------------------------------------------------- from __future__ import annotations from pathlib import Path import numpy as np import pytest import soundfile as sf import torch from huggingface_hub.errors import GatedRepoError from .demo import main as demo_main from .model import ( CONTEXT_LENGTH, DECODE_SEQ_LEN, HEAD_DIM, NUM_KEY_VALUE_HEADS, NUM_LAYERS, PREFILL_SEQ_LEN, SAMPLE_RATE, NeuTTSNano, build_attention_mask, empty_kv_cache, ) # The NeuTTS-Nano backbone is a gated HF repo. Skip rather than fail on machines # whose token has not accepted the NeuTTS Open License. _GATED_SKIP = "NeuTTS-Nano weights are HF-gated for this token" def _load_or_skip() -> NeuTTSNano: try: return NeuTTSNano.from_pretrained() except GatedRepoError: pytest.skip(_GATED_SKIP) def _run(model: NeuTTSNano, seq_len: int) -> list[torch.Tensor]: ids = torch.zeros((1, seq_len), dtype=torch.int32) position_ids = torch.arange(seq_len).reshape(1, seq_len) cos, sin = model.backbone.embedding.get_embedding(position_ids) caches = empty_kv_cache(CONTEXT_LENGTH - seq_len) with torch.no_grad(): return model.backbone( ids, build_attention_mask(seq_len, seq_len), cos, sin, *caches ) @pytest.mark.slow @pytest.mark.parametrize("seq_len", [PREFILL_SEQ_LEN, DECODE_SEQ_LEN]) def test_backbone_graph_forward(seq_len: int) -> None: model = _load_or_skip() out = _run(model, seq_len) vocab_size = model.backbone.llm_config.vocab_size assert out[0].shape == (1, seq_len, vocab_size) assert torch.isfinite(out[0]).all() # logits plus one key and one value per layer assert len(out) == 1 + 2 * NUM_LAYERS for layer in range(NUM_LAYERS): key, value = out[1 + 2 * layer], out[2 + 2 * layer] assert key.shape == (NUM_KEY_VALUE_HEADS, 1, HEAD_DIM, seq_len) assert value.shape == (NUM_KEY_VALUE_HEADS, 1, seq_len, HEAD_DIM) assert torch.isfinite(key).all() and torch.isfinite(value).all() @pytest.mark.slow def test_graphs_share_one_source() -> None: """Both graphs must describe the same weights at two sequence lengths.""" model = _load_or_skip() assert model.graph_names == [model.prefill_graph, model.decode_graph] assert model.shared_source_model prefill = model.get_graph_input_spec(model.prefill_graph) decode = model.get_graph_input_spec(model.decode_graph) # Same inputs, same cache contract; only the sequence dimension differs. assert set(prefill) == set(decode) assert prefill["input_ids"][0] == (1, PREFILL_SEQ_LEN) assert decode["input_ids"][0] == (1, DECODE_SEQ_LEN) assert model.get_graph_output_spec(model.prefill_graph) == ( model.get_graph_output_spec(model.decode_graph) ) @pytest.mark.slow def test_shared_source_serializes_once(tmp_path: Path) -> None: """One .pt must serve both graphs, and be traced at the longest shape.""" model = _load_or_skip() path = model.serialize_graph(model.decode_graph, tmp_path) assert path.is_file() traced = torch.jit.load(str(path)) inputs = [ torch.from_numpy(v[0]) for v in model.get_graph_sample_inputs(model.prefill_graph).values() ] with torch.no_grad(): out = traced(*inputs) assert out[0].shape[1] == PREFILL_SEQ_LEN @pytest.mark.slow def test_decode_matches_prefill() -> None: """Cached decode of the last token must match prefilling the whole window. Predicting what follows tokens ``0..n-1`` is done two ways: one prefill over all ``n``, versus a prefill over ``0..n-2`` (left-padded) followed by a single-token decode. Agreement exercises the cache layout, the sliding window and the mask together. """ model = _load_or_skip() backbone = model.backbone rng = np.random.default_rng(seed=0) ids = torch.from_numpy( rng.integers(low=0, high=1000, size=(1, PREFILL_SEQ_LEN)).astype(np.int32) ) split = PREFILL_SEQ_LEN - 1 cos, sin = backbone.embedding.get_embedding( torch.arange(PREFILL_SEQ_LEN).reshape(1, -1) ) with torch.no_grad(): whole = backbone( ids, build_attention_mask(PREFILL_SEQ_LEN, PREFILL_SEQ_LEN), cos, sin, *empty_kv_cache(CONTEXT_LENGTH - PREFILL_SEQ_LEN), ) # Left-pad by one so the first ``split`` real tokens stay right-aligned. padded = torch.cat([torch.zeros((1, 1), dtype=torch.int32), ids[:, :split]], dim=1) cos, sin = backbone.embedding.get_embedding(torch.tensor([[0, *range(split)]])) with torch.no_grad(): partial = backbone( padded, build_attention_mask(PREFILL_SEQ_LEN, split), cos, sin, *empty_kv_cache(CONTEXT_LENGTH - PREFILL_SEQ_LEN), ) # Drop the pad position's cache entry, then right-align in the decode buffer. caches = empty_kv_cache(CONTEXT_LENGTH - 1) for layer in range(NUM_LAYERS): caches[2 * layer][..., -split:] = partial[1 + 2 * layer][..., 1:] caches[2 * layer + 1][:, :, -split:, :] = partial[2 + 2 * layer][:, :, 1:, :] cos, sin = backbone.embedding.get_embedding(torch.tensor([[split]])) with torch.no_grad(): step = backbone( ids[:, split:], build_attention_mask(1, PREFILL_SEQ_LEN), cos, sin, *caches, ) torch.testing.assert_close(step[0][:, 0], whole[0][:, split], atol=1e-3, rtol=1e-3) def _spectral_centroid_std(wav: np.ndarray, sample_rate: int) -> float: """Standard deviation of the spectral centroid over time, in Hz. Separates speech from a stationary drone. librosa would do this in one call but imports numba, so compute it with numpy. """ frame, hop = 1024, 256 n = 1 + max(0, (len(wav) - frame) // hop) idx = np.arange(frame)[None, :] + hop * np.arange(n)[:, None] mag = np.abs(np.fft.rfft(wav[idx] * np.hanning(frame), axis=1)) freqs = np.fft.rfftfreq(frame, 1.0 / sample_rate) centroid = (mag * freqs).sum(axis=1) / (mag.sum(axis=1) + 1e-10) return float(centroid.std()) @pytest.mark.slow def test_demo(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: # The demo fixes its own sampling seed, so this is deterministic. Seeding from # here would not work: upstream's infer() calls torch.manual_seed itself. monkeypatch.chdir(tmp_path) try: demo_main(is_test=True) except GatedRepoError: pytest.skip(_GATED_SKIP) output = tmp_path / "neutts_nano_output.wav" assert output.is_file() wav, sample_rate = sf.read(output, dtype="float32", always_2d=True) wav = wav.mean(axis=1) assert sample_rate == SAMPLE_RATE assert np.isfinite(wav).all() # Guard against a silent / degenerate waveform passing as success. assert len(wav) > SAMPLE_RATE assert np.abs(wav).max() > 1e-2 # A broken cache or mask collapses generation into a short repeating token # loop, which decodes to a stationary drone: real speech measures >1300 Hz of # centroid spread, a greedy-collapsed loop 47 Hz. assert _spectral_centroid_std(wav, sample_rate) > 300.0