MNIST数据集训练避坑指南:从数据加载到模型保存,PyTorch新手常踩的5个雷
MNIST数据集训练避坑指南:从数据加载到模型保存,PyTorch新手常踩的5个雷
第一次用PyTorch跑MNIST手写数字识别,本以为跟着教程敲代码就能轻松搞定,结果从数据加载到模型保存,处处是坑。明明代码和教程里一模一样,偏偏报错不断:CUDA内存不足、DataLoader卡住不动、模型保存后加载出错、预测结果全错……这些问题我都遇到过,今天就把这些坑一个个填平。
1. 数据加载的隐藏陷阱
1.1 下载速度慢到怀疑人生
很多教程直接使用 torchvision.datasets.MNIST 自动下载数据集,但国内用户经常会遇到下载速度极慢甚至失败的情况。这时候可以手动下载MNIST数据集,放到指定目录。
手动下载MNIST数据集步骤:
- 访问 MNIST官网 下载四个.gz文件
- 在项目目录下创建
data/MNIST/raw文件夹 - 将下载的文件放入raw文件夹
- 代码中设置
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"的错误。
解决方案:
- 统一设备设置
- 保存模型时记录设备信息
# 训练时
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 预测结果全错的常见原因
当你的模型在训练集表现良好,但预测自己手写的数字却全错时,可能是以下原因:
- 预处理不一致 :训练时做了Normalize,预测时没做
- 颜色通道反转 :MNIST是黑底白字,而你画的是白底黑字
- 图像尺寸问题 :没有正确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错误时,可以尝试以下方法:
- 减小batch size(从256降到64或32)
- 使用梯度累积:多次小batch的前向后向,再更新参数
- 混合精度训练:减少显存占用
# 梯度累积示例
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 验证模型是否真的在学习
训练初期,可以通过以下方法验证模型是否正常工作:
- 在训练集的一个小batch上过拟合(应该能达到100%准确率)
- 检查损失是否在下降
- 检查随机权重和训练后权重的变化
# 在小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()
更多推荐



所有评论(0)