新闻详情

新闻详情

首页 / 资讯中心 / 详情

PyTorch MNIST手写数字识别实战:从数据加载到模型训练全流程

发布时间:2026/10/1 17:44:37来源:尧图网络
PyTorch MNIST手写数字识别实战:从数据加载到模型训练全流程
MNIST手写数字识别可以说是深度学习圈子的“Hello World”也是我第一次认真用PyTorch跑通完整训练链路时做的项目。这个任务本身不难数据是28x28的灰度图一共10个类别用PyTorch实现起来一晚上就能见到不错的效果但它几乎囊括了深度学习项目最核心的闭环数据集加载、数据预处理、模型定义、损失计算、梯度回传、参数更新、模型评估与保存。我刚开始学的时候也踩了不少坑比如torchvision下载MNIST时总是404、Windows下DataLoader的无故报错、环境版本不匹配导致安装失败等等。这篇就把我从零开始跑通MNIST的全过程整理出来包括环境搭建、数据处理、模型设计、完整代码和避坑记录。无论你是第一次接触PyTorch还是想快速了解一个完整项目长什么样这篇都适合你照着抄。1. 项目背景与整体思路拆解1.1 为什么选MNIST一个经典项目的不可替代价值MNIST数据集来自美国国家标准与技术研究院全称是Modified National Institute of Standards and Technology database包含60000张训练图片和10000张测试图片内容都是0到9的手写数字。每张图片是28x28像素的灰度图像素值范围从0到255背景偏黑、数字偏白。这个数据集被用到烂大街但它依然是无数人踏入深度学习领域的第一个台阶原因也很简单。第一数据量适中且类别均衡。6万个样本足够训练一个不大不小的模型又不会像ImageNet那样动辄几百GB普通笔记本CPU都能在十分钟内完整跑一轮训练。第二图像尺寸小。28x28意味着模型输入维度非常低全连接网络直接把784个像素拉平就能用卷积网络也能很轻松地设计成几层结构。第三任务目标清晰给定一张图输出它是0到9中的哪一个数字。这是最基础的图像分类任务没有目标检测、没有分割、没有多任务适合把精力集中在理解pipeline上。我个人的体会是MNIST的定位不只是一个“练习题”它更像一个调试工具。很多研究论文和开源项目在做模型验证时都会先用MNIST跑一遍确认代码逻辑没毛病再迁移到更复杂的数据集。所以你把MNIST搞明白了后面换到CIFAR-10、ImageNet自己的数据本质上都是同一套流程的变体数据加载、模型forward、loss回传、评估。这也是为什么我强烈建议新手不要跳过MNIST直接上手复杂项目你一上来就搞ResNet、Transformer大概率会在环境配置和调试上耗掉好几天打击自信心。1.2 技术栈选型为什么是PyTorch而非其他框架初学者经常问一个问题PyTorch和TensorFlow到底选哪个我的答案是如果是从零开始学深度学习选PyTorch理由很直接。PyTorch的编程风格非常Pythonic它把张量计算、自动求导、神经网络模块都组织得很符合直觉。尤其重要的是动态计算图机制你写的每一行张量运算都会被记录成一个动态图方便你在调试时用断点查看中间变量的值。相比之下早期TensorFlow的静态图方式对新手非常不友好你需要用tf.Session()去“运行”一个定义好的图中间想看某个张量到底是什么还得调sess.run非常别扭。虽然现在TensorFlow也支持了eager execution但社区生态和各类论文的实现仍然明显偏向PyTorch。具体到MNIST这个项目PyTorch生态里的torchvision库直接内置了MNIST数据集的下载和读取接口一行代码就能拿到训练集和测试集。再加上nn.Module模型封装、optim包里的各种优化器、DataLoader自动批处理和打乱整个项目代码可以写得非常紧凑且容易读。还有一点很重要PyTorch在学术圈和工业界的占有率都很高你现在学这套框架后面看论文源码、跑开源项目、部署模型都要用到属于投入一次、长期受益的选择。这里我也顺便解释下“PyTorch基础框架”到底由哪些东西组成方便你建立整体感。最底层是torch.Tensor负责存储数据和做矩阵运算然后是torch.autograd通过自动微分记录每次运算的梯度再往上是torch.nn提供了神经网络层、损失函数、激活函数等接着是torch.optim实现了SGD、Adam等优化算法最后是torch.utils.data负责数据集封装、批处理和采样。这五个模块构成了PyTorch日常99%的工作场景MNIST项目就是把它们串起来各自走一遍流程。2. 环境搭建与常见坑位2.1 版本对应Python、PyTorch、CUDA到底该怎么匹配环境搭建是新手最容易卡住的地方很多报错回头查都是版本不匹配导致的。PyTorch、Python和CUDA的版本不能乱配官方在PyTorch官网给出了每个版本对应的组合你最好先看一下自己的环境再动手装。如果你是纯CPU环境那非常简单直接pip install torch torchvision就能装到最新的CPU版本。很多人问“安装pytorch是不是必须装有GPU”答案是否定的。对于MNIST这种小任务CPU跑全连接网络一个epoch也就几秒跑CNN一个epoch在十几秒到几十秒之间完全可接受。你完全可以在没有独立显卡的电脑上用CPU先把逻辑跑通之后再考虑用GPU加速。这个“先跑通再加速”的顺序对于新手来说非常推荐。如果你的电脑有NVIDIA显卡想装GPU版本先运行nvidia-smi看看驱动支持的CUDA版本再根据这个去装对应的PyTorch。以我常用的组合为例PyTorch版本Python版本CUDA版本推荐说明1.13.13.8-3.1111.7稳定适合老项目2.0.13.8-3.1111.7 / 11.8生态成熟兼容性好2.1.x3.8-3.1111.8 / 12.1新特性较多2.2.x3.8-3.1112.1较新适合新环境安装命令一般长这样pip install torch2.0.1cu118 torchvision0.15.2cu118 --extra-index-url https://download.pytorch.org/whl/cu118如果你不想精确指定版本也可以直接去PyTorch官网的Get Started页面选操作系统和CUDA版本它会直接生成安装命令。装完之后一定要验证一下import torch print(torch.__version__) print(torch.cuda.is_available())如果torch.cuda.is_available()返回True那么你的GPU环境基本没问题了。如果返回False多半是装的CPU版或者CUDA版本和驱动不匹配。关于AMD显卡我在WSL里也帮朋友折腾过一段时间。简单说Windows下的PyTorch官方对A卡支持很有限需要走RockML或DirectML的实验性方案WSL里虽然能通过ROCm跑一部分模型但配置复杂、版本要求苛刻非常不适合新手在入门阶段折腾。AMD用户我建议先用CPU跑通项目等理解清楚了整个流程再考虑性能优化的问题不要在环境上打击自己的耐心。2.2 torchvision下载MNIST总是404的解决方案这可以说是MNIST项目中最常见的一个坑了搜索频率特别高。原因在于torchvision的datasets.MNIST接口默认从yann.lecun.com/exdb/mnist/下载原始数据文件这个托管在个人站点的资源有时候不稳定或者因为网络环境原因无法正常访问于是报出HTTP Error 404: Not Found。你可能会觉得奇怪明明数据接口就在torchvision里为什么还会下载失败因为这个接口只是提供了一个下载功能数据并不随torchvision包一起分发而是运行时现场从外部地址拉取。解决办法不外乎三种我逐个说清楚。第一种是手动下载数据处理文件到本地目录。MNIST原始数据其实是四个gz压缩包分别是对应的训练集图像、训练集标签、测试集图像、测试集标签。你找到你项目里设置的root路径在它下面建立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然后在代码里把downloadTrue保留torchvision会先检查raw目录下是否已有这四个文件如果有就直接解压处理不再走网络下载流程。如果实在无法访问原始地址也可以从一些镜像源把对应的gz文件下载回来放到这个目录下本质一样。第二种方式是先跑一次完整下载过程把成功下载的一部分保留下来。有时候404不是一次性的而是某个文件下载失败其余文件都已经成功落盘。你只要重复运行几次代码尽量让四个gz文件都完整落地之后再配合第一种方案的目录结构就能绕过去。这个方法适合网络间歇性失败的情况。第三种是在代码里直接设置downloadTrue但预先处理好processed文件。其实torchvision在raw目录有gz文件后会尝试本地解压生成processed下的pt文件。如果已经存在MNIST/processed/training.pt和test.pt那么即使raw里没有gz文件也能直接加载。很多镜像站点会提供打包好的MNIST数据里面已经包含processed文件你直接放到对应目录就行。我个人最推荐第一种方案手动下载四个gz文件放进raw目录干净利落还不依赖别人的打包是否完整。2.3 WSL和GPU环境下的PyTorch配置经验很多人的Windows机器性能不错但更喜欢在WSL2里面跑深度学习项目这样做的好处是命令行为主隔离环境干净而且能直接调用宿主机的NVIDIA GPU。PyTorch在WSL2里的安装方式和原生Linux几乎一致但有几个坑需要注意。首先WSL2里面不需要再单独安装NVIDIA显卡驱动。驱动是Windows宿主和WSL2共享的你只需要保证Windows里的驱动版本够新即可。然后在WSL2里运行nvidia-smi如果能正确显示显卡信息和驱动版本说明CUDA可见性没问题。接下来正常pip安装GPU版PyTorch就行不需要装所谓的WSL专用CUDA toolkit因为PyTorch的CUDA运行时是自带在wheel包里的。其次很多人习惯在WSL2中使用Anaconda这没问题。建议单独创建虚拟环境比如conda create -n mnist python3.10然后conda activate mnist之后在虚拟环境中用pip装torch。注意不要混用conda的pytorch频道和pip的版本容易把环境搞乱。用pip从官方源安装清晰可控。还有一点容易忽略WSL2默认的网络是NAT模式如果你在Windows侧配置了代理WSL2里访问外网可能会不通这也会导致下载MNIST失败。遇到这种情况优先使用前面说的手动下载方案或者检查WSL里的网络配置是否正常访问外部地址。另外WSL2里使用DataLoader时num_workers0是最稳妥的设置成大于0偶尔会遇到文件系统相关的问题。3. 数据预处理与DataLoader细节3.1 MNIST数据集的正确打开方式Dataset与transformsPyTorch里加载MNIST的代码非常简单核心就是一行torchvision.datasets.MNIST。但这里有几个参数如果不理解后续也容易出bug。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)先说root参数它指定数据集存放的根目录我习惯用./data这样所有数据都集中在项目下的data文件夹里方便打包和清理。trainTrue表示加载训练集False表示测试集。downloadTrue的作用是如果本地没有对应数据就自动下载但这个问题前面说过经常404建议按我第二小节的方法手动处理。transform是整个数据预处理的关键。这里我用了两个步骤。第一步是transforms.ToTensor()它会把PIL图像或numpy数组转成PyTorch张量同时把像素值从0到255缩放到0到1之间。这一步很重要如果你忘了转Tensor模型会直接拿PIL图像对象参与计算一定会报错。第二步是transforms.Normalize((0.1307,), (0.3081,))作用是把像素值从0到1的范围再标准化到接近标准正态分布均值为0.1307标准差为0.3081这两个值是MNIST官方统计出来的全量像素均值和方法差你可以直接借用。为什么需要标准化直观解释是神经网络对输入特征的尺度很敏感。原始灰度图的像素值分布跨度大且不均匀如果不做标准化模型在训练初期梯度会上下震荡收敛速度变慢甚至出现loss不降的情况。标准化之后每个像素分布对齐优化过程会平稳很多。如果以后你自己去处理图片数据集也要养成先算均值和标准差再Normalize的习惯这已经是深度学习流程里的标准操作。还有一个细节需要留意MNIST原始数据是PIL图像Transform在Dataset返回样本时发生作用所以每次迭代拿到的已经是预处理后的Tensor。你可以打印一下train_dataset[0][0].shape会看到torch.Size([1, 28, 28])第一个维度是channel灰度图只有1个通道。3.2 DataLoader参数选择的策略与原因有了Dataset之后下一步就是把它传给DataLoader。DataLoader的作用是把数据集自动切分成一个又一个批次并在训练时按需打乱顺序、多进程加载。from torch.utils.data import DataLoader train_loader DataLoader(train_dataset, batch_size64, shuffleTrue, num_workers0, pin_memoryTrue) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse, num_workers0, pin_memoryTrue)batch_size取值64是常见经验值MNIST这种小图任务64或128都能跑得很快。如果你显存不大32也是可以的。batch_size并不是越大越好太大的batch会导致单次参数更新方向平均化可能收敛到尖锐极小值泛化性变差。太小的batch比如1梯度噪声大收敛不稳定。64是一个很好的折中。shuffleTrue只在训练集使用目的是让每个epoch里样本顺序都不同避免模型学到顺序相关的假规律。测试集不需要shuffleshuffleFalse让评估过程更稳定也方便分析错分样本。num_workers这个参数在Windows下要格外注意。它控制用几个子进程加载数据在Linux下设置2到4能明显提升数据读取速度但在Windows下如果你写的是脚本直接运行num_workers大于0时有可能会报错因为Windows不像Linux那样通过fork创建子进程它会重新导入主模块容易引发冻结问题。对于MNIST这种小数据量num_workers0完全够用根本不会成为瓶颈我不建议新手在这上面折腾。pin_memoryTrue的意思是当数据要传输到GPU时先锁页内存能小幅提升拷贝速度。如果你用CPU训练这个参数设置True没有坏处但在GPU场景下优势更明显。我写了一段简单的数据预览代码方便你确认数据已经正确加载import matplotlib.pyplot as plt import torchvision images, labels next(iter(train_loader)) print(一个batch的shape:, images.shape) # torch.Size([64, 1, 28, 28]) print(label列表长度:, len(labels)) # 64 # 显示前8张图 grid torchvision.utils.make_grid(images[:8], nrow4) plt.imshow(grid.permute(1, 2, 0).numpy()) plt.show()运行这段代码你能直观看到一批手写数字的长相。确认图像和标签对应得上再进入模型设计阶段。4. 模型构建与训练实现4.1 从全连接到卷积两个模型的取舍在MNIST上最省事的是直接用一个全连接网络MLP。把28x28的图片拉平成784维向量经过几层线性变换加ReLU激活最后输出10个类别的logits。这样的模型代码很短训练也快准确率能到97%左右。但当你跑通之后我很建议再试一个简单CNN因为CNN在图像任务上的优势实在太明显了。先看MLP的实现import torch.nn as nn import torch.nn.functional as F class MLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(28 * 28, 256) self.fc2 nn.Linear(256, 128) self.fc3 nn.Linear(128, 10) def forward(self, x): x x.view(-1, 28 * 28) x F.relu(self.fc1(x)) x F.relu(self.fc2(x)) x self.fc3(x) return x注意forward第一行x.view(-1, 28*28)把四维张量变成二维-1表示自动推导batch维度。这个过程就是“展平”如果你忘记做这一步后面的线性层会直接报维度不匹配的错误。再看CNN版本class CNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, stride1, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, stride1, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) self.dropout nn.Dropout(0.25) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x x.view(-1, 64 * 7 * 7) x F.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x这个CNN的结构是两层卷积加池化再接两个全连接。第一个卷积层输入1通道灰度图输出32个特征图第二个卷积层把32个特征图升到64个。每次卷积后接ReLU和2x2的MaxPool2d特征图尺寸从28x28变14x14再变7x7。最后展平得到64773136维向量经过全连接层输出10类。为什么CNN在图像分类上效果更好核心原因是卷积的局部连接和参数共享。卷积核只关注一个小邻域的特征比如边缘、拐角、纹理然后在整个图上滑动同一个核这样既能捕捉局部空间相关性又大幅减少了参数数量。全连接网络则是每个像素和每个神经元全相连一方面参数量爆炸另一方面没有利用图像的二维空间结构。在MNIST这种相对简单的任务上MLP能达到97%CNN能轻松超过99%这个差距在更复杂的图像上会变得更大。两个模型的参数量也有直观对比。用模型的parameters()加总一下MLP大概是20多万个参数上面这个CNN大概是16万个参数。CNN参数更少准确率反而更高这就是卷积结构在图像任务中的优势体现。4.2 训练循环核心逻辑loss、优化器和评估深度学习训练的每个epoch基本就是一个固定套路清空梯度、前向传播、计算损失、反向传播、更新参数。很多人一开始不太理解为什么每次要先调用optimizer.zero_grad()这是因为PyTorch的梯度是累积的如果不清空上一batch的梯度会叠加到当前batch上导致参数更新方向错误。所以这个步骤不能省。损失函数我选用nn.CrossEntropyLoss。这里有个隐藏细节需要注意这个损失函数已经把LogSoftmax和NLLLoss整合在一起了也就是说模型最后一层输出的原始logits可以直接喂给它不需要在模型里手动加softmax。如果你在模型的最后一层加了nn.Softmax反而会造成信息损失因为softmax会把输出压缩到0到1之间再丢给CrossEntropyLoss时它内部还要再做一次log数值上会出现重复计算的问题。优化器我常用Adam学习率设为1e-3。Adam的优势在于它对学习率的敏感度比SGD低不少自带动量机制在MNIST这种任务上收敛快且稳定。如果你是第一次跑1e-3基本不用调就能得到不错的结果。评估流程里准确率是一个最直观的指标。在每个batch的输出中用torch.max(outputs, dim1)拿到预测的类别索引然后和真实标签比较累计正确预测数最后除以总数。注意在评估时要用model.eval()切到评估模式并用torch.no_grad()包裹避免计算图被构建出来浪费内存。我用文字描述一下完整训练流程对于每一个epoch遍历train_loader里所有的batch把图像和标签搬到设备上CPU或GPU调用optimizer.zero_grad()清空梯度用模型算出outputs计算loss调用loss.backward()回传梯度再调用optimizer.step()更新参数。每跑完一个epoch在测试集上评估一次准确率打印出来。这样一个epoch接一个epoch你会发现准确率从90%左右逐渐升到99%以上。4.3 完整可运行代码从数据到保存模型这里我给出一个可以直接复制运行的完整代码包含数据加载、CNN模型、训练循环、验证和模型保存。所有细节我都加了注释你可以一边跑一边对照理解。import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader from torchvision import datasets, transforms import os class CNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, stride1, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, stride1, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) self.dropout nn.Dropout(0.25) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x x.view(-1, 64 * 7 * 7) x F.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x def load_data(): 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) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue, num_workers0, pin_memoryTrue) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse, num_workers0, pin_memoryTrue) return train_loader, test_loader def evaluate(model, loader, device): 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() return correct / total def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(使用设备:, device) train_loader, test_loader load_data() model CNN().to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) epochs 8 for epoch in range(1, epochs 1): model.train() total_loss 0 for images, labels in train_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() train_acc evaluate(model, train_loader, device) test_acc evaluate(model, test_loader, device) print(fEpoch {epoch:02d} | Loss {total_loss / len(train_loader):.3f} | Train Acc {train_acc:.4f} | Test Acc {test_acc:.4f}) torch.save(model.state_dict(), mnist_cnn.pth) print(模型已保存到 mnist_cnn.pth) if __name__ __main__: main()这段代码在CPU上跑8个epoch大约需要几分钟准确率一般能到99%左右。在GPU上会更快每个epoch只需要几秒。注意我把入口函数用if __name__ __main__保护起来这样在Windows下配合num_workers0不会出现多进程问题也方便你在交互式环境里逐段调试。训练结束后会生成一个mnist_cnn.pth权重文件这就是训练好的模型。以后要使用的话加载方式如下model CNN() model.load_state_dict(torch.load(mnist_cnn.pth)) model.eval()注意加载完也要调用model.eval()因为Dropout层在训练和推理时的行为不同。5. 常见问题与避坑实录5.1 新手高频报错速查表我把这个项目里最常见的报错和排查方法整理成了表格你可以直接对照查找。报错或现象可能原因解决思路HTTP Error 404: Not Found 或 WinError 10054MNIST数据源临时不可用或网络中断手动下载四个gz文件放入data/MNIST/raw再重新运行代码FileNotFoundError: MNIST/raw/xxx 不存在数据没有下载成功且download未触发检查root路径确保downloadTrue或手动放置数据RuntimeError: DataLoader worker (pid xxx) exited unexpectedlyWindows下num_workers大于0导致的问题设置num_workers0并把主逻辑放入ifnamemain训练时loss始终在0.69左右徘徊模型输出有问题或者没有数据标准化确认transform里有ToTensor和Normalize尝试把学习率调低ValueError: Expected input batch_size to match target dimension标签和输出维度对不上检查最后一层输出维度是否为10检查DataLoader返回的标签形态CUDA out of memory单batch显存占用过大减小batch_size如果只是跑MNISTCPU就足够torch.cuda.is_available()返回False装了CPU版PyTorch或CUDA驱动版本不匹配重装对应GPU版先跑nvidia-smi确认驱动信息显示“Module torch has no attribute device”代码里torch.device写错了或torch版本过老确认安装的是完整版PyTorch而不是某些阉割包准确率一直上不去停在90%左右用了太浅的模型或epoch不够多试试CNN结构适当增加训练轮数这里我最想强调的是第一个404问题它真的困住了很多人。如果你手头网络无法直接访问数据源就不要一直重试downloadTrue了手动下载文件放到指定目录是最快最稳妥的路。另一个容易被忽略的是Windows下如果在Jupyter Notebook里跑num_workers0的DataLoader有时候报错不明显但就是不输出我建议所有Windows用户一律先设num_workers0等以后上Linux服务器再调大也不迟。5.2 如何从训练日志中判断模型状态训练日志里最核心的两个数值就是train loss和test acc。观察它们的相对趋势能帮你判断模型是欠拟合还是过拟合。如果train loss一直在下降但test acc也同步上升这是理想状态说明模型在学到有用的特征。如果train loss降得很低比如0.01但test acc停滞甚至下降那大概率是过拟合了。模型把训练集记得太死对没见过的测试数据反而表现不好。解决办法包括增加Dropout、增大数据集、使用数据增强、提前停止训练或降低模型复杂度。反过来如果train loss下降很慢test acc也很低可能是欠拟合。此时通常的做法是增加模型容量、加深网络层数、增大训练轮数或者看看学习率是否设置得太低导致收敛缓慢。在MNIST这个项目上因为数据集简单、模型容量也不算小我观察到的现象通常是第1个epoch结束时test acc已经能达到97%以上后面几个epoch缓慢爬升到99%。如果遇到第1个epoch准确率就卡在90%附近不上不下先检查预处理有没有出问题比如是否忘了Normalize或者模型最后一层输出维度是不是写成了784。很多时候不是模型设计有问题而是数据流向里的一个小bug。还有一个小建议训练时把每一步的loss打印出来观察而不是只等epoch结束。如果发现loss在某个值附近震荡不下降先不要急着改网络结构试着把学习率减小10倍再观察几十个step很多时候问题就出在学习率上。我用Adam跑MNIST默认1e-3就可以但换到别的数据上可能要重新调。5.3 入门阶段不得不说的几个习惯项目跑通之后我希望你能养成几个后续会受益的习惯。第一个是固定随机种子。在训练前设置torch.manual_seed(42)保证每次跑的初始条件一致这样调参时你才确定效果差异来自代码改动而不是随机性。第二个是定期保存checkpoint不只是训练完保存最终模型。每隔几个epoch保存一次比如mnist_cnn_epoch8.pth这样一旦后面训练出问题你还能回到之前效果较好的版本。第三个是可视化中间结果。把测试集里预测错的样本打印出来看看模型究竟输错在哪有时候一眼就能发现数据标签或预处理的问题。我见过不少人跑完MNIST就算结束了我觉得其实可以再往前走一步把模型输出层的10个类别置信度用softmax打印出来挑一张看起来特别潦草的手写数字看看模型给出的概率分布是什么样的。你会发现即使预测正确模型对相近数字比如4和9的置信度也会接近这能帮你直观理解“分类置信度”这个概念。之后再去看FashionMNIST、CIFAR-10都会顺畅很多。6. 几个值得尝试的后期扩展方向当你在MNIST上跑出99%以上的准确率后我个人建议你按下面的方向逐层递进把这次项目经验的边界拓宽一些。第一个方向是换数据集。和MNIST结构几乎一样但更有挑战性的是FashionMNIST同样28x28灰度图但内容变成了T恤、裤子、鞋等时尚品类分类难度比数字高一些能逼着你在模型上做一些小改动比如加深网络层数。然后是CIFAR-10图像变成32x32彩色图类别是飞机、汽车、鸟等这时你会发现简单的两层小CNN开始吃力需要引入更现代的网络结构。第二个方向是给模型增加数据增强。MNIST本身可以用随机旋转、随机平移、随机权重抖动等方式扩充样本让模型更鲁棒。torchvision的transforms里提供了RandomRotation、RandomAffine等现成接口玩起来很有乐趣。你会发现随着数据增强力度加大训练集上的准确率会下降但测试集准确率可能反而上升这就是数据多样性的价值。第三个方向是尝试可视化卷积核的特征图。可以把模型第一层卷积的权重打印成图像看看它学了哪些边缘或纹理模式也可以把某一层的特征图画出来观察激活情况。这个过程对理解“卷积神经网络到底学到了什么”特别有帮助比看一堆数学公式直观得多。第四个方向是保存推理脚本做一个小型可交互演示。把自己的手写数字画到28x28的网格里用训练好的模型实时预测。这种从“跑通教程”到“完成一个能交互的小工具”的转变会在很大程度上提升你对项目掌控的信心。这些都是我在MNIST跑通之后陆续做过的事每一样都不复杂但对夯实基础很有用。项目可以从MNIST毕业但你对深度学习数据流和训练机制的掌握才刚刚开始。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

