music3lab / tests /test_audio_continuation.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
5.37 kB
"""Expected-red CPU contract for prompt-free arbitrary-WAV continuation.
This deliberately imports the future continuous-latent implementation. It is
not a native AR/RVQ continuation test and it makes no long-song claim.
"""
from __future__ import annotations
import wave
from pathlib import Path
import pytest
import torch
from music3lab.editing.audio_continuation import ( # expected red until implemented
AudioContinuationError,
AudioContinuationProjection,
continue_audio,
load_audio_tail,
)
SAMPLE_RATE = 44_100
SAMPLES = 44_032
OVERLAP = 1_024
CAPABILITY = "continuous-latent local continuation; not native AR/token continuation or long-song generalization"
def _source(samples: int = SAMPLES) -> torch.Tensor:
time = torch.arange(samples, dtype=torch.float32) / SAMPLE_RATE
return torch.stack((0.25 * torch.sin(2 * torch.pi * 220 * time), 0.2 * torch.cos(2 * torch.pi * 330 * time))).unsqueeze(0)
def _write_pcm(path: Path, audio: torch.Tensor, sample_rate: int) -> None:
channels = audio.shape[0]
pcm = (audio.clamp(-1, 1).transpose(0, 1) * 32767).to(torch.int16).numpy().tobytes()
with wave.open(str(path), "wb") as output:
output.setnchannels(channels)
output.setsampwidth(2)
output.setframerate(sample_rate)
output.writeframes(pcm)
class _Projector:
capability = CAPABILITY
def __init__(self) -> None:
self.seen: list[torch.Tensor] = []
def __call__(self, source_audio_tail: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
self.seen.append(source_audio_tail.clone())
return torch.ones(1, 128, 86), torch.ones(1, 86, 2048)
class _Renderer:
def __init__(self, audio: torch.Tensor) -> None:
self.audio, self.calls = audio, 0
def __call__(self, encoded_context: torch.Tensor, conditioning: torch.Tensor) -> torch.Tensor:
self.calls += 1
assert encoded_context.shape == (1, 128, 86)
assert conditioning.shape == (1, 86, 2048)
return self.audio.clone()
def test_load_audio_tail_normalizes_recent_mono_and_stereo_wav_to_music3_shape(tmp_path: Path) -> None:
mono = _source(SAMPLES // 2)[0, :1]
stereo = _source(SAMPLES)[0]
mono_path, stereo_path = tmp_path / "mono.wav", tmp_path / "stereo.wav"
_write_pcm(mono_path, mono, SAMPLE_RATE // 2)
_write_pcm(stereo_path, stereo, SAMPLE_RATE)
for path in (mono_path, stereo_path):
tail = load_audio_tail(path)
assert tail.shape == (1, 2, SAMPLES)
assert tail.dtype == torch.float32 and bool(torch.isfinite(tail).all())
assert tail.abs().max() > 0
def test_projection_receives_only_audio_tail_and_exposes_continuous_geometry() -> None:
projector, source = _Projector(), _source()
projection = AudioContinuationProjection(projector)(source)
assert len(projector.seen) == 1 and torch.equal(projector.seen[0], source)
assert projection.encoded_context.shape == (1, 128, 86)
assert projection.conditioning.shape == (1, 86, 2048)
assert projection.capability == CAPABILITY
def test_continuation_rejects_text_captured_tokens_and_future_audio_inputs() -> None:
source, projector, renderer = _source(), _Projector(), _Renderer(_source())
forbidden = {
"text": "continue this song",
"captured_c0": torch.zeros(1, 86, 2048),
"native_tokens": torch.zeros(1, 8, 86, dtype=torch.long),
"right_audio": _source(),
}
for name, value in forbidden.items():
with pytest.raises(AudioContinuationError, match=name):
continue_audio(source, projector=projector, renderer=renderer, **{name: value})
def test_continuation_renders_one_section_and_preserves_pre_overlap_source_bit_exactly() -> None:
source, projector = _source(), _Projector()
renderer = _Renderer(torch.flip(_source(), dims=(-1,)))
result = continue_audio(source, projector=projector, renderer=renderer, overlap_samples=OVERLAP)
assert renderer.calls == 1
assert result.shape == (1, 2, SAMPLES + SAMPLES - OVERLAP)
assert torch.equal(result[:, :, : SAMPLES - OVERLAP], source[:, :, : SAMPLES - OVERLAP])
assert result[:, :, -SAMPLES:].shape == (1, 2, SAMPLES)
def test_continuation_rejects_silent_nonfinite_or_source_copy_renderer_output() -> None:
source, projector = _source(), _Projector()
invalid = (torch.zeros_like(source), torch.full_like(source, float("nan")), source)
for rendered in invalid:
with pytest.raises(AudioContinuationError, match="finite|silent|copy"):
continue_audio(source, projector=projector, renderer=_Renderer(rendered), overlap_samples=OVERLAP)
def test_both_audio_conditioning_beats_zero_and_unrelated_injected_metric_baselines() -> None:
source, projector, rendered = _source(), _Projector(), torch.flip(_source(), dims=(-1,))
scores = {"both": 0.91, "zero": 0.40, "unrelated": 0.31}
result = continue_audio(
source,
projector=projector,
renderer=_Renderer(rendered),
overlap_samples=OVERLAP,
baseline_conditions={"zero": torch.zeros(1, 86, 2048), "unrelated": -torch.ones(1, 86, 2048)},
metric=lambda name, _: scores[name],
)
assert result.metrics == scores
assert result.metrics["both"] > result.metrics["zero"] > result.metrics["unrelated"]
assert result.capability == CAPABILITY