Layer Normalization: Stabilizing Transformer Training

Michael BrenndoerferUpdated June 11, 202552 min read

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.

Historical Context: The Road to Stable Deep Learning

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.

In[3]:
Code
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 shift
Out[4]:
Console
Activation 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 x=[x1,x2,…,xd]\mathbf{x} = [x_1, x_2, \ldots, x_d] representing one token, we calculate its mean:

μ=1d∑i=1dxi\mu = \frac{1}{d} \sum_{i=1}^{d} x_i

where:

  • μ\mu: the arithmetic mean of all dd features in this token's representation
  • dd: the hidden dimension (e.g., 768 for BERT-base, 4096 for LLaMA-7B)
  • xix_i: the value of the ii-th feature

This tells us the "center of mass" of the representation. If μ=2.5\mu = 2.5, the features are shifted toward positive values; if μ=−1.3\mu = -1.3, 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:

σ2=1d∑i=1d(xi−μ)2\sigma^2 = \frac{1}{d} \sum_{i=1}^{d} (x_i - \mu)^2

where:

  • σ2\sigma^2: the variance, measuring how much the features deviate from their mean
  • (xi−μ)2(x_i - \mu)^2: the squared deviation of each feature from the mean

Squaring ensures that positive and negative deviations don't cancel out. If σ2\sigma^2 is large, the features are spread out; if small, they're tightly clustered. We'll use the standard deviation σ=σ2\sigma = \sqrt{\sigma^2} to rescale the values to unit variance.

Notice that we use the population variance (dividing by dd) rather than the sample variance (dividing by d−1d - 1). 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 dd values in front of us. Using dd rather than d−1d - 1 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:

x^i=xi−μσ2+ϵ\hat{x}_i = \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}}

where:

  • x^i\hat{x}_i: the normalized value of the ii-th feature
  • ϵ\epsilon: a tiny constant (typically 10−510^{-5}) 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 x^i\hat{x}_i has zero mean and approximately unit variance across the dd features.

The subtraction by μ\mu ensures the output is centered at zero, removing any bias in the representation. The division by σ2+ϵ\sqrt{\sigma^2 + \epsilon} 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:

yi=γi⋅x^i+βiy_i = \gamma_i \cdot \hat{x}_i + \beta_i

where:

  • yiy_i: the final output for the ii-th feature
  • γi\gamma_i: a learned scale parameter for feature ii (initialized to 1)
  • βi\beta_i: a learned shift parameter for feature ii (initialized to 0)

These parameters are learned during training, just like weights and biases. If the network discovers that feature ii should have mean 3.5 and standard deviation 2.0, it can learn γi=2.0\gamma_i = 2.0 and βi=3.5\beta_i = 3.5 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 γ=1\gamma = 1 and β=0\beta = 0 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, γ\gamma and β\beta 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 x\mathbf{x} of dimension dd into:

