ARTICLE DETAIL

资讯详情

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

PyTorch实战:MNIST手写数字识别CNN模型从训练到推理全流程

PyTorch实战:MNIST手写数字识别CNN模型从训练到推理全流程 简介这份资源面向深度学习入门者与计算机视觉初学者围绕MNIST手写数字识别这一经典任务提供从数据理解到卷积神经网络训练落地的完整实践材料。包内共6个文件以2个Python脚本、2张PNG图表、1份TXT说明和1个H5权重文件为主压缩包约2.19MB体积轻巧便于快速上手。其中脚本覆盖CNN模型的构建、训练与评估流程权重文件可直接加载已训练模型省去重复训练的时间与算力两张图表分别呈现损失与准确率变化曲线以及测试样本预测效果便于直观判断模型性能说明文档则交代运行方式与使用注意事项。已有4620人学习下载适合希望掌握TensorFlow下图像分类基本流程、理解卷积层与池化层作用、并动手复现高准确率识别模型的读者参考。1. MNIST 手写数字识别从数据集到 CNN 模型落地的完整路径MNIST 手写数字识别是深度学习入门最经典的实战场景也是卷积神经网络原理最容易验证的试验田。很多人第一次跑通深度学习模型就是从 MNIST 数据集开始的。但真正把它做完整——数据加载不出错、CNN 结构设计合理、训练过程可监控、模型能保存并直接推理——中间有不少容易翻车的地方。比如 torchvision 下载 MNIST 会 404 这个问题几乎每个国内开发者都遇到过。这篇文章面向想用 PyTorch 跑通手写数字识别的新手和需要快速复现基线模型的工程师从数据集介绍、环境配置、CNN 结构设计、训练调参到模型保存与推理每一步都给出可抄作业的代码和参数说明。读完你手里会有一个训练好的模型文件拿一张手写数字图片就能直接预测。2. MNIST 数据集与 PyTorch 环境先把地基打牢2.1 MNIST 数据集到底长什么样MNIST 全称 Modified National Institute of Standards and Technology database由 Yann LeCun 等人整理发布。训练集 60000 张测试集 10000 张每张是 28×28 像素的灰度图对应 0 到 9 十个类别。图片是黑底白字像素值范围 0 到 255数字大致居中但位置和笔画粗细有差异。这个数据集之所以经久不衰是因为它足够小——整个数据集压缩后不到 12MBCPU 上也能在几分钟内跑完一轮训练同时又足够真实——手写数字的类内差异明显能有效检验模型的泛化能力。用 PyTorch 加载 MNIST 的标准做法是通过 torchvision.datasets。但这里有个高频踩坑点torchvision 默认从境外源下载国内网络环境下大概率超时或返回 404。解决办法有两种一是手动下载四个 gz 文件放到指定目录二是修改下载源。我一般会提前把文件准备好避免训练脚本跑到一半卡在下载上。import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 定义预处理转张量 归一化 transform transforms.Compose([ transforms.ToTensor(), # 将 PIL 图像转为 tensor像素值缩放到 [0,1] transforms.Normalize((0.1307,), (0.3081,)) # MNIST 全局均值和标准差 ]) # 加载训练集和测试集 train_dataset datasets.MNIST( root./data, # 数据存放路径 trainTrue, # 训练集 downloadTrue, # 首次运行设为 True之后可改为 False transformtransform ) test_dataset datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) # 构建 DataLoader train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_loader DataLoader(test_dataset, batch_size1000, shuffleFalse)这段代码里有两个参数值得展开说。Normalize((0.1307,), (0.3081,))里的两个数字是 MNIST 训练集的全局像素均值和标准差用它们做归一化能让输入分布更接近标准正态加速收敛。batch_size64是训练时的批大小显存或内存不够就降到 32想更快跑完可以升到 128但学习率也要相应调整。测试集的 batch_size 设成 1000 是为了一次性算完所有测试样本的准确率减少循环次数。2.2 环境配置与依赖版本PyTorch 生态更新快版本不匹配是另一个常见翻车点。截至我写这篇文章时比较稳的组合是 Python 3.9 到 3.11、PyTorch 2.0 以上、torchvision 0.15 以上。如果你用 GPU 训练CUDA 版本要和 PyTorch 安装命令里的 cu 版本对应。CPU 训练 MNIST 完全够用一轮大约 10 到 20 秒十轮下来两三分钟。安装命令按官方推荐的方式走# CPU 版本 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu # GPU 版本以 CUDA 11.8 为例 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118装完之后用一行代码验证import torch print(torch.__version__) print(torch.cuda.is_available()) # GPU 用户应返回 True如果torch.cuda.is_available()返回 False先检查驱动版本再检查安装命令里的 CUDA 版本是否和本机匹配。CPU 用户看到 False 是正常的不影响后续训练。注意不要混用 conda 和 pip 安装 PyTorch容易出现动态库冲突。选一种方式装到底。3. 卷积神经网络结构设计与 PyTorch 实现3.1 为什么 CNN 比全连接网络更适合手写数字识别全连接网络处理 28×28 图像时要把 784 个像素拉成一维向量。这样做丢失了空间结构信息——相邻像素的关系、笔画的局部模式都被打散了。卷积神经网络通过卷积核在图像上滑动天然保留了空间局部性。一个 3×3 的卷积核能捕捉边缘、拐角这类低级特征多层堆叠后能组合出数字的笔画结构。池化层则负责降维和提供一定的平移不变性让模型对数字位置的微小偏移不那么敏感。具体到 MNIST一个典型的 CNN 结构是两层卷积加池化后面接全连接分类头。这个规模在 MNIST 上能轻松达到 99% 以上的测试准确率参数量不到 50 万训练和推理都很快。再深的网络在这个数据集上收益递减反而容易过拟合。3.2 用 PyTorch 定义 CNN 模型下面是我常用的一个 CNN 结构两层卷积、两层池化、两层全连接代码简洁且效果稳定。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_channels1, out_channels32, kernel_size3, padding1) # 第二层卷积输入 32 通道输出 64 通道卷积核 3x3 self.conv2 nn.Conv2d(in_channels32, out_channels64, kernel_size3, padding1) # 池化层2x2 最大池化 self.pool nn.MaxPool2d(kernel_size2, stride2) # 全连接层经过两次池化后特征图大小为 7x7通道数 64 self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) # 10 个类别 self.dropout nn.Dropout(0.25) # 防止过拟合 def forward(self, x): # 第一层卷积 - ReLU - 池化 x self.pool(F.relu(self.conv1(x))) # 输出: [batch, 32, 14, 14] # 第二层卷积 - ReLU - 池化 x self.pool(F.relu(self.conv2(x))) # 输出: [batch, 64, 7, 7] # 展平 x x.view(-1, 64 * 7 * 7) # 全连接 Dropout x F.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x逐层拆解一下。conv1用 32 个 3×3 卷积核padding1保证输出尺寸和输入一致都是 28×28。经过第一次 2×2 池化后变成 14×14。conv2用 64 个 3×3 卷积核再经过第二次池化变成 7×7。展平后是 64×7×73136 维接一个 128 维的全连接层最后输出 10 维对应十个数字类别。Dropout(0.25)在训练时随机丢弃 25% 的神经元是防止过拟合的常规手段。参数调整建议如果训练准确率远高于测试准确率把 dropout 提高到 0.5如果欠拟合把 fc1 的 128 改成 256或者再加一层卷积。3.3 训练循环与关键参数设置训练循环是整条链路里最容易出细节问题的地方。损失函数用交叉熵优化器用 Adam学习率从 1e-3 开始试。import torch 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) optimizer optim.Adam(model.parameters(), lr1e-3) scheduler StepLR(optimizer, step_size3, gamma0.5) # 每 3 轮学习率减半 criterion nn.CrossEntropyLoss() def train(model, device, train_loader, optimizer, epoch): model.train() 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() # 更新参数 if batch_idx % 100 0: print(fEpoch {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)}] f Loss: {loss.item():.4f}) def test(model, device, test_loader): model.eval() test_loss 0 correct 0 with torch.no_grad(): for data, target in test_loader: data, target data.to(device), target.to(device) output model(data) test_loss criterion(output, target).item() pred output.argmax(dim1, keepdimTrue) correct pred.eq(target.view_as(pred)).sum().item() test_loss / len(test_loader) acc 100. * correct / len(test_loader.dataset) print(fTest Loss: {test_loss:.4f}, Accuracy: {acc:.2f}%) return acc # 主训练流程 for epoch in range(1, 11): train(model, device, train_loader, optimizer, epoch) acc test(model, device, test_loader) scheduler.step() # 更新学习率几个关键参数说明。lr1e-3是 Adam 在 MNIST 上的常用起点如果 loss 震荡明显就降到 5e-4。StepLR每 3 轮把学习率乘 0.5后期微调时步长更小有助于收敛到更优的点。optimizer.zero_grad()必须放在前向传播之前否则梯度会累加。model.eval()和torch.no_grad()在测试阶段必须加上前者关闭 dropout后者节省显存和计算。训练 10 轮后测试准确率通常能到 99% 以上。如果只有 97% 左右检查归一化参数是否写对、学习率是否过大、batch_size 是否太小导致梯度噪声大。4. 模型保存、加载与单张图片推理4.1 保存训练好的模型文件训练完成后把模型参数保存到磁盘。PyTorch 推荐只保存 state_dict不保存整个模型对象这样加载时更灵活。# 保存模型参数 torch.save(model.state_dict(), mnist_cnn.pth) print(模型已保存为 mnist_cnn.pth) # 加载模型参数需要先实例化模型结构 loaded_model MnistCNN().to(device) loaded_model.load_state_dict(torch.load(mnist_cnn.pth, map_locationdevice)) loaded_model.eval() print(模型加载完成)map_locationdevice的作用是让模型在加载时自动映射到当前设备避免在 CPU 上加载 GPU 保存的模型时报错。保存的文件大约 1.2MB非常轻量。4.2 用单张图片做推理实际使用时你手里可能是一张手机拍的手写数字照片或者从测试集里抽出来的一张图。推理流程是读入图片、转灰度、缩放到 28×28、转张量、归一化、送入模型。from PIL import Image import numpy as np def predict_image(image_path, model, device): # 读取图片并转为灰度 img Image.open(image_path).convert(L) # 缩放到 28x28 img img.resize((28, 28), Image.LANCZOS) # 转为 numpy 数组并归一化到 [0,1] img_array np.array(img, dtypenp.float32) / 255.0 # 应用与训练时相同的归一化 img_array (img_array - 0.1307) / 0.3081 # 转为 tensor增加 batch 和 channel 维度 img_tensor torch.tensor(img_array).unsqueeze(0).unsqueeze(0).to(device) # 推理 with torch.no_grad(): output model(img_tensor) pred output.argmax(dim1).item() prob torch.softmax(output, dim1).max().item() return pred, prob # 使用示例 pred, prob predict_image(my_digit.png, loaded_model, device) print(f预测数字: {pred}, 置信度: {prob:.4f})这里有个容易忽略的细节训练时用了Normalize((0.1307,), (0.3081,))推理时也必须用同样的均值和标准差做归一化否则输入分布和训练时不一致准确率会明显下降。另外如果图片是白底黑字需要先反色因为 MNIST 是黑底白字。提示用手机拍照做推理时先用图像处理把数字区域裁剪出来并居中效果会好很多。直接拿整张照片缩放成 28×28数字会太小模型很难识别。5. 避坑与排查MNIST 训练中最容易翻车的五个地方5.1 torchvision 下载 MNIST 报 404 或超时现象运行datasets.MNIST(downloadTrue)时卡住然后抛出 URLError 或 HTTPError 404。原因torchvision 默认从境外服务器下载国内网络访问不稳定或者官方镜像地址变更导致旧版本 torchvision 的下载链接失效。解决手动下载四个文件——train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz放到./data/MNIST/raw/目录下然后把download参数改为 False。或者升级 torchvision 到最新版下载源可能已更新。5.2 训练 loss 不下降或震荡严重现象训练几个 epoch 后 loss 一直在 2.3 左右徘徊或者上下剧烈跳动。原因学习率过大是最常见的原因。Adam 默认 lr1e-3 在 MNIST 上通常没问题但如果数据预处理有误——比如忘了归一化——输入值在 0 到 255 之间梯度会爆炸。解决先检查transforms.Normalize是否加了。如果加了还震荡把学习率降到 1e-4 试试。另外确认optimizer.zero_grad()在正确的位置。5.3 测试准确率远低于训练准确率现象训练集准确率 99.8%测试集只有 97%。原因过拟合。模型记住了训练样本的细节泛化能力差。解决提高 dropout 比例到 0.5或者加 L2 正则化在 optimizer 里设weight_decay1e-4。数据增强也是有效手段比如随机旋转 ±10 度、随机平移几个像素。5.4 GPU 显存不足报 CUDA out of memory现象batch_size 设大了训练一开始就报显存不够。原因MNIST 模型很小但如果你把 batch_size 设成 1024 甚至更大中间激活值占用的显存会线性增长。解决把 batch_size 降到 64 或 32。MNIST 上 batch_size 对最终准确率影响不大小批量反而有正则化效果。5.5 保存的模型加载后预测结果全一样现象加载模型后推理所有图片都预测成同一个数字。原因保存时用了torch.save(model)保存整个模型对象加载时模型结构或设备不匹配参数没有正确恢复。解决统一用torch.save(model.state_dict(), path)保存参数加载时先实例化模型结构再load_state_dict。加载后调一下model.eval()。6. 把 MNIST 模型用起来从测试集评估到自定义图片推理的完整验证训练完模型、保存好文件之后真正让它产生价值的是推理环节。我一般会做两件事来验证模型是否可靠一是在测试集上跑一遍完整的混淆矩阵看看哪些数字容易混淆二是拿自己手写的数字拍照做端到端测试。先看混淆矩阵的代码from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt def get_all_preds(model, loader, device): all_preds [] all_labels [] model.eval() with torch.no_grad(): for data, target in loader: data data.to(device) output model(data) preds output.argmax(dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(target.numpy()) return np.array(all_preds), np.array(all_labels) preds, labels get_all_preds(loaded_model, test_loader, device) cm confusion_matrix(labels, preds) sns.heatmap(cm, annotTrue, fmtd, cmapBlues) plt.xlabel(Predicted) plt.ylabel(True) plt.show()跑完这个你会看到大部分数字的识别率都在 99% 以上但 4 和 9、3 和 8、5 和 6 之间偶尔会有混淆。这是手写数字本身的模糊性导致的不是模型结构的问题。如果想进一步提升可以针对混淆对做数据增强比如把 4 和 9 的样本做轻微旋转和弹性形变。自定义图片推理的完整流程我习惯封装成一个脚本传入图片路径就输出预测结果和置信度。前面第 4 章已经给了核心函数这里补充一个批量处理的版本import os from pathlib import Path def batch_predict(image_dir, model, device): results [] for img_path in Path(image_dir).glob(*.png): pred, prob predict_image(str(img_path), model, device) results.append((img_path.name, pred, prob)) print(f{img_path.name}: 预测{pred}, 置信度{prob:.4f}) return results # 批量推理示例 batch_predict(./my_digits/, loaded_model, device)这里有个血泪经验手机拍的照片直接缩放成 28×28 效果很差因为 MNIST 的数字是居中且笔画粗细均匀的。我一般会先用 OpenCV 做自适应二值化再找数字轮廓的外接矩形裁剪出来居中放到 28×28 的画布上。这一步预处理做得好推理准确率能从 70% 提到 95% 以上。最后说一个模型文件复用的技巧。如果你在多个项目里都要用这个 MNIST 模型可以把模型定义和加载逻辑封装成一个类初始化时自动加载权重class MnistPredictor: def __init__(self, model_pathmnist_cnn.pth, devicecpu): self.device torch.device(device) self.model MnistCNN().to(self.device) self.model.load_state_dict( torch.load(model_path, map_locationself.device) ) self.model.eval() def predict(self, image_path): return predict_image(image_path, self.model, self.device) # 使用 predictor MnistPredictor(mnist_cnn.pth) print(predictor.predict(test_digit.png))这样在任何 Python 脚本里两行代码就能调用模型不用重复写加载逻辑。我自己的习惯是每个训练好的模型都配一个这样的封装类放在项目根目录的models/文件夹下用的时候直接 import。MNIST 虽然简单但把这套流程跑通之后换成 Fashion-MNIST、CIFAR-10 甚至自定义数据集结构都是类似的——改一下输入通道数、类别数和卷积核数量就行。希望帮到你。本文还有配套的精品资源点击获取
返回列表