Loading [MathJax]/extensions/TeX/AMSmath.js
前往小程序,Get更优阅读体验!
立即前往
首页
学习
活动
专区
圈层
工具
发布
首页
学习
活动
专区
圈层
工具
MCP广场
社区首页 >专栏 >【现代深度学习技术】现代循环神经网络01:门控循环单元(GRU)

【现代深度学习技术】现代循环神经网络01:门控循环单元(GRU)

作者头像
Francek Chen
发布于 2025-05-02 13:22:08
发布于 2025-05-02 13:22:08
15800
代码可运行
举报
运行总次数:0
代码可运行
深度学习 (DL, Deep Learning) 特指基于深层神经网络模型和方法的机器学习。它是在统计

  在通过时间反向传播中,我们讨论了如何在循环神经网络中计算梯度,以及矩阵连续乘积可以导致梯度消失或梯度爆炸的问题。下面我们简单思考一下这种梯度异常在实践中的意义:

  • 我们可能会遇到这样的情况:早期观测值对预测所有未来观测值具有非常重要的意义。考虑一个极端情况,其中第一个观测值包含一个校验和,目标是在序列的末尾辨别校验和是否正确。在这种情况下,第一个词元的影响至关重要。我们希望有某些机制能够在一个记忆元里存储重要的早期信息。如果没有这样的机制,我们将不得不给这个观测值指定一个非常大的梯度,因为它会影响所有后续的观测值。
  • 我们可能会遇到这样的情况:一些词元没有相关的观测值。例如,在对网页内容进行情感分析时,可能有一些辅助HTML代码与网页传达的情绪无关。我们希望有一些机制来跳过隐状态表示中的此类词元。
  • 我们可能会遇到这样的情况:序列的各个部分之间存在逻辑中断。例如,书的章节之间可能会有过渡存在,或者证券的熊市和牛市之间可能会有过渡存在。在这种情况下,最好有一种方法来重置我们的内部状态表示。

  在学术界已经提出了许多方法来解决这类问题。其中最早的方法是“长短期记忆”(long-short-term memory,LSTM),我们将在下一节中讨论。门控循环单元(gated recurrent unit,GRU)是一个稍微简化的变体,通常能够提供同等的效果,并且计算的速度明显更快。由于门控循环单元更简单,我们从它开始解读。

一、门控隐状态

  门控循环单元与普通的循环神经网络之间的关键区别在于:前者支持隐状态的门控。这意味着模型有专门的机制来确定应该何时更新隐状态,以及应该何时重置隐状态。这些机制是可学习的,并且能够解决了上面列出的问题。例如,如果第一个词元非常重要,模型将学会在第一次观测之后不更新隐状态。同样,模型也可以学会跳过不相关的临时观测。最后,模型还将学会在需要的时候重置隐状态。下面我们将详细讨论各类门控。

(一)重置门和更新门

  我们首先介绍重置门(reset gate)和更新门(update gate)。我们把它们设计成

区间中的向量,这样我们就可以进行凸组合。重置门允许我们控制“可能还想记住”的过去状态的数量;更新门将允许我们控制新状态中有多少个是旧状态的副本。

  我们从构造这些门控开始。图1描述了门控循环单元中的重置门和更新门的输入,输入是由当前时间步的输入和前一时间步的隐状态给出。两个门的输出是由使用sigmoid激活函数的两个全连接层给出。

图1 在门控循环单元模型中计算重置门和更新门

  我们来看一下门控循环单元的数学表达。对于给定的时间步

,假设输入是一个小批量

(样本个数

,输入个数

),上一个时间步的隐状态是

(隐藏单元个数

)。那么,重置门

和更新门

的计算如下所示:

其中,

是权重参数,

是偏置参数。请注意,在求和过程中会触发广播机制(请参阅广播机制)。我们使用sigmoid函数(如多层感知机概述中介绍的)将输入值转换到区间

(二)候选隐状态

  接下来,让我们将重置门

中的常规隐状态更新机制集成,得到在时间步

候选隐状态(candidate hidden state)

其中,

是权重参数,

是偏置项,符号

是Hadamard积(按元素乘积)运算符。在这里,我们使用tanh非线性激活函数来确保候选隐状态中的值保持在区间

中。

  与公式

相比,式(2)中的

的元素相乘可以减少以往状态的影响。每当重置门

中的项接近

时,我们恢复一个普通的循环神经网络。对于重置门

