Files
VoiceStudio/tests/test_mlx_audio_sample_rate.py

74 lines
2.8 KiB
Python

"""MLX results must match the adapter's declared rate before stitching/export."""
import importlib
from types import SimpleNamespace
import numpy as np
import pytest
@pytest.fixture
def backend(monkeypatch):
cls = importlib.import_module("services.tts_backend").MLXAudioBackend
be = cls.__new__(cls)
be._model_id = "mlx-community/Dia-1.6B"
be._sr = 24000
monkeypatch.setattr(be, "_ensure_loaded", lambda: None)
return be
def _tone(rate):
return np.sin(2 * np.pi * 440 * np.arange(rate) / rate).astype(np.float32)
@pytest.mark.parametrize("rate", [24000, 44100, 48000])
@pytest.mark.parametrize("fallback", [False, True])
def test_generated_audio_keeps_its_duration_and_pitch(backend, rate, fallback):
def generate(**kwargs):
if fallback and "voice" in kwargs:
raise TypeError("unexpected keyword argument 'voice'")
yield SimpleNamespace(audio=_tone(rate), sample_rate=rate)
backend._model = SimpleNamespace(generate=generate)
declared_rate = backend.sample_rate # callers can read this before generate
wav = backend.generate("Hello", voice="speaker")
assert backend.sample_rate == declared_rate
assert wav.shape == (1, declared_rate)
spectrum = np.abs(np.fft.rfft(wav[0].numpy()))
assert np.fft.rfftfreq(wav.shape[-1], 1 / declared_rate)[spectrum.argmax()] == 440
def test_each_piece_is_resampled_before_concatenation(backend):
def generate(**kwargs):
for rate in (44100, 24000):
yield SimpleNamespace(audio=_tone(rate), sample_rate=rate)
backend._model = SimpleNamespace(generate=generate)
assert backend.generate("Two pieces").shape == (1, 48000)
def test_raw_audio_without_rate_metadata_keeps_the_existing_contract(backend):
audio = _tone(24000)
backend._model = SimpleNamespace(generate=lambda **kwargs: iter([audio]))
np.testing.assert_array_equal(backend.generate("Hello")[0].numpy(), audio)
@pytest.mark.parametrize("fallback", [False, True])
@pytest.mark.parametrize("rates", [(44100, 44100), (44100, 44100, 48000, 48000)])
def test_contiguous_chunks_share_resampling_context(backend, fallback, rates):
import torch
import torchaudio
pieces = [_tone(rate)[:101] for rate in rates]
def generate(**kwargs):
if fallback and "voice" in kwargs:
raise TypeError("unsupported voice")
for audio, rate in zip(pieces, rates):
yield SimpleNamespace(audio=audio, sample_rate=rate)
backend._model = SimpleNamespace(generate=generate)
expected = []
for start in range(0, len(rates), 2):
joined = torch.from_numpy(np.concatenate(pieces[start:start + 2]))
expected.append(torchaudio.functional.resample(joined, rates[start], 24000))
np.testing.assert_allclose(backend.generate("chunks", voice="speaker")[0].numpy(),
torch.cat(expected).numpy(), atol=1e-6)