IT数码 购物 网址 头条 软件 日历 阅读 图书馆
TxT小说阅读器
↓语音阅读,小说下载,古典文学↓
图片批量下载器
↓批量下载图片,美女图库↓
图片自动播放器
↓图片自动播放器↓
一键清除垃圾
↓轻轻一点,清除系统垃圾↓
开发: C++知识库 Java知识库 JavaScript Python PHP知识库 人工智能 区块链 大数据 移动开发 嵌入式 开发工具 数据结构与算法 开发测试 游戏开发 网络协议 系统运维
教程: HTML教程 CSS教程 JavaScript教程 Go语言教程 JQuery教程 VUE教程 VUE3教程 Bootstrap教程 SQL数据库教程 C语言教程 C++教程 Java教程 Python教程 Python3教程 C#教程
数码: 电脑 笔记本 显卡 显示器 固态硬盘 硬盘 耳机 手机 iphone vivo oppo 小米 华为 单反 装机 图拉丁
 
   -> 人工智能 -> pytorch lstm -> 正文阅读

[人工智能]pytorch lstm

1、输入input:数据维度是 (seq, batch, feature),即序列长度、batch_size、每个时刻特征数量。

2、output, (hn, cn) = nn.LSTM(input_size, hidden_size, num_layers)

input_size:每时刻输入特征数量

hidden_size:隐藏层特征数量,决定着输出每一时刻的特征维度

num_layers:隐藏层数量

output:输出结果,维度为(seq, batch_size,hidden_size),相比于输入而言,输出的时间长度和batch_size都不变,只有特征维度发生变化,与设置的隐藏状态数一致。

hn:最后一个时刻的隐藏状态,每一层都会有一个最后时刻的隐藏状态,因此它的维度为(num_layers,batch_size,hidden_size)。

cn:最后一个时刻的记忆单元状态,每一层都会有一个最后时刻的状态,因此它的维度也为(num_layers,batch_size,hidden_size)。

关于output和hn,output为最后一层所有时刻的隐藏状态输出,而hn为最后一个时刻的所有层的隐藏状态输出。因此,可以得到output[-1] = hn[-1]。

3、注意事项,与CNN不同的是batch_size在中间,而CNN通常Batch_size在第一个维度,可通过batch_first=True设置与CNN保持一致。

5、双向LSTM,设置bidirectional=True。与单向lstm区别在于,output的特征维度增加一倍,维度变为(seq, batch_size,2*hidden_size),层数也相当于增加了一倍,即每一层都会相应增加一层逆向推理的层,那么hn和cn的一个维度加倍,即(2*num_layers,batch_size,hidden_size)。由于双向lstm先从左到右再从右到左,因此hn的输出反而与output的第一个维度相等,output[0,:,hidden_size:] = hn[-1]。

4、示例代码

import numpy as np
from torch import nn
import torch

x = torch.randn(13, 7, 3)#seq length 13, batch size 7, features 3
lstm = nn.LSTM(3, 5, 2)
output, (hn, cn) = lstm(x)
print('input shape: ', x.shape)          #13,7,3
print('output shape: ', output.shape)    #13,7,5
print('hn shape: ', hn.shape)            #2,7,5
print('cn shape: ', cn.shape)            #2,7,5
#print(output[-1]) 
#print(hn[-1])
print(output[-1] == hn[-1])              #True, 二者相等

#batch_first=True
x = torch.randn(13, 7, 3)#seq length 7, batch size 13, features 3
lstm = nn.LSTM(3, 5, 2, batch_first=True) #batch size 13
output, (hn, cn) = lstm(x)
print('input shape: ', x.shape)          #13,7,3
print('output shape: ', output.shape)    #13,7,5
print('hn shape: ', hn.shape)            #2,13,5
print('cn shape: ', cn.shape)            #2,13,5
#print(output[:,-1,:])                    #在序列维度上的最后一维
#print(hn[-1,])
print(output[:,-1,:] == hn[-1])         #True, 二者相等

#bidirectional=True
x = torch.randn(13, 7, 3)#seq length 13, batch size 7, features 3
lstm = nn.LSTM(3, 5, 2, bidirectional=True)
output, (hn, cn) = lstm(x)
print('input shape: ', x.shape)          #13,7,3
print('output shape: ', output.shape)    #13,7,10
print('hn shape: ', hn.shape)            #4,7,5
print('cn shape: ', cn.shape)            #4,7,5
print(output[-1]) 
print(hn[-1])
print(output[0, :, 5:] == hn[-1])      #True, 二者相等,由于双向lstm先从左到右再从右到左,因此hn的输出反而与output的第一个维度相等。

  人工智能 最新文章
2022吴恩达机器学习课程——第二课(神经网
第十五章 规则学习
FixMatch: Simplifying Semi-Supervised Le
数据挖掘Java——Kmeans算法的实现
大脑皮层的分割方法
【翻译】GPT-3是如何工作的
论文笔记:TEACHTEXT: CrossModal Generaliz
python从零学(六)
详解Python 3.x 导入(import)
【答读者问27】backtrader不支持最新版本的
上一篇文章      下一篇文章      查看所有文章
加:2021-08-20 15:05:55  更:2021-08-20 15:06:24 
 
开发: C++知识库 Java知识库 JavaScript Python PHP知识库 人工智能 区块链 大数据 移动开发 嵌入式 开发工具 数据结构与算法 开发测试 游戏开发 网络协议 系统运维
教程: HTML教程 CSS教程 JavaScript教程 Go语言教程 JQuery教程 VUE教程 VUE3教程 Bootstrap教程 SQL数据库教程 C语言教程 C++教程 Java教程 Python教程 Python3教程 C#教程
数码: 电脑 笔记本 显卡 显示器 固态硬盘 硬盘 耳机 手机 iphone vivo oppo 小米 华为 单反 装机 图拉丁

360图书馆 购物 三丰科技 阅读网 日历 万年历 2025年1日历 -2025/1/12 0:43:16-

图片自动播放器
↓图片自动播放器↓
TxT小说阅读器
↓语音阅读,小说下载,古典文学↓
一键清除垃圾
↓轻轻一点,清除系统垃圾↓
图片批量下载器
↓批量下载图片,美女图库↓
  网站联系: qq:121756557 email:121756557@qq.com  IT数码