ArGrigorov's picture
Upload folder using huggingface_hub
e9c8366 verified
Raw
History Blame Contribute Delete
8.56 kB
"""presets — format string → Quantizer kwargs.
Each preset maps a short format name to the kwargs that configure the unified
Quantizer for that specific quantization format.
Usage:
q = Quantizer(**FORMAT_PRESETS["nvfp4"])
# or via dispatch:
quantize_model(model, format="nvfp4")
"""
from __future__ import annotations
FORMAT_PRESETS: dict[str, dict] = {
# ---- Basic integer ----
"int8": {
"value_bits": 8, "value_repr": "int", "scale_mode": "per-channel",
},
"int4": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_size": 64,
},
"int2": {
"value_bits": 2, "value_repr": "int", "scale_mode": "per-group",
"group_size": 64,
},
"w4a16": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_size": 128, "quantizes_input": False,
},
"w4a4": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_size": 64, "quantizes_input": True,
},
"w8a8": {
"value_bits": 8, "value_repr": "int", "scale_mode": "per-channel",
"quantizes_input": True,
},
# ---- LUT / NF4 ----
"nf4": {
"value_bits": 4, "value_repr": "nf4_lut", "scale_mode": "per-group",
"group_size": 64, "double_quant": True, "block_size": 256,
},
"nf4_dq": {
"value_bits": 4, "value_repr": "nf4_lut", "scale_mode": "per-group",
"group_size": 64, "double_quant": False,
},
# ---- FP4 / NVFP4 / MXFP4 ----
"nvfp4": {
"value_bits": 4, "value_repr": "fp4_e2m1", "scale_mode": "per-group",
"scale_dtype": "fp8_e4m3", "group_size": 16, "quantizes_input": True,
},
"nvfp4_wo": { # weight-only NVFP4
"value_bits": 4, "value_repr": "fp4_e2m1", "scale_mode": "per-group",
"scale_dtype": "fp8_e4m3", "group_size": 16, "quantizes_input": False,
},
"mxfp4": {
"value_bits": 4, "value_repr": "fp4_e2m1", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
# ---- FP6 / MXFP6 / NVFP6 ----
"mxfp6_e3m2": {
"value_bits": 6, "value_repr": "fp6_e3m2", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"mxfp6_e2m3": {
"value_bits": 6, "value_repr": "fp6_e2m3", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"nvfp6_e3m2": {
"value_bits": 6, "value_repr": "fp6_e3m2", "scale_mode": "per-group",
"scale_dtype": "fp8_e4m3", "group_size": 16,
},
"nvfp6_e2m3": {
"value_bits": 6, "value_repr": "fp6_e2m3", "scale_mode": "per-group",
"scale_dtype": "fp8_e4m3", "group_size": 16,
},
# ---- FP8 / MXFP8 / NVFP8 ----
"fp8": {
"value_bits": 8, "value_repr": "fp8_e4m3", "scale_mode": "per-channel",
},
"fp8_e5m2": {
"value_bits": 8, "value_repr": "fp8_e5m2", "scale_mode": "per-channel",
},
"fp8_w8a8": {
"value_bits": 8, "value_repr": "fp8_e4m3", "scale_mode": "per-channel",
"quantizes_input": True,
},
"mxfp8_e4m3": {
"value_bits": 8, "value_repr": "fp8_e4m3", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"mxfp8_e5m2": {
"value_bits": 8, "value_repr": "fp8_e5m2", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"nvfp8_e4m3": {
"value_bits": 8, "value_repr": "fp8_e4m3", "scale_mode": "per-group",
"scale_dtype": "fp8_e4m3", "group_size": 16,
},
"nvfp8_e5m2": {
"value_bits": 8, "value_repr": "fp8_e5m2", "scale_mode": "per-group",
"scale_dtype": "fp8_e4m3", "group_size": 16,
},
# ---- MXINT (INT + E8M0 block scale) ----
"mxint2": {
"value_bits": 2, "value_repr": "int", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"mxint4": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"mxint6": {
"value_bits": 6, "value_repr": "int", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
"mxint8": {
"value_bits": 8, "value_repr": "int", "scale_mode": "per-group",
"scale_dtype": "e8m0", "group_size": 32,
},
# ---- INT6 (plain integer, no block scale) ----
"int6": {
"value_bits": 6, "value_repr": "int", "scale_mode": "per-group",
"group_size": 64,
},
# ---- Ternary / Binary ----
"ternary": {
"value_bits": 2, "value_repr": "ternary", "scale_mode": "per-channel",
},
"binary": {
"value_bits": 1, "value_repr": "binary", "scale_mode": "per-channel",
},
# ---- GGUF k-quants ----
"q8_0": {
"value_bits": 8, "value_repr": "int", "scale_mode": "per-group",
"group_size": 32,
},
"q4_0": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_size": 32,
},
"q4_k": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_mode": "super-block-nested", "group_size": 256,
},
"q5_k": {
"value_bits": 5, "value_repr": "int", "scale_mode": "per-group",
"group_mode": "super-block-nested", "group_size": 256,
},
"q6_k": {
"value_bits": 6, "value_repr": "int", "scale_mode": "per-group",
"group_mode": "super-block-nested", "group_size": 256,
},
"q2_k": {
"value_bits": 2, "value_repr": "int", "scale_mode": "per-group",
"group_mode": "super-block-nested", "group_size": 256,
},
"q3_k": {
"value_bits": 3, "value_repr": "int", "scale_mode": "per-group",
"group_mode": "super-block-nested", "group_size": 256,
},
# ---- Advanced PTQ ----
"gptq": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_size": 0, "error_compensation": "gptq-hessian",
},
"awq": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_size": 128, "activation_aware": "awq",
},
"smoothquant": {
"value_bits": 8, "value_repr": "int", "scale_mode": "per-channel",
"activation_aware": "smoothquant", "alpha": 0.5, "quantizes_input": True,
},
"llmint8": {
"value_bits": 8, "value_repr": "int", "scale_mode": "per-channel",
"outlier_threshold": 6.0,
},
# ---- Codebook / VQ ----
"codebook": {
"value_repr": "codebook", "codebook_source": "kmeans",
"codebook_size": 16, "scale_mode": "per-channel",
},
"vq": {
"value_repr": "codebook", "codebook_source": "vq",
"codebook_size": 256, "vq_group_size": 2, "scale_mode": "per-channel",
},
"aqlm": {
"value_repr": "codebook", "codebook_source": "vq",
"codebook_size": 256, "vq_group_size": 2, "scale_mode": "per-channel",
},
# ---- QuIP ----
"quip": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-channel",
"rotation": "random",
},
"quip_hadamard": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-channel",
"rotation": "hadamard",
},
# ---- Outlier / Prune ----
"outlier": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-channel",
"outlier_threshold": None, # auto 3σ
},
"outlier2": {
"value_bits": 2, "value_repr": "int", "scale_mode": "per-channel",
"outlier_threshold": None,
},
"prune": {
"prune_mode": "ratio", "prune_ratio": 0.5,
},
"prune_structured": {
"prune_mode": "structured", "prune_ratio": 0.25,
},
# ---- QAT ----
"lsq": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-channel",
"learnable": True,
},
# ---- Clustered (magnitude-binned) ----
"clustered": {
"value_bits": 4, "value_repr": "int", "scale_mode": "per-group",
"group_mode": "magnitude-binned", "group_size": 64, "block_size": 4,
},
# ---- Passthrough ----
"none": {
"value_repr": "none",
},
"fp16": {
"value_repr": "none",
},
}
def get_preset(format: str, **overrides) -> dict:
"""Get preset kwargs, with optional overrides."""
if format not in FORMAT_PRESETS:
raise ValueError(f"Unknown format {format!r}. Available: {sorted(FORMAT_PRESETS)}")
kwargs = dict(FORMAT_PRESETS[format])
kwargs.update(overrides)
return kwargs