循环神经网络(五):LSTM(长短期记忆网络)

LSTM (长短期记忆网络) 原理解析笔记

1. 什么是LSTM?

  • LSTM (Long Short-Term Memory Network) 是循环神经网络(RNN)的一种变体
  • 主要解决长期依赖问题(即模型能够记住长期信息)
  • 与GRU相比,LSTM参数更多,通常在处理更复杂的长期依赖任务时表现更好

2. LSTM 总体结构

核心概念:

  • 长期记忆:存储在 细胞状态(cell state),记为 C
  • 短期记忆:存储在 隐藏状态(hidden state),记为 h

LSTM 通过三个“门”结构来控制信息的流动和记忆的更新:

  1. 遗忘门(Forget Gate)
  2. 输入门(Input Gate)
  3. 输出门(Output Gate)

如下图所示:
在这里插入图片描述


3. LSTM 执行步骤详解

(1) 遗忘门(Forget Gate)

在这里插入图片描述

作用:决定从长期记忆(cell state)中丢弃哪些信息
公式

ft = σ ( Wf ⋅ [ ht − 1 , xt ] + bf )

  • ft 是一个介于0~1之间的值(通过Sigmoid函数得到)
  • 0 表示“完全遗忘”,1 表示“完全保留”

(2) 输入门(Input Gate)

在这里插入图片描述

作用:决定哪些新信息将被存储到细胞状态中

分为两部分:

  • it:输入重要性系数(缩放因子)

    it = σ ( Wi ⋅ [ ht − 1 , xt ] + bi )

  • C~t:候选更新向量(表示新的信息)

    C~t = tanh ( Wc ⋅ [ ht − 1 , xt ] + bc )

最终输入向量为:

输入向量=itC~t


(3) 更新长期记忆(Cell State Update)

在这里插入图片描述

作用:结合遗忘门和输入门,更新细胞状态

Ct= ftCt −1+ itC~t


(4) 输出门(Output Gate)

在这里插入图片描述

作用:基于当前细胞状态,决定输出什么信息(即更新隐藏状态)

  • 输出比例系数:

    ot = σ ( Wo ⋅ [ht−1,xt] + bo)

  • 更新隐藏状态:

    ht = ot ⋅ tanh(Ct)


4. LSTM 变体

(1) 窥视孔连接(Peephole Connections)

在这里插入图片描述

  • 在计算门控信号(ft,it,ot)时,除了使用 ht−1xt,还引入细胞状态 Ct−1Ct

(2) 耦合遗忘门与输入门

在这里插入图片描述

  • 使用 1−ft 代替输入门 it,使遗忘与输入形成互补关系

(3) GRU(门控循环单元)

  • LSTM的简化版本,将遗忘门和输入门合并为“更新门”
  • 只有两个门:更新门和重置门

5. 总结对比

组件作用输出范围计算公式
遗忘门 ft控制历史信息的保留比例(0, 1)σ(Wf[ht−1,xt]+bf)
输入门 it控制新信息的重要性(0, 1)σ(Wi[ht−1,xt]+bi)
输出门 ot控制输出信息的比例(0, 1)σ(Wo[ht−1,xt]+bo)
候选记忆C~t生成候选更新信息(-1, 1)tanh(Wc[ht−1,xt]+bc)

6.代码实现:

import torch
from torch import nn

class LSTMCell(nn.Module):
    def __init__(self,input_size,hidden_size):
        super().__init__()
        self.hidden_size = hidden_size
        self.linear_ft = nn.Linear(input_size + hidden_size, hidden_size)
        self.linear_it = nn.Linear(input_size + hidden_size, hidden_size)
        self.linear_ct = nn.Linear(input_size + hidden_size, hidden_size)
        self.linear_ot = nn.Linear(input_size + hidden_size, hidden_size)


        self.sigmoid = nn.Sigmoid()
        self.tanh = nn.Tanh()

    def forward(self, x,c=None,h=None):
        _x = torch.concat([h,x],dim=1)
        #遗忘门
        ft = self.sigmoid(self.linear_ft(_x))
        #输入门
        it = self.sigmoid(self.linear_it(_x))
        ct = self.tanh(self.linear_ct(_x))
        #更新长期记忆
        c = ft*c + it*ct
        #输出门
        ot = self.sigmoid(self.linear_ot(_x))
        h = ot*torch.tanh(c)
        return c,h

class LSTM(nn.Module):
    def __init__(self,input_size,hidden_size):
        super().__init__()
        self.cell = LSTMCell(input_size,hidden_size)
        self.hidden_size = hidden_size
        self.fc_out = nn.Linear(hidden_size,hidden_size)


    def forward(self, x,c=None,h=None):
        N,L,input_size = x.shape

        if h is None:
            h = torch.zeros(N,self.hidden_size)
        if c is None:
            c = torch.zeros(N,self.hidden_size)

        outputs = []

        for i in range(L):
            c, h = self.cell(x[:,i],c,h)
            out = self.fc_out(h)
            outputs.append(out)

        outputs = torch.stack(outputs,dim=1)
        return outputs,c,h

if __name__ == '__main__':
    x = torch.rand(5,6,10)
    model = LSTM(10,20)
    y,c,h = model(x)
    print(y.shape)
    print(c.shape)
    print(h.shape)
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值