中所有接近

的项,候选隐状态是以

作为输入的多层感知机的结果。因此,任何预先存在的隐状态都会被重置为默认值。

  图2说明了应用重置门之后的计算流程。

图2 在门控循环单元模型中计算候选隐状态

(三)隐状态

  上述的计算结果只是候选隐状态,我们仍然需要结合更新门

的效果。这一步确定新的隐状态

在多大程度上来自旧的状态

和新的候选状态

。更新门

仅需要在

之间进行按元素的凸组合就可以实现这个目标。这就得出了门控循环单元的最终更新公式:

  每当更新门

接近

时,模型就倾向只保留旧状态。此时,来自

的信息基本上被忽略,从而有效地跳过了依赖链条中的时间步

。相反,当

接近

时,新的隐状态

就会接近候选隐状态

。这些设计可以帮助我们处理循环神经网络中的梯度消失问题,并更好地捕获时间步距离很长的序列的依赖关系。例如,如果整个子序列的所有时间步的更新门都接近于

,则无论序列的长度如何,在序列起始时间步的旧隐状态都将很容易保留并传递到序列结束。

  图3说明了更新门起作用后的计算流。

图3 计算门控循环单元模型中的隐状态

总之,门控循环单元具有以下两个显著特征:

  • 重置门有助于捕获序列中的短期依赖关系;
  • 更新门有助于捕获序列中的长期依赖关系。

二、从零开始实现

  为了更好地理解门控循环单元模型,我们从零开始实现它。首先,我们读取循环神经网络的从零开始实现中使用的时间机器数据集:

代码语言:javascript
代码运行次数:0
运行
AI代码解释
复制
import torch
from torch import nn
from d2l import torch as d2l

batch_size, num_steps = 32, 35
train_iter, vocab = d2l.load_data_time_machine(batch_size, num_steps)
(一)初始化模型参数

  下一步是初始化模型参数。我们从标准差为

的高斯分布中提取权重,并将偏置项设为

,超参数num_hiddens定义隐藏单元的数量,实例化与更新门、重置门、候选隐状态和输出层相关的所有权重和偏置。

代码语言:javascript
代码运行次数:0
运行
AI代码解释
复制
def get_params(vocab_size, num_hiddens, device):
    num_inputs = num_outputs = vocab_size

    def normal(shape):
        return torch.randn(size=shape, device=device)*0.01

    def three():
        return (normal((num_inputs, num_hiddens)),
                normal((num_hiddens, num_hiddens)),
                torch.zeros(num_hiddens, device=device))

    W_xz, W_hz, b_z = three()  # 更新门参数
    W_xr, W_hr, b_r = three()  # 重置门参数
    W_xh, W_hh, b_h = three()  # 候选隐状态参数
    # 输出层参数
    W_hq = normal((num_hiddens, num_outputs))
    b_q = torch.zeros(num_outputs, device=device)
    # 附加梯度
    params = [W_xz, W_hz, b_z, W_xr, W_hr, b_r, W_xh, W_hh, b_h, W_hq, b_q]
    for param in params:
        param.requires_grad_(True)
    return params
(二)定义模型

  现在我们将定义隐状态的初始化函数init_gru_state。与循环神经网络的从零开始实现中定义的init_rnn_state函数一样,此函数返回一个形状为(批量大小,隐藏单元个数)的张量,张量的值全部为零。

代码语言:javascript
代码运行次数:0
运行
AI代码解释
复制
def init_gru_state(batch_size, num_hiddens, device):
    return (torch.zeros((batch_size, num_hiddens), device=device), )

  现在我们准备定义门控循环单元模型,模型的架构与基本的循环神经网络单元是相同的,只是权重更新公式更为复杂。

代码语言:javascript
代码运行次数:0
运行
AI代码解释
复制
def gru(inputs, state, params):
    W_xz, W_hz, b_z, W_xr, W_hr, b_r, W_xh, W_hh, b_h, W_hq, b_q = params
    H, = state
    outputs = []
    for X in inputs:
        Z = torch.sigmoid((X @ W_xz) + (H @ W_hz) + b_z)
        R = torch.sigmoid((X @ W_xr) + (H @ W_hr) + b_r)
        H_tilda = torch.tanh((X @ W_xh) + ((R * H) @ W_hh) + b_h)
        H = Z * H + (1 - Z) * H_tilda
        Y = H @ W_hq + b_q
        outputs.append(Y)
    return torch.cat(outputs, dim=0), (H,)
