標籤:屬性 https 元素 dde text alt loss pac rnn
torch.nn.utils.rnn.pack_padded_sequence()
這裡的pack,理解成壓緊比較好。 將一個 填充過的變長序列 壓緊。(填充時候,會有冗餘,所以壓緊一下)
其中pack的過程為:(注意pack的形式,不是按行壓,而是按列壓)
(下面方框內為PackedSequence對象,由data和batch_sizes組成)
輸入的形狀可以是(T×B×* )。T是最長序列長度,B是batch size,*代表任意維度(可以是0)。如果batch_first=True的話,那麼相應的 input size 就是 (B×T×*)。
Variable中儲存的序列,應該按序列長度的長短排序,長的在前,短的在後。即input[:,0]代表的是最長的序列,input[:, B-1]儲存的是最短的序列。
NOTE: 只要是維度大於等於2的input都可以作為這個函數的參數。你可以用它來打包labels,然後用RNN的輸出和打包後的labels來計算loss。通過PackedSequence對象的.data屬性可以擷取 Variable。
參數說明:
- input (Variable) – 變長序列 被填充後的 batch
- lengths (list[int]) –
Variable 中 每個序列的長度。
- batch_first (bool, optional) – 如果是
True,input的形狀應該是B*T*size。
傳回值:
一個PackedSequence 對象。
torch.nn.utils.rnn.pad_packed_sequence()
填充packed_sequence。
上面提到的函數的功能是將一個填充後的變長序列壓緊。 這個操作和pack_padded_sequence()是相反的。把壓緊的序列再填充回來。
返回的Varaible的值的size是 T×B×*, T 是最長序列的長度,B 是 batch_size,如果 batch_first=True,那麼傳回值是B×T×*。
Batch中的元素將會以它們長度的逆序排列。
參數說明:
- sequence (PackedSequence) – 將要被填充的 batch
- batch_first (bool, optional) – 如果為True,返回的資料的格式為
B×T×*。
傳回值: 一個tuple,包含被填充後的序列,和batch中序列的長度列表
一個例子:
輸出:
此時PackedSequence對象輸入RNN後,輸出RNN的還是PackedSequence對象
參考:
https://www.cnblogs.com/lindaxin/p/8052043.html
https://pytorch.org/docs/stable/nn.html?highlight=pack_padded_sequence#torch.nn.utils.rnn.pack_padded_sequence
Pytorch中的RNN之pack_padded_sequence()和pad_packed_sequence()