Papers
Topics
Authors
Recent
Search
2000 character limit reached

Rank-Aware Spectral Bounds on Attention Logits for Stable Low-Precision Training

Published 21 Feb 2026 in cs.LG and cs.AI | (2602.18851v1)

Abstract: Attention scores in transformers are bilinear forms Sij=xi<sup>⊤</sup>Mxj/dhS_{ij} = x_i<sup>\top</sup> M x_j / \sqrt{d_h} whose maximum magnitude governs overflow risk in low-precision training. We derive a \emph{rank-aware concentration inequality}: when the interaction matrix M=W<sup>Q</sup>W<sup>K⊤M = W<sup>Q</sup> W<sup>{K\top} has rank r≪dr \ll d, tail probabilities for max⁡i,j∣Sij∣\max_{i,j}|S_{ij}| decay as exp⁡(−d<sup>2α<sup>2/(γr))\exp(-d<sup>{2}α<sup>{2}/(γr)) rather than exp⁡(−dα<sup>2)\exp(-dα<sup>{2}), where $γ&gt; 1$ is a typicality parameter. For transformer attention where r=dhr = d_h, this yields $8$--28×28\times tighter concentration than rank-agnostic bounds in modern architectures. We apply this result to FP8 training, deriving \emph{geometry-aware scale factors} that provide principled overflow guarantees without observing activations. The method computes per-layer scales from the spectral norm ∣W<sup>Q</sup>W<sup>K⊤∣2|W<sup>Q</sup> W<sup>{K\top}|_2 via implicit power iteration, includes a grouped query attention formulation that avoids key expansion, and remains compatible with fused attention kernels. Across GPT-2 XL to Llama-2-70B, geometry-aware scaling eliminates overflows in transient scenarios where delayed scaling fails, while achieving comparable downstream MMLU accuracy.

Authors (1)

Summary

  • The paper derives a rank-aware concentration inequality for attention logits that improves tail exponents by 8–28× across GPT-2 XL, Mistral-7B, and Llama models.
  • The method uses implicit power iteration and expansion-free grouped-query attention to estimate per-layer spectral norms with 1.0–4.3% overhead while remaining compatible with fused kernels.
  • Geometry-aware scaling produced zero overflows during loading, checkpoint resumption, and 100× learning-rate spikes, while auto-calibration restored downstream MMLU accuracy to 33.6% versus 28.1% for conservative scaling.

Motivation and problem statement

FP8 training with the E4M3 format offers a representable range of only ±448, compared to ±65,504 for FP16, so attention logits that exceed the scaled range produce quantization overflows and NaN values that corrupt training. The paper frames the calibration of per-tensor scale factors as a two-objective problem: safety, meaning a bounded overflow probability Pr⁡(max⁡i,j∣Sij∣>s⋅Rmax⁡)≤δ∗\Pr(\max_{i,j}|S_{ij}| > s \cdot R_{\max}) \leq \delta^*, and utilization, meaning maximal use of the FP8 dynamic range. These objectives conflict, and existing methods resolve them reactively. Delayed scaling (Micikevicius et al., 2022) derives scales from a 16-step history buffer of past activation maxima; current scaling observes max⁡∣St∣\max|S_t| each iteration but requires materializing the full L×LL \times L score matrix, which is incompatible with fused kernels such as FlashAttention (Sigrist, 2022, Dao, 2023).

The paper identifies a specific failure mode of delayed scaling, termed history staleness: whenever weight dynamics outpace the history buffer — checkpoint loading, resumption from saved states (which typically omit scaling state), or learning-rate transitions — the scale factor is computed from statistics irrelevant to current weights. Prior to this work, no method was simultaneously transient-safe and fused-kernel-compatible.

Spectral interaction bound

The core observation is that attention scores are bilinear forms Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h} with M=WQWK⊤M = W^Q W^{K\top}. Under pre-LN architectures, LayerNorm or RMSNorm constrains token norms to ∥xi∥2≈d\|x_i\|_2 \approx \sqrt{d}, yielding the worst-case bound