LayerNorm(x)=γ⊙x−μσ2+ϵ+β\text{LayerNorm}(\mathbf{x}) = \gamma \odot \frac{\mathbf{x} - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta

where:

  • x=[x1,x2,…,xd]\mathbf{x} = [x_1, x_2, \ldots, x_d]: the input vector representing one token's hidden state
  • μ=1d∑i=1dxi\mu = \frac{1}{d} \sum_{i=1}^{d} x_i: the mean across all features
  • σ2=1d∑i=1d(xi−μ)2\sigma^2 = \frac{1}{d} \sum_{i=1}^{d} (x_i - \mu)^2: the variance across all features
  • ϵ\epsilon: a small stability constant (typically 10−510^{-5} or 10−610^{-6})
  • γ=[γ1,γ2,…,γd]\gamma = [\gamma_1, \gamma_2, \ldots, \gamma_d]: learned scale parameters, initialized to ones
  • β=[β1,β2,…,βd]\beta = [\beta_1, \beta_2, \ldots, \beta_d]: learned shift parameters, initialized to zeros
  • ⊙\odot: 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 ii 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.

Out[5]:
Visualization
Diagram showing batch normalization highlighting a column across batch samples.
Batch normalization computes statistics across samples (blue column). For each feature, the mean and variance are estimated over all batch samples. This requires a meaningful batch size and produces statistics that change across sequence positions.
Diagram showing layer normalization highlighting a row across features.
Layer normalization computes statistics across features (orange row). For each token, the mean and variance are computed over all feature dimensions. Statistics are exact and independent of batch size or other samples.

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.

In[6]:
Code
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.

In[7]:
Code
# 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)
Out[8]:
Console
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.

Out[9]:
Visualization
Histogram showing feature values with positive mean around 2-3.
Feature distribution before layer normalization. The 8 features of a single token show varied values with a non-zero mean (dashed line) and non-unit variance. Each feature contributes to the overall distribution.
Histogram showing feature values centered at zero with unit spread.
Feature distribution after layer normalization. The same token's features are now centered at zero with unit variance. The transformation standardizes each token independently.

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:

x=[3.0,  1.0,  5.0,  7.0]\mathbf{x} = [3.0, \; 1.0, \; 5.0, \; 7.0]

Step 1: Compute the mean.

μ=3.0+1.0+5.0+7.04=16.04=4.0\mu = \frac{3.0 + 1.0 + 5.0 + 7.0}{4} = \frac{16.0}{4} = 4.0

Step 2: Compute the variance.

σ2=(3.0−4.0)2+(1.0−4.0)2+(5.0−4.0)2+(7.0−4.0)24\sigma^2 = \frac{(3.0 - 4.0)^2 + (1.0 - 4.0)^2 + (5.0 - 4.0)^2 + (7.0 - 4.0)^2}{4} σ2=1.0+9.0+1.0+9.04=20.04=5.0\sigma^2 = \frac{1.0 + 9.0 + 1.0 + 9.0}{4} = \frac{20.0}{4} = 5.0

So the standard deviation is σ=5.0≈2.236\sigma = \sqrt{5.0} \approx 2.236.

Step 3: Normalize each feature. Using ϵ=10−5≈0\epsilon = 10^{-5} \approx 0 for simplicity:

x^1=3.0−4.02.236≈−0.447\hat{x}_1 = \frac{3.0 - 4.0}{2.236} \approx -0.447 x^2=1.0−4.02.236≈−1.342\hat{x}_2 = \frac{1.0 - 4.0}{2.236} \approx -1.342 x^3=5.0−4.02.236≈+0.447\hat{x}_3 = \frac{5.0 - 4.0}{2.236} \approx +0.447 x^4=7.0−4.02.236≈+1.342\hat{x}_4 = \frac{7.0 - 4.0}{2.236} \approx +1.342

Verify: The mean of [−0.447,−1.342,0.447,1.342][-0.447, -1.342, 0.447, 1.342] is zero. The variance is 0.200+1.800+0.200+1.8004=1.0\frac{0.200 + 1.800 + 0.200 + 1.800}{4} = 1.0. The normalization is correct.

Step 4: Apply learned parameters. If γ=[1,1,1,1]\gamma = [1, 1, 1, 1] and β=[0,0,0,0]\beta = [0, 0, 0, 0], the output equals the normalized values. If instead γ=[2,2,2,2]\gamma = [2, 2, 2, 2] and β=[1,1,1,1]\beta = [1, 1, 1, 1], the output would be [−0.894+1,−2.684+1,0.894+1,2.684+1]=[0.106,−1.684,1.894,3.684][-0.894 + 1, -2.684 + 1, 0.894 + 1, 2.684 + 1] = [0.106, -1.684, 1.894, 3.684], 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.

In[10]:
Code
# 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)
Out[11]:
Console
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 γ\gamma (scale) and β\beta (shift) solve this problem.

For each feature dimension ii, the final output is:

yi=γi⋅x^i+βiy_i = \gamma_i \cdot \hat{x}_i + \beta_i

where:

  • yiy_i: the final output for the ii-th feature
  • x^i\hat{x}_i: the normalized value (zero mean, unit variance)
  • γi\gamma_i: the learned scale for feature ii, which controls the spread of values
  • βi\beta_i: the learned shift for feature ii, which controls the center of the distribution

If the network learns γi=σoriginal\gamma_i = \sigma_{\text{original}} and βi=μoriginal\beta_i = \mu_{\text{original}}, 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 γ\gamma 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 γ\gamma vector encodes exactly this per-feature preference.

The β\beta 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 β\beta, shifting the normalized distribution to a mean that works well with the downstream computation.

In[12]:
Code
# 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)
Out[13]:
Console
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 γ\gamma and β\beta values. Features with larger γ\gamma values have wider distributions, while β\beta 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.

Out[14]:
Visualization
Bar chart comparing feature output means with identity parameters versus custom gamma and beta values.
Effect of learned gamma and beta parameters on output distributions. Each bar shows the mean output value for one feature. With identity parameters (gamma=1, beta=0), all features have mean near zero. With custom parameters, each feature has its own center (controlled by beta) and spread (controlled by gamma). This shows the per-dimension flexibility that layer normalization provides.

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.

In[15]:
Code
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)
Out[16]:
Console
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 γ\gamma and bias corresponds to β\beta. 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:

In[17]:
Code
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 x

Modern architectures like GPT-2, GPT-3, and LLaMA use "pre-norm" placement, where normalization comes before the sublayer:

In[18]:
Code
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 x

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

