Why Post-Norm Transformers Collapse: Attention Amplification and Gradient Repair Failure

arXiv:2608.09417 · cs.LG · Submitted 2026-08-11 · Read on arXiv

Xingjian Wang, Qingyu Han, Xiaodong Luo, Yin Zhang

The Chinese University of Hong Kong, Shenzhen · Shenzhen Research Institute of Big Data

cs.LG

Submitted: 2026-08-11

Updated: 2026-08-12

License: http://creativecommons.org/licenses/by/4.0/

Importance score: 50/100

The gist: The paper "Why Post-Norm Transformers Collapse: Attention Amplification and Gradient Repair Failure" by Xingjian Wang, Qingyu Han, Xiaodong Luo, and Yin Zhang provides a two-stage analysis of rank

Terminology

Summary

The paper Why Post-Norm Transformers Collapse: Attention Amplification and Gradient Repair Failure by Xingjian Wang, Qingyu Han, Xiaodong Luo, and Yin Zhang provides a two-stage analysis of rank collapse in Post-Norm decoder-only Transformers, using token similarity as a scalar state variable.

Central Message: The paper states that "Post-Norm rank collapse is not caused by a single instability, but by the combination of two mechanisms: attention amplification, which drives token similarity upward at initialization, and RMSNorm-induced gradient shrinkage, which reduces the gradient that reaches earlier layers during collapse."

Stage I: Forward Similarity Amplification (Section 3.1): At initialization, causal attention acts approximately as a prefix-averaging operator that increases token similarity across depth. The paper shows that a single attention sublayer already increases it in expectation, and derives a closed-form one-step change: Theorem 3.2 gives ∆attn(s, t1) = sf1(t1)/(1 + sf2(t1)), where the amount of attention s = nd2σ2W/∥X1∥2F. The paper notes that "for t1 ∈ [1/n, 1), every factor in f1(t1), f2(t1) is strictly positive, so ∆attn(s, t1) > 0: the one-step attention contribution is positive. The SwiGLU FFN branch provides only a small damping effect, with Theorem 3.4 showing ∆FFN(ξ, t2) = ξ(t2 − 1/n)(m(ρ(t2)) − m(1))/(1 + ξm(1)), and the remark states ξm(1) is about 0.011, so the damping effect is small."

Stage II: Backward Repair Incapacity (Section 3.2): Once training enters a high-similarity regime, growth of pre-normalization residual norms makes the RMSNorm backward factor contractive. Theorem 3.6 shows that when tsim(Ylk) = tsim(Xlk) = 1, the gradient contraction factor is c(y, α) = α√d/∥y∥2, and "whenever ∥ykl∥22/d > (αkl)2, the corresponding factor c(ykl, αkl) falls below one, and the backward gradient contracts from Xk+1 to Xk. The paper explains the Pre-Norm difference: In Post-Norm, the RMS Jacobian multiplies both the sublayer and the skip-connection gradient. As a result, when the residual norm grows, Post-Norm shrinks the whole backward signal, whereas Pre-Norm shrinks only part of it."

Collapsed Network Characterization (Section 3.3): The paper characterizes properties of a collapsed network. Theorem 3.8 shows: (i) Lower Bound. If ∥XH − Π1XH∥F ≤ ϵ, then LCE(Fθ(t), y) ≥ Lfreq(y) − (2/√n)∥Wlm∥2ϵ. This means the best achievable probability distribution is the frequency distribution, yielding a relatively high loss floor. Part (ii) shows if the model output equals to pfreq(y) at every position, then for all sublayer such that tsim(Xlk) = 1, the gradients with respect to the parameter matrices in these layers are zero.

Experimental Verification (Section 4): Experiments on 48-layer decoder-only Transformers (d=512, dff=1536, 4 heads, 180M parameters) trained on C4 with sequence length 2048 match predictions. The paper reports: attention produces a positive one-step increment with strong dependence on layer index, and removing the prefix-averaging component suppresses the similarity growth at initialization. For the backward stage, measurements show the curves for higher-similarity layers bend downward sharply in the transition window, showing a rapid drop in √d/∥ykl∥2, and gradient norms in early layers decrease by orders of magnitude, while gradient norms in later layers change much less. Finally, collapsed training runs stay near the predicted frequency loss.

