观文听傑

返回

上一篇用特征金字塔解决了图像中的空间尺度问题:同一张图的不同位置可以并行计算。但文本、语音和传感器流多了一条不能随意打乱的时间轴:「银行」出现在句尾时,它的含义可能取决于很早以前的「河岸」或「存款」。

固定窗口只能看最近的 KK 个输入;增大 KK 又会让参数量、内存和边界处理随窗口改变。循环神经网络(Recurrent Neural Network, RNN)的核心想法是:不把全部历史拼进输入,而是用一个定长隐状态(Hidden State)逐步压缩过去。

本文只追问三件紧密相连的事:隐状态怎样更新,同一个单元怎样沿时间展开,以及这条长链为何会忘记远处信息。

01 固定窗口的不足究竟在哪里?#

假设每个时刻的特征是 xtRDx_t\in\mathbb{R}^{D},序列长度是 TT。一个长度为 KK 的多层感知机(Multilayer Perceptron, MLP)需要拼接:

zt=[xtK+1;;xt]RKDz_t=[x_{t-K+1};\ldots;x_t]\in\mathbb{R}^{KD}

这会带来三个直接问题:

  1. tKt-K 以前的证据必然不可见;
  2. 修改 KK 会改变第一层参数形状,不能直接处理任意长序列;
  3. 同一种局部模式出现在不同位置时,普通 MLP 不会自动共享对应参数。

RNN 将历史压缩为 ht1RHh_{t-1}\in\mathbb{R}^{H},再与当前 xtx_t 一起生成 hth_t。因此输入接口始终是 D+HD+H,不随已经看过多少步而变。

固定窗口 K=3: x1  x2  x3  x4  x5  x6
                              └─── [x4,x5,x6] ──► 预测
                 x1,x2,x3 的细节不可达

RNN:            x1 ─► h1 ─► h2 ─► h3 ─► h4 ─► h5 ─► h6
                         ▲     ▲     ▲     ▲     ▲
                         x2    x3    x4    x5    x6
                 历史以隐状态的形式向右传递
text

02 一个 RNN 单元究竟计算什么?#

最经典的 Elman RNN 在时刻 tt 计算:

at=Wxhxt+bxh+Whhht1+bhha_t=W_{xh}x_t+b_{xh}+W_{hh}h_{t-1}+b_{hh} ht=tanh(at)h_t=\tanh(a_t)

其中:

  • xtRDx_t\in\mathbb{R}^{D}:当前输入;
  • ht1,htRHh_{t-1},h_t\in\mathbb{R}^{H}:上一步与当前隐状态;
  • WxhRH×DW_{xh}\in\mathbb{R}^{H\times D}:输入到隐状态的权重;
  • WhhRH×HW_{hh}\in\mathbb{R}^{H\times H}:隐状态到隐状态的循环权重;
  • bxh,bhhRHb_{xh},b_{hh}\in\mathbb{R}^{H}:两条仿射路径的偏置;
  • tanh\tanh:将每维压到 (1,1)(-1,1) 的双曲正切激活。

若要在每一步输出 CC 类分数,再加一个读出层:

ot=Whyht+by,WhyRC×Ho_t=W_{hy}h_t+b_y,\qquad W_{hy}\in\mathbb{R}^{C\times H}

隐状态是「内部记忆」,读出是「任务答案」。二者不要混为一个变量:同一串 hth_t 可以支持序列分类、逐步标注或下一步预测。

03 「循环」怎样变成可求导的时间展开?#

代码里只有一组 Wxh,WhhW_{xh},W_{hh},计算图里却会出现 TT 个时间节点。这称为时间展开(Unrolling Through Time):

                  同一组 W_xh, W_hh 被重复使用

h0 ──► [ RNN cell ] ──► h1 ──► [ RNN cell ] ──► h2 ──► [ RNN cell ] ──► h3
           ▲                            ▲                            ▲
           x1                           x2                           x3
           │                            │                            │
          o1                           o2                           o3 ──► loss