Out[19]:
Visualization
Diagram comparing pre-norm and post-norm transformer block data flow paths.
Comparison of pre-norm and post-norm transformer block architectures. In post-norm (left), the residual addition happens before normalization, meaning the main signal path passes through the normalization operation. In pre-norm (right), normalization happens before the sublayer, leaving the residual path clean. Pre-norm has become the dominant design in modern large language models due to more stable gradient flow.

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 xix_i, it affects its own normalized output, the mean μ\mu, and the variance σ2\sigma^2, which in turn affects every output element. This coupling makes the gradient computation more involved.

Given the loss LL and the upstream gradient ∂L∂y\frac{\partial L}{\partial y} (the gradient flowing back from later layers), we need to compute three gradients: ∂L∂x\frac{\partial L}{\partial x} for backpropagation, and ∂L∂γ\frac{\partial L}{\partial \gamma} and ∂L∂β\frac{\partial L}{\partial \beta} for updating the learnable parameters.

Gradients for Learnable Parameters

The output of layer normalization is yi=γi⋅x^i+βiy_i = \gamma_i \cdot \hat{x}_i + \beta_i, where x^i\hat{x}_i is the normalized input. Since this is a simple affine transformation, the gradients follow directly from the chain rule.

For the scale parameter γi\gamma_i:

∂L∂γi=∑n∂L∂yn,i⋅x^n,i\frac{\partial L}{\partial \gamma_i} = \sum_{n} \frac{\partial L}{\partial y_{n,i}} \cdot \hat{x}_{n,i}

where:

  • ∂L∂γi\frac{\partial L}{\partial \gamma_i}: the gradient of the loss with respect to the ii-th scale parameter
  • nn: an index over all tokens (across batch and sequence dimensions)
  • ∂L∂yn,i\frac{\partial L}{\partial y_{n,i}}: the upstream gradient for the ii-th feature of the nn-th token
  • x^n,i\hat{x}_{n,i}: the normalized value of the ii-th feature for the nn-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 βi\beta_i:

∂L∂βi=∑n∂L∂yn,i\frac{\partial L}{\partial \beta_i} = \sum_{n} \frac{\partial L}{\partial y_{n,i}}

where:

  • ∂L∂βi\frac{\partial L}{\partial \beta_i}: the gradient of the loss with respect to the ii-th shift parameter

This is simply the sum of upstream gradients, since βi\beta_i adds directly to the output.

Gradient for Input

The gradient with respect to input is more involved because x^i\hat{x}_i depends on xix_i in three ways: directly, through the mean μ\mu, and through the variance σ2\sigma^2. Applying the chain rule carefully yields:

∂L∂xi=γiσ(∂L∂yi−1d∑j=1d∂L∂yj−x^id∑j=1d∂L∂yjx^j)\frac{\partial L}{\partial x_i} = \frac{\gamma_i}{\sigma} \left( \frac{\partial L}{\partial y_i} - \frac{1}{d}\sum_{j=1}^{d}\frac{\partial L}{\partial y_j} - \frac{\hat{x}_i}{d}\sum_{j=1}^{d}\frac{\partial L}{\partial y_j}\hat{x}_j \right)

where:

  • ∂L∂xi\frac{\partial L}{\partial x_i}: the gradient of the loss with respect to the ii-th input element
  • γi\gamma_i: the learned scale parameter for the ii-th feature
  • σ=σ2+ϵ\sigma = \sqrt{\sigma^2 + \epsilon}: the standard deviation (with epsilon for stability)
  • dd: the feature dimension (number of elements in the input vector)
  • x^i\hat{x}_i: the normalized input, equal to (xi−μ)/σ(x_i - \mu) / \sigma

Let's break down the three terms inside the parentheses:

  1. Direct contribution ∂L∂yi\frac{\partial L}{\partial y_i}: The gradient that would flow if normalization were a simple scaling operation.

  2. Mean correction −1d∑j=1d∂L∂yj-\frac{1}{d}\sum_{j=1}^{d}\frac{\partial L}{\partial y_j}: Accounts for how changing xix_i affects μ\mu, which affects all outputs. This term subtracts the average gradient, centering the gradient distribution.

  3. Variance correction −x^id∑j=1d∂L∂yjx^j-\frac{\hat{x}_i}{d}\sum_{j=1}^{d}\frac{\partial L}{\partial y_j}\hat{x}_j: Accounts for how changing xix_i affects σ2\sigma^2, which scales all outputs. This term is proportional to the normalized value x^i\hat{x}_i, 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 x^i\hat{x}_i, the normalized value at position ii. Features that are far from the normalized mean (large ∣x^i∣|\hat{x}_i|) 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.

In[20]:
Code
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
In[21]:
Code
# 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)
Out[22]:
Console
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 10−710^{-7} 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.

