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

Advertisement

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.

Quantize points inside one FlashAttention-style tile stepQ tileper-block scale sQK tilesmoothed: K - mean(K)Low-precision GEMM 1INT8 or FP8, exact int32/fp32 accDequant + softmaxS * sQ * sK, online, FP32SP tilein [0, 1]; keep FP16 or FP8GEMM 2: P x VFP16 or FP8, careful accumulatorV tileper-channel or per-blockOutput accumulatorrescaled each tile, FP32GEMM 1 tolerates INT8 or INT4 once K is smoothed; GEMM 2 is where naive INT8 breaks.Optional: rotate Q and K by the same orthogonal matrix first (FlashAttention-3 FP8).
The two GEMMs inside a tile step. Scores are dequantized with one multiply before an FP32 online softmax; the probability tile then feeds a second GEMM whose precision and accumulator need more care.
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 / l

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

Advertisement

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 j

So 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

GranularityScales per headAccuracyOverhead
Per tensor1Lowest; any outlier sets the scaleNone
Per block (one tile of tokens)N / block sizeGood with smoothingOne multiply per tile
Per tokenNBetter for Q and KScales must be applied per row and column
Per thread (SageAttention2, INT4)Several per tileBest for INT4Kernel-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

OptionSpeedRiskUse when
FP16/BF16 attentionBaselineNoneShort contexts, decode, sensitive models
INT8 QK + FP16 PVHigh on GPUs with fast INT8Low with K smoothingPrefill and diffusion on most recent GPUs
INT4 QK + FP8 PVHigherModerate; per-thread scales neededThroughput-critical inference after per-layer checks
FP8 both GEMMs + rotationHigh on FP8 tensor coresModerateHopper-class serving and training forward passes

What to do next

  1. Profile one long-context prefill or diffusion step and measure what fraction of time is attention. Below about a third, look elsewhere first.
  2. Capture Q, K and V from real prompts and run the simulator per layer, with and without K smoothing. Note the worst layer.
  3. Try a published quantized attention kernel as a drop-in replacement on that model and compare task metrics, not just cosine similarity.
  4. Keep the worst layers, often the first and last, in 16-bit attention if they fail your threshold.
  5. Validate at your longest production sequence length and with your real masks.
  6. For decode-heavy serving, quantize the KV cache instead and keep the attention math in 16 bits.
Key takeaway: Quantizing attention means running Q times K-transpose and P times V on low-precision tensor cores inside a FlashAttention-style loop. It pays off when attention dominates, as in long prefills and diffusion, and not in memory-bound decode. The first GEMM tolerates INT8 and even INT4 once K is smoothed by subtracting its token mean, which leaves softmax exactly unchanged, or once Q and K are rotated by a shared orthogonal matrix. The second GEMM is fragile because P is mostly small probabilities, so it stays in FP16 or uses FP8 with careful accumulation. Simulate per layer on real activations and check the worst layer before shipping.