AI 工程基础体系 · 第 60/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

序列模型:RNN、LSTM、GRU、Teacher Forcing 与长依赖

序列模型处理的对象不是一组相互独立的样本,而是有顺序关系的数据:

x1,x2,,xTx_1,x_2,\ldots,x_T

其中 xtx_t 是第 tt 个时间步的输入,TT 是序列长度。文本中的词元、语音中的声学帧、传感器的时间序列、用户行为事件和视频帧都可以采用这种表示。

序列建模的难点不只是“输入长度可变”,更重要的是:当前输出可能依赖很久以前的信息。例如:

  • 句子后半部分的代词可能依赖前半部分的实体;
  • 传感器当前状态可能受到数百个时间步之前的控制信号影响;
  • 生成模型在第 tt 步生成的词会影响第 t+1t+1 步;
  • 序列中不同位置的重要信息可能相隔很远。

RNN、LSTM 和 GRU 都通过维护某种“状态”来处理历史信息;Teacher Forcing 则解决序列生成训练时如何提供前一步输入的问题。长依赖问题贯穿这几个概念:状态如何保存历史、梯度如何穿过历史、训练和推理时状态输入是否一致,决定了模型能否真正利用远距离信息。


一、序列模型的基本形式

1. 从独立样本到条件分布

普通分类模型通常假设每个样本独立:

p(yx)p(y\mid x)

序列模型则经常建模联合概率:

p(y1,y2,,yTx)=t=1Tp(yty<t,x)p(y_1,y_2,\ldots,y_T\mid x) = \prod_{t=1}^{T}p(y_t\mid y_{<t},x)

其中:

  • y<t=(y1,,yt1)y_{<t}=(y_1,\ldots,y_{t-1})
  • p(yty<t,x)p(y_t\mid y_{<t},x) 表示在已知此前输出和条件输入后预测第 tt 个输出;
  • 这个分解来自概率链式法则,并不要求模型一定使用 RNN。

在语言模型中,通常令 xx 为空,得到:

p(y1,,yT)=t=1Tp(yty<t)p(y_1,\ldots,y_T) = \prod_{t=1}^{T}p(y_t\mid y_{<t})

训练目标是最大化真实序列的对数似然,等价于最小化交叉熵:

L=t=1Tlogpθ(yty<t)\mathcal{L} = -\sum_{t=1}^{T}\log p_\theta(y_t\mid y_{<t})

如果不同样本长度不同,还需要对 padding 位置进行掩码,否则模型会把填充符号当成真实目标参与训练。

2. 隐状态的作用

RNN 类模型不会把完整历史 x1,,xtx_1,\ldots,x_t 直接全部传给当前计算,而是将历史压缩为隐状态:

ht=fθ(xt,ht1)h_t = f_\theta(x_t,h_{t-1})

然后由隐状态产生输出:

ot=gθ(ht)o_t = g_\theta(h_t)

这里的 hth_t 是模型在时间步 tt 的内部状态。它不是“数据库中的历史记录”,而是一个固定维度的连续向量,因此存在信息压缩和遗忘。

单向序列模型的数据流可以表示为:

flowchart LR
    X1[x1] --> R1[状态更新]
    H0[h0] --> R1
    R1 --> H1[h1]
    H1 --> Y1[y1]

    X2[x2] --> R2[状态更新]
    H1 --> R2
    R2 --> H2[h2]
    H2 --> Y2[y2]

    X3[x3] --> R3[状态更新]
    H2 --> R3
    R3 --> H3[h3]
    H3 --> Y3[y3]

关键路径是 ht1hth_{t-1}\rightarrow h_t。如果这条路径能够长期保存有效信息,模型就可能学习长依赖;如果信息或梯度在这条路径上快速衰减,模型就只能利用局部上下文。


二、基础 RNN:递归状态与时间反向传播

1. Vanilla RNN 的递推公式

最基本的 RNN,也常称为 vanilla RNN 或 Elman RNN,可以写成:

at=Wxhxt+Whhht1+bha_t = W_{xh}x_t + W_{hh}h_{t-1}+b_h

ht=ϕ(at)h_t = \phi(a_t)

ot=Whyht+byo_t = W_{hy}h_t+b_y

对于分类或语言模型,通常再使用:

pt=softmax(ot)p_t=\operatorname{softmax}(o_t)

变量含义如下:

  • xtRdxx_t\in\mathbb{R}^{d_x}:当前输入;
  • htRdhh_t\in\mathbb{R}^{d_h}:当前隐状态;
  • WxhRdh×dxW_{xh}\in\mathbb{R}^{d_h\times d_x}:输入到隐状态的权重;
  • WhhRdh×dhW_{hh}\in\mathbb{R}^{d_h\times d_h}:前一状态到当前状态的循环权重;
  • WhyW_{hy}:隐状态到输出的权重;
  • ϕ\phi:非线性函数,经典实现常用 tanh,也可以使用其他激活函数。

