↓ Skip to main content
  1. Posts/

从 RNN 到 LSTM,再到 GRU:循环网络真正改变的是什么?

·3008 words·7 mins
Note: This article is available in Chinese only. 本文暂无英文版本。 View original

最近复习 RNN、LSTM 和 GRU,我把它们的状态更新公式放在一起看。有一处值得单独想想:candidate 和 input gate 最后只是相乘,为什么要分成两组参数来学?从这个问题往回看,才能看清几个 gate 各自控制了什么。

Vanilla RNN:问题出在哪里
#

最朴素的 Vanilla RNN 的状态更新方程非常直观:

$$ h_t = \tanh(W_x x_t + W_h h_{t-1} + b) $$
Vanilla RNN 的状态更新路径

对比 LSTM 和 GRU 时,我关注的是历史状态经过哪条路径传到下一步:

  1. 历史信息全部被压缩在 \( h_t \) 中:无论是五步前还是五十步前的信息,都必须挤在这一个 hidden state 向量里。
  2. 每一步经历完全相同的状态转换(state transition):无论当前 timestep 遇到了什么输入,历史状态 \( h_{t-1} \) 都会无条件地与 \( W_h \) 相乘。
  3. 参数 \( W_h \) 跨 timestep 共享:这是循环神经网络的本质,但这也意味着网络对历史的“处理方式”是静态的。
  4. 缺乏显式的控制机制:vanilla RNN 没有任何机制来显式决定哪些历史应该保留、哪些应该忘记、或者当前时刻应该写入什么新信息。

这种机制在长序列反向传播(BPTT)时很容易导致梯度消失(gradient vanishing)或梯度爆炸(gradient exploding)。沿时间展开后,对应的是 Jacobian 的连乘:

$$ \frac{\partial h_t}{\partial h_{t-k}} = \prod_j \frac{\partial h_j}{\partial h_{j-1}} $$

每一次状态传递都会连乘一次 \( W_h \) 和激活函数的导数。如果连乘的项大多小于 1,梯度很快就会衰减到零,导致模型根本无法学习到远距离的依赖关系。

LSTM:多出来的 cell state
#

为了解决这个问题,LSTM 被提了出来。先看它维护的状态。Vanilla RNN 主要维护一个状态:

$$ h_t $$

而 LSTM 维护了两个状态:

$$ c_t,\quad h_t $$

这里把 \( c_t \)(cell state)理解为更适合长期传播的 memory path,而把 \( h_t \) 理解为当前的 hidden 或 output state。需要提醒的是,“长期记忆”和“短期记忆”只是为了方便理解的比喻,并不是严谨的数学语义。

LSTM 核心的状态更新公式如下:

$$ c_t = f_t \odot c_{t-1} + i_t \odot \tilde c_t $$

旧状态经过 forget gate 保留一部分,candidate 经过 input gate 写入一部分,两者相加得到新的 cell state。

其中:

  • \( f_t \)(forget gate):决定旧信息 \( c_{t-1} \) 保留多少;
  • \( i_t \)(input gate):决定新信息写多少;
  • \( \tilde c_t \)(candidate memory):准备写进去的新内容。

在计算出当前时刻的 cell state \( c_t \) 后,LSTM 还需要对外输出当前的 hidden state:

$$ h_t = o_t \odot \tanh(c_t) $$

这里的 \( o_t \)(output gate)控制着内部的 \( c_t \) 有多少会被暴露为当前的 \( h_t \)。

LSTM 内部结构详解

为什么 candidate 和 input gate 要分开
#

输入阶段的两个变量分别使用 sigmoid 和 tanh:

一个是真正的 gate:

$$ i_t = \sigma(W_i x_t + U_i h_{t-1} + b_i) $$

另一个是 candidate memory,它并不是一个 gate,而是包含了实际内容:

$$ \tilde c_t = \tanh(W_c x_t + U_c h_{t-1} + b_c) $$

