Skip to content

实战项目:手写数字识别 —— 项目需求与设计

从这一篇开始,我们用 5 篇文章完成一个完整项目:训练一个能识别手写数字(0~9)的模型。这是深度学习界的 "Hello World",但麻雀虽小五脏俱全:数据管道、模型设计、训练循环、验证调参、模型导出,一个不缺。中途我们会故意引入两个新手必踩的错误,带你完整走一遍"报错 → 定位 → 修复"的排错流程。

项目需求

  • 输入:28×28 灰度手写数字图片(MNIST 公开数据集,60000 张训练图 + 10000 张测试图)
  • 输出:0~9 的数字类别 + 置信度
  • 指标:测试集准确率 ≥ 97%(全连接网络的合理水平)
  • 交付:能加载保存的模型权重,对任意一张图片给出预测的脚本

为什么选 MNIST

  1. 数据自带torchvision 一行下载,免去数据清洗的干扰,聚焦框架本身。
  2. 难度适中:太简单(如拟合直线)体现不出网络价值,太难(如 ImageNet)CPU 跑不动。
  3. 社区基准海量:你的结果随时可以对标全网方案,便于自查。

技术选型

环节选择理由
模型3 层全连接网络(MLP)先跑通全流程;进阶篇再换 CNN 对比提升
损失CrossEntropyLoss多分类标配,输入 logits + int64 标签
优化器AdamW, lr=1e-3默认首选,几乎不用调
批大小64CPU/GPU 都友好的起点
轮数5 epochMLP 在 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. 模型设计 —— 搭网络,踩第 1 个坑:形状不匹配
  3. 训练循环 —— 训练 + 验证,踩第 2 个坑:梯度忘清零
  4. 评估与导出 —— 测试集验收、保存/加载模型、单图预测

建议:跟着敲,不要复制粘贴。"报错—修好"的肌肉记忆只有亲手踩过才长得住。