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.

Advertisement

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)) / 127

Two 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.

Advertisement

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.

Activation quantization sites in one decoder block (W8A8 example)Residual streambf16, never quantizedRMSNorm+ fused quantizeQKV GEMMint8 x int8 -> int32Attentionsoftmax in fp; KV opt. int8/fp8Q1dequantO-proj GEMMQ2Residual addbf16dequantRMSNorm+ fused quantizeGate/Up GEMMQ3SiLU(g) * uheavy-tailed outputDown GEMMhardest inputQ4Residual addbf16dequantQ1-Q4: quantize points for linear-layer inputs. Each is ideally fused into the kernel that produces the tensor.Dequant: applied in the GEMM epilogue as s_a (per token) x s_w (per channel), with bias, before writing bf16.Kept in high precision: residual stream, norms, softmax, and usually the logits.Q4 sees SwiGLU products with the widest dynamic range; it is where most W8A8 and W4A8 accuracy is lost.
Quantize points Q1-Q4 sit at linear-layer inputs. The residual stream, norms and softmax stay in bf16 or fp32.
SiteTypical treatmentWhy
Q1: QKV input (after norm)int8 or fp8, per tokenSystematic outlier channels from the norm's scale
Q2: O-proj inputint8 or fp8, per tokenUsually well behaved
Q3: gate/up input (after norm)int8 or fp8, per tokenSame outlier pattern as Q1
Q4: down-proj inputint8 or fp8, per token; first site to keep in bf16 if accuracy dropsProduct of SiLU and up-projection has very heavy tails
Q x K and P x VOften left in bf16; fp8 in some attention kernelsSoftmax output spans many orders of magnitude
KV cacheSeparate decision, see KV cache quantizationMemory-bound in decode, different trade-off
Residual stream, norms, logitsNever quantized in practiceErrors 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 weights

The 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.

Unfused vs fused activation quantization: the same math, a third of the memory trafficUnfusedNorm kernelAbsmax kernelQuantize kernelGEMMDequantbf16bf16+sint8int32FusedNorm + row absmax + quantizeone read of x, writes int8 + per-token scaleGEMM + epilogueint32 acc -> x s_a x s_w + bias -> bf16int8, s_aEach round trip through HBM costs bandwidth. In decode, activations are tiny and the win is launch count;in prefill, activations are large and the win is bytes moved.
The quantize step joins the kernel that produced the activation; the dequantize step joins the GEMM epilogue.

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

ChoiceGainCost
Weight-only (W4A16, W8A16)No activation risk; best for memory-bound decodeNo low-precision tensor-core compute
W8A8 int8, per-tokenInt8 tensor cores, half the activation bytesOutliers must be handled; Q4 often sensitive
W8A8 fp8, per-tokenMore tolerant of outliers, simple castNeeds hardware with fp8 tensor cores
Static per-tensor scalesNo run-time reductionClipping on out-of-distribution inputs
Mixed: sensitive sites in bf16Recovers most accuracyExtra kernels and a less uniform graph

What to do next

  1. Build the numpy reference GEMM above and use it as the ground truth for any kernel you adopt.
  2. Run the per-site SQNR hook on a few hundred representative prompts and list the ten worst sites.
  3. Check exported scale shapes to confirm per-token activation scales are actually used everywhere.
  4. Keep the residual stream, norms, softmax and logits in high precision.
  5. Try fp8 or bf16 at the worst down-projection sites before reaching for more complex methods.
  6. Evaluate on your real task with long inputs, not only perplexity on short text.
Key takeaway: Activation quantization is a datapath decision: quantize the inputs of linear layers, keep the residual stream and softmax in high precision, and let the GEMM run on integers with the scales applied once in the epilogue. Per-token scales, fused into the norm or activation kernel, cost almost nothing and remove most clipping. When accuracy drops, measure each site, fix the few that hurt, and verify against a reference implementation.