欢迎光临
我们一直在努力

PyTorch 实现 MNIST 手写数字识别学习笔记

最近在学习了PyTorch的基础知识,讲了MNIST手写数字识别这个经典例子。这个项目麻雀虽小五脏俱全,包含了数据加载、模型搭建、训练和测试的完整流程。本文我会把代码拆开揉碎,用大白话讲解每一步在做什么,以及那些容易踩坑的地方。如果你也是刚入门深度学习,不妨跟着走一遍。


1. 项目背景

MNIST数据集包含7万张手写数字图片,其中6万张用于训练,1万张用于测试。图片是28×28的灰度图,数字已经居中,预处理很简单。我们的目标就是训练一个神经网络,让它能认出图片里写的是0‑9中的哪个数字。


2. 环境准备

首先导入必要的库:

import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets
from torchvision.transforms import ToTensor
import matplotlib.pyplot as plt

  • torch:PyTorch核心库

  • nn:神经网络模块,包含各种层和损失函数

  • DataLoader:数据加载器,负责批量打包数据

  • datasets:torchvision中的数据集工具,可以直接下载MNIST

  • ToTensor:把PIL图像或numpy数组转换成张量(tensor),并归一化到[0,1]


3. 下载并加载数据

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

这里做了几件事:

  • 从网上下载MNIST数据集到本地data文件夹(如果已存在就不会重复下载)

  • train=True表示加载训练集(6万张),train=False加载测试集(1万张)

  • transform=ToTensor():把图片转换成PyTorch张量,并且像素值从0‑255缩放到0‑1之间,方便神经网络处理

小知识: 为什么要把数据变成张量?因为PyTorch的模型只能处理张量,张量可以放在GPU上加速计算,而numpy数组只能在CPU上跑。


4. 看看数据长什么样

训练之前先可视化几张图片,确认数据没问题:

figure = plt.figure()
for i in range(9):
img, label = training_data[i] # 取出第i个样本:img为图像张量,label为对应数字标签
figure.add_subplot(3, 3, i+1) # 创建3行3列子图,选中第i+1个子画布
plt.title(label) # 设置子图标题为图片真实标签
plt.axis("off") # 关闭坐标轴,不显示刻度边框
plt.imshow(img.squeeze(), cmap="gray") # 将张量绘制为图片
a = img.squeeze() # 去除张量中维度为1的通道维度

plt.show() # 把画布整体渲染弹出显示

 img原始shape:[1,28,28],1代表灰度图通道数;squeeze()会删除大小等于1的维度,得到[28,28]
 imshow无法处理带单通道的三维张量,所以需要squeeze降维
 cmap="gray" 指定灰度色彩映射,保证图片以黑白灰度形式展示


5. 创建DataLoader

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

DataLoader的作用是把数据集切分成一个个小批量(batch),本案例每个batch包含64张图片。

  • 减少内存占用:不需要一次性把全部图片加载到内存

  • 提高训练速度:每次参数更新仅使用一小批样本,计算效率更高

  • 引入随机性:默认打乱样本顺序,有助于提升模型泛化能力

查看单批数据的维度:

# 遍历测试集dataloader,查看一个batch的数据维度,只取第一批就break,不完整遍历整个数据集
for X, y in test_dataloader:
# X:一批图片张量,格式 [N, C, H, W] N批次大小、C通道数、H图片高、W图片宽
print(f"Shape of X [N, C, H, W]: {X.shape}")
# y:这批样本对应的标签,dtype打印标签的数据类型
print(f"Shape of y: {y.shape} {y.dtype}")
break # 只看第一个batch的形状,直接跳出循环,避免打印全部数据

输出结果:X形状[64, 1, 28, 28],代表64张图片,单张1通道,高28、宽28;y形状[64],对应64个样本的数字标签。


6. 选择设备

device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"Using {device} device")

根据硬件自动选择计算设备:

  • NVIDIA显卡:使用cuda

  • 苹果M系列芯片:使用mps

  • 其余环境:使用cpu

重要提醒: 模型与输入数据必须处于同一个设备,后面通过model.to(device)、X.to(device)完成迁移。


7. 构建神经网络模型

本项目使用简单的全连接网络(多层感知机MLP):

class NeuralNetwork(nn.Module): # 继承PyTorch内置的nn.Module父类
def __init__(self):
super().__init__() # 调用父类nn.Module的构造函数
self.flatten = nn.Flatten() # 把28×28的图片拉平成一维向量
self.hidden1 = nn.Linear(28*28, 128) # 输入784个神经元,输出128个
self.hidden2 = nn.Linear(128, 256) # 第二层隐藏层
self.out = nn.Linear(256, 10) # 输出层,对应10个数字

