PyTorch 手写数字识别实战:从数据加载到模型训练(零基础详解版)

1. 引言

手写数字识别是深度学习入门最经典的实战项目之一,它基于 MNIST 数据集,目标是把 28×28 像素的灰度图片正确分类为 0 到 9 十个数字。本文结合一段完整的 PyTorch 代码,从数据加载、数据加载器、设备选择、神经网络构建、训练与测试等环节,系统讲解深度学习的基础概念,帮助读者理解模型是如何一步步学会识别数字的。

如果你只学过 Python 基础语法,不用担心。本文会尽量用通俗的语言解释每一个概念,包括什么是张量、什么是神经网络、什么是训练。你只需要跟着代码一步步看下去,就能理解整个流程。

2. 环境准备与依赖库

在开始之前,需要确保本机已安装 PyTorch 及其配套库。本文代码主要依赖以下三个核心库:

  • torch:PyTorch 主库,提供张量计算、神经网络模块和自动求导能力。
  • torchvision:PyTorch 的视觉工具库,提供常用数据集、图像变换和预训练模型。
  • torchaudio:PyTorch 的音频工具库,本文主要用于验证安装完整性。

这三个库的关系可以这样理解:torch 是核心引擎,负责所有数学计算;torchvision 是专门为图像任务准备的工具箱,里面已经封装好了 MNIST 等常用数据集;torchaudio 则是音频相关的工具箱,本文用不到它的功能,只是顺便验证一下安装是否完整。

可以通过以下命令验证三个库是否安装成功:

python 复制代码
import torch
import torchvision
import torchaudio
print(torch.__version__)
print(torchvision.__version__)
print(torchaudio.__version__)

这里用到了 Python 的 import 语句,它的作用是把别人写好的功能模块引入到当前代码中,这样我们就能直接使用 torch 提供的各种函数和类了。print 是 Python 内置的输出函数,会把括号里的内容打印到屏幕上。torch.version 是 torch 模块内部保存的版本号字符串,打印出来就能确认安装是否成功。

如果能够正常输出版本号,说明环境已就绪,可以继续后续步骤。

3. 数据准备:MNIST 数据集

MNIST 是深度学习领域最经典的入门数据集,包含 60000 张训练图片和 10000 张测试图片,每张图片都是 28×28 像素的灰度图,对应 0 到 9 十个数字类别。torchvision 提供了直接下载和加载 MNIST 的接口,使用非常方便。

这里先解释两个概念:训练集和测试集。训练集是给模型学习用的数据,就像学生做练习题;测试集是模型学完之后用来检验学习效果的数据,就像期末考试。模型在训练阶段从未见过测试集,这样才能真实反映它的学习能力。

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

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

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

这段代码调用了 datasets.MNIST 这个函数来创建数据集对象。在 Python 中,函数调用就是执行一段预先写好的逻辑,这里 MNIST 函数会负责下载数据、读取图片、整理成统一的格式。等号左边的 training_data 和 test_data 是变量名,用来保存函数返回的结果,之后我们可以通过这两个变量来访问数据集。

这里有几个关键参数需要理解:

  • root:数据集保存的本地目录,首次运行会自动下载。
  • train:为 True 时加载训练集,为 False 时加载测试集。
  • transform:对图片应用的预处理操作,ToTensor() 会把 PIL 图片或 NumPy 数组转换为张量,并把像素值从 0 到 255 归一化到 0 到 1 之间。
  • download:如果本地不存在数据集,则自动从网络下载。

关于 transform 参数,可以这样理解:原始图片在计算机里是以 0 到 255 的整数存储的,数值越大表示越亮。ToTensor() 会做两件事:一是把图片转换成 PyTorch 能处理的张量格式,二是把所有像素值除以 255,变成 0 到 1 之间的小数。这样做的好处是数值范围更稳定,模型训练时更容易收敛。

4. 数据可视化:查看样本图片

