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 逐行解读训练过程
这段代码是整个训练的核心,我们逐行拆解:
model.train():切换到训练模式。有些层(如 Dropout、BatchNorm)在训练和测试时的行为不同,PyTorch 通过这个开关来区分。optimizer.zero_grad():每次更新参数前,必须把上一次计算的梯度清零。否则 PyTorch 会默认累加梯度,导致参数更新方向错误。loss.backward():反向传播的核心。它根据损失值,自动计算每个参数对损失的「贡献」(即梯度)。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 深度学习项目。让我们回顾一下学到的核心内容:
- 数据准备 :使用
torchvision.datasets加载 MNIST 数据集,用DataLoader分批加载数据。 - 模型搭建 :用
nn.Module定义全连接网络,理解前向传播的过程。 - 训练流程:掌握「前向传播 → 计算损失 → 反向传播 → 更新参数」的完整循环。
- 模型评估:在测试集上评估模型的真实表现。
- 模型保存与加载:把训练好的模型持久化,方便后续使用。
思考题
- 如果把隐藏层的神经元数量从 128 改成 32,准确率会有什么变化?为什么?
- 如果去掉 ReLU 激活函数,模型还能正常工作吗?
- 如果学习率设置得过大(比如 1.0),训练过程会发生什么?
下一步学习建议
- 尝试用 CNN(卷积神经网络)替代全连接网络,挑战 99% 以上的准确率。
- 学习 PyTorch 的
Dataset和DataLoader自定义数据加载方式。 - 尝试用 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()