Part of Language AI Handbook
Explains how layer normalization enables stable transformer training by normalizing.
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
Layer Normalization
Batch normalization transformed how we train deep feedforward networks, but it stumbles when applied to transformers. The batch dimension becomes problematic: batch sizes vary, sequences have different lengths, and the statistics computed across a batch of diverse sentences lack semantic coherence. Layer normalization, introduced by Ba et al. (2016), sidesteps these issues entirely by normalizing across features rather than across the batch. This seemingly simple change made layer normalization the default normalization technique for transformers, from the original "Attention is All You Need" architecture to modern large language models.
In this chapter, we'll explore why layer normalization works so well for transformers, how it differs from batch normalization in both computation and behavior, and the subtle implementation details that affect training stability. We'll also examine how the placement of layer normalization within transformer blocks affects learning dynamics, a design choice that has evolved significantly since the original transformer architecture.
To appreciate why layer normalization matters, you need to first understand what can go wrong without it. Deep neural networks are sensitive creatures. The gradient signal that propagates through dozens or hundreds of layers is amplified and attenuated by countless multiplication operations. When activations at one layer are large, subsequent layers receive outsized gradients that cause unstable weight updates. When activations are small, gradients vanish before they can reach the earliest layers. Without intervention, training can become brittle or slow and may fail catastrophically.
Normalization techniques solve this by actively constraining the statistical distribution of activations at each layer. Rather than hoping that careful initialization and learning rate schedules will keep activations well-behaved throughout training, normalization directly resets the distribution at predefined points in the network. The question is not whether to normalize, but how. Different normalization strategies make different tradeoffs between computational cost, statistical quality, and architectural compatibility. Batch normalization was the answer for convolutional networks. Layer normalization turned out to be the answer for transformers.
The key intuition is that transformers process each token as an independent entity. When the model applies attention, each token gathers context from others, but each token's representation still flows through the layers as a distinct vector. It makes sense, then, to normalize each representation individually, based on its own statistics, rather than blending its statistics with those of other samples in the batch. Layer normalization does exactly this, and the result is a normalization scheme that works consistently regardless of batch size, sequence length, or what other samples happen to be processed at the same time.
The problem of training instability in deep networks predates transformers by decades. Early neural networks in the 1980s and 1990s were often limited to just two or three layers, partly because deeper networks were nearly impossible to train reliably. Vanishing and exploding gradients, first formally analyzed by Hochreiter in 1991 and Bengio et al. in 1994, explained why: gradients either shrank exponentially as they propagated backward, starving early layers of learning signal, or grew exponentially, causing catastrophic weight updates.
Batch normalization, introduced by Ioffe and Szegedy in 2015, was a breakthrough for convolutional networks, enabling stable training of very deep architectures. However, its reliance on batch statistics made it unsuitable for recurrent models and transformers. Layer normalization (Ba et al., 2016) solved this by shifting from batch statistics to per-sample statistics. The original transformer paper (Vaswani et al., 2017) adopted layer normalization from the outset, and every major language model since, including BERT, GPT, T5, and LLaMA, has built on this foundation. More recently, variants like RMSNorm and Pre-LayerNorm have refined the approach further, trading some statistical richness for computational efficiency.
Why Batch Normalization Fails for Transformers
Before diving into layer normalization, it's worth understanding exactly why batch normalization doesn't work well for sequence models. The core issue is that batch normalization computes statistics across the batch dimension, assuming that each position in a layer sees similar data across samples.
In a transformer processing sentences of varying lengths, each position in the sequence represents something different. Position 0 might be "The" in one sentence and "Scientists" in another. Position 50 might be a verb in one sentence, a noun in another, and padding in a third. Computing a mean and variance across these semantically unrelated positions produces statistics that don't reflect any meaningful property of the data.
Think of it this way: batch normalization operates like a quality control station on an assembly line. It measures the typical behavior of a particular station across many products passing through it, then adjusts each product to conform to that average. This works well when all products are similar, as in image processing where every image has the same spatial structure. But in language, the "products" (tokens) at each position are wildly different across samples. A batch normalization layer at position 10 has to average over "the" from one sentence, "president" from another, and a padding token from a third. The resulting statistics are meaningless.
There are two additional practical problems. First, batch normalization requires a reasonably large batch size to produce stable statistics. With small batches (which are common when training large models on long sequences), the estimated mean and variance become noisy, degrading training. Second, batch normalization maintains running statistics during training that are used at inference time. These running statistics must match the test-time distribution, which can be tricky to ensure and makes model deployment more complex.
Layer normalization sidesteps all of these problems by computing statistics entirely within each sample, independently of the batch. Each token's representation is normalized using only that token's own features. The statistics are exact, not estimated; they don't depend on batch size; and there are no running statistics to maintain. Training and inference use the same computation.
import torch
# Simulate a batch of sequences with varying content
batch_size = 4
seq_len = 8
hidden_dim = 16
# Different sequences have very different activation patterns
activations = torch.randn(batch_size, seq_len, hidden_dim)
# Exaggerate differences between sequences
activations[0] *= 0.5 # First sequence has small activations
activations[1] *= 3.0 # Second sequence has large activations
activations[2] += 5.0 # Third sequence has positive shift
activations[3] -= 5.0 # Fourth sequence has negative shiftActivation statistics per sequence (averaged across positions and features): Sequence 0: mean= 0.041, std=0.481 Sequence 1: mean= 0.135, std=3.056 Sequence 2: mean= 5.027, std=0.976 Sequence 3: mean= -5.042, std=0.935 Batch statistics at position 0: Mean range: [-1.47, 1.68] Var range: [8.38, 28.55]
The batch statistics are dominated by the extreme sequences, and these statistics change dramatically between positions. Layer normalization avoids this problem entirely by computing statistics within each sample independently, treating each token's representation as a self-contained unit to normalize.
Notice that the statistics at position 0 span a wide range because different sequences have wildly different activation magnitudes. Any normalization based on these cross-sample statistics would apply an inappropriate correction to every token in the batch. Layer normalization avoids this problem by never mixing statistics across samples.
The Layer Normalization Formula
To understand layer normalization, let's start with a fundamental question: what does it mean for a neural network layer to have "unstable" activations, and how can we fix it?
The Problem: Activation Scale Drift
Imagine a token's hidden representation as a vector of 768 numbers (a typical transformer dimension). During training, these numbers can drift: some become very large, others very small, and their collective distribution shifts unpredictably. This creates two problems. First, downstream layers must constantly adapt to changing input statistics, making learning inefficient. Second, when values grow too large or too small, gradients either explode or vanish, destabilizing training entirely.
The solution is elegant: before each token's representation moves to the next layer, we transform it to have a predictable, standardized distribution. Specifically, we want the 768 features to have zero mean and unit variance. This "resets" the scale at every layer, preventing drift from accumulating.
The key insight is that we don't need to look outside the current token to achieve this. The token's own 768 values carry all the information we need to compute a mean and a standard deviation. Subtract the mean, divide by the standard deviation, and you have a normalized representation that is stable regardless of how the upstream weights evolved during training.
In practice, this normalization is applied at multiple points within each transformer block. This ensures that neither the attention mechanism nor the feed-forward network ever sees wildly out-of-distribution inputs. The cumulative effect over many layers is dramatic: networks with layer normalization can train stably with learning rates that would cause catastrophic divergence in networks without it.
Step 1: Finding the Center
The first step is computing where the current distribution is centered. Given a hidden state vector representing one token, we calculate its mean:
where:
- : the arithmetic mean of all features in this token's representation
- : the hidden dimension (e.g., 768 for BERT-base, 4096 for LLaMA-7B)
- : the value of the -th feature
This tells us the "center of mass" of the representation. If , the features are shifted toward positive values; if , they lean negative. The goal is to shift this center to zero.
Step 2: Measuring the Spread
Next, we need to know how spread out the values are. A representation where all values cluster tightly around the mean is very different from one where values are scattered widely. We capture this with variance:
where:
- : the variance, measuring how much the features deviate from their mean
- : the squared deviation of each feature from the mean
Squaring ensures that positive and negative deviations don't cancel out. If is large, the features are spread out; if small, they're tightly clustered. We'll use the standard deviation to rescale the values to unit variance.
Notice that we use the population variance (dividing by ) rather than the sample variance (dividing by ). This is a deliberate choice: we're not trying to estimate the variance of an underlying population; we're computing the exact variance of the values in front of us. Using rather than gives us an unbiased estimator for the squared deviation within this specific vector, which is what we need for a deterministic normalization operation.
Step 3: The Normalization Transform
With mean and variance in hand, we can now standardize each feature:
where:
- : the normalized value of the -th feature
- : a tiny constant (typically ) added to prevent division by zero if variance is extremely small
This two-part transformation is exactly what you'd do to standardize any dataset: subtract the mean (centering at zero), then divide by the standard deviation (scaling to unit variance). The result has zero mean and approximately unit variance across the features.
The subtraction by ensures the output is centered at zero, removing any bias in the representation. The division by ensures that features with large variance are scaled down to unit range, and features with small variance are scaled up. This compression of the dynamic range is what prevents activations from exploding across layers.
Step 4: Restoring Flexibility with Learnable Parameters
Layer normalization need not be restrictive. Forcing every representation to have exactly zero mean and unit variance might seem limiting: what if the optimal representation for some layer needs a different distribution?
The solution is to add learnable parameters that can undo the normalization if needed:
where:
- : the final output for the -th feature
- : a learned scale parameter for feature (initialized to 1)
- : a learned shift parameter for feature (initialized to 0)
These parameters are learned during training, just like weights and biases. If the network discovers that feature should have mean 3.5 and standard deviation 2.0, it can learn and to recover that distribution. This means layer normalization never reduces the network's representational power: it starts from a stable baseline but can learn any distribution it needs.
The initialization of and is deliberate. At the start of training, layer normalization acts as a pure standardization with no modification: the output is exactly the normalized input. As training progresses, and adjust to whatever distribution the subsequent layers find most useful. This means the network starts with the stability benefits of normalization and gradually adapts to the specific distribution requirements of the task, without ever losing the ability to represent the original data.
The Complete Formula
Putting all the pieces together, layer normalization maps a hidden state vector of dimension into:
where:
- : the input vector representing one token's hidden state
- : the mean across all features
- : the variance across all features
- : a small stability constant (typically or )
- : learned scale parameters, initialized to ones
- : learned shift parameters, initialized to zeros
- : element-wise multiplication
The formula reads naturally: subtract the mean, divide by the standard deviation (with a safety epsilon), then apply a learned scale and shift.
Why Features, Not Samples?
The key insight that distinguishes layer normalization from batch normalization is the dimension over which we compute statistics. Batch normalization asks: "What's the typical value of feature across all samples in this batch?" Layer normalization asks: "What's the typical value across all features for this particular token?"
For transformers, the layer normalization approach is far more natural. Each token is processed independently, and we want stable statistics regardless of what other tokens or samples happen to be in the batch. This independence also means layer normalization works identically during training and inference, with no need for running statistics or batch size considerations.
To make this concrete, consider a transformer with hidden dimension 512. Batch normalization would compute 512 separate statistics (one per feature dimension), each estimated from however many tokens appear at the same position across the batch. Layer normalization computes 2 statistics (one mean and one variance) from the 512 features of a single token. The batch normalization statistics are estimates that improve with larger batches; the layer normalization statistics are exact computations that require no batch at all. During inference when you process one sequence at a time, layer normalization requires no special handling; batch normalization in inference mode must use stored running statistics from training, which may not accurately represent the test distribution.


