PyTorch torch.nn.ConvTranspose2d Function
PyTorch torch.nn Reference Manual
torch.nn.ConvTranspose2dIt is a 2D transposed convolution in PyTorch, also known as deconvolution or upsampling convolution.
It is used to upsample feature maps and is a key component of generative networks and segmentation networks.
Function Definition
torch.nn.ConvTranspose2d(in_channels, out_channels, kernel_size, stride=1, padding=0, output_padding=0, groups=1, bias=True, dilation=1)
Parameters
output_padding: extra output padding
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# Transposed convolution: upsampling
deconv = nn.ConvTranspose2d(in_channels=64, out_channels=64, kernel_size=2, stride=2)
x = torch.randn(1, 64, 16, 16)
output = deconv(x)
print(Input:, x.shape, -> Output:, output.shape)
import torch.nn as nn
# Transposed convolution: upsampling
deconv = nn.ConvTranspose2d(in_channels=64, out_channels=64, kernel_size=2, stride=2)
x = torch.randn(1, 64, 16, 16)
output = deconv(x)
print(Input:, x.shape, -> Output:, output.shape)
Example 2: Generative Network
Example
import torch
import torch.nn as nn
# Simplified DCGAN generator
class Generator(nn.Module):
def __init__(self, latent_dim=100):
super(Generator, self).__init__()
self.fc = nn.Linear(latent_dim, 512 * 4 * 4)
self.deconv = nn.Sequential(
nn.ConvTranspose2d(512, 256, 4, 2, 1), # 4->8
nn.BatchNorm2d(256),
nn.ReLU(),
nn.ConvTranspose2d(256, 128, 4, 2, 1), # 8->16
nn.BatchNorm2d(128),
nn.ReLU(),
nn.ConvTranspose2d(128, 64, 4, 2, 1), # 16->32
nn.BatchNorm2d(64),
nn.ReLU(),
nn.ConvTranspose2d(64, 3, 4, 2, 1), # 32->64
nn.Tanh()
)
def forward(self, x):
x = self.fc(x).view(-1, 512, 4, 4)
return self.deconv(x)
gen = Generator()
z = torch.randn(1, 100)
img = gen(z)
print(Noise:, z.shape, -> Image:, img.shape)
import torch.nn as nn
# Simplified DCGAN generator
class Generator(nn.Module):
def __init__(self, latent_dim=100):
super(Generator, self).__init__()
self.fc = nn.Linear(latent_dim, 512 * 4 * 4)
self.deconv = nn.Sequential(
nn.ConvTranspose2d(512, 256, 4, 2, 1), # 4->8
nn.BatchNorm2d(256),
nn.ReLU(),
nn.ConvTranspose2d(256, 128, 4, 2, 1), # 8->16
nn.BatchNorm2d(128),
nn.ReLU(),
nn.ConvTranspose2d(128, 64, 4, 2, 1), # 16->32
nn.BatchNorm2d(64),
nn.ReLU(),
nn.ConvTranspose2d(64, 3, 4, 2, 1), # 32->64
nn.Tanh()
)
def forward(self, x):
x = self.fc(x).view(-1, 512, 4, 4)
return self.deconv(x)
gen = Generator()
z = torch.randn(1, 100)
img = gen(z)
print(Noise:, z.shape, -> Image:, img.shape)
Example 3: Segmentation Network Upsampling
Example
import torch
import torch.nn as nn
# U-Net decoder part
decoder = nn.Sequential(
nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2),
nn.Conv2d(128, 128, 3, padding=1),
nn.ReLU(),
nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2),
nn.Conv2d(64, 64, 3, padding=1),
nn.ReLU()
)
x = torch.randn(1, 256, 8, 8)
output = decoder(x)
print(Input:, x.shape, -> Output:, output.shape)
import torch.nn as nn
# U-Net decoder part
decoder = nn.Sequential(
nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2),
nn.Conv2d(128, 128, 3, padding=1),
nn.ReLU(),
nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2),
nn.Conv2d(64, 64, 3, padding=1),
nn.ReLU()
)
x = torch.randn(1, 256, 8, 8)
output = decoder(x)
print(Input:, x.shape, -> Output:, output.shape)
Example 4: Output Calculation with stride=2
Example
import torch
import torch.nn as nn
# Different configurations
configs = [
(1, 1, 0), # stride=1, padding=0
(2, 1, 0), # stride=2, padding=0
(2, 1, 1), # stride=2, padding=1
]
x = torch.randn(1, 64, 4, 4)
for stride, k, p in configs:
deconv = nn.ConvTranspose2d(64, 64, kernel_size=k, stride=stride, padding=p)
out = deconv(x)
print(f"k={k}, s={stride}, p={p}: {x.shape} -> {out.shape}")
import torch.nn as nn
# Different configurations
configs = [
(1, 1, 0), # stride=1, padding=0
(2, 1, 0), # stride=2, padding=0
(2, 1, 1), # stride=2, padding=1
]
x = torch.randn(1, 64, 4, 4)
for stride, k, p in configs:
deconv = nn.ConvTranspose2d(64, 64, kernel_size=k, stride=stride, padding=p)
out = deconv(x)
print(f"k={k}, s={stride}, p={p}: {x.shape} -> {out.shape}")
Use Cases
- Generative networks: GAN、VAE
- Semantic segmentation: U-Net
- Upsampling: replace pooling
Note: Transposed convolution is not the inverse operation of convolution, but just a way of upsampling.
Other extensions