Skip to main content
Back to timeline
arXivSource publication:

Quantizing softmax inside attention: whether pretraining works from scratch depends on the backward rule, with detached row-extrema gradients diverging late and MinMax plus Weight-STE trailing softmax by 0.89 nats at 2.5B tokens

Synopsis

This work studies replacing exact softmax with a quantized softmax during pretraining (K-interval attention, which approximates the exponential with K+1 grid values), derives the corresponding backward rules including calibration derivatives, and compares per-row grid calibration (MinMax versus a fixed window, FWM), interpolation (LERP) versus hard rounding (Nearest), and placement of a straight-through surrogate before normalization (Weight-STE) or after it (Prob-STE) in pretraining experiments matched on model, data and optimizer; it finds that detaching the row extrema leaves the forward unchanged but causes a delayed divergence after 25–30M tokens ending 0.65–3.07 nats above the matched run, that under hard rounding MinMax with Weight-STE ends 0.

AI-generated editorial illustration: Pretraining Transformers with Quantized Softmax in Attention

Interpretation

The same forward computation can fail to train solely because the backward rule is incomplete: detaching the row extrema leaves the forward numerically unchanged but drops two Jacobian terms and violates the zero-sum identity implied by shift invariance, so training tracks softmax for 25–30M tokens, then diverges and ends 0.65–3.07 nats above the matched run with full calibration gradients, replicated on two further seeds. Subtracting the row maximum had been treated as a numerical convenience whose gradient cancels exactly, and that habit carried over to a calibrated grid; this work puts the calibration statistics back on the autograd tape, gives closed-form calibration derivatives and the zero-sum identity, and shows that detaching injects a spurious common-mode component along a direction that cannot change the loss. The full-gradient LERP run shows no detectable difference from softmax (95% CI includes zero), while the same forward with the detached backward ends 0.65–3.07 nats higher; the failure is delayed, appearing after 25M tokens and separating near 30M, and replicates on seeds 7 and 42; Nearest fails at every K against its matched run.

Projecting the detached gradients onto the zero-sum subspace does not repair the result: the projection removes the common-mode component exactly, yet that run ends worse than plain detach, whereas retaining only the maximum-dependent calibration gradient recovers almost the whole full-versus-detach gap and retaining only the minimum-dependent gradient recovers almost none of it. This separates the explanation 'the zero-sum identity is broken' from 'the row-maximum channel is missing', showing that zero-sum alone does not explain the outcomes and locating the dominant channel on the maximum side. The P1 projection run is bitwise equal to detach in the forward and ends worse; the maximum-only backward ends within 0.015 nats of full on both seeds, while the minimum-only backward recovers almost nothing; fixed-upstream decompositions at 20M and 40M agree.

Under hard rounding (Nearest), calibration policy and surrogate placement interact strongly: MinMax with Weight-STE at K=4 ends 0.89 nats above softmax at 250M tokens and 0.89 nats at 2.5B; moving the surrogate to the probability level (Prob-STE) or replacing MinMax with the fixed window FWM each removes most of the deficit, and changing both adds little more, so the interaction is large and positive. A reconstruct-then-normalize operator creates a weight-level versus probability-level surrogate choice that shares one hard forward but evaluates the normalization Jacobian at different points; previously the straight-through estimator in Gumbel-softmax had only one place to sit. The 250M-token balanced factorial (calibration by surrogate) estimates main effects and the interaction; the direction holds at 1B parameters and on all five seeds; averaged over five seeds FWM reduces NLL by 0.067 nats relative to MinMax and Prob-STE by 0.021 relative to Weight-STE, with a Kendall's tau of 0.91 for the condition ordering.

Coarse deterministic reconstruction need not prevent near-baseline performance: LERP at the tested K ends within about 0.01 nats of softmax at 2.5B tokens and FWM–Weight at K=16 likewise, but a small NLL gap does not establish downstream equivalence, since every large-gap condition scores below softmax on every benchmark configuration while conditions within 0.01 nats show small, task-dependent differences in both directions. The work reframes 'can an approximate softmax be used for pretraining' from a pointwise error-budget question into a joint evaluation of calibration, reconstruction and backward rule, and observes that the training budget changes the relative standing of hard rounding and interpolation. Fifteen runs at 2.5B tokens, twelve runs at 1B parameters and 100M tokens, and 130 runs on five seeds; downstream ordering is preserved on WikiText-103, PTB and C4 (Spearman 0.952–0.996), all 35 large-gap comparisons across six benchmarks in seven configurations are negative with 27 reaching significance, and 8 of the near-baseline group's 35 comparisons reach significance (six negative, two positive).

Perspective

The results are aimed at system designers who replace the softmax normalization operator during pretraining, and apply to GPT-2-style 124M and 1B parameter models, the FineWebEdu-3B and WikiText-103 corpora, and training budgets up to 2.5B tokens, with the attention operator running in fp32 and TF32 matmuls enabled during training. They let follow-up work compare calibration policies and surrogate placements under one forward, and treat the row-maximum gradient as a channel that must be retained; several LERP and FWM configurations serve as near-baseline candidates for further downstream and hardware evaluation.

Suites A–C use one seed each, and suite D's seed spread applies only to its own setting; the tail-policy comparison is single-seed and not cost-matched, the channel ablation covers two seeds, and the strict-forward retraining and reduced-precision baselines are one seed each at 100M tokens. The two surrogates share the same hard forward mathematically, but the historical implementations are not bitwise identical, with a largest mean validation NLL difference of about 0.001 nats on frozen endpoints and single-block differences up to 0.0125 nats, so the surrogate contrast carries a forward-rounding confound. The window width tau came from an inference-time screen on frozen weights whose final stage used the WikiText-103 validation split later used for suites C and D, and it is not shown optimal for training. The range–resolution feedback hypothesis is not identified by any intervention, and geometric diagnostics moved before the loss in the detach failure but after it in the formal MinMax–Weight runs, so they are not a universal early warning. The near-baseline downstream group shows small differences of both signs, which with one seed per condition cannot be attributed to the operator rather than to the run.

Sources