Papers
arxiv:2609.33591

Pretraining Transformers with Quantized Softmax in Attention

Published on Sep 27
· Submitted by
Shangzhen Zhu
on Sep 30
Authors:
,

Abstract

Low-precision Transformer systems increasingly quantize attention matrix multiplications, while softmax often remains at higher precision. During pretraining, an approximate softmax changes the gradients that train the model as well as its forward computation. We study this interaction with K-interval attention, which approximates the exponential using K+1 grid values. We vary per-row grid calibration, interpolation versus hard rounding, and the placement of a straight-through surrogate relative to normalization. We derive the corresponding backward rules, including calibration derivatives, and compare these choices in pretraining experiments matched on model, data, and optimizer. Detaching the row extrema leaves the forward computation unchanged but produces a delayed increase in validation loss. With hard rounding at K=4, min-max calibration and a pre-normalization surrogate incur a large loss gap; changing either choice substantially reduces it. At 124M parameters and 2.5B training tokens, fixed-window calibration with a post-normalization surrogate yields a validation loss gap of +0.019 nats relative to softmax at K=4, and with a pre-normalization surrogate yields +0.004 nats at K=16.

Community

Paper author Paper submitter

When you approximate softmax during pretraining, the approximation defines the learning rule. We study which backward choices make a quantized softmax trainable from scratch.

Key findings (GPT-2-style 124M/1B, up to 2.5B tokens, 5-seed replication):

  • Same forward, delayed failure. Detaching the row max/min, a harmless habit for exact softmax and standard in fake-quant observers, leaves the forward unchanged but drops calibration gradient terms. Training tracks softmax for ~25–30M tokens, then diverges and ends 0.65–3.07 nats worse. Keeping the gradient through the row max alone recovers almost all of the gap.
  • Calibration × surrogate placement interact strongly. With hard rounding at K=4, min–max calibration + a weight-level straight-through estimator ends +0.89 nats behind softmax at 2.5B tokens. Changing either choice (a fixed window anchored at the row max, or moving the STE after normalization) removes most of the deficit. The direction holds at 1B parameters and on all 5 seeds.
  • Coarse quantization can still train well. Fixed-window + post-normalization STE: +0.019 nats at K=4; fixed-window at K=16: +0.004 nats (124M, 2.5B tokens). Interpolated reconstruction at K=4–32 ends within 0.005 nats of softmax.
  • Downstream transfer is graded: large NLL gaps show up on every benchmark, but a small NLL gap does not guarantee parity on every task.

Takeaway: if you quantize softmax for training, evaluate calibration, reconstruction, and backward rule together, and consider the gradient through the row max of attention score matrix.

Companion inference-side paper (softmax approximation in frozen LLMs + a FlashAttention-4 kernel): https://arxiv.org/abs/2609.33586

Code will be released. Questions welcome!

Sign up or log in to comment

Get this paper in your agent:

hf papers read 2609.33591
Don't have the latest CLI?
curl -LsSf https://hf.co/cli/install.sh | bash

Models citing this paper 0

No model linking this paper

Cite arxiv.org/abs/2609.33591 in a model README.md to link it from this page.

Datasets citing this paper 0

No dataset linking this paper

Cite arxiv.org/abs/2609.33591 in a dataset README.md to link it from this page.

Spaces citing this paper 0

No Space linking this paper

Cite arxiv.org/abs/2609.33591 in a Space README.md to link it from this page.

Collections including this paper 0

No Collection including this paper

Add this paper to a collection to link it from this page.