十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

PyTorch实战指南:MNIST手写数字识别训练全流程与调优

PyTorch实战指南:MNIST手写数字识别训练全流程与调优 1. 环境准备动手前先把pytorch装明白1.1 版本匹配的逻辑不是越新越好很多新手上来就先装最新的pytorch结果一路报错到怀疑人生。我在实际项目中踩过不少坑这里直接给大家一套我验证过的稳妥方案。先说结论如果你用python 3.10.11装pytorch 2.8.0配CUDA 12.1这套组合实测下来兼容性非常稳。为什么是这个组合因为pytorch的版本和CUDA版本之间有严格的对应关系不是随便配就能跑的。CUDA版本太低新版pytorch根本用不了GPU加速CUDA版本太高显卡驱动又可能不支持。这套组合是官方文档里明确标注的稳定搭配我自己在训练MNIST项目时用过无数次从来没有因为版本问题卡住过。安装方式我推荐用pip而不是conda。虽然conda在管理环境方面确实方便但pytorch的pip安装包更新更及时而且对CUDA版本的支持也更精确。用conda装的时候经常会遇到版本依赖冲突尤其是你机器上还有其他深度学习框架的时候那真是剪不断理还乱。pip就干净利落多了pip install torch2.8.0 torchvision0.19.0 torchaudio2.8.0 --index-url https://download.pytorch.org/whl/cu121注意这里的cu121就是CUDA 12.1的标识。torchvision这个库很多人会忽略但它内置了MNIST这类经典数据集的下载接口我们今天就用得上所以必须一起装。装完以后验证一下是否真的能用GPUimport torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU mode)这一步不能省。要是torch.cuda.is_available()返回False那后面训练MNIST虽然也能跑但速度会慢不少。如果我这段代码跑出来是True恭喜你GPU环境OK了。1.2 CPU训练也能跑但GPU是分水岭MNIST数据集本身很小每张图片只有28x28像素总共6万张训练图片1万张测试图片。说实话这样一个数据集用纯CPU训练也不是不行但体验就是两回事了。我用CPU训过这个模型一个epoch大概要1分多钟10个epoch下来就是十几分钟。用GPU之后一个epoch几秒钟就完事全程不到半分钟。这就是为什么环境准备阶段要花心思把GPU环境配好而不是随便装个CPU版就开干。训练模型的时间成本很多新手意识不到有多重要。你后面要调参、改网络结构每一次改动都得重新训练如果单次训练就要十分钟起那你调参的热情会被迅速消磨干净。另外多说一句如果你是用笔记本跑GPU训练时风扇会狂转这是正常的。但如果温度过高导致训练中断或者蓝屏那就是散热问题了建议把batch size调小一点给GPU降压。2. MNIST数据集神经网络界的Hello World2.1 这个数据集到底长什么样MNIST全称是Modified National Institute of Standards and Technology Database就是美国国家标准与技术研究院的手写数字数据库。它收集了大量人手写的0到9的数字图片每张图片被归一化到28x28像素灰度图像素值范围是0到255。说句实话这个数据集放到现在已经有点“老掉牙”了各种模型在上面都能刷到99%以上的准确率。但为什么每个学深度学习的人都要用它练手因为它足够简单、足够标准、足够容易验证。你搭的神经网络是不是有问题跑到MNIST上一看准确率就一目了然。就像学编程先写Hello World一样你不可能上来就写操作系统内核。MNIST的数据结构我给大家拆一下训练集有60000张图片和对应的标签测试集有10000张图片和对应的标签。标签就是0到9的数字也就是这张图片里手写体是几个几。在pytorch里我们用torchvision.datasets.MNIST这个接口来加载它会自动帮你下载、解压、转换成tensor。2.2 torchvision下载404的终极解法我用torchvision.datasets.MNIST训练的时候经常遇到下载404的问题。这个问题的根源是mnist数据集官方源在某些网络环境下无法访问或者访问超时。网上很多教程对这一块要么避而不谈要么让你用代理都不是特别让人满意的方案。我给大家一个我自己在项目中实际验证过的解法。首先说清楚这个404不是pytorch版本问题是网络环境问题。解决办法是在datasets.MNIST外面手动指定数据源from torchvision.datasets import MNIST # 下载失败时手动指定镜像源 train_dataset MNIST( root./data, trainTrue, transformtransforms.ToTensor(), downloadTrue, # 重点指定可用的镜像源 # 如果上面的download仍报404, 则先手动下载好四个文件放到 ./data/MNIST/ Raw 目录 )更稳的方法其实就是绕开自动下载自己手动下载文件。MNIST官网或者一些高校的镜像站提供了四个文件训练集的图片和标签测试集的图片和标签文件名分别是train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz。你把它们手动下载下来解压之后放到./data/MNIST/raw/目录下再设置downloadFalsetorchvision就会直接读取本地文件不会再走网络。2.3 DataLoader到底在干什么数据加载这一步最核心的组件是DataLoader。很多教程就直接丢一行代码说“这是加载数据”根本没讲清楚背后的机制。我用自己的话来讲一下。DataLoader做的事就是“批量取数据”。训练神经网络的时候我们不会一次把全部6万张图片都扔进去因为内存根本扛不住而且梯度下降的效率也不高。我们通常用“小批量随机梯度下降”也就是每次随机取一小批图片比如64张或128张算完梯度更新一次参数然后再取下一批。这个“一小批”就是batch。shuffleTrue这个参数也很关键。如果不打乱数据的顺序每次训练看到的都是排列整齐的数据模型可能会学到数据集内部的秩序特征而不是真正的手写数字特征。打个比方如果你学英语的时候永远按照字母表顺序背单词那你背到Z的时候可能已经把A忘光了而且考试的时候单词也不会按顺序出题。数据乱序就是要让模型每一次“见到”的上一个数字和下一个数字之间毫无规律逼它去学真正重要的特征。num_workers参数是开几个子进程来读取数据。CPU是4核的你可以设成4是8核的可以设成8。不要设太大否则会抢占训练资源反而拖慢速度。我实测下来num_workers2在大多数机器上表现不错。3. 神经网络设计从零搭一个能用的小模型3.1 全连接网络还是卷积神经网络听到“神经网络”这个词很多人的第一反应就是“深度学习一定得用卷积神经网络CNN”。这个理解不算错但在MNIST这个任务上我会先从一个全连接网络讲起。为什么因为全连接网络的结构最简单代码直白非常利于理解神经网络的核心计算流程输入层 - 隐藏层 - 输出层 - 反向传播更新参数。你把这个流程吃透了后面再加卷积层、池化层、dropout层都只是在这个基础上加积木而已。我做了一次对比实验一个简单的全连接网络两层结构第一层把784维输入映射到128维第二层把128维映射到10维输出。经过5个epoch训练后测试集准确率大约是97%。而一个简单的CNN两层卷积加一个全连接层同样的训练轮数下准确率能到99%。两者都能用但差距是实实在在的。如果你今天是第一次接触pytorch我的建议是先搭全连接网络跑通整个流程看到准确率稳步上升然后再改成CNN感受一下特征提取能力带来的提升。这个由简到繁的过程比一上来就抄一个ResNet的代码要有效得多。3.2 每一层大小和参数量的计算全连接网络的输入一定是784维这个数字怎么来的28x28784就是图片拉成一维向量的长度。很多人不知道这个784的来龙去脉只知道是784就拿来用。现在我说明白之后你以后再看到类似的数据比如CIFAR-10是32x32x3就能举一反三了。第一层隐藏层我把维度设为128。这个数字不是随便定的我试过64、128、256三组对比128在准确率和训练速度之间取得了最好的平衡。64表达能力偏弱准确率会掉一截256提升有限但训练时间明显变长。这种对比实验其实花不了多少时间但我建议新手不要跳过它这是个很宝贵的感受过程。第二层是输出层维度必须是10。因为我们要识别0到9十个数字每一维代表预测为对应数字的概率所有维度加起来等于1softmax层做的。模型的参数量可以算一下第一层权重矩阵是784x128加上128个偏置第二层权重矩阵是128x10加上10个偏置。总计约10.1万个参数。这个量级的参数量在现代深度学习里可以说轻如鸿毛但足以完成手写数字识别任务也从侧面说明MNIST任务本身并不复杂。CNN版本的模型设计我会这样做class CNN(nn.Module): def __init__(self): super(CNN, self).__init__() self.conv1 nn.Conv2d(1, 32, 3, padding1) # 28x28 - 28x28 self.conv2 nn.Conv2d(32, 64, 3, padding1) # 28x28 - 28x28 self.pool nn.MaxPool2d(2, 2) # 28x28 - 14x14 self.fc1 nn.Linear(64 * 14 * 14, 128) self.fc2 nn.Linear(128, 10) self.relu nn.ReLU() def forward(self, x): x self.pool(self.relu(self.conv1(x))) x self.pool(self.relu(self.conv2(x))) x x.view(-1, 64 * 14 * 14) x self.relu(self.fc1(x)) x self.fc2(x) return x这里padding1的作用是保持卷积后尺寸不变否则经过一次3x3卷积28x28会变成26x26经过两次之后再池化尺寸就乱了。3.3 激活函数为什么非用ReLU不可很多教程在代码里直接写了F.relu(x)但从不解释为什么。我在这里把话说透如果不加激活函数不管你的神经网络有多少层它实际上都等价于一层线性变换永远学不了非线性关系。手写数字识别这个任务本质上是非线性的。“一个像素亮不亮”和“这个数字是几”之间不是简单的直线关系而是极其复杂的函数映射。激活函数给神经网络注入了非线性表达能力让它可以拟合任意复杂的函数映射这是神经网络最核心的法宝。为什么选ReLU而不是经典的sigmoid或者tanh我在实践中体会最大的原因有两个。第一ReLU的计算非常简单函数式是max(0, x)正向传播和反向传播都快第二sigmoid和tanh在输入绝对值大的时候导数趋于0梯度会消失模型训练不动。ReLU在正区间导数恒为1能有效缓解梯度消失问题。当然ReLU也有“神经元死亡”的问题即某个神经元所有输入都为负梯度永远为0这个神经元就废了。应对方案是调小学习率或者换LeakyReLU。我在MNIST项目里用ReLU没遇到过严重的神经元死亡所以先用它够用。4. 核心代码实现训练流程一步步走4.1 关键组件与完整代码核心代码其实分成这么几块数据准备、模型定义、损失函数与优化器定义、训练循环、评估循环。我来逐块拆开讲。首先是最关键的“损失函数”和“优化器”。MNIST是分类问题损失函数我用交叉熵损失nn.CrossEntropyLoss()。这个损失函数衡量的是“模型预测的概率分布”和“真实标签分布”之间的距离数字越小代表预测越准。它其实内部已经包含了softmax操作所以模型最后一层输出原始逻辑值logits就好不需要手动再套一个softmax。优化器我用torch.optim.Adam。至于为什么不用最基础的SGD原因也很直接Adam自带自适应学习率每个参数的学习率会根据梯度大小自动调整。这意味着它对初始学习率不那么敏感对新手极其友好。我用SGD的话还得手动调momentum、调学习率衰减策略麻烦。当然SGD也不是没有优势它最终的收敛效果在一些场景下比Adam好有正则化效应。但在MNIST这个任务上Adam又快又稳直接选它没毛病。这里先给出完整的训练代码import torch import torch.nn as nn import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader # 数据预处理 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载数据集 train_dataset datasets.MNIST(root./data, trainTrue, transformtransform, downloadTrue) test_dataset datasets.MNIST(root./data, trainFalse, transformtransform, downloadTrue) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse) # 定义模型 class NeuralNet(nn.Module): def __init__(self): super(NeuralNet, self).__init__() self.fc1 nn.Linear(28*28, 128) self.fc2 nn.Linear(128, 10) def forward(self, x): x x.view(-1, 28*28) x torch.relu(self.fc1(x)) x self.fc2(x) return x model NeuralNet().cuda() # 有GPU就用cuda criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) # 训练循环 for epoch in range(10): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.cuda(), labels.cuda() outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() avg_loss running_loss / len(train_loader) print(fEpoch [{epoch1}/10], Loss: {avg_loss:.4f})这段代码跑起来你会看到loss从最初的2.3左右一路下降到第10个epoch结束大约是0.1左右。loss下降说明模型在持续学习这是训练正常的最直接信号。4.2zero_grad的意义和踩坑训练代码中有一行不可少的代码optimizer.zero_grad()这行代码的价值很多人没讲清楚。pytorch的梯度是“累积”的。什么意思就是每次backward()计算出来的梯度会累加到参数的grad属性上不会自动清零。如果不清零下一轮batch的梯度就会和上一轮的梯度叠加导致参数更新方向完全错乱loss可能不降反升。我第一次写训练代码的时候就忘了这一行结果loss像过山车一样忽上忽下怎么调学习率都没用。后来加上了zero_grad()问题瞬间解决。所以这是一个“只差一行代码效果天差地别”的经典案例。4.3 评估模型的正确方式训练完成后我们要在测试集上评估模型效果。MNIST测试集有1万张图模型从没见过这些图在测试集上的准确率才真正代表模型的泛化能力。评估代码比训练代码简单核心是torch.no_grad()这个上下文管理器model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.cuda(), labels.cuda() outputs model(images) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() print(fTest Accuracy: {100 * correct / total:.2f}%)这里有几个关键点。model.eval()是切换到测试模式会关闭dropout和batch normalization在训练时的行为。torch.no_grad()告诉pytorch“我这边只做前向推理不需要计算梯度”能省下大量内存和计算时间。torch.max(outputs.data, 1)返回每一行最大的值和对应的索引索引就是模型预测的数字。跑完这个评估全连接网络下准确率大概在97%左右CNN则在98.5%以上。如果准确率低于95%那大概率是训练有问题或者模型结构有问题建议回头检查训练过程的loss有没有真正下降。5. 训练细节调优从97%到99%的差距在哪里5.1 数据标准化的作用我前面给的代码里有一行transforms.Normalize((0.1307,), (0.3081,))这两个数字不少教程直接抄过来但没人讲为什么是这两个。这个0.1307是MNIST训练集所有像素的均值0.3081是标准差。因为原始图片是0到255之间的整数直接送进网络的话数值分布范围太大不利于梯度下降。做了标准化之后数据均值变成0方差变成1模型收敛会更快更稳。如果你不标准化直接用原始像素值去训练模型也能收敛但可能会慢一些。这个差距在小数据集上不明显但标准化是一个良好习惯尤其在处理图像数据时应该成为固定动作。不同数据集的均值和标准差是不同的比如CIFAR-10就有三通道的均值和标准差。所以这两个数不是万能的用其他数据集时要重新计算。5.2 batch size和学习率的联动关系batch size和学习率这两个参数不是独立的它们之间存在着内部联系。我的经验是batch size越大模型梯度方向越稳定学习率可以适当调大batch size越小噪声越大学习率要相应调小否则容易震荡。原因很简单batch size大每次计算的梯度更接近整个数据集的真实梯度方向走得更稳步子可以迈大点batch size小每次梯度方向抖动大步子大了容易跑偏。在MNIST上我用batch size64配合学习率0.001Adam效果最好。如果batch size改成256我一般会把学习率提到0.002batch size改成16学习率降到0.0005。这个频谱关系大家可以记一下以后换数据集调参会少走很多弯路。5.3 训练轮数到底该训练多少个epoch很多新手会问训练多少个epoch合适。答案是看验证集的准确率不再提升了就停这个技巧叫“早停法”。我在MNIST任务上简单实验了一下10个epoch全连接网络准确率已经稳定在97%左右后面再加20个epoch也提升不了多少了。CNN会稍微好一点15-20个epoch能到99%。如果训练过程中loss已经很低但测试准确率停滞不前再训练下去也没有意义反而可能过拟合测试准确率反而下降。判断是否过拟合的办法很简单训练准确率远高于测试准确率比如训练99.9%测试96%就是过拟合信号。过拟合的意思就是模型把训练集的“个性特征”背下来了却没有抓住数字的“共性规律”。这时候需要加正则化Dropout、L2、增强数据或者减小模型容量而不是无限堆epoch。6. 可视化与结果分析模型训练好坏的直观判断6.1 损失曲线的画法和解读训练过程中我强烈建议把loss记录下来画成曲线。这一步对调试模型特别有价值因为它能直观反映模型有没有在学习。示例代码import matplotlib.pyplot as plt losses [] # 在训练循环中记录: losses.append(avg_loss) plt.plot(range(1, len(losses)1), losses, markero) plt.xlabel(Epoch) plt.ylabel(Loss) plt.title(Training Loss Curve) plt.grid(True) plt.show()正常的loss曲线应该是这样的第1个epoch从2.3左右急剧下降到0.5左右后续epoch缓慢下降并最终趋于平稳。如果loss曲线是平的模型压根没在学习就赶紧检查代码如果loss先降后升模型在过拟合了如果loss在一个高位反复震荡模型不收敛学习率大概率调大了。6.2 看几张预测错误的图比看准确率更有收获准确率只是一个冷冰冰的数字真正有价值的是那些被模型分错的样本。它们能告诉我们模型错在了哪里为什么错。我把“预测错误”的图片集中打印出来发现一个规律错误的样本大多是手写质量很差的数字比如一个7写得潦草至极上半截看起来像1连人眼都分辨困难。模型犯错是可以理解的。但如果你发现一些正常人眼都能轻易认出的数字被模型分错那说明模型还有改进空间比如增加训练轮数、加深网络、做数据增强等。这里有个很实用的操作展示错误样本的代码片段import matplotlib.pyplot as plt wrong_images [] wrong_labels [] wrong_preds [] model.eval() with torch.no_grad(): for images, labels in test_loader: images, labels images.cuda(), labels.cuda() outputs model(images) _, predicted torch.max(outputs, 1) is_wrong predicted ! labels if is_wrong.any(): wrong_images.extend(images[is_wrong].cpu().numpy()) wrong_labels.extend(labels[is_wrong].cpu().numpy()) wrong_preds.extend(predicted[is_wrong].cpu().numpy()) # 展示前5张错误样本 for i in range(min(5, len(wrong_images))): plt.imshow(wrong_images[i].squeeze(), cmapgray) plt.title(fTrue: {wrong_labels[i]}, Pred: {wrong_preds[i]}) plt.show()6.3 混淆矩阵看清每个数字的识别bin准确率99%只是整体结论但具体到每个数字识别难度是完全不同的。这时候要用混淆矩阵它是一张10x10的表格行代表真实标签列代表预测标签交叉点的数字表示有多少个真实类别为i的样本被预测成了j。比如混淆矩阵的第7行第1列数字比较大说明7经常被识别成1这两个数字在书写上确实很像模型的混淆是可以理解的。如果矩阵的非对角线元素整体都比较大说明整个模型还有很大的提升空间。pytorch中统计混淆矩阵很简单import numpy as np confusion np.zeros((10, 10), dtypeint) with torch.no_grad(): for images, labels in test_loader: images, labels images.cuda(), labels.cuda() outputs model(images) _, predicted torch.max(outputs, 1) for t, p in zip(labels.cpu().numpy(), predicted.cpu().numpy()): confusion[t][p] 1 print(confusion)拿到混淆矩阵之后对角线上的数字应该是压倒性的优势非对角线上的小数字就是模型的错误分布。我实测发现模型在数字“9”和“4”上的错误率相对高一些因为这两个数字的手写体在倾斜度、收笔方式上有很多相似之处。7. 常见问题与坑位指南7.1 “CUDNN_STATUS_NOT_INITIALIZED”这类异常我一开始跑代码的时候经常遇到RuntimeError: cuDNN error: CUDNN_STATUS_NOT_INITIALIZED。这个问题多半是显存不够或者GPU被其他进程占满了。解决办法是重启内核把之前没释放的显存清掉。用nvidia-smi看一下GPU占用情况如果发现显存被占满但自己什么都没跑那就是上次内核崩溃的僵尸进程直接kill -9干掉它。7.2 在jupyter里模型一直占着显存一运行训练代码显存就满然后代码结束显存也不释放。这个问题的根源是pytorch的缓存机制它为了加速下一次计算会保留一部分显存不还给系统。解决办法是在代码末尾加一行import torch torch.cuda.empty_cache()不过这个命令只是清空“缓存”如果有实际变量还持有显存是清不掉的。真正要释放显存得把所有模型和张量都del掉。7.3 下载MNIST数据集时卡住或404如果downloadTrue时卡住很久没反应大概率是网络连不上官方源。解法就是前面讲的手动下载文件后放到本地目录再设置downloadFalse。有朋友照着网上的方法从某个镜像下载后代码还是不认原因多半是目录姿势不对。注意文件一定要放在./data/MNIST/raw/这个目录下文件名不能改动后缀.gz不要手动解压成原始文件torchvision读取时自己会解压。我见过好几个人把文件解压后放进目录反而报错就是这个原因。7.4 模型加载时shape不匹配有朋友跑完训练保存了模型下次加载继续做迁移学习结果报shape不匹配维度对不上。这个问题一般是保存的模型输入输出和自己重新定义的类不一致导致的。比如你保存时输入是2828784新建模型时用了128128的输入维度当然不匹配。建议用torch.save(model.state_dict(), mnist_model.pt)加载时用model.load_state_dict(torch.load(mnist_model.pt))并且保证两个模型定义完全一样。7.5 常见问题速查表问题现象可能原因优先排查方向torch.cuda.is_available()为FalseCUDA和pytorch版本不匹配按cu121版本重装pytorchloss不下降忘记zero_grad检查optimizer.zero_grad()是否存在训练集准确率高但测试集低过拟合加dropout或减小模型容量显存不足batch_size太大调小batch_size或换小模型MNIST下载404网络问题手动下载后本地加载准确率只有10%左右模型输出维度不对确认输出层维度是否为108. 最后的个人实践体会写了这么多最后分享几点我在实操中反复体会到的收获。MNIST这个项目虽然小但它涵盖了深度学习的完整流程数据准备、模型设计、损失函数选择、优化器配置、训练、评估、可视化、调优。把这套流程跑通了你再看其他数据集的代码基本一眼就能看懂大概不会再有无从下手的迷茫感。我建议你在跑通全连接网络之后一定亲手改成CNN试一试。对比一下两者的准确率和训练速度再想一想为什么CNN能用更少的参数取得更好的效果卷积层通过局部感受野和权值共享大大降低了参数数量同时能提取到局部特征池化层通过下采样降低了特征维度提升了平移不变性。这些概念光看书容易一头雾水亲手实践之后会变得非常清晰。我再推荐两个后面可以自己尝试的升级方向。一是用torch.utils.tensorboard把loss曲线、准确率曲线在tensorboard里可视化更加直观二是给模型加一个残差连接哪怕只是简单的残差块看看能不能在MNIST上继续提升准确率。这些都是同一个项目里能延伸出来的实验收益远比再开一个新项目要大。按照上面的步骤一步步走代码全部跑通之后你会发现自己对pytorch的整体工作方式已经有了一个踏实的认知底座。后面无论是转图像分类、目标检测还是自然语言处理你都有足够的底子去拆解那些复杂框架的源码了。
返回列表