神经网络是怎么学的:深入反向传播与梯度下降

神经网络不过是一堆数字——成百万上亿个按层组织起来的参数(权重和偏置)。层的不同排布给你不同的架构(CNN、Transformer 等等),但归根到底它们都是决定网络行为的参数。未经训练的网络输出随机噪声。训练好的网络能识别人脸、翻译语言、写代码。区别就在这些参数的取值。

训练就是为这些参数找到正确取值的过程。网络做出一个预测,衡量自己错得多离谱,然后调整参数,好让下一次错得少一些。判断每个参数该往哪个方向调的算法叫反向传播,真正把调整落实下去的算法叫梯度下降。两者合起来,就是几乎全部现代深度学习背后的引擎。

在最简单不过的例子上,它长这样——一个模型学着拟合一条直线。点几下 Step 看看:

step 0 · loss = 28.00

每一次点击都跑了一轮同样的算法,而这个算法训练着每一个神经网络——从最小的课堂例子到 GPT。模型做出预测,衡量自己错得多离谱(红色虚线),算出该往哪个方向调参数,再轻轻推一下。就这样。重复足够多次,直线就锁到数据上了。

这篇文章要讲的,正是这些步骤内部究竟发生了什么。我们会从一个神经元出发,一路搭到多层网络——每一步都配上交互组件和 Python 代码,好让你直接拿概念做实验。

这是一篇长文——不是给你一口气读完的。分段来看:先从基础开始(什么是神经元、损失怎么算),玩玩那些组件直到豁然开朗,等准备好了再回来看微积分和链式法则。每一节都建立在前一节之上,所以如果哪里觉得不清楚,值得往回翻一节重读。

神经网络就是一堆参数

神经网络由构成,而每一层由神经元构成。神经元是计算的最小单位——它接收一些输入,做一个简单的计算,产生一个输出。把很多神经元并排放好,就得到一层。把几层首尾相接——一层的输出喂进下一层的输入——就得到一个神经网络。整个网络,无论多大,都只是这些相同的小零件重复并连接起来。所以要理解整体,我们可以从理解一个神经元开始。

我们用下面的组件玩一玩,建立直觉。 这是一个有 3 个输入的单神经元。 左边你控制输入值,右边你设置神经元的参数—— 它的权重向量、偏置和激活函数。你改动任何东西,图示都会实时更新。

w₁x₁ + w₂x₂ + w₃x₃ + b → activation(sum) → output
0.5·1 + 0.5·0 + 0.5·1 + 0.0 = 1.0 → perceptron(1.0) = 1
inputs
1.0
0.0
1.0
neuron parameters
0.5
0.5
0.5
0.0

在这个组件上你能看到,神经元把每个输入乘以对应的权重,连同偏置一起加起来, 再把结果送过一个激活函数

所以从数学上说,我们做的是这个:

output=f(w1x1+w2x2+w3x3+b)\text{output} = f(w_1 x_1 + w_2 x_2 + w_3 x_3 + b)

其中 ff 是激活函数。每个输入 xix_i 乘以对应的权重 wiw_i,乘积与偏置 bb 一起求和,结果再送过激活函数 ff神经元做的全部就是这些——相乘、求和、激活。

试几种配置来建立直觉:

  • 权重控制某个输入有多重要。 把 x₁ 设为 1.0,其余设为 0。现在拖动 w₁——输出会直接跟着变。把 w₁ 设为 0,那个输入就被彻底忽略,无论它的值是多少。
  • 负权重起抑制作用。 设 x₁ = 1.0、w₁ = -1.5,其余为 0。加权和变成负的——在感知机激活下,输出是 0。这个神经元在主动压制那个输入。
  • 偏置平移决策边界。 当所有输入都是 0 时,只有偏置决定这个和。正偏置意味着即使没有输入神经元也会发放。负偏置意味着输入必须先「压过」它,神经元才会激活。
  • 激活函数塑造输出的形状。 从 Perceptron(硬性的 0/1)切换到 Sigmoid——现在输出是 0 到 1 之间的平滑值。再试试 ReLU——它把正值原样放行,把负值截成零。

激活函数为什么重要?没有它们,每个神经元都只是一个线性函数(乘一乘、加一加),而把线性函数叠起来仍然是线性函数——不管你加多少层。激活函数引入了非线性,正是它让神经网络能学到曲线、边缘和复杂模式,而不只是直线。事实上,只要神经元足够多、激活是非线性的,神经网络几乎能逼近任意函数——这被称为通用逼近定理。让这成为可能的,正是激活函数。想看网络如何用这些非线性变换扭曲、折叠输入空间,直到复杂模式变得可分,Chris Olah 的 Neural Networks, Manifolds, and Topology 有一份漂亮的可视化讲解。

权重和偏置就是神经元的参数——它需要学出来的那些值。神经元为每个输入存一个权重——在上面的组件里,就是一个 3 元素的向量。权重决定每个输入有多重要,偏置把结果整体上移或下移。这些参数合起来定义了神经元对什么有反应。换一组权重和偏置,同一个神经元就会在输入中检测出完全不同的模式。

由于「相乘再求和」就是点积,这通常写成向量形式:

output=f(wx+b)\text{output} = f(\mathbf{w} \cdot \mathbf{x} + b)

这里 x\mathbf{x} 是所有输入组成的向量(列表)——例如 [x1,x2,x3][x_1, x_2, x_3]—— 而 w\mathbf{w} 是所有权重组成的向量——例如 [w1,w2,w3][w_1, w_2, w_3]。 点积 wx\mathbf{w} \cdot \mathbf{x} 把每一对相乘再把结果加起来:w1x1+w2x2+w3x3w_1 x_1 + w_2 x_2 + w_3 x_3

在 Python 里,可能长这样:

import numpy as np

class Neuron:
    def __init__(self, n_inputs):
        self.w = np.random.randn(n_inputs)   # w — weight vector, e.g. [w₁, w₂, w₃]
        self.b = 0.0                         # b — bias

    def forward(self, x):                    # x — input vector, e.g. [x₁, x₂, x₃]
        z = np.dot(self.w, x) + self.b       # w · x + b — dot product + bias
        return max(0, z)                     # f(z) — activation function (ReLU)

# Create a neuron with 3 inputs and run it
neuron = Neuron(3)
output = neuron.forward(np.array([1.0, 0.5, 0.7]))

一层就是一堆神经元

把很多这样的神经元摞在一起,你就得到一层。再把几层拼起来,就得到一个神经网络。 这里有个小的——2 个输入、两个各含 3 个神经元的隐藏层、1 个输出。之所以叫「隐藏」,是因为你只看得到进去的输入和出来的输出——中间那些层是网络内部的,从外面看不见:

x₁x₂yinputlayer 1layer 2output

上面我们看到,每个神经元都为它的所有输入持有一组权重,外加一个偏置,并包在激活函数里:

output=f(w1x1+w2x2+w3x3+b)\text{output} = f(w_1 x_1 + w_2 x_2 + w_3 x_3 + b)

一层通常把它们统一存在一个权重矩阵WW)里——每个神经元占一行权重。 对上图那个网络,layer1.W 是一个 3×2 的矩阵——3 个神经元,每个 2 个权重,因为每个神经元接收 2 个输入:

layer1.W = [[ 0.4, -0.2],    ← neuron 0: weights for x₁, x₂
            [ 0.1,  0.7],    ← neuron 1: weights for x₁, x₂
            [-0.3,  0.5]]    ← neuron 2: weights for x₁, x₂

对单个神经元,我们有一个权重向量与输入的点积:f(wx+b)f(\mathbf{w} \cdot \mathbf{x} + b)。对一整层,我们把所有权重向量堆成矩阵 WW,把所有偏置堆成向量 b\mathbf{b},于是同一个运算一次就作用到所有神经元上:

output=f(Wx+b)\text{output} = f(W\mathbf{x} + \mathbf{b})

计算 WxW\mathbf{x} 时,WW 的每一行都与输入做点积——那就是一个神经元的加权和。矩阵乘法一次运算就把它们全做完了。

single neuron (dot product)w₁w₂w₃w·x₁x₂x₃x=outone row → one outputstack 3 neuronsfull layer (matrix multiply)0.4-0.2← n₀0.10.7← n₁-0.30.5← n₂W@x₁x₂x=w₀·xw₁·xw₂·x← neuron 0← neuron 1← neuron 2

在 Python 里,长这样:

import numpy as np

class Layer:
    def __init__(self, n_inputs, n_neurons):
        # W is a matrix where each ROW is one neuron's weights.
        # Shape: (n_neurons, n_inputs) — so W[0] is neuron 0's weights,
        # W[1] is neuron 1's weights, etc.
        self.W = np.random.randn(n_neurons, n_inputs)
        self.b = np.zeros(n_neurons)  # b — bias vector, one per neuron

    def forward(self, x):
        # W @ x multiplies every neuron's weight row by the input,
        # computing all dot products at once
        return np.maximum(0, self.W @ x + self.b)  # ReLU activation

# Build the network from the diagram above
layer1 = Layer(2, 3)   # 2 inputs  → 3 neurons  (6 weights + 3 biases = 9)
layer2 = Layer(3, 3)   # 3 inputs  → 3 neurons  (9 weights + 3 biases = 12)
output = Layer(3, 1)   # 3 inputs  → 1 neuron   (3 weights + 1 bias   = 4)

# Forward pass — each layer's output feeds into the next
x = np.array([0.5, 0.8])
h1 = layer1.forward(x)       # input → hidden layer 1
h2 = layer2.forward(h1)      # hidden layer 1 → hidden layer 2
y  = output.forward(h2)      # hidden layer 2 → output

当我们计算 W @ x 时,每一行都与输入做点积——一层里的所有神经元在一次运算里就算完了。这就是神经网络为什么用线性代数,也是 GPU(为矩阵运算而生)为什么能让训练变快的原因。

每个权重、每个偏置都是一个参数。把上面那个网络里的数一数——第 1 层有 9 个,第 2 层有 12 个,输出层有 4 个——总共 25 个参数。这是个极小的网络。GPT-2 大约有 15 亿参数;GPT-3 有 1750 亿。

关于缩放定律的研究表明,随着模型规模、训练数据和算力的增加,模型性能往往可预测地变好——这也是这个领域不断把这些数字往上推的原因。不过也有迹象表明,单纯堆参数正在遇到边际收益递减,重心正转向更好的训练数据、更高效的架构,以及推理、思维链这类能从既有模型规模里榨出更多东西的技术。

感受一下涉及的算力:训练 GPT-3 大约需要 3 × 10²³ 次运算。按每秒十亿次算,那是单处理器上的 1000 万年。数千块 GPU 并行,把它压缩到了几周。

网络是怎么学的?

训练一个神经网络,意味着为所有这些参数找到正确的取值—— 每一层矩阵里的每个权重,每一层向量里的每个偏置。 训练开始时,这些参数被初始化成很小的随机值——网络字面意义上什么都不知道。训练就是不断迭代调整这些随机数,直到它们能给出有用的预测。

每次迭代都会重复的通用训练循环长这样:

阶段它做什么
1. 前向传播把输入送过网络的每一层,乘以权重并施加激活,产生一个预测
2. 计算损失用损失函数(例如 MSE)把预测与真实目标值作比较,把所有误差压缩成一个数字——模型错得有多离谱?
3. 反向传播用链式法则在网络中反向行走,为每个参数算出一个梯度——它该往哪个方向动,动多少?
4. 梯度下降把每个参数减去其梯度的一小部分(学习率),把整个网络朝更低的损失推一把

这四个阶段分成两大块。前向传播(阶段 1)把输入送过网络得到预测。当模型被部署做推理——在生产环境里给出预测——时,只有前向传播在跑。权重不会被更新。

训练(阶段 2、3、4)是预测之后发生的一切:衡量误差、判断每个参数该往哪调、把更新落实下去。这三个阶段协同工作,也是本文的重点。

这个过程有两个方向——数据向前流动产生预测,然后梯度向后流动更新权重:

  Forward pass: each layer receives activations, passes output →

              activations    activations    activations
  Input ─────────▶ Layer 1 ─────────▶ Layer 2 ─────────▶ Output ──▶ Loss


  Backward pass: each layer receives gradient signal, passes it ←

              gradients      gradients      gradients
  Input ◀───────── Layer 1 ◀───────── Layer 2 ◀───────── Output ◀── ∂L
              ↓ ∂L/∂W₁           ↓ ∂L/∂W₂           ↓ ∂L/∂W₃
           (own weight       (own weight         (own weight
            gradients)        gradients)          gradients)

注意其中的对称:前向传播时,每一层从前一层收到激活值,并把自己的输出往前传。反向传播时,每一层从后一层收到梯度信号,并把它往后传。两个方向上,每一层都需要邻居给的输入才能干活。

一个简单例子:拟合一条直线

为了理解这一切怎么运作,我们把一切精简到最简单的网络:一个神经元、一个输入、一个权重、一个偏置。一旦你看清梯度下降在 2 个参数(一个权重和一个偏置)上是怎么工作的,跳到 25 个或者 250 亿个,就只是规模问题了。

回忆一下单个神经元算的是什么:f(wx+b)f(w \cdot x + b),其中 xx 是输入向量,ww 是权重向量。把它精简到一个输入并忽略激活函数,你就得到 f(x)=wx+bf(x) = wx + b——你在学校学过的那个基本线性函数。这不是简化;在经由激活函数引入非线性之前,神经元做的字面上就是这件事。 所以训练单个神经元去拟合一条直线,是这个问题最纯粹的版本。

假设有人给你五个点,说它们来自某个线性函数 f(x)=wx+bf(x) = wx + b,并让你找出 wb

x(输入)-2-1012
y(输出)-3-1135

这跟你在学校做的事正好相反。 在代数课上,人家给你一个方程比如 y=3x+5y = 3x + 5,让你「解出 xx」——函数参数 w=3w=3b=5b=5 是已知的,你要找的是输入 xx。而我们这里不是在解 xx。 我们在求的是定义这个函数本身的参数——ww(斜率,也就是这条线有多陡)和 bb(截距,也就是它在 y 轴上的交点)。

可既然已经有了数据,为什么还要找一个函数?因为整件事的重点是处理你从没见过的输入。如果有人问「x = 1.5 时输出是多少?」而 1.5 不在你的数据里,表格帮不上忙。但如果你已经发现底层函数是 f(x) = 2x + 1,你立刻就能回答:4。这就是泛化——在没见过的新输入上做出正确预测的能力。

那我们把这些点画到图上(绿色圆点),试着靠拖动滑块手动找出 wb。 用损失值来指引你——拖一拖,看看能不能把损失降到零。你会发现 w=2w = 2b=1b = 1 能让损失归零——那正是生成这些数据的参数,也就是 y=2x+1y = 2x + 1

0.0
0.0

我们一直在用损失来指引搜索——但它到底是什么?损失是一个数字,它告诉你模型总体上错得有多离谱。损失高时,预测离数据很远。损失为零时,模型拟合得完美。

拖动滑块时,注意那些红色虚线——它们是每个数据点上各自的误差,显示预测与真实值差了多少。我们把每个点的误差算作 error=predictionactualerror = prediction - actual

训练模型时,我们需要一个数字告诉我们预测总体上错得多离谱。那就是损失——一个把所有误差收拢成一个分数的函数。机器学习里有很多损失函数,各自适合不同任务:

  • 均方误差(MSE)——用于回归(预测数值)。把每个误差平方后求平均。
  • 交叉熵——用于分类(预测类别)。衡量预测的概率离真实标签有多远。
  • 平均绝对误差(MAE)——像 MSE,但用绝对值而非平方,对离群值不那么敏感。