Implementing Layer Normalization from Scratch
Now that we understand the formula conceptually, let's translate it into code. Building layer normalization from scratch will solidify our understanding and reveal the implementation details that matter in practice.
The Forward Pass
Our implementation follows the mathematical steps exactly: compute mean, compute variance, normalize, then apply the learnable transformation.
def layer_norm_forward(x, gamma, beta, eps=1e-5):
"""
Layer normalization forward pass.
Args:
x: Input tensor of shape (batch, seq_len, hidden_dim)
gamma: Scale parameters of shape (hidden_dim,)
beta: Shift parameters of shape (hidden_dim,)
eps: Small constant for numerical stability
Returns:
Normalized output with same shape as input
"""
# Step 1: Compute mean across the feature dimension (last axis)
mu = x.mean(dim=-1, keepdim=True)
# Step 2: Compute variance across the feature dimension
var = x.var(dim=-1, keepdim=True, unbiased=False)
# Step 3: Normalize to zero mean and unit variance
x_norm = (x - mu) / torch.sqrt(var + eps)
# Step 4: Apply learnable affine transformation
out = gamma * x_norm + beta
return out, (x, x_norm, mu, var, gamma, eps)The dim=-1 argument tells PyTorch to compute statistics across the last dimension (features), which is exactly what layer normalization requires. The keepdim=True preserves the dimension for broadcasting during subtraction and division. Without keepdim=True, the mean would have shape (batch, seq_len) and could not be subtracted from x with shape (batch, seq_len, hidden_dim) without an explicit unsqueeze operation.
Notice also that we cache the intermediate values (x, x_norm, mu, var, gamma, eps) in the return value. This cache is needed for the backward pass: computing gradients through layer normalization requires the normalized values and statistics from the forward pass. Modern frameworks like PyTorch compute this automatically through autograd, but understanding the forward-backward relationship helps you reason about memory usage in large models.
A Worked Example
Let's trace through layer normalization with concrete numbers to see exactly what happens at each step.
# Create a simple example: 2 samples, 4 tokens each, 8 features per token
batch_size = 2
seq_len = 4
hidden_dim = 8
# Generate input with non-standard distribution (mean ~2, varied std)
x = torch.randn(batch_size, seq_len, hidden_dim) * 3 + 2
# Initialize gamma=1 and beta=0 (identity affine transform)
gamma = torch.ones(hidden_dim)
beta = torch.zeros(hidden_dim)
# Apply layer normalization
output, cache = layer_norm_forward(x, gamma, beta)Input statistics (per token): Shape: torch.Size([2, 4, 8]) First token mean: 1.410, std: 2.718 Second token mean: 0.350, std: 2.670 Output statistics (per token): First token mean: 0.000000, std: 1.069044 Second token mean: 0.000000, std: 1.069044
The input tokens have varying means (around 2-3) and standard deviations (around 2-4). This reflects the non-standard distribution we created. After layer normalization, each token has mean neededly zero and standard deviation neededly one. The tiny deviations from exactly 0 and 1 are floating-point precision artifacts, not algorithmic issues.
Let's visualize this transformation to see exactly how layer normalization reshapes the feature distribution.


