Weight-only quantization shrinks a model, but the matrix multiplies still run in 16-bit floating point: weights are expanded back to bf16 before or inside the GEMM. To use the int8 or fp8 tensor cores, and to halve the bytes moved for activations, the other operand has to be quantized too. That operand is the activation, the tensor flowing between layers, and it is a much less cooperative thing to quantize.
This article treats activation quantization as a datapath: where in a transformer block the quantize and dequantize steps sit, the integer arithmetic a quantized GEMM actually performs, how production kernels fuse those steps so they cost almost nothing, and how to find which site is losing your accuracy. Choosing between dynamic and static scales is covered in dynamic vs static quantization, and scale granularity in quantization granularity; this page builds on both.
Why activations are harder than weights
Weights are fixed. You can inspect every value offline, spend minutes choosing scales, and even adjust the weights themselves to compensate for rounding. Activations exist only at run time, depend on the input, and change from one token to the next.
Two empirical facts about large language models make this worse. First, a handful of hidden channels carry values tens of times larger than the rest, consistently across tokens; these systematic outliers appear in the inputs to the attention and MLP projections. Second, a few tokens, such as the first token or delimiters, carry extremely large values in a few specific channels. One scale per tensor must cover the largest value, so the ordinary values are rounded onto a few integer levels; with per-token scales, that token's row gets a huge maximum and the damage stays in that row.
A small example shows the size of the effect. Suppose a row of activations is mostly within plus or minus 2 but one outlier channel holds 60. Symmetric int8 with one scale gives s = 60 / 127, about 0.47. A value of 0.3 becomes round(0.3 / 0.47) = 1, which dequantizes to 0.47, an error of more than 50 percent. Without the outlier, s = 2 / 127, about 0.016, and the same value is represented to within about 3 percent. Everything in activation quantization, from per-token scales to SmoothQuant and rotations, is a way of dealing with that one outlier.
The quantizer itself
An affine quantizer maps a real value x to an integer q with a scale s and a zero point z:
q = clamp(round(x / s) + z, qmin, qmax) # quantize
x_hat = s * (q - z) # dequantize
# symmetric int8 (z = 0), the common choice for activations feeding int8 GEMMs:
s = max(abs(x)) / 127Two kinds of error compete. Rounding error is at most s/2 per value and shrinks as the scale shrinks. Clipping error appears when the scale is set below the true maximum so values outside the range saturate. Calibration methods trade the two; per-token dynamic scaling largely removes clipping by measuring each row's maximum at run time.
Symmetric quantization is preferred for GEMM inputs because a zero point adds work inside the matrix multiply, shown in the next section. Asymmetric quantization helps for one-sided distributions, such as the output of a ReLU, but those are rare in modern transformer blocks, where most linear-layer inputs come from a normalization layer and are roughly centred.
Where the quantize sites are
A decoder block does not quantize everything. The diagram marks the four standard sites, the inputs to the four linear layers, and the tensors that stay in high precision.
| Site | Typical treatment | Why |
|---|---|---|
| Q1: QKV input (after norm) | int8 or fp8, per token | Systematic outlier channels from the norm's scale |
| Q2: O-proj input | int8 or fp8, per token | Usually well behaved |
| Q3: gate/up input (after norm) | int8 or fp8, per token | Same outlier pattern as Q1 |
| Q4: down-proj input | int8 or fp8, per token; first site to keep in bf16 if accuracy drops | Product of SiLU and up-projection has very heavy tails |
| Q x K and P x V | Often left in bf16; fp8 in some attention kernels | Softmax output spans many orders of magnitude |
| KV cache | Separate decision, see KV cache quantization | Memory-bound in decode, different trade-off |
| Residual stream, norms, logits | Never quantized in practice | Errors accumulate across layers and directly change token choice |
The residual stream deserves emphasis. Every block adds its output to it, so an error there propagates through all later layers. Quantizing a block's inputs is safe because the error stops at the block's output; quantizing the residual stream is not.
The integer GEMM, written out
Consider one output element of a linear layer, y = sum over k of x_k times w_k, with activations quantized per token (scale s_a, zero point z_a) and weights per output channel (scale s_w, symmetric). Substituting the dequantization formula gives
y = sum_k s_a (qa_k - z_a) * s_w qw_k
= s_a * s_w * ( sum_k qa_k * qw_k - z_a * sum_k qw_k )
\_______________/ \___________/
int32 GEMM per-channel constant,
on tensor cores precomputed from weightsThe expensive term is a pure integer dot product, accumulated in int32. The scales come out of the sum entirely, so they are applied once per output element in the epilogue. With a symmetric activation quantizer the correction term disappears; with an asymmetric one it costs a precomputed per-channel sum times the per-token zero point, cheap but easy to forget.
The accumulator is safe: each product is at most 127 times 127 = 16,129 in magnitude, and int32 holds about 2.1 billion, so a dot product of up to about 133,000 terms cannot overflow even in the worst case, far beyond any hidden size in use. A numpy reference implementation is the right tool for testing a kernel against:
import numpy as np
def quantize_per_token(x): # x: [tokens, K] float32
s = np.abs(x).max(axis=1, keepdims=True) / 127.0
s = np.maximum(s, 1e-8) # all-zero rows
q = np.clip(np.rint(x / s), -127, 127).astype(np.int8)
return q, s.astype(np.float32)
def quantize_per_channel(w): # w: [N, K] float32
s = np.abs(w).max(axis=1, keepdims=True) / 127.0
q = np.clip(np.rint(w / s), -127, 127).astype(np.int8)
return q, s.astype(np.float32)
def int8_linear(x, w, bias=None):
qa, sa = quantize_per_token(x)
qw, sw = quantize_per_channel(w)
acc = qa.astype(np.int32) @ qw.astype(np.int32).T # [tokens, N] int32
y = acc.astype(np.float32) * sa * sw.T # epilogue: per token x per channel
return y if bias is None else y + bias
def sqnr_db(ref, approx):
noise = np.sum((ref - approx) ** 2)
return 10 * np.log10(np.sum(ref ** 2) / max(noise, 1e-30))
rng = np.random.default_rng(0)
x = rng.standard_normal((16, 4096)).astype(np.float32)
x[:, 7] *= 40 # one outlier channel
w = (rng.standard_normal((1024, 4096)) * 0.02).astype(np.float32)
print(sqnr_db(x @ w.T, int8_linear(x, w)))Running variants of this harness, with the outlier channel removed, with per-tensor instead of per-token scales, or with a smoothing factor applied, is the fastest way to build intuition for how much each technique buys on your own activation statistics.
Fusion: making quantization nearly free
Written naively, activation quantization adds kernels: one to find each row's maximum, one to divide and round, and one to dequantize the output. Each reads and writes the full activation tensor through high-bandwidth memory. Production kernels fold these steps into their neighbours.
On the input side, the RMSNorm kernel already holds each row in registers or shared memory, so it can compute the row's absolute maximum, write int8 values and write one float scale per token in the same pass. For Q4, the SiLU-and-multiply kernel does the same. The GEMM then reads int8, half the bytes of bf16.
On the output side, the GEMM's epilogue, the code that runs after the int32 accumulation and before the store, multiplies by the per-token scale and per-channel weight scale, adds the bias and writes bf16. Some kernels go further and quantize the output for the next GEMM directly, but that only works when the next consumer is another linear layer, which in a transformer block it rarely is, because a residual add or norm sits in between.
Kernel pseudocode for the fused norm and quantize step:
# one program instance per token row
x = load(row, K) as fp32
rms = sqrt(mean(x * x) + eps)
y = x / rms * gamma # RMSNorm
amax = max(abs(y)) # row reduction, already in registers
s = max(amax, 1e-8) / 127
store(q_out[row], clamp(round(y / s), -127, 127) as int8)
store(scale_out[row], s)Per-token dynamic scaling therefore costs one extra reduction over data already in registers. That is why it has become the default for int8 and fp8 activations in LLM serving, and static per-tensor scales are mostly reserved for hardware or runtimes that cannot afford the reduction.
FP8 activations in the same datapath
With fp8 the structure is the same but the quantizer is a cast. E4M3, the format used for activations in inference, has a maximum normal value of 448 and spaces its levels logarithmically, so small values keep relative precision that int8 would round away. A scale is still applied so the tensor's largest value lands near 448, either per tensor from calibration or per token at run time. The GEMM accumulates in higher precision and the epilogue applies the scales just as in the int8 case. Format details and training use are in FP8 formats; for this datapath the practical difference is that fp8 tolerates the remaining outliers better, which often lets Q4 stay quantized where int8 would have to fall back to bf16.
Finding the site that hurts
When a W8A8 model loses accuracy, the useful question is which site, in which layer. Measure the signal-to-quantization-noise ratio at every linear input with forward pre-hooks, which PyTorch provides as register_forward_pre_hook on any module:
import torch
def fake_quant_per_token(x):
s = x.abs().amax(dim=-1, keepdim=True).clamp_min(1e-8) / 127
return (x / s).round().clamp(-127, 127) * s
def site_sqnr(model, batches):
stats, hooks = {}, []
def make_hook(name):
def hook(module, args):
x = args[0].float()
noise = (x - fake_quant_per_token(x)).pow(2).sum()
sig = x.pow(2).sum()
s, n = stats.get(name, (0.0, 0.0))
stats[name] = (s + sig.item(), n + noise.item())
return hook
for name, m in model.named_modules():
if isinstance(m, torch.nn.Linear):
hooks.append(m.register_forward_pre_hook(make_hook(name)))
with torch.no_grad():
for b in batches:
model(**b)
for h in hooks:
h.remove()
return sorted(((10 * torch.log10(torch.tensor(s / max(n, 1e-30))).item(), k)
for k, (s, n) in stats.items()))Sort ascending and the worst sites come first. A typical result shows most sites above about 30 dB and a few down-projection inputs, often in the first and last layers, well below. Fix those first: keep them in bf16, apply smoothing or rotation, or switch them to fp8. Then confirm on a task evaluation, because SQNR measures the tensor, not the model's output.
Failure modes
- Silent per-tensor fallback: a runtime that does not support per-token scales for one layer type quietly uses a per-tensor scale. Inspect the scale tensor shapes in the exported model.
- Stale static scales: calibrated on short English prompts, served long code or other languages; activations exceed the range and clip.
- Missing zero-point correction: an asymmetric activation quantizer paired with a kernel that assumes symmetric input gives a constant per-channel bias error.
- All-zero rows: padding tokens give a zero scale and a division by zero unless the scale is clamped.
- Quantizing the residual stream: errors accumulate across layers and appear as drift on long generations.
- Prefill and decode mismatch: different kernels for the two phases with different quantization schemes, so outputs differ depending on prompt length.
Trade-offs
| Choice | Gain | Cost |
|---|---|---|
| Weight-only (W4A16, W8A16) | No activation risk; best for memory-bound decode | No low-precision tensor-core compute |
| W8A8 int8, per-token | Int8 tensor cores, half the activation bytes | Outliers must be handled; Q4 often sensitive |
| W8A8 fp8, per-token | More tolerant of outliers, simple cast | Needs hardware with fp8 tensor cores |
| Static per-tensor scales | No run-time reduction | Clipping on out-of-distribution inputs |
| Mixed: sensitive sites in bf16 | Recovers most accuracy | Extra kernels and a less uniform graph |
What to do next
- Build the numpy reference GEMM above and use it as the ground truth for any kernel you adopt.
- Run the per-site SQNR hook on a few hundred representative prompts and list the ten worst sites.
- Check exported scale shapes to confirm per-token activation scales are actually used everywhere.
- Keep the residual stream, norms, softmax and logits in high precision.
- Try fp8 or bf16 at the worst down-projection sites before reaching for more complex methods.
- Evaluate on your real task with long inputs, not only perplexity on short text.