Bmax⁡=∥WQWK⊤∥2⋅ddh.B_{\max} = \|W^Q W^{K\top}\|_2 \cdot \frac{d}{\sqrt{d_h}}.

This bound depends only on current weights, not activations, enabling predictive calibration. A corollary establishes that the interaction bound is never looser than the naive submultiplicative bound ∥WQ∥2∥WK∥2\|W^Q\|_2 \|W^K\|_2, and is strictly tighter unless the top right singular vectors of WQW^Q and WKW^K coincide — a measure-zero event under random initialization. This matters because bounds derived from individual weight norms, as in MOSS (2511.05811), can substantially overestimate logit magnitudes and waste dynamic range.

The worst-case bound assumes perfect alignment of inputs with the top singular direction, which is exponentially unlikely in high dimensions. The paper therefore introduces a calibration factor max⁡∣St∣\max|S_t|0 giving max⁡∣St∣\max|S_t|1, and develops probabilistic machinery to select max⁡∣St∣\max|S_t|2.

Rank-aware concentration inequality

The central theoretical contribution is a two-stage conditioning argument exploiting max⁡∣St∣\max|S_t|3. First, the probability that any key projection onto the row space of max⁡∣St∣\max|S_t|4 is atypical is bounded via a Beta-distribution Chernoff argument, giving max⁡∣St∣\max|S_t|5. Conditioned on typical keys, Lévy's lemma applied with the tightened Lipschitz constant max⁡∣St∣\max|S_t|6, where max⁡∣St∣\max|S_t|7, yields max⁡∣St∣\max|S_t|8. The tail exponent improves over the rank-agnostic bound by a factor max⁡∣St∣\max|S_t|9: 8× for GPT-2 XL, 14× for Mistral-7B, 18× for Llama-2-13B, and 28× for Llama-2-70B. Union bounds over all L×LL \times L0 pairs and all L×LL \times L1 heads then yield closed-form selection rules for L×LL \times L2 and L×LL \times L3 given a target failure probability L×LL \times L4; at L×LL \times L5 and L×LL \times L6, L×LL \times L7 ranges from 0.074 (GPT-2 XL) down to 0.018 (Llama-2-70B), reflecting stronger concentration in larger models.

Two assumptions bear directly on this result. The analysis idealizes post-normalization tokens as i.i.d. uniform on the sphere, treating diagonal (L×LL \times L8) terms as independent bilinear forms — conservative, since quadratic forms concentrate more tightly. For RoPE architectures, rotations are orthogonal so the worst-case bound holds rigorously, but the tighter interaction bound holds only under an empirically verified condition that RoPE rotations do not systematically align with the singular subspaces of L×LL \times L9 and Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}0; this was checked across all layers of Mistral-7B, Llama-2-13B, and Llama-2-70B but is not proven in general.

Efficient estimation and GQA formulation

Computing Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}1 by full SVD costs Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}2 and would require forming a 67-million-entry matrix per layer for Llama-2-70B. Instead, power iteration computes matrix-vector products implicitly via Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}3 at cost Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}4 per iteration, never forming Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}5. Persistent vectors are maintained across steps with one update per forward pass during steady-state training, and five iterations at cold start. The paper argues that warm-started power iteration underestimates the spectral norm by at most the growth factor Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}6 between steps, and that the margin above Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}7 absorbs moderate underestimation; a controlled test with a Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}8 weight spike confirms same-forward-pass adaptation.

For grouped query attention, an implicit formulation avoids expanding Sij=xi⊤Mxj/dhS_{ij} = x_i^\top M x_j / \sqrt{d_h}9 (32 MB per layer on Mistral-7B): forward products use block replication of small intermediate vectors, backward products use group summation, and a proposition proves convergence to the identical spectral norm. Because spectral norms vary by up to 19.5× across layers within a model, scales are computed per layer as M=WQWK⊤M = W^Q W^{K\top}0 with M=WQWK⊤M = W^Q W^{K\top}1.

