music3lab / tests /test_inversion_phase2.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
18.3 kB
from __future__ import annotations
import inspect
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
from music3lab.checkpoint_audit_schema import (
BASE_ID,
CONVERTER_SHA256,
DIFFUSERS_REVISION,
MODEL_REVISION,
)
from music3lab.inversion import (
EvaluatorMetrics,
LossConfig,
MetricsPayload,
ObjectiveSnapshot,
OptimizationTraceEntry,
OptimizedExperiment,
TracePayload,
_replay_final_audio_batch,
_verify_inversion_session_with_authorities,
build_experiment_artifacts,
evaluate_thresholds,
evaluator_metrics,
initialize_restarts,
inversion_loss,
load_inversion_config,
publish_inversion_session,
verify_inversion_session,
)
from music3lab.vocoder import (
SHORT_ORACLE_RUN_ID,
VOCODER_SOURCE_SHA256,
DifferentiableFrozenVocoder,
FlowVocoderLatents,
FlowVocoderOracle,
FrozenVocoderLoadReport,
OracleChunk,
OracleDescriptor,
module_state_sha256,
tensor_sha256,
)
ROOT = Path(__file__).resolve().parents[1]
CONFIG = ROOT / "configs" / "inversion-v1.yaml"
def _loss_config() -> LossConfig:
return LossConfig(
charbonnier_epsilon=1e-3,
stft_epsilon=1e-7,
stft_center=False,
definition_version="mrstft-center-false-unscaled-snr-v1",
stft_fft_sizes=(16, 32),
stft_hop_sizes=(4, 8),
envelope_windows=(3, 7),
weights={
"waveform_charbonnier": 1.0,
"mrstft": 0.15,
"mid_side_charbonnier": 0.25,
"multiscale_envelope": 0.15,
"latent_prior": 1e-6,
},
)
class _ToyDecoder(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.projection = torch.nn.Conv1d(2, 2, 1, bias=False)
with torch.no_grad():
self.projection.weight.copy_(
torch.tensor([[[0.85], [0.15]], [[-0.10], [0.90]]])
)
for parameter in self.parameters():
parameter.requires_grad_(False)
def forward(self, latent: torch.Tensor) -> torch.Tensor:
projected = self.projection(latent)
return torch.tanh(
torch.nn.functional.interpolate(
projected,
scale_factor=8,
mode="linear",
align_corners=False,
)
)
def _report() -> FrozenVocoderLoadReport:
return FrozenVocoderLoadReport.create(
base_id=BASE_ID,
model_revision=MODEL_REVISION,
diffusers_revision=DIFFUSERS_REVISION,
converter_sha256=CONVERTER_SHA256,
implementation_sha256=VOCODER_SOURCE_SHA256,
project_git_commit="4" * 40,
project_source_sha256="5" * 64,
project_git_dirty=False,
dav_file_sha256="1" * 64,
converted_vocoder_file_sha256="2" * 64,
raw_tensor_count=548,
raw_numel=122_904_034,
mapped_tensor_count=121,
known_unmapped_tensor_count=427,
unknown_tensor_count=0,
target_tensor_count=121,
exact_target_count=121,
mapping_semantic_digest="3" * 64,
)
def _tiny_adapter() -> DifferentiableFrozenVocoder:
from diffusers.models.autoencoders.minimax_music3_vocoder import (
MiniMaxMusic3Vocoder,
)
model = MiniMaxMusic3Vocoder(
latent_channels=4,
decoder_input_dim=8,
decoder_hidden_dim=16,
upsampling_ratios=(2,),
sampling_rate=44_100,
)
return DifferentiableFrozenVocoder(model, _report())
def _target_audio() -> torch.Tensor:
pattern = torch.where(
torch.arange(44_032) % 2 == 0,
torch.tensor(0.25),
torch.tensor(-0.25),
).float()
return torch.stack((pattern, pattern), dim=0).unsqueeze(0).contiguous()
class _EvidenceDecoder(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.config = SimpleNamespace(latent_channels=128)
self.register_buffer("target", _target_audio().to(torch.bfloat16))
self.zero_audio_gradient = False
self.zero_audio_gradient_restart = None
def forward(self, latent: torch.Tensor) -> torch.Tensor:
if self.zero_audio_gradient:
latent_source = latent.detach()
elif self.zero_audio_gradient_restart is not None:
mask = torch.ones(
(latent.shape[0], 1, 1),
device=latent.device,
dtype=latent.dtype,
)
mask[self.zero_audio_gradient_restart] = 0
latent_source = latent.detach() + (
latent - latent.detach()
) * mask
else:
latent_source = latent
control = (
latent_source[:, :1, :4].float().mean(dim=(1, 2), keepdim=True)
* 0.05
)
return self.target.float().expand(
latent.shape[0], -1, -1
) + control
def _evidence_adapter() -> DifferentiableFrozenVocoder:
return DifferentiableFrozenVocoder(_EvidenceDecoder(), _report())
def _oracle() -> FlowVocoderOracle:
latent = torch.zeros((1, 128, 86), dtype=torch.bfloat16)
audio = _target_audio()
descriptor = OracleDescriptor.create(
oracle_kind="short",
case_name="short_parity",
run_id=SHORT_ORACLE_RUN_ID,
base_id=BASE_ID,
contract_id="7" * 64,
run_manifest_file_sha256="8" * 64,
run_manifest_semantic_digest="9" * 64,
state_manifest_file_sha256="a" * 64,
state_manifest_semantic_digest="b" * 64,
latent_hop_length=512,
sampling_rate=44_100,
chunks=(
OracleChunk(
index=0,
latent_key="chunks.0.final_latents",
latent_shape=(1, 128, 86),
latent_dtype="float32",
latent_content_sha256=tensor_sha256(latent),
latent_length=86,
crop_left_latent_frames=0,
crop_right_latent_frames=0,
crop_left_samples=0,
crop_right_samples=0,
),
),
expected_audio_shape=(1, 2, 44_032),
expected_audio_file_sha256="d" * 64,
expected_audio_content_sha256=tensor_sha256(audio),
)
return FlowVocoderOracle(
descriptor=descriptor,
latents=(FlowVocoderLatents(latent),),
expected_audio=audio,
)
def _fake_result(experiment_id: str) -> OptimizedExperiment:
loaded = load_inversion_config(CONFIG)
experiment = next(
item
for item in loaded.config.experiments
if item.experiment_id == experiment_id
)
oracle = _oracle()
oracle_latent = oracle.latents[0].tensor
initial_latents = initialize_restarts(
experiment,
seeds=loaded.config.execution.restart_seeds,
shape=(4, 128, 86),
oracle_latent=oracle_latent,
)
steps = tuple(range(0, experiment.steps + 1, 10))
trajectory_latents = torch.stack(
tuple(initial_latents * (1.0 - step / experiment.steps) for step in steps)
).contiguous()
target_audio = oracle.expected_audio
adapter = _evidence_adapter()
with torch.no_grad():
trajectory_audio = torch.stack(tuple(
adapter(
FlowVocoderLatents(value.to(dtype=torch.bfloat16))
).float().clamp(-1.0, 1.0)
for value in trajectory_latents
)).contiguous()
entries = []
for position, step in enumerate(steps):
snapshot = ObjectiveSnapshot.from_loss(
inversion_loss(
trajectory_audio[position],
target_audio,
trajectory_latents[position],
loaded.config.loss,
)
)
gradient = (0.0,) * 4 if step == 0 else (0.001,) * 4
entries.append(
OptimizationTraceEntry(
step=step,
objective=snapshot,
gradient_max_abs=gradient,
gradient_mean_abs=gradient,
gradient_norm=gradient,
)
)
trace = TracePayload.create(
experiment_id=experiment_id,
restart_seeds=loaded.config.execution.restart_seeds,
optimization_steps=experiment.steps,
trace_interval_steps=loaded.config.execution.trace_interval_steps,
entries=tuple(entries),
)
initial_audio = trajectory_audio[0].clone()
final_audio = trajectory_audio[-1].clone()
final_latents = trajectory_latents[-1].clone()
evaluator = evaluator_metrics(
final_audio[0:1],
target_audio,
final_latents[0:1],
oracle_latent,
config=loaded.config.loss,
include_latent=experiment_id == "P2-E1",
)
initial_objectives = trace.entries[0].objective.objective
final_objectives = trace.entries[-1].objective.objective
gate = evaluate_thresholds(
torch.tensor(initial_objectives, dtype=torch.float64),
torch.tensor(final_objectives, dtype=torch.float64),
0,
evaluator,
experiment.thresholds,
)
metrics = MetricsPayload.create(
experiment_id=experiment_id,
selected_restart_index=0,
selected_restart_seed=101,
selection_criterion="final_optimization_objective",
selection_tie_rule="lowest_restart_index",
evaluator_computed_after_selection=True,
initial_objective=initial_objectives[0],
final_objective=final_objectives[0],
all_initial_objectives=initial_objectives,
all_final_objectives=final_objectives,
evaluator=evaluator,
thresholds=experiment.thresholds,
threshold_evaluation=gate,
status="FEASIBILITY_PASS",
high_fidelity_status="HIGH_FIDELITY_PASS",
quality_claim="high_fidelity",
)
weight_hash = module_state_sha256(_evidence_adapter().model)
return OptimizedExperiment(
experiment,
initial_latents,
final_latents,
initial_audio,
final_audio,
target_audio,
trajectory_latents,
trajectory_audio,
trace,
metrics,
weight_hash,
weight_hash,
1.0,
123,
456,
)
def test_strict_preregistration_and_tiers_are_frozen() -> None:
loaded = load_inversion_config(CONFIG)
assert loaded.config.execution.restart_seeds == (101, 103, 107, 109)
assert [item.experiment_id for item in loaded.config.experiments] == [
"P2-E1",
"P2-E2",
]
e2 = loaded.config.experiments[1].thresholds
assert e2.minimum_median_objective_improvement_fraction == 0.40
assert e2.minimum_correlation == 0.50
assert e2.high_fidelity_minimum_si_sdr_db == 20
assert e2.high_fidelity_minimum_unscaled_snr_db == 18
assert e2.high_fidelity_maximum_loudness_error_db == 0.5
assert e2.high_fidelity_maximum_stereo_correlation_error == 0.05
assert loaded.config.loss.stft_center is False
assert loaded.config.loss.definition_version == "mrstft-center-false-unscaled-snr-v1"
def test_loss_identity_beats_distortion_for_reconstruction_terms(
monkeypatch: pytest.MonkeyPatch,
) -> None:
config = _loss_config()
original_stft = torch.stft
centers = []
def observed_stft(*args, **kwargs):
centers.append(kwargs.get("center"))
return original_stft(*args, **kwargs)
monkeypatch.setattr(torch, "stft", observed_stft)
target = torch.zeros((1, 2, 96))
latent = torch.zeros((1, 2, 12))
identity = inversion_loss(target, target, latent, config)
distorted = inversion_loss(torch.ones_like(target), target, torch.ones_like(latent), config)
for name in (
"waveform_charbonnier",
"mrstft",
"mid_side_charbonnier",
"multiscale_envelope",
"objective",
):
assert getattr(identity, name).item() < getattr(distorted, name).item()
assert centers == [False] * 8
def test_restart_initialization_is_cpu_deterministic_and_seed_independent() -> None:
loaded = load_inversion_config(CONFIG)
experiment = loaded.config.experiments[1]
kwargs = {
"experiment": experiment,
"shape": (4, 2, 12),
"oracle_latent": torch.zeros((1, 2, 12)),
}
first = initialize_restarts(seeds=(101, 103, 107, 109), **kwargs)
second = initialize_restarts(seeds=(101, 103, 107, 109), **kwargs)
reordered = initialize_restarts(
seeds=(109, 101, 103, 107), **kwargs
)
assert first.dtype is torch.float32
assert torch.equal(first, second)
assert torch.equal(first[0], reordered[1])
assert torch.equal(first[3], reordered[0])
assert not torch.equal(first[0], first[1])
def test_frozen_official_vocoder_has_nonzero_input_gradient_only() -> None:
adapter = _tiny_adapter()
latent = torch.randn((1, 4, 3), requires_grad=True)
before = {
name: value.detach().clone()
for name, value in adapter.model.state_dict().items()
}
output = adapter(FlowVocoderLatents(latent))
output.square().mean().backward()
assert latent.grad is not None
assert torch.isfinite(latent.grad).all()
assert torch.count_nonzero(latent.grad)
assert all(
not parameter.requires_grad and parameter.grad is None
for parameter in adapter.model.parameters()
)
assert all(
torch.equal(before[name], value)
for name, value in adapter.model.state_dict().items()
)
def test_exact_decoder_replay_preserves_four_restart_batch_geometry() -> None:
class BatchSensitiveDecoder(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.config = SimpleNamespace(latent_channels=128)
self.register_buffer(
"anchor", torch.zeros((), dtype=torch.bfloat16)
)
def forward(self, latent: torch.Tensor) -> torch.Tensor:
value = self.anchor + latent.shape[0] / 10
return value.expand(latent.shape[0], 2, 32).contiguous()
adapter = DifferentiableFrozenVocoder(BatchSensitiveDecoder(), _report())
latents = torch.zeros((4, 128, 2), dtype=torch.float32)
with torch.no_grad():
expected = adapter(
FlowVocoderLatents(latents.to(torch.bfloat16))
).float().clamp(-1.0, 1.0)
wrong_geometry = adapter(
FlowVocoderLatents(latents[:1].to(torch.bfloat16))
).float().clamp(-1.0, 1.0)
assert not torch.equal(wrong_geometry, expected[:1])
_replay_final_audio_batch(adapter, latents, expected)
with pytest.raises(RuntimeError, match="batch replay differs"):
_replay_final_audio_batch(adapter, latents[:1], expected[:1])
def test_toy_optimizer_lowers_real_inversion_objective() -> None:
decoder = _ToyDecoder()
target_latent = torch.zeros((1, 2, 12))
target = decoder(target_latent).detach()
master = torch.nn.Parameter(torch.ones((1, 2, 12)))
optimizer = torch.optim.Adam([master], lr=0.1)
initial = inversion_loss(decoder(master), target, master, _loss_config()).objective.item()
for _ in range(8):
optimizer.zero_grad()
loss = inversion_loss(decoder(master), target, master, _loss_config()).objective.mean()
loss.backward()
optimizer.step()
final = inversion_loss(decoder(master), target, master, _loss_config()).objective.item()
assert final < initial
def test_batch_gate_and_high_fidelity_are_distinct() -> None:
loaded = load_inversion_config(CONFIG)
thresholds = loaded.config.experiments[1].thresholds
feasible_only = EvaluatorMetrics(
waveform_mae=0.05,
waveform_rmse=0.06,
correlation=0.6,
si_sdr_db=0,
unscaled_snr_db=0,
loudness_error_db=1,
stereo_correlation_error=0.2,
mrstft=1,
multiscale_envelope=1,
latent_rmse=None,
)
gate = evaluate_thresholds(
torch.ones(4),
torch.tensor((0.4, 0.5, 0.6, 0.7)),
0,
feasible_only,
thresholds,
)
assert gate.median_objective_improvement_pass
assert gate.feasibility_pass
assert not gate.high_fidelity_pass
def test_atomic_artifacts_verify_and_reject_tamper_and_mode(tmp_path: Path) -> None:
loaded = load_inversion_config(CONFIG)
adapter = _evidence_adapter()
oracle = _oracle()
built = tuple(
build_experiment_artifacts(
_fake_result(experiment_id),
loaded_config=loaded,
adapter=adapter,
oracle=oracle,
device_name="NVIDIA H100 80GB HBM3",
device_capability=(9, 0),
cuda_runtime="13.0",
)
for experiment_id in ("P2-E1", "P2-E2")
)
output = tmp_path / "session"
verified = publish_inversion_session(
output_root=output,
loaded_config=loaded,
adapter=adapter,
oracle=oracle,
experiments=built,
)
assert verified.manifest.all_feasibility_pass
assert verified.manifest.all_high_fidelity_pass
assert len(verified.file_paths) == 17
assert (output.stat().st_mode & 0o777) == 0o755
assert ((output / "P2-E1").stat().st_mode & 0o777) == 0o700
trace = output / "P2-E1" / "trace.json"
original = trace.read_bytes()
trace.write_bytes(b"tampered")
with pytest.raises(RuntimeError, match="size|hash"):
_verify_inversion_session_with_authorities(
output,
loaded_config=loaded,
adapter=adapter,
oracle=oracle,
)
trace.write_bytes(original)
trace.chmod(0o666)
with pytest.raises(RuntimeError, match="mode"):
_verify_inversion_session_with_authorities(
output,
loaded_config=loaded,
adapter=adapter,
oracle=oracle,
)
def test_public_verifier_requires_live_authority_paths() -> None:
parameters = inspect.signature(verify_inversion_session).parameters
assert tuple(parameters) == (
"output_root",
"config_path",
"snapshot",
"base_manifest",
"diffusers_root",
"phase0_artifacts",
)
assert not {
"expected_config_digest",
"expected_adapter_digest",
"expected_oracle_digest",
} & set(parameters)