这里的“循环”不是指模型一次性把整个序列重新输入,而是同一组参数在每个时间步重复使用。参数共享使模型可以处理不同长度的序列,但也意味着所有时间步必须共用相同的状态更新规则。

2. 一个完整的标量算例

为了看清状态如何传播,考虑只有一个输入维度和一个隐状态维度的 RNN:

ht=tanh(0.5xt+0.8ht1)h_t=\tanh(0.5x_t+0.8h_{t-1})

h0=0h_0=0,输入序列为:

x1=1,x2=0,x3=0x_1=1,\quad x_2=0,\quad x_3=0

则:

a1=0.5×1+0.8×0=0.5a_1=0.5\times1+0.8\times0=0.5

h1=tanh(0.5)0.4621h_1=\tanh(0.5)\approx0.4621

第二步:

a2=0.5×0+0.8×0.46210.3697a_2=0.5\times0+0.8\times0.4621\approx0.3697

h2=tanh(0.3697)0.3537h_2=\tanh(0.3697)\approx0.3537

第三步:

a3=0.8×0.35370.2829a_3=0.8\times0.3537\approx0.2829

h3=tanh(0.2829)0.2756h_3=\tanh(0.2829)\approx0.2756

即使后续输入全为零,h1h_1 中的信息仍然会影响 h2h_2h3h_3,但影响逐步减弱。这正是递归状态的基本工作方式:历史信息不会自动以独立字段保存,而是通过后续状态更新间接存在。

3. 时间反向传播与梯度乘积

训练 RNN 时,损失可能来自多个时间步:

L=t=1TLt\mathcal{L}=\sum_{t=1}^{T}\mathcal{L}_t

要更新早期时间步的参数,梯度必须沿时间方向反向传播。对于某个较早状态 hkh_k,来自较晚状态 hTh_T 的梯度包含多个雅可比矩阵的乘积:

hThk=t=k+1Ththt1\frac{\partial h_T}{\partial h_k} = \prod_{t=k+1}^{T} \frac{\partial h_t}{\partial h_{t-1}}

对于:

ht=tanh(Whhht1+Wxhxt+b)h_t=\tanh(W_{hh}h_{t-1}+W_{xh}x_t+b)

有:

htht1=diag(1tanh2(at))Whh\frac{\partial h_t}{\partial h_{t-1}} = \operatorname{diag}\left(1-\tanh^2(a_t)\right)W_{hh}

因此:

hThk=t=k+1T[diag(1tanh2(at))Whh]\frac{\partial h_T}{\partial h_k} = \prod_{t=k+1}^{T} \left[ \operatorname{diag}(1-\tanh^2(a_t))W_{hh} \right]

每一项的范数如果平均小于 1,乘积会趋近于 0,形成梯度消失;如果平均大于 1,乘积可能快速变大,形成梯度爆炸

标量反例

假设为了观察数量级,忽略 tanh 的导数,只考虑标量循环权重 ww

ht=wht1h_t=w h_{t-1}

那么:

hThk=wTk\frac{\partial h_T}{\partial h_k}=w^{T-k}

w=0.5w=0.5Tk=20T-k=20 时:

0.5209.54×1070.5^{20}\approx9.54\times10^{-7}

早期状态几乎收不到来自后期的梯度。

w=1.2w=1.2Tk=20T-k=20 时:

1.22038.341.2^{20}\approx38.34

梯度则可能放大几十倍。真实 RNN 中还存在矩阵方向、激活函数饱和和多个损失项,因此数值不一定正好如此,但“长链式乘积导致不稳定”的因果关系不变。

4. 梯度消失和信息遗忘不是同一件事

这两个现象相关,但不能混为一谈:

  • 梯度消失:训练信号难以从后面的时间步传回前面的参数或状态;
  • 信息遗忘:前面的输入在状态更新过程中逐渐不再影响当前状态。

如果梯度消失,模型通常很难学会“应该保留某条长距离信息”的参数;但即使梯度没有严重消失,固定维度状态也可能因为容量不足而丢失信息。

反过来,某些任务本来就不需要保存全部历史。对短序列分类,vanilla RNN 可能足够,而且参数和计算开销通常小于更复杂的门控模型。

5. 梯度裁剪能解决什么,不能解决什么

梯度裁剪常用于抑制梯度爆炸。例如:

loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()

clip_grad_norm_ 会在梯度整体范数超过阈值时按比例缩小梯度。它可以避免一次更新过大、减少训练发散,但不能恢复已经消失的梯度,也不能让 RNN 自动获得更大的记忆容量。

因此,梯度裁剪是数值稳定措施,不是长依赖的根本解决方案。


三、LSTM:用细胞状态和门控控制记忆

1. LSTM 的两个状态

LSTM,即 Long Short-Term Memory,使用两个状态:

  • hth_t:隐状态,通常用于当前输出和传递短期表示;
  • ctc_t:细胞状态,主要承担长期记忆通道。

对输入 xtx_t 和上一时刻状态 (ht1,ct1)(h_{t-1},c_{t-1}),LSTM 计算四个量:

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)

