PyTorch torch.nn.BatchNorm2d Function
PyTorch torch.nn Reference Manual
torch.nn.BatchNorm2dIt is a module in PyTorch for two-dimensional batch normalization.
Batch normalization accelerates training and stabilizes convergence by normalizing the input of a layer, and is one of the most commonly used techniques in modern deep neural networks.
Function Definition
torch.nn.BatchNorm2d(num_features, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
Parameter Description:
num_features(int): Number of input channels C.eps(float): Value added to the denominator for numerical stability. Default is 1e-5.momentum(float): Used for computing running mean and variance. Default is 0.1.affine(bool): Whether to use learnable affine parameters (gamma and beta). Default is True.track_running_stats(bool): Whether to track running statistics. Default is True.
Attributes:
weight(Tensor): Learnable scaling parameter gamma, shape (num_features,).bias(Tensor): Learnable shift parameter beta, shape (num_features,).running_mean(Tensor): Running mean, not updated during training.running_var(Tensor): Running variance, not updated during training.
Mathematical Principle
Batch normalization normalizes each channel independently:
y = (x - E[x]) / sqrt(Var[x] + eps) * gamma + beta
Here, gamma and beta are learnable parameters that allow the network to restore its expressive capability.
Usage Examples
Example 1: Basic Usage
Using batch normalization after a convolutional layer:
Example
import torch.nn as nn
# Create a batch normalization layer: 32 channels
bn = nn.BatchNorm2d(num_features=32)
# Print parameters
print("gamma (weight):", bn.weight.shape)
print("beta (bias):", bn.bias.shape)
print("running_mean:", bn.running_mean.shape)
print("running_var:", bn.running_var.shape)
# Create input: batch=4, channels=32, height=16, width=16
input_tensor = torch.randn(4, 32, 16, 16)
# Forward pass
output = bn(input_tensor)
print("nInput mean (per channel):", input_tensor.mean(dim=(0, 2, 3))[:5].tolist())
print("Output mean (per channel):", output.mean(dim=(0, 2, 3))[:5].tolist())
print("nInput shape:", input_tensor.shape)
print("Output shape:", output.shape)
The output result is:
gamma (weight): torch.Size([32]) beta (bias): torch.Size([32]) running_mean: torch.Size([32]) running_var: torch.Size([32]) 输入均值 (按通道): tensor([ 0.0234, 0.0456, -0.0123, -0.0345, 0.0567]) 输出均值 (按通道): tensor([ 0., 0., 0., 0., 0.]) 输入形状: torch.Size([4, 32, 16, 16]) 输出形状: torch.Size([4, 32, 16, 16])
During training, the output is normalized to mean 0 and variance 1 (followed by gamma and beta transformation).
Example 2: Training vs Evaluation Mode
Batch normalization behaves differently during training and evaluation:
Example
import torch.nn as nn
bn = nn.BatchNorm2d(num_features=16)
# Training mode
bn.train()
print("Training mode - requires grad:", bn.weight.requires_grad)
# Simulate training
for _ in range(10):
x = torch.randn(8, 16, 8, 8)
output = bn(x)
print("First 5 running_mean after training:", bn.running_mean[:5].tolist())
# Evaluation mode
bn.eval()
print("nEvaluation mode - requires grad:", bn.weight.requires_grad)
# Use running stats during evaluation
x = torch.randn(4, 16, 8, 8)
output = bn(x)
print("Output shape during evaluation:", output.shape)
Example 3: Complete CNN Example
Typical batch normalization CNN structure:
Example
import torch.nn as nn
class BNConvNet(nn.Module):
def __init__(self, num_classes=10):
super(BNConvNet, self).__init__()
# Convolution + batch normalization + activation + pooling
self.block1 = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(),
nn.MaxPool2d(2, 2) # 32 -> 16
)
self.block2 = nn.Sequential(
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(2, 2) # 16 -> 8
)
# Global average pooling
self.gap = nn.AdaptiveAvgPool2d(1)
self.classifier = nn.Linear(64, num_classes)
def forward(self, x):
x = self.block1(x)
x = self.block2(x)
x = self.gap(x)
x = x.view(x.size(0), -1)
x = self.classifier(x)
return x
model = BNConvNet()
input_image = torch.randn(2, 3, 32, 32)
output = model(input_image)
print("Input shape:", input_image.shape)
print("Output shape:", output.shape)
# Print parameters of the first BN layer
print("nFirst BN layer gamma:", model.block1[1].weight[:5].tolist())
print("First BN layer beta:", model.block1[1].bias[:5].tolist())
Example 4: Not Using affine Parameters
Disable learnable parameters:
Example
import torch.nn as nn
# Batch normalization without learnable parameters
bn_no_affine = nn.BatchNorm2d(16, affine=False)
print("Has weight:", bn_no_affine.weight is not None)
print("Has bias:", bn_no_affine.bias is not None)
# It still performs normalization
x = torch.randn(4, 16, 8, 8)
output = bn_no_affine(x)
print("nOutput shape:", output.shape)
Common Questions
Q1: Should batch normalization be placed before or after ReLU?
Both approaches are effective. The original paper places it after convolution and before activation; in practice, it is also often placed after activation.
Q2: What to do when it doesn't work well with small batch sizes?
- Use
GroupNorminstead of - Use
LayerNorm - Increase batch size
- Adjust the momentum parameter
Q3: Why switch to eval mode during evaluation?
During training, batch statistics are used; during evaluation, running statistics are used. Forgetting to switch will lead to inconsistent outputs.
Use Cases
nn.BatchNorm2dMain application scenarios include:
- Accelerate training: Allows using a larger learning rate
- Stabilize convergence: Reduces internal covariate shift
- Regularization: Provides a slight regularization effect
- Image classification: Used by almost all modern CNNs
Note: Batch normalization requires a sufficient batch size during training; small batches may lead to unstable statistics.
Other Extensions