"""Versioned, isolated LoopQ export artifact schema.""" from __future__ import annotations from collections.abc import Mapping from pathlib import Path from typing import Any import torch from .numerical_contract import numerical_contract, validate_numerical_contract FORMAT = "loopq_calibrated_model" FORMAT_VERSION = 2 def build_export_artifact( *, model: Mapping[str, Any], mapping_validation: Mapping[str, Any], activation_bits: int, las_state: Mapping[str, Any], shared_transform_states: Mapping[str, Any], selected_transform_states: Mapping[str, Any], cta_state: Mapping[str, Any], calibration: Mapping[str, Any], ) -> dict[str, Any]: if activation_bits not in (4, 8): raise ValueError("LoopQ exports support W4A4 or W4A8 only") artifact = { "format": FORMAT, "format_version": FORMAT_VERSION, "method": "LoopQ", "model": dict(model), "mapping_validation": dict(mapping_validation), "quantization": { "weight_bits": 4, "activation_bits": activation_bits, "format": "symmetric_uniform_rtn", "group_size": 32, "numerical_contract": numerical_contract( activation_bits, dynamic_las=bool(las_state.get("dynamic_clip", False)) ), }, "components": { "las": dict(las_state), "shared_transforms": dict(shared_transform_states), "selected_loop_transforms": dict(selected_transform_states), "cta": dict(cta_state), }, "calibration": dict(calibration), } validate_export_artifact(artifact) return artifact def validate_export_artifact(artifact: Mapping[str, Any]) -> None: required = { "format", "format_version", "method", "model", "mapping_validation", "quantization", "components", "calibration", } missing = required.difference(artifact) if missing: raise ValueError(f"LoopQ export is missing fields: {sorted(missing)}") if artifact["format"] != FORMAT or artifact["format_version"] != FORMAT_VERSION: raise ValueError("unsupported LoopQ export format/version; legacy artifacts require their original runtime") if artifact["method"] != "LoopQ": raise ValueError("artifact method must be LoopQ") quantization = artifact["quantization"] expected = {"weight_bits": 4, "format": "symmetric_uniform_rtn", "group_size": 32} if any(quantization.get(key) != value for key, value in expected.items()): raise ValueError("artifact violates the LoopQ symmetric group-32 RTN contract") if quantization.get("activation_bits") not in (4, 8): raise ValueError("artifact activation precision must be 4 or 8 bits") calibration = artifact["calibration"] ablation = calibration.get("ablation") if ablation not in {None, "no_las", "no_slt", "no_cta"}: raise ValueError("unknown LoopQ component ablation") components = artifact["components"] validate_numerical_contract( quantization.get("numerical_contract"), quantization["activation_bits"], dynamic_las=bool(components["las"].get("dynamic_clip", False)), ) recorded = calibration.get("numerical_contract") if recorded is not None and recorded != quantization["numerical_contract"]: raise ValueError("calibration and export numerical contracts differ") if bool(components["las"].get("shared_across_loops", False)) != (ablation == "no_las"): raise ValueError("LAS sharing mode does not match the declared ablation") if bool(components["cta"].get("enabled", True)) != (ablation != "no_cta"): raise ValueError("CTA mode does not match the declared ablation") if ablation == "no_slt" and components["selected_loop_transforms"]: raise ValueError("no-SLT artifact must not contain selected transforms") calibrated_bits = calibration.get("provenance", {}).get("arguments", {}).get("activation_bits") if calibrated_bits is not None and calibrated_bits != quantization["activation_bits"]: raise ValueError("export activation bits differ from calibration precision") serialized_text = repr(artifact).lower() forbidden = ("residualquant", "momentumquant", "nvfp4", "anchor_schedule") if any(token in serialized_text for token in forbidden): raise ValueError("non-LoopQ quantization policy found in isolated artifact") def save_export_artifact(path: str | Path, artifact: Mapping[str, Any]) -> None: validate_export_artifact(artifact) torch.save(dict(artifact), Path(path)) def load_export_artifact(path: str | Path) -> dict[str, Any]: artifact = torch.load(Path(path), map_location="cpu", weights_only=True) validate_export_artifact(artifact) return artifact