参数量不随 T 增长;中间激活和反向路径却随 T 增长。
text

共享参数让模型学的是「遇到一个新输入时如何更新记忆」,而不是「第 17 个位置专用什么权重」。不过,展开后的计算必须按 t=1,2,,Tt=1,2,\ldots,T 依次发生,因为 hth_t 依赖 ht1h_{t-1}。这也是基础 RNN 不如卷积和 Transformer 容易在时间维并行的根源。

04 用三个标量手算一次记忆更新#

为了只看清数据流,令 D=H=1D=H=1h0=0h_0=0,并设:

Wxh=0.5,quadWhh=0.5,quadbxh=bhh=0W_{xh}=0.5,quad W_{hh}=0.5,quad b_{xh}=b_{hh}=0

输入序列为 (x1,x2,x3)=(1,0,1)(x_1,x_2,x_3)=(1,0,1)。逐步计算:

h1=tanh(0.5×1+0.5×0)0.4621h_1=\tanh(0.5\times1+0.5\times0)\approx0.4621 h2=tanh(0.5×0+0.5×0.4621)0.2270h_2=\tanh(0.5\times0+0.5\times0.4621)\approx0.2270 h3=tanh(0.5×1+0.5×0.2270)0.5466h_3=\tanh(0.5\times1+0.5\times0.2270)\approx0.5466

x2=0x_2=0 时,h2h_2 仍非零,说明第一步的影响已通过状态传到后面;但 0.46210.4621 经一步已衰减到 0.22700.2270

若最后的二分类读出为 z=2h30.5z=2h_3-0.5,则:

z0.5932,σ(z)0.6441z\approx0.5932,\qquad \sigma(z)\approx0.6441

这个概率不是某个输入单独给出的,而是前三步按顺序压入 h3h_3 后的结果。如果交换 x1x_1x3x_3,即使元素集合不变,路径和结果也可能改变。

05 放回 batch 后,每个张量是什么形状?#

采用 batch-first 约定时:

名称形状含义
x[N,T,D]NN 个序列,每个 TT 步,每步 DD
x[:, t, :][N,D]整个 batch 在时刻 tt 的输入
h0[L·R,N,H]LL 层、RR 个方向的初始状态
output[N,T,R·H]最后一层在每个时刻的状态
h_n[L·R,N,H]每层、每个方向的最终状态
logits[N,C][N,T,C]序列级或时间步级的分类分数

这里 L=\text{num_layers};单向时 R=1R=1,双向时 R=2R=2。单层单向 RNN 本体的参数量是:

HD+H2+2HHD+H^2+2H

它不含读出层,也不随 TT 增长。但训练时为反向传播保存的中间状态大致按 O(NTH)O(NTH) 增长。

06 不调用 RNN 封装,先写出循环本体#

下面实现单层、单向、tanh RNN。注意参数只创建一次,循环里重用它们。

这段代码没有隐藏任何时间操作:states[t] 就是 ht+1h_{t+1}h_n 只是最后一个状态的分层接口形状。

07 用 PyTorch 2.13 官方 API 验证同一个计算#

PyTorch 2.13 当前的 torch.nn.RNN 接收 input_sizehidden_sizenum_layersnonlinearitybatch_firstdropoutbidirectional 等参数。下面把手写模型的参数复制给官方实现,而不是只相信形状一样:

官方契约中,output 包含最后一层的每个时刻,h_n 包含每层、每个方向的最终状态。单层单向时 output[:, -1] == h_n[0];双向、多层或变长 batch 中不要无条件照搬这条索引。

还有三个 API 细节容易被误读:

  • dropout>0 只作用在相邻 RNN 层之间,最后一层之后不用;单层 RNN 不会因此获得时间步间 dropout。
  • bidirectional=True 会让输出最后一维变为 2H2H;反向分支使用未来输入,不适用于严格在线预测。
  • 某些 cuDNN/CUDA 组合的 RNN 运算存在已知非确定性;复现问题时要记录软硬件版本和确定性设置。

08 一个序列分类器的训练与推理数据流#

