"Quantizing attention" can mean three different things. You can quantize the weights of the Q, K, V and output projections, which is ordinary weight quantization. You can store the KV cache in fewer bits, which is a memory optimisation covered in KV cache quantization. Or you can run the two attention matrix multiplications, Q times K-transpose and P times V, on low-precision tensor cores. This article is about the third. It matters most for long-sequence prefill and for diffusion models, where attention dominates the compute.
You will see where attention FLOPs go, why the two matrix multiplications need different treatment, the smoothing and rotation tricks, what published kernels do, and how to measure the error on your own model.
When attention compute is worth quantizing
For one head, Q times K-transpose costs 2·N²·d FLOPs and P times V costs another 2·N²·d, where N is sequence length and d is head dimension. The linear layers cost about 2 FLOPs per parameter per token, so they grow with N while attention grows with N².
Take a layer shaped like a common 8B model: width 4096, 32 heads of dimension 128, an MLP of 14,336, and grouped-query attention ignored for simplicity. The projections and MLP hold about 4·4096² + 3·4096·14,336 ≈ 243 million parameters, so at N = 32,768 tokens they cost about 2 × 243M × 32,768 ≈ 1.6 × 10¹³ FLOPs. Full bidirectional attention costs 4 × 32,768² × 128 × 32 ≈ 1.76 × 10¹³. Causal masking halves that. So at 32k tokens attention is already as expensive as everything else in the layer, and at 128k it is several times more.
Decode is different. Generating one token computes attention for a single query against the cached keys, which takes a few FLOPs per byte read, so the step is limited by memory bandwidth. Low-precision tensor cores do not help a memory-bound step. Storing the cache in fewer bits does. Use attention-compute quantization for prefill, training forward passes and diffusion, and use KV-cache quantization for decode.
The datapath and its quantize points
Modern attention kernels follow the FlashAttention structure: tile the queries, stream key and value tiles through on-chip memory, keep a running softmax, and never write the N × N score matrix to memory (see FlashAttention). Quantization adds scale factors to that loop. It does not change its structure.
k_mean = mean(K, over tokens) # computed once per head, outside the loop
Qi, sQ = quant_int8(Qi_tile) # per-block scale
m, l, acc = -inf, 0, zeros(Br, d) # online-softmax state, FP32
for j in key_tiles:
Kj, sK = quant_int8(K[j] - k_mean)
S = int8_gemm(Qi, Kj.T) * (sQ * sK) * softmax_scale
m_new = max(m, rowmax(S))
P = exp(S - m_new) # FP32
l = l * exp(m - m_new) + rowsum(P)
acc = acc * exp(m - m_new) + gemm_16bit(P.to(fp16), V[j]) # or FP8 with scales
m = m_new
O_i = acc / lTwo properties make this cheap. Dequantizing the first GEMM costs one multiply per score, because a per-block scale for Q and a per-block scale for K factor out of the dot product. The numerical risk sits in two places: quantizing Q and K, and the second GEMM.
Q times K: channel outliers and the smoothing trick
Keys in trained transformers commonly have a few channels with a large offset shared by every token. With a per-block scale, the outlier channel sets the scale for the whole tile, and every other channel is squeezed into a handful of integer levels.
SageAttention's fix rests on a property of softmax. Subtracting the mean key from every key changes every score in a query's row by the same constant, and softmax ignores constant shifts:
softmax(q · (k_j - k_mean)) = softmax(q · k_j - q · k_mean) = softmax(q · k_j) # q · k_mean is the same for every jSo K - mean(K), with the mean taken over tokens, gives exactly the same attention output and removes the shared offset before quantization. The cost is one reduction over K per head, which can be fused into the kernel that produces K. SageAttention2 adds a smoothing step for Q as well (the correction term no longer cancels in softmax, so it has to be added back) and pushes Q and K down to INT4 with per-thread scales, which are matched to how the tensor core distributes tile fragments across threads.
Worked example: what smoothing buys at INT8
Suppose one channel of K carries an offset of about 8 with a spread of about ±1, and the other channels sit within ±1. With a per-block INT8 scale, the largest magnitude in the tile is about 9, so the scale is 9 / 127 ≈ 0.071. A channel that varies within ±0.3 then gets 0.6 / 0.071 ≈ 8 distinct levels, which is close to 3-bit precision for most of the information in the tile.
After subtracting the mean, the largest magnitude is about 1, the scale becomes 1 / 127 ≈ 0.0079, and the same ±0.3 channel gets about 76 levels. That is a ninefold finer grid for a single subtraction, and the attention output is mathematically unchanged. This is the same idea that SmoothQuant applies to linear-layer activations, but here it is exact rather than a rescaling that has to be folded into weights.
P times V: why the second GEMM is different
The probability matrix P is not like other activations. Its values lie between 0 and 1, most are close to zero, and a few dominate each row. Quantizing it to INT8 with a fixed scale of 1/127 rounds most small probabilities to zero, and their combined weight can be large in long contexts. SageAttention measured this directly: quantizing P and V naively to INT8 gave a worst-case cosine similarity of 56.40% across layers, against 99.99% when the product stayed in FP16.
The published answers differ. SageAttention keeps P times V in FP16 and uses an FP16 accumulator, which it reports as accurate and faster than FP32 accumulation on the GPUs it targets. SageAttention2 quantizes P and V to FP8, which suits P because a floating-point format spends its levels where small values live. It adds a two-level accumulation scheme to limit the error from the reduced-precision accumulator. FlashAttention-3 runs both GEMMs in FP8 on Hopper with FP32 softmax and accumulators, and it transposes V in the kernel because the FP8 tensor-core instruction needs V contiguous along the sequence dimension.
Granularity and rotation
| Granularity | Scales per head | Accuracy | Overhead |
|---|---|---|---|
| Per tensor | 1 | Lowest; any outlier sets the scale | None |
| Per block (one tile of tokens) | N / block size | Good with smoothing | One multiply per tile |
| Per token | N | Better for Q and K | Scales must be applied per row and column |
| Per thread (SageAttention2, INT4) | Several per tile | Best for INT4 | Kernel-specific layout |
Rotation is the other tool. FlashAttention-3's FP8 path multiplies Q and K by the same random orthogonal matrix M before quantizing. Because (QM)(KM)ᵀ = QKᵀ, the scores are unchanged, but the rotation spreads each outlier across all channels, so no single channel dominates the scale. M is built from a random ±1 diagonal and a Hadamard matrix, so applying it costs O(d log d) per vector instead of a full matrix multiply. On synthetic data where 0.1% of entries were outliers at ten times the standard deviation, the paper reports RMSE of 9.1e-3 with block quantization and rotation, against 2.4e-2 for per-tensor FP8, a 2.6× improvement. The same idea, applied to weights and activations, underpins rotation-based quantization.
Simulate the error before you touch a kernel
Kernels are hard to debug, and numerics are easy to simulate. The code below fake-quantizes Q, K and optionally P and V in PyTorch, compares the result with full-precision attention, and reports cosine similarity and relative L1 error, the metrics the SageAttention papers use.
import torch
def quant_sym(x, bits, block):
# Symmetric fake-quant with one scale per block of `block` tokens (all channels).
B, H, N, D = x.shape
qmax = 2 ** (bits - 1) - 1
xb = x.reshape(B, H, N // block, block, D)
scale = xb.abs().amax(dim=(-2, -1), keepdim=True).clamp(min=1e-8) / qmax
q = torch.round(xb / scale).clamp(-qmax, qmax)
return (q * scale).reshape(B, H, N, D)
def attention(q, k, v):
s = (q @ k.transpose(-1, -2)) / q.shape[-1] ** 0.5
return torch.softmax(s.float(), dim=-1) @ v.float()
def attention_quant(q, k, v, qk_bits=8, block=64, smooth_k=True, pv="fp16"):
if smooth_k:
k = k - k.mean(dim=-2, keepdim=True) # same shift for every key: softmax unchanged
qd, kd = quant_sym(q, qk_bits, block), quant_sym(k, qk_bits, block)
s = (qd @ kd.transpose(-1, -2)) / q.shape[-1] ** 0.5
p = torch.softmax(s.float(), dim=-1)
if pv == "fp16": # on CPU, emulate: (p.half().float() @ v.half().float())
return (p.half() @ v.half()).float()
if pv == "int8": # naive: P scaled by 1/127, V per block
pq = torch.round(p * 127).clamp(0, 127) / 127
return pq @ quant_sym(v, 8, block)
return p @ v.float()
def report(ref, out):
cos = torch.nn.functional.cosine_similarity(ref.flatten(), out.flatten(), dim=0)
rel_l1 = (ref - out).abs().sum() / ref.abs().sum()
return f"cos={cos.item():.5f} relL1={rel_l1.item():.4f}"torch.manual_seed(0)
B, H, N, D = 1, 8, 1024, 128
dev = "cuda" if torch.cuda.is_available() else "cpu"
q, k, v = (torch.randn(B, H, N, D, device=dev) for _ in range(3))
k[..., 7] += 8.0 # a channel-wise offset in K, the pattern seen in real models
ref = attention(q, k, v)
for smooth in (False, True):
for pv in ("fp16", "int8"):
out = attention_quant(q, k, v, smooth_k=smooth, pv=pv)
print(f"smooth_k={smooth!s:5} pv={pv:5}", report(ref, out))On this synthetic data, smoothing K roughly halves the relative L1 error of the INT8 Q times K path (expect roughly 2.5% falling to 1.3%). Naive INT8 P times V is far worse, with cosine similarity near 0.71, because attention over 1,024 random keys is close to uniform and almost every probability is below 1/254, so it rounds to zero. Treat this as a smoke test only. The real step is to register forward hooks on each attention layer, capture Q, K and V from a few hundred real prompts, and print the worst layer rather than the mean. A model can look fine on average while one layer near the input or output loses the information that matters.
What the published kernels claim
Every figure here comes from the authors' own measurements, so check it on your own hardware and model. SageAttention reports about 2.1× the speed of FlashAttention2 and 2.7× that of xformers, with almost no end-to-end metric loss across language, image and video models. It explains the choice of INT8 over FP8 by noting that INT8 matrix multiplies on several common GPUs run four times faster than FP16. SageAttention2 reports about 3× FlashAttention2 and 4.5× xformers on an RTX 4090, and roughly matches FlashAttention-3 FP8 speed on Hopper with higher accuracy. FlashAttention-3 reports FP8 throughput close to 1.2 PFLOPs/s on H100.
Most of these kernels are packaged as drop-in replacements for a scaled-dot-product-attention call, which makes trials cheap. Not all of them support every mask, head dimension or backward pass, so read the support matrix first.
Failure modes
- Smoothing in the wrong place. Subtracting the mean over channels instead of tokens, or computing it on a different sequence than the one attended over, changes the output.
- Outliers in Q, not K. Some layers carry the problem in queries, where K smoothing does not help. Per-token scales, Q smoothing or rotation are the remedies.
- Masks and padding. Padding tokens with large values inflate block scales, and the mean of K should cover only valid tokens.
- Long-context drift. Error in P times V grows with the number of small probabilities, so a kernel validated at 4k tokens can fail at 128k. Validate at your longest length.
- Averaged metrics. One fragile layer hidden in an average. Report the worst layer and keep it in higher precision if needed.
Trade-offs
| Option | Speed | Risk | Use when |
|---|---|---|---|
| FP16/BF16 attention | Baseline | None | Short contexts, decode, sensitive models |
| INT8 QK + FP16 PV | High on GPUs with fast INT8 | Low with K smoothing | Prefill and diffusion on most recent GPUs |
| INT4 QK + FP8 PV | Higher | Moderate; per-thread scales needed | Throughput-critical inference after per-layer checks |
| FP8 both GEMMs + rotation | High on FP8 tensor cores | Moderate | Hopper-class serving and training forward passes |
What to do next
- Profile one long-context prefill or diffusion step and measure what fraction of time is attention. Below about a third, look elsewhere first.
- Capture Q, K and V from real prompts and run the simulator per layer, with and without K smoothing. Note the worst layer.
- Try a published quantized attention kernel as a drop-in replacement on that model and compare task metrics, not just cosine similarity.
- Keep the worst layers, often the first and last, in 16-bit attention if they fail your threshold.
- Validate at your longest production sequence length and with your real masks.
- For decode-heavy serving, quantize the KV cache instead and keep the attention math in 16 bits.