ot=σ(Woxt+Uoht1+bo)o_t=\sigma(W_o x_t+U_o h_{t-1}+b_o)

c~t=tanh(Wcxt+Ucht1+bc)\tilde{c}_t=\tanh(W_c x_t+U_c h_{t-1}+b_c)

其中:

  • iti_t:输入门,控制新信息写入多少;
  • ftf_t:遗忘门,控制旧细胞状态保留多少;
  • oto_t:输出门,控制细胞状态暴露给隐状态多少;
  • c~t\tilde c_t:候选记忆;
  • σ(z)=1/(1+ez)\sigma(z)=1/(1+e^{-z}),输出范围为 (0,1)(0,1)

状态更新为:

ct=ftct1+itc~tc_t=f_t\odot c_{t-1}+i_t\odot\tilde{c}_t

ht=ottanh(ct)h_t=o_t\odot\tanh(c_t)

这里的 \odot 表示逐元素乘法。

2. LSTM 为什么更容易保留长依赖

细胞状态的更新是加法路径:

ct=ftct1+itc~tc_t=f_t\odot c_{t-1}+i_t\odot\tilde{c}_t

对上一时刻细胞状态求导:

ctct1=ft\frac{\partial c_t}{\partial c_{t-1}}=f_t

连续传播 kk 个时间步时:

cTck=t=k+1Tft\frac{\partial c_T}{\partial c_k} = \prod_{t=k+1}^{T}f_t

当遗忘门 ftf_t 接近 1 时,这个乘积可以较长时间保持稳定。与普通 RNN 的“每一步都经过固定非线性变换”相比,LSTM 可以学习一条近似恒等的记忆通道。

但这不是无条件保证:

  • ftf_t 如果长期小于 1,梯度仍会衰减;
  • ctc_t 仍是固定维度,容量有限;
  • 门控计算本身也可能训练不稳定;
  • LSTM 只能通过顺序递推访问历史,长序列推理难以完全并行。

“LSTM 能解决梯度消失”更准确的说法是:它通过门控和加法状态路径显著缓解了普通 RNN 的长距离优化问题,而不是在任意任务和任意长度上保证没有梯度消失。

3. LSTM 的逐步记忆例子

假设某个维度上:

ct1=10,ft=0.9,it=0.2,c~t=1c_{t-1}=10,\quad f_t=0.9,\quad i_t=0.2,\quad \tilde c_t=1

则:

ct=0.9×10+0.2×1=9.2c_t=0.9\times10+0.2\times1=9.2

如果下一步没有新内容且模型决定继续保留:

ct+1=0.99×9.2+0×c~t+1=9.108c_{t+1}=0.99\times9.2+0\times\tilde c_{t+1}=9.108

这说明模型可以把旧信息逐渐衰减,而不是每个时间步都强制用新输入覆盖它。实际模型会为每个维度、每个时间步分别计算门值,因此不同特征可以有不同的保留时间。

4. PyTorch 中的 LSTM 状态

PyTorch 的 torch.nn.LSTM 返回:

output, (h_n, c_n) = lstm(x, (h_0, c_0))

在常见的 batch_first=True 配置下:

  • x 形状为 [batch_size, seq_len, input_size]
  • output 形状为 [batch_size, seq_len, hidden_size * num_directions]
  • h_n 形状为 [num_layers * num_directions, batch_size, hidden_size]
  • c_n 形状与 h_n 相同。

output[:, -1, :] 是最后一个时间步的输出,但它不一定等于所有情况下的最终语义表示:

  • 双向 LSTM 的最后位置包含正向和反向信息;
  • 多层 LSTM 的 output 来自最后一层;
  • 如果序列进行了 padding,最后一个物理位置可能是填充,而不是最后一个有效位置。

因此,变长序列需要结合长度信息、pack_padded_sequence 或显式 mask 处理,不能无条件取 [:, -1, :]


四、GRU:用更少门控压缩 LSTM 结构

1. GRU 的基本公式

GRU,即 Gated Recurrent Unit,只有一个主要状态 hth_t,不单独维护 LSTM 的 ctc_t。一种常见公式为:

zt=σ(Wzxt+Uzht1+bz)z_t=\sigma(W_zx_t+U_zh_{t-1}+b_z)

rt=σ(Wrxt+Urht1+br)r_t=\sigma(W_rx_t+U_rh_{t-1}+b_r)

h~t=tanh(Whxt+Uh(rtht1)+bh)\tilde h_t= \tanh(W_hx_t+U_h(r_t\odot h_{t-1})+b_h)

ht=(1zt)ht1+zth~th_t=(1-z_t)\odot h_{t-1}+z_t\odot\tilde h_t

其中:

  • ztz_t:更新门,决定新候选状态与旧状态的混合比例;
  • rtr_t:重置门,决定生成候选状态时使用多少旧状态;
  • h~t\tilde h_t:候选状态;
  • hth_t:新的隐藏状态。

