| """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) |
|
|