def forward(self, x):
x = self.flatten(x) # [batch, 1, 28, 28] -> [batch, 784]
x = self.hidden1(x) # [batch, 784] -> [batch, 128]
x = torch.sigmoid(x) # 激活函数
x = self.hidden2(x) # [batch, 128] -> [batch, 256]
x = torch.sigmoid(x) # 激活函数
x = self.out(x) # [batch, 256] -> [batch, 10]
return x

逐层解释:

  • nn.Flatten():将[batch, 1, 28, 28]转为[batch, 784],把图片像素展平为一维,满足全连接层输入要求。
  • nn.Linear():全连接层,执行y = xW^T + b运算,神经元数量可自定义。
  • torch.sigmoid(x):激活函数,引入非线性;若无激活函数,多层网络等价于单层线性模型,学习能力受限。常用替代还有ReLU、tanh。
  • 输出层输出10个logits得分,得分下标最大即为预测数字。
  • 为什么需要隐藏层? 输入直接连接输出属于简单线性模型,无法学习复杂特征。隐藏层用来提取笔画、边缘等底层特征,再组合为高级特征,完成分类。

    实例化模型并迁移到设备:

    model = NeuralNetwork().to(device) # 把模型权重迁移到指定设备(cuda/mps/cpu)
    print(model)


    8. 训练函数

    def train(dataloader, model, loss_fn, optimizer):
    model.train() # 切换到训练模式
    batch_size_num = 1 # 统计 训练的batch数量
    for X, y in dataloader:
    X, y = X.to(device), y.to(device)
    # 前向传播
    pred = model(X)
    # 计算损失
    loss = loss_fn(pred, y)
    # 反向传播
    optimizer.zero_grad() # 梯度清零
    loss.backward() # 计算梯度
    optimizer.step() # 更新参数
    # 打印损失
    if batch_size_num % 100 == 0:
    loss_value = loss.item()
    print(f"loss: {loss_value:>7f} [number:{batch_size_num}]")
    batch_size_num += 1

    关键点解析:

    • model.train():开启训练模式,部分层(Dropout、BatchNorm)训练、测试行为不一样,养成书写习惯。
    • pred = model(X):自动调用forward(),执行前向传播,不要手动写model.forward(X)。
    • loss_fn(pred, y):计算预测值与真实标签之间的损失。
    • optimizer.zero_grad():梯度清零;PyTorch默认梯度累加,每个batch训练前必须清零,否则参数更新异常。
    • loss.backward():反向传播,自动求解各可训练参数的梯度。
    • optimizer.step():依据梯度更新网络权重。

    9. 测试函数

    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(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}")

    注意点:

    • model.eval():切换评估模式。
    • torch.no_grad():测试阶段关闭梯度计算,节省内存、加速推理。
    • pred.argmax(1):按行取最大值索引,得到预测数字。
    • 布尔张量转为浮点型,求和统计样本预测正确的总数量。

    10. 损失函数和优化器

    loss_fn = nn.CrossEntropyLoss() #创建交叉熵损失函数对象
    optimizer = torch.optim.SGD(model.parameters(), lr=0.01)#创建一个优化器,SGD为随机梯度下降算法

    • 损失函数:CrossEntropyLoss交叉熵损失,多用于多分类任务。内部自动完成softmax,模型输出直接传logits即可,无需额外添加softmax层。
    • 优化器:SGD随机梯度下降,lr=0.01为学习率。学习率代表参数更新步长;学习率过大容易震荡不收敛,过小训练速度慢。工程中Adam使用更加广泛。

    补充说明:交叉熵先将输出分数转为概率,取真实类别对应概率做负对数运算;概率越接近1,损失数值越小。


    11. 开始训练

    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)

    设置10轮epoch,一个epoch代表完整遍历一遍全部训练集。示例代码只在全部训练结束后执行一次测试;训练过程每100个batch打印损失,损失逐步下降代表模型在学习。


    12. 完整代码

    将上述所有代码按顺序复制运行,注意检查缩进与变量名。


    13. 总结与思考

    通过该项目完整走完深度学习标准流程:

  • 加载数据并预处理
  • 定义模型结构
  • 选择损失函数和优化器
  • 循环训练:前向传播 → 计算损失 → 反向传播 → 更新参数
  • 在测试集上评估性能
  • 常见踩坑:

    • 设备不匹配:模型在GPU,数据在CPU直接报错,数据、模型必须统一to(device)。
    • 忘记梯度清零,损失不下降、来回震荡。
    • CrossEntropyLoss输入不需要手动加softmax,额外添加会影响效果。

    改进方向:

    • 使用卷积神经网络CNN替换全连接网络,进一步提升识别准确率。
    • 将sigmoid替换为ReLU激活函数。
    • 更换Adam优化器,调试学习率。
    • 引入数据增强(旋转、平移),提升模型泛化能力。

    希望这篇文章能帮你理清PyTorch的基本用法。如果还有疑问,欢迎在评论区交流。

    赞(0)
    未经允许不得转载:171主机测评 » PyTorch 实现 MNIST 手写数字识别学习笔记
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址