因为我们在拟合一条直线——这类任务叫回归(预测一个连续数值)——我们用了均方误差(MSE):把每个误差平方,然后全部求平均。平方做了两件事——让所有误差变正(这样它们不会互相抵消),并且对大误差的惩罚远重于小误差:

MSE=1ni=1n(yprediyactuali)2=(ypred1yactual1)2+(ypred2yactual2)2++(yprednyactualn)2nMSE = \frac{1}{n} \sum_{i=1}^{n} (y_{\text{pred}_i} - y_{\text{actual}_i})^2 = \frac{(y_{\text{pred}_1} - y_{\text{actual}_1})^2 + (y_{\text{pred}_2} - y_{\text{actual}_2})^2 + \cdots + (y_{\text{pred}_n} - y_{\text{actual}_n})^2}{n}

把这个公式用到我们的数据上。假设你当前的猜测是 w=3w = 3b=1b = 1,也就是 f(x)=3x+1f(x) = 3x + 1。对我们那 5 个数据点,我们逐一算出预测、误差(差多远)和误差的平方:

xxyactualy_{\text{actual}}ypred=3x+1y_{\text{pred}} = 3x + 1误差误差²
-2-3-5-24
-1-1-2-11
01100
13411
25724
平均 →损失 = 2.0

把每个误差平方(这样负值不会抵消),然后求平均。结果是一个数字:损失 2.0 意味着我们的预测平均偏离约 1.4(21.4\sqrt{2} \approx 1.4)。越大表示拟合越差,零表示完美。当我们在上面的组件里把 w=2w = 2b=1b = 1 设好时,损失降到零,因为那正是生成数据的参数。

对我们这个简单的双参数函数,我们靠手动就把参数找出来了,但想象一下 25 个参数的情形,更别提几百万个。稍后我们会把本文的一切用到一个真实任务上——训练一个网络识别手写数字,它有超过 10 万个参数。 在那个规模上手动调参是不可能的——没有人能在百万维的空间里摸索。我们需要一种系统的办法:看着损失,用数学算出该把每个参数往哪个方向推,才能让它变小。 这正是反向传播梯度下降合力做的事:反向传播算出每个参数该往哪调, 梯度下降沿那个方向迈一小步。我们不断重复这个过程,直到损失最小。

对简单的线性模型来说,有一个精确公式——正规方程——一步就给出完美的 wb,不需要迭代。但它只对线性模型有效。一旦有了非线性激活、多层结构和几百万参数,就没有公式了。梯度下降不需要公式——它只需要衡量误差、算出该往哪推。它对任何可微模型都有效,这也是它成为通用训练算法的原因。

一步步走完训练循环

先来看看这个自动化的算法是怎么工作的。 下面的组件让你为我们的任务一步步跑反向传播和梯度下降,把一切看在眼里:

  • 左图:数据点(绿色圆点)和模型基于当前 wwbb 给出的预测直线(蓝色)。红色虚线显示每个点上的误差——预测与真实值之差。这些误差被平方并求平均,得到 MSE 损失。
  • 右图:每一步的损失曲线——你就是靠它来监控训练是否奏效。曲线稳步下降说明模型在学;如果它走平或者飙升,就有什么需要调整了。

Step 执行一次梯度下降更新,或者点 Step x10 一口气跑十次。

Backpropagation

-3.0
3.0

Gradient Descent

0.10

Computation (step 0)

一步之后,直线仍然是错的,但错得少一些了。再跑一次。再一次。每一步都让误差更小、梯度更小,直线也一点点爬向目标。 展开「Computation」部分可以看到每一步的数字。现在先把学习率保持在默认值(0.1)——它是干什么的、怎么选,下一节再讲。

不断点 Step 的同时看看右图——损失一开始陡降(模型很快纠正了最严重的错误,因为梯度很大),接近答案时就走平了(误差更小意味着梯度更小,所以每一步做的事更少)。这叫收敛——模型稳定到正确的参数上。

这一切发生的同时,学习率自始至终没变过。记住,梯度是 dw = 2 * mean(error * x)——它是从误差算出来的。模型越接近正确答案,误差越小,梯度也就越小,于是更新量 lr * dw 也变小。学习率没变,但步子会自动变小,因为可纠正的误差本就少了。这种自我调节的行为是梯度下降的标志:在最要紧的地方走得快,然后不需要你动手就小心翼翼地精修。

每点一次「Step」,就跑了一整轮训练迭代——对应我们前面描述的训练循环的四个步骤,在 Python 里长这样:

 # 1. forward pass
y_pred = w * x + b

# 2. loss computation
error  = y_pred - y
loss   = np.mean(error ** 2)

# 3. backpropagation
dw = 2 * np.mean(error * x)
db = 2 * np.mean(error)

# 4. gradient descent
w = w - lr * dw
b = b - lr * db

步骤 1 和 2 是前向传播和损失。前向传播把每个输入送过 y = wx + b 得到预测(每个输入的预测值可以在 Computation 面板里看到)。接着损失计算衡量我们错得有多离谱——每个预测与真实值之差,平方后求平均,得到一个数字(MSE)。MSE 怎么工作,我们在上一节已经讲过。

步骤 3 和 4 才是学习发生的地方。反向传播(步骤 3)为每个参数算出一个梯度——该往哪推、推多少。梯度下降(步骤 4)把这些梯度用上,从每个参数里减去一小部分(学习率)。在扎进复杂度大多集中的反向传播之前,我们先简单聊聊学习率,也就是步骤 4 里的那个 lr

计算学习率

在上面的计算部分里可以看到,梯度下降是这样更新参数的:

w = w - lr * dw
b = b - lr * db

梯度(dwdb)告诉我们每个参数该往哪个方向动,以及相对于其他参数动多少。但实际上该迈多远?这由学习率lr)控制——它在应用之前把每个梯度缩放一下。

方向总是对的——但步长可能是错的。如果 lr 太小,每一步几乎挪不动,训练要跑到天荒地老。如果 lr 太大,你会越过最小值,落到另一侧,那里的损失反而更糟。把学习率想成一个「自信程度」的旋钮:

  • 太小(试试 0.01)——每一步都极小。模型一点点蹭向答案,要几百步才到。安全,但慢得难受。
  • 刚刚好(试试 0.1)——模型迈着自信的步子,20–30 步就收敛。损失先快速下降,然后精修。
  • 稍微偏大(试试 0.5)——模型越过最小值,在两侧来回弹。但每次越过都落得离底部更近,那里梯度更小,所以弹幅逐渐收缩,最终仍然收敛——只是路径呈锯齿形,步数比 lr = 0.1 多。
  • 太大(试试 1.0)——越冲越极端。每一步都落在离最小值很远的地方,那里梯度仍然很大,于是引出下一个大步。它也许还能收敛,但抖得厉害,也很浪费。
  • 大得离谱(试试 1.5)——越冲极端到每一步落点都比上一步更糟。梯度变大而不是变小,于是下一步更大——一个把损失螺旋推高的正反馈回路。这叫发散

自己试试——改一下学习率再点 Step x10 看效果:

0.10

「正确」的学习率没有公式。实践中,大多数人从一个常见默认值(0.001 或 0.0001)起步,使用像 Adam 这样的自适应优化器,它会根据每个参数梯度的历史表现自动调整该参数的步长,并配上学习率调度:一开始大(大步靠近),训练过程中逐渐变小(小步精修)。几乎所有人用的都是 Adam 或它的变体,而不是朴素的梯度下降。

这些全都是同一个 4 阶段训练循环之上的改良。核心算法没有变。