The histograms make the transformation crystal clear. Before normalization, the feature values are scattered around a positive mean with varied spread. After normalization, they're centered at zero with approximately unit variance. This happens independently for every token in the sequence.
Notice that each token is normalized independently: the first token's statistics don't affect the second token's normalization. This independence is precisely what makes layer normalization suitable for transformers, where tokens must be processed in parallel and sequences have variable lengths.
Numeric Walk-Through
To build deep intuition, let's work through a minimal example by hand before letting the code do it. Consider a single token with just four features:
Step 1: Compute the mean.
Step 2: Compute the variance.
So the standard deviation is .
Step 3: Normalize each feature. Using for simplicity:
Verify: The mean of is zero. The variance is . The normalization is correct.
Step 4: Apply learned parameters. If and , the output equals the normalized values. If instead and , the output would be , a distribution with mean 1 and standard deviation 2. The network can express this transformation freely through the learned parameters.
This walk-through reveals something important: layer normalization is not a mere computational trick; it is a principled recentering and rescaling operation that the network can tune at training time. The initialization puts every token in a standard reference frame; the learned parameters then shift it to wherever the subsequent layer finds most useful.
# Verify the hand-calculated example
x_manual = torch.tensor([[3.0, 1.0, 5.0, 7.0]]) # shape (1, 4)
gamma_manual = torch.ones(4)
beta_manual = torch.zeros(4)
mu_manual = x_manual.mean(dim=-1, keepdim=True)
var_manual = x_manual.var(dim=-1, keepdim=True, unbiased=False)
x_norm_manual = (x_manual - mu_manual) / torch.sqrt(var_manual + 1e-5)Manual walk-through verification: Input: [3.0, 1.0, 5.0, 7.0] Mean: 4.0000 Variance: 5.0000 Std dev: 2.2361 Normalized: [-0.4472, -1.3416, 0.4472, 1.3416] Normalized mean: 0.00000000 Normalized std: 0.999999
The computed values match the hand calculation exactly (up to floating-point precision from the epsilon term). Each normalized value has the expected sign and magnitude, and the normalized vector has mean zero and variance one.
The Role of Learnable Parameters
After normalizing activations to zero mean and unit variance, we've effectively forced all features into a standardized distribution. But what if the network needs some features to have a larger spread, or to be centered around a non-zero value? The learnable parameters (scale) and (shift) solve this problem.
For each feature dimension , the final output is:
where:
- : the final output for the -th feature
- : the normalized value (zero mean, unit variance)
- : the learned scale for feature , which controls the spread of values
- : the learned shift for feature , which controls the center of the distribution
If the network learns and , it can completely undo the normalization and recover the original distribution. This means layer normalization can never hurt the network's representational capacity; in the worst case, it learns to bypass itself entirely. In practice, the network finds an intermediate setting that benefits from stable optimization while still representing the patterns it needs.
This "can always undo itself" property is worth examining more carefully. Suppose a particular layer in the network has learned to encode certain semantic properties through the absolute scale of its activations. Without layer normalization, those large activations pass through unchanged. With layer normalization, they get normalized down to unit variance. But then can scale them back up. The network does not lose information; it just processes it differently. The training process can discover, for instance, that one feature dimension should carry a signal with large variance (indicating strong activation) while another should operate at a smaller scale. The vector encodes exactly this per-feature preference.
The vector serves a complementary role. Downstream computations like softmax, activation functions, and biased projections all shift values around. If a particular layer normalization is followed by a ReLU, having outputs that start from a non-zero center may help more of the features remain active. The network can learn this through , shifting the normalized distribution to a mean that works well with the downstream computation.
# Demonstrate how gamma and beta affect the output
gamma_custom = torch.tensor([2.0, 0.5, 1.5, 0.8, 1.0, 3.0, 0.3, 2.5])
beta_custom = torch.tensor([1.0, -1.0, 0.0, 2.0, -0.5, 0.5, 0.0, -2.0])
output_custom, _ = layer_norm_forward(x, gamma_custom, beta_custom)Feature-wise comparison (first token): Feature gamma beta Output mean ---------------------------------------- 0 2.0 1.0 1.663 1 0.5 -1.0 -0.945 2 1.5 0.0 0.672 3 0.8 2.0 2.192 4 1.0 -0.5 -0.485 5 3.0 0.5 -0.862 6 0.3 0.0 -0.143 7 2.5 -2.0 -2.539
The output distribution for each feature is controlled by its corresponding and values. Features with larger values have wider distributions, while shifts the center. This per-feature control lets different dimensions of the representation operate at different scales, which matters in transformers because attention heads and feature dimensions may need different dynamic ranges.

