Download tests/test_external_encoder_interim.py from coolpoodle/music3lab: direct link, hf CLI and curl.
- Browser
- Download file 5.86 kB
-
https://huggingface.co/coolpoodle/music3lab/resolve/main/tests/test_external_encoder_interim.py
- Command line
-
hf download hf://coolpoodle/music3lab/tests/test_external_encoder_interim.py
-
curl -L -o test_external_encoder_interim.py https://huggingface.co/coolpoodle/music3lab/resolve/main/tests/test_external_encoder_interim.py
5.86 kB
| from __future__ import annotations | |
| import os | |
| from pathlib import Path | |
| import pytest | |
| import torch | |
| import music3lab.codec.flow_encoder as flow_encoder | |
| from music3lab.codec.external_finetune import ( | |
| InterimExternalDataConfig, | |
| load_external_finetune_config, | |
| load_interim_source_splits, | |
| ) | |
| from music3lab.codec.flow_encoder import load_flow_encoder_config | |
| from music3lab.inversion_v2 import mid_side_nmse | |
| ROOT = Path(__file__).resolve().parents[1] | |
| CONFIG = ROOT / "configs" / "external-flow-encoder-interim-678-v1.yaml" | |
| def test_interim_config_versions_counts_without_redefining_final_api() -> None: | |
| loaded = load_external_finetune_config(CONFIG) | |
| assert ( | |
| loaded.config.schema_version | |
| == "music3lab.external-flow-encoder-finetune.interim678.v1" | |
| ) | |
| assert isinstance(loaded.config.data, InterimExternalDataConfig) | |
| assert loaded.config.data.split_counts == { | |
| "train": 542, | |
| "validation": 68, | |
| "heldout": 68, | |
| } | |
| assert loaded.config.data.dataset_status == "not_2k_final" | |
| assert loaded.config.training.steps == 12000 | |
| assert loaded.config.training.validation_interval_steps == 500 | |
| assert loaded.config.loss.external_ruler_weight == pytest.approx(0.70) | |
| assert loaded.config.loss.teacher_weight == pytest.approx(0.30) | |
| assert loaded.config.loss.teacher_ruler_weight == pytest.approx(0.05) | |
| assert loaded.config.loss.external_prior_weight == pytest.approx(0.005) | |
| def test_real_frozen_interim_assignment_is_loaded_without_resplitting() -> None: | |
| value = os.environ.get("MINIMAX_INTERIM_678_ROOT") | |
| if value is None: | |
| pytest.skip("MINIMAX_INTERIM_678_ROOT is not set") | |
| root = Path(value).resolve(strict=True) | |
| loaded = load_external_finetune_config(CONFIG) | |
| assert isinstance(loaded.config.data, InterimExternalDataConfig) | |
| splits = load_interim_source_splits( | |
| root / "splits.json", | |
| loaded.config.data, | |
| verify_audio_files=False, | |
| ) | |
| assert {name: len(rows) for name, rows in splits.items()} == { | |
| "train": 542, | |
| "validation": 68, | |
| "heldout": 68, | |
| } | |
| all_rows = [row for rows in splits.values() for row in rows] | |
| assert len({row.source_id for row in all_rows}) == 678 | |
| assert len({row.canonical_sha256 for row in all_rows}) == 678 | |
| assert all( | |
| Path(row.canonical_path).parent == (root / "files") | |
| for row in all_rows | |
| if row.canonical_path is not None | |
| ) | |
| def test_interim_loader_rejects_even_reformatted_assignment(tmp_path: Path) -> None: | |
| value = os.environ.get("MINIMAX_INTERIM_678_ROOT") | |
| if value is None: | |
| pytest.skip("MINIMAX_INTERIM_678_ROOT is not set") | |
| source = Path(value).resolve(strict=True) | |
| (tmp_path / "manifest.jsonl").write_bytes((source / "manifest.jsonl").read_bytes()) | |
| (tmp_path / "summary.json").write_bytes((source / "summary.json").read_bytes()) | |
| data = (source / "splits.json").read_text(encoding="utf-8") | |
| (tmp_path / "splits.json").write_text(data + " ", encoding="utf-8") | |
| loaded = load_external_finetune_config(CONFIG) | |
| assert isinstance(loaded.config.data, InterimExternalDataConfig) | |
| with pytest.raises(ValueError, match="split file SHA-256"): | |
| load_interim_source_splits( | |
| tmp_path / "splits.json", | |
| loaded.config.data, | |
| verify_audio_files=False, | |
| ) | |
| def test_nonfinite_ruler_names_component_bad_row_and_exact_values( | |
| monkeypatch: pytest.MonkeyPatch, | |
| ) -> None: | |
| finite = lambda *args, **kwargs: torch.tensor([0.25, 0.5]) | |
| monkeypatch.setattr(flow_encoder, "target_energy_normalized_time_nmse", finite) | |
| monkeypatch.setattr( | |
| flow_encoder, | |
| "complex_stft_nmse", | |
| lambda *args, **kwargs: torch.tensor([0.25, float("nan")]), | |
| ) | |
| monkeypatch.setattr(flow_encoder, "mrstft_magnitude_distance", finite) | |
| monkeypatch.setattr(flow_encoder, "mid_side_nmse", finite) | |
| monkeypatch.setattr(flow_encoder, "relative_envelope_distance", finite) | |
| config = load_flow_encoder_config( | |
| ROOT / "configs" / "flow-encoder-v1.yaml" | |
| ).config | |
| audio = torch.zeros(2, 2, 2048) | |
| with pytest.raises(FloatingPointError) as observed: | |
| flow_encoder.audio_ruler_components(audio, audio, config.loss) | |
| message = str(observed.value) | |
| assert "complex_stft_nmse" in message | |
| assert "row_finite=[true,false]" in message | |
| assert "values=[0.25, nan]" in message | |
| assert "time_nmse" not in message | |
| def test_mid_side_absolute_epsilon_handles_silence_without_changing_ordinary_targets() -> None: | |
| silent_target = torch.zeros(1, 2, 32) | |
| nonzero_prediction = torch.full((1, 2, 32), 0.125) | |
| silent_score = mid_side_nmse( | |
| nonzero_prediction, | |
| silent_target, | |
| floor_fraction=1e-4, | |
| epsilon=1e-8, | |
| ) | |
| assert bool(torch.isfinite(silent_score).all()) | |
| target = torch.stack( | |
| ( | |
| torch.linspace(-0.75, 0.75, 32), | |
| torch.linspace(0.5, -0.25, 32), | |
| ) | |
| ).unsqueeze(0) | |
| predicted = target * 0.8 + 0.03 | |
| pred_mid = 0.5 * (predicted[:, 0] + predicted[:, 1]) | |
| pred_side = 0.5 * (predicted[:, 0] - predicted[:, 1]) | |
| target_mid = 0.5 * (target[:, 0] + target[:, 1]) | |
| target_side = 0.5 * (target[:, 0] - target[:, 1]) | |
| total = target.square().mean(dim=(1, 2)) | |
| legacy_terms = [] | |
| for pred_part, target_part in ( | |
| (pred_mid, target_mid), | |
| (pred_side, target_side), | |
| ): | |
| numerator = (pred_part - target_part).square().mean(dim=1) | |
| denominator = torch.maximum( | |
| target_part.square().mean(dim=1), 1e-4 * total | |
| ) | |
| legacy_terms.append(numerator / denominator) | |
| legacy = torch.stack(legacy_terms).mean(dim=0) | |
| corrected = mid_side_nmse(predicted, target, floor_fraction=1e-4, epsilon=1e-8) | |
| assert torch.equal(corrected, legacy) | |