LSTM & GRU

LSTM & GRU

LSTM模型

LSTM介绍

LSTM(Long Short-Term Memory,长短时记忆网络)是传统RNN的变体,与经典RNN相比能够有效捕捉长序列之间的语义关联,缓解梯度消失或爆炸现象。LSTM的结构比传统RNN更复杂,其核心结构可以分为四个部分:

  1. 遗忘门:决定从细胞状态中丢弃哪些信息
  2. 输入门:决定哪些新信息被存储在细胞状态中
  3. 细胞状态:长期记忆的载体,贯穿整个序列
  4. 输出门:决定基于细胞状态的输出内容

LSTM的内部结构

1. LSTM整体结构

LSTM在每个时间步的输入包括:当前输入\(x_t\)、上一个时间步的隐藏状态\(h_{t-1}\)和细胞状态\(C_{t-1}\)。输出包括:当前隐藏状态\(h_t\)和当前细胞状态\(C_t\)

31

2. 遗忘门

遗忘门决定从细胞状态中丢弃哪些信息。它通过一个sigmoid层实现,输出一个介于0和1之间的值,表示每个信息应该保留的程度。

结构图:

32

结构分析:

  • 将当前时间步输入\(x_t\)与上一个时间步隐藏状态\(h_{t-1}\)拼接,得到\([h_{t-1}, x_t]\)
  • 通过一个全连接层(权重矩阵\(W_f\)和偏置\(b_f\))进行线性变换
  • 通过sigmoid函数\(\sigma\)进行激活,得到遗忘门门值\(f_t\)
  • \(f_t\)将作用于上一时间步的细胞状态\(C_{t-1}\),决定保留多少过往信息
  • 值接近0表示"完全忘记",值接近1表示"完全保留"

3. 输入门

输入门决定哪些新信息被存储在细胞状态中。它包含两部分:一个sigmoid层决定更新哪些值,一个tanh层创建新的候选值向量。

结构图:

34

结构分析:

  • 输入门门值\(i_t\)的计算方式与遗忘门类似,决定哪些信息需要更新
  • \(\tilde{C}_t\)是候选细胞状态,包含当前时间步的新信息
  • 整个输入门决定了当前时间步有多少新信息需要存储到细胞状态中

4. 细胞状态更新

细胞状态是LSTM中的长期记忆单元,它贯穿整个序列,只进行少量线性操作,信息可以很容易地在其中保持不变地流动。

结构图:

35

结构分析:

  • \(f_t \odot C_{t-1}\):遗忘门决定保留多少旧信息
  • \(i_t \odot \tilde{C}_t\):输入门决定添加多少新信息
  • \(\odot\)表示逐元素相乘(Hadamard积)
  • 细胞状态更新是线性的,有助于缓解梯度消失问题

5. 输出门

输出门基于细胞状态决定最终的输出。它使用一个sigmoid层决定输出哪些部分,然后通过tanh处理细胞状态并与sigmoid输出相乘。

结构图:

37

结构分析:

  • 输出门门值\(o_t\)决定细胞状态的哪些部分将作为输出
  • \(\tanh(C_t)\)将细胞状态的值压缩到[-1, 1]范围
  • \(h_t\)作为当前时间步的隐藏状态输出,并传递给下一个时间步

LSTM总结

LSTM核心思想:

  1. 通过门控机制控制信息的流动
  2. 细胞状态作为"传送带",可以在序列中传递不变的信息
  3. 三个门(遗忘门、输入门、输出门)共同决定信息的保留、更新和输出

LSTM前向传播完整公式:

\[\begin{aligned} f_t &= \sigma(W_f \cdot [h_{t-1}, x_t] + b_f) \\ i_t &= \sigma(W_i \cdot [h_{t-1}, x_t] + b_i) \\ \tilde{C}_t &= \tanh(W_C \cdot [h_{t-1}, x_t] + b_C) \\ C_t &= f_t \odot C_{t-1} + i_t \odot \tilde{C}_t \\ o_t &= \sigma(W_o \cdot [h_{t-1}, x_t] + b_o) \\ h_t &= o_t \odot \tanh(C_t) \end{aligned} \]

