| import math |
|
|
| import pytest |
| import torch |
| import torch.nn.functional as Fn |
|
|
| import kernels |
|
|
| rp = kernels.get_kernel("phanerozoic/resample-poly", version=1, |
| trust_remote_code=True) |
|
|
| |
| PAIRS = [(44100, 16000), (48000, 16000), (16000, 24000), (22050, 16000), |
| (8000, 16000), (44100, 22050), (48000, 44100), (16000, 16000)] |
|
|
|
|
| def _sig(T, seed=0, B=1): |
| g = torch.Generator().manual_seed(seed) |
| return torch.randn(B, T, generator=g, dtype=torch.float32) |
|
|
|
|
| def _ref(x, plan): |
| """torchaudio's formulation: conv1d over the padded-to-common-width bank.""" |
| L, M, width, kern = plan.L, plan.M, plan.width, plan.kernel |
| B, T = x.shape |
| pad = Fn.pad(x, (width, width + M)) |
| r = Fn.conv1d(pad[:, None], kern[:, None, :], stride=M) |
| r = r.transpose(1, 2).reshape(B, -1) |
| return r[:, :plan.out_len(T)] |
|
|
|
|
| |
| |
| |
|
|
| @pytest.mark.parametrize("orig,new", PAIRS) |
| def test_matches_padded_bank(orig, new): |
| plan = rp.Resampler(orig, new) |
| x = _sig(orig // 2, seed=orig + new) |
| got = plan(x) |
| ref = _ref(x, plan) |
| assert got.shape == ref.shape, (got.shape, ref.shape) |
| scale = ref.abs().max().clamp(min=1e-6) |
| rel = (got - ref).abs().max() / scale |
| assert rel < 2e-5, f"{orig}->{new} max rel {rel.item():.3e}" |
|
|
|
|
| @pytest.mark.parametrize("width", [1, 4, 6, 16]) |
| def test_matches_across_filter_widths(width): |
| plan = rp.Resampler(44100, 16000, lowpass_filter_width=width) |
| x = _sig(22050, seed=width) |
| scale = _ref(x, plan).abs().max().clamp(min=1e-6) |
| assert ((plan(x) - _ref(x, plan)).abs().max() / scale) < 2e-5 |
|
|
|
|
| def test_kaiser_window(): |
| plan = rp.Resampler(48000, 16000, method="sinc_interp_kaiser") |
| x = _sig(24000, seed=3) |
| scale = _ref(x, plan).abs().max().clamp(min=1e-6) |
| assert ((plan(x) - _ref(x, plan)).abs().max() / scale) < 2e-5 |
|
|
|
|
| |
| |
| |
|
|
| def test_support_is_a_small_fraction_of_the_padded_bank(): |
| plan = rp.Resampler(44100, 16000) |
| assert plan.L == 160 and plan.M == 441 |
| assert plan.density < 0.15, plan.density |
| |
| kept = int(plan.len.sum()) |
| assert int((plan.kernel != 0).sum()) == kept |
|
|
|
|
| def test_dropped_entries_are_exactly_zero(): |
| """Clamping past the filter half width sets the window to cos(pi/2)^2, so |
| the trimmed region is exact rather than thresholded.""" |
| for orig, new in [(44100, 16000), (16000, 24000)]: |
| plan = rp.Resampler(orig, new) |
| for i in range(plan.L): |
| row = plan.kernel[i] |
| s = int(plan.start[i]) + plan.width |
| n = int(plan.len[i]) |
| outside = torch.cat([row[:s], row[s + n:]]) |
| assert outside.abs().max() == 0.0 |
|
|
|
|
| |
| |
| |
|
|
| def test_identity_rate_passes_the_band_but_is_not_the_identity(): |
| """At equal rates the bank is a single phase, but rolloff still puts a |
| low-pass at 0.99 of Nyquist. A tone in the band comes through untouched; |
| white noise does not, because its top percent is removed by design.""" |
| sr = 16000 |
| plan = rp.Resampler(sr, sr) |
| assert plan.L == 1 and plan.M == 1 |
|
|
| t = torch.arange(sr, dtype=torch.float32) / sr |
| tone = torch.sin(2 * math.pi * 1000.0 * t).unsqueeze(0) |
| y = plan(tone) |
| assert y.shape == tone.shape |
| e = 64 |
| assert (y[:, e:-e] - tone[:, e:-e]).abs().max() < 5e-3 |
|
|
| noise = _sig(8000, seed=5) |
| assert (plan(noise)[:, e:-e] - noise[:, e:-e]).abs().max() > 1e-3 |
|
|
|
|
| def test_sine_survives_downsampling(): |
| """A tone well below the new Nyquist must come through at its own frequency |
| and amplitude.""" |
| sr, new_sr, f = 48000, 16000, 440.0 |
| n = sr |
| t = torch.arange(n, dtype=torch.float32) / sr |
| x = torch.sin(2 * math.pi * f * t).unsqueeze(0) |
| y = rp.Resampler(sr, new_sr)(x) |
| e = 512 |
| mid = y[0, e:-e] |
| assert abs(mid.abs().max().item() - 1.0) < 0.02 |
| spec = torch.fft.rfft(mid * torch.hann_window(mid.numel())) |
| peak = int(spec.abs().argmax()) |
| got_f = peak * new_sr / mid.numel() |
| assert abs(got_f - f) < 5.0, got_f |
|
|
|
|
| def test_above_nyquist_is_attenuated(): |
| """A tone above the new Nyquist must be filtered out, not aliased down.""" |
| sr, new_sr = 48000, 16000 |
| n = sr |
| t = torch.arange(n, dtype=torch.float32) / sr |
| x = torch.sin(2 * math.pi * 11000.0 * t).unsqueeze(0) |
| y = rp.Resampler(sr, new_sr)(x) |
| e = 512 |
| assert y[0, e:-e].abs().max() < 0.05, y[0, e:-e].abs().max().item() |
|
|
|
|
| def test_dc_preserved(): |
| x = torch.full((1, 8000), 0.25) |
| y = rp.Resampler(44100, 16000)(x) |
| e = 64 |
| assert (y[0, e:-e] - 0.25).abs().max() < 5e-3 |
|
|
|
|
| |
| |
| |
|
|
| @pytest.mark.parametrize("orig,new", [(44100, 16000), (16000, 48000)]) |
| def test_output_length(orig, new): |
| plan = rp.Resampler(orig, new) |
| for T in (1000, 4321, 44100): |
| assert plan(_sig(T, seed=T)).shape[-1] == plan.out_len(T) |
|
|
|
|
| def test_unbatched_drops_batch_dim(): |
| plan = rp.Resampler(44100, 16000) |
| x = _sig(4410, seed=7) |
| flat = plan(x[0]) |
| assert flat.dim() == 1 |
| assert torch.equal(flat, plan(x)[0]) |
|
|
|
|
| def test_batch_invariance(): |
| plan = rp.Resampler(48000, 16000) |
| x = _sig(4800, seed=8, B=5) |
| batched = plan(x) |
| for i in range(5): |
| assert torch.equal(batched[i], plan(x[i:i + 1])[0]), f"row {i}" |
|
|
|
|
| def test_deterministic(): |
| plan = rp.Resampler(44100, 16000) |
| x = _sig(4410, seed=9) |
| assert torch.equal(plan(x), plan(x)) |
|
|
|
|
| def test_thread_count_invariance(): |
| plan = rp.Resampler(44100, 16000) |
| x = _sig(441000, seed=10) |
| prev = torch.get_num_threads() |
| try: |
| torch.set_num_threads(1) |
| a = plan(x) |
| torch.set_num_threads(max(2, prev)) |
| b = plan(x) |
| finally: |
| torch.set_num_threads(prev) |
| assert torch.equal(a, b) |
|
|
|
|
| def test_path_selection_matches_the_measured_crossover(): |
| """Both paths must agree; the plan picks between them on taps per output, |
| since the fused reduction is fixed cost per sample and only pays above |
| roughly 16 taps.""" |
| fused = rp.Resampler(48000, 16000) |
| bank = rp.Resampler(16000, 24000) |
| assert fused.fused and fused.taps_per_output >= 16.0 |
| assert not bank.fused and bank.taps_per_output < 16.0 |
|
|
| |
| x = _sig(16000, seed=31) |
| plan = rp.Resampler(48000, 16000) |
| direct = plan(_sig(48000, seed=32)) |
| plan.fused = False |
| fallback = plan(_sig(48000, seed=32)) |
| plan.fused = True |
| scale = fallback.abs().max().clamp(min=1e-6) |
| assert ((direct - fallback).abs().max() / scale) < 2e-5 |
|
|
|
|
| def test_rejects_bad_arguments(): |
| with pytest.raises(ValueError): |
| rp.Resampler(0, 16000) |
| with pytest.raises(ValueError): |
| rp.Resampler(44100, 16000, lowpass_filter_width=0) |
| with pytest.raises(ValueError): |
| rp.Resampler(44100, 16000, method="nope") |
| with pytest.raises(ValueError): |
| rp.Resampler(44100, 16000)(torch.randn(2, 3, 4)) |
|
|