不同论文和框架可能将 ztz_t 的语义写成“保留旧状态的比例”,对应的插值公式会变成:

ht=ztht1+(1zt)h~th_t=z_t\odot h_{t-1}+(1-z_t)\odot\tilde h_t

这两种写法只是门的定义约定不同,阅读公式或实现时必须同时检查门的定义和状态更新式,不能只看变量名。

2. GRU 的长依赖路径

在上述公式中:

ht=(1zt)ht1+zth~th_t=(1-z_t)\odot h_{t-1}+z_t\odot\tilde h_t

如果 ztz_t 接近 0,则:

htht1h_t\approx h_{t-1}

旧状态可以沿近似恒等路径继续传递。这与 LSTM 的细胞状态通道有类似效果,只是 GRU 将记忆和输出合并在一个状态中。

GRU 相比 LSTM 通常具有:

  • 更少的状态;
  • 更少的门控结构;
  • 较低的参数量和计算量;
  • 更简单的接口。

但“参数更少”不等于“总是更好”。任务可能需要将长期记忆和当前输出分开控制,此时 LSTM 的独立 ctc_thth_t 可能更合适。具体效果取决于数据规模、序列长度、优化设置和任务本身。

3. RNN、LSTM 和 GRU 的结构差异

模型 状态 主要机制 长依赖处理
Vanilla RNN hth_t 非线性递推 容易受到梯度消失或爆炸影响
LSTM ht,cth_t,c_t 输入门、遗忘门、输出门 通过细胞状态的加法路径保留信息
GRU hth_t 更新门、重置门 通过状态插值保留或替换信息

这些结构都仍然是顺序递推:计算第 tt 步需要第 t1t-1 步的状态。因此,门控改善的是记忆和优化,不会改变 RNN 的时间依赖结构。


五、Teacher Forcing:训练时使用真实前一步输出

1. 自回归解码的输入来源

在生成任务中,解码器通常根据此前的输出生成下一个词元:

p(yty<t,x)p(y_t\mid y_{<t},x)

以目标序列:

<BOS> 我 喜欢 机器 学习 <EOS>

为例,训练输入和目标通常错开一个位置:

时间步 解码器输入 目标输出
1 <BOS>
2 喜欢
3 喜欢 机器
4 机器 学习
5 学习 <EOS>

Teacher Forcing 指的是:在训练第 tt 步时,把真实目标 yt1y_{t-1} 作为解码器输入,而不是把模型上一步预测的 y^t1\hat y_{t-1} 作为输入。

因此训练阶段计算的是:

LTF=t=1Tlogpθ(yty<t真实,x)\mathcal{L}_{\text{TF}} = -\sum_{t=1}^{T} \log p_\theta(y_t\mid y_{<t}^{\text{真实}},x)

这种方式可以并行计算所有时间步的输出,因为整个真实目标序列都已知。

2. 推理时为什么不能使用 Teacher Forcing

推理时没有真实目标序列,模型只能使用自己的输出:

y^tpθ(y^<t,x)\hat y_t\sim p_\theta(\cdot\mid \hat y_{<t},x)

或者使用贪心解码:

y^t=argmaxvpθ(vy^<t,x)\hat y_t=\arg\max_v p_\theta(v\mid \hat y_{<t},x)

所以训练和推理的状态输入不同:

flowchart TD
    A[编码器输入 x] --> B[解码器初始状态]
    B --> C{训练还是推理}

    C -->|Teacher Forcing| D[输入真实 y(t-1)]
    C -->|自回归推理| E[输入模型预测 y_hat(t-1)]

    D --> F[预测 y(t)]
    E --> F
    F --> G[计算损失或输出]
    G --> H[进入下一时间步]

训练期间模型看到的是“正确历史”,推理期间看到的是“可能已经错误的历史”。这会造成 exposure bias,通常译为暴露偏差或训练—推理不一致。

3. 暴露偏差的具体传播

假设真实目标是:

我 喜欢 机器 学习

模型在推理时第一步错误生成了:

我 需要 机器 学习

那么第二步输入变成了“需要”,而不是训练中经常看到的“喜欢”。即使模型本身在正确上下文中能预测“机器”,它也未必学会处理“需要”这个错误上下文。

错误可能继续累积:

y^1y1p(y^2y^1)p(y2y1)\hat y_1\neq y_1 \Rightarrow p(\hat y_2\mid \hat y_1)\neq p(y_2\mid y_1)

这不是 Teacher Forcing “错误”导致的,而是最大似然训练和自回归推理的条件分布不同所导致的结构性问题。

4. Scheduled Sampling 的边界

一种常见尝试是逐渐减少 Teacher Forcing:训练时以一定概率使用真实前一步,以另一概率使用模型预测。这类方法通常称为 scheduled sampling。

它的直觉是让训练分布更接近推理分布,但需要注意:

  • 使用模型采样结果作为下一步输入时,采样操作通常不可微;
  • 训练目标不再简单等同于标准最大似然;
  • 采样概率、衰减策略和随机性会影响训练稳定性;
  • 它不是所有生成任务的通用改进。

