从零理解 PyTorch:用神经网络完成 MNIST 手写数字识别

最近学习了一些深度学习和 PyTorch 框架的基础知识,并尝试使用 MNIST 手写数字数据集完成一个简单的图像分类任务。虽然手写数字识别是深度学习入门中非常经典的案例,但它包含了神经网络训练的完整流程:数据集加载、模型构建、前向传播、损失计算、反向传播、参数更新以及模型测试。

通过这个项目,我逐渐理解了神经网络并不是一个神秘的"黑盒子",它本质上是由许多矩阵运算、激活函数和参数更新组成的计算过程。只要把每个步骤拆开来看,整个模型就会变得清晰很多。

一、MNIST 手写数字数据集

MNIST 是一个非常经典的手写数字数据集,里面包含大量 0 到 9 的灰度图像。每张图片的大小都是 28×28 像素,图片中的每个像素都可以用一个数值表示。

在训练模型之前,首先需要下载并加载数据集。PyTorch 的 torchvision 提供了非常方便的数据集接口,可以直接使用 datasets.MNIST 加载训练集和测试集。

复制代码
from torchvision import datasets
from torchvision.transforms import ToTensor

training_data = datasets.MNIST(
    root="data",
    train=True,
    download=True,
    transform=ToTensor()
)

test_data = datasets.MNIST(
    root="data",
    train=False,
    download=True,
    transform=ToTensor()
)

这里的 root 表示数据保存的位置,train=True 表示加载训练集,train=False 表示加载测试集,download=True 表示如果本地没有数据就自动下载,transform=ToTensor() 则会把图片转换为 PyTorch 可以处理的张量。

图片原本是二维结构,但全连接神经网络通常需要一维向量作为输入。因此,一张 28×28 的图片最终会被转换成长度为 784 的向量。

例如,一张图片可以表示为:

复制代码
28 × 28 = 784

这 784 个数字就是图片的像素特征。模型要做的,就是根据这些像素值判断图片对应的是哪个数字。

二、DataLoader 的作用

数据集准备好之后,还需要使用 DataLoader 按照批次读取数据。训练时通常不会一次性把所有图片都送入模型,而是将数据划分成多个 batch。

复制代码
from torch.utils.data import DataLoader

batch_size = 64

train_dataloader = DataLoader(
    training_data,
    batch_size=batch_size
)

test_dataloader = DataLoader(
    test_data,
    batch_size=batch_size
)

这里设置 batch_size=64,表示模型每次读取 64 张图片。使用 batch 训练有很多好处。

第一,可以减少内存占用。如果一次性加载全部数据,可能会占用大量内存。第二,使用小批量数据可以提高训练效率。第三,每个 batch 产生的梯度具有一定随机性,有助于模型跳出一些不理想的状态。

在实际训练中,每一个 batch 通常包含两部分内容:

复制代码
X, y

其中 X 是输入图片,y 是对应的真实标签。例如,y 可能是数字 3,表示这张图片真实代表数字 3。

可以通过下面的代码查看一个 batch 的形状:

复制代码
for X, y in train_dataloader:
    print("Shape of X:", X.shape)
    print("Shape of y:", y.shape)
    break

如果 batch size 是 64,那么输入数据的形状通常是:

复制代码
X: [64, 1, 28, 28]
y: [64]

其中 64 表示图片数量,1 表示灰度通道,28 和 28 分别表示图片的高度和宽度。

三、神经网络的基本结构

一个神经网络通常由输入层、隐藏层和输出层组成。

对于 MNIST 手写数字识别任务来说,输入层接收 784 个像素值,隐藏层负责提取和组合特征,输出层输出 10 个类别的结果。

模型结构可以简单表示为:

复制代码
输入层:784 个节点
隐藏层:128 个节点
隐藏层:256 个节点
输出层:10 个节点

这里的 10 个输出节点分别对应数字 0 到 9。模型输出的不是直接的数字,而是每个类别对应的分数。分数最高的那个位置,就可以作为模型最终的预测结果。

在 PyTorch 中,自定义神经网络一般需要继承 nn.Module

复制代码
import torch
from torch import nn

class NeuralNetwork(nn.Module):
    def __init__(self):
        super().__init__()

        self.flatten = nn.Flatten()
        self.hidden1 = nn.Linear(28 * 28, 128)
        self.hidden2 = nn.Linear(128, 256)
        self.output = nn.Linear(256, 10)

    def forward(self, x):
        x = self.flatten(x)
        x = self.hidden1(x)
        x = torch.relu(x)
        x = self.hidden2(x)
        x = torch.relu(x)
        x = self.output(x)
        return x

