学习VLA第3天:训练一个最简单的神经网络

1. 引言:从零开始,训练你的第一个神经网络

这是「学习 VLA」系列的第 3 天。前面的内容已经对视觉-语言-动作(Vision-Language-Action)模型有了初步认识。今天暂时放下复杂的多模态大模型,回到最基础的一步:亲手训练一个最简单的神经网络

VLA 模型再强大,底层也是由一个个基础神经网络模块堆叠而成。先把「砖块」烧好,才能盖起高楼。今天用 PyTorch 在经典的 MNIST 手写数字数据集上,训练一个简单的全连接网络,目标是把准确率做到 95% 以上。

第一次接触 PyTorch 也没关系。这篇文章会一步步拆解整个过程,从环境准备、数据加载,到模型搭建、训练与评估,每个环节都会讲清楚「是什么、为什么、怎么做」。学完今天的内容,就拥有了亲手训练一个神经网络的能力,这也是通往 VLA 世界的坚实第一步。

1.1 什么是 MNIST?

MNIST 是机器学习领域最经典的「Hello World」数据集。它包含了 6 万张训练图片1 万张测试图片 ,每张图片都是一个 28×28 像素 的手写数字(0 到 9),灰度图。

简单来说,任务就是:给计算机看一张手写数字的图片,让它告诉我们这个数字是几

1.2 什么是全连接网络?

全连接网络(Fully Connected Network)也叫多层感知机(MLP),是最基础的神经网络结构。它的核心思想是:把输入数据「摊平」成一维向量,然后通过一层层的矩阵乘法(线性变换)和激活函数(非线性变换),逐步提取特征,最终输出分类结果。

虽然现在卷积神经网络(CNN)在图像任务上更常用,但全连接网络结构简单、易于理解,非常适合作为入门第一个项目。

1.3 为什么目标定在 95%?

对于 MNIST 数据集,95% 的准确率是一个「跳一跳就够得着」的目标:它比随机猜测(10%)高出太多,但又不需要复杂的网络结构或调参技巧。用即将搭建的简单全连接网络,只需要训练几个 epoch 就能轻松达到,非常适合建立信心。

下面是今天要完成的整体流程:
#mermaid-svg-78APA6G4Ejz3tEqS{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-78APA6G4Ejz3tEqS .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-78APA6G4Ejz3tEqS .error-icon{fill:#552222;}#mermaid-svg-78APA6G4Ejz3tEqS .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-78APA6G4Ejz3tEqS .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-78APA6G4Ejz3tEqS .marker{fill:#333333;stroke:#333333;}#mermaid-svg-78APA6G4Ejz3tEqS .marker.cross{stroke:#333333;}#mermaid-svg-78APA6G4Ejz3tEqS svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-78APA6G4Ejz3tEqS p{margin:0;}#mermaid-svg-78APA6G4Ejz3tEqS .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-78APA6G4Ejz3tEqS .cluster-label text{fill:#333;}#mermaid-svg-78APA6G4Ejz3tEqS .cluster-label span{color:#333;}#mermaid-svg-78APA6G4Ejz3tEqS .cluster-label span p{background-color:transparent;}#mermaid-svg-78APA6G4Ejz3tEqS .label text,#mermaid-svg-78APA6G4Ejz3tEqS span{fill:#333;color:#333;}#mermaid-svg-78APA6G4Ejz3tEqS .node rect,#mermaid-svg-78APA6G4Ejz3tEqS .node circle,#mermaid-svg-78APA6G4Ejz3tEqS .node ellipse,#mermaid-svg-78APA6G4Ejz3tEqS .node polygon,#mermaid-svg-78APA6G4Ejz3tEqS .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-78APA6G4Ejz3tEqS .rough-node .label text,#mermaid-svg-78APA6G4Ejz3tEqS .node .label text,#mermaid-svg-78APA6G4Ejz3tEqS .image-shape .label,#mermaid-svg-78APA6G4Ejz3tEqS .icon-shape .label{text-anchor:middle;}#mermaid-svg-78APA6G4Ejz3tEqS .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-78APA6G4Ejz3tEqS .rough-node .label,#mermaid-svg-78APA6G4Ejz3tEqS .node .label,#mermaid-svg-78APA6G4Ejz3tEqS .image-shape .label,#mermaid-svg-78APA6G4Ejz3tEqS .icon-shape .label{text-align:center;}#mermaid-svg-78APA6G4Ejz3tEqS .node.clickable{cursor:pointer;}#mermaid-svg-78APA6G4Ejz3tEqS .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-78APA6G4Ejz3tEqS .arrowheadPath{fill:#333333;}#mermaid-svg-78APA6G4Ejz3tEqS .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-78APA6G4Ejz3tEqS .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-78APA6G4Ejz3tEqS .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-78APA6G4Ejz3tEqS .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-78APA6G4Ejz3tEqS .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-78APA6G4Ejz3tEqS .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-78APA6G4Ejz3tEqS .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-78APA6G4Ejz3tEqS .cluster text{fill:#333;}#mermaid-svg-78APA6G4Ejz3tEqS .cluster span{color:#333;}#mermaid-svg-78APA6G4Ejz3tEqS div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-78APA6G4Ejz3tEqS .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-78APA6G4Ejz3tEqS rect.text{fill:none;stroke-width:0;}#mermaid-svg-78APA6G4Ejz3tEqS .icon-shape,#mermaid-svg-78APA6G4Ejz3tEqS .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-78APA6G4Ejz3tEqS .icon-shape p,#mermaid-svg-78APA6G4Ejz3tEqS .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-78APA6G4Ejz3tEqS .icon-shape .label rect,#mermaid-svg-78APA6G4Ejz3tEqS .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-78APA6G4Ejz3tEqS .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-78APA6G4Ejz3tEqS .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-78APA6G4Ejz3tEqS :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 环境准备
数据加载
模型搭建
训练循环
结果评估
保存模型
可视化展示

