How to Implement Layer Normalization in Transformer Blocks: Pre-Norm Architecture Explained
Layer Normalization stabilizes transformer training by normalizing inputs before attention and MLP sub-layers, using residual connections to preserve gradient flow.
Implementing layer normalization in transformer blocks is essential for training deep language models effectively. In the FareedKhan-dev/train-llm-from-scratch repository, the implementation follows the GPT-style "Pre-Norm" pattern, where normalization occurs before rather than after each sub-layer. This design choice prevents gradient vanishing and allows for deeper network architectures.
Understanding the Pre-Norm Architecture
Modern transformer blocks rely on Pre-Layer Normalization (Pre-Norm) to maintain stable hidden-state dynamics throughout deep stacks. Unlike the original "Post-Norm" design that normalizes after sub-layers, Pre-Norm applies nn.LayerNorm to the input of each sub-layer, then adds the sub-layer output back to the original input via a residual connection.
This ordering ensures that gradients flow directly through the residual pathway without being scaled by the normalization parameters, enabling the training of models with hundreds of layers.
LayerNorm Implementation in transformer_block.py
The core implementation resides in src/models/transformer_block.py, where two LayerNorm modules normalize inputs before the attention and MLP components respectively.
Instantiating LayerNorm Modules
Following the super().__init__() call, the Block class initializes two normalization layers using PyTorch's built-in implementation:
self.ln1 = nn.LayerNorm(n_embed) # Normalization before attention
self.ln2 = nn.LayerNorm(n_embed) # Normalization before MLP
These instances normalize across the embedding dimension (n_embed) independently for each token in the sequence. The normalization computes mean and variance across the last dimension, ensuring consistent activation distributions regardless of batch statistics.
Forward Pass with Residual Connections
In the forward method, the normalized tensors feed into each sub-layer, with residual connections preserving the original signal:
x = x + self.attn(self.ln1(x)) # Attention sub-layer with pre-norm
x = x + self.mlp(self.ln2(x)) # MLP sub-layer with pre-norm
This pattern—LayerNorm → Sub-layer → Residual Addition—appears in both lines 44 and 46 of src/models/transformer_block.py. The self.attn reference points to src/models/attention.py, while self.mlp corresponds to src/models/mlp.py, both receiving normalized inputs that stabilize their internal computations.
Building Layer Normalization from Scratch
While nn.LayerNorm provides optimized performance, implementing the algorithm manually offers educational value and customization flexibility. A custom implementation computes mean and variance across the feature dimension, then applies learnable affine parameters.
Custom LayerNorm Class
This implementation mirrors PyTorch's functionality while exposing the internal statistics:
import torch
import torch.nn as nn
class SimpleLayerNorm(nn.Module):
"""Manual LayerNorm reducing over the last dimension."""
def __init__(self, normalized_shape, eps: float = 1e-5, elementwise_affine: bool = True):
super().__init__()
if isinstance(normalized_shape, int):
normalized_shape = (normalized_shape,)
self.normalized_shape = normalized_shape
self.eps = eps
self.elementwise_affine = elementwise_affine
if self.elementwise_affine:
self.weight = nn.Parameter(torch.ones(normalized_shape))
self.bias = nn.Parameter(torch.zeros(normalized_shape))
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Compute statistics over embedding dimension
mean = x.mean(dim=-1, keepdim=True)
var = x.var(dim=-1, unbiased=False, keepdim=True)
# Normalize with numerical stability
x_norm = (x - mean) / torch.sqrt(var + self.eps)
# Apply learnable affine transform
if self.elementwise_affine:
x_norm = x_norm * self.weight + self.bias
return x_norm
To use this custom implementation within the transformer block, simply replace the standard initialization:
self.ln1 = SimpleLayerNorm(n_embed)
self.ln2 = SimpleLayerNorm(n_embed)
Both approaches produce mathematically identical outputs, but the custom class allows direct access to mean and variance tensors for debugging or architectural research.
Complete Usage Example
The following example demonstrates how to instantiate and run the transformer block as implemented in the repository:
import torch
from src.models.transformer_block import Block
# Configuration
batch_size, seq_len = 2, 5
embed_dim = 32
n_heads = 4
context_len = seq_len
# Create input tensor
x = torch.randn(batch_size, seq_len, embed_dim)
# Initialize block with built-in LayerNorm
block = Block(n_head=n_heads, n_embed=embed_dim, context_length=context_len)
# Forward pass through transformer block
output = block(x)
print(f"Input shape: {x.shape}")
print(f"Output shape: {output.shape}")
For experimentation with the custom LayerNorm implementation, extend the base class:
class CustomBlock(Block):
def __init__(self, n_head, n_embed, context_length):
super().__init__(n_head, n_embed, context_length)
# Override with custom normalization
self.ln1 = SimpleLayerNorm(n_embed)
self.ln2 = SimpleLayerNorm(n_embed)
custom_block = CustomBlock(n_head=n_heads, n_embed=embed_dim, context_length=context_len)
output = custom_block(x)
Summary
- Pre-Norm architecture places LayerNorm before attention and MLP sub-layers in
src/models/transformer_block.py, specifically at lines 28 and 30. - Residual connections immediately follow each sub-layer (lines 44 and 46), adding the normalized sub-layer output back to the input tensor.
- Per-token normalization occurs across the embedding dimension (
n_embed), maintaining independent statistics for each sequence position. - The repository uses
nn.LayerNormfor production efficiency, but the mathematical implementation involves mean/variance computation over the last dimension followed by affine transformation. - Supporting modules in
src/models/attention.pyandsrc/models/mlp.pyreceive pre-normalized inputs, ensuring stable gradient flow throughout deep stacks.
Frequently Asked Questions
What is the difference between Pre-Norm and Post-Norm in transformers?
Pre-Norm applies layer normalization before the attention and feed-forward sub-layers, while Post-Norm applies it after. According to the train-llm-from-scratch source code, Pre-Norm prevents gradient vanishing in deep models by allowing gradients to flow directly through residual connections without passing through normalization layers. Post-Norm can cause training instability in very deep networks because gradients must propagate through the normalization parameters.
Why does layer normalization use the embedding dimension rather than the batch dimension?
LayerNorm normalizes across the embedding dimension (feature dimension) for each token independently, unlike BatchNorm which normalizes across the batch dimension. In transformer blocks, this per-token normalization ensures consistent activation distributions regardless of batch size or sequence composition. The implementation in src/models/transformer_block.py specifies nn.LayerNorm(n_embed), confirming that normalization occurs across the hidden size dimension for every position in the sequence.
Can I remove layer normalization from transformer blocks to speed up training?
Removing layer normalization is not recommended for deep transformer architectures. The normalization layers in src/models/transformer_block.py stabilize activations and enable the residual connections to function effectively. Without LayerNorm,深层网络的梯度会迅速消失或爆炸,导致训练无法收敛。While shallow networks might train without normalization, any model with more than a few layers requires LayerNorm to maintain trainable gradients.
How do the learnable parameters in LayerNorm affect model capacity?
The weight (gamma) and bias (beta) parameters in nn.LayerNorm provide an affine transformation after normalization, allowing the network to learn the optimal mean and variance for each layer. In the custom SimpleLayerNorm implementation, these are initialized to ones and zeros respectively, but they update during backpropagation to rescale normalized outputs. These parameters add minimal computational overhead (only two vectors of size n_embed per LayerNorm layer) but significantly increase model flexibility by allowing each layer to adjust its activation distribution.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →