深度学习之门控循环单元(GRU)
GRU(Gated Recurrent Unit,门控循环单元)是 LSTM 的简化兄弟。它把 LSTM 的三个门压缩成两个门——更新门与重置门——同时去掉了独立的细胞状态,让隐藏状态自己兼任”记忆载体”与”输出”双重角色。这种”做减法”的设计让它在效果与 LSTM 相当的前提下,参数更少、训练更快,长期活跃在 seq2seq 机器翻译、语音识别、时间序列预测等任务里。即使 Transformer 出现后,GRU 的门控思想仍然是理解循环神经网络的重要一环。

从 LSTM 到 GRU:为什么要简化?
GRU 的出发点是”用更少的门达到相近的长程记忆能力”,它把 LSTM 中功能互补的遗忘门与输入门合并成一个更新门,再用一个重置门控制候选新状态对旧记忆的依赖程度。
普通 RNN 的问题已经讲过:每一步隐藏状态都做一次非线性变换,梯度在远距离上迅速衰减。LSTM 的回应是增加一条几乎不衰减的细胞状态线,并配三个门分别控制遗忘、写入和读出。这条设计确实有效,但也让 LSTM 拥有了六个核心公式、三套独立权重和不少工程调试细节。

2014 年,Cho 等人在机器翻译的 seq2seq 论文中提出 GRU 时,他们观察到两个现象:
- LSTM 的遗忘门和输入门通常此消彼长——旧信息丢弃越多,新信息写入往往越多,二者高度互补;
- 输出门在很多任务里并不是必需的,因为隐藏状态本身就可以作为输出。
于是他们设计出一个更紧凑的结构:把遗忘门与输入门合并为”更新门”,让 $z_t$ 同时决定”旧记忆保留多少”与”新候选写入多少”;再增加”重置门” $r_t$,控制生成候选状态时要把旧状态忽略到什么程度。
随后的实证研究支持了这种简化。Greff 等人在 2015 年、2017 年的系统实验中发现,去掉 peephole、耦合遗忘门与输入门、甚至去掉输出门的 LSTM 变体,在标准任务上与标准 LSTM 差异微小。Jozefowicz 等人则用架构搜索在一万种 RNN 变体中寻找最优结构,结果显示 GRU 始终排在前列。这些研究表明:LSTM 中相当一部分结构是冗余的,GRU 以更小的代价获得了相似的能力。
| 年份 | 事件 | 意义 |
| 1997 | Hochreiter & Schmidhuber 提出 LSTM | 用细胞状态 + 门控机制解决 RNN 长程依赖问题 |
| 2014 | Cho et al. 在 seq2seq 论文中提出 GRU | 把三个门简化为两个门,参数减少,效果仍接近 LSTM |
| 2014 | Chung et al. 系统评估 GRU | 在音乐、语音、自然语言任务上证明 GRU 与 LSTM 互有胜负 |
| 2015 | Greff et al. / Jozefowicz et al. 变体比较 | 验证门控结构存在冗余,GRU 是稳健的选择 |
| 2017 | Vaswani et al. 提出 Transformer | 自注意力逐渐取代循环结构,但门控思想仍被继承 |
两个门如何工作
更新门同时控制”旧记忆保留多少”与”新信息写入多少”,重置门则控制”在生成新内容时先忘掉多少旧信息”。两者共同决定了当前步的隐藏状态如何由旧状态与新输入演化而来。
先用一个类比建立直觉。把隐藏状态 $h$ 想象成一本实时更新的笔记。普通 RNN 每一步都把整本笔记重写一遍;LSTM 用一本主笔记(细胞状态)加三本副本来管理读写;GRU 则只派两个管理员:更新门决定”这一页保留多少旧内容、覆写多少新内容”,重置门决定”在写新内容时,先擦掉多少旧笔记的干扰”。因为 $z_t$ 和 $1-z_t$ 天然互补,GRU 不需要像 LSTM 那样分别训练一个遗忘门和一个输入门。

