Zhanbo's Blog
Back to home

CS546 Learning Notes:从梯度裁剪到 LSTM、GRU,串起一次参数更新

Algorithm/Learning Notes
2026-09-13
11 min read
0 reads
AIMachine Learning

In brief

上次学到 RNN 的 BPTT,知道梯度沿时间反传时会连续相乘,也算出了 0.25=0.000320.2^5=0.00032。但看到 gradient clipping 时,我发现自己还没有真正理解:梯度爆炸到底有什么问题,为什么要裁剪? 这次从这个问题出发,逐步理解 LSTM 和 GRU 的门控机制,最后把 BPTT 与 Optimizer 的分工接起来。最有收获的地方,是分清几组容易混淆的概念:记忆和输出、门值和参数、沿路径相乘和多路径相加。

1. 梯度爆炸:为什么大梯度可能把训练带偏?

先看不带动量、也不带 weight decay 的普通 SGD:

wnew=wηg,g=Lww_{\text{new}}=w-\eta g,\qquad g=\frac{\partial L}{\partial w}

ww 是参数,η\eta 是学习率,gg 是梯度。假设 w=1w=1η=0.1\eta=0.1

梯度更新后的参数更新幅度
2210.1×2=0.81-0.1\times2=0.80.20.2
1000100010.1×1000=991-0.1\times1000=-99100100

我当时算出两次更新幅度相差 99.899.8,也意识到第二种情况一步跨得太远。

梯度描述当前位置附近的变化。沿负梯度方向走足够小的一步,通常能降低可微 loss;走很远却没有这个保证。因此,大梯度配上固定学习率,可能让参数跨过合适的位置,反而增大 loss。

RNN 为什么容易出现这种情况?用一个去掉输入和激活函数的标量例子:

ht=wht1,htht1=wh_t=w h_{t-1},\qquad \frac{\partial h_t}{\partial h_{t-1}}=w

若末端梯度为 11w=3w=3,反传三步得到:

Lh0=1×33=27\frac{\partial L}{\partial h_0}=1\times3^3=27

经过十步,同样的连乘因子是 310=590493^{10}=59049。这里算的是对 hidden state 的梯度;计算参数梯度时,也会用到这些反传信号,所以参数梯度可能随之变大。

“时间步增加”指反向路径跨过更多序列位置,不是说训练轮数增加就必然爆炸。实际向量 RNN 涉及 Jacobian 连乘,也不是所有路径都会持续放大。这个标量例子展示的是一种机制。

2. Gradient clipping:限制长度,保留方向

我一开始把裁剪理解成“过滤异常值”。更准确的说法是:限制过大梯度带来的影响,而不是丢掉梯度。 大梯度本身也不一定说明数据有问题。

把待裁剪的梯度整理成向量 gg,设最大长度为 c>0c>0,L2 norm clipping 为:

