news 2026/9/3 6:16:36

从MNIST到自定义数据集:CNN手写数字识别完整流程拆解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从MNIST到自定义数据集:CNN手写数字识别完整流程拆解

前阵子有个朋友问我一个很现实的问题:他在MNIST上把手写数字识别准确率刷到99%了,还在课程报告里贴了混淆矩阵,结果换了一批公司内部扫描的手写表单,模型直接崩了,很多数字连“像不像数字”都判断不出来。我问他,你是真的把数据通路吃透了,还是只把别人的网络结构搭起来跑了一遍?他想了很久说:入口是现成的,标签是现成的,预处理也是现成的,我好像只做了“运行”这件事。

这个回答其实很典型。很多人第一次接触基于卷积神经网络CNN的手写数字识别时,都是在MNIST上跑通,然后误以为自己掌握了整个流程。真正把数据集换掉,才发现事情没那么简单。这也是我觉得这个项目值得写一篇完整拆解的原因:它同时支持MNIST数据集和普通数据集,而且允许重新训练。单看标题,你可能以为这是又一个手写数字识别demo;拆开看,它更像是一条从公开数据集迁移到自定义数据集的完整路径。

读这类源码时,我最关心的问题不是“它用了什么网络结构”,而是“如果我现在手头有一批自己的图片,我能照着它的流程重新走一遍吗”。如果能,说明你已经具备把一个图像分类小项目跑完的能力;如果不能,说明你还没有碰到这个领域真正麻烦的部分。

1. 先想清楚:这个系统真正教你的是哪一层能力

1.1 源码之外,真正决定结果的是数据通路

拿到一份手写数字识别源码时,新手第一件事通常是找到 CNN 网络定义,看看有几层卷积、用了什么池化、最终接了几分类。这个动作没有错,但它会带来一个错觉:在网络结构里看懂卷积、池化和全连接,就等于学会了这个项目。

真正决定最终识别效果的,往往不是网络结构本身,而是数据通路。MNIST 数据集是从原始图片怎么变成 28x28 灰度图、像素值怎么归一化、标签怎么和图片一一对应、训练集和验证集怎么划分、类别数量是多少,这些环节看似基础,却直接决定训练能不能收敛、模型能不能泛化到新数据。

这个项目把“普通数据集”加进来之后,数据通路的问题就藏不住了。如果只有 MNIST,你可以用官方脚本直接下载数据,图片和标签都是处理好的。换成普通数据集时,你得自己准备文件夹结构、自己处理图片尺寸、自己检查标签是否对齐。而这恰恰是真实项目里每天都会遇到的事。

1.2 三层能力框架:数据层、模型层、工程层

我会把这类图像分类项目拆成三个层次,分别对应三套能力:

  • 数据层:图片加载、格式统一、归一化、数据增强、训练集/验证集/测试集划分、标签检查。
  • 模型层:网络结构设计、损失函数、优化器、学习率、批大小、训练轮数等参数调整。
  • 工程层:训练过程监控、模型保存与加载、验证评估、错误样本分析、迁移到新数据集、结果复现。

很多人拿到开源项目后,只碰了模型层,数据层和工程层都靠现成脚本绕过。于是,MNIST 上表现很好,换数据就崩。这篇文章后续所有讨论,都会围绕这三层展开。你可以先判断一下自己卡在哪一层,再看下面几节。

读源码时,你可以拿这三层去对照:先找数据加载脚本,再找网络定义,最后找训练保存与评估模块。如果某个模块找不到,说明项目可能只覆盖了其中一层,你需要自己补上。别小看这个对照过程,它比单纯看网络结构更能暴露你对整个系统的理解盲区。

读源码时,优先去追“数据怎么进模型、模型怎么存下来、换数据后怎么跑通”,而不是只看网络结构图。

2. 为什么MNIST跑通不等于能处理普通数据集

2.1 MNIST的“友好”,建立在很多隐藏假设上

MNIST 是入门深度学习最经典的数据集,但它的“友好”恰恰会掩盖很多工程细节。它由 70000 张 28x28 灰度手写数字图片组成,图片已经完成切割、缩放、居中、灰度归一化,每个数字都有明确的标签,类别从 0 到 9 均匀分布。

这些条件意味着什么?意味着你拿到手的数据是干净的、整齐的、可预测的。你不用处理图片尺寸不一致,不用纠结通道是灰度还是彩色,不用担心背景噪声,不用检查标签文件是否错位。你只需要把数据喂给网络,调几个训练参数,就能获得比较高的准确率。

问题是,真实应用里几乎不存在这样的数据。手写表单可能是扫描件,可能是手机拍照,可能是一张纸上多个数字混在一起,也可能数字变形严重、笔画断裂、背景有表格线或水印。数据一变,所有隐藏假设都会失效。

2.2 换成普通数据集,最先要补的是什么