With identity parameters (gamma=1, beta=0), all feature means are near zero as expected. With custom parameters, each feature shifts to its corresponding beta value. This shows how the learnable parameters give each dimension independent control over its output distribution.
PyTorch's LayerNorm
PyTorch provides a built-in nn.LayerNorm that handles all these details efficiently. Let's verify our implementation matches PyTorch's behavior.
import torch.nn as nn
# Create PyTorch LayerNorm
pytorch_ln = nn.LayerNorm(hidden_dim, eps=1e-5)
# Initialize with same parameters
with torch.no_grad():
pytorch_ln.weight.fill_(1.0) # gamma
pytorch_ln.bias.fill_(0.0) # beta
# Compare outputs
pytorch_output = pytorch_ln(x)
our_output, _ = layer_norm_forward(x, gamma, beta)Maximum difference between implementations: 1.19e-07 Outputs match: True
The outputs match within floating-point precision, confirming our implementation is correct.
PyTorch exposes the learnable parameters through the weight and bias attributes. weight corresponds to and bias corresponds to . This naming can be confusing because "weight" typically refers to the parameters of a linear transformation, but here it refers to the per-feature scale. The reason for this naming is historical: when PyTorch first implemented layer normalization, it followed the convention established by batch normalization, where the parameters are called weight and bias to indicate that they scale and shift the normalized values.
When you inspect a trained transformer model and look at the layer normalization parameters, you'll often find that the weight parameters are very close to 1 and the bias parameters are very close to 0, especially in the earlier layers. This suggests that those layers don't need to deviate much from the standard normalized distribution. In later layers, or in attention-specific normalizations, the parameters may drift further from their initialization, indicating that those layers benefit from a different distribution.
Layer Normalization in Transformers
In transformer architectures, layer normalization appears in two key locations: after the attention mechanism and after the feed-forward network. The original transformer used "post-norm" placement, where normalization comes after the residual connection:
class PostNormTransformerBlock(nn.Module):
"""Transformer block with post-normalization (original architecture)."""
def __init__(self, d_model, n_heads, d_ff, dropout=0.1):
super().__init__()
self.attention = nn.MultiheadAttention(
d_model, n_heads, batch_first=True
)
self.ff = nn.Sequential(
nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model)
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
# Attention with residual, then normalize
attn_out, _ = self.attention(x, x, x)
x = self.norm1(x + self.dropout(attn_out))
# FFN with residual, then normalize
ff_out = self.ff(x)
x = self.norm2(x + self.dropout(ff_out))
return xModern architectures like GPT-2, GPT-3, and LLaMA use "pre-norm" placement, where normalization comes before the sublayer:
class PreNormTransformerBlock(nn.Module):
"""Transformer block with pre-normalization (modern architecture)."""
def __init__(self, d_model, n_heads, d_ff, dropout=0.1):
super().__init__()
self.attention = nn.MultiheadAttention(
d_model, n_heads, batch_first=True
)
self.ff = nn.Sequential(
nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model)
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
# Normalize, then attention with residual
normed = self.norm1(x)
attn_out, _ = self.attention(normed, normed, normed)
x = x + self.dropout(attn_out)
# Normalize, then FFN with residual
ff_out = self.ff(self.norm2(x))
x = x + self.dropout(ff_out)
return xPre-Norm vs. Post-Norm: Why the Placement Matters
The difference in placement has clear implications for gradient flow and training dynamics. Understanding why requires tracing the residual stream through a transformer.
In the post-norm design, the residual connection adds the sublayer output to its input, and the sum is then normalized. The residual stream at any depth carries unnormalized activations. When the gradients flow backward, they pass through the normalization operation and then split between the residual path and the sublayer path. The residual path provides a direct gradient highway from the loss to the input embeddings. However, the normalization after each residual addition can interfere with this highway: the gradient must pass through the normalization Jacobian, which can compress or distort the signal.
In the pre-norm design, normalization happens before the sublayer, and the residual connection bypasses the normalization entirely. The residual stream is the raw, unnormalized activations accumulating across layers. The gradient flowing through the residual path does not pass through any normalization operations, giving it a clean highway all the way back to the first layer. This is particularly beneficial in very deep networks, where the clean gradient signal allows earlier layers to receive meaningful updates even during the early stages of training.
Think of the residual stream as a highway and the sublayers as on-ramps. In post-norm, there is a tollbooth (normalization) on the main highway after each on-ramp merge. The tollbooth ensures the merged traffic stays within bounds, but it also introduces bottlenecks and can disrupt the flow of fast-moving traffic (gradients). In pre-norm, the tollbooth sits on the on-ramp before the merge, so the main highway remains unimpeded. The traffic entering from the on-ramp gets regulated, but the existing highway flow is never disrupted.
This explains a well-documented empirical observation: post-norm transformers typically require careful learning rate warmup schedules, while pre-norm transformers can be trained with constant or cosine-decayed learning rates from the start. The warmup in post-norm is essentially giving the network time to settle into a state where the normalized residual stream is stable enough for larger gradient updates. Pre-norm avoids this need by maintaining a stable gradient highway from the beginning of training.
The tradeoff is representation quality at the final layer. In post-norm, the output of the last layer is normalized, putting all outputs in a consistent scale before they reach the final projection layer. In pre-norm, the output is the raw residual accumulation, which may have a non-standard scale. Practitioners have found various ways to address this, such as adding a final layer normalization before the output projection, a practice adopted by GPT-2 and later models.

