ARTICLE DETAIL

资讯详情

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

PyTorch入门实战:MNIST手写数字识别全流程详解

PyTorch入门实战:MNIST手写数字识别全流程详解 直接说结论这篇文章写给两类人。第一类是刚入门深度学习、被各种术语劝退的新手想找一条从零开始能跑通的路第二类是已经跑过一些例子但想系统搞清楚数据加载、模型构建、训练流程这些环节为什么这么写的人。我不打算给你堆一堆学术名词而是用一套完整的、能从零跑到结果的项目代码把整个链路拆开讲透。mnist这个数据集在计算机视觉领域的地位差不多等同于编程界的Hello World你把它彻底吃透后续再去看卷积神经网络、图像分类、目标检测这些方向底子就扎实了。先说清楚这篇文章能解决什么问题从torchvision下载mnist数据集顺带解决下载404的问题到用pytorch搭建神经网络模型再到训练、评估、可视化全部走一遍。看完之后你不仅有一个能跑的分类器更关键的是理解每一行代码背后的设计逻辑。整个过程不需要GPU纯CPU也能在几分钟内拿到97%以上的准确率这对入门来说非常友好。1. 内容整体设计与思路拆解1.1 为什么选择mnist作为入门项目mnist全称是Modified National Institute of Standards and Technology database由Yann LeCun等人整理发布。它包含0到9共10个类别的手写数字灰度图像训练集有60000张测试集有10000张每张图片的分辨率是28x28像素。这个数据集最大的优点是干净、小、标准统一你不需要像处理真实业务数据那样花大量时间做清洗、标注、格式转换可以聚焦在模型本身的学习上。同时它又足够有代表性涵盖了完整的分类识别流程数据预处理、模型构建、损失函数设计、梯度优化、效果评估。这些流程是所有深度学习任务共通的骨架所以我在文章里刻意没有用更花哨的卷积网络而是先用最简单的全连接神经网络把流程打通。等你理解了骨架后面再换上更复杂的模型结构只是替换中间的一个模块而已整个框架不需要大改。1.2 技术方案选型全连接神经网络为什么够用很多新手上来就直奔卷积神经网络CNN这其实是个误区。CNN相对于传统全连接网络的核心优势在于能提取空间局部特征但mnist图片尺寸只有28x28还是灰度图像素点之间虽然也有空间关系但数字的辨识主要依赖笔画结构全连接网络通过足够多的参数也能学会这些特征。实测下来一个三层的全连接网络隐藏层分别设512和256个神经元在mnist上就能轻松达到97%左右的准确率。这个数据足以说明问题入门阶段没必要把模型复杂度拉满。用全连接网络的另一个好处是代码更直观。pytorch的nn.Linear就是做矩阵乘法加偏置你只需要理解输入维度、输出维度、激活函数这几个概念就能把网络结构搭出来。等这个流程跑通了我再建议你去试试CNN用对比的心态去看两种网络在mnist上的表现差异和训练速度差异那种理解深度是直接抄一个CNN代码无法比的。1.3 项目文件结构与整体流程我在实际做这个项目时强烈建议你把代码拆分成清晰的模块而不是所有东西都堆在一个脚本里。但考虑到入门读者的需求这篇文章的主体代码我会保持在一个notebook或者一个python文件内能跑通同时会在关键位置用注释标注清楚每个区块的职责。整个流程可以概括为准备环境 - 加载数据 - 数据预处理 - 定义模型 - 定义损失函数和优化器 - 训练循环 - 测试评估 - 可视化结果这个顺序就是pytorch项目的标准工作流。你以后接任何深度学习任务哪怕是目标检测、文本分类骨架都是这样的变的只是数据格式、模型结构、损失函数这几个环节。2. 环境准备与依赖安装2.1 版本选择与安装命令先说环境。pytorch的安装是很多新手第一道坎尤其是涉及到GPU版本的时候。我的建议是第一步先去pytorch官网pytorch.org的Get Started页面选择你的操作系统、包管理器pip还是conda、CUDA版本官网会给出对应的安装命令。如果你的电脑没有NVIDIA显卡或者在macOS上就选CPU版本先用CPU跑通流程完全没问题。这里给出一套参考命令。如果你用conda管理环境推荐创建一个独立的环境避免依赖冲突conda create -n mnist_demo python3.10 conda activate mnist_demo pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu如果电脑有NVIDIA GPU并且已经装好了CUDA那么把最后的参数换成对应的CUDA版本比如CUDA 12.1的安装命令是pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121安装完成后在python环境里验证一下是否成功import torch import torchvision print(torch.__version__) print(torchvision.__version__) print(torch.cuda.is_available())如果最后一行输出True说明pytorch能识别到你的GPU输出False也不影响跑mnist只是会慢一点。注意不要直接pip install torch不指定镜像源这样很可能装到CPU版本虽然也能用但如果你有GPU就浪费了算力。而且不同版本的pytorch和CUDA之间有兼容矩阵官网的配置页已经替你算好了直接用官方给的命令最稳妥。2.2 torchvision下载mnist时遇到404错误的解决办法这里单独拿出一节因为这正是我最初跑这个项目时踩过的一个典型的坑。torchvision自带的datasets.MNIST接口会自动下载数据集但很多人在这一步遇到404错误或者下载速度极其缓慢。原因在于数据集托管在Yann LeCun的个人网站yann.lecun.com/exdb/mnist/有时会因为网络环境或者网站服务器问题访问不了。而torchvision下载默认用的是这个源一旦下载失败就会报类似404 Not Found的错误。解决办法有几种方法一手动下载数据集文件然后放到指定的目录。需要下载以下四个文件train-images-idx3-ubyte.gz训练图像train-labels-idx1-ubyte.gz训练标签t10k-images-idx3-ubyte.gz测试图像t10k-labels-idx1-ubyte.gz测试标签下载后解压放到项目的./data/MNIST/raw/目录下代码运行时会自动检测到文件已经存在跳过下载步骤。方法二使用国内的镜像源比如某些云厂商维护的数据集镜像站。我尝试过清华源等方式但这种方式有个风险就是镜像站的地址会变化不保证长期有效。# 如果已经手动下载好代码里这样写就能直接使用 train_dataset datasets.MNIST( root./data, trainTrue, transformtransforms.ToTensor(), downloadFalse # 改为False使用本地文件 )方法三如果是公司内网有代理环境的可以尝试配置代理但这种方案涉及具体网络环境普适性不强。所以综合来看推荐方法一一次手动下载终身复用。3. 数据集的加载与预处理3.1 深入理解mnist数据集结构先把mnist的数据格式搞清楚后续代码才不会懵。每个样本是一个28x28的灰度图像像素值范围是0到2550表示黑色背景255表示白色笔画。标签是一个0到9的整数表示图片中手写的数字。torchvision在加载时返回的dataset对象每个元素是一个元组(image, label)其中image是PIL Image对象label是int类型。在交给模型之前需要做转换。我们通常用transforms.Compose把多个预处理步骤组合起来transform transforms.Compose([ transforms.ToTensor(), # PIL Image转Tensor像素值缩放到[0,1] transforms.Normalize((0.1307,), (0.3081,)) # mnist官方推荐的均值和标准差 ])这里的关键点是ToTensor操作。它会将原本维度为(H, W)的PIL图像转换为维度为(C, H, W)的Tensor即(1, 28, 28)并且像素值从0-255映射到0-1之间。这一步不做的话模型输入数据范围差异太大不利于梯度下降收敛。Normalize操作是围绕均值和标准差做的标准化tensor_normalized (tensor - mean) / std。在多个数据集上这个均值和标准差需要自己计算但mnist太经典了全网统一用0.1307和0.3081这两个值就行。这组数值的意思是mnist全部训练集图像的像素均值约为0.1307标准差约为0.3081。3.2 DataLoader的机制与参数选择Dataset负责管理数据样本而DataLoader负责在训练时按批次取出数据并且支持打乱、并行加载等操作。这里面的几个参数值得你仔细体会batch_size 64 train_loader DataLoader( train_dataset, batch_sizebatch_size, shuffleTrue, num_workers2 ) test_loader DataLoader( test_dataset, batch_size1000, shuffleFalse, num_workers2 )为什么训练集要shuffleTrue而测试集shuffleFalse因为训练时我们希望通过随机打乱让每个mini-batch的数据分布尽可能接近整体分布避免模型只看到某个类别的连续数据导致梯度更新方向有偏。测试时不需要打乱因为我们是在评估模型已学到的能力数据顺序不影响结果同时保留顺序便于排查问题。batch_size的选择是一个权衡。64是我在这个项目上的推荐值。批大小太小比如1或者4梯度更新太频繁训练不稳定且慢批大小太大比如1024虽然单步计算效率高但容易收敛到平坦的极小值泛化能力反而可能下降。在mnist任务上64到256都在合理区间想减少训练时间可以用128。num_workers表示用几个子进程加载数据。在Windows上如果设置为大于0的值有时会报错那就设置为0在Linux或者macOS上设置为CPU核心数的一半左右比较合适。这里设置为2对新手来说更省心。3.3 数据可视化与样本检查拿到数据后先别急着训模型习惯上我都会先看看数据长什么样这是排查问题的第一步。如果你加载出来的图片是反色的、模糊的、或者标签对不上这时候发现成本最低。画图的代码很简单import matplotlib.pyplot as plt # 取一个batch的数据 images, labels next(iter(train_loader)) # images.shape: torch.Size([64, 1, 28, 28]) # 画一个4x4的网格 fig, axes plt.subplots(4, 4, figsize(8, 8)) for i in range(16): ax axes[i // 4][i % 4] ax.imshow(images[i].squeeze(), cmapgray) ax.set_title(fLabel: {labels[i].item()}) ax.axis(off) plt.tight_layout() plt.show()这里有个细节images[i].squeeze()把(1, 28, 28)压缩成(28, 28)matplotlib才能正常显示灰度图。cmapgray指定灰度色彩映射不加的话默认是viridis彩色映射看起来会误导你对数据的判断。我见过不少新手因为忘了这两步盯着花里胡哨的彩图一头雾水其实数据本身没任何问题。4. 神经网络模型的构建4.1 从零手写一个全连接网络把模型这块吃透是整个项目的核心。我在这里用最接近数学定义的方式实现网络结构然后再给你看pytorch更简洁的写法。先看手写版本import torch.nn as nn import torch.nn.functional as F class NeuralNet(nn.Module): def __init__(self, input_size784, hidden1_size512, hidden2_size256, num_classes10): super(NeuralNet, self).__init__() self.fc1 nn.Linear(input_size, hidden1_size) self.fc2 nn.Linear(hidden1_size, hidden2_size) self.fc3 nn.Linear(hidden2_size, num_classes) def forward(self, x): x x.view(-1, 784) # 形状从 [batch, 1, 28, 28] 变为 [batch, 784] x F.relu(self.fc1(x)) x F.relu(self.fc2(x)) x self.fc3(x) # 输出层不加激活直接给Logits return x逐行解释。nn.Linear(input_size, hidden1_size)做的事情是y xW^T b其中W是权重矩阵形状为(hidden1_size, input_size)b是偏置向量。输入784维输出512维这一层总共的参数数量是784 * 512 512 401408个。整个网络的参数量你可以自己算一下大概在67万左右这个规模的网络对mnist来说已经相当充裕了。forward函数定义了数据的前向传播路径。注意x.view(-1, 784)这一步-1表示自动推断这一维的大小因为我们传入的batch大小是64那么这里就会被自动推断为64结果就是[64, 784]的张量。这一步是把二维图片展平成一维向量。很多新手在这里容易报维度错误核心原因就是没有理解pytorch中张量维度的变化。F.relu是激活函数。这里补充一下激活函数的作用如果每一层都只做线性变换那么不管叠加多少层本质还是一个线性模型根本无法学习非线性的决策边界。ReLU的公式是max(0, x)计算简单梯度不会像sigmoid那样在两端趋近于0导致梯度消失是目前全连接网络和卷积网络里默认使用的激活函数。输出层没有加激活函数这里非常关键。因为后面我们用的损失函数nn.CrossEntropyLoss()内部已经包含了softmax操作。如果我们在这里提前加softmax会导致softmax被计算两次虽然数值上不一定错得很离谱但会降低数值稳定性也会影响梯度计算。这是新手特别容易犯的错误。4.2 用nn.Sequential快速搭建如果你理解了上面每一层的含义就会发现在实际项目中我们可以用更简洁的方式表达同样的结构class NeuralNetV2(nn.Module): def __init__(self): super(NeuralNetV2, self).__init__() self.net nn.Sequential( nn.Linear(784, 512), nn.ReLU(), nn.Linear(512, 256), nn.ReLU(), nn.Linear(256, 10) ) def forward(self, x): x x.view(-1, 784) return self.net(x)nn.Sequential把多个层按顺序串联起来前一个的输出自动作为后一个的输入。这种方式代码更短但调试时不容易在中间层插入打印语句。我的建议是项目初期用第一种写法逻辑更透明确认没问题后可以用第二种写法代码更简洁。4.3 损失函数与优化器的选择逻辑损失函数衡量的是模型预测和真实标签之间的差距优化器决定如何根据这个差距更新模型的参数。这里选型背后有明确的逻辑。分类任务用交叉熵损失函数criterion nn.CrossEntropyLoss()交叉熵为什么适合分类它背后的信息论含义是度量两个概率分布之间的距离。我们对每个样本的输出是10个类别的logits未归一化的分数CrossEntropyLoss内部先做softmax把logits转换为概率分布然后计算真实标签对应的负对数概率。如果模型对正确类别的置信度接近1损失接近0如果置信度低损失就大。相比均方误差MSE交叉熵在分类任务上收敛更快而且梯度形式更利于反向传播。优化器选择Adamoptimizer torch.optim.Adam(model.parameters(), lr0.001)Adam结合了Momentum和RMSProp的优点能够自适应地为每个参数调整学习率。它的收敛速度在大多数任务上都要比原生的随机梯度下降SGD快对新手来说容错也更高不太需要精细调学习率。我实测mnist上用Adam学习率0.001基本不需要额外的学习率调整策略就能收敛得很好。这里补充一下学习率的概念。学习率决定了每一步参数更新的幅度。如果学习率是0.1参数会大幅跳跃容易震荡不收敛如果是0.00001参数更新太慢训练要很久。Adam lr0.001是经过无数实践验证的安全组合先用它跑通再考虑其他调参技巧。5. 训练循环与模型评估5.1 训练一个epoch的完整代码训练循环是pytorch项目中最机械但也最重要的部分。我先把代码放出来然后逐段解释def train_one_epoch(model, train_loader, criterion, optimizer, epoch): model.train() running_loss 0.0 correct 0 total 0 for batch_idx, (images, labels) in enumerate(train_loader): # 数据维度检查 # images: [batch_size, 1, 28, 28] # labels: [batch_size] # 前向传播 outputs model(images) loss criterion(outputs, labels) # 反向传播 optimizer.zero_grad() loss.backward() optimizer.step() # 统计 running_loss loss.item() _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() if (batch_idx 1) % 200 0: print(fEpoch [{epoch1}], Batch [{batch_idx1}], Loss: {loss.item():.4f}) epoch_loss running_loss / len(train_loader) epoch_acc 100.0 * correct / total print(fEpoch [{epoch1}] Training Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%) return epoch_loss, epoch_acc分几个关键点说。model.train()这行代码常被忽略但它很重要。它会将模型切换到训练模式影响BatchNorm和Dropout层的行为。对于我们现在这个全连接网络没有BN和Dropout效果等同于什么都不做但养成习惯总归没错因为换到复杂模型时它就不一样了。optimizer.zero_grad()在每次前向传播前把梯度清零。这个操作位置有两个可选项一是在每个batch训练之前整体清零二是在损失计算之后、backward之前清零。效果基本相同但推荐放在forward之前逻辑更清晰。新手最容易犯的错误是忘记清零梯度导致梯度累加模型参数更新方向混乱loss出现诡异波动。你可能听说过pytorch和tensorflow有个核心区别是tensorflow默认自动更新参数而pytorch默认不会自动清空梯度所以这个细节在pytorch里尤其重要。loss.backward()就是反向传播也就是根据损失函数计算每个参数对应的梯度值。optimizer.step()根据计算好的梯度和优化器内部状态更新模型参数。损失下降的核心原理是这样的前向传播计算损失反向传播计算梯度优化器沿梯度的反方向更新参数。这三个操作构成一个完整的训练step。每走一步参数就向着让loss更小的方向靠近一点。5.2 测试评估的细节训练集准确率再高也不能代表模型泛化能力好我们真正关心的是模型在没见过的数据测试集上的表现。测试代码和训练代码很像但有几个关键差异def evaluate(model, test_loader, criterion): model.eval() test_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: outputs model(images) loss criterion(outputs, labels) test_loss loss.item() _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() test_loss / len(test_loader) test_acc 100.0 * correct / total print(fTest Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.2f}%) return test_loss, test_acc带decoder注意model.eval()和torch.no_grad()。model.eval()切换模型到评估模式影响BN和Dropout行为。测试时BatchNorm会使用累计的running mean而不是当前batch的统计量Dropout会关闭随机失活用全部神经元。torch.no_grad()关闭梯度计算图。测试阶段我们不需要计算梯度因为不需要更新参数。关闭梯度记录可以显著减少内存消耗和计算量。没有这个上下文管理器测试会变慢而且在某些复杂模型上可能因为梯度图累积导致内存溢出。这里有个细节torch.max(outputs.data, 1)返回两个值第一个是最大值第二个是最大值的索引下标。predicted保存的就是这个下标它与labels对比就能统计正确个数。5.3 主循环训练多个epoch有了训练函数和测试函数主循环就非常简洁了epochs 5 train_losses [] test_losses [] train_accs [] test_accs [] for epoch in range(epochs): train_loss, train_acc train_one_epoch(model, train_loader, criterion, optimizer, epoch) test_loss, test_acc evaluate(model, test_loader, criterion) train_losses.append(train_loss) test_losses.append(test_loss) train_accs.append(train_acc) test_accs.append(test_acc)这个任务5个epoch已经完全足够了。我实测的结果第一个epoch结束训练集准确率就能到90%左右5个epoch之后测试集准确率稳定在97%到98%。如果你把网络再加宽一些或者做一次数据增强98.5%以上也是可以达到的。5.4 训练过程中的可视化与指标分析光看一个最终准确率其实你学不到太多东西。把训练过程中的loss和accuracy画出来你能直观看到模型的收敛过程以及是否出现异常plt.figure(figsize(12, 4)) plt.subplot(1, 2, 1) plt.plot(range(1, epochs1), train_losses, labelTrain Loss) plt.plot(range(1, epochs1), test_losses, labelTest Loss) plt.xlabel(Epoch) plt.ylabel(Loss) plt.legend() plt.subplot(1, 2, 2) plt.plot(range(1, epochs1), train_accs, labelTrain Accuracy) plt.plot(range(1, epochs1), test_accs, labelTest Accuracy) plt.xlabel(Epoch) plt.ylabel(Accuracy (%)) plt.legend() plt.tight_layout() plt.show()怎么看这张图如果train loss持续下降而test loss上升到某个点后回升说明模型开始过拟合此时应该减少epoch或者增加正则化如果train和test的loss都一直下降说明还有继续训练的空间如果train loss就迟迟降不下来问题大概率出在学习率设置、数据预处理或者模型结构上。养成查看训练曲线的习惯之后你会发现调参不再像玄学而是有迹可循的过程。6. 模型预测与结果分析6.1 识别新样本的完整流程训练好的模型最终要用于推断。实际使用中我们不会每次都重新训练模型而是保存模型参数然后在需要时加载。pytorch保存和加载模型的推荐做法是# 保存模型参数推荐 torch.save(model.state_dict(), mnist_model.pth) # 加载模型参数 model NeuralNet() model.load_state_dict(torch.load(mnist_model.pth)) model.eval()这里强调一下model.state_dict()保存的只是参数不含网络结构。加载时必须先构建一个相同结构的模型对象再载入参数。而torch.save(model, ...)虽然能直接保存整个模型但存在兼容性和安全性的潜在问题不推荐在正式项目中使用。对单张图片做预测的完整代码如下from PIL import Image import numpy as np # 读取图片并转换为28x28灰度 image Image.open(digit.png).convert(L) image image.resize((28, 28)) # 转换为Tensor并做同样的预处理 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) input_tensor transform(image) # 形状: [1, 28, 28] input_tensor input_tensor.unsqueeze(0) # 增加batch维度: [1, 1, 28, 28] # 预测 with torch.no_grad(): output model(input_tensor) prediction torch.argmax(output, dim1).item() print(f预测结果: {prediction})关键点在于unsqueeze(0)这一步。模型训练时接受的输入是[batch_size, 1, 28, 28]单张图片自然没有batch维度所以手动加一个让数据形状和训练时保持一致。这是做推理时最容易出错的地方。6.2 用混淆矩阵深入分析模型表现准确率97%听起来不错但不够细。我们得知道模型在哪些数字上容易犯错这就是混淆矩阵的价值。混淆矩阵是一个10x10的矩阵第i行第j列表示真实类别是i、但被预测成类别j的样本数量。from sklearn.metrics import confusion_matrix import seaborn as sns model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in test_loader: outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) plt.figure(figsize(10, 8)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues) plt.xlabel(Predicted) plt.ylabel(True) plt.title(Confusion Matrix) plt.show()我跑出来的结果显示最容易混淆的是4和9、7和2、3和8这几对。原因也不难理解这些数字在书写体中的确有相似之处特别是4的上半部分如果写得比较开很容易被识别成9。这里可以延伸出一个思路通过混淆矩阵定位模型的弱点再有针对性地补充训练数据或者设计更合适的网络结构。真实项目中这类分析往往是提升模型效果的突破口。6.3 模型错误样本的直观展示比混淆矩阵更直观的是把预测错误的样本直接画出来。我自己每次训练完都会做这个步骤它往往能带来一些意外的洞察model.eval() misclassified [] mis_preds [] mis_labels [] with torch.no_grad(): for images, labels in test_loader: outputs model(images) _, preds torch.max(outputs, 1) incorrect_indices (preds ! labels).nonzero(as_tupleTrue)[0] for idx in incorrect_indices: misclassified.append(images[idx]) mis_preds.append(preds[idx].item()) mis_labels.append(labels[idx].item()) if len(misclassified) 20: break fig, axes plt.subplots(4, 5, figsize(12, 10)) for i, (img, pred, true) in enumerate(zip(misclassified[:20], mis_preds, mis_labels)): ax axes[i // 5][i % 5] ax.imshow(img.squeeze(), cmapgray) ax.set_title(fTrue: {true}, Pred: {pred}, colorred) ax.axis(off) plt.tight_layout() plt.show()你会惊讶地发现有些错误连人眼都很难分辨。比如一个人写的7顶部没有横杠看起来就是1或者一个极潦草的2笔画完全连在一起跟8的处理结果很像。这说明mnist虽然说是简单数据集但真实世界的手写变体还是保留了一定的难度。理解了这一点你就不会因为模型没到99%就觉得是自己代码写错了。7. 常见问题与排查技巧实录7.1 我在训练mnist时遇到的典型问题第一个问题是维度不匹配。RuntimeError: mat1 and mat2 shapes cannot be multiplied。这个报错出现在我第一次把数据和模型对接时。原因就是我忘了做view(-1, 784)的操作直接把[64, 1, 28, 28]的四维张量传给nn.Linear(784, 512)。报错信息看起来挺吓人其实核心就一句话Linear层期望的输入特征数是784但实际传进来的形状对不上。排查方式是打印x.shape逐层看数据形状变化。第二个问题是loss下降缓慢或不下降。我把学习率调到0.01结果发现训练loss在0.3附近来回震荡降不下去。这是学习率过大的典型表现梯度更新步长太大参数在最优解附近反复横跳。后来调回0.001loss曲线才顺利下降。这里想提醒你遇到loss问题不要怀疑模型代码有bug先检查学习率和数据预处理。第三个问题是运行速度很慢。大概率是num_workers设置过大或者没有合理利用GPU。我自己的经历是在一台4核CPU的旧笔记本上训练CPU版本跑一个epoch大概要30秒5个epoch就要2分半看着不急但也不快。如果你配置了GPU务必在初始化模型后加一句model model.to(device)同时把数据也移到GPU上device torch.device(cuda if torch.cuda.is_available() else cpu) model NeuralNet().to(device) # 训练时 images, labels images.to(device), labels.to(device)完整跑通CPU版本后再加设备迁移逻辑是更稳妥的路线。7.2 mnist分类问题的关键参数速查表我整理了一些在mnist上效果较好的参数范围给不同目标的读者做参考。参数推荐值说明batch_size64可以尝试32、128效果差异不大learning_rate0.001Adam优化器下的默认安全值epochs55轮即可达到97%加大到10轮收益很小隐藏层512-256这个配置性价比最高激活函数ReLU优先选择收敛快损失函数CrossEntropyLoss分类任务的标准选择优化器Adam新手首选自适应学习率如果你追求更高的精度有一个简单有效的思路数据增强。虽然mnist本身已经很规范但不妨试一下在训练时对图片做轻微旋转、平移# 训练集数据增强 train_transform transforms.Compose([ transforms.RandomAffine(degrees10, translate(0.1, 0.1)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])这个方法在mnist上能把准确率往上推0.3到0.5个百分点。它的原理是人为增加训练样本的多样性让模型对轻微形变更鲁棒减少过拟合。7.3 内存占用与训练效率的优化建议mnist本身数据量小即便不使用GPU内存也完全不是瓶颈。但为了让代码在更复杂的数据集上也能跑得顺利有几个习惯建议你从现在就养成。第一用torch.no_grad()包住测试流程避免梯度图累积。第二及时释放缓存在一个epoch训练结束后有时可以用torch.cuda.empty_cache()释放GPU缓存。第三尽量使用enumerate(train_loader)而不是range(len(train_loader))再通过索引取数据后者会增加代码复杂度和出错概率。还有一点如果使用DataLoader时报告了BrokenPipeError这通常出现在Windows系统上常见原因是在代码执行完毕后子进程还没完全退出。解决方案是在主代码外面套一层if __name__ __main__: main()这可以确保子进程在正确的时机退出。这个坑在Windows上非常常见但网上很多教程都没提我在实际使用中踩过几次后就特别注意这一点了。8. 后续进阶方向与扩展建议mnist项目做完之后你可以沿几个方向继续深入每个方向都会涉及新的技术点。第一个方向是把全连接网络升级为卷积神经网络。只需要改模型定义部分数据加载、训练循环这些代码都不用动。我建议你尝试用两层卷积加池化再加全连接层的结构对比CNN和全连接网络在mnist上的表现差异。CNN通常能到99%以上而且模型参数量不一定更大这就是特征提取能力的差异。你亲手改一遍代码体会会非常深。第二个方向是更换数据集。mnist搞明白了可以试一下FashionMNIST这个数据集同样是60k张28x28灰度图但内容是衣服、鞋子、包包等时尚品类比mnist更难一些图像特征更加多样全连接网络做它准确率会掉到90%以下这才有挑战性。你会发现代码基本不用改只改一行数据集类名就能跑这就是框架抽象带来的便利。第三个方向是深入研究训练细节。比如给模型添加Dropout层防止过拟合、尝试不同的优化器、实现学习率衰减策略。这些技巧在mnist上的收益可能不大但在真实数据集上往往就是90%和95%准确率的分水岭。我自己的感觉是mnist项目最大的价值不在于教会你做一个数字识别器而在于把一个深度学习项目的完整生命周期走了一遍。数据怎么加载、模型怎么设计、怎么训练、怎么评估、怎么诊断问题这套方法论是放之四海而皆准的。你今天花几个小时把这个项目跑通、吃透后面再遇到任何更复杂的问题无非是在这个骨架的某些环节上换更复杂的模块。基础打得牢后面的路就走得稳。
返回列表