PyTorch torch.nn.LayerNorm Function

PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual


torch.nn.LayerNormis the layer normalization module in PyTorch.

Unlike batch normalization, layer normalization normalizes along the feature dimension of a single sample and does not depend on batch size.

Function Definition

torch.nn.LayerNorm(normalized_shape, eps=1e-05, elementwise_affine=True)

Parameter Description:

  • normalized_shape(int or list): The dimensions to normalize.
  • eps(float): Epsilon for numerical stability. Default is 1e-5.
  • elementwise_affine(bool): Whether to use learnable scale and shift. Default is True.

Mathematical Principle

Layer normalization formula:

y = (x - E[x]) / sqrt(Var[x] + eps) * gamma + beta

Difference from batch normalization: Layer normalization computes the mean and variance over the last dimension of the features.


Usage Examples

Example 1: Basic Usage

Perform layer normalization on features:

Example

import torch
import torch.nn as nn

# Layer normalization: normalize the last dimension
ln = nn.LayerNorm(normalized_shape=10)

# Input: batch=4, feature dimension=10
x = torch.randn(4, 10)

# Forward propagation
output = ln(x)

print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Original input first row:", x[0].tolist())
print("First row after normalization:", output[0].tolist())

Example 2: Multi-dimensional Input

Handling 3D or 4D input:

Example

import torch
import torch.nn as nn

# Normalize over sequence dimension: (batch, seq, features)
ln_seq = nn.LayerNorm(normalized_shape=64)

# 3D input
x_3d = torch.randn(2, 10, 64)
output_3d = ln_seq(x_3d)
print("3D input:", x_3d.shape, "-> output:", output_3d.shape)

# 4D input (e.g., images): (batch, height, width, channels)
# LayerNorm normalizes the last channel dimension
ln_channel = nn.LayerNorm(normalized_shape=128)
x_4d = torch.randn(2, 8, 8, 128)
output_4d = ln_channel(x_4d)
print("4D input:", x_4d.shape, "-> output:", output_4d.shape)

Example 3: Use in Transformer

Typical LayerNorm usage:

Example

import torch
import torch.nn as nn

class TransformerBlock(nn.Module):
    def __init__(self, d_model, nhead):
        super(TransformerBlock, self).__init__()
        self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.ffn = nn.Sequential(
            nn.Linear(d_model, d_model * 4),
            nn.GELU(),
            nn.Linear(d_model * 4, d_model)
        )

    def forward(self, x):
        # Self-attention with residual
        attn_out, _ = self.self_attn(x, x, x)
        x = self.norm1(x + attn_out)

        # FFN with residual
        ffn_out = self.ffn(x)
        x = self.norm2(x + ffn_out)

        return x

# Test
block = TransformerBlock(d_model=512, nhead=8)
x = torch.randn(4, 100, 512)  # (batch, seq, d_model)
output = block(x)

print("Input shape:", x.shape)
print("Output shape:", output.shape)

Example 4: Without Learnable Parameters

Pure normalization without scale and shift:

Example

import torch
import torch.nn as nn

# Without learnable parameters
ln = nn.LayerNorm(16, elementwise_affine=False)

x = torch.randn(4, 16)
output = ln(x)

# No weight and bias
print("Has weight:", hasattr(ln, 'weight'))
print("Has bias:", hasattr(ln, 'bias'))
print("Output shape:", output.shape)

Comparison of Normalization Methods

Method Normalization dimension Batch dependency Applicable scenarios
BatchNorm Batch dimension Yes CNN, stable batch
LayerNorm Feature dimension no Transformer、RNN
InstanceNorm Channel + spatial no Style transfer
GroupNorm Channel grouping no Small batch scenarios

Frequently Asked Questions

Q1: What is the difference between LayerNorm and BatchNorm?

LayerNorm does not depend on batch size, making it suitable for sequence models and scenarios where batch size varies greatly.

Q2: How to choose normalized_shape?

Usually choose the feature dimension, such as 768 or 1024 in BERT.

Q3: Why does Transformer use LayerNorm?

Transformer input sequence length is variable, and LayerNorm is more stable.


Use Cases

nn.LayerNormMain application scenarios include:

  • Transformer architecture: BERT, GPT, etc.
  • Recurrent neural networks: LSTM、GRU
  • Variable-length sequence processing: Batch size is not fixed

Tip: LayerNorm is a standard component of Transformer, placed after the residual connection (Post-LN) or before it (Pre-LN).


PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual

Other extensions