拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch实战:MNIST手写数字识别CNN模型从数据加载到99.5%准确率

PyTorch实战:MNIST手写数字识别CNN模型从数据加载到99.5%准确率

简介:这份资源面向深度学习入门者与计算机视觉初学者,围绕MNIST手写数字识别这一经典基准任务,提供从数据理解到卷积神经网络训练落地的完整实践材料。包内共6个文件,以2个Python脚本、2张PNG图表、1份TXT说明和1个H5权重文件为主:脚本对应CNN模型的构建、训练与评估流程,图表分别呈现损失与准确率曲线以及测试样本预测效果,权重文件可直接加载已训练模型,省去重复训练开销,说明文档则交代运行方式与使用要点。压缩包约2.19MB,轻量易取。目前已有4620人学习下载,热度较高。读者可借此掌握MNIST数据预处理、卷积与池化特征提取、模型训练监控及权重保存等关键环节,并对照预测结果与曲线图排查欠拟合、学习率设置等问题,快速完成一次可复现的图像分类实战。

1. MNIST 手写数字识别:从数据集到 CNN 模型落地的完整路径

MNIST 手写数字识别几乎是每个深度学习入门者绕不开的第一个实战项目,也是卷积神经网络最经典的教学场景。但很多人卡在第一步:torchvision 下载 MNIST 会 404,数据集拿不到,后面的一切都无从谈起。这篇文章要解决的就是这个问题——从数据集的获取与理解开始,一步步用 CNN 搭建一个能跑通、能收敛、能保存模型文件的手写数字识别系统。不管你是刚配好 PyTorch 环境的新手,还是想找一个干净 baseline 的熟手,这里给出的代码和参数配置都可以直接复现。整个方案不依赖任何在线下载,数据集和训练好的模型文件都会给出明确的落地方式,训练完成后模型可以直接用于推理。

2. MNIST 数据集:结构、加载与离线获取方案

2.1 MNIST 到底长什么样

MNIST 全称 Modified National Institute of Standards and Technology database,包含 70000 张灰度手写数字图片,其中 60000 张为训练集,10000 张为测试集。每张图片是 28×28 像素,像素值范围 0 到 255,对应数字 0 到 9 共十个类别。训练集来自 250 个不同书写者,测试集来自另外 250 个书写者,这种书写者不重叠的设计保证了测试集能真实反映模型的泛化能力。

从数据分布来看,十个类别的样本数量大致均衡,训练集中每个数字大约 6000 张左右,测试集中每个数字大约 1000 张。这个均衡性意味着你不需要额外做类别加权或重采样,直接训练就能得到比较稳定的结果。但要注意,MNIST 的图片已经过尺寸归一化和居中处理,数字基本位于图像中央,这降低了识别难度,也意味着在 MNIST 上表现好的模型未必能直接迁移到真实场景的手写识别任务上。

数据集的原始格式是 IDX 文件,训练图像文件名为 train-images-idx3-ubyte,训练标签为 train-labels-idx1-ubyte,测试集对应 t10k-images-idx3-ubyte 和 t10k-labels-idx1-ubyte。IDX 格式的头部包含魔数和维度信息,读取时需要先解析头部再按字节偏移取像素数据。不过用 PyTorch 的话,torchvision.datasets.MNIST 已经封装好了这些细节,你只需要关注根目录和 download 参数。

2.2 torchvision 下载 404 的根因与离线加载

很多人第一次用 torchvision 下载 MNIST 时会遇到 404 错误,这不是你的网络问题,而是 torchvision 默认的下载源地址已经失效。原始下载链接指向的是 yann.lecun.com 的旧路径,该路径早已不可访问。torchvision 在新版本中虽然更新了镜像地址,但部分版本仍然会命中失效链接。

解决思路有三种:第一种是手动下载四个 IDX 文件放到指定目录,然后设置 download=False;第二种是修改 torchvision 的下载 URL 指向可用镜像;第三种是直接用本地文件读取,绕过 torchvision 的下载逻辑。我一般推荐第一种,最稳定也最可控。

手动下载后,目录结构必须严格符合 torchvision 的预期:

# MNIST 数据集的目录结构要求 # 根目录/ # └── MNIST/ # └── raw/ # ├── train-images-idx3-ubyte # ├── train-labels-idx1-ubyte # ├── t10k-images-idx3-ubyte # └── t10k-labels-idx1-ubyte

把四个文件放到raw/目录下之后,用下面的代码加载:

import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 定义预处理:转为张量并归一化 transform = transforms.Compose([ transforms.ToTensor(), # 将 PIL 图像转为 [0,1] 范围的张量,形状 [1,28,28] transforms.Normalize((0.1307,), (0.3081,)) # MNIST 全局均值和标准差 ]) # 加载训练集,download=False 表示使用本地文件 train_dataset = datasets.MNIST( root='./data', # 数据集根目录 train=True, # 训练集 transform=transform, # 预处理 download=False # 关键:不触发在线下载 ) # 加载测试集 test_dataset = datasets.MNIST( root='./data', train=False, transform=transform, download=False ) # 创建 DataLoader train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False) print(f"训练集样本数: {len(train_dataset)}") print(f"测试集样本数: {len(test_dataset)}")

这里的Normalize((0.1307,), (0.3081,))两个参数是 MNIST 训练集的全局像素均值和标准差,这是社区公认的经验值。归一化的目的是让输入数据分布更接近标准正态,加速收敛。如果你不做归一化,模型也能训练,但收敛速度会慢不少,而且对学习率更敏感。

batch_size设为 64 是训练时的常用值,测试时设为 1000 可以一次性处理更多样本,减少评估时间。shuffle=True只在训练集上开启,测试集不需要打乱。

2.3 数据可视化与质量检查

在正式训练之前,建议先做一次数据可视化,确认加载的数据没有错位、标签和图像对应正确。这一步花不了两分钟,但能帮你排除掉很多低级错误。

import matplotlib.pyplot as plt import numpy as np # 取一个 batch 的数据 images, labels = next(iter(train_loader)) # 显示前 8 张图片 fig, axes = plt.subplots(2, 4, figsize=(10, 5)) for i, ax in enumerate(axes.flat): # images[i] 形状是 [1,28,28],squeeze 去掉通道维度 ax.imshow(images[i].squeeze(), cmap='gray') ax.set_title(f'Label: {labels[i].item()}') ax.axis('off') plt.tight_layout() plt.show() # 检查数据分布 unique, counts = np.unique(train_dataset.targets.numpy(), return_counts=True) for u, c in zip(unique, counts): print(f"数字 {u}: {c} 张")

如果可视化出来的图片是清晰的数字,标签也对应正确,说明数据加载没问题。如果图片是全黑或者噪声,大概率是归一化参数用错了,或者 IDX 文件损坏。另外注意images[i].squeeze()这一步,因为 DataLoader 返回的张量形状是[batch, channel, height, width],单张图片需要去掉 channel 维度才能用imshow正常显示。

3. 卷积神经网络设计:从 LeNet 到适配 MNIST 的现代结构

3.1 为什么 CNN 比全连接网络更适合图像

手写数字识别本质上是一个图像分类任务。如果用全连接网络处理 28×28 的图片,首先要把图像展平成 784 维向量,这一步就丢失了像素之间的空间关系。相邻像素的关联信息在展平后被彻底打散,网络只能靠权重矩阵去重新学习这些关系,参数量大且效率低。

卷积神经网络通过三个核心机制解决这个问题:局部连接、权重共享和池化。局部连接意味着每个卷积核只关注输入的一个小区域,比如 3×3 或 5×5,这符合图像中相邻像素高度相关的先验知识。权重共享意味着同一个卷积核在整张图上滑动,检测相同的特征模式,这大幅减少了参数量。池化则通过下采样降低特征图的空间尺寸,同时保留最显著的特征响应,增强平移不变性。

具体到 MNIST,一个简单的 CNN 通常包含两到三个卷积层、两到三个池化层,最后接全连接层输出 10 类概率。卷积层负责提取边缘、角点、笔画等局部特征,越深的卷积层提取的特征越抽象。池化层逐步降低空间分辨率,让网络从关注“哪个位置有什么笔画”过渡到“整体是什么数字”。

3.2 一个可复现的 CNN 结构定义

下面这个网络结构是我在 MNIST 上反复用过的,结构不复杂但效果稳定,测试集准确率能到 99% 以上。整个网络由两个卷积块和一个分类头组成。

import torch.nn as nn import torch.nn.functional as F class MNISTCNN(nn.Module): def __init__(self): super(MNISTCNN, self).__init__() # 第一个卷积块:输入 1 通道,输出 32 通道,卷积核 3x3 self.conv1 = nn.Conv2d(in_channels=1, out_channels=32, kernel_size=3, padding=1) self.bn1 = nn.BatchNorm2d(32) # 批归一化,加速收敛 self.conv2 = nn.Conv2d(in_channels=32, out_channels=32, kernel_size=3, padding=1) self.bn2 = nn.BatchNorm2d(32) self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2) # 28x28 -> 14x14 # 第二个卷积块:32 通道 -> 64 通道 self.conv3 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, padding=1) self.bn3 = nn.BatchNorm2d(64) self.conv4 = nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3, padding=1) self.bn4 = nn.BatchNorm2d(64) self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2) # 14x14 -> 7x7 # 全连接分类头 self.fc1 = nn.Linear(64 * 7 * 7, 128) self.dropout = nn.Dropout(0.5) # 防止过拟合 self.fc2 = nn.Linear(128, 10) # 输出 10 个类别 def forward(self, x): # 第一个卷积块 x = F.relu(self.bn1(self.conv1(x))) x = F.relu(self.bn2(self.conv2(x))) x = self.pool1(x) # 第二个卷积块 x = F.relu(self.bn3(self.conv3(x))) x = F.relu(self.bn4(self.conv4(x))) x = self.pool2(x) # 展平后送入全连接层 x = x.view(x.size(0), -1) # [batch, 64*7*7] x = F.relu(self.fc1(x)) x = self.dropout(x) x = self.fc2(x) return x # 实例化并查看参数量 model = MNISTCNN() total_params = sum(p.numel() for p in model.parameters()) print(f"模型总参数量: {total_params:,}")

这个结构的设计逻辑是这样的:第一个卷积块用两个 3×3 卷积层堆叠,感受野等效于一个 5×5 卷积,但参数量更少且非线性更强。每个卷积层后面接 BatchNorm 和 ReLU,BatchNorm 的作用是稳定训练过程中的分布,允许使用更大的学习率。池化层用 2×2 最大池化,每次把空间尺寸减半。

第二个卷积块把通道数从 32 提升到 64,增加特征表达能力。经过两次池化后,特征图尺寸从 28×28 降到 7×7,通道数为 64,所以展平后是 64×7×7=3136 维。全连接层先降到 128 维,再输出 10 维。Dropout 设 0.5 是为了防止全连接层过拟合,因为 3136 到 128 的权重矩阵参数量很大。

整个模型参数量大约在 42 万左右,对于 MNIST 来说绰绰有余。如果你想要更轻量的模型,可以减少通道数或者去掉一个卷积层,但准确率可能会下降零点几个百分点。

3.3 训练循环与关键参数设置

训练循环的代码看起来模板化,但里面的参数设置直接决定模型能不能收敛、收敛多快。下面这份训练代码是我常用的配置。

import torch.optim as optim from torch.optim.lr_scheduler import StepLR # 设备选择 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = MNISTCNN().to(device) # 优化器:AdamW,学习率 1e-3,权重衰减 1e-4 optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4) # 学习率调度:每 5 个 epoch 学习率乘以 0.5 scheduler = StepLR(optimizer, step_size=5, gamma=0.5) # 损失函数:交叉熵 criterion = nn.CrossEntropyLoss() # 训练参数 epochs = 15 best_acc = 0.0 for epoch in range(epochs): model.train() # 切换到训练模式 running_loss = 0.0 correct = 0 total = 0 for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) optimizer.zero_grad() # 清空梯度 output = model(data) # 前向传播 loss = criterion(output, target) # 计算损失 loss.backward() # 反向传播 optimizer.step() # 更新参数 running_loss += loss.item() pred = output.argmax(dim=1) # 取概率最大的类别 correct += pred.eq(target).sum().item() total += target.size(0) scheduler.step() # 更新学习率 train_acc = correct / total avg_loss = running_loss / len(train_loader) # 每个 epoch 结束后在测试集上评估 model.eval() test_correct = 0 test_total = 0 with torch.no_grad(): for data, target in test_loader: data, target = data.to(device), target.to(device) output = model(data) pred = output.argmax(dim=1) test_correct += pred.eq(target).sum().item() test_total += target.size(0) test_acc = test_correct / test_total print(f"Epoch {epoch+1}/{epochs} | Loss: {avg_loss:.4f} | " f"Train Acc: {train_acc:.4f} | Test Acc: {test_acc:.4f}") # 保存最佳模型 if test_acc > best_acc: best_acc = test_acc torch.save(model.state_dict(), 'best_mnist_cnn.pth') print(f" -> 模型已保存,当前最佳准确率: {best_acc:.4f}") print(f"训练完成,最佳测试准确率: {best_acc:.4f}")

几个关键参数的解释:优化器选 AdamW 而不是 SGD,是因为 AdamW 对学习率不那么敏感,新手不用花太多时间调参就能得到不错的结果。学习率 1e-3 是 Adam 系列在 MNIST 上的常用起点,配合 StepLR 每 5 个 epoch 减半,能让模型在后期精细调整权重。权重衰减 1e-4 是轻量正则化,配合 Dropout 一起抑制过拟合。

model.train()和model.eval()的切换很重要。训练模式下 BatchNorm 使用当前 batch 的统计量,Dropout 会随机丢弃神经元;评估模式下 BatchNorm 使用全局统计量,Dropout 关闭。如果你忘了切换,测试结果会不稳定。

保存模型时用state_dict()而不是整个模型对象,这样加载时更灵活,也不依赖原始的类定义路径。最佳模型的保存策略是每次测试准确率提升就覆盖保存,这样训练结束后你拿到的一定是测试集上表现最好的那个版本。

4. 训练过程排查:损失不降、准确率震荡与过拟合的应对

4.1 损失不下降的常见原因

训练时最让人头疼的就是 loss 一动不动,或者下降极慢。根据我的踩坑经验,原因通常集中在几个地方。

第一,学习率太大或太小。学习率太大会导致 loss 震荡甚至发散,表现为 loss 变成 NaN 或者在一个高值附近反复跳。学习率太小则 loss 下降极慢,十几个 epoch 过去还在高位。判断方法是打印每个 batch 的 loss,如果前几个 batch 的 loss 就异常大或者 NaN,基本可以确定是学习率问题。解决方法是把学习率降到 1e-4 或 1e-5 再试。

第二,数据没有归一化或者归一化参数错误。MNIST 的像素值原始范围是 0 到 255,如果不做归一化直接输入网络,梯度会非常大,训练很难稳定。用Normalize((0.1307,), (0.3081,))是经过验证的配置,不要随意改。

第三,标签和图像错位。这种情况比较隐蔽,loss 会下降但准确率上不去,或者准确率在随机水平附近徘徊。用前面给的可视化代码检查一下,确认图像和标签对应正确。

第四,网络结构有 bug。比如卷积层输出通道数和下一层输入通道数不匹配,或者全连接层的输入维度算错了。PyTorch 会在前向传播时报错,但如果你用了view或reshape且维度算错,可能不会报错但结果完全错误。建议在定义完模型后,用一个随机张量做一次前向传播,检查输出形状是否符合预期。

4.2 测试准确率震荡的排查思路

测试准确率在每个 epoch 之间上下跳动是正常现象,但如果震荡幅度超过 1% 就需要排查了。

一个常见原因是 BatchNorm 在训练和评估时的行为差异。如果训练集和测试集的分布差异较大,BatchNorm 的全局统计量可能不准确。MNIST 的训练集和测试集来自不同书写者,分布本身有差异,但通常不会导致大幅震荡。如果你自己划分了验证集,确保验证集的预处理和训练集完全一致。

另一个原因是学习率调度不合理。StepLR 的 step_size 和 gamma 需要根据总 epoch 数调整。如果总 epoch 只有 10 个,step_size 设为 5 意味着只衰减一次,后期学习率可能还是偏大。可以改成 CosineAnnealingLR,让学习率平滑下降。

还有一个容易被忽略的点是 DataLoader 的 shuffle。训练集必须开启 shuffle,否则每个 batch 的样本分布可能严重偏斜,导致梯度方向不稳定。测试集不需要 shuffle,但如果你在测试时也开了 shuffle,评估结果本身不会变,只是顺序不同。

4.3 过拟合与欠拟合的判断和处理

过拟合的典型表现是训练准确率持续上升,但测试准确率停滞甚至下降。在 MNIST 上,如果训练准确率到了 99.9% 而测试准确率只有 98%,说明模型开始记住训练样本的细节了。

处理过拟合的手段按优先级排序:增加 Dropout 比例、增加权重衰减、减少模型参数量、增加数据增强。Dropout 从 0.5 提到 0.6 或 0.7 通常有效,但太高会导致欠拟合。权重衰减从 1e-4 提到 1e-3 也可以,但要注意不要压得太狠。数据增强方面,MNIST 可以做的有随机旋转 ±10 度、随机平移 10% 以内、随机缩放 0.9 到 1.1 倍。这些增强模拟了手写数字的自然变化,能显著提升泛化能力。

欠拟合的表现是训练准确率和测试准确率都低,且 loss 下降缓慢。原因通常是模型容量不够或者训练不够充分。可以增加卷积层通道数、增加全连接层宽度、延长训练 epoch 数。MNIST 本身比较简单,一个两层的 CNN 就足够,欠拟合更多出现在你用了过于简单的全连接网络时。

4.4 显存不足与低显存运行技巧

虽然 MNIST 模型很小,但如果你同时开了多个实验或者用了很大的 batch_size,也可能遇到显存不足。降低 batch_size 是最直接的方法,从 64 降到 32 或 16。另外可以在训练循环中加torch.cuda.empty_cache(),但这会拖慢训练速度,不建议频繁调用。

如果显存实在紧张,可以把模型放到 CPU 上训练。MNIST 的模型在 CPU 上训练一个 epoch 大概几十秒,15 个 epoch 也就十几分钟,完全可以接受。用torch.device('cpu')即可,代码不需要其他改动。

还有一个技巧是用混合精度训练,PyTorch 的torch.cuda.amp可以自动把部分计算转为 float16,显存占用能减少一半左右。但对于 MNIST 这种小模型,收益不大,反而增加了代码复杂度,不太推荐。

5. 模型保存、加载与推理:让训练结果真正可用

5.1 保存格式的选择与加载验证

训练完成后,你手里会有一个best_mnist_cnn.pth文件。这个文件只包含模型的参数,不包含网络结构。加载时需要先实例化模型类,再加载参数。

# 加载模型进行推理 model = MNISTCNN() # 先实例化结构 model.load_state_dict(torch.load('best_mnist_cnn.pth', map_location='cpu')) model.eval() # 切换到评估模式 # 验证加载是否正确 correct = 0 total = 0 with torch.no_grad(): for data, target in test_loader: output = model(data) pred = output.argmax(dim=1) correct += pred.eq(target).sum().item() total += target.size(0) print(f"加载后模型测试准确率: {correct/total:.4f}")

map_location='cpu'的作用是让模型可以在没有 GPU 的机器上加载。如果你在 GPU 上训练并保存,在 CPU 上加载时不加这个参数会报错。加上之后,模型参数会被映射到 CPU 上。

加载后一定要用测试集验证一遍准确率,确认和保存时的最佳准确率一致。如果不一致,可能是保存时用了model.state_dict()但加载时用了model.load_state_dict()之外的方式,或者模型结构定义有变化。

5.2 单张图片推理的完整流程

实际使用模型时,你面对的是一张张单独的图片,而不是打包好的 DataLoader。单张推理需要手动做预处理,确保输入格式和训练时一致。

from PIL import Image import torchvision.transforms as transforms def predict_single_image(image_path, model, device='cpu'): """对单张手写数字图片进行预测""" # 预处理:和训练时保持一致 transform = transforms.Compose([ transforms.Grayscale(num_output_channels=1), # 确保单通道 transforms.Resize((28, 28)), # 统一尺寸 transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载图片 image = Image.open(image_path) image_tensor = transform(image).unsqueeze(0) # 增加 batch 维度 image_tensor = image_tensor.to(device) # 推理 model.eval() with torch.no_grad(): output = model(image_tensor) probabilities = torch.softmax(output, dim=1) predicted = output.argmax(dim=1).item() confidence = probabilities[0][predicted].item() return predicted, confidence # 使用示例 pred, conf = predict_single_image('test_digit.png', model) print(f"预测数字: {pred}, 置信度: {conf:.4f}")

这里有几个容易翻车的点。第一,unsqueeze(0)必须加,因为模型期望的输入形状是[batch, channel, height, width],单张图片只有[channel, height, width],不加 batch 维度会报错。第二,预处理必须和训练时完全一致,包括 Grayscale、Resize、ToTensor 和 Normalize 的顺序和参数。第三,如果图片是白底黑字,而 MNIST 是黑底白字,需要做反色处理,否则模型会把数字识别成背景。

5.3 批量推理与结果导出

如果你有一批图片需要批量识别,可以复用 DataLoader 的逻辑,或者手动组 batch。

import os from torch.utils.data import Dataset, DataLoader class CustomImageDataset(Dataset): def __init__(self, image_dir, transform=None): self.image_dir = image_dir self.transform = transform self.image_files = [f for f in os.listdir(image_dir) if f.endswith(('.png', '.jpg', '.jpeg'))] def __len__(self): return len(self.image_files) def __getitem__(self, idx): img_path = os.path.join(self.image_dir, self.image_files[idx]) image = Image.open(img_path) if self.transform: image = self.transform(image) return image, self.image_files[idx] # 批量推理 batch_transform = transforms.Compose([ transforms.Grayscale(num_output_channels=1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) dataset = CustomImageDataset('my_digits/', transform=batch_transform) loader = DataLoader(dataset, batch_size=32, shuffle=False) results = [] model.eval() with torch.no_grad(): for images, filenames in loader: outputs = model(images) preds = outputs.argmax(dim=1) for fname, pred in zip(filenames, preds): results.append((fname, pred.item())) # 导出结果 with open('predictions.csv', 'w') as f: f.write('filename,prediction\n') for fname, pred in results: f.write(f'{fname},{pred}\n') print(f"共处理 {len(results)} 张图片,结果已保存到 predictions.csv")

批量推理的关键是保持预处理一致,并且注意文件名和预测结果的对应关系。导出 CSV 时建议用 UTF-8 编码,避免中文文件名乱码。

6. 把 MNIST 模型推到 99.5% 以上:数据增强与集成技巧

6.1 数据增强的实操配置

MNIST 的测试集准确率从 99% 提到 99.5% 以上,数据增强是最有效的手段之一。手写数字的自然变化包括轻微旋转、平移、缩放和笔画粗细变化,用 torchvision 的 transforms 可以模拟这些变化。

from torchvision import transforms # 训练集增强配置 train_transform = transforms.Compose([ transforms.RandomAffine( degrees=10, # 随机旋转 ±10 度 translate=(0.1, 0.1), # 随机平移 ±10% scale=(0.9, 1.1), # 随机缩放 0.9 到 1.1 倍 shear=5 # 随机错切 ±5 度 ), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 测试集不做增强,只做归一化 test_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])

RandomAffine的参数需要控制幅度。旋转超过 ±15 度可能会让 6 和 9 变得难以区分,平移超过 15% 可能把数字移出画面。缩放范围 0.9 到 1.1 是安全的,错切 5 度以内也比较温和。这些参数是我试过多次后觉得比较平衡的配置,再大就容易伤害准确率了。

用了数据增强后,训练 epoch 数需要相应增加,因为每个 epoch 看到的样本都是不同的变体。原来 15 个 epoch 可能不够,建议加到 20 到 25 个。同时学习率可以稍微调大一点,因为增强后的数据分布更广,模型需要更强的更新力度。

6.2 模型集成:多个模型的投票策略

单个模型到了 99.4% 左右再往上提就很困难了,这时候可以用模型集成。最简单的做法是训练多个结构不同或初始化不同的模型,推理时对它们的输出概率取平均。

def ensemble_predict(models, data_loader, device='cpu'): """多个模型投票预测""" all_probs = [] for model in models: model.eval() model_probs = [] with torch.no_grad(): for data, _ in data_loader: data = data.to(device) output = model(data) probs = torch.softmax(output, dim=1) model_probs.append(probs) all_probs.append(torch.cat(model_probs, dim=0)) # 平均所有模型的概率 avg_probs = torch.stack(all_probs, dim=0).mean(dim=0) return avg_probs.argmax(dim=1) # 假设你训练了 3 个模型 model1 = MNISTCNN() model1.load_state_dict(torch.load('model1.pth', map_location='cpu')) model2 = MNISTCNN() model2.load_state_dict(torch.load('model2.pth', map_location='cpu')) model3 = MNISTCNN() model3.load_state_dict(torch.load('model3.pth', map_location='cpu')) # 集成预测 predictions = ensemble_predict([model1, model2, model3], test_loader)

集成的效果取决于模型之间的差异性。如果三个模型结构完全一样、初始化种子也一样,集成不会带来提升。要保证差异性,可以换不同的随机种子、不同的网络结构(比如一个用 3×3 卷积,一个用 5×5 卷积)、或者不同的数据增强配置。三个模型的集成通常能把准确率从 99.4% 提到 99.5% 到 99.6%,再往上就需要更多模型或者更复杂的集成策略了。

6.3 学习率预热与余弦退火

学习率调度对最终准确率的影响经常被低估。StepLR 简单有效,但 CosineAnnealingLR 配合 Warmup 往往能拿到更好的结果。

from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4) # 前 2 个 epoch 线性预热,之后余弦退火 warmup = LinearLR(optimizer, start_factor=0.1, total_iters=2) cosine = CosineAnnealingLR(optimizer, T_max=18, eta_min=1e-5) scheduler = SequentialLR(optimizer, schedulers=[warmup, cosine], milestones=[2])

预热的目的是让模型在训练初期不要因为学习率太大而震荡。start_factor=0.1表示从 1e-4 开始,线性增加到 1e-3。余弦退火则让学习率按余弦曲线平滑下降到 1e-5,避免 StepLR 那种突然减半带来的损失波动。

这套调度策略配合数据增强,在 MNIST 上跑到 99.6% 以上是比较稳的。但要注意,准确率到了这个级别,每次训练的波动可能在 0.05% 左右,不要因为一次没到 99.6% 就反复调参,先检查数据增强和调度配置是否一致。

6.4 我踩过的坑和最后几条建议

第一个坑是数据增强用在了测试集上。测试集必须用和训练时相同的归一化,但绝对不能加随机旋转和平移,否则评估结果没有意义。我见过有人把增强配置直接套在测试集上,准确率掉到 90% 以下还找不到原因。

第二个坑是模型集成时忘了切换 eval 模式。如果模型还在 train 模式,BatchNorm 会用当前 batch 的统计量,Dropout 也会随机丢弃,集成结果的方差会很大。每个模型在推理前都要model.eval()。

第三个坑是保存模型时只保存了 state_dict,但加载时用了torch.load直接加载整个模型对象。如果模型类定义发生了变化,加载会失败。统一用 state_dict 是最稳妥的做法。

最后一个建议:MNIST 是一个入门数据集,在上面刷到 99.6% 已经接近上限了。如果你要做真实场景的手写数字识别,建议尽早换到 EMNIST 或 USPS 数据集上验证,那些数据集的分布更接近实际应用。MNIST 上表现好的模型不一定能直接迁移,但训练和调参的思路是通用的。希望帮到你。

本文还有配套的精品资源,点击获取

返回列表