2. 环境准备:先把「厨房」搭好

在开始写代码之前,需要先准备好开发环境。就像做饭前要先准备好锅碗瓢盆一样。

2.1 安装 PyTorch

打开终端(Terminal),输入以下命令安装 PyTorch:

bash 复制代码
pip install torch torchvision

小贴士 :如果有 NVIDIA 显卡,想用 GPU 加速训练,可以到 PyTorch 官网 选择对应的 CUDA 版本安装命令。不过对于 MNIST 这个任务,CPU 也完全够用,训练时间通常只要几分钟。

2.2 验证安装

安装完成后,在 Python 环境中运行以下代码,确认 PyTorch 能正常导入:

python 复制代码
import torch
print("PyTorch 版本:", torch.__version__)
print("CUDA 是否可用:", torch.cuda.is_available())

如果能看到版本号输出,说明环境已经准备好了!

3. 数据准备:认识我们的「食材」

在训练模型之前,首先要拿到数据。PyTorch 的 torchvision 库帮我们封装好了 MNIST 数据集的下载和加载,非常方便。

3.1 加载数据集

python 复制代码
import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

# 定义数据预处理:把 PIL 图片转成 Tensor,并归一化到 [0, 1]
transform = transforms.Compose([
    transforms.ToTensor(),               # 将图片从 [0, 255] 转为 [0, 1] 的 Tensor
    transforms.Normalize((0.1307,), (0.3081,))  # 用 MNIST 的均值和标准差做标准化
])

# 下载并加载训练集
train_dataset = datasets.MNIST(
    root='./data',       # 数据保存路径
    train=True,          # 加载训练集
    transform=transform, # 应用预处理
    download=True        # 如果本地没有就自动下载
)

# 下载并加载测试集
test_dataset = datasets.MNIST(
    root='./data',
    train=False,         # 加载测试集
    transform=transform,
    download=True
)

print(f"训练集大小: {len(train_dataset)}")
print(f"测试集大小: {len(test_dataset)}")

3.2 数据加载器(DataLoader)

DataLoader 是 PyTorch 提供的数据加载工具,它会把数据集自动分成一个个小批次(batch),并在训练时自动打乱顺序,避免模型学到数据的排列规律。

python 复制代码
batch_size = 64

train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)

print(f"训练集批次数量: {len(train_loader)}")
print(f"测试集批次数量: {len(test_loader)}")

为什么用 batch? 如果一次性把所有 6 万张图片都喂给模型,内存会爆掉,而且计算效率很低。分批训练(Mini-batch Gradient Descent)是深度学习中最常用的折中方案:每次用一小批数据计算梯度并更新参数,既省内存,又能让训练过程更稳定。

3.3 看一眼数据长什么样

在训练之前,不妨先看看数据长什么样,做到「心中有数」:

python 复制代码
import matplotlib.pyplot as plt

# 取一个批次的数据
images, labels = next(iter(train_loader))
print(f"一个批次图片的维度: {images.shape}")  # [64, 1, 28, 28]
print(f"一个批次标签的维度: {labels.shape}")  # [64]

# 显示前 6 张图片
fig, axes = plt.subplots(1, 6, figsize=(12, 3))
for i in range(6):
    axes[i].imshow(images[i].squeeze(), cmap='gray')
    axes[i].set_title(f"标签: {labels[i].item()}")
    axes[i].axis('off')
plt.show()

运行这段代码,会看到 6 张手写数字图片和它们对应的标签。这就是模型的「学习素材」。

4. 搭建模型:设计我们的「大脑」

现在到了最核心的部分:搭建神经网络模型。使用 torch.nn 模块来定义网络结构。

4.1 网络结构设计

全连接网络包含三层:

输入维度 输出维度 激活函数
输入层(展平) 28×28=784 784 -
隐藏层 1 784 128 ReLU
隐藏层 2 128 64 ReLU
输出层 64 10 -

设计思路

  • 输入层:每张图片是 28×28 的二维矩阵,需要把它「展平」成一维的 784 维向量,才能输入全连接层。
  • 隐藏层:中间的两层负责提取特征。第一层把 784 维压缩到 128 维,第二层再压缩到 64 维。维度逐层递减,强迫模型学习最重要的特征。
  • 输出层:最终输出 10 个数值,分别对应数字 0~9 的「得分」。得分最高的那个数字就是模型的预测结果。
  • 激活函数 :ReLU(Rectified Linear Unit)是最常用的激活函数,公式是 f(x) = max(0, x)。它给网络引入非线性,让模型能学习更复杂的模式。

4.2 用 PyTorch 实现

python 复制代码
import torch.nn as nn
import torch.nn.functional as F

class SimpleNN(nn.Module):
    def __init__(self):
        super(SimpleNN, self).__init__()
        # 定义三个全连接层
        self.fc1 = nn.Linear(28 * 28, 128)  # 输入层 -> 隐藏层1
        self.fc2 = nn.Linear(128, 64)       # 隐藏层1 -> 隐藏层2
        self.fc3 = nn.Linear(64, 10)        # 隐藏层2 -> 输出层

    def forward(self, x):
        # 前向传播:数据从输入到输出的流动过程
        x = x.view(-1, 28 * 28)   # 展平:把 [batch, 1, 28, 28] 变成 [batch, 784]
        x = F.relu(self.fc1(x))   # 第一层 + ReLU 激活
        x = F.relu(self.fc2(x))   # 第二层 + ReLU 激活
        x = self.fc3(x)           # 输出层(不加激活,后面用交叉熵损失)
        return x

# 创建模型实例
model = SimpleNN()
print(model)

运行后会输出模型的结构信息,可以清楚地看到每一层的参数配置。

