MNIST运行卷积自编码器报MaxUnpool2d维度不匹配错误求解
问题原因
报错来自两个核心逻辑错误,和卷积、池化的层尺寸参数无关:
- 编码器前向传播流程写错:
maxpool2的输出没有传入conv3层,反而直接把maxpool2之前的14×14特征传入conv3再做池化,导致三个池化层输出的索引尺寸分别是14×14、7×7、3×3,和预期的尺寸流完全错位。 - 解码器结构不对称、索引传参顺序错误:解码器的上采样(Unmaxpool)和反卷积顺序和编码器不匹配,传入的池化索引顺序和上采样步骤不对应,导致执行
unmaxpool2时,待上采样的特征是14×14尺寸,但拿到的索引是7×7尺寸,直接触发形状不匹配报错。
保留3卷积+3最大池化结构的前提下,MNIST数据集(输入尺寸1×28×28)的正确尺寸流如下:
- 编码器路径:输入28×28 → Conv1(尺寸不变)→ MaxPool1(下采样2倍到14×14)→ Conv2(尺寸不变)→ MaxPool2(下采样2倍到7×7)→ Conv3(尺寸不变)→ MaxPool3(下采样2倍到3×3,7整除2得3)
- 解码器路径:输入3×3特征 → Unmaxpool对应MaxPool3(上采样到7×7)→ ConvT1(尺寸不变)→ Unmaxpool对应MaxPool2(上采样到14×14)→ ConvT2(尺寸不变)→ Unmaxpool对应MaxPool1(上采样到28×28)→ ConvT3(输出1通道,尺寸不变),最终输出和原输入尺寸完全一致。
修复后可运行代码
import torch import torch.nn as nn import torchvision from torch.utils.data import DataLoader from tqdm import tqdm class ConvolutionEncoder(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1,32,3,stride=1,padding=1,dilation=1) self.maxpool1 = nn.MaxPool2d(2,padding=0,return_indices=True) self.conv2 = nn.Conv2d(32,32,3,stride=1,padding=1,dilation=1) self.maxpool2 = nn.MaxPool2d(2,padding=0,return_indices=True) self.conv3 = nn.Conv2d(32,32,3,stride=1,padding=1,dilation=1) self.maxpool3 = nn.MaxPool2d(2,padding=0,return_indices=True) self.relu = nn.ReLU() self.maxpool1_index = None self.maxpool2_index = None self.maxpool3_index = None def forward(self, image): temp = self.conv1(image) temp = self.relu(temp) temp, maxpool1_index = self.maxpool1(temp) temp = self.conv2(temp) temp = self.relu(temp) temp, maxpool2_index = self.maxpool2(temp) temp = self.conv3(temp) temp = self.relu(temp) feature_map, maxpool3_index = self.maxpool3(temp) self.maxpool1_index = maxpool1_index self.maxpool2_index = maxpool2_index self.maxpool3_index = maxpool3_index return feature_map class ConvolutionDecoder(nn.Module): def __init__(self): super().__init__() self.unmaxpool1 = nn.MaxUnpool2d(2,padding=0) self.conv1 = nn.ConvTranspose2d(32,32,3,stride=1,padding=1,dilation=1) self.unmaxpool2 = nn.MaxUnpool2d(2,padding=0) self.conv2 = nn.ConvTranspose2d(32,32,3,stride=1,padding=1,dilation=1) self.unmaxpool3 = nn.MaxUnpool2d(2,padding=0) self.conv3 = nn.ConvTranspose2d(32,1,3,stride=1,padding=1,dilation=1) self.relu = nn.ReLU() self.sigmoid = nn.Sigmoid() def forward(self, feature_map, maxpool_index): # 索引顺序和编码器下采样顺序逆序对应:[maxpool3索引, maxpool2索引, maxpool1索引] # 3*3特征上采样2倍默认输出6*6,需指定output_size为7*7匹配后续层尺寸 temp = self.unmaxpool1(feature_map, maxpool_index[0], output_size=[feature_map.shape[0],32,7,7]) temp = self.relu(temp) temp = self.conv1(temp) temp = self.unmaxpool2(temp, maxpool_index[1]) temp = self.relu(temp) temp = self.conv2(temp) temp = self.unmaxpool3(temp, maxpool_index[2]) temp = self.relu(temp) temp = self.conv3(temp) reconstructed_image = self.sigmoid(temp) return reconstructed_image class ConvolutionAutoencoder(nn.Module): def __init__(self): super().__init__() self.encoder = ConvolutionEncoder() self.decoder = ConvolutionDecoder() def forward(self, image): feature_map = self.encoder(image) reconstructed_image = self.decoder(feature_map, [self.encoder.maxpool3_index, self.encoder.maxpool2_index, self.encoder.maxpool1_index]) return reconstructed_image # 加载数据集 train_data_transformed = torchvision.datasets.MNIST(root="/MNIST", train=True, download=True,transform=torchvision.transforms.ToTensor()) train_dataloader = DataLoader(train_data_transformed, batch_size=1024, shuffle=True) test_data_transformed = torchvision.datasets.MNIST(root="/MNIST", train=False, download=True,transform=torchvision.transforms.ToTensor()) test_dataloader = DataLoader(test_data_transformed, batch_size=1024) conv_autoencoder = ConvolutionAutoencoder() conv_optimizer = torch.optim.AdamW(conv_autoencoder.parameters(), lr=1e-3) conv_MSELoss = nn.MSELoss() epochs = 10 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") conv_autoencoder.to(device) for conv_epoch_idx in tqdm(range(epochs)): conv_autoencoder.train() train_loss = 0 for conv_batch_idx, (imgs, _) in enumerate(train_dataloader): imgs = imgs.to(device) conv_optimizer.zero_grad() reconstructed = conv_autoencoder(imgs) loss = conv_MSELoss(reconstructed, imgs) loss.backward() conv_optimizer.step() train_loss += loss.item() print(f"Epoch {conv_epoch_idx+1}, Train Loss: {train_loss/len(train_dataloader):.6f}")
补充说明
- 移除了原编码器、解码器中存储中间特征的
self.middle列表,避免训练时中间特征一直被类属性引用导致内存泄漏、梯度计算异常。 - 第一次Unmaxpool操作必须指定
output_size参数:3×3的特征经过核为2的MaxUnpool默认输出尺寸是6×6,和MaxPool2对应的7×7输入尺寸差1,会触发第二次形状报错,指定输出尺寸即可解决。 - 原训练循环缺失了损失计算、反向传播、优化器更新参数的步骤,补全后即可正常在MNIST上训练。
内容的提问来源于stack exchange,提问作者vesii
相关产品推荐
相关产品推荐