g~={cg2g,g2>cg,g2c\tilde g= \begin{cases} \dfrac{c}{\|g\|_2}g,&\|g\|_2>c\\ g,&\|g\|_2\le c \end{cases}

例如 g=[3,4]g=[3,4],长度为 32+42=5\sqrt{3^2+4^2}=5。阈值为 22 时:

g~=25[3,4]=[1.2,1.6]\tilde g=\frac25[3,4]=[1.2,1.6]

所有分量乘以同一个正数,所以方向保持不变,长度缩成 22。如果直接改成 [2,2][2,2],比例从 3:43:4 变成 1:11:1,方向就不同了;那不是这里的 norm clipping。

课程 Lecture 5 第 44 页介绍了这个方法,所引用的 Pascanu 等人的论文提出用 gradient norm clipping 应对梯度爆炸。

为什么不直接调低学习率?

我最初的想法是:“调低学习率只能压住前面,后面还是会爆炸。”这里需要修正:小学习率会缩小每次更新,不只是前几步。真正的区别在于,它连正常大小的梯度更新也一起缩小。

做法梯度为 22 时的更新幅度梯度为 10001000 时的更新幅度
学习率 0.0010.0010.0020.00211
学习率 0.10.1,裁剪阈值 550.20.20.50.5

以上仍限定普通 SGD。裁剪只在超过阈值时介入,小学习率则统一缩小更新,两者也可以配合使用。

裁剪发生在梯度算完之后,不会取消反传中的连乘。它也解决不了梯度消失:0.000010.00001 小于阈值 55,裁剪后仍是 0.000010.00001

3. LSTM:保留旧记忆,也能写入新内容

梯度消失不只是“步子太小”。后面的 loss 传到早期位置时信号太弱,会让模型难以学会利用很久以前的信息。

LSTM 维护 cell state ctc_t,并通过门控制记忆更新:

ct=ftct1+itc~tc_t=f_t\odot c_{t-1}+i_t\odot\tilde c_t
符号含义
ct1c_{t-1}旧记忆
ftf_tforget gate,旧记忆保留多少
c~t\tilde c_t候选记忆,准备写入的内容
iti_tinput gate,候选内容写入多少

\odot 是逐元素相乘。实际这些量通常都是向量,每个分量有自己的门值;以下先用一个分量手算。

旧记忆为 88,遗忘门为 0.750.75,先保留 66。候选记忆为 0.80.8,输入门为 0.50.5

ct=0.75×8+0.5×0.8=6.4c_t=0.75\times8+0.5\times0.8=6.4

这里 ctc_t 可以超过 11;经过 tanh 的候选记忆分量则在 (1,1)(-1,1) 内。这个 88 的例子假设记忆已经在此前积累形成。

门值是人为固定的吗?

我接着问:“输入门的数值怎么定义?遗忘门是固定值吗?”

它们由输入、上下文和可训练参数共同算出:

it=σ(Wixt+Uiht1+bi)i_t=\sigma(W_i x_t+U_i h_{t-1}+b_i) ft=σ(Wfxt+Ufht1+bf)f_t=\sigma(W_f x_t+U_f h_{t-1}+b_f)

σ\sigma 是 sigmoid,把有限输入映射到 (0,1)(0,1),作为通过比例。比如 sigmoid 的输入分别为 2,0,2-2,0,2 时,输出约为 0.12,0.5,0.880.12,0.5,0.88

这里采用列向量约定:若输入维度为 dd、隐藏维度为 mm,则 xtRdx_t\in\mathbb R^dht1,ct,it,ftRmh_{t-1},c_t,i_t,f_t\in\mathbb R^mWiRm×dW_i\in\mathbb R^{m\times d}UiRm×mU_i\in\mathbb R^{m\times m}biRmb_i\in\mathbb R^m;其他门同理。

各时间步共享同一套参数,但输入与上下文不同,所以门值可以不同。 训练更新的是这些参数;门值是每次前向计算的结果。

4. 输出门:记住了,不等于现在就呈现出来

LSTM 还有 output gate:

ot=σ(Woxt+Uoht1+bo)o_t=\sigma(W_o x_t+U_o h_{t-1}+b_o) ht=ottanh(ct)h_t=o_t\odot\tanh(c_t)

我能算出 0.25×0.8=0.20.25\times0.8=0.2,却又问了一次:“这里的 hth_t 到底代表什么?”

这次区分清楚了:ctc_t 是内部记忆,hth_t 是当前 hidden state。hth_t 可以送到输出层做预测,也参与下一步的门值与候选内容计算;ctc_t 则沿记忆路径继续传递。标准 LSTM 的这些状态关系可对照 PyTorch LSTM 公式

例如语言模型可以计算 logitst=Wyht+by\text{logits}_t=W_yh_t+b_y,再经过 Softmax 预测下一个词。hth_t 本身不是词,也不是概率。

输出门接近 00 时,对应的 hth_t 分量接近零,但 ctc_t 不会因此直接被清空。想保留某条记忆、暂时不呈现、也不写入新内容,可以让对应的:

ft1,it0,ot0f_t\approx1,\qquad i_t\approx0,\qquad o_t\approx0

我还纠正了一个说法:“输出门打开,输出就按照旧记忆来。”实际上,输出门作用于更新后的 ctc_t。刚才记忆已经变成 6.46.4,所以输出门为 11 的理想情况对应 tanh(6.4)\tanh(6.4),不是 tanh(6)\tanh(6)

如果希望记忆基本原样保留,除了 ft1f_t\approx1,还需要新增项接近零,例如 it0i_t\approx0。只让遗忘门接近 11,仍可能继续写入内容。

5. 为什么 LSTM 能缓解梯度消失?

沿着旧记忆直接传到新记忆的路径,每个分量的直接局部导数是对应的 ftf_t 分量。因此当连续五步都为 0.990.99 时,这条路径的梯度因子是:

0.9950.9510.99^5\approx0.951

相比 0.25=0.000320.2^5=0.00032,信号保留得多得多。我的理解是:“ftf_t 接近 11 更能保留旧信息,每步乘 0.20.2 就会逐渐遗忘。”还要补上反向传播:梯度也更容易沿这条直接路径传回早期位置。

这里的“直接”很重要。完整计算图还包含门值和候选内容经 ht1h_{t-1} 产生的依赖,不能把完整导数一概写成遗忘门连乘。LSTM 提供了更有利的传播路径,但不保证梯度永远不消失或不爆炸。

6. GRU:用更新门混合旧状态与候选状态

GRU 不单独维护 ctc_t,而是更新 hth_t。本文采用 PyTorch GRU 的约定

ht=ztht1+(1zt)nth_t=z_t\odot h_{t-1}+(1-z_t)\odot n_t

ztz_t 是 update gate,ntn_t 是候选状态。这里 ztz_t 越接近 11,越保留旧状态;越接近 00,越采用候选状态。

有些资料让 ztz_t 乘候选状态,记号含义正好相反。阅读时要看公式,不能只记门的名字。

例如旧状态为 0.80.8,候选状态为 0.20.2zt=0.75z_t=0.75

ht=0.75×0.8+0.25×0.2=0.65h_t=0.75\times0.8+0.25\times0.2=0.65

我当时总结为“要么保留更多旧状态,要么保留更多新输入”。这里需要把“新输入”改成“候选状态”:它不是原始 xtx_t,而是经过计算的内容,本身也可能包含历史信息。

7. Reset gate 与 forget gate 的区别在哪里?

GRU 的 reset gate rtr_t 出现在候选状态的生成过程。继续采用 PyTorch 写法:

nt=tanh(Winxt+bin+rt(Whnht1+bhn))n_t=\tanh\left(W_{in}x_t+b_{in}+r_t\odot(W_{hn}h_{t-1}+b_{hn})\right)

它乘在旧状态经过变换后的贡献上。rt0r_t\approx0 时,候选内容少参考历史;rt1r_t\approx1 时,允许这部分历史贡献通过。PyTorch 的 reset 放置位置与部分原始实现不同,本文始终使用上面这一版。

“重置”这个名字让我最初猜反了,以为想重新开始就该接近 11。看位置才知道,这里数值表示通过比例。

我随后问:“它和 LSTM 的遗忘门是不是作用相似?”二者都控制历史信息,但位置不同:

控制的位置
LSTM ftf_t旧记忆直接保留到新记忆的路径
GRU rtr_t生成候选状态时,旧状态的贡献
GRU ztz_t最终混合时,旧状态的保留比例

所以在“直接保留旧状态”这一点上,本文约定里的 ztz_t 更像 ftf_t

两个组合帮助我真正分清了它们:rt0,zt1r_t\approx0,z_t\approx1 时,候选内容少参考历史,但最终仍主要保留旧状态;rt1,zt0r_t\approx1,z_t\approx0 时,最终主要采用候选状态,但候选状态仍可能包含加工后的历史信息。允许参考历史,不等于保证原样保留历史。

8. BPTT:什么时候相乘,什么时候相加?

门控参数由任务 loss 学习,不需要给每个门额外标注正确数值。但 BPTT 是不是每反传一步,就更新一次那一步的参数?

通常不是。各时间步共享参数,BPTT 收集并累加梯度贡献,之后再由 optimizer 更新。

问答中给出两个时间步对同一参数的贡献 0.30.30.1-0.1,我误算成了乘积 0.03-0.03。其实应该是:

0.3+(0.1)=0.20.3+(-0.1)=0.2

规则仍然是:沿同一条路径,局部导数相乘;多条路径汇到同一参数,贡献相加。

即使只有第二步的 loss,共享参数 ww 也有不同路径:

h1=F(x1,h0;w),h2=F(x2,h1;w)h_1=F(x_1,h_0;w),\qquad h_2=F(x_2,h_1;w)

一次使用通过 wh2L2w\rightarrow h_2\rightarrow L_2 影响 loss,另一次通过 wh1h2L2w\rightarrow h_1\rightarrow h_2\rightarrow L_2 影响它。分别沿路径求导,再相加。

“不同时间步是不同路径”是一个入门理解。更准确地说,不同时间步的参数使用形成不同路径,而每次使用还可能影响多个后续 loss。所谓某一步的梯度贡献,必须已经包含相关后续传播,不能漏算或重复计算。

9. Optimizer:拿到梯度之后,真正修改参数

BPTT 负责算梯度,Optimizer 负责利用梯度更新参数。最简单的例子仍是普通 SGD:

wnew=wηgw_{\text{new}}=w-\eta g

三个时间步的贡献为 0.3,0.1,0.40.3,-0.1,0.4,总梯度是 0.60.6。若参数为 11、学习率为 0.10.1

wnew=10.1×0.6=0.94w_{\text{new}}=1-0.1\times0.6=0.94

下一轮若梯度为 0.4-0.4

wnew=0.940.1×(0.4)=0.98w_{\text{new}}=0.94-0.1\times(-0.4)=0.98

负梯度让参数增加。更新后的同一套参数,会用于下一次前向计算的各时间步。这里的计算对应 SGD 不启用动量、weight decay 等额外项的情况

Momentum 和 Adam 会进一步利用历史信息决定更新,这次只建立了 Optimizer 的职责与基本 SGD 计算,还没有展开这些算法。

10. 这次串起来的训练流程

一次常见的训练迭代可以顺着读下来:清理上一轮梯度 → 前向计算状态、门值和预测 → 计算任务 loss → BPTT 计算并累加共享参数的梯度 → 必要时做 clipping → Optimizer 更新参数。

实际也有跨 batch 累积梯度、截断 BPTT 等安排,因此不必把“一整条序列结束才更新”当作唯一做法。核心分工不变:求导和更新参数是不同步骤。

这轮我最想记住的三点是:裁剪限制过大的梯度,不能恢复消失的信号;门控让模型学习怎样保留和使用历史,但门值不等于参数;反传时沿路径相乘,汇总共享参数的贡献时相加。

下次再看到一条复杂公式,我会先问:这个量控制的是哪一条路径?它是当前算出来的状态,还是训练要更新的参数?这比只记住门的名字更可靠。

ZB

Zhanbo Chen

Java Backend & AI Agent Developer

Back to home
Comments