Conclusion: The paper concludes that "At initialization, causal attention increases token similarity, while the feed-forward correction is much smaller. During training, growth of the pre-RMS residual norm reduces the gradient that reaches layers with smaller indices, making the high-similarity state difficult to repair. Once collapse occurs, the training loss stays near the frequency loss."

Improvements for AI systems

Improvements to AI Systems:

  1. Adaptive Normalization Placement: Implement a dynamic normalization strategy that switches between Post-Norm and Pre-Norm based on measured token similarity and residual norm growth during training. The improved system can monitor the contraction factor c(y, α) = α√d/∥y∥2 in real time and automatically adjust normalization placement or scaling to prevent gradient shrinkage in early layers, avoiding rank collapse without sacrificing Post-Norm's training stability benefits.

  2. Gradient Repair Mechanism: Add a targeted gradient amplification module that detects when the backward contraction factor falls below 1 (indicating gradient repair failure) and compensates by scaling gradients for earlier layers proportionally to the residual norm growth. The improved system can maintain effective gradient flow to early layers even when residual norms grow large, enabling successful training of very deep Post-Norm transformers (e.g., 100+ layers) that would otherwise collapse.

  3. Similarity-Aware Initialization: Modify weight initialization schemes to account for the forward attention amplification effect. The system can pre-compute the expected one-step similarity increment ∆attn(s, t1) and adjust initial attention scaling (s = nd2σ2W/∥X1∥2F) to keep token similarity below a critical threshold (e.g., t1 < 0.5) during early training, preventing the cascade into the high-similarity regime where repair becomes impossible.

  4. Collapse Early-Warning System: Integrate a real-time monitor that tracks token similarity across layers and the ratio √d/∥y∥2 for each sublayer. The improved system can issue early warnings when similarity approaches 1 or when the contraction factor drops below 1, allowing for proactive interventions (e.g., gradient clipping, residual scaling, or layer-specific learning rate adjustments) before irreversible collapse occurs.

  5. Frequency-Loss-Aware Regularization: Add a regularization term that penalizes the model's output distribution from converging toward the frequency distribution pfreq(y), as characterized by the lower bound LCE ≥ Lfreq(y) − (2/√n)∥Wlm∥2ϵ. The improved system can maintain output diversity and avoid the high loss floor by explicitly discouraging token-wise uniformity, even when internal representations become highly similar.

  6. Residual Norm Budgeting: Implement a training scheduler that constrains the growth of pre-normalization residual norms ∥y∥2 relative to √d, keeping the backward contraction factor c(y, α) above a safety threshold (e.g., > 0.9). The improved system can allocate norm budgets across layers to ensure gradient signals reach all layers uniformly, preventing the sharp gradient decay observed in early layers during collapse.

Abstract

Deep decoder-only Transformers often replace the original Post-Norm architecture with Pre-Norm variants because Post-Norm training is highly sensitive to warmup and learning rate under conventional initialization schemes. Although prior work has identified rank collapse and gradient vanishing as related symptoms, it remains poorly understood how causal attention creates high-similarity representations and why training dynamics fail to repair them. We give a two-stage analysis of Post-Norm rank collapse using token similarity as a scalar state variable. First, at initialization, causal attention acts approximately as a prefix-averaging operator that increases token similarity across depth, while the SwiGLU branch contributes only a smaller damping effect. Second, once training enters a high-similarity regime, growth of pre-normalization residual norms makes the RMSNorm backward factor contractive; under mild conditions, gradients to earlier layers decay geometrically. As a complementary result, we characterize the properties of a collapsed network: its best predictor is frequency distribution with relatively high loss floor, and gradients in collapsed layers vanish at frequency distribution. Experiments on 48-layer decoder-only Transformers trained on C4 dataset match the predicted initialization-time similarity growth and collapse-time gradient contraction, and show that collapsed runs stay near the predicted frequency loss. Together, these results distinguish the forward similarity amplification and backward repair incapacity in Post-Norm collapse, while also characterizing the behavior of collapsed networks.

Sources

Related papers