File size: 5,963 Bytes
b9dc61d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | """CPU contract for an in-domain captured-Music3 token-prediction experiment.
This is deliberately not a native Music3 tokenizer and it makes no claim for
arbitrary external music. Its only eligible teachers are paired captured
Music3 waveforms, tokens, priming rows, and re-emitted waveforms.
"""
from __future__ import annotations
import pytest
import torch
from music3lab.codec import native_token_adapter as adapter
SAMPLES = 44_032
FRAMES = 25
SPLIT_SEEDS = {
"train": tuple(range(1000, 1064)),
"validation": tuple(range(2000, 2016)),
"heldout": tuple(range(3000, 3016)),
}
def _sidecar(split: str) -> adapter.TokenTeacherSidecar:
seeds = SPLIT_SEEDS[split]
count = len(seeds)
audio = torch.zeros(count, 2, SAMPLES)
audio[:, 0, 0] = torch.tensor(seeds, dtype=audio.dtype)
semantic = torch.arange(count * FRAMES, dtype=torch.int64).reshape(count, FRAMES) % 16_384
residual = torch.arange(count * FRAMES * 7, dtype=torch.int64).reshape(count, FRAMES, 7) % 1_024
priming = torch.arange(count * 8, dtype=torch.int64).reshape(count, 8) % 1_024
return adapter.TokenTeacherSidecar(
split=split,
seeds=seeds,
audio=audio,
semantic=semantic,
residual=residual,
priming=priming,
reemitted_audio=audio.clone(),
source_domain="captured_music3",
)
def test_teacher_sidecars_have_exact_64_16_16_seed_pairing_and_reemission() -> None:
sidecars = tuple(_sidecar(name) for name in ("train", "validation", "heldout"))
report = adapter.validate_token_teacher_sidecars(sidecars)
assert report.split_counts == {"train": 64, "validation": 16, "heldout": 16}
assert report.seed_sets == SPLIT_SEEDS
assert report.source_domain == "captured_music3_only"
missing = _sidecar("train")
missing.reemitted_audio = None
with pytest.raises(ValueError, match="reemitted_audio"):
adapter.validate_token_teacher_sidecars((missing,))
mismatched = _sidecar("train")
mismatched.reemitted_audio[0, 0, 0] += 1
with pytest.raises(ValueError, match="re-emitted audio"):
adapter.validate_token_teacher_sidecars((mismatched,))
def test_waveform_adapter_emits_all_codebook_logits_and_per_code_cross_entropy() -> None:
batch = _sidecar("validation")
model = adapter.CapturedMusic3TokenAdapter(adapter.TokenAdapterConfig(hidden_channels=8))
semantic_logits, residual_logits = model(batch.audio[:2])
assert semantic_logits.shape == (2, FRAMES, 16_384)
assert residual_logits.shape == (2, FRAMES, 7, 1_024)
loss = adapter.token_cross_entropy(
semantic_logits,
residual_logits,
batch.semantic[:2],
batch.residual[:2],
)
assert set(loss) == {"semantic", "residual", "total"}
assert loss["total"].isfinite() and loss["total"] > 0
loss["total"].backward()
assert any(parameter.grad is not None for parameter in model.parameters())
def test_modes_teacher_forcing_and_predicted_tokens_are_range_valid_and_renderable() -> None:
batch = _sidecar("validation")
modes = adapter.per_position_mode_tokens(batch.semantic, batch.residual)
assert modes.shape == (FRAMES, 8)
adapter.validate_predicted_tokens(modes.unsqueeze(0))
backend = _TinyCapturedFuser()
captured_fused = backend.expected(batch.semantic[:1], batch.residual[:1], batch.priming[:1])
replayed = adapter.teacher_force_captured_rows(
backend, batch.semantic[:1], batch.residual[:1], batch.priming[:1]
)
assert torch.equal(replayed, captured_fused)
assert backend.sampling_rng is None
logits = torch.full((1, FRAMES, 16_384), -20.0)
logits[..., 7] = 20.0
residual_logits = torch.full((1, FRAMES, 7, 1_024), -20.0)
residual_logits[..., 9] = 20.0
predicted = adapter.predict_tokens(logits, residual_logits)
adapter.validate_predicted_tokens(predicted)
rendered = adapter.render_predicted_tokens(_TinyRenderer(), predicted)
assert rendered.shape == (1, 2, SAMPLES)
capability = adapter.describe_capability()
assert capability["native_tokenizer"] is False
assert capability["arbitrary_external_music"] is False
assert capability["domain"] == "captured_music3_only"
class _TinyCapturedFuser:
sampling_rng = "not-called"
@staticmethod
def expected(semantic: torch.Tensor, residual: torch.Tensor, priming: torch.Tensor) -> torch.Tensor:
rows = torch.cat((semantic.unsqueeze(-1), residual), dim=-1).to(torch.float32)
return rows.mean(dim=-1, keepdim=True).expand(-1, -1, 32_768).contiguous()
def teacher_force(self, semantic: torch.Tensor, residual: torch.Tensor, priming: torch.Tensor, *, sampling_rng=None) -> torch.Tensor:
self.sampling_rng = sampling_rng
return self.expected(semantic, residual, priming)
class _TinyRenderer:
def render(self, tokens: torch.Tensor) -> torch.Tensor:
assert tokens.shape[1:] == (FRAMES, 8)
return torch.zeros(tokens.shape[0], 2, SAMPLES)
def test_frozen_native_token_config_and_cli_are_cpu_safe() -> None:
import os
import subprocess
import sys
from pathlib import Path
loaded = adapter.load_native_token_config(
Path(__file__).parents[1] / "configs/native-token-adapter-v1.yaml"
)
assert loaded.config.teacher.split_seeds() == SPLIT_SEEDS
assert loaded.config.native_tokenizer is False
result = subprocess.run(
[
sys.executable,
"-B",
str(Path(__file__).parents[1] / "scripts/run_native_token_adapter.py"),
"--help",
],
check=True,
capture_output=True,
text=True,
env={
"PATH": os.environ["PATH"],
"PYTHONPATH": str(Path(__file__).parents[1] / "src"),
"CUDA_VISIBLE_DEVICES": "-1",
"PYTHONDONTWRITEBYTECODE": "1",
},
)
assert "not a native tokenizer" in " ".join(result.stdout.split())
|