7、长短期记忆网络(LSTM)

发布时间:2026/8/13 21:01:37

7、长短期记忆网络(LSTM)
1 LSTM基本结构假如现在有一个需求根据现有文本预测以下一个词语比如天上的云朵漂浮在通过间隔不远的位置就可以预测出来词语是天上但对于其他一些句子可能需要被预测的词语前100个词语前此时由于间隔非常大随着时间的间隔的增加会导致真实的预测结果对结果的影响变得非常小而无法非常好的进行预测RNN中的长期依赖问题long-Term Dependencies为了解决RNN中时间上的梯度消失机器学习领域发展出了长短时记忆单元LSTM通过门的开关实现时间上记忆功能并防止梯度消失。LSTM是一种特殊的RNN(循环神经网络)可以学习长期依赖信息一个LSTM的单元就是下图中的一个绿色方框中的内容1.1 LSTM的网络结构可以发现RNN只有一个传递状态hth_tht​LSTM有两个传输状态一个ctc_tct​cell state和一个hth_tht​hidden state。通常输出的ctc_tct​是上一个状态传过来的加上一些数值而hth_tht​则在不同节点下往往会有很大的区别。​ LSTM的核心单元细胞中的状态也就是上图中最上面的那根线。但是如果只有上面那一根线那么没有办法实现信息的增加或者删除所以LSTM是通过一个叫做门的结构实现门可以选择让信息通过或者不通过。这个门主要通过sigmoid和点乘pointwise multiplication实现的。sigmoid的取值范围在01之间如果接近0表示不让任何信息通过如果接近1表示所有的信息都会通过。各个门各司其职每个门通常使用Sigmoid 函数作为激活函数激活后的值处在0和1之间故方便控制 “门” 的开启和关闭输入门决定Z能走多远遗忘门决定记忆单元的值是否刷新或者重置输出门则决定最后的能否被输出。1.2 通俗理解LSTM的三个门门(Gate)是一种可选地让信息通过的方式LSTM有三个门用于保护和控制细胞的状态。LSTM内部主要有三个阶段1.忘记阶段这个阶段主要是对上一个节点传进来的输入进行选择性忘记。简单来说就是会 “忘记不重要的记住重要的”。具体来说是通过计算得到的ftf_tft​来作为忘记门控来控制上一个状态Ct−1C_{t-1}Ct−1​的哪些需要留哪些需要忘。2.选择记忆阶段这个阶段将这个阶段的输入有选择性地进行“记忆”。主要是会对输入xtx_txt​进行选择记忆。哪些重要则着重记录下来哪些不重要则少记一些。当前的输入内容由前面计算得到的表示。而选择的门控信号则是由iti_tit​来进行控制将上面两步得到的结果相加即可得到传输给下一个状态Ctft×Ct−1it×Ct^C_tf_t×C_{t-1}i_t×\hat{C_t}Ct​ft​×Ct−1​it​×Ct​^​3.输出阶段这个阶段将决定哪些将会被当成当前状态的输出。主要是通过oto_tot​来进行控制的。并且还对上一阶段得到的CtC_tCt​进行了放缩通过一个tanh激活函数进行变化与普通RNN类似输出yty_tyt​往往最终也是通过hth_tht​变化得到。2 LSTM训练过程2.1 计算过程第一步决定我们要从细胞状态中丢弃什么信息 该决定由被称为**“遗忘门”**的Sigmoid层实现。它查看ht−1h_{t-1}ht−1​(前一个输出)和xtx_txt​(当前输入)并为单元格状态Ct−1C_{t-1}Ct−1​(上一个状态)中的每个数字输出0和1之间的数字1代表完全保留而0代表彻底删除。第二步决定我们要在细胞状态中存储什么信息。 首先称为“输入门”的Sigmoid层决定更新哪些值。 接下来一个tanh层创建候选向量Ct^\hat{C_t}Ct​^​该向量将会被加到细胞的状态中。 在下一步中我们将结合这两个向量来创建更新值。第三步更新状态值CtC_tCt​。我们将上一个状态值Ct−1C_{t-1}Ct−1​乘以ftf_tft​以此表达期待忘记的部分。之后我们将得到的值加上it∗Ct^i_t∗\hat{C_t}it​∗Ct​^​这个得到的是新的候选值 按照我们决定更新每个状态值的多少来衡量。最后我们需要决定我们要输出什么。 此输出将基于我们的细胞状态但将是一个过滤版本。 首先我们运行一个Sigmoid层它决定了我们要输出的细胞状态的哪些部分 然后我们将单元格状态通过tanh将值规范化到−1和1 之间并将其乘以Sigmoid门的输出至此我们输出了我们决定的那些部分。2.2 LSTM训练算法框架LSTM的训练算法仍然是反向传播算法对于这个算法我们已经非常熟悉了。主要有下面三个步骤前向计算每个神经元的输出值对于LSTM来说即上述ft、it、Ct、ot、htf_t、i_t、C_t、o_t、h_tft​、it​、Ct​、ot​、ht​五个向量的值。计算方法已经在上一节中描述过了。反向计算每个神经元的误差项值。与循环神经网络一样LSTM误差项的反向传播也是包括两个方向一个是沿时间的反向传播即从当前ttt时刻开始计算每个时刻的误差项一个是将误差项向上一层传播。根据相应的误差项计算每个权重的梯度。3 LSTM优缺点3.1 LSTM优点CNN并不完全适用于学习时间序列因此会需要各种辅助性处理且效果也不一定好。面对对时间序列敏感的问题和任务RNN(如LSTM)通常会比较合适。RNN用于序列数据并且有了一定的记忆效应RNN可以视为一个所有层共享同样权值的深度前馈神经网络。它很难学习并长期保存信息。为了解决这个问题一个增大网络存储的想法随之产生。采用了特殊隐式单元的LSTM便是为了长期的保存输入。一种称作记忆细胞的特殊单元类似累加器和门控神经元它在下一个时间步长将拥有一个权值并联接到自身拷贝自身状态的真实值和累积的外部信号但这种自联接是由另一个单元学习并决定何时清除记忆内容的乘法门控制的解决了RNN在长序列训练过程中存在的梯度消失和梯度爆炸的问题。3.2 LSTM缺点并行处理上存在劣势。与一些最新的网络相对效果一般RNN的梯度问题在LSTM及其变种里面得到了一定程度的解决但还是不够。它可以处理100个量级的序列而对于1000个量级或者更长的序列则依然会显得很棘手计算费时。每一个LSTM的cell里面都意味着有4个全连接层(MLP)如果LSTM的时间跨度很大并且网络又很深这个计算量会很大很耗时。4 基于Pytorch的LSTM代码实现下面我们就用一个简单的小例子来说明如何使用Pytorch来构建LSTM模型。我们使用正弦函数和余弦函数来构造时间序列而正余弦函数之间是成导数关系所以我们可以构造模型来学习正弦函数与余弦函数之间的映射关系通过输入正弦函数的值来预测对应的余弦函数的值。正弦函数和余弦函数对应关系图如下图所示可以看到每一个函数曲线上每一个正弦函数的值都对应一个余弦函数值。但其实如果只关心正弦函数的值本身而不考虑当前值所在的时间那么正弦函数值和余弦函数值不是一一对应关系。例如当t2.5t2.5t2.5和t6.8t6.8t6.8时sin(t)0.5sin(t)0.5sin(t)0.5但在这两个不同的时刻cos(t)cos(t)cos(t)的值却不一样也就是说如果不考虑时间同一个正弦函数值可能对应了不同的几个余弦函数值。对于传统的神经网络来说它仅仅基于当前的输入来预测输出对于这种同一个输入可能对应多个输出的情况不再适用。我们取正弦函数的值作为LSTM的输入来预测余弦函数的值。基于Pytorch来构建LSTM模型采用1个输入神经元1个输出神经元16个隐藏神经元作为LSTM网络的构成参数平均绝对误差LMSE作为损失误差使用Adam优化算法来训练LSTM神经网络。基于Anaconda和Python3.6的完整代码如下# -*- coding:UTF-8 -*-importnumpyasnpimporttorchfromtorchimportnnimportmatplotlib.pyplotasplt# Define LSTM Neural NetworksclassLstmRNN(nn.Module): Parameters - input_size: feature size - hidden_size: number of hidden units - output_size: number of output - num_layers: layers of LSTM to stack def__init__(self,input_size,hidden_size1,output_size1,num_layers1):super().__init__()self.lstmnn.LSTM(input_size,hidden_size,num_layers)# utilize the LSTM model in torch.nnself.forwardCalculationnn.Linear(hidden_size,output_size)defforward(self,_x):x,_self.lstm(_x)# _x is input, size (seq_len, batch, input_size)s,b,hx.shape# x is output, size (seq_len, batch, hidden_size)xx.view(s*b,h)xself.forwardCalculation(x)xx.view(s,b,-1)returnxif__name____main__:# create databasedata_len200tnp.linspace(0,12*np.pi,data_len)sin_tnp.sin(t)cos_tnp.cos(t)datasetnp.zeros((data_len,2))dataset[:,0]sin_t dataset[:,1]cos_t datasetdataset.astype(float32)# plot part of the original datasetplt.figure()plt.plot(t[0:60],dataset[0:60,0],labelsin(t))plt.plot(t[0:60],dataset[0:60,1],labelcos(t))plt.plot([2.5,2.5],[-1.3,0.55],r--,labelt 2.5)# t 2.5plt.plot([6.8,6.8],[-1.3,0.85],m--,labelt 6.8)# t 6.8plt.xlabel(t)plt.ylim(-1.2,1.2)plt.ylabel(sin(t) and cos(t))plt.legend(locupper right)# choose dataset for training and testingtrain_data_ratio0.5# Choose 80% of the data for testingtrain_data_lenint(data_len*train_data_ratio)train_xdataset[:train_data_len,0]train_ydataset[:train_data_len,1]INPUT_FEATURES_NUM1OUTPUT_FEATURES_NUM1t_for_trainingt[:train_data_len]# test_x train_x# test_y train_ytest_xdataset[train_data_len:,0]test_ydataset[train_data_len:,1]t_for_testingt[train_data_len:]# ----------------- train -------------------train_x_tensortrain_x.reshape(-1,5,INPUT_FEATURES_NUM)# set batch size to 5train_y_tensortrain_y.reshape(-1,5,OUTPUT_FEATURES_NUM)# set batch size to 5# transfer data to pytorch tensortrain_x_tensortorch.from_numpy(train_x_tensor)train_y_tensortorch.from_numpy(train_y_tensor)# test_x_tensor torch.from_numpy(test_x)lstm_modelLstmRNN(INPUT_FEATURES_NUM,16,output_sizeOUTPUT_FEATURES_NUM,num_layers1)# 16 hidden unitsprint(LSTM model:,lstm_model)print(model.parameters:,lstm_model.parameters)loss_functionnn.MSELoss()optimizertorch.optim.Adam(lstm_model.parameters(),lr1e-2)max_epochs10000forepochinrange(max_epochs):outputlstm_model(train_x_tensor)lossloss_function(output,train_y_tensor)loss.backward()optimizer.step()optimizer.zero_grad()ifloss.item()1e-4:print(Epoch [{}/{}], Loss: {:.5f}.format(epoch1,max_epochs,loss.item()))print(The loss value is reached)breakelif(epoch1)%1000:print(Epoch: [{}/{}], Loss:{:.5f}.format(epoch1,max_epochs,loss.item()))# prediction on training datasetpredictive_y_for_traininglstm_model(train_x_tensor)predictive_y_for_trainingpredictive_y_for_training.view(-1,OUTPUT_FEATURES_NUM).data.numpy()# torch.save(lstm_model.state_dict(), model_params.pkl) # save model parameters to files# ----------------- test -------------------# lstm_model.load_state_dict(torch.load(model_params.pkl)) # load model parameters from fileslstm_modellstm_model.eval()# switch to testing model# prediction on test datasettest_x_tensortest_x.reshape(-1,5,INPUT_FEATURES_NUM)# set batch size to 5, the same value with the training settest_x_tensortorch.from_numpy(test_x_tensor)predictive_y_for_testinglstm_model(test_x_tensor)predictive_y_for_testingpredictive_y_for_testing.view(-1,OUTPUT_FEATURES_NUM).data.numpy()# ----------------- plot -------------------plt.figure()plt.plot(t_for_training,train_x,g,labelsin_trn)plt.plot(t_for_training,train_y,b,labelref_cos_trn)plt.plot(t_for_training,predictive_y_for_training,y--,labelpre_cos_trn)plt.plot(t_for_testing,test_x,c,labelsin_tst)plt.plot(t_for_testing,test_y,k,labelref_cos_tst)plt.plot(t_for_testing,predictive_y_for_testing,m--,labelpre_cos_tst)plt.plot([t[train_data_len],t[train_data_len]],[-1.2,4.0],r--,labelseparation line)# separation lineplt.xlabel(t)plt.ylabel(sin(t) and cos(t))plt.xlim(t[0],t[-1])plt.ylim(-1.2,4)plt.legend(locupper right)plt.text(14,2,train,size15,alpha1.0)plt.text(20,2,test,size15,alpha1.0)plt.show()训练的过程如下该模型在训练集和测试集上的结果如下图中红色虚线的左边表示该模型在训练数据集上的表现右边表示该模型在测试数据集上的表现。可以看到使用LSTM构建训练模型我们可以仅仅使用正弦函数在 t 时刻的值作为输入来准确预测 t 时刻的余弦函数值不用额外添加当前的时间信息、速度信息等。5 LSTM变体5.1 双向LSTM单向的RNN是根据前面的信息推出后面的但有时候只看前面的词是不够的可能需要预测的词语和后面的内容也相关那么此时需要一种机制能够让模型不仅能够从前往后的具有记忆。此时双向LSTM可以解决这个问题。由于是双向LSTM所以每个方向的LSTM都会有一个输出最终的输出会有2部分所以往往需要concat的操作。在单向LSTM中output最后一个time step的输出和最后一层隐藏状态hnh_nhn​的输出相同那么双向LSTM呢双向LSTM中output按照正反计算结果的顺序在最后一个维度进行拼接正向第一个time step输出拼接反向的最后一个time step输出hidden state按照得到的结果在第0个维度进行拼接正向第一层之后接着是反向第一层正向第二层之后接着是反向第二层。。。前向LSTM中output最后一个time step的输出和最后一层 前向传播 隐藏状态h_n的输出相同后向LSTM中output最后一个time step的输出和最后一层 后向传播 隐藏状态h_n的输出相同5.2 PeepholeLSTM就是计算输入门、遗忘门和输出门 的时候我们不仅仅考虑h和x还将C考虑进来5.3 coupled LSTM输入门和遗忘门二合一5.4 Conv LS可以看到conv LSTM中也使用了peephole LSTM的结构——cell部分也用于遗忘门和输入门的计算于是我们有如下的计算流程在这里*表示 卷积操作 ●表示哈达玛积另一种convLSTM的理解方法是我们普通的LSTM可以看成最后两个维度都是1 的ConvLSTM其中卷积核大小为1×1