假设任务是判断一段长度固定的传感器序列是否异常,每步 D=8D=8 个特征,最后输出 C=2C=2 类。

x [N,T,8]


RNN(input_size=8, hidden_size=32, batch_first=True)

    ├─ output [N,T,32]   逐步标注时使用
    └─ h_n [1,N,32]


         Linear(32,2)


         logits [N,2] ─► CrossEntropyLoss(logits, y[N])
text

cross_entropy 直接接收未经 softmax 的 logits。当前 clip_grad_norm_ 会将全部参数梯度视为一个连接向量计算总范数,就地修改梯度,并返回裁剪前的总范数。因此应当在 backward() 之后、step() 之前调用,并把返回值写入训练日志。

09 时间反向传播为何是一串乘法?#

将普通反向传播应用到展开图,就得到时间反向传播(Backpropagation Through Time, BPTT)。若损失 L\mathcal L 只依赖最后状态 hTh_T,较早状态 hkh_k 收到的信号是:

Lhk=LhTt=k+1Ththt1\frac{\partial \mathcal L}{\partial h_k} = \frac{\partial \mathcal L}{\partial h_T} \prod_{t=k+1}^{T} \frac{\partial h_t}{\partial h_{t-1}}

tanh RNN:

htht1=diag(1ht2)Whh\frac{\partial h_t}{\partial h_{t-1}} = \operatorname{diag}(1-h_t^2)W_{hh}

因此长程信号会反复乘上 WhhW_{hh}tanh 导数。直觉上:

loss
  │ × J_T
 h_T
  │ × J_(T-1)
h_(T-1)
  │ × J_(T-2)       J_t = diag(1 - h_t²) W_hh

  │ × J_(k+1)
 h_k        每穿过一步就再乘一次
text

在标量、线性化的极简情况中,若每步局部导数都近似 0.50.5,跨 10 步只剩:

0.5109.77×1040.5^{10}\approx9.77\times10^{-4}

跨 50 步则约为 8.88×10168.88\times10^{-16},这就是梯度消失(Vanishing Gradient):模型几乎收不到「应该修改很早状态」的信号。反之,若连乘方向的增益持续大于 1,梯度会指数增长,形成梯度爆炸(Exploding Gradient)。

tanh 饱和时 1ht21-h_t^2 接近 0,会进一步截断梯度。因此「状态数值没有变成 0」并不能证明长程依赖正在被学习;还必须检查对早期输入和状态的梯度。

10 用一段最小代码观察梯度乘积#

下面先去掉输入和激活,只保留 ht=wht1h_t=wh_{t-1},便可直接验证 hT/h0=wT\partial h_T/\partial h_0=w^T。这不是完整 RNN,而是隔离「连乘」的调试实验。

真实 RNN 中的雅可比是矩阵,不能只看某个权重元素是否小于 1。更可操作的方法是:在人工长程任务上改变延迟步数,记录梯度总范数、分层梯度、早期输入梯度与准确率随距离的曲线。

11 完整 BPTT 与截断 BPTT 分别做了什么?#

完整 BPTT 的逻辑是:

h = h0
states = []
for t = 1 ... T:
    h = rnn_cell(x[t], h)   # 前向时保留计算图
    states.append(h)

loss = task_loss(states, targets)
loss.backward()             # 从后向前穿过全部 T 步
clip_grad_norm(parameters)
optimizer.step()
text

当连续流太长时,常把它切成长度 KK 的块,每块传入上一块的数值状态,但用 detach() 切断跨块计算图。这是截断时间反向传播(Truncated BPTT):

h = None
for x_chunk, y_chunk in stream:  # x_chunk: [N,K,D]
    optimizer.zero_grad(set_to_none=True)
    output, h = model.rnn(x_chunk, h)
    logits = model.head(output)  # [N,K,C]
    loss = nn.functional.cross_entropy(
        logits.reshape(-1, logits.shape[-1]),
        y_chunk.reshape(-1),
    )
    loss.backward()
    nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    optimizer.step()
    h = h.detach()