两者的分工是:

  • \( \tilde c_t \) 决定了写什么。
  • \( i_t \) 决定了写多少。

既然最后都是要把这两项做 element-wise 乘法(\( i_t \odot \tilde c_t \)),为什么不能直接用一个变量表示最终的 update,而非要拆开呢?

用一个函数直接生成 update 也是一种设计。LSTM 选择拆开,是在 parameterization 中显式区分“候选内容生成”和“是否写入、写多少”。网络可以独立学习如何产生有用的特征(content),以及在当前上下文中这些特征是否有用(control)。

LSTM 为什么有利于梯度传播
#

沿着 cell state 的直接保留路径来看:

$$ c_t = f_t \odot c_{t-1} + \cdots $$
LSTM 的状态更新路径

只看这条直接路径,逐元素的导数由 forget gate 给出;下面省略 gate 经由 hidden state 产生的间接依赖:

$$ \frac{\partial c_t}{\partial c_{t-1}} \approx f_t $$

当跨越多个 timestep 时:

$$ \frac{\partial c_t}{\partial c_{t-k}} \approx \prod_j f_j $$

当模型学到了 \( f_j \approx 1 \)(即网络决定保留历史记忆)时,梯度在反向传播时就相当于乘上了一连串的 1。这意味着,cell state 提供了一条比 vanilla RNN 更容易保持梯度的路径。

当然,这不能声称 LSTM “彻底解决”了梯度消失问题。更准确的说法是,它通过引入线性相加的 memory path 和显式的保留机制,显著改善了长期信息和梯度的传播条件。

GRU 合并了哪些状态和控制
#

到了 GRU,网络不再维护独立的 cell state:

$$ c_t $$

它退回到了像 vanilla RNN 那样,只有:

$$ h_t $$

因此,在 GRU 中不存在 \( c_t \to h_t \) 的这一层转换,自然也就没有了 LSTM 中用来控制这层暴露的 output gate。相应地,保留旧状态和写入 candidate 的方式也变了。

Update gate:把 forget 和 input 做了耦合
#

GRU 引入了 update gate \( z_t \),它的状态更新方程长这样:

$$ h_t = (1 - z_t) \odot h_{t-1} + z_t \odot \tilde h_t $$

这里是在 old state 和 candidate 之间做逐元素 interpolation(插值)。

在 LSTM 中,\( f_t \) 和 \( i_t \) 是独立的,你可以既保留旧信息又写入新信息。而在 GRU 中,\( 1-z_t \) 和 \( z_t \) 天然满足相加等于 1:

$$ (1 - z_t) + z_t = 1 $$

这意味着,GRU 用更强的结构约束换来了更简单的 state update。当更新门倾向于保留历史时,新信息的写入就会被相应抑制。需要提醒的是,不同资料对 \( z_t \) 的定义可能反过来(例如用 \( z_t \) 乘历史,\( 1-z_t \) 乘候选),阅读时需要核对它乘的是 old state 还是 candidate。

GRU 的状态更新路径

Reset gate 如何影响历史状态的使用
#

先回到 Vanilla RNN 生成新状态的过程:

$$ \tilde h_t = \tanh(W_x x_t + W_h h_{t-1} + b) $$

每一步生成新状态时,历史状态 \( h_{t-1} \) 都是无条件、完整地参与计算的。

而 GRU 引入了 reset gate \( r_t \):

$$ r_t = \sigma(W_r [x_t, h_{t-1}] + b_r) $$

然后在生成 candidate 时:

$$ \tilde h_t = \tanh(W_x x_t + W_h (r_t \odot h_{t-1})) $$

Reset gate 控制的是:生成 candidate 时,要参考多少历史。

GRU 结构详解

这里有个容易卡住的问题:\( h_{t-1} \) 参与多少,难道不能直接由 \( W_h \) 的权重来决定吗?

