music3lab / tests /test_captured_continuation.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
5.23 kB
"""CPU contract for continuation from authenticated Music3 capture only."""
from __future__ import annotations
from dataclasses import replace
from pathlib import Path
import pytest
import torch
from music3lab.editing.continue_capture import (
Capture, ChunkGeometry, extend_capture, load_verified_capture, rebuild_cache,
replay_fused_to_audio, splice_original, teacher_force_capture,
)
class TinyLM:
def __init__(self, generated=(31, 32, 999)):
self.generated, self.rows = iter(generated), []
def prefill(self, rows):
self.rows.extend(row.clone() for row in rows)
return torch.stack([row[:2].float() for row in rows]), {"sequence": rows.clone()}
def step(self, cache):
token = next(self.generated)
row = torch.tensor((token, token + 1, token + 2))
self.rows.append(row)
return row[:2].float(), token, {"sequence": torch.cat((cache["sequence"], row[None]))}
class TinyDecoder:
def __call__(self, fused, geometry):
return fused.T.contiguous()
def capture():
rows = torch.tensor(((20, 21, 22), (23, 24, 25), (26, 27, 28), (29, 30, 31)))
fused = torch.tensor(((10., 11.), (20., 21.), (23., 24.), (26., 27.), (29., 30.)))
geometry = ChunkGeometry(200, 100, 172, 86, 258)
return Capture(torch.tensor((10, 11, 12)), rows, fused, TinyDecoder()(fused, geometry), geometry, 999, 6, "verified_music3_capture")
def test_verified_capture_loader_rejects_arbitrary_wav(tmp_path: Path):
expected, seen = capture(), []
assert load_verified_capture(tmp_path, verifier=lambda p: (seen.append(p), expected)[1]) is expected
assert seen == [tmp_path]
with pytest.raises(ValueError, match="verified|phase0|capture"):
load_verified_capture(tmp_path, verifier=lambda p: replace(expected, source_kind="arbitrary_wav"))
def test_teacher_forcing_exactly_replays_priming_and_rows_without_rng():
source, lm = capture(), TinyLM()
before = torch.get_rng_state().clone()
replay = teacher_force_capture(lm, source)
assert torch.equal(torch.get_rng_state(), before)
assert torch.equal(torch.stack(lm.rows), torch.cat((source.priming_row[None], source.token_rows)))
assert torch.equal(replay.fused_hidden, source.fused_hidden)
assert replay.cache["sequence"].shape == (5, 3)
def test_cache_rebuild_validates_captured_tokens_and_fused_hidden():
source = capture()
rebuilt = rebuild_cache(TinyLM(), source)
assert torch.equal(rebuilt["sequence"], torch.cat((source.priming_row[None], source.token_rows)))
with pytest.raises(ValueError, match="hidden|token|mismatch"):
rebuild_cache(TinyLM(), replace(source, token_rows=source.token_rows + 1))
with pytest.raises(ValueError, match="hidden|mismatch"):
rebuild_cache(TinyLM(), replace(source, fused_hidden=source.fused_hidden + 1))
def test_replay_keeps_captured_chunk_geometry_and_source_audio_exact():
source = capture()
assert torch.equal(replay_fused_to_audio(TinyDecoder(), source.fused_hidden, source.geometry), source.waveform)
assert (source.geometry.chunk_frames, source.geometry.chunk_hop, source.geometry.overlap_latent_frames) == (200, 100, 172)
assert (source.geometry.crop_left_latent_frames, source.geometry.crop_right_latent_frames) == (86, 258)
def test_extension_is_chronological_grows_cache_and_stops_at_eoa():
source = capture()
result = extend_capture(TinyLM((31, 32, 999)), TinyDecoder(), source, requested_frames=5)
assert result.termination_reason == "audio_end" and result.generated_token_ids == (31, 32)
assert torch.equal(result.capture.token_rows[:4], source.token_rows)
assert result.capture.fused_hidden[:, 0].tolist() == [10., 20., 23., 26., 29., 31., 32.]
assert result.cache["sequence"].shape == (7, 3)
def test_extension_honors_max_duration_and_rejects_wrong_cache():
source = replace(capture(), max_frames=5)
result = extend_capture(TinyLM((31, 32)), TinyDecoder(), source, requested_frames=8)
assert result.termination_reason == "max_frames" and result.generated_token_ids == (31,)
cache = rebuild_cache(TinyLM(), source)
cache["sequence"] = cache["sequence"][:-1]
with pytest.raises(ValueError, match="cache|sequence|mismatch"):
extend_capture(TinyLM(), TinyDecoder(), source, requested_frames=1, cache=cache)
def test_equal_power_splice_keeps_prefix_and_expected_length():
original = torch.stack((torch.arange(12.), torch.arange(100., 112.)))
extension = torch.stack((torch.arange(1000., 1008.), torch.arange(2000., 2008.)))
output = splice_original(original, extension, overlap_samples=4)
assert output.shape == (2, 16) and torch.equal(output[:, :8], original[:, :8])
assert not torch.equal(output[:, 8:12], original[:, 8:12])
assert not torch.equal(output[:, 8:12], extension[:, :4])
with pytest.raises(ValueError, match="overlap"):
splice_original(original, extension, overlap_samples=0)
def test_extension_refuses_non_capture_wav():
with pytest.raises(ValueError, match="arbitrary WAV|verified_music3_capture|native token"):
extend_capture(TinyLM(), TinyDecoder(), replace(capture(), source_kind="arbitrary_wav"), requested_frames=1)