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 , where each entry captures how the gradient of the loss with respect to parameter changes as parameter moves. The largest eigenvalue of , often called the sharpness , determines how large a learning rate is safe. Specifically, gradient descent is guaranteed to decrease loss only when:
where is the learning rate and is the largest eigenvalue of the Hessian. When the optimizer encounters a region where 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 layers, each backward pass involves 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 , the gradient of the loss with respect to the first layer involves a product of the form:
where is the output of the final layer. If the spectral norm of each weight matrix is , then this product can grow as , 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 where is the sequence length, and if , 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:
where is the gradient at step , is the first moment decay, and is the second moment decay. The second moment decay means that the influence of a gradient from 1000 steps ago is reduced by a factor of . 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 .
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:
where:
- : the gradient vector containing all parameter gradients at the current step
- : the gradient component for parameter , computed by backpropagation
- : 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 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 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 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.
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
)

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.
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()# 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()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 , gradient clipping transforms the gradient vector as follows:
where:
- : the raw gradient vector computed by backpropagation
- : the maximum allowed gradient norm (a hyperparameter, typically 1.0)
- : the clipped gradient used for the parameter update
- : 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 with norm and a clip threshold of . Global norm clipping rescales to , which points in exactly the same direction as the original gradient. Component clipping would produce , 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 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.
# 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()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.

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:
where and are the mean and standard deviation of loss over the last steps, and is a sensitivity threshold (typically 3-5).
This formulation is essentially a statistical process control chart adapted for training monitoring. The window 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 is too short, the baseline is noisy and you get false positives. If is too long, the baseline does not adapt and you miss spikes during the early high-loss phase.
The threshold balances sensitivity against false alarms. Using 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, or works better for long runs.
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())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)

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:
- Maintain checkpoints at regular intervals (every 500-1000 steps is typical for large models).
- When a spike is detected, roll back to the most recent checkpoint before the spike.
- Optionally skip or re-weight the batch that triggered the spike.
- 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:
where is the residual stream at layer , is layer normalization, and is the attention or FFN sub-layer. The gradient of the loss with respect to is:
where:
- : input to layer (the residual stream value)
- : the attention or FFN sub-layer function applied after LayerNorm
- : 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 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 , the residual stream after layers has variance approximately , 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 , where 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 , where is the default initialization standard deviation. This ensures that the total variance contributed to the residual stream across all layers remains approximately constant:
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 in its denominator to prevent division by zero:
where:
- : the bias-corrected second moment estimate (variance of recent gradients)
- : the bias-corrected first moment estimate (mean of recent gradients)
- : a small constant for numerical stability (default in most frameworks)
- : the learning rate
The default 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 before underflow to zero. If is small (which happens early in training or for parameters that receive weak gradient signals), then can be close to or below the float16 precision floor, making the 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 to or even . 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.
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)
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:
For very large models with large head dimensions , or for long sequences where many tokens contribute to the dot product sum, the pre-softmax logits 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:
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 issuesStep 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
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)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 characterizes how large updates are relative to parameter magnitudes. Values above 0.1 are often associated with instability.
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) 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

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.
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
)

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 has absorbed a large gradient squared value. Because , this inflated estimate decays slowly:
where is the typical second moment value. With and a typical training run, the spike's influence on is not halved until about 693 steps later ( such that ). During those 693 steps, Adam's adaptive step size for the affected parameters is artificially suppressed, because the denominator 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 has:
- Current value:
- Typical gradient: (normal)
- Adam state before spike: , (so )
- Learning rate: ,
The normal update magnitude is:
That is a small, reasonable update. Now a spike batch arrives with gradient (50x normal). After one Adam step:
The update magnitude at step is:
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:
So even 300 steps after the spike, is still 3x its pre-spike value, and the adaptive step size for this parameter is still suppressed by a factor of . 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 at 1.0 (assuming a clip threshold of 1.0), the spike's contribution to and is:
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:
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 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 , or
- Set for float32 training; increase to for mixed-precision
- Apply gradient clipping at 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 to 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
Reference
Citation details
Cite or share this article.
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 HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

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