从 MNIST 切到普通数据集,我认为最先要补的是下面四件事:

  1. 统一图片尺寸和通道。CNN 的输入层尺寸是固定的,普通数据集的图片可能是 224x224、1200x900、16 位 PNG,甚至是 RGB 彩色图。你必须先把所有图片缩放到同一个尺寸,并确定模型输入通道是 1 还是 3。
  2. 确认标签来源。要么按文件夹名生成标签,要么用一个 label 文件记录每张图片对应的类别。实际项目里,标签文件错位、漏行、空行的问题非常常见。
  3. 划分训练集、验证集、测试集。MNIST 已经划分好了,普通数据集需要你自己划分。划分时还要保证类别分布尽量均衡,避免某个类别全在训练集,另一个类别全在测试集。
  4. 检查数据分布。如果普通数据集中 0 到 9 的数量差异很大,训练时模型会对样本多的类别产生偏向。这一点在数字识别里容易被忽视,因为很多人默认手写数字一定是十个类别数量均衡的。

这些准备工作没有一个需要高深理论,但每一项都会直接影响训练结果。如果一个项目只支持 MNIST,你永远碰不到这些问题。

2.3 所以“支持普通数据集+重新训练”才是这个项目的关键

我一直觉得,判断一个入门级图像识别项目是“demo”还是“可用流程”,就看三件事:能不能换数据集、能不能重新训练、能不能保存模型后单独做预测。

这个项目同时满足这三条,所以它不是一道“跑通 MNIST 就结束”的演示题,而是一条可以复用到自定义数据的流程。重新训练意味着你必须自己决定数据加载方式、训练参数、验证逻辑和模型保存位置,这些动作看似繁琐,却是把深度学习能力转化为实际工程交付能力的关键一步。

3. 从数据集到CNN模型,一条完整可复现的落地流程

3.1 环境准备与目录约定

在开始之前,先确认环境。常见使用的 MATLAB 版本中,R2020a 之后对深度学习工具箱的支持相对稳定,训练和验证流程都比较完整。要注意,这里的版本信息只是我平时使用的参考,具体落地前要以你本机安装版本为准。

这个工具箱提供了数据存储、网络搭建、训练和评估等相关能力。如果是 GPU 训练,还需要确认本地是否有可用显卡,以及 GPU 相关的并行计算环境是否正常。

目录结构建议按下面这样约定:

project/ ├── data/ │ ├── mnist/ # MNIST 原始图片或 mat 格式数据 │ └── mydata/ # 自己的普通数据集 │ ├── train/ │ │ ├── 0/ │ │ ├── 1/ │ │ └── ... │ └── test/ │ ├── 0/ │ ├── 1/ │ └── ... ├── models/ # 保存训练好的网络 ├── scripts/ # 数据准备、训练、评估脚本 └── results/ # 混淆矩阵、错误样本等输出

这样组织的好处是,数据、模型、脚本、结果分开,后续重新训练时不会因为路径混乱而找不到训练结果。

3.2 用imageDatastore加载数据

在常见实践里,MATLAB 处理图片分类一般会先用imageDatastore把文件夹里的图片读进来,自动把文件夹名作为标签。然后通过splitEachLabel划分训练集、验证集和测试集。代码结构类似下面:

imds = imageDatastore('data/mydata/train', ... 'IncludeSubfolders', true, ... 'LabelSource', 'foldernames'); [imdsTrain, imdsVal, imdsTest] = splitEachLabel(imds, ... 0.7, 0.15, 0.15, 'randomized'); disp(countEachLabel(imdsTrain));

countEachLabel这一步很重要。它能让你在训练前就看到每个类别有多少张图片,避免某个类别一张图都没有或严重失衡。

接下来用augmentedImageDatastore把图片统一缩放到网络输入尺寸。如果你的输入层是 28x28 单通道,可以写成:

inputSize = [28 28 1]; augTrain = augmentedImageDatastore(inputSize, imdsTrain); augVal = augmentedImageDatastore(inputSize, imdsVal); augTest = augmentedImageDatastore(inputSize, imdsTest);

这里要注意,augmentedImageDatastore输出的数据已经做了尺寸对齐,适合直接扔给trainNetwork使用。如果你读取的普通数据集是彩色 RGB 图,而网络输入要求单通道,需要提前把图片转换成灰度图,或者在输入层指定 3 个通道。

3.3 搭建一个能跑通的小型CNN

手写数字识别不需要一上来就堆 ResNet 这种重网络。一个小型 CNN 通常就能在 MNIST 上取得不错的效果,也能在普通数据集上验证流程是否完整。我给出的建议是先从一个简单的 3 层卷积网络开始:

layers = [ imageInputLayer([28 28 1]) convolution2dLayer(3, 16, 'Padding', 'same') batchNormalizationLayer reluLayer maxPooling2dLayer(2, 'Stride', 2) convolution2dLayer(3, 32, 'Padding', 'same') batchNormalizationLayer reluLayer maxPooling2dLayer(2, 'Stride', 2) fullyConnectedLayer(10) softmaxLayer classificationLayer ];

