PyTorch torch.nn.InstanceNorm2d Function
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 channelsaffine: 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)
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)
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())
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.
Other Extensions