nn.Flatten() 的作用是将输入图片展开。输入形状从 [batch_size, 1, 28, 28] 变成 [batch_size, 784]

nn.Linear 表示全连接层。第一层把 784 个输入特征映射到 128 个神经元,第二层把 128 个特征映射到 256 个神经元,最后一层把 256 个特征映射到 10 个输出类别。

四、为什么需要激活函数?

如果神经网络的每一层都只是线性变换,那么即使堆叠很多层,最终得到的仍然只是一个线性函数。这样的模型表达能力非常有限,无法处理复杂的图像特征。

激活函数的作用,就是在网络中加入非线性能力,让模型能够学习更加复杂的规律。

常见的激活函数包括 Sigmoid、Tanh、ReLU 和 Leaky ReLU。

1. Sigmoid 函数

Sigmoid 的输出范围是 0 到 1,表达式为:

复制代码
σ(x) = 1 / (1 + e^(-x))

它的曲线呈现出类似"S"形。当输入值特别大或特别小时,函数会逐渐趋于饱和,导数接近 0。

这会带来一个问题:在反向传播过程中,多个接近 0 的导数连续相乘后,最终的梯度会变得非常小,这就是梯度消失。

2. Tanh 函数

Tanh 的输出范围是 -1 到 1,且以 0 为中心。它比 Sigmoid 更适合某些需要零中心输出的场景,但当输入值绝对值较大时,同样会出现梯度变小的问题。

3. ReLU 函数

ReLU 的表达式非常简单:

复制代码
f(x) = max(0, x)

当输入大于 0 时,输出等于输入;当输入小于等于 0 时,输出为 0。

ReLU 在正数区域的导数为 1,因此可以有效缓解梯度消失问题,而且计算速度非常快,是目前深度神经网络中最常用的激活函数之一。

不过,ReLU 也存在一个问题。当某个神经元长期接收到负数输入时,它的输出可能一直为 0,梯度也一直为 0,这种现象被称为"神经元死亡"。

4. Leaky ReLU

Leaky ReLU 对 ReLU 做了改进。在负数区域不再完全输出 0,而是保留一个很小的斜率。

复制代码
f(x) = x,x > 0
f(x) = αx,x ≤ 0

这样即使输入为负数,也能保留一部分梯度,从而降低神经元死亡的概率。

五、梯度消失和梯度爆炸

神经网络训练的核心是反向传播。在反向传播过程中,模型会根据链式法则计算损失函数对每一个参数的梯度。

假设网络有多层结构,那么前面层的梯度通常需要乘以后面多层的导数和权重。如果这些乘数大部分都小于 1,连续相乘后梯度就会越来越小,最终导致梯度消失。

梯度消失会带来以下问题:

  1. 前面层的参数几乎不再更新;
  2. 模型训练速度变慢;
  3. 深层网络难以学习有效特征;
  4. 损失函数下降不明显。

相反,如果连续相乘的因素大部分大于 1,梯度就可能越来越大,最终出现梯度爆炸。梯度爆炸会导致参数数值迅速变大,损失函数出现 NaN,训练过程无法继续。

解决这些问题的方法包括使用 ReLU 等激活函数、合理初始化权重、使用 BatchNorm、调整学习率以及进行梯度裁剪等。

六、损失函数的作用

模型输出预测结果后,还需要知道这个结果到底好不好。损失函数就是用来衡量预测结果与真实标签之间差距的工具。

对于多分类问题,通常使用交叉熵损失函数:

复制代码
loss_fn = nn.CrossEntropyLoss()

交叉熵损失会比较模型对每个类别的预测分数和真实标签之间的差异。模型预测越准确,损失越小;模型预测越偏离真实标签,损失越大。

需要注意的是,使用 CrossEntropyLoss 时,模型最后一层通常不需要手动添加 Softmax。因为交叉熵损失函数内部已经完成了相关计算。

七、优化器与梯度下降

损失函数只能告诉我们模型当前表现如何,但不能自动修改模型参数。优化器的作用,就是根据梯度更新网络中的权重和偏置。

例如,可以使用 SGD 优化器:

复制代码
optimizer = torch.optim.SGD(
    model.parameters(),
    lr=1e-3
)

这里的 lr 是学习率。它决定了每次参数更新的幅度。

参数更新的基本公式为:

复制代码
θ = θ - η × ∇L(θ)

其中,θ 表示模型参数,η 表示学习率,∇L(θ) 表示损失函数关于参数的梯度。

