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

PyTorch实现的自编码器无法学习FashionMNIST数据集问题排查

自编码器重建FashionMNIST效果异常问题排查

问题复现

你使用的全连接自编码器结构为:

  • 编码器:784维输入→512维全连接+ReLU→30维瓶颈层+ReLU
  • 解码器:30维输入→512维全连接+ReLU→784维输出+Sigmoid
    你已完成图像灰度化、0-1区间归一化,刻意控制网络深度避免恒等映射,但训练1000轮后模型输出完全失真,无法还原输入图像。
    对应实现代码如下:
import torch
import torchvision as tv
import torchvision.transforms as transforms
import matplotlib.pyplot as plt
from torch import nn
import os
from torchviz import make_dot
transforms = tv.transforms.Compose([tv.transforms.Grayscale(num_output_channels=1)])
trainset = tv.datasets.FashionMNIST(root='./data', train=True,
                                        download=True, transform=transforms)
PATH = './ae.pth'
data = trainset.data.float()
data = data/255
plt.imshow(trainset.data[0], cmap = 'gray')
plt.show()

class NeuralNetwork(nn.Module):
    def __init__(self):
        super(NeuralNetwork, self).__init__()
        self.flatten = nn.Flatten()
        self.encode = nn.Sequential(
            nn.Linear(28*28, 512),
            nn.ReLU(),
            nn.Linear(512, 30),
            nn.ReLU()
        )
        self.decode = nn.Sequential(
            nn.Linear(30, 512),
            nn.ReLU(),
            nn.Linear(512, 28*28),
            nn.Sigmoid()
        )

    def forward(self, x):
        x = self.flatten(x)
        encoded = self.encode(x)
        decoded = self.decode(encoded)
        return decoded
if(os.path.exists(PATH)):
    print("Loading data on cpu")
    device = torch.device('cpu')
    model = NeuralNetwork()
    model.load_state_dict(torch.load(PATH, map_location=device))

else:
    device = "cuda" if torch.cuda.is_available() else "cpu"
    data = data.to(device)
    print(f"Using device = {device}")
    model = NeuralNetwork().to(device)

    lossFn  = nn.BCELoss()

    optimizer = torch.optim.SGD(model.parameters(), lr = 1e-3)

    for epoch in range(1000):
        print("Epoch = ", epoch)
        optimizer.zero_grad()
        outputs = model(data)
        loss = lossFn(outputs, data.reshape(-1, 784))
        loss.backward()
        optimizer.step()

    torch.save(model.state_dict(), PATH)
    data = data.to("cpu")
    model = model.to("cpu")

pred = model(data)
pred = pred.reshape(-1, 28, 28)
plt.imshow(pred.detach().numpy()[0], cmap = 'gray')
plt.show()

核心问题原因

  • 训练逻辑错误,收敛严重不足:你直接将全部60000张训练集一次性作为输入喂给网络,采用全批量梯度下降,同时设置的SGD学习率仅为1e-3。全批量梯度下降每次参数更新依赖全量样本的平均梯度,本身收敛速度远慢于小批量训练,搭配极小的学习率,1000轮迭代远不足以让模型收敛到有效状态,参数基本停留在初始化附近,输出自然是无意义的灰度块。
  • 优化器配置不合理:使用无动量的基础SGD优化器,在当前训练设置下极易卡在损失平台期,进一步拖慢收敛速度。
  • 代码存在冗余和变量冲突:你定义了灰度转换的transform但从未实际生效(后续直接读取trainset.data手动做归一化,没有调用transform逻辑),同时自定义的transforms变量直接覆盖了导入的torchvision.transforms模块,属于潜在代码隐患。
  • 补充说明:你的网络结构本身没有问题,30维瓶颈层配合两层全连接完全可以完成FashionMNIST的简单重建,不存在容量不足或者容易形成恒等映射的问题——恒等映射只有在瓶颈层维度大于等于输入维度、网络容量足够且无正则约束时才可能出现,当前30维远小于784维输入,本身就存在强制信息压缩,不会出现恒等映射问题。

修复方案

  • 替换全批量训练为小批量训练:使用DataLoader封装数据集,设置batch size为64或128,开启样本打乱,逐批次完成前向传播、损失计算和参数更新,参考代码如下:
from torch.utils.data import DataLoader
# 注意删除之前覆盖transforms模块的自定义变量,避免冲突
train_transform = tv.transforms.Compose([tv.transforms.ToTensor()])
trainset = tv.datasets.FashionMNIST(root='./data', train=True, download=True, transform=train_transform)
trainloader = DataLoader(trainset, batch_size=64, shuffle=True)
# 训练循环修改为
for epoch in range(50):
    total_loss = 0
    for batch_data, _ in trainloader:
        batch_data = batch_data.float().to(device).reshape(-1, 784)
        optimizer.zero_grad()
        outputs = model(batch_data)
        loss = lossFn(outputs, batch_data)
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
    print(f"Epoch {epoch}, Loss: {total_loss/len(trainloader):.4f}")
  • 调整优化器配置:要么将SGD学习率提升至0.1并设置momentum=0.9,要么直接更换为Adam优化器,初始学习率设为1e-3即可获得更快的收敛速度。
  • 清理冗余代码:删除未生效的transform定义,修改自定义变量名避免覆盖导入的模块,消除潜在冲突。

按上述方案修改后,训练30-50轮即可得到清晰的重建结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 08:12:25