跳转到内容

基础原理第 13 篇,共 13 篇

反向传播与参数更新

从 loss 的梯度追到共享参数,区分 forward、backward 和 optimizer step

更新于 阅读约 2 分钟

Forward 计算预测和 loss。Backward 计算 loss 对参数的导数。Optimizer 使用这些导数更新参数。这是三个不同的动作。

先读:CE 的 p−q、模型参数的位置。

设一个位置 hidden 为行向量 h,LM Head 权重 W 使用数学布局 [D,Vocab],z=hW。若 g=∂L/∂z,则:

∂L∂W=h⊤g,∂L∂h=gW⊤\frac{\partial L}{\partial W}=h^\top g,\qquad \frac{\partial L}{\partial h}=gW^\top

g 从 CE 得到,接着回传经过 final Norm、各层 Block 和 Embedding。共享 W 在三个位置都被使用,所以 W 的总梯度累加各位置的贡献;均值归约也要计入相应缩放。

因果 mask 禁止未来位置的信息进入当前预测。它不会禁止合法历史路径上的梯度回传。训练不经过 argmax 来计算 CE:argmax 把连续分数变成离散选择,不提供这里需要的可用梯度路径。

假设某个标量参数 θ=0.3,当前批次得到 ∂L/∂θ=−0.2。最简单的 SGD、学习率 η=0.1:

θnew=θ−η∂L∂θ=0.3−0.1(−0.2)=0.32\theta_{new}=\theta-\eta\frac{\partial L}{\partial\theta}=0.3-0.1(-0.2)=0.32

该例不是完整模型训练。真实参数同时接收多个位置/样本的梯度;AdamW 还使用动量、平方梯度估计和独立的权重衰减,更新不能简单解释成“直接把正确概率加一点”。一步更新也不保证每个样本的 loss 都下降。

optimizer.zero_grad() # 清理旧梯度,除非有意做梯度累积
outputs = model(...) # forward;输入/标签 shift 由接口约定决定
loss = outputs.loss
loss.backward() # 计算并累加到 parameter.grad
optimizer.step() # 根据梯度更新参数

backward() 本身通常不改参数。连续多次 backward 可以累积梯度;若忘记清理,含义会改变。使用梯度累积时应同时处理 loss 缩放。低精度、梯度裁剪、激活重算和优化器状态是进一步的工程细节,不是理解链式法则的前置。

普通生成用固定参数做 forward。model.eval() 关闭 Dropout 等训练行为;no_grad() / inference_mode() 控制自动求导记录。二者用途不同,LayerNorm 不使用运行均值/方差。

代码入口:Karpathy ng-video-lecture、The Annotated Transformer。