如果学习率太大,模型可能会越过最优点,甚至导致损失越来越大。如果学习率太小,模型虽然能够慢慢收敛,但训练时间会变得很长。

除了 SGD,还可以使用 Adam:

复制代码
optimizer = torch.optim.Adam(
    model.parameters(),
    lr=1e-3
)

Adam 会根据历史梯度自动调整不同参数的学习步长,通常收敛速度比较快,适合快速实验。

八、训练循环

训练循环是整个项目中最核心的部分。一个完整的训练过程通常包含以下步骤:

  1. 从 DataLoader 中读取一个 batch;
  2. 将数据传入模型;
  3. 计算预测结果;
  4. 根据预测结果计算损失;
  5. 清空之前累积的梯度;
  6. 进行反向传播;
  7. 使用优化器更新参数。

代码如下:

复制代码
def train(dataloader, model, loss_fn, optimizer):
    model.train()

    for X, y in dataloader:
        pred = model(X)
        loss = loss_fn(pred, y)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

这里的 model.train() 表示进入训练模式。某些网络层,例如 Dropout 和 BatchNorm,在训练模式和测试模式下的行为不同,所以训练前需要调用这个函数。

optimizer.zero_grad() 用来清除上一个 batch 保存的梯度。PyTorch 默认会累积梯度,如果不清零,梯度就会不断叠加,影响参数更新。

loss.backward() 会自动计算损失函数关于模型参数的梯度。

optimizer.step() 会根据计算出来的梯度更新模型参数。

九、如何评估模型?

训练完成后,需要使用测试集评估模型的准确率。

复制代码
def test(dataloader, model, loss_fn):
    model.eval()

    test_loss = 0
    correct = 0

    with torch.no_grad():
        for X, y in dataloader:
            pred = model(X)
            test_loss += loss_fn(pred, y).item()
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()

    test_loss /= len(dataloader)
    accuracy = correct / len(dataloader.dataset)

    print("Test Error:")
    print(f"Accuracy: {accuracy * 100:.2f}%")
    print(f"Average loss: {test_loss:.4f}")

model.eval() 表示进入测试模式。

torch.no_grad() 表示测试时不计算梯度。因为测试阶段不需要更新参数,所以关闭梯度可以节省内存并提升运行速度。

pred.argmax(1) 表示在每一行输出分数中找到最大值的位置。这个位置就是模型认为最可能的类别。

例如:

复制代码
模型输出:[1.2, -0.3, 4.5, 0.8, ...]

如果最大值 4.5 位于下标 2,那么模型预测的数字就是 2。

十、完整训练流程

模型、损失函数和优化器准备好之后,就可以开始训练。

复制代码
model = NeuralNetwork()
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(
    model.parameters(),
    lr=1e-3
)

epochs = 10

for t in range(epochs):
    print(f"Epoch {t + 1}")
    train(train_dataloader, model, loss_fn, optimizer)
    test(test_dataloader, model, loss_fn)

print("Done!")

epochs 表示整个训练集被模型重复学习多少遍。训练轮数太少,模型可能还没有学会足够的特征;训练轮数太多,又可能出现过拟合。

在每个 epoch 结束后,可以观察训练损失和测试准确率。如果训练损失不断下降,说明模型正在学习。如果训练集准确率很高,但测试集准确率没有提升,可能说明模型已经过拟合。

十一、CPU 和 GPU 的选择

如果计算机中有 NVIDIA GPU,可以使用 CUDA 加速训练。

复制代码
device = (
    "cuda"
    if torch.cuda.is_available()
    else "cpu"
)

model = NeuralNetwork().to(device)

训练时,也需要把输入数据和标签移动到同一个设备:

复制代码
X = X.to(device)
y = y.to(device)

模型和数据必须位于同一个设备上,否则会出现设备不匹配的错误。

在使用 GPU 时,训练速度通常会明显提升,尤其是当模型结构较复杂、数据量较大时。不过,对于 MNIST 这样的简单任务,CPU 也可以完成训练。

十二、训练中常见的问题

1. 忘记清零梯度

如果不调用:

复制代码
optimizer.zero_grad()

那么每个 batch 的梯度会不断累积,导致参数更新异常。

2. 忘记切换模型模式

训练时需要:

复制代码
model.train()

测试时需要:

复制代码
model.eval()

如果忘记切换,Dropout 和 BatchNorm 等层可能产生不正确的结果。

3. 输出层维度错误

MNIST 有 10 个类别,因此输出层必须是:

复制代码
nn.Linear(256, 10)

如果输出类别数量不正确,模型就无法与标签对应。