为什么输出层不加激活函数? 因为后面要使用 CrossEntropyLoss(交叉熵损失),它内部已经包含了 Softmax 操作。如果在输出层提前做了 Softmax,反而会导致计算不稳定。

下面是这个全连接网络的数据流向图:
#mermaid-svg-kbAmXXHA6QPzyaZ0{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .error-icon{fill:#552222;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .marker.cross{stroke:#333333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 p{margin:0;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .cluster-label text{fill:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .cluster-label span{color:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .cluster-label span p{background-color:transparent;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .label text,#mermaid-svg-kbAmXXHA6QPzyaZ0 span{fill:#333;color:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .node rect,#mermaid-svg-kbAmXXHA6QPzyaZ0 .node circle,#mermaid-svg-kbAmXXHA6QPzyaZ0 .node ellipse,#mermaid-svg-kbAmXXHA6QPzyaZ0 .node polygon,#mermaid-svg-kbAmXXHA6QPzyaZ0 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .rough-node .label text,#mermaid-svg-kbAmXXHA6QPzyaZ0 .node .label text,#mermaid-svg-kbAmXXHA6QPzyaZ0 .image-shape .label,#mermaid-svg-kbAmXXHA6QPzyaZ0 .icon-shape .label{text-anchor:middle;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .rough-node .label,#mermaid-svg-kbAmXXHA6QPzyaZ0 .node .label,#mermaid-svg-kbAmXXHA6QPzyaZ0 .image-shape .label,#mermaid-svg-kbAmXXHA6QPzyaZ0 .icon-shape .label{text-align:center;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .node.clickable{cursor:pointer;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .arrowheadPath{fill:#333333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-kbAmXXHA6QPzyaZ0 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-kbAmXXHA6QPzyaZ0 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-kbAmXXHA6QPzyaZ0 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .cluster text{fill:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .cluster span{color:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-kbAmXXHA6QPzyaZ0 rect.text{fill:none;stroke-width:0;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .icon-shape,#mermaid-svg-kbAmXXHA6QPzyaZ0 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .icon-shape p,#mermaid-svg-kbAmXXHA6QPzyaZ0 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .icon-shape .label rect,#mermaid-svg-kbAmXXHA6QPzyaZ0 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-kbAmXXHA6QPzyaZ0 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-kbAmXXHA6QPzyaZ0 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-kbAmXXHA6QPzyaZ0 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 输入图片

28×28
展平

784 维
全连接层 1

784 → 128
ReLU 激活
全连接层 2

128 → 64
ReLU 激活
输出层

64 → 10
Softmax

分类结果

5. 训练准备:定好「游戏规则」

模型搭好了,接下来要定义「怎么学」------也就是损失函数和优化器。

5.1 损失函数(Loss Function)

损失函数用来衡量模型的预测结果和真实标签之间的差距。使用交叉熵损失(Cross Entropy Loss),它是分类任务中最常用的损失函数:

python 复制代码
criterion = nn.CrossEntropyLoss()

通俗理解:交叉熵损失会「惩罚」模型------当模型预测错误时,损失值会很大;预测正确时,损失值会很小。训练的目标就是不断降低这个损失值。

5.2 优化器(Optimizer)

优化器负责根据损失值更新模型的参数。使用随机梯度下降(SGD)优化器:

python 复制代码
import torch.optim as optim

learning_rate = 0.01
optimizer = optim.SGD(model.parameters(), lr=learning_rate)

通俗理解 :优化器就像是一个「导航员」,它根据损失值告诉模型参数该往哪个方向调整、调整多少。learning_rate(学习率)控制每次调整的步长------太大容易「走过头」,太小则「走得太慢」。

5.3 设备选择:用 CPU 还是 GPU?

python 复制代码
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)
print(f"使用设备: {device}")

这段代码会自动检测是否有 GPU,如果有就用 GPU 加速,否则用 CPU。

6. 训练循环:让模型「学起来」

万事俱备,现在开始训练!训练过程的核心是一个循环:前向传播 → 计算损失 → 反向传播 → 更新参数

6.1 训练一个 epoch 的代码

python 复制代码
def train_one_epoch(model, train_loader, criterion, optimizer, device):
    model.train()  # 切换到训练模式
    total_loss = 0
    correct = 0
    total = 0

    for images, labels in train_loader:
        # 把数据放到指定设备(CPU/GPU)
        images, labels = images.to(device), labels.to(device)

        # 1. 前向传播:把图片输入模型,得到预测结果
        outputs = model(images)

        # 2. 计算损失:比较预测结果和真实标签
        loss = criterion(outputs, labels)

        # 3. 反向传播:计算梯度
        optimizer.zero_grad()  # 先把之前的梯度清零,否则会累加
        loss.backward()        # 反向传播,计算每个参数的梯度

        # 4. 更新参数:沿着梯度下降的方向调整参数
        optimizer.step()

        # 统计损失和准确率
        total_loss += loss.item()
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

    avg_loss = total_loss / len(train_loader)
    accuracy = 100 * correct / total
    return avg_loss, accuracy

6.2 逐行解读训练过程

这段代码是整个训练的核心,我们逐行拆解:

  1. model.train():切换到训练模式。有些层(如 Dropout、BatchNorm)在训练和测试时的行为不同,PyTorch 通过这个开关来区分。
  2. optimizer.zero_grad():每次更新参数前,必须把上一次计算的梯度清零。否则 PyTorch 会默认累加梯度,导致参数更新方向错误。
  3. loss.backward():反向传播的核心。它根据损失值,自动计算每个参数对损失的「贡献」(即梯度)。
  4. optimizer.step():根据计算好的梯度,沿着梯度下降的方向更新所有参数。

训练的本质:重复以上四步成千上万次,模型的参数会逐渐调整到「能正确识别手写数字」的状态。