计算反向传播

在上面的计算部分里可以看到,反向传播是这样算梯度的:

dw = 2 * np.mean(error * x)
db = 2 * np.mean(error)

这两行里塞了不少东西。为什么 dw 要把 error 乘上 x,而 db 不用?那个 2 从哪来的?mean 又和这有什么关系?我们一步步拆开。

记住,损失是从预测算出来的,而预测取决于 wb。 归根到底,损失是参数的函数——改变 wb,损失就变。 在下面的组件里拖动 wb,看着损失怎么变——白点沿曲线移动,精确显示你在损失地形上的位置:

-3.0
3.0

试着拖 w——点沿左边的曲线移动,但右边的曲线会重新成形。为什么?右图问的是「对每一个可能的 b,损失是多少?」——其中 w 固定为滑块给出的值。当你改变 w,你就改变了那个固定值,从而改变了每个 b 处的误差,得到一条完全不同的曲线。反过来也一样:拖 b,左边的曲线就重新成形。w 的最佳取值取决于 b 在哪,反之亦然——它们是耦合的。

梯度公式来自对这些曲线求导数——衡量当你把每个参数推动一丁点时,损失变化多少。 所以 2 * mean(error * x) 不过就是损失函数对 w 的导数。

要理解我们如何从损失函数走到 2 * mean(error * x),需要三个层层递进的概念:

  1. 导数——衡量一个函数如何变化意味着什么
  2. 链式法则——当函数被串起来时,导数怎么算
  3. 偏导数与梯度——如何同时处理多个参数

到最后,我们会精确追溯这个公式里每一部分的来历。先从导数究竟是什么讲起。

我们刚看到的公式——2 * mean(error * x)——是针对线性模型配 MSE 损失的特例。不同的损失函数和架构会给出不同的梯度公式——但底层的数学原理始终相同。对我们这个简单模型,公式可以手推;对有几百万参数的深度网络,PyTorch、TensorFlow 这类框架会用 autograd(自动微分)自动求导。

导数:某一点上的斜率

导数回答一个问题:如果我把这个输入推动一丁点,输出会变多少? 把它想成曲线在某一点的斜率。如果你站在山坡上,导数告诉你脚下的地有多陡——以及往哪边是下坡。

拿一个简单函数比如 f(x) = x²。沿曲线拖动点 x,看看斜率和导数怎么变:

At x=1.0 derivative is 2.0
1.0
0.80

Computation

在组件里把 x 设为 2、dx 设为 0.5。黄线dxdx)是给输入的一个推动,绿线dfdf)是输出随之变化了多少。图下方的 Computation 部分展示了它们怎样组合成导数。

首先,我们在这一点上求值:f(2)=4f(2) = 4。然后把输入推动 dx 再求一次:f(2.5)=6.25f(2.5) = 6.25。差值告诉我们输出变化了多少:df=6.254=2.25df = 6.25 - 4 = 2.25。除以推动量得到变化率:df/dx=2.25/0.5=4.5df/dx = 2.25 / 0.5 = 4.5

这个比值(4.5)近似等于 x = 2 处的导数——它给出的是速率:在这一点上,输出变化的速度大约是输入的 4 倍。 它不正好是 4,因为 dx = 0.5 仍是个不小的推动。现在把 dx 调小——试着拖到 0.1

  • f(2)=4f(2) = 4
  • f(2.1)=4.41f(2.1) = 4.41
  • df=0.41df = 0.41
  • df/dx=0.41/0.1=4.1df/dx = 0.41 / 0.1 = 4.1——更接近 4

随着 dx 变小,这个比值收敛到精确的导数。

这就是全部的想法——导数就是当 dxdx 收缩趋近于零时 df/dxdf/dx 所趋向的值:某一点上精确的变化率。

一般公式看着复杂,其实正是我们刚才做的事:

f(x)=limdx0f(x+dx)f(x)dxf'(x) = \lim_{dx \to 0} \frac{f(x + dx) - f(x)}{dx}

f(x+dx)f(x)f(x + dx) - f(x) 是输出的变化(dfdf)。除以 dxdx 得到比值。limdx0\lim_{dx \to 0} 这一部分只是说「让 dxdx 收缩趋近于零」——正是你用滑块做的事,看着比值收敛到精确值。

f(x)=x2f(x) = x^2,我们可以推一推:

f(x+dx)=(x+dx)2=x2+2xdx+dx2f(x + dx) = (x + dx)^2 = x^2 + 2x \cdot dx + dx^2 f(x+dx)f(x)=2xdx+dx2f(x + dx) - f(x) = 2x \cdot dx + dx^2 f(x+dx)f(x)dx=2x+dx\frac{f(x + dx) - f(x)}{dx} = 2x + dx

dx0dx \to 0 时,就只剩 2x2x。所以 dfdx=2x\frac{df}{dx} = 2x

这正是微积分的意义所在——它免去了挑一个 dxdx 的必要。组件展示了为什么你可以信任那个精确公式:不管你挑哪个 dxdx,随着它变小,比值都趋向 2x2x。所以我们跳过「让它变小」的过程,直接用 2x2x

如果你想对导数建立更深的理解,3Blue1Brown 的 The Essence of Calculus 是现存最好的讲解。整个系列都值得看——它建立起教科书常常跳过的那份直觉。

链式法则:求复合函数的导数

我们知道怎么求 f(x)=x2f(x) = x^2 这类简单函数的导数。但当一个函数的输出喂进另一个函数时会怎样? 那叫函数复合——而我们的计算正是这样:

y_pred = w * x + b             # prediction
error  = y_pred - y            # how far off
loss   = np.mean(error ** 2)   # squared error, averaged

我们想弄清楚怎么推动 w 才能减小损失,但只有 f1f_1 含有 w。 我们要算损失对 w 的导数,可 lossf3f_3)并不以 w 为参数——它接收的是 error。 而 errorf2f_2)也不接收 w——它接收 y_pred。只有 y_predf1f_1)最终接收 w

所以可以看到,从 w 算到损失并不是一个函数,而是三个函数组成的链条,每一个都把自己的输出送进下一个:

wf1y_predf2errorf3error2w \xrightarrow{f_1} y\_pred \xrightarrow{f_2} error \xrightarrow{f_3} error^2

写清楚就是:

  • f1(w)=wx+bf_1(w) = w \cdot x + b——模型的预测
  • f2(y_pred)=y_predyf_2(y\_pred) = y\_pred - y——我们差了多远
  • f3(error)=error2f_3(error) = error^2——误差的平方(我们要最小化的东西)

损失就是 f3(f2(f1(w)))f_3(f_2(f_1(w)))——三个函数一层套一层。

每个单独函数的导数我们都会求,但怎么把它们合起来得到整条链的导数? 答案是链式法则:把各段的局部导数相乘。

d(loss)dw=f1f2f3\frac{d(\text{loss})}{dw} = f'_1 \cdot f'_2 \cdot f'_3

回忆一下导数公式:

f(x)=limdx0f(x+dx)f(x)dxf'(x) = \lim_{dx \to 0} \frac{f(x + dx) - f(x)}{dx}

分子 f(x+dx)f(x)f(x + dx) - f(x) 是输出的变化——dfdf。分母 dxdx 是输入的变化。所以整体就是 dfdx\frac{df}{dx}——「ff 的变化除以 xx 的变化」。这是写 f(x)f'(x) 的另一种方式。输出在上,输入在下:

  • f1f'_1:输出是 y_predy\_pred,输入是 wwd(y_pred)dw\frac{d(y\_pred)}{dw}
  • f2f'_2:输出是 errorerror,输入是 y_predy\_predd(error)d(y_pred)\frac{d(\text{error})}{d(y\_pred)}
  • f3f'_3:输出是 error2error^2,输入是 errorerrord(error2)d(error)\frac{d(\text{error}^2)}{d(\text{error})}

