| """Expected-red CPU contract for learned arbitrary-WAV local continuation. |
| |
| This deliberately specifies a continuous-latent, local path only. It does |
| not claim native Music3 token/AR continuation or long-song conditioning. |
| """ |
| from __future__ import annotations |
|
|
| import torch |
| import pytest |
| from torch import nn |
|
|
| from music3lab.editing.learned_audio_continuation import ( |
| CAPABILITY, |
| ContinuationRecord, |
| analyze_tail, |
| build_consecutive_windows, |
| compose_continuation, |
| conditional_velocity, |
| evaluate_continuation, |
| flow_velocity_target, |
| freeze_for_continuation_training, |
| train_continuation_step, |
| ) |
|
|
|
|
| SAMPLES = 44_032 |
| OVERLAP = 1_024 |
|
|
|
|
| def _track(windows: int = 3) -> torch.Tensor: |
| """Stereo samples whose value identifies its absolute source position.""" |
| values = torch.arange(windows * SAMPLES, dtype=torch.float32) |
| return torch.stack((values, -values)) |
|
|
|
|
| class _RecordingAdapter(nn.Module): |
| def __init__(self) -> None: |
| super().__init__() |
| self.projector = nn.Linear(128, 128, bias=False) |
| self.lora_scale = nn.Parameter(torch.tensor(0.0)) |
| self.calls: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] |
|
|
| def forward( |
| self, |
| context_latent: torch.Tensor, |
| noisy_target_latent: torch.Tensor, |
| t: torch.Tensor, |
| ) -> torch.Tensor: |
| self.calls.append((context_latent.detach().clone(), noisy_target_latent.detach().clone(), t.detach().clone())) |
| return noisy_target_latent + self.lora_scale * context_latent |
|
|
|
|
| def test_consecutive_source_exclusive_windows_are_deterministic_and_causal() -> None: |
| records = [ |
| ContinuationRecord(source_id="song-a", audio=_track(), split_id="train"), |
| ContinuationRecord(source_id="song-b", audio=_track(), split_id="heldout"), |
| ] |
| first = build_consecutive_windows(records, context_samples=SAMPLES, target_samples=SAMPLES) |
| second = build_consecutive_windows(list(reversed(records)), context_samples=SAMPLES, target_samples=SAMPLES) |
|
|
| assert [(item.source_id, item.split_id) for item in first] == [(item.source_id, item.split_id) for item in second] |
| assert len(first) == 4 |
| for item in first: |
| assert item.context.shape == item.target.shape == (2, SAMPLES) |
| assert item.context.data_ptr() != item.target.data_ptr() |
| assert torch.equal(item.target[:, :-1], item.context[:, 1:] + 1) is False |
| assert torch.equal(item.target[:, :], item.context[:, -1:].expand_as(item.target)) is False |
| assert torch.equal(item.target[:, 0], item.context[:, -1] + torch.tensor([1.0, -1.0])) |
| assert item.split_id == ("train" if item.source_id == "song-a" else "heldout") |
| with torch.no_grad(): |
| item = first[0] |
| context = item.context.clone() |
| item.target.zero_() |
| assert torch.equal(item.context, context) |
| duplicate = ContinuationRecord(source_id="song-a", audio=_track(), split_id="heldout") |
| with pytest.raises(ValueError, match="source_id.*split"): |
| build_consecutive_windows([records[0], duplicate], context_samples=SAMPLES, target_samples=SAMPLES) |
|
|
|
|
| def test_analysis_strips_only_terminal_exact_zeros_and_composition_preserves_source() -> None: |
| source = _track(2).unsqueeze(0) |
| source[:, :, -3:] = 0 |
| tail = analyze_tail(source) |
| assert torch.equal(tail, source[:, :, :-3]) |
|
|
| rendered = torch.full((1, 2, SAMPLES), 0.25) |
| composed = compose_continuation(source, rendered, overlap_samples=OVERLAP) |
| preserved = source.shape[-1] - OVERLAP |
| assert torch.equal( |
| composed[:, :, :preserved], |
| source[:, :, :preserved], |
| ) |
| assert composed.shape[-1] == source.shape[-1] + rendered.shape[-1] - OVERLAP |
|
|
|
|
| def test_conditional_flow_receives_only_context_noise_and_time_and_freezes_base_modules() -> None: |
| adapter = _RecordingAdapter() |
| encoder, flow, vocoder = nn.Linear(1, 1), nn.Linear(1, 1), nn.Linear(1, 1) |
| freeze_for_continuation_training(adapter, encoder=encoder, flow=flow, vocoder=vocoder) |
| assert all(not parameter.requires_grad for module in (encoder, flow, vocoder) for parameter in module.parameters()) |
| assert {name for name, parameter in adapter.named_parameters() if parameter.requires_grad} == {"projector.weight", "lora_scale"} |
|
|
| context, noise, t = torch.ones(2, 128, 86), torch.zeros(2, 128, 86), torch.tensor([0.2, 0.8]) |
| velocity = conditional_velocity(adapter, context_latent=context, noisy_target_latent=noise, t=t) |
| assert velocity.shape == noise.shape |
| seen_context, seen_noise, seen_t = adapter.calls[-1] |
| assert torch.equal(seen_context, context) and torch.equal(seen_noise, noise) and torch.equal(seen_t, t) |
|
|
|
|
| def test_flow_target_is_data_minus_noise_and_context_changes_same_noise_prediction() -> None: |
| noise, target = torch.full((1, 128, 86), 3.0), torch.full((1, 128, 86), 7.0) |
| assert torch.equal(flow_velocity_target(target_latent=target, noise_latent=noise), target - noise) |
|
|
| adapter = _RecordingAdapter() |
| adapter.lora_scale.data.fill_(1.0) |
| fixed_noise, t = torch.zeros_like(noise), torch.ones(1) |
| left = conditional_velocity(adapter, context_latent=torch.zeros_like(noise), noisy_target_latent=fixed_noise, t=t) |
| changed = conditional_velocity(adapter, context_latent=torch.ones_like(noise), noisy_target_latent=fixed_noise, t=t) |
| assert not torch.equal(left, changed) |
|
|
|
|
| def test_tiny_conditional_learning_step_improves_without_unfreezing_frozen_modules() -> None: |
| adapter = _RecordingAdapter() |
| encoder, flow, vocoder = nn.Linear(1, 1), nn.Linear(1, 1), nn.Linear(1, 1) |
| freeze_for_continuation_training(adapter, encoder=encoder, flow=flow, vocoder=vocoder) |
| frozen_before = [parameter.detach().clone() for module in (encoder, flow, vocoder) for parameter in module.parameters()] |
| context = torch.ones(2, 128, 86) |
| noise, target, t = torch.zeros_like(context), torch.ones_like(context), torch.full((2,), 0.5) |
| optimizer = torch.optim.SGD((parameter for parameter in adapter.parameters() if parameter.requires_grad), lr=0.2) |
| initial = train_continuation_step(adapter, optimizer, context_latent=context, noise_latent=noise, target_latent=target, t=t) |
| final = initial |
| for _ in range(12): |
| final = train_continuation_step(adapter, optimizer, context_latent=context, noise_latent=noise, target_latent=target, t=t) |
| assert final < initial |
| frozen_after = [parameter.detach() for module in (encoder, flow, vocoder) for parameter in module.parameters()] |
| assert all(torch.equal(before, after) for before, after in zip(frozen_before, frozen_after, strict=True)) |
|
|
|
|
| def test_terminal_evaluation_uses_same_noise_baselines_hidden_target_and_seam_floors() -> None: |
| adapter = _RecordingAdapter() |
| result = evaluate_continuation( |
| adapter=adapter, |
| context_latent=torch.ones(1, 128, 86), |
| hidden_next_window=torch.full((1, 2, SAMPLES), 0.5), |
| unrelated_context_latent=-torch.ones(1, 128, 86), |
| repeat_tail_audio=torch.zeros(1, 2, SAMPLES), |
| noise_latent=torch.zeros(1, 128, 86), |
| source_audio=_track(2).unsqueeze(0), |
| overlap_samples=OVERLAP, |
| ) |
| assert CAPABILITY == "continuous-latent local continuation; non-native-token; not long-song" |
| assert result.same_noise is True |
| assert set(result.ruler) == {"conditional", "zero_context", "unrelated_context", "repeat_tail"} |
| assert result.ruler["conditional"] < result.ruler["zero_context"] |
| assert result.ruler["conditional"] < result.ruler["unrelated_context"] |
| assert result.ruler["conditional"] < result.ruler["repeat_tail"] |
| assert result.composed_seam["conditional"] <= result.seam_floor |
| assert result.composed_seam["conditional"] < result.composed_seam["repeat_tail"] |
|
|