6.3 评估函数

训练完每个 epoch 后,我们需要在测试集上评估模型的真实表现:

python 复制代码
def evaluate(model, test_loader, device):
    model.eval()  # 切换到评估模式
    correct = 0
    total = 0

    with torch.no_grad():  # 评估时不需要计算梯度,节省内存
        for images, labels in test_loader:
            images, labels = images.to(device), labels.to(device)
            outputs = model(images)
            _, predicted = torch.max(outputs.data, 1)
            total += labels.size(0)
            correct += (predicted == labels).sum().item()

    accuracy = 100 * correct / total
    return accuracy

torch.no_grad() 是什么? 在评估阶段,我们不需要更新参数,也就不需要计算梯度。用 torch.no_grad() 包裹代码可以关闭梯度计算,既节省内存又加快速度。

6.4 完整训练循环

python 复制代码
num_epochs = 5

for epoch in range(1, num_epochs + 1):
    train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
    test_acc = evaluate(model, test_loader, device)

    print(f"Epoch [{epoch}/{num_epochs}] "
          f"训练损失: {train_loss:.4f} | "
          f"训练准确率: {train_acc:.2f}% | "
          f"测试准确率: {test_acc:.2f}%")

运行这段代码,你会看到类似下面的输出:

复制代码
Epoch [1/5] 训练损失: 0.3124 | 训练准确率: 90.12% | 测试准确率: 93.45%
Epoch [2/5] 训练损失: 0.1567 | 训练准确率: 95.23% | 测试准确率: 95.89%
Epoch [3/5] 训练损失: 0.1123 | 训练准确率: 96.45% | 测试准确率: 96.52%
Epoch [4/5] 训练损失: 0.0876 | 训练准确率: 97.12% | 测试准确率: 96.98%
Epoch [5/5] 训练损失: 0.0712 | 训练准确率: 97.56% | 测试准确率: 97.23%

可以看到,从第二个 epoch 开始,测试准确率就已经超过 95% 了!到第五个 epoch,准确率稳定在 97% 左右。


下面是训练循环的完整流程:
#mermaid-svg-HiJC3ioZeDO41fWZ{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-HiJC3ioZeDO41fWZ .error-icon{fill:#552222;}#mermaid-svg-HiJC3ioZeDO41fWZ .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-HiJC3ioZeDO41fWZ .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-HiJC3ioZeDO41fWZ .marker{fill:#333333;stroke:#333333;}#mermaid-svg-HiJC3ioZeDO41fWZ .marker.cross{stroke:#333333;}#mermaid-svg-HiJC3ioZeDO41fWZ svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-HiJC3ioZeDO41fWZ p{margin:0;}#mermaid-svg-HiJC3ioZeDO41fWZ .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ .cluster-label text{fill:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ .cluster-label span{color:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ .cluster-label span p{background-color:transparent;}#mermaid-svg-HiJC3ioZeDO41fWZ .label text,#mermaid-svg-HiJC3ioZeDO41fWZ span{fill:#333;color:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ .node rect,#mermaid-svg-HiJC3ioZeDO41fWZ .node circle,#mermaid-svg-HiJC3ioZeDO41fWZ .node ellipse,#mermaid-svg-HiJC3ioZeDO41fWZ .node polygon,#mermaid-svg-HiJC3ioZeDO41fWZ .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-HiJC3ioZeDO41fWZ .rough-node .label text,#mermaid-svg-HiJC3ioZeDO41fWZ .node .label text,#mermaid-svg-HiJC3ioZeDO41fWZ .image-shape .label,#mermaid-svg-HiJC3ioZeDO41fWZ .icon-shape .label{text-anchor:middle;}#mermaid-svg-HiJC3ioZeDO41fWZ .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-HiJC3ioZeDO41fWZ .rough-node .label,#mermaid-svg-HiJC3ioZeDO41fWZ .node .label,#mermaid-svg-HiJC3ioZeDO41fWZ .image-shape .label,#mermaid-svg-HiJC3ioZeDO41fWZ .icon-shape .label{text-align:center;}#mermaid-svg-HiJC3ioZeDO41fWZ .node.clickable{cursor:pointer;}#mermaid-svg-HiJC3ioZeDO41fWZ .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-HiJC3ioZeDO41fWZ .arrowheadPath{fill:#333333;}#mermaid-svg-HiJC3ioZeDO41fWZ .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-HiJC3ioZeDO41fWZ .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-HiJC3ioZeDO41fWZ .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-HiJC3ioZeDO41fWZ .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-HiJC3ioZeDO41fWZ .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-HiJC3ioZeDO41fWZ .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-HiJC3ioZeDO41fWZ .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-HiJC3ioZeDO41fWZ .cluster text{fill:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ .cluster span{color:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-HiJC3ioZeDO41fWZ .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-HiJC3ioZeDO41fWZ rect.text{fill:none;stroke-width:0;}#mermaid-svg-HiJC3ioZeDO41fWZ .icon-shape,#mermaid-svg-HiJC3ioZeDO41fWZ .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-HiJC3ioZeDO41fWZ .icon-shape p,#mermaid-svg-HiJC3ioZeDO41fWZ .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-HiJC3ioZeDO41fWZ .icon-shape .label rect,#mermaid-svg-HiJC3ioZeDO41fWZ .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-HiJC3ioZeDO41fWZ .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-HiJC3ioZeDO41fWZ .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-HiJC3ioZeDO41fWZ :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 是



开始训练
读取一个 batch 数据
前向传播

计算预测结果
计算损失

CrossEntropyLoss
反向传播

计算梯度
更新参数

optimizer.step
还有数据?
评估测试集
达到 epoch 数?
训练完成

7. 结果可视化:看看模型学得怎么样

训练完成后,我们来看看模型的实际预测效果。

