ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

MNIST手写数字识别实战:PyTorch搭建三层全连接网络

MNIST手写数字识别实战:PyTorch搭建三层全连接网络 简介这是一份面向机器学习与图像识别初学者的完整实现资源用纯Python和NumPy从零搭建三层全连接神经网络不依赖TensorFlow、PyTorch等深度学习框架完成MNIST手写数字分类任务适合课程设计、实验复现以及希望透过底层代码理解前向传播、反向传播和梯度更新的读者。压缩包共16个文件约9.51MB以3个Python源码和8个txt数据文件为主体其余为PyCharm工程配置xml/iml与编译缓存pyctxt文件存放训练特征、测试样本、手写数字标签及网络各层的权重与偏置等参数可逐项核对网络中间结果。资源还附带了将MNIST图片批量转换为txt的预处理代码方便改造自定义输入格式便于后续调试与扩展。从模型搭建、数据加载到参数保存均有对应脚本整体结构清晰已有3670人学习下载适合希望脱离高级框架、亲手实现并验证神经网络细节的开发者。1. 先纠正一个拼写歧义minist 就是 MNIST1.1 这个数据集到底是什么如果你在搜索引擎里敲下“minist 图像分类”大概率会看到一行红色提示你是不是想找 MNIST这个问题我遇到过而且我敢说十个初学者里有八个第一次都会把 MNIST 敲成 minist。这个手写数字识别任务几乎是所有人接触全连接神经网络时的第一个完整项目。MNIST 数据集来自美国人口普查局的员工和美国高中学生手写的数字经过采集、尺寸归一化后整理成 28×28 的灰度图。训练集有 6 万张测试集有 1 万张一共 10 个类别对应数字 0 到 9。每张图片都是单通道灰度像素值范围在 0 到 255背景大面积是黑色有效信息集中在中心区域。正因为数据量小、任务简单它对新手极其友好下载只需要几十 MBCPU 上跑一个简单的全连接网络几个 epoch 也就几分钟。在图像分类领域里MNIST 通常被当成“Hello World”。任何一个全连接神经网络、卷积神经网络甚至最新的图像分类模型在正式落到复杂场景之前都会先拿这个数据集验证网络是否搭得对、训练流程是否通、超参数设置是否合理。别看它简单跑通了这一步后面换到 CIFAR、ImageNet 或者其他森林图像分类项目时核心训练链路几乎不变变的只是网络结构、数据加载方式和更多的调参细节。这个项目解决的核心问题就是让模型看一张 28×28 的手写数字图片输出它是 0 到 9 中哪一个数字的预测结果。1.2 “三层全连接网络”里的三层到底算哪三层这是我在复现时遇到的一个容易让人晕的问题。“三层全连接网络”这个说法在不同的教材里指代并不完全一致。有的老师把“输入层-隐藏层-输出层”称为三层网络也就是只有一个隐藏层有的工程实现里把可学习的全连接层数量作为层数像我用 fc1、fc2、fc3 三个全连接层也叫三层全连接网络。这种命名混乱在深度学习项目里很常见我觉得最重要不是纠结叫法而是搞清楚代码里到底堆了几层 Linear。我这次采用的结构是 784-128-64-10。784 是输入的 28×28 像素展平后的维度128 和 64 是两个隐藏层的神经元数量10 是输出类别数。严格按前面的说法就是两个隐藏层加一个输出层一共三个全连接层。这个结构在 MNIST 上已经能跑到 97%~98% 左右。为什么选择这个结构以及每个数字背后的理由后面会展开聊。另外有一点值得注意MNIST 里的手写数字虽然只有 10 类但同一个数字在不同人笔下差异非常大比如 0 可能写得很扁、7 可能带横杠所以模型必须学到一定程度的抽象能力这也是为什么不能只用一层线性层的原因。2. 环境准备与数据加载复现时最容易翻车的地方2.1 依赖安装与项目结构我用的框架是 PyTorch搭配 torchvision 来下载和管理数据集。安装命令很简单pip install torch torchvision numpy matplotlib没有 GPU 完全不影响直接装 CPU 版本就行。我实测下来这个项目在普通笔记本 CPU 上跑 5 个 epoch 大约需要 2 到 3 分钟时间主要花在数据读取和矩阵运算上。项目本身不需要复杂的工程结构一个 Python 脚本就能跑通目录里会自动生成一个data/文件夹存放 MNIST 原始文件。如果你遇到下载特别慢或者屡次失败的情况可以手动去数据集官网下载四个 gz 文件放到data/MNIST/raw/对应目录下这样 torchvision 会识别到已有文件不再重复下载。2.2 数据预处理ToTensor 和 Normalize 缺一不可数据加载部分最容易忽略的是 transform。我使用的代码是这样from torchvision import transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])ToTensor()会把 PIL 图片从 0~255 的整数像素值转为 0~1 的浮点张量同时把维度从(28, 28)变成(1, 28, 28)新增的那个维度是通道数灰度图只有一个通道。Normalize((0.1307,), (0.3081,))使用 MNIST 全量训练集的均值和标准差把数据分布拉回到 0 附近、标准差调整到 1 附近。新手最容易漏掉后面这一步结果也能训练但收敛速度和最终准确率会差一些。这个问题我实测过同样是 5 个 epoch不加归一化的测试准确率大概停在 94%加了之后能到 97% 以上。原因是全连接网络对输入特征的尺度很敏感输入数据分布偏移会导致梯度更新不稳定网络需要额外花几个 epoch 去“适应”输入尺度。所以图像分类项目的预处理环节优先级非常高它的影响往往比换一个更复杂的网络结构还大。加载部分的代码from torch.utils.data import DataLoader from torchvision import datasets 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) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse)test_loader不设shuffleTrue是因为测试集不需要打乱顺序保持原始顺序方便后续输出混淆矩阵和可视化错误样本。另外如果你在 Windows 上运行DataLoader里的num_workers最好保持默认 0或者把训练代码写进if __name__ __main__:里否则多进程加载数据时容易报错。2.3 训练前先看形状避免第一个报错出现在最不该出现的地方新手最常见的报错之一是mat1 and mat2 shapes cannot be multiplied出现这个报错的原因很直接DataLoader返回的一个 batch 形状是[64, 1, 28, 28]而全连接层期望输入是[batch, 784]。你必须把图片向量展平也就是把每个样本从 28×28 的二维结构变成 784 个像素的一维向量。我建议在写网络之前先做一步检查images, labels next(iter(train_loader)) print(images.shape, labels.shape) # 期望输出: torch.Size([64, 1, 28, 28]) torch.Size([64])然后随机挑几张图用matplotlib画出来确认图片和标签是对应关系比如标签是 7图里确实是个手写的 7。这一步虽然简单但对排查后续问题非常有用尤其是在你刚开始接触这个领域、对张量维度还没有形成直觉的时候。3. 网络结构逐层拆解为什么是 784-128-64-103.1 全连接层到底做了什么计算全连接层做的事可以概括成一句话把输入向量的每个元素与当前层每个神经元进行加权求和再加上偏置公式就是 y xW^T b。第一层输入是 784 维输出是 128 维意味着需要学习一个形状为(784, 128)的权重矩阵再加上 128 个偏置fc2 是(128, 64)加偏置fc3 是(64, 10)加偏置。整个网络的可训练参数量大约是 784×128128 100480加上 128×6464 8256再加上 64×1010 650合计约 109386 个参数。这个规模对现代算力来说微不足道所以 CPU 也能轻松跑。参数量的意义在于它决定了网络的表达能力。参数太少模型容易欠拟合连训练集都学不好参数太多又会容易过拟合训练集接得住但测试集表现下滑。109K 参数处理 6 万张训练图片算是一个比较适中偏小的容量选择这也是 MNIST 用全连接网络好调的原因之一。如果你把两个隐藏层都改成 256参数量会直接翻倍到 230K 以上但测试准确率未必有明显提升因为数据的复杂度没有高到需要那么多参数。3.2 隐藏层神经元数量怎么选隐藏层神经元数量没有绝对标准更多是经验和实验的结合。128 作为第一层隐藏层宽度是 MNIST 项目里很常见的起点第二层再降到 64形成一种“先扩展再压缩”的信息处理方式。784 维输入本身像素冗余很高全连接层可以把相邻像素的冗余信息合并逐步抽象出更高级的特征。你也可以试 256-128 或 512-128参数量会变大训练时间会变长但测试准确率不一定会跟着涨。我建议在入门阶段把隐藏层控制在 128 和 64先把训练链路跑通再回去调宽度这样排查问题的时候变量少很多。代码实现部分其实非常简洁import torch import torch.nn as nn class ThreeLayerNet(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(784, 128) self.fc2 nn.Linear(128, 64) self.fc3 nn.Linear(64, 10) def forward(self, x): x x.view(x.size(0), -1) x torch.relu(self.fc1(x)) x torch.relu(self.fc2(x)) x self.fc3(x) return x注意forward里第一行做了展平用x.size(0)取 batch 维度-1的意思是把剩下的 1×28×28 自动展平成 784。这个写法比直接写死x.view(-1, 784)更通用换到其他尺寸的输入图片时不用改代码。3.3 激活函数和初始化ReLU 与默认初始化没有激活函数时多个线性层叠加仍然等价于一个线性层无论堆多少层表达能力都非常有限。ReLU 的引入让网络变成非线性这样才能拟合更复杂的决策边界。ReLU 在负数部分直接置零正数部分保留计算简单而且能缓解梯度消失问题。这里要特别提醒输出层不要加 ReLU。因为之后会用交叉熵损失PyTorch 的CrossEntropyLoss内部自带 softmax 操作需要的是每个类别的原始 logits而 ReLU 会把负值裁掉破坏 logits 的分布导致训练效果明显变差。这个坑不少人都踩过表现就是 loss 下降得很慢准确率始终上不去。权重的初始化 PyTorch 的nn.Linear默认使用 Kaiming 均匀初始化配合 ReLU 激活函数在多数情况下已经够用不需要手动干预。如果哪天你用到了 Sigmoid 作隐藏层激活就得考虑换成 Xavier 初始化否则深层网络很容易梯度消失。这个三层小网络对初始化并不敏感但理解这个原理有助于你迁移到更复杂的图像分类模型。4. 训练过程的核心参数损失函数、优化器与学习率4.1 交叉熵损失和 Adam 优化器为什么是默认组合训练分类网络本质上是在最小化损失函数。CrossEntropyLoss做的事情是把模型的 10 个 logits 经过 softmax 转成概率分布再计算预测分布与真实标签的交叉熵。模型预测越接近正确答案损失越低。多分类问题几乎都用这个损失函数原因在于它训练出来的概率分布有明确含义模型不仅告诉你是哪一个数字还告诉你它对每个数字的置信程度。优化器我一开始选了 Adam。Adam 的优点是自适应学习率对初始学习率的敏感度比 SGD 低很多对深度学习新手来说比较省心。如果你想要更好的泛化性能可以试试 SGD 加上 momentum但调参会更麻烦一点。两个优化器在 MNIST 上都能达到 97% 以上不用太纠结。真正需要注意的是学习率这个下面会专门讲。4.2 一次完整的训练循环训练循环的代码是import torch.optim as optim model ThreeLayerNet() criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) model.train() for epoch in range(5): running_loss 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() print(fEpoch {epoch1}, loss: {running_loss / len(train_loader):.4f})这段代码里最容易被忽略的是optimizer.zero_grad()。PyTorch 的梯度默认是累加的如果不手动清零下一轮 batch 的梯度会叠加到上一轮上导致 loss 乱跳甚至不收敛。我见过很多“loss 怎么都不降”的问题最后发现只是少了这一行。如果你希望结果可复现在代码最前面加一句torch.manual_seed(42)固定随机种子。评估函数我单独写了一个def evaluate(model, loader): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: outputs model(images) _, predicted torch.max(outputs, dim1) total labels.size(0) correct (predicted labels).sum().item() return correct / total print(fTest accuracy: {evaluate(model, test_loader) * 100:.2f}%)model.eval()和with torch.no_grad()的作用分别是关闭训练模式下的随机失活行为、关闭梯度计算。虽然这个三层小网络里没有 dropout 和 BatchNorm但养成这种习惯很重要否则以后切到复杂模型时很容易因为漏了这两行得到完全不可信的验证结果。4.3 学习率、batch size 和 epoch 的实测对比我用相同的网络结构跑了几组常见超参数结果大致如下。因为随机种子不同会有波动所以这里的数字只作为趋势参考学习率batch sizeepoch测试准确率约备注0.0164597.2%loss 后期抖动明显0.00164597.8%比较平滑推荐0.0001641097.9%收敛慢需要更多轮次0.00132597.9%单轮稍慢梯度噪声大0.001128597.5%大 batch 收敛稳定但泛化略降从趋势能看出来学习率 0.01 太大loss 后期会有明显抖动说明模型在最优参数附近来回震荡学习率 0.0001 太小5 个 epoch 不够充分需要把训练轮数翻倍才能追上来。batch size 32 和 64 差距不大128 时收敛稳一点但测试准确率略降这是大 batch 容易收敛到平坦度较低极值点的常见现象。对 MNIST 和这个小网络来说lr0.001、batch size64、epoch5 是我比较推荐的起点。5. 结果分析与避坑记录从失败到 98% 的调试路径5.1 我的第一次失败准确率卡在 92% 附近的排查过程我先说一个印象很深的教训。第一次实现时我把数据 transform 写成了只做ToTensor()没有做 Normalize结果训练 5 个 epoch 测试准确率卡在 92%怎么调学习率都上不去。我一度以为是网络层数不够把隐藏层改成 256 和 128结果不升反降训练速度还慢了不少。后来排查了很久才发现问题出在特征归一化上。原因前面提过0~1 的像素分布虽然比 0~255 好很多但和网络初始化时假定的零均值分布仍有偏移导致梯度更新方向在小数据上不够稳定。加上 MNIST 的均值和标准差后同样设置直接冲到 97.8%。这件事让我意识到图像分类项目里如果准确率一直上不去第一件事应该检查数据预处理而不是急着改网络结构。排查顺序我建议这样先看数据形状对不对再看预处理是否完整然后看训练 loss 是否正常下降最后才去动网络结构和超参数。顺序反了经常会为了一个错误原因浪费大量时间。5.2 如果训练准确率高但测试低过拟合怎么处理全连接网络有 109K 参数对 MNIST 这种相对简单的数据集来说容量已经偏高了所以训练时间一长就会出现轻微过拟合。一个典型信号是训练集准确率到了 99.8%测试集只有 97.2%。解决方案是先不要在前期加太多正则化把基模型跑通后再在 fc1 和 fc2 后添加 dropoutp 设为 0.2 左右测试准确率往往会有 0.2 到 0.5 个百分点的提升。另一种思路是增大训练数据的多样性比如对图片做随机平移和旋转但 MNIST 已经相对标准化数据增强带来的提升没有大图像数据集那么明显我建议把重点放在网络结构和正则化上。要注意 dropout 和model.train()、model.eval()的配合训练时 dropout 随机丢弃神经元评估时必须切回 eval 模式让所有神经元都参与推理否则每次预测结果都会因为随机性而不同。5.3 不要把 98% 当成终点看混淆矩阵和错误样本准确率只是一个汇总数字真正让模型继续变好的是看它错在哪里。我写了一个简单的错误收集逻辑wrong [] for images, labels in test_loader: outputs model(images) preds torch.argmax(outputs, dim1) mask preds ! labels for img, true, pred in zip(images[mask], labels[mask], preds[mask]): wrong.append((img, true.item(), pred.item()))把wrong里的样本按真实标签和预测标签分组统计我发现最容易混淆的数字对是 4 和 9、3 和 8、7 和 2。这些数字在人类手写时也长得很像模型把一部分置信度分配到了错误的类别上属于正常现象。看错误样本的另一个价值是定位数据问题比如如果某些图片本身被错误标注模型再强也不可能分对这时候要考虑清理训练数据。你可以用matplotlib打印一个 10×10 的错误图片网格标上真实标签和预测标签很快就能找到规律。这个习惯帮我省了很多时间建议你在任何图像分类项目里都保留这套调试链路。5.4 全连接网络的边界与后续扩展空间全连接网络在 MNIST 上能做到 98% 左右但再往上就非常吃力了因为输入像素被展平后空间结构完全没有利用相邻像素之间的关系也丢失了。这也是为什么最新的图像分类模型大多基于卷积神经网络或者 Transformer 架构。不过这不代表全连接网络没有价值相反先跑通这个项目你会清楚地理解数据流、梯度回传、损失函数这些所有模型共有的核心机制。之后再切换到图像分类算法里更复杂的结构至少不会因为基础概念不熟而寸步难行。项目实操结束时我自己最大的收获不是那 98% 的准确率而是建立了一套 debug 习惯先看数据形状、再确认预处理、然后调网络结构最后才动手调参。我给每个项目建了固定的实验记录模板把数据预处理、网络结构、超参数和测试结果四样东西写在一起下次调参直接看历史记录就能定位问题。这个习惯帮我省下了大量重复实验的时间。如果你也在复现这个项目我希望这篇能帮你少走一点弯路尽快跑到结果。本文还有配套的精品资源点击获取
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表