因此,不能把 scheduled sampling 简化为“Teacher Forcing 比例越低越好”。实际还可能使用序列级目标、最小风险训练、强化学习目标或基于偏好的优化,但这些方法改变了训练目标,必须单独评估。

5. Teacher Forcing 与并行计算

Teacher Forcing 允许解码器在训练时一次接收:

decoder_input = target[:, :-1]

并同时产生:

logits = decoder(decoder_input, hidden)
target_next = target[:, 1:]

因为第 tt 步输入直接来自真实数据,而不是必须等待第 t1t-1 步模型生成结果,所以训练可以在时间维度上批量计算。推理仍然需要一步一步生成。

这也是训练速度和推理速度可能差异很大的原因之一。


六、一个可运行的 PyTorch Teacher Forcing 示例

下面的示例训练一个 GRU 编码器—解码器,将长度为 5 的数字序列反转。例如:

输入:1 4 2 8 3
目标:3 8 2 4 1

为了让解码器知道何时开始和结束,词表包含:

  • 0<PAD>
  • 1<BOS>
  • 2<EOS>
  • 312:数字 09

目标序列实际表示为:

<BOS> 3 8 2 4 1 <EOS>
import random
import torch
from torch import nn

torch.manual_seed(7)
random.seed(7)

PAD, BOS, EOS = 0, 1, 2
DIGIT_OFFSET = 3
VOCAB_SIZE = 13


def make_batch(batch_size=128, seq_len=5):
    # 输入是数字 token,范围为 [3, 12]
    src = torch.randint(
        low=DIGIT_OFFSET,
        high=DIGIT_OFFSET + 10,
        size=(batch_size, seq_len),
    )

    # 反转数字顺序,并加上 BOS/EOS
    reversed_digits = torch.flip(src, dims=[1])
    bos = torch.full((batch_size, 1), BOS, dtype=torch.long)
    eos = torch.full((batch_size, 1), EOS, dtype=torch.long)
    target = torch.cat([bos, reversed_digits, eos], dim=1)
    return src, target


class Seq2SeqGRU(nn.Module):
    def __init__(self, vocab_size, emb_dim=32, hidden_dim=64):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, emb_dim, padding_idx=PAD)
        self.encoder = nn.GRU(
            input_size=emb_dim,
            hidden_size=hidden_dim,
            batch_first=True,
        )
        self.decoder = nn.GRU(
            input_size=emb_dim,
            hidden_size=hidden_dim,
            batch_first=True,
        )
        self.output = nn.Linear(hidden_dim, vocab_size)

    def forward(self, src, target, teacher_forcing=True):
        # src: [B, S]
        # target: [B, T]
        _, hidden = self.encoder(self.embedding(src))
        # hidden: [1, B, H]

        # 第一个解码输入固定为 BOS
        decoder_input = target[:, :1]
        logits_steps = []

        for t in range(1, target.size(1)):
            decoder_output, hidden = self.decoder(
                self.embedding(decoder_input),
                hidden,
            )
            # decoder_output: [B, 1, H]
            step_logits = self.output(decoder_output)
            # step_logits: [B, 1, V]
            logits_steps.append(step_logits)

            predicted = step_logits.argmax(dim=-1)
            if teacher_forcing:
                # 使用真实的 target[t] 作为下一步输入
                decoder_input = target[:, t:t + 1]
            else:
                # 推理时使用模型自己的预测
                decoder_input = predicted

        return torch.cat(logits_steps, dim=1)
        # 返回 [B, T-1, V]


def train():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = Seq2SeqGRU(VOCAB_SIZE).to(device)
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
    criterion = nn.CrossEntropyLoss(ignore_index=PAD)

    model.train()
    for step in range(1000):
        src, target = make_batch()
        src, target = src.to(device), target.to(device)

        # logits 的时间范围对应 target[:, 1:]
        logits = model(src, target, teacher_forcing=True)
        loss = criterion(
            logits.reshape(-1, VOCAB_SIZE),
            target[:, 1:].reshape(-1),
        )

        optimizer.zero_grad(set_to_none=True)
        loss.backward()

        # RNN 类模型常用梯度裁剪抑制梯度爆炸
        nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        optimizer.step()

        if (step + 1) % 200 == 0:
            print(f"step={step + 1}, loss={loss.item():.4f}")

    return model, device


@torch.no_grad()
def greedy_decode(model, src, device, max_len=7):
    model.eval()
    _, hidden = model.encoder(model.embedding(src.to(device)))

    decoder_input = torch.full(
        (src.size(0), 1),
        BOS,
        dtype=torch.long,
        device=device,
    )
    result = []

    for _ in range(max_len):
        output, hidden = model.decoder(
            model.embedding(decoder_input),
            hidden,
        )
        logits = model.output(output[:, -1:, :])
        token = logits.argmax(dim=-1)
        result.append(token)
        decoder_input = token

        if (token == EOS).all():
            break

    return torch.cat(result, dim=1)