7.1 随机展示预测结果

python 复制代码
import matplotlib.pyplot as plt
import numpy as np

def show_predictions(model, test_loader, device, num_images=10):
    model.eval()
    images, labels = next(iter(test_loader))
    images, labels = images.to(device), labels.to(device)

    with torch.no_grad():
        outputs = model(images)
        _, predicted = torch.max(outputs, 1)

    # 显示前 num_images 张图片及其预测结果
    fig, axes = plt.subplots(2, 5, figsize=(12, 5))
    axes = axes.flatten()
    for i in range(num_images):
        img = images[i].cpu().squeeze().numpy()
        axes[i].imshow(img, cmap='gray')
        color = 'green' if predicted[i] == labels[i] else 'red'
        axes[i].set_title(f"预测: {predicted[i].item()} | 真实: {labels[i].item()}", color=color)
        axes[i].axis('off')
    plt.tight_layout()
    plt.show()

show_predictions(model, test_loader, device)

绿色标题表示预测正确,红色表示预测错误。你会看到绝大多数图片都被正确识别了。

7.2 绘制训练曲线

python 复制代码
# 记录训练过程中的数据
train_losses = []
train_accs = []
test_accs = []

for epoch in range(1, num_epochs + 1):
    train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
    test_acc = evaluate(model, test_loader, device)

    train_losses.append(train_loss)
    train_accs.append(train_acc)
    test_accs.append(test_acc)

# 绘制损失曲线
plt.figure(figsize=(12, 4))

plt.subplot(1, 2, 1)
plt.plot(range(1, num_epochs + 1), train_losses, marker='o')
plt.xlabel('Epoch')
plt.ylabel('训练损失')
plt.title('训练损失曲线')
plt.grid(True)

# 绘制准确率曲线
plt.subplot(1, 2, 2)
plt.plot(range(1, num_epochs + 1), train_accs, marker='o', label='训练准确率')
plt.plot(range(1, num_epochs + 1), test_accs, marker='s', label='测试准确率')
plt.xlabel('Epoch')
plt.ylabel('准确率 (%)')
plt.title('准确率曲线')
plt.legend()
plt.grid(True)

plt.tight_layout()
plt.show()

通过曲线图,你可以直观地看到:随着训练的进行,损失不断下降,准确率不断上升,模型在「越学越好」。


8. 保存与加载模型:把「学习成果」存下来

训练好的模型不能每次用都重新训练一遍,我们需要把它保存到磁盘上。

8.1 保存模型

python 复制代码
# 保存整个模型(包括结构和参数)
torch.save(model, 'mnist_model_full.pth')

# 推荐:只保存模型参数(更轻量,更灵活)
torch.save(model.state_dict(), 'mnist_model_weights.pth')
print("模型已保存!")

两种保存方式的区别 :保存整个模型文件更大,而且加载时对代码结构有依赖;只保存参数(state_dict)更轻量,加载时需要先定义好相同的模型结构。推荐使用第二种方式

8.2 加载模型

python 复制代码
# 方式一:加载整个模型
model_loaded = torch.load('mnist_model_full.pth')
model_loaded.eval()

# 方式二:先定义模型结构,再加载参数(推荐)
model = SimpleNN()
model.load_state_dict(torch.load('mnist_model_weights.pth'))
model.to(device)
model.eval()
print("模型加载成功!")

下面是模型保存与加载的两种方式对比:
#mermaid-svg-6ee8XpBZ8G1QQABW{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-6ee8XpBZ8G1QQABW .error-icon{fill:#552222;}#mermaid-svg-6ee8XpBZ8G1QQABW .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-6ee8XpBZ8G1QQABW .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-6ee8XpBZ8G1QQABW .marker{fill:#333333;stroke:#333333;}#mermaid-svg-6ee8XpBZ8G1QQABW .marker.cross{stroke:#333333;}#mermaid-svg-6ee8XpBZ8G1QQABW svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-6ee8XpBZ8G1QQABW p{margin:0;}#mermaid-svg-6ee8XpBZ8G1QQABW .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW .cluster-label text{fill:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW .cluster-label span{color:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW .cluster-label span p{background-color:transparent;}#mermaid-svg-6ee8XpBZ8G1QQABW .label text,#mermaid-svg-6ee8XpBZ8G1QQABW span{fill:#333;color:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW .node rect,#mermaid-svg-6ee8XpBZ8G1QQABW .node circle,#mermaid-svg-6ee8XpBZ8G1QQABW .node ellipse,#mermaid-svg-6ee8XpBZ8G1QQABW .node polygon,#mermaid-svg-6ee8XpBZ8G1QQABW .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-6ee8XpBZ8G1QQABW .rough-node .label text,#mermaid-svg-6ee8XpBZ8G1QQABW .node .label text,#mermaid-svg-6ee8XpBZ8G1QQABW .image-shape .label,#mermaid-svg-6ee8XpBZ8G1QQABW .icon-shape .label{text-anchor:middle;}#mermaid-svg-6ee8XpBZ8G1QQABW .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-6ee8XpBZ8G1QQABW .rough-node .label,#mermaid-svg-6ee8XpBZ8G1QQABW .node .label,#mermaid-svg-6ee8XpBZ8G1QQABW .image-shape .label,#mermaid-svg-6ee8XpBZ8G1QQABW .icon-shape .label{text-align:center;}#mermaid-svg-6ee8XpBZ8G1QQABW .node.clickable{cursor:pointer;}#mermaid-svg-6ee8XpBZ8G1QQABW .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-6ee8XpBZ8G1QQABW .arrowheadPath{fill:#333333;}#mermaid-svg-6ee8XpBZ8G1QQABW .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-6ee8XpBZ8G1QQABW .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-6ee8XpBZ8G1QQABW .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-6ee8XpBZ8G1QQABW .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-6ee8XpBZ8G1QQABW .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-6ee8XpBZ8G1QQABW .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-6ee8XpBZ8G1QQABW .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-6ee8XpBZ8G1QQABW .cluster text{fill:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW .cluster span{color:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-6ee8XpBZ8G1QQABW .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-6ee8XpBZ8G1QQABW rect.text{fill:none;stroke-width:0;}#mermaid-svg-6ee8XpBZ8G1QQABW .icon-shape,#mermaid-svg-6ee8XpBZ8G1QQABW .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-6ee8XpBZ8G1QQABW .icon-shape p,#mermaid-svg-6ee8XpBZ8G1QQABW .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-6ee8XpBZ8G1QQABW .icon-shape .label rect,#mermaid-svg-6ee8XpBZ8G1QQABW .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-6ee8XpBZ8G1QQABW .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-6ee8XpBZ8G1QQABW .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-6ee8XpBZ8G1QQABW :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 加载方式
保存方式
方式一

