20. Quantization
In this chapter
- Why fewer bits per weight means faster decode, with the arithmetic to predict how much.
- Symmetric integer quantization derived from scratch: scales, codes, the error bound and granularity (per tensor, per row, per group).
- Packing two 4-bit values per byte, and the metadata cost that "4-bit" leaves out.
- Calibration: choosing clipping by output error, not weight error.
- A W4A16 Triton kernel that dequantizes in registers, and the landscape beyond round-to-nearest: GPTQ, AWQ, SmoothQuant, FP8, NF4 and KV-cache quantization.
You will build
engine/quant.py (INT8/INT4 groupwise quantization, packing, QuantLinear) and a W4A16 matmul kernel in engine/kernels/triton_quant.py. You'll quantize your Qwen3 engine and measure size, quality and speed.
Time: 5-7 hours. GPU: optional (needed for the speed measurements and the bitsandbytes workflow).
Fewer bytes, faster tokens
Decode reads every weight once per token, so its ceiling is bandwidth ÷ bytes per token (Chapter 10). Halve the bytes and you double the ceiling. For an 8B-parameter model:
| format | bits per weight (incl. scales) | weights | ceiling at 273 GB/s (DGX Spark) | at 1,008 GB/s (RTX 4090) |
|---|---|---|---|---|
| BF16 | 16 | 16.4 GB | 17 tokens/s | 61 tokens/s |
| INT8, per row | ~8 | 8.2 GB | 33 tokens/s | 123 tokens/s |
| INT4, groups of 128 | ~4.1 | 4.3 GB | 63 tokens/s | 234 tokens/s |
Quantization also decides what fits. A 32B model in BF16 needs 64 GB, more than an RTX 4090’s 24 GB; in INT4 it needs about 17 GB. For the sparse models of Part VII, with hundreds of billions of parameters, low-bit weights are the difference between running and not running.
The price is error: every weight becomes the nearest value on a coarse grid. The whole subject is about spending a fixed bit budget where it matters most.
Symmetric quantization, from scratch
Take a group of weights $w_1, \dots, w_n$ and a bit width $b$. A signed $b$-bit integer can hold $-2^{b-1} \dots 2^{b-1}-1$; we use the symmetric range $[-Q, Q]$ with $Q = 2^{b-1}-1$ (127 for INT8, 7 for INT4) so that zero is exact and the grid is symmetric. Map the largest magnitude $a = \max_i |w_i|$ to $Q$:
$$ s = \frac{a}{Q}, \qquad q_i = \operatorname{clamp}!\Big(\operatorname{round}\big(w_i / s\big), -Q, Q\Big), \qquad \hat w_i = s, q_i . $$
The integers $q_i$ (the codes) are stored, along with one floating-point scale $s$ per group. Without clamping (the range covers every weight), rounding moves each value by at most half a step:
$$ |\hat w_i - w_i| \le \tfrac{s}{2} = \frac{a}{2Q}. $$
Worked example. $w = [-2, -1, 0, 1, 2]$ at 4 bits: $Q = 7$, $s = 2/7 \approx 0.2857$. Dividing gives $[-7, -3.5, 0, 3.5, 7]$. PyTorch rounds halves to even, so the codes are $[-7, -4, 0, 4, 7]$ and $\hat w = [-2, -1.143, 0, 1.143, 2]$. The error on $\pm 1$ is $0.143 = s/2$, exactly at the bound. (Rounding conventions differ between libraries, which is one reason to test against the specific runtime you’ll use.)
Granularity: who shares a scale
The bound says the error is proportional to the largest magnitude in the group. One large weight coarsens the grid for everything that shares its scale. So the choice of groups matters as much as the bit width:
- Per tensor: one scale for the whole matrix. A single outlier anywhere ruins it. Fine for INT8 activations with care, poor for low-bit weights.
- Per row (per output channel): one scale per row of the
[out, in]weight. Standard for INT8 weights. - Per group: one scale per
gconsecutive inputs within a row, typically $g = 32$-$128$. Standard for INT4.
def quantize_int8_rows(weight):
"""One scale per output row: scale = max|w| / 127, q = round(w / scale). (Your engine: Chapter 20)"""
maximum = weight.float().abs().amax(dim=-1, keepdim=True)
scale = torch.where(maximum > 0, maximum / 127.0, torch.ones_like(maximum))
return (weight.float() / scale).round().clamp(-127, 127).to(torch.int8), scale
def dequantize_int8_rows(q, scale):
return q.float() * scale
def quantize_groupwise(weight, bits=4, group_size=128, clip=1.0):
"""Symmetric signed codes with one scale per (row, group of group_size inputs). (Your engine: Chapter 20)
weight [O, I] -> codes int8 [O, I] in [-(2^(b-1)-1), 2^(b-1)-1], scales [O, ceil(I/g)].
clip < 1 shrinks the range: large outliers saturate, everything else gets finer steps.
"""
if weight.ndim != 2 or bits not in (2, 3, 4, 8) or group_size < 1 or not 0 < clip <= 1:
raise ValueError("Expected a matrix, 2/3/4/8 bits, positive group size and 0 < clip <= 1")
out_features, in_features = weight.shape
groups = -(-in_features // group_size)
padded = F.pad(weight.float(), (0, groups * group_size - in_features))
blocks = padded.view(out_features, groups, group_size)
limit = 2 ** (bits - 1) - 1
maximum = blocks.abs().amax(-1, keepdim=True) * clip
scales = torch.where(maximum > 0, maximum / limit, torch.ones_like(maximum))
codes = (blocks / scales).round().clamp(-limit, limit).to(torch.int8)
return codes.view(out_features, -1)[:, :in_features], scales.squeeze(-1)
def dequantize_groupwise(codes, scales, group_size):
out_features, in_features = codes.shape
expanded = scales.repeat_interleave(group_size, dim=1)[:, :in_features]
return codes.float() * expanded
Two edge cases the code handles: a group of all zeros gets scale 1 (any scale works, and 0/0 doesn’t); the last group of a row may be short when in_features isn’t a multiple of g, so the row is padded for the computation and the padding is dropped.
The metadata the name leaves out
“4-bit” counts the codes only. With groups of $g$ inputs and a 16-bit scale per group, the true cost is $4 + 16/g$ bits per weight: 4.125 for $g = 128$, 4.5 for $g = 32$. Asymmetric schemes add a zero point per group. Layers usually left in high precision (embeddings, the LM head, MoE routers, norms) add more. For Qwen3-0.6B, whose tied embedding is 26% of the parameters, “4-bit” weights come to well over 4 bits on average. Always report measured bytes.
Packing
PyTorch has no 4-bit dtype, so an int8 tensor holding values in $[-7, 7]$ still costs 8 bits each. Pack two codes per byte: the first in the low nibble, the second in the high nibble, each in 4-bit two’s complement (values 8-15 of a nibble mean -8 to -1):
def pack_int4(codes):
"""Two signed 4-bit values per byte, two's complement, first value in the low nibble."""
values = codes.to(torch.int16)
if bool(((values < -8) | (values > 7)).any()):
raise ValueError("INT4 values must lie in [-8, 7]")
flat = values.flatten()
if flat.numel() % 2:
flat = torch.cat((flat, flat.new_zeros(1)))
nibbles = flat & 0xF
return (nibbles[0::2] | (nibbles[1::2] << 4)).to(torch.uint8)
def unpack_int4(packed, shape):
"""Inverse of pack_int4: nibbles 8..15 decode to -8..-1."""
nibbles = torch.stack((packed & 0xF, packed >> 4), dim=-1).flatten().to(torch.int16)
signed = torch.where(nibbles >= 8, nibbles - 16, nibbles)
count = 1
for size in shape:
count *= size
return signed[:count].to(torch.int8).reshape(shape)
inline std::pair<std::vector<int8_t>, std::vector<float>> quantize_row(const std::vector<float>& w, int bits, size_t group) {
float limit = float((1 << (bits - 1)) - 1);
std::vector<int8_t> codes;
std::vector<float> scales;
for (size_t g = 0; g < w.size(); g += group) {
size_t end = std::min(w.size(), g + group);
float m = 0;
for (size_t i = g; i < end; ++i) m = std::max(m, std::fabs(w[i]));
float scale = m > 0 ? m / limit : 1.f;
scales.push_back(scale);
for (size_t i = g; i < end; ++i) codes.push_back(int8_t(std::clamp(std::round(w[i] / scale), -limit, limit)));
}
return {codes, scales};
}
inline std::vector<uint8_t> pack_int4(const std::vector<int8_t>& c) { // first value in the low nibble
std::vector<uint8_t> out((c.size() + 1) / 2, 0);
for (size_t i = 0; i < c.size(); ++i) out[i / 2] |= uint8_t((c[i] & 0x0f) << (4 * (i % 2)));
return out;
}
inline int8_t unpack_nibble(uint8_t byte, int high) {
int n = (byte >> (4 * high)) & 0x0f;
return int8_t(n >= 8 ? n - 16 : n); // two's complement
}
#![allow(unused)]
fn main() {
/// Quantize one row in groups of `group` values: scale = max|w| / (2^(bits-1) - 1).
pub fn quantize_row(w: &[f32], bits: u32, group: usize) -> (Vec<i8>, Vec<f32>) {
let limit = ((1 << (bits - 1)) - 1) as f32;
let mut codes = Vec::with_capacity(w.len());
let mut scales = vec![];
for chunk in w.chunks(group) {
let max = chunk.iter().fold(0.0f32, |m, v| m.max(v.abs()));
let scale = if max > 0.0 { max / limit } else { 1.0 };
scales.push(scale);
codes.extend(chunk.iter().map(|v| (v / scale).round().clamp(-limit, limit) as i8));
}
(codes, scales)
}
/// Two signed 4-bit codes per byte, first code in the low nibble (two's complement).
pub fn pack_int4(codes: &[i8]) -> Vec<u8> {
codes.chunks(2).map(|p| {
let lo = (p[0] as u8) & 0x0f;
let hi = (*p.get(1).unwrap_or(&0) as u8) & 0x0f;
lo | (hi << 4)
}).collect()
}
pub fn unpack_int4(packed: &[u8], count: usize) -> Vec<i8> {
let nibble = |n: u8| if n >= 8 { n as i8 - 16 } else { n as i8 };
packed.iter().flat_map(|b| [nibble(b & 0x0f), nibble(b >> 4)]).take(count).collect()
}
}
Packed layouts are a contract between the quantizer and the kernel, and every library chooses differently: nibble order, interleaving for vectorized loads, and transposition for tensor-core fragment layouts. A GPTQ checkpoint packs eight 4-bit values into an int32; Marlin kernels reorder them again for fast loading. A checkpoint is only usable by a runtime that implements its exact format.
Calibration: optimize what you care about
Clipping (shrinking $a$ by a factor clip < 1) saturates the largest weights but gives everything else a finer grid. Whether that helps depends on how the weights are used. A large weight attached to an input feature that’s almost always near zero barely affects the output, so it’s worth clipping; a large weight on an active feature isn’t.
So choose quantization parameters by output error on representative inputs (calibration data), not by weight error:
def choose_clip(weight, calibration_inputs, bits=4, group_size=128, candidates=(1.0, 0.95, 0.9, 0.85, 0.8, 0.7)):
"""Pick the clip that minimizes OUTPUT error on calibration activations, not weight error."""
reference = calibration_inputs.float() @ weight.float().T
best = None
for clip in candidates:
codes, scales = quantize_groupwise(weight, bits, group_size, clip)
approx = calibration_inputs.float() @ dequantize_groupwise(codes, scales, group_size).T
error = ((approx - reference).norm() / reference.norm().clamp_min(1e-12)).item()
if best is None or error < best[0]:
best = (error, clip)
return best[1], best[0]
The lab plants exactly that situation: a 16 × 128 weight whose input feature 0 has weights 12× larger than the rest, while that feature is nearly silent in the activations. It picks the clip on 256 calibration rows and reports error on 256 different held-out rows:
python model_workflows.py quantize
{"bits": 8, "group": 128, "clip": 0.95, "weight_error": 0.03627, "held_out_output_error": 0.02119, ...}
{"bits": 8, "group": 32, "clip": 1.0, "weight_error": 0.00874, "held_out_output_error": 0.01161, ...}
{"bits": 4, "group": 128, "clip": 0.7, "weight_error": 0.28655, "held_out_output_error": 0.28023, ...}
{"bits": 4, "group": 64, "clip": 0.7, "weight_error": 0.25714, "held_out_output_error": 0.22791, ...}
{"bits": 4, "group": 32, "clip": 0.85, "weight_error": 0.17493, "held_out_output_error": 0.19074, ...}
Three lessons: 4-bit chose aggressive clipping because the outlier is harmless to clip; smaller groups help most at low bit widths; and at group 32 the output error is larger than the weight error, while at group 64 it’s smaller. Weight error is not the objective. Keep calibration and evaluation data separate, or you’ll fool yourself.
A quantized linear layer
class QuantLinear(nn.Module):
"""Weight-only quantized linear layer: integer weights in memory, higher-precision math. (Your engine: Chapter 20)
The reference forward dequantizes the whole matrix each call, which saves memory but not
time. A real W4A16 kernel (kernels/triton_quant.py) dequantizes tiles in registers so the
weights cross the memory bus as 4-bit values.
"""
def __init__(self, linear, bits=4, group_size=128, clip=1.0):
super().__init__()
self.in_features, self.out_features = linear.in_features, linear.out_features
self.bits, self.group_size = bits, group_size
codes, scales = quantize_groupwise(linear.weight.data, bits, group_size, clip)
if bits == 4:
self.register_buffer("packed", pack_int4(codes))
else:
self.register_buffer("packed", codes)
self.register_buffer("scales", scales.to(linear.weight.dtype))
self.bias = None if linear.bias is None else nn.Parameter(linear.bias.data.clone(), requires_grad=False)
self.compute_dtype = linear.weight.dtype
def codes(self):
shape = (self.out_features, self.in_features)
return unpack_int4(self.packed, shape) if self.bits == 4 else self.packed
def dequantized_weight(self):
"""(Your engine: Chapter 20)"""
return dequantize_groupwise(self.codes(), self.scales.float(), self.group_size).to(self.compute_dtype)
def forward(self, x):
"""(Your engine: Chapter 20)"""
return F.linear(x, self.dequantized_weight(), self.bias)
def storage_bytes(self):
return sum(t.numel() * t.element_size() for t in (self.packed, self.scales))
def quantize_model(model, bits=4, group_size=128, skip=("lm_head", "head", "gate", "router", "shared_expert_gate")):
"""Replace every nn.Linear with a QuantLinear, in place, unless a component of its name is in
`skip` (output heads and MoE routers are small and sensitive, so they stay in high precision)."""
replaced = 0
for name, module in list(model.named_modules()):
for child_name, child in list(module.named_children()):
full = f"{name}.{child_name}" if name else child_name
if isinstance(child, nn.Linear) and not set(full.split(".")) & set(skip):
setattr(module, child_name, QuantLinear(child, bits, group_size))
replaced += 1
return replaced
quantize_model replaces every nn.Linear in place, except layers whose names mark them as sensitive: the LM head (its errors land directly on the logits) and MoE routers (small, and a flipped expert choice changes the computation discretely). On the tiny Qwen3 test model in FP32:
python run.py quant --bits 4
{"bits": 4, "bytes_before": 15739904, "bytes_after": 3124224,
"logit_error": {"max_abs": 0.1275, "mean_abs": 0.0178, "relative_l2": 0.0396}, "same_top1": 0.953}
{"bits": 8, "bytes_before": 15739904, "bytes_after": 4959232,
"logit_error": {"max_abs": 0.0075, "mean_abs": 0.00098, "relative_l2": 0.0022}, "same_top1": 1.0}
(The 8-bit line comes from --bits 8; both runs are from a laptop CPU.) The 8-bit model’s size isn’t a quarter of FP32’s because the embedding stays unquantized. At 4 bits, 95% of positions keep the same top token, with 4% relative logit error. Real models are more robust than this tiny random one, but measure, don’t assume.
Speed needs a kernel
QuantLinear.forward dequantizes the whole matrix to BF16 and calls an ordinary matmul. That saves memory but not time: the full-precision weights are materialized and read again every call, so decode is slower than BF16. The speedup only appears when a kernel reads the packed bytes and dequantizes them in registers, so weights cross the memory bus at 4 bits. That’s a W4A16 kernel: 4-bit weights, 16-bit activations, 16/32-bit math.
@triton.jit
def w4a16_kernel(x_ptr, packed_ptr, scales_ptr, y_ptr, M, N, K, groups,
GROUP: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
"""(Your engine: Chapter 20)"""
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)
half = tl.arange(0, BLOCK_K // 2) # byte index inside the K tile
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)
# Packed row n holds K/2 bytes; byte j carries code 2j (low nibble) and 2j+1 (high nibble).
byte_cols = k0 // 2 + half
packed = tl.load(packed_ptr + rn[:, None] * (K // 2) + byte_cols[None, :],
mask=(rn[:, None] < N) & (byte_cols[None, :] < K // 2), other=0).to(tl.int32)
low, high = packed & 0xF, (packed >> 4) & 0xF
codes = tl.reshape(tl.join(low, high), (BLOCK_N, BLOCK_K)) # interleave back to k order
codes = tl.where(codes >= 8, codes - 16, codes) # two's complement nibble
scale = tl.load(scales_ptr + rn * groups + k0 // GROUP, mask=rn < N, other=0.0).to(tl.float32)
w = codes.to(tl.float32) * scale[:, None] # [BLOCK_N, BLOCK_K]
acc += tl.dot(x.to(tl.float32), tl.trans(w), input_precision="ieee") # FP32, not TF32
tl.store(y_ptr + rm[:, None] * N + rn[None, :], acc.to(y_ptr.dtype.element_ty),
mask=(rm[:, None] < M) & (rn[None, :] < N))
def w4a16_matmul(x, packed, scales, group_size, out_features, block_m=16, block_n=32, block_k=32):
"""x [M, K]; packed uint8 [N * K / 2] (row-major, from quant.pack_int4); scales [N, K / group]."""
check_device(x, packed, scales)
m, k = x.shape
if k % 2 or group_size % block_k or k % block_k:
raise ValueError("K must be a multiple of BLOCK_K, and BLOCK_K must divide the group size")
y = torch.empty((m, out_features), device=x.device, dtype=x.dtype)
grid = (triton.cdiv(m, block_m), triton.cdiv(out_features, block_n))
w4a16_kernel[grid](x.contiguous(), packed, scales.contiguous(), y, m, out_features, k, scales.shape[1],
GROUP=group_size, BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k)
return y
It’s the Chapter 14 matmul with the B-tile loaded differently. Each program loads BLOCK_K / 2 bytes per weight row, splits each byte into its two nibbles, interleaves them back into $k$ order with tl.join and a reshape, sign-extends, and multiplies by the group’s scale. Requiring BLOCK_K to divide the group size means one scale per row per K-tile. As in Chapter 15, input_precision="ieee" keeps the FP32 tl.dot exact on the GPU instead of TF32.
What production W4A16 kernels (Marlin, Machete, AWQ’s GEMV, ExLlama) add: weight layouts pre-shuffled so a 128-bit load fills a tensor-core fragment directly, dequantization via bit tricks instead of shifts and multiplies, and split-K for the batch-1 decode shape. At batch 1 a good INT4 kernel approaches 3.5-3.9× BF16 decode speed; at large batches the kernel becomes compute-bound and the advantage shrinks, because the dequantization arithmetic is extra work.
Beyond round-to-nearest
Round-to-nearest with good groups is a strong baseline at 8 bits and a reasonable one at 4. The research methods address different parts of the problem:
| method | idea | needs |
|---|---|---|
| LLM.int8() (Dettmers et al., 2022) | A few activation features have huge outliers; compute those columns in FP16 and the rest in INT8 | mixed-precision kernels |
| SmoothQuant (Xiao et al., 2022) | Move activation outliers into the weights with a per-channel scale: $XW^\top = (X S^{-1})(W S)^\top$, exact before rounding | calibration; enables W8A8 |
| GPTQ (Frantar et al., 2022) | Quantize weights one column at a time and adjust the remaining columns to compensate the error, using second-order statistics of calibration inputs | calibration; a few minutes per model |
| AWQ (Lin et al., 2023) | Protect the ~1% of weight channels attached to large activations by scaling them up before quantization | calibration; very fast |
| NF4 (QLoRA, Dettmers et al., 2023) | A 16-value non-uniform grid matched to normally distributed weights; a table lookup to dequantize | used for frozen fine-tuning bases |
| FP8 (E4M3, E5M2) | 8-bit floating point with per-tensor or per-block scales; native tensor-core support on Hopper and later | W8A8 with almost no accuracy loss |
| MXFP4 / NVFP4 | 4-bit floats with a shared scale per block of 16-32 values, native on Blackwell | hardware support |
| QAT | Simulate quantization during training so the model adapts to it | training compute |
The scaling identity behind SmoothQuant and AWQ is worth remembering: for a diagonal $S$, $XW^\top = (XS^{-1})(WS)^\top$ exactly. It moves magnitude between activations and weights without changing the function, so you can choose where the outliers live before you round.
Quantizing the KV cache
At long contexts, the KV cache dominates memory traffic (Chapter 16). Quantizing it to FP8 or INT8, with a scale per head or per token, halves the traffic of attention during decode. Keys are more sensitive than values (errors pass through the softmax’s exponential), and keys often have outlier channels, so KV quantization usually uses per-channel scales for K and per-token scales for V (KIVI, 2024). Chapter 16’s last stretch exercise implements a simple version.
Quantize a real checkpoint
For Qwen3-0.6B, the companion script compares BF16 against bitsandbytes NF4 or LLM.int8() on held-out instruction data, then saves and reloads the quantized model to confirm the output survives the round trip:
uv pip install -r workflow-requirements.txt # transformers, safetensors, peft
uv pip install 'bitsandbytes>=0.48' # check platform support first
python quantize_checkpoint.py --model-dir models/Qwen3-0.6B --bits 4 \
--valid-file data/sft-valid.jsonl --output runs/qwen3-nf4
It reports the response loss of both models, the KL divergence between their next-token distributions, and parameter bytes. These are what to look at when judging a quantization: loss and KL on your own data first, then task scores and some generations read by a human.
Warning
bitsandbytes ships compiled CUDA binaries. If its CUDA version doesn’t match PyTorch’s, import fails or 4-bit layers silently fall back. The bitsandbytes installation guide explains
BNB_CUDA_VERSION; Chapter 22 shows how to diagnose it. The exported checkpoint runs in Transformers + bitsandbytes, not in your engine’s loader, which expects ordinary floating-point tensors.
An evaluation protocol
- Fix the evaluation set, token IDs and decoding settings before quantizing anything.
- Report held-out loss (or perplexity) and the KL from the original model, then task metrics.
- Measure memory (weights, KV cache, peak) and speed separately for prefill and decode, with batch size, lengths, GPU, backend and warmup recorded.
- Re-evaluate after every transformation that follows (merging adapters, editing, exporting): quantization doesn’t commute with them.
Build it
Engine milestone 20: quantization. Implement quantize_int8_rows, quantize_groupwise, QuantLinear.dequantized_weight and QuantLinear.forward in engine/quant.py (packing, clip search and quantize_model are provided), and w4a16_kernel in engine/kernels/triton_quant.py.
pytest tests/test_ch20_quant.py
python run.py quant --impl engine --bits 4
The tests check the half-step error bound for every bit width, short tail groups, zero groups, INT4 packing at both ends of the range, that a quantized Qwen3 shrinks and stays close, and your Triton kernel against dequantize-then-matmul. Then quantize your Qwen3-0.6B engine to 8 and 4 bits and record held-out loss and size in your notes.
Stretch exercises
- ★ Plot held-out output error against bits per weight (including scales) for group sizes 16-256 on a real Qwen3-0.6B layer. Where:
experiments/ch20.py(create it), usingengine.quant.quantize_groupwiseanddequantize_groupwise. - ★★ Make
QuantLinearcall your W4A16 kernel on CUDA, and measure decode tokens/s for Qwen3-0.6B against BF16 in graph mode (Chapter 19). Where:QuantLinear.forwardinengine/quant.py, callingengine.kernels.triton_quant.w4a16_matmul. - ★★ Implement AWQ’s core: for one linear layer, search a per-input-channel scale $s_j = \bar{|x_j|}^{\alpha}$ over $\alpha \in [0, 1]$ that minimizes calibration output error after 4-bit quantization of $W \operatorname{diag}(s)$. Where: add a calibration/search helper in
engine/quant.py; Chapter 39’sengine/formats/awq.pyprovides the later integration. - ★★★ Implement GPTQ for one layer: accumulate $H = 2X^\top X$ from calibration inputs, then quantize column by column, spreading each column’s error over the remaining columns with the Cholesky factor of $H^{-1}$. Compare with round-to-nearest at 3 and 4 bits. Where: add a one-layer GPTQ helper in
engine/quant.py; Chapter 39’sengine/formats/gptq.pyprovides the later integration.
Check your understanding
- Why does halving the bytes per weight roughly double the decode ceiling, but barely change large-batch prefill speed?
- Why do smaller groups reduce error, and what do they cost?
- Why can a clip below 1.0 reduce output error while increasing weight error?
- Why doesn’t dequantize-then-matmul make decode faster?
- Why are the LM head and MoE routers often left unquantized?
Going deeper
- GPU Mode L7 (Advanced quantization: weight-only and dynamic quantization in torchao, with Triton kernels), L33 (BitBLAS: mixed-precision kernels for arbitrary bit widths), L30 (Quantized training).
- PMPP Appendix A (Numerical considerations): the low-precision floating-point formats (FP16, BF16, FP8, block-scaled formats) and how GPUs implement them.
- Papers: LLM.int8() (2208.07339), SmoothQuant (2211.10438), GPTQ (2210.17323), AWQ (2306.00978), QLoRA/NF4 (2305.14314), KIVI (2402.02750); Micikevicius et al., FP8 Formats for Deep Learning (2022).
gpt-fast’squantize.py(int8 and int4 weight-only with GPTQ in a few hundred lines); the Marlin kernel write-up; Maarten Grootendorst’s A Visual Guide to Quantization.