if __name__ == "__main__":
    model, device = train()

    src, target = make_batch(batch_size=1)
    prediction = greedy_decode(model, src, device)

    print("src      :", src.tolist())
    print("target   :", target.tolist())
    print("predict  :", prediction.cpu().tolist())

代码中的关键时间对齐

如果目标为:

<BOS> 8 2 5 1 <EOS>

则:

target[:, :-1]

是解码器输入:

<BOS> 8 2 5 1

而:

target[:, 1:]

是监督目标:

8 2 5 1 <EOS>

模型输出的第一个位置必须预测 8,不是重新预测 <BOS>。因此交叉熵使用 target[:, 1:] 是时间对齐成立的关键。

示例中的训练循环使用 Teacher Forcing;greedy_decode 则完全不使用目标序列,而是在每一步将 argmax 结果送回解码器。这两个路径正好展示了训练与推理的输入差异。

该示例只用于说明机制,不代表实际翻译或生成系统的完整实现。它还没有处理:

  • 变长序列;
  • padding mask;
  • beam search;
  • 训练集和验证集划分;
  • 词表外词元;
  • 混合精度和分布式训练;
  • 生成重复、非法长度和安全过滤。

七、长依赖:模型需要保存什么,梯度需要经过什么

1. 长依赖的形式化描述

设当前目标 yTy_T 依赖早期输入 xkx_k,其中 TkT-k 很大。如果模型要正确预测 yTy_T,则需要使:

I(yT;xkxk+1:T)I(y_T;x_k\mid x_{k+1:T})

保持可利用,或者至少让状态 hTh_T 中保留与 xkx_k 有关的信息。这里 I(;)I(\cdot;\cdot\mid\cdot) 表示条件互信息,可理解为:在已知中间信息后,早期输入对当前目标还剩多少有用信息。

对 RNN 来说,早期信息必须经过:

xkhkhk+1hTx_k\rightarrow h_k\rightarrow h_{k+1} \rightarrow\cdots\rightarrow h_T

同时训练信号还要沿反方向传播:

LThThT1hk\mathcal{L}_T\rightarrow h_T\rightarrow h_{T-1} \rightarrow\cdots\rightarrow h_k

长依赖因此有两个独立困难:

  1. 前向记忆问题:信息是否被状态保存;
  2. 反向优化问题:损失是否能有效指导早期参数。

LSTM 和 GRU 主要改善第二条路径,同时通过门控改善第一条路径;它们并没有消除固定状态容量和顺序访问的限制。

2. 一个长依赖任务与反例

考虑序列分类任务:

标记 A,随后有 100 个无关标记,最后判断是否出现过 A

理想模型需要在看到 A 时写入一个状态标记,然后在中间 100 个时间步保持它,最后读取该标记。

普通 RNN 可能出现以下失败:

  • 状态经过大量更新后,A 的影响趋近于零;
  • 模型只根据末尾局部模式猜测;
  • 训练集上序列较短时表现良好,序列变长后准确率明显下降。

这是一个典型的长度外推反例:模型并不是“理解了记忆规则”,而是记住了训练长度范围中的局部统计模式。

反过来,长依赖也不意味着所有远距离信息都应该保留。如果序列中包含大量噪声,始终保留全部历史会降低有效容量。遗忘门、更新门的作用不仅是“记得更久”,也是学习“哪些信息应该被丢弃”。

3. Truncated BPTT 的影响

对非常长的序列,完整的 Backpropagation Through Time(BPTT,时间反向传播)需要保存所有时间步的中间激活,显存和计算开销随 TT 增长。

工程上常把序列切成长度为 KK 的片段,仅在每个片段内反向传播。这个方法称为截断 BPTT:

hidden = None

for chunk in sequence_chunks:
    output, hidden = rnn(chunk, hidden)

    loss = compute_loss(output)
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()

    # 保留数值状态,但切断它与上一片段计算图的连接
    hidden = hidden.detach()

detach() 的含义不是把 hidden 清零,而是让它保留当前数值、停止梯度继续穿过之前的片段。

因此,截断 BPTT 会产生一个明确边界:

  • 前向状态可以跨片段继续传递;
  • 梯度不能跨越片段传回更早时间步。

如果真实依赖长度明显大于 KK,模型可能在前向上看似携带了历史,却无法通过训练得到正确的长程更新。增大 KK 可以扩大反向范围,但会增加显存和计算成本,不能无条件增大。


八、双向 RNN 与因果约束

双向 RNN 使用两个方向的递推:

ht=f(xt,ht1)\overrightarrow{h_t} = f(x_t,\overrightarrow{h_{t-1}})

ht=f(xt,ht+1)\overleftarrow{h_t} = f(x_t,\overleftarrow{h_{t+1}})

然后拼接:

ht=[ht;ht]h_t= [\overrightarrow{h_t};\overleftarrow{h_t}]

这样,第 tt 个位置可以同时利用左侧和右侧上下文,适合:

  • 离线序列标注;
  • 文本分类;
  • 已完整采集的语音或传感器序列分析。

