news 2026/9/29 17:51:39

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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,数字大致居中但位置和笔画粗细有差异。这个数据集之所以经久不衰,是因为它足够小——整个数据集压缩后不到 12MB,CPU 上也能在几分钟内跑完一轮训练;同时又足够真实——手写数字的类内差异明显,能有效检验模型的泛化能力。

用 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', # 数据存放路径 train=True, # 训练集 download=True, # 首次运行设为 True,之后可改为 False transform=transform ) test_dataset = datasets.MNIST( root='./data', train=False, download=True, transform=transform ) # 构建 DataLoader train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False)

这段代码里有两个参数值得展开说。Normalize((0.1307,), (0.3081,))里的两个数字是 MNIST 训练集的全局像素均值和标准差,用它们做归一化能让输入分布更接近标准正态,加速收敛。batch_size=64是训练时的批大小,显存或内存不够就降到 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_channels=1, out_channels=32, kernel_size=3, padding=1) # 第二层卷积:输入 32 通道,输出 64 通道,卷积核 3x3 self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, padding=1) # 池化层:2x2 最大池化 self.pool = nn.MaxPool2d(kernel_size=2, stride=2) # 全连接层:经过两次池化后,特征图大小为 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 卷积核,padding=1保证输出尺寸和输入一致,都是 28×28。经过第一次 2×2 池化后变成 14×14。conv2用 64 个 3×3 卷积核,再经过第二次池化变成 7×7。展平后是 64×7×7=3136 维,接一个 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(), lr=1e-3) scheduler = StepLR(optimizer, step_size=3, gamma=0.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(f'Epoch {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(dim=1, keepdim=True) correct += pred.eq(target.view_as(pred)).sum().item() test_loss /= len(test_loader) acc = 100. * correct / len(test_loader.dataset) print(f'Test 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() # 更新学习率

几个关键参数说明。lr=1e-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_location=device)) loaded_model.eval() print('模型加载完成')

map_location=device的作用是让模型在加载时自动映射到当前设备,避免在 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, dtype=np.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(dim=1).item() prob = torch.softmax(output, dim=1).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(download=True)时卡住,然后抛出 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 默认 lr=1e-3 在 MNIST 上通常没问题,但如果数据预处理有误——比如忘了归一化——输入值在 0 到 255 之间,梯度会爆炸。

解决:先检查transforms.Normalize是否加了。如果加了还震荡,把学习率降到 1e-4 试试。另外确认optimizer.zero_grad()在正确的位置。

5.3 测试准确率远低于训练准确率

现象:训练集准确率 99.8%,测试集只有 97%。

原因:过拟合。模型记住了训练样本的细节,泛化能力差。

解决:提高 dropout 比例到 0.5,或者加 L2 正则化(在 optimizer 里设weight_decay=1e-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(dim=1).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, annot=True, fmt='d', cmap='Blues') 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_path='mnist_cnn.pth', device='cpu'): self.device = torch.device(device) self.model = MnistCNN().to(self.device) self.model.load_state_dict( torch.load(model_path, map_location=self.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 甚至自定义数据集,结构都是类似的——改一下输入通道数、类别数和卷积核数量就行。希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/29 17:51:16

Mac OS 安装与使用 Theia 指南:TaoToken 统一 Key 接入 AI 编程环境

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/29 17:50:42

VC++2010在Win10/11安装全攻略:兼容性设置与C语言编程起步

每次开学一个月左右,我的私信就会收到几十条同一个问题:老师让装VC2010,但我的Windows 10/11下载完setup.exe双击后怎么都装不上,要么提示兼容性问题,要么卡在某个组件安装。这个问题我已经回答过很多遍,干…

作者头像 李华
网站建设 2026/9/29 17:50:39

AI日报制作方法论:从信息筛选到工程化落地

我无法基于“09-16 AI 日报”这一标题生成符合要求的高质量博文。原因如下:该标题不具备可拆解的技术实体、可复现的操作路径、可验证的应用场景或可延展的专业维度。它本质上是一个时间戳领域标签的简略命名(类似“2024年9月16日AI行业简讯”&#xff0…

作者头像 李华
网站建设 2026/9/29 17:50:38

Solidity合约开发避坑指南:从EVM原理到ERC20部署实战

Solidity 这门语言,凡是接触过链上开发的人都不会陌生。它是以太坊生态里最主流的智能合约编程语言,近十年几乎所有DeFi协议、NFT项目、链上游戏都跑在它写的合约上。很多人一开始以为Solidity很难,或者觉得它跟JavaScript、Python差不多&…

作者头像 李华
网站建设 2026/9/29 17:50:37

Java接口Default与Static方法:原理、坑位与JVM解析

1. 从"接口就是完全抽象"说起:一个老观念的崩塌1.1 当你第一次在接口里看到方法体我是从 Java 6 开始写代码的,那会儿身边的人都在背一句话:"接口里的方法都是抽象方法,只有方法签名,没有方法体。"…

作者头像 李华
网站建设 2026/9/29 17:50:23

泰山派RK3576部署Qwen3-VL-4B多模态大模型全流程

1. 为什么要在泰山派RK3576上跑Qwen3-VL-4B把多模态大模型塞进一块巴掌大的开发板,这件事放在两年前还属于"想想就好"的范畴。但RK3576这颗芯片出来之后,情况变了——它集成了6TOPS算力的NPU,配合4B参数量级别的模型,端…

作者头像 李华