PyTorch中如何使Conv2d与ConvTranspose2d保持输入输出形状一致?
解决PyTorch中Conv2d与ConvTranspose2d形状不匹配的问题
问题根源
Conv2d和ConvTranspose2d的输出尺寸由kernel_size、stride、padding、output_padding参数共同决定,默认参数下无法保证输入输出形状一致,尤其是输入尺寸无法被stride整除时,形状差异会更明显。
针对任意height/width的解决方案
要让两者互为逆操作(输入输出形状完全一致),可通过手动计算匹配参数或利用PyTorch内置特性实现:
方案1:手动计算参数匹配
先明确两个层的尺寸计算公式:
- Conv2d输出尺寸:
H_out = floor((H_in + 2*pad - kernel) / stride) + 1 W_out = floor((W_in + 2*pad - kernel) / stride) + 1 - ConvTranspose2d输出尺寸:
H_out = (H_in - 1)*stride - 2*pad + kernel + output_pad W_out = (W_in - 1)*stride - 2*pad + kernel + output_pad
要让ConvTranspose2d的输出等于原始输入尺寸,需根据Conv2d的输出反推output_padding:
output_pad_H = original_H - ((conv_H_out -1)*stride - 2*conv_pad + kernel) output_pad_W = original_W - ((conv_W_out -1)*stride - 2*conv_pad + kernel)
对应代码修改如下:
import torch data = torch.rand(1,1,16,26) original_h, original_w = data.shape[2], data.shape[3] # 定义Conv2d,可自定义kernel和stride conv = torch.nn.Conv2d(1,1,kernel_size=3, stride=2, padding=0) conv_out = conv(data) conv_out_h, conv_out_w = conv_out.shape[2], conv_out.shape[3] # 计算所需的output_padding output_pad_h = original_h - ((conv_out_h -1)*conv.stride[0] - 2*conv.padding[0] + conv.kernel_size[0]) output_pad_w = original_w - ((conv_out_w -1)*conv.stride[1] - 2*conv.padding[1] + conv.kernel_size[1]) # 定义匹配的ConvTranspose2d conv_trans = torch.nn.ConvTranspose2d(1,1,kernel_size=3, stride=2, padding=0, output_padding=(output_pad_h, output_pad_w)) trans_out = conv_trans(conv_out) print(trans_out.shape) # torch.Size([1, 1, 16, 26]),与原始输入形状一致
方案2:使用padding='same'简化设置
PyTorch 1.10+支持padding='same'参数,Conv2d使用该参数时,输出尺寸为ceil(input_size / stride)。配合ConvTranspose2d的padding='same'和计算出的output_padding,可快速实现形状匹配:
import torch data = torch.rand(1,1,16,26) original_h, original_w = data.shape[2], data.shape[3] conv = torch.nn.Conv2d(1,1,kernel_size=3, stride=2, padding='same') conv_out = conv(data) conv_out_h, conv_out_w = conv_out.shape[2], conv_out.shape[3] # 计算output_padding output_pad_h = original_h - (conv_out_h -1)*conv.stride[0] output_pad_w = original_w - (conv_out_w -1)*conv.stride[1] conv_trans = torch.nn.ConvTranspose2d(1,1,kernel_size=3, stride=2, padding='same', output_padding=(output_pad_h, output_pad_w)) trans_out = conv_trans(conv_out) print(trans_out.shape) # torch.Size([1, 1, 16, 26])
方案3:PixelShuffle替代(适用于stride为2的幂次场景)
如果是上采样类分割任务,可结合普通卷积与PixelShuffle,避免手动计算参数的繁琐,仅适用于stride为2的幂次的情况:
import torch # 先通过卷积降维,再用PixelShuffle上采样 conv = torch.nn.Conv2d(1,4,kernel_size=3, stride=2, padding=1) pixel_shuffle = torch.nn.PixelShuffle(2) data = torch.rand(1,1,16,26) conv_out = conv(data) trans_out = pixel_shuffle(conv_out) print(trans_out.shape) # torch.Size([1, 1, 16, 26])
关键注意事项
output_padding的取值范围是0到stride-1,超出范围会报错。- 需针对height和width分别计算
output_padding,不能统一设置单一值。 - 若使用分组卷积或
dilation参数,需调整尺寸计算公式以匹配参数。
内容的提问来源于stack exchange,提问作者vivian
相关产品推荐
相关产品推荐