In[23]:
Code
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)
Out[24]:
Visualization
Line plot showing stable activation mean near zero and std near one across 10 layers with LayerNorm.
Activation statistics with layer normalization across 10 transformer layers. The mean (blue) stays near zero and the standard deviation (red) remains stable around 1. This shows the consistent normalization that enables reliable training in deep networks.
Line plot showing activation mean and std that drift away from normalized values across 10 layers without LayerNorm.
Activation statistics without layer normalization across 10 transformer layers. Without normalization, activations may drift in mean and variance as they propagate through layers, a behavior that becomes more pronounced with larger networks or more training steps.

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 (ϵ\epsilon) appears in the denominator of the normalization formula:

x^i=xi−μσ2+ϵ\hat{x}_i = \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}}

where:

  • x^i\hat{x}_i: the normalized value for the ii-th feature
  • xix_i: the original input value
  • μ\mu: the mean across all features
  • σ2\sigma^2: the variance across all features
  • ϵ\epsilon: a small constant added to prevent division by zero

The purpose of ϵ\epsilon is to ensure numerical stability. If all input values are identical (or nearly so), the variance σ2\sigma^2 approaches zero. Without ϵ\epsilon, we would divide by zero, producing infinity or NaN. Adding a small positive constant like 10−510^{-5} 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 6×10−86 \times 10^{-8}. 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 10−410^{-4} or 10−310^{-3}) 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.

In[25]:
Code
# 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]
Out[26]:
Console
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 10−810^{-8}), 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 10−510^{-5} 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 10−510^{-5} to 10−610^{-6}) 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:

In[27]:
Code
# 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)
Out[28]:
Console
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 γ=1\gamma = 1 and β=0\beta = 0 is appropriate in most cases, but some practitioners have found that initializing γ\gamma to a smaller value (like 0.10.1 or even 00) at the start of training can help with particularly deep or wide models. The motivation is that smaller γ\gamma 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 γ\gamma and β\beta parameters of layer normalization are typically excluded from weight decay regularization. This is because weight decay penalizes large parameter values, and penalizing γ\gamma 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 LL with hidden dimension dd and NN layers, this adds O(N⋅L⋅2)O(N \cdot L \cdot 2) 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.

Out[29]:
Visualization
Line chart showing training loss curves with and without layer normalization over training steps.
Simulated training loss comparison between a model with layer normalization and a model without normalization. The model with layer normalization converges more smoothly and reaches a lower final loss, illustrating the practical training stability benefits that made layer normalization standard in transformer architectures.

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 dd, each layer normalization requires computing a mean (sum of dd elements) and variance (sum of dd squared differences), then normalizing all dd 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 γ\gamma and β\beta parameters add 2d2d 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 γ\gamma values have been observed to grow to extremely large magnitudes, effectively disabling the normalization for those feature dimensions. Monitoring the distribution of γ\gamma 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 γ\gamma 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 like 1e-4 to avoid underflow issues when variance is very small.

  • elementwise_affine: Whether to include learnable γ\gamma and β\beta parameters. Default is True. Setting to False removes 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 (γ\gamma) and shift (β\beta) 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: γ\gamma and β\beta 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 (10−410^{-4} 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

Question 1 of 80 of 8 completed
What dimension does layer normalization compute statistics across?

Comments

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

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2025layernormalization-2, author = {Michael Brenndoerfer}, title = {Layer Normalization: Stabilizing Transformer Training}, year = {2025}, url = {https://mbrenndoerfer.com/writing/layer-normalization-transformers-implementation}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-09-27} }
APAAcademic
Michael Brenndoerfer (2025). Layer Normalization: Stabilizing Transformer Training. Retrieved from https://mbrenndoerfer.com/writing/layer-normalization-transformers-implementation
MLAAcademic
Michael Brenndoerfer. "Layer Normalization: Stabilizing Transformer Training." 2026. Web. September 27, 2026. <https://mbrenndoerfer.com/writing/layer-normalization-transformers-implementation>.
CHICAGOAcademic
Michael Brenndoerfer. "Layer Normalization: Stabilizing Transformer Training." Accessed September 27, 2026. https://mbrenndoerfer.com/writing/layer-normalization-transformers-implementation.
HARVARDAcademic
Michael Brenndoerfer (2025) 'Layer Normalization: Stabilizing Transformer Training'. Available at: https://mbrenndoerfer.com/writing/layer-normalization-transformers-implementation (Accessed: September 27, 2026).
SimpleBasic
Michael Brenndoerfer (2025). Layer Normalization: Stabilizing Transformer Training. https://mbrenndoerfer.com/writing/layer-normalization-transformers-implementation

About the author

Continue with the full handbook

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

Explore Language AI Handbook
Newsletter

Stay up to date

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

No spam, unsubscribe anytime.

or

Join the community

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