Appearance
实战项目:手写数字识别 —— 项目需求与设计
从这一篇开始,我们用 5 篇文章完成一个完整项目:训练一个能识别手写数字(0~9)的模型。这是深度学习界的 "Hello World",但麻雀虽小五脏俱全:数据管道、模型设计、训练循环、验证调参、模型导出,一个不缺。中途我们会故意引入两个新手必踩的错误,带你完整走一遍"报错 → 定位 → 修复"的排错流程。
项目需求
- 输入:28×28 灰度手写数字图片(MNIST 公开数据集,60000 张训练图 + 10000 张测试图)
- 输出:0~9 的数字类别 + 置信度
- 指标:测试集准确率 ≥ 97%(全连接网络的合理水平)
- 交付:能加载保存的模型权重,对任意一张图片给出预测的脚本
为什么选 MNIST
- 数据自带:
torchvision一行下载,免去数据清洗的干扰,聚焦框架本身。 - 难度适中:太简单(如拟合直线)体现不出网络价值,太难(如 ImageNet)CPU 跑不动。
- 社区基准海量:你的结果随时可以对标全网方案,便于自查。
技术选型
| 环节 | 选择 | 理由 |
|---|---|---|
| 模型 | 3 层全连接网络(MLP) | 先跑通全流程;进阶篇再换 CNN 对比提升 |
| 损失 | CrossEntropyLoss | 多分类标配,输入 logits + int64 标签 |
| 优化器 | AdamW, lr=1e-3 | 默认首选,几乎不用调 |
| 批大小 | 64 | CPU/GPU 都友好的起点 |
| 轮数 | 5 epoch | MLP 在 MNIST 上 5 轮即可到 96%+ |
最终项目结构(先睹为快)
text
mnist-pytorch/
├── data/ # MNIST 原始数据(自动下载,勿提交 git)
├── model.py # 模型定义
├── train.py # 训练入口:数据 + 循环 + 验证
├── predict.py # 推理入口:加载权重,预测单张图
├── requirements.txt # torch / torchvision
└── weights/best.pt # 训练产出的最优权重(保存/加载验证)开发节奏
按以下 4 篇顺序推进,每篇结尾都有可运行的代码和预期输出:
- 数据准备 —— 下载、变换、组批,先确认"进模型的东西长什么样"
- 模型设计 —— 搭网络,踩第 1 个坑:形状不匹配
- 训练循环 —— 训练 + 验证,踩第 2 个坑:梯度忘清零
- 评估与导出 —— 测试集验收、保存/加载模型、单图预测
建议:跟着敲,不要复制粘贴。"报错—修好"的肌肉记忆只有亲手踩过才长得住。