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. 看看数据长什么样

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

python 复制代码
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张图片。

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

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

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

查看单批数据的维度:

python 复制代码
# 遍历测试集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):

python 复制代码
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

逐层解释:

  1. nn.Flatten():将[batch, 1, 28, 28]转为[batch, 784],把图片像素展平为一维,满足全连接层输入要求。
  2. nn.Linear():全连接层,执行y = xW^T + b运算,神经元数量可自定义。
  3. torch.sigmoid(x):激活函数,引入非线性;若无激活函数,多层网络等价于单层线性模型,学习能力受限。常用替代还有ReLU、tanh。
  4. 输出层输出10个logits得分,得分下标最大即为预测数字。

为什么需要隐藏层? 输入直接连接输出属于简单线性模型,无法学习复杂特征。隐藏层用来提取笔画、边缘等底层特征,再组合为高级特征,完成分类。

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

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

8. 训练函数

python 复制代码
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. 测试函数

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(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. 损失函数和优化器

python 复制代码
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. 总结与思考

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

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

常见踩坑:

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

改进方向:

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

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

相关推荐
其实防守也摸鱼11 分钟前
智能体推荐:精选 AI Agent 工具与实战指南
运维·开发语言·人工智能·学习·web安全·自动化
always_TT11 分钟前
【Python 字符串格式化:format() 方法】
android·开发语言·python
风123456789~16 分钟前
【架构设计】3.2 信息化系统的典型应用 3/6
笔记·系统架构设计
小义_16 分钟前
JDK 深度解析
java·linux·开发语言·python·面试
SendTomo19 分钟前
send.wang:基于浏览器WebRTC实现无客户端文件互传
网络·python·网络协议·webrtc·p2p
β添砖java24 分钟前
深度学习24物体检测算法R-CNN、SSD、YOLO
人工智能·深度学习
向哆哆27 分钟前
打击罂粟种植检测数据集分享(适用于YOLO系列深度学习分类检测任务)
深度学习·yolo·目标检测·分类
咖啡星人k28 分钟前
2026 AI Agent 智能体:ReAct 让 AI 边想边做、把大任务拆成小步骤(MonkeyCode 实战)
人工智能·深度学习·机器学习·语言模型·自然语言处理
风123456789~29 分钟前
【架构设计】第3章 信息系统基础知识 1/6
笔记·系统架构设计