A 4-bit model only gets faster if a kernel reads the 4-bit weights directly. Quantizing the checkpoint and then dequantizing the whole tensor to FP16 before a normal matmul saves disk and memory, but it is slower than not quantizing at all: you read the small weights, write the big ones, and read them again. The speed lives in a kernel that streams packed integers from memory, turns them into FP16 inside registers, and feeds tensor cores without the expanded weights ever touching DRAM.
The architecture of such kernels, including W4A16 versus W4A4, the bit tricks for dequantization and the decode-to-prefill crossover, is covered in INT4 kernel architecture. This page is the practitioner's loop around it: write a working W4A16 kernel in Triton, test it against the correct reference, benchmark it against the bandwidth limit, and see which gaps the production kernels close.
The job, stated precisely
A linear layer computes C = A W, where A holds activations of shape [M, K], W is [K, N], and M is the number of tokens in flight. In W4A16 the weights are stored as 4-bit integers with one FP16 scale per group of G consecutive values along K (G = 128 is the common choice), while activations stay in FP16. The kernel must compute the same result as multiplying A by the dequantized weights, w = (q - 8) times s for the symmetric format used here, with FP32 accumulation, to within floating-point rounding.
Why bother? During autoregressive decode M is small: one token per sequence, times the batch. Each weight is loaded from memory and used for only M multiply-adds, so arithmetic intensity is tiny and the GPU spends its time waiting on DRAM. Time is roughly bytes divided by bandwidth. A 4096 by 11008 projection is 90 MB in FP16 and about 23 MB in INT4 plus scales; at an illustrative 2 TB/s that is roughly 45 microseconds against 11.5. That ratio, close to 4x, is the entire prize, and it shrinks as M grows and the matmul becomes compute-bound. The INT4 format article covers the error side of the same trade.
The data contract
Before writing a kernel, fix the layout, because a kernel is only correct for one layout. Here eight 4-bit values that are consecutive along K share one int32, with value k at bits 4(k mod 8) through 4(k mod 8)+3. The packed tensor is [K/8, N]; scales are [K/G, N]. Values are stored offset by 8 so they fit 0..15. Real formats differ in pack axis, nibble order, zero points and interleaving, as the bit packing article lays out, and many production kernels permute this layout further at load time. Treat the packer and the dequantizer as the specification, and keep them next to the kernel.
import torch
def quantize_pack(w, G=128):
# w: [K, N] float (the transposed weight, so K is the reduction axis)
K, N = w.shape
assert K % G == 0 and K % 8 == 0
wg = w.float().reshape(K // G, G, N)
s = (wg.abs().amax(dim=1) / 7.0).clamp_min(1e-8) # [K//G, N], symmetric
q = torch.clamp(torch.round(wg / s[:, None, :]), -8, 7) + 8 # stored as 0..15
q = q.reshape(K // 8, 8, N).to(torch.int32)
packed = torch.zeros(K // 8, N, dtype=torch.int32, device=w.device)
for i in range(8): # k = 8*row + i lives at bits 4i..4i+3
packed |= q[:, i, :] << (4 * i)
return packed, s.half()
def dequantize(packed, s, G=128):
K = packed.shape[0] * 8
shifts = torch.arange(8, device=packed.device, dtype=torch.int32) * 4
q = (packed[:, None, :] >> shifts[None, :, None]) & 0xF # [K//8, 8, N]
q = q.reshape(K, -1).float() - 8.0
return (q * s.float().repeat_interleave(G, dim=0)).half()The eighth nibble, shifted by 28, sets the int32 sign bit whenever it is 8 or more. That is harmless only if every unpacking path masks with 0xF after the arithmetic right shift.
A W4A16 kernel in Triton
Triton lets you write a tiled GPU kernel in Python while the compiler handles shared-memory staging, software pipelining and tensor-core instruction selection. Each program instance computes one BM by BN tile of C and loops over K in steps of BK.
import triton
import triton.language as tl
@triton.autotune(
configs=[
triton.Config({"BM": 16, "BN": 64, "BK": 64}, num_warps=4, num_stages=3),
triton.Config({"BM": 16, "BN": 128, "BK": 128}, num_warps=4, num_stages=4),
triton.Config({"BM": 64, "BN": 128, "BK": 64}, num_warps=8, num_stages=3),
triton.Config({"BM": 128, "BN": 128, "BK": 32}, num_warps=8, num_stages=3),
],
key=["M", "N", "K"],
)
@triton.jit
def w4a16_gemm(a_ptr, qw_ptr, s_ptr, c_ptr, M, N, K,
stride_am, stride_ak, stride_qk, stride_qn, stride_sk, stride_sn,
stride_cm, stride_cn, G: tl.constexpr,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr):
pid_m, pid_n = tl.program_id(0), tl.program_id(1)
rm = pid_m * BM + tl.arange(0, BM)
rn = pid_n * BN + tl.arange(0, BN)
rk = tl.arange(0, BK)
shifts = (rk % 8) * 4 # nibble position of each k
acc = tl.zeros((BM, BN), dtype=tl.float32)
for k0 in range(0, K, BK): # assumes K % BK == 0
kk = k0 + rk
a = tl.load(a_ptr + rm[:, None] * stride_am + kk[None, :] * stride_ak,
mask=rm[:, None] < M, other=0.0)
q = tl.load(qw_ptr + (kk[:, None] // 8) * stride_qk + rn[None, :] * stride_qn,
mask=rn[None, :] < N, other=0)
q = (q >> shifts[:, None]) & 0xF # unpack in registers
s = tl.load(s_ptr + (kk[:, None] // G) * stride_sk + rn[None, :] * stride_sn,
mask=rn[None, :] < N, other=0.0)
w = ((q.to(tl.float32) - 8.0) * s.to(tl.float32)).to(tl.float16)
acc += tl.dot(a, w) # tensor cores, fp32 accumulate
tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :] * stride_cn, acc.to(tl.float16),
mask=(rm[:, None] < M) & (rn[None, :] < N))
def w4a16(a, packed, s, G=128):
M, K = a.shape
N = packed.shape[1]
c = torch.empty(M, N, dtype=torch.float16, device=a.device)
grid = lambda meta: (triton.cdiv(M, meta["BM"]), triton.cdiv(N, meta["BN"]))
w4a16_gemm[grid](a, packed, s, c, M, N, K,
a.stride(0), a.stride(1), packed.stride(0), packed.stride(1),
s.stride(0), s.stride(1), c.stride(0), c.stride(1), G=G)
return cRead the inner loop as four steps. Load a BM by BK tile of activations, masking rows beyond M. Load packed row kk // 8 for the same K range, so each int32 serves eight consecutive k and the shift picks the nibble. Unpack, centre and scale in registers. Then tl.dot runs on tensor cores and accumulates in FP32.
The autotune decorator benchmarks each configuration the first time it sees a new (M, N, K) and caches the winner: small BM for decode, large for prefill.
What this simple kernel leaves on the table
It is correct, but it is not a production kernel. Knowing why is the useful part.
- Wasted tensor-core rows at decode. tl.dot needs BM of at least 16, so with M = 1 fifteen of sixteen rows are padding. The memory traffic is still right, which is why it helps at all, but production stacks often use a separate GEMV-style kernel or split-K for the smallest batches.
- No split-K. With M = 1 and N = 4096, a BN of 128 launches only 32 programs, too few to fill a large GPU. Splitting K across programs and reducing partial sums afterwards restores parallelism.
- Dequantization in the critical loop. Shifts, masks, conversions and multiplies all run on CUDA cores between tensor-core calls. At decode the GPU is idle anyway; at prefill this work competes with the matmul.
- Layout is not tuned for the hardware. Production kernels repack weights offline so that one wide load lands each thread exactly the nibbles its tensor-core fragment needs, and use bit tricks that convert several nibbles to FP16 at once.
Two well-known kernels show what closing these gaps buys. Marlin targets compute capability 8.0 and above; its README says Ampere or Ada, not yet optimized for Hopper, and reports close to the ideal 4x speedup up to batch sizes of about 16 to 32 tokens at group size 128. Machete, described in Red Hat's October 2024 announcement as Marlin's successor for Hopper, is built on CUTLASS and CuTe and uses Hopper's TMA and WGMMA instructions to overlap the weight upconversion with data movement and compute; vLLM 0.6.2 and later use it for w4a16 and w8a16 models. Use those in production; write your own to understand them or to support a layout they do not.
Testing: compare against the dequantized weights
The common mistake is comparing against the original FP16 weights, which mixes expected quantization error with kernel error that should be near zero. Compare against A times dequantize(packed) in FP32; any real disagreement is then a kernel bug.
import pytest
SHAPES = [(1, 4096, 4096), (7, 4096, 11008), (33, 11008, 4096), (512, 4096, 4096), (1, 4096, 4160)]
@pytest.mark.parametrize("M,K,N", SHAPES)
def test_w4a16_matches_dequantized_reference(M, K, N):
torch.manual_seed(0)
w = torch.randn(K, N, device="cuda") * 0.02
a = torch.randn(M, K, device="cuda", dtype=torch.float16)
packed, s = quantize_pack(w)
ref = a.float() @ dequantize(packed, s).float() # same quantized weights, fp32 math
out = w4a16(a, packed, s).float()
err = (out - ref).abs().max() / ref.abs().max()
assert err < 2e-3, f"kernel error {err:.2e}" # kernel bugs, not quantization error
assert torch.isfinite(out).all()Choose shapes deliberately. M = 1 and odd values such as 7 or 33 exercise the row masks; an N such as 4160 that is not a multiple of the largest BN exercises the column masks. Include your model's exact projection shapes, because autotuning picks a different configuration for each and a bug can live in only one. Separately, check quantization error against a budget agreed with the model owner; the evaluation methodology article covers setting it.
Benchmarking against the roofline
Microseconds alone do not tell you whether a kernel is good. Convert them into effective bandwidth, bytes the kernel must move divided by time, and compare with what the GPU can sustain. A well-written decode kernel reaches a large fraction of peak DRAM bandwidth; one that reaches a third of it has a problem worth finding.
from triton.testing import do_bench
def bench(M, K, N, G=128):
w = torch.randn(K, N, device="cuda") * 0.02
a = torch.randn(M, K, device="cuda", dtype=torch.float16)
packed, s = quantize_pack(w, G)
w16 = w.half()
t4 = do_bench(lambda: w4a16(a, packed, s, G)) # milliseconds
t16 = do_bench(lambda: a @ w16)
bytes4 = K * N // 2 + s.numel() * 2 + M * K * 2 + M * N * 2
print(f"M={M:5d} int4 {t4*1e3:7.1f} us fp16 {t16*1e3:7.1f} us "
f"speedup {t16/t4:4.2f}x int4 eff. bandwidth {bytes4/(t4*1e-3)/1e9:6.0f} GB/s")
for M in (1, 4, 16, 64, 256, 1024):
bench(M, 4096, 11008)do_bench handles warm-up, repeats and synchronisation. Run on an idle GPU with the real projection shapes. Expect a curve shaped like the table below; the numbers are illustrative, not measurements.
| Tokens M | Bound by | Expected INT4 vs FP16 | What to check if it is worse |
|---|---|---|---|
| 1-16 | weight bytes | approaching 3-4x | effective bandwidth, grid size, split-K |
| 16-128 | mixed | falling towards 1.5-2x | BM choice, autotune key, occupancy |
| 512+ | tensor-core math | about 1x or slightly slower | dequant overhead; consider FP16 path |
The crossover is the operational fact that matters. Measure your own, on your own GPU, rather than borrowing someone else's.
Worked example: one projection, end to end
Take the 4096 by 11008 up-projection of a 7B-class model. Quantize with G = 128: the packed tensor is 512 by 11008 int32 values, about 22.5 MB, plus 32 by 11008 FP16 scales, about 0.7 MB. Run the tests, then the benchmark. At M = 1 the INT4 path should clearly beat FP16; if effective bandwidth sits far below the datasheet figure, print the autotuner's chosen configuration, because a small BN with few programs can starve the memory system. At M = 1024 expect FP16 to catch up. Finally, swap the kernel into one transformer block and compare logits on a fixed prompt against the dequantized-reference model: layer tests catch indexing bugs, model tests catch layout mismatches.
Integrating a kernel into serving
- Repack once, at load, never per call, and verify one layer after repacking.
- Warm the autotuner before traffic by running every shape the scheduler can produce.
- Capture decode with CUDA graphs, since launch overhead rivals kernel time at M = 1.
- Keep a logged fallback to dequantize-then-matmul for unsupported shapes or GPUs.
Failure modes
- Layout mismatch. The checkpoint packs along N, the kernel expects K; the logit test catches it.
- Missing mask after shift. Sign-extended nibbles produce huge values for some weights only.
- Wrong scale indexing. Using k // 8 instead of k // G for scales passes tests where G equals 8 and fails everywhere else.
- FP16 overflow in the epilogue. Large activations overflow when the FP32 accumulator is cast; keep outlier layers in higher precision.
- Benchmarking the cache. Small weights fit in L2 and look impossibly fast; benchmark real sizes.
What to do next
- Write down your format's layout as a packer and dequantizer pair, and treat them as the specification.
- Implement the Triton kernel above for your layout and make the shape-parametrised test pass against the dequantized reference.
- Benchmark your model's real projection shapes from M = 1 to M = 1024 and record effective bandwidth and the INT4-versus-FP16 crossover.
- Compare against the production kernel your serving stack uses (Marlin on Ampere or Ada, Machete on Hopper in vLLM) on the same shapes.
- Add a model-level logit comparison to CI for every kernel or layout change.
- Route prefill through whichever path wins above the crossover, and warm all shapes before taking traffic.