Bi-LSTM介绍

Bi-LSTM(双向LSTM)是LSTM的一种扩展,它同时处理正向和反向两个方向的序列信息,从而更好地捕捉上下文依赖关系。

工作原理:

  1. 一个LSTM处理正向序列(从左到右)
  2. 另一个LSTM处理反向序列(从右到左)
  3. 将两个LSTM的输出在特定维度上拼接,作为最终输出

结构图:

38

结构分析:

  • 对于输入序列"我爱中国",正向LSTM按"我"→"爱"→"中"→"国"顺序处理
  • 反向LSTM按"国"→"中"→"爱"→"我"顺序处理
  • 将两个方向的输出拼接,得到每个时间步的最终表示
  • 这种结构能够同时捕捉前向和后向的语义依赖

Bi-LSTM输出计算:

\[\begin{aligned} \overrightarrow{h_t} &= \text{LSTM}(x_t, \overrightarrow{h_{t-1}}) \\ \overleftarrow{h_t} &= \text{LSTM}(x_t, \overleftarrow{h_{t+1}}) \\ h_t &= [\overrightarrow{h_t}; \overleftarrow{h_t}] \end{aligned} \]

优点: 能够捕捉更完整的上下文信息
缺点: 参数数量和计算复杂度增加一倍

PyTorch中LSTM的使用

在PyTorch中,可以通过torch.nn.LSTM类构建LSTM模型。

参数说明:

  • input_size:输入特征维度
  • hidden_size:隐藏状态维度
  • num_layers:LSTM层数
  • bidirectional:是否使用双向LSTM

输入输出维度:

  • 输入张量:(sequence_length, batch_size, input_size)
  • 隐藏状态:(num_layers * num_directions, batch_size, hidden_size)
  • 细胞状态:(num_layers * num_directions, batch_size, hidden_size)

示例代码:

# 定义LSTM:输入维度5,隐藏状态维度6,2层LSTM
rnn = nn.LSTM(input_size=5, hidden_size=6, num_layers=2)

# 输入数据:序列长度1,批量大小3,输入维度5
input = torch.randn(1, 3, 5)

# 初始隐藏状态和细胞状态:层数*方向数=2*1=2,批量大小3,隐藏状态维度6
h0 = torch.randn(2, 3, 6)
c0 = torch.randn(2, 3, 6)

# 前向传播
output, (hn, cn) = rnn(input, (h0, c0))

print(f'输出形状: {output.shape}')  # (1, 3, 6)
print(f'最后隐藏状态形状: {hn.shape}')  # (2, 3, 6)
print(f'最后细胞状态形状: {cn.shape}')  # (2, 3, 6)

LSTM优缺点

优点:

  1. 能够有效缓解梯度消失问题,适合处理长序列
  2. 门控机制可以学习长期依赖关系
  3. 在多种序列任务上表现优异

缺点:

  1. 结构复杂,参数数量多
  2. 训练时间较长,计算资源需求高
  3. 在某些任务上可能过拟合

GRU模型

GRU介绍

GRU(Gated Recurrent Unit,门控循环单元)是LSTM的一种简化变体,同样能够有效捕捉长序列之间的语义关联,缓解梯度消失问题。与LSTM相比,GRU结构更简单,计算效率更高。

GRU的核心结构包含两个门:

  1. 更新门:决定多少过去信息被传递到未来
  2. 重置门:决定多少过去信息被忽略

GRU的内部结构

1. GRU整体结构

GRU在每个时间步的输入包括:当前输入\(x_t\)和上一个时间步的隐藏状态\(h_{t-1}\)。输出为当前隐藏状态\(h_t\)

结构图:

gru

结构分析:

  • \(z_t\)(更新门):决定多少过去信息传递到未来,值越大保留越多历史信息
  • \(r_t\)(重置门):决定多少过去信息被忽略,值越小忽略越多历史信息
  • 两个门都使用sigmoid激活函数,输出值在0到1之间

2. 更新门和重置门

GRU使用两个门控机制:更新门控制历史信息的保留程度,重置门控制历史信息对当前候选状态的影响程度。

结构图:

gru2

3. 候选隐藏状态