(三)训练与预测

  训练和预测的工作方式与循环神经网络的从零开始实现完全相同。训练结束后,我们分别打印输出训练集的困惑度,以及前缀“time traveler”和“traveler”的预测序列上的困惑度。

代码语言:javascript
代码运行次数:0
运行
AI代码解释
复制
vocab_size, num_hiddens, device = len(vocab), 256, d2l.try_gpu()
num_epochs, lr = 500, 1
model = d2l.RNNModelScratch(len(vocab), num_hiddens, device, get_params, init_gru_state, gru)
d2l.train_ch8(model, train_iter, vocab, lr, num_epochs, device)

  高级API包含了前文介绍的所有配置细节,所以我们可以直接实例化门控循环单元模型。这段代码的运行速度要快得多,因为它使用的是编译好的运算符而不是Python来处理之前阐述的许多细节。

代码语言:javascript
代码运行次数:0
运行
AI代码解释
复制
num_inputs = vocab_size
gru_layer = nn.GRU(num_inputs, num_hiddens)
model = d2l.RNNModel(gru_layer, len(vocab))
model = model.to(device)
d2l.train_ch8(model, train_iter, vocab, lr, num_epochs, device)

小结

  • 门控循环神经网络可以更好地捕获时间步距离很长的序列上的依赖关系。
  • 重置门有助于捕获序列中的短期依赖关系。
  • 更新门有助于捕获序列中的长期依赖关系。
  • 重置门打开时,门控循环单元包含基本循环神经网络;更新门打开时,门控循环单元可以跳过子序列。
本文参与 腾讯云自媒体同步曝光计划,分享自作者个人站点/博客。
原始发表:2025-05-01,如有侵权请联系 cloudcommunity@tencent.com 删除

本文分享自 作者个人站点/博客 前往查看

如有侵权,请联系 cloudcommunity@tencent.com 删除。

本文参与 腾讯云自媒体同步曝光计划  ,欢迎热爱写作的你一起参与!

评论
登录后参与评论
暂无评论
推荐阅读
编辑精选文章
换一批
【现代深度学习技术】现代循环神经网络02:长短期记忆网络(LSTM)
深度学习 (DL, Deep Learning) 特指基于深层神经网络模型和方法的机器学习。它是在统计机器学习、人工神经网络等算法模型基础上,结合当代大数据和大算力的发展而发展出来的。深度学习最重要的技术特征是具有自动提取特征的能力。神经网络算法、算力和数据是开展深度学习的三要素。深度学习在计算机视觉、自然语言处理、多模态数据分析、科学探索等领域都取得了很多成果。本专栏介绍基于PyTorch的深度学习算法实现。 【GitCode】专栏资源保存在我的GitCode仓库:https://gitcode.com/Morse_Chen/PyTorch_deep_learning。
Francek Chen
2025/05/02
2810
【现代深度学习技术】现代循环神经网络02:长短期记忆网络(LSTM)
【深度学习实验】循环神经网络(五):基于GRU的语言模型训练(包括自定义门控循环单元GRU)
get_params 函数用于初始化模型的参数。它接受三个参数:vocab_size 表示词汇表的大小,num_hiddens 表示隐藏单元的数量,device 表示模型所在的设备(如 CPU 或 GPU)。
Qomolangma
2024/07/30
4110
【深度学习实验】循环神经网络(五):基于GRU的语言模型训练(包括自定义门控循环单元GRU)
循环神经网络——中篇【深度学习】【PyTorch】【d2l】
来杯Sherry
2023/09/19
3960
循环神经网络——中篇【深度学习】【PyTorch】【d2l】
【现代深度学习技术】现代循环神经网络03:深度循环神经网络
深度学习 (DL, Deep Learning) 特指基于深层神经网络模型和方法的机器学习。它是在统计机器学习、人工神经网络等算法模型基础上,结合当代大数据和大算力的发展而发展出来的。深度学习最重要的技术特征是具有自动提取特征的能力。神经网络算法、算力和数据是开展深度学习的三要素。深度学习在计算机视觉、自然语言处理、多模态数据分析、科学探索等领域都取得了很多成果。本专栏介绍基于PyTorch的深度学习算法实现。
Francek Chen
2025/05/03
1320
【现代深度学习技术】现代循环神经网络03:深度循环神经网络
【现代深度学习技术】现代循环神经网络04:双向循环神经网络
深度学习 (DL, Deep Learning) 特指基于深层神经网络模型和方法的机器学习。它是在统计机器学习、人工神经网络等算法模型基础上,结合当代大数据和大算力的发展而发展出来的。深度学习最重要的技术特征是具有自动提取特征的能力。神经网络算法、算力和数据是开展深度学习的三要素。深度学习在计算机视觉、自然语言处理、多模态数据分析、科学探索等领域都取得了很多成果。本专栏介绍基于PyTorch的深度学习算法实现。 【GitCode】专栏资源保存在我的GitCode仓库:https://gitcode.com/Morse_Chen/PyTorch_deep_learning。
Francek Chen
2025/05/04
1100
【现代深度学习技术】现代循环神经网络04:双向循环神经网络
【现代深度学习技术】循环神经网络05:循环神经网络的从零开始实现
  本节将根据循环神经网络中的描述, 从头开始基于循环神经网络实现字符级语言模型。 这样的模型将在H.G.Wells的时光机器数据集上训练。 和前面语言模型和数据集中介绍过的一样, 我们先读取数据集。