读这张结构图时关注四个要点:
- 顶部粗线是隐藏状态 $h$ 的”传送带”,从左边的 $h_{t-1}$ 进来,经过融合后变成 $h_t$ 传给下一步;
- 更新门和重置门共享同一个拼接输入 $[h_{t-1}, x_t]$,但各自学各自的权重矩阵 $W_z$ 与 $W_r$;
- 重置门输出 $r_t$ 与 $h_{t-1}$ 逐元素相乘,得到”被过滤后的旧记忆”,再与 $x_t$ 拼接送入候选 $\tanh$;
- 最终融合是旧状态与新候选的加权平均:$h_t = (1-z_t) \odot h_{t-1} + z_t \odot \tilde h_t$。
四个公式逐条拆解
GRU 的全部计算可以写成四个公式:
$$\begin{aligned} z_t &= \sigma(W_z \cdot [h_{t-1}, x_t] + b_z) \\ r_t &= \sigma(W_r \cdot [h_{t-1}, x_t] + b_r) \\ \tilde h_t &= \tanh(W_h \cdot [r_t \odot h_{t-1}, x_t] + b_h) \\ h_t &= (1 – z_t) \odot h_{t-1} + z_t \odot \tilde h_t \end{aligned}$$
逐条解释它们的职责:
- 更新门 $z_t$:它用 sigmoid 输出一个 0~1 之间的向量,每一维都是”更新比例”。当某维接近 0 时,$h_t$ 在该维几乎原样复制 $h_{t-1}$;当接近 1 时,$h_t$ 在该维几乎完全采用候选 $\tilde h_t$。因为保留比例是 $1-z_t$,写入比例是 $z_t$,二者天然互补,不会出现”既忘记又写入”的冲突。
- 重置门 $r_t$:它也输出 0~1 向量,但不直接参与最终融合,而是控制候选状态 $\tilde h_t$ 在生成时依赖旧状态的程度。$r_t \approx 0$ 时,候选状态几乎只看当前输入 $x_t$;$r_t \approx 1$ 时,旧状态充分参与新内容生成。这相当于在写新笔记前,先用橡皮擦把旧笔记擦掉一部分。
- 候选状态 $\tilde h_t$:通过 $\tanh$ 输出 -1~1 范围的向量,代表”如果重置门完全放开,当前步应该写入什么新内容”。注意它的输入不是 $[h_{t-1}, x_t]$,而是 $[r_t \odot h_{t-1}, x_t]$,这是 GRU 与 LSTM 的关键差异之一。
- 最终隐藏状态 $h_t$:简单的加权融合。$z_t$ 同时扮演了 LSTM 中遗忘门和输入门的双重角色,这也是 GRU 能省掉一个门的原因。