候选隐藏状态\(\tilde{h}_t\)基于重置门处理后的历史信息和当前输入计算得到。

计算公式:

\[\tilde{h}_t = \tanh(W_h \cdot [r_t \odot h_{t-1}, x_t] + b_h) \]

结构分析:

  • \(r_t \odot h_{t-1}\):重置门控制历史信息对当前计算的影响
  • 如果\(r_t\)接近0,则忽略历史信息,候选状态主要基于当前输入
  • 如果\(r_t\)接近1,则保留历史信息,候选状态基于历史信息和当前输入

4. 隐藏状态更新

最终隐藏状态是更新门控制的候选状态和历史状态的组合。

计算公式:

\[h_t = (1 - z_t) \odot h_{t-1} + z_t \odot \tilde{h}_t \]

结构分析:

  • \((1 - z_t) \odot h_{t-1}\):保留的历史信息
  • \(z_t \odot \tilde{h}_t\):加入的新信息
  • \(z_t\)接近0时,主要保留历史信息
  • \(z_t\)接近1时,主要使用新信息

GRU总结

GRU核心思想:

  1. 使用两个门(更新门和重置门)控制信息流动
  2. 更新门平衡历史信息和当前信息
  3. 重置门控制历史信息对当前计算的影响

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} \]

GRU与LSTM对比:

特性 LSTM GRU
门数量 3个(遗忘门、输入门、输出门) 2个(更新门、重置门)
参数数量 较多 较少(约LSTM的75%)
计算复杂度 较高 较低
细胞状态 有单独细胞状态 无单独细胞状态
性能 在处理非常长的序列时可能更优 在大多数任务上与LSTM相当

Bi-GRU介绍

Bi-GRU(双向GRU)与Bi-LSTM类似,同时处理正向和反向序列,以捕捉更完整的上下文信息。

工作原理:

  1. 一个GRU处理正向序列
  2. 另一个GRU处理反向序列
  3. 将两个GRU的输出拼接作为最终输出

Bi-GRU输出计算:

\[\begin{aligned} \overrightarrow{h_t} &= \text{GRU}(x_t, \overrightarrow{h_{t-1}}) \\ \overleftarrow{h_t} &= \text{GRU}(x_t, \overleftarrow{h_{t+1}}) \\ h_t &= [\overrightarrow{h_t}; \overleftarrow{h_t}] \end{aligned} \]

PyTorch中GRU的使用

在PyTorch中,可以通过torch.nn.GRU类构建GRU模型。

参数说明:

  • input_size:输入特征维度
  • hidden_size:隐藏状态维度
  • num_layers:GRU层数
  • bidirectional:是否使用双向GRU

示例代码:

# 定义GRU:输入维度5,隐藏状态维度6,2层GRU
rnn = nn.GRU(input_size=5, hidden_size=6, num_layers=2)

# 输入数据:序列长度1,批量大小3,输入维度5
input = torch.randn(1, 3, 5)

# 初始隐藏状态:层数=2,批量大小3,隐藏状态维度6
h0 = torch.randn(2, 3, 6)

# 前向传播
output, hn = rnn(input, h0)

print(f'输出形状: {output.shape}')  # (1, 3, 6)
print(f'最后隐藏状态形状: {hn.shape}')  # (2, 3, 6)

GRU优缺点

优点:

  1. 结构简单,参数数量少,计算效率高
  2. 在大多数任务上与LSTM性能相当
  3. 训练速度较快,内存占用较少

缺点:

  1. 对于非常长的序列,可能不如LSTM有效
  2. 门控机制相对简单,可能无法学习某些复杂模式
  3. 与所有RNN变体一样,无法并行计算,处理长序列时效率较低

总结

LSTM和GRU都是RNN的重要变体,通过门控机制有效缓解了传统RNN的梯度消失问题,能够学习长期依赖关系。GRU作为LSTM的简化版本,在大多数任务中表现相当且计算效率更高,是现代深度学习中的常用选择。在实际应用中,应根据具体任务需求、数据特性和计算资源进行选择。

posted @ 2026-01-21 16:33  xggx  阅读(104)  评论(0)    收藏  举报