2.5 循环与序列模型 2.5.0 循环的公共机制(家族公共部分)
一句话
用一个隐状态 h 把过去的信息带到现在,每个时间步用同一组权重更新它:h_t = f(h_{t−1}, x_t)。
展开 :把循环按时间展开,就是一个很深的前馈网络,每一层共享同一组权重。T 步序列 = T 层网络(上图上半)。
隐状态就是学出来的「状态」 :控制论里,状态是「知道它就足以预测未来、不必回看历史」的那组变量(马尔可夫性)。传感器的单帧观测往往不满足这一点(看一帧图不知道物体在往哪动),RNN 的隐状态把历史压缩进一个固定大小的向量,补上缺失的信息。
输入输出的几种对应
形式
例子
多对一
整段序列 → 一个类别(情感分类、动作识别)
一对多
一个输入 → 一段序列(图像描述)
多对多,同步
每步输入对应每步输出(逐帧标注、语音帧分类)
多对多,异步
读完整段再输出另一段(翻译,见 2.1.4 Seq2Seq)
BPTT 与梯度消失/爆炸 :反向传播沿时间展开后的图走一遍,叫 BPTT(Backpropagation Through Time)。损失对很久以前的隐状态求导,要连乘很多个雅可比矩阵:
如果每一步都有 ‖J_t‖₂ ≤ q < 1(‖·‖₂ 是最大奇异值;比如 W_hh 的最大奇异值乘上 |f′| 的上界小于 1),这个乘积的范数至多是 q 的 (T − k) 次方,随距离指数衰减(梯度消失,学不到长程依赖);如果一直在放大,则可能指数增长(梯度爆炸,训练发散),但最大奇异值大于 1 本身不保证一定爆炸,还要看梯度的方向。
问题
对策
梯度爆炸
梯度裁剪
梯度消失
门控结构(LSTM、GRU),正交初始化,或者换用注意力、SSM
长序列显存不够
截断 BPTT:每 k 步截断一次梯度,隐状态照常往后传
堆叠与双向 :多层 RNN 把下层每个时间步的输出当作上层同一时间步的输入;双向 RNN 用两个方向的 RNN 分别从头到尾、从尾到头读,把两边的隐状态拼起来,每个位置同时看到过去和未来。对还没到来的未来帧,整段双向编码要等数据到齐,不能逐帧实时输出;但只对截至当前时刻的历史窗口做双向编码,或者只往后多看有限几帧、接受一点延迟,都可以在线用,代价是延迟和重复计算。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 import torchtorch.manual_seed(0 ) Q, _ = torch.linalg.qr(torch.randn(32 , 32 )) for rho in [0.9 , 1.0 , 1.1 ]: x = torch.randn(100 , 32 , requires_grad=True ) h = torch.zeros(32 ) for t in range (100 ): h = rho * Q @ h + x[t] h.sum ().backward() g = x.grad.norm(dim=1 ) print (f"ρ = {rho} :最后一步对第 99 / 50 / 0 步输入的梯度 {g[99 ]:.1 e} / {g[50 ]:.1 e} / {g[0 ]:.1 e} " ) rnn = torch.nn.RNN(32 , 32 , batch_first=True ) x = torch.randn(1 , 100 , 32 , requires_grad=True ) out, _ = rnn(x) out[0 , -1 ].sum ().backward() g = x.grad[0 ].norm(dim=1 ) print ("nn.RNN:对第 99 / 95 / 90 / 80 步输入的梯度" , [f"{g[i]:.1 e} " for i in (99 , 95 , 90 , 80 )])
2.5.1 LSTM 长短期记忆网络 〔图:LSTM〕
一句话
在 RNN 里加一条「细胞状态」传送带和三个门;细胞状态用加法更新,梯度可以沿它几乎无损地传很远。
图中 LSTM 的隐藏单元画成「记忆单元」(带圈的蓝色圆)。Hochreiter & Schmidhuber 1997 年提出,遗忘门是 Gers 等 2000 年加上的。
为什么能记得久 :沿细胞状态这条直通路,c_t 对 c_{t−1} 的导数就是遗忘门 f_t(逐元素),不经过权重矩阵和 tanh。门值本身也依赖 h_{t−1},所以完整的导数还有经过各个门的路径,但直通路这一项只受 f 控制:网络学会让 f 接近 1,梯度沿细胞状态往回传时就几乎不衰减。门由 sigmoid 输出 0 到 1 之间的「开度」,开关是学出来的、随输入变化的。
实用细节 :常见的手动初始化是把遗忘门的偏置设成 1,让训练一开始倾向于「记住」;这不是 PyTorch 的默认初始化,下面的代码也保持默认(PyTorch 的门偏置是 bias_ih 和 bias_hh 两份相加,手动设置时按两者之和算)。参数量是同尺寸普通 RNN 的 4 倍。nn.LSTM 在 NVIDIA GPU 上会调用 cuDNN 的融合实现,比手写的逐步循环快得多;CPU 和 Mac 的 MPS 不走 cuDNN。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 import torchimport torch.nn as nndef lstm_step (x, h, c, W_ih, W_hh, b_ih, b_hh ): i, f, g, o = (x @ W_ih.T + b_ih + h @ W_hh.T + b_hh).chunk(4 , dim=-1 ) i, f, o = torch.sigmoid(i), torch.sigmoid(f), torch.sigmoid(o) c = f * c + i * torch.tanh(g) h = o * torch.tanh(c) return h, c cell = nn.LSTMCell(3 , 5 ) x, h, c = torch.randn(2 , 3 ), torch.randn(2 , 5 ), torch.randn(2 , 5 ) h1, c1 = lstm_step(x, h, c, cell.weight_ih, cell.weight_hh, cell.bias_ih, cell.bias_hh) h2, c2 = cell(x, (h, c)) print ("手写 = nn.LSTMCell:" , torch.allclose(h1, h2, atol=1e-6 ) and torch.allclose(c1, c2, atol=1e-6 ))def adding_problem (B, T=50 ): vals = torch.rand(B, T) marks = torch.zeros(B, T) marks[torch.arange(B), torch.randint(0 , T // 2 , (B,))] = 1 marks[torch.arange(B), torch.randint(T // 2 , T, (B,))] = 1 return torch.stack([vals, marks], -1 ), (vals * marks).sum (1 , keepdim=True ) class SeqRegressor (nn.Module): def __init__ (self, kind, hid=64 ): super ().__init__() self.rnn = {"RNN" : nn.RNN, "LSTM" : nn.LSTM, "GRU" : nn.GRU}[kind](2 , hid, batch_first=True ) self.head = nn.Linear(hid, 1 ) def forward (self, x ): out, _ = self.rnn(x) return self.head(out[:, -1 ]) for kind in ["RNN" , "LSTM" , "GRU" ]: torch.manual_seed(0 ) model = SeqRegressor(kind) opt = torch.optim.Adam(model.parameters(), lr=3e-3 ) for step in range (1500 ): x, y = adding_problem(64 ) loss = nn.functional.mse_loss(model(x), y) opt.zero_grad(); loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0 ) opt.step() x, y = adding_problem(2000 ) with torch.no_grad(): print (f"{kind:4s} 测试 MSE: {nn.functional.mse_loss(model(x), y).item():.4 f} " ) print ("瞎猜(永远输出 1.0)的 MSE ≈" , round (1 / 6 , 4 ))
普通 RNN 一直停在瞎猜的水平;GRU 和 LSTM 在几百步后突然「开窍」,损失下降两到三个数量级。
典型例子
语音识别:2015 年前后 Google 语音搜索用 LSTM + CTC
机器翻译:GNMT(2016)是 8 层 LSTM 编码器 + 8 层解码器
手写识别与生成:Graves 2013 用 LSTM + 混合密度输出生成逼真的手写笔迹
OpenAI Five(Dota 2,2018–2019):主体是一个 4096 单元的 LSTM
时间序列预测:Amazon 的 DeepAR
机器人:力/力矩估计、World Models 的记忆模块(MDN-RNN,见 2.10.2)
优势 :经验上能记住几百步以前的信息(具体多远取决于任务和训练设置);梯度稳定;流式推理每步 O(1);数据量中等时比 Transformer 稳
局限 :顺序计算,训练不能沿时间并行,长序列很慢;参数是普通 RNN 的 4 倍;上千步的超长依赖仍不如注意力
适合的数据 :有长程依赖的一维序列:语音、手写轨迹、传感器时间序列、中等长度的文本;中小数据量的时序回归和分类
2.5.2 GRU 门控循环单元与双向 RNN 〔图:GRU〕
一句话
LSTM 的简化版:把遗忘门和输入门合成一个更新门,去掉单独的细胞状态;参数少四分之一,效果多数时候差不多。
图中 GRU 的隐藏单元画成「另一种记忆单元」(带三角的蓝色圆)。Cho 等 2014 年提出。下面是 PyTorch 的写法:
z 接近 1 时直接沿用旧状态,梯度也能沿这条路直通。
上面的候选状态是 PyTorch(cuDNN)的写法:先对旧状态做线性变换,再乘重置门。Cho 2014 原论文的顺序相反,先用重置门筛选旧状态,再做线性变换:
矩阵乘法和逐元素相乘不能交换顺序,所以两种写法一般不等价。更新门的写法 Cho 2014 和 PyTorch 一致;有些文献(如 Chung 2014)把 z 和 1 − z 反过来写,这只是约定不同,可以等价转换。
和 LSTM 怎么选 :大量对比实验(Chung 2014、Jozefowicz 2015)表明两者在多数任务上相近。GRU 参数少、快,小数据和嵌入式场景常选它;LSTM 有独立的细胞状态,极长依赖时略有优势。
双向 RNN :正向读一遍、反向读一遍,每个位置拼接两个方向的隐状态。适合整段序列都已知的任务(分类、标注、语音识别的编码器)。需要逐帧实时输出时不能等还没到来的未来帧;可以只对截至当前时刻的历史窗口做双向编码,或者只往后多看有限几帧、接受一点延迟(语音识别里的 latency-controlled BLSTM 就是这种做法),代价是延迟和重复计算。
变长序列 :一个 batch 里序列长短不一时,先右侧补零,再用 pack_padded_sequence 告诉 RNN 每条的真实长度,补零的位置不会参与计算,也不会污染最终状态。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 import torchimport torch.nn as nnfrom torch.nn.utils.rnn import pack_padded_sequencecount = lambda m: sum (p.numel() for p in m.parameters()) print ("参数量 RNN:" , count(nn.RNN(64 , 128 )), " GRU:" , count(nn.GRU(64 , 128 )), " LSTM:" , count(nn.LSTM(64 , 128 )))def gru_step (x, h, cell ): gi = x @ cell.weight_ih.T + cell.bias_ih gh = h @ cell.weight_hh.T + cell.bias_hh i_r, i_z, i_n = gi.chunk(3 , -1 ) h_r, h_z, h_n = gh.chunk(3 , -1 ) r, z = torch.sigmoid(i_r + h_r), torch.sigmoid(i_z + h_z) n = torch.tanh(i_n + r * h_n) return (1 - z) * n + z * h cell = nn.GRUCell(3 , 5 ) x, h = torch.randn(2 , 3 ), torch.randn(2 , 5 ) print ("手写 = nn.GRUCell:" , torch.allclose(gru_step(x, h, cell), cell(x, h), atol=1e-6 ))class GRUClassifier (nn.Module): def __init__ (self, in_dim, hid, n_cls ): super ().__init__() self.gru = nn.GRU(in_dim, hid, num_layers=2 , batch_first=True , bidirectional=True , dropout=0.1 ) self.head = nn.Linear(2 * hid, n_cls) def forward (self, x, lengths ): packed = pack_padded_sequence(x, lengths, batch_first=True , enforce_sorted=False ) _, h_n = self.gru(packed) return self.head(torch.cat([h_n[-2 ], h_n[-1 ]], -1 )) model = GRUClassifier(6 , 64 , 5 ).eval () x = torch.randn(3 , 100 , 6 ) lengths = torch.tensor([100 , 73 , 41 ]) out1 = model(x, lengths) x[1 , 73 :] = 99.0 print ("补零区域被改动后输出不变:" , torch.allclose(out1, model(x, lengths)))
典型例子
实时语音降噪:RNNoise 用几层小 GRU,在普通 CPU 上实时运行
世界模型:Dreamer 系列的 RSSM 用 GRU 维护确定性的隐状态,在隐空间里「想象」未来
机器人高频控制:一些视触觉策略在低频的大模型之外,再挂一个小 GRU 跑几十 Hz 的快速修正
双向 GRU/LSTM:命名实体识别等序列标注、语音识别编码器
优势 :比 LSTM 少约 25% 参数、更快,多数任务效果持平;小数据上不易过拟合;适合嵌入式实时运行
局限 :与 LSTM 相同:不能沿时间并行,超长依赖不如注意力
适合的数据 :与 LSTM 相同;尤其是算力和延迟受限的实时序列:音频流、高频传感器、控制回路、世界模型的递推状态
2.5.3 状态空间模型 SSM 与 Mamba
一句话
用控制论里的线性状态空间方程做序列建模:训练时可以像卷积一样并行,推理时像 RNN 一样每步 O(1);Mamba 让方程的参数随输入变化,学会选择性地记住或忘掉。
连续时间的状态空间方程 (和控制论里的完全是同一个东西):
离散化 (步长 Δ,零阶保持):
D x_t 是输入直接到输出的直通项;下面第一个验证「递推 = 卷积」的例子取 D = 0,后面的 Mamba 例子保留可学习的 D(初始化为 1)。这里用零阶保持离散化。Mamba 对 A 项用零阶保持,B 项用简化的 B̄ ≈ ΔB(官方实现和下面的代码都这样);S4 原论文用的是双线性离散化。
两种算法,结果相同 (A、B、C 不随时间变化时):
递推 :一步步更新 h,用于推理,每步 O(1)、显存恒定
卷积 :初始状态为 0 时,展开递推得 y = K * x,卷积核 K = (C B̄, C Ā B̄, C ² B̄, …),训练时整条序列一次并行算完(初始状态 h₀ 不为 0 时,第 t 步的输出还要加上 C Āᵗ h₀)
S4(Gu 2021) :给 A 一个特殊的初始化(HiPPO,让状态能最优地压缩历史),并设计了快速计算 K 的方法。它是第一个解决 Long Range Arena 里 Path-X(长度 16384 的序列)的模型。
Mamba(Gu & Dao 2023) :让 B、C、Δ 都由当前输入 x_t 算出来,于是模型可以按内容决定:Δ 大时状态被当前输入大幅改写(「记住这个」),Δ 小时状态几乎不变(「跳过这个」)。这和 LSTM 的门是同一类思想。参数随输入变化后不能再写成卷积,Mamba 用硬件友好的并行扫描算法在 GPU 上高效计算。
Mamba block :输入投影成两路,一路过短的因果卷积和 SiLU 后进选择性 SSM,另一路过 SiLU 当门,两者相乘后投影回原维度,外面套残差。
相关路线 :Mamba-2(2024)证明了 SSM 和线性注意力之间的对偶关系;RWKV、RetNet、线性注意力,以及 xLSTM 里的 mLSTM 单元,都能训练时沿时间并行、推理时递推;xLSTM 的另一种单元 sLSTM 保留了对上一步隐藏状态的依赖,沿时间仍要逐步算,整个 xLSTM 能不能沿时间并行要看用了哪种单元;Jamba 等把 Transformer 层和 Mamba 层混合使用。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 import torchtorch.manual_seed(0 ) N, T, dt = 4 , 30 , 0.1 A = -torch.diag(torch.rand(N) + 0.1 ) B, C = torch.randn(N, 1 ), torch.randn(1 , N) A_bar = torch.matrix_exp(dt * A) B_bar = torch.linalg.solve(A, (A_bar - torch.eye(N)) @ B) x = torch.randn(T) h, y_rec = torch.zeros(N, 1 ), [] for t in range (T): h = A_bar @ h + B_bar * x[t] y_rec.append((C @ h).squeeze()) y_rec = torch.stack(y_rec) K = torch.stack([(C @ torch.linalg.matrix_power(A_bar, k) @ B_bar).squeeze() for k in range (T)]) y_conv = torch.stack([(K[:t + 1 ].flip(0 ) * x[:t + 1 ]).sum () for t in range (T)]) print ("递推结果 = 卷积结果:" , torch.allclose(y_rec, y_conv, atol=1e-5 ))print ("卷积核前 5 项(随距离衰减 = 渐渐遗忘):" , K[:5 ].numpy().round (3 ))
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 import torchimport torch.nn as nnimport torch.nn.functional as Fclass SelectiveSSM (nn.Module): """Mamba 的选择性 SSM。为了看清原理用朴素的逐步扫描;官方实现用并行扫描的 CUDA kernel""" def __init__ (self, d, n_state=16 ): super ().__init__() self.A_log = nn.Parameter(torch.log(torch.arange(1 , n_state + 1 ).float ()).repeat(d, 1 )) self.x_proj = nn.Linear(d, 2 * n_state + 1 ) self.dt_proj = nn.Linear(1 , d) self.D = nn.Parameter(torch.ones(d)) def forward (self, x ): A = -torch.exp(self.A_log) n = A.size(1 ) B_t, C_t, dt = self.x_proj(x).split([n, n, 1 ], dim=-1 ) dt = F.softplus(self.dt_proj(dt)) h = x.new_zeros(x.size(0 ), x.size(2 ), n) ys = [] for t in range (x.size(1 )): dA = torch.exp(dt[:, t, :, None ] * A) h = dA * h + dt[:, t, :, None ] * B_t[:, t, None , :] * x[:, t, :, None ] ys.append((h * C_t[:, t, None , :]).sum (-1 )) return torch.stack(ys, 1 ) + x * self.D class MambaBlock (nn.Module): def __init__ (self, d, expand=2 ): super ().__init__() di = expand * d self.norm = nn.LayerNorm(d) self.in_proj = nn.Linear(d, 2 * di) self.conv = nn.Conv1d(di, di, 4 , padding=3 , groups=di) self.ssm = SelectiveSSM(di) self.out_proj = nn.Linear(di, d) def forward (self, x ): u, gate = self.in_proj(self.norm(x)).chunk(2 , -1 ) u = self.conv(u.transpose(1 , 2 ))[..., :x.size(1 )].transpose(1 , 2 ) return x + self.out_proj(self.ssm(F.silu(u)) * F.silu(gate)) def adding_problem (B, T=30 ): vals, marks = torch.rand(B, T), torch.zeros(B, T) marks[torch.arange(B), torch.randint(0 , T // 2 , (B,))] = 1 marks[torch.arange(B), torch.randint(T // 2 , T, (B,))] = 1 return torch.stack([vals, marks], -1 ), (vals * marks).sum (1 , keepdim=True ) torch.manual_seed(0 ) model = nn.Sequential(nn.Linear(2 , 32 ), MambaBlock(32 )) head = nn.Linear(32 , 1 ) opt = torch.optim.Adam(list (model.parameters()) + list (head.parameters()), lr=3e-3 ) for step in range (500 ): x, y = adding_problem(64 ) loss = F.mse_loss(head(model(x)[:, -1 ]), y) opt.zero_grad(); loss.backward(); opt.step() x, y = adding_problem(2000 ) with torch.no_grad(): print (f"Mamba 测试 MSE: {F.mse_loss(head(model(x)[:, -1 ]), y).item():.4 f} (瞎猜 ≈ 0.167)" )
典型例子
长序列基准:S4 在 Long Range Arena 上大幅领先当时的高效 Transformer
语言模型:Mamba-3B 的效果接近两倍参数量的 Transformer;Jamba 混合 Transformer、Mamba 和 MoE
基因组:DNA 序列动辄几十万碱基,Caduceus 等 Mamba 类模型可以直接处理
音频:SaShiMi 等直接建模原始波形
机器人与控制:高频传感器流、长时程的观测历史,推理时显存恒定这一点很有用
优势 :训练可并行,推理每步 O(1)、显存不随长度增长;复杂度对序列长度是线性(Mamba 的扫描)或近线性(S4 的 FFT 卷积)的,能处理几万到上百万长度;选择性机制让它能按内容决定记什么
局限 :固定大小的状态要压缩全部历史,精确回忆长上下文中任意一段(比如原样复制前文)不如注意力;生态、工具链和经验积累不如 Transformer;纯 SSM 在部分任务上需要和注意力层混合
适合的数据 :超长序列:基因组、长音频、高频传感器流、长时间序列;对推理延迟和显存敏感的流式场景
2.5.4 普通 RNN 〔图:RNN〕
一句话
最基本的循环网络:新的隐状态 = tanh(当前输入的线性变换 + 上一步隐状态的线性变换)。
图中 RNN 的隐藏单元画成「循环单元」(蓝色圆,带自连接的小环)。常见的是 Elman 网络(1990),把隐状态反馈回来;另一种 Jordan 网络(1986)把上一步的输出反馈回来。
能力和问题都在 2.5.0 讲过:结构简单、流式友好,但梯度随距离指数衰减,经验上常常只记得住十几步(取决于任务和训练设置,不是硬上限)。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 import torchimport torch.nn as nnclass VanillaRNN (nn.Module): """手写递推,看清 h_t 怎么依赖 h_{t-1}""" def __init__ (self, in_dim, hid_dim, out_dim ): super ().__init__() self.W_xh = nn.Linear(in_dim, hid_dim) self.W_hh = nn.Linear(hid_dim, hid_dim, bias=False ) self.W_hy = nn.Linear(hid_dim, out_dim) self.hid_dim = hid_dim def forward (self, x ): h = x.new_zeros(x.size(0 ), self.hid_dim) outs = [] for t in range (x.size(1 )): h = torch.tanh(self.W_xh(x[:, t]) + self.W_hh(h)) outs.append(self.W_hy(h)) return torch.stack(outs, dim=1 ) torch.manual_seed(0 ) mine, ref = VanillaRNN(1 , 16 , 1 ), nn.RNN(1 , 16 , batch_first=True ) with torch.no_grad(): ref.weight_ih_l0.copy_(mine.W_xh.weight); ref.bias_ih_l0.copy_(mine.W_xh.bias) ref.weight_hh_l0.copy_(mine.W_hh.weight); ref.bias_hh_l0.zero_() x = torch.randn(4 , 10 , 1 ) print ("手写 = nn.RNN:" , torch.allclose(mine(x), mine.W_hy(ref(x)[0 ]), atol=1e-6 ))t = torch.linspace(0 , 20 * torch.pi, 2000 ) seqs = torch.sin(t).unfold(0 , 51 , 10 ).unsqueeze(-1 ) x, y = seqs[:, :-1 ], seqs[:, 1 :] opt = torch.optim.Adam(mine.parameters(), lr=1e-2 ) for step in range (300 ): loss = nn.functional.mse_loss(mine(x), y) opt.zero_grad(); loss.backward(); opt.step() print (f"正弦单步预测 MSE: {loss.item():.6 f} " )
典型例子
字符级语言模型:Karpathy 2015 年的 char-rnn 学会生成莎士比亚风格的文本、Linux 源码风格的代码(char-rnn 支持普通 RNN、LSTM、GRU,博客里的例子用的是多层 LSTM)
简单的时间序列预测、短序列分类
教学:理解循环、BPTT、梯度消失的最小例子
优势 :结构最简单;任意长度的序列都能处理,参数量与长度无关;流式推理每步 O(1)
局限 :梯度消失/爆炸,经验上常常只能记住十几步;训练不能沿时间并行;实际项目中基本被 LSTM、GRU、Transformer、SSM 取代
适合的数据 :只有短程依赖的序列:简单传感器流、短时间序列;教学