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