File size: 6,700 Bytes
90884df | 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 | """CPU contract for an oracle-free learned flow-latent encoder.
The test decoder is intentionally tiny. It maps required continuous Music3
flow shape [B,128,86] to [B,2,44032], never RVQ/native tokens.
"""
from __future__ import annotations
from dataclasses import asdict
from pathlib import Path
import pytest
import torch
from torch import nn
from torch.nn import functional as F
from music3lab.codec import learned_flow_encoder as flow
SAMPLES = 44_032
SEEDS = (101, 103, 107, 109)
class TinyFrozenFlowDecoder(nn.Module):
def __init__(self) -> None:
super().__init__()
self.gain = nn.Parameter(torch.tensor([1.25, 0.75]), requires_grad=False)
def forward(self, z: torch.Tensor) -> torch.Tensor:
assert z.shape[1:] == (128, 86)
audio = F.interpolate(z[:, :2], size=SAMPLES, mode="linear", align_corners=False)
return audio * self.gain.view(1, 2, 1)
def _latents(count: int = 12) -> torch.Tensor:
time = torch.linspace(-1.0, 1.0, 86)
values = []
for i in range(count):
z = torch.zeros(128, 86)
z[0] = (i + 1) / count * time
z[1] = .25 * torch.cos((i + 1) * torch.pi * time)
values.append(z)
return torch.stack(values)
def _examples(count: int = 12) -> tuple[flow.AudioExample, ...]:
with torch.no_grad():
audio = TinyFrozenFlowDecoder()(_latents(count))
return tuple(flow.AudioExample(flow.deterministic_audio_id(a), a.clone()) for a in audio)
def _config() -> flow.LearnedFlowEncoderConfig:
return flow.LearnedFlowEncoderConfig(
latent_channels=128, latent_frames=86, audio_channels=2, audio_samples=SAMPLES,
hidden_channels=8, learning_rate=.03, batch_size=3, seeds=SEEDS,
)
def test_exact_geometry_frozen_decoder_and_encoder_gradients() -> None:
encoder, decoder = flow.LearnedFlowEncoder(_config()), TinyFrozenFlowDecoder()
audio = torch.randn(2, 2, SAMPLES)
z = encoder(audio)
assert z.shape == (2, 128, 86)
flow.audio_reconstruction_loss(decoder(z), audio).backward()
assert all(p.grad is None for p in decoder.parameters())
assert any(p.grad is not None and p.grad.abs().sum() > 0 for p in encoder.parameters())
def test_deterministic_ids_and_four_seed_splits_are_complete_disjoint_and_leak_free() -> None:
examples = _examples()
changed = examples[0].audio.clone()
changed[0, 0] += 1e-4
assert flow.deterministic_audio_id(changed) != examples[0].sample_id
first = flow.build_four_seed_splits(examples, seeds=SEEDS, validation_fraction=.25)
assert first == flow.build_four_seed_splits(examples, seeds=SEEDS, validation_fraction=.25)
assert tuple(first) == SEEDS
all_ids = {x.sample_id for x in examples}
for seed, split in first.items():
assert seed in SEEDS and split.train_ids and split.validation_ids
assert set(split.train_ids).isdisjoint(split.validation_ids)
assert set(split.train_ids) | set(split.validation_ids) == all_ids
assert split.train_ids == tuple(sorted(split.train_ids))
assert split.validation_ids == tuple(sorted(split.validation_ids))
with pytest.raises((TypeError, ValueError), match="latent|unknown|field"):
flow.AudioExample(examples[0].sample_id, examples[0].audio, latent_target=torch.zeros(128, 86))
def test_loss_identity_order_and_heldout_metric_are_audio_only() -> None:
decoder, target = TinyFrozenFlowDecoder(), _examples(2)[0].audio.unsqueeze(0)
exact = flow.audio_reconstruction_loss(target, target)
distorted = flow.audio_reconstruction_loss(target + .1, target)
inverted = flow.audio_reconstruction_loss(-target, target)
assert exact.item() == pytest.approx(0.0, abs=1e-12)
assert distorted > exact and inverted > distorted
metric = flow.evaluate_heldout(flow.LearnedFlowEncoder(_config()), decoder, _examples(4), torch.zeros(1, 128, 86))
assert set(metric) == {"encoder_loss", "prior_mean_loss", "improvement_fraction", "sample_count"}
assert metric["sample_count"] == 4 and metric["encoder_loss"] >= 0 and metric["prior_mean_loss"] > 0
assert metric["improvement_fraction"] == pytest.approx(1 - metric["encoder_loss"] / metric["prior_mean_loss"])
assert "token" not in " ".join(metric).lower()
def test_short_cpu_training_reduces_heldout_loss_below_prior_mean() -> None:
torch.manual_seed(7)
examples, config, decoder = _examples(), _config(), TinyFrozenFlowDecoder()
split = flow.build_four_seed_splits(examples, seeds=SEEDS, validation_fraction=.25)[101]
encoder = flow.LearnedFlowEncoder(config)
before = flow.evaluate_heldout(encoder, decoder, split.validation_examples(examples), torch.zeros(1, 128, 86))
run = flow.train_encoder(encoder, decoder, split.training_examples(examples), split.validation_examples(examples), config, steps=36)
after = flow.evaluate_heldout(encoder, decoder, split.validation_examples(examples), torch.zeros(1, 128, 86))
assert run.training_losses[0] > run.training_losses[-1]
assert after["encoder_loss"] < before["encoder_loss"]
assert after["encoder_loss"] < after["prior_mean_loss"] and after["improvement_fraction"] > 0
def test_optional_refinement_is_bounded_and_never_claims_native_tokens() -> None:
decoder, target = TinyFrozenFlowDecoder(), _examples(1)[0].audio.unsqueeze(0)
initial = torch.zeros(1, 128, 86)
result = flow.optional_refine_latents(decoder, initial, target, steps=4, learning_rate=.1, max_steps=4)
assert result.latents.shape == (1, 128, 86)
assert 0 <= result.steps_used <= 4 and result.final_loss <= result.initial_loss
assert result.kind == "continuous_flow_latent" and "token" not in result.kind
with pytest.raises((TypeError, ValueError), match="max_steps|bound"):
flow.optional_refine_latents(decoder, initial, target, steps=5, learning_rate=.1, max_steps=4)
def test_checkpoint_contains_only_encoder_config_and_provenance(tmp_path: Path) -> None:
config, encoder = _config(), flow.LearnedFlowEncoder(_config())
path = tmp_path / "learned-flow-encoder.pt"
provenance = {"base_commit": "9606448", "dataset_ids": tuple(x.sample_id for x in _examples()), "split_seeds": SEEDS, "decoder_identity": "frozen-test-decoder"}
flow.save_encoder_checkpoint(path, encoder, config, provenance)
payload = torch.load(path, map_location="cpu", weights_only=False)
assert set(payload) == {"schema_version", "encoder_state_dict", "config", "provenance"}
assert payload["config"] == asdict(config) and payload["provenance"] == provenance
serialized = repr(payload).lower()
for word in ("decoder_state", "latent_target", "native", "rvq"):
assert word not in serialized
|