这里需要分清权重和 gate 的输出。\( W_h \) 是跨 timestep 共享的静态模型参数:

$$ W_h^{(1)} = W_h^{(2)} = \cdots = W_h^{(T)} $$

它确实能学习到“历史通常应该如何被处理”,但它无法针对当前 timestep 的具体输入,动态决定“这一次我要不要依赖历史”。

而 \( r_t = \sigma(W_r [x_t, h_{t-1}]) \) 中,虽然 \( W_r \) 也是共享参数,但输入 \( x_t \) 和状态 \( h_{t-1} \) 是每一步都在变的。因此:

$$ r_1, r_2, \ldots, r_t $$

是可以动态变化的。共享的是生成 gate 的参数,gate 的输出仍然依赖每一步的输入和状态。

并且,这些 gate 的输出并不是一个 scalar,而是 hidden_size 维度的向量。这意味着模型实际上可以在**每一维(per-dimension)**上,动态地控制信息流的开与关。

清楚区分 GRU 的两个 gate
#

GRU 的两个门可以这样区分:

  • reset gate (\( r_t \)):新东西怎么生成,要参考多少历史?
  • update gate (\( z_t \)):新东西已经生成了,最终采用多少?

Gate 的线性变换可以一起计算
#

LSTM 要计算 4 个不同的项(3 个门加上 1 个 candidate),数学上可以写成四个独立的线性变换:

14 × Linear(input_size + hidden_size, hidden_size)

但在工程实现(如 PyTorch 或 CUDA 算子)中,为了最大化 GPU 上 GEMM(通用矩阵乘法)的效率,通常会把这四组完全独立的参数拼成一个大矩阵一次性计算,算完之后再通过 chunk 沿 feature 维度拆开。

下面用 PyTorch 写一个示意实现:

合并 gate 线性变换的 LSTMCell 示意实现
 1import torch
 2import torch.nn as nn
 3
 4class LSTMCellDemo(nn.Module):
 5    def __init__(self, input_size, hidden_size):
 6        super().__init__()
 7        self.hidden_size = hidden_size
 8
 9        # 工程优化:将 4 个 gate 的权重合并成一个大矩阵,提升 GPU GEMM 效率
10        self.weight_ih = nn.Linear(input_size, 4 * hidden_size)
11        self.weight_hh = nn.Linear(hidden_size, 4 * hidden_size)
12
13    def forward(self, x, state):
14        h_prev, c_prev = state
15
16        # 1. 大矩阵乘法:一次性计算所有 gate 的 pre-activation
17        gates = self.weight_ih(x) + self.weight_hh(h_prev)
18
19        # 2. 沿 feature 维度拆分为 4 份
20        ingate, forgetgate, cellgate, outgate = gates.chunk(4, dim=1)
21
22        # 3. 分别应用激活函数 (显式的 Control 与 Content 生成)
23        i_t = torch.sigmoid(ingate)       # input gate: 写多少
24        f_t = torch.sigmoid(forgetgate)   # forget gate: 留多少
25        c_tilde = torch.tanh(cellgate)    # candidate: 写什么
26        o_t = torch.sigmoid(outgate)      # output gate: 暴露多少
27
28        # 4. 核心状态更新路径
29        c_t = (f_t * c_prev) + (i_t * c_tilde)
30        h_t = o_t * torch.tanh(c_t)
31
32        return h_t, c_t

四个线性变换能合并成一次矩阵乘,是因为它们在计算前彼此毫无依赖。这也从工程侧印证了前面的 content / control 解耦。

小结
#

把三个更新式放在一起,差别就落在了具体路径上:LSTM 给 cell state 留出独立的保留路径,分别控制保留、写入和输出;GRU 只维护 hidden state,用 update gate 耦合保留与写入,用 reset gate 控制生成 candidate 时使用多少历史。

这些 gate 控制的是状态如何更新。Encoder–decoder 如何读取已有状态,是下一篇讨论的问题。

Related