"""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