You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

无池化PyTorch卷积自编码器训练MNIST时维度错误求助

问题根源分析

  1. 编码器通道不匹配错误:编码器序列倒数第二层nn.Conv2d(128, 256, kernel_size=5, stride=1)的输入通道设置错误,前一层输出通道为256,这层输入通道应改为256,否则会直接触发通道维度不匹配的RuntimeError。
  2. 编码器输出形状不符合要求:当前编码器的输出无法达到你需要的(256,16,1,1)(默认第一个维度为batch size),多次无padding的5x5卷积仅能缩小空间维度,既无法压缩到1x1,也没有将通道数映射到嵌入维度16。
  3. 训练循环损失计算错误:自编码器的训练目标是重构输入图像,损失计算应使用模型输出与输入特征features对比,而非与分类标签labels对比。

修正后的代码实现

1. 修正自编码器结构

调整编码器卷积层的通道参数,添加最终卷积层将通道映射到16,并通过stride参数替代池化压缩空间维度;解码器对应调整输入通道,逐步恢复原始图像尺寸:

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor

class AutoEncoderCNN(nn.Module):
    def __init__(self, embedding_dim=16):
        super(AutoEncoderCNN, self).__init__()
        self.encoder = nn.Sequential(
            # 输入(1,28,28) → 输出(16,24,24)
            nn.Conv2d(1, 16, kernel_size=5, stride=1),
            nn.ReLU(),
            # (16,24,24) → (32,20,20)
            nn.Conv2d(16, 32, kernel_size=5, stride=1),
            nn.ReLU(),
            # (32,20,20) → (64,16,16)
            nn.Conv2d(32, 64, kernel_size=5, stride=1),
            nn.ReLU(),
            # (64,16,16) → (128,12,12)
            nn.Conv2d(64, 128, kernel_size=5, stride=1),
            nn.ReLU(),
            # (128,12,12) → (256,8,8)
            nn.Conv2d(128, 256, kernel_size=5, stride=1),
            nn.ReLU(),
            # 修正通道错误+用stride=2缩小维度 → (256,3,3)
            nn.Conv2d(256, 256, kernel_size=5, stride=2),
            nn.ReLU(),
            # 映射到嵌入维度16,压缩空间到1x1 → (16,1,1)
            nn.Conv2d(256, embedding_dim, kernel_size=3, stride=1),
            nn.ReLU()
        )
        self.decoder = nn.Sequential(
            # 输入(16,1,1) → 转置卷积恢复维度 → (256,3,3)
            nn.ConvTranspose2d(embedding_dim, 256, kernel_size=3, stride=1),
            nn.ReLU(),
            # (256,3,3) → (256,8,8),output_padding修正stride=2的维度对齐
            nn.ConvTranspose2d(256, 256, kernel_size=5, stride=2, output_padding=1),
            nn.ReLU(),
            # (256,8,8) → (128,12,12)
            nn.ConvTranspose2d(256, 128, kernel_size=5, stride=1),
            nn.ReLU(),
            # (128,12,12) → (64,16,16)
            nn.ConvTranspose2d(128, 64, kernel_size=5, stride=1),
            nn.ReLU(),
            # (64,16,16) → (32,20,20)
            nn.ConvTranspose2d(64, 32, kernel_size=5, stride=1),
            nn.ReLU(),
            # (32,20,20) → (16,24,24)
            nn.ConvTranspose2d(32, 16, kernel_size=5, stride=1),
            nn.ReLU(),
            # (16,24,24) → (1,28,28),恢复原始图像尺寸
            nn.ConvTranspose2d(16, 1, kernel_size=5, stride=1),
            nn.Sigmoid()      
        )

    def encode(self, x):
        x = self.encoder(x)
        return x
        
    def decode(self, x):
        x = self.decoder(x)
        return x
        
    def forward(self, x):
        x = self.encoder(x)
        x = self.decoder(x)
        return x

2. 修正训练循环

调整损失计算的目标为输入特征,同时统一优化器的使用:

def train(model, data_loader, opt, n_epochs):
    losses = []  
    i=0
    for epoch in range(n_epochs):
        running_loss = 0.0
        for features, labels in data_loader:       
            # 前向传播
            labels_pred = model(features) 
            # 修正:用输入图像作为目标计算重构损失
            loss = loss_function(labels_pred, features)
            losses.append(loss.item())
            
            # 梯度更新流程
            opt.zero_grad()
            loss.backward()
            opt.step()

            running_loss += loss.item()
            if i % 10 == 9:    
                print('[Epoque : %d, iteration: %5d] loss: %.3f'% 
                      (epoch + 1, i + 1, running_loss / 10))
                running_loss = 0.0
            i+=1   
    print('Entrainement terminé')
    return losses

关键说明

  • 通道修正:将编码器倒数第二层的输入通道改为256,匹配前一层输出通道,解决通道不匹配错误。
  • 维度控制:用stride=2的卷积替代池化缩小空间维度,最后用3x3卷积将空间维度压缩到1x1,同时将通道数映射到嵌入维度16,满足输出形状要求。
  • 解码器对齐:解码器使用转置卷积逐步恢复空间维度,output_padding参数修正stride=2时的维度偏差,确保最终输出与输入尺寸一致。
  • 损失逻辑修正:自编码器核心是重构输入,损失计算必须基于输入图像而非分类标签。

内容的提问来源于stack exchange,提问作者Rayzzen

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.03 06:56:55