from __future__ import annotations import json from pathlib import Path import pytest import torch from pydantic import ValidationError from safetensors.torch import save_file from music3lab.checkpoint_audit import ( CONVERTER_SHA256, ConverterAuthority, DuplicateJSONKeyError, archive_members, audit_safetensors_index, converter_authority_for_tests, nested_search_values, plan_dav_key, plan_flow_key, plan_qwen_key, safe_load_pth, select_tensor_mapping, strict_json_loads, ) from music3lab.checkpoint_audit_schema import CoverageSummary, IndexAudit from music3lab.checkpoint_audit_runner import TensorCollection, _mapping_row def _dav_keys() -> list[str]: keys = ["dec_in_proj.weight", "dec_in_proj.bias"] for model_index in (0, 6): keys += [f"decoder.model.{model_index}.{suffix}" for suffix in ("weight_g", "weight_v", "bias")] keys.append("decoder.model.5.alpha") for block in range(1, 5): keys.append(f"decoder.model.{block}.block.0.alpha") keys += [f"decoder.model.{block}.block.1.{suffix}" for suffix in ("weight_g", "weight_v", "bias")] for unit in (2, 3, 4): keys += [ f"decoder.model.{block}.block.{unit}.block.0.alpha", f"decoder.model.{block}.block.{unit}.block.2.alpha", ] for inner in (1, 3): keys += [ f"decoder.model.{block}.block.{unit}.block.{inner}.{suffix}" for suffix in ("weight_g", "weight_v", "bias") ] return keys def _flow_keys() -> list[str]: keys = [ "cond_layer_logits", "cond_layer_scale", "latent_conditioners.0.weight", "latent_conditioners.0.bias", ] keys += [ "diffusion_transformer." + value for value in ( "timestep_features.weight", "to_timestep_embed.0.weight", "to_timestep_embed.0.bias", "to_timestep_embed.2.weight", "to_timestep_embed.2.bias", "preprocess_conv.weight", "postprocess_conv.weight", "transformer.project_in.weight", "transformer.project_out.weight", ) ] tails = ( "pre_norm.gamma", "pre_norm.beta", "self_attn.to_qkv.weight", "self_attn.to_out.weight", "ff_norm.gamma", "ff_norm.beta", "ff.ff.0.proj.weight", "ff.ff.0.proj.bias", "ff.ff.2.weight", "ff.ff.2.bias", ) for layer in range(36): keys += [ f"diffusion_transformer.transformer.layers.{layer}.{tail}" for tail in tails ] return keys def _qwen_audio_keys() -> list[str]: keys = [ "model.audio_extra_embedding.weight", "model.audio_decoder.projection.weight", "model.audio_decoder.pos_embedding.weight", "model.audio_decoder.norm.weight", ] keys += [f"model.audio_decoder.audio_heads.{index}.weight" for index in range(7)] tails = ( "input_layernorm.weight", "post_attention_layernorm.weight", "self_attn.q_proj.weight", "self_attn.k_proj.weight", "self_attn.v_proj.weight", "self_attn.o_proj.weight", "mlp.gate_proj.weight", "mlp.up_proj.weight", "mlp.down_proj.weight", ) for layer in range(4): keys += [f"model.audio_decoder.layers.{layer}.{tail}" for tail in tails] return keys def test_exact_converter_plans_close_required_source_and_target_counts(): authority = converter_authority_for_tests() dav = [plan_dav_key(key, authority) for key in _dav_keys()] assert len(dav) == 121 assert all(item.status == "mapped" for item in dav) assert sum(len(item.targets) for item in dav) == 121 known_dav = [ *[f"encoder.synthetic.{index}" for index in range(119)], "mean_proj.weight", "mean_proj.bias", "logs_proj.weight", "logs_proj.bias", *[f"flow.synthetic.{index}" for index in range(304)], ] assert len(known_dav) == 427 assert all(plan_dav_key(key, authority).status == "unmapped_known" for key in known_dav) flow = [plan_flow_key(key, authority) for key in _flow_keys()] assert len(flow) == 373 assert all(item.status == "mapped" for item in flow) assert sum(len(item.targets) for item in flow) == 445 assert sum(item.transform == "chunk(dim=0,parts=3)" for item in flow) == 36 rotary = plan_flow_key( "diffusion_transformer.transformer.rotary_pos_emb.inv_freq", authority ) assert rotary.status == "unmapped_known" audio_keys = _qwen_audio_keys() assert len(audio_keys) == 47 qwen_keys = [*audio_keys, *[f"model.synthetic_lm.{index}.weight" for index in range(399)]] qwen = [plan_qwen_key(key, authority) for key in qwen_keys] assert len(qwen) == 446 assert all(item.status == "mapped" for item in qwen) assert sum(len(item.targets) for item in qwen) == 446 assert sum(item.targets[0].component == "rvq_depth_decoder" for item in qwen) == 47 assert sum(item.targets[0].component == "language_model" for item in qwen) == 399 def test_qkv_plan_records_ordered_split_and_exact_converter_evidence(): plan = plan_flow_key( "diffusion_transformer.transformer.layers.17.self_attn.to_qkv.weight", converter_authority_for_tests(), ) assert plan.transform == "chunk(dim=0,parts=3)" assert [item.chunk_index for item in plan.targets] == [0, 1, 2] assert [item.key for item in plan.targets] == [ "transformer_blocks.17.attn.to_q.weight", "transformer_blocks.17.attn.to_k.weight", "transformer_blocks.17.attn.to_v.weight", ] assert plan.evidence[0].file_sha256 == CONVERTER_SHA256 assert (plan.evidence[0].line_start, plan.evidence[0].line_end) == (60, 78) def test_mapping_rules_are_gated_on_pinned_converter_identity(): wrong = ConverterAuthority(Path("/tmp/diffusers"), "0" * 40, "0" * 64) with pytest.raises(RuntimeError, match="pinned converter"): plan_dav_key("dec_in_proj.weight", wrong) def test_prefix_similarity_is_not_consumption_evidence(): authority = converter_authority_for_tests() assert plan_dav_key("decoder.looks_similar.weight", authority).status == "unknown" assert plan_flow_key("diffusion_transformer.transformer.layers.0.fake.weight", authority).status == "unknown" assert plan_qwen_key("model.audio_decoder.layers.0.fake.weight", authority).status == "unknown" def test_coverage_equations_and_duplicate_targets_are_fail_closed(): with pytest.raises(ValidationError, match="raw coverage equation"): CoverageSummary( source_kind="dav", raw_total=2, raw_numel=1, raw_nbytes=4, mapped=1, unmapped_known=0, unknown=0, converted_target_total=1, covered_targets=1, orphan_targets=0, duplicate_targets=0, qkv_splits=0, ) with pytest.raises(ValidationError, match="duplicate target"): CoverageSummary( source_kind="dav", raw_total=1, raw_numel=1, raw_nbytes=4, mapped=1, unmapped_known=0, unknown=0, converted_target_total=1, covered_targets=1, orphan_targets=0, duplicate_targets=1, qkv_splits=0, ) def _write_index(root: Path, weight_map: dict[str, str], total_size: int) -> None: (root / "model.safetensors.index.json").write_text( json.dumps({"metadata": {"total_size": total_size}, "weight_map": weight_map}), encoding="utf-8", ) def test_safetensors_index_validates_every_owner_and_records_size_contradiction(tmp_path): save_file({"a": torch.arange(4, dtype=torch.float32)}, tmp_path / "model-00001-of-00002.safetensors") save_file({"b": torch.arange(3, dtype=torch.bfloat16)}, tmp_path / "model-00002-of-00002.safetensors") _write_index( tmp_path, { "a": "model-00001-of-00002.safetensors", "b": "model-00002-of-00002.safetensors", }, total_size=999, ) audit, metadata = audit_safetensors_index( tmp_path, "model.safetensors.index.json", name="raw_qwen" ) assert audit.ownership_valid is True assert audit.total_size_matches is False assert audit.declared_filename_shard_count == 2 assert audit.actual_shard_count == audit.indexed_shard_count == 2 assert metadata["a"].nbytes == 16 assert metadata["b"].nbytes == 6 def test_safetensors_index_rejects_wrong_owner_or_orphan_shard(tmp_path): save_file({"a": torch.zeros(1)}, tmp_path / "model-00001-of-00002.safetensors") save_file({"b": torch.zeros(1)}, tmp_path / "model-00002-of-00002.safetensors") _write_index( tmp_path, {"a": "model-00002-of-00002.safetensors", "b": "model-00002-of-00002.safetensors"}, total_size=8, ) audit, _ = audit_safetensors_index( tmp_path, "model.safetensors.index.json", name="adversarial" ) assert audit.ownership_valid is False assert audit.orphan_shards == ("model-00001-of-00002.safetensors",) assert "a" in audit.missing_keys and "a" in audit.orphan_keys def test_strict_index_json_rejects_duplicate_and_traversal_paths(tmp_path): with pytest.raises(DuplicateJSONKeyError): strict_json_loads(b'{"weight_map":{"a":"x","a":"y"},"metadata":{}}', label="index") (tmp_path / "model.safetensors.index.json").write_text( '{"metadata":{"total_size":0},"weight_map":{"a":"../escape.safetensors"}}', encoding="utf-8", ) with pytest.raises(ValueError, match="normalized relative"): audit_safetensors_index(tmp_path, "model.safetensors.index.json", name="bad") def test_safe_pth_load_never_falls_back_to_unsafe_pickle(tmp_path): marker = tmp_path / "executed" class Evil: def __reduce__(self): return (marker.write_text, ("unsafe",)) path = tmp_path / "malicious.pth" torch.save({"payload": Evil()}, path) with pytest.raises(ValueError, match="safely load"): safe_load_pth(path) assert not marker.exists() def test_pytorch_zip_prefix_is_not_a_state_key_or_module(tmp_path): path = tmp_path / "62000_generator.pth" torch.save({"encoder.weight": torch.ones(2)}, path) members = archive_members(path) assert members assert all(item.startswith("62000_generator/") for item in members) checkpoint = safe_load_pth(path) state_path, state = select_tensor_mapping(checkpoint) assert state_path == "$" assert tuple(state) == ("encoder.weight",) assert not any("generator" in key for key in state) def test_nested_checkpoint_search_includes_wrappers_keys_and_scalar_values(): checkpoint = { "wrapper": { "encoder_config": {"analysis_mode": "latent"}, "weight": torch.ones(1), } } values = set(nested_search_values(checkpoint)) assert ("$.wrapper.encoder_config", "$.wrapper.encoder_config") in values assert ("$.wrapper.encoder_config.analysis_mode", "latent") in values assert not any("tensor(" in value for _path, value in values) def test_mapped_tensor_records_exact_mismatch_without_reclassifying_mapping(tmp_path): save_file( {"dec_in_proj.weight": torch.ones(2, dtype=torch.float32)}, tmp_path / "diffusion_pytorch_model.safetensors", ) target = TensorCollection.single("vocoder", tmp_path) row = _mapping_row( source_kind="dav", source_file="dav.pth", source_shard=None, key="dec_in_proj.weight", tensor=torch.zeros(2, dtype=torch.float32), plan=plan_dav_key( "dec_in_proj.weight", converter_authority_for_tests() ), targets={"vocoder": target}, ) assert row.status == "mapped" assert row.targets[0].equality == "mismatch" compared = row.targets[0] assert compared.source_transformed_metadata == compared.metadata assert compared.source_transformed_sha256 != compared.target_sha256 tampered = compared.model_dump(mode="json") tampered["equality"] = "exact" with pytest.raises(ValidationError, match="content hashes"): type(compared).model_validate(tampered) exact = _mapping_row( source_kind="dav", source_file="dav.pth", source_shard=None, key="dec_in_proj.weight", tensor=torch.ones(2, dtype=torch.float32), plan=plan_dav_key( "dec_in_proj.weight", converter_authority_for_tests() ), targets={"vocoder": target}, ).targets[0] assert exact.equality == "exact" assert exact.source_transformed_sha256 == exact.target_sha256 from music3lab.checkpoint_audit_render import publish_audit_outputs from music3lab.checkpoint_audit_schema import ( ArchitectureAudit, CapabilityAssessment, ) from music3lab.manifests import atomic_write_bytes def _failed_audit() -> ArchitectureAudit: coverage = tuple( CoverageSummary( source_kind=kind, raw_total=0, raw_numel=0, raw_nbytes=0, mapped=0, unmapped_known=0, unknown=0, converted_target_total=0, covered_targets=0, orphan_targets=0, duplicate_targets=0, qkv_splits=0, ) for kind in ("dav", "flow", "qwen", "all") ) return ArchitectureAudit.create( audit_status="FAIL", generated_by_commit="0" * 40, project_source_sha256="1" * 64, project_git_dirty=True, base_id="2" * 64, identities=(), mappings=(), coverage=coverage, indexes=(), searches=(), contradictions=(), capabilities=CapabilityAssessment( native_wav_to_rvq="BLOCKED", dav_continuous_analysis="CANDIDATE", native_reason="missing executable encoder", dav_reason="candidate weights only", limitations=("synthetic fixture",), ), archive_observations=(), ) def test_audit_publication_is_canonical_and_markdown_is_report_derived(tmp_path): audit = _failed_audit() json_path = tmp_path / "audit.json" markdown_path = tmp_path / "audit.md" architecture_path = tmp_path / "ARCHITECTURE_AUDIT.md" architecture_json_path = tmp_path / "ARCHITECTURE_AUDIT.json" hashes = publish_audit_outputs( audit, json_path=json_path, markdown_path=markdown_path, architecture_path=architecture_path, architecture_json_path=architecture_json_path, ) assert ArchitectureAudit.model_validate_json(json_path.read_bytes()) == audit assert json_path.read_bytes() == architecture_json_path.read_bytes() assert markdown_path.read_bytes() == architecture_path.read_bytes() assert audit.semantic_digest.encode() in markdown_path.read_bytes() assert set(hashes) == { "json", "architecture_json", "markdown", "architecture" } def test_audit_publication_rolls_back_every_written_output_on_failure(tmp_path): audit = _failed_audit() paths = [ tmp_path / "audit.json", tmp_path / "audit.md", tmp_path / "ARCHITECTURE_AUDIT.md", tmp_path / "ARCHITECTURE_AUDIT.json", ] for index, path in enumerate(paths): path.write_bytes(f"prior-{index}".encode()) priors = {path: path.read_bytes() for path in paths} calls = 0 def failing_writer(path, data, *, mode): nonlocal calls calls += 1 if calls == 2: raise OSError("injected publication failure") atomic_write_bytes(path, data, mode=mode) with pytest.raises(OSError, match="injected"): publish_audit_outputs( audit, json_path=paths[0], markdown_path=paths[1], architecture_path=paths[2], writer=failing_writer, architecture_json_path=paths[3], ) assert {path: path.read_bytes() for path in paths} == priors def test_checkpoint_audit_cli_refuses_visible_cuda(monkeypatch, capsys): from music3lab.checkpoint_audit_cli import main monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "0") assert main([]) == 2 payload = json.loads(capsys.readouterr().out) assert payload["status"] == "ERROR" assert "CUDA_VISIBLE_DEVICES=-1" in payload["error"] def test_stage_specific_sample_rate_evidence_is_pinned() -> None: from music3lab.checkpoint_audit_runner import ( SGLANG_FILES, _sglang_evidence, ) constants = "sglang_omni/models/minimax_music3/constants.py" acoustic = "sglang_omni/models/minimax_music3/acoustic.py" assert SGLANG_FILES[constants] == ( "325fcbb1c59eefff7e7957e2ffd2f6c0d84b3930d945436e10b248fc6cd2a62f" ) assert SGLANG_FILES[acoustic] == ( "98649dd669564b2e73fa67262fa571a20ce808bb55aa10b370fa3083dd3a8b6c" ) rates = _sglang_evidence(constants, 18, 19, "stage-specific rates") resample = _sglang_evidence(acoustic, 55, 58, "explicit resample") assert (rates.line_start, rates.line_end) == (18, 19) assert (resample.line_start, resample.line_end) == (55, 58) assert rates.file_sha256 == SGLANG_FILES[constants] assert resample.file_sha256 == SGLANG_FILES[acoustic]