Francek Chen
2025/04/22
1110
【现代深度学习技术】循环神经网络05:循环神经网络的从零开始实现
【深度学习实验】循环神经网络(三):门控制——自定义循环神经网络LSTM(长短期记忆网络)模型
LSTM(长短期记忆网络)是一种循环神经网络(RNN)的变体,用于处理序列数据。它具有记忆单元和门控机制,可以有效地捕捉长期依赖关系。
Qomolangma
2024/07/30
1.5K0
【深度学习实验】循环神经网络(三):门控制——自定义循环神经网络LSTM(长短期记忆网络)模型
【现代深度学习技术】现代循环神经网络07:序列到序列学习(seq2seq)
  正如我们在机器翻译与数据集中看到的,机器翻译中的输入序列和输出序列都是长度可变的。为了解决这类问题,我们在编码器-解码器架构中设计了一个通用的”编码器-解码器“架构。本节,我们将使用两个循环神经网络的编码器和解码器,并将其应用于序列到序列(sequence to sequence,seq2seq)类的学习任务。
Francek Chen
2025/05/06
1970
【现代深度学习技术】现代循环神经网络07:序列到序列学习(seq2seq)
动手学深度学习(十二) NLP循环神经网络进阶
RNN存在的问题:梯度较容易出现衰减或爆炸(BPTT) ⻔控循环神经⽹络:捕捉时间序列中时间步距离较⼤的依赖关系 RNN:
致Great
2020/02/25
4730
动手学深度学习(十二)  NLP循环神经网络进阶
【现代深度学习技术】循环神经网络04:循环神经网络
是隐状态(hidden state),也称为隐藏变量(hidden variable),它存储了到时间步
Francek Chen
2025/04/20
1840
【现代深度学习技术】循环神经网络04:循环神经网络
循环神经网络——下篇【深度学习】【PyTorch】【d2l】
设计多个隐藏层,目的是为了获取更多的非线性性。深度循环神经网络需要大量的调参(如学习率和修剪) 来确保合适的收敛,模型的初始化也需要谨慎。
来杯Sherry
2023/09/19
4750
循环神经网络——下篇【深度学习】【PyTorch】【d2l】
用pytorch写个RNN 循环神经网络
持续创作,加速成长!这是我参与「掘金日新计划 · 10 月更文挑战」的第4天,点击查看活动详情
程序猿川子
2022/10/21
1K0
三步理解--门控循环单元(GRU),TensorFlow实现。
版权声明:本文为博主原创文章,遵循 CC 4.0 by-sa 版权协议,转载请附上原文出处链接和本声明。
mantch
2019/08/29
1.3K0
三步理解--门控循环单元(GRU),TensorFlow实现。
【AI前沿】深度学习基础:循环神经网络(RNN)
循环神经网络(RNN)与传统的前馈神经网络(如多层感知器和卷积神经网络)不同,RNN具有内存能力,能够在处理当前输入时保留之前的信息。这使得RNN特别适合处理序列数据,如文本、语音和时间序列等。
屿小夏
2024/07/13
3830
【AI前沿】深度学习基础:循环神经网络(RNN)
【深度学习实验】循环神经网络(四):基于 LSTM 的语言模型训练
【深度学习实验】循环神经网络(一):循环神经网络(RNN)模型的实现与梯度裁剪_QomolangmaH的博客-CSDN博客
Qomolangma
2024/07/30
4220
【深度学习实验】循环神经网络(四):基于 LSTM 的语言模型训练
【现代深度学习技术】循环神经网络06:循环神经网络的简洁实现
深度学习 (DL, Deep Learning) 特指基于深层神经网络模型和方法的机器学习。它是在统计机器学习、人工神经网络等算法模型基础上,结合当代大数据和大算力的发展而发展出来的。深度学习最重要的技术特征是具有自动提取特征的能力。神经网络算法、算力和数据是开展深度学习的三要素。深度学习在计算机视觉、自然语言处理、多模态数据分析、科学探索等领域都取得了很多成果。本专栏介绍基于PyTorch的深度学习算法实现。 【GitCode】专栏资源保存在我的GitCode仓库:https://gitcode.com/Morse_Chen/PyTorch_deep_learning。
Francek Chen
2025/04/26
1150
【现代深度学习技术】循环神经网络06:循环神经网络的简洁实现
深度学习基础入门篇-序列模型[11]:循环神经网络 RNN、长短时记忆网络LSTM、门控循环单元GRU原理和应用详解
生活中,我们经常会遇到或者使用一些时序信号,比如自然语言语音,自然语言文本。以自然语言文本为例,完整的一句话中各个字符之间是有时序关系的,各个字符顺序的调换有可能变成语义完全不同的两句话,就像下面这个句子:
汀丶人工智能
2023/05/24
1.3K0
深度学习基础入门篇-序列模型[11]:循环神经网络 RNN、长短时记忆网络LSTM、门控循环单元GRU原理和应用详解
动手学深度学习(十一) NLP循环神经网络
本节介绍循环神经网络,下图展示了如何基于循环神经网络实现语言模型。我们的目的是基于当前的输入与过去的输入序列,预测序列的下一个字符。循环神经网络引入一个隐藏变量
致Great
2020/02/25
7820
动手学深度学习(十一)  NLP循环神经网络
循环神经网络入门基础
例如 “Cats average 15 hours of sleep a day”
timerring
2023/07/05
2990
循环神经网络入门基础
从零开始学Pytorch(十一)之ModernRNN
• 重置⻔有助于捕捉时间序列⾥短期的依赖关系; • 更新⻔有助于捕捉时间序列⾥⻓期的依赖关系。
墨明棋妙27
2022/09/23
4390
推荐阅读
【现代深度学习技术】现代循环神经网络02:长短期记忆网络(LSTM)
2810
【深度学习实验】循环神经网络(五):基于GRU的语言模型训练(包括自定义门控循环单元GRU)
4110
循环神经网络——中篇【深度学习】【PyTorch】【d2l】
3960
【现代深度学习技术】现代循环神经网络03:深度循环神经网络
1320
【现代深度学习技术】现代循环神经网络04:双向循环神经网络
1100
【现代深度学习技术】循环神经网络05:循环神经网络的从零开始实现
1110
【深度学习实验】循环神经网络(三):门控制——自定义循环神经网络LSTM(长短期记忆网络)模型
1.5K0
【现代深度学习技术】现代循环神经网络07:序列到序列学习(seq2seq)
1970
动手学深度学习(十二) NLP循环神经网络进阶
4730
【现代深度学习技术】循环神经网络04:循环神经网络
1840
循环神经网络——下篇【深度学习】【PyTorch】【d2l】
4750
用pytorch写个RNN 循环神经网络
1K0
三步理解--门控循环单元(GRU),TensorFlow实现。
1.3K0
【AI前沿】深度学习基础:循环神经网络(RNN)
3830
【深度学习实验】循环神经网络(四):基于 LSTM 的语言模型训练
4220
【现代深度学习技术】循环神经网络06:循环神经网络的简洁实现
1150
深度学习基础入门篇-序列模型[11]:循环神经网络 RNN、长短时记忆网络LSTM、门控循环单元GRU原理和应用详解
1.3K0
动手学深度学习(十一) NLP循环神经网络
7820
循环神经网络入门基础
2990
从零开始学Pytorch(十一)之ModernRNN
4390
相关推荐
【现代深度学习技术】现代循环神经网络02:长短期记忆网络(LSTM)
更多 >
领券
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档