File size: 5,863 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
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
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)