4. 输入形状错误

全连接层需要二维输入,通常形状是:

复制代码
[batch_size, 784]

如果直接把 [batch_size, 1, 28, 28] 输入到线性层,就会出现维度不匹配。因此需要使用 Flatten 展开图片。

5. 学习率不合适

学习率太大时,损失可能不下降;学习率太小时,训练过程会非常缓慢。遇到这种情况,可以尝试调整学习率,或者更换优化器。

十三、从全连接网络到卷积神经网络

虽然全连接网络可以完成 MNIST 识别,但它没有充分利用图像的空间结构。图片中的相邻像素往往具有很强的关联性,而全连接层会把所有像素看成相互独立的特征。

卷积神经网络通过卷积核提取局部特征,例如边缘、线条和纹理,再通过多层卷积逐渐组合成更加复杂的结构。

对于手写数字来说,模型可能先学习简单的横线和竖线,再学习圆弧、交叉和数字整体形状。相比全连接网络,卷积神经网络通常更适合图像任务,也能减少参数数量。

后续可以将模型改成如下结构:

复制代码
class CNN(nn.Module):
    def __init__(self):
        super().__init__()

        self.network = nn.Sequential(
            nn.Conv2d(1, 32, kernel_size=3),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(32, 64, kernel_size=3),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Flatten(),
            nn.Linear(64 * 5 * 5, 10)
        )

    def forward(self, x):
        return self.network(x)

卷积网络的结构虽然比全连接网络复杂,但核心训练流程并没有改变,仍然是前向传播、损失计算、反向传播和参数更新。

十四、这次学习的收获

通过这个手写数字识别项目,我对深度学习的理解从"知道一些概念"变成了"能够把概念写成代码"。

以前看到神经网络结构图时,常常只知道输入层、隐藏层和输出层的名称,却不清楚每一层到底做了什么。实际编写模型后,我发现每一层其实都对应具体的张量变换。

输入图片进入模型后,首先被展平为向量;线性层通过权重矩阵和偏置完成特征变换;激活函数增加非线性能力;输出层给出各个类别的分数;损失函数衡量预测结果与真实标签之间的差异;反向传播则计算每个参数应该如何调整。

整个过程可以概括为:

复制代码
数据输入 → 前向传播 → 计算损失 → 反向传播 → 更新参数

模型就是在这个循环中不断学习,逐渐提高预测准确率。

十五、总结

PyTorch 为深度学习提供了非常灵活和方便的开发方式。通过 DatasetDataLoader,可以快速准备训练数据;通过继承 nn.Module,可以自由设计网络结构;通过自动微分机制,可以自动完成梯度计算;通过优化器,可以方便地更新模型参数。

MNIST 手写数字识别虽然是一个入门案例,但它完整展示了深度学习项目的基本流程。理解这个案例之后,再去学习卷积神经网络、目标检测、图像分割和自然语言处理,就会有更加扎实的基础。

这次学习也让我认识到,深度学习不仅仅是调用几个函数,更重要的是理解每一个组件背后的原理。只有知道数据如何流动、参数如何更新、损失如何变化,才能真正理解模型为什么有效,以及出现问题时应该从哪里排查。

接下来,我准备继续学习卷积神经网络,并尝试加入数据增强、BatchNorm、Dropout 和学习率调度等技术。通过不断实验和调整模型结构,让模型从简单的数字识别任务逐步走向更加复杂的图像分类任务。

相关推荐
吴佳浩1 小时前
大模型是怎么来的:从数据到 Foundation Model
人工智能·llm·agent
IT_陈寒1 小时前
Redis Pipeline用错竟比不用还慢,这个坑我帮你踩过了
前端·人工智能·后端
Henry-SAP1 小时前
AI新闻精选:聚焦前沿科技
人工智能·erp
庖丁AI1 小时前
PDF 转 Markdown 工具怎么选?AI 知识库和 RAG 场景要注意什么
人工智能·pdf·文档解析
AI推荐率1 小时前
为每种主体关系建立标准句:企业归属关系实操清单
人工智能
小小程序猴11 小时前
企业AI培训怎么选?一个四层加权评估模型的实现:需求定位 + 深潜解析
人工智能
人工智能时代 准备好了吗1 小时前
产品搜索别名应该怎么用?
人工智能
火山引擎开发者社区1 小时前
多行业专家招募|参与访谈最高拿 4000元/h 现金报酬!
人工智能
大模型码小白1 小时前
数据可视化:AI处理多维数据的HTML5可视化方案
大数据·前端·javascript·人工智能·机器学习·信息可视化·html5