欧美bbbwbbbw肥妇,免费乱码人妻系列日韩,一级黄片

對pytorch中不定長序列補齊的操作

 更新時間:2021年05月31日 08:42:58   作者:XJTU-Qidong  
這篇文章主要介紹了對pytorch中不定長序列補齊的操作,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教

第二種方法通常是在load一個batch數(shù)據(jù)時, 在collate_fn中進行補齊的.

以下給出兩種思路:

第一種思路是比較容易想到的, 就是對一個batch的樣本進行遍歷, 然后使用np.pad對每一個樣本進行補齊.

for unit in data:
        mask = np.zeros(max_length)
        s_len = len(unit[0])    # calculate the length of sequence in each unit
        mask[: s_len] = 1
        unit[0] = np.pad(unit[0], (0, max_length - s_len), 'constant', constant_values=(0, 0))
        mask_batch.append(mask)

但是這種方法在batch size很大的情況下會很慢, 因為使用for循環(huán)進行了遍歷. 我在實際用的時候, 當batch_size=128時, 一個batch的加載時間甚至是一個batch訓練時間的幾倍!

因此, 我想到如何并行地對序列進行補齊. 第二種方法的思路就是使用torch中自帶的pad_sequence來并行補齊.

batch_sequence = list(map(lambda x: torch.tensor(x[findex]), x_data))
batch_data[feat] = torch.nn.utils.rnn.pad_sequence(batch_sequence).T

可以看到這里使用pad_sequence一次性對整個batch進行補齊. 下面對這個函數(shù)進行詳細說明.

pad_sequence詳解

from torch.utils.rnn import pad_sequence
a = torch.ones(10)
b = torch.ones(6)
c = torch.ones(20)
abc = pad_sequence([a,b,c])  # shape(20, 3)

注意這個函數(shù)接收的是一個元素為tensor的列表, 而不是tensor.

最終, 這個函數(shù)會將所有tensor轉(zhuǎn)換為tensor矩陣#shape(max_length, batch_size). 因此, 在使用完后通常還需要轉(zhuǎn)置一下.

補充:PyTorch中用于RNN變長序列填充函數(shù)的簡單使用

1、PyTorch中RNN變長序列的問題   

RNN在處理變長序列時有它的優(yōu)勢。在分批處理變長序列問題時,每個序列的長度往往不會完全相等,因此針對一個batch中序列長度不一的情況,需要對某些序列進行PAD(填充)操作,使得一個batch內(nèi)的序列長度相等。   

PyTorch中的pack_padded_sequence和pad_packed_sequence可處理上述問題,以下用一個示例演示這兩個函數(shù)的簡單使用方法。

2、填充函數(shù)簡介

“壓縮”函數(shù):用于將填充后的序列tensor進行壓縮,方便RNN處理

pack_padded_sequence(input, lengths, batch_first=False, enforce_sorted=True)

(1)input->被“壓縮”的tensor,維度一般為[batch_size,_max_seq_len[,embedding_size]]或者[max_seq_len,batch_size[,embedding_size]]

若input維度為:[batch_size,_max_seq_len[,embedding_size]]

要將batch_first設置為True,這表示input的第一個維度為batch的數(shù)量

若input維度為:[max_seq_len,batch_size[,embedding_size]]

要將batch_first設置為False(默認值),這表示input的第一個維度不是batch的數(shù)量

(2)lengths->lengths參數(shù)表示一個batch中序列真實長度,類型為列表,在例子中詳細說明

(3)batch_first->表示batch的數(shù)量是否在input的第一維度,默認值為False

(4)enforce_sorted->input中的會自動按照lengths的情況進行排序,默認值為

“解壓”函數(shù):該函數(shù)與"壓縮函數(shù)"相對應,經(jīng)“壓縮函數(shù)”處理的輸入經(jīng)過RNN得到的最終結(jié)果可以利用該函數(shù)進行“解壓”

pad_packed_sequence(sequence, batch_first=False, padding_value=0.0, total_length=None):

(1)sequence->壓縮函數(shù)處理過的input經(jīng)RNN后得到的結(jié)果

(2)batch_first->與“壓縮”函數(shù)中的batch_first一致

(3)padding_value->序列進行填充時使用的索引,默認為0

(4)total_length->暫略

3、PyTorch代碼示例

代碼如下(示例):

