File size: 4,823 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
"""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