Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

39. Quantized checkpoints: GPTQ, AWQ, FP8, FP4 and trellis codes

In this chapter

  • GPTQ: quantizing a layer column by column while correcting the remaining columns for each error, using the inverse Hessian of the layer's output error.
  • AWQ: scaling up the weights that large activations multiply before quantizing, and folding the scale into the previous operation.
  • The checkpoint formats behind "quant_method": "gptq" | "awq" | "fp8", read and written, and served through the same layer as GGUF.
  • Low-precision floating point: block-scaled FP8, and the FP4 microscaling formats MXFP4 and NVFP4.
  • Trellis-coded quantization (QTIP, ExLlamaV3's EXL3): coding whole sequences of weights with a Viterbi search, and why it beats any scalar code at 2 bits.

You will build

gptq_quantize and unpack_gptq (engine/formats/gptq.py), search_scale, apply_awq and unpack_awq (awq.py), quantize_blocks (fp8.py), quantize_mxfp4 (mx.py), viterbi_encode (trellis.py), load_quantized (hf_quant.py) and the FP8 kernel.

Time: 8-10 hours. GPU: not needed.

Beyond round-to-nearest

Chapter 20 quantized each weight to the nearest point of a grid, with one scale per group, and noted that the methods people actually ship do better. This chapter builds them. They answer three different questions:

questionmethods
Given the grid, which grid point should each weight get? Not always the nearest.GPTQ
Can we change the problem so that rounding hurts less?AWQ (scale salient channels), QuIP#/QTIP (rotate), SmoothQuant
Is a uniform integer grid even the right code?FP8, MXFP4, NVFP4 (floating-point grids); trellis codes (code sequences, not values)

And for each, a checkpoint format that the engine must read, because users download quantized checkpoints far more often than they quantize models themselves. All of them, like GGUF, repack into Chapter 38’s AffineQuantLinear (code × scale + offset per group), except the floating-point formats, which get their own layer.

Throughout, the measure of damage is the output, not the weights: Chapter 20’s choose_clip already minimized output error on calibration activations, and that idea is the seed of everything here.

GPTQ

Quantize a layer’s weight $W$ (rows are outputs, columns are inputs) to minimize the squared error of its output on calibration inputs $X$:

$$ \min_{\hat W} ; \lVert X \hat W^\top - X W^\top \rVert^2 . $$

The error is a quadratic function of each row of $\hat W - W$, with the same Hessian for every row: $H = 2 X^\top X$. Its diagonal says how strongly each input is activated; its off-diagonal entries say which inputs move together. That second part is what round-to-nearest ignores. If inputs $i$ and $j$ are correlated, an error on weight $i$ can be partly cancelled by adjusting weight $j$.

GPTQ (Frantar et al., 2022), built on Optimal Brain Quantization, quantizes the columns one at a time and, after each, updates every column not yet quantized to absorb the error:

$$ e = \frac{w_i - \text{quant}(w_i)}{[H^{-1}]{ii}}, \qquad w{j} \leftarrow w_j - e , [H^{-1}]_{ij} \quad \text{for every later } j . $$

The inverse Hessian is computed once, through a Cholesky factorization (with a little damping on the diagonal for stability), and the updates are applied lazily: inside a block of 128 columns each error updates the block, and the rest of the matrix is updated with one matmul per block. That makes GPTQ fast enough to quantize a 70-billion-parameter model on one GPU in a few hours.

@torch.no_grad()
def gptq_quantize(weight, hessian, bits=4, group_size=128, act_order=False, damp=0.01, block=128):
    """weight [out, in], hessian [in, in] = 2 X^T X / n -> (codes [out, in], scales, zeros [out, groups], g_idx [in]).
    (Your engine: Chapter 39)

    Columns are processed in blocks of `block`; inside a block every column's error is pushed
    onto the block's remaining columns at once, and onto the columns after the block in one
    matmul when the block is done ("lazy batch updates"). With act_order, columns are visited
    from the largest Hessian diagonal (the most activated input) down.
    """
    W = weight.float().clone()
    H = hessian.float().clone()
    cols = W.shape[1]
    dead = torch.diag(H) == 0                     # inputs that are never active: their weights don't matter
    H[dead, dead] = 1
    W[:, dead] = 0
    perm = torch.argsort(torch.diag(H), descending=True) if act_order else torch.arange(cols)
    W, H = W[:, perm], H[perm][:, perm]
    H += damp * torch.diag(H).mean() * torch.eye(cols)
    Hinv = torch.linalg.cholesky(torch.cholesky_inverse(torch.linalg.cholesky(H)), upper=True)
    codes = torch.zeros_like(W)
    groups = -(-cols // group_size)
    scales, zeros = torch.zeros(W.shape[0], groups), torch.zeros(W.shape[0], groups)
    top = 2 ** bits - 1
    for i1 in range(0, cols, block):
        i2 = min(i1 + block, cols)
        W1, Err = W[:, i1:i2].clone(), torch.zeros(W.shape[0], i2 - i1)
        H1 = Hinv[i1:i2, i1:i2]
        for i in range(i2 - i1):
            col = i1 + i
            if col % group_size == 0:              # a new group: its range from the CURRENT (corrected) weights
                g = col // group_size
                ahead = torch.cat((W1[:, i:], W[:, i2:]), 1)[:, :group_size]
                scales[:, g], zeros[:, g] = find_params(ahead, bits)
            g = col // group_size
            q = (W1[:, i] / scales[:, g] + zeros[:, g]).round().clamp(0, top)
            codes[:, col] = q
            error = (W1[:, i] - scales[:, g] * (q - zeros[:, g])) / H1[i, i]
            W1[:, i:] -= error[:, None] * H1[i, i:][None, :]           # correct the rest of the block
            Err[:, i] = error
        W[:, i2:] -= Err @ Hinv[i1:i2, i2:]                             # and everything after it
    inverse = torch.argsort(perm)
    g_idx = (torch.arange(cols) // group_size)[inverse]                 # group of each ORIGINAL column
    return codes[:, inverse].to(torch.uint8), scales, zeros, g_idx.to(torch.int32)

Two details matter in practice:

  • Group ranges come from the corrected weights. When the column loop enters a new group of 128, the group’s scale and zero point are computed from the weights as GPTQ has already adjusted them, not from the original ones.
  • Act-order (desc_act in configs) visits columns from the most activated (largest $H_{ii}$) to the least, so the important columns are quantized while the most freedom remains to compensate them. It helps quality, at a cost the format must carry: groups no longer cover contiguous columns, so the checkpoint stores g_idx, each input column’s group.

The format

qweight  int32 [in / 8, out]      eight 4-bit codes along the INPUT dimension per int32, low bits first
qzeros   int32 [groups, out / 8]  zero points along the OUTPUT dimension, stored minus one
scales   f16   [groups, out]
g_idx    int32 [in]
W[o, i] = scales[g_idx[i], o] * (q[i, o] - (qzeros[g_idx[i], o] + 1))

The “stored minus one” is a historical quirk of AutoGPTQ’s first packing code that every reader must replicate (newer tools call the corrected layout gptq_v2). To serve it, to_runtime sorts the input columns by g_idx once at load time, so that groups become contiguous, and permutes the activations to match at run time. ExLlama’s kernels do the same, permuting the previous layer’s output instead, which is free.

def pack_gptq(codes, scales, zeros, g_idx):
    """Our [out, in] codes -> the checkpoint tensors (4-bit)."""
    q = codes.T.to(torch.int64)                                   # [in, out]
    qweight = torch.zeros(q.shape[0] // 8, q.shape[1], dtype=torch.int64)
    for j in range(8):
        qweight |= q[j::8] << (4 * j)
    z = (zeros.T.to(torch.int64) - 1) & 0xF                       # [groups, out], stored minus 1
    qzeros = torch.zeros(z.shape[0], z.shape[1] // 8, dtype=torch.int64)
    for j in range(8):
        qzeros |= z[:, j::8] << (4 * j)
    return {"qweight": as_int32(qweight), "qzeros": as_int32(qzeros), "scales": scales.T.to(torch.float16).contiguous(),
            "g_idx": g_idx.to(torch.int32)}


def as_int32(t):
    """32 packed bits held in an int64 -> int32 with the same bits (two's complement)."""
    return ((t + 2 ** 31) % 2 ** 32 - 2 ** 31).to(torch.int32)


def unpack_gptq(qweight, qzeros, scales, g_idx):
    """Checkpoint tensors -> (codes [out, in], scales [out, groups], zeros [out, groups], g_idx).  (Your engine: Chapter 39)"""
    shifts = torch.arange(0, 32, 4, dtype=torch.int32)
    q = (qweight[:, None, :] >> shifts[None, :, None]) & 0xF                  # [in/8, 8, out]
    codes = q.reshape(-1, qweight.shape[1]).T                                # [out, in]
    z = (qzeros[:, :, None] >> shifts[None, None, :]) & 0xF                  # [groups, out/8, 8]
    zeros = (z.reshape(qzeros.shape[0], -1) + 1).T.float()                   # stored minus 1
    return codes.to(torch.uint8), scales.T.float(), zeros, g_idx

AWQ

GPTQ corrects errors after the fact. AWQ (Lin et al., 2023) prevents some of them. Its observation: about 1% of a layer’s input channels carry activations an order of magnitude larger than the rest (Chapter 20’s outliers), and errors in the weights those channels multiply dominate the output error. Scale those weights up by $s > 1$ before quantizing: their quantization step stays the same in absolute terms (the group’s range barely changes), so their relative error shrinks. Then divide the activations by $s$, which leaves the product exactly unchanged before rounding:

$$ X W^\top = (X ,\text{diag}(s)^{-1}) , (W ,\text{diag}(s))^\top . $$

The division is free, because it folds into whatever produced $X$: an RMSNorm’s weight (for Q/K/V and gate/up, which read the norms’ outputs), or the previous linear layer’s output rows (V for O, up for down). The scale is $s = \overline{|x|}^{,\alpha}$ per input channel, with $\alpha$ chosen per group of layers by a 20-point grid search on calibration data. Since $\alpha = 0$ is plain round-to-nearest, the search can only help:

@torch.no_grad()
def search_scale(weights, inputs, bits=4, group_size=128, grid=20):
    """Best per-input-channel scale for the linears in `weights` (all reading `inputs`).  (Your engine: Chapter 39)

    For alpha in 0, 1/grid, ..., (grid-1)/grid: s = mean|x|^alpha (normalized), quantize W * s,
    undo the scale, and measure the output error on the calibration inputs. alpha = 0 is plain
    round-to-nearest, so the search can only help.
    """
    x = inputs.reshape(-1, inputs.shape[-1]).float()
    x_mean = x.abs().mean(0)
    reference = [x @ w.float().T for w in weights]
    best = (float("inf"), torch.ones_like(x_mean), 0.0)
    for step in range(grid):
        alpha = step / grid
        s = x_mean.clamp_min(1e-4) ** alpha
        s = s / (s.max() * s.min()).sqrt()
        error = 0.0
        for w, ref in zip(weights, reference):
            approx = (x / s) @ fake_quantize(w.float() * s, bits, group_size).T
            error += (approx - ref).pow(2).mean().item()
        if error < best[0]:
            best = (error, s, alpha)
    return best[1], best[2]
@torch.no_grad()
def apply_awq(model, calibration_ids, bits=4, group_size=128):
    """Scale and quantize every Qwen3 layer in place, AutoAWQ's way.  (Your engine: Chapter 39)

    Scaling pairs (previous op -> linears that read its output):
      input_layernorm -> q, k, v        post_attention_layernorm -> gate, up
      v_proj (rows)   -> o_proj         up_proj (rows)           -> down_proj
    Each layer's inputs are captured from a forward pass over the calibration tokens.
    Returns {linear name: AffineQuantLinear} after replacing the modules.
    """
    captured = {}

    def hook(name):
        def save(module, args, output):
            captured.setdefault(name, []).append(args[0].detach().reshape(-1, args[0].shape[-1]))
        return save

    handles = []
    for i, layer in enumerate(model.model.layers):
        for name in ("self_attn.q_proj", "self_attn.o_proj", "mlp.gate_proj", "mlp.down_proj"):
            handles.append(layer.get_submodule(name).register_forward_hook(hook(f"{i}.{name}")))
    model(calibration_ids)
    for h in handles:
        h.remove()
    replaced = {}
    for i, layer in enumerate(model.model.layers):
        attn, mlp = layer.self_attn, layer.mlp
        x = lambda name: torch.cat(captured[f"{i}.{name}"])                      # noqa: E731
        pairs = [(layer.input_layernorm, None, [attn.q_proj, attn.k_proj, attn.v_proj], x("self_attn.q_proj")),
                 (layer.post_attention_layernorm, None, [mlp.gate_proj, mlp.up_proj], x("mlp.gate_proj")),
                 (None, attn.v_proj, [attn.o_proj], x("self_attn.o_proj")),
                 (None, mlp.up_proj, [mlp.down_proj], x("mlp.down_proj"))]
        for norm, previous, linears, inputs in pairs:
            if previous is attn.v_proj and attn.v_proj.out_features != attn.o_proj.in_features:
                continue                                    # GQA: v's rows feed several heads' columns
            s, _ = search_scale([l.weight for l in linears], inputs, bits, group_size)
            if norm is not None:
                norm.weight.div_(s.to(norm.weight.dtype))
            else:
                previous.weight.div_(s[:, None].to(previous.weight.dtype))
            for linear in linears:
                linear.weight.mul_(s[None, :].to(linear.weight.dtype))
    for i, layer in enumerate(model.model.layers):
        for name in ("self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj", "self_attn.o_proj",
                     "mlp.gate_proj", "mlp.up_proj", "mlp.down_proj"):
            parent_name, _, child = name.rpartition(".")
            parent = layer.get_submodule(parent_name)
            linear = getattr(parent, child)
            codes, scale, offset = quantize_to_affine(linear.weight, bits, group_size)
            quantized = AffineQuantLinear(pack_nibbles(codes), scale, offset, group_size, True,
                                          dtype=linear.weight.dtype)
            setattr(parent, child, quantized)
            replaced[f"model.layers.{i}.{name}"] = quantized
    return replaced

AutoAWQ’s checkpoint packs eight codes along the output dimension in an interleaved order, 0 2 4 6 1 3 5 7, chosen to match its CUDA kernel’s register layout. That’s the kind of detail a format reader can only get from the source, and only verify with a test that pins it down (output 1’s code sits in bits 16-19):

def pack_awq(codes, zeros):
    """[rows, cols] codes 0..15 -> int32 [rows, cols / 8] in AWQ's interleaved order."""
    c = codes.to(torch.int64)
    packed = torch.zeros(c.shape[0], c.shape[1] // 8, dtype=torch.int64)
    for j, o in enumerate(AWQ_ORDER):
        packed |= c[:, o::8] << (4 * j)
    return ((packed + 2 ** 31) % 2 ** 32 - 2 ** 31).to(torch.int32)


def unpack_awq(packed):
    """(Your engine: Chapter 39)"""
    out = torch.zeros(packed.shape[0], packed.shape[1] * 8, dtype=torch.uint8)
    for j, o in enumerate(AWQ_ORDER):
        out[:, o::8] = ((packed >> (4 * j)) & 0xF).to(torch.uint8)
    return out


def to_checkpoint(layer):
    """AffineQuantLinear (from apply_awq) -> AutoAWQ GEMM tensors."""
    codes = layer.weight_codes.T                                        # [in, out]
    scales = layer.scales.T                                             # [groups, out]
    zeros = (-layer.offsets / layer.scales).round().T                   # offset = -zero * scale
    return {"qweight": pack_awq(codes, None), "qzeros": pack_awq(zeros.clamp(0, 15), None),
            "scales": scales.to(torch.float16).contiguous()}


def from_checkpoint(qweight, qzeros, scales, group_size, bias=None, dtype=torch.bfloat16):
    codes = unpack_awq(qweight).T                                       # [out, in]
    zeros = unpack_awq(qzeros).T.float()                                # [out, groups]
    s = scales.T.float()
    return AffineQuantLinear(pack_nibbles(codes), s, -zeros * s, group_size, True, bias=bias, dtype=dtype)

FP8 checkpoints

FP8 E4M3 has 3 mantissa bits and 4 exponent bits: 448 is its largest value, and its grid is relative: fine near zero, coarse near the top. With one scale per block that keeps each block’s largest value near 448, it reaches nearly BF16 quality at half the memory. DeepSeek-V3 was trained in FP8, and Qwen3 ships FP8 checkpoints, with one FP32 scale per 128 × 128 weight block (weight_scale_inv, which despite its name multiplies the stored values):

def quantize_blocks(weight, block=128):
    """[out, in] -> (fp8 weight, scale_inv [ceil(out/b), ceil(in/b)]) with one scale per block.  (Your engine: Chapter 39)"""
    out_f, in_f = weight.shape
    rows, cols = -(-out_f // block), -(-in_f // block)
    padded = F.pad(weight.float(), (0, cols * block - in_f, 0, rows * block - out_f))
    tiles = padded.reshape(rows, block, cols, block)
    amax = tiles.abs().amax(dim=(1, 3)).clamp_min(1e-12)
    scale = amax / E4M3_MAX
    q = (tiles / scale[:, None, :, None]).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn)
    return q.reshape(rows * block, cols * block)[:out_f, :in_f].contiguous(), scale


def dequantize_blocks(weight, scale_inv, block=128):
    s = scale_inv.repeat_interleave(block, 0)[:weight.shape[0]].repeat_interleave(block, 1)[:, :weight.shape[1]]
    return weight.float() * s

On Hopper and newer GPUs the matmul can run in FP8 on both sides: activations are quantized on the fly with one scale per token per 128 input channels, so each 128-wide slice of the product has one activation scale and one weight scale. w8a8_reference computes exactly what such a kernel computes, so you can measure its error on a CPU:

def quantize_activations(x, group=128):
    """Per token, per 128 input channels: the FP8 activations a W8A8 GEMM consumes."""
    n, k = x.shape
    g = x.float().reshape(n, k // group, group)
    scale = g.abs().amax(-1).clamp_min(1e-12) / E4M3_MAX
    return (g / scale[..., None]).to(torch.float8_e4m3fn).reshape(n, k), scale


def w8a8_reference(x, weight, scale_inv, block=128):
    """What an FP8 x FP8 tensor-core GEMM computes: both operands rounded to E4M3, products
    accumulated in FP32, each 128-wide slice of K rescaled by its activation and weight scales."""
    xq, xs = quantize_activations(x, block)
    out = torch.zeros(x.shape[0], weight.shape[0])
    for kb in range(x.shape[1] // block):
        a = xq[:, kb * block:(kb + 1) * block].float() * xs[:, kb:kb + 1]
        w = weight[:, kb * block:(kb + 1) * block].float() * scale_inv[:, kb].repeat_interleave(block)[:weight.shape[0], None]
        out += a @ w.T
    return out

The engine’s FP8 kernel is the portable W8A16 variant, which widens FP8 tiles to FP32 in registers and works on every GPU (and in the Triton interpreter):

@triton.jit
def fp8_block_kernel(x_ptr, w_ptr, s_ptr, y_ptr, M, N, K, s_cols,
                     BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
    """W8A16: FP8 weight tiles are widened in registers and scaled by their block's factor.  (Your engine: Chapter 39)"""
    pid_m, pid_n = tl.program_id(0), tl.program_id(1)
    rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for k0 in range(0, K, BLOCK_K):
        rk = k0 + tl.arange(0, BLOCK_K)
        x = tl.load(x_ptr + rm[:, None] * K + rk[None, :], mask=(rm[:, None] < M) & (rk[None, :] < K), other=0.0)
        w_ok = (rn[:, None] < N) & (rk[None, :] < K)
        w = tl.load(w_ptr + rn[:, None] * K + rk[None, :], mask=w_ok, other=0.0).to(tl.float32)
        s = tl.load(s_ptr + (rn[:, None] // 128) * s_cols + rk[None, :] // 128, mask=w_ok, other=0.0)
        acc += tl.dot(x.to(tl.float32), tl.trans(w * s), input_precision="ieee")
    tl.store(y_ptr + rm[:, None] * N + rn[None, :], acc.to(y_ptr.dtype.element_ty),
             mask=(rm[:, None] < M) & (rn[None, :] < N))


def fp8_block_matmul(x, weight, scale_inv, block_m=16, block_n=32, block_k=32):
    check_device(x, weight, scale_inv)
    m, k = x.shape
    n = weight.shape[0]
    y = torch.empty((m, n), device=x.device, dtype=x.dtype)
    grid = (triton.cdiv(m, block_m), triton.cdiv(n, block_n))
    fp8_block_kernel[grid](x.contiguous(), weight.contiguous(), scale_inv.float().contiguous(), y, m, n, k,
                           scale_inv.shape[1], BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k)
    return y

FP4: MXFP4 and NVFP4

Blackwell GPUs multiply 4-bit floats natively. E2M1, with 2 exponent bits and 1 mantissa bit, has only 16 values: $0, 0.5, 1, 1.5, 2, 3, 4, 6$ and their negatives. With so few, the scales do most of the work, and two formats differ exactly there:

formatblockscalebits / weightused by
MXFP4 (OCP Microscaling)32a power of two (one E8M0 byte)4.25gpt-oss’s MoE weights, GGUF
NVFP416FP8 E4M3 per block × FP32 per tensor4.5NVIDIA’s Blackwell-optimized checkpoints

A power-of-two scale can be off from the ideal by almost a factor of 2, wasting most of a binade; NVFP4’s FP8 scale is accurate to about 6% and its blocks are half as long, which is why it loses less.

def quantize_mxfp4(w, layout="interleaved"):
    """[..., K] (K % 32 == 0) -> (blocks uint8 [..., K/32, 16], scales uint8 [..., K/32]).  (Your engine: Chapter 39)

    The shared exponent puts the block's largest magnitude in E2M1's top binade:
    e = floor(log2(amax)) - 2, since E2M1's largest value is 6 = 1.5 * 2^2.
    """
    x = w.float().reshape(*w.shape[:-1], -1, 32)
    amax = x.abs().amax(-1).clamp_min(2.0 ** -126)
    exponent = (torch.floor(torch.log2(amax)) - 2).clamp(-127, 127)
    codes = to_e2m1(x / torch.exp2(exponent)[..., None])
    if layout == "interleaved":
        packed = codes[..., 0::2] | (codes[..., 1::2] << 4)
    else:
        packed = codes[..., :16] | (codes[..., 16:] << 4)
    return packed, (exponent + 127).to(torch.uint8)


def dequantize_mxfp4(blocks, scales, layout="interleaved"):
    low, high = (blocks & 0xF).long(), (blocks >> 4).long()
    if layout == "interleaved":
        codes = torch.stack((low, high), -1).flatten(-2)
    else:
        codes = torch.cat((low, high), -1)
    values = E2M1[codes] * torch.exp2(scales.float() - 127)[..., None]
    return values.flatten(-2)
def quantize_nvfp4(w):
    """[..., K] (K % 16 == 0) -> (codes uint8 [..., K/16, 8], block scales fp8 [..., K/16], tensor scale)."""
    x = w.float().reshape(*w.shape[:-1], -1, 16)
    tensor_scale = x.abs().amax().clamp_min(1e-12) / (448.0 * 6.0)        # E4M3's max x E2M1's max
    block_scale = (x.abs().amax(-1) / 6.0 / tensor_scale).clamp(max=448.0).to(torch.float8_e4m3fn)
    scale = block_scale.float() * tensor_scale
    codes = to_e2m1(torch.where(scale[..., None] > 0, x / scale[..., None].clamp_min(1e-30), torch.zeros_like(x)))
    return codes[..., 0::2] | (codes[..., 1::2] << 4), block_scale, tensor_scale


def dequantize_nvfp4(packed, block_scale, tensor_scale):
    codes = torch.stack(((packed & 0xF).long(), (packed >> 4).long()), -1).flatten(-2)
    return (E2M1[codes] * (block_scale.float() * tensor_scale)[..., None]).flatten(-2)

The two codes in each byte appear in two orders in the wild: gpt-oss interleaves them (elements $2j$ and $2j+1$), GGUF splits each block into halves (elements $j$ and $j + 16$). The test checks the second against Chapter 38’s GGUF decoder.

Trellis codes: QTIP and EXL3

At 2 bits per weight, every scalar code fails: four levels can’t represent a weight distribution well, however they’re placed. The best possible 2-bit scalar code for Gaussian data, the Lloyd-Max quantizer, still has a mean squared error of 0.118 per unit variance. Vector quantization does better by coding several weights jointly, but a codebook for 8 weights at 2 bits has $2^{16}$ entries to search and to keep in fast memory.

Trellis-coded quantization codes a long sequence with no stored codebook at all. QTIP (Tseng et al., 2024), which ExLlamaV3’s EXL3 format is built on, uses a bitshift trellis:

  • a state is an $L$-bit integer, and each weight consumes $K$ new bits: $s_{t+1} = ((s_t \ll K) ,|, b_t) \bmod 2^L$;
  • every state stands for a value, computed from its bits by a tiny hash. QTIP’s “1MAD” does one multiply-add mod $2^{32}$ and sums the four bytes of the result, and a sum of four roughly uniform bytes is roughly Gaussian;
  • a group of 256 weights is coded by an $L$-bit start state and 255 $K$-bit symbols: $K$ bits per weight.
def code_1mad(states):
    """A pseudo-Gaussian value for each L-bit state: one multiply-add mod 2^32, then the sum of
    the four bytes of the result (a sum of 4 roughly uniform bytes is roughly normal)."""
    x = (states.to(torch.int64) * 34038481 + 76625530) % (2 ** 32)
    total = sum((x >> (8 * i)) & 0xFF for i in range(4))
    return (total.double() - 510.0) / 147.8


def codebook(L):
    values = code_1mad(torch.arange(2 ** L))
    return ((values - values.mean()) / values.std()).float()            # exactly zero mean, unit variance

Which sequence of symbols reproduces the weights best? All $2^{L + 255K}$ sequences are candidates, but the cost decomposes along the path, so the best one is a shortest path through a graph of $2^L$ states per step: the Viterbi algorithm. Each state can be reached from exactly $2^K$ predecessors (those whose low $L - K$ bits equal its high bits), so each step is a vectorized minimum over a [2^L, 2^K] table:

@torch.no_grad()
def viterbi_encode(x, L=8, K=2):
    """x [T] (unit variance) -> (start state, K-bit symbols [T-1], reconstruction [T]).  (Your engine: Chapter 39)

    cost[s] = least squared error of any path ending in state s. A state s at step t can only
    follow the 2^K states p whose low L - K bits are s's high bits: p = (s >> K) + j * 2^(L-K).
    """
    table = codebook(L)
    states = torch.arange(2 ** L)
    T = x.shape[0]
    cost = (x[0] - table) ** 2                                          # any start state
    back = torch.zeros((T, 2 ** L), dtype=torch.long)
    preds = (states >> K)[:, None] + (torch.arange(2 ** K) << (L - K))[None, :]       # [S, 2^K]
    for t in range(1, T):
        options = cost[preds]                                           # [S, 2^K]
        best, which = options.min(-1)
        back[t] = preds[states, which]
        cost = best + (x[t] - table) ** 2
    path = torch.zeros(T, dtype=torch.long)
    path[-1] = cost.argmin()
    for t in range(T - 1, 0, -1):
        path[t - 1] = back[t, path[t]]
    symbols = path[1:] & (2 ** K - 1)                                   # the K new bits of each step
    return int(path[0]), symbols, table[path]

Decoding needs no search and no table in memory: a shift, an OR and the hash per weight, which a GPU computes faster than it could load a codebook entry. The test checks Viterbi’s optimality by brute force on a small trellis (every one of the $2^4 \times 2^5$ codes), and that on Gaussian data its error is well below the Lloyd-Max optimum: about 0.08 at $L = 8$ and 0.07 at $L = 10$, against Lloyd-Max’s 0.118. QTIP reports about 0.07 at $L = 16$ with its tuned codes.

The codebook is shaped for Gaussian data, and real weight rows aren’t Gaussian: they have outliers and different scales. A random Hadamard transform (QuIP#, Tseng et al., 2024) fixes that. Multiplying a group by $H D / \sqrt{n}$ (a Hadamard matrix and random signs) is orthogonal, so it’s undone exactly at run time, and it spreads every outlier across all 256 coordinates, which then look like i.i.d. Gaussian samples:

def random_hadamard(n, seed=0):
    """An orthogonal n x n matrix H D / sqrt(n): Sylvester Hadamard times random signs."""
    if n & (n - 1):
        raise ValueError("n must be a power of two")
    h = torch.ones(1, 1)
    while h.shape[0] < n:
        h = torch.cat((torch.cat((h, h), 1), torch.cat((h, -h), 1)), 0)
    signs = torch.where(torch.rand(n, generator=torch.Generator().manual_seed(seed)) < 0.5, -1.0, 1.0)
    return h * signs[None, :] / math.sqrt(n)


@torch.no_grad()
def quantize_trellis(weight, L=8, K=2, group=256, seed=0):
    """Rotate each row's groups to look Gaussian, code each group with the trellis, rotate back.
    Returns the dequantized weight (what a decoder would produce) and the bits per weight."""
    rows, cols = weight.shape
    rotation = random_hadamard(group, seed)
    x = weight.float().reshape(rows, -1, group) @ rotation.T            # incoherence processing
    scale = x.pow(2).mean(-1, keepdim=True).sqrt().clamp_min(1e-12)
    out = torch.empty_like(x)
    for r in range(rows):
        for g in range(x.shape[1]):
            _, _, recon = viterbi_encode(x[r, g] / scale[r, g], L, K)
            out[r, g] = recon * scale[r, g]
    bits = (L + K * (group - 1) + 16) / group                            # start state, symbols, an f16 scale
    return (out @ rotation).reshape(rows, cols), bits

This is the algorithm. ExLlamaV3 adds its engineering: $L = 16$, a hash called “3INST” chosen for GPU instruction throughput, 16 × 16 tiles laid out for tensor cores, per-tensor bit widths from 1 to 8 chosen by sensitivity, and its own file format. This chapter’s code doesn’t read EXL3 files; it shows why EXL3 can make a 2.5-bit model usable where round-to-nearest can’t.

Loading and writing checkpoints

All three Hugging Face formats share a structure: config.json’s quantization_config names the method, and each quantized linear layer’s tensors sit under its name. One loader streams the tensors and builds each quantized layer as soon as all of its parts have arrived:

@torch.no_grad()
def load_quantized(directory, device="cpu", dtype=torch.bfloat16):
    """A Qwen3 (dense) model from a GPTQ, AWQ or FP8 checkpoint.  (Your engine: Chapter 39)

    Tensors are streamed; a quantized layer is built as soon as all of its parts have arrived,
    then put in place of the model's nn.Linear.
    """
    from ..loaders import assign, read_config
    from ..qwen3 import Qwen3, Qwen3Config
    from ..safetensors_io import snapshot_tensors
    raw = read_config(directory)
    q = raw.get("quantization_config") or {}
    method = q.get("quant_method")
    if method not in PARTS:
        raise ValueError(f"Unsupported quant_method {method!r} (gptq, awq or fp8)")
    cfg = Qwen3Config.from_hf(raw)
    with torch.device("meta"):
        model = Qwen3(cfg)
    model = model.to_empty(device=device).to(dtype)
    if cfg.tie_word_embeddings:
        model.lm_head.weight = model.model.embed_tokens.weight
    params = dict(model.named_parameters())
    linears = {name for name, m in model.named_modules() if isinstance(m, torch.nn.Linear) and name != "lm_head"}
    pending = {}
    for name, value in snapshot_tensors(directory):
        module, _, part = name.rpartition(".")
        if module in linears and (part in PARTS[method] or part == "bias"):
            pending.setdefault(module, {})[part] = value
            if all(p in pending[module] for p in PARTS[method]):
                parts = pending.pop(module)
                layer = build(method, parts, q, dtype)
                parent, _, child = module.rpartition(".")
                setattr(model.get_submodule(parent), child, layer.to(device))
            continue
        if name == "lm_head.weight" and cfg.tie_word_embeddings:
            continue
        if name not in params:
            raise ValueError(f"Unexpected tensor {name}")
        assign(params[name], value, name)
    if pending:
        raise ValueError(f"Incomplete quantized layers: {sorted(pending)[:3]}")
    return model.eval()


def build(method, parts, q, dtype):
    bias = parts.get("bias")
    if method == "gptq":
        if q.get("bits", 4) != 4:
            raise ValueError("Only 4-bit GPTQ is implemented")
        codes, scales, zeros, g_idx = gptq.unpack_gptq(parts["qweight"], parts["qzeros"], parts["scales"], parts["g_idx"])
        return gptq.to_runtime(codes, scales, zeros, g_idx, q["group_size"], bias, dtype)
    if method == "awq":
        return awq.from_checkpoint(parts["qweight"], parts["qzeros"], parts["scales"], q["group_size"], bias, dtype)
    if q.get("weight_block_size", [128, 128]) != [128, 128]:
        raise ValueError("Only 128 x 128 FP8 blocks are implemented")
    return fp8.FP8BlockLinear(parts["weight"], parts["weight_scale_inv"], bias, dtype)

and one writer quantizes a dense model and produces the same layout, so a model you fine-tuned in Chapter 22 can ship as GPTQ, AWQ or FP8:

@torch.no_grad()
def save_quantized(model, directory, method, calibration_ids=None, group_size=128, act_order=False):
    """Quantize a dense Qwen3 and write it as a checkpoint that load_quantized (and vLLM's and
    transformers' loaders for the same method) can read."""
    from ..safetensors_io import save_file
    directory = Path(directory)
    directory.mkdir(parents=True, exist_ok=True)
    raw = {k: v for k, v in vars(model.cfg).items()}
    raw["model_type"] = "qwen3"
    tensors = {}
    inputs = {}
    if method == "gptq":
        handles = [m.register_forward_hook(lambda mod, a, o, n=n: inputs.setdefault(n, []).append(a[0].reshape(-1, a[0].shape[-1])))
                   for n, m in model.named_modules() if isinstance(m, torch.nn.Linear) and n != "lm_head"]
        model(calibration_ids)
        for h in handles:
            h.remove()
    if method == "awq":
        replaced = awq.apply_awq(model, calibration_ids, 4, group_size)
    for name, module in model.named_modules():
        if name == "lm_head" or not isinstance(module, (torch.nn.Linear, awq.AffineQuantLinear)):
            continue
        if method == "gptq":
            H = gptq.hessian_from(torch.cat(inputs[name]))
            codes, scales, zeros, g_idx = gptq.gptq_quantize(module.weight, H, 4, group_size, act_order)
            tensors.update({f"{name}.{k}": v for k, v in gptq.pack_gptq(codes, scales, zeros, g_idx).items()})
        elif method == "awq":
            tensors.update({f"{name}.{k}": v for k, v in awq.to_checkpoint(replaced[name]).items()})
        else:
            weight, scale_inv = fp8.quantize_blocks(module.weight)
            tensors[f"{name}.weight"], tensors[f"{name}.weight_scale_inv"] = weight, scale_inv
    for name, p in model.named_parameters():
        module = name.rpartition(".")[0]
        if module == "lm_head" and model.cfg.tie_word_embeddings:
            continue
        if not any(key.startswith(module + ".") for key in tensors) or module in ("lm_head", ""):
            tensors.setdefault(name, p.detach().contiguous())
    raw["quantization_config"] = ({"quant_method": "fp8", "weight_block_size": [128, 128], "activation_scheme": "dynamic"}
                                  if method == "fp8" else
                                  {"quant_method": method, "bits": 4, "group_size": group_size, "desc_act": act_order,
                                   "sym": False, **({"version": "gemm", "zero_point": True} if method == "awq" else {})})
    (directory / "config.json").write_text(json.dumps(raw))
    save_file(tensors, directory / "model.safetensors")

The tests write and reload each format and check that the result stays close to the original model. They can’t check compatibility with the reference tools themselves, AutoGPTQ, AutoAWQ and vLLM’s FP8 loader, because those need downloads this book’s test environment doesn’t allow. The layouts follow those tools’ source code and are pinned by the format tests; Appendix F records what was and wasn’t verified.

Run it

python run.py quantformats --steps 300

The command trains a small Qwen3 (2 layers, 256 wide) on The Verdict as bytes, then quantizes every linear layer with each method and measures the KL divergence of the quantized model’s next-token distributions from the original’s on held-out text. KL is llama.cpp’s preferred metric because, unlike perplexity, it can’t be improved by noise: a small overfit model’s held-out perplexity sometimes drops when quantization perturbs it, which says nothing about quality.

{"format": "bf16 reference", "bits": 16, "held_out_perplexity": 21.63}
{"format": "RTN int4, g64", "bits": 4.5, "kl_nats": 0.02238}
{"format": "GPTQ int4, g64", "bits": 4.5, "kl_nats": 0.00588}
{"format": "GPTQ int4, g64, act-order", "bits": 4.5, "kl_nats": 0.00505}
{"format": "AWQ int4, g64", "bits": 4.5, "kl_nats": 0.02201}
{"format": "FP8 E4M3, 128x128 blocks", "bits": 8.0, "kl_nats": 0.00216}
{"format": "MXFP4", "bits": 4.25, "kl_nats": 0.04254}
{"format": "NVFP4", "bits": 4.5, "kl_nats": 0.02723}
{"format": "trellis (QTIP-style), L=8", "bits": 2.09, "kl_nats": 0.29767}
{"format": "trellis (QTIP-style), L=8", "bits": 3.08, "kl_nats": 0.08815}
{"format": "RTN int2, g64 (for comparison)", "bits": 2.5, "kl_nats": 0.52499}

GPTQ cuts round-to-nearest’s damage by a factor of four at the same size, and act-order helps a little more. AWQ barely beats round-to-nearest here, and that’s informative rather than disappointing: this small model, trained for a few minutes, has none of the large outlier channels that real LLMs develop and that AWQ was designed to protect (the test on synthetic outliers shows the effect). FP8 is nearly lossless. NVFP4 beats MXFP4, its finer scales worth more than their quarter-bit. At about 2 bits, trellis codes are much better than 2.5-bit round-to-nearest, while still well behind 4 bits: on this tiny model every bit matters; on large models, whose weights are more redundant, QTIP and EXL3 at 2-3 bits come much closer to the original.

Build it

Engine milestone 39: quantized checkpoints. Implement gptq_quantize and unpack_gptq (engine/formats/gptq.py); search_scale, apply_awq and unpack_awq (engine/formats/awq.py); quantize_blocks (engine/formats/fp8.py); quantize_mxfp4 (engine/formats/mx.py); viterbi_encode (engine/formats/trellis.py); load_quantized (engine/formats/hf_quant.py); and fp8_block_kernel (engine/kernels/triton_formats.py). The packers, NVFP4, Hadamard rotation, decoders and the writer are provided.

pytest tests/test_ch39_quantized_checkpoints.py
python run.py quantformats --impl engine

The tests check GPTQ against round-to-nearest with and without act-order, GPTQ’s packing round trip and its act-order run-time permutation, AWQ on synthetic outlier channels and its interleaved packing, saving and loading GPTQ, AWQ and FP8 checkpoints of a Qwen3, FP8 blocks with the kernel and the W8A8 reference, MXFP4 in both byte orders (one against the GGUF decoder), NVFP4, Viterbi’s optimality by brute force, and trellis codes against the Lloyd-Max optimum.

Stretch exercises

  1. ★★ Measure every method on a real model: quantize Qwen3-0.6B with RTN, GPTQ (with and without act-order) and AWQ, using 128 calibration sequences of 2,048 tokens, and compare KL and a downstream accuracy (Chapter 44’s harness). Which layers suffer most? Where: experiments/ch39.py (create it), importing engine.formats.gptq, engine.formats.awq and engine.evaluation.
  2. ★★ Write a Triton kernel that decodes trellis codes on the fly: each program replays the bitshift state for its tile’s weights and applies the inverse Hadamard transform to the activations instead of the weights ($x H^\top$ is cheap with a fast Walsh-Hadamard transform). Where: add a trellis kernel in engine/kernels/triton_formats.py; connect it to engine/formats/trellis.py.
  3. ★★★ Combine them: GPTQ’s column-by-column error feedback with trellis codes instead of scalar rounding (QTIP’s “BlockLDLQ”). Measure the improvement over plain trellis coding at 2 bits. Where: add a feedback-based trellis quantizer in engine/formats/trellis.py, using engine.formats.gptq.
  4. ★ Load gpt-oss-20b’s MXFP4 expert tensors (*_blocks, *_scales) with dequantize_mxfp4(layout="interleaved") and check a few values against the BF16 conversion that Hugging Face publishes. Where: experiments/ch39.py (create it), importing engine.formats.mx.dequantize_mxfp4.

Check your understanding

  1. What information does GPTQ use that round-to-nearest ignores, and how does a quantization error in one column change the others?
  2. Why does act-order need g_idx, and how does the engine serve such a checkpoint without scattered groups?
  3. Why does scaling a weight column up before quantization reduce its relative error, and why is dividing the activations by the same factor free?
  4. Why does FP8 need block scales at all, given its exponent bits?
  5. MXFP4 and NVFP4 both store E2M1 values. Why does NVFP4 lose less?
  6. Why can a trellis code beat the best scalar code at the same number of bits? What makes it cheap to decode?
  7. Why is KL divergence a better quantization metric than perplexity for a small model?

Going deeper

  • Frantar, Ashkboos, Hoefler, Alistarh, GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers (ICLR 2023); Frantar and Alistarh, Optimal Brain Compression (2022).
  • Lin et al., AWQ: Activation-aware Weight Quantization for LLM Compression and Acceleration (MLSys 2024); Xiao et al., SmoothQuant (ICML 2023).
  • Tseng et al., QuIP#: Even Better LLM Quantization with Hadamard Incoherence and Lattice Codebooks (ICML 2024) and QTIP: Quantization with Trellises and Incoherence Processing (NeurIPS 2024); the ExLlamaV3 repository’s documentation of EXL3.
  • DeepSeek-AI, DeepSeek-V3 Technical Report (2024), §3.3 on FP8 training with fine-grained block scaling; the Open Compute Project’s Microscaling Formats (MX) Specification v1.0 (2023); NVIDIA’s NVFP4 documentation.
  • AutoGPTQ / GPTQModel, AutoAWQ and vLLM’s vllm/model_executor/layers/quantization/ for the reference formats and kernels (Marlin, Machete, the FP8 block GEMMs).