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)