GRU vs LSTM:谁更强?
在大多数任务上,GRU 与 LSTM 效果相当,但 GRU 参数更少、训练更快;当序列特别长或需要更稳定的记忆边界时,LSTM 仍是更稳健的选择。
两者的差异首先体现在结构上。LSTM 维护两套状态:细胞状态 $C_t$ 负责长期记忆,隐藏状态 $h_t$ 负责输出;GRU 只有一套 $h_t$,长期记忆和当前输出由同一个向量承担。LSTM 用三个门分别控制遗忘、写入和读出;GRU 用更新门同时控制保留与写入,用重置门控制候选生成。
| 维度 | LSTM | GRU |
| 门数量 | 3(遗忘门、输入门、输出门) | 2(更新门、重置门) |
| 细胞状态 | 有独立的 $C_t$ | 无,$h_t$ 身兼二职 |
| 每个时间步输出 | $h_t$ 与 $C_t$ | 只有 $h_t$ |
| 参数量 | $4 \cdot H \cdot (H + I)$ | $3 \cdot H \cdot (H + I)$ |
| 相对参数量 | 100% | 约 75% |
| 训练速度 | 较慢 | 较快 |
| 长程依赖稳定性 | 理论上更稳 | 多数任务足够 |
实证层面有几项经典结论值得记住。Greff 等人 2015/2017 年在大量序列任务上系统比较了 LSTM 的多种变体,发现 peephole 连接、输出门等组件虽然增加了表达能力,但在统计上并未显著优于更简化的版本。Chung 等人 2014 年直接在音乐、语音、自然语言任务上对比 GRU 与 LSTM,结果两者互有胜负,没有一个压倒性胜出。Jozefowicz 等人 2015 年用架构搜索评估了一万种 RNN 变体,发现 GRU 是表现最好的结构之一;他们还找到了一个优于 LSTM 的门控变体,说明搜索空间远未穷尽。
用 PyTorch 从零实现
手写一个 GRU 单元的关键是”把更新门、重置门、候选状态的线性变换一次性算完再切分”,思路与手写 LSTM 完全一致。下面的实现与 PyTorch 的 nn.GRUCell 对齐,参数形状为 (hidden_size, hidden_size + input_size),把输入 $[h, x]$ 拼接后做矩阵乘法。
import torch
import torch.nn as nn
import math
class GRUCell(nn.Module):
"""手写 GRU 单元:与 nn.GRUCell 行为对齐"""
def __init__(self, input_size, hidden_size):
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
# 更新门、重置门、候选状态共享拼接输入 [h; x]
self.W_z = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_z = nn.Parameter(torch.Tensor(hidden_size))
self.W_r = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_r = nn.Parameter(torch.Tensor(hidden_size))
self.W_h = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_h = nn.Parameter(torch.Tensor(hidden_size))
self.reset_parameters()
def reset_parameters(self):
stdv = 1.0 / math.sqrt(self.hidden_size)
for w in (self.W_z, self.W_r, self.W_h):
nn.init.uniform_(w, -stdv, stdv)
for b in (self.b_z, self.b_r, self.b_h):
nn.init.uniform_(b, -stdv, stdv)
def forward(self, x, h_prev):
hx = torch.cat([h_prev, x], dim=-1) # [h_{t-1}; x_t]
z = torch.sigmoid(hx @ self.W_z.T + self.b_z) # 更新门
r = torch.sigmoid(hx @ self.W_r.T + self.b_r) # 重置门
rhx = torch.cat([r * h_prev, x], dim=-1) # [r_t ⊙ h_{t-1}; x_t]
h_tilde = torch.tanh(rhx @ self.W_h.T + self.b_h) # 候选状态
h = (1 - z) * h_prev + z * h_tilde # 融合
return h
# 与 PyTorch 官方 GRUCell 对齐性检查
gru_ours = GRUCell(input_size=4, hidden_size=8)
gru_torch = nn.GRUCell(input_size=4, hidden_size=8)
# 把权重改成一致即可逐元素比较;这里仅展示结构等价
x = torch.randn(2, 4)
h = torch.randn(2, 8)
print("ours:", gru_ours(x, h).shape)
print("torch:", gru_torch(x, h).shape)
PyTorch 的 nn.GRU 已经封装了多层、双向、dropout 等功能,日常建模直接用它即可:
import torch.nn as nn
gru = nn.GRU(
input_size=10,
hidden_size=32,
num_layers=2,
batch_first=True,
bidirectional=True,
dropout=0.3,
)
# x: [batch, seq_len, input_size]
# out: [batch, seq_len, hidden_size * 2]
# h_n: [num_layers * 2, batch, hidden_size]
out, h_n = gru(x)
下面用一个能直观检验长程记忆能力的小任务:模型需要把序列第一个 token 的值复制到输出端。这个任务对普通 RNN 来说序列稍长就会失败,但 GRU 可以靠更新门把首 token 一直保留到最后。
import torch
import torch.nn as nn
import torch.optim as optim
def make_batch(batch_size, seq_len, dim=4):
"""每个样本的第一个 token 是目标,其余是随机噪声。"""
x = torch.randn(batch_size, seq_len, dim)
y = x[:, 0, :] # 目标:复制第一个 token
return x, y
class SimpleGRU(nn.Module):
def __init__(self, input_size=4, hidden_size=16, output_size=4):
super().__init__()
self.gru = nn.GRU(input_size, hidden_size, batch_first=True)
self.fc = nn.Linear(hidden_size, output_size)
def forward(self, x):
out, _ = self.gru(x) # out: [B, T, H]
return self.fc(out[:, -1, :]) # 取最后时刻预测首 token
model = SimpleGRU()
optimizer = optim.Adam(model.parameters(), lr=1e-2)
loss_history = []
max_epochs = 500
patience = 30
best_loss = float('inf')
no_improve = 0
for epoch in range(max_epochs):
x, y = make_batch(batch_size=32, seq_len=40, dim=4)
pred = model(x)
loss = nn.MSELoss()(pred, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_history.append(loss.item())
if loss.item() < best_loss - 1e-4:
best_loss = loss.item()
no_improve = 0
else:
no_improve += 1
if (epoch + 1) % 50 == 0:
print(f"epoch {epoch+1:03d}, loss={loss.item():.6f}")
if no_improve >= patience:
print(f"early stop at epoch {epoch+1}, best_loss={best_loss:.6f}")
break
else:
print("warning: reached max_epochs without convergence")
这个例子的训练可观测性包括:每 50 轮打印一次损失、损失不再下降时触发提前停止、最大轮次未收敛时给出警告。你可以把 seq_len 从 40 改到 100 或 200,观察 GRU 仍能保持稳定收敛,而普通 RNN 会迅速失效。
参考文献 / 扩展阅读
- Cho, K., van Merriënboer, B., Gulcehre, C., Bahdanau, D., Bougares, F., Schwenk, H., & Bengio, Y. (2014). Learning Phrase Representations using RNN Encoder–Decoder for Statistical Machine Translation. arXiv:1406.1078.
- Chung, J., Gulcehre, C., Cho, K., & Bengio, Y. (2014). Empirical Evaluation of Gated Recurrent Neural Networks on Sequence Modeling. arXiv:1412.3555.
- Hochreiter, S., & Schmidhuber, J. (1997). Long Short-Term Memory. Neural Computation, 9(8), 1735–
- Greff, K., Srivastava, R. K., Koutník, J., Steunebrink, B. R., & Schmidhuber, J. (2017). LSTM: A Search Space Odyssey. IEEE Transactions on Neural Networks and Learning Systems, 28(10), 2222–
- Jozefowicz, R., Zaremba, W., & Sutskever, I. (2015). An Empirical Exploration of Recurrent Network Architectures. Proceedings of ICML 2015.
- PyTorch 官方文档:https://pytorch.org/docs/stable/generated/torch.nn.GRU.html
- GRU Recurrent Neural Networks – A Smart Way to Predict Sequences in Python | Towards Data Science





