A 4-bit weight is only useful if some kernel can multiply with it faster than with the 16-bit original. A careful quantisation run through a generic dequantise-then-matmul path is smaller but no faster. The kernel is where the memory saving is turned into speed.
This article explains how INT4 kernels are put together from first principles: the roofline argument that justifies them, the three families (W4A16, W4A8 and W4A4) and the hardware each needs, how weights are stored and packed, how a handful of bit operations turn packed nibbles into fp16 values, how the load pipeline hides that work, and why the best kernel for one token is not the best for a thousand. The internals of one specific kernel are covered in the Marlin deep dive; how the 4-bit values are chosen in the first place is covered in the AWQ and GPTQ article.
Why INT4 kernels exist: the roofline
During decode an LLM generates one token per sequence per step, so every weight matrix is read from memory to be multiplied by a handful of activation vectors. Each weight byte is used for only a few floating point operations, so the step is limited by memory bandwidth, not arithmetic. The lower bound on time per token is simply bytes of weights divided by bandwidth.
Worked example for an 8-billion-parameter model on a GPU with 3.35 TB/s of memory bandwidth (the H100 SXM figure). In fp16 the weights are 16.1 GB, so a single-sequence decode step cannot take less than about 4.8 ms. In INT4 with one fp16 scale per group of 128 weights, the average cost is 4.125 bits per weight, so 4.1 GB and a floor of about 1.24 ms. That factor of nearly four is the whole prize, and a kernel collects it only if it reads packed bytes and keeps everything else hidden under those reads.
The picture changes with batch size. For a W4A16 GEMM of M tokens against an N by K weight matrix, the work is 2MNK flops and the traffic is dominated by the 0.52 bytes per weight, so arithmetic intensity is about 3.9M flops per byte. The same GPU's ridge point, dense fp16 throughput over bandwidth, is about 989 TFLOPS / 3.35 TB/s, roughly 295 flops per byte. The kernel therefore stops being memory-bound at around M = 76, before counting dequantisation, which is not free and pulls the crossover lower. Beyond that point the tensor cores are the limit and 4-bit weights no longer save time, only memory.
Three kernel families and the hardware each needs
The name tells you what the tensor cores multiply. In W4A16, weights are stored in 4 bits and activations stay in fp16 or bf16; the kernel dequantises weights to 16 bits in registers and uses ordinary fp16 tensor-core instructions. This works on any GPU with fp16 tensor cores and is what GPTQ and AWQ checkpoints use. In W4A8, weights are expanded to 8-bit integers or fp8 and activations are quantised to 8 bits, so the multiply runs at 8-bit rates, which helps the compute-bound, large-batch regime; QServe's W4A8 design is an example. In W4A4, both operands are 4-bit integers and the multiply uses INT4 tensor-core instructions directly.
W4A4 depends on hardware support that has changed between generations. Ampere and Turing tensor cores support INT4 operands. Hopper (H100) tensor cores do not: they support fp16, bf16, tf32, fp8, int8 and fp64. Blackwell adds 4-bit floating point formats, which are a different thing from INT4 integers, with their own scaling schemes. So an INT4 kernel on Hopper is necessarily a mixed-input kernel that dequantises to a format the tensor cores accept. vLLM's Machete kernel is an example built for Hopper on CUTLASS, using the Tensor Memory Accelerator and warpgroup MMA instructions.
| Family | Tensor core math | Wins at | Needs |
|---|---|---|---|
| W4A16 | fp16/bf16 | Decode, small batches | Fast in-register dequant |
| W4A8 | int8 or fp8 | Larger batches, prefill | Activation quantisation, 8-bit MMA |
| W4A4 | int4 | Compute-bound, if accuracy holds | INT4 tensor cores (not on Hopper) |
How the weights are stored
A 4-bit weight is a nibble, so eight fit in one 32-bit word. Alongside the packed words sit the quantisation parameters: usually one fp16 scale per group of 64 or 128 input channels per output column, and for asymmetric schemes a zero point per group, often itself packed to 4 bits. The granularity article covers how group size trades accuracy for overhead, and the symmetric versus asymmetric article covers zero points.
The packing order is not the natural one. Kernels reorder weights offline so that the bytes each thread loads are exactly the values it needs for its tensor-core fragment, and so that the dequant trick below emits values in the right order. The reference below uses the interleave [0, 2, 4, 6, 1, 3, 5, 7]: nibble position 0 holds element 0, position 1 holds element 2, and so on. Real kernels layer a further tile-level permutation on top, which is why a checkpoint packed for one kernel cannot be fed to another without repacking.
import numpy as np
ORDER = [0, 2, 4, 6, 1, 3, 5, 7] # interleave so the lop3 trick emits values in order
def quantize_g128(w, group=128):
"""Asymmetric 4-bit, one fp16 scale and zero per group of `group` input channels.
w: [K, N] fp32. Returns q [K, N] uint8 in 0..15, scales/zeros [K//group, N]."""
K, N = w.shape
g = w.reshape(K // group, group, N)
lo, hi = g.min(axis=1), g.max(axis=1)
scale = np.maximum(hi - lo, 1e-8) / 15.0
zero = np.round(-lo / scale)
q = np.clip(np.round(g / scale[:, None] + zero[:, None]), 0, 15).astype(np.uint8)
return q.reshape(K, N), scale.astype(np.float16), zero.astype(np.float16)
def pack_rows(q):
"""Pack 8 consecutive K entries of each column into one uint32, interleaved."""
K, N = q.shape
q = q.reshape(K // 8, 8, N).astype(np.uint32)
out = np.zeros((K // 8, N), dtype=np.uint32)
for pos, idx in enumerate(ORDER):
out |= q[:, idx, :] << np.uint32(4 * pos)
return out
def dequant_reference(q, scale, zero, group=128):
K, N = q.shape
g = q.reshape(K // group, group, N).astype(np.float32)
return ((g - zero[:, None]) * scale[:, None]).reshape(K, N)
Dequantising in registers with bit tricks
Converting a nibble to fp16 the obvious way, shift, mask, integer-to-float convert, costs several instructions per value, and there are many values per tile. The standard trick avoids the convert entirely. The fp16 number 1024.0 has bit pattern 0x6400, and at that exponent one unit in the last place is exactly 1. OR a 4-bit value q into the low mantissa bits and the result, read as fp16, is 1024 + q. Subtract 1024 and you have q as fp16.
Two halves fit in a 32-bit register, so one masked OR fills a half2. The mask 0x000f000f picks nibble positions 0 and 4; the mask 0x00f000f0 picks positions 1 and 5 but leaves them shifted by four bits, giving 1024 + 16q, which one fused multiply-add by 1/16 minus 64 fixes. Shifting the word right by 8 and repeating covers positions 2, 6, 3 and 7. That is why the interleave exists: with it, the four results come out as element pairs (0,1), (2,3), (4,5) and (6,7). On NVIDIA GPUs the mask and OR combine into one lop3 instruction.
// Eight 4-bit values (one uint32, packed in ORDER) -> four half2, then scale and zero.
__device__ inline void dequant8(uint32_t q, half2 s, half2 z, half2 out[4]) {
const uint32_t LO = 0x000f000f, HI = 0x00f000f0, EX = 0x64006400; // 0x6400 = 1024.0h
uint32_t r[4];
// lop3 with immediate 0xEA computes (a & b) | c in one instruction
asm("lop3.b32 %0, %1, %2, %3, 0xEA;" : "=r"(r[0]) : "r"(q), "n"(LO), "n"(EX));
asm("lop3.b32 %0, %1, %2, %3, 0xEA;" : "=r"(r[1]) : "r"(q), "n"(HI), "n"(EX));
q >>= 8;
asm("lop3.b32 %0, %1, %2, %3, 0xEA;" : "=r"(r[2]) : "r"(q), "n"(LO), "n"(EX));
asm("lop3.b32 %0, %1, %2, %3, 0xEA;" : "=r"(r[3]) : "r"(q), "n"(HI), "n"(EX));
const half2 k1024 = __float2half2_rn(1024.f);
const half2 k16th = __float2half2_rn(1.f / 16.f);
const half2 kn64 = __float2half2_rn(-64.f);
#pragma unroll
for (int i = 0; i < 4; ++i) {
half2 v = *reinterpret_cast<half2*>(&r[i]);
v = (i & 1) ? __hfma2(v, k16th, kn64) // (1024 + 16q)/16 - 64 = q
: __hsub2(v, k1024); // (1024 + q) - 1024 = q
out[i] = __hmul2(__hsub2(v, z), s); // (q - zero) * scale
}
}The scale and zero multiply can be folded further: precompute zero * scale per group and use one fused multiply-add. Symmetric formats that store q + 8 subtract 1032 on the low path and 72 on the high path, or fold the 8 into the zero point. The budget is a few instructions per eight weights, small enough to hide under loads.
The load pipeline
Because the kernel is memory-bound at small batch, its structure is about keeping loads in flight. Packed weights move from HBM to shared memory with asynchronous copies, several tiles ahead, so the next tile is arriving while the current one is computed. From shared memory, threads load packed words into registers in a pattern that matches the tensor-core fragment layout; offline reordering exists to make those loads wide and conflict-free. Dequantisation happens in registers immediately before the MMA, and the 16-bit values never touch memory.
Scales and zeros are loaded once per group and reused across all columns of the tile, so their cost is small with groups of 128. Activations are loaded once per tile and reused against every weight, which is why they can stay 16-bit at small batch without hurting the memory budget.
One token versus a thousand
At M = 1 the problem is a matrix-vector product. There are few output tiles relative to the number of streaming multiprocessors, so kernels split the K dimension across thread blocks, a technique called split-K, and add the partial sums at the end. At large M, the prefill of a long prompt, the problem is an ordinary compute-bound GEMM, and the best choice is often not a 4-bit kernel at all: some servers dequantise the whole layer once into a fp16 buffer and call a dense GEMM, or use a W4A8 kernel so the multiply itself is faster.
Serving engines therefore dispatch by shape. The same layer may use a split-K W4A16 kernel for decode steps, a tiled mixed-input kernel for mid-sized batches, and a dense path for prefill.
Act-order, zero points and the formats that cost extra
GPTQ's act-order option quantises columns in order of activation magnitude, which improves accuracy but means that consecutive input channels no longer share a quantisation group. A naive kernel then gathers a different scale for every row. Kernels that support it usually permute the weight rows once at load time so groups are contiguous again, and apply the matching permutation to the activations on every call. That gather is cheap next to a matmul but it is not free, and not every kernel supports it.
Zero points add one subtract per value and a second small tensor to load. Very small groups, such as 32, raise the per-weight overhead from 4.125 to 4.5 bits with fp16 scales. Bf16 models add a wrinkle: the magic-number trick is written for fp16 bit patterns, so a bf16 path needs its own constants or a different conversion. Check which of these a kernel supports before choosing a quantisation format, not after.
Failure modes
- Silent fallback. A layer whose shape the fast kernel does not support, often an output dimension that is not a multiple of the tile or a group size that does not divide K, falls back to a slow generic path. Throughput drops and nothing errors. Log which kernel each layer uses at load.
- Wrong packing order. A checkpoint packed for one kernel and read by another runs at full speed and produces fluent-looking garbage or small but real accuracy loss. Validate layer outputs against a dequantised reference.
- Zero point convention. Off-by-one conventions, such as symmetric values stored with an offset of 8 versus a zero tensor, shift every weight by one step. Perplexity rises modestly, which is easy to miss.
- Tensor-parallel splits across groups. Sharding K across GPUs must cut on group boundaries and shard scales and permutation indexes consistently.
- Wrong regime. A kernel tuned for batch 1 can be slower than fp16 at batch 256. Benchmark across the batch range.
# Per-layer validation: the kernel must match the dequantised reference, at your batch sizes.
import torch
def check_layer(kernel, x_shapes, packed, scales, zeros, w_ref, atol=2e-2):
for m in x_shapes: # e.g. [1, 8, 32, 128, 1024]
x = torch.randn(m, w_ref.shape[0], dtype=torch.float16, device="cuda")
y_ref = (x.float() @ w_ref.float()).half() # w_ref = dequant_reference(...) in fp16
y = kernel(x, packed, scales, zeros)
err = (y.float() - y_ref.float()).abs().max().item()
assert err < atol * y_ref.float().abs().max().item(), f"M={m}: max err {err}"
Operating INT4 kernels in production
Profile decode steps and look at achieved memory bandwidth as a percentage of peak for the linear layers. A good W4A16 kernel at small batch should be a large fraction of peak; a low figure means dequant or pipeline stalls are exposed. Measure time to first token separately; prefill uses a different kernel. Keep a dense fp16 or fp8 baseline in the benchmark suite, so you notice when a new GPU, driver or engine version changes which path is fastest.
What to do next
- Compute your model's decode floor: weight bytes at 4.125 bits divided by your GPU's bandwidth, and compare it with measured time per token.
- Confirm which tensor-core formats your GPU supports before choosing W4A4, W4A8 or W4A16.
- Log the kernel chosen for every quantised layer at load and fail the deployment on unexpected fallbacks.
- Validate each layer against a dequantised reference at batch sizes 1, 8, 32 and your prefill size.
- Benchmark tokens per second across your real concurrency range, not just batch 1.
- Pick group size, zero points and act-order with the kernel's supported formats in hand.