JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
4.82 kB
"""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