python

这里传给下一块的是状态的数值,所以推理上下文仍连续;被切断的是梯度路径,因此训练信号最多跨 KK 步。detach() 不是解决长程依赖的方法,而是用优化偏差换内存和吞吐。若每个 chunk 都把 h 重置为零,连前向记忆也一起丢了;若从不 detach(),图会持续增长,第二次反向还可能触发已释放图错误。

12 变长序列为什么不能直接拿最后一列?#

一个 batch 的真实长度可能是 [7,4,2]。补零后张量是 [3,7,D],但 output[:, -1] 对后两个样本对应的是 padding 位置,不是真实末尾。至少有两种正确路径:

  1. 用长度索引每个样本的 output[n, length[n]-1],同时确保 padding 不污染后续计算;
  2. 用打包序列(Packed Sequence)让 RNN 跳过 padding。

PyTorch 2.13 当前的 pack_padded_sequencebatch_first=True 时接收 [N,T,*];若 lengths 是张量,它必须位于 CPU。enforce_sorted=False 可接受未按长度降序排列的 batch。

打包只解决无效 padding 计算与末尾定位,不会解决梯度消失。做逐步损失时还要明确标签的 padding 掩码(Mask)或同样打包标签,不能让补零位置进入损失均值。

13 一条可执行的 RNN 调试路径#

  1. 先固定形状契约。 在入口断言 x=[N,T,D]lengths=[N]h0=[L·R,N,H];不要用变量名 batch 同时代指数据和 batch size。
  2. 用三步序列对齐手写实现。 将手写单元参数复制到 nn.RNN,比较全部 output,而不是只比较最后一项。
  3. 做顺序敏感性测试。 对同一批样本分别输入原序列与时间反转序列;若任务应依赖顺序而输出始终相同,检查维度是否被错误求和或展平。
  4. 做记忆探针。 构造“第一位决定标签,中间全是噪声”的任务,将 TT 从 5 增到 100,画准确率和早期输入梯度随距离的变化。
  5. 记录裁剪前梯度。 clip_grad_norm_ 返回的总范数才显示爆炸是否发生;只记录裁剪后范数会把所有异常伪装成阈值。
  6. 区分数值状态与计算图。 流式推理用 inference_mode() 并显式传递状态;截断训练在 chunk 边界 detach();独立样本之间必须重置状态。
  7. 先过拟合一个极小 batch。 若 8 个短序列都不能把训练损失压低,优先检查标签对齐、最后状态索引、损失输入和 zero_grad,不要先加层数。

14 最常见的“能运行,但序列语义错了”#

  • [N,T,D] 送给默认 batch_first=False 模型会把 NN 当时间、TT 当 batch,形状有时仍能通过。
  • 认为 batch_first 也会改变 h_n 于是错误地读取 h_n[:, -1],拿到的是最后一个样本而不是最后一层。
  • 变长序列直接取 output[:, -1] 短样本读到 padding 后状态。
  • 在线任务使用双向 RNN。 离线验证很好,上线时却需要尚未到达的未来输入。
  • 把 batch 之间的状态无条件复用。 若样本互不相关,这会把上一位用户的信息泄漏给下一位用户。
  • 切 chunk 时每次清零状态。 这把有效上下文上限硬性改成 chunk 长度。
  • 长期不 detach() 连续流的计算图与显存不断增长,或在重复反向时出错。
  • 只裁剪梯度却不记录裁剪比例。 模型可能每一步都撞上阈值,训练看似稳定,实际更新方向长期失真。
  • 分类前先 softmax 再传给交叉熵。 破坏数值稳定的 logits 接口,并改变梯度。
  • h_n 当成人类可读摘要。 隐状态坐标由任务共同学习,单维数值通常没有稳定语义。

