Why Post-Norm Transformers Collapse: Attention Amplification and Gradient Repair Failure
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:
-
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.
-
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.
-
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.
-
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.
-
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.
-
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
- Layer Normalization
- Post-LayerNorm Is Back: Stable, ExpressivE, and Deep
- Dynamical Isometry and a Mean Field Theory of RNNs: Gating Enables Signal Propagation in Recurrent Neural Networks
- An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale
- Exact Attention Sensitivity and the Geometry of Transformer Stability
- Representation Degeneration Problem in Training Natural Language Generation Models
- RealFormer: Transformer Likes Residual Attention
- Training Compute-Optimal Large Language Models
- Scaling Laws for Neural Language Models
- Transformers Get Stable: An End-to-End Signal Propagation Theory for Language Models
- Deep Neural Networks as Gaussian Processes
- Decoupled Weight Decay Regularization
- Gaussian Process Behaviour in Wide Deep Neural Networks
- Transformers without Tears: Improving the Normalization of Self-Attention
- Resurrecting the sigmoid in deep learning through dynamical isometry: theory and practice
- Mind the Gap: a Spectral Analysis of Rank Collapse and Signal Propagation in Attention Layers
- Deep Information Propagation
- NormFormer: Improved Transformer Pretraining with Extra Normalization
- Llama 2: Open Foundation and Fine-Tuned Chat Models
- Dynamical Isometry and a Mean Field Theory of CNNs: How to Train 10,000-Layer Vanilla Convolutional Neural Networks
Related papers
- Polynomial-Augmented Neural Networks (PANNs) with Weak Orthogonality Constraints for Enhanced Function and PDE Approximation
- AIRL-S: Unifying Reinforcement Learning and Search-Based Test-Time Scaling via Adversarial Inverse Reinforcement Learning
- Transformers as Bayesian In-Context Experimenters: Smoothness-Adaptive Efficient ATE Estimation
- Convergence issues in Relational Concept Analysis based on AOC-posets
- Beliefs Beyond Posteriors: Local-Consistency Optimisation for Bayesian Neural Networks
- Understanding Diffusion Models via Ratio-Based Function Approximation with SignReLU Networks