| |
| |
| |
| |
| 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, |
| ) |
|
|
| |
| |
| _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() |
|
|
| |
| 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) |
| |
| 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), |
| ) |
|
|
| |
| 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), |
| ) |
|
|
| |
| 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: |
| |
| |
| 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() |
| |
| assert len(wav) > SAMPLE_RATE |
| assert np.abs(wav).max() > 1e-2 |
|
|
| |
| |
| |
| assert _spectral_centroid_std(wav, sample_rate) > 300.0 |
|
|