RNN 与 LSTM 入门
2017 年之前,序列模型是要携带状态的。这一篇把它拆开:算出这个状态的循环、几十步之后就把它杀死的那个导数连乘、为了让这个连乘活下来而造的门,以及终结了这一整族模型的那些量。
一次一个 token
循环网络只有一个循环和一份记忆。这一页后面所有东西都是它的推论。
它从左往右读,只留一个向量,每一步把这个向量和新 token 混在一起再压一下:h_t = tanh(W_h·h_{t−1} + W_x·x_t + b)。我们让它跑在一维上,每一帧都能用手算验。
8 个 token是给定的,所以一开始就是灰的;已经算出来的状态在它们上面,循环此刻正在算的那个状态就在你手里。拖滑块把 t 往前走:
注意滑块右边什么都没有。那些格子是虚线,因为 h_4 不是网络藏起来不给你,而是根本没人算过。这就是那条不变式,值得写成一句断言:每一步开头,h_t 只是 x_0 … x_t 的函数,跟后面的东西无关。
那个压缩不是装饰。它把任何预激活值折进 −1 到 1 的开区间,而你站的那一点上它的斜率,正是下一节从头到尾要讲的东西。拖着点沿曲线走:
看那条切线。z = 0 处它最陡,读数给出 tanh' = 1.000;到 z = 2 曲线已经平到 0.071。饱和的单元就是输出不再回应输入的单元 —— 也是梯度不再往回走的单元。
W_h 只有一个。同一个数把每个状态乘到下一个状态上,这就是一套权重能读任意长句子的原因。动一下它,8 个一起动 —— 试试 :
共享这个权重让参数量跟序列长度无关:循环一个 d×d 矩阵、输入一个、一个偏置。1.40 时最后那个状态读作 0.841; 时它读作 0.310,正好是 tanh(0.8 × 0.4),只剩自己那个 token。
答案要再过一个矩阵,才从那些状态里出来。分类只从最后一个状态读一个输出,序列标注则每个位置读一个。切换看看:
不管哪一种,读出头看到的只有它底下那个状态。第一个 token 对最后那个答案的所有贡献,都得穿过 7 次W_h 和 7 次压缩才能到。这趟路对它做了什么,就是下一节。
梯度为什么死在循环里
训练要沿着链往回走,每跳一次乘一次。一堆小于 1 的数乘起来不是一个小数,而是根本不成其为数。
要从最后一个 token 学到东西,损失必须走到第一个。它沿同一条链倒着走,每一跳都被乘上 ∂h_t/∂h_{t−1} = tanh'(z_t) · W_h。灰色那条曲线就是它前一半的斜率,而那根柱子是整整一跳 —— 斜率乘上那个权重:
斜率只在 z = 0 处等于 1,两边都往下掉,所以那根柱子处处小于 W_h,而且只要这个单元在干活就小得多。这是一跳。让梯度连着往回走 8 跳:
7 跳就把它从 1 带到 0.057。没有任何东西出错 —— 这里每一个因子都是一个正常网络的正常导数。麻烦只在于它们有 7 个,而句子比 8 个词长。
那就单看这些因子,放到 40 步的链上。虚线是 × 1,缩小和放大的分界。拖 W_h,再换一下激活函数:
因为 tanh' 永远不超过 1,因子就永远不超过 W_h;又因为这些状态从不停在 0 上,它严格小于 W_h:默认的 0.90 上最大的是 0.884。把 W_h 推到 ,柱子还是够不到那条虚线 —— 饱和就是刹车,它把峰值压在 0.616。切到 ReLU,刹车没了:它导数在开着的地方正好是 1。
把这些因子乘到一起 —— 这次放到 200 步的链上 —— 就是真正到达 h_0 的那个数。纵轴是对数的,每一条网格线比下面那条高 24 个数量级,所以一个恒定的因子画出来是一条直线:
在默认设置上,这个积在第 161 跳穿过下面那条线(float32 最小的正规数),终点是 9.8e-48。切到 ReLU 并设 ,它反过来往上爬,在第 189 跳穿过上面那条线。越过它的 float32 就是无穷,碰到无穷的损失就是 NaN。
在一维里刹车总是赢;在多维里,W_h 最大的奇异值可以作用在一个没人饱和过的方向上,积是真的会跑掉的。标准的护栏是 Pascanu、Mikolov 和 Bengio 2013 年那个缩放,而这一步裁剪前的范数是 84.0,阈值是 5 —— 把它拖回阈值以下:
因为裁剪是天花板不是地板,它对另一半毫无作用。消失是安静地失败的:不报错、不出 NaN、损失照常下降,模型悄悄学会只根据最后几个 token 预测 —— 只有那几个的梯度活着走完了全程。
让这个乘积活下来的门
LSTM 并没有把乘法变小。它另修了一条几乎什么都没有的路。
Hochreiter 和 Schmidhuber 1997 年的答案是再加一份记忆 c,它的更新是 c_t = f · c_{t−1} + i · ĉ_t —— 乘一个数,再加一下。c_{t−1} 和 c_t 之间没有权重矩阵,也没有压缩函数。
这张图跟 §01 只差一个元件。token 行一样,格子一样,只有连接线不同:细胞态走的是一条素线,而这一步写进去的东西从下面来。往前走一遍:
把这条连接线跟 §01 那个箭头比一比。那边一跳是 tanh'(z)·W_h,两个网络很难控制的东西;这边是遗忘门,一个网络每一步专门算出来的数。这一次替换就是全部的主意。
两个门都是比例,所以一次更新就是两段首尾相接的长度:旧状态活下来的那份,再加候选被写进去的那份,以及被扔掉的余数。先动遗忘门,再动输入门:
把它设成 ,被扔掉的那段余数就没了:细胞态变成一个纯累加器。设成 0,旧状态一步就没了 —— 一句话说完、下一句跟它毫无关系的时候,网络要的正是这个。
这能解掉 §02 的原因是:从 c_0 到 c_k 只有一条路,而那条路是一串遗忘门的积 —— f^k,别的什么都没有。对照的是玫红那条,上一节那个普通 RNN 的积,画在同一根线性轴上。拖 f:
f = 0.95 时,100 步之后还剩 0.006 —— 很小,但是一个浮点数放得下、优化器用得上的数。那个 RNN 的积第 9 步就跌到 0.01 以下,所以它后面一路贴着横轴。而且遗忘门不是架构里的常数:网络会学它,逐维度、逐步地学。
经典的坑就在这里。初始化时权重接近 0,所以 f 就是它的偏置说了算,而偏置为 0 意味着细胞每一步把自己砍一半。这条曲线是 20 步之后 c_0 还剩多少,横轴是那个偏置 —— 纵轴取对数,好把 8 个数量级放进来:
默认这一帧就是那个 bug。bias = 0 给出 f = 0.500,20 步之后只剩 9.5e-07 —— 一个天生就带着它要解决的梯度消失的 LSTM。拖到 ,同样 20 步还留着 0.002。写 b_f = 0,它就毫无缘由地训不好;写 b_f = 1 —— Jozefowicz、Zaremba 和 Sutskever,2015 —— 它就不会。
门从来没解决的两件事
门控解决了梯度。终结循环时代的那两个性质,它一个也没碰。
第一个是容量。细胞态活得再好,它也只是一个定宽向量,门不会让它变大一点。
下面这个状态宽 512 个数,永远不变;上面是它要总结的那些 token,而第一个 token 分到的那份就是左边那一块。往右拖,把序列拉长:
4 个 token 的时候,每一份是 128 个数。到 时只有 2.00,上面那排刻度也糊成了一片纹理。这不是可以调掉的 bug —— 这就是「用定长向量总结前文」这句话的含义,也是为什么问一个 LSTM 300 个 token 之前的具体细节,它会答得像模像样而不是答对。
第二个更糟,而且它讲的是机器不是模型。行是位置,列是墙钟步,某个位置的活儿发生时它那格才亮。把步往前拖,看看这张格子里到底有多少在忙:
只有对角线能亮。位置 5 得等位置 4 算完才能开始,所以 8 个位置要花8 个顺序步,走到最后一步,64 格里也只有 8 格干过活。GPU 有上万条通道;把不同序列打成 batch 能填掉一些,但一个序列内部的活儿是一条链,而链没有宽度。
GRU(Cho 等,2014)是流行的省钱版:两个门而不是三个,而且输入门不自由 —— 它被强制成 1 − f,于是留下什么和写进什么变成同一个决定。拖一下:
因为两段长度必须凑满整条轨道,GRU 没法在同一步里既保住旧状态又使劲写;LSTM 可以,而那正是多出来的那个门买到的唯一东西。它省下的是一个门的矩阵 —— 在同一个固定预算上,一根条对另外两根:
每种 cell 都是整数个 2d² + d 块:RNN 一块、GRU 三块、LSTM 四块 —— d = 1024 时是 2.10 M、6.29 M 和 8.39 M。上面两页并不因此改变:状态还是那一个定长向量,循环还是循环。
把循环扔掉换来了什么
三个量变了,而且每一个都是 Vaswani 等 2017 年那篇 Table 1 里的一列。
自注意力一次性从所有位置算出每个位置:给每一对打分、softmax 一下、加权求和。没有状态要携带,所以任何一个位置都不用等别人算完才开始。
这是 §04 那张格子,加了个开关。循环层每列只点亮一格;注意力一次点亮一整列。先把序列拉长,再切换层的类型:
顺序深度就是有东西的列数:左边是 n,右边是1,滑块能到的每一个 n 上都成立。这就是「GPU 只用了百分之几的宽度」和「GPU 跑满」之间的差别 —— 也是 2017 年那个 base Transformer 能在 8 块 P100 上 12 小时训完,而它打败的那些 LSTM 系统要在多得多的硬件上跑好几天的原因。
第二个量是距离。每一级台阶是两个 token 之间的一次乘法,所以这条路的高度就是它们之间隔着几次乘法。先拖目标 token,再切换层的类型:
切一下层的类型,看那个高度塌下去。循环在 token 0 和 token k 之间放了 k 次乘法 —— 正是 §02 讲的那个积 —— 而注意力在那儿只放一条边,无论 k 是多少。最大路径长度从 O(n) 变成 O(1),衰减也就跟着没了。
第三个是解码器被允许看什么。RNN 编码器只交出最后一个状态,所以整个源句必须挤过一个向量;注意力让每个输出位置读到每一个输入位置。切换一下连接方式:
12 个源 token 时,瓶颈不管源有多长都只给解码器 512 个数;直连给它 6,144 个,24 个 token 时是 12,288 个。 Bahdanau、Cho 和 Bengio 2014 年看到了这一点,把注意力接在 RNN 上; 2017 年把注意力留下、把 RNN 删了。
这些都不是白拿的。循环层做 n·d² 次乘加 —— n 步、每步一个 d×d 矩阵 —— 而注意力层做 n²·d 次,每对一次打分。下面两根轴都是对数的,所以两条都是直线,注意力那条更陡。拖 d:
陡的那条在 n = d 处穿过缓的那条,那里两边代价相等。 时,注意力在 512 token 以下更便宜、以上更贵。
速查,以及回来的路
cell 的完整写法、这个循环要多少钱,以及实际训练到底往回走多远。
# keep · write · what · show f = sigmoid(W_f[h,x] + b_f) i = sigmoid(W_i[h,x] + b_i) g = tanh(W_g[h,x] + b_g) o = sigmoid(W_o[h,x] + b_o) # the only path back to c_0 c = f * c + i * g h = o * tanh(c)
这个循环不会被训到头。沿时间反向传播要留下链上每个激活值,所以训练时它被切成一个窗口:只有窗口里的跳拿得到梯度,切口以外的一切什么都拿不到。拖那个窗口:
这一刀切掉的不是信息 —— 前向依然把状态运过去 —— 切掉的是归因:窗口之前的东西不会因为末端的误差被追责。
为什么这个界是紧的
真正决定架构选型的是顺序深度和路径长度,两者都是 n 对 1。注意力用每层 n²·d 买下它们。
上线服务的时候,这笔账会变成钱。解码器要为它见过的每个位置留着KV cache,而循环状态不管多长都是那么大。把上下文拖过三个数量级:
在 上,一个 32 层、宽 4096 的 fp16 模型要留64 GB,对面是256 KB —— 差 262,144 倍,还没算权重就占掉一张 80 GB 加速卡的大半。分组查询注意力把它砍到四分之一。
所以循环正在那个交点告诉我们的地方回来: Mamba(Gu 和 Dao,2023)和 RWKV 跑的是循环状态,但更新被安排成训练时能并行算。对今天 LLM 在服务的窗口,还是注意力那三个胜利说了算 —— 而这一页值得带走的东西比它们都老:只有一次乘法的路不会衰减。