Training Stability: Loss Spikes, Gradient Norms & Debugging

Michael BrenndoerferJanuary 28, 202660 min read

Part of Language AI Handbook

Detect and prevent training instability in deep learning. Topics include loss spikes, gradient norm monitoring, gradient clipping.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

Training Stability

Training a large language model is not a smooth, monotonic descent toward a good solution. At any moment, loss can spike dramatically, gradients can explode, or training can stall entirely without warning. These instabilities waste compute, corrupt checkpoints, and, in the worst case, require restarting runs from scratch. Understanding why they happen and how to prevent them is one of the most practically important skills in modern deep learning.

This chapter focuses on training stability: the set of phenomena that cause training to behave erratically, the diagnostic signals that reveal problems early, and the techniques that keep training on track. We cover loss spikes and their root causes, gradient norm monitoring as an early warning system, and a toolkit of stabilization techniques from gradient clipping to architectural choices. We also develop a debugging workflow for diagnosing instabilities when they occur.

By the end, you will understand what each technique does and why instability happens in the first place, which makes the solutions feel inevitable rather than arbitrary.

Why Training Becomes Unstable

Before reaching for solutions, it is worth understanding the underlying mechanics. Training instability is not a single phenomenon but a family of related problems, each with its own signature and cause.

The Loss Landscape Is Not Smooth

Neural network loss landscapes are high-dimensional and highly non-convex. Most of the time, gradient descent moves through these landscapes smoothly, making incremental progress. But occasionally, the optimizer wanders into regions where curvature is extremely high: a small step in parameter space produces a large change in loss. These high-curvature regions cause loss spikes.

The curvature of the loss landscape is described mathematically by the Hessian matrix H\mathbf{H}, where each entry Hij=∂2L∂θi∂θjH_{ij} = \frac{\partial^2 \mathcal{L}}{\partial \theta_i \partial \theta_j} captures how the gradient of the loss with respect to parameter ii changes as parameter jj moves. The largest eigenvalue of H\mathbf{H}, often called the sharpness λmax⁡\lambda_{\max}, determines how large a learning rate is safe. Specifically, gradient descent is guaranteed to decrease loss only when:

η<2λmax⁡\eta < \frac{2}{\lambda_{\max}}

where η\eta is the learning rate and λmax⁡\lambda_{\max} is the largest eigenvalue of the Hessian. When the optimizer encounters a region where λmax⁡\lambda_{\max} is much larger than expected, the effective learning rate becomes too large for that region, and a single update can push parameters past the minimum and far up the other side of the curvature.

This explains why loss spikes often happen suddenly after many stable steps. The optimizer can spend thousands of iterations in well-behaved regions, then stumble into a sharp curvature cliff that sends loss skyrocketing. The transition is not gradual because the loss landscape itself has abrupt changes in curvature that have no visible warning in the loss curve.

One insight from recent research on sharpness-aware minimization (SAM) is that the loss landscapes of large language models tend to contain sharp minima, meaning regions where the loss is low but the curvature is high. These sharp minima are not just bad for generalization; they are also fragile during training because the same property that makes them sharp makes them prone to instability. Flat minima, where a wide basin surrounds the optimal point, are more stable because the region of safe learning rates is much larger.

Gradient Accumulation Through Depth

In deep networks, gradients flow backward through many layers. As covered in earlier chapters on backpropagation and vanishing gradients, this can lead to gradients either shrinking exponentially (vanishing) or growing exponentially (exploding). Exploding gradients are the more acute stability threat because they cause immediate, dramatic parameter updates.

The fundamental issue is that in a network with LL layers, each backward pass involves LL matrix multiplications. If the singular values of the weight matrices are consistently above 1, gradients compound multiplicatively and can reach astronomically large values before a single parameter update is applied. For a simple linear network with weight matrices W1,W2,…,WL\mathbf{W}_1, \mathbf{W}_2, \ldots, \mathbf{W}_L, the gradient of the loss with respect to the first layer involves a product of the form:

∂L∂W1∝WLTWL−1T⋯W2T∂L∂hL\frac{\partial \mathcal{L}}{\partial \mathbf{W}_1} \propto \mathbf{W}_L^T \mathbf{W}_{L-1}^T \cdots \mathbf{W}_2^T \frac{\partial \mathcal{L}}{\partial \mathbf{h}_L}

where hL\mathbf{h}_L is the output of the final layer. If the spectral norm of each weight matrix is σ>1\sigma > 1, then this product can grow as σL−1\sigma^{L-1}, which becomes enormous for deep networks. This exponential amplification is precisely the exploding gradient problem.

Recurrent neural networks are especially vulnerable because the same weight matrix is applied repeatedly for each time step. As covered in the BPTT chapter, gradients scale as (WT)T(\mathbf{W}^T)^T where TT is the sequence length, and if ∥W∥>1\|\mathbf{W}\| > 1, the gradient norm grows exponentially with sequence length. This is one historical reason why RNNs were notoriously hard to train on long sequences before gradient clipping became standard practice.

Modern transformers are less susceptible to the pure exponential blowup because residual connections provide gradient highways that bypass layers, and layer normalization prevents activation magnitudes from drifting. But they are not immune. Very deep transformers (48+ layers), long sequence lengths, and the presence of attention softmax can still produce gradient explosions under the right conditions.

Batch Statistics and Outlier Samples

A less discussed but important source of instability is data heterogeneity. When a batch contains outlier samples with very high loss (rare but extreme examples), the gradient from that batch carries an unusually large signal. If the learning rate is tuned for typical batches, this outsized gradient can cause a destabilizing update.

This is particularly relevant in language modeling, where certain sequences (very long documents, repeated tokens, unusual characters) can produce loss values far outside the typical distribution. Consider a batch that happens to sample a document consisting of thousands of repetitions of a single rare token. The model will assign very low probability to each token given its context (because the pattern is unusual), resulting in a very high cross-entropy loss. The gradient from this batch is dominated by this single anomalous document and is not representative of the general distribution the model should learn.

These outlier batches are often the direct trigger for loss spikes that practitioners observe when training on large, diverse web-scraped corpora. The spike is not a fundamental property of the model or optimizer; it is an artifact of sampling a particularly difficult batch. Data quality filtering therefore improves training stability rather than merely making the training process easier to manage.

Optimizer State Corruption

Modern optimizers like Adam maintain momentum estimates for each parameter. These estimates are exponential moving averages of past gradients, and they allow the optimizer to move smoothly even when individual gradient estimates are noisy. But they also carry history: a period of large gradients can corrupt momentum estimates in ways that take many steps to decay.

Adam's first and second moment estimates are defined as:

mt=β1mt−1+(1−β1)gtvt=β2vt−1+(1−β2)gt2\begin{aligned} m_t &= \beta_1 m_{t-1} + (1 - \beta_1) g_t \\ v_t &= \beta_2 v_{t-1} + (1 - \beta_2) g_t^2 \end{aligned}

where gtg_t is the gradient at step tt, β1≈0.9\beta_1 \approx 0.9 is the first moment decay, and β2≈0.999\beta_2 \approx 0.999 is the second moment decay. The second moment decay β2=0.999\beta_2 = 0.999 means that the influence of a gradient from 1000 steps ago is reduced by a factor of 0.9991000≈0.370.999^{1000} \approx 0.37. So a corrupt gradient signal persists in the optimizer state for roughly 1000 steps before it decays substantially.

When a loss spike occurs, the Adam optimizer's momentum buffers absorb large gradient values. Even after the spike resolves, the inflated momentum can continue causing larger-than-appropriate updates for many subsequent steps. This creates a cascading instability where a single bad batch can degrade training for hundreds of iterations. Practitioners who simply continue training after a spike often see the loss partially recover and then remain high for a long time, which is exactly the signature of corrupted optimizer state slowly decaying.

Numerical Precision Issues

A fourth, often underappreciated cause of instability is numerical precision. Modern large model training uses mixed-precision arithmetic, combining float32 for some computations with float16 or bfloat16 for others. Float16 has a maximum representable value of 65504, and values larger than this cause overflow to infinity. The underflow boundary is even more constraining for small values, with subnormal numbers starting below approximately 6×10−56 \times 10^{-5}.

In practice, this means that attention logits for very long sequences (where the dot-product values can be large), or softmax outputs for peaked distributions, or intermediate activations in very deep networks, can easily overflow float16 representation. The result is NaN (not a number) values that propagate forward through the network and produce NaN gradients, which corrupt optimizer state catastrophically.

The bfloat16 format, used by default in many modern TPU and GPU training setups, has the same exponent range as float32 (so no overflow risk for typical values) but only 7 bits of mantissa instead of float32's 23 bits. This makes bfloat16 much more resistant to overflow, but it can still produce instability through accumulated rounding errors in numerically sensitive operations like layer normalization.

Understanding these four sources of instability, sharp loss landscape curvature, exploding gradients through depth, outlier batches, and numerical precision limitations, gives you a principled framework for diagnosing problems and choosing appropriate interventions.

Gradient Norm Monitoring