15 它会在哪些场景失败?#

  • 很长的精确记忆。 vanilla RNN 很难把早期一个比特可靠保留几百步,BPTT 的连乘会让学习信号消失或爆炸。
  • 必须保留大量细节。 固定 HH 的状态是瓶颈;长文档、长音频或多事件流会竞争有限容量。
  • 需要大规模并行训练。 hth_tht1h_{t-1} 的依赖限制了时间维并行,长序列吞吐通常不如卷积或自注意力结构。
  • 不规则时间间隔。 普通 RNN 默认相邻步间隔等价;医疗记录或事件流需要显式加入时间差,甚至使用连续时间模型。
  • 分布漂移的流式状态。 状态会积累旧分布影响;缺少重置、超时和会话边界时,错误可跨很长时间传播。
  • 需要解释具体证据位置。 单一最终状态不直接告诉人们答案主要来自哪个时间步,必须另加探针、注意力或归因分析。

梯度裁剪只能限制爆炸,不能把已接近 0 的梯度放大成有用信号;正交初始化可改善早期训练,也不能保证 tanh 长链永久保真。这些是缓解手段,不是结构性保证。

16 与相近序列方法的边界#

方法怎样读取历史时间维并行长程信息的主要瓶颈
固定窗口 MLP拼接最近 KKKK 之外绝对不可见
一维因果卷积局部卷积核逐层扩大感受野感受野由深度、卷积核和 dilation 决定
vanilla RNN单一隐状态逐步递推状态瓶颈与雅可比连乘
长短期记忆网络 LSTM门控单元状态与隐状态门控仍可能饱和,且顺序计算仍存在
门控循环单元 GRU更紧凑的更新/重置门状态与 LSTM 类似但状态接口更少
Transformer每个位置直接聚合其他位置训练时高标准全局注意力的时间/显存随 T2T^2 增长

vanilla RNN 的价值不只在于今天是否是最强模型。它把“状态、共享转移、时间展开、BPTT”放进最小系统;LSTM、GRU、状态空间模型和自回归推理都在不同程度上继承或改造这些问题。

17 今天真正需要记住什么?#

  1. RNN 用 ht=f(xt,ht1)h_t=f(x_t,h_{t-1}) 将任意长历史压入定长状态;参数不随序列长度增长,但激活内存和顺序计算会增长。
  2. 循环单元在代码中只有一份,沿时间展开后成为共享参数的长计算图;outputh_n 表达不同接口。
  3. BPTT 让远处梯度反复乘以循环雅可比;乘积持续收缩会消失,持续放大会爆炸。
  4. 梯度裁剪、截断 BPTT 和打包变长序列分别处理爆炸、图长度和 padding,它们不能互相替代。
  5. 调试 RNN 要同时检查时间顺序、长度、状态边界与梯度随距离的变化,不能只看最终损失。

18 思考题与小练习#

  1. D=H=1D=H=1Wxh=1W_{xh}=1Whh=0.25W_{hh}=0.25h0=0h_0=0,手算输入 (1,1,0)(1,1,0) 的三个隐状态。再将输入反转为 (0,1,1)(0,1,1),解释为何元素相同而最终状态不同。
  2. TransparentRNN 扩展为返回每一步预激活 ata_t,对 T{5,20,50}T\in\{5,20,50\} 保留早期状态梯度并画范数。分别尝试循环权重缩放为 0.5、1.0 和 1.5,观察 tanh 饱和如何改变纯线性结论。
  3. 构造真实长度 [7,4,2] 的 batch,分别用长度索引和 pack_padded_sequence 取得最终状态,验证两者一致;然后故意使用 output[:, -1],定位短样本偏差从哪一步开始。

相关工作#

19 下一篇预告#

vanilla RNN 已经给了历史一条通路,却让信息和梯度每一步都必须穿过同一个非线性变换。下一篇将继续追问:长短期记忆网络(Long Short-Term Memory, LSTM)与门控循环单元(Gated Recurrent Unit, GRU)怎样用门控加法路径决定何时写入、保留和遗忘。

固定窗口为何记不住更早的信息?RNN 的隐状态、时间展开与梯度链
https://zwjcode.cn/blog/rnn-hidden-state-time-unrolling-gradient-memory
作者
发布于 2026年9月5日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。