Download loopq_quantization/scripts/loopq/export.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 4.82 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/export.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/export.py
-
curl -L -o export.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/export.py
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 | |