The most important diagnostic tool for training stability is gradient norm monitoring. Before parameters are updated, measuring the total magnitude of the gradient vector provides an early warning of impending instability.

What Gradient Norm Tells You

The gradient norm is the Euclidean (L2) norm of the concatenated gradient vector across all parameters:

∥g∥=∑igi2\|\mathbf{g}\| = \sqrt{\sum_{i} g_i^2}

where:

  • g\mathbf{g}: the gradient vector containing all parameter gradients at the current step
  • gig_i: the gradient component for parameter ii, computed by backpropagation
  • ∥g∥\|\mathbf{g}\|: the overall magnitude of the gradient signal before any update is applied

Under normal training conditions, the gradient norm fluctuates within a predictable range, varying from batch to batch but staying within roughly an order of magnitude of its typical value. A sudden spike in gradient norm, say 10x or 100x the recent average, is a strong signal that either the optimizer has entered a high-curvature region or the current batch contains outlier samples.

The relationship between gradient norm and parameter update magnitude makes this diagnostic especially valuable. A standard SGD update applies a parameter change of η∥g∥\eta \|\mathbf{g}\| in the gradient direction. With Adam, the relationship is more complex due to adaptive scaling, but the gradient norm still captures the raw signal strength before adaptation. Monitoring gradient norm lets you catch problems that Adam's adaptation would otherwise obscure until they manifest as a loss spike.

Think of gradient norm as the "pressure" on the system at each step. Normal training has low, steady pressure. Loss spikes have extremely high pressure that can break the system if not controlled. You want to watch the pressure gauge, not wait for the pipe to burst.

Why Gradient Norm Precedes Loss Spikes

An important property of gradient norm as a diagnostic is that it often increases before the loss spike becomes visible in the loss curve. This makes it a leading indicator rather than a concurrent one.

The reason is causal: the large gradient norm at step tt causes a large parameter update, which sends parameters into a worse region of the loss landscape. The higher loss from that region only becomes visible at step t+1t+1 and subsequent steps. So by monitoring gradient norm in real time, you can sometimes intervene (by reducing learning rate or triggering a checkpoint rollback) before the loss spike fully materializes.

This predictive relationship is not perfect. Some loss spikes arise from data issues (an extremely high-loss batch) where the gradient norm at that step is large and the loss spike is simultaneous, not subsequent. But for architecture-driven instabilities and optimizer state problems, gradient norm often gives you a one-step warning.

Typical Gradient Norm Trajectories

In practice, gradient norms follow recognizable patterns across different training phases and health states:

  • Early training: Gradient norms are typically high and variable. The model is far from convergence, loss landscape curvature is less predictable, and the optimizer is making large exploratory steps. This is normal and expected.
  • Mid training: Gradient norms settle into a relatively stable range. Updates become more consistent as the optimizer finds a good region of parameter space.
  • Spike events: Sudden increases (often 5-100x the running average) indicating problematic updates. These warrant immediate investigation.
  • Gradual drift: A slow upward trend can indicate the optimizer is leaving a stable region. This is more insidious than sudden spikes because it can go unnoticed for many steps and suggests the learning rate may need to be reduced.
  • Near-zero norms: If gradient norms drop very close to zero and stay there, the model may be stuck in a region of flat loss landscape or the learning rate is too small to produce meaningful updates.
In[4]:
Code
import numpy as np
import torch

np.random.seed(42)
torch.manual_seed(42)


def simulate_grad_norms(n_steps=500, spike_steps=None, spike_magnitude=10.0):
    """Simulate gradient norm trajectory with optional spike events."""
    if spike_steps is None:
        spike_steps = []
    base_norms = np.abs(np.random.normal(1.5, 0.3, n_steps))
    # Add realistic autocorrelation
    for i in range(1, n_steps):
        base_norms[i] = 0.7 * base_norms[i - 1] + 0.3 * base_norms[i]
    # Inject spikes
    for step in spike_steps:
        base_norms[step] = base_norms[step] * spike_magnitude
        if step + 1 < n_steps:
            base_norms[step + 1] = base_norms[step + 1] * (
                spike_magnitude * 0.4
            )
    return base_norms


# Simulate healthy and unstable trajectories
healthy_norms = simulate_grad_norms(500)
unstable_norms = simulate_grad_norms(
    500, spike_steps=[150, 280, 420], spike_magnitude=12.0
)
Out[5]:
Visualization
Line plot showing gradient norm values hovering around 1.5 with small fluctuations over 500 steps.
Healthy gradient norm trajectory over 500 training steps. The norm fluctuates in a stable band around 1.5, with no extreme spikes, indicating the optimizer is navigating a well-behaved region of the loss landscape throughout training.
Line plot showing gradient norm with three tall spikes reaching above 15 against a stable baseline of 1.5.
Unstable gradient norm trajectory with three spike events at steps 150, 280, and 420. Each spike represents a sudden 12x increase in gradient magnitude, which would cause destabilizing parameter updates without intervention. The partial recovery visible after each spike reflects the decaying autocorrelation in the simulation.

The contrast between these trajectories illustrates what to watch for in practice. The healthy run shows normal variance, never straying far from its mean. The unstable run's spikes stand out immediately and would alert any practitioner reviewing monitoring dashboards. In real training runs, you log the gradient norm every N steps and watch the monitoring dashboard for these patterns.

Logging Gradient Norms in PyTorch

Implementing gradient norm monitoring requires a single line between loss.backward() and optimizer.step(). PyTorch's clip_grad_norm_ function both optionally clips gradients and returns the pre-clipping norm as a scalar tensor. Using it with max_norm=float("inf") gives you the norm computation as a free side effect without modifying gradients.

In[6]:
Code
import torch
import torch.nn as nn
import torch.optim as optim