Gradient Flow Through Layer Normalization
Understanding how gradients flow through layer normalization is essential for debugging training issues and understanding why normalization stabilizes training.
The backward pass through layer normalization is more complex than a simple element-wise operation because each output depends on all inputs through the mean and variance computation. When we change a single input , it affects its own normalized output, the mean , and the variance , which in turn affects every output element. This coupling makes the gradient computation more involved.
Given the loss and the upstream gradient (the gradient flowing back from later layers), we need to compute three gradients: for backpropagation, and and for updating the learnable parameters.
Gradients for Learnable Parameters
The output of layer normalization is , where is the normalized input. Since this is a simple affine transformation, the gradients follow directly from the chain rule.
For the scale parameter :
where:
- : the gradient of the loss with respect to the -th scale parameter
- : an index over all tokens (across batch and sequence dimensions)
- : the upstream gradient for the -th feature of the -th token
- : the normalized value of the -th feature for the -th token
Intuitively, this sums up how much each token's normalized value contributed to the loss through this scale parameter.
For the shift parameter :
where:
- : the gradient of the loss with respect to the -th shift parameter
This is simply the sum of upstream gradients, since adds directly to the output.
Gradient for Input
The gradient with respect to input is more involved because depends on in three ways: directly, through the mean , and through the variance . Applying the chain rule carefully yields:
where:
- : the gradient of the loss with respect to the -th input element
- : the learned scale parameter for the -th feature
- : the standard deviation (with epsilon for stability)
- : the feature dimension (number of elements in the input vector)
- : the normalized input, equal to
Let's break down the three terms inside the parentheses:
-
Direct contribution : The gradient that would flow if normalization were a simple scaling operation.
-
Mean correction : Accounts for how changing affects , which affects all outputs. This term subtracts the average gradient, centering the gradient distribution.
-
Variance correction : Accounts for how changing affects , which scales all outputs. This term is proportional to the normalized value , meaning inputs far from the mean get larger corrections.
This formula reveals something important: the gradient for each input element depends on the gradients of all other elements through the mean and variance terms. This coupling helps distribute gradient information across features, which can improve training stability.
The mean correction term subtracts a constant from every gradient. The effect is similar to centering the gradients, preventing them from all shifting in the same direction. If the loss pushes all features to increase, the mean correction tempers this. This ensures the relative ordering of features matters more than their absolute magnitude. This is analogous to how the forward pass removes the mean from the input; the backward pass removes the mean from the gradient.
The variance correction term is more subtle. It is proportional to , the normalized value at position . Features that are far from the normalized mean (large ) receive a larger variance correction. This makes sense: those features have a larger influence on the variance computation, so their gradient needs a larger adjustment to account for that influence. In practice, this term prevents individual features from receiving gradients that are systematically larger simply because they happen to have larger normalized values.
def layer_norm_backward(dout, cache):
"""
Layer normalization backward pass.
Args:
dout: Upstream gradient of shape (batch, seq_len, hidden_dim)
cache: Values from forward pass
Returns:
dx: Gradient with respect to input
dgamma: Gradient with respect to scale
dbeta: Gradient with respect to shift
"""
x, x_norm, mu, var, gamma, eps = cache
d = x.shape[-1]
# Gradients for learnable parameters (sum over batch and sequence)
dgamma = (dout * x_norm).sum(dim=(0, 1))
dbeta = dout.sum(dim=(0, 1))
# Gradient for normalized input
dx_norm = dout * gamma
# Gradient for input (the complex part)
std = torch.sqrt(var + eps)
# Three terms in the gradient
term1 = dx_norm / std
term2 = dx_norm.mean(dim=-1, keepdim=True) / std
term3 = (dx_norm * x_norm).mean(dim=-1, keepdim=True) * x_norm / std
dx = term1 - term2 - term3
return dx, dgamma, dbeta# Verify against PyTorch autograd
x_test = torch.randn(2, 4, 8, requires_grad=True)
gamma_test = torch.ones(8, requires_grad=True)
beta_test = torch.zeros(8, requires_grad=True)
# Forward pass with our implementation
output, cache = layer_norm_forward(x_test, gamma_test, beta_test)
# Create fake upstream gradient
dout = torch.randn_like(output)
# Our backward pass
dx_ours, dgamma_ours, dbeta_ours = layer_norm_backward(dout, cache)
# PyTorch autograd backward
output.backward(dout)Gradient comparison with PyTorch autograd: dx max difference: 7.15e-07 dgamma max difference: 0.00e+00 dbeta max difference: 0.00e+00
The differences are on the order of or smaller, well within floating-point precision. This confirms our manual backward pass implementation correctly computes the gradients that PyTorch's autograd produces automatically.
Visualizing Layer Normalization's Effect
Let's visualize how layer normalization transforms the activation distribution during a forward pass through multiple transformer blocks.
class StackedTransformerBlocks(nn.Module):
"""Stack of transformer blocks for visualization."""
def __init__(self, d_model, n_heads, d_ff, n_layers, use_layernorm=True):
super().__init__()
self.use_layernorm = use_layernorm
self.layers = nn.ModuleList(
[
PreNormTransformerBlock(d_model, n_heads, d_ff)
for _ in range(n_layers)
]
)
if not use_layernorm:
# Replace LayerNorm with identity
for layer in self.layers:
layer.norm1 = nn.Identity()
layer.norm2 = nn.Identity()
def forward_with_activations(self, x):
"""Return activations after each layer."""
activations = [x.detach().clone()]
for layer in self.layers:
x = layer(x)
activations.append(x.detach().clone())
return activations
# Create models with and without layer normalization
d_model, n_heads, d_ff, n_layers = 64, 4, 256, 10
model_with_ln = StackedTransformerBlocks(
d_model, n_heads, d_ff, n_layers, use_layernorm=True
)
model_without_ln = StackedTransformerBlocks(
d_model, n_heads, d_ff, n_layers, use_layernorm=False
)
# Forward pass with larger initial variance to amplify divergence without LayerNorm
x = torch.randn(4, 16, d_model) * 2.0
acts_with_ln = model_with_ln.forward_with_activations(x)
acts_without_ln = model_without_ln.forward_with_activations(x)

