LSTM 与 GRU 单元

PyTorch 深度学习进阶

Michal Oleszak

Machine Learning Engineer

短期记忆问题

  • RNN 单元通过隐藏状态维护记忆
  • 该记忆非常短期
  • 两种更强的单元可解决此问题:
    • LSTM(长短期记忆)单元
    • GRU(门控循环单元)单元

循环神经元示意图。时间步 2 接收输入 h2 与 x2,并输出 y2 与 h3。

PyTorch 深度学习进阶

RNN 单元

RNN 单元示意图。

  • 两个输入:
    • 当前输入 x
    • 上一隐藏状态 h
  • 两个输出:
    • 当前输出 y
    • 下一隐藏状态 h
PyTorch 深度学习进阶

LSTM 单元

LSTM 单元示意图。

  • 输出 hy 相同
  • 三个输入与输出(两个隐藏状态):

    • h:短期状态
    • c:长期状态
  • 三个"门":

    • 记忆门(Forget gate):从长期记忆中移除
    • 输入门(Input gate):写入长期记忆
    • 输出门(Output gate):当前时间步的输出
PyTorch 深度学习进阶

PyTorch 中的 LSTM

class Net(nn.Module):
    def __init__(self, input_size):
        super().__init__()

self.lstm = nn.LSTM( input_size=1, hidden_size=32, num_layers=2, batch_first=True, ) self.fc = nn.Linear(32, 1)
def forward(self, x): h0 = torch.zeros(2, x.size(0), 32) c0 = torch.zeros(2, x.size(0), 32)
out, _ = self.lstm(x, (h0, c0))
out = self.fc(out[:, -1, :]) return out
  • __init__():
    • nn.RNN 替换为 nn.LSTM
  • forward():
    • 新增另一个隐藏状态 c
    • ch 用零初始化
    • 将两个隐藏状态一并传入 lstm
PyTorch 深度学习进阶

GRU 单元

GRU 单元示意图。

  • LSTM 的简化版本
  • 只有一个隐藏状态
  • 无输出门
PyTorch 深度学习进阶

PyTorch 中的 GRU

class Net(nn.Module):
    def __init__(self, input_size):
        super().__init__()

self.gru = nn.GRU( input_size=1, hidden_size=32, num_layers=2, batch_first=True, ) self.fc = nn.Linear(32, 1)
def forward(self, x): h0 = torch.zeros(2, x.size(0), 32) out, _ = self.gru(x, h0) out = self.fc(out[:, -1, :]) return out
  • __init__():
    • nn.RNN 替换为 nn.GRU
  • forward():
    • 使用 gru
PyTorch 深度学习进阶

应使用 RNN、LSTM 还是 GRU?

  • RNN 现已较少使用
  • GRU 比 LSTM 更简单,计算更少
  • 相对性能依用例而异
  • 两者都试试并比较

LSTM 与 GRU 单元示意图。

PyTorch 深度学习进阶

Passons à la pratique !

PyTorch 深度学习进阶

Preparing Video For Download...