MNIST数据集训练避坑指南:从数据加载到模型保存,PyTorch新手常踩的5个雷

第一次用PyTorch跑MNIST手写数字识别,本以为跟着教程敲代码就能轻松搞定,结果从数据加载到模型保存,处处是坑。明明代码和教程里一模一样,偏偏报错不断:CUDA内存不足、DataLoader卡住不动、模型保存后加载出错、预测结果全错……这些问题我都遇到过,今天就把这些坑一个个填平。

1. 数据加载的隐藏陷阱

1.1 下载速度慢到怀疑人生

很多教程直接使用 torchvision.datasets.MNIST 自动下载数据集,但国内用户经常会遇到下载速度极慢甚至失败的情况。这时候可以手动下载MNIST数据集,放到指定目录。

手动下载MNIST数据集步骤:

  1. 访问 MNIST官网 下载四个.gz文件
  2. 在项目目录下创建 data/MNIST/raw 文件夹
  3. 将下载的文件放入raw文件夹
  4. 代码中设置 download=False
# 正确的数据加载方式(考虑国内网络环境)
import os
from torchvision import datasets, transforms

# 创建数据目录(如果不存在)
os.makedirs('./data/MNIST/raw', exist_ok=True)

# 加载数据集(设置download=False)
train_dataset = datasets.MNIST(
    root='./data', 
    train=True, 
    transform=transforms.ToTensor(),
    download=False  # 因为我们已手动下载
)

1.2 transforms.Normalize的参数陷阱

很多新手直接复制代码,却不知道 transforms.Normalize 的参数需要根据数据集特性设置。MNIST是单通道图像,均值标准差应该这样设置:

# 正确的Normalize参数设置
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST的标准均值和标准差
])

注意:这里的(0.1307,)和(0.3081,)是MNIST数据集的标准统计值,直接使用这些值能获得更好的训练效果。

2. 模型训练中的常见错误

2.1 维度不匹配:view操作的坑

在CNN中,从卷积层过渡到全连接层时,需要用 view 改变张量形状。这里最常见的错误是算错尺寸。

# 错误的view操作
x = x.view(x.size(0), -1)  # 如果前面层输出尺寸计算错误,这里会报错

# 正确的做法是先计算卷积后的尺寸
# 对于MNIST的28x28输入,经过两次2x2池化后是7x7
# 如果卷积核和padding设置不同,这个值会变化
def forward(self, x):
    x = self.conv_layers(x)
    x = x.view(x.size(0), -1)  # 确保这里的-1等于全连接层的输入尺寸
    x = self.fc_layers(x)
    return x

2.2 GPU/CPU转换遗漏

当你在GPU上训练模型,却在CPU上加载测试时,会遇到"Expected tensor to be on device X but got device Y"的错误。

解决方案:

  1. 统一设备设置
  2. 保存模型时记录设备信息
# 训练时
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

# 测试时
model.load_state_dict(torch.load('model.pth', map_location=device))

3. 模型保存与加载的坑

3.1 保存整个模型 vs 只保存参数

新手常犯的错误是混淆两种保存方式:

保存方式 代码示例 优点 缺点
整个模型 torch.save(model, 'model.pth') 简单直接 文件大,可能与环境不兼容
仅参数 torch.save(model.state_dict(), 'params.pth') 文件小,灵活 需要重建模型结构

推荐做法:

# 保存
torch.save({
    'epoch': epoch,
    'model_state_dict': model.state_dict(),
    'optimizer_state_dict': optimizer.state_dict(),
    'loss': loss,
}, 'checkpoint.pth')

# 加载
checkpoint = torch.load('checkpoint.pth')
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
epoch = checkpoint['epoch']
loss = checkpoint['loss']

3.2 预测结果全错的常见原因

当你的模型在训练集表现良好,但预测自己手写的数字却全错时,可能是以下原因:

  1. 预处理不一致 :训练时做了Normalize,预测时没做
  2. 颜色通道反转 :MNIST是黑底白字,而你画的是白底黑字
  3. 图像尺寸问题 :没有正确resize到28x28

修正方案:

# 预测前的正确预处理
def preprocess_custom_image(image):
    # 1. 反色处理(如果是白底黑字)
    image = 255 - image
    
    # 2. 归一化到[0,1]
    image = image / 255.0
    
    # 3. 应用与训练相同的Normalize
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    return transform(image)

4. 性能优化技巧

4.1 解决DataLoader卡住问题

当使用多进程DataLoader时,可能会出现程序卡住的情况。这是因为Windows和Linux/macOS的多进程实现方式不同。

解决方案:

# 设置num_workers为0(Windows)
train_loader = DataLoader(
    dataset=train_dataset,
    batch_size=64,
    shuffle=True,
    num_workers=0  # Windows下设为0
)

# 或者在主模块中添加保护
if __name__ == '__main__':
    # 你的训练代码

4.2 内存不足的应对策略

遇到CUDA out of memory错误时,可以尝试以下方法:

  1. 减小batch size(从256降到64或32)
  2. 使用梯度累积:多次小batch的前向后向,再更新参数
  3. 混合精度训练:减少显存占用
# 梯度累积示例
accumulation_steps = 4
optimizer.zero_grad()

for i, (inputs, labels) in enumerate(train_loader):
    outputs = model(inputs)
    loss = criterion(outputs, labels)
    loss = loss / accumulation_steps  # 归一化损失
    loss.backward()
    
    if (i+1) % accumulation_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

5. 调试与验证技巧

5.1 数据可视化检查

在训练前,先检查数据是否正确加载和预处理:

import matplotlib.pyplot as plt

# 显示一个batch的图像
images, labels = next(iter(train_loader))
img = images[0].numpy().squeeze()
label = labels[0].item()

plt.imshow(img, cmap='gray')
plt.title(f'Label: {label}')
plt.show()

5.2 验证模型是否真的在学习

训练初期,可以通过以下方法验证模型是否正常工作:

  1. 在训练集的一个小batch上过拟合(应该能达到100%准确率)
  2. 检查损失是否在下降
  3. 检查随机权重和训练后权重的变化
# 在小batch上过拟合测试
small_dataset = torch.utils.data.Subset(train_dataset, range(100))
small_loader = DataLoader(small_dataset, batch_size=10)

for epoch in range(20):
    for images, labels in small_loader:
        # 训练代码
    # 计算准确率
    # 应该很快达到100%

5.3 学习率与优化器选择

MNIST相对简单,但优化器选择仍会影响训练速度:

优化器 适用场景 典型学习率
SGD 简单任务 0.01-0.1
Adam 默认选择 0.001
AdamW 更稳定 0.001

学习率测试代码:

# 学习率范围测试
lr_finder = LRFinder(model, optimizer, criterion)
lr_finder.range_test(train_loader, end_lr=1, num_iter=100)
lr_finder.plot()  # 找到损失下降最快的区间
lr_finder.reset()
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