#!/usr/bin/env python3 """Audit every Mellum INT4 projection and scale on CPU without loading the model.""" from __future__ import annotations import argparse import hashlib import json import math import struct from pathlib import Path SOURCE_REVISION = "92ddae9fc7665e9f801d141d2e5a6b2caf2460c4" def tensor_files(directory: Path) -> dict: """Read tensor headers only and check that every indexed shard is complete.""" index_path = directory / "model.safetensors.index.json" index = json.loads(index_path.read_text()) if index_path.exists() else None filenames = sorted(set(index["weight_map"].values())) if index else ["model.safetensors"] result = {} for filename in filenames: path = directory / filename with path.open("rb") as stream: raw = stream.read(8) if len(raw) != 8: raise ValueError(f"Truncated safetensors file: {filename}") header_size = struct.unpack(" 32 * 1024 * 1024: raise ValueError(f"Unexpectedly large safetensors header: {filename}") header = json.loads(stream.read(header_size)) start = 8 + header_size for name, spec in header.items(): if name == "__metadata__": continue begin, end = spec["data_offsets"] if begin < 0 or end < begin or start + end > path.stat().st_size: raise ValueError(f"Truncated/invalid tensor payload: {name}") if name in result: raise ValueError(f"Duplicate tensor name: {name}") if index is not None and index["weight_map"].get(name) != filename: raise ValueError(f"Index/shard mismatch: {name}") result[name] = (path, spec, start) if index is not None and set(result) != set(index["weight_map"]): raise ValueError("Shard headers do not match the complete tensor index") return result def require_tensor(tensors, name, dtype, shape): if name not in tensors: raise ValueError(f"Required tensor missing: {name}") _, spec, _ = tensors[name] if spec["dtype"] not in dtype or spec["shape"] != shape: raise ValueError(f"Unexpected tensor dtype/shape: {name}: {spec}") widths = {"BF16": 2, "F16": 2, "F32": 4, "I32": 4} if spec["data_offsets"][1] - spec["data_offsets"][0] != math.prod(shape) * widths[spec["dtype"]]: raise ValueError(f"Tensor byte size disagrees with shape: {name}") def positive_scale_count(tensor) -> int: import numpy as np path, spec, start = tensor begin, end = spec["data_offsets"] count = 0 dtype = {"BF16": " 0)): raise ValueError(f"Nonpositive or nonfinite group scales in {path.name}") count += values.size remaining -= len(raw) return count def tensor_sha256(tensor) -> str: path, spec, start = tensor begin, end = spec["data_offsets"] digest = hashlib.sha256() with path.open("rb") as stream: stream.seek(start + begin) remaining = end - begin while remaining: data = stream.read(min(8 * 1024 * 1024, remaining)) if not data: raise ValueError("Truncated tensor during hashing") digest.update(data) remaining -= len(data) return digest.hexdigest() def audit(directory: Path, source: Path | None = None) -> dict: config = json.loads((directory / "config.json").read_text()) expected_config = {"model_type": "mellum", "num_hidden_layers": 28, "hidden_size": 2304, "moe_intermediate_size": 896, "num_experts": 64, "num_experts_per_tok": 8, "num_attention_heads": 32, "num_key_value_heads": 4, "head_dim": 128, "sliding_window": 1024, "max_position_embeddings": 131072} if any(config.get(k) != v for k, v in expected_config.items()): raise ValueError("Export architecture differs from pinned Mellum2.1") if config.get("layer_types") != ["full_attention" if i % 4 == 3 else "sliding_attention" for i in range(28)]: raise ValueError("Export lost the Mellum mixed-attention layout") quant = config.get("quantization_config", {}) if quant.get("quant_method") != "compressed-tensors" or quant.get("format") != "pack-quantized": raise ValueError("Require compressed-tensors pack-quantized W4A16 export") groups = quant.get("config_groups", {}) if len(groups) != 1: raise ValueError("Expected one explicit symmetric group-32 quantization scheme") group = next(iter(groups.values())) weights = group.get("weights", {}) if any(weights.get(k) != v for k, v in {"num_bits": 4, "type": "int", "symmetric": True, "group_size": 32, "strategy": "group", "dynamic": False}.items()): raise ValueError("Export is not symmetric static group-32 INT4") if group.get("input_activations") is not None or group.get("output_activations") is not None: raise ValueError("Explicit activation metadata bypasses the required MoE WNA16 adapter") tensors = tensor_files(directory) if any("mtp" in name.lower() for name in tensors): raise ValueError("Unexpected MTP tensors; this source checkpoint has no native head") projections = {} for layer in range(28): for expert in range(64): prefix = f"model.layers.{layer}.mlp.experts.{expert}" for name, shape in (("gate_proj", (896, 2304)), ("up_proj", (896, 2304)), ("down_proj", (2304, 896))): projections[f"{prefix}.{name}"] = shape for name, shape in (("q_proj", (4096, 2304)), ("k_proj", (512, 2304)), ("v_proj", (512, 2304)), ("o_proj", (2304, 4096))): projections[f"model.layers.{layer}.self_attn.{name}"] = shape require_tensor(tensors, f"model.layers.{layer}.mlp.gate.weight", ["BF16"], [64, 2304]) for norm in ("input_layernorm", "post_attention_layernorm"): require_tensor(tensors, f"model.layers.{layer}.{norm}.weight", ["BF16"], [2304]) for norm in ("q_norm", "k_norm"): require_tensor(tensors, f"model.layers.{layer}.self_attn.{norm}.weight", ["BF16"], [128]) expected_packed = {name + ".weight_packed" for name in projections} actual_packed = {name for name in tensors if name.endswith(".weight_packed")} if expected_packed != actual_packed: raise ValueError(f"Packed projection coverage mismatch: missing={len(expected_packed-actual_packed)}, " f"unexpected={len(actual_packed-expected_packed)}") scales = 0 for name, (out_dim, in_dim) in projections.items(): require_tensor(tensors, name + ".weight_packed", ["I32"], [out_dim, in_dim // 8]) require_tensor(tensors, name + ".weight_scale", ["BF16", "F16", "F32"], [out_dim, in_dim // 32]) if name + ".weight" in tensors: raise ValueError(f"Unexpected duplicate full-precision quantized projection: {name}") scales += positive_scale_count(tensors[name + ".weight_scale"]) require_tensor(tensors, "model.norm.weight", ["BF16"], [2304]) heads = ("lm_head.weight", "model.embed_tokens.weight") for name in heads: require_tensor(tensors, name, ["BF16"], [98304, 2304]) hashes = {name: tensor_sha256(tensors[name]) for name in heads} if source is not None: original = tensor_files(source) for name in heads: if hashes[name] != tensor_sha256(original[name]): raise ValueError(f"Preserved BF16 vocabulary tensor changed: {name}") for filename in ("chat_template.jinja", "generation_config.json"): if (source / filename).exists() and (directory / filename).read_bytes() != (source / filename).read_bytes(): raise ValueError(f"Export changed the source {filename}") return {"status": "passed", "scope": "export structure/scales only; serving quality unverified", "source_revision": SOURCE_REVISION, "expert_projections": 5376, "attention_projections": 112, "positive_finite_scales": scales, "bf16_router_layers": 28, "mtp_tensors": 0, "head_sha256": hashes, "heads_compared_with_source": source is not None, "payload_bytes": sum(v[1]["data_offsets"][1] - v[1]["data_offsets"][0] for v in tensors.values())} def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("output_dir", type=Path) parser.add_argument("--source-dir", type=Path) parser.add_argument("--report", type=Path) args = parser.parse_args() report = audit(args.output_dir, args.source_dir) rendered = json.dumps(report, indent=2) + "\n" if args.report: args.report.write_text(rendered) print(rendered, end="") if __name__ == "__main__": main()