from __future__ import annotations import wave from pathlib import Path import pytest import torch from torch import nn from music3lab.editing.audio_prepend import ( AudioPrependError, SuffixProjection, SuffixProjector, compose_prepend, load_wav_window, prepend_audio, ) N, F, C, Z, E = 44032, 86, 2, 128, 2048 def write_wav(path: Path, rate: int, channels: int) -> None: with wave.open(str(path), "wb") as handle: handle.setnchannels(channels) handle.setsampwidth(2) handle.setframerate(rate) handle.writeframes(torch.arange(rate * channels, dtype=torch.int16).numpy().tobytes()) class ProjectorFixture: def __init__(self) -> None: self.calls: list[torch.Tensor] = [] def __call__(self, suffix: torch.Tensor) -> SuffixProjection: self.calls.append(suffix.clone()) return SuffixProjection( torch.full((suffix.shape[0], Z, F), 7.0), torch.zeros((suffix.shape[0], F, E)), ) class VelocityFixture: def __init__(self, value: float = 1.0) -> None: self.value = value self.calls: list[torch.Tensor] = [] def __call__(self, latent: torch.Tensor, condition: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor: self.calls.append(latent.clone()) return torch.full_like(latent, self.value) class DecoderFixture: def __call__(self, latent: torch.Tensor, samples: int) -> torch.Tensor: return torch.full((latent.shape[0], C, samples), 0.25) def suffix() -> torch.Tensor: return torch.linspace(-0.9, 0.9, N).repeat(C, 1).unsqueeze(0) def run(**kwargs): return prepend_audio( suffix_audio=kwargs.pop("suffix_audio", suffix()), sample_rate=kwargs.pop("sample_rate", 44100), projector=kwargs.pop("projector", ProjectorFixture()), flow=kwargs.pop("flow", VelocityFixture()), decoder=kwargs.pop("decoder", DecoderFixture()), noise=kwargs.pop("noise", torch.zeros(1, Z, F)), **kwargs, ) def test_wav_mono_stereo_normalization_and_format_rejection(tmp_path: Path) -> None: mono, stereo, wrong = (tmp_path / name for name in ("mono.wav", "stereo.wav", "wrong.wav")) write_wav(mono, 44100, 1) write_wav(stereo, 44100, 2) write_wav(wrong, 48000, 2) left, right = load_wav_window(mono), load_wav_window(stereo) assert left.shape == right.shape == (1, C, N) assert torch.equal(left[:, 0], left[:, 1]) assert not torch.equal(right[:, 0], right[:, 1]) with pytest.raises(AudioPrependError, match="44.1"): load_wav_window(wrong) def test_suffix_is_the_only_conditioner_call_and_shapes_are_fixed() -> None: projector = ProjectorFixture() source = suffix() result = run(projector=projector, suffix_audio=source) assert len(projector.calls) == 1 and torch.equal(projector.calls[0], source) assert result.prefix_latent.shape == (1, Z, F) assert result.conditioning.shape == (1, F, E) assert result.generated_prefix.shape == (1, C, N) @pytest.mark.parametrize("name", ("hidden_prefix", "captured_c0", "native_tokens", "text", "global_cache")) def test_forbidden_inputs_are_rejected(name: str) -> None: with pytest.raises(AudioPrependError, match="suffix-only"): run(**{name: object()}) def test_noise_to_data_sign_and_full_prefix_generation() -> None: velocity = VelocityFixture(1.0) result = run(flow=velocity, steps=2) assert len(velocity.calls) == 2 assert torch.all(velocity.calls[0] == 0) assert torch.all(velocity.calls[1] == 0.5) assert torch.all(result.prefix_latent == 1.0) def test_equal_power_compositor_preserves_source_after_transition() -> None: source = suffix() generated = torch.full_like(source, 0.25) result = compose_prepend(generated, source, transition_samples=1024) assert result.shape[-1] == 2 * N - 1024 assert torch.equal(result[:, :, N:], source[:, :, 1024:]) assert torch.equal(result[:, :, : N - 1024], generated[:, :, : N - 1024]) for bad in (0, 1025): with pytest.raises(AudioPrependError, match="1024"): compose_prepend(generated, source, transition_samples=bad) class TinyEncoder(nn.Module): def __init__(self) -> None: super().__init__() self.weight = nn.Parameter(torch.ones(())) def forward(self, audio: torch.Tensor) -> torch.Tensor: return audio.mean(dim=-1, keepdim=True).mean(dim=1, keepdim=True).expand(-1, Z, F) * self.weight def test_projector_freezes_encoder_and_gradients_stop_at_suffix_latent() -> None: encoder = TinyEncoder() projector = SuffixProjector(encoder, width=16, layers=1) output = projector(suffix()).conditioning output.square().mean().backward() assert encoder.weight.requires_grad is False and encoder.weight.grad is None assert any(parameter.grad is not None for name, parameter in projector.named_parameters() if not name.startswith("encoder.")) def test_deterministic_adjacent_pair_direction_and_baseline_inventory() -> None: from music3lab.editing.audio_prepend_runner import BASELINES, make_prepend_pair previous = torch.full((C, N), -0.5) following = torch.full((C, N), 0.5) prefix, suffix_value = make_prepend_pair(previous, following) assert torch.equal(prefix, previous) and torch.equal(suffix_value, following) assert BASELINES == ("zero_condition", "unrelated_suffix", "repeat_future", "roll_future", "silence") def test_cpu_focused_batch_keeps_target_out_of_projector() -> None: projector = ProjectorFixture() hidden = torch.randn(2, C, N) following = torch.randn(2, C, N) run(projector=projector, suffix_audio=following, noise=torch.randn(2, Z, F)) assert torch.equal(projector.calls[0], following) assert not torch.equal(projector.calls[0], hidden)