An optional auto-M=WQWK⊤M = W^Q W^{K\top}2 mode targets steady-state fine-tuning: after a burn-in phase collecting slack ratios M=WQWK⊤M = W^Q W^{K\top}3, M=WQWK⊤M = W^Q W^{K\top}4 is set to the 99.99th percentile times a safety multiplier and frozen. The paper is explicit that auto-M=WQWK⊤M = W^Q W^{K\top}5 forfeits the Proposition guarantee, since the factor is empirical rather than derived from the selection rule, and relies on the burn-in distribution being representative; it also requires materializing the score matrix during burn-in, though this overhead is under 0.1% of total compute.

Empirical results

Across GPT-2 XL (1.5B), Mistral-7B, Llama-2-13B, and Llama-2-70B, geometry-aware scaling achieves zero overflows in all three transient scenarios where delayed scaling fails:

Scenario Delayed scaling Geometry-aware
Pretrained loading 100% of layers overflow; max scaled logits up to 9498 0 overflows; max ≤ 196
Checkpoint resumption 1–4 overflowing steps per model 0 overflows
100× LR spike 3–5 overflowing steps; NaNs terminate GPT-2 XL training 0 overflows

Forward-pass overhead is +1.0% (GPT-2 XL), +1.9% (Llama-2-13B), and +4.3% (Llama-2-70B); on Mistral-7B the implicit GQA formulation yields −5.3%, i.e., faster than the baseline, though the paper concedes this negative result is implementation-specific.

Downstream evaluation fine-tunes Llama-2-13B for 3000 steps on MMLU STEM. All three methods converge to similar training loss (~0.012), yet conservative spectral scaling degrades accuracy to 28.1% versus 33.0% for delayed scaling, because 0.5% FP8 utilization introduces excessive quantization noise — evidence that training loss alone does not predict downstream quality. Auto-M=WQWK⊤M = W^Q W^{K\top}6 tightens M=WQWK⊤M = W^Q W^{K\top}7 by 83× (from 0.03 to 0.00036), raising utilization to 31.2% and achieving 33.6% accuracy with zero overflows. The paper explicitly does not claim statistical significance for the 0.6-point improvement over delayed scaling, noting the small training set (295 examples); the delayed baseline's 68 overflows were handled by clamping, without which NaN propagation would have terminated training.

Limitations and open questions

Several caveats are stated plainly in the paper. The spherical-token assumption is an idealization of post-normalization directions, and while argued to be conservative, the guarantees rest on it. The tighter interaction bound under RoPE is validated empirically rather than proven. The transient-robustness argument for power iteration bounding underestimation by the growth factor M=WQWK⊤M = W^Q W^{K\top}8 is informal, supported by stress tests rather than a formal tracking bound. Auto-M=WQWK⊤M = W^Q W^{K\top}9 provides no worst-case guarantee and presumes distributional stability after calibration. Evaluation covers fine-tuning-scale experiments (300–3000 steps); whether the calibration framework sustains zero overflows over full pretraining runs, and how ∥xi∥2≈d\|x_i\|_2 \approx \sqrt{d}0 behaves under much longer sequences where the union bound over ∥xi∥2≈d\|x_i\|_2 \approx \sqrt{d}1 pairs grows, remain open questions the paper does not address.

Conclusion

The paper derives a rank-aware concentration inequality showing that attention-logit tails decay as ∥xi∥2≈d\|x_i\|_2 \approx \sqrt{d}2 rather than ∥xi∥2≈d\|x_i\|_2 \approx \sqrt{d}3 when the query-key interaction matrix has low rank, yielding 8–28× tighter bounds on modern architectures. Combined with implicit power iteration and an expansion-free GQA formulation, this produces the first FP8 attention calibration method that is simultaneously transient-safe and compatible with fused attention kernels, demonstrated with zero overflows from GPT-2 XL through Llama-2-70B and preserved downstream MMLU accuracy via auto-∥xi∥2≈d\|x_i\|_2 \approx \sqrt{d}4 calibration.

Paper to Video (Beta)

No one has generated a video about this paper yet.

Whiteboard

No one has generated a whiteboard explanation for this paper yet.

Open Problems

We haven't generated a list of open problems mentioned in this paper yet.