第 4 步 · 你的第一个视觉网络
MNIST 手写数字识别MNIST Handwritten Digit Recognition
用 CNN 在 6 万张小图上练手,感受卷积如何把像素变成「认数字」的能力。
13 分钟
阅读 + 实操
1 个
交互演示
中级
难度
MNIST 手写数字识别 · 交互演示
准确率0.99
迭代5000
批次64
训练仪表 — 每轮结束记录一个点(琥珀 = 训练 Loss·左轴,绿 = 测试准确率·右轴)
Epoch
0
0
训练 Loss
—
—
测试准确率
—
—
混淆矩阵(行 = 真实数字,列 = 预测数字 · 绿 = 认对,红 = 认错 · 每 2 轮刷新)
第一层权重模板 — 32 个神经元各自学到的 8×8「探测器」(蓝 = 正权重·想要笔画,红 = 负权重·想要空白)
手写画板 — 用鼠标 / 手指写一个数字,点「预测」让网络识别
降采样 8×8(网络真正"看到"的输入)
训练几轮后再来考考它
为什么学这步?
MNIST 是深度视觉的「Hello World」。它足够小能在笔记本跑通,又能完整体现 CNN 训练全流程。
📌 发生了什么
- 卷积提取局部笔画特征
- 池化压缩并增强不变性
- 全连接输出类别概率
⚠️ 常见陷阱
- 太小网络学不动
- 学习率不当训练震荡
- 不归一化像素收敛慢
✅ 本章小结
- CNN 是图像分类主力
- MNIST 验证端到端流程
- 准确率可作基线对照
📐 前向传播与交叉熵
MNIST 把前面学的拼成完整流程:前向传播算预测,交叉熵算损失,反向传播算梯度。
前向传播逐层计算激活,每层是线性变换加非线性激活,输入像素经多层映射到 10 维类别得分。
输出层用 softmax 把 logits 转成概率分布,再与独热标签 y 计算交叉熵损失。
交叉熵对 logits 的梯度恰好是 ŷ−y,简洁且数值稳定,配合 softmax 自然衔接反向传播。
🎛 批次大小对比
批次大小决定每次更新的样本数,影响梯度噪声、显存与收敛速度。
| 批次大小 | 表现 | 结果 |
|---|---|---|
| 16 | 梯度噪声大,正则效应强 | 训练慢但泛化好 |
| 64 | 噪声与效率均衡 | 推荐 |
| 256 | 梯度估计准,但单步慢 | 收敛平稳 |
💡 大 batch 需配大学习率并做预热(warmup);小 batch 显存友好、常带来更好泛化。
💻 MNIST 训练循环
mini-batch SGD 训练循环 Python 实现:
# MNIST 训练循环(mini-batch SGD)
for epoch in range(epochs):
shuffle(train_data) # 打乱样本
for i in range(0, n, batch_size):
batch = train_data[i:i + batch_size]
logits, loss = forward(batch) # 前向 + 交叉熵
grads = backward(loss) # 反向传播
update_params(grads, lr) # SGD 更新
acc = evaluate(test_data) # 每 epoch 评估
print('epoch', epoch, 'acc', acc)
📚 参考文献与延伸阅读
- LeCun et al., Gradient-based learning applied to document recognition, Proc. IEEE (1998) — MNIST 数据集与 LeNet
- Goodfellow, Bengio & Courville, Deep Learning (2016), §6.2 Back-Propagation + §8.1 Mini-Batch SGD
- THE MNIST DATABASE of handwritten digits — LeCun 官方 MNIST 页面
- Bottou, Curtis & Nocedal (2018), Optimization Methods for Large-Scale Machine Learning, SIAM Review — mini-batch 理论
📝 课后练习
检验你的理解——答对为止