保存整个模型
文件较大

依赖代码结构
方式二

只保存参数 state_dict
文件轻量

灵活通用
直接加载整个模型
torch.load
先定义模型结构
再加载参数

9. 完整代码汇总

为了方便你直接运行,这里把完整代码整合在一起:

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

# ========== 1. 数据准备 ==========
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))
])

train_dataset = datasets.MNIST(root='./data', train=True, transform=transform, download=True)
test_dataset = datasets.MNIST(root='./data', train=False, transform=transform, download=True)

batch_size = 64
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)

# ========== 2. 定义模型 ==========
class SimpleNN(nn.Module):
    def __init__(self):
        super(SimpleNN, self).__init__()
        self.fc1 = nn.Linear(28 * 28, 128)
        self.fc2 = nn.Linear(128, 64)
        self.fc3 = nn.Linear(64, 10)

    def forward(self, x):
        x = x.view(-1, 28 * 28)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

# ========== 3. 训练准备 ==========
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = SimpleNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.01)

# ========== 4. 训练与评估函数 ==========
def train_one_epoch(model, train_loader, criterion, optimizer, device):
    model.train()
    total_loss = 0
    correct = 0
    total = 0

    for images, labels in train_loader:
        images, labels = images.to(device), labels.to(device)

        outputs = model(images)
        loss = criterion(outputs, labels)

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

        total_loss += loss.item()
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

    avg_loss = total_loss / len(train_loader)
    accuracy = 100 * correct / total
    return avg_loss, accuracy

def evaluate(model, test_loader, device):
    model.eval()
    correct = 0
    total = 0

    with torch.no_grad():
        for images, labels in test_loader:
            images, labels = images.to(device), labels.to(device)
            outputs = model(images)
            _, predicted = torch.max(outputs.data, 1)
            total += labels.size(0)
            correct += (predicted == labels).sum().item()

    return 100 * correct / total

# ========== 5. 训练循环 ==========
num_epochs = 5
for epoch in range(1, num_epochs + 1):
    train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
    test_acc = evaluate(model, test_loader, device)
    print(f"Epoch [{epoch}/{num_epochs}] "
          f"训练损失: {train_loss:.4f} | "
          f"训练准确率: {train_acc:.2f}% | "
          f"测试准确率: {test_acc:.2f}%")

# ========== 6. 保存模型 ==========
torch.save(model.state_dict(), 'mnist_model_weights.pth')
print("训练完成,模型已保存!")

10. 进阶优化:如何进一步提升准确率?

我们已经达到了 95% 的目标,但如果你想让模型表现更好,可以尝试以下方法:

10.1 调整超参数

方法 说明 预期效果
增加训练轮数 从 5 个 epoch 增加到 10~20 个 准确率可提升到 98% 左右
调整学习率 尝试 0.001、0.005 等更小的值 训练更稳定,但需要更多轮数
增大隐藏层 把 128 改成 256 或 512 模型容量更大,但注意过拟合
使用 Adam 优化器 optim.Adam(model.parameters()) 收敛更快,对学习率不敏感

10.2 添加正则化

python 复制代码
# 在模型中加入 Dropout 层
class SimpleNNWithDropout(nn.Module):
    def __init__(self):
        super(SimpleNNWithDropout, self).__init__()
        self.fc1 = nn.Linear(28 * 28, 128)
        self.fc2 = nn.Linear(128, 64)
        self.fc3 = nn.Linear(64, 10)
        self.dropout = nn.Dropout(0.2)  # 随机丢弃 20% 的神经元

    def forward(self, x):
        x = x.view(-1, 28 * 28)
        x = F.relu(self.fc1(x))
        x = self.dropout(x)
        x = F.relu(self.fc2(x))
        x = self.dropout(x)
        x = self.fc3(x)
        return x

Dropout 在训练时随机「关闭」一部分神经元,迫使网络学习更鲁棒的特征,能有效防止过拟合。

10.3 尝试卷积神经网络(CNN)

如果想把准确率推到 99% 以上,可以尝试 CNN。CNN 能更好地利用图片的二维空间结构,是图像识别的主流方案。不过那是下一个进阶话题了。


11. 总结与思考

恭喜你!通过这篇文章,你已经完成了第一个完整的 PyTorch 深度学习项目。让我们回顾一下学到的核心内容:

  1. 数据准备 :使用 torchvision.datasets 加载 MNIST 数据集,用 DataLoader 分批加载数据。
  2. 模型搭建 :用 nn.Module 定义全连接网络,理解前向传播的过程。
  3. 训练流程:掌握「前向传播 → 计算损失 → 反向传播 → 更新参数」的完整循环。
  4. 模型评估:在测试集上评估模型的真实表现。
  5. 模型保存与加载:把训练好的模型持久化,方便后续使用。

