恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
PyTorch实战:用CNN训练MNIST手写数字识别与模型保存
首页
资讯中心
/
PyTorch实战:用CNN训练MNIST手写数字识别与模型保存
PyTorch实战:用CNN训练MNIST手写数字识别与模型保存
发布时间:2026/10/10 14:50:56
简介MNIST手写数字识别是深度学习的经典入门场景。这一压缩包面向正在学习神经网络与TensorFlow的开发者完整展示了如何用卷积神经网络CNN训练手写数字识别模型涵盖数据预处理、模型搭建、训练评估与保存调用等环节。资源共6个文件以Python脚本、模型权重、训练可视化图表和说明文档为主压缩后仅2.19MB轻量便于下载。其中h5权重文件可直接加载训练好的模型省去重新训练的时间两张png图片分别展示损失与准确率变化曲线以及预测效果帮助直观理解训练质量read_me说明txt则对运行方式和常见问题做了交代。目前已有4621人学习下载适合深度学习初学者作为图像分类实践的第一份完整资料也适合需要快速搭建CNN基线模型的研究者直接复用。1. 为什么 MNIST 过了这么多年依然是深度学习入门的第一道坎MNIST 是每个做神经网络的人几乎都绕不过去的数据集。它由 70000 张 28×28 的灰度手写数字图片组成其中 60000 张用于训练、10000 张用于测试内容覆盖了 0 到 9 这十个类别。很多人在看完吴恩达的深度学习课程或者《动手学深度学习》之后第一个上手的实战项目就是它——不是因为简单而是因为它足够干净不需要清洗、不需要标注、不需要纠结类别不平衡你只需要把模型建好剩下的时间全花在网络结构和训练技巧上。这篇文章要做的是带你用卷积神经网络CNN把这个任务完整做一遍从数据加载、网络设计、训练调参到最后的模型保存与加载。标题里提到的训练好的模型文件我会告诉你它是什么格式、怎么存、怎么在推理时重新加载而不是丢给你一个黑匣子。读完你应该能自己复现出 99% 以上的测试准确率并且明白每一步为什么要这么写。2. 理解手写数字识别从像素到十维向量的全流程2.1 一张 28×28 的灰度图在 CNN 眼里是什么一张 MNIST 图片在送入网络之前本质上是一个 28×28×1 的矩阵每个元素的值在 0 到 255 之间代表像素亮度。在 PyTorch 里我们用torchvision.datasets.MNIST加载后一般会立刻做两步预处理转成 Tensor再归一化到 [0, 1] 或 [-1, 1] 区间。from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_dataset datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform)这里的0.1307和0.3081是 MNIST 全量数据的全局均值和标准差。把它们写死是一个很常见的做法很多开源项目里都直接用这两个值。Normalize的作用是把像素分布拉到接近标准正态分布这样梯度更新会更平稳。注意一个细节ToTensor()会把原本的 H×W 数组变成 C×H×W 的 Tensor并自动把像素值缩放到 [0, 1]。如果你自己用 OpenCV 读图片这一步就不存在了得手动做img.astype(np.float32) / 255.0而且要记得在 channel 维度上补一个轴。很多人第一次在 DataLoader 里报错就是因为在transform和自定义 Dataset 之间没有把维度对齐。2.2 卷积层、池化层与全连接层各自的职责CNN 之所以在图像任务上碾压普通的前馈神经网络是因为它用卷积核在局部区域做特征提取。一个卷积核本质上是一个小矩阵比如 3×3 或者 5×5它在整张图上滑动对局部像素做加权求和。这种操作有两个好处参数共享同一个核用在全图和局部感知每个输出只看输入的一个邻域。池化层的职责是降采样。最常见的是 2×2 最大池化它把每 2×2 的区域压缩成一个最大值让特征图尺寸减半。这样做能扩大感受野同时也让模型对轻微的偏移和形变更鲁棒。全连接层则位于网络的尾部把卷积层输出的高维特征图展平成向量然后映射到十个类别的得分上。以 LeNet 为蓝本一个典型的手写数字识别网络是import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(inplaceTrue), nn.Linear(128, 10), ) def forward(self, x): return self.classifier(self.features(x))这块网络设计的关键是输入是 1 通道经过两次卷积 池化后变成 64 通道、7×7 的特征图所以全连接层输入维度是64 * 7 * 7 3136。这个数字不是玄学而是算出来的28 → 28padding1 保持尺寸→ 14池化→ 14 → 7池化。如果你想调整池化 stride 或者卷积核大小这个数值要重新算否则Linear层会直接报维度不匹配的错误。2.3 从输入到输出一次完整的前向传播过程前向传播的过程可以用一句话概括图像经过卷积和池化逐步变成更抽象的特征图最后经过全连接层输出十个 logits再用交叉熵损失函数计算预测与真实标签的差距。比如输入一张7网络输出的可能是一个十维向量其中对应7的那个维度得分最高。这就是热搜词里说的人脸识别图像进入神经网络到输出高维度向量的过程的基本逻辑MNIST 只是把这个过程简化到了十维。理解这一步很重要因为后面所有训练技巧——学习率调整、Dropout、数据增强——都是在优化这个特征提取的过程。3. 用 PyTorch 完整训一个 CNN训练循环与代码拆解3.1 DataLoader 配置batch size、shuffle 与 num_workers数据加载是深度学习项目里最容易被低估的环节。常见的做法是用DataLoader把 Dataset 包一层让它在迭代时自动组合 batch、打乱顺序。MNIST 太小不需要分布式训练但 batch size 和 worker 数量的选择依然有讲究。from torch.utils.data import DataLoader BATCH_SIZE 64 train_loader DataLoader( train_dataset, batch_sizeBATCH_SIZE, shuffleTrue, num_workers2, pin_memoryTrue, ) test_loader DataLoader( test_dataset, batch_size256, shuffleFalse, num_workers2, pin_memoryTrue, )batch_size64是 MNIST 上非常经典的选择它在梯度更新频率和单次迭代的计算开销之间取了一个平衡。shuffleTrue只在训练集上开测试时不需要打乱顺序因为准确率计算与样本顺序无关。pin_memoryTrue在 GPU 训练时可以加速 CPU 到 GPU 的数据拷贝CPU 训练时这个参数没有副作用保留即可。关于num_workers在 Windows 上设置大于 0 时偶尔会遇到子进程相关的问题。如果训练脚本频繁崩溃且错误堆栈指向 DataLoader先把它改成 0 跑一次。绝大多数 MNIST 场景下num_workers0或2在速度上感知不到明显差距因为没有必要为了这一点性能去折腾多进程调度的问题。3.2 训练循环损失函数、优化器与准确率跟踪训练循环本身的写法高度模式化但每一行都有实际含义。先从核心部分开始损失函数用CrossEntropyLoss优化器用 Adam 或 SGD每个 epoch 里遍历 train_loader 做前向传播、反向传播和参数更新。import torch import torch.nn as nn import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleCNN().to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) EPOCHS 10 for epoch in range(1, EPOCHS 1): model.train() running_loss 0.0 correct 0 total 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() train_acc 100.0 * correct / total avg_loss running_loss / len(train_loader) print(fEpoch {epoch}: loss{avg_loss:.4f}, train_acc{train_acc:.2f}%)这里有几个细节值得说明。optimizer.zero_grad()必须在backward()之前调用否则梯度会累加到上一轮的参数上导致更新方向错乱。torch.max(outputs, 1)返回每个样本在十个类别上得分最高的索引这就是模型的预测结果。model.train()和model.eval()会影响 Dropout 和 BatchNorm 的行为虽然当前这个网络没有用这两种层但把它写上是一个好习惯防止后续加层时忘记切换模式。3.3 验证集评估model.eval 与 no_grad 的必要性如果只把训练集上的准确率打印出来你看到的会是一个漂亮但不可信的数字。MNIST 的真正考点在测试集。每次训练循环结束后需要单独在 test_loader 上计算准确率而且整个过程必须关闭梯度记录。def evaluate(model, loader): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() acc 100.0 * correct / total print(fTest Accuracy: {acc:.2f}%) return acctorch.no_grad()的作用是告诉 PyTorch 不需要构建计算图。在推理时梯度没有意义强行保留的话会额外占用大量显存甚至导致 OOM。model.eval()与no_grad()是两件事前者改变层的行为如 Dropout 关闭、BatchNorm 用统计均值后者关闭自动求导。它们通常一起出现但各有各的职责不要混为一谈。4. 训练完怎么把模型带走保存、加载与推理验证4.1 state_dict 与完整模型的保存差异训练结束后你手里最值钱的东西不是代码而是模型参数。PyTorch 里保存模型有两种常见做法只存state_dict或者把整个模型对象序列化。torch.save(model.state_dict(), mnist_cnn.pt) torch.save(model, mnist_cnn_full.pt)第一行保存的是模型的权重和偏置张量。第二行保存的是包含网络结构在内的完整对象。前者的文件更小、跨版本兼容性更好是工程上的首选。后者虽然读回来更省事——不需要手动构建网络类——但严格依赖原本的 Python 环境和类定义路径一旦类名改了或者文件路径变了就容易加载失败。加载state_dict时的标准动作是先实例化模型再 load 参数最后切到推理模式model SimpleCNN().to(device) model.load_state_dict(torch.load(mnist_cnn.pt, map_locationdevice)) model.eval()map_locationdevice很关键。如果你在 GPU 上训练保存的权重张量默认在 cuda 设备上拿到一台没有 GPU 的机器上直接torch.load会报错。加上这个参数加载时才会自动映射到 CPU。4.2 用训练好的模型对单张图片做预测模型文件在手上最终要落到一个实际动作上给一张图片输出一个数字。下面是一段完整的推理代码可以直接拷贝使用。from PIL import Image import torchvision.transforms as transforms import torch def predict_image(model, image_path): image Image.open(image_path).convert(L) transform transforms.Compose([ transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): output model(tensor) prob torch.softmax(output, dim1) pred torch.argmax(prob, dim1).item() return pred, prob[0][pred].item() pred, confidence predict_image(model, ./test_digit.png) print(fPredicted digit: {pred}, confidence: {confidence:.4f})unsqueeze(0)把单张图片的 shape 从[1, 28, 28]变成[1, 1, 28, 28]多出来的那一维是 batch 维度。模型在训练时看到的是四维输入推理时也必须保持同样形状。softmax把 logits 转成概率分布方便看置信度。MNIST 上训练良好的模型对清晰手写数字的置信度通常高于 0.99如果低于 0.9建议检查图片是否反色、是否有黑边、数字是否居中。4.3 模型文件的验证清单别拿训练集准确率骗自己拿到一个训练好的模型文件后第一步不是直接部署而是验证它的真实表现。验证清单按经验排这么几条固定随机种子后重新加载模型在测试集上跑一次准确率看是否和训练日志里记录的一致抽十几张图片人工看一眼包括那些被预测错的样本因为测试准确率只告诉你一个数字不告诉你错在哪里再把模型输出 logits 和概率分布打印出来确认没有 NaN 和明显的数值异常。任何一步不对都应该怀疑模型文件本身有问题——可能是保存时没选state_dict、加载时 map_location 写错、或者训练时用了和推理时不同的预处理方式。5. MNIST 训练中的常见问题与避坑指南5.1 torchvision 下载 MNIST 报 404 或连接超时这是国内开发者几乎绕不过去的一个坑。downloadTrue时 torchvision 会尝试从 Yann LeCun 维护的官方地址下载 MNIST而这个地址在部分地区响应很慢甚至直接超时热搜里的torchvision下载mnist会404说的就是这个问题。现象脚本卡在Downloading http://yann.lecun.com/exdb/mnist/...很久不动然后报URLError或超时错误。原因官方服务器在国外网络链路不稳定并非代码问题。解决 方案有两种。一是手动下载四个压缩包train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz解压后放到./data/MNIST/raw/目录并保证文件名与 torchvision 预期一致然后设置downloadFalse。二是找一个国内可达的镜像源修改 torchvision 的下载逻辑但一般不建议直接改第三方库源码。最省心的做法是找一个已经导出了 MNIST 图片文件或 numpy 数组格式的版本自己写 Dataset 读取。5.2 训练多次准确率到不了 99%学习率与过拟合问题MNIST 的一个特征是小训练集 60000 张图片在 CNN 面前很快会被背下来。如果训练多次后测试准确率始终在 98% 附近徘徊观察训练集准确率是否已经接近 100%。如果两者差距大说明过拟合。常见的补救手段按优先级排加 Dropout在池化层之后、全连接层之前dropout 概率 0.2 到 0.5、调低学习率Adam 下从 1e-3 调到 lr_scheduler 的 StepLR 或 ReduceLROnPlateau、加 L2 正则化Adam 的weight_decay参数从 0 改成 1e-4。这三件事做完测试准确率通常能稳定越过 99%。不要一上来就换更大的网络MNIST 太小大网络反而更容易过拟合热搜里的深度学习l2正则化pytorch代码在 MNIST 上就是设一个weight_decay参数的事。5.3 训练过程中 loss 变成 NaN 或梯度爆炸这是一个一看就知道完了但不知道为什么完了的问题。现象loss 在前几个 epoch 正常下降某一步突然变成nan之后永不恢复。原因一般有两种一是学习率太大导致梯度更新迈过了最优区域权重发散二是数据没有归一化喂进去的像素值范围不对导致早期激活值过大。对应的解决方法是把学习率从 1e-3 降到 1e-4 重跑一次检查transform里是否真的做了Normalize——如果你加载的是自己构造的数据集而不是 torchvision 的 MNIST这一步最容易漏。还有一个隐蔽的原因是torch.set_printoptions打开了大数显示后误以为数值异常实际上只是 logits 变成了几千这不算问题softmax 之后照样能正常工作。5.4 同一份代码在别人的机器上准确率不一样这不是代码 bug而是随机性导致的正常现象。PyTorch 里涉及随机初始化的地方很多模型权重初始化、DataLoader 打乱顺序、CUDA 上的卷积核选择。即使设置torch.manual_seed(0)在 GPU 上依然不一定能完全复现结果。MNIST 上这个差异大概是 ±0.1%即 99.2% 和 99.3% 的区别。如果你的目标是复现论文里的结果就固定住所有随机种子并记录 GPU 型号如果只是评估模型能力跑三次取中间值通常是一个更务实的手段。6. 把准确率再往上推一步一个小技巧与最终验证到这里模型可能已经达到 99% 的测试准确率但还想再往上走半个百分点往往比从零到 99% 更费功夫。一个对 MNIST 非常有效的做法是输出十个类别的混淆矩阵看看这 0.5% 的错误到底去哪了一个有代表性的实操是把预测错误的样本连同真实标签和图片像素保存下来逐个看。import matplotlib.pyplot as plt import numpy as np def show_misclassified(model, loader, num_samples10): model.eval() misclassified [] with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) preds torch.argmax(outputs, dim1) mask preds ! labels if mask.any(): for img, label, pred in zip(images[mask], labels[mask], preds[mask]): misclassified.append((img.cpu(), label.item(), pred.item())) fig, axes plt.subplots(1, min(num_samples, len(misclassified)), figsize(15, 3)) for idx, (img, label, pred) in enumerate(misclassified[:num_samples]): axes[idx].imshow(img.squeeze(), cmapgray) axes[idx].set_title(flabel{label}, pred{pred}) axes[idx].axis(off) plt.show()MNIST 模型最常见的错误是 4↔9 混淆、7↔2 混淆以及 3↔8 混淆这是因为这些数字在低分辨率下的形态确实接近。如果把错误图片打出来看到的是这类合理错误基本可以认为网络已把信息压榨得差不多了。这时候还想继续提点一个简单的操作是凑一批由不同初始化种子训练出的模型对它们的 softmax 输出做平均。三种子训出来的单模型准确率都在 99% 左右平均后往往能到 99.3% 附近代价是推理速度变成原来的三倍但对 MNIST 来说几乎可以忽略。回到模型文件这件事上我的个人习惯是保存state_dict时把验证准确率、epoch 数、batch size 和优化器参数写到一个同名 JSON 文件里。因为三个月后你自己可能都不记得这个.pt是用 Adam 还是 SGD 训练的更别说别人拿到文件时的困惑。把生成条件记录清楚比把代码写漂亮更重要这也是一条希望帮到你的实践建议。本文还有配套的精品资源点击获取