但双向模型不适合严格的实时因果预测,因为反向状态依赖未来输入。若系统要求在时刻 tt 只能使用 xtx_{\le t},使用双向 RNN 会造成未来信息泄漏。

例如,训练数据中用完整句子的双向编码器预测当前位置标签,离线评测可能很高;部署到实时输入流后,未来词元不存在,性能会显著下降。这不是模型随机失效,而是训练和服务时的数据可见范围不一致。


九、RNN 类模型与 Transformer 的长依赖边界

Transformer 通常通过注意力机制直接建立任意两个位置之间的连接。若序列长度为 TT,位置 ii 可以直接注意位置 jj,路径长度不再必须经过所有中间时间步。

RNN 类模型和 Transformer 的差异可以概括为:

维度 RNN/LSTM/GRU Transformer
历史访问方式 通过递归状态压缩 通过注意力直接访问上下文
训练时间并行性 时间方向受递推限制 通常可并行处理训练序列
推理状态 可压缩为隐藏状态 自回归生成仍需维护缓存或历史表示
长依赖路径 可能经过很多状态转移 可建立较短的直接连接
长序列成本 顺序计算,但状态紧凑 注意力可能带来较高的序列长度成本
流式处理 自然适合 需要专门设计缓存和窗口机制

Transformer 并不意味着“长依赖自动解决”:

  • 注意力仍受上下文长度和表示容量限制;
  • 长序列的计算、显存和成本可能很高;
  • 训练数据中的长依赖必须有足够监督;
  • 推理时生成错误仍可能累积;
  • 位置编码、截断和检索策略会影响可用上下文。

因此,RNN/LSTM/GRU 仍适合低延迟流式场景、小模型、资源受限设备和明确的递归状态任务。若需要对长上下文中的任意位置进行灵活访问,注意力机制通常更直接。


十、序列训练中的数据、损失和评测问题

1. Padding 与 mask

批处理通常要求同一批样本具有相同长度,因此短序列会补 <PAD>。若目标为:

样本 A:8 2 <EOS> <PAD> <PAD>
样本 B:4 9 1 6 <EOS>

直接对全部位置计算交叉熵会让模型学习预测 padding。应使用:

criterion = nn.CrossEntropyLoss(ignore_index=PAD)

或者手动构造有效位置 mask。输入端的 padding 也不能被当成真实事件参与状态更新;在使用 RNN 编码变长序列时,应根据实际长度使用 pack_padded_sequence,或确保后续池化和损失计算忽略无效位置。

2. 训练指标与生成指标不是一回事

Teacher Forcing 下的 token-level cross-entropy 测量的是:

p(yty<t真实,x)p(y_t\mid y_{<t}^{\text{真实}},x)

但实际生成测量的是:

p(y^ty^<t,x)p(\hat y_t\mid \hat y_{<t},x)

前者较低,并不保证后者生成质量高。评测时至少应区分:

  • Teacher Forcing loss;
  • 自回归 token accuracy;
  • 序列级完全匹配率;
  • 任务相关指标;
  • 不同长度区间上的指标;
  • 错误历史条件下的鲁棒性。

对于记忆任务,还应按依赖距离分桶评测。例如分别统计需要记住 10、50、100、500 个时间步的准确率,才能看出模型到底学到了多长的依赖。

3. 防止序列泄漏

时间序列和用户行为数据不能随意随机切分。若同一个实体或同一时间段的相邻事件同时出现在训练集和测试集,模型可能通过重复模式获得虚高结果。

更可靠的切分方式取决于任务:

  • 预测未来:按时间切分;
  • 用户泛化:按用户或实体切分;
  • 新设备泛化:按设备切分;
  • 长度泛化:训练和测试使用不同序列长度区间。

模型、数据、评测和部署的可见信息范围必须一致,否则长依赖结论不可信。


十一、常见失败表现与诊断路径

1. 训练损失下降,长序列准确率不升

可能原因包括:

  • 训练序列太短,模型没有学习长依赖;
  • 训练时使用全量 Teacher Forcing,推理时错误累积;
  • 截断 BPTT 长度小于任务依赖长度;
  • padding 或目标偏移处理错误;
  • 模型通过局部模式投机完成训练集任务。

诊断时可以做长度分桶、打乱中间无关片段、改变关键早期标记、检查梯度范数和隐藏状态相似度。

2. 梯度范数突然变大或变成 NaN

应依次检查:

  1. 输入是否包含非法数值;
  2. 学习率是否过高;
  3. logits 和 loss 是否出现 infnan
  4. 是否需要梯度裁剪;
  5. 序列长度是否突然扩大;
  6. 混合精度训练是否存在溢出;
  7. padding、mask 和标签范围是否正确。

标签用于 CrossEntropyLoss 时必须是合法的整数类别索引;若词表大小为 VV,标签应位于 [0, V-1],除非使用的是专门的忽略索引。

3. LSTM 或 GRU 比 RNN 差

