PyTorch torch.nn.BatchNorm2d Function

PyTorch torch.nn 参考手册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
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
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
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
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?

  • UseGroupNorminstead of
  • UseLayerNorm
  • 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.


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

Other Extensions