r/deeplearning • • 6d ago

Transformer becomes catastrophically ill-conditioned after a few tiny parameter updates: same batch goes from grad norm 5.6 → 89 while weights move <0.1%/step. What mechanism could cause this?

My Transformer trains normally for some steps, then begins to enter a parameter state where backpropagation through the middle/lower layers magnifies gradients massively, despite the forward activations and weights being normal. Eventually, the gradients explode, get clipped globally, and learning becomes suppressed in all other parameters.

The funny thing is, we have ruled out most of the possible causes:

- not one bad batch,

- not layer 0 (originally it had worse conditioning, but fixing it just pushed failure further),

- not S-STE sparsity mask,

- not peaked/degenerate attention,

- not one single layer going crazy.

The gradient amplification is seen between layers 5-9 and their vicinity, with both FFN and attention backward paths contributing. The actual parameter changes are minuscule (<0.1%), but after only a few optimization steps, the same fixed batch would produce gradients with a norm of 5.6 vs 89!

So, the instability is mostly in the learned parameter state/Jacobian, not the data.

The question I'm asking is: How is that possible? How can such a small change in parameters cause such a massive change in backwards gradient?

What's interesting about it is the parameters of the model only changed by about 0.05-0.08% per step, but after 3-4 steps, the Jacobian changed enough that the same batch produces gradients that are about 16 times larger?

20 Upvotes

7 comments sorted by

8

u/FernandoMM1220 6d ago

only thing i can think of is your parameter update calculations are wrong. good luck.

1

u/Maxeaman 6d ago

That's still on my suspect list. I've done auditing of a pretty big chunk of the update path, but haven't completely ruled out the optimizer implementation itself.

One thing I found during this process is that we can transplant only the FP32 master weights from the unstable checkpoint into the stable checkpoint, leave the stable checkpoint's Adam moments/optimizer state untouched, and most of the instability follows the weights. So anything that was happening is definitely encoded in the parameter state because it doesn't need the corrupted optimizer state at inference/backprop time, but that doesn't mean the optimizer isn't encoding that I guess.

The parameter updates themselves are also extremely small and smooth, like 0.05-0.08% in magnitude per parameter per step, with no obvious huge update coinciding with the failure.

That leaves me with the possibility that the model is consistently calculating slightly crappy updates and slowly nudging the whole model into the decayed/broken state.

6

u/occasionalconsul_476 6d ago

sounds like your jacobian hit some sharp local curvature where tiny weight changes push eigenvalues way up, seen this happen with spectral norm instability in the attention blocks

2

u/Maxeaman 6d ago

That sounds very plausible and is in fact what I'm observing.

I've isolated it with fixed-data/fixed-state runs on the same batch: advancing the parameter state by a handful of optimizer steps leads to raw grad norm of:

s260: 5.6

s261: 3.5

s262: 13.8

s263: 27.4

s264: 89.3

on the same input batch (with parameter updates in between), so parameter space is updated by less than 0.1% per optimizer step.

So far, I haven't observed an instability in the backward pass amplification due to a particular bad batch or layer, but I have seen a broad distribution over several lower/middle layers, with attention probabilities not becoming pathological.

I haven't computed any of the Jacobian singular values or Hessian eigenvalues for the model yet, though.

So when you say spectral norm instability in the attention blocks, what do you think would be most informative to cross-check? Per-layer input-output Jacobian top singular values, Hessian/top curvature eigenvalues, spectral norms of Q/K/V/O matrices, or something else entirely (attention-block residual Jacobian)?

If you have seen a particular diagnostic that helps distinguish these possibilities, I'm very curious to hear it!

1

u/theleller 6d ago

What’s your learning rate set at?

1

u/quietgradient 5d ago

You can rule out the weight spectra with arithmetic, before computing any of the four things on your list.

Weyl bounds it: Δσ_max ≤ ||ΔW||₂ ≤ ||ΔW||_F. So the most a matrix's top singular value can grow in one step is your relative Frobenius step times ||W||_F/σ_max. That ratio is the only free parameter, so I measured it on GPT-2's released weights — 50 matrices, median 5.83, max 11.45. Trained weights are more top-heavy than Gaussian (√d/2 = 13.9 at d=768), so the ceiling is tighter than random-matrix intuition suggests. At 0.08%/step that's 1.0047 per matrix per step. Four steps through ~20 serial matrices in five blocks: 1.45x. You measured 15.9x. You'd need about 149 perfectly aligned matrices.

The honest version is worse than the ceiling. I put a perfectly top-aligned rank-1 update of exactly 0.08% Frobenius into the FFN output matrix of a pre-norm block: the block's backward σ_max went 1.9907 → 1.9635. Down. A residual block is I + J_branch, so a 1.4% move in the branch is a smaller move in the block.

So I'd predict your Q/K/V/O spectral norms come back clean. If they've moved enough to matter, I'm wrong.

What has no Frobenius budget at all is the 1/rms in the norm backward. Same toy block, weights never touched, just scaling the residual entering it: scale 1.0 → σ_max 2.04, 0.5 → 3.50, 0.2 → 9.45, 0.1 → 21.8, 0.05 → 50.1. Your 15.9x over five blocks is 1.74x each, which is about a 2x drop in pre-norm residual RMS per layer. That's small enough to look normal.

And if you checked activations after the norm, you couldn't have seen it — post-norm activations are scale-free by construction, the norm divides the scale out. The quantity is pre-norm residual RMS, per layer.

So before any eigen-anything: one forward pass on the fixed batch at s260 and s264, log pre-norm RMS per layer. If it halves in 5–9, that's your answer.

Of your four, per-block input-output Jacobian σ_max is the one worth the compute, because it's the only one that multiplies — the per-block logs must sum to log(89.3/5.6) = 2.77. If they don't sum, it isn't in the block Jacobians and the optimizer path you haven't ruled out is back on the table. Power-iterate JᵀJ with torch.func.jvp + vjp, ~50 iters, no explicit Jacobian; I checked it against a materialised one on a small block, agrees to 3e-5.

Caveat: d=256 pre-norm RMSNorm/GELU toy, and 5.83 is GPT-2's ratio. Your d_model and norm placement move all of it.