2.5 循环与序列模型

2.5.0 循环的公共机制(家族公共部分)

<center>2.5.0 循环展开
2.5.0 循环展开
<center>2.5.0 循环展开
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 torch

# 线性 RNN h_t = W h_{t-1} + x_t,W = ρ·Q(Q 正交),梯度大小恰好按 ρ 的幂变化
torch.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]:.1e} / {g[50]:.1e} / {g[0]:.1e}")

# 真实的 tanh RNN(PyTorch 默认初始化):梯度同样随距离迅速衰减
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]:.1e}" for i in (99, 95, 90, 80)])
<center>RNN梯度随距离衰减
RNN梯度随距离衰减
<center>RNN梯度随距离衰减
RNN梯度随距离衰减

2.5.1 LSTM 长短期记忆网络 〔图:LSTM〕

<center>2.5.1 LSTM
2.5.1 LSTM
<center>2.5.1 LSTM
2.5.1 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 torch
import torch.nn as nn

# 1) 按公式手写一步 LSTM,和 nn.LSTMCell 对比(PyTorch 的门顺序:输入、遗忘、候选、输出)
def 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))

# 2) 长程依赖测试(adding problem):50 步序列里只有两个位置被标记,目标是这两个数之和
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) # (B, T, hid)
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():.4f}")
print("瞎猜(永远输出 1.0)的 MSE ≈", round(1 / 6, 4))
<center>adding问题训练曲线
adding问题训练曲线
<center>adding问题训练曲线
adding问题训练曲线

普通 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〕

<center>2.5.2 GRU
2.5.2 GRU
<center>2.5.2 GRU
2.5.2 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 torch
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence

count = 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)))
# 1 组 / 3 组 / 4 组权重

# 手写一步 GRU,和 nn.GRUCell 对比(PyTorch 的门顺序:重置、更新、候选)
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))

# 实用写法:双向两层 GRU 编码一批长度不同的序列(比如时长不同的力/触觉信号片段)
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): # x: (B, T_max, in_dim),右侧补零
packed = pack_padded_sequence(x, lengths, batch_first=True, enforce_sorted=False)
_, h_n = self.gru(packed) # h_n: (层数×2, B, hid)
return self.head(torch.cat([h_n[-2], h_n[-1]], -1)) # 最后一层的正向、反向最终状态

model = GRUClassifier(6, 64, 5).eval()
x = torch.randn(3, 100, 6) # 3 段 6 维信号,补齐到 100 步
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

<center>2.5.3 SSM 与 Mamba
2.5.3 SSM 与 Mamba
<center>2.5.3 SSM 与 Mamba
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 torch

# S4 的出发点:同一个线性状态空间模型,递推算和卷积算结果相同
torch.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): # 递推:推理时用,每步 O(1)
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 torch
import torch.nn as nn
import torch.nn.functional as F

class 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)) # (d, N)
self.x_proj = nn.Linear(d, 2 * n_state + 1) # 从输入算出 B_t、C_t、Δ_t → 「选择性」
self.dt_proj = nn.Linear(1, d)
self.D = nn.Parameter(torch.ones(d))

def forward(self, x): # x: (B, T, d)
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)) # (B, T, d):Δ 大 → 写入当前输入;Δ 小 → 保持旧状态
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) # Ā_t = exp(Δ_t A)
h = dA * h + dt[:, t, :, None] * B_t[:, t, None, :] * x[:, t, :, None] # h_t = Ā_t h_{t-1} + Δ_t B_t x_t(B 项用简化的 ΔB,和官方实现一致)
ys.append((h * C_t[:, t, None, :]).sum(-1)) # y_t = C_t h_t
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): # 与 LSTM 一节相同的测试;朴素扫描慢,序列缩短到 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():.4f}(瞎猜 ≈ 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〕

<center>2.5.4 RNN
2.5.4 RNN
<center>2.5.4 RNN
2.5.4 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 torch
import torch.nn as nn

class 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): # x: (B, T, in_dim)
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) # (B, T, out_dim)

# 和 nn.RNN 对比:拷贝同样的权重,输出一致
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 个点,预测第 t+1 个点
t = torch.linspace(0, 20 * torch.pi, 2000)
seqs = torch.sin(t).unfold(0, 51, 10).unsqueeze(-1) # (195, 51, 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():.6f}")

典型例子

  • 字符级语言模型:Karpathy 2015 年的 char-rnn 学会生成莎士比亚风格的文本、Linux 源码风格的代码(char-rnn 支持普通 RNN、LSTM、GRU,博客里的例子用的是多层 LSTM)
  • 简单的时间序列预测、短序列分类
  • 教学:理解循环、BPTT、梯度消失的最小例子
  • 优势:结构最简单;任意长度的序列都能处理,参数量与长度无关;流式推理每步 O(1)
  • 局限:梯度消失/爆炸,经验上常常只能记住十几步;训练不能沿时间并行;实际项目中基本被 LSTM、GRU、Transformer、SSM 取代
  • 适合的数据:只有短程依赖的序列:简单传感器流、短时间序列;教学