从 MNIST 手写数字识别出发,一步步拆解「学习」的原理
你看到一张手写的"5",大脑瞬间就能识别它。但对计算机来说,这只是一堆 0-255 之间的像素值。神经网络做的事情,本质上就是找到一个数学函数,把 784 个像素值映射到 10 个数字的概率分布上。
flowchart LR
A --> B
B --> C
整个 MNIST 识别问题,可以浓缩为一句话:学习 = 找一个好的函数 f(x),使得对于输入 x(图片),输出 f(x)(数字)尽可能接近真实答案。神经网络的所有复杂结构,都是为了表达这个函数而存在。
一张 28x28 的灰度图片,本质上是一个 784 维的向量。每一位的取值范围是 [0, 255]。但这意味着:现实世界的物体被抽象成了高维空间中的一个点。
60000 张训练图片,就是 60000 个点,它们在 784 维空间中分布。同一数字的图片,在这个高维空间中应该"聚集"在一起,形成一个"簇"。神经网络的任务,就是找到这些簇之间的"分界线"。
为什么要把像素值从 [0, 255] 除以 255 变成 [0, 1]?这不仅仅是数值计算的需要,更是让不同维度处于同一尺度。
想象你用"身高(厘米)"和"体重(公斤)"来预测一个人的健康状况。身高数值在 150-200 之间,体重在 40-100 之间。如果不归一化,身高的数值波动会"淹没"体重的影响。归一化就是把所有科目的成绩都换算成百分制,让每一维都有同等的话语权。
归一化的本质是坐标变换:将数据从一个坐标系映射到另一个坐标系,使得优化问题的条件数(condition number)更小,梯度下降收敛更快。从几何上看,就是把一个"扁长"的误差曲面变成"圆润"的,梯度方向更直接指向最低点。
每张图片都有一个标签(0-9),这就是监督学习的定义:我们给模型"出题"的同时也给了"标准答案"。学习的过程,就是不断调整参数,让模型的输出越来越接近标准答案。
其中 x 是输入(图片),y 是标签(真实数字),n 是样本数量。训练集有 60000 个样本,测试集有 10000 个样本。
我们的模型有三层:输入层(784维)→ 隐藏层(128个神经元)→ 输出层(10个神经元)。每一层都做两件事:线性变换 + 非线性激活。
flowchart TD
subgraph Input["输入层"]
I1["x1"] & I2["x2"] & I3["..."] & I784["x784"]
end
subgraph Hidden["隐藏层 (ReLU)"]
H1["h1"] & H2["h2"] & H3["..."] & H128["h128"]
end
subgraph Output["输出层 (Softmax)"]
O1["0"] & O2["1"] & O3["..."] & O10["9"]
end
I1 --> H1 & H2
I784 --> H128
H1 --> O1 & O10
H128 --> O1 & O10
style Input fill:#f0f4ff,stroke:#007aff
style Hidden fill:#f0f9ff,stroke:#5ac8fa
style Output fill:#f0fdf4,stroke:#34c759
前向传播就是从输入到输出的计算过程。每一层都可以写成一个数学公式:
其中 W 是权重矩阵,b 是偏置向量。W 的形状决定了"神经元之间的连接强度",也就是这个网络"记住"的知识。
| 参数 | 形状 | 含义 | 参数量 |
|---|---|---|---|
| W(1) | 784 x 128 | 输入层到隐藏层的权重 | 100,352 |
| b(1) | 128 | 隐藏层偏置 | 128 |
| W(2) | 128 x 10 | 隐藏层到输出层的权重 | 1,280 |
| b(2) | 10 | 输出层偏置 | 10 |
| 总计 | - | - | 101,770 |
为什么需要 ReLU?如果没有激活函数,无论多少层,最终都等价于一层线性变换。激活函数引入了非线性,才让神经网络有了"万能逼近"的能力。
数学上已经证明:一个宽度足够大的单隐藏层神经网络,可以逼近任意连续函数。这就是为什么神经网络如此强大——它理论上可以学习任何输入到输出的映射关系,只要给它足够多的神经元和足够好的训练数据。
输出层的 10 个神经元,经过 Softmax 之后变成了 10 个概率值,加起来等于 1。Softmax 的数学定义是:
想象 10 个学生考试,分数有高有低。Softmax 就像把分数换算成"胜率":分数最高的,胜率最大,但其他人也不是零概率。指数函数的作用是放大差距——高分的概率会被指数放大,让模型更"自信"。
如何知道模型学得好不好?我们需要一个量化指标,这就是损失函数。对于分类问题,最常用的是交叉熵损失:
直觉理解:如果真实类别是 5,模型给 5 的概率越高,损失越小;给 5 的概率越低,损失越大。损失就是模型"错误程度"的数学度量。
训练的目标是最小化损失函数。怎么最小化?梯度下降法:计算损失对每个参数的偏导数(梯度),然后沿着梯度的反方向走一小步。
其中 η(eta)是学习率,决定每一步走多大。
想象你蒙着眼睛站在一座山上,想要走到山谷最低处。你怎么做?用脚感受脚下最陡的方向,然后往那个方向走一步。重复这个过程,直到你感觉不到坡度了——你就到了谷底。梯度就是"最陡的下坡方向",学习率就是每一步迈多大。
怎么高效计算每个参数的梯度?答案是反向传播,它的数学基础是微积分中的链式法则。
flowchart LR
A["输入 x"] -->|前向传播| B["隐藏层"] --> C["输出 y"] --> D["损失 L"]
D -->|反向传播| C
C -->|梯度回传| B
B -->|梯度回传| A
style A fill:#f0f4ff,stroke:#007aff
style B fill:#f0f9ff,stroke:#5ac8fa
style C fill:#f0fdf4,stroke:#34c759
style D fill:#fff8e1,stroke:#ff9500
链式法则说:∂L/∂x = ∂L/∂y · ∂y/∂x。反向传播就是从输出端往回走,每一步都把"下游的梯度"乘以"本地的梯度",就像多米诺骨牌一样,误差信号一路传回到每一个参数。这使得计算梯度的复杂度和前向传播差不多——这是深度学习能跑起来的关键!
我们不每次只用一张图片更新参数(太慢),也不用全部 60000 张(太占内存),而是用Mini-batch:每次取一小批(比如 32 张)来计算梯度。
一个 Epoch 就像把整本书看了一遍。Batch 就像一次看几页然后做个总结。看 5 遍书(5 epochs),每遍分很多个小批次,这样既高效又稳定。批次太小了容易"走弯路"(噪声大),太大了又"走不动"(计算慢)。32 是个经验上不错的折中。
训练集上表现好不代表真的学会了。模型可能只是"背下了所有答案",遇到新题就不会了。这叫过拟合(overfitting)。
训练集就像平时做的练习题,测试集就像期末考试。真正的能力不是你做过多少题,而是你能不能做对没见过的题。如果一个学生把练习题答案都背下来了,考试遇到新题就完蛋——这就是过拟合。我们用测试集来检验模型的"真本事"。
准确率 = 预测正确的样本数 / 总样本数。我们的模型在测试集上达到了约 93.5% 的准确率,意味着 100 张图片里约 93-94 张能认对。
模型输出的是 10 个数字的概率分布。我们取概率最大的那个作为最终预测(argmax)。但概率分布本身也很有价值——它告诉我们模型"有多确定"。
如果模型说"这是 5 的概率是 99%",说明它很确定;如果说"45% 是 5,40% 是 3",说明它也拿不准。不确定性本身就是信息——在真实应用中,我们可以对低置信度的样本交给人来审核,这就是人机协作的基础。
让我们退一步,从最高的抽象层次来看。我们的神经网络有 101,770 个参数。每一组参数值,对应一个函数。所有可能的参数组合,构成了一个 101,770 维的"参数空间"。
训练的过程,就是在这个超高维空间中,从一个随机的起点出发,沿着梯度的方向,一步一步走到一个"损失最小"的位置。每一个位置,都对应一种"知识状态"。
| 层次 | 概念 | 数学对象 | 直觉 |
|---|---|---|---|
| 第一层 | 数据 | 高维空间中的点集 | 经验、观察、事实 |
| 第二层 | 模型 | 参数化的函数族 | 假设、结构、可能性的范围 |
| 第三层 | 学习 | 在参数空间中优化 | 从经验中提炼规律 |
有趣的是,人工神经网络的学习方式,和人脑的学习有深刻的相似性:
也许,学习的本质,无论生物的还是人工的,都是同一个数学原理的不同实现——通过反馈来调整内部状态,使得对外界的预测越来越准确。
开头那句爱因斯坦的话:"Everything should be made as simple as possible. But not simpler." 神经网络的设计完美诠释了这句话:
学习,就是从数据中提炼结构,
用数学的语言,书写世界的规律。
以下是用纯 NumPy 实现的 MNIST 神经网络,没有任何深度学习框架依赖。每一行代码都对应着上面的数学公式,让你看到"理论如何落地"。
点击下载按钮,保存为 mnist_nn.py
只需要 NumPy 和 Matplotlib
一键运行,自动下载 MNIST 数据
原图使用 TensorFlow/Keras,封装了所有数学细节。但从零用 NumPy 实现,能让你亲手写出每一个矩阵乘法、每一次梯度计算。当你看到 np.dot(X.T, dz1) 时,你看到的不再是黑盒,而是链式法则在你指尖流动。