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