训练

反向传播

链式法则:从损失一路回到每个权重

模型
GPT-2 small 1.24 亿个梯度
句子
5 个预测 损失 4.263
画出
第 1–8 维 共 3,072 × 768

全部步骤

  1. 责任倒着流

    GPT-2 small 读入 “The cat sat on the floor”,为它的 5 个下一词元猜测打分:平均损失 4.263。训练需要知道它 1.24 亿个权重中的每一个一旦变动,损失会怎样变化。反向传播从损失倒推回输入,一趟就得到全部。

    损失 4.263 · 一次反向传播得到 1.24 亿个梯度

  2. 线性层:dW = Xᵀ · dY

    取第 12 块的最后一个矩阵,y = x · W。它的梯度是 dW = Xᵀ · dY:每个权重的责任是它的输入乘以到达它输出处的梯度,再对 5 个位置求和。把鼠标停在一格上可查看它的求和。

    dW = Xᵀ · dY · 对 5 个位置求和

  3. 往下传:dX = dY · Wᵀ

    同一层还把责任传给它的输入,dX = dY · Wᵀ,然后一路往下:经过 GELU(乘以它的斜率)、升维投影、LayerNorm 和注意力。每个权重矩阵都在途中得到自己的 dW。残差相加把梯度直接复制,越过每个子层。

    每一层:留下 dW,把 dX 往下传

  4. 反向的因果掩码

    位置 i 的损失只能触及 i 及之前的词元,因为注意力从不往后看。每一行是一个位置的损失;每一列表示它对那个词元嵌入的影响有多强。

    i 处的损失 → j ≤ i 处的嵌入

  5. 穿过十二个块

    梯度到达了每一个块。从第 12 块到第 1 块,残差流上的梯度大小几乎没有缩小,因为每个块都是往残差流上相加,而梯度原样穿过这些加法流回去。只有在嵌入处它变大了,因为 LayerNorm 对这些很小的向量所除的 σ 很小。

    残差相加让梯度保持有力

  6. 扰动一下来核对

    核对方法:把一个权重扰动 ±0.0001,让模型前向跑两次,看损失怎么变。得到的斜率与反向传播的结果吻合到六位数字。这样每个权重要跑两趟;反向传播只花大约两趟的代价就得到全部 1.24 亿个。

    (L(w + ε) − L(w − ε)) / 2ε = ∂L/∂w

代码

logits = model(ids[:, :-1])
loss = F.cross_entropy(logits.flatten(0, 1), ids[:, 1:].flatten())
loss.backward()                  # fills .grad of every parameter, last layer first
# what it does for each linear layer y = x @ W + b, given dy = ∂L/∂y:
W.grad += x.T @ dy               # summed over positions (and the batch)
b.grad += dy.sum(0)
dx = dy @ W.T                    # passed down to the layer below
# check one weight numerically
w[i] += eps; up = loss_fn(); w[i] -= 2 * eps; down = loss_fn(); w[i] += eps
assert abs((up - down) / (2 * eps) - w.grad[i]) < 1e-6

延伸阅读