
简介基于Python深度学习实现MNIST手写数据集识别的完整项目包面向计算机、电子信息工程、数学等专业的大学生适用于课程设计、期末大作业或毕业设计阶段的参考资料。项目以卷积神经网络为核心代码组织清晰包含ConvNet、layers、functions、mnist等Python模块以及训练与测试所需的原始数据集能够帮助读者理解图像分类、卷积运算、参数训练与模型评估的完整流程。压缩包共18个文件主要文件类型包括py源码、pyc编译文件、json配置、idx格式二进制数据以及pkl模型文件整体大小约19.77MB目录结构便于按模块对照学习。同时包内还提供了VSCode相关配置文件可快速搭建Python运行环境节省调试时间pyc文件可用于快速复现运行结果结合py源码则可深入分析每一层网络实现。目前已有511人学习下载适合具备一定Python和深度学习基础、能够自行调试并扩展功能的读者作为参考。1. 一个老掉牙的数据集为什么今天还要亲手跑一遍MNIST 手写数据集识别几乎是每个做深度学习的人第一个亲手跑通的项目也是很多入门课程的作业标配。别小看这个“老掉牙”的例子哪怕你是让 Codex 这类工具帮你把网络代码生成出来最终要在本地跑通、拿到 99% 以上的测试准确率该踩的坑一个都少不掉数据下载 404、归一化漏写、显卡显存报错、训练曲线死活不上涨这些都比写网络结构更磨人。这篇笔记要讲清楚的就是一件事拿到一份“源码数据”的 MNIST 识别资源后从环境准备到模型落地完整复现并理解每一步在干什么。适合初学深度学习的同学照着做也适合已经跑通但说不清参数为什么这样设的从业者回头补课。2. 先把环境铺平Python 装到能用torchvision 下载 MNIST 报 404 的手动解法2.1 Python 与 PyTorch 安装别在第一步就把自己劝退MNIST 识别的技术栈很简单Python 加 PyTorch 就够。深度学习的编程语言选择上Python 基本是唯一需要认真考虑的主流选项因为数据加载、模型定义、训练可视化这些生态都在 Python 这边。安装时我一般不建议新手直接裸装 Python然后拿 pip 一个一个补包而是先装 Anaconda 或 Miniconda用 conda 建一个独立环境。这样做的好处是后面换项目时不会把依赖搞成一锅粥比如你同时跑 NLP 和 CV依赖版本冲突是常有的事。conda create -n mnist python3.10 -y conda activate mnist pip install torch torchvision matplotlib scikit-learn逻辑说明第一行创建名为 mnist 的独立环境并指定 Python 3.10第二行激活第三行把训练和可视化需要的包一次装齐。这里用 pip 而不是 conda 装 PyTorch是因为 PyTorch 官方对 pip 的预编译包支持最及时conda 渠道有时候版本滞后。参数说明Python 版本选 3.10 而不是最新的 3.12是因为 PyTorch 对 3.10 的预编译 wheel 覆盖最成熟能省掉很多“装了但 import 报错”的折腾。如果你电脑有 NVIDIA 显卡建议去 PyTorch 官网按 CUDA 版本生成安装命令没有独立显卡就装 CPU 版MNIST 这种小任务 CPU 一样能几分钟训完不必为了它折腾 CUDA 环境。装完后打开终端输入python -c import torch; print(torch.__version__)能输出版本号就算通了。这一步卡住的人最多常见原因是网络问题导致 wheel 下载失败解决方案是换国内 pip 镜像源这个属于装机基本功不多展开。2.2 torchvision 下载 MNIST 报 404手动备好四件套跑 MNIST 最常见的第一道坎是datasets.MNIST(downloadTrue)时直接报HTTP Error 404: Not Found或者进度条卡在某个百分比不动。这是因为 torchvision 的 MNIST 下载地址在不同版本之间换过老版本指向原始站点新版指向对象存储域名而你的网络环境访问其中某个地址很可能失败。很多人在这里反复重试浪费时间。我现在的做法是干脆不走自动下载手动把数据备好让 torchvision 只做本地加载。需要准备四个 gz 压缩包文件名是固定的训练图像train-images-idx3-ubyte.gz、训练标签train-labels-idx1-ubyte.gz、测试图像t10k-images-idx3-ubyte.gz、测试标签t10k-labels-idx1-ubyte.gz。从任何能访问的源拿到这四个文件后放到项目的data/MNIST/raw/目录下mkdir -p data/MNIST/raw # 将四个 .gz 文件放入 data/MNIST/raw 后执行 cd data/MNIST/raw gzip -dkf train-images-idx3-ubyte.gz gzip -dkf train-labels-idx1-ubyte.gz gzip -dkf t10k-images-idx3-ubyte.gz gzip -dkf t10k-labels-idx1-ubyte.gz # 回项目根目录删掉可能存在的旧缓存 rm -rf ../processed逻辑说明gzip -dkf中的-d是解压-k保留原 gz 文件-f强制覆盖已存在的同名文件。最后删掉 processed 目录很关键因为 torchvision 第一次加载成功后会生成train.pt和test.pt缓存如果这份缓存是之前下载不完整时生成的模型会读到残缺数据表现就是训练准确率诡异。参数说明这四个文件解压后是二进制格式不是图片文件夹。torchvision 的 MNIST 类在downloadFalse时会直接读取raw/目录下这些解压后的文件找不到就会报错。所以解压是必须的只放 gz 不 gzip 解压是很多人踩的坑。做完这一步代码里把downloadTrue改成downloadFalse加载就再也不会碰网络了from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST( root./data, trainTrue, transformtransform, downloadFalse ) test_dataset datasets.MNIST( root./data, trainFalse, transformtransform, downloadFalse )逻辑说明这段代码定义了两个数据集对象训练集和测试集都套用同一个 transform 流程。root./data指向数据根目录torchvision 会在其下自动找MNIST/raw/。downloadFalse表示只用本地文件不再尝试联网。参数说明Normalize((0.1307,), (0.3081,))是 MNIST 官方统计出的像素均值和标准差作用是让输入分布接近标准正态分布让网络收敛更稳。这两个值不要自己随便算直接用这个公认值就行。2.3 拿到“源码数据”压缩包后先按这个结构整理解压 rar 包后不管原作者怎么命名目录我的习惯是先重排成统一结构再动手。这样做的好处是后面所有命令、脚本路径都基于同一个约定不用每次猜。常见的组织方式是这样mnist_project/ ├─ data/ │ └─ MNIST/ │ ├─ raw/ # 手动准备的四个数据文件所在 │ └─ processed/ # torchvision 自动生成的缓存可删 ├─ models/ # 训练好的权重、导出的模型文件 ├─ train.py # 数据加载 模型定义 训练循环 └─ predict.py # 加载权重做单张识别逻辑说明train.py负责从数据加载到训练保存的全过程predict.py只负责推理两者分开是工程上的好习惯。models/目录专门存.pt或.pth权重文件这样训练脚本和推理脚本不会互相污染。参数说明这个结构里没有把源代码拆成多个模块因为 MNIST 这种规模的工程单文件反而更好读。等模型结构复杂了再拆models.py、utils.py不迟。新手不要一上来就过度设计目录维护成本比代码本身还高。3. 把图片喂给网络之前理解 MNIST 的数据格式与 DataLoader 的三个关键参数3.1 MNIST 原始的图像格式28×28 灰度图与标签映射MNIST 数据集包含 6 万张训练图和 1 万张测试图每张图是 28×28 像素的灰度图像素值范围 0 到 2550 是黑色背景255 是白色笔迹。标签是 0 到 9 的整数这张图是哪个数字标签就是几。数据在磁盘上是二进制文件图像和标签分开存储所以第一步一定是交给 torchvision 的 dataset 类解析不要自己去读二进制。加载后可以先用 Matplotlib 看一眼原始长什么样确认数据没问题再进训练import matplotlib.pyplot as plt from torch.utils.data import DataLoader train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) images, labels next(iter(train_loader)) img images[0].squeeze() # (1, 28, 28) - (28, 28) print(pixel range:, img.min().item(), img.max().item()) plt.imshow(img, cmapgray) plt.title(flabel: {labels[0].item()}) plt.axis(off) plt.show()逻辑说明next(iter(train_loader))取出一个 batchimages[0]是 batch 里第一张图shape 是(1, 28, 28)1 是通道数灰度图只有一个通道。squeeze()去掉通道维变成(28, 28)才能被 Matplotlib 正常显示。参数说明打印像素范围时你会看到不是 0 到 1而是负数和大于 1 的数这是归一化后的正常现象不影响训练。显示出来的图像可能比原图“灰”一些因为归一化把像素均值拉到了 0 附近视觉上会变暗。3.2 DataLoader 里的三个参数batch_size、shuffle、num_workers数据准备好后训练时不会一张一张喂而是打包成 batch 喂给网络。这一步由 DataLoader 完成三个参数决定了训练的速度和稳定性from torch.utils.data import DataLoader train_loader DataLoader( train_dataset, batch_size64, shuffleTrue, num_workers0, ) test_loader DataLoader( test_dataset, batch_size128, shuffleFalse, num_workers0, )逻辑说明训练集shuffleTrue会让每个 epoch 的样本顺序重新打乱避免模型学到批次内的顺序关联。测试集不需要打乱shuffleFalse即可。num_workers0表示用主进程加载数据Windows 下最稳妥后面避坑章节会细说。参数说明batch_size64是 MNIST 的常见选择太小收敛慢太大会让梯度更新过于平滑。显卡显存小的改成 32 也完全没问题。测试集batch_size128比训练集大是因为推理不需要反向传播不占梯度内存。一个容易被忽略的点DataLoader 返回的 images 和 labels 是成对的images的 shape 是(batch_size, 1, 28, 28)labels是(batch_size,)。很多报错都来自 shape 对不上比如全连接层输入维度算错就是因为没意识到图像是四维张量不是二维矩阵。3.3 为什么 MNIST 不需要做强数据增强图像分类项目里数据增强是标配随机裁剪、翻转、色彩抖动能显著提升泛化能力。但 MNIST 是个例外我一般不做强增强。原因很直接MNIST 是工整的、居中的手写数字语义对方向敏感。把 6 水平翻转可能变成 9把 7 旋转 180 度可能变成 L这类变换对数字识别是破坏性的而不是增强性的。常见的做法是只做轻微扰动比如 2 像素以内的随机平移或者 ±10 度以内的小角度旋转。但说实话对 MNIST 来说收益非常有限因为数据集本身已经足够干净、足够大。很多人的经验是与其花时间调增强不如把网络结构多写一层卷积提升更明显。如果一定要加我建议用transforms.RandomAffine做轻量增强幅度控制住transform_train transforms.Compose([ transforms.RandomAffine(degrees10, translate(0.05, 0.05)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])逻辑说明RandomAffine的degrees10表示随机旋转 ±10 度translate(0.05, 0.05)表示水平和垂直方向随机平移最多 5% 像素。这个幅度对数字语义几乎没有破坏又能增加一点多样性。参数说明degrees超过 20 时3、7、9 这类数字的识别率会明显下降因为旋转后和别的数字字形接近了。新手容易一上来就抄 ImageNet 的增强策略在 MNIST 上反而翻车。4. 用 PyTorch 写一个能到 99% 的 CNN网络结构、训练循环与关键参数4.1 网络结构设计两层卷积加两层全连接每一层 shape 变化都算清楚MNIST 的经典网络结构不复杂两层卷积池化加两层全连接就能到 99%。我一般习惯先定义模型类再写训练循环。模型的代码长这样import torch import torch.nn as nn class MnistCNN(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2), ) self.classifier nn.Sequential( nn.Linear(64 * 7 * 7, 128), nn.ReLU(inplaceTrue), nn.Dropout(0.2), nn.Linear(128, 10), ) def forward(self, x): x self.features(x) x x.view(x.size(0), -1) out self.classifier(x) return out逻辑说明features部分负责提取图像特征卷积核 3×3padding1 保证卷积不改变图像尺寸池化层把尺寸减半。classifier部分把特征图拉平后做分类。view(x.size(0), -1)把四维特征图(batch, 64, 7, 7)拉平成二维(batch, 3136)这样才能进入全连接层。参数说明Conv2d(1, 32, kernel_size3, padding1)的第一个参数 1 是输入通道数灰度图是 1RGB 图就是 3。32 是输出通道数相当于用 32 个卷积核提取 32 种特征。第二个卷积层把 32 通道升到 64 通道特征更丰富。padding1配合kernel_size3才能保持尺寸不变去掉 padding 每过一层卷积图像就小一圈后面的全连接输入维度就要重新算。每一层的 shape 变化可以用一张表算清楚这在写代码前就应该心里有数层输入 shape输出 shape说明Conv2d(1, 32, 3, padding1)(1, 28, 28)(32, 28, 28)通道数变 32尺寸不变MaxPool2d(2)(32, 28, 28)(32, 14, 14)尺寸减半Conv2d(32, 64, 3, padding1)(32, 14, 14)(64, 14, 14)通道数变 64MaxPool2d(2)(64, 14, 14)(64, 7, 7)尺寸再减半Linear(64×7×7, 128)(batch, 3136)(batch, 128)拉平后全连接Linear(128, 10)(batch, 128)(batch, 10)输出 10 个类别的得分这张表的价值在于全连接层的输入维度64 * 7 * 7是从前面层层推导出来的不是拍脑袋。改网络结构时先改表再改代码能少踩很多 shape 报错。4.2 训练循环五个步骤一个都不能少模型定义好后训练循环是核心。PyTorch 的训练循环模式是固定的五个操作按顺序执行前向传播算输出、计算损失、清零梯度、反向传播、更新参数。代码长这样import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model MnistCNN().to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) scheduler optim.lr_scheduler.StepLR(optimizer, step_size3, gamma0.5) EPOCHS 10 for epoch in range(1, EPOCHS 1): model.train() total_loss, correct, total 0.0, 0, 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() total_loss loss.item() * images.size(0) correct (outputs.argmax(dim1) labels).sum().item() total labels.size(0) train_acc correct / total avg_loss total_loss / total model.eval() test_correct, test_total 0, 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) test_correct (outputs.argmax(dim1) labels).sum().item() test_total labels.size(0) test_acc test_correct / test_total print(fepoch {epoch}: loss {avg_loss:.4f} | train_acc {train_acc:.4f} | test_acc {test_acc:.4f})逻辑说明model.train()和model.eval()切换模型状态train 模式下 Dropout 生效eval 模式下 Dropout 关闭。optimizer.zero_grad()每次迭代前把梯度清零不清零的话梯度会累加。loss.backward()计算梯度optimizer.step()更新参数这五个操作的顺序在 PyTorch 里是固定的写反了模型就训不动。参数说明criterion用的是交叉熵损失多分类的标准选择。optimizer用 Adamlr1e-3是默认学习率MNIST 上表现稳定。scheduler是学习率调度器每 3 个 epoch 把学习率乘以 0.5作用是后期减小步长、让损失更精准地落到低点。EPOCHS10对 MNIST 足够跑 15 个也不会明显更好只是浪费时间。4.3 模型的参数量到底是多少42 万参数不是 42 万 MB训练时很多人会好奇模型有多大常听到的疑问是“深度学习里的 parameter 应该不是 mb 吧” 这个疑问是对的。parameter 的单位是“个”不是字节。上面的 MnistCNN 参数量大约 42 万个可以用一行代码确认total_params sum(p.numel() for p in model.parameters()) print(total params:, total_params)逻辑说明numel()返回每个参数张量的元素个数把所有参数张量的元素数加起来就是总参数量。打印出来的结果约等于 42 万这就是模型“有多少个可学习的参数”。参数说明42 万参数在 float32 精度下占用的显存是421642 × 4字节约 1.7 MB。注意这 1.7 MB 只是权重本身。训练时还需要额外存储梯度和中间激活值实际显存占用大概是权重的几倍不同 batch_size 差别很大。所以 MNIST 训练时显存占用一般也就几百 MB4 GB 显存的旧显卡也能轻松跑。很多入门者把参数量和显存混淆看到“7B 参数”就以为要 7 GB 显存实际上 7B 参数在 float32 下光权重就要 28 GB两者完全是两个概念。4.4 训练日志怎么看loss 下降、准确率停滞与过拟合信号训练过程的打印结果不是拿来看一眼就完事的每一行都有含义。正常的训练曲线是第一个 epoch 结束时 loss 从初始的 2.3 左右掉到 0.3 附近train_acc 到 95% 上下第二个 epoch loss 继续降到 0.1 以下test_acc 到 98% 左右之后每轮提升幅度越来越小到第 8 轮以后基本稳定在 99% 以上。如果发现 train_acc 接近 100% 但 test_acc 停在 97% 附近这是典型的过拟合信号。原因通常是模型容量对 MNIST 来说偏大训练集太简单网络开始“背题”。解决方案不是换模型而是给 classifier 里的 Dropout 加大比例从 0.2 调到 0.5或者减小训练轮数。反过来如果 train_acc 和 test_acc 都很低比如都在 90% 以下说明模型欠拟合优先检查数据加载和归一化而不是急着调网络结构。如果 loss 出现先降后升的“V 型”反转且 test_acc 同步下跌多半是学习率太大导致权重震荡。这时看 scheduler 是否生效或者直接把初始lr从1e-3改到1e-4重训。记住一个原则训练日志是模型健康的体温计只看最终准确率不看过程就等于开盲盒。5. MNIST 识别避坑指南五个真实翻车现场与排查路径5.1 测试准确率卡在 97% 上不去问题在归一化而不是网络现象网络结构照着抄训练轮数也不少但 test_acc 始终在 97% 附近晃怎么加层都上不去。原因transform 里只写了ToTensor()漏了Normalize那行。像素只缩放到 0 到 1没有做标准化导致网络输入分布与最优解之间有明显偏移收敛变慢且上限受限。MNIST 对归一化的敏感度不如 ImageNet 高所以不是“完全训不动”而是“差一口气”。解决补上Normalize((0.1307,), (0.3081,))保持训练集和测试集用完全相同的 transform。改完通常能直接拉到 99% 以上。检查方法很简单打印第一张图的像素范围如果最小值接近 0、最大值接近 1说明归一化没生效。5.2 训练 loss 在降、测试准确率却停滞过拟合与学习率过大怎么分辨现象训练集 loss 一路下降train_acc 到 99.9%但 test_acc 一直在 98% 附近徘徊甚至最后两轮还跌了一点。原因两个可能。一是模型记住了训练集样本泛化能力不足典型过拟合二是学习率后期偏大权重在最优解附近来回震荡测试集上的表现波动。解决先加 Dropout把classifier里的nn.Dropout(0.2)改成0.5重训。如果加了 Dropout 还是老样子再检查 schedulerStepLR(step_size3, gamma0.5)是否在跑学习率后期有没有降下来。我自己的经验是MNIST 上大多数“训练好、测试差”的情况调一下 Dropout 就能解决不需要动网络结构。5.3 torchvision 下载 MNIST 反复 404缓存文件是幕后黑手现象downloadTrue第一次报 404网上搜教程改成手动下载后数据文件也放进去了还是报错或者训练时准确率极低。原因data/MNIST/processed/目录下残留了之前下载失败时生成的残缺缓存train.pt和test.pt。torchvision 每次加载时优先读 processed 缓存只要文件存在就不检查完整性直接拿来用。解决手动下载四个 gz 文件并解压后一定要删掉 processed 目录再运行。让 torchvision 重新从 raw 文件生成缓存。这个动作我已经养成肌肉记忆了每次手动更新 MNIST 数据的第一件事就是rm -rf data/MNIST/processed。5.4 Windows 上 num_workers 大于 0 直接崩溃BrokenPipeError 的解法现象代码在 Linux 上跑得好好的换到 Windows 上一运行就报BrokenPipeError或DataLoader worker (pidxxxx) exited unexpectedly有时候还会卡死。原因Windows 下多进程数据加载的启动方式和 Linux 不同需要主代码有if __name__ __main__:保护。如果训练脚本直接写在模块顶层Windows 的多进程 worker 会重复执行顶层代码导致进程冲突。解决把训练循环包进main()函数并在脚本末尾加if __name__ __main__: main()。如果加了保护还崩溃直接把num_workers改为 0用主进程加载数据。MNIST 数据集很小num_workers0带来的速度损失可以忽略换来的是稳定。5.5 两次训练结果不一样随机种子不是玄学是可复现的关键现象同一份代码、同一个数据集跑两次 test_acc 分别是 99.1% 和 99.4%有时候差的还不止 0.3%。原因模型权重初始化、DataLoader 的 shuffle、CUDA 算子都存在随机性。这不叫 bug是深度学习的默认行为。但如果项目要求结果可复现比如做实验对比就必须手动固定随机种子。解决在训练脚本开头加一段固定种子的代码import random import numpy as np def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed(42)逻辑说明torch.manual_seed固定 CPU 上的随机数生成器torch.cuda.manual_seed_all固定所有 GPU 的随机数生成器random.seed和np.random.seed分别固定 Python 原生随机和 NumPy 随机。四行一起才能覆盖几乎所有随机源。参数说明固定种子并不能保证两次训练结果严格一致因为某些 CUDA 算子本身有非确定性。但能把方差压到很小的范围比如 0.1% 以内。对 MNIST 这种任务我更建议关注多次运行的平均水平而不是追求某一轮的精确复现。6. 最后一里路用混淆矩阵纠错把模型导出成可交付的 TorchScript6.1 混淆矩阵与错误样本定位 4/9、7/2 这类冤家对头模型训完不是终点还得知道它错在哪。准确率 99% 意味着每 100 张图错 1 张那 1 张长什么样、被错认成了什么混淆矩阵能给答案。在测试集上跑一遍完整预测生成矩阵import numpy as np from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt model.eval() preds_all, labels_all [], [] with torch.no_grad(): for images, labels in test_loader: images images.to(device) preds model(images).argmax(dim1) preds_all.extend(preds.cpu().tolist()) labels_all.extend(labels.tolist()) cm confusion_matrix(labels_all, preds_all) ConfusionMatrixDisplay(cm, display_labelsrange(10)).plot(cmapBlues) plt.show()逻辑说明confusion_matrix(labels_all, preds_all)的返回值是 10×10 的矩阵第 i 行第 j 列表示“真实类别是 i、预测成 j”的样本数。对角线越大越好非对角线上的数字就是具体的错误模式。参数说明看混淆矩阵时重点关注非对角线中数值较大的位置。MNIST 上最常见的混淆对是 4 和 9、7 和 2、3 和 8因为字形相近。如果发现某个对角线值异常小比如 5 被错认成 3 的有几十个说明模型对这两个数字的区分能力弱可以考虑针对性增加这类样本或者提高输入分辨率。再进一步把预测错误的样本画出来肉眼看看到底是模型蠢还是数据本身有问题preds_np np.array(preds_all) labels_np np.array(labels_all) err_idx np.where(preds_np ! labels_np)[0] fig, axes plt.subplots(1, 5, figsize(12, 3)) for i, idx in enumerate(err_idx[:5]): img, label test_dataset[idx] img img.squeeze().numpy() img img * 0.3081 0.1307 # 反归一化还原显示范围 axes[i].imshow(img, cmapgray) axes[i].set_title(flabel{label}, pred{preds_np[idx]}) axes[i].axis(off) plt.show()逻辑说明np.where(preds_np ! labels_np)返回所有预测错误的索引取前 5 个画出来。反归一化那行是重点训练时减均值除方差显示时要把变换逆回去否则图像灰蒙蒙的看不清。参数说明画错误样本时经常会看到两种情况一种是确实连人眼都难分辨的潦草写法属于数据本身的噪声模型错了不冤另一种是人眼一看就是某个数字、模型却认错说明特征提取还有短板值得回去调网络。这个判断方法百试百灵。6.2 导出 TorchScript 模型脱离训练代码也能跑推理训练完把权重存下来只是第一步真正交付给其他程序用时需要把模型导出成不带训练逻辑的独立文件。PyTorch 的 TorchScript 是标准方案训练时用的 dataset、transform、optimizer 全都不需要了一个.pt文件就能完成推理model.eval() model_cpu model.to(cpu) example torch.randn(1, 1, 28, 28) traced_model torch.jit.trace(model_cpu, example) traced_model.save(models/mnist_cnn_script.pt) print(saved to models/mnist_cnn_script.pt)逻辑说明torch.jit.trace用一张示例输入跑一遍前向计算把网络结构和参数一起固化成一个独立的计算图。导出前必须model.eval()因为 trace 会固化当前模型状态如果还在 train 模式Dropout 会被固化进去推理时行为就不对了。参数说明example的 shape 必须是(1, 1, 28, 28)和训练时的输入一致。导出后可以用一行加载代码验证文件能不能独立跑通loaded_model torch.jit.load(models/mnist_cnn_script.pt) loaded_model.eval() with torch.no_grad(): logits loaded_model(example) pred logits.argmax(dim1).item() print(prediction:, pred)逻辑说明torch.jit.load加载导出的文件不再需要MnistCNN类定义也不依赖 torchvision 的 transform。logits是 10 个类别的得分argmax(dim1)取最大得分的下标就是预测结果。参数说明推理时必须放在torch.no_grad()里省掉梯度计算的开销。这里返回的是 logits 而不是概率如果后续逻辑需要置信度要先过一层softmax。导出验证通过后这个.pt文件就可以交给任何部署环境了这是我们工程落地的最后一环。这几年跑 MNIST 我最大的体会是这个数据集看着简单但它覆盖了深度学习项目从数据到部署的每一个关键环节。数据手动下载、归一化、网络 shape 推导、训练日志判断、模型导出每一环都在为将来跑更大更复杂的任务做铺垫。先把 MNIST 每一步的原理和坑都吃透后面遇到 CIFAR、ImageNet至少不会在同样的地方栽两遍希望这些经验能帮你少走一段弯路。本文还有配套的精品资源点击获取