思考题

  • 如果把隐藏层的神经元数量从 128 改成 32,准确率会有什么变化?为什么?
  • 如果去掉 ReLU 激活函数,模型还能正常工作吗?
  • 如果学习率设置得过大(比如 1.0),训练过程会发生什么?

下一步学习建议

  • 尝试用 CNN(卷积神经网络)替代全连接网络,挑战 99% 以上的准确率。
  • 学习 PyTorch 的 DatasetDataLoader 自定义数据加载方式。
  • 尝试用 Matplotlib 可视化模型的权重,看看网络到底「学」到了什么。

深度学习的世界才刚刚开始,希望这篇文章能成为你探索之路的坚实起点。动手实践是最好的学习方式,快去运行代码,亲眼见证模型从「一无所知」到「准确识别手写数字」的神奇过程吧!

附录:完整的代码

python 复制代码
# -*- coding: utf-8 -*-
"""
第二章练习 1:用最简单的神经网络(MLP)在 MNIST 上训练
=====================================================
练习要求:
  1. 用 PyTorch 实现一个 3 层 MLP(Linear(784,128) → ReLU → Linear(128,64) → ReLU → Linear(64,10))
  2. 在 MNIST 数据集上训练,达到 95%+ 的测试准确率
  3. 用交叉熵损失 + Adam 优化器
  4. 观察训练 loss 曲线和测试准确率

本脚本实现了完整的训练流程:
  - 数据下载/预处理(MNIST)
  - 模型定义(3 层 MLP)
  - 训练循环(前向 → 损失 → 反向传播 → 参数更新)
  - 测试评估
  - 训练曲线可视化 + 预测样例可视化
  - 模型保存与加载

运行方式(在项目根目录 vla-chapter-02-exercises 下):
  .venv\\Scripts\\python code\\train_mnist.py
"""
import os
import time
import argparse

import matplotlib
matplotlib.use("Agg")  # 无界面环境下保存图片
import matplotlib.pyplot as plt

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

# ---------- 路径配置 ----------
BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
DATA_DIR = os.path.join(BASE_DIR, "data")
OUT_DIR = os.path.join(BASE_DIR, "outputs")
os.makedirs(DATA_DIR, exist_ok=True)
os.makedirs(OUT_DIR, exist_ok=True)


# ---------- 1. 模型定义:最简单的 3 层 MLP ----------
class SimpleMLP(nn.Module):
    """
    3 层全连接网络(MLP)。
    输入: 784 维 (28*28 像素展平)
    隐藏层1: 128 神经元 + ReLU
    隐藏层2: 64 神经元 + ReLU
    输出层: 10 神经元(对应 0-9 十个数字),不接激活,交给 CrossEntropyLoss
    """

    def __init__(self, input_dim=784, hidden1=128, hidden2=64, num_classes=10):
        super().__init__()
        self.fc1 = nn.Linear(input_dim, hidden1)
        self.fc2 = nn.Linear(hidden1, hidden2)
        self.fc3 = nn.Linear(hidden2, num_classes)

    def forward(self, x):
        # 展平: (batch, 1, 28, 28) -> (batch, 784)
        x = x.view(x.size(0), -1)
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        x = self.fc3(x)  # 输出 10 个 logits
        return x


# ---------- 2. 数据准备 ----------
def get_dataloaders(batch_size=128):
    """下载 MNIST 并构造 DataLoader(含归一化)"""
    transform = transforms.Compose([
        transforms.ToTensor(),                 # 0~255 -> 0~1,并转为张量
        transforms.Normalize((0.1307,), (0.3081,)),  # 用 MNIST 官方均值和标准差归一化
    ])
    train_set = datasets.MNIST(DATA_DIR, train=True, download=True, transform=transform)
    test_set = datasets.MNIST(DATA_DIR, train=False, download=True, transform=transform)

    train_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True, num_workers=0)
    test_loader = DataLoader(test_set, batch_size=batch_size, shuffle=False, num_workers=0)
    return train_loader, test_loader


# ---------- 3. 训练一个 epoch ----------
def train_one_epoch(model, loader, criterion, optimizer, device):
    model.train()
    total_loss, correct, total = 0.0, 0, 0
    for images, labels in loader:
        images, labels = images.to(device), labels.to(device)

        # --- 前向传播 ---
        outputs = model(images)
        loss = criterion(outputs, labels)

        # --- 反向传播 ---
        optimizer.zero_grad()   # 清空上一轮的梯度
        loss.backward()         # 自动求梯度(链式法则)
        optimizer.step()        # 参数更新(梯度下降)

        # --- 统计 ---
        total_loss += loss.item() * images.size(0)
        _, preds = outputs.max(1)
        correct += (preds == labels).sum().item()
        total += labels.size(0)

    return total_loss / total, correct / total


# ---------- 4. 测试评估 ----------
def evaluate(model, loader, criterion, device):
    model.eval()
    total_loss, correct, total = 0.0, 0, 0
    with torch.no_grad():  # 测试阶段不计算梯度,省显存加速
        for images, labels in loader:
            images, labels = images.to(device), labels.to(device)
            outputs = model(images)
            loss = criterion(outputs, labels)
            total_loss += loss.item() * images.size(0)
            _, preds = outputs.max(1)
            correct += (preds == labels).sum().item()
            total += labels.size(0)
    return total_loss / total, correct / total


