PyTorch torch.nn.GroupNorm Function
PyTorch torch.nn Reference Manual
torch.nn.GroupNormIt is the group normalization module in PyTorch.
It divides channels into groups for normalization, does not depend on batch size, and is suitable for small batches or variable-length sequences.
Function Definition
torch.nn.GroupNorm(num_groups, num_channels, eps=1e-05, affine=True)
Parameters:
num_groups: Number of groupsnum_channels: Number of channels, must be divisible by num_groups
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# GroupNorm: 4 groups, 32 channels
gn = nn.GroupNorm(num_groups=4, num_channels=32)
# Input: (batch, channels, height, width)
x = torch.randn(8, 32, 16, 16)
output = gn(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
import torch.nn as nn
# GroupNorm: 4 groups, 32 channels
gn = nn.GroupNorm(num_groups=4, num_channels=32)
# Input: (batch, channels, height, width)
x = torch.randn(8, 32, 16, 16)
output = gn(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
Example 2: Special Case - InstanceNorm
Example
import torch
import torch.nn as nn
# InstanceNorm is a special case where num_groups=num_channels
inorm = nn.InstanceNorm2d(32)
# Equivalent to GroupNorm, num_groups=32, num_channels=32
gn_equiv = nn.GroupNorm(32, 32)
x = torch.randn(4, 32, 16, 16)
print("InstanceNorm2d:", inorm(x).mean().item())
print("GroupNorm:", gn_equiv(x).mean().item())
import torch.nn as nn
# InstanceNorm is a special case where num_groups=num_channels
inorm = nn.InstanceNorm2d(32)
# Equivalent to GroupNorm, num_groups=32, num_channels=32
gn_equiv = nn.GroupNorm(32, 32)
x = torch.randn(4, 32, 16, 16)
print("InstanceNorm2d:", inorm(x).mean().item())
print("GroupNorm:", gn_equiv(x).mean().item())
Example 3: GroupNorm in ResNet
Example
import torch
import torch.nn as nn
# ResNet structure uses GroupNorm
class ConvBlock(nn.Module):
def __init__(self, in_ch, out_ch, num_groups=8):
super(ConvBlock, self).__init__()
self.conv = nn.Conv2d(in_ch, out_ch, 3, padding=1)
self.gn = nn.GroupNorm(num_groups, out_ch)
self.relu = nn.ReLU()
def forward(self, x):
return self.relu(self.gn(self.conv(x)))
block = ConvBlock(64, 128)
x = torch.randn(2, 64, 32, 32)
output = block(x)
print("Input:", x.shape, "-> Output:", output.shape)
import torch.nn as nn
# ResNet structure uses GroupNorm
class ConvBlock(nn.Module):
def __init__(self, in_ch, out_ch, num_groups=8):
super(ConvBlock, self).__init__()
self.conv = nn.Conv2d(in_ch, out_ch, 3, padding=1)
self.gn = nn.GroupNorm(num_groups, out_ch)
self.relu = nn.ReLU()
def forward(self, x):
return self.relu(self.gn(self.conv(x)))
block = ConvBlock(64, 128)
x = torch.randn(2, 64, 32, 32)
output = block(x)
print("Input:", x.shape, "-> Output:", output.shape)
Example 4: Effect of Different num_groups
Example
import torch
import torch.nn as nn
# 32 channels, using different numbers of groups
for num_groups in [1, 2, 4, 8, 32]:
gn = nn.GroupNorm(num_groups, 32)
x = torch.randn(4, 32, 8, 8)
out = gn(x)
print(f"groups={num_groups}, output mean: {out.mean().item():.4f}")
import torch.nn as nn
# 32 channels, using different numbers of groups
for num_groups in [1, 2, 4, 8, 32]:
gn = nn.GroupNorm(num_groups, 32)
x = torch.randn(4, 32, 8, 8)
out = gn(x)
print(f"groups={num_groups}, output mean: {out.mean().item():.4f}")
Use Cases
- Small batch: still works when batch=1
- Video/3D: BatchNorm is unstable
- MobileNet: GroupNorm replaces BatchNorm
Note: num_channels must be divisible by num_groups.
Other extensions