ARTICLE DETAIL

资讯详情

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

PyTorch实现MNIST手写数字识别:从数据集到模型部署的完整指南

PyTorch实现MNIST手写数字识别:从数据集到模型部署的完整指南 简介一份基于机器学习方法的MNIST手写数字识别项目面向计算机相关专业的毕设、课设或机器学习入门者完整演示了SVM、决策树、KNN、朴素贝叶斯四种经典算法的实现与准确率对比。项目使用Python 3.6编写代码、数据集、结果图分目录存放结构清晰便于学习。压缩包共19个文件包含4个Python源码、2组MNIST标准数据集idx格式、8张准确率对比图与分类效果图另附已训练好的模型文件、决策树可视化dot文件及README说明文档整体大小约11.04MB。目前已有538人学习下载适合需要参考多算法手写数字识别完整流程、进行算法横向对比或作为课程设计/毕业设计基础的同学使用。1. 先定目标MNIST 手写数字识别到什么水平才算达标把 MNIST 手写数字识别做到 99% 以上在今天已经不是能力证明而是一套机器学习工程链路的体检报告。数据怎么进内存、模型参数怎么算、训练时 loss 为什么抖、验证集和测试集的区别在哪里这些真正决定项目能不能复现的问题全都会在这个看似简单的任务上暴露出来。这篇博文要解决的不只是给你一段能跑的源代码而是把基于机器学习方法的 MNIST 手写数字识别中从数据集下载到模型落地这一段路里所有常见的坑挨个踩平。无论你是刚开始写分类器还是在准备期末项目或者只是想找一个干净的数据集练手这条链路都值得完整走一遍。2. 数据是根MNIST 数据集下载、读取与预处理2.1 torchvision 下载 MNIST 报 404 的本地化方案MNIST 数据集以二进制 idx 格式分发torchvision 的datasets.MNIST类通过trainTrue/False区分训练集和测试集。理论上downloadTrue会自动从源站拉取四个.gz文件但实际运行时经常遇到torchvision下载mnist会404的报错根因是源站文件路径迁移、CDN 配置变更或学校、公司网络对境外域名访问不稳定。最可靠的做法是把下载环节从代码里抽出来手工准备原始文件。先在项目目录下建立固定结构mkdir -p data/MNIST/raw # 将以下四个文件放入 data/MNIST/raw 目录 # train-images-idx3-ubyte.gz # train-labels-idx1-ubyte.gz # t10k-images-idx3-ubyte.gz # t10k-labels-idx1-ubyte.gz这四个文件分别对应训练图像、训练标签、测试图像、测试标签。torchvision在downloadFalse时会检查data/MNIST/raw下是否存在同名文件存在就直接解压读取不再发起网络请求。我一般会把这一步写进 README避免团队成员重复踩同一个坑。2.2 归一化、批大小与两个 DataLoader拿到原始文件后读取逻辑完全交给datasets.MNIST。需要注意的是一份标准 transform。import torch import torchvision from torch.utils.data import DataLoader transform torchvision.transforms.Compose([ torchvision.transforms.ToTensor(), torchvision.transforms.Normalize((0.1307,), (0.3081,)) ]) train_set torchvision.datasets.MNIST( root./data, trainTrue, downloadFalse, transformtransform ) val_set torchvision.datasets.MNIST( root./data, trainFalse, downloadFalse, transformtransform )(0.1307,)和(0.3081,)是 MNIST 官方统计的像素均值与标准差单通道灰度图所以是一维。做归一化之后每个像素的分布近似为均值为 0、方差为 1 的标准正态分布这能显著加快梯度下降的收敛速度是机器学习入门阶段最容易漏掉的一步。DataLoader 的参数直接决定训练效率和随机性train_loader DataLoader( train_set, batch_size128, shuffleTrue, num_workers2, pin_memoryTrue ) val_loader DataLoader( val_set, batch_size256, shuffleFalse, num_workers2, pin_memoryTrue )训练集必须shuffleTrue否则每个 epoch 内样本顺序固定模型会按类别批次更新权重收敛曲线会出现周期性波动。验证集不需要 shuffle。pin_memoryTrue适用于 GPU 训练能把主机内存锁页减少 Host 到 Device 的拷贝时间纯 CPU 训练时pin_memory可以不开。2.2.1 transform 的取舍与数据增强MNIST 数字本身存在轻微形变训练集随机增强能提升泛化能力。但增强不是越猛越好9 和 6 旋转超过 20 度后人眼都无法区分模型强行学习这种样本只会降低上限。我一般对训练集使用 7 度以内的随机旋转加上 10% 以内的平移train_transform torchvision.transforms.Compose([ torchvision.transforms.RandomAffine(degrees7, translate(0.1, 0.1)), torchvision.transforms.ToTensor(), torchvision.transforms.Normalize((0.1307,), (0.3081,)) ])验证集和测试集只用归一化不做任何随机变换否则评估结果会随每次运行而抖动无法稳定对比模型好坏。三份数据的定位也不一样用途数据范围是否参与反向传播典型用途训练集60000 张是更新权重验证集10000 张否早停、调超参、选模型测试集10000 张否最后只跑一次的最终报告很多人拿到数据集后只拆 train 和 test忽略验证集最后用 test 反复调参测试集就成了验证集的替代品报告出来的准确率会虚高。标准的做法是 test 集只能碰一次。3. 模型与源代码从全连接网络到 CNN 的参数选择3.1 全连接基线784 维向量到 10 个类别的映射MNIST 单张图片是 28×28 的灰度图展平后是 784 维向量。最朴素的机器学习分类器是逻辑回归但工程上大家很少手写它而是直接用一个不带隐藏层的nn.Linear(784, 10)替代。要衡量模型结构带来的收益我会先搭一个带单隐藏层的全连接网络作为基线import torch.nn as nn class FCN(nn.Module): def __init__(self, hidden128, dropout0.2): super().__init__() self.net nn.Sequential( nn.Flatten(), nn.Linear(784, hidden), nn.BatchNorm1d(hidden), nn.ReLU(inplaceTrue), nn.Dropout(dropout), nn.Linear(hidden, 10) ) def forward(self, x): return self.net(x)hidden128是我习惯的起点。隐藏层宽度从 128 加到 512MNIST 准确率大约只能提升 0.2 到 0.3 个百分点但参数量从 10 万涨到 40 万训练时间和过拟合风险都在增加。BatchNorm1d 放在全连接层之后、ReLU 之前能让每层输入分布稳定这里model.train()和model.eval()切换时必须严格因为 BN 在训练时用 batch 统计量验证时用滚动均值忘掉eval()会导致验证准确率异常波动。3.2 CNN 的参数量与特征层设计全连接网络把每个像素单独对待忽略了相邻像素之间的空间关系。卷积神经网络通过共享卷积核用远少于全连接的参数量提取局部特征这也是手写数字识别接近饱和的常见做法。一个可以在 MNIST 上稳定跑到 99.2% 以上的结构如下import torch class CNN(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, 3, padding1), # (1,28,28) - (32,28,28) nn.ReLU(inplaceTrue), nn.Conv2d(32, 32, 3, padding1), # - (32,28,28) nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # - (32,14,14) nn.Conv2d(32, 64, 3, padding1), # - (64,14,14) nn.ReLU(inplaceTrue), nn.Conv2d(64, 64, 3, padding1), # - (64,14,14) nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # - (64,7,7) ) self.classifier nn.Linear(64 * 7 * 7, 10) def forward(self, x): x self.features(x) return self.classifier(torch.flatten(x, 1))注释里的 shape 变化是调试 CNN 最重要的参照。输入单通道 28×28第一层卷积因为padding1保持尺寸不变MaxPool2d 把空间尺寸减半经过两次池化后从 28×28 降到 7×7通道数从 1 扩到 64展平后得到 3136 维特征向量。卷积层的参数量计算公式是输入通道 × 输出通道 × 核宽 × 核高 输出通道。第一层 Conv2d 的参数只有1×32×3×332320个最后一层 Linear 反而占了大头。整体对比网络可训练参数量MNIST 期望 test acc适用理由Linear 784→107850约 92%基线参考FCN hidden128约 10 万97.5% ~ 98.5%特征线性组合CNN 上图约 6.6 万99.2% ~ 99.4%空间局部特征 参数共享CNN 参数量只有全连接网络的六成准确率反而高出一个多百分点这就是卷积核共享权值带来的结构性优势。用sum(p.numel() for p in model.parameters())可以随时打印实际参数量不需要靠估算。3.3 损失函数、优化器与梯度穿过的路径多分类问题标准选择是nn.CrossEntropyLoss()。这个损失函数内部对网络输出做 log-softmax再接负对数似然标签用整数索引即可不需要手动做 one-hot 编码。相比 MSE交叉熵在分类问题上梯度形状更合理不会因为 softmax 输出接近 0 或 1 时梯度消失。优化器我推荐从 Adam 起步optimizer torch.optim.Adam(model.parameters(), lr1e-3, weight_decay1e-4)Adam 的默认学习率在 MNIST 这种小数据集上表现稳定。weight_decay等价于在损失函数后面加 L2 正则项能轻微约束权重幅度。追求更高精度时再换SGD(lr0.01, momentum0.9)配余弦退火两者在孩子数据集上的差距通常不到 0.1%不值得在入门阶段花太多时间纠结。梯度从损失函数出发依次穿过最后一层 Linear、展平操作、池化层、卷积层。ReLU 的作用是给网络注入非线性同时避免 sigmoid/tanh 在深层网络中容易出现的梯度衰减如果发现训练 loss 长时间不降第一步就应该检查网络深处是不是堆叠了过多非线性层且缺少 BN。4. 训练与评估MNIST 识别的收敛、过拟合与指标4.1 最小训练循环的骨架代码训练一个 epoch 的标准循环可以收敛为以下函数后续换模型、换数据集只需要改model和loader的入口。def train_epoch(model, loader, criterion, optimizer, device): model.train() total_loss, correct, total 0.0, 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) correct (outputs.argmax(dim1) labels).sum().item() total labels.size(0) return total_loss / total, correct / totaloutputs.argmax(dim1)是逐样本取 10 个类别中得分最高的类别对应模型预测结果。loss.item()取出标量值用于累计但注意.item()会中断梯度图所以必须在backward()之后调用。images.size(0)是当前 batch 的样本数用它对 loss 做加权平均避免最后一个 batch 样本数不满时统计失真。验证函数必须关闭梯度torch.no_grad() def evaluate(model, loader, device): model.eval() correct, total 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) correct (outputs.argmax(dim1) labels).sum().item() total labels.size(0) return correct / totaltorch.no_grad()告诉框架不需要为验证阶段构建计算图内存占用降低推理速度也能快 20% 以上。每个 epoch 结束后先跑验证再决定是否保存模型是训练脚本的基本盘。4.2 早停、学习率调度与模型保存MNIST 的 CNN 在 5 万个训练样本上不容易严重过拟合但验证 loss 依然会经历先降后升的过程。我用ReduceLROnPlateau做学习率衰减用验证 loss 做早停scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.5, patience3, verboseTrue ) best_loss float(inf) patience_counter 0 for epoch in range(30): train_loss, train_acc train_epoch(model, train_loader, criterion, optimizer, device) val_acc evaluate(model, val_loader, device) # 用一个完整的验证集 loss 计算早停状态 val_loss compute_val_loss(model, criterion, val_loader, device) scheduler.step(val_loss) if val_loss best_loss - 1e-4: best_loss val_loss patience_counter 0 torch.save(model.state_dict(), best_mnist.pth) else: patience_counter 1 if patience_counter 8: break print(fepoch{epoch} train_acc{train_acc:.4f} val_acc{val_acc:.4f} lr{optimizer.param_groups[0][lr]:.2e})patience3表示连续 3 个 epoch 验证 loss 不下降就把学习率减半。1e-4的容差避免微小抖动触发送药。早停条件设 8 个 epoch 不改善就终止并把最佳权重保存在磁盘上这样即使训练过程被中断也保留最优状态。学习率调度是整个训练过程里最值得盯的参数。如果你发现验证准确率在 98% 左右反复横跳八成不是模型结构问题而是学习率过大导致参数在最优解附近震荡。这时候观察打印出的lr是否真的在每 3 个 epoch 后减半比怀疑数据增强更有效。4.3 用 classification_report 和混淆矩阵看真实泛化能力准确率对 MNIST 这种类别均衡的数据集够用但要看模型具体在哪些数字上犯错必须落到每个类别的精确率和召回率from sklearn.metrics import classification_report, confusion_matrix import numpy as np y_true, y_pred [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: outputs model(images) y_true.extend(labels.numpy()) y_pred.extend(outputs.argmax(dim1).numpy()) print(classification_report(y_true, y_pred, digits4)) cm confusion_matrix(y_true, y_pred) print(cm)classification_report会输出每个数字的 precision、recall、f1-score。MNIST 上最容易出现的混淆对是 4/9、3/8、7/9这些数字外形相似正是 CNN 最后那 0.5% 准确率的主要失分点。混淆矩阵的对角线数字代表正确分类数量非对角线cm[i][j]表示真实类别 i 被预测成 j 的样本数哪两个数字互相干扰一眼就能看出来。5. 最后 1%误分类分析、鲁棒性测试与 ONNX 导出5.1 把错题本打印出来看模型在什么情况下犯错验证集上误分类样本通常集中在两种一种是本身潦草到人眼都难分辨另一种是数据标注错误。我的做法是筛出预测置信度最高的前 20 个错分样本排成网格逐一对照真实标签def collect_misclassified(model, loader, device, k20): model.eval() mistakes [] with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) probs torch.softmax(model(images), dim1) pred probs.argmax(dim1) for i in range(images.size(0)): if pred[i] ! labels[i]: mistakes.append((images[i], labels[i], pred[i], probs[i].max().item())) mistakes.sort(keylambda x: x[3], reverseTrue) return mistakes[:k]置信度最高的 10 个错分样本如果大部分是标注问题说明模型已经学到了数据分布里真正可分的部分如果高置信度样本本身字形清楚那就是模型结构或增强策略出问题了优先检查训练集与验证集的预处理是否一致。5.2 用旋转和偏移压测鲁棒性反推数据增强是否有效验证集只跑一次是纪律但在报告之前可以额外做一次鲁棒性压测。把验证集整体旋转 15 度再评估如果准确率从 99.3% 掉到 85% 以下说明模型对旋转敏感训练时的RandomAffine(degrees7)还不够。我一般会准备两套极端测试旋转 15 度、平移 0.2 个图像宽度对比同一模型在原始验证集和扰动验证集上的准确率差值。这个差值就是模型的鲁棒性裕度。差值小于 1% 说明增强策略合理差值过大先把增强角度调大再训练一轮。5.3 导出 ONNX在 C# 侧用 ONNX Runtime 做推理训练完成不代表项目结束。如果业务侧是 C# 服务常见做法不是用 C# 重写网络结构而是把 PyTorch 模型导出成 ONNX再用 ONNX Runtime 加载model.eval() dummy torch.randn(1, 1, 28, 28) torch.onnx.export( model, dummy, mnist.onnx, input_names[images], output_names[logits], dynamic_axes{images: {0: batch}} )dynamic_axes把 batch 维度设为动态这样 C# 侧可以任意指定一次推理的图片数量不需要模型输入固定为 1。C# 里只需要InferenceSession(mnist.onnx)就能加载并推理后续所有更新都集中在训练侧部署侧代码几乎不用动。把旋转测试、误分类抽样这些脚本写进 CI每次模型更新后自动跑一遍对比新旧版本在错题样本上的表现这才是 MNIST 手写数字识别项目从演示走向可维护的关键一步。本文还有配套的精品资源点击获取
返回列表