# Create by leslie_miao on 2020/11/1
import torch
import torch.nn as nn
d_model = 10 # 詞嵌入的維度
hidden_size = 20 # lstm隱藏層單元數(shù)量
layer_num = 1 # lstm層數(shù)
# 輸入inputs,維度為[batch_size,max_seq_len]=[3,4],其中0代表填充
# 該input包含3個序列,每個序列的真實長度分別為: 4 3 2
inputs = torch.tensor([[1,2,3,4],[1,2,3,0],[1,2,0,0]])
embedding = nn.Embedding(5,d_model)
# 獲取詞嵌入后的inputs 當前inputs的維度為[batch_size,max_seq_len,d_model]=[3,4,10]
inputs = embedding(inputs)
# 查看inputs的維度
print(inputs.size())
# print: torch.Size([3, 4, 10])
# 利用“壓縮”函數(shù)對inputs進行壓縮處理,[4,3,2]分別為inputs中序列的真實長度,batch_first=True表示inputs的第一維是batch_size
inputs = nn.utils.rnn.pack_padded_sequence(inputs,lengths=[4,3,2],batch_first=True)
# 查看經(jīng)“壓縮”函數(shù)處理過的inputs的維度
print(inputs[0].size())
# print: torch.Size([9, 10])
# 定義RNN網(wǎng)絡
network = nn.LSTM(input_size=d_model,hidden_size=hidden_size,batch_first=True,num_layers=layer_num)
# 初始化RNN相關(guān)門參數(shù)
c_0 = torch.zeros((layer_num,3,hidden_size))
h_0 = torch.zeros((layer_num,3,hidden_size)) # [rnn層數(shù),batch_size,hidden_size]
# inputs經(jīng)過RNN網(wǎng)絡后得到的結(jié)果outputs
output,(h_n,c_n) = network(inputs,(h_0,c_0))
#查看未經(jīng)“解壓函數(shù)”處理的outputs維度
print(output[0].size())
# print: torch.Size([9, 20])
# 利用“解壓函數(shù)”對outputs進行解壓操作,其中batch_first設置與“壓縮函數(shù)相同”,padding_value為0
output = nn.utils.rnn.pad_packed_sequence(output,batch_first=True,padding_value=0)
# 查看經(jīng)“解壓函數(shù)”處理的outputs維度
print(output[0].size())
# print:torch.Size([3, 4, 20])

總結(jié)

介紹了PyTorch中兩個應用于RNN變長序列填充的函數(shù)pack_padded_sequence和 pad_packed_sequence的簡單使用方法,歡迎指正交流!

相關(guān)文章

  • Python將8位的圖片轉(zhuǎn)為24位的圖片實現(xiàn)方法

    Python將8位的圖片轉(zhuǎn)為24位的圖片實現(xiàn)方法

    這篇文章主要介紹了Python將8位的圖片轉(zhuǎn)為24位的圖片的實現(xiàn)代碼,非常不錯,具有一定的參考借鑒價值,需要的朋友可以參考下
    2018-10-10
  • 實例探究Python以并發(fā)方式編寫高性能端口掃描器的方法

    實例探究Python以并發(fā)方式編寫高性能端口掃描器的方法

    端口掃描器就是向一批端口上發(fā)送請求來檢測端口是否打開的程序,這里我們以實例探究Python以并發(fā)方式編寫高性能端口掃描器的方法
    2016-06-06
  • Form表單及django的form表單的補充

    Form表單及django的form表單的補充

    這篇文章主要介紹了Form表單及django的form表單的補充,文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友可以參考下
    2019-07-07
  • 深入淺析pycharm中 Make available to all projects的含義

    深入淺析pycharm中 Make available to all projects的含義

    這篇文章主要介紹了pycharm中 Make available to all projects的含義,本文給大家介紹的非常詳細,對大家的學習或工作具有一定的參考借鑒價值,需要的朋友可以參考下
    2020-09-09
  • 對python:threading.Thread類的使用方法詳解

    對python:threading.Thread類的使用方法詳解

    今天小編就為大家分享一篇對python:threading.Thread類的使用方法詳解,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2019-01-01
  • opencv實現(xiàn)圖像幾何變換

    opencv實現(xiàn)圖像幾何變換

    這篇文章主要為大家詳細介紹了opencv實現(xiàn)圖像幾何變換,文中示例代碼介紹的非常詳細,具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2021-03-03
  • python中星號變量的幾種特殊用法

    python中星號變量的幾種特殊用法

    不知道大家知不知道在Python中,星號除了用于乘法數(shù)值運算和冪運算外,還有一種特殊的用法"在變量前添加單個星號或兩個星號",實現(xiàn)多參數(shù)的傳入或變量的拆解,本文將詳細介紹"星號參數(shù)"的用法。有需要的可以參考借鑒。
    2016-09-09
  • opencv 圖像濾波(均值,方框,高斯,中值)

    opencv 圖像濾波(均值,方框,高斯,中值)

    這篇文章主要介紹了opencv 圖像濾波(均值,方框,高斯,中值),文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友們下面隨著小編來一起學習學習吧
    2020-07-07
  • Python爬蟲實現(xiàn)自動登錄、簽到功能的代碼

    Python爬蟲實現(xiàn)自動登錄、簽到功能的代碼

    這篇文章主要介紹了Python爬蟲實現(xiàn)自動登錄、簽到功能的代碼,本文通過圖文并茂的形式給大家介紹的非常詳細,對大家的學習或工作具有一定的參考借鑒價值,需要的朋友可以參考下
    2020-08-08
  • Python實現(xiàn)問題回答小游戲

    Python實現(xiàn)問題回答小游戲

    這篇文章主要介紹了利用Python制作一個簡單的知識競賽小游戲,可以實現(xiàn)回答問題功能,文中的示例代碼介紹詳細,感興趣的同學快跟隨小編一起學習吧
    2021-12-12

最新評論