用这套记法,链式法则展开成:

d(loss)dw=f1f2f3=d(y_pred)dwd(error)d(y_pred)d(error2)d(error)\frac{d(\text{loss})}{dw} = f'_1 \cdot f'_2 \cdot f'_3 = \frac{d(y\_pred)}{dw} \cdot \frac{d(\text{error})}{d(y\_pred)} \cdot \frac{d(\text{error}^2)}{d(\text{error})}

为什么是相乘?因为每个函数都嵌在下一个里面——一个的输出成了另一个的输入。把它想成一串推动:如果你把 w 推动一丁点,y_pred 就变化 x 倍于这个推动量。然后 error 变化 1 倍于 y_pred 的变化量。再然后 error² 变化 2·error 倍于 error 的变化量。链条上每一环都在缩放这个推动——而缩放是靠相乘复合起来的。

组合函数的三种方式

把两个函数 f(x)f(x)g(x)g(x) 组合起来有三种基本方式,每一种都有自己那套导数如何组合的规则:

  1. 相加h(x)=f(x)+g(x)h(x) = f(x) + g(x)——导数相加。如果 ff 变化 3、gg 变化 5,那么和变化 8。这就是和法则h(x)=f(x)+g(x)h'(x) = f'(x) + g'(x)

  2. 相乘h(x)=f(x)g(x)h(x) = f(x) \cdot g(x)——这更复杂些,因为两个因子都会变。这就是乘积法则h(x)=f(x)g(x)+f(x)g(x)h'(x) = f'(x) \cdot g(x) + f(x) \cdot g'(x)。你得分别考虑每个函数在另一个保持不变时的变化。

  3. 复合(嵌套)h(x)=f(g(x))h(x) = f(g(x))——gg 的输出喂进 ff。导数相乘。这就是链式法则h(x)=f(g(x))g(x)h'(x) = f'(g(x)) \cdot g'(x)。给 xx 的一个推动先被 gg' 缩放,缩放后的变化再被 ff' 缩放一次。

我们的损失计算是一个复合——f3(f2(f1(w)))f_3(f_2(f_1(w)))——所以我们把导数相乘。要是这些函数是相加或相乘组合的,我们就得用相应的规则。实践中,神经网络三种都用:相加(偏置项)、相乘(权重乘输入)和复合(层与层相互喂入)。反向传播会为每个运算套用对应的那条规则。

好,那我们来算算这条损失链的整体导数。我们展示过怎么求 x2x^2 的导数,得到了 2x2x。 同样的办法对更简单的函数也适用:ax+bax + b 的导数就是 aa(一个常数倍),xcx - c 的导数是 11(减去一个常数不改变变化率)。这让每个单独导数的计算都很直白:

  • 要算 f1=d(y_pred)dwf'_1 = \frac{d(y\_pred)}{dw},我们利用「ax+bax + b 的导数是 aa」。因为 y_pred=wx+by\_pred = w \cdot x + b,导数就是 x
  • 要算 f2=d(error)d(y_pred)f'_2 = \frac{d(\text{error})}{d(y\_pred)},我们利用「xcx - c 的导数是 11」。因为 error=y_predyerror = y\_pred - y,导数就是 1
  • 要算 f3=d(error2)d(error)f'_3 = \frac{d(\text{error}^2)}{d(\text{error})},我们利用「x2x^2 的导数是 2x2x」。因为函数是 error2error^2,导数就是 2 · error

于是我们得到:

d(loss)dw=f1f2f3=x1(2error)=2errorx\frac{d(\text{loss})}{dw} = f'_1 \cdot f'_2 \cdot f'_3 = x \cdot 1 \cdot (2 \cdot error) = 2 \cdot error \cdot x

对单个数据点来说,这就是 2errorx2 \cdot error \cdot x

用真实数字走一遍。取 w = 3, b = 1,数据点 x = 2, y = 5

w = 3
  ↓  × x = ×2
y_pred = 3·2 + 1 = 7
  ↓  × 1
error = 7 - 5 = 2
  ↓  × 2·error = ×4
error² = 4

链式法则:f1f2f3=214=8f'_1 \cdot f'_2 \cdot f'_3 = 2 \cdot 1 \cdot 4 = 8。意思是如果我们把 w 推动 1,这个数据点的平方误差会变化 8。

但我们有 5 个数据点,不是一个。既然 MSE 把所有点的平方误差取平均,那导数也要取平均。对每个点我们算 2errorx2 \cdot error \cdot x

xxyyypred=3x+1y_{pred} = 3x + 1errorerror2errorx2 \cdot error \cdot x
-2-3-5-22(2)(2)=82 \cdot (-2) \cdot (-2) = 8
-1-1-2-12(1)(1)=22 \cdot (-1) \cdot (-1) = 2
0110200=02 \cdot 0 \cdot 0 = 0
1341211=22 \cdot 1 \cdot 1 = 2
2572222=82 \cdot 2 \cdot 2 = 8

求平均:8+2+0+2+85=4\frac{8 + 2 + 0 + 2 + 8}{5} = 4。所以 dw = 4——梯度告诉我们增大 w 会让损失上升,因此我们应当减小它。(确实如此,真值是 w = 2,比我们猜的 3 要小。)

用数学记法就是:

d(loss)dw=1ni=1n2errorixi=21ni=1nerrorixi\frac{d(\text{loss})}{dw} = \frac{1}{n} \sum_{i=1}^{n} 2 \cdot error_i \cdot x_i = 2 \cdot \frac{1}{n} \sum_{i=1}^{n} error_i \cdot x_i

在 Python 里:

dw = 2 * np.mean(error * x)

下面的组件让你为每个数据点追踪这条链。点击不同的 x= 按钮,看看局部导数怎么变——注意链式法则对每个点给出的值都不同,因为 xerror 不一样:

Tracing d(loss)/dw for data point:
w
y_pred = w·x + b
error = y_pred - y
error²
-3.0
3.0

db 来说链条是一样的,只是 f1f'_1 不同: 因为 y_pred=wx+by\_pred = w \cdot x + b,对 b 的导数就是 1(而不是 x)。所以:

d(loss)db=f1f2f3=11(2error)=2error\frac{d(\text{loss})}{db} = f'_1 \cdot f'_2 \cdot f'_3 = 1 \cdot 1 \cdot (2 \cdot error) = 2 \cdot error

在我们的 Python 代码里长这样:

db = 2 * np.mean(error)

这就是链式法则为什么重要:只要损失是通过一连串运算算出来的(而它总是如此),你就需要它来回溯每个参数是怎么影响最终结果的。对我们这个双参数模型,链条有 3 步。一个深度神经网络可能有几百步——每层一步——但原理完全一致:每一层不过是复合里多出的一个函数、多出的一个要相乘的局部导数。

把导数用到我们的损失函数上

既然知道了怎么求单个导数、怎么用链式法则把它们组合起来,就把这套知识用到我们的问题上。 我们想最小化的那条「曲线」就是我们的损失函数——我们选来衡量误差的那个公式。 我们对这个函数了如指掌:

loss(w,b)=1ni=1n(wxi+byi)2loss(w, b) = \frac{1}{n}\sum_{i=1}^{n}(w \cdot x_i + b - y_i)^2

loss = np.mean((w * x + b - y) ** 2)

我们不知道的是,wb 取什么值时它最小。导数帮我们弄清楚这一点: 它告诉我们,如果把 w 增大一丁点,损失是上升还是下降?有多快?

数据(xy)是固定的——那是我们的训练数据。如果我们暂时也把 b 固定住(比如 b = 3),那么损失就成了只关于 w 的函数,可以画成一条简单曲线。比如在 w = 0 处:

y_pred = w * x + b                  # 0 * [-2,-1,0,1,2] + 3 = [3, 3, 3, 3, 3]
error  = y_pred - y                 # [3,3,3,3,3] - [-3,-1,1,3,5] = [6, 4, 2, 0, -2]
loss   = np.mean(error ** 2)        # mean([36, 16, 4, 0, 4]) = 12.0

这给了我们曲线上的一个点:(w=0, loss=12)。对从 -5 到 5 的每个 w 都这么做一遍(保持 b = 3 不变),我们就得到完整图景——损失作为 w 单独的函数:

0.0(b = 3.0 fixed)

x 轴是 w(我们正在测试的参数值),y 轴是损失(模型在那个 w 上错得多离谱)。结果是一条抛物线——它的最低点在 w = 2,那里损失降到零。拖动 w 滑块并看下面的计算——对上面训练集里的 5 个数据点(x = [-2, -1, 0, 1, 2]),它逐一算出预测、误差和平方误差,然后把它们平均成一个损失值。那就是曲线上的白点。(-5 到 5 这个范围是随意选的——只要够宽,能看出形状并包含最小值就行。我们本可以从 -100 扫到 100,但那样曲线会缩得太小,细节看不清。)

b 我们可以做同样的事——这次固定 w = 2,让 b 从 -5 变到 5:

3.0(w = 2.0 fixed)

同样的抛物线形状,只是现在 x 轴是 b。最小值在 b = 1,那里损失降到零。合起来,w = 2b = 1 正是生成我们数据的那对参数——y = 2x + 1

wb 两次计算,我们都在整个范围内的许多取值上算了损失,才画出完整曲线。这有助于建立直觉——但在实践中,你绝不会这么干。只有 2 个参数时,穷举每种组合微不足道。可真实的神经网络有几百万参数。要画出损失地形,你得把它们在所有组合下全试一遍——代价高到不可想象。这正是我们需要导数的原因:与其把整条曲线画出来找最小值,不如在单个点上算出斜率,然后往下坡迈一步。我们从来看不到全貌。我们只是感受脚下的地面。

偏导数

注意我们刚才做了什么:为了理解损失如何依赖 w,我们冻住 b,只变 w。为了理解它如何依赖 b,我们冻住 w,只变 b。这正是偏导数的含义——在其他参数保持不变的前提下,损失对某一个参数的导数:

  • ∂loss/∂w——推动 w 时损失如何变化(b 被冻住)
  • ∂loss/∂b——推动 b 时损失如何变化(w 被冻住)

上面那两条曲线,各自都是沿着一个参数在损失地形上切出的一片。那条曲线在任意点的斜率就是偏导数:

dw = 2 * np.mean(error * x)        # ∂loss/∂w — how loss changes with w
db = 2 * np.mean(error)            # ∂loss/∂b — how loss changes with b

下面的组件把上面看过的两条曲线合在一起,现在展示的是偏导数在起作用。左图变动 w(固定 b),右图变动 b(固定 w)。在每张图上,白点是你当前所在的位置,蓝色虚线是切线(它的斜率就是偏导数),绿色箭头指出往哪边走能减小损失。

-3.0(b = 3.0 fixed)
3.0(w = -3.0 fixed)

拖动滑块,看看会怎样:

  • 远离最小值时——曲线很陡,切线倾斜得厉害,导数是个大数。梯度下降在这里迈大步。
  • 靠近最小值时——曲线变平,切线几乎水平,导数接近零。步子变得极小——模型在精修。
  • 在最小值处——切线完全水平。导数为零。无处可去——你到了。

注意一件有意思的事:当你拖 w 时,右图会重新成形——反过来也一样。为什么?

左图问的是:「对每一个可能的 w,损失是多少?」——其中 b 固定为滑块给出的值。当你拖 w,你只是在挑要站在那条曲线的哪一点上。曲线本身不会变,因为 b 没变。

但右图问的是:「对每一个可能的 b,损失是多少?」——其中 w 是固定的。当你拖 w,你改变了用于计算右边曲线上每个点的那个固定 w。不同的 w 意味着每个 b 处的误差不同,也就意味着一条完全不同的曲线。(反过来也是——拖 b 会让左图重新成形,而右图上只是那个点在动。)

这就是为什么我们要在更新任何东西之前,用同一批误差把两个偏导数都算出来。推动 w 的最佳方向取决于 b 当前在哪,反之亦然——所以先测两个斜率,再同时移动两个参数。

从导数到梯度

梯度不过是把所有偏导数捆在一起的向量:[dw, db]。它指向损失上升最陡的方向。所以我们朝相反方向走——这就是更新规则里为什么是减法:w = w - lr * dw

上面那两张图其实只是一个三维曲面的二维切片。有两个参数时,我们可以把完整的损失地形可视化出来——一个轴是 w,另一个轴是 b,高度是损失。尽管我们的模型是线性的(y = wx + b),损失函数却是二次的——一个碗状——因为 MSE 把误差平方了。不同的损失函数(比如交叉熵)会给出不同形状的地形。这个碗被称为凸的损失地形——只有一个山谷,所以无论从哪里出发,每一个下坡方向都通向同一个底部。试着从不同起点点 Step (both)——你总会落在 w ≈ 2, b ≈ 1。深度神经网络的地形更复杂,是非凸的,有多个山谷,但同样的梯度下降算法在实践中仍然管用得出奇。

-3.0
3.0

试着分别点 Step (w) 和 Step (b)——你会看到点每次只沿一个轴移动,在碗里走出阶梯状的路径。 然后试试 Step (both)——这才是真正的梯度下降在做的事,一次更新两个参数。你可以拖动来旋转曲面,从不同角度看它。

黄色箭头是梯度向量——它显示下一步会往哪走。它把 dwdb 合成一个方向:「往这边走,损失下降得最快。」当你点 Step (w) 或 Step (b) 时,你只是在沿这个向量的一个分量移动。当你点 Step (both) 时,你沿着完整的箭头走。

你可能注意到箭头主要沿着 w 轴。那是因为两个梯度并不相等——在起点处 dw = -20,而 db = 4w 的分量大了 5 倍,所以由它主导方向。 梯度下降并不在所有方向上等量移动; 它按照损失对每个参数的敏感程度成比例地移动。 这里损失沿 w 陡峭得多,所以 w 被优先纠正。

为什么损失对 w 的敏感度远高于对 b?答案在梯度公式里。对照一下——dw = 2 * mean(error * x)db = 2 * mean(error)。注意关键差别:dw 把每个误差乘上了对应的 x 值,而 db 只用误差本身。我们的 x 取值在 -2 到 2 之间,所以当模型错得很离谱(误差大)而且输入也大时,乘积 error * x 会变得很大。偏置的梯度 db 只是对误差本身求平均——没有乘 x——所以天然更小。

这是一个普遍性质,并非我们这个玩具例子独有。在任何神经网络里,有些参数对损失的影响大于另一些。梯度自动抓住了这一点——梯度大的参数得到大更新,梯度小的参数得到小更新。每个参数被纠正的幅度,正比于它对误差的贡献。这就是梯度下降之所以高效的原因:它不会把力气浪费在那些本来就已经接近正确的参数上。

走出碗外:深度网络为何非凸

我们这个玩具损失之所以是个完美的碗,有一个具体原因:模型 y = wx + b 对其参数是线性的,而把线性模型的误差平方,总会得到二次的损失——而任何二次函数都是单个凸的山谷。只要加上一个带非线性激活的隐藏层,这个保证就没了。

有两样东西破坏了凸性。第一是嵌套。在真实网络里,一个参数并不直接接触损失——它坐在某个激活函数里面、坐在下一层加权和里面、坐在那个激活函数里面,如此层层套下去。哪怕最简单的单隐藏层网络算的也是类似这样的东西

