MNIST手写数字识别实战:基于Keras的深度学习入门与调参指南
发布时间:2026/9/28 12:20:43 锦皓数字建站

先说明一下自己的经历这个 MNIST 手写数字识别几乎是我见过生命力最强的入门项目。这几年在小厂带过团队也帮不少新人改过代码不管大家学的框架是 PyTorch 还是 TensorFlow第一个真正“跑起来有成就感”的 Demo十个里有八个都是它。网上相关的 python 教程和 keras 源码可以说多如牛毛但很多文章要么只甩一段代码要么把概念讲得玄乎其玄真正能让人一遍跑通、还能理解背后每个步骤为什么这么写的反而不多。这篇文章我打算换个方式不贴那种“复制就能用”的完整脚本就完事而是把我自己从零搭这个项目时的思考过程、踩过的坑、调参时的心得全部摊开来讲。你会看到为什么选 Keras 而不是纯手写网络为什么数据要做归一化而不是直接丢进去为什么损失函数偏偏用交叉熵乃至 404 下载失败这类问题怎么处理。目标只有一个让你看完之后不只是跑通 MNIST而是能自己动手改网络结构、调整参数真正理解深度学习项目从数据到模型的完整链条。这个项目适合的人群很广哪怕你刚学完 python 基础语法对机器学习还是一头雾水只要敢动手敲代码一个下午就能见到效果已经有一定基础但只会照着别人的教程跑不知道每行代码在干嘛的朋友也能从中找到不少“原来如此”的瞬间。1. 为什么人人都拿 MNIST 练手项目核心价值拆解1.1 MNIST 数据集到底是个什么东西MNIST 全称是 Modified National Institute of Standards and Technology 数据库最早来自美国国家标准与技术研究院收集的手写数字样本经过一系列预处理之后变成了我们现在看到的样子一共 70000 张图片其中 60000 张是训练集10000 张是测试集每一张都是 28x28 像素的灰度图内容就是 0 到 9 中的一个手写数字。你可能会觉得“就这这也太简单了”但恰恰是这种简单让它成了深度学习领域的“单元测试”。一张图片 784 个像素值加上一个 0 到 9 的标签数据量不大不小训练速度快可视化又方便。无论是刚接触神经网络的新手还是想快速验证一个新想法是否可行的研究员MNIST 都是最顺手的工具。我经常把它比作编程领域的“Hello World”只不过 Hello World 教你的是语法MNIST 教你的是整个深度学习的标准工作流。这个数据集还有一个特点值得注意它已经被清理得非常干净了。图片尺寸统一、数字居中、背景基本为黑色这意味着你可以把大部分精力放在模型本身而不是在数据清洗上。等到以后接触真实业务数据你会发现别说统一尺寸了光是把乱七八糟的格式整理干净就能耗掉你一半时间。1.2 为什么偏偏选 Keras Python 这个组合选 Keras 做这个项目不是因为它是最强大的框架恰恰相反是因为它足够简单直接。Keras 在设计之初就把“让深度学习平民化”当作核心理念它的 API 风格非常接近人类的自然思维。你想加一层全连接网络就是一个Dense想让模型开始学习就是compile加fit。这种“所见即所得”的体验对刚接触神经网络的人特别友好可以把你从繁琐的底层计算中解放出来专注于理解模型结构和训练流程。至于 Python那就更不用多说了。它现在的生态已经强大到几乎成了深度学习的代名词NumPy、Matplotlib、Pandas 这些库和 Keras 配合得天衣无缝。更关键的是Python 的学习曲线比较平缓哪怕你只掌握了列表、字典、循环和函数这几个基础知识就已经能看懂 MNIST 训练代码的大部分内容了。这也是为什么我一直建议初学者直接走 Python 路线学深度学习而不是一上来就碰那些底层框架。还有一个容易被忽略的点社区生态。Keras 背靠 Google文档完善遇到问题一搜基本都有答案热词里那些“keras安装教程”“python入门教程”几乎都是围绕这条链路展开的。一个项目能不能跑通很多时候不是看你水平多高而是看你遇到问题时能不能快速找到解决方案Keras 庞大的用户群体就是你最强的后援。2. 环境准备从零把深度学习跑起来2.1 Python 与 Keras 的依赖关系很多人一上来就在终端敲pip install keras结果装完发现根本没法用或者运行时报一堆错。这背后的核心原因是Keras 本身不是一个独立运行的框架它需要依赖后端引擎来完成真正的数值计算。早年间 Keras 可以选 Theano 或 TensorFlow 作为后端现在主流基本只剩 TensorFlow 一家所以最省心的做法是直接安装完整的 TensorFlow里面已经自带了 Keras API。换句话说你在代码里写的from tensorflow import keras和你单独安装的keras包完全是两回事。前者是 TensorFlow 官方封装的高级 API后者是独立的 Keras 库两者混用容易引发版本冲突。我自己刚入行时就吃过这个亏装了独立的 keras 包又装了 tensorflow结果模型训练时一会儿报找不到模块一会儿报版本不对折腾了整整一个晚上。后来学乖了直接用pip install tensorflow里面什么都有从此再没出过这种问题。关于 Python 版本建议装 3.9 或者 3.10。太老的版本比如 3.6 或 3.7对新版 TensorFlow 支持不好太新的版本比如 3.12 或 3.13反而可能因为依赖库还没适配而报错。TensorFlow 2.10 到 2.16 这个区间搭配 Python 3.9 或 3.10是见过最多、最稳定的组合也是个经过大量实践验证的配置方案。2.2 安装容易踩的坑版本、镜像源真正动手安装的时候还有一个让无数新手抓狂的问题下载速度。TensorFlow 的完整包体积不小如果直接使用默认的 PyPI 源在国内的网络环境下经常会出现超时、断流、甚至卡住一动不动的情况。这时候就需要切换镜像源。常见的做法是临时指定清华大学或阿里云的镜像地址下载速度会从几十 KB/s 直接飙升到几 MB/s。pip install tensorflow -i https://pypi.tuna.tsinghua.edu.cn/simple如果你希望以后每次安装都默认走镜像源可以配置全局源而不是每次敲一长串参数。我自己习惯用阿里云的源稳定性我个人体感最好清华源更新快但偶尔会有同步延迟。哪个源出问题就换另一个这是很常见的环境切换思路。安装完成后建议顺手验证一下版本确保一切正常import tensorflow as tf print(tf.__version__) print(tf.config.list_physical_devices())能正常输出版本号就说明基础环境已经通了。如果你用的是 Nvidia 显卡还可以进一步配置 GPU 版本训练速度会快很多。但对 MNIST 这个项目来说CPU 训练其实也完全够用毕竟 60000 张 28x28 的小图片普通笔记本跑几十秒就能完成一个 epoch没有必要一开始就折腾 CUDA、cuDNN 这些让人头皮发麻的配置。2.3 验证环境是否能用环境装好之后不要急着写完整代码先跑一个极简的验证脚本确认 TensorFlow 能正常导入、计算图能正常建立。我见过不少朋友装完环境满怀激情地把完整代码一贴结果报错信息密密麻麻根本分不清是环境问题还是代码问题。这种排查体验非常消耗信心。最简单的验证方式就是创建一个极小的张量做个简单的矩阵加法import tensorflow as tf a tf.constant([[1.0, 2.0], [3.0, 4.0]]) b tf.constant([[1.0, 1.0], [1.0, 1.0]]) print((a b).numpy())能输出[[2. 3.][4. 5.]]就说明 TensorFlow 的 Op 可以正常运转。这一步验证看似简单但能把环境问题快速隔离出来。如果这一步都报错那你需要回头检查 Python 版本是否匹配、TensorFlow 是否装成了别的架构比如装了 arm64 版却在 x86 机器上跑以及是否有多个 Python 环境互相干扰。尤其最后这一点Anaconda 用户最容易中招conda 里一个 Python系统里还有一个 Pythonpip命令对应的库到底装到了哪个环境里很多时候自己都搞不清楚。验环境时建议直接在 Python 交互式命令行里跑这样至少能确认“当前这个 Python 环境”的依赖是可用的。3. 数据加载与预处理细节决定成败3.1 用 Keras 内置接口加载 MNISTKeras 提供了一个极其方便的数据加载接口keras.datasets.mnist.load_data()。一行代码就能把训练集和测试集全部下载并解析好返回的格式直接就是 NumPy 数组。对于 MNIST 这种入门项目来说完全不需要自己去网上找数据文件、手动解压解析这是 Keras 能让人“零负担入门”的重要原因之一。from tensorflow import keras (x_train, y_train), (x_test, y_test) keras.datasets.mnist.load_data()首次运行时它会自动从网上下载mnist.npz文件大小约 11 MB下载完成后会缓存在本地之后再次运行就直接从本地读取不会重复下载。默认缓存位置在用户主目录下的~/.keras/datasets/目录里。但这里其实藏着一个常见的坑也是网络热搜里“torchvision 下载 mnist 会 404”这类问题出现的根源数据文件的下载是联网进行的如果网络不稳定或者数据源的下载地址访问受限就很容易卡住或者报错。关于这个问题怎么排查我会在后面专门用一整节来详细说这里先卖个关子。数据加载成功后你可以先看看它的形状建立直观印象print(x_train.shape) # (60000, 28, 28) print(y_train.shape) # (60000,) print(x_test.shape) # (10000, 28, 28)60000 张训练图片每一张是 28 行 28 列的二维数组标签就是对应的数字。你可以用 Matplotlib 随机画几张看看这一步非常有利于建立直观印象。import matplotlib.pyplot as plt plt.figure(figsize(10, 4)) for i in range(10): plt.subplot(2, 5, i 1) plt.imshow(x_train[i], cmapgray) plt.title(flabel: {y_train[i]}) plt.axis(off) plt.tight_layout() plt.show()如果你之前学过用 NumPy 造数组、用 Matplotlib 画图这里会有一种非常奇妙的贯通感深度学习的数据本质上就是 NumPy 数组只不过维度比普通表格数据多了一两层而已。3.2 归一化处理先把数值拉回同一起跑线这一步是整个预处理里最关键、也最容易被新手跳过的一步。原始图片的像素值是 0 到 255 的整数而神经网络的激活函数比如后面要用的 sigmoid、tanh、relu对输入数值的范围非常敏感。如果你直接把 0 到 255 的原始值丢给网络数值大的像素会在计算时产生过大的加权和导致神经元迅速进入饱和区梯度变得极小训练几乎停滞。生活化地理解这件事就好比你要对比两个人的身高和体重一个量纲是米、一个量纲是千克数值大小完全不具备可比性。归一化就是先把所有输入拉到同一个尺度上让网络能够在各个维度上平等地学习和更新权重。MNIST 数据最常用的归一化方式非常简单就是直接除以 255x_train x_train.astype(float32) / 255.0 x_test x_test.astype(float32) / 255.0注意这里有个隐藏细节先用astype(float32)把整数转成浮点数再做除法。如果你直接对整数数组做除法NumPy 会默认做整除结果一堆 0 和 1数据信息几乎全部丢失。我第一次写这个项目的时候就犯过这个错误看到训练准确率一直上不去想了半天才发现是数据类型的问题。还有一个很多人会纠结的问题只除 255够不够其实对 MNIST 这种本身就比较“干净”的数据集来说足够了。像素值除以 255 之后会落在 0 到 1 之间网络的输入分布比较友好。更高级一些的做法是标准化也就是减去均值再除以标准差让数据分布接近标准正态分布。不过对于入门项目先做好除以 255 这步就行等以后处理更复杂的图像数据集时再系统学习更精细的预处理方案。3.3 从二维图片到一维向量理解数据的形状变化你可能已经注意到原始图片是 28x28 的二维数组但模型里的全连接层Dense接收的输入通常是一维向量。怎么办两种思路。第一种是手动把二维数组展平比如用 NumPy 的reshape或 Keras 的Flatten层。Flatten层的核心作用就是把 (28, 28) 的数据变成 (784,) 的一维向量。这个操作不涉及任何学习参数纯粹是改变数据的组织方式。第二种是想办法让网络直接处理二维结构这就是卷积神经网络CNN做的事。卷积层通过滑动窗口的方式可以保留图片的空间结构信息效果通常比单纯的全连接网络更好。不过在这个入门项目里我们先聚焦全连接网络先把最基本的流程跑通等理解了全连接网络是怎么学习的再上手 CNN 会更从容。这里也建议大家花点时间理解一下Flatten层的意义。很多人在代码里看到Flatten就机械地写上不知道它为什么存在。举一个直观对比28x28 的图片展平后第 0 个像素是左上角那个点第 783 个像素是右下角那个点原本的二维位置关系被打散了。全连接网络从这个一维向量中学习像素之间的组合模式进而分辨数字。这个思路简单粗暴但在 MNIST 这种图片主体位置相对固定的数据集上效果已经相当不错了。4. 模型构建、训练与评估从原理到代码4.1 搭建最简单的全连接网络环境准备好了数据也预处理好了接下来就是核心环节定义模型结构。Keras 提供了两种常见的模型定义方式顺序模型Sequential和函数式 API。对于 MNIST 这种简单网络Sequential就够了。from tensorflow.keras import layers model keras.Sequential([ layers.Input(shape(28, 28)), layers.Flatten(), layers.Dense(128, activationrelu), layers.Dense(10, activationsoftmax) ])这个结构看起来简单每一层都有它存在的道理。Input层声明输入数据的形状是 28x28也就是一张图片Flatten把二维变成一维第一个Dense(128, activationrelu)是一个有 128 个神经元的全连接层负责从 784 个输入特征中学习和组合出更有意义的特征表达最后一个Dense(10, activationsoftmax)输出 10 个类别的概率分布哪个数字的概率最大模型就预测哪张图片是哪个数字。先解释一下为什么中间层选 128 个神经元。这个数字不是拍脑袋定的在 MNIST 这个数据规模上128 个神经元是一个在表达能力和计算量之间比较平衡的选择。你可以试试把 128 改成 32 或 64模型表达能力会下降准确率可能会有轻微下降改成 512 或 1024训练时间变长但准确率提升的空间已经很小甚至可能因为过拟合导致测试集表现变差。这种“够用就好”的选参思路在工程实践里非常重要。再解释一下激活函数。relu是目前隐藏层用得最多的激活函数它的计算非常简单输入大于 0 就原样输出小于等于 0 就输出 0。这种非线性的截断操作让网络能够拟合复杂的函数关系。早期人们喜欢用sigmoid或tanh但它们在深层网络中容易导致梯度消失训练速度很慢。relu天然不存在这个问题正半轴的梯度始终是 1训练时收敛速度快很多。我在实际项目中对比过同一套模型分别用 relu 和 sigmoid 做隐藏层激活relu 的训练时间往往能比 sigmoid 快一半以上最终准确率也更高。4.2 训练参数为什么这么选损失函数、优化器、批量大小模型搭好了还需要回答一个关键问题模型怎么才算“学得好”这就是损失函数和优化器的职责。训练时每个样本会得到一个预测值损失函数负责衡量预测值和真实标签之间的差距优化器则负责根据这个差距去更新网络的权重。model.compile( optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy] )这里没有用最常见的categorical_crossentropy而是用了带sparse前缀的版本。区别在于标签的编码方式如果你的标签是独热编码如数字 3 变成[0, 0, 0, 1, 0, 0, 0, 0, 0, 0]用categorical_crossentropy如果你的标签还是整数比如 3 就是 3就用sparse_categorical_crossentropy。MNIST 的标签默认是整数所以直接用sparse版本就行。很多新手一上来用categorical_crossentropy然后报维度不匹配的错就是这个原因。优化器选择adam是因为它在实践中的表现非常稳定。Adam 可以理解成“加了动量并且能自适应调整学习率”的梯度下降算法它会在训练过程中根据每个参数的历史梯度信息自动为不同参数调整合适的学习步长。相比传统的 SGDAdam 对学习率的敏感度更低即使你初始学习率设得不太合适它也能比较稳健地收敛。对入门项目来说这是容错率最高的选择。训练时还有一个重要参数批量大小batch size。默认可以用 32意思是一次取 32 张图片计算损失、更新一次权重。为什么不是一个一个地更新batch size 为 1因为单个样本的梯度噪声太大训练过程会非常不稳定为什么不是一次性把所有 60000 张图都算完再更新batch size 为 60000因为那样一次迭代的计算量太大而且容易陷入局部最优。用 32 或 64 这样的小批量既能让梯度估计更稳定又能保证训练速度是实践中最常用的折中方案。设置epochs时5 到 10 个回合在这个数据集上已经能看到不错的效果。一个 epoch 意味着把所有训练图片完整过一遍。你可能想“那我多跑几个 epoch 是不是更好”不一定。训练轮数太多模型可能把训练数据里的噪声和个例都背下来了也就是过拟合反而导致测试集准确率下降。后面我会专门演示怎么看这个现象。4.3 训练过程可视化与评估fit方法执行之后Keras 会在控制台实时打印每个 epoch 结束时的训练损失和准确率。如果配置了验证集还会同步打印验证集上的表现。history model.fit( x_train, y_train, batch_size32, epochs5, validation_split0.1 )validation_split0.1表示从训练数据里再切出 10%6000 张作为验证集训练过程中模型不会在这部分数据上更新权重而是等每个 epoch 结束时用它们来衡量模型对“没见过的数据”的表现。这个做法可以帮你实时监控模型是否开始过拟合。训练完成之后最直接的评价方式就是看测试集上的表现。测试集是模型完全没见过的 10000 张图片它在测试集上的准确率才是真正可信的“成绩单”test_loss, test_acc model.evaluate(x_test, y_test) print(f测试集准确率: {test_acc:.4f})按照我自己的经验这个简单的两层全连接网络训练 5 个 epoch测试准确率通常能达到 97% 以上。听起来不错但距离人类识别手写数字的水平还有差距这也为后面引入 CNN 和改进技巧留下了空间。除了跑出一个准确率数字我还强烈建议大家把预测结果可视化出来哪怕是只为满足一下自己的好奇心。下面这段代码会随机选几张测试图片对比真实标签和模型预测的标签import numpy as np predictions model.predict(x_test) plt.figure(figsize(10, 4)) for i in range(10): idx np.random.randint(0, len(x_test)) plt.subplot(2, 5, i 1) plt.imshow(x_test[idx], cmapgray) predicted_label np.argmax(predictions[idx]) true_label y_test[idx] color green if predicted_label true_label else red plt.title(fpred:{predicted_label} true:{true_label}, colorcolor) plt.axis(off) plt.tight_layout() plt.show()预测正确的用绿色标题预测错误的用红色标题。实际跑下来你会发现即便是 97% 准确率的模型犯的错也往往集中在一些“人看着都费劲”的图片上字迹潦草、形变严重的数字。这个观察能帮你理解一个事实模型的错误模式和数据本身密切相关并不是简单的“不够聪明”。5. 实战中的常见问题与排查实录5.1 数据集下载失败、404 错误怎么办这是所有新手几乎都会遇到的一道坎也是“torchvision 下载 mnist 会 404”这类问题被反复搜索的原因。具体表现是第一次运行load_data()时程序卡在下载阶段很久然后报出 URL 错误、连接超时或者说找不到文件。如果你遇到这种情况先别急着怀疑自己的代码大概率只是网络问题。keras.datasets.mnist.load_data()默认从 Google 存储等外部地址下载文件有些网络环境下这个地址并不稳定。解决方案有几种按推荐顺序排第一手动下载文件后放到本地缓存目录。先找到你系统里的~/.keras/datasets/目录在 Windows 下通常是C:\Users\你的用户名\.keras\datasets然后找一个网络畅通的环境比如用浏览器直接访问或者通过其他途径下载mnist.npz文件把它放到这个目录下。之后再次运行load_data()程序发现本地已经有文件就会直接跳过下载步骤不再外网访问。这个方案的通用性最好因为 Keras 会优先读取本地缓存。第二临时构造数据加载函数自己从 mnist 的源文件里读取。这种方案稍微复杂一点需要额外安装python-mnist之类的第三方库来解析 IDX 格式的原始文件。一般情况下用第一种方案就够了我不建议入门阶段把时间花在数据格式解析上。第三如果你是做项目想完全离线使用可以把下载好的mnist.npz放到项目目录里然后自己写一个加载函数用 NumPy 的load方法去读这样就不依赖 Keras 的缓存机制了。这个方案可控性最强但多出来的代码复杂度其实并没必要。我自己的建议是直接使用方案一手动放置文件。网络这个问题说到底是环境问题不值得为了“验证下载功能”硬耗时间把时间花在模型本身才是正事。5.2 shape 不匹配、准确率不升、过拟合迹象第一个高频报错是shape不匹配。最常见的情况是定义了Input(shape(28, 28))但训练数据传进去却变成了(60000,)或者其他形状。使用keras.datasets.mnist.load_data()时拿到的x_train形状确实是(60000, 28, 28)正常情况下不会出错。但如果你自己改了数据加载方式或者中间做了reshape导致维度丢失就容易出现这个问题。排查思路很直接在fit之前把x_train.shape打印出来看一眼确认数据和模型的输入层对得上。第二个常见现象是训练损失下降但准确率几乎不变或者从一开始就卡在某个值比如 10%。这种情况 80% 是数据预处理出了问题——我印象最深的是整数除法把 255 除掉后所有像素值都变成 0 或 1模型学到了个寂寞。另外还有标签数据格式的问题比如 y 变成了浮点类型也会导致损失计算异常。排查时优先检查x_train.dtype和y_train.dtype再把归一化前后的像素值各打印一个出来看看。第三个现象比较隐蔽就是过拟合。表现为训练准确率越来越高甚至到 99% 以上但验证集准确率停滞不前甚至开始下降。这说明模型开始把训练样本的细节特征背下来了。针对 MNIST 这个项目最简单的对策是降低中间层的神经元数量或者增加Dropout层随机丢弃一部分神经元的输出再或者增加训练数据量做平移、旋转增强。这些方法每一个都可以单独写一篇长文这里先留个印象过拟合不是模型的“缺陷”而是模型和数据之间关系的一种表现工程上有很多手段去缓解它。5.3 参数调整方向和技巧速查下面这个表是我在实际调参过程中总结出来的一套快速自查思路不一定适合所有场景但对 MNIST 这类入门项目非常管用。现象最可能的原因尝试的调整方向训练 loss 不降学习率过大或过小调低或调高 adam 的学习率尝试 1e-3 到 1e-4训练准确率低归一化被跳过或写错检查像素值范围确认是除以 255 后的浮点数验证准确率远低于训练过拟合增加 Dropout减小网络容量增加数据增强预测结果全是同一个数字标签编码和损失不匹配确认用的是 sparse_categorical_crossentropy损失变为 NaN学习率过大降低学习率或检查数据里是否出现 NaN 值训练时间过长中间层神经元太多把 Dense 层神经元从 256 降到 128 或 64这张表的意义不在于让你背下来而是给你一个排查问题的切入点。深度学习项目的“玄学”感很大程度上来自参数之间互相牵连一张表没法覆盖所有情况但能让你在迷失方向时有一个相对可靠的抓手。另外还有一个小技巧训练时把validation_split加上哪怕只是验证一下都能让你在训练过程中提前发现大量问题。很多人喜欢训练完直接看测试集成绩但这样你就没法区分是训练过程出问题还是模型泛化能力不足。验证集的存在就是为了帮你把这个过程拆开来看。写在后面的一点个人体会MNIST 手写数字识别这个项目我前前后后带过不少人跑过也在不同机器、不同环境、不同框架下复现过。每次有人问我“入门深度学习该做什么”我的答案从来没变过先跑通 MNIST而且要亲手把每一行代码敲进去而不是复制粘贴。因为在跑通这个项目的过程中你会真正接触到一个模型从数据到结果的全部关键节点加载数据、观察数据、预处理、定义网络、选择损失函数、设置优化器、迭代训练、评估效果、发现问题、调整参数。这一步流程一但你完整走了一遍后面无论换成 CIFAR-10、还是真实业务里的图像分类核心套路都是相通的。反之如果一上来就去啃大而复杂的模型反而容易陷入无尽的环境配置和报错大海里消耗掉最初那点学习的热情。最后再送一个我个人的小经验跑通之后别急着关掉随便改点东西试试。把Dense(128)改成Dense(16)或者把relu换成sigmoid或者把adam换成sgd看看结果有什么变化。这种“破坏性实验”比照着标准代码跑十遍更有收获。模型是个会反馈的黑盒子你喂它不同的设定它就还你不同的结果这个过程本身就是最好的老师。
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。