music3lab / tests /test_external_encoder_finetune.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
6.82 kB
from __future__ import annotations
import copy
import pytest
import torch
from torch import nn
from music3lab.codec.external_finetune import (
ExternalSource,
FineTuneMetrics,
WarmupCosineSchedule,
deterministic_guarded_crop,
evaluate_gates,
load_source_splits,
mixed_batch_loss,
select_checkpoint,
serialize_source_splits,
split_external_sources,
)
USER22 = tuple(f"user22-{index:02d}" for index in range(22))
def _sources() -> tuple[ExternalSource, ...]:
"""Two thousand admissible sources plus user music that must stay OOD."""
return tuple(
ExternalSource(source_id=f"external-{index:04d}", samples=44100 * 12)
for index in range(2000)
) + tuple(
ExternalSource(source_id=source_id, samples=44100 * 12, user_provided=True)
for source_id in USER22
)
def test_source_exclusive_1600_200_200_split_and_guarded_crops_are_deterministic() -> None:
first = split_external_sources(_sources(), seed=73)
second = split_external_sources(_sources(), seed=73)
assert {name: len(rows) for name, rows in first.items()} == {
"train": 1600,
"validation": 200,
"heldout": 200,
}
assert first == second
owners = {
row.source_id: name for name, rows in first.items() for row in rows
}
assert len(owners) == 2000
assert not set(USER22) & set(owners)
crop = deterministic_guarded_crop(
first["train"][0], split="train", epoch=3, item_index=7, seed=73
)
assert crop == deterministic_guarded_crop(
first["train"][0], split="train", epoch=3, item_index=7, seed=73
)
assert crop.samples == 44032
assert 44100 * 5 <= crop.start_sample
assert crop.start_sample + crop.samples <= 44100 * 7
assert crop.source_id in owners and owners[crop.source_id] == "train"
def test_manifest_loader_rejects_user_audio_from_train_or_checkpoint_selection(tmp_path) -> None:
manifest = split_external_sources(_sources(), seed=9)
path = tmp_path / "sources.json"
path.write_text(serialize_source_splits(manifest), encoding="utf-8")
loaded = load_source_splits(path)
assert loaded == manifest
assert not set(USER22) & {
row.source_id for rows in loaded.values() for row in rows
}
with pytest.raises(ValueError, match="user-provided"):
split_external_sources(
_sources(), seed=9, train_source_ids=("external-0000", USER22[0])
)
def test_mixed_loss_uses_six_external_two_music3_and_exact_weights() -> None:
teacher_nmse = torch.tensor(2.0)
teacher_ruler = torch.tensor(3.0)
external_ruler = torch.tensor(5.0)
prior = torch.tensor(7.0)
result = mixed_batch_loss(
teacher_nmse=teacher_nmse,
teacher_ruler=teacher_ruler,
external_ruler=external_ruler,
external_prior=prior,
external_count=6,
music3_count=2,
)
assert result.external_count == 6
assert result.music3_count == 2
assert float(result.total) == pytest.approx(
0.70 * 5.0 + 0.30 * (2.0 + 0.05 * 3.0) + 0.005 * 7.0
)
assert result.teacher_weight == pytest.approx(0.30)
assert result.external_weight == pytest.approx(0.70)
def test_only_encoder_receives_gradients_and_frozen_decoder_is_unchanged() -> None:
encoder = nn.Linear(2, 2, bias=False)
decoder = nn.Linear(2, 2, bias=False)
for parameter in decoder.parameters():
parameter.requires_grad_(False)
decoder_before = copy.deepcopy(decoder.state_dict())
teacher_audio = torch.tensor([[1.0, -1.0], [0.5, 0.25]])
external_audio = torch.tensor([[0.25, -0.5]]).repeat(6, 1)
predicted = encoder(torch.cat((teacher_audio, external_audio), dim=0))
rendered = decoder(predicted)
result = mixed_batch_loss(
teacher_nmse=(predicted[:2] - teacher_audio).square().mean(),
teacher_ruler=(rendered[:2] - teacher_audio).square().mean(),
external_ruler=(rendered[2:] - external_audio).square().mean(),
external_prior=predicted[2:].square().mean(),
external_count=6,
music3_count=2,
)
result.total.backward()
assert all(parameter.grad is not None for parameter in encoder.parameters())
assert all(parameter.grad is None for parameter in decoder.parameters())
assert decoder.state_dict().keys() == decoder_before.keys()
assert all(torch.equal(value, decoder_before[key]) for key, value in decoder.state_dict().items())
def test_warmup_cosine_endpoints_are_frozen() -> None:
schedule = WarmupCosineSchedule(
warmup_steps=10,
total_steps=100,
maximum_learning_rate=1e-3,
minimum_learning_rate=1e-5,
)
assert schedule(0) == pytest.approx(0.0)
assert schedule(10) == pytest.approx(1e-3)
assert schedule(100) == pytest.approx(1e-5)
assert schedule(50) < schedule(10)
def test_synthetic_training_improves_external_ruler_without_teacher_regression() -> None:
scale = nn.Parameter(torch.zeros(()))
optimizer = torch.optim.SGD([scale], lr=0.15)
baseline = FineTuneMetrics(teacher_ruler=1.0, external_ruler=1.1025)
for _ in range(50):
optimizer.zero_grad()
loss = mixed_batch_loss(
teacher_nmse=(scale - 1.0).square(),
teacher_ruler=(scale - 1.0).square(),
external_ruler=(scale - 1.05).square(),
external_prior=scale.square(),
external_count=6,
music3_count=2,
)
loss.total.backward()
optimizer.step()
candidate = FineTuneMetrics(
teacher_ruler=float((scale.detach() - 1.0).square()),
external_ruler=float((scale.detach() - 1.05).square()),
)
gate = evaluate_gates(baseline=baseline, candidate=candidate)
assert gate.external_ruler_improved is True
assert gate.teacher_regression_fraction <= 0.05
assert gate.teacher_regression_within_limit is True
assert gate.passes is True
def test_heldout_gates_and_user22_ood_report_cannot_select_a_checkpoint() -> None:
baseline = FineTuneMetrics(teacher_ruler=1.0, external_ruler=4.0)
checkpoints = {
"step-010": FineTuneMetrics(teacher_ruler=1.02, external_ruler=2.0),
"step-020": FineTuneMetrics(teacher_ruler=1.01, external_ruler=2.1),
}
selected = select_checkpoint(
baseline=baseline,
validation=checkpoints,
heldout={"step-010": FineTuneMetrics(teacher_ruler=4.0, external_ruler=0.1)},
user22_ood={"step-010": FineTuneMetrics(teacher_ruler=0.0, external_ruler=0.0)},
)
assert selected.name == "step-010"
assert selected.selection_split == "validation"
assert selected.heldout_gate is not None
assert selected.user22_ood_report is not None
assert selected.user22_ood_report.influenced_selection is False