PyTorch torch.nn.ConvTranspose2d Function

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

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)

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)

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

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.


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

Other extensions