With layer normalization, activations maintain stable statistics throughout the network. Without it, activations can drift, though the effect depends heavily on initialization. During training, weight updates can shift activation statistics dramatically, so this stability helps keep them anchored.
The stability plots also reveal an important practical consequence of pre-norm design: the residual accumulation means the standard deviation at the output of the model can grow slightly over layers, even with layer normalization in place. This is because each layer adds its output to the residual stream without normalizing the stream itself; normalization only happens on the input to each sublayer. Well-designed pre-norm models account for this by including a final layer normalization before the output head. This ensures the last representation is in a well-defined scale before the projection to vocabulary logits.
Epsilon: A Small but Critical Detail
The epsilon parameter () appears in the denominator of the normalization formula:
where:
- : the normalized value for the -th feature
- : the original input value
- : the mean across all features
- : the variance across all features
- : a small constant added to prevent division by zero
The purpose of is to ensure numerical stability. If all input values are identical (or nearly so), the variance approaches zero. Without , we would divide by zero, producing infinity or NaN. Adding a small positive constant like ensures the denominator is always positive.
The choice of epsilon can affect numerical stability, especially with mixed-precision (FP16) training where very small values may underflow.
In mixed-precision training, activations are stored as 16-bit floating-point numbers, which have a minimum representable positive value of approximately . If the variance of a representation is smaller than this threshold, the FP16 representation will round it to zero, causing the same division-by-zero problem that epsilon was designed to prevent. For FP16 training, it is common to use a larger epsilon (around or ) to ensure stability even when the FP16 representation truncates small variance values. This concern is practical: several practitioners have traced training instabilities in large language models to epsilon values that were appropriate for FP32 but insufficient for FP16.
# Demonstrate epsilon's role with near-constant input
near_constant = torch.ones(1, 4, 8) * 5.0
near_constant[0, 0, 0] = 5.001 # Tiny variation
def test_epsilon(x, eps):
"""Test layer normalization with different epsilon values."""
mu = x.mean(dim=-1, keepdim=True)
var = x.var(dim=-1, keepdim=True, unbiased=False)
try:
x_norm = (x - mu) / torch.sqrt(var + eps)
return x_norm.std().item(), "OK"
except Exception as e:
return float("nan"), str(e)
epsilons = [0, 1e-12, 1e-8, 1e-5, 1e-3]Effect of epsilon on near-constant input: Input variance: 1.25e-07 Epsilon Output std Status ---------------------------------------- 0e+00 NaN/Inf Division issue 1e-12 0.507998 OK 1e-08 0.486255 OK 1e-05 0.052836 OK 1e-03 0.005312 OK
The input variance is extremely small (around ), which means we're dividing by a very small number. With epsilon = 0, the output standard deviation explodes because we're essentially dividing by nearly zero. As epsilon increases, the output becomes more stable. The standard choice of strikes a balance: it's large enough to prevent numerical issues but small enough not to distort the normalization when variance is reasonably sized. A reasonable epsilon value (typically to ) provides a safety net without affecting normal computations.
Layer Normalization with Different Normalized Shapes
PyTorch's nn.LayerNorm accepts a normalized_shape parameter that controls which dimensions are normalized. For transformers, we typically normalize over the feature dimension only:
# Different normalized shapes
x = torch.randn(2, 4, 8) # (batch, seq_len, features)
# Normalize over features only (most common for transformers)
ln_features = nn.LayerNorm(8)
# Normalize over sequence and features
ln_seq_features = nn.LayerNorm([4, 8])
# Normalize over entire sample (batch, sequence, features)
ln_all = nn.LayerNorm([4, 8])
out_features = ln_features(x)
out_seq_features = ln_seq_features(x)LayerNorm with normalized_shape=(8,) - normalize over features: Output shape: torch.Size([2, 4, 8]) Each token normalized independently Token 0,0 mean: 0.000000 Token 0,1 mean: -0.000000 LayerNorm with normalized_shape=(4, 8) - normalize over seq+features: Output shape: torch.Size([2, 4, 8]) Each sample normalized as a whole Sample 0 mean: 0.000000
Notice the difference: with normalized_shape=(8,), each individual token has zero mean (Token 0,0 and Token 0,1 both have mean approximately 0). With normalized_shape=(4, 8), the entire sample is normalized together, so individual tokens may have non-zero means but the sample as a whole has zero mean.
Normalizing over features only is the standard choice for transformers because it treats each token independently, matching the autoregressive nature of language models and allowing the model to process variable-length sequences.
The choice of normalized_shape also has practical implications during inference. When you normalize over the feature dimension only, you can process a single token at a time in autoregressive generation, and layer normalization works exactly as during training. If you were to normalize over the sequence dimension, you would need to know the full sequence before normalizing any token, which breaks the causal autoregressive property required for efficient generation.
In Practice: Using Layer Normalization Effectively
Understanding layer normalization in theory is one thing; applying it correctly in real models requires a few additional considerations that are easy to overlook.
Parameter initialization and scale. The default initialization of and is appropriate in most cases, but some practitioners have found that initializing to a smaller value (like or even ) at the start of training can help with particularly deep or wide models. The motivation is that smaller values reduce the signal amplitude in early layers, giving the optimizer more room to maneuver before gradients become large. The T5 paper and some other work explore this idea in the context of large-scale pretraining.
Interaction with weight decay. The and parameters of layer normalization are typically excluded from weight decay regularization. This is because weight decay penalizes large parameter values, and penalizing away from zero would interfere with the normalization's ability to recover the original distribution when needed. Most modern deep learning frameworks and training recipes exclude all normalization parameters from weight decay by default, but it's worth verifying this in your own training setup.
Combining with other normalization. Some architectures combine layer normalization with other normalization techniques. For instance, Batch Normalization followed by Layer Normalization is occasionally used in hybrid architectures that process both image patches and text tokens. When combining normalizations, be careful about which dimension each normalization operates on: if two normalizations operate on the same dimension, the second one may undo some of the statistical correction applied by the first.
Memory considerations. Layer normalization requires storing the mean and variance for each token during the forward pass, so they can be used in the backward pass. For a model processing sequences of length with hidden dimension and layers, this adds stored scalars per sample. In practice, this memory cost is negligible compared to the activations themselves, but it is worth keeping in mind when debugging out-of-memory errors in very long sequence models.
Gradient accumulation and normalization. When using gradient accumulation (accumulating gradients over multiple mini-batches before applying an update), layer normalization behaves correctly because it computes statistics independently for each sample. In contrast, batch normalization would see artificially small batches at each forward pass, producing noisier statistics. This is one reason why large language model training almost exclusively uses layer normalization: it remains well-behaved under the gradient accumulation strategies necessary for training on long sequences.

