PyTorch torch.nn.GroupNorm Function

PyTorch torch.nn 参考手册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 groups
  • num_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)

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())

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)

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}")

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.


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

Other extensions