门控结构不是无条件优势。可能原因包括:

  • 数据量不足,复杂模型更难训练;
  • hidden size、学习率或 dropout 未重新调节;
  • 任务很短,门控带来的能力没有价值;
  • 初始化或状态重置不正确;
  • 评测中错误使用了 padding 后的最后位置;
  • 双向模型违反了因果约束。

应先确认输入输出形状、状态维度、序列长度和目标对齐,再比较模型;否则“RNN、LSTM、GRU 谁更好”的结论没有可比性。

4. 隐状态跨样本污染

如果连续处理多个互不相关的序列,却把上一批的隐藏状态直接传给下一批,模型会把不同样本错误连接起来。默认情况下,批次之间应重新初始化状态,或者明确地按照同一条连续流维护状态。

在线流式场景则相反:可以跨窗口传递状态,但必须规定:

  • 什么时候重置状态;
  • 用户、设备或会话是否隔离;
  • 服务重启后状态如何恢复;
  • 并发请求是否可能共享错误状态;
  • 状态是否包含敏感信息,是否需要权限控制和生命周期清理。

状态本身也是生产数据的一部分,不能只把它看成模型内部细节。


十二、生产取舍:模型、状态、数据流和成本

在生产系统中,序列模型至少涉及四类状态:

  1. 模型参数状态:权重、词表和配置;
  2. 请求状态:当前序列的隐藏状态或解码历史;
  3. 数据状态:序列偏移、时间戳、版本和标签;
  4. 评测状态:切分规则、长度分桶和指标版本。

流式推理通常具有如下数据流:

flowchart LR
    E[事件输入] --> V[校验与排序]
    V --> S[会话状态存储]
    S --> M[RNN/LSTM/GRU]
    M --> P[预测或生成]
    P --> Q[质量与安全检查]
    Q --> O[输出]
    M --> S

几个故障路径尤其需要明确:

  • 乱序事件:如果时间步顺序错误,递归状态会沿错误历史更新;
  • 重复事件:同一事件处理两次会改变状态,结果通常不可逆;
  • 服务重启:丢失隐藏状态可能导致会话中断;
  • 并发竞争:同一会话的两个请求同时更新状态,会产生非确定性;
  • 状态过期:长期保留会话状态可能造成内存、隐私和权限风险;
  • 模型版本切换:旧模型产生的隐藏状态未必能被新模型正确解释。

因此,线上状态需要会话键、顺序控制、超时策略、版本绑定和恢复方案。对于无法可靠持久化的状态,可以在故障后重新读取最近一段原始序列并重放;但重放窗口长度会影响恢复成本和恢复后的预测一致性。

成本方面,RNN 的时间计算复杂度通常随序列长度线性增长,但时间维度难以完全并行;LSTM 和 GRU 每一步门控更多,单位时间步计算高于 vanilla RNN。Transformer 训练更易并行,却可能在长上下文上消耗更多显存。实际成本应同时计算:

  • 训练时的序列长度和反向保存;
  • 推理时的每 token 延迟;
  • 并发请求数量;
  • 状态存储与恢复;
  • 失败重试和数据重放;
  • 评测所需的长序列覆盖率。

单看参数量不能准确推断系统成本。


十三、如何选择和验证

可以按照任务约束进行选择,而不是先固定某一种模型:

  • 序列很短、结构简单:vanilla RNN 可能已经足够;
  • 需要更稳妥地处理较长依赖:优先尝试 LSTM 或 GRU;
  • 参数量和延迟较敏感:GRU 常是合理的起点;
  • 需要独立控制长期记忆与当前输出:LSTM 提供更细的结构;
  • 需要离线利用未来上下文:考虑双向结构;
  • 需要严格实时预测:使用单向模型,并验证没有未来信息泄漏;
  • 需要灵活访问很长上下文:比较带注意力或 Transformer 的方案;
  • 需要流式、小内存、可恢复状态:重点评估递归状态的生命周期和故障恢复。

验证时应至少包含四组实验:

  1. 长度实验:逐步增加序列长度;
  2. 依赖距离实验:改变关键信息与目标之间的间隔;
  3. 教师信号实验:比较 Teacher Forcing 训练与自回归推理;
  4. 故障实验:检查乱序、重复、截断、重启和状态重置后的行为。

最终要回答的不是“哪个模型最先进”,而是:

  • 模型是否真正保存了所需历史;
  • 梯度是否能够训练这条记忆路径;
  • 训练输入和推理输入是否一致;
  • 序列长度变化时性能是否稳定;
  • 状态在生产系统中是否可隔离、可恢复、可控成本地运行。

RNN 提供了最直接的递归状态机制;LSTM 和 GRU 通过门控改善记忆与梯度传递;Teacher Forcing 提高了序列训练效率,却引入训练—推理不一致;长依赖则要求同时审视前向信息保存、反向梯度传播、序列截断和评测切分。理解这几条因果链,才能正确判断一个序列模型是在记忆历史,还是只是在利用短期统计模式。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。