#!/usr/bin/env python3 """Fail-closed structural acceptance for the all-SQG GLM-5.2 checkpoint.""" from __future__ import annotations import argparse import hashlib import json import os from pathlib import Path from typing import Any FIRST_ROUTED = 3 LAST_MTP = 78 EXPERTS = 256 PROJECTIONS = ("gate_proj", "up_proj", "down_proj") NONROUTED = ( "mlp.shared_experts.gate_proj", "mlp.shared_experts.up_proj", "mlp.shared_experts.down_proj", "self_attn.q_b_proj", "self_attn.o_proj", ) ROUTED_MATRICES = 76 * 256 * 3 NONROUTED_K6_MATRICES = 76 * 5 TOTAL_SQG_MATRICES = ROUTED_MATRICES + NONROUTED_K6_MATRICES DIRECT_BF16_MATRICES = 837 OFFICIAL_SOURCE_TENSORS = 59585 SCHEMA = "glm52-full-sqg-native-checkpoint-acceptance-v1" def load(path: Path) -> dict[str, Any]: value = json.loads(path.read_text(encoding="utf-8")) if not isinstance(value, dict): raise TypeError(f"JSON root must be an object: {path}") return value def canonical_sha256(value: Any) -> str: return hashlib.sha256( json.dumps(value, sort_keys=True, separators=(",", ":")).encode() ).hexdigest() def replacement_weights() -> tuple[set[str], set[str]]: routed = { f"model.layers.{layer}.mlp.experts.{expert}.{projection}.weight" for layer in range(FIRST_ROUTED, LAST_MTP + 1) for expert in range(EXPERTS) for projection in PROJECTIONS } nonrouted = { f"model.layers.{layer}.{suffix}.weight" for layer in range(FIRST_ROUTED, LAST_MTP + 1) for suffix in NONROUTED } return routed, nonrouted def validate_sidecars(model: Path, routed_weights: set[str]) -> None: for layer in range(FIRST_ROUTED, LAST_MTP + 1): path = model / f"r7-experts-layer-{layer:03d}.json" value = load(path) bit_map = value.get("bit_map") if not isinstance(bit_map, dict): raise ValueError(f"layer {layer}: SQG bit map is absent") expected = { name.removesuffix(".weight") for name in routed_weights if name.startswith(f"model.layers.{layer}.mlp.experts.") } values = tuple(int(bits) for bits in bit_map.values()) if ( set(bit_map) != expected or (values.count(3), values.count(4), len(values)) != (384, 384, 768) or value.get("codebook") not in {"SQG", "sqg_xor_cheb_t12"} or value.get("allow_a16_fallback") is not False ): raise ValueError(f"layer {layer}: native SQG K3/K4 sidecar differs") def validate(args: argparse.Namespace) -> dict[str, Any]: model = args.model.resolve() source = load(args.bf16_index.resolve()) final_index = load(model / "model.safetensors.index.json") source_map = source.get("weight_map") final_map = final_index.get("weight_map") if not isinstance(source_map, dict) or not isinstance(final_map, dict): raise ValueError("source/final SafeTensors index is malformed") if len(source_map) != OFFICIAL_SOURCE_TENSORS: raise ValueError("official BF16 source tensor census differs from 59,585") routed, nonrouted = replacement_weights() replaced = routed | nonrouted if len(routed) != ROUTED_MATRICES or len(nonrouted) != NONROUTED_K6_MATRICES: raise AssertionError("programmed all-SQG census differs") if not replaced <= set(source_map): raise ValueError("official BF16 source lacks a required SQG matrix") ignored = set(source_map) - replaced if len(ignored) != DIRECT_BF16_MATRICES: raise ValueError("direct-BF16 tensor census differs from 837") forbidden = sorted(name for name in final_map if name.endswith(".mcg")) if forbidden: raise ValueError(f"final checkpoint contains {len(forbidden)} MCG marker tensors") if replaced & set(final_map): raise ValueError("final index retains replaced BF16 matrix weights") missing_ignored = ignored - set(final_map) if missing_ignored: raise ValueError(f"final index lost {len(missing_ignored)} direct BF16 ignored tensors") prefixes = {name.removesuffix(".weight") for name in replaced} sqg = {name.removesuffix(".sqg") for name in final_map if name.endswith(".sqg")} trellis = { name.removesuffix(".trellis") for name in final_map if name.endswith(".trellis") } if sqg != prefixes or not prefixes <= trellis: raise ValueError("final SQG marker/trellis matrix census differs") if len(sqg) != TOTAL_SQG_MATRICES: raise ValueError("final SQG matrix total differs") allowed_quantized = { name for name in final_map if name.endswith((".trellis", ".suh", ".svh", ".sqg")) } shared_h = { f"model.layers.{layer}.mlp.experts.r7_shared.{field}" for layer in range(FIRST_ROUTED, LAST_MTP + 1) for field in ("gate_up_suh", "down_svh") } if not shared_h <= set(final_map): raise ValueError("final checkpoint lacks topology-neutral shared-H tensors") allowed_quantized.update(shared_h) unexpected = set(final_map) - ignored - allowed_quantized if unexpected: raise ValueError( f"final checkpoint contains {len(unexpected)} non-SQG, non-ignored tensors" ) validate_sidecars(model, routed) k6 = load(model / "SQG_K6_NONROUTED.json") entries = k6.get("matrices") if ( k6.get("complete") is not True or not isinstance(entries, dict) or set(entries) != {name.removesuffix(".weight") for name in nonrouted} or any(int(item.get("bits", -1)) != 6 for item in entries.values()) ): raise ValueError("380-matrix native SQG K6 manifest differs") hessian = load(args.hessian_dataset_receipt.resolve()) commit = hessian.get("commit") revision = hessian.get("revision") if ( hessian.get("complete") is not True or hessian.get("public") is not True or hessian.get("repo") != "brandonmusic/GLM-5.2-BMM-Law-SQG-Hessians" or hessian.get("phase") != "final_publication_verified" or hessian.get("main_verified") is not True or hessian.get("accepted_tag_verified") is not True or revision != "main" or not isinstance(commit, str) or len(commit) != 40 ): raise ValueError("final public Hessian dataset binding is absent") manifest = load(model / "FULL_SQG_NATIVE_MANIFEST.json") if ( manifest.get("complete") is not True or manifest.get("official_bf16_revision") != args.bf16_revision or manifest.get("hessian_dataset", {}).get("repo") != hessian["repo"] or manifest.get("hessian_dataset", {}).get("revision") != revision or manifest.get("hessian_dataset", {}).get("commit") != commit or manifest.get("census") != { "routed_k3_k4_matrices": ROUTED_MATRICES, "nonrouted_k6_matrices": NONROUTED_K6_MATRICES, "total_sqg_matrices": TOTAL_SQG_MATRICES, "mcg_marker_tensors": 0, "retained_replaced_bf16_weights": 0, "direct_bf16_ignored_tensors": len(ignored), "official_source_tensor_total": OFFICIAL_SOURCE_TENSORS, } ): raise ValueError("full-SQG native materialization manifest differs") construction = manifest.get("construction", {}) if construction and ( construction.get("checkpoint_tensor_layout") != "whole_logical_tensor" or construction.get("tensor_parallel_slicing") is not False or construction.get("acceptance_runtime") != "tp1_pp8" ): raise ValueError("final checkpoint is not bound to whole-tensor TP1/PP8 loading") result: dict[str, Any] = { "schema": SCHEMA, "complete": True, "passed": True, "model": str(model), "official_bf16_revision": args.bf16_revision, "hessian_dataset_commit": commit, "hessian_dataset_revision": revision, "census": manifest["census"], "ignored_tensor_domain_sha256": canonical_sha256(sorted(ignored)), "sqg_matrix_domain_sha256": canonical_sha256(sorted(prefixes)), } result["acceptance_id"] = canonical_sha256(result) return result def main() -> int: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model", type=Path, required=True) parser.add_argument("--bf16-index", type=Path, required=True) parser.add_argument("--bf16-revision", required=True) parser.add_argument("--hessian-dataset-receipt", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) args = parser.parse_args() result = validate(args) args.output.parent.mkdir(parents=True, exist_ok=True) temporary = args.output.with_name(f".{args.output.name}.tmp-{os.getpid()}") with temporary.open("x", encoding="utf-8") as handle: json.dump(result, handle, indent=2, sort_keys=True) handle.write("\n") handle.flush() os.fsync(handle.fileno()) os.replace(temporary, args.output) print(json.dumps(result, sort_keys=True)) return 0 if __name__ == "__main__": raise SystemExit(main())