最近在重新跟 MIT 6.S191 和 CS336,发现自己处在一个有点尴尬的状态:RNN 的 forward 随时能默写,hidden state 一步一步往前推,没什么好想的。但如果有人接着问"那它到底怎么 backward",我能说出来的只剩几个词——BPTT、沿时间反向传播、梯度消失——再往下就模糊了。
Forward 记得很清楚,backward 停在术语层面。
training loop 那篇重新理过 forward / backward / optimizer 的职责划分,RNN state transition 那篇理过 hidden state、output、prediction 的区别和 weight sharing,next-token prediction 那篇解释了为什么每个 timestep 都能产生一个 loss。这篇把这几块拼起来,回答一个问题:
RNN 到底是怎么训练的?所谓 Backpropagation Through Time,到底"特殊"在哪里?
一句话版本:它并不特殊。RNN 没有发明新的反向传播算法,把 recurrent computation 沿时间展开以后,它仍然是一张普通的 computation graph。真正特殊的地方只有两处——同一套参数会在不同 timestep 被重复使用,hidden state 建立了跨 timestep 的 dependency。
回到最普通的 training loop#
先把最普通的一轮训练放回来:
1optimizer.zero_grad()
2pred = model(x)
3loss = loss_fn(pred, y)
4loss.backward()
5optimizer.step()loss.backward() 沿 computation graph 反向遍历,把 loss 对每个 parameter 的偏导写进 .grad;optimizer.step() 拿着这些 gradient 更新参数。两件事的边界在 training loop 那篇里写过,这里不重复。
RNN 用的是同一个循环。没有 rnn.backward(),也没有另一套 optimizer。所以问题不能停在"RNN 怎么训练",得问得再具体一点:
当同一个 cell 沿时间不断重复时,computation graph 会变成什么样?gradient 又是怎么沿这张图流动的?
参数共享:一份 parameter,多个 use site#
从最简单的 recurrence 开始:
$$h_t = f(W_h h_{t-1} + W_x x_t)$$沿时间展开 3 步,就得到一张小图:
1h0 → h1 → h2 → h3
2 W W W这里有三件容易糊在一起的事:
- \(h_1, h_2, h_3\) 是不同时刻的 activation,数值不同,在 graph 上是不同的 node;
- 每个 timestep 是一次独立的计算,有自己的输入和输出;
- 但三处计算用的是同一个 parameter tensor \(W\)。
真实模型里没有 \(W_1, W_2, W_3\) 三份参数,也不存在"第一步的 W"这种对象。只有一份 \(W\)。(严格说,\(W_h\) 和 \(W_x\) 是两份不同的参数,但各自都在所有 timestep 复用;这里为了简洁只画一个 \(W\),结论对两份参数都成立。)
这个区分在 forward 时看起来像文字游戏,到 backward 时就是关键。我后来习惯用一个说法:
parameter identity 和 parameter use 是两回事。参数空间里只有一份 \(W\),但 computation graph 上存在多个 \(W\) 的 use site。
可以记成 W@t1、W@t2、W@t3——同一个对象,在图上被用了三次。
RNN state transition 那篇从建模角度讲过 weight sharing 的意义:transition rule 在时间上不变,参数量不随序列长度增长。这篇关心的是它在 backward 方向上的后果。
标量 demo:gradient 为什么会自动累加#
先完全离开 RNN,看一个标量乘法:
1import torch
2
3w = torch.tensor(2.0, requires_grad=True)
4x1 = torch.tensor(1.0)
5x2 = torch.tensor(2.0)
6x3 = torch.tensor(3.0)
7
8y = w * x1 + w * x2 + w * x3
9y.backward()
10print(w.grad)结果是 tensor(6.)。这里没有新数学:\(y = w x_1 + w x_2 + w x_3\),所以 \(\frac{\partial y}{\partial w} = x_1 + x_2 + x_3 = 6\)。
重点看 autograd 怎么处理它。forward 时,同一个 leaf tensor w 在 graph 上出现了三次;backward 时,每个 use site 都会产生一份 gradient contribution;三份 contribution 指向的是同一个 leaf,于是被累加进 w.grad。
同一个 leaf 被使用多次,autograd 会把所有指向它的 gradient contribution 加起来。 这就是 shared parameter 上 gradient 行为的全部机制。
从加法迁移到 recurrence#
把加法换成 recurrence:
1import torch
2w = torch.tensor(2.0, requires_grad=True)
3h0 = torch.tensor(1.0)
4h1 = w * h0
5h2 = w * h1
6h3 = w * h2
7
8for h in (h1, h2, h3): h.retain_grad()
9h3.backward()
10print(h1.grad, h2.grad, h3.grad, w.grad) # 4 2 1 12打印的四个值依次是 h1 / h2 / h3 / w 的 gradient。可以手算核对:\(h_1 = 2\),\(h_2 = 4\),\(h_3 = 8\)。
h3.grad = 1,因为它就是这次的 loss;h2.grad = ∂h3/∂h2 = w = 2;h1.grad = ∂h3/∂h2 · ∂h2/∂h1 = w² = 4;w.grad = 12。
这里需要解释 retain_grad()。PyTorch 默认只为 leaf tensor 保留 .grad;h1, h2, h3 是计算产生的中间结果(non-leaf tensor),backward 走完以后 gradient 不会留在它们身上。想观察中间节点的 gradient,必须显式调用 retain_grad()。h0 本身是 leaf,但它的 requires_grad=False,所以不参与 gradient 计算。
w.grad = 12 可以拆成三个 use site 的贡献:
- 经过 \(h_1\):\(\frac{\partial h_3}{\partial h_1} \cdot \frac{\partial h_1}{\partial w} = 4 \times h_0 = 4\);
- 经过 \(h_2\):\(\frac{\partial h_3}{\partial h_2} \cdot \frac{\partial h_2}{\partial w} = 2 \times h_1 = 4\);
- 经过 \(h_3\):\(\frac{\partial h_3}{\partial w} = h_2 = 4\)。
三者相加正好是 12,和直接对 \(h_3 = w^3 h_0\) 求导得到的 \(3w^2 h_0\) 一致。
假设每个 timestep 有自己的 W#
把"同一个 W 被用了三次"和"每个 timestep 有独立 W"两种情况放在一起,累加关系会更直观:
shared parameter vs 假设的 W1/W2/W3
1import torch
2
3def run(shared):
4 if shared:
5 w1 = w2 = w3 = torch.tensor(2.0, requires_grad=True)
6 else:
7 w1 = torch.tensor(2.0, requires_grad=True)
8 w2 = torch.tensor(2.0, requires_grad=True)
9 w3 = torch.tensor(2.0, requires_grad=True)
10
11 h0 = torch.tensor(1.0)
12 h3 = w3 * (w2 * (w1 * h0))
13 h3.backward()
14
15 if shared:
16 return w1.grad.item()
17 return [w1.grad.item(), w2.grad.item(), w3.grad.item()]
18
19shared_grad = run(shared=True)
20separate_grads = run(shared=False)
21
22print(shared_grad) # 12.0
23print(separate_grads) # [4.0, 4.0, 4.0]
24print(sum(separate_grads)) # 12.0假设版本里,每个 \(W_t\) 单独拿到的 gradient 都是 4;它们相加等于 shared 版本里那唯一一份 \(W\) 的 gradient 12。
写成公式:
$$\frac{\partial L}{\partial W} = \sum_t \frac{\partial L}{\partial W_t}$$这里的 \(W_t\) 只是分析用的记号,表示"假设第 t 个 timestep 有独立参数时,它单独会拿到多少 gradient"。真实模型里没有 W1/W2/W3,autograd 也不需要真的复制参数,它只是把所有 use site 的 contribution 累加到同一个 leaf 上。
一个容易混淆的点:这里说的 accumulation 发生在一次 backward 内部,是同一个 parameter 的多个 use site 产生的 contribution 求和。它和 training loop 里跨 mini-batch 累加 .grad 不是一回事——后者发生在多次 backward() 之间,虽然利用的都是"往 .grad 里加"这个行为。
BPTT:普通 backprop,只是图沿时间展开#
回到 RNN。
Forward 的时候,随着 sequence 一步步输入,computation graph 已经沿时间建好了。loss.backward() 做的事情只有一件:沿着这张已经存在的图反向遍历。
所以:
Backpropagation Through Time 并没有发明新的 backward。它就是对 time-unrolled computation graph 做普通的 backprop。
对 autograd 来说,“time” 不是一个特殊概念。它看到的只有依赖关系:
1A depends on B
2B depends on C
3C depends on D至于 B、C、D 是不同层的 neuron,还是同一个 cell 在不同 timestep 的 state,autograd 不关心。反向遍历时它只做一件事:把 upstream gradient 乘以当前 node 的 local derivative,继续往前传。
“Through Time” 这个名字,是人从 RNN 的语义角度给这张图起的。图上的 backward 本身,和普通 feedforward 网络没有区别。它唯一"特殊"的地方,还是前面那两条:图上的参数被复用了多次,图上的节点沿时间形成了依赖链。
为什么 future loss 会训练 earlier hidden state#
前面的 demo 里 loss 就是最后一个 state。真实的语言模型不是这样:每个 timestep 的 hidden state 都会经过 output projection 得到一个 prediction,每个 prediction 都产生一个 loss:
$$L = L_1 + L_2 + \cdots + L_T$$看 \(h_1\)。它不只影响 \(L_1\)。后面所有的计算都以 \(h_1\) 为起点:
1h1 → h2 → L2
2h1 → h2 → h3 → L3所以反向传回 \(h_1\) 的 gradient,是所有这些路径贡献的叠加:
$$\frac{\partial L}{\partial h_1} = \frac{\partial L_1}{\partial h_1} + \frac{\partial L_2}{\partial h_1} + \frac{\partial L_3}{\partial h_1}$$每一项的来源不太一样。\(\frac{\partial L_1}{\partial h_1}\) 来自当前时刻的 prediction;\(\frac{\partial L_2}{\partial h_1}\) 是 \(L_2\) 的 gradient 先传到 \(h_2\),再通过 state transition 的 backward 继续传到 \(h_1\);\(L_3\) 同理,只是路径更长。
hidden state 不只是当前 prediction 的中间变量,它还是未来 computation 的 state。 因此未来预测做错以后,future loss 必须能够告诉 earlier hidden state:你当时应该保存什么信息,才能让我现在预测正确。这就是 recurrent model 里的 temporal credit assignment——future loss 在监督 earlier state。
从标量导数到 Jacobian#
上面说"gradient 继续往前传",在数学上到底是什么?
先从标量开始,因为一维情况下不用引入 Jacobian 也能看清结构。取最简单的一维 recurrence:
$$h_t = w h_{t-1}$$每一步的局部导数都是 \(w\):
$$\frac{\partial h_t}{\partial h_{t-1}} = w$$跨两个 timestep,链式法则把局部导数乘起来:
$$\frac{\partial h_3}{\partial h_0} = \frac{\partial h_3}{\partial h_2} \frac{\partial h_2}{\partial h_1} \frac{\partial h_1}{\partial h_0} = w^3$$标量的情况到这里就完了。但真实 hidden state 是 vector。当函数从 vector 映射到 vector 时,一个标量导数不够用了——输入的每个分量都可能影响输出的每个分量。描述这件事的对象是 Jacobian matrix。
Jacobian 可以这样理解:它是多维函数在当前点附近的"变化传播规则"。 输入有一个小扰动 \(\Delta x\),输出端的扰动近似为:
$$\Delta y \approx J \Delta x$$直白一点说,它是一张表:哪个输入分量会影响哪个输出分量,以及影响多大。
具体到 RNN 的 state transition:
$$h_t = \tanh(W_h h_{t-1} + W_x x_t + b)$$每个 timestep 的局部 Jacobian 是:
$$J_t = \frac{\partial h_t}{\partial h_{t-1}}$$它大致受到两部分影响:recurrent weight \(W_h\) 和 activation 的导数 \(\tanh'(\cdot)\)。关于 Jacobian 本身的推导,从局部线性化到 RNN 那篇已经从直觉上写了一遍,这里不重复。这篇只需要它的一种用途:描述跨时间的 gradient 传播。
Jacobian 连乘与 vanishing / exploding#
把每个 timestep 的局部 Jacobian 乘起来:
$$\frac{\partial h_T}{\partial h_0} = J_T J_{T-1} \cdots J_1$$这个式子说的是:long-range gradient propagation 本质上是很多局部 Jacobian 的连续复合。
标量直觉很好记:\(0.5^{100}\) 趋近 0,\(1.5^{100}\) 是个天文数字。用 10 步的标量 recurrence 就能直接看到这个效应:
1import torch
2def grad_after(w_val, steps=10):
3 w = torch.tensor(w_val, requires_grad=True)
4 h = torch.tensor(1.0)
5 for _ in range(steps): h = w * h
6 h.backward()
7 return w.grad.item()
8
9print(grad_after(0.5), grad_after(1.5)) # 0.0195... 384.43...同样只走 10 步,\(w=0.5\) 时 gradient 已经衰减到 0.02 以下,\(w=1.5\) 时放大到接近 400。步数再多一些,差距会拉开得更夸张。矩阵世界里,每个 \(J_t\) 都可能在某些方向上压缩扰动、在另一些方向上放大扰动,连续很多次之后,这种效应会指数积累。越早的 state,梯度要穿过的路径越长,情况越不利——这就是 vanishing / exploding gradient 的来源。
对 vanilla RNN 来说,tanh 让事情更麻烦一点:
$$\tanh'(x) = 1 - \tanh^2(x)$$它的最大值是 1,并且在 saturation 区域接近 0。也就是说,梯度每一步不仅要经过 recurrent transformation,还要乘上一个不大于 1、很多时候接近 0 的因子。这让 long-range gradient propagation 更难。
需要避免一个过度简化:不能简单地说"\(W_h\) 小于 1 就梯度消失、大于 1 就爆炸"。对矩阵 RNN 来说这个判断太粗糙了,一个 Jacobian 在不同方向上的行为可能完全不同,真正要看的是它在整条连乘里的复合效果。但结构上的结论是确定的——传播路径越长,复合后的 Jacobian 越难保持稳定。详细一点的讨论在 那篇 Jacobian 笔记里。
Truncated BPTT:截断的是 gradient,不一定是 state#
完整 BPTT 面对长 sequence 有几个现实问题:
- computation graph 很长,backward 要沿整条链走;
- forward 的 activation 要保存到 backward 用,内存随序列长度增长;
- 那条 Jacobian 连乘也长得难以优化。
工程上的做法是把 sequence 切成若干 chunk,每个 chunk 内部正常 forward / backward,chunk 之间把 graph 切断:
1[chunk 1] [chunk 2] [chunk 3]在 PyTorch 里,这个切断通常就是一行:
1h = h.detach()detach() 保留 hidden state 的数值,但把它从当前 computation graph 上摘下来。下一个 chunk 拿到的 h,数值上仍然是连续的 state,forward 行为不中断;但 backward 走到它就会停住——从 autograd 的视角看,它已经是一个新的 leaf,不再依赖之前的图。
state continuity 不等于 gradient continuity。
- hidden state 的数值可以继续传到下一个 chunk;
- 但 future loss 的 gradient 再也回不到之前的 chunk,之前的 state 收不到来自未来的监督。
所以 TBPTT 的 tradeoff 不是简单的"效果 vs 稳定性",更准确的说法是:
1long-range credit assignment
2 vs
3memory / computation / optimization difficultyTBPTT 的 window length 本质上定义了一个 optimization horizon,或者说 credit-assignment horizon:只有落在这个窗口内的 future loss,才能参与对应 state 的训练。
顺带澄清一下 TBPTT 和 vanishing gradient 的关系。截断确实会让 Jacobian 链变短,训练通常更稳定,但它的首要目的不是"解决梯度消失"。TBPTT 不修复长距离梯度,它只是把超出窗口的依赖从训练目标里拿掉了。
Gradient clipping#
如果训练中出现 exploding gradient,最常见的处理是 clipping:
1torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)clip_grad_norm_ 通常不是逐元素把 gradient clamp 到某个区间,而是看整个 gradient vector 的 norm。如果 norm 超过阈值,就把整个 vector 按比例缩小。比如:
1g = [6, 8]
2||g|| = 10
3max_norm = 5
4=> g becomes [3, 4]缩放后方向不变,模长回到阈值以内。gradient clipping 主要缓解 exploding gradient;它不能解决 vanishing gradient,也不能真正解决 long-range dependency。 它更像一个防止单步更新过大的安全阀。
LSTM / GRU 想解决的是什么#
到这里,vanilla RNN 真正困难的地方就清楚了。不是"它没有 hidden state",而是信息和梯度都必须沿 recurrent chain 一步一步穿过大量 nonlinear transformation。每一步都要经过 \(W_h\) 和 \(\tanh\) 的局部 Jacobian,走几十步以后就很难保持稳定。
LSTM / GRU 的设计价值之一,就是在 state 传播的路径上构造更适合长距离信息和 gradient 流动的通路——比如让某些方向上的 Jacobian 更接近恒等映射。具体怎么用 gate 做到这一点,留到后面单独写。
结论#
BPTT 不是一种新的 backprop。它只是普通 backprop 作用在 time-unrolled computation graph 上;“time” 对 autograd 来说不是特殊概念。
RNN 只有一份 parameter,但 graph 上有多个 use site。每个 use site 产生一份 gradient contribution,最终累加到同一个 leaf 上。
hidden state 是未来 computation 的 state,所以 future loss 会监督 earlier state;长距离梯度传播是一串局部 Jacobian 的复合,vanishing / exploding 是这个结构的自然结果。TBPTT 切断的是 gradient continuity 而不是 state continuity,clipping 只负责控制爆炸。
参考资料#
- MIT 6.S191: Introduction to Deep Learning — https://introtodeeplearning.com/
- Stanford CS336: Language Modeling from Scratch — https://stanford-cs336.github.io/spring2025/
- PyTorch autograd mechanics — https://pytorch.org/docs/stable/notes/autograd.html
- PyTorch clip_grad_norm_ — https://pytorch.org/docs/stable/generated/torch.nn.utils.clip_grad_norm_.html
- 从 Forward 到 Optimizer:再看一遍 Training Loop
- 重新理解 RNN:Hidden State、Output 和 Prediction 的区别
- Next-Token Prediction:语言模型的训练目标从哪来
- 从局部线性化到 RNN:理解 Jacobian 连乘