深度学习入门:从结构图到函数表达式解析简单循环神经网络、GRU和LSTM

1次阅读
没有评论

共计 2763 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

背景介绍

循环神经网络(RNN)是处理序列数据的利器,比如自然语言、时间序列等。传统的神经网络难以捕捉序列中的时间依赖关系,而 RNN 通过引入循环连接,使得信息可以在时间步之间传递。然而,基础的 RNN 存在梯度消失和梯度爆炸的问题,难以学习长距离依赖。为了解决这些问题,门控循环单元(GRU)和长短期记忆网络(LSTM)应运而生。

深度学习入门:从结构图到函数表达式解析简单循环神经网络、GRU 和 LSTM

结构对比

简单循环神经网络(RNN)

RNN 的单元结构非常简单,主要由一个隐藏状态 $h_t$ 组成,它在每个时间步接收当前输入 $x_t$ 和上一个时间步的隐藏状态 $h_{t-1}$,并通过一个非线性激活函数(如 tanh)生成新的隐藏状态 $h_t$。

  • 结构图示意

    x_t -> [RNN Cell] -> h_t
            ^
            |
           h_{t-1}

  • 组件作用

  • $h_t$:当前时间步的隐藏状态,携带了序列的历史信息。
  • $x_t$:当前时间步的输入。

门控循环单元(GRU)

GRU 引入了两个门控机制:更新门($z_t$)和重置门($r_t$),用于控制信息的流动。

  • 结构图示意

    x_t -> [GRU Cell] -> h_t
            ^
            |
           h_{t-1}

  • 组件作用

  • 更新门 $z_t$:决定有多少历史信息需要保留。
  • 重置门 $r_t$:决定有多少历史信息需要忽略。
  • 候选隐藏状态 $\tilde{h}_t$:基于当前输入和重置后的历史信息生成。

长短期记忆网络(LSTM)

LSTM 比 GRU 更复杂,引入了三个门控机制:输入门($i_t$)、遗忘门($f_t$)和输出门($o_t$),以及一个细胞状态($C_t$)。

  • 结构图示意

    x_t -> [LSTM Cell] -> h_t
            ^
            |
           h_{t-1}, C_{t-1}

  • 组件作用

  • 输入门 $i_t$:决定有多少新信息需要存入细胞状态。
  • 遗忘门 $f_t$:决定有多少历史信息需要遗忘。
  • 输出门 $o_t$:决定有多少细胞状态信息需要输出到隐藏状态。
  • 细胞状态 $C_t$:长期记忆的载体。

数学表达

RNN 的数学表达式

$$
h_t = \tanh(W_{hx} x_t + W_{hh} h_{t-1} + b_h)
$$

  • $W_{hx}$:输入到隐藏层的权重矩阵。
  • $W_{hh}$:隐藏层到隐藏层的权重矩阵。
  • $b_h$:隐藏层的偏置项。

GRU 的数学表达式

$$
z_t = \sigma(W_{zx} x_t + W_{zh} h_{t-1} + b_z)
$$
$$
r_t = \sigma(W_{rx} x_t + W_{rh} h_{t-1} + b_r)
$$
$$
\tilde{h}t = \tanh(W) + b_h)
$$
$$
h_t = z_t \odot h_{t-1} + (1 – z_t) \odot \tilde{h}_t
$$} x_t + W_{hh} (r_t \odot h_{t-1

  • $\sigma$:sigmoid 函数,输出在 0 到 1 之间,用于门控。
  • $\odot$:逐元素乘法。

LSTM 的数学表达式

$$
i_t = \sigma(W_{ix} x_t + W_{ih} h_{t-1} + b_i)
$$
$$
f_t = \sigma(W_{fx} x_t + W_{fh} h_{t-1} + b_f)
$$
$$
o_t = \sigma(W_{ox} x_t + W_{oh} h_{t-1} + b_o)
$$
$$
\tilde{C}t = \tanh(W + b_C)
$$
$$
C_t = f_t \odot C_{t-1} + i_t \odot \tilde{C}_t
$$
$$
h_t = o_t \odot \tanh(C_t)
$$} x_t + W_{Ch} h_{t-1

  • $i_t, f_t, o_t$:输入门、遗忘门、输出门。
  • $\tilde{C}_t$:候选细胞状态。
  • $C_t$:当前细胞状态。

代码实现

RNN 实现(PyTorch)

import torch
import torch.nn as nn

class SimpleRNN(nn.Module):
    def __init__(self, input_size, hidden_size):
        super(SimpleRNN, self).__init__()
        self.hidden_size = hidden_size
        self.rnn = nn.RNN(input_size, hidden_size, batch_first=True)

    def forward(self, x):
        h0 = torch.zeros(1, x.size(0), self.hidden_size)  # 初始隐藏状态
        out, _ = self.rnn(x, h0)
        return out

GRU 实现(PyTorch)

class GRU(nn.Module):
    def __init__(self, input_size, hidden_size):
        super(GRU, self).__init__()
        self.hidden_size = hidden_size
        self.gru = nn.GRU(input_size, hidden_size, batch_first=True)

    def forward(self, x):
        h0 = torch.zeros(1, x.size(0), self.hidden_size)  # 初始隐藏状态
        out, _ = self.gru(x, h0)
        return out

LSTM 实现(PyTorch)

class LSTM(nn.Module):
    def __init__(self, input_size, hidden_size):
        super(LSTM, self).__init__()
        self.hidden_size = hidden_size
        self.lstm = nn.LSTM(input_size, hidden_size, batch_first=True)

    def forward(self, x):
        h0 = torch.zeros(1, x.size(0), self.hidden_size)  # 初始隐藏状态
        c0 = torch.zeros(1, x.size(0), self.hidden_size)  # 初始细胞状态
        out, _ = self.lstm(x, (h0, c0))
        return out

应用场景

  • RNN:适用于短序列任务,计算资源有限时可以考虑。
  • GRU:比 LSTM 参数更少,训练更快,适合中等长度序列。
  • LSTM:适用于长序列任务,如机器翻译、语音识别。

避坑指南

  1. 梯度消失 / 爆炸
  2. 使用梯度裁剪(torch.nn.utils.clip_grad_norm_)。
  3. 选择合适的初始化方法(如 Xavier 初始化)。

  4. 过拟合

  5. 使用 Dropout(nn.Dropout)。
  6. 增加正则化(L2 正则化)。

  7. 长序列处理

  8. 优先选择 LSTM 或 GRU。
  9. 考虑使用注意力机制(Transformer)。

思考题

  1. 为什么 LSTM 和 GRU 能缓解梯度消失问题?
  2. 如何在 RNN 中实现双向处理(Bidirectional RNN)?
  3. 除了序列数据,RNN 还能用于哪些类型的任务?

希望这篇文章能帮助你理解 RNN、GRU 和 LSTM 的核心概念和实现方法。如果有任何问题,欢迎留言讨论!

正文完
 0
评论(没有评论)