相关新闻

Intelli IDEA:Cannot connect to already running IDE instance. Process xxx is still running的原因及解决方法

Intelli IDEA:Cannot connect to already running IDE instance. Process xxx is still running的原因及解决方法

2026/8/13 20:51:37

1.查看自己的idea以及版本这里用WebStorm2023.2 为例,环境为windows;2. 原因是jetbrain启动时候会在 C:\Users\lingm\AppData\Roaming\JetBrains\WebStorm2023.2 目录下创建一个.lock文件,这个文件打开里面记录的是当前idea的进程号,每次idea…

关于STM32(或GD32)ADC的理解及使用

关于STM32(或GD32)ADC的理解及使用

2026/8/13 20:51:37

强烈推荐野火的资料:零角度玩转STM32-F103指南。相关ADC的部分,看得出来他们很用心去写的。下面主要以野火的资料和代码进行整理。 STM32的最大转换速率为1Mhz,也就是转换时间是1us,不要让ADC的时钟超过14M(查看时钟树…

为什么选择Giffy Dialog?Flutter动画对话框的性能与设计优势

为什么选择Giffy Dialog?Flutter动画对话框的性能与设计优势

2026/8/13 20:51:37

为什么选择Giffy Dialog?Flutter动画对话框的性能与设计优势 【免费下载链接】giffy_dialog A Flutter package for a quick and handy giffy dialog. 项目地址: https://gitcode.com/gh_mirrors/gi/giffy_dialog Giffy Dialog是一款专为Flutter开发者打造的…

TransUNet完整训练教程:从零开始掌握医学图像分割的终极指南

TransUNet完整训练教程:从零开始掌握医学图像分割的终极指南

2026/8/13 22:01:49

TransUNet完整训练教程:从零开始掌握医学图像分割的终极指南 【免费下载链接】TransUNet This repository includes the official project of TransUNet, presented in our paper: TransUNet: Transformers Make Strong Encoders for Medical Image Segmentation. …

做电商网站建设大作业:从零基础到上线,那些踩坑与成长的故事

做电商网站建设大作业:从零基础到上线,那些踩坑与成长的故事

2026/8/13 22:01:49

说实话,接到“电子商务网站建设大作业”这个课题的时候,我的第一反应不是兴奋,而是深深的头大。在这个万物皆可电商的时代,似乎每个人都想开网店,每个人都想搞流量,但真当你自己双手沾泥去搭建一个完整的电商平台时,才发现这背后涉及的复杂度远超想象。它不仅仅是一门课…

Vue.js对象操作全解析:从基础访问到响应式合并实战

Vue.js对象操作全解析:从基础访问到响应式合并实战

2026/8/13 22:01:49

1. 从日常开发痛点说起:为什么需要掌握对象操作?在Vue.js项目中,无论是处理从后端API返回的复杂JSON数据,还是在组件内部管理响应式状态,JavaScript对象都是我们打交道最频繁的数据结构之一。我见过不少开发者&#xf…

二叉树中序遍历:原理、实现与工程实践

二叉树中序遍历:原理、实现与工程实践

2026/8/13 22:01:49

1. 中序遍历的核心概念与应用场景中序遍历(In-order Traversal)是二叉树遍历的三种基本方式之一,它的遍历顺序遵循"左子树-根节点-右子树"的原则。这种遍历方式之所以重要,是因为它能以升序方式输出二叉搜索树&#xff…

微信机器人终极指南:30分钟打造你的智能聊天助手

微信机器人终极指南:30分钟打造你的智能聊天助手

2026/8/13 22:01:49

微信机器人终极指南:30分钟打造你的智能聊天助手 【免费下载链接】wechat-bot 🤖 Multi-platform IM AI Agent for Telegram, WhatsApp, Lark, and WeChat. Connects ChatGPT / Claude / Kimi / DeepSeek / Ollama / Pi for auto-replies, community ana…

Python数据可视化进阶:matplotlib画出专业图表的8个实战技巧

Python数据可视化进阶:matplotlib画出专业图表的8个实战技巧

2026/8/13 21:51:48

文章目录环境信息前言:matplotlib能画图,但离"专业"还差这8个技巧技巧一:中文字体 负号,90%的人第一张图就卡在这技巧二:双Y轴,让两个不同量纲的数据同框对比技巧三:子图布局&#x…

比较好的亚太EMBA,问了6位校友师资差别真的挺大

比较好的亚太EMBA,问了6位校友师资差别真的挺大

2026/8/13 11:01:28

比较好的亚太EMBA核心差异先看什么?对于希望兼顾工作与系统管理能力提升的亚太区高管而言,筛选匹配度高的EMBA项目时,师资配置是决定学习体验与实际收获的核心要素之一。我们结合3-4个公开信息透明、办学历史较长的亚太区主流EMBA项目特点&am…

备考3个月对比6份资料 海外游学的亚洲EMBA面试注意点

备考3个月对比6份资料 海外游学的亚洲EMBA面试注意点

2026/8/11 8:44:43

备考海外游学的亚洲EMBA面试,核心要围绕项目国际化设计逻辑、个人跨文化管理经验匹配度两个维度准备,避免把游学模块等同于普通旅游参访的认知偏差。不少备考者花3个月对比6份资料,却容易忽略面试官对“国际视野落地能力”的考察——比如香港…

比较好的国内EMBA,问了二十位校友聊透人脉价值

比较好的国内EMBA,问了二十位校友聊透人脉价值

2026/8/13 17:17:06

比较好的国内EMBA核心差异体现在哪些方面?比较好的国内EMBA的核心长期价值,很大程度上依托于校友网络的连接质量与资源生态的活跃度,这也是不少高管在择校时优先考量的因素。我们结合3-4个市场关注度较高的项目公开信息,从课程、师…

电商毛利率别再手动算了!2026年3种自动分析工具实测对比

电商毛利率别再手动算了!2026年3种自动分析工具实测对比

2026/8/13 0:00:21

一、开篇:毛利率——电商运营最该盯但最难盯的指标 电商运营中有一个指标,几乎所有老板都会问,但几乎所有运营都回答得不够确定——毛利率。不是"店铺毛利率",而是"每条链接的毛利率""每个品类的毛利率…

15-SaaS系统灰度发布:滚动更新、金丝雀发布、不停机迭代

15-SaaS系统灰度发布:滚动更新、金丝雀发布、不停机迭代

2026/8/13 0:00:21

15-SaaS系统灰度发布:滚动更新、金丝雀发布、不停机迭代 一、为什么需要不停机发布? 传统发布方式:停服务 → 替换包 → 启服务。在内部系统里勉强能用,但在SaaS系统中是灾难。 我们的无人售货柜SaaS平台服务全国几千台设备&#…

17-线上Bug热修复流程:紧急分支、补丁合并、版本快速回退方案

17-线上Bug热修复流程:紧急分支、补丁合并、版本快速回退方案

2026/8/13 0:00:21

17-线上Bug热修复流程:紧急分支、补丁合并、版本快速回退方案 前言 大家好,我是黒漂技术佬。 线上出 Bug 这种事,就像你正吃着火锅唱着歌,突然接到电话说"柜子门打不开了"。炸不炸?慌不慌?别急&a…

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

2026/8/8 5:07:31

一天写完毕业论文在2026年已不再是天方夜谭。2026年最炸裂、实测能大幅提速的AI论文写作工具,覆盖选题构思、文献整理、内容生成、格式排版等核心场景,真正帮你高效搞定论文难题。 一、全流程王者:一站式搞定论文全链路(一天定稿首…

导师推荐!2026最新AI论文工具测评与实用推荐

导师推荐!2026最新AI论文工具测评与实用推荐

2026/8/9 13:42:46

2026年真正好用的AI论文工具,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。 一、…

告别游戏崩溃:XCOM 2模组管理器的智能革命

告别游戏崩溃:XCOM 2模组管理器的智能革命

2026/8/8 2:30:15

告别游戏崩溃:XCOM 2模组管理器的智能革命 【免费下载链接】xcom2-launcher The Alternative Mod Launcher (AML) is a replacement for the default game launchers from XCOM 2 and XCOM Chimera Squad. 项目地址: https://gitcode.com/gh_mirrors/xc/xcom2-lau…