PyTorch torch.nn.InstanceNorm2d Function

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


torch.nn.InstanceNorm2dIt is the instance normalization module in PyTorch.

It normalizes each channel of each sample independently, commonly used in style transfer.

Function Definition

torch.nn.InstanceNorm2d(num_features, eps=1e-05, momentum=0.1, affine=False, track_running_stats=False)

Parameters

  • num_features: number of channels
  • affine: whether to use learnable parameters

Usage Examples

Example 1: Basic Usage

Example

import torch
import torch.nn as nn

inorm = nn.InstanceNorm2d(64)

# Input
x = torch.randn(4, 64, 16, 16)
output = inorm(x)

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

Example 2: Style Transfer

Example

import torch
import torch.nn as nn

# Style transfer networks commonly use InstanceNorm
class StyleNet(nn.Module):
    def __init__(self):
        super(StyleNet, self).__init__()
        self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
        self.in1 = nn.InstanceNorm2d(32)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.in2 = nn.InstanceNorm2d(64)

    def forward(self, x):
        x = self.conv1(x)
        x = self.in1(x)
        x = torch.relu(x)
        x = self.conv2(x)
        x = self.in2(x)
        return x

net = StyleNet()
x = torch.randn(1, 3, 256, 256)
output = net(x)

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

Example 3: Comparison with BatchNorm

Example

import torch
import torch.nn as nn

bn = nn.BatchNorm2d(32)
inorm = nn.InstanceNorm2d(32)

x = torch.randn(4, 32, 8, 8)

print("BatchNorm mean:", bn(x).mean(dim=(0, 2, 3))[:5].tolist())
print("InstanceNorm mean:", inorm(x).mean(dim=(0, 2, 3))[:5].tolist())

Use Cases

  • Style Transfer: NIN、AdaIN
  • Texture Synthesis
  • Small batch: suitable for batch=1

Note: Default affine=False, contains no learnable parameters.


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

Other Extensions