固定窗口为何记不住更早的信息?RNN 的隐状态、时间展开与梯度链
从变长序列的固定窗口局限出发,手算 RNN 隐状态,拆解时间反向传播的梯度乘积,并用 PyTorch 2.13 验证张量契约与工程边界。
上一篇用特征金字塔解决了图像中的空间尺度问题:同一张图的不同位置可以并行计算。但文本、语音和传感器流多了一条不能随意打乱的时间轴:「银行」出现在句尾时,它的含义可能取决于很早以前的「河岸」或「存款」。
固定窗口只能看最近的 个输入;增大 又会让参数量、内存和边界处理随窗口改变。循环神经网络(Recurrent Neural Network, RNN)的核心想法是:不把全部历史拼进输入,而是用一个定长隐状态(Hidden State)逐步压缩过去。
本文只追问三件紧密相连的事:隐状态怎样更新,同一个单元怎样沿时间展开,以及这条长链为何会忘记远处信息。
01 固定窗口的不足究竟在哪里?#
假设每个时刻的特征是 ,序列长度是 。一个长度为 的多层感知机(Multilayer Perceptron, MLP)需要拼接:
这会带来三个直接问题:
- 以前的证据必然不可见;
- 修改 会改变第一层参数形状,不能直接处理任意长序列;
- 同一种局部模式出现在不同位置时,普通 MLP 不会自动共享对应参数。
RNN 将历史压缩为 ,再与当前 一起生成 。因此输入接口始终是 ,不随已经看过多少步而变。
固定窗口 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
历史以隐状态的形式向右传递text02 一个 RNN 单元究竟计算什么?#
最经典的 Elman RNN 在时刻 计算:
其中:
- :当前输入;
- :上一步与当前隐状态;
- :输入到隐状态的权重;
- :隐状态到隐状态的循环权重;
- :两条仿射路径的偏置;
- :将每维压到 的双曲正切激活。
若要在每一步输出 类分数,再加一个读出层:
隐状态是「内部记忆」,读出是「任务答案」。二者不要混为一个变量:同一串 可以支持序列分类、逐步标注或下一步预测。
03 「循环」怎样变成可求导的时间展开?#
代码里只有一组 ,计算图里却会出现 个时间节点。这称为时间展开(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 个位置专用什么权重」。不过,展开后的计算必须按 依次发生,因为 依赖 。这也是基础 RNN 不如卷积和 Transformer 容易在时间维并行的根源。
04 用三个标量手算一次记忆更新#
为了只看清数据流,令 ,,并设:
输入序列为 。逐步计算:
时, 仍非零,说明第一步的影响已通过状态传到后面;但 经一步已衰减到 。
若最后的二分类读出为 ,则:
这个概率不是某个输入单独给出的,而是前三步按顺序压入 后的结果。如果交换 和 ,即使元素集合不变,路径和结果也可能改变。
05 放回 batch 后,每个张量是什么形状?#
采用 batch-first 约定时:
| 名称 | 形状 | 含义 |
|---|---|---|
x | [N,T,D] | 个序列,每个 步,每步 维 |
x[:, t, :] | [N,D] | 整个 batch 在时刻 的输入 |
h0 | [L·R,N,H] | 层、 个方向的初始状态 |
output | [N,T,R·H] | 最后一层在每个时刻的状态 |
h_n | [L·R,N,H] | 每层、每个方向的最终状态 |
logits | [N,C] 或 [N,T,C] | 序列级或时间步级的分类分数 |
这里 L=\text{num_layers};单向时 ,双向时 。单层单向 RNN 本体的参数量是:
它不含读出层,也不随 增长。但训练时为反向传播保存的中间状态大致按 增长。
06 不调用 RNN 封装,先写出循环本体#
下面实现单层、单向、tanh RNN。注意参数只创建一次,循环里重用它们。
import torch
from torch import nn
class TransparentRNN(nn.Module):
def __init__(self, input_size: int, hidden_size: int) -> None:
super().__init__()
self.hidden_size = hidden_size
self.weight_ih = nn.Parameter(torch.empty(hidden_size, input_size))
self.weight_hh = nn.Parameter(torch.empty(hidden_size, hidden_size))
self.bias_ih = nn.Parameter(torch.zeros(hidden_size))
self.bias_hh = nn.Parameter(torch.zeros(hidden_size))
nn.init.xavier_uniform_(self.weight_ih)
nn.init.orthogonal_(self.weight_hh)
def forward(
self,
x: torch.Tensor, # [N, T, D]
h0: torch.Tensor | None = None, # [1, N, H]
) -> tuple[torch.Tensor, torch.Tensor]:
assert x.ndim == 3
n, time_steps, _ = x.shape
if h0 is None:
h = x.new_zeros(n, self.hidden_size) # [N, H]
else:
assert h0.shape == (1, n, self.hidden_size)
h = h0[0]
states = []
for t in range(time_steps):
h = torch.tanh(
x[:, t] @ self.weight_ih.T
+ self.bias_ih
+ h @ self.weight_hh.T
+ self.bias_hh
) # [N, H]
states.append(h)
output = torch.stack(states, dim=1) # [N, T, H]
h_n = h.unsqueeze(0) # [1, N, H]
return output, h_n
x = torch.tensor([[[1.0], [0.0], [1.0]]]) # [N=1,T=3,D=1]
model = TransparentRNN(input_size=1, hidden_size=1)
with torch.no_grad():
model.weight_ih.fill_(0.5)
model.weight_hh.fill_(0.5)
output, h_n = model(x)
expected = torch.tensor([[[0.4621], [0.2270], [0.5466]]])
assert output.shape == (1, 3, 1)
assert h_n.shape == (1, 1, 1)
assert torch.allclose(output, expected, atol=1e-4)
assert torch.allclose(output[:, -1], h_n[0])python这段代码没有隐藏任何时间操作:states[t] 就是 ,h_n 只是最后一个状态的分层接口形状。
07 用 PyTorch 2.13 官方 API 验证同一个计算#
PyTorch 2.13 当前的 torch.nn.RNN ↗ 接收 input_size、hidden_size、num_layers、nonlinearity、batch_first、dropout 和 bidirectional 等参数。下面把手写模型的参数复制给官方实现,而不是只相信形状一样:
import torch
from torch import nn
torch.manual_seed(7)
x = torch.randn(2, 4, 3) # [N=2,T=4,D=3]
manual = TransparentRNN(input_size=3, hidden_size=5)
official = nn.RNN(
input_size=3,
hidden_size=5,
num_layers=1,
nonlinearity="tanh",
batch_first=True,
bidirectional=False,
)
with torch.no_grad():
official.weight_ih_l0.copy_(manual.weight_ih)
official.weight_hh_l0.copy_(manual.weight_hh)
official.bias_ih_l0.copy_(manual.bias_ih)
official.bias_hh_l0.copy_(manual.bias_hh)
manual_output, manual_hn = manual(x)
output, h_n = official(x)
assert output.shape == (2, 4, 5)
assert h_n.shape == (1, 2, 5)
torch.testing.assert_close(output, manual_output)
torch.testing.assert_close(h_n, manual_hn)python官方契约中,output 包含最后一层的每个时刻,h_n 包含每层、每个方向的最终状态。单层单向时 output[:, -1] == h_n[0];双向、多层或变长 batch 中不要无条件照搬这条索引。
还有三个 API 细节容易被误读:
dropout>0只作用在相邻 RNN 层之间,最后一层之后不用;单层 RNN 不会因此获得时间步间 dropout。bidirectional=True会让输出最后一维变为 ;反向分支使用未来输入,不适用于严格在线预测。- 某些 cuDNN/CUDA 组合的 RNN 运算存在已知非确定性;复现问题时要记录软硬件版本和确定性设置。
08 一个序列分类器的训练与推理数据流#
假设任务是判断一段长度固定的传感器序列是否异常,每步 个特征,最后输出 类。
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])textimport torch
from torch import nn
class SequenceClassifier(nn.Module):
def __init__(self, input_size: int = 8, hidden_size: int = 32) -> None:
super().__init__()
self.rnn = nn.RNN(
input_size=input_size,
hidden_size=hidden_size,
batch_first=True,
)
self.head = nn.Linear(hidden_size, 2)
def forward(self, x: torch.Tensor) -> torch.Tensor:
assert x.ndim == 3 and x.shape[-1] == self.rnn.input_size
_, h_n = self.rnn(x) # [1,N,H]
return self.head(h_n[-1]) # [N,2]
model = SequenceClassifier()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
x = torch.randn(16, 40, 8) # [N,T,D]
y = torch.randint(0, 2, size=(16,)) # [N], int64
model.train()
optimizer.zero_grad(set_to_none=True)
logits = model(x) # [N,C]
loss = nn.functional.cross_entropy(logits, y)
loss.backward()
total_norm = nn.utils.clip_grad_norm_(
model.parameters(), max_norm=1.0, error_if_nonfinite=True
)
optimizer.step()
model.eval()
with torch.inference_mode():
probabilities = model(x[:2]).softmax(dim=-1) # [2,2]
predictions = probabilities.argmax(dim=-1) # [2]pythoncross_entropy 直接接收未经 softmax 的 logits。当前 clip_grad_norm_ ↗ 会将全部参数梯度视为一个连接向量计算总范数,就地修改梯度,并返回裁剪前的总范数。因此应当在 backward() 之后、step() 之前调用,并把返回值写入训练日志。
09 时间反向传播为何是一串乘法?#
将普通反向传播应用到展开图,就得到时间反向传播(Backpropagation Through Time, BPTT)。若损失 只依赖最后状态 ,较早状态 收到的信号是:
对 tanh RNN:
因此长程信号会反复乘上 和 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在标量、线性化的极简情况中,若每步局部导数都近似 ,跨 10 步只剩:
跨 50 步则约为 ,这就是梯度消失(Vanishing Gradient):模型几乎收不到「应该修改很早状态」的信号。反之,若连乘方向的增益持续大于 1,梯度会指数增长,形成梯度爆炸(Exploding Gradient)。
tanh 饱和时 接近 0,会进一步截断梯度。因此「状态数值没有变成 0」并不能证明长程依赖正在被学习;还必须检查对早期输入和状态的梯度。
10 用一段最小代码观察梯度乘积#
下面先去掉输入和激活,只保留 ,便可直接验证 。这不是完整 RNN,而是隔离「连乘」的调试实验。
import torch
def gradient_across_time(weight: float, steps: int) -> float:
h0 = torch.tensor(1.0, requires_grad=True)
h = h0
for _ in range(steps):
h = weight * h
h.backward()
assert h0.grad is not None
return h0.grad.item()
assert abs(gradient_across_time(0.5, 10) - 0.5**10) < 1e-9
assert abs(gradient_across_time(1.5, 10) - 1.5**10) < 1e-4
for steps in (1, 5, 10, 50):
print(steps, gradient_across_time(0.5, steps))python真实 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当连续流太长时,常把它切成长度 的块,每块传入上一块的数值状态,但用 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这里传给下一块的是状态的数值,所以推理上下文仍连续;被切断的是梯度路径,因此训练信号最多跨 步。detach() 不是解决长程依赖的方法,而是用优化偏差换内存和吞吐。若每个 chunk 都把 h 重置为零,连前向记忆也一起丢了;若从不 detach(),图会持续增长,第二次反向还可能触发已释放图错误。
12 变长序列为什么不能直接拿最后一列?#
一个 batch 的真实长度可能是 [7,4,2]。补零后张量是 [3,7,D],但 output[:, -1] 对后两个样本对应的是 padding 位置,不是真实末尾。至少有两种正确路径:
- 用长度索引每个样本的
output[n, length[n]-1],同时确保 padding 不污染后续计算; - 用打包序列(Packed Sequence)让 RNN 跳过 padding。
PyTorch 2.13 当前的 pack_padded_sequence ↗ 在 batch_first=True 时接收 [N,T,*];若 lengths 是张量,它必须位于 CPU。enforce_sorted=False 可接受未按长度降序排列的 batch。
import torch
from torch import nn
from torch.nn.utils.rnn import pack_padded_sequence
rnn = nn.RNN(input_size=8, hidden_size=16, batch_first=True)
x_padded = torch.randn(3, 7, 8) # [N=3,T_max=7,D=8]
lengths = torch.tensor([7, 4, 2]) # CPU int tensor
packed = pack_padded_sequence(
x_padded,
lengths.cpu(),
batch_first=True,
enforce_sorted=False,
)
packed_output, h_n = rnn(packed)
assert h_n.shape == (1, 3, 16) # each sample's real final statepython打包只解决无效 padding 计算与末尾定位,不会解决梯度消失。做逐步损失时还要明确标签的 padding 掩码(Mask)或同样打包标签,不能让补零位置进入损失均值。
13 一条可执行的 RNN 调试路径#
- 先固定形状契约。 在入口断言
x=[N,T,D]、lengths=[N]、h0=[L·R,N,H];不要用变量名batch同时代指数据和 batch size。 - 用三步序列对齐手写实现。 将手写单元参数复制到
nn.RNN,比较全部output,而不是只比较最后一项。 - 做顺序敏感性测试。 对同一批样本分别输入原序列与时间反转序列;若任务应依赖顺序而输出始终相同,检查维度是否被错误求和或展平。
- 做记忆探针。 构造“第一位决定标签,中间全是噪声”的任务,将 从 5 增到 100,画准确率和早期输入梯度随距离的变化。
- 记录裁剪前梯度。
clip_grad_norm_返回的总范数才显示爆炸是否发生;只记录裁剪后范数会把所有异常伪装成阈值。 - 区分数值状态与计算图。 流式推理用
inference_mode()并显式传递状态;截断训练在 chunk 边界detach();独立样本之间必须重置状态。 - 先过拟合一个极小 batch。 若 8 个短序列都不能把训练损失压低,优先检查标签对齐、最后状态索引、损失输入和
zero_grad,不要先加层数。
14 最常见的“能运行,但序列语义错了”#
- 把
[N,T,D]送给默认batch_first=False。 模型会把 当时间、 当 batch,形状有时仍能通过。 - 认为
batch_first也会改变h_n。 于是错误地读取h_n[:, -1],拿到的是最后一个样本而不是最后一层。 - 变长序列直接取
output[:, -1]。 短样本读到 padding 后状态。 - 在线任务使用双向 RNN。 离线验证很好,上线时却需要尚未到达的未来输入。
- 把 batch 之间的状态无条件复用。 若样本互不相关,这会把上一位用户的信息泄漏给下一位用户。
- 切 chunk 时每次清零状态。 这把有效上下文上限硬性改成 chunk 长度。
- 长期不
detach()。 连续流的计算图与显存不断增长,或在重复反向时出错。 - 只裁剪梯度却不记录裁剪比例。 模型可能每一步都撞上阈值,训练看似稳定,实际更新方向长期失真。
- 分类前先
softmax再传给交叉熵。 破坏数值稳定的 logits 接口,并改变梯度。 - 把
h_n当成人类可读摘要。 隐状态坐标由任务共同学习,单维数值通常没有稳定语义。
15 它会在哪些场景失败?#
- 很长的精确记忆。 vanilla RNN 很难把早期一个比特可靠保留几百步,BPTT 的连乘会让学习信号消失或爆炸。
- 必须保留大量细节。 固定 的状态是瓶颈;长文档、长音频或多事件流会竞争有限容量。
- 需要大规模并行训练。 对 的依赖限制了时间维并行,长序列吞吐通常不如卷积或自注意力结构。
- 不规则时间间隔。 普通 RNN 默认相邻步间隔等价;医疗记录或事件流需要显式加入时间差,甚至使用连续时间模型。
- 分布漂移的流式状态。 状态会积累旧分布影响;缺少重置、超时和会话边界时,错误可跨很长时间传播。
- 需要解释具体证据位置。 单一最终状态不直接告诉人们答案主要来自哪个时间步,必须另加探针、注意力或归因分析。
梯度裁剪只能限制爆炸,不能把已接近 0 的梯度放大成有用信号;正交初始化可改善早期训练,也不能保证 tanh 长链永久保真。这些是缓解手段,不是结构性保证。
16 与相近序列方法的边界#
| 方法 | 怎样读取历史 | 时间维并行 | 长程信息的主要瓶颈 |
|---|---|---|---|
| 固定窗口 MLP | 拼接最近 步 | 高 | 之外绝对不可见 |
| 一维因果卷积 | 局部卷积核逐层扩大感受野 | 高 | 感受野由深度、卷积核和 dilation 决定 |
| vanilla RNN | 单一隐状态逐步递推 | 低 | 状态瓶颈与雅可比连乘 |
| 长短期记忆网络 LSTM | 门控单元状态与隐状态 | 低 | 门控仍可能饱和,且顺序计算仍存在 |
| 门控循环单元 GRU | 更紧凑的更新/重置门状态 | 低 | 与 LSTM 类似但状态接口更少 |
| Transformer | 每个位置直接聚合其他位置 | 训练时高 | 标准全局注意力的时间/显存随 增长 |
vanilla RNN 的价值不只在于今天是否是最强模型。它把“状态、共享转移、时间展开、BPTT”放进最小系统;LSTM、GRU、状态空间模型和自回归推理都在不同程度上继承或改造这些问题。
17 今天真正需要记住什么?#
- RNN 用 将任意长历史压入定长状态;参数不随序列长度增长,但激活内存和顺序计算会增长。
- 循环单元在代码中只有一份,沿时间展开后成为共享参数的长计算图;
output与h_n表达不同接口。 - BPTT 让远处梯度反复乘以循环雅可比;乘积持续收缩会消失,持续放大会爆炸。
- 梯度裁剪、截断 BPTT 和打包变长序列分别处理爆炸、图长度和 padding,它们不能互相替代。
- 调试 RNN 要同时检查时间顺序、长度、状态边界与梯度随距离的变化,不能只看最终损失。
18 思考题与小练习#
- 令 、、、,手算输入 的三个隐状态。再将输入反转为 ,解释为何元素相同而最终状态不同。
- 将
TransparentRNN扩展为返回每一步预激活 ,对 保留早期状态梯度并画范数。分别尝试循环权重缩放为 0.5、1.0 和 1.5,观察tanh饱和如何改变纯线性结论。 - 构造真实长度
[7,4,2]的 batch,分别用长度索引和pack_padded_sequence取得最终状态,验证两者一致;然后故意使用output[:, -1],定位短样本偏差从哪一步开始。
相关工作#
- Elman (1990), Finding Structure in Time ↗:展示简单循环网络如何通过上下文单元学习序列结构。
- Rumelhart, Hinton & Williams (1986), Learning Representations by Back-propagating Errors ↗:系统阐述多层网络的反向传播学习机制。
- Werbos (1990), Backpropagation Through Time: What It Does and How to Do It ↗:给出时间展开网络的反向传播分析与实践说明。
- Bengio, Simard & Frasconi (1994), Learning Long-Term Dependencies with Gradient Descent Is Difficult ↗:分析梯度方法学习长程依赖的根本困难。
- Pascanu, Mikolov & Bengio (2013), On the Difficulty of Training Recurrent Neural Networks ↗:从几何角度分析消失/爆炸梯度并讨论范数裁剪。
19 下一篇预告#
vanilla RNN 已经给了历史一条通路,却让信息和梯度每一步都必须穿过同一个非线性变换。下一篇将继续追问:长短期记忆网络(Long Short-Term Memory, LSTM)与门控循环单元(Gated Recurrent Unit, GRU)怎样用门控加法路径决定何时写入、保留和遗忘。