# ---------- 5. 可视化 ----------
def plot_curves(train_losses, train_accs, test_losses, test_accs):
    epochs = range(1, len(train_losses) + 1)
    fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))

    axes[0].plot(epochs, train_losses, "o-", label="Train Loss", color="#2563eb")
    axes[0].plot(epochs, test_losses, "s-", label="Test Loss", color="#7c3aed")
    axes[0].set_title("Loss Curve")
    axes[0].set_xlabel("Epoch"); axes[0].set_ylabel("Loss")
    axes[0].legend(); axes[0].grid(alpha=0.3)

    axes[1].plot(epochs, [a * 100 for a in train_accs], "o-", label="Train Acc", color="#059669")
    axes[1].plot(epochs, [a * 100 for a in test_accs], "s-", label="Test Acc", color="#d97706")
    axes[1].set_title("Accuracy Curve")
    axes[1].set_xlabel("Epoch"); axes[1].set_ylabel("Accuracy (%)")
    axes[1].legend(); axes[1].grid(alpha=0.3)
    axes[1].axhline(95, color="red", linestyle="--", alpha=0.6, label="95% target")

    plt.tight_layout()
    path = os.path.join(OUT_DIR, "training_curves.png")
    plt.savefig(path, dpi=150, bbox_inches="tight")
    plt.close()
    print(f"[可视化] 训练曲线已保存 -> {path}")
    return path


def plot_predictions(model, test_set, device, num=16):
    """从测试集挑几张图展示真实标签与预测标签"""
    model.eval()
    fig, axes = plt.subplots(2, 8, figsize=(16, 4))
    axes = axes.flatten()
    shown = 0
    with torch.no_grad():
        idx = 0
        while shown < num and idx < len(test_set):
            img, label = test_set[idx]
            img_4d = img.unsqueeze(0).to(device)
            out = model(img_4d)
            pred = out.argmax(1).item()
            axes[shown].imshow(img.squeeze(), cmap="gray")
            color = "green" if pred == label else "red"
            axes[shown].set_title(f"True:{label} Pred:{pred}", color=color, fontsize=10)
            axes[shown].axis("off")
            shown += 1
            idx += 1
    plt.tight_layout()
    path = os.path.join(OUT_DIR, "predictions.png")
    plt.savefig(path, dpi=150, bbox_inches="tight")
    plt.close()
    print(f"[可视化] 预测样例已保存 -> {path}")
    return path


# ---------- 6. 主流程 ----------
def main():
    parser = argparse.ArgumentParser(description="第二章练习1:MNIST MLP 训练")
    parser.add_argument("--epochs", type=int, default=5)
    parser.add_argument("--batch-size", type=int, default=128)
    parser.add_argument("--lr", type=float, default=1e-3)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

    torch.manual_seed(args.seed)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"[环境] 设备: {device}")

    # 数据
    train_loader, test_loader = get_dataloaders(args.batch_size)
    print(f"[数据] 训练集: {len(train_loader.dataset)} 张, 测试集: {len(test_loader.dataset)} 张")

    # 模型 / 损失 / 优化器
    model = SimpleMLP().to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=args.lr)

    # 参数统计
    total_params = sum(p.numel() for p in model.parameters())
    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
    print(f"[模型] SimpleMLP 参数量: {total_params:,} (可训练 {trainable:,})")

    # 训练
    train_losses, train_accs, test_losses, test_accs = [], [], [], []
    print(f"\n开始训练 {args.epochs} 个 epoch ...\n")
    t0 = time.time()
    for epoch in range(1, args.epochs + 1):
        tr_loss, tr_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
        te_loss, te_acc = evaluate(model, test_loader, criterion, device)
        train_losses.append(tr_loss); train_accs.append(tr_acc)
        test_losses.append(te_loss); test_accs.append(te_acc)
        print(f"Epoch {epoch:02d}/{args.epochs:02d} | "
              f"Train Loss {tr_loss:.4f} Acc {tr_acc*100:.2f}% | "
              f"Test Loss {te_loss:.4f} Acc {te_acc*100:.2f}%")
    total_time = time.time() - t0
    print(f"\n训练完成,总耗时 {total_time:.1f}s")

    # 可视化
    plot_curves(train_losses, train_accs, test_losses, test_accs)
    plot_predictions(model, train_loader.dataset, device)

    # 保存模型
    ckpt_path = os.path.join(OUT_DIR, "mnist_mlp.pt")
    torch.save({
        "model_state_dict": model.state_dict(),
        "test_acc": test_accs[-1],
        "args": vars(args),
    }, ckpt_path)
    print(f"[保存] 模型已保存 -> {ckpt_path}")
    print(f"\n最终测试准确率: {test_accs[-1]*100:.2f}% "
          f"({'✅ 达到 95%+ 目标' if test_accs[-1] >= 0.95 else '⚠️ 未达 95%,可增加 epoch 或调参'})")


if __name__ == "__main__":
    main()
相关推荐
数字护盾(和中)1 小时前
和中科技剖析 EDR 绕过全链路,AMSI、ETW 规避技术与防御对策
运维·网络·人工智能·科技·安全·web安全
土星云SaturnCloud1 小时前
高速服务区AI视觉全场景方案:安全管控+运营提效+服务升级,土星云边缘算力赋能智慧交通
服务器·人工智能·ai·边缘计算
开源量化GO1 小时前
最新AI辅助量化表达:先理清规则,再按需求选工具
人工智能·python
天天代码码天天1 小时前
不用 ONNX Runtime,也不用 OpenCV:我用纯 C 做了一个 PP-OCR 专用推理 Runtime
人工智能
u1301301 小时前
AI 日报(2026年08月25日)
人工智能
2601_957879331 小时前
Seedance排队太久怎么办?2026免排队AI视频工具与替代方案怎么选
大数据·人工智能·深度学习
IvanCodes1 小时前
RAG 实战教程(四):GraphRAG 查询实战,本地检索、全局检索与 DRIFT Search
人工智能·agent
能源革命1 小时前
AI+能源前沿20260825
人工智能·能源
m4Rk_1 小时前
【论文阅读】Agent 记忆机制(50):PACE——按下一步预测价值动态分配历史记忆粒度
论文阅读·人工智能·学习·开源·github