f(xi)=β0+k=1Kβkg ⁣(wk0+j=1pwkjxij),f(x_i) = \beta_0 + \sum_{k=1}^{K} \beta_k \, g\!\left(w_{k0} + \sum_{j=1}^{p} w_{kj} x_{ij}\right),

其中每个隐藏单元 kk 都把自己的加权和裹进非线性的 gg 里。损失 12i(yif(xi))2\frac{1}{2}\sum_i (y_i - f(x_i))^2 现在是权重的深度复合函数,而复合函数会弯——它们长出隆起和凹陷,而不是一个干净的山谷。

第二是对称性。隐藏单元是可互换的:把单元 1 和单元 2 对调(连同它们的权重),网络算出的是完全相同的函数完全相同的损失。有 KK 个隐藏单元时,这样的重新编号有 K!K! 种,所以每个解都伴随着一群散落在地形各处的同卵双胞胎。一个有许多同样优秀的最低点的损失函数,按定义就不是单个碗。

后果就是局部极小值——比周围一切都低、却不是全局最低点的山谷。梯度下降永远只从起点往下走,所以原则上它可能落进某个局部极小值就停下,找到了一个解,但不是最好的那个,即全局极小值

在实践中这远没有听起来那么要紧。在真实网络所处的极高维空间里,糟糕的局部极小值很少见——大多数梯度消失的点其实是鞍点(某些方向下坡、另一些方向上坡),梯度下降会从旁边滑过去;而确实存在的那些局部极小值,往往几乎和全局的一样好。小批量更新带来的噪声(下一节)也会把参数摇晃得足以逃出浅坑。这就是为什么同一个简单算法,在纸面上看似无望的地形上依然管用得出奇。两个标准习惯能让它保持诚实:慢慢学——逐步拟合,一旦验证损失开始上升就立刻停下——以及正则化——加上惩罚项(比如 L2/岭项)把权重往零拉,顺带把地形抹平。

随机梯度下降

到目前为止,我们一直用全部训练数据一次性计算梯度。每次你在上面的组件里点「Step」,dw = 2 * mean(error * x) 都会为我们那 5 个数据点分别算出 2 * error * x,再把它们平均成一个梯度:

dw = 2 * mean(error * x)
   = 2 * mean([-24.00, -7.00, 0.00, -3.00, -16.00])
   = -20.00

5 个点的话,这不算什么。用上全部数据的好处是,平均后的梯度指向可能的最佳方向——每个数据点都有一票,所以没有哪个离群点能把这次更新带偏。但真实数据集有几百万甚至几十亿条样本。为每一条都算一次梯度、全部平均之后才迈出一步,代价太高。想象一个有十亿条样本的数据集——你得处理完全部十亿条,才能把 wb 更新哪怕一次。那才一步。然后下一步再来一遍。每一步给你的梯度都非常准,但你在两次更新之间要等到天荒地老。

解决办法很简单:别一次用完所有数据。 把训练样本打乱,然后切成小组——小批量(mini-batch)。 在第一个小批量上跑完那套流程:算预测、误差、损失、梯度,然后更新参数 ——只用那几条样本。接着换下一个小批量,如此继续。

这就是随机梯度下降(SGD)。「随机」不过是指那次随机打乱。更新规则一样,只是求和范围从整个数据集换成了小批量:

w=wlr1BiBlossiww = w - lr \cdot \frac{1}{|B|} \sum_{i \in B} \frac{\partial \text{loss}_i}{\partial w}

其中 BB 是当前的小批量。每个小批量给出的是真实梯度的一个带噪估计——它不会指向精确正确的方向,但大体是对的。在一个 epoch 的过程中,每条训练样本都会有所贡献,噪声也就被平均掉了。

当你把每条样本都过了一遍,那就是一个 epoch。再打乱一次,开始下一个 epoch。 这就是你在训练日志里看到「epoch」的原因——每个 epoch 意味着模型把数据集里每条样本恰好看过一次。

对我们那 5 个数据点、批大小为 2 的情形,长这样:

Epoch打乱后的数据批 1批 2批 3
1[0, 2, -1, -2, 1](0, 2)(-1, -2)(1)
2[2, -2, 1, 0, -1](2, -2)(1, 0)(-1)
3[-1, 1, -2, 2, 0](-1, 1)(-2, 2)(0)

每个批次都跑完整的 4 步循环(前向传播 → 损失 → 反向传播 → 梯度下降),所以每个 epoch 做 3 次参数更新,而不是 1 次。到每个 epoch 结束时,每个数据点恰好被用过一次——但顺序每次都不同,这能防止模型把序列背下来。

下面的组件在我们这个 5 点数据集上并排跑两种方法,方便你直接对比。两者从相同的参数(w = -3, b = 3)出发,用相同的学习率。每点一次「Step」,两种方法各做一次参数更新。 蓝线(全批量)每一步都用全部 5 个点。 橙线(小批量)只用 batch_size 个点 ——橙色圆圈标出用的是哪些,epoch 进度条则追踪在数据集里走到了哪。

2
0.10

点几下「Step」,看看右边的损失曲线。蓝色曲线(全批量)平滑下降——每一步都用上全部数据,所以梯度始终指向最佳方向。 橙色曲线(小批量)则是锯齿状的。

默认启用了固定顺序,所以每次运行的批次都一样,这个锯齿图案是可复现的。取消勾选就会在每个 epoch 随机打乱——橙色曲线每次都不同,但整体行为是一样的。

点过前几步就能看出原因:第 1→2 步损失下降(这个批次恰好给了个好梯度),但第 2→3 步损失上升了——那个批次把参数拉向了对它自己那些点有利、却伤害了其他点的方向。然后第 3→4 步又降下来。这很正常:每个批次只看到数据的一个切片,所以有些步会过冲甚至走错方向。经过很多步之后,这些误差互相抵消,模型照样收敛。

全批量用更少的步数就收敛,因为每一步都用上全部数据——那何必用小批量?把我们 5 个数据点上做 3 次参数更新的代价比一比:

  • 全批量(3 步):每步都用全部 5 个点。那是 3 × 5 = 15 次数据点计算。每个点被处理 3 次。
  • 批大小为 2 的 SGD(3 步 = 1 个 epoch):每步只用 2 个点(最后一批只有 1 个)。那是 2 + 2 + 1 = 5 次数据点计算。每个点被处理一次。

两者都做了 3 次参数更新,但 SGD 用了少 3 倍的计算量。更新更带噪,但在规模上省下的量是巨大的——十亿条样本、批大小 1000 时,一个 epoch 就给你一百万次更新,而每条样本只被处理了一次。全批量则要为每一次这样的更新处理完全部十亿条。组件展示不出这个代价差异(两者都不过是点一下),但在真实规模下,小批量尽管步数更多,按墙钟时间却更快抵达答案。

你也可以试试不同的批大小,看看行为差异:

  • 批大小 = 1——噪声最大,每步只用一个点。损失曲线剧烈震荡,但仍会收敛。5 步 = 1 个 epoch(每个点看过一次)。这就是最原始的「随机」梯度下降。
  • 批大小 = 2——噪声更小,每个 epoch 需要 3 步(2 + 2 + 1 余下)。这更接近实践中的用法。
  • 批大小 = 5——那就是把我们全部数据放进一个批次,因此与全批量梯度下降完全相同。两条线完全重合。1 步 = 1 个 epoch。

实践中,32、64 或 256 这样的批大小很常见。权衡在于:批越小,每个 epoch 的步数越多(更带噪,但每步更快);批越大,步数越少(更平滑,但每步计算更多)。小批量带来的噪声其实可能有帮助——它能防止模型卡在浅的局部极小值里。

跨越多层的链式法则

还记得反向传播如何用链式法则把局部导数相乘吗? 我们的例子是最简单的情形:一个神经元,一个权重 w、一个偏置 b、一个输入:

wbxnŷerrlossinput1 neuronoutput∂loss/∂w: x · 1 · 2·error ∂loss/∂b: 1 · 1 · 2·error

从每个参数到损失的链条有 3 环 (f1f2f3)(f'_1 \cdot f'_2 \cdot f'_3),在 Python 里长这样:

y_pred = w * x + b                  # f1: prediction
error  = y_pred - y                 # f2: how far off
loss   = error ** 2                 # f3: squared error

dw = x * 1 * (2 * error)           # chain rule for w: f'1 · f'2 · f'3
db = 1 * 1 * (2 * error)           # chain rule for b: same chain, different f'1

但真实网络每层有多个神经元,每个都有自己的权重。梯度计算本身没变——我们仍然用链式法则为每个单独的权重算偏导数,和之前一样。区别在于规模。对离损失更远的权重,链条变得更长——要穿过更多层,要相乘的导数更多。而且它还变得更宽——一个神经元的输出可以喂进下一层的许多神经元,所以梯度必须把所有这些路径的贡献加起来。

看看这个两层网络。在两个按钮之间切换,看看权重所在的位置如何改变梯度路径:

w₁v₁x₁x₂h₁h₂h₃g₁g₂g₃ŷlossinputlayer 1layer 2output

Gradient for v₁(第 2 层) 默认是激活的——v₁ 是 h₁ 到 g₁ 这条连接上的权重。它的梯度链很短:只有那一条连接,然后 g₁→ŷ→损失。注意只有 h₁→g₁ 这条连接亮起,而不是 h₂→g₁ 或 h₃→g₁。为什么?因为当我们对 v₁ 求偏导时,g₁ 的其他输入是保持不变的——它们乘的是别的权重,不会出现在 v₁ 的导数里。这跟单神经元模型里的原理一样:w·x + bw 的导数就是 x——另一个参数 b 不会出现。这里同理:v₁·h₁ + v₂·h₂ + v₃·h₃ 对 v₁ 的导数就是 h₁。所以梯度是:

lossv1=h1局部导数lossg1从输出 → 到损失\frac{\partial \text{loss}}{\partial v_1} = \underbrace{h_1}_{\text{局部导数}} \cdot \underbrace{\frac{\partial \text{loss}}{\partial g_1}}_{\text{从输出 → 到损失}}

看这个公式的另一种方式是把 lossg1\frac{\partial \text{loss}}{\partial g_1} 展开——既然损失就是 error²,而误差直接来自输出,这就归结为 2 · error

lossv1=h1局部导数lossg1从输出 → 到损失=h1输入2error来自损失\frac{\partial \text{loss}}{\partial v_1} = \underbrace{h_1}_{\text{局部导数}} \cdot \underbrace{\frac{\partial \text{loss}}{\partial g_1}}_{\text{从输出 → 到损失}} = \underbrace{h_1}_{\text{输入}} \cdot \underbrace{2 \cdot error}_{\text{来自损失}}

现在点 Gradient for w₁ (layer 1)——w₁ 是 x₁→h₁ 这条连接上的权重。链条开头一样(一条连接),但接下来 h₁ 的输出会喂进第 2 层的每个神经元——g₁、g₂ 和 g₃。w₁ 的一个变化要先波及它们全部,才能抵达输出和损失。梯度必须把这三条路径的贡献加起来

lossw1=x1relu(z1)(v1lossg1+v2lossg2+v3lossg3)对穿过第 2 层的所有路径求和\frac{\partial \text{loss}}{\partial w_1} = x_1 \cdot \text{relu}'(z_1) \cdot \underbrace{\left(v_1 \cdot \frac{\partial \text{loss}}{\partial g_1} + v_2 \cdot \frac{\partial \text{loss}}{\partial g_2} + v_3 \cdot \frac{\partial \text{loss}}{\partial g_3}\right)}_{\text{对穿过第 2 层的所有路径求和}}

把每个 lossgi\frac{\partial \text{loss}}{\partial g_i} 展开成 2 · error(跟前面一样——损失就是 error²):

=x1relu(z1)(v12error+v22error+v32error)= x_1 \cdot \text{relu}'(z_1) \cdot \left(v_1 \cdot 2 \cdot error + v_2 \cdot 2 \cdot error + v_3 \cdot 2 \cdot error\right)

所以链条不只是更长(要穿过更多层),还更宽(每层要累加的路径更多)。网络更深意味着链条更长;层更宽意味着每条链上的路径更多。

这跟我们用 f1f2f3f'_1 \cdot f'_2 \cdot f'_3 时是同一个链式法则原理。输出神经元算的是 ŷ = v₁·h₁ + v₂·h₂ + b。当我们问「h₁ 变化时 ŷ 怎么变?」,局部导数就是 v₁——正如在单神经元模型里,w·x + bx 的导数是 w。所以 v₁ 在这里出现,并不是作为被优化的参数,而是作为连接 h₁ 到 ŷ 的那个函数的局部导数。链条上每一环都贡献自己的局部导数,我们把它们全乘起来——只不过其中一个导数恰好是另一层的权重罢了。这就是与单层模型的关键区别:更靠前那些层的梯度,必须穿过后面所有层的权重。

一旦拿到所有梯度,更新和之前一样——从每个权重里减去 lr × 梯度

# Output layer (short chain)
W2 = W2 - lr * dW2          # 2 weights

# Hidden layer (long chain — gradients passed through W2)
W1 = W1 - lr * dW1          # 4 weights (2×2 matrix)

下面把两层里各一个权重的梯度计算并排放在一起:

# Gradient for v₁ (output layer) — short chain
dv1 = h1 * 2 * error
#     ↑    ↑
#     │    └── from loss
#     └── local derivative: ∂ŷ/∂v₁ = h₁

# Gradient for w₁₁ (hidden layer) — longer chain
dw11 = x1 * relu_deriv(z1) * v1 * 2 * error
#      ↑    ↑                 ↑    ↑
#      │    │                 │    └── from loss
#      │    │                 └── passes through output layer weight
#      │    └── activation derivative
#      └── local derivative: ∂z₁/∂w₁₁ = x₁

输出层的权重 v₁ 链条里有 2 项。 隐藏层的权重 w₁₁ 有 4 项——它必须穿过激活函数 和输出层的权重 v₁ 才能抵达损失。 你每加一层,就把之前每一层的链条再拉长一次乘法。 到了 100 层,第一层的梯度就是 100 多项的乘积。 把这么多数字乘在一起会发生什么?

梯度消失问题

如果每层的局部导数都是 0.5,那么经过 100 层之后梯度就是 0.51000.5^{100}——小到实际上等于零。最前面的那些层拿不到任何有用的梯度信号。它们学不到东西。这就是梯度消失问题,它困扰了深度网络数十年。

反过来同样糟糕:如果局部导数大于 1,梯度就会爆炸——大到让更新变得极不稳定。

下面的组件让你亲眼看到这一点。把局部导数拖到 1 以下,看着梯度在向后流经各层时衰减为无——然后试试大于 1 的值,看它爆炸:

8
0.50

这就是深度学习长期停滞的原因——在 sigmoid 激活下(其导数总是小于 1),超过几层的网络里梯度就消失了。解决它的那些突破包括:

  • ReLU 激活——它的导数不是 0 就是 1,所以梯度穿过它时不会缩小
  • 残差连接(skip connection)——给梯度一条绕过若干层的捷径,免得它必须逐层相乘穿过每一层
  • 批归一化——把流经每层的数值保持在表现良好的范围内,防止导数持续缩小或增长

这些都是让局部导数保持接近 1 的办法,这样即便跨越几百层,链式法则的乘积也不会消失或爆炸。

准备好把这一切付诸实践了吗?在下一篇文章里,我们会在 MNIST 上构建并训练一个真正的神经网络——用 NumPy 从零写出前向传播、反向传播和梯度下降,然后与 Keras 实现作对比。