最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

Pytorch中關于RNN輸入和輸出的形狀總結(jié)

 更新時間:2023年06月15日 08:35:29   作者:會唱歌的豬233  
這篇文章主要介紹了Pytorch中關于RNN輸入和輸出的形狀總結(jié),具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教

Pytorch對RNN輸入和輸出的形狀總結(jié)

個人對于RNN的一些總結(jié)。

RNN的輸入和輸出

RNN的經(jīng)典圖如下所示


各個參數(shù)的含義

  • Xt: t時刻的輸入,形狀為[batch_size, input_dim]。對于整個RNN來說,總的X輸入為[seq_len, batch_size, input_dim],具體如何理解batch_size和seq_len在下面有說明。
  • St: t時刻隱藏層的狀態(tài),也有時用ht表示,形狀為[batch_size, hidden_size],St=f(U·Xt+W·St-1),通過W和U矩陣的映射,將embedding后的Xt和上一狀態(tài)St-1轉(zhuǎn)為St
  • Ot: t時刻的輸出,Ot=g(V·St),形狀為[batch_size, hidden_size],總的為輸出O為[seq_len, batch_size, hidden_size]

Pytorch中的使用

Pytorch中RNN函數(shù)如下

RNN的主要參數(shù)如下

nn.RNN(input_size, hidden_size, num_layers=1, bias=True)

參數(shù)解釋

  • input_size: 輸入特征的維度,一般rnn中輸入的是詞向量,那么就為embedding-dim
  • hidden_size: 隱藏層神經(jīng)元的個數(shù),或者也叫輸出的維度
  • num_layers: 隱藏層的個數(shù),默認為1

output=輸出O, 隱藏狀態(tài)St,其中輸出O=[time_step, batch_size, hidden_size],St為t時刻的隱藏層狀態(tài)

理解RNN中的batch_size和seq_len

深度學習中采用mini-batch的方法進行迭代優(yōu)化,在CNN中batch的思想較容易理解,一次輸入batch個圖片,進行迭代。但是RNN中引入了seq_len(time_step), 理解較為困難,下面是我自己的一些理解。

首先假如我有五句話,作為訓練的語料。

sentences = ["i like dog", "i love coffee", "i hate milk", "i like music", "i hate you"]

那么在輸入RNN之前要先進行embedding,比如one-hot encoding,容易得到這里的embedding-dim為9.

那么輸入的sentences可以表示為如下方式

t=0t=1t=2
batch1ilikedog
batch2ilovecoffee
batch3ihatemilk
batch4ilikemusic
batch5ihateyou

那么在RNN的訓練中。

  • t=0時, 輸入第一個batch[i, i, i, i, i]這里用字符表示,其實應該是對應的one-hot編碼。
  • t=1時,輸入第二個batch[like, love, hate, like, hate]
  • t=2時,輸入第三個batch[dog, coffee, milk, music, you]

那么對應的時間t來說,RNN需要對先后輸入的batch_size個字符進行前向計算迭代,得到輸出。

Pytorch雙向RNN隱藏層和輸出層結(jié)果拆分

1 RNN隱藏層和輸出層結(jié)果的形狀

從Pytorch官方文檔可以得到,對于批量化輸入的RNN來講,其隱藏層的shape為(num_directions*num_layers, batch_size, hidden_size)。

其輸出的shape為(seq_len, batch_size, D*hidden_size)。

2 雙向RNN情況下,隱藏層和輸出層結(jié)果拆分

當采用雙向RNN時,其輸出的結(jié)果包含正向和反向兩個方向輸出的結(jié)果。

2.1 輸出層結(jié)果拆分

其中對于輸出output來講,從官方文檔我們可以得到,其拆分正向和反向兩個方向結(jié)果的方法為:

output.shape = (seq_len, batch_size, num_directions*hidden_size)

output.view(seq_len, batch, num_directions, hidden_size)

其中,對于(num_directions)方向維度,正向和反向的維度值分別為??0???和??1?。

2.2 隱藏層結(jié)果拆分

而對于隱藏層,包括初始值h_0以及最終輸出h_n,也都包含兩個方向的隱藏狀態(tài),但是其拆分方式跟輸出層不一樣。

方法如下:

h_0, h_n.shape = (num_directions*num_layers, batch_size, hidden_size)

h_0, h_n.view(num_layers, num_directions, batch_size, hidden_size)

可以從簡單單層雙向RNN的輸出結(jié)果來驗證,此時RNN的輸出結(jié)果與最后一層的隱藏層結(jié)果是一樣的。

import torch
import torch.nn as nn
if __name__ == "__main__":
    # input_size: 3, hidden_size: 5, num_layers: 3
    BiRNN_Net = nn.RNN(3, 5, 3, bidirectional=True, batch_first=True)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    # batch_size: 1, seq_len: 1, input_size: 3
    inputs = torch.zeros(1, 1, 3, device=device)
    # state: (num_directions*num_layers, batch_size, hidden_size)
    state = torch.randn(6, 1, 5, device=device)
    BiRNN_Net.to(device)
    output, hidden = BiRNN_Net(inputs, state)
    output_re = output.reshape((1, 1, 2, 5))
    hidden_re = hidden.reshape((3, 2, 1, 5))
    print(output)
    print(output_re)
    print(hidden)
    print(hidden_re)

輸出結(jié)果可以看出,隱藏層的結(jié)果是優(yōu)先num_layers網(wǎng)絡層數(shù)這一個維度來構成的。