class SimpleTransformerBlock(nn.Module):
    """Minimal transformer block for demonstrating gradient monitoring."""

    def __init__(self, d_model=128, n_heads=4):
        super().__init__()
        self.attention = nn.MultiheadAttention(
            d_model, n_heads, batch_first=True
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.ffn = nn.Sequential(
            nn.Linear(d_model, d_model * 4),
            nn.GELU(),
            nn.Linear(d_model * 4, d_model),
        )
        self.norm2 = nn.LayerNorm(d_model)

    def forward(self, x):
        attn_out, _ = self.attention(x, x, x)
        x = self.norm1(x + attn_out)
        ffn_out = self.ffn(x)
        return self.norm2(x + ffn_out)


torch.manual_seed(42)
model = SimpleTransformerBlock(d_model=128, n_heads=4)
optimizer = optim.AdamW(model.parameters(), lr=1e-3)
criterion = nn.MSELoss()
In[7]:
Code
# Training loop with gradient norm monitoring
grad_norm_log = []
loss_log = []

for step in range(100):
    # Generate synthetic batch
    x = torch.randn(16, 32, 128)  # batch=16, seq_len=32, d_model=128
    target = torch.randn(16, 32, 128)

    optimizer.zero_grad()
    output = model(x)
    loss = criterion(output, target)
    loss.backward()

    # Compute gradient norm BEFORE clipping
    total_norm = torch.nn.utils.clip_grad_norm_(
        model.parameters(), max_norm=float("inf")
    )
    grad_norm_log.append(total_norm.item())
    loss_log.append(loss.item())

    optimizer.step()
Out[8]:
Console
Steps completed: 100
Mean gradient norm: 0.1987
Max gradient norm: 0.2163
Min gradient norm: 0.1877
Final loss: 1.6492

The key insight here is that clip_grad_norm_ with max_norm=float("inf") computes and returns the gradient norm without modifying any gradients. This gives you a clean measurement at every step for free, as a side effect of computing what you need for clipping anyway. In a production training loop, you would send this value to your monitoring system (Weights and Biases, TensorBoard, or a custom logging database) and alert if it exceeds a threshold.

Gradient Norm Statistics for Threshold Setting

One practical question is: what gradient norm is too high? The answer is not absolute but relative. A gradient norm of 10.0 might be normal for one model and catastrophic for another, depending on the model's scale, architecture, and typical training dynamics.

The right approach is to establish a baseline by running the first 1,000-2,000 steps without any interventions (or with only a very loose clip threshold like 100.0), recording the gradient norm at every step, and computing the empirical distribution. From this distribution, you can set alert thresholds at, say, the 99th percentile for "high gradient norm" and 3 standard deviations above the mean for "spike detected." This adaptive approach is more reliable than picking a fixed threshold before seeing any training dynamics.

Gradient Clipping

Gradient clipping is the most widely deployed stability technique in large model training. The idea is straightforward: if the gradient norm exceeds a threshold, rescale all gradients so the norm equals the threshold.

How Gradient Clipping Works

Given a maximum allowed norm cc, gradient clipping transforms the gradient vector g\mathbf{g} as follows:

g^={gif ∥g∥≤cc⋅g∥g∥if ∥g∥>c\hat{\mathbf{g}} = \begin{cases} \mathbf{g} & \text{if } \|\mathbf{g}\| \leq c \\ c \cdot \dfrac{\mathbf{g}}{\|\mathbf{g}\|} & \text{if } \|\mathbf{g}\| > c \end{cases}

where:

  • g\mathbf{g}: the raw gradient vector computed by backpropagation
  • cc: the maximum allowed gradient norm (a hyperparameter, typically 1.0)
  • g^\hat{\mathbf{g}}: the clipped gradient used for the parameter update
  • g∥g∥\dfrac{\mathbf{g}}{\|\mathbf{g}\|}: the unit vector in the gradient direction, preserving directional information

The critical property of this transformation is that it preserves the direction of the gradient while constraining its magnitude. When the gradient is already within the threshold, nothing changes. When it exceeds the threshold, the optimizer still moves in the correct direction but takes a smaller step. No information about which direction to update is lost; only the step size is bounded.

This is the essential difference between gradient clipping and simply capping individual gradient components at some value. Clipping by individual component changes direction: a large component in one dimension gets reduced while others do not, which can send the optimizer in a substantially different direction. Clipping by global norm maintains the optimizer's directional intent.

Consider a gradient vector g=[4.0,3.0]\mathbf{g} = [4.0, 3.0] with norm ∥g∥=5.0\|\mathbf{g}\| = 5.0 and a clip threshold of c=1.0c = 1.0. Global norm clipping rescales to g^=[0.8,0.6]\hat{\mathbf{g}} = [0.8, 0.6], which points in exactly the same direction as the original gradient. Component clipping would produce g^=[1.0,1.0]\hat{\mathbf{g}} = [1.0, 1.0], which is at 45 degrees and points in a completely different direction. This is why global norm clipping is always preferred over component-wise clipping.

Historical Context

Gradient clipping for recurrent networks was formalized by Pascanu, Mikolov, and Bengio in a 2013 paper titled "On the difficulty of training recurrent neural networks." They showed analytically that the loss field of RNNs has cliff-like curvature in directions corresponding to exploding gradients, and that clipping is the appropriate response. The main insight from that paper was that the cliff structure is predictable: it occurs when gradients follow a direction that corresponds to the largest eigenvector of the Hessian, and clipping prevents the optimizer from taking a catastrophically large step in that direction.

The technique was subsequently adopted for transformer training and has become standard in virtually every large language model training codebase, from GPT-2 through GPT-4, LLaMA, and beyond.

Choosing the Clipping Threshold

The clipping threshold cc is a hyperparameter that requires empirical tuning. Common values are 1.0 for most tasks and 0.5 for particularly unstable setups. The choice involves a tradeoff: a threshold that is too small clips gradients even when they are legitimate and informative, slowing convergence. A threshold that is too large provides no protection against spikes.

A practical approach is to monitor the gradient norm for the first few thousand steps with no clipping (or with a very loose threshold of 100.0) and observe the typical range. Setting the threshold at roughly the 95th percentile of the observed distribution allows normal gradients to pass through unmodified while clipping the true outliers.

You can also tune the clipping threshold during training by observing what fraction of steps are clipped. If fewer than 5% of steps are clipped, the threshold is not restricting legitimate gradients. If more than 20-30% of steps are clipped, the threshold may be too aggressive and could be impeding learning. A well-calibrated threshold clips only the outlier tail, leaving normal gradient dynamics untouched.

In[9]:
Code
# Demonstrate gradient clipping effect
torch.manual_seed(42)
model_clipped = SimpleTransformerBlock(d_model=128, n_heads=4)
model_unclipped = SimpleTransformerBlock(d_model=128, n_heads=4)

# Copy weights so both start identical
model_unclipped.load_state_dict(model_clipped.state_dict())

opt_clipped = optim.AdamW(model_clipped.parameters(), lr=1e-3)
opt_unclipped = optim.AdamW(model_unclipped.parameters(), lr=1e-3)

clip_threshold = 1.0
clipped_norms = []
unclipped_norms = []

for step in range(80):
    x = torch.randn(16, 32, 128)
    target = torch.randn(16, 32, 128)

    # Unclipped model
    opt_unclipped.zero_grad()
    loss_u = criterion(model_unclipped(x), target)
    loss_u.backward()
    raw_norm = torch.nn.utils.clip_grad_norm_(
        model_unclipped.parameters(), max_norm=float("inf")
    )
    unclipped_norms.append(raw_norm.item())
    opt_unclipped.step()

    # Clipped model (same batch)
    opt_clipped.zero_grad()
    loss_c = criterion(model_clipped(x), target)
    loss_c.backward()
    clipped_norm = torch.nn.utils.clip_grad_norm_(
        model_clipped.parameters(), max_norm=clip_threshold
    )
    clipped_norms.append(clipped_norm.item())
    opt_clipped.step()
Out[10]:
Console
Steps where clipping activated: 0 of 80 (0.0%)
Mean raw gradient norm: 0.1944
Max raw gradient norm: 0.2130
Mean clipped gradient norm: 0.1944

The fraction of steps where clipping activates tells you whether your threshold is well-calibrated. If clipping activates on more than 20-30% of steps, the threshold is too aggressive and may be impeding learning. In this run it never activates because every pre-clipping norm is well below 1.0, so the example illustrates the pass-through case rather than an actual clipping event.

Out[11]:
Visualization
Line plot of two overlapping pre-clipping gradient-norm trajectories near 0.2, both well below a horizontal threshold at 1.0.
Pre-clipping gradient norms from otherwise identical clipped and unclipped runs over 80 training steps. All norms remain near 0.2, well below the 1.0 threshold, so clipping activates on zero steps and the two trajectories overlap. This is the pass-through behavior of norm clipping: ordinary updates remain unchanged.

Loss Spike Detection and Recovery

Gradient norm monitoring helps you anticipate problems, but loss spikes still occur in practice. Knowing how to detect them reliably and recover efficiently is essential for long training runs.

Detecting Loss Spikes

A loss spike is a sudden increase in training loss that is substantially larger than normal batch-to-batch variance. The challenge is distinguishing a damaging spike from ordinary noise. A reliable detection approach uses a rolling baseline:

spike detected if: Lt>μrecent+k⋅σrecent\text{spike detected if: } \mathcal{L}_t > \mu_{\text{recent}} + k \cdot \sigma_{\text{recent}}

where μrecent\mu_{\text{recent}} and σrecent\sigma_{\text{recent}} are the mean and standard deviation of loss over the last WW steps, and kk is a sensitivity threshold (typically 3-5).

This formulation is essentially a statistical process control chart adapted for training monitoring. The window WW should be long enough to capture a stable baseline (50-200 steps is typical) but short enough to adapt to the gradual downward trend in loss as training progresses. If WW is too short, the baseline is noisy and you get false positives. If WW is too long, the baseline does not adapt and you miss spikes during the early high-loss phase.

The threshold kk balances sensitivity against false alarms. Using k=3k=3 gives a threshold at 3 standard deviations, which under a normal distribution would trigger on about 0.1% of steps by chance. For a training run of 100,000 steps, that is still 100 false alarms. In practice, k=4k=4 or k=5k=5 works better for long runs.

In[12]:
Code
def detect_loss_spikes(loss_history, window=50, threshold_k=3.5):
    """Detect loss spikes using rolling statistics.

    Returns a list of (step, loss_value) for detected spikes.
    """
    spikes = []
    for i in range(window, len(loss_history)):
        recent = loss_history[i - window : i]
        mu = np.mean(recent)
        sigma = np.std(recent)
        if sigma > 0 and loss_history[i] > mu + threshold_k * sigma:
            spikes.append((i, loss_history[i]))
    return spikes


# Simulate training loss with injected spikes
np.random.seed(42)
n_steps = 300
base_loss = 2.0 * np.exp(-np.linspace(0, 3, n_steps)) + 0.3
noise = np.random.normal(0, 0.05, n_steps)
simulated_loss = base_loss + noise

# Inject realistic loss spikes
spike_configs = [(100, 3.5), (180, 2.8), (240, 4.1)]
for step_idx, magnitude in spike_configs:
    simulated_loss[step_idx] *= magnitude
    simulated_loss[step_idx + 1] *= 1.8  # Partial recovery lag

detected_spikes = detect_loss_spikes(simulated_loss.tolist())
Out[13]:
Console
Total training steps: 300
Detected spike events: 3
  Step 100: loss=3.369 (baseline=1.258, ratio=2.7x)
  Step 180: loss=1.848 (baseline=0.734, ratio=2.5x)
  Step 240: loss=1.805 (baseline=0.541, ratio=3.3x)
Out[14]:
Visualization
Line plot of training loss over 300 steps with red triangle markers indicating detected spike events at steps 100, 180, and 240.
Training loss trajectory with automated spike detection using a rolling 50-step window and 3.5-sigma threshold. Detected spikes (marked with red triangles) correspond to sudden loss increases that exceed the expected variance band (orange shaded region), cleanly distinguishing genuine instabilities from normal batch-to-batch fluctuations. The rolling mean (orange line) adapts to the downward trend in loss without being distorted by the spikes themselves.

Recovery Strategies

When a loss spike occurs, the default behavior of most training scripts is to continue training, hoping the optimizer self-corrects. For small spikes this often works. For large spikes (loss returning to near-initialization levels), it usually does not, because the optimizer's momentum state has been corrupted.

The most reliable recovery approach is checkpoint rollback with optimizer state reset:

  1. Maintain checkpoints at regular intervals (every 500-1000 steps is typical for large models).
  2. When a spike is detected, roll back to the most recent checkpoint before the spike.
  3. Optionally skip or re-weight the batch that triggered the spike.
  4. Consider reducing the learning rate slightly (10-20%) before resuming, to allow gradual re-entry into the training trajectory.

The reason optimizer state matters here is subtle but important. Rolling back model parameters without also rolling back optimizer state (momentum buffers) means the optimizer has incorrect expectations about the loss landscape geometry. The momentum carries the "memory" of a corrupted gradient signal, and this can cause additional instability immediately after recovery. Always checkpoint and restore both model parameters and optimizer state together.

A more sophisticated recovery strategy used in some large-scale training runs involves partial rollback: instead of reverting to the most recent checkpoint, you revert only the parameters of the layers that showed the largest gradient norms during the spike. This preserves progress made in stable parts of the network while resetting the layers that were most affected by the instability. This approach requires per-layer gradient monitoring infrastructure but can save significant compute when spikes happen late in training.

Stability Techniques

Beyond gradient clipping, several architectural and algorithmic choices improve training stability. Understanding these techniques and why they work allows you to apply them appropriately rather than blindly copying configuration choices from other projects.

Layer Normalization Placement

Normalization layers are powerful stabilizers because they prevent the magnitude of activations from drifting unboundedly through forward passes. As covered in the layer normalization chapter, LayerNorm operates independently per sample, making it compatible with variable-length sequences and small batch sizes, unlike BatchNorm.

The placement of layer normalization has a substantial effect on stability. Two conventions exist:

  • Post-norm: LayerNorm applied after the residual connection, as in the original "Attention Is All You Need" paper. The residual path passes through unnormalized, which can lead to large activation magnitudes in deep models.
  • Pre-norm: LayerNorm applied before the self-attention or FFN sub-layer, inside the residual path. The residual connection accumulates clean (unnormalized) values, giving a stable gradient highway through depth.

Pre-norm transformers are substantially more stable for deep architectures. GPT-2 and later models use pre-norm by default. With post-norm, it is common to need aggressive learning rate warmup (thousands of steps) to avoid early instability. Pre-norm tolerates larger learning rates and shorter warmup schedules.

The mathematical reason is that with pre-norm, gradient flow through the residual connection is direct and unobstructed. In a pre-norm transformer, each layer computes:

x(l+1)=x(l)+f(LN(x(l)))\mathbf{x}^{(l+1)} = \mathbf{x}^{(l)} + f\bigl(\text{LN}(\mathbf{x}^{(l)})\bigr)

where x(l)\mathbf{x}^{(l)} is the residual stream at layer ll, LN(⋅)\text{LN}(\cdot) is layer normalization, and f(⋅)f(\cdot) is the attention or FFN sub-layer. The gradient of the loss with respect to x(l)\mathbf{x}^{(l)} is:

∂L∂x(l)=∂L∂x(l+1)⋅(1+∂f(LN(x(l)))∂x(l))\frac{\partial \mathcal{L}}{\partial \mathbf{x}^{(l)}} = \frac{\partial \mathcal{L}}{\partial \mathbf{x}^{(l+1)}} \cdot \left(1 + \frac{\partial f(\text{LN}(\mathbf{x}^{(l)}))}{\partial \mathbf{x}^{(l)}}\right)

where:

  • x(l)\mathbf{x}^{(l)}: input to layer ll (the residual stream value)
  • f(⋅)f(\cdot): the attention or FFN sub-layer function applied after LayerNorm
  • LN(⋅)\text{LN}(\cdot): layer normalization applied before the sub-layer
  • The "1 +" term: the direct residual path that always passes gradient through unchanged

The "1 +" in the gradient equation means the residual connection provides an additive path that is always present regardless of the sub-layer's behavior. In pre-norm, this path carries the unchanged residual stream values, creating a gradient highway from output to input that is independent of depth. In post-norm, this path goes through an additional normalization operation that can attenuate gradients in early training.

A recent variant, RMSNorm (root mean square normalization), provides similar stability benefits to LayerNorm but with reduced computational cost. It removes the mean subtraction step of standard LayerNorm, relying only on scale normalization. LLaMA and many subsequent models use RMSNorm with pre-norm placement, combining the stability benefits of both techniques.

Residual Connections and Initialization Scaling

Residual connections, introduced to enable very deep network training, also contribute directly to stability by providing gradient highways through depth. But they interact with initialization in a subtle way that large model practitioners have learned to address carefully.

In a transformer with LL layers, if each sub-layer adds a residual contribution with unscaled weights, the variance of the residual stream grows with depth. Specifically, if each sub-layer output has variance σf2\sigma_f^2, the residual stream after LL layers has variance approximately Lσf2L \sigma_f^2, growing linearly with depth. For models with 24, 48, or 96 layers, this can produce activations with large magnitudes even at initialization, before any training.

The solution adopted in GPT-2 and variants is to scale down the initialization of the final projection in each sub-layer by a factor of 12L\frac{1}{\sqrt{2L}}, where LL is the number of transformer layers. Specifically, the output projection of the attention block and the second linear layer of the FFN block are initialized with standard deviation σ2L\frac{\sigma}{\sqrt{2L}}, where σ\sigma is the default initialization standard deviation. This ensures that the total variance contributed to the residual stream across all layers remains approximately constant:

Var[x(L)]≈Var[x(0)]+L⋅σf22L=Var[x(0)]+σf22\text{Var}[\mathbf{x}^{(L)}] \approx \text{Var}[\mathbf{x}^{(0)}] + L \cdot \frac{\sigma_f^2}{2L} = \text{Var}[\mathbf{x}^{(0)}] + \frac{\sigma_f^2}{2}

This is initialization scale invariance with respect to depth: adding more layers does not change the activation magnitudes at initialization, which is a desirable property for numerical stability.

Learning Rate Scheduling and Warmup

As discussed in the learning rate warmup and cosine schedule chapters, the learning rate schedule has a direct effect on stability. The warmup phase improves performance and is also required for stable training in many architectures.

Early training is the most fragile phase. Model weights are far from their final values, the loss landscape is poorly characterized, and the optimizer has no momentum history to guide it. A large learning rate at initialization sends the optimizer on large random walks that can easily land in regions of extreme curvature.

Learning rate warmup addresses this by starting with an extremely small learning rate (often 1-10% of the peak value) and linearly increasing it over the first few thousand steps. This allows the optimizer to take small, cautious steps while building up accurate momentum estimates. By the time the learning rate reaches its peak, the model is in a better-conditioned region of parameter space and the optimizer's momentum reflects the actual loss landscape geometry.

The warmup duration is typically expressed as a fixed number of steps (1000-2000 for small models, 4000-10000 for large models) rather than as a fraction of total training. The reason is that the required warmup duration scales with model depth and initial instability, not with the total number of training steps. A 175B parameter model requires roughly the same warmup duration regardless of whether you train it for 100B or 300B tokens.

Post-warmup learning rate decay also contributes to long-term stability. As the model approaches a good region of parameter space, reducing the learning rate prevents the optimizer from overshooting and allows fine-grained convergence. Cosine decay is preferred because it is smooth (no discontinuities that could cause sudden instability) and naturally brings the learning rate to near-zero at the end of training.

Adam Epsilon and Numerical Stability

The Adam optimizer has a small constant ϵ\epsilon in its denominator to prevent division by zero:

θt=θt−1−ηv^t+ϵ⋅m^t\theta_t = \theta_{t-1} - \frac{\eta}{\sqrt{\hat{v}_t} + \epsilon} \cdot \hat{m}_t

where:

  • v^t\hat{v}_t: the bias-corrected second moment estimate (variance of recent gradients)
  • m^t\hat{m}_t: the bias-corrected first moment estimate (mean of recent gradients)
  • ϵ\epsilon: a small constant for numerical stability (default 10−810^{-8} in most frameworks)
  • η\eta: the learning rate

The default ϵ=10−8\epsilon = 10^{-8} works well for standard float32 training. But for mixed-precision training with float16, this value is often too small relative to the precision of float16 arithmetic. Float16 can represent values as small as approximately 6×10−56 \times 10^{-5} before underflow to zero. If v^t\hat{v}_t is small (which happens early in training or for parameters that receive weak gradient signals), then v^t\sqrt{\hat{v}_t} can be close to or below the float16 precision floor, making the ϵ\epsilon in the denominator effectively invisible. The result is effective division by near-zero, which causes enormous parameter updates.

A common fix for instability in mixed-precision training is to increase ϵ\epsilon to 10−610^{-6} or even 10−510^{-5}. This makes the effective step size slightly less adaptive for parameters with very small second moments, but significantly more numerically stable. Google's Adafactor optimizer, designed for extreme-scale models, uses an even larger effective epsilon to ensure stability.

In[15]:
Code
def run_training_with_epsilon(epsilon, n_steps=200, lr=1e-3, use_fp16=False):
    """Train a model with a specific Adam epsilon and return loss trajectory."""
    torch.manual_seed(42)
    model_eps = SimpleTransformerBlock(d_model=64, n_heads=4)
    dtype = torch.float16 if use_fp16 else torch.float32
    model_eps = model_eps.to(dtype=dtype)
    opt = optim.Adam(model_eps.parameters(), lr=lr, eps=epsilon)
    loss_history = []

    for step in range(n_steps):
        x = torch.randn(8, 16, 64, dtype=dtype)
        target = torch.randn(8, 16, 64, dtype=dtype)
        opt.zero_grad()
        try:
            out = model_eps(x)
            loss = criterion(out.float(), target.float())
            loss.backward()
            opt.step()
            loss_history.append(loss.item())
        except RuntimeError:
            loss_history.append(float("nan"))

    return loss_history


epsilon_values = [1e-8, 1e-7, 1e-6]
epsilon_results = {}
for eps in epsilon_values:
    epsilon_results[eps] = run_training_with_epsilon(eps, use_fp16=False)
Out[16]:
Visualization
Line plot showing training loss for three epsilon settings (1e-8, 1e-7, 1e-6) converging at similar rates over 200 steps.
Training loss curves for Adam with three different epsilon values over 200 steps in standard float32 mode. All three settings converge similarly under float32, showing that the choice of epsilon matters less in full precision. The critical distinction emerges in float16 mixed-precision training, where smaller epsilon values can cause numerical underflow in the second moment estimates, leading to instability that does not manifest in this float32 comparison.

Weight Decay and Regularization

Weight decay, as covered in the AdamW chapter, constrains parameter magnitudes by penalizing large weights. From a stability perspective, this has an important secondary effect: it prevents any individual weight from growing to extreme values that would dominate the gradient signal in subsequent steps. Bounded weights mean bounded activation magnitudes (given normalized inputs), which in turn means bounded gradients.

AdamW's decoupled weight decay is more effective for stability than the L2 regularization form of Adam. In standard Adam with L2 regularization, the regularization signal is mixed with the gradient signal and then adaptively scaled, which weakens its effect. When a parameter has a small second moment estimate (that is, it has received weak gradients recently), Adam applies a very large adaptive scaling factor. If L2 regularization is included in the gradient, this large factor also applies to the regularization term, which can cause the regularization to suddenly push parameters far from their current values. AdamW applies weight decay directly to parameters before the gradient update, avoiding this interaction entirely.

QK Normalization in Attention

A more recent stability technique, adopted in models like PaLM 2 and some variants of GPT-4, is query-key normalization in the attention mechanism. The attention computation involves dot products between query and key vectors:

Attention(Q,K,V)=softmax ⁣(QKTdk)V\text{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{softmax}\!\left(\frac{\mathbf{Q}\mathbf{K}^T}{\sqrt{d_k}}\right)\mathbf{V}

For very large models with large head dimensions dkd_k, or for long sequences where many tokens contribute to the dot product sum, the pre-softmax logits QKT/dk\mathbf{Q}\mathbf{K}^T / \sqrt{d_k} can still reach large magnitudes. When these logits are very large, the softmax becomes extremely peaked (a near-one-hot distribution), the gradient of the softmax with respect to the logits approaches zero everywhere, and training slows dramatically or becomes unstable.

QK normalization applies LayerNorm or L2 normalization to the query and key vectors before computing the dot product. This bounds the range of attention logits independent of sequence length or model scale, giving a structural guarantee against attention logit explosion. The technique adds almost no computational overhead but provides meaningful stability benefits for large models trained on long sequences.

Debugging Training Instability

When instability occurs, the debugging process requires methodical isolation. A spike in training loss could stem from many causes, and the symptoms often look similar regardless of root cause.

A Systematic Debugging Workflow

Effective debugging follows a structured approach that eliminates hypotheses one by one.

Step 1: Isolate the timeline. Identify exactly when the instability began by examining gradient norm logs alongside loss. If the gradient norm spikes before the loss spike, the cause is an optimizer issue. If loss spikes without a gradient norm precursor, the cause may be a data issue or numerical overflow.

Step 2: Examine the triggering batch. If your training loop logs data indices, retrieve the batch that preceded the spike and inspect it. In language modeling, look for sequences with unusual token distributions, very long lengths, or repeated patterns. These outlier batches are common triggers.

Step 3: Check for numerical issues. NaN or Inf values in gradients or activations always indicate a numerical problem. Adding NaN checks after loss computation and after backpropagation immediately narrows the search:

In[17]:
Code
def check_for_nan(model, step, loss):
    """Check for NaN or Inf values in loss, gradients, and parameters."""
    issues = []

    if torch.isnan(loss) or torch.isinf(loss):
        issues.append(f"Step {step}: NaN/Inf in loss (value={loss.item()})")

    for name, param in model.named_parameters():
        if param.grad is not None:
            if torch.any(torch.isnan(param.grad)):
                issues.append(f"Step {step}: NaN gradient in {name}")
            if torch.any(torch.isinf(param.grad)):
                issues.append(f"Step {step}: Inf gradient in {name}")
        if torch.any(torch.isnan(param.data)):
            issues.append(f"Step {step}: NaN weight in {name}")

    return issues

Step 4: Bisect the architecture. If numerical issues are found in specific layers, reduce the model to just those layers and reproduce the problem with a minimal example. This makes the root cause much easier to identify and fix.

Step 5: Check initialization. Poorly initialized models are fragile in early training. Re-running the first 1000 steps with a different random seed can determine whether instability is systematic (same architecture, same data, same behavior) or stochastic (different seed resolves it). Systematic instability requires architectural or hyperparameter fixes; stochastic instability may simply require better luck or a more reliable initialization scheme.

Layer-Wise Gradient Analysis

When the global gradient norm shows a spike, it does not tell you which layers are responsible. Layer-wise gradient analysis decomposes the gradient norm by component, revealing which parts of the network are generating the instability. This is valuable for architectural debugging.

In a transformer, the most common sources of gradient spikes are:

  • The embedding layer, which can receive very large gradients when the vocabulary is large and some tokens are rare
  • The final projection (language model head), which connects the full model width to the vocabulary size and amplifies gradients proportionally
  • Attention output projections, especially in early layers where gradients from all subsequent layers accumulate
In[18]:
Code
def layer_gradient_analysis(model):
    """Compute per-layer gradient norms for architectural debugging."""
    layer_norms = {}
    for name, param in model.named_parameters():
        if param.grad is not None:
            layer_norms[name] = param.grad.norm(2).item()
    return layer_norms


# Run a single forward-backward pass and analyze layer gradients
torch.manual_seed(42)
debug_model = SimpleTransformerBlock(d_model=128, n_heads=4)
debug_opt = optim.AdamW(debug_model.parameters(), lr=1e-3)

x_debug = torch.randn(16, 32, 128)
target_debug = torch.randn(16, 32, 128)
debug_opt.zero_grad()
out_debug = debug_model(x_debug)
loss_debug = criterion(out_debug, target_debug)
loss_debug.backward()

layer_norms = layer_gradient_analysis(debug_model)
Out[19]:
Console
Loss: 2.0107

Per-layer gradient norms:
Layer                                               Grad Norm
--------------------------------------------------------------
norm2.weight                                         0.177542
ffn.2.weight                                         0.058973
ffn.0.weight                                         0.030297
norm2.bias                                           0.016122
attention.in_proj_weight                             0.012976
attention.out_proj.weight                            0.012630
attention.out_proj.bias                              0.007798
ffn.2.bias                                           0.007755
norm1.bias                                           0.007732
norm1.weight                                         0.007320
attention.in_proj_bias                               0.004454
ffn.0.bias                                           0.002660

This output reveals the relative contribution of each layer to the total gradient signal. If one layer consistently shows gradient norms 10-100x larger than others, that is the component deserving architectural attention.

Monitoring Infrastructure

Effective debugging requires proper monitoring infrastructure set up before problems occur, not after. The minimum viable monitoring for a large training run includes:

  • Loss curve: Training loss (and validation loss if available) at every step or every N steps
  • Gradient norm: Computed and logged before each optimizer step
  • Parameter norms: The Frobenius norm of each weight matrix layer, logged every 100-200 steps. Rapidly growing parameter norms indicate the weight decay coefficient is too small.
  • Learning rate: The current learning rate at each step (especially during warmup and decay phases)
  • Gradient-to-parameter ratio: The ratio ∥g∥/∥θ∥\|\mathbf{g}\| / \|\boldsymbol{\theta}\| characterizes how large updates are relative to parameter magnitudes. Values above 0.1 are often associated with instability.
In[20]:
Code
def compute_monitoring_metrics(model, loss, step):
    """Compute comprehensive monitoring metrics for one training step."""
    metrics = {"step": step, "loss": loss.item()}

    # Gradient norm
    total_grad_norm_sq = 0.0
    for param in model.parameters():
        if param.grad is not None:
            total_grad_norm_sq += param.grad.data.norm(2).item() ** 2
    metrics["grad_norm"] = total_grad_norm_sq**0.5

    # Parameter norm (Frobenius across all params)
    total_param_norm_sq = 0.0
    for param in model.parameters():
        total_param_norm_sq += param.data.norm(2).item() ** 2
    metrics["param_norm"] = total_param_norm_sq**0.5

    # Gradient-to-parameter ratio
    if metrics["param_norm"] > 0:
        metrics["grad_param_ratio"] = (
            metrics["grad_norm"] / metrics["param_norm"]
        )
    else:
        metrics["grad_param_ratio"] = float("nan")

    return metrics


# Run a monitored training loop
torch.manual_seed(42)
monitored_model = SimpleTransformerBlock(d_model=128, n_heads=4)
monitored_opt = optim.AdamW(
    monitored_model.parameters(), lr=1e-3, weight_decay=0.01
)
monitoring_log = []

for step in range(60):
    x = torch.randn(16, 32, 128)
    target = torch.randn(16, 32, 128)
    monitored_opt.zero_grad()
    out = monitored_model(x)
    loss_val = criterion(out, target)
    loss_val.backward()
    metrics = compute_monitoring_metrics(monitored_model, loss_val, step)
    torch.nn.utils.clip_grad_norm_(monitored_model.parameters(), max_norm=1.0)
    monitored_opt.step()
    monitoring_log.append(metrics)
Out[21]:
Console
  Step     Loss  Grad Norm  Param Norm  g/p Ratio
----------------------------------------------------
     0   2.0107     0.1917     26.5670    0.00722
    10   1.9674     0.1889     26.5752    0.00711
    20   1.9495     0.1889     26.5815    0.00711
    30   1.9310     0.1885     26.6000    0.00709
    40   1.8933     0.1896     26.6370    0.00712
    50   1.8609     0.1924     26.7158    0.00720
Out[22]:
Visualization
Dual-axis line plot showing gradient norm near 0.19 and gradient-to-parameter ratio near 0.007 over 60 steps, with both dipping and then rising slightly.
Gradient norm and gradient-to-parameter ratio over 60 monitored training steps. The two metrics follow the same shallow dip and late rise because the parameter norm changes only slightly; the gradient norm stays between roughly 0.188 and 0.200, far below the 1.0 clipping threshold, while the ratio remains near 0.007.
Line plot showing parameter norm increasing modestly from about 26.57 to 26.84 over 60 training steps.
Parameter norm over 60 training steps with AdamW weight decay. It increases modestly from about 26.57 to 26.84: weight decay opposes parameter growth, but the task-gradient updates are larger than the decay contribution in this short run.

Worked Example: Diagnosing a Simulated Instability

Let us walk through a realistic scenario: a small transformer model that encounters instability and recovers. This consolidates the monitoring, detection, and recovery concepts into a concrete workflow that mirrors what practitioners encounter in real large-scale training runs.

The scenario involves a small language model trained on random token sequences, with two "hard batches" injected at known steps to amplify the backward signal. We run two versions, one without gradient clipping and one with clipping. The logged loss is divided by the artificial scale, while PyTorch's clipping function returns the norm measured before rescaling. The comparison therefore shows how similar loss trajectories can hide different update rules at the injected steps.

In[23]:
Code
class MiniLM(nn.Module):
    """Small language model for training stability demonstration."""

    def __init__(self, vocab_size=100, d_model=64, n_layers=4, n_heads=4):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, d_model)
        self.layers = nn.ModuleList(
            [SimpleTransformerBlock(d_model, n_heads) for _ in range(n_layers)]
        )
        self.norm = nn.LayerNorm(d_model)
        self.head = nn.Linear(d_model, vocab_size)

    def forward(self, x):
        h = self.embedding(x)
        for layer in self.layers:
            h = layer(h)
        return self.head(self.norm(h))


def run_stability_scenario(use_clipping=True, clip_val=1.0, seed=42):
    """Run a training scenario with or without gradient clipping."""
    torch.manual_seed(seed)
    np.random.seed(seed)

    model_s = MiniLM(vocab_size=100, d_model=64, n_layers=4, n_heads=4)
    opt_s = optim.Adam(model_s.parameters(), lr=5e-3, eps=1e-8)
    ce_loss = nn.CrossEntropyLoss()

    losses_s = []
    grad_norms_s = []

    for step in range(200):
        # Occasional "hard" batch simulating outlier data
        if step in [60, 120]:
            seq_len = 64
            batch = torch.randint(0, 100, (8, seq_len))
            target_tok = torch.randint(0, 100, (8 * seq_len,))
            scale = 5.0  # amplify loss signal
        else:
            seq_len = 32
            batch = torch.randint(0, 100, (8, seq_len))
            target_tok = torch.randint(0, 100, (8 * seq_len,))
            scale = 1.0

        opt_s.zero_grad()
        logits = model_s(batch).reshape(-1, 100)
        loss_s = ce_loss(logits, target_tok) * scale
        loss_s.backward()

        if use_clipping:
            g_norm = torch.nn.utils.clip_grad_norm_(
                model_s.parameters(), max_norm=clip_val
            )
        else:
            g_norm = torch.nn.utils.clip_grad_norm_(
                model_s.parameters(), max_norm=float("inf")
            )

        grad_norms_s.append(g_norm.item())
        losses_s.append(loss_s.item() / scale)  # log true loss, not scaled
        opt_s.step()

    return losses_s, grad_norms_s


losses_no_clip, gnorms_no_clip = run_stability_scenario(use_clipping=False)
losses_with_clip, gnorms_with_clip = run_stability_scenario(
    use_clipping=True, clip_val=1.0
)
Out[24]:
Visualization
Line plot showing nearly overlapping clipped and unclipped loss trajectories with hard-batch markers at steps 60 and 120.
Logged cross-entropy loss for clipped and unclipped runs over 200 steps, with hard batches marked at steps 60 and 120. Because the artificial hard-batch scale is removed before logging, the two independently trained runs follow nearly identical noisy declines and show no persistent loss spike at either marker.
Line plot showing pre-clipping gradient norms for both runs spiking to about 1.4 at the two hard-batch steps.
Pre-clipping gradient norms for clipped and unclipped runs over 200 steps. Both returned norm series spike to about 1.4 at the hard batches because PyTorch reports the norm before rescaling; in the clipped run, the gradients used for the optimizer update are nevertheless rescaled to the 1.0 threshold.
Out[25]:
Console
Summary: No Clipping vs. Gradient Clipping (threshold=1.0)
Metric                              No Clip    With Clip
--------------------------------------------------------
Mean loss                            4.6297       4.6293
Final loss                           4.6135       4.6138
Max gradient norm                    1.3862       1.3862
Mean gradient norm                   0.4079       0.4095

The results are deliberately more modest than the stylized failure case discussed earlier. The two runs have nearly identical mean and final losses, and both report the same maximum norm because clip_grad_norm_ returns the pre-clipping value. The operational difference is in the update: at steps 60 and 120, the clipped run rescales gradients above 1.0 before optimizer.step(), while the unclipped run applies them at full magnitude. In this short seeded example, that safeguard does not produce a visible loss advantage.

Understanding Loss Spike Anatomy: A Numerical Walkthrough

To build deeper intuition, it is worth walking through the exact sequence of events that happens at the parameter level during a loss spike. This makes the abstract mechanisms concrete and clarifies why the interventions in this chapter are effective.

The Spike as a Three-Phase Event

A loss spike is not instantaneous; it unfolds over several steps. Understanding each phase clarifies what can be done at each point.

Phase 1: The trigger step. A batch arrives with unusually high loss. This could be an outlier document, a sharp curvature cliff in the landscape, or a sudden change in the effective learning rate. The gradient computed from this batch has a norm that is much larger than the running average, say 50x. The optimizer applies this gradient to all parameters.

Phase 2: Parameter displacement. After the large update, model parameters are in a different region of the loss landscape than they were a step earlier. If the update was large enough, they have moved past a local minimum and are now in a high-loss region. This is when the loss curve shows the spike: the next batch sees a higher loss than before the trigger.

Phase 3: Optimizer state corruption. Even if the parameters are eventually restored (by gradient descent pulling them back toward a good region), the Adam optimizer's second moment estimate vtv_t has absorbed a large gradient squared value. Because β2=0.999\beta_2 = 0.999, this inflated estimate decays slowly:

vt+k=β2kvt+(1−β2k)vˉv_{t+k} = \beta_2^k v_t + (1 - \beta_2^k) \bar{v}

where vˉ\bar{v} is the typical second moment value. With β2=0.999\beta_2 = 0.999 and a typical training run, the spike's influence on vtv_t is not halved until about 693 steps later (kk such that 0.999k=0.50.999^k = 0.5). During those 693 steps, Adam's adaptive step size for the affected parameters is artificially suppressed, because the denominator v^t+ϵ\sqrt{\hat{v}_t} + \epsilon is larger than it should be. The model is effectively frozen in those parameters while they slowly "forget" the spike.

A Concrete Numerical Example

Let us trace through the numbers for a single parameter to make this vivid. Suppose a parameter θ\theta has:

  • Current value: θt=0.5\theta_t = 0.5
  • Typical gradient: g∼0.1g \sim 0.1 (normal)
  • Adam state before spike: mt=0.1m_t = 0.1, vt=0.01v_t = 0.01 (so vt=0.1\sqrt{v_t} = 0.1)
  • Learning rate: η=0.001\eta = 0.001, ϵ=10−8\epsilon = 10^{-8}

The normal update magnitude is:

Δθ=ηvt+ϵ⋅mt≈0.0010.1⋅0.1=0.001\Delta\theta = \frac{\eta}{\sqrt{v_t} + \epsilon} \cdot m_t \approx \frac{0.001}{0.1} \cdot 0.1 = 0.001

That is a small, reasonable update. Now a spike batch arrives with gradient gspike=5.0g_{\text{spike}} = 5.0 (50x normal). After one Adam step:

mt+1=0.9×0.1+0.1×5.0=0.59vt+1=0.999×0.01+0.001×25.0=0.034899\begin{aligned} m_{t+1} &= 0.9 \times 0.1 + 0.1 \times 5.0 = 0.59 \\ v_{t+1} &= 0.999 \times 0.01 + 0.001 \times 25.0 = 0.034899 \end{aligned}

The update magnitude at step t+1t+1 is:

Δθt+1=0.0010.034899⋅0.59≈0.0010.1868⋅0.59≈0.00316\Delta\theta_{t+1} = \frac{0.001}{\sqrt{0.034899}} \cdot 0.59 \approx \frac{0.001}{0.1868} \cdot 0.59 \approx 0.00316

That is 3.16x larger than the normal update: the momentum-inflated first moment drives a much larger step. Three hundred steps later, the second moment has decayed to approximately:

vt+300≈0.999300×0.034899+(1−0.999300)×0.01≈0.027+0.0026≈0.0296v_{t+300} \approx 0.999^{300} \times 0.034899 + (1 - 0.999^{300}) \times 0.01 \approx 0.027 + 0.0026 \approx 0.0296

So even 300 steps after the spike, vtv_t is still 3x its pre-spike value, and the adaptive step size for this parameter is still suppressed by a factor of 3≈1.73\sqrt{3} \approx 1.73. The parameter is learning at 58% of its normal rate for 300+ steps after a single bad gradient. At scale, with thousands of affected parameters, this explains the "loss plateau" that practitioners often observe after a spike even when the loss appears to have recovered.

Gradient clipping prevents all of this from happening. By capping gspikeg_{\text{spike}} at 1.0 (assuming a clip threshold of 1.0), the spike's contribution to mtm_t and vtv_t is:

mt+1=0.9×0.1+0.1×1.0=0.19vt+1=0.999×0.01+0.001×1.0=0.011\begin{aligned} m_{t+1} &= 0.9 \times 0.1 + 0.1 \times 1.0 = 0.19 \\ v_{t+1} &= 0.999 \times 0.01 + 0.001 \times 1.0 = 0.011 \end{aligned}

This is only slightly above normal. The optimizer state remains well-behaved, and training continues smoothly after the clipped step.

The Stability-Quality Tradeoff

Not every large gradient signal is a problem. This is worth exploring in depth because a naive reading of this chapter might suggest that stability should be maximized at all costs, which is incorrect.

When Large Gradients Are Informative

Some large gradients are large for good reasons. When a model encounters an unusually difficult example that it is poorly calibrated for, the gradient from that example carries a strong, useful learning signal. Clipping that gradient prevents the model from fully incorporating the lesson.

Consider training a language model on domain-specific scientific text. The model, trained predominantly on general web text, initially assigns very low probability to specialized terminology and technical phrasing. When the training corpus introduces scientific papers, the gradients from those documents are large, showing the model's large prediction error. These large gradients are not errors; they are the mechanism by which the model learns the new domain. Aggressively clipping them slows this adaptation.

This is one reason why the optimal clipping threshold is not zero, and why models trained with no clipping (if they can train stably) sometimes outperform models trained with aggressive clipping on in-distribution test sets, while being worse on the out-of-distribution metrics. The stability-accuracy frontier depends on the data distribution and the model's starting point.

Soft Clipping as a Middle Ground

Standard gradient clipping applies a hard constraint: gradients above the threshold are rescaled to exactly the threshold. A softer alternative applies a smooth nonlinear transformation that preserves small gradients exactly, applies light smoothing to moderate gradients, and more aggressively bounds large gradients. One such formulation uses:

g^i=gi⋅cmax⁡(c,∥g∥)\hat{g}_i = g_i \cdot \frac{c}{\max(c, \|\mathbf{g}\|)}

which is just the standard clipping formula, but the key insight is that the "smoothness" can be achieved by using a smaller clip threshold in combination with a higher learning rate. The net effect is equivalent to using a larger clip threshold with a lower learning rate, but the framing as soft clipping makes the tradeoff explicit.

In practice, practitioners rarely use formulations more exotic than standard global norm clipping, because the tuning cost of additional hyperparameters rarely pays off. Standard clipping with a well-calibrated threshold and monitored clip frequency is sufficient for almost all applications.

Loss Spike Tolerance During Training

Not every spike requires intervention. Small spikes (2-3x the rolling average) in the middle of a long stable training run often resolve themselves within 10-20 steps as the optimizer's momentum returns to normal. Intervening with a checkpoint rollback for every small fluctuation is unnecessarily disruptive and wastes compute.

A reasonable policy for a practical training system is:

  • Small spikes (loss 2-5x rolling average for 1-3 steps): monitor, no intervention
  • Medium spikes (loss 5-20x rolling average or persisting more than 10 steps): reduce learning rate by 10-20% for 500 steps, then resume
  • Large spikes (loss 20x+ rolling average or divergence to infinity): checkpoint rollback with optimizer state reset

This tiered approach avoids the overhead of frequent rollbacks while ensuring that catastrophic instabilities are addressed before they waste significant compute. The thresholds depend on model size and training stage; early in training, higher variance is normal and intervention thresholds should be set looser.

Data Quality and Its Effect on Stability

The role of data quality in training stability is substantial but often underemphasized relative to architectural and optimizer choices. This section addresses the data-centric view of stability.

Why Web-Scraped Data Is Unstable

Large language model pretraining relies on web-scraped corpora containing text from billions of web pages. These corpora are extraordinarily diverse, which is a feature for coverage. But they also contain a long tail of document types that produce unusually high loss during training.

Documents that are problematic from a stability perspective tend to have one or more of these properties:

  • High repetition: Documents consisting of repeated phrases or paragraphs produce high-entropy gradients because the model cannot rely on context to predict the next token. The gradient is dominated by all-or-nothing predictions.
  • Encoding artifacts: Documents with non-UTF-8 encoding that has been force-converted produce token sequences that look like random byte sequences. The cross-entropy loss for random sequences is much higher than for natural language, producing proportionally large gradients.
  • Code with many special characters: Dense code (especially assembly, compiled bytecode, or minified JavaScript) has very different token statistics from natural language, producing high loss if the model was mostly trained on natural language.
  • Template-generated text: Spam, boilerplate, and SEO-optimized text often contain long runs of repeated content with very predictable structure, which the model over-learns quickly and then shows very high loss on the small distinctive tokens.

The practical fix is to filter these document types during data preprocessing, before they enter the training corpus. The most effective filtering signals are:

  • Perplexity under a small reference language model (very high perplexity indicates unusual text)
  • Repetition ratio (fraction of n-grams that appear more than once in the same document)
  • Character diversity (very low character diversity indicates encoding artifacts or repeated text)

Curriculum Learning for Stability

Curriculum learning is the practice of ordering training examples from easier to harder, rather than presenting them in random order. From a stability perspective, curriculum learning is appealing because it ensures that the model encounters well-represented, high-frequency patterns first, building a strong foundation before seeing outlier data.

The empirical evidence for curriculum learning improving stability is suggestive but mixed. Some large-scale training experiments have found that starting with documents of moderate complexity and gradually introducing outlier documents reduces early training instability. Others have found no significant benefit, perhaps because the gradient clipping and warmup techniques already handle the outlier batches adequately.

A more targeted version of curriculum learning, sometimes called "data pacing," focuses specifically on the outlier tail rather than full ordering. The idea is to hold back the top 5-10% highest-loss documents for the second half of training, when the model is better calibrated and can learn more reliably from difficult examples. This approach has the advantage of being easy to implement (filter documents by initial loss under a small model) and addresses specifically the data-driven cause of stability problems.

Dynamic Data Filtering During Training

A powerful but computationally expensive technique is dynamic data filtering: computing the loss for each document in a candidate batch before applying the gradient, and discarding documents whose loss exceeds a threshold. This ensures that no single document with extreme loss corrupts the gradient estimate.

The computational cost is one forward pass per batch to compute per-document losses. For large models, this doubles the per-step compute. But for models where stability is critical and compute is available, it can substantially reduce spike frequency. Some practitioners implement a softer version: re-weight documents by the inverse of their loss (so high-loss documents contribute less gradient signal) rather than discarding them entirely, preserving some learning signal from difficult documents.

Stability in Practice: A Reference Checklist

Bringing together the techniques from this chapter, here is a practical reference for building a stable training configuration:

Architecture choices:

  • Use pre-norm (LayerNorm before sub-layers) rather than post-norm for transformers with more than 12 layers
  • Scale down output projections by 1/2L1/\sqrt{2L} at initialization to prevent residual stream variance growth
  • Consider QK normalization for models with very large heads or trained on very long sequences

Optimizer choices:

  • Use AdamW with β1=0.9\beta_1 = 0.9, β2=0.95\beta_2 = 0.95 or 0.9990.999
  • Set ϵ=10−8\epsilon = 10^{-8} for float32 training; increase to 10−610^{-6} for mixed-precision
  • Apply gradient clipping at c=1.0c = 1.0 as a starting point; tune based on observed clip frequency
  • Use non-zero weight decay (typically 0.01-0.1) to prevent unbounded parameter growth

Scheduling choices:

  • Warm up the learning rate linearly for 1000-4000 steps depending on model size
  • Use cosine decay after warmup for smooth long-term stability
  • Avoid sudden learning rate changes after warmup (no step functions)

Monitoring choices:

  • Log gradient norm at every step
  • Log per-layer gradient norms every 100 steps for architectural debugging
  • Log parameter norms every 200 steps to detect weight decay failure
  • Set up automated spike detection with rolling window statistics
  • Checkpoint every 500-1000 steps with full optimizer state

Data choices:

  • Filter documents with extreme token repetition, very high loss, or encoding artifacts
  • Consider curriculum learning (easier examples first) to avoid outlier batches in early training
  • Log the data indices of batches that precede loss spikes for post-hoc analysis

Limitations and Practical Considerations

Gradient clipping is effective but not a complete solution to training instability. It addresses the symptom (large gradient steps) without necessarily addressing the cause (a difficult loss landscape, a bad data distribution, or suboptimal hyperparameters). A model that requires constant aggressive clipping, say more than 50% of steps hitting the threshold, is signaling a deeper problem that deserves investigation rather than simply continuing to clip. When clipping activates this frequently, the practical effect is that you have reduced the effective learning rate substantially, and it may be more principled to simply lower the nominal learning rate and relax the clip threshold, giving you more transparent control over the training dynamics.

The interaction between gradient clipping and adaptive optimizers deserves careful thought. Adam already scales gradients by an estimate of their recent variance; adding norm clipping on top creates a system where the effective learning rate is constrained by two separate mechanisms. For practitioners, this means that the Adam learning rate and the clipping threshold are not fully independent hyperparameters. Aggressive clipping effectively reduces the learning rate for large-gradient steps, which may require a compensating increase in the base learning rate. This interdependence makes joint tuning more important than tuning each parameter in isolation. A useful mental model is to think of gradient clipping as a "circuit breaker" for outlier batches, not as a substitute for appropriate learning rate scheduling.

Loss spike frequency is also a function of dataset quality in ways that are easy to underestimate. Web-scraped text corpora (used for most large language model pretraining) contain a long tail of extremely unusual documents: very long repetitive sequences, documents consisting entirely of symbols, encoding artifacts, and spam. Even with gradient clipping, these sequences can cause persistent elevation in the loss for many steps after the batch is processed, because the optimizer state carries the signal forward. Data filtering pipelines that remove or downweight outlier documents are a complementary and often underinvested stability technique. A rough heuristic: filtering out the top 0.1% of highest-loss documents from a training corpus can reduce spike frequency without noticeably reducing data diversity. In absolute terms, for a training corpus of 1 trillion tokens, this corresponds to removing about 1 billion tokens from the most problematic documents, which is a very small fraction but disproportionately impacts stability.

There is also a tension between stability and expressiveness that deserves acknowledgment. Many stability techniques, including gradient clipping, weight decay, and conservative learning rate schedules, are forms of regularization that limit how aggressively the model can update in response to any single batch. This is the right tradeoff for most training scenarios. But sometimes the gradient norm is large because the model is making an important learning step from a difficult batch. Clipping that gradient prevents the model from fully capturing that signal. This is one reason why stability and final model quality are sometimes in tension, and why finding the right balance (rather than maximizing stability at all costs) is the practitioner's goal.

Stability techniques scale differently with model size. The interventions in this chapter, gradient clipping, pre-norm architectures, learning rate warmup, and optimizer hyperparameter tuning, work at any model scale. But their relative importance changes with scale. For small models (under 1B parameters), stability is easy to achieve and most runs succeed without careful tuning. For models above 10B parameters, stability becomes the dominant operational concern. Long runs (300B+ tokens) have more exposure to rare outlier batches that can trigger instabilities. Very deep models (48+ layers) are more sensitive to initialization and normalization choices. The techniques in this chapter become more valuable and more necessary as model scale increases, which is why they have received so much attention in the era of large-scale pretraining.

Finally, the monitoring infrastructure outlined in this chapter is not optional for serious training runs. It is easy to set up and dramatically reduces the time to diagnose and fix instabilities when they occur. Gradient norm logs can also be valuable retroactively: by examining the gradient norm history at the time a loss spike began, you can often distinguish between a data problem (single-step spike with no precursor), an optimizer state problem (gradual gradient norm increase before a spike), and an architectural problem (gradient norms that are uniformly extreme from early in training). Each of these root causes has different remedies, and the monitoring record is the evidence that distinguishes them. Treating monitoring as a first-class concern, not an afterthought, is one of the most cost-effective investments you can make in a large training project.

Summary

Training stability is an active concern throughout the lifetime of any large model training run. The key ideas from this chapter are:

  • Loss spikes arise from high-curvature loss landscape regions, exploding gradients through depth, outlier data batches, and optimizer state corruption. Each cause has a different signature and different remedy.
  • Gradient norm monitoring is the most important diagnostic tool. A sudden increase in gradient norm precedes many loss spikes and allows intervention before damage occurs. It is a one-line addition to any training loop.
  • Gradient clipping by global norm is the primary stability technique. It preserves gradient direction while bounding step magnitude, and it is standard practice in all large language model training. Common threshold: 1.0.
  • Architectural choices matter for stability. Pre-norm transformers are significantly more stable than post-norm for deep models. Residual initialization scaling prevents activation variance from growing with depth.
  • Adam hyperparameters affect stability directly. Increasing epsilon from 10−810^{-8} to 10−610^{-6} is often necessary in mixed-precision training. AdamW's decoupled weight decay provides more consistent regularization than L2 regularization in Adam.
  • Data quality is an underinvested stability lever. Filtering outlier documents from pretraining corpora reduces spike frequency without reducing diversity.
  • Recovery from spikes requires rolling back both model parameters and optimizer state together. Rolling back parameters alone leaves corrupted momentum that causes immediate re-instability.
  • Debugging instability requires monitoring infrastructure set up proactively, and a systematic workflow for isolating causes when they occur. Layer-wise gradient analysis is the key tool for architectural debugging.

The next chapter covers mixed-precision training, which introduces its own stability considerations and motivates why the techniques in this chapter are especially critical when working with float16 arithmetic.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about training stability, gradient norm monitoring, and debugging instabilities.

Training Stability Quiz

Question 1 of 80 of 8 completed
What is the primary cause of a loss spike during neural network training?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026trainingstability, author = {Michael Brenndoerfer}, title = {Training Stability: Loss Spikes, Gradient Norms & Debugging}, year = {2026}, url = {https://mbrenndoerfer.com/writing/training-stability-loss-spikes-gradient-norm-debugging}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-10-06} }
APAAcademic
Michael Brenndoerfer (2026). Training Stability: Loss Spikes, Gradient Norms & Debugging. Retrieved from https://mbrenndoerfer.com/writing/training-stability-loss-spikes-gradient-norm-debugging
MLAAcademic
Michael Brenndoerfer. "Training Stability: Loss Spikes, Gradient Norms & Debugging." 2026. Web. October 6, 2026. <https://mbrenndoerfer.com/writing/training-stability-loss-spikes-gradient-norm-debugging>.
CHICAGOAcademic
Michael Brenndoerfer. "Training Stability: Loss Spikes, Gradient Norms & Debugging." Accessed October 6, 2026. https://mbrenndoerfer.com/writing/training-stability-loss-spikes-gradient-norm-debugging.
HARVARDAcademic
Michael Brenndoerfer (2026) 'Training Stability: Loss Spikes, Gradient Norms & Debugging'. Available at: https://mbrenndoerfer.com/writing/training-stability-loss-spikes-gradient-norm-debugging (Accessed: October 6, 2026).
SimpleBasic
Michael Brenndoerfer (2026). Training Stability: Loss Spikes, Gradient Norms & Debugging. https://mbrenndoerfer.com/writing/training-stability-loss-spikes-gradient-norm-debugging

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.