Limitations and Impact
Layer normalization has become ubiquitous in transformer architectures, but it's not without drawbacks. Understanding its limitations helps you apply it wisely and know when alternatives like RMSNorm might be preferable.
The primary computational overhead comes from computing statistics for every token at every layer. For a model with hidden dimension , each layer normalization requires computing a mean (sum of elements) and variance (sum of squared differences), then normalizing all elements. While these operations are memory-bandwidth bound rather than compute-bound on modern GPUs, they still add up in models with hundreds of layers. Profiling studies of large transformer models have found that layer normalization can account for 5-10% of total training time in some configurations. This motivated the development of RMSNorm (Root Mean Square Layer Normalization), which removes the mean-centering step and computes only the root mean square of the features, reducing the operation to a single pass through the data instead of two.
The learned and parameters add parameters per layer normalization, which is negligible compared to attention and FFN parameters but contributes to model complexity. More importantly, these parameters can be a source of numerical issues when they grow very large or approach zero, requiring careful initialization and sometimes explicit constraints. In very large models trained for many steps, some values have been observed to grow to extremely large magnitudes, effectively disabling the normalization for those feature dimensions. Monitoring the distribution of norms during training can serve as a useful diagnostic for training stability.
Layer normalization also introduces a subtle form of coupling between features that can affect interpretability. Because each feature is normalized relative to the others, the absolute activation value of any single feature becomes less meaningful. This makes it harder to interpret individual neurons or feature dimensions in isolation. If you're performing mechanistic interpretability studies on a transformer, you need to account for the fact that what you see in the residual stream is the raw, unnormalized representation, while what the attention and FFN layers process is the normalized version. The two can differ significantly when the learned parameters are large.
A less-discussed limitation is the interaction between layer normalization and the residual stream in very deep pre-norm transformers. Since the residual path bypasses all layer normalizations, the residual stream can accumulate activations that grow with depth. At extreme depths (hundreds of layers), this growth can cause the unnormalized residual stream to have a much larger magnitude than the normalized inputs to each sublayer, leading to a situation where the residual additions dominate the sublayer outputs. Some research has explored remedies like scaled initialization (where sublayer output weights are initialized to smaller values to counteract this growth) or architectural changes like sandwich-normalization (normalizing both the input and the output of each sublayer).
Despite these limitations, layer normalization's impact on transformer training stability cannot be overstated. Before normalization techniques were widely adopted, training deep networks required careful learning rate tuning, extensive warmup periods, and often failed entirely for very deep models. Layer normalization enables stable training with higher learning rates, reduces sensitivity to initialization, and allows models to scale to unprecedented depths. The original transformer used layer normalization, and every major language model since has relied on some form of normalization to train successfully.
The success of layer normalization has also spurred research into alternatives. RMSNorm, which we'll cover in the next chapter, removes the mean-centering step to improve computational efficiency while maintaining most of the stability benefits. Other variants like Power Normalization and Fixup Normalization explore different approaches to the same stability problem, though none has achieved the widespread adoption of layer normalization in language models.
Key Parameters
When using nn.LayerNorm in PyTorch, understanding the key parameters helps you configure it correctly for your architecture:
-
normalized_shape: The shape of the input over which to normalize. For transformers, this is typically the hidden dimension
d_model(e.g., 768, 1024). You can also pass a list like[seq_len, d_model]to normalize over multiple dimensions, though normalizing over features only is the standard choice. -
eps: The epsilon value added to the denominator for numerical stability. Default is
1e-5, which works well for most cases. For mixed-precision (FP16) training, you may need a larger value like1e-4to avoid underflow issues when variance is very small. -
elementwise_affine: Whether to include learnable and parameters. Default is
True. Setting toFalseremoves the learnable parameters, reducing model size slightly but limiting the network's ability to learn optimal feature scales. Some research has found that removing the affine transform causes minimal performance degradation in very large models, suggesting that the normalization itself, not the learnable parameters, is primarily responsible for the stability benefits.
Summary
Layer normalization is a standard component of transformer architectures because it stabilizes training in deep models. Unlike batch normalization, which computes statistics across the batch dimension, layer normalization operates on each sample independently. This makes it well-suited for variable-length sequences and small batch sizes.
The core operation normalizes each token's representation to zero mean and unit variance, then applies learned scale () and shift () parameters to recover representational flexibility. This simple transformation stabilizes activations throughout the network, prevents gradient issues during training, and reduces sensitivity to initialization. The backward pass through layer normalization distributes gradient information across all features through the mean and variance coupling, which helps prevent individual features from dominating the gradient signal.
The placement of layer normalization within transformer blocks has evolved from post-norm (original transformer) to pre-norm (modern language models). Pre-norm preserves a clean gradient highway through the residual stream, enabling training without extensive warmup schedules and making it the default choice for large-scale pretraining. The tradeoff is that the residual stream can grow with depth, a problem addressed in practice by a final layer normalization before the output head.
Key takeaways:
- Feature-wise normalization: Layer normalization computes mean and variance across the feature dimension, treating each token independently
- Learnable parameters: and allow the network to undo normalization when beneficial, preserving representational capacity
- Placement matters: Pre-norm (normalize before sublayer) has become the modern standard, improving gradient flow in deep networks
- Epsilon for stability: A small constant prevents division by zero with near-constant inputs; larger values ( or more) may be needed for FP16 training
- No batch dependency: Works with any batch size, including single samples during inference
- Gradient coupling: The shared mean and variance computation means every feature's gradient depends on all other features, distributing the learning signal across the representation
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about layer normalization in transformers.
Layer Normalization
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!