LLM Agent实战:用回形针思想实验剖析奖励函数与AI安全对齐 2026/10/1 20:17:53

LLM Agent实战:用回形针思想实验剖析奖励函数与AI安全对齐

最近我把“paperclip”(回形针)这个词丢进了一个技术实验里,结果收获远比想象中大。你可能在物理课上听过那句经典提问——为什么不能用一枚回形针的价钱把人造卫星送上太空,也可能在AI安全讨论里看到过“回形针最大化工”这个思想…

阅读更多 →
万集科技与中国人寿宣讲会:技术岗与业务岗的求职攻略 2026/10/1 20:17:47

万集科技与中国人寿宣讲会:技术岗与业务岗的求职攻略

1. 宣讲会前,先把这两家企业的底牌摸清楚10月24日这场宣讲会,来的两家单位放在一起看挺有意思。万集科技是典型的硬科技赛道选手,做智能交通、激光雷达、ETC设备起家;中国人寿北京分公司则是金融保险领域的头部机构。一个偏技术研…

阅读更多 →
OpenClaw多agent数字分身接入飞书:TaoToken统一Key配置与验证 2026/10/1 20:17:47

OpenClaw多agent数字分身接入飞书:TaoToken统一Key配置与验证

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

阅读更多 →
九坤量化开源IQuest-Coder-V1:代码大模型“流式”训练实战拆解与TaoToken接入 2026/10/1 20:17:46

九坤量化开源IQuest-Coder-V1:代码大模型“流式”训练实战拆解与TaoToken接入

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

阅读更多 →
6个神级Skill,让Agent原地开挂:字幕变动画、图片变3D、长内容变知识库|TaoToken统一Key实战 2026/10/1 20:17:46

6个神级Skill,让Agent原地开挂:字幕变动画、图片变3D、长内容变知识库|TaoToken统一Key实战

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

阅读更多 →
从float到inf:技术科普如何用推演链传递思考能力 2026/10/1 20:17:46

从float到inf:技术科普如何用推演链传递思考能力

半夜两点四十分,我本来已经关掉电脑准备睡觉。屏幕突然亮了一下,是推荐流里弹出一条视频:“为什么float(1e308)在Python里会变成inf?这个细节我查了一下午”。 说实话,这种标题我平时根本不会点——太像技术营销号的标…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