这个结构里,第一个卷积层用 3x3 卷积核,输出 16 个特征图;第二个卷积层输出 32 个特征图。中间插入batchNormalizationLayer是为了让网络在训练时更稳定,reluLayer负责引入非线性,maxPooling2dLayer降低特征图尺寸。最后接一个 10 分类的全连接层,因为手写数字识别固定是 0 到 9 十个类别。

如果是普通数据集,类别数量不一定是 10,那么fullyConnectedLayer的节点数就要改成实际类别数。这是换数据集后最容易出错的地方之一。

3.4 训练配置与参数理解

模型定义好后,用trainingOptions配置训练参数。常见参数如下:

参数常见取值作用
solver'adam''sgdm'优化器。adam 更容易适应不同数据集,sgdm 在数据量较大时也常用
InitialLearnRate1e-21e-4初始学习率,决定参数更新步长
MiniBatchSize163264每个批次图片数量,影响显存占用和收敛稳定性
MaxEpochs102050训练轮数。轮数过多可能导致过拟合
ValidationDataaugVal验证集,用于观察训练过程中的泛化变化
ValidationFrequency几十或几百每多少次迭代执行一次验证
Plots'training-progress'实时显示损失和准确率曲线
ExecutionEnvironment'auto'自动选择 GPU 或 CPU

训练命令本身很简单:

options = trainingOptions('adam', ... 'InitialLearnRate', 1e-3, ... 'MiniBatchSize', 32, ... 'MaxEpochs', 20, ... 'ValidationData', augVal, ... 'ValidationFrequency', 30, ... 'Plots', 'training-progress', ... 'ExecutionEnvironment', 'auto'); net = trainNetwork(augTrain, layers, options);

我可以给一个很直接的建议:第一次跑,不要追求多高的准确率,先用一个小批量数据把流程跑通。比如每个类别只用几十张图片,训练 5 到 10 轮,确认没有报错,再逐步增加数据量和训练轮数。单次跑通只能说明流程没有断,不能说明模型已经可用。

3.5 验证:不能只看训练集准确率

训练结束后,很多人看训练曲线里准确率很高,就认为任务完成了。这是新手最容易踩的坑。训练集准确率高可能只是模型记住了训练数据,真正要评估的是验证集和测试集表现。

常见评估写法如下:

predLabels = classify(net, augTest); trueLabels = imdsTest.Labels; acc = mean(predLabels == trueLabels); figure; confusionchart(trueLabels, predLabels);

confusionchart会生成混淆矩阵。它能告诉你:哪些数字被错误识别成了别的数字。比如 4 和 9、3 和 8 之间经常混淆,这些信息比一个笼统的准确率更有用。

模型训练好后要保存下来,否则下次使用还要重新训练:

save('models/my_cnn.mat', 'net');

加载时用 `load('models

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/3 6:15:19

基于51单片机的功率因数校正系统设计与实现

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/3 6:15:18

Matlab车牌字符分割实战:从预处理到智能分割

简介:本资源是一套面向MATLAB初学者与图像处理学习者的蓝色车牌字符分割完整实现方案,聚焦车牌识别流程中的关键环节——字符区域精准切分,适用于课程设计、毕业设计及算法验证等实践场景。压缩包共42个文件,含24个核心MATLAB脚本…

作者头像 李华
网站建设 2026/9/3 6:15:04

LLM能做技术选型吗?实验设计与偏差分析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/3 6:14:47

基于Transformer的事件抽取实战:从ACE2005到工业应用

简介:本资源是一套基于Transformer架构的预训练模型在ACE2005事件抽取任务上的完整实现方案,面向NLP方向的研究生、算法工程师及进阶学习者,聚焦于事件触发识别、论元角色分类等核心子任务,适用于学术研究复现、竞赛基线构建与工业…

作者头像 李华
网站建设 2026/9/3 6:13:14

基于STM32的水质检测系统实战:PH/TDS/温度监测与嵌入式开发详解

简介:本资源是一套基于STM32F103系列单片机开发的水质检测系统完整源码工程,面向嵌入式初学者、课程设计学生及毕业设计需求者,解决PH值、TDS(总溶解固体)与水温三项核心水质参数的实时采集、处理与显示问题&#xff0…

作者头像 李华
网站建设 2026/9/3 6:13:01

while循环遇上软件测试:从语法到接口自动化实战

你是否遇到过这样的困境:面试把软件测试流程、测试用例设计背得滚瓜烂熟,入职后打开项目却不知道从哪开始测;或者跟着教程学会了 Python 语法、while 循环,真到写自动化测试脚本时,又发现循环根本用不上。很多准备入行…

作者头像 李华