tensor([[[ 0.3939, -0.9160, ?0.5054, ?0.2949, -0.5225, ?0.0533, ?0.4197,
? ? ? ? ? -0.7200, -0.1262, -0.7975]]], device='cuda:0',
? ? ? ?grad_fn=<CudnnRnnBackward0>)
tensor([[[[ 0.3939, -0.9160, ?0.5054, ?0.2949, -0.5225],
? ? ? ? ? [ 0.0533, ?0.4197, -0.7200, -0.1262, -0.7975]]]], device='cuda:0',
? ? ? ?grad_fn=<ReshapeAliasBackward0>)
tensor([[[-0.2606, ?0.5410, -0.2663, ?0.6418, -0.2902]],
? ? ? ? [[ 0.1367, ?0.7222, -0.3051, -0.6410, -0.3062]],
? ? ? ? [[ 0.2433, ?0.3287, -0.4809, -0.1782, -0.5582]],
? ? ? ? [[ 0.4824, -0.8529, ?0.7604, ?0.8508, -0.1902]],
? ? ? ? [[ 0.3939, -0.9160, ?0.5054, ?0.2949, -0.5225]],
? ? ? ? [[ 0.0533, ?0.4197, -0.7200, -0.1262, -0.7975]]], device='cuda:0',
? ? ? ?grad_fn=<CudnnRnnBackward0>)
tensor([[[[-0.2606, ?0.5410, -0.2663, ?0.6418, -0.2902]],
? ? ? ? ?[[ 0.1367, ?0.7222, -0.3051, -0.6410, -0.3062]]],
? ? ? ? [[[ 0.2433, ?0.3287, -0.4809, -0.1782, -0.5582]],
? ? ? ? ?[[ 0.4824, -0.8529, ?0.7604, ?0.8508, -0.1902]]],
? ? ? ? [[[ 0.3939, -0.9160, ?0.5054, ?0.2949, -0.5225]],
? ? ? ? ?[[ 0.0533, ?0.4197, -0.7200, -0.1262, -0.7975]]]], device='cuda:0',
? ? ? ?grad_fn=<ReshapeAliasBackward0>)

總結(jié)

以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。

相關文章

  • pycharm如何為函數(shù)插入文檔注釋

    pycharm如何為函數(shù)插入文檔注釋

    這篇文章主要介紹了pycharm如何為函數(shù)插入文檔注釋,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-02-02
  • 解決Django migrate No changes detected 不能創(chuàng)建表的問題

    解決Django migrate No changes detected 不能創(chuàng)建表的問題

    今天小編就為大家分享一篇解決Django migrate No changes detected 不能創(chuàng)建表的問題,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-05-05
  • python中的字符串切割 maxsplit

    python中的字符串切割 maxsplit

    這篇文章主要介紹了python中的字符串切割 maxsplit,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-12-12
  • 如何用python抓取B站數(shù)據(jù)

    如何用python抓取B站數(shù)據(jù)

    今天介紹一個獲取B站數(shù)據(jù)的Python擴展庫-bilibili_api,對此感興趣的同學,可以實驗一下
    2021-05-05
  • Python之dict(或?qū)ο?與json之間的互相轉(zhuǎn)化實例

    Python之dict(或?qū)ο?與json之間的互相轉(zhuǎn)化實例

    今天小編就為大家分享一篇Python之dict(或?qū)ο?與json之間的互相轉(zhuǎn)化實例,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-06-06
  • 從零學Python之入門(五)縮進和選擇

    從零學Python之入門(五)縮進和選擇

    空白在Python中是重要的。事實上行首的空白是重要的。它稱為縮進。在邏輯行首的空白(空格和制表符)用來決定邏輯行的縮進層次,從而用來決定語句的分組。
    2014-05-05
  • flask庫中sessions.py的使用小結(jié)

    flask庫中sessions.py的使用小結(jié)

    在Flask中Session是一種用于在不同請求之間存儲用戶數(shù)據(jù)的機制,Session默認是基于客戶端Cookie的,但數(shù)據(jù)會經(jīng)過加密簽名,防止篡改,下面就來具體介紹一下如何使用
    2025-07-07
  • python數(shù)據(jù)結(jié)構算法分析

    python數(shù)據(jù)結(jié)構算法分析

    這篇文章主要介紹了python數(shù)據(jù)結(jié)構算法分析,在python的數(shù)據(jù)結(jié)構的章節(jié)中,我們上次學習到了python面向?qū)ο蟮乃枷?,即我們想用程序來實現(xiàn)一個東西,我們需是用對象的特征來描述我們想構建的對象。感興趣的小伙伴可以查看下面內(nèi)容</P><P>
    2021-12-12
  • 詳解在Python程序中使用Cookie的教程

    詳解在Python程序中使用Cookie的教程

    這篇文章主要介紹了詳解在Python程序中使用Cookie的教程,Cookie在無論哪種語言的網(wǎng)絡編程學習當中都是重要的知識點,需要的朋友可以參考下
    2015-04-04
  • Selenium(Python web測試工具)基本用法詳解

    Selenium(Python web測試工具)基本用法詳解

    這篇文章主要介紹了Selenium(Python web測試工具)基本用法,結(jié)合實例形式分析了Selenium的基本安裝、簡單使用方法及相關操作技巧,需要的朋友可以參考下
    2018-08-08

最新評論

唐河县| 京山县| 井研县| 南雄市| 新和县| 孙吴县| 都兰县| 斗六市| 昭平县| 明光市| 瑞金市| 抚宁县| 永福县| 桦南县| 柞水县| 酒泉市| 丹巴县| 池州市| 汉川市| 丹阳市| 永顺县| 关岭| 团风县| 上栗县| 万盛区| 新巴尔虎左旗| 西乡县| 黑山县| 诸暨市| 木兰县| 永兴县| 抚顺市| 延川县| 乌拉特前旗| 德清县| 耿马| 侯马市| 岳阳市| 夏津县| 神木县| 邯郸市|