| """quantizer — ONE parameterized quantizer for all formats.
|
|
|
| All 25+ quantization formats are configurations of this single class.
|
| Parameters control value representation, scale granularity, grouping,
|
| error compensation, rotation, outlier handling, activation awareness,
|
| codebook, pruning, learnability — every aspect.
|
|
|
| Format presets (see presets.py) map format strings to kwargs:
|
| "int4" → value_bits=4, value_repr="int", scale_mode="per-group", group_size=64
|
| "nvfp4" → value_bits=4, value_repr="fp4_e2m1", scale_dtype="fp8_e4m3", group_size=16, ...
|
| "q4_k" → value_bits=4, value_repr="int", group_mode="super-block-nested", group_size=256, ...
|
| etc.
|
|
|
| learnable=True → all tensor-parameters become nn.Parameter (QAT: latent
|
| weights, scale, group boundaries, codebook, rotation, zero_point, d-scales).
|
| STE through round/sign/argmin.
|
|
|
| chunked dequant: dequantize_weight(qw, compute_dtype, slice=(start,end))
|
| reconstructs only output rows [start:end] — for QuantizedModule chunked forward.
|
| """
|
|
|
| from __future__ import annotations
|
|
|
| from typing import Any
|
|
|
| import torch
|
| import torch.nn as nn
|
|
|
| from agiws_neural_quant.base import QuantizedWeight, QuantizedActivation, _compute_dtype_to_torch
|
| from agiws_neural_quant.ternary import ternarize_tensor
|
| from agiws_neural_quant.training.ste import STEQuantize
|
| from agiws_neural_quant import kquant as _kquant
|
|
|
|
|
|
|
|
|
|
|
|
|
| from agiws_neural_quant.nf4.nf4 import NF4_LUT, quantize_nf4, dequantize_nf4, pack_nf4, unpack_nf4
|
| from agiws_neural_quant.nf4.double_quant import double_quantize_scales_2d, dequantize_scales_2d
|
| from agiws_neural_quant.fp8 import FP8_E4M3_LUT, FP8_E5M2_LUT, quantize_fp8, dequantize_fp8
|
| from agiws_neural_quant.fp4 import FP4_E2M1_LUT, quantize_fp4, dequantize_fp4, pack_fp4, unpack_fp4, E8M0_LUT
|
| from agiws_neural_quant.fp6 import FP6_E3M2_LUT, FP6_E2M3_LUT, quantize_fp6, dequantize_fp6, pack_fp6, unpack_fp6
|
|
|
|
|
| def _kquant_bits_to_fmt(bits: int) -> str | None:
|
| """Map value_bits → GGUF k-quant format name for super-block-nested."""
|
| return {2: "q2_k", 3: "q3_k", 4: "q4_k", 5: "q5_k", 6: "q6_k",
|
| 8: "q8_0", 7: "q4_0"}.get(bits)
|
|
|
|
|
|
|
|
|
|
|
|
|
| class Quantizer:
|
| """Parameterized quantizer for all formats.
|
|
|
| Every quantization format is a configuration of these parameters.
|
| """
|
|
|
| def __init__(
|
| self,
|
| value_bits: int = 8,
|
| value_repr: str = "int",
|
| scale_mode: str = "per-channel",
|
| scale_dtype: str = "fp32",
|
| group_mode: str = "contiguous",
|
| group_size: int = 0,
|
| symmetric: bool = True,
|
| double_quant: bool = False,
|
| block_size: int = 256,
|
| rotation: str = "none",
|
| outlier_threshold: float | None = None,
|
| error_compensation: str = "none",
|
| activation_aware: str = "none",
|
| alpha: float = 0.5,
|
| codebook_source: str = "kmeans",
|
| codebook_size: int = 16,
|
| vq_group_size: int = 2,
|
| prune_mode: str = "none",
|
| prune_ratio: float = 0.5,
|
| quantizes_input: bool = False,
|
| learnable: bool = False,
|
| compute_dtype: str = "fp32",
|
| activation_scale_mode: str = "per-tensor",
|
| residual_levels: int = 1,
|
| residual_codebook_size: int = 0,
|
| num_heads: int = 0,
|
| head_dim: int = 0,
|
| ):
|
| self.value_bits = value_bits
|
| self.value_repr = value_repr
|
| self.scale_mode = scale_mode
|
| self.scale_dtype = scale_dtype
|
| self.group_mode = group_mode
|
| self.group_size = group_size
|
| self.symmetric = symmetric
|
| self.double_quant = double_quant
|
| self.block_size = block_size
|
| self.rotation = rotation
|
| self.outlier_threshold = outlier_threshold
|
| self.error_compensation = error_compensation
|
| self.activation_aware = activation_aware
|
| self.alpha = alpha
|
| self.codebook_source = codebook_source
|
| self.codebook_size = codebook_size
|
| self.vq_group_size = vq_group_size
|
| self.prune_mode = prune_mode
|
| self.prune_ratio = prune_ratio
|
| self.quantizes_input = quantizes_input
|
| self.learnable = learnable
|
| self.compute_dtype = compute_dtype
|
| self.activation_scale_mode = activation_scale_mode
|
| self.residual_levels = residual_levels
|
| self.residual_codebook_size = residual_codebook_size
|
| self.num_heads = num_heads
|
| self.head_dim = head_dim
|
| self._validate()
|
|
|
| def _validate(self):
|
| """Validate parameter combinations. Raise ValueError on invalid combos."""
|
| valid_reprs = {"int", "nf4_lut", "fp4_e2m1", "fp6_e3m2", "fp6_e2m3",
|
| "fp8_e4m3", "fp8_e5m2", "codebook", "binary", "ternary",
|
| "none", "outlier", "prune"}
|
| if self.value_repr not in valid_reprs:
|
| raise ValueError(
|
| f"value_repr={self.value_repr!r} not in {sorted(valid_reprs)}"
|
| )
|
| valid_scale_modes = {"per-tensor", "per-channel", "per-group",
|
| "per-block", "per-head"}
|
| if self.scale_mode not in valid_scale_modes:
|
| raise ValueError(
|
| f"scale_mode={self.scale_mode!r} not in {sorted(valid_scale_modes)}"
|
| )
|
| valid_group_modes = {"contiguous", "super-block-nested", "magnitude-binned"}
|
| if self.group_mode not in valid_group_modes:
|
| raise ValueError(
|
| f"group_mode={self.group_mode!r} not in {sorted(valid_group_modes)}"
|
| )
|
| if self.scale_mode == "per-head" and self.head_dim <= 0:
|
| raise ValueError(
|
| "per-head scale_mode requires head_dim > 0"
|
| )
|
| if self.scale_mode in ("per-group", "per-block") and self.group_size == 0:
|
|
|
| pass
|
| if self.residual_levels < 1:
|
| raise ValueError(
|
| f"residual_levels must be >= 1, got {self.residual_levels}"
|
| )
|
| if self.residual_levels > 1 and self.value_repr != "codebook":
|
| raise ValueError(
|
| "residual_levels > 1 only supported with value_repr='codebook'"
|
| )
|
| if self.value_bits < 1 or self.value_bits > 16:
|
| raise ValueError(
|
| f"value_bits must be 1-16, got {self.value_bits}"
|
| )
|
| valid_scale_dtypes = {"fp32", "fp8_e4m3", "e8m0"}
|
| if self.scale_dtype not in valid_scale_dtypes:
|
| raise ValueError(
|
| f"scale_dtype={self.scale_dtype!r} not in {sorted(valid_scale_dtypes)}"
|
| )
|
| valid_act_scales = {"per-tensor", "per-token", "per-group", "per-channel"}
|
| if self.activation_scale_mode not in valid_act_scales:
|
| raise ValueError(
|
| f"activation_scale_mode={self.activation_scale_mode!r} not in {sorted(valid_act_scales)}"
|
| )
|
|
|
| def to_config(self) -> dict[str, Any]:
|
| """Serialize the full Quantizer configuration to a dict.
|
|
|
| The dict can be passed to Quantizer(**config) or from_config(config)
|
| to reconstruct an identical Quantizer. Used by QuantizedModule
|
| to_dict/from_dict (v3_hybrid_state save/load via save_model/load_model
|
| and convert_model).
|
| """
|
| return {
|
| "value_bits": self.value_bits,
|
| "value_repr": self.value_repr,
|
| "scale_mode": self.scale_mode,
|
| "scale_dtype": self.scale_dtype,
|
| "group_mode": self.group_mode,
|
| "group_size": self.group_size,
|
| "symmetric": self.symmetric,
|
| "double_quant": self.double_quant,
|
| "block_size": self.block_size,
|
| "rotation": self.rotation,
|
| "outlier_threshold": self.outlier_threshold,
|
| "error_compensation": self.error_compensation,
|
| "activation_aware": self.activation_aware,
|
| "alpha": self.alpha,
|
| "codebook_source": self.codebook_source,
|
| "codebook_size": self.codebook_size,
|
| "vq_group_size": self.vq_group_size,
|
| "prune_mode": self.prune_mode,
|
| "prune_ratio": self.prune_ratio,
|
| "quantizes_input": self.quantizes_input,
|
| "learnable": self.learnable,
|
| "compute_dtype": self.compute_dtype,
|
| "activation_scale_mode": self.activation_scale_mode,
|
| "residual_levels": self.residual_levels,
|
| "residual_codebook_size": self.residual_codebook_size,
|
| "num_heads": self.num_heads,
|
| "head_dim": self.head_dim,
|
| }
|
|
|
| @classmethod
|
| def from_config(cls, config: dict[str, Any]) -> "Quantizer":
|
| """Reconstruct a Quantizer from a to_config() dict."""
|
| return cls(**config)
|
|
|
| @property
|
| def n_levels(self) -> int:
|
| """Max quantized value for symmetric int."""
|
| n = (1 << (self.value_bits - 1)) - 1
|
| return max(n, 1)
|
|
|
| @property
|
| def min_val(self) -> int:
|
| return -self.n_levels if self.symmetric else -self.n_levels
|
|
|
| @property
|
| def max_val(self) -> int:
|
| return self.n_levels if self.symmetric else self.n_levels - 1
|
|
|
| def _get_lut(self) -> torch.Tensor:
|
| """Get the LUT for the value representation."""
|
| if self.value_repr == "nf4_lut":
|
| return NF4_LUT
|
| elif self.value_repr == "fp4_e2m1":
|
| return FP4_E2M1_LUT
|
| elif self.value_repr == "fp6_e3m2":
|
| return FP6_E3M2_LUT
|
| elif self.value_repr == "fp6_e2m3":
|
| return FP6_E2M3_LUT
|
| elif self.value_repr == "fp8_e4m3":
|
| return FP8_E4M3_LUT
|
| elif self.value_repr == "fp8_e5m2":
|
| return FP8_E5M2_LUT
|
| return None
|
|
|
| def _compute_scale(self, W: torch.Tensor) -> torch.Tensor:
|
| """Compute per-channel/per-tensor scale for the weight."""
|
| if self.scale_mode == "per-tensor":
|
| max_abs = W.abs().amax().clamp(min=1e-8)
|
| return max_abs / self.n_levels
|
| elif self.scale_mode == "per-channel":
|
| if W.dim() > 1:
|
| reduce_dims = tuple(range(1, W.dim()))
|
| max_abs = W.abs().amax(dim=reduce_dims).clamp(min=1e-8)
|
| else:
|
| max_abs = W.abs().clamp(min=1e-8)
|
| return max_abs / self.n_levels
|
| elif self.scale_mode == "per-head":
|
|
|
|
|
| hd = self.head_dim if self.head_dim > 0 else (
|
| W.shape[-1] // self.num_heads if self.num_heads > 0 else 1)
|
| if W.dim() > 1:
|
| flat = W.reshape(W.shape[0], -1) if W.dim() > 2 else W
|
| else:
|
| flat = W.reshape(1, -1)
|
| pad = (hd - (flat.shape[1] % hd)) % hd if hd > 0 else 0
|
| if pad > 0:
|
| flat = torch.nn.functional.pad(flat, (0, pad))
|
| num_heads_eff = flat.shape[1] // hd if hd > 0 else 1
|
| grouped = flat.reshape(flat.shape[0], num_heads_eff, hd)
|
| max_abs = grouped.abs().amax(dim=2).clamp(min=1e-8)
|
| return max_abs / self.n_levels
|
| elif self.scale_mode in ("per-group", "per-block"):
|
| gs = self.group_size if self.group_size > 0 else 1
|
| if W.dim() > 1:
|
| flat = W.reshape(W.shape[0], -1) if W.dim() > 2 else W
|
| else:
|
| flat = W.reshape(1, -1)
|
| pad = (gs - (flat.shape[1] % gs)) % gs
|
| if pad > 0:
|
| flat = torch.nn.functional.pad(flat, (0, pad))
|
| num_groups = flat.shape[1] // gs
|
| grouped = flat.reshape(flat.shape[0], num_groups, gs)
|
| max_abs = grouped.abs().amax(dim=2).clamp(min=1e-8)
|
| return max_abs / self.n_levels
|
| else:
|
| raise ValueError(f"Unknown scale_mode: {self.scale_mode}")
|
|
|
| def _apply_rotation(self, W: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| """Apply rotation (QuIP). Returns (rotated_weight, rotation_matrix)."""
|
| if self.rotation == "none":
|
| return W, None
|
| n = W.shape[1] if W.dim() > 1 else W.shape[0]
|
| if self.rotation == "hadamard" and (n & (n - 1)) == 0:
|
| H = torch.ones(1, 1, dtype=torch.float32)
|
| while H.shape[0] < n:
|
| H = torch.cat([torch.cat([H, H], dim=1), torch.cat([H, -H], dim=1)], dim=0)
|
| Q = H / (n ** 0.5)
|
| elif self.rotation == "random":
|
| A = torch.randn(n, n, dtype=torch.float32)
|
| Q, R = torch.linalg.qr(A)
|
| d = torch.diagonal(R).sign()
|
| Q = Q * d.unsqueeze(0)
|
| else:
|
| return W, None
|
| W_rot = W @ Q.to(W.device)
|
| return W_rot, Q.to(torch.float32)
|
|
|
|
|
|
|
| def quantize_weight(self, W: torch.Tensor) -> QuantizedWeight:
|
| """Quantize weight tensor → QuantizedWeight (packed buffers + meta)."""
|
| W = W.detach().float()
|
| original_shape = list(W.shape)
|
|
|
|
|
| if self.prune_mode != "none":
|
| return self._quantize_prune(W, original_shape)
|
|
|
|
|
| W_proc, Q = self._apply_rotation(W)
|
|
|
|
|
| if self.outlier_threshold is not None:
|
| return self._quantize_outlier(W_proc, original_shape, Q)
|
|
|
|
|
| if self.activation_aware == "awq":
|
| W_proc = self._apply_awq_scaling(W_proc)
|
|
|
|
|
| if self.value_repr == "int":
|
| return self._quantize_int(W_proc, original_shape, Q)
|
| elif self.value_repr == "nf4_lut":
|
| return self._quantize_nf4(W_proc, original_shape, Q)
|
| elif self.value_repr == "fp4_e2m1":
|
| return self._quantize_fp4(W_proc, original_shape, Q)
|
| elif self.value_repr in ("fp6_e3m2", "fp6_e2m3"):
|
| return self._quantize_fp6(W_proc, original_shape, Q)
|
| elif self.value_repr in ("fp8_e4m3", "fp8_e5m2"):
|
| return self._quantize_fp8(W_proc, original_shape, Q)
|
| elif self.value_repr == "codebook":
|
| return self._quantize_codebook(W_proc, original_shape, Q)
|
| elif self.value_repr == "binary":
|
| return self._quantize_binary(W_proc, original_shape, Q)
|
| elif self.value_repr == "ternary":
|
| return self._quantize_ternary(W_proc, original_shape, Q)
|
| elif self.value_repr == "none":
|
| return self._quantize_none(W_proc, original_shape, Q)
|
| else:
|
| raise ValueError(f"Unknown value_repr: {self.value_repr}")
|
|
|
|
|
|
|
| def dequantize_weight(
|
| self,
|
| qw: QuantizedWeight,
|
| compute_dtype: str = "fp32",
|
| slice: tuple[int, int] | None = None,
|
| ) -> torch.Tensor:
|
| """Reconstruct float weight from QuantizedWeight.
|
|
|
| Args:
|
| qw: the weight container.
|
| compute_dtype: target dtype.
|
| slice: optional (start, end) for chunked dequant (output dim).
|
| """
|
| t = _compute_dtype_to_torch(compute_dtype)
|
| meta = qw.weight_meta
|
| repr_ = meta.get("value_repr", self.value_repr)
|
| original_shape = meta.get("original_shape")
|
|
|
|
|
| if slice is not None and repr_ in ("int", "nf4_lut", "fp4_e2m1", "fp6_e3m2", "fp6_e2m3", "fp8_e4m3", "fp8_e5m2"):
|
| w = self._dequant_slice(qw, slice)
|
| else:
|
| w = self._dequant_full(qw)
|
|
|
|
|
| if meta.get("has_rotation", False):
|
| Q = qw.weight_buffers.get("rotation_Q")
|
| if Q is not None:
|
| w = w @ Q.to(w.dtype).T
|
|
|
|
|
| if original_shape is not None and list(w.shape) != original_shape and slice is None:
|
| w = w.reshape(*original_shape)
|
|
|
| return w.to(t)
|
|
|
| def _dequant_full(self, qw: QuantizedWeight) -> torch.Tensor:
|
| """Full dequant (no slicing)."""
|
| meta = qw.weight_meta
|
| repr_ = meta.get("value_repr", self.value_repr)
|
| buf = qw.weight_buffers
|
|
|
| if repr_ == "int":
|
|
|
| if meta.get("group_mode") == "magnitude-binned":
|
| weight_int = buf["weight_int"].to(torch.float32)
|
| bin_idx = buf["bin_idx"].long()
|
| bin_scales = buf["bin_scales"].to(torch.float32)
|
| gs = meta.get("group_size", self.group_size)
|
| n_bins = meta.get("n_bins", 4)
|
| out_features = weight_int.shape[0]
|
| in_padded = weight_int.shape[1]
|
| num_groups = in_padded // gs if gs > 0 else 1
|
| wi_g = weight_int.reshape(out_features, num_groups, gs)
|
| bi_g = bin_idx.reshape(out_features, num_groups, gs)
|
| scale_sel = torch.gather(bin_scales.unsqueeze(2).expand(*bi_g.shape, bin_scales.shape[-1]), 3, bi_g.unsqueeze(-1).long()).squeeze(-1)
|
| w = (wi_g * scale_sel).reshape(out_features, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
|
|
| if meta.get("group_mode") == "super-block-nested":
|
| fmt = meta.get("kquant_fmt", _kquant_bits_to_fmt(meta.get("value_bits", 0)))
|
| qd = {k: v for k, v in buf.items() if isinstance(v, torch.Tensor)}
|
| qd["in_features"] = meta.get("in_features", 0)
|
| qd["in_padded"] = meta.get("in_padded", 0)
|
| qd["out_features"] = meta.get("out_features", 0)
|
| return _kquant.dequantize_blocks(qd, fmt)
|
| weight_int = buf["weight_int"].to(torch.float32)
|
|
|
| scale_dtype = meta.get("scale_dtype", "fp32")
|
| if "scale_e8m0" in buf:
|
| e8m0_code = buf["scale_e8m0"]
|
| scale = E8M0_LUT.to(e8m0_code.device)[e8m0_code.long()].to(torch.float32)
|
| elif "scale_fp8" in buf:
|
| fp8_codes = buf["scale_fp8"]
|
| ws2 = buf.get("weight_scale_2", torch.tensor(1.0, dtype=torch.float32))
|
| scale = dequantize_fp8(fp8_codes, torch.ones(fp8_codes.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| scale = scale * ws2.to(torch.float32)
|
| else:
|
| scale = buf["scale"].to(torch.float32)
|
| if meta.get("scale_mode") == "per-channel" and scale.numel() > 1:
|
| reshape = [1] * weight_int.dim()
|
| reshape[0] = weight_int.shape[0]
|
| return weight_int * scale.reshape(reshape)
|
| elif meta.get("scale_mode") == "per-head" and meta.get("head_dim", 0) > 0:
|
| hd = meta.get("head_dim", self.head_dim)
|
| if weight_int.dim() > 1:
|
| out_features = weight_int.shape[0]
|
| in_padded = weight_int.shape[1]
|
| num_heads_eff = in_padded // hd if hd > 0 else 1
|
| grouped = weight_int.reshape(out_features, num_heads_eff, hd)
|
| scale_exp = scale.unsqueeze(2).expand_as(grouped)
|
| w = (grouped * scale_exp).reshape(out_features, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
| elif meta.get("scale_mode") in ("per-group", "per-block") and meta.get("group_size", 0) > 0:
|
| gs = meta.get("group_size", self.group_size)
|
| if weight_int.dim() > 1:
|
| out_features = weight_int.shape[0]
|
| in_padded = weight_int.shape[1]
|
| num_groups = in_padded // gs if gs > 0 else 1
|
| grouped = weight_int.reshape(out_features, num_groups, gs)
|
| scale_exp = scale.unsqueeze(2).expand_as(grouped)
|
| w = (grouped * scale_exp).reshape(out_features, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
|
|
| if scale.numel() > 1 and scale.dim() == 1 and weight_int.dim() > 1:
|
| reshape = [1] * weight_int.dim()
|
| reshape[0] = weight_int.shape[0]
|
| return weight_int * scale.reshape(reshape)
|
| return weight_int * scale
|
|
|
| elif repr_ == "nf4_lut":
|
| weight_packed = buf["weight_packed"]
|
| idx = unpack_nf4(weight_packed)
|
| if meta.get("use_double_quant", False):
|
| in_padded = idx.shape[1]
|
| gs = meta["group_size"]
|
| num_groups = in_padded // gs
|
| scale = dequantize_scales_2d(
|
| buf["scale_packed"], buf["block_scale"],
|
| block_size=meta["block_size"], num_groups=num_groups,
|
| )
|
| else:
|
| scale = buf["scale"]
|
| gs = meta["group_size"]
|
| w = dequantize_nf4(idx, scale, group_size=gs, in_features=meta.get("in_features"))
|
| return w
|
|
|
| elif repr_ == "fp4_e2m1":
|
| weight_packed = buf["weight_packed"]
|
| idx = unpack_fp4(weight_packed)
|
| if meta.get("scale_dtype") == "fp8_e4m3":
|
| fp8_codes = buf["scale_fp8"].view(torch.uint8)
|
| norm_scale = dequantize_fp8(fp8_codes, torch.ones(fp8_codes.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| weight_scale_2 = buf["weight_scale_2"]
|
| scale = norm_scale * weight_scale_2.to(torch.float32)
|
| elif meta.get("scale_dtype") == "e8m0":
|
| e8m0_code = buf["scale_e8m0"]
|
| scale = E8M0_LUT.to(e8m0_code.device)[e8m0_code.long()].to(torch.float32)
|
| else:
|
| scale = buf["scale"]
|
| gs = meta["group_size"]
|
| w = dequantize_fp4(idx, scale, group_size=gs, in_features=meta.get("in_features"))
|
| return w
|
|
|
| elif repr_ in ("fp6_e3m2", "fp6_e2m3"):
|
| lut = FP6_E3M2_LUT if repr_ == "fp6_e3m2" else FP6_E2M3_LUT
|
| weight_packed = buf["weight_packed"]
|
| codes = unpack_fp6(weight_packed)
|
| if meta.get("scale_dtype") == "fp8_e4m3":
|
| fp8_codes = buf["scale_fp8"].view(torch.uint8)
|
| norm_scale = dequantize_fp8(fp8_codes, torch.ones(fp8_codes.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| weight_scale_2 = buf["weight_scale_2"]
|
| scale = norm_scale * weight_scale_2.to(torch.float32)
|
| elif meta.get("scale_dtype") == "e8m0":
|
| e8m0_code = buf["scale_e8m0"]
|
| scale = E8M0_LUT.to(e8m0_code.device)[e8m0_code.long()].to(torch.float32)
|
| else:
|
| scale = buf["scale"]
|
| gs = meta["group_size"]
|
| w = dequantize_fp6(codes, scale, lut, group_size=gs, in_features=meta.get("in_features"))
|
| return w
|
|
|
| elif repr_ in ("fp8_e4m3", "fp8_e5m2"):
|
| lut = FP8_E4M3_LUT if repr_ == "fp8_e4m3" else FP8_E5M2_LUT
|
| packed = buf["weight_packed"]
|
| codes = packed.view(torch.uint8)
|
|
|
| if "scale_e8m0" in buf:
|
| e8m0_code = buf["scale_e8m0"]
|
| scale = E8M0_LUT.to(e8m0_code.device)[e8m0_code.long()].to(torch.float32)
|
| elif "scale_fp8" in buf:
|
| fp8_codes = buf["scale_fp8"].view(torch.uint8)
|
| norm_scale = dequantize_fp8(fp8_codes, torch.ones(fp8_codes.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| weight_scale_2 = buf["weight_scale_2"]
|
| scale = norm_scale * weight_scale_2.to(torch.float32)
|
| else:
|
| scale = buf["scale"]
|
| if meta.get("scale_mode") == "per-group" and meta.get("group_size", 0) > 0:
|
| gs = meta["group_size"]
|
| out_f = codes.shape[0]
|
| in_padded = codes.shape[1]
|
| num_groups = in_padded // gs
|
| grouped_codes = codes.reshape(out_f, num_groups, gs)
|
| w_norm = lut.to(codes.device)[grouped_codes.long()]
|
| scale_exp = scale.unsqueeze(2).expand_as(w_norm)
|
| w = (w_norm * scale_exp).reshape(out_f, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
| if scale.numel() == 1:
|
| w = lut.to(codes.device)[codes.long()] * scale.to(torch.float32)
|
| else:
|
| reshape = [1] * codes.dim()
|
| reshape[0] = codes.shape[0]
|
| w = lut.to(codes.device)[codes.long()] * scale.to(torch.float32).reshape(reshape)
|
| return w
|
|
|
| elif repr_ == "codebook":
|
| scale = buf["scale"].to(torch.float32)
|
|
|
| n_res = meta.get("residual_levels", 1)
|
| if n_res > 1 and "codebooks" in buf:
|
| codebooks = buf["codebooks"].to(torch.float32)
|
| cb_sizes = buf["codebook_sizes"].long()
|
| indices_pl = buf["indices_per_level"].long()
|
| w_norm = torch.zeros_like(indices_pl[0].to(torch.float32))
|
| for lvl in range(n_res):
|
| actual_K = cb_sizes[lvl].item()
|
| cb = codebooks[lvl, :actual_K]
|
| idx = indices_pl[lvl]
|
| w_norm = w_norm + cb.to(idx.device)[idx]
|
| if scale.dim() == 1 and scale.numel() > 1:
|
| w = w_norm * scale.unsqueeze(1)
|
| else:
|
| w = w_norm * scale
|
| return w[:, :meta.get("in_features", w.shape[1])]
|
|
|
| indices = buf["indices"].long()
|
| codebook = buf["codebook"].to(torch.float32)
|
| if meta.get("codebook_source") == "vq":
|
| gs = meta.get("vq_group_size", 2)
|
| vecs = codebook[indices]
|
| w = vecs.reshape(indices.shape[0], -1)
|
| if scale.dim() == 1:
|
| w = w * scale.unsqueeze(1)
|
| else:
|
| w_norm = codebook[indices]
|
| if scale.dim() == 1 and scale.numel() > 1:
|
| w = w_norm * scale.unsqueeze(1)
|
| else:
|
| w = w_norm * scale
|
| return w[:, :meta.get("in_features", w.shape[1])]
|
|
|
| elif repr_ == "binary":
|
| binary = buf["weight_binary"].to(torch.float32)
|
| scale = buf["scale"].to(torch.float32)
|
| if scale.numel() == 1:
|
| return binary * scale
|
| reshape = [1] * binary.dim()
|
| reshape[0] = binary.shape[0]
|
| return binary * scale.reshape(reshape)
|
|
|
| elif repr_ == "ternary":
|
| ternary = buf["weight_ternary"].to(torch.float32)
|
| scale = buf["scale"].to(torch.float32)
|
| if scale.numel() == 1:
|
| return ternary * scale
|
| reshape = [1] * ternary.dim()
|
| reshape[0] = ternary.shape[0]
|
| return ternary * scale.reshape(reshape)
|
|
|
| elif repr_ == "none":
|
| return buf["weight_fp"].to(torch.float32)
|
|
|
| elif repr_ == "outlier":
|
| dense_int = buf["dense_int"].to(torch.float32)
|
| scale = buf["scale"].to(torch.float32)
|
| sm = meta.get("scale_mode", self.scale_mode)
|
| if sm == "per-channel" and scale.numel() > 1 and dense_int.dim() > 1:
|
| w = dense_int * scale.unsqueeze(1)
|
| else:
|
| w = dense_int * scale
|
| outlier_indices = buf["outlier_indices"]
|
| outlier_values = buf["outlier_values"]
|
| if outlier_indices.numel() > 0:
|
| w_flat = w.flatten()
|
| w_flat[outlier_indices.long()] = outlier_values.to(torch.float32)
|
| w = w_flat.reshape(dense_int.shape)
|
| return w
|
|
|
| elif repr_ == "prune":
|
| return buf["weight_pruned"].to(torch.float32) * buf["mask"].to(torch.float32)
|
|
|
| raise ValueError(f"Unknown value_repr in dequant: {repr_}")
|
|
|
| def _dequant_slice(self, qw: QuantizedWeight, slc: tuple[int, int]) -> torch.Tensor:
|
| """Dequant only output rows [start:end] — for chunked forward."""
|
| start, end = slc
|
| meta = qw.weight_meta
|
| repr_ = meta.get("value_repr", self.value_repr)
|
| buf = qw.weight_buffers
|
|
|
| if repr_ == "int":
|
|
|
| if meta.get("group_mode") == "magnitude-binned":
|
| weight_int = buf["weight_int"][start:end].to(torch.float32)
|
| bin_idx = buf["bin_idx"][start:end].long()
|
| bin_scales = buf["bin_scales"][start:end].to(torch.float32)
|
| gs = meta.get("group_size", self.group_size)
|
| out_features = weight_int.shape[0]
|
| in_padded = weight_int.shape[1]
|
| num_groups = in_padded // gs if gs > 0 else 1
|
| wi_g = weight_int.reshape(out_features, num_groups, gs)
|
| bi_g = bin_idx.reshape(out_features, num_groups, gs)
|
| scale_sel = torch.gather(bin_scales.unsqueeze(2).expand(*bi_g.shape, bin_scales.shape[-1]), 3, bi_g.unsqueeze(-1).long()).squeeze(-1)
|
| w = (wi_g * scale_sel).reshape(out_features, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
|
|
| if meta.get("group_mode") == "super-block-nested":
|
| fmt = meta.get("kquant_fmt", _kquant_bits_to_fmt(meta.get("value_bits", 0)))
|
| out_total = meta.get("out_features", 0)
|
| qd = {k: (v[start:end] if isinstance(v, torch.Tensor) and v.dim() > 0 and v.shape[0] == out_total else v)
|
| for k, v in buf.items() if isinstance(v, torch.Tensor) and k != "rotation_Q"}
|
| qd["in_features"] = meta.get("in_features", 0)
|
| qd["in_padded"] = meta.get("in_padded", 0)
|
| qd["out_features"] = end - start
|
| return _kquant.dequantize_blocks(qd, fmt)
|
| weight_int = buf["weight_int"][start:end].to(torch.float32)
|
|
|
| if "scale_e8m0" in buf:
|
| e8m0_full = buf["scale_e8m0"]
|
| scale_full = E8M0_LUT.to(e8m0_full.device)[e8m0_full.long()].to(torch.float32)
|
| scale = scale_full[start:end]
|
| elif "scale_fp8" in buf:
|
| fp8_full = buf["scale_fp8"].view(torch.uint8)
|
| norm_full = dequantize_fp8(fp8_full, torch.ones(fp8_full.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| ws2 = buf.get("weight_scale_2", torch.tensor(1.0, dtype=torch.float32))
|
| scale = norm_full[start:end] * ws2.to(torch.float32)
|
| else:
|
| scale_full = buf["scale"].to(torch.float32)
|
| out_total = meta.get("out_features", 0)
|
| if scale_full.shape[0] == out_total and out_total > 0:
|
| scale = scale_full[start:end]
|
| else:
|
| scale = scale_full
|
| sm = meta.get("scale_mode", self.scale_mode)
|
| if sm == "per-channel" or (sm in ("per-group", "per-block") and meta.get("group_size", 0) == 0):
|
| s = scale if scale.numel() > 1 else scale
|
| return weight_int * s.unsqueeze(1) if s.numel() > 1 else weight_int * s
|
| elif sm == "per-head" and meta.get("head_dim", 0) > 0:
|
| hd = meta.get("head_dim", self.head_dim)
|
| out_f = weight_int.shape[0]
|
| in_padded = weight_int.shape[1]
|
| num_heads_eff = in_padded // hd if hd > 0 else 1
|
| grouped = weight_int.reshape(out_f, num_heads_eff, hd)
|
| scale_exp = scale.unsqueeze(2).expand_as(grouped)
|
| w = (grouped * scale_exp).reshape(out_f, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
| elif sm in ("per-group", "per-block") and meta.get("group_size", 0) > 0:
|
| gs = meta.get("group_size", self.group_size)
|
| out_f = weight_int.shape[0]
|
| in_padded = weight_int.shape[1]
|
| num_groups = in_padded // gs if gs > 0 else 1
|
| grouped = weight_int.reshape(out_f, num_groups, gs)
|
| s = scale if scale.dim() > 1 else scale
|
| scale_exp = s.unsqueeze(2).expand_as(grouped)
|
| w = (grouped * scale_exp).reshape(out_f, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
| return weight_int * scale
|
|
|
| elif repr_ == "nf4_lut":
|
| weight_packed = buf["weight_packed"][start:end]
|
| idx = unpack_nf4(weight_packed)
|
| if meta.get("use_double_quant", False):
|
| gs = meta["group_size"]
|
| in_padded = idx.shape[1]
|
| num_groups = in_padded // gs
|
| scale_packed = buf["scale_packed"][start:end]
|
| block_scale = buf["block_scale"][start:end]
|
| scale = dequantize_scales_2d(scale_packed, block_scale, block_size=meta["block_size"], num_groups=num_groups)
|
| else:
|
| scale = buf["scale"][start:end] if buf["scale"].dim() > 1 else buf["scale"]
|
| gs = meta["group_size"]
|
| return dequantize_nf4(idx, scale, group_size=gs, in_features=meta.get("in_features"))
|
|
|
| elif repr_ == "fp4_e2m1":
|
| weight_packed = buf["weight_packed"][start:end]
|
| idx = unpack_fp4(weight_packed)
|
| if meta.get("scale_dtype") == "fp8_e4m3":
|
| fp8_codes = buf["scale_fp8"][start:end].view(torch.uint8)
|
| norm_scale = dequantize_fp8(fp8_codes, torch.ones(fp8_codes.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| scale = norm_scale * buf["weight_scale_2"].to(torch.float32)
|
| elif meta.get("scale_dtype") == "e8m0":
|
| e8m0_code = buf["scale_e8m0"][start:end]
|
| scale = E8M0_LUT.to(e8m0_code.device)[e8m0_code.long()].to(torch.float32)
|
| else:
|
| scale = buf["scale"][start:end] if buf["scale"].dim() > 1 else buf["scale"]
|
| gs = meta["group_size"]
|
| return dequantize_fp4(idx, scale, group_size=gs, in_features=meta.get("in_features"))
|
|
|
| elif repr_ in ("fp6_e3m2", "fp6_e2m3"):
|
| lut = FP6_E3M2_LUT if repr_ == "fp6_e3m2" else FP6_E2M3_LUT
|
| packed = buf["weight_packed"][start:end]
|
| codes = unpack_fp6(packed)
|
| if "scale_e8m0" in buf:
|
| e8m0_full = buf["scale_e8m0"]
|
| scale_full = E8M0_LUT.to(e8m0_full.device)[e8m0_full.long()].to(torch.float32)
|
| scale = scale_full[start:end]
|
| elif "scale_fp8" in buf:
|
| fp8_full = buf["scale_fp8"].view(torch.uint8)
|
| norm_full = dequantize_fp8(fp8_full, torch.ones(fp8_full.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| ws2 = buf.get("weight_scale_2", torch.tensor(1.0, dtype=torch.float32))
|
| scale = norm_full[start:end] * ws2.to(torch.float32)
|
| else:
|
| scale = buf["scale"][start:end] if buf["scale"].numel() > 1 else buf["scale"]
|
| gs = meta.get("group_size", 32)
|
| out_f = codes.shape[0]
|
| in_padded = codes.shape[1]
|
| num_groups = in_padded // gs
|
| grouped = codes.reshape(out_f, num_groups, gs)
|
| w_norm = lut.to(codes.device)[grouped.long()]
|
| scale_exp = scale.unsqueeze(2).expand_as(w_norm)
|
| w = (w_norm * scale_exp).reshape(out_f, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
|
|
| elif repr_ in ("fp8_e4m3", "fp8_e5m2"):
|
| lut = FP8_E4M3_LUT if repr_ == "fp8_e4m3" else FP8_E5M2_LUT
|
| packed = buf["weight_packed"][start:end]
|
| codes = packed.view(torch.uint8)
|
| if "scale_e8m0" in buf:
|
| e8m0_full = buf["scale_e8m0"]
|
| scale_full = E8M0_LUT.to(e8m0_full.device)[e8m0_full.long()].to(torch.float32)
|
| scale = scale_full[start:end]
|
| elif "scale_fp8" in buf:
|
| fp8_full = buf["scale_fp8"].view(torch.uint8)
|
| norm_full = dequantize_fp8(fp8_full, torch.ones(fp8_full.shape[0], dtype=torch.float32), FP8_E4M3_LUT)
|
| ws2 = buf.get("weight_scale_2", torch.tensor(1.0, dtype=torch.float32))
|
| scale = norm_full[start:end] * ws2.to(torch.float32)
|
| else:
|
| scale = buf["scale"]
|
| if meta.get("scale_mode") == "per-group" and meta.get("group_size", 0) > 0:
|
| gs = meta["group_size"]
|
| out_f = codes.shape[0]
|
| in_padded = codes.shape[1]
|
| num_groups = in_padded // gs
|
| grouped = codes.reshape(out_f, num_groups, gs)
|
| w_norm = lut.to(codes.device)[grouped.long()]
|
| scale_exp = scale.unsqueeze(2).expand_as(w_norm)
|
| w = (w_norm * scale_exp).reshape(out_f, in_padded)
|
| return w[:, :meta.get("in_features", in_padded)]
|
| if scale.numel() == 1:
|
| return lut.to(codes.device)[codes.long()] * scale.to(torch.float32)
|
| s = scale[start:end] if scale.dim() > 1 else scale
|
| return lut.to(codes.device)[codes.long()] * s.to(torch.float32).unsqueeze(1) if s.numel() > 1 else lut.to(codes.device)[codes.long()] * s.to(torch.float32)
|
|
|
|
|
| w_full = self._dequant_full(qw)
|
| return w_full[start:end]
|
|
|
|
|
|
|
| def _compute_activation_scale(self, x_f: torch.Tensor, lut: torch.Tensor | None) -> torch.Tensor:
|
| """Compute activation scale according to activation_scale_mode.
|
|
|
| Args:
|
| x_f: [batch, features] or [batch, seq, features] activation tensor.
|
| lut: LUT for the value representation (None for int/ternary/binary).
|
|
|
| Returns:
|
| scale tensor broadcastable to x_f for dequant:
|
| per-tensor: scalar [1]
|
| per-token: [batch] or [batch, seq] (per-row max)
|
| per-group: [batch, num_groups] or [batch, seq, num_groups]
|
| per-channel: [features] (per-column max)
|
| """
|
| max_lut = lut.abs().amax().clamp(min=1e-12) if lut is not None else None
|
| n = self.n_levels if max_lut is None else max_lut
|
| asm = self.activation_scale_mode
|
|
|
| if asm == "per-tensor":
|
| return (x_f.abs().amax().clamp(min=1e-8) / n).reshape(1)
|
|
|
| if asm == "per-token":
|
|
|
| reduce_dim = x_f.dim() - 1
|
| return (x_f.abs().amax(dim=reduce_dim).clamp(min=1e-8) / n)
|
|
|
| if asm == "per-channel":
|
|
|
| reduce_dims = tuple(range(x_f.dim() - 1))
|
| return (x_f.abs().amax(dim=reduce_dims).clamp(min=1e-8) / n)
|
|
|
| if asm == "per-group":
|
| gs = self.group_size if self.group_size > 0 else 32
|
| feat = x_f.shape[-1]
|
| pad = (gs - (feat % gs)) % gs
|
| if pad > 0:
|
| x_padded = torch.nn.functional.pad(x_f, (0, pad))
|
| else:
|
| x_padded = x_f
|
| num_groups = x_padded.shape[-1] // gs
|
| grouped = x_padded.reshape(*x_padded.shape[:-1], num_groups, gs)
|
| max_abs = grouped.abs().amax(dim=-1).clamp(min=1e-8)
|
| return max_abs / n
|
|
|
| raise ValueError(f"Unknown activation_scale_mode: {asm}")
|
|
|
| def _apply_activation_scale(self, x_f: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
| """Broadcast scale to x_f shape for element-wise division."""
|
| asm = self.activation_scale_mode
|
| if asm == "per-tensor":
|
| return x_f / scale
|
| if asm == "per-token":
|
|
|
| shape = list(scale.shape) + [1]
|
| return x_f / scale.reshape(shape)
|
| if asm == "per-channel":
|
| return x_f / scale
|
| if asm == "per-group":
|
|
|
| gs = self.group_size if self.group_size > 0 else 32
|
| scale_exp = scale.unsqueeze(-1).expand(*scale.shape, gs)
|
| scale_exp = scale_exp.reshape(*scale_exp.shape[:-2], -1)
|
| feat = x_f.shape[-1]
|
| return x_f / scale_exp[..., :feat]
|
| raise ValueError(f"Unknown activation_scale_mode: {asm}")
|
|
|
| def _dequant_activation(self, x_quant: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
| """Multiply quantized activation by scale with proper broadcasting."""
|
| asm = self.activation_scale_mode
|
| if asm == "per-tensor":
|
| return x_quant * scale
|
| if asm == "per-token":
|
| shape = list(scale.shape) + [1]
|
| return x_quant * scale.reshape(shape)
|
| if asm == "per-channel":
|
| return x_quant * scale
|
| if asm == "per-group":
|
| gs = self.group_size if self.group_size > 0 else 32
|
| scale_exp = scale.unsqueeze(-1).expand(*scale.shape, gs)
|
| scale_exp = scale_exp.reshape(*scale_exp.shape[:-2], -1)
|
| feat = x_quant.shape[-1]
|
| return x_quant * scale_exp[..., :feat]
|
| raise ValueError(f"Unknown activation_scale_mode: {asm}")
|
|
|
| def quantize_input(self, x: torch.Tensor, qw: QuantizedWeight) -> QuantizedActivation:
|
| """Quantize activations -> QuantizedActivation for ALL value_repr.
|
|
|
| Supports activation_scale_mode: per-tensor / per-token / per-group /
|
| per-channel. Dequantizes immediately (stores fake-quantized data in
|
| 'data' buffer) since there is no int matmul kernel — dequantize_input
|
| just returns the stored fake-quantized tensor.
|
| """
|
| if not self.quantizes_input:
|
| return QuantizedActivation()
|
|
|
| x_f = x.detach().float()
|
| buffers = {}
|
| meta = {
|
| "activation_scale_mode": self.activation_scale_mode,
|
| "value_repr": self.value_repr,
|
| }
|
|
|
| if self.activation_aware == "smoothquant":
|
| s = qw.weight_meta.get("smoothing_s")
|
| if s is None:
|
| s = torch.ones(x_f.shape[-1], dtype=torch.float32)
|
| x_f = x_f / s.to(torch.float32).unsqueeze(0)
|
| meta["smoothing_s"] = s
|
|
|
| lut = self._get_lut()
|
| scale = self._compute_activation_scale(x_f, lut)
|
| x_norm = self._apply_activation_scale(x_f, scale)
|
|
|
| if self.value_repr == "int":
|
| n = self.n_levels
|
| if self.symmetric:
|
| q = torch.clamp(torch.round(x_norm), min=-n, max=n)
|
| else:
|
| q = torch.clamp(torch.round(x_norm), min=-n, max=n - 1)
|
| buffers["data"] = self._dequant_activation(q.to(torch.float32), scale).to(torch.float32)
|
|
|
| elif self.value_repr in ("fp8_e4m3", "fp8_e5m2"):
|
| lut_act = FP8_E4M3_LUT if self.value_repr == "fp8_e4m3" else FP8_E5M2_LUT
|
| lut_max = lut_act.abs().amax().item()
|
| x_norm_c = x_norm.clamp(-lut_max, lut_max)
|
| diff = x_norm_c.unsqueeze(-1) - lut_act.to(x_norm_c.device)
|
| codes = diff.abs().argmin(dim=-1).to(torch.uint8)
|
| w_norm = lut_act.to(codes.device)[codes.long()].to(torch.float32)
|
| buffers["data"] = self._dequant_activation(w_norm, scale)
|
|
|
| elif self.value_repr == "fp4_e2m1":
|
| lut_act = FP4_E2M1_LUT
|
| lut_max = lut_act.abs().amax().item()
|
| x_norm_c = x_norm.clamp(-lut_max, lut_max)
|
| diff = x_norm_c.unsqueeze(-1) - lut_act.to(x_norm_c.device)
|
| idx = diff.abs().argmin(dim=-1)
|
| w_norm = lut_act.to(idx.device)[idx].to(torch.float32)
|
| buffers["data"] = self._dequant_activation(w_norm, scale)
|
|
|
| elif self.value_repr in ("fp6_e3m2", "fp6_e2m3"):
|
| lut_act = FP6_E3M2_LUT if self.value_repr == "fp6_e3m2" else FP6_E2M3_LUT
|
| lut_max = lut_act.abs().amax().item()
|
| x_norm_c = x_norm.clamp(-lut_max, lut_max)
|
| diff = x_norm_c.unsqueeze(-1) - lut_act.to(x_norm_c.device)
|
| idx = diff.abs().argmin(dim=-1)
|
| w_norm = lut_act.to(idx.device)[idx].to(torch.float32)
|
| buffers["data"] = self._dequant_activation(w_norm, scale)
|
|
|
| elif self.value_repr == "nf4_lut":
|
| lut_act = NF4_LUT
|
| lut_max = lut_act.abs().amax().item()
|
| x_norm_c = x_norm.clamp(-lut_max, lut_max)
|
| diff = x_norm_c.unsqueeze(-1) - lut_act.to(x_norm_c.device)
|
| idx = diff.abs().argmin(dim=-1)
|
| w_norm = lut_act.to(idx.device)[idx].to(torch.float32)
|
| buffers["data"] = self._dequant_activation(w_norm, scale)
|
|
|
| elif self.value_repr == "ternary":
|
|
|
|
|
| threshold = x_f.abs().mean().clamp(min=1e-8) * 0.5
|
| ternary = torch.where(x_f.abs() < threshold, torch.zeros_like(x_f),
|
| torch.sign(x_f))
|
| buffers["data"] = self._dequant_activation(ternary.to(torch.float32), scale).to(torch.float32)
|
|
|
| elif self.value_repr == "binary":
|
|
|
| binary = torch.sign(x_f)
|
| buffers["data"] = self._dequant_activation(binary.to(torch.float32), scale).to(torch.float32)
|
|
|
| elif self.value_repr == "codebook":
|
|
|
| codebook = qw.weight_buffers.get("codebook")
|
| if codebook is not None:
|
| cb = codebook.to(torch.float32)
|
| diff = x_norm.unsqueeze(-1) - cb.to(x_norm.device)
|
| indices = diff.abs().argmin(dim=-1)
|
| w_norm = cb.to(indices.device)[indices].to(torch.float32)
|
| buffers["data"] = self._dequant_activation(w_norm, scale)
|
| else:
|
|
|
| cb = self._kmeans_1d(x_norm.flatten(), self.codebook_size)
|
| diff = x_norm.unsqueeze(-1) - cb.to(x_norm.device)
|
| indices = diff.abs().argmin(dim=-1)
|
| w_norm = cb.to(indices.device)[indices].to(torch.float32)
|
| buffers["data"] = self._dequant_activation(w_norm, scale)
|
|
|
| else:
|
|
|
| buffers["data"] = x_f
|
|
|
| buffers["scale"] = scale.to(torch.float32)
|
| return QuantizedActivation(activation_buffers=buffers, activation_meta=meta)
|
|
|
| def dequantize_input(self, qa: QuantizedActivation, compute_dtype: str = "fp32") -> torch.Tensor:
|
| """Reconstruct activations from QuantizedActivation.
|
|
|
| Since quantize_input stores the fake-quantized (dequantized) data
|
| directly, this just returns it in the target compute_dtype.
|
| """
|
| if not self.quantizes_input or "data" not in qa.activation_buffers:
|
| if "data" in qa.activation_buffers:
|
| return qa.activation_buffers["data"].to(_compute_dtype_to_torch(compute_dtype))
|
| return None
|
| return qa.activation_buffers["data"].to(_compute_dtype_to_torch(compute_dtype))
|
|
|
|
|
|
|
| def storage_bytes(self, qw: QuantizedWeight) -> int:
|
| total = 0
|
| for buf in qw.weight_buffers.values():
|
| if buf is None:
|
| continue
|
| total += buf.numel() * buf.element_size()
|
| return total
|
|
|
| def info(self) -> dict[str, Any]:
|
| return {
|
| "repr": self.value_repr,
|
| "bits": self.value_bits,
|
| "scale": self.scale_mode,
|
| "group": self.group_size if self.group_size > 0 else "-",
|
| "w": True,
|
| "a": self.quantizes_input,
|
| "learnable": self.learnable,
|
| "residual_levels": self.residual_levels,
|
| "codebook_size": self.codebook_size if self.value_repr == "codebook" else "-",
|
| "scale_dtype": self.scale_dtype,
|
| "group_mode": self.group_mode,
|
| "activation_scale": self.activation_scale_mode,
|
| "num_heads": self.num_heads if self.scale_mode == "per-head" else "-",
|
| "head_dim": self.head_dim if self.scale_mode == "per-head" else "-",
|
| }
|
|
|
|
|
|
|
| def _quantize_int(self, W, original_shape, Q=None):
|
| """Uniform integer quantization (int2/3/4/8)."""
|
| n = self.n_levels
|
| if W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| elif W.dim() == 1:
|
| flat = W.reshape(1, -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
|
|
|
|
| if self.group_mode == "super-block-nested":
|
| fmt = _kquant_bits_to_fmt(self.value_bits)
|
| if fmt is None:
|
| raise ValueError(
|
| f"super-block-nested not supported for value_bits={self.value_bits}"
|
| )
|
| qd = _kquant.quantize_blocks(flat, fmt)
|
| meta = {
|
| "value_repr": "int", "value_bits": self.value_bits,
|
| "scale_mode": self.scale_mode, "group_mode": "super-block-nested",
|
| "group_size": qd["block_size"], "kquant_fmt": fmt,
|
| "symmetric": self.symmetric,
|
| "in_features": in_f, "in_padded": qd["in_padded"],
|
| "out_features": out_f, "original_shape": original_shape,
|
| "ndim": len(original_shape), "has_rotation": Q is not None,
|
| }
|
| buffers = {k: v for k, v in qd.items()
|
| if isinstance(v, torch.Tensor)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
|
|
|
|
|
|
| if (self.group_mode == "magnitude-binned"
|
| and self.scale_mode in ("per-group", "per-block") and self.group_size > 0):
|
| n_bins = self.block_size if self.block_size and self.block_size > 1 else 4
|
| gs = self.group_size
|
| pad = (gs - (in_f % gs)) % gs
|
| if pad > 0:
|
| flat = torch.nn.functional.pad(flat, (0, pad))
|
| in_padded = flat.shape[1]
|
| num_groups = in_padded // gs
|
| grouped = flat.reshape(out_f, num_groups, gs)
|
| abs_g = grouped.abs()
|
|
|
|
|
| quantiles = torch.linspace(1.0 / n_bins, 1.0 - 1.0 / n_bins, n_bins - 1,
|
| device=grouped.device)
|
| boundaries = torch.quantile(abs_g, quantiles, dim=2).permute(1, 2, 0)
|
|
|
|
|
| bin_idx = (abs_g.unsqueeze(-1) > boundaries.unsqueeze(2)).sum(dim=-1)
|
| bin_idx = bin_idx.clamp(max=n_bins - 1).to(torch.int16)
|
|
|
| bin_scales = torch.zeros(out_f, num_groups, n_bins, dtype=torch.float32, device=grouped.device)
|
| for b in range(n_bins):
|
| mask_b = (bin_idx == b)
|
| if mask_b.any():
|
| max_b = (abs_g * mask_b).amax(dim=2)
|
| bin_scales[..., b] = torch.where(mask_b.any(dim=2), max_b / n,
|
| torch.full_like(max_b, 1e-8))
|
| scale_sel = torch.gather(bin_scales.unsqueeze(2).expand(out_f, num_groups, gs, n_bins), 3, bin_idx.unsqueeze(-1).long()).squeeze(-1)
|
| weight_int = torch.clamp(torch.round(grouped / scale_sel), min=self.min_val, max=self.max_val).to(torch.int8)
|
| weight_int = weight_int.reshape(out_f, in_padded)
|
| meta = {
|
| "value_repr": "int", "value_bits": self.value_bits,
|
| "scale_mode": self.scale_mode, "group_mode": "magnitude-binned",
|
| "group_size": gs, "n_bins": n_bins, "symmetric": self.symmetric,
|
| "scale_dtype": "fp32",
|
| "in_features": in_f, "in_padded": in_padded, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "has_rotation": Q is not None,
|
| }
|
| buffers = {
|
| "weight_int": weight_int,
|
| "bin_idx": bin_idx.reshape(out_f, in_padded).to(torch.int8),
|
| "bin_scales": bin_scales.to(torch.float32),
|
| }
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
|
|
| if self.scale_mode == "per-head" and self.head_dim > 0:
|
| hd = self.head_dim
|
| pad = (hd - (in_f % hd)) % hd
|
| if pad > 0:
|
| flat = torch.nn.functional.pad(flat, (0, pad))
|
| in_padded = flat.shape[1]
|
| num_heads_eff = in_padded // hd
|
| grouped = flat.reshape(out_f, num_heads_eff, hd)
|
| max_abs = grouped.abs().amax(dim=2).clamp(min=1e-8)
|
| scale = max_abs / n
|
| scale_exp = scale.unsqueeze(2).expand_as(grouped)
|
| weight_int = torch.clamp(torch.round(grouped / scale_exp), min=self.min_val, max=self.max_val).to(torch.int8)
|
| weight_int = weight_int.reshape(out_f, in_padded)
|
|
|
| elif self.scale_mode in ("per-group", "per-block") and self.group_size > 0:
|
| gs = self.group_size
|
| pad = (gs - (in_f % gs)) % gs
|
| if pad > 0:
|
| flat = torch.nn.functional.pad(flat, (0, pad))
|
| in_padded = flat.shape[1]
|
| num_groups = in_padded // gs
|
| grouped = flat.reshape(out_f, num_groups, gs)
|
| max_abs = grouped.abs().amax(dim=2).clamp(min=1e-8)
|
| scale = max_abs / n
|
| scale_exp = scale.unsqueeze(2).expand_as(grouped)
|
| weight_int = torch.clamp(torch.round(grouped / scale_exp), min=self.min_val, max=self.max_val).to(torch.int8)
|
| weight_int = weight_int.reshape(out_f, in_padded)
|
| elif self.scale_mode == "per-channel" or (self.scale_mode in ("per-group", "per-block") and self.group_size == 0):
|
| if flat.dim() > 1:
|
| reduce_dims = tuple(range(1, flat.dim()))
|
| max_abs = flat.abs().amax(dim=reduce_dims).clamp(min=1e-8)
|
| else:
|
| max_abs = flat.abs().clamp(min=1e-8)
|
| scale = max_abs / n
|
| reshape = [1] * flat.dim()
|
| reshape[0] = flat.shape[0]
|
| weight_int = torch.clamp(torch.round(flat / scale.reshape(reshape)), min=self.min_val, max=self.max_val).to(torch.int8)
|
| in_padded = in_f
|
| else:
|
| max_abs = flat.abs().amax().clamp(min=1e-8)
|
| scale = (max_abs / n).reshape(1)
|
| weight_int = torch.clamp(torch.round(flat / scale), min=self.min_val, max=self.max_val).to(torch.int8)
|
| in_padded = in_f
|
|
|
| meta = {
|
| "value_repr": "int", "value_bits": self.value_bits, "scale_mode": self.scale_mode,
|
| "group_size": self.group_size, "symmetric": self.symmetric,
|
| "scale_dtype": getattr(self, "scale_dtype", "fp32"),
|
| "head_dim": self.head_dim, "num_heads": self.num_heads,
|
| "in_features": in_f, "in_padded": in_padded, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "has_rotation": Q is not None,
|
| }
|
|
|
| if self.scale_dtype == "e8m0" and self.scale_mode in ("per-group", "per-block"):
|
| log2_s = torch.log2(scale.clamp(min=1e-38))
|
| e8m0_code = torch.round(log2_s).to(torch.int32) + 127
|
| e8m0_code = e8m0_code.clamp(0, 255).to(torch.uint8)
|
| buffers = {"weight_int": weight_int, "scale_e8m0": e8m0_code}
|
| elif self.scale_dtype == "fp8_e4m3" and self.scale_mode in ("per-group", "per-block"):
|
| weight_scale_2 = scale.amax().clamp(min=1e-12).reshape(1)
|
| scale_norm = scale / weight_scale_2
|
| fp8_codes, _ = quantize_fp8(scale_norm, FP8_E4M3_LUT, scale=None)
|
| buffers = {"weight_int": weight_int, "scale_fp8": fp8_codes.to(torch.uint8),
|
| "weight_scale_2": weight_scale_2.to(torch.float32)}
|
| else:
|
| buffers = {"weight_int": weight_int, "scale": scale.to(torch.float32)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
|
|
| if self.error_compensation == "gptq-hessian":
|
| self._apply_gptq_compensation_simple(weight_int, scale, flat, in_f, in_padded)
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _apply_gptq_compensation_simple(self, weight_int, scale, W_flat, in_features, in_padded):
|
| """GPTQ data-free greedy column compensation (H≈I).
|
|
|
| For each column i: quantize → compute error → push error to column i+1.
|
| This is simpler than full GPTQ (which uses Hessian) but still distributes
|
| quantization error across subsequent columns.
|
| """
|
| n = self.n_levels
|
| min_v, max_v = self.min_val, self.max_val
|
| W_work = W_flat.clone()
|
|
|
| for i in range(min(in_features, weight_int.shape[1])):
|
| col = W_work[:, i]
|
| if scale.numel() == 1:
|
| s_i = scale
|
| elif scale.dim() == 1:
|
|
|
| s_i = scale
|
| else:
|
|
|
| gs = self.group_size if self.group_size > 0 else 1
|
| group_idx = i // gs
|
| s_i = scale[:, group_idx] if scale.dim() > 1 else scale
|
|
|
| if s_i.numel() == 1:
|
| q_col = torch.clamp(torch.round(col / s_i), min_v, max_v)
|
| deq_col = q_col.to(torch.float32) * s_i
|
| else:
|
| q_col = torch.clamp(torch.round(col / s_i), min_v, max_v)
|
| deq_col = q_col.to(torch.float32) * s_i
|
| weight_int[:, i] = q_col.to(torch.int8)
|
| err = col - deq_col
|
|
|
| if i + 1 < weight_int.shape[1]:
|
| W_work[:, i + 1] -= err
|
|
|
| def _quantize_nf4(self, W, original_shape, Q=None):
|
| """NF4 LUT quantization."""
|
| if W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
| gs = self.group_size if self.group_size > 0 else 64
|
| idx, scale = quantize_nf4(flat, group_size=gs)
|
| weight_packed = pack_nf4(idx)
|
| buffers = {"weight_packed": weight_packed}
|
| meta = {
|
| "value_repr": "nf4_lut", "group_size": gs, "in_features": in_f,
|
| "in_padded": idx.shape[1], "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "use_double_quant": self.double_quant, "has_rotation": Q is not None,
|
| }
|
| if self.double_quant:
|
| scale_packed, block_scale = double_quantize_scales_2d(scale, block_size=self.block_size)
|
| buffers["scale_packed"] = scale_packed
|
| buffers["block_scale"] = block_scale
|
| meta["block_size"] = self.block_size
|
| else:
|
| buffers["scale"] = scale.to(torch.float32)
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_fp4(self, W, original_shape, Q=None):
|
| """FP4 E2M1 quantization (NVFP4 or MXFP4 style)."""
|
| if W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
| gs = self.group_size if self.group_size > 0 else 16
|
| idx, scale_fp32 = quantize_fp4(flat, group_size=gs)
|
| weight_packed = pack_fp4(idx)
|
| meta = {
|
| "value_repr": "fp4_e2m1", "group_size": gs, "in_features": in_f,
|
| "in_padded": idx.shape[1], "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "scale_dtype": self.scale_dtype, "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_packed": weight_packed}
|
| if self.scale_dtype == "fp8_e4m3":
|
|
|
| max_lut = FP4_E2M1_LUT.abs().amax().clamp(min=1e-12)
|
| weight_scale_2 = (scale_fp32.amax() / max_lut).clamp(min=1e-12).reshape(1)
|
| scale_norm = scale_fp32 / weight_scale_2
|
| fp8_codes, _ = quantize_fp8(scale_norm, FP8_E4M3_LUT, scale=None)
|
| buffers["scale_fp8"] = fp8_codes.to(torch.uint8)
|
| buffers["weight_scale_2"] = weight_scale_2.to(torch.float32)
|
| elif self.scale_dtype == "e8m0":
|
|
|
| log2_s = torch.log2(scale_fp32.clamp(min=1e-38))
|
| e8m0_code = torch.round(log2_s).to(torch.int32) + 127
|
| e8m0_code = e8m0_code.clamp(0, 255).to(torch.uint8)
|
| buffers["scale_e8m0"] = e8m0_code
|
| else:
|
| buffers["scale"] = scale_fp32.to(torch.float32)
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_fp6(self, W, original_shape, Q=None):
|
| """FP6 (E3M2/E2M3) quantization with optional E8M0 or FP8 block scale."""
|
| lut = self._get_lut()
|
| if W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| elif W.dim() == 1:
|
| flat = W.reshape(1, -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
| gs = self.group_size if self.group_size > 0 else 32
|
| codes, scale_fp32 = quantize_fp6(flat, lut, group_size=gs)
|
| packed = pack_fp6(codes)
|
| meta = {
|
| "value_repr": self.value_repr, "group_size": gs, "in_features": in_f,
|
| "in_padded": codes.shape[1], "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "scale_dtype": self.scale_dtype, "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_packed": packed}
|
| if self.scale_dtype == "fp8_e4m3":
|
| max_lut = lut.abs().amax().clamp(min=1e-12)
|
| weight_scale_2 = (scale_fp32.amax() / max_lut).clamp(min=1e-12).reshape(1)
|
| scale_norm = scale_fp32 / weight_scale_2
|
| fp8_codes, _ = quantize_fp8(scale_norm, FP8_E4M3_LUT, scale=None)
|
| buffers["scale_fp8"] = fp8_codes.to(torch.uint8)
|
| buffers["weight_scale_2"] = weight_scale_2.to(torch.float32)
|
| elif self.scale_dtype == "e8m0":
|
| log2_s = torch.log2(scale_fp32.clamp(min=1e-38))
|
| e8m0_code = torch.round(log2_s).to(torch.int32) + 127
|
| e8m0_code = e8m0_code.clamp(0, 255).to(torch.uint8)
|
| buffers["scale_e8m0"] = e8m0_code
|
| else:
|
| buffers["scale"] = scale_fp32.to(torch.float32)
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_fp8(self, W, original_shape, Q=None):
|
| """FP8 (E4M3/E5M2) quantization.
|
|
|
| Supports:
|
| - per-tensor / per-channel fp32 scale (standard FP8)
|
| - per-group E8M0 block scale (MXFP8)
|
| - per-group FP8 E4M3 scale (NVFP8)
|
| """
|
| lut = self._get_lut()
|
| if W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| elif W.dim() == 1:
|
| flat = W.reshape(1, -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
|
|
|
|
| if self.scale_dtype in ("e8m0", "fp8_e4m3") and self.group_size > 0:
|
| gs = self.group_size
|
| pad = (gs - (in_f % gs)) % gs
|
| if pad > 0:
|
| flat = torch.nn.functional.pad(flat, (0, pad))
|
| in_padded = flat.shape[1]
|
| num_groups = in_padded // gs
|
| grouped = flat.reshape(out_f, num_groups, gs)
|
| max_lut = lut.abs().amax().clamp(min=1e-12)
|
| scale_fp32 = (grouped.abs().amax(dim=2) / max_lut).clamp(min=1e-8)
|
| scale_exp = scale_fp32.unsqueeze(2).expand_as(grouped)
|
| w_norm = grouped / scale_exp
|
| w_norm = w_norm.clamp(-lut.abs().amax().item(), lut.abs().amax().item())
|
| diff = w_norm.unsqueeze(-1) - lut.to(w_norm.device)
|
| codes = diff.abs().argmin(dim=-1).to(torch.uint8)
|
| codes = codes.reshape(out_f, in_padded)
|
| packed = codes.view(torch.int8)
|
| meta = {
|
| "value_repr": self.value_repr, "group_size": gs,
|
| "in_features": in_f, "in_padded": in_padded, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "scale_mode": "per-group", "scale_dtype": self.scale_dtype,
|
| "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_packed": packed}
|
| if self.scale_dtype == "e8m0":
|
| log2_s = torch.log2(scale_fp32.clamp(min=1e-38))
|
| e8m0_code = torch.round(log2_s).to(torch.int32) + 127
|
| e8m0_code = e8m0_code.clamp(0, 255).to(torch.uint8)
|
| buffers["scale_e8m0"] = e8m0_code
|
| else:
|
| weight_scale_2 = scale_fp32.amax().clamp(min=1e-12).reshape(1)
|
| scale_norm = scale_fp32 / weight_scale_2
|
| fp8_codes, _ = quantize_fp8(scale_norm, FP8_E4M3_LUT, scale=None)
|
| buffers["scale_fp8"] = fp8_codes.to(torch.uint8)
|
| buffers["weight_scale_2"] = weight_scale_2.to(torch.float32)
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
|
|
| if self.scale_mode == "per-tensor":
|
| max_lut = lut.abs().amax().clamp(min=1e-12)
|
| scale = (flat.abs().amax() / max_lut).clamp(min=1e-12).reshape(1)
|
| else:
|
| scale = None
|
| codes, scale = quantize_fp8(flat, lut, scale)
|
| packed = codes.view(torch.int8)
|
| meta = {
|
| "value_repr": self.value_repr, "in_features": in_f,
|
| "out_features": out_f, "original_shape": original_shape,
|
| "ndim": len(original_shape), "scale_mode": self.scale_mode,
|
| "scale_dtype": "fp32", "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_packed": packed, "scale": scale.to(torch.float32)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_codebook(self, W, original_shape, Q=None):
|
| """Codebook/VQ quantization with optional residual levels.
|
|
|
| When residual_levels > 1, a cascade of codebooks is built:
|
| level 1: quantize w_norm → residual = w_norm - dequant(1)
|
| level 2: quantize residual → residual2 = residual - dequant(2)
|
| ...
|
| dequant = sum of all levels.
|
| """
|
| if W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| elif W.dim() == 1:
|
| flat = W.reshape(1, -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
| K = self.codebook_size
|
| n_levels_res = self.residual_levels if self.residual_levels > 1 else 1
|
| K_next = self.residual_codebook_size if self.residual_codebook_size > 0 else K
|
|
|
| if self.scale_mode == "per-channel":
|
| scale = flat.abs().amax(dim=1).clamp(min=1e-8)
|
| w_norm = flat / scale.unsqueeze(1)
|
| else:
|
| scale = flat.abs().amax().clamp(min=1e-8).reshape(1)
|
| w_norm = flat / scale
|
|
|
|
|
| if self.codebook_source == "vq":
|
| gs = self.vq_group_size
|
| pad = (gs - (in_f % gs)) % gs
|
| if pad > 0:
|
| w_norm = torch.nn.functional.pad(w_norm, (0, pad))
|
| in_padded = w_norm.shape[1]
|
| num_vectors = in_padded // gs
|
| channel_scale = flat.abs().amax(dim=1).clamp(min=1e-8)
|
| vectors_norm = (flat / channel_scale.unsqueeze(1)).reshape(out_f, num_vectors, gs)
|
| all_vecs = vectors_norm.reshape(-1, gs)
|
| codebook = self._kmeans_vectors(all_vecs, K, gs)
|
| diff = vectors_norm.unsqueeze(2) - codebook.unsqueeze(0).unsqueeze(0)
|
| dist = (diff ** 2).sum(dim=3)
|
| indices = dist.argmin(dim=2).to(torch.int32)
|
| meta = {
|
| "value_repr": "codebook", "codebook_source": "vq",
|
| "vq_group_size": gs, "K": K, "residual_levels": 1,
|
| "in_features": in_f, "in_padded": in_padded, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "has_rotation": Q is not None,
|
| }
|
| buffers = {"indices": indices, "codebook": codebook.to(torch.float32),
|
| "scale": channel_scale.to(torch.float32)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
|
|
| codebooks_list = []
|
| indices_list = []
|
| residual = w_norm.clone()
|
| for lvl in range(n_levels_res):
|
| K_lvl = K if lvl == 0 else K_next
|
| if self.codebook_source == "kmeans":
|
| cb = self._kmeans_1d(residual.flatten(), K_lvl)
|
| else:
|
| cb = torch.linspace(-1, 1, K_lvl, dtype=torch.float32)
|
| diff = residual.unsqueeze(2) - cb.unsqueeze(0).unsqueeze(0)
|
| idx = diff.abs().argmin(dim=2).to(torch.int32)
|
|
|
| deq_lvl = cb.to(residual.device)[idx.long()]
|
| residual = residual - deq_lvl
|
| codebooks_list.append(cb.to(torch.float32))
|
| indices_list.append(idx)
|
|
|
| if n_levels_res == 1:
|
|
|
| meta = {
|
| "value_repr": "codebook", "codebook_source": self.codebook_source,
|
| "K": K, "residual_levels": 1,
|
| "in_features": in_f, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "scale_mode": self.scale_mode, "has_rotation": Q is not None,
|
| }
|
| buffers = {"indices": indices_list[0],
|
| "codebook": codebooks_list[0],
|
| "scale": scale.to(torch.float32)}
|
| else:
|
|
|
| max_K = max(cb.shape[0] for cb in codebooks_list)
|
| codebooks_padded = torch.zeros(n_levels_res, max_K, dtype=torch.float32)
|
| codebook_sizes = torch.zeros(n_levels_res, dtype=torch.int32)
|
| for lvl, cb in enumerate(codebooks_list):
|
| codebooks_padded[lvl, :cb.shape[0]] = cb
|
| codebook_sizes[lvl] = cb.shape[0]
|
| meta = {
|
| "value_repr": "codebook", "codebook_source": self.codebook_source,
|
| "K": K, "residual_levels": n_levels_res,
|
| "K_next": K_next, "max_K": max_K,
|
| "in_features": in_f, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "scale_mode": self.scale_mode, "has_rotation": Q is not None,
|
| }
|
| buffers = {
|
| "indices_per_level": torch.stack(indices_list, dim=0),
|
| "codebooks": codebooks_padded,
|
| "codebook_sizes": codebook_sizes,
|
| "scale": scale.to(torch.float32),
|
| }
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| @staticmethod
|
| def _kmeans_1d(data: torch.Tensor, K: int, n_iters: int = 20) -> torch.Tensor:
|
| """1D k-means: return K cluster centers sorted ascending."""
|
| data = data.detach().float().flatten()
|
| if data.numel() == 0:
|
| return torch.zeros(K, dtype=torch.float32)
|
| quantiles = torch.linspace(0, 1, K + 1, device=data.device)[1:-1]
|
| centers = torch.quantile(data, quantiles).to(torch.float32)
|
| if centers.numel() < K:
|
| centers = torch.linspace(data.min(), data.max(), K, dtype=torch.float32, device=data.device)
|
| for _ in range(n_iters):
|
| diff = data.unsqueeze(1) - centers.unsqueeze(0)
|
| assign = diff.abs().argmin(dim=1)
|
| for k in range(K):
|
| mask = assign == k
|
| if mask.any():
|
| centers[k] = data[mask].mean()
|
| centers, _ = torch.sort(centers)
|
| return centers.to(torch.float32)
|
|
|
| @staticmethod
|
| def _kmeans_vectors(data: torch.Tensor, K: int, dim: int, n_iters: int = 20) -> torch.Tensor:
|
| """K-means on D-dimensional vectors. Returns [K, dim] centers."""
|
| N = data.shape[0]
|
| if N == 0:
|
| return torch.zeros(K, dim, dtype=torch.float32)
|
| idx = torch.randperm(N)[:K].to(data.device)
|
| centers = data[idx].clone().to(torch.float32)
|
| if centers.shape[0] < K:
|
| extra = torch.randn(K - centers.shape[0], dim, dtype=torch.float32) * 0.01
|
| centers = torch.cat([centers, extra])
|
| for _ in range(n_iters):
|
| diff = data.unsqueeze(1) - centers.unsqueeze(0)
|
| dist = (diff ** 2).sum(dim=2)
|
| assign = dist.argmin(dim=1)
|
| for k in range(K):
|
| mask = assign == k
|
| if mask.any():
|
| centers[k] = data[mask].mean(dim=0)
|
| return centers
|
|
|
| def _quantize_binary(self, W, original_shape, Q=None):
|
| """Binary {-1, +1} quantization."""
|
| if self.scale_mode == "per-tensor":
|
| scale = W.abs().mean().clamp(min=1e-8).reshape(1)
|
| else:
|
| reduce_dims = tuple(range(1, W.dim())) if W.dim() > 1 else ()
|
| if reduce_dims:
|
| scale = W.abs().mean(dim=reduce_dims).clamp(min=1e-8)
|
| else:
|
| scale = W.abs().mean().clamp(min=1e-8).reshape(1)
|
| binary = torch.sign(W).to(torch.int8)
|
| binary = torch.where(binary == 0, torch.tensor(1, dtype=torch.int8), binary)
|
| meta = {
|
| "value_repr": "binary", "scale_mode": self.scale_mode,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_binary": binary, "scale": scale.to(torch.float32)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_ternary(self, W, original_shape, Q=None):
|
| """Ternary {-1, 0, +1} quantization (BitNet 1.58)."""
|
| ternary, scale = ternarize_tensor(W, scale_mode=self.scale_mode)
|
| meta = {
|
| "value_repr": "ternary", "scale_mode": self.scale_mode,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_ternary": ternary.to(torch.int8), "scale": scale.to(torch.float32)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_none(self, W, original_shape, Q=None):
|
| """Passthrough — no quantization (fp16/fp32 weights)."""
|
| meta = {
|
| "value_repr": "none", "original_shape": original_shape,
|
| "ndim": len(original_shape), "has_rotation": Q is not None,
|
| }
|
| buffers = {"weight_fp": W.to(torch.float32)}
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_outlier(self, W, original_shape, Q=None):
|
| """Outlier-aware mixed precision (LLM.int8 / SpQR).
|
|
|
| Handles 1D weights (LayerNorm/Embedding-bias) by treating them as a
|
| single output row, consistent with _quantize_int / _quantize_codebook.
|
| """
|
| if W.dim() == 1:
|
| flat = W.reshape(1, -1)
|
| elif W.dim() > 2:
|
| flat = W.reshape(W.shape[0], -1)
|
| else:
|
| flat = W
|
| out_f, in_f = flat.shape
|
| if self.outlier_threshold is not None:
|
| threshold = self.outlier_threshold
|
| else:
|
| abs_w = flat.abs()
|
| threshold = (abs_w.mean() + 3 * abs_w.std()).item()
|
| outlier_mask = flat.abs() > threshold
|
| W_dense = flat.clone()
|
| W_dense[outlier_mask] = 0
|
| if self.scale_mode == "per-channel" and flat.dim() > 1:
|
| max_abs = W_dense.abs().amax(dim=1).clamp(min=1e-8)
|
| else:
|
| max_abs = W_dense.abs().amax().clamp(min=1e-8)
|
| scale = max_abs / self.n_levels
|
| if self.scale_mode == "per-channel" and flat.dim() > 1:
|
| dense_int = torch.clamp(torch.round(W_dense / scale.unsqueeze(1)), min=self.min_val, max=self.max_val).to(torch.int8)
|
| else:
|
| dense_int = torch.clamp(torch.round(W_dense / scale), min=self.min_val, max=self.max_val).to(torch.int8)
|
| outlier_flat = outlier_mask.flatten()
|
| outlier_indices = torch.where(outlier_flat)[0].to(torch.int64)
|
| if outlier_indices.numel() > 0:
|
| outlier_values = flat.flatten()[outlier_indices].to(torch.float16)
|
| else:
|
| outlier_values = torch.zeros(0, dtype=torch.float16)
|
| meta = {
|
| "value_repr": "outlier", "value_bits": self.value_bits,
|
| "threshold": threshold, "in_features": in_f, "out_features": out_f,
|
| "original_shape": original_shape, "ndim": len(original_shape),
|
| "num_outliers": outlier_indices.numel(), "has_rotation": Q is not None,
|
| }
|
| buffers = {
|
| "dense_int": dense_int,
|
| "scale": scale.to(torch.float32) if self.scale_mode == "per-channel" else scale.to(torch.float32).reshape(1),
|
| "outlier_indices": outlier_indices, "outlier_values": outlier_values,
|
| }
|
| if Q is not None:
|
| buffers["rotation_Q"] = Q
|
| return QuantizedWeight(weight_buffers=buffers, weight_meta=meta)
|
|
|
| def _quantize_prune(self, W, original_shape):
|
| """Structural pruning (magnitude/ratio/structured)."""
|
| if self.prune_mode == "magnitude":
|
| threshold = self.outlier_threshold if self.outlier_threshold is not None else 0.01
|
| mask = W.abs() > threshold
|
| elif self.prune_mode == "ratio":
|
| abs_w = W.abs().flatten()
|
| k = int(abs_w.numel() * self.prune_ratio)
|
| if k > 0:
|
| threshold_val = torch.kthvalue(abs_w, k).values.item()
|
| else:
|
| threshold_val = 0.0
|
| mask = W.abs() > threshold_val
|
| elif self.prune_mode == "structured":
|
| if W.dim() > 1:
|
| channel_mag = W.abs().mean(dim=tuple(range(1, W.dim())))
|
| k = int(W.shape[0] * self.prune_ratio)
|
| if k > 0:
|
| threshold_val = torch.kthvalue(channel_mag, k).values.item()
|
| else:
|
| threshold_val = 0.0
|
| channel_mask = channel_mag > threshold_val
|
| mask = channel_mask.unsqueeze(1).expand_as(W).to(torch.bool)
|
| else:
|
| mask = W.abs() > 0
|
| else:
|
| mask = torch.ones_like(W, dtype=torch.bool)
|
| pruned = W * mask.to(torch.float32)
|
| meta = {
|
| "value_repr": "prune", "prune_mode": self.prune_mode,
|
| "prune_ratio": self.prune_ratio, "original_shape": original_shape,
|
| "ndim": len(original_shape), "sparsity": float((~mask).float().mean().item()),
|
| }
|
| return QuantizedWeight(
|
| weight_buffers={"weight_pruned": pruned.to(torch.float32), "mask": mask.to(torch.uint8)},
|
| weight_meta=meta,
|
| )
|
|
|
| def _apply_awq_scaling(self, W: torch.Tensor) -> torch.Tensor:
|
| """AWQ: amplify salient channels before quantization."""
|
| out_f, in_f = W.shape if W.dim() > 1 else (W.shape[0], W.numel())
|
| act_scale = W.abs().mean(dim=0).clamp(min=1e-8) if W.dim() > 1 else W.abs().mean().clamp(min=1e-8).reshape(1)
|
| w_scale = W.abs().mean(dim=0).clamp(min=1e-8) if W.dim() > 1 else W.abs().mean().clamp(min=1e-8).reshape(1)
|
| salience = act_scale * w_scale
|
| s = (salience / salience.mean().clamp(min=1e-8)).clamp(0.5, 2.0)
|
| return W * s.unsqueeze(0) if W.dim() > 1 else W * s |