在训练之前,先可视化一部分样本,直观感受数据集的内容。下面的代码使用 matplotlib 在一个窗口中绘制 9 张训练图片及其标签:

python 复制代码
from matplotlib import pyplot as plt

figure = plt.figure()
for i in range(9):
    img, label = training_data[i]
    figure.add_subplot(3, 3, i + 1)
    plt.title(label)
    plt.axis("off")
    plt.imshow(img.squeeze(), cmap="gray")
plt.show()

这段代码用到了 Python 的 for 循环。range(9) 会生成 0 到 8 这 9 个整数,循环体里的代码会依次执行 9 次,每次 i 的值分别是 0、1、2......8。training_datai 表示取出数据集中的第 i 张图片,返回两个值:img 是图片数据,label 是对应的数字标签。这种同时给两个变量赋值的写法叫元组解包,是 Python 中很常见的语法。

figure.add_subplot(3, 3, i + 1) 的作用是把整个绘图窗口划分成 3 行 3 列共 9 个小格子,第 i + 1 次循环就在第 i + 1 个格子里画图。plt.title(label) 给当前小图加上标题,显示这个数字是几。plt.imshow 负责把图片数据显示成图像,plt.show 最后把整个窗口显示出来。

这里有两个重要的知识点:

  • img.squeeze():从张量中去掉维度为 1 的轴。MNIST 单张图片的形状是 (1, 28, 28),第一个维度表示通道数,灰度图只有 1 个通道,squeeze() 后变成 (28, 28),才能被 imshow 正常显示。
  • cmap="gray":指定使用灰度色彩映射来显示图像,否则 matplotlib 会使用默认的彩色映射。

关于维度,可以这样理解:一张彩色图片有三个通道(红、绿、蓝),而灰度图只有一个通道。PyTorch 习惯用 (通道数, 高度, 宽度) 这样的顺序来表示图片,所以单张灰度图的形状是 (1, 28, 28)。squeeze() 会把长度为 1 的那个维度去掉,变成 (28, 28),这样 matplotlib 才能正确地把二维数组显示成图像。

5. 数据加载器:DataLoader

深度学习训练通常不会把整个数据集一次性送入模型,而是把数据分成多个小批次(batch)逐批处理。DataLoader 就是负责完成这个任务的工具。

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

train_dataloader = DataLoader(training_data, batch_size=64)
test_dataloader = DataLoader(test_data, batch_size=64)

batch_size 表示每个批次包含的样本数量,这里设置为 64,即每 64 张图片打包成一个批次。使用小批次训练有两大好处:

  • 减少内存占用:不需要一次性把全部数据加载到内存或显存中。
  • 提高训练速度:每次用一小批数据计算梯度并更新参数,比逐条样本更新更高效。

可以这样理解 DataLoader 的作用:它像一个自动分拣员,把 60000 张训练图片按每 64 张一包分好,训练时我们每次取一包来用。这样既不会因为数据太多撑爆内存,又能通过批量计算提高效率。

可以通过下面的代码查看一个批次数据的形状:

python 复制代码
for X, y in test_dataloader:
    print(f"Shape of X [N, C, H, W]: {X.shape}")
    print(f"Shape of y: {y.shape} {y.dtype}")
    break

这里的 for 循环会遍历 DataLoader 中的每一个批次。每次循环,X 是一个批次的所有图片,y 是这个批次所有图片对应的标签。break 语句表示只取第一个批次就跳出循环,因为我们只是想看一眼数据的形状,不需要全部遍历。

print(f"Shape of X N, C, H, W: {X.shape}") 用到了 Python 的 f-string 格式化语法。f 前缀表示这个字符串里的大括号 {X.shape} 会被替换成 X.shape 的实际值。X.shape 是张量的形状属性,返回一个元组,表示每个维度的大小。

输出中 X 的形状为 64, 1, 28, 28,分别表示批次大小、通道数、图片高度和图片宽度;y 的形状为 64,表示每个样本对应的数字标签。

6. 设备选择:CPU、CUDA 与 MPS

PyTorch 支持在多种硬件设备上运行。训练前需要判断当前环境可用的设备,并把模型和数据都移动到该设备上。

python 复制代码
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"Using {device} device")

这行代码用到了 Python 的三元表达式,它是一种简写形式的 if-else 语句。整行代码可以翻译成:如果 torch.cuda.is_available() 为真,就把 device 设为 "cuda";否则再看 torch.backends.mps.is_available() 是否为真,为真就设为 "mps";如果都不满足,就设为 "cpu"。这种嵌套的三元表达式虽然看起来复杂,但逻辑和普通的 if-else 完全一样。

这段代码的优先级是:

  • cuda:如果检测到 NVIDIA 显卡且 CUDA 可用,优先使用 GPU 加速。
  • mps:如果使用的是苹果 M 系列芯片,则使用 MPS 后端调用 GPU。
  • cpu:如果以上都不可用,则回退到 CPU 训练。

需要特别说明的是,CUDA 驱动的作用是让 PyTorch 能够执行 CUDA 指令,而 CUDA 再通过 GPU 指令集去控制显卡硬件。模型参数和每个批次的数据都必须先移动到 GPU 上,才能进行 GPU 加速计算。

可以这样理解:CPU 是电脑的通用大脑,什么都能算但速度一般;GPU 是专门为并行计算设计的,特别适合深度学习这种大量重复的矩阵运算。CUDA 是 NVIDIA 提供的一套软件接口,让程序能调用 GPU 的计算能力。如果你的电脑没有 NVIDIA 显卡,就老老实实用 CPU 训练,虽然慢一些,但结果是一样的。

7. 构建神经网络模型

本文使用一个简单的全连接神经网络来完成分类任务。模型通过继承 nn.Module 来定义,这是 PyTorch 中所有神经网络模型的基类。

在深入代码之前,先解释两个 Python 面向对象的概念。类是创建对象的模板,对象是类的具体实例。继承表示一个类可以复用另一个类已有的功能。这里的 NeuralNetwork 类继承了 nn.Module,意味着它自动拥有了 nn.Module 提供的各种能力,比如参数管理、设备迁移等,我们只需要专注于定义网络的结构。

python 复制代码
from torch import nn

class NeuralNetwork(nn.Module):
    def __init__(self):
        super().__init__()
        self.flatten = nn.Flatten()
        self.hidden1 = nn.Linear(28 * 28, out_features=128)
        self.hidden2 = nn.Linear(128, 256)
        self.out = nn.Linear(256, 10)

    def forward(self, x):
        x = self.flatten(x)
        x = self.hidden1(x)
        x = torch.sigmoid(x)
        x = self.hidden2(x)
        x = torch.sigmoid(x)
        x = self.out(x)
        return x

model = NeuralNetwork().to(device)
print(model)

这段代码定义了一个类。init 是 Python 类的构造函数,在创建对象时自动调用,用来初始化对象的属性。super().init() 是调用父类 nn.Module 的构造函数,确保父类的初始化逻辑也执行了。self 表示对象自身,self.flatten 等就是在给对象添加属性。

forward 方法定义了数据在网络中的流向。当你调用 model(x) 时,PyTorch 会自动调用 forward 方法,把 x 作为参数传进去。方法体里的每一行都在对数据进行一次变换,最后 return x 把结果返回给调用者。

这个模型包含以下几个关键部分:

  • nn.Flatten:把 28×28 的二维图片展开成一维向量,长度为 784,作为全连接层的输入。
  • nn.Linear:全连接层,第一个参数是输入神经元个数,第二个参数是输出神经元个数。hidden1 把 784 维映射到 128 维,hidden2 把 128 维映射到 256 维,out 层把 256 维映射到 10 维。
  • 输出层维度:输出必须是 10,因为手写数字一共有 10 个类别,每个输出对应一个数字的得分。
  • 激活函数:sigmoid 是非线性激活函数,它的作用是给网络引入非线性能力。如果没有激活函数,多层线性变换叠加后仍然等价于一层线性变换,无法拟合复杂的数据分布。除了 sigmoid,常用的激活函数还有 ReLU 和 tanh。
  • forward 方法:定义了数据在网络中的流向,即前向传播过程。这个方法名不能修改,当调用 model(x) 时,PyTorch 会自动调用 forward。

关于全连接层,可以这样理解:每一层都像是一个加工车间,输入是一批数字,经过加权求和和偏置调整后,输出一批新的数字。hidden1 接收 784 个输入数字,通过内部的计算产生 128 个输出数字;hidden2 再把这 128 个数字加工成 256 个;最后 out 层把 256 个数字加工成 10 个,分别代表 0 到 9 这十个数字的得分。得分最高的那个数字,就是模型认为图片最可能对应的数字。

关于激活函数,可以这样理解:如果没有激活函数,无论堆多少层线性变换,最终结果都等价于一层线性变换,模型的学习能力会非常有限。sigmoid 函数会把任意实数压缩到 0 到 1 之间,给网络引入非线性,让模型能够拟合更复杂的模式。

8. 损失函数与优化器

训练神经网络需要两个核心组件:损失函数和优化器。

python 复制代码
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

损失函数用于衡量模型预测结果与真实标签之间的差距。本文使用交叉熵损失函数(CrossEntropyLoss),它非常适合多分类任务。损失值越小,说明模型的预测越接近真实标签,模型训练得越好。

可以这样理解损失函数:模型对一张图片会输出 10 个得分,比如 0.1, 0.2, 0.3, 0.05, 0.15, 0.05, 0.05, 0.05, 0.05, 0.0,如果真实标签是 2,那么理想情况下第 2 个得分应该最高。损失函数会计算预测结果和真实标签之间的差距,差距越大,损失值越大。训练的目标就是不断减小这个损失值。

优化器负责根据损失值更新模型的参数。本文使用随机梯度下降算法(SGD),其中 lr 是学习率,表示每次参数更新的步长。学习率的选择很关键:学习率过大可能导致训练震荡甚至发散,学习率过小则收敛速度过慢。model.parameters() 返回模型中所有需要训练的参数,优化器会基于这些参数的梯度来更新它们。

可以这样理解优化器:模型内部有很多参数(权重和偏置),这些参数决定了模型的预测结果。优化器的作用就是根据损失值的大小,不断微调这些参数,让损失值越来越小。学习率决定了每次调整的幅度:步子迈得太大容易走过头,步子太小又走得太慢。

9. 训练函数

训练过程的核心是让模型在训练数据上不断迭代,通过反向传播更新参数,使损失值逐渐下降。

python 复制代码
def train(dataloader, model, loss_fn, optimizer):
    model.train()
    batch_size_num = 1
    for X, y in dataloader:
        X, y = X.to(device), y.to(device)
        pred = model.forward(X)
        loss = loss_fn(pred, y)

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

        loss_value = loss.item()
        if batch_size_num % 100 == 0:
            print(f"loss: {loss_value:>7f}  [number:{batch_size_num}]")
        batch_size_num += 1

这里用 def 关键字定义了一个函数。函数是 Python 中组织代码的基本单位,把一段可复用的逻辑封装起来。train 函数接收四个参数:dataloader 是数据加载器,model 是神经网络模型,loss_fn 是损失函数,optimizer 是优化器。函数体里的代码会在调用 train(...) 时执行。

函数体里有一个 for 循环,会遍历 dataloader 中的每一个批次。每次循环,X 是一个批次的所有图片,y 是对应的标签。batch_size_num 是一个计数器,从 1 开始,每处理完一个批次就加 1,用来统计已经处理了多少个批次。

训练函数中的每一步都有明确的含义:

  • model.train():把模型切换到训练模式。在训练模式下,模型中的参数会被更新,某些层(如 Dropout、BatchNorm)的行为也会与测试模式不同。
  • 数据移动到设备:X 和 y 都需要通过 .to(device) 移动到 GPU 或 CPU 上,才能参与计算。
  • 前向传播:model.forward(X) 计算模型的预测输出。实际上调用 model(X) 也会自动执行 forward,因为父类 nn.Module 已经实现了这个调用逻辑。
  • 计算损失:loss_fn(pred, y) 计算预测结果与真实标签之间的交叉熵损失。
  • 梯度清零:optimizer.zero_grad() 把上一次迭代累积的梯度清零。如果不清零,梯度会不断累加,导致参数更新错误。
  • 反向传播:loss.backward() 自动计算损失对每个参数的梯度。
  • 参数更新:optimizer.step() 根据计算出的梯度更新网络参数。

关于梯度,可以这样理解:梯度是一个数学概念,表示损失值对每个参数的敏感程度。如果某个参数稍微变大一点,损失值会下降很多,说明这个参数应该往变大的方向调整。反向传播就是利用链式法则,从输出层开始,一层一层地计算出每个参数的梯度。PyTorch 的自动求导功能帮我们完成了这些复杂的数学计算,我们只需要调用 loss.backward() 和 optimizer.step() 即可。

这里体现了深度学习训练的核心循环:前向传播计算损失,反向传播计算梯度,优化器更新参数。每处理一个批次的数据,就完成一次这样的循环。

10. 测试函数

训练完成后,需要在测试集上评估模型的泛化能力,即模型在未见过的数据上的表现。

python 复制代码
def test(dataloader, model, loss_fn):
    size = len(dataloader.dataset)
    num_batches = len(dataloader)
    model.eval()
    test_loss, correct = 0, 0
    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            pred = model.forward(X)
            test_loss += loss_fn(pred, y).item()
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()
    test_loss /= num_batches
    correct /= size
    print(f"Test result: \n Accuracy: {(100 * correct)}%, Avg loss: {test_loss}")

len(dataloader.dataset) 返回测试集的总样本数,len(dataloader) 返回批次数。test_loss 和 correct 是两个累加器,分别用来累计总损失和预测正确的样本数。with torch.no_grad(): 是一个上下文管理器,它告诉 PyTorch 在这个代码块里不需要计算梯度,从而节省内存和计算时间。

测试函数有几个关键点:

  • model.eval():把模型切换到评估模式,此时模型不会更新参数,某些层的行为也会切换到推理模式。
  • torch.no_grad():在测试阶段关闭梯度计算,因为不需要反向传播,这样可以节省内存并加快计算速度。
  • pred.argmax(1):在预测结果的第 1 个维度上取最大值对应的索引,即模型认为最可能的数字类别。
  • 准确率计算:把预测正确的样本数除以总样本数,得到模型在测试集上的准确率。

关于 argmax,可以这样理解:模型对一张图片会输出 10 个得分,argmax(1) 会找出这 10 个得分中最大的那个对应的索引。比如输出 0.1, 0.8, 0.05, ...,最大的是 0.8,索引是 1,说明模型认为这张图片最可能是数字 1。然后把这个预测结果和真实标签 y 比较,相等就说明预测正确。

11. 完整训练流程

把以上各个部分组合起来,就构成了完整的训练流程。首先创建损失函数和优化器,然后训练一轮,最后进行多轮迭代训练并在测试集上评估。

python 复制代码
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

train(train_dataloader, model, loss_fn, optimizer)

epochs = 10
for t in range(epochs):
    print(f"Epoch {t + 1}\n-------------------")
    train(train_dataloader, model, loss_fn, optimizer)
print('Done')

test(test_dataloader, model, loss_fn)

这里引入了一个重要的概念:epoch(轮次)。一个 epoch 表示模型完整地遍历了一遍训练集。本文设置了 10 个 epoch,即模型会反复学习训练数据 10 遍。每经过一个 epoch,模型通常会对数据有更好的拟合,损失值逐渐下降,准确率逐渐提升。

可以这样理解 epoch:就像学生复习功课,只看一遍不一定记得住,需要反复复习多遍才能掌握。每一遍复习就是一次 epoch。for t in range(epochs) 会循环 10 次,每次调用 train 函数让模型完整学习一遍训练集。print(f"Epoch {t + 1}\n-------------------") 用来打印当前是第几轮,方便观察训练进度。

训练过程中,可以观察到每个批次输出的 loss 值。随着训练的进行,loss 会呈现下降趋势,这说明模型正在逐步学习到数字图片的特征。

12. 深度学习核心概念总结

通过这个手写数字识别项目,可以串联起深度学习的几个核心概念:

  • 张量:PyTorch 中的基本数据结构,可以理解为多维数组。图片、标签、模型参数都以张量的形式存在。
  • 数据集与数据加载器:数据集负责存储和管理样本,数据加载器负责按批次把数据送入模型。
  • 神经网络:由多个线性层和非线性激活函数堆叠而成的函数逼近器,通过调整参数来拟合输入到输出的映射关系。
  • 前向传播与反向传播:前向传播计算预测结果和损失,反向传播通过链式法则计算每个参数的梯度。
  • 损失函数:衡量预测与真实标签的差距,是模型优化的目标函数。
  • 优化器:根据梯度更新参数,常用的有 SGD、Adam 等。
  • 训练模式与评估模式:训练时更新参数,评估时只做预测,两者通过 model.train() 和 model.eval() 切换。
  • 设备管理:通过 .to(device) 把模型和数据移动到 GPU 或 CPU 上,充分利用硬件加速能力。

这些概念环环相扣:张量是数据的基本载体,数据集和数据加载器负责把数据组织好,神经网络是学习的核心,损失函数告诉我们学得好不好,优化器负责改进,前向传播和反向传播构成了学习的循环。

13. 总结

本文通过一个完整的 PyTorch 手写数字识别项目,系统讲解了从数据加载、数据可视化、数据加载器、设备选择、模型构建、损失函数与优化器,到训练与测试的完整流程。这段代码虽然结构简单,却涵盖了深度学习训练的全部核心环节,是理解更复杂模型的重要基础。

对于只学过 Python 基础语法的读者,建议先理解几个关键点:import 是引入功能模块,变量用来保存数据,for 循环用来重复执行代码,def 用来定义函数,class 用来定义类。在此基础上,再逐步理解张量、神经网络、损失函数这些深度学习概念,就会容易很多。

建议读者在理解每一行代码含义的基础上,动手修改 batch_size、学习率、隐藏层神经元个数等超参数,观察它们对训练速度和准确率的影响,从而加深对深度学习原理的理解。

相关推荐
TaoMetrix6 小时前
TaoMetrix:AI 硬件创业公司如何做好原型制造
人工智能·制造
青 春 记 忆6 小时前
零基础入门python30:Flask个人账本从空目录运行与阶段验收
python·flask·后端开发
智圣新创016 小时前
面向多级组织协同场景 高校第二课堂一站式管理中枢落地全场景实操答疑
大数据·人工智能
cd_949217216 小时前
AI纹理和Substance Painter手绘纹理哪个更适合游戏资产制作?
人工智能·游戏·substance painter
sali-tec6 小时前
C# 基于OpenCv的视觉工作流-章106-差值追踪
图像处理·人工智能·opencv·算法·计算机视觉
Shockang7 小时前
用 AI 打造高品质 Web 应用
人工智能
l1258657 小时前
# LangGraph Memory机制深度解析:短期记忆与长期记忆的工程实践
前端·人工智能·python·langchain·bootstrap
手写码匠7 小时前
Dify 多 Agent 工具权限与安全沙箱实战:让智能体“有能力,但不越权“
人工智能·深度学习·算法·aigc
ZGIAI7 小时前
ZGI Skill Loop:给工具调用设边界
人工智能·架构