【PyTorch】torch.nn.utils.rnn.pad_packed_sequence() 函数: 打包后的变长序列还原为填充形式
·
torch.nn.utils.rnn.pad_packed_sequence()
torch.nn.utils.rnn.pad_packed_sequence() 是 PyTorch 中用于将 打包后的变长序列(PackedSequence)还原为填充形式 的函数。它通常配合 pack_padded_sequence() 使用,用于从 RNN(如 LSTM、GRU)的输出中恢复原始 shape,方便进一步处理(如分类、损失计算等)。
1. 应用背景
在变长序列处理时,我们使用:
pad_sequence():将变长序列对齐(padding);pack_padded_sequence():打包序列,跳过 padding 提高 RNN 训练效率;pad_packed_sequence():将 RNN 的输出从 PackedSequence 解包回 padded 格式。
2. 语法
torch.nn.utils.rnn.pad_packed_sequence(sequence, batch_first=False, padding_value=0.0, total_length=None)
| 参数 | 说明 |
|---|---|
sequence | 一个 PackedSequence 对象(通常是 RNN 的输出) |
batch_first | 是否返回 (batch_size, max_seq_len, *) 格式(默认为 False) |
padding_value | 用于填充的数值(默认 0.0) |
total_length | 返回的最大长度(可用于强制指定输出长度) |
3. 示例:与 RNN 搭配使用
import torch
import torch.nn as nn
from torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence
# 3 个变长序列
seq1 = torch.tensor([1.0, 2.0, 3.0])
seq2 = torch.tensor([4.0, 5.0])
seq3 = torch.tensor([6.0])
# 填充并打包
sequences = [seq1, seq2, seq3]
padded = pad_sequence(sequences, batch_first=True) # shape: (3, 3)
lengths = torch.tensor([3, 2, 1])
packed = pack_padded_sequence(padded, lengths, batch_first=True, enforce_sorted=False)
# 输入 RNN
rnn = nn.RNN(input_size=1, hidden_size=4, batch_first=True)
packed_input = packed.data.unsqueeze(-1) # 添加 input_size=1
packed_output, _ = rnn(packed)
# 解包
output, output_lengths = pad_packed_sequence(packed_output, batch_first=True)
print("还原后的输出形状:", output.shape) # e.g. (3, 3, 4)
print("每个序列的有效长度:", output_lengths) # tensor([3, 2, 1])
4. 参数解析
total_length
如果你在使用 pack_padded_sequence() 时输入的是 被裁剪过的序列(如 padding 后又被截断),你可以用 total_length 强制还原到指定长度:
output, _ = pad_packed_sequence(packed_output, batch_first=True, total_length=5)
5. 常见搭配函数
| 函数 | 说明 |
|---|---|
pad_sequence() | 将变长序列 padding 成固定长度张量 |
pack_padded_sequence() | 将 padding 后的序列打包 |
pad_packed_sequence() | 将 PackedSequence 解包为 padding 格式 |
6. 应用场景
- NLP:处理不同长度句子输入输出
- 语音:处理不等长音频序列
- 时间序列建模:对 RNN 输出进行进一步处理,如 attention、分类等
7. 总结
| 步骤 | 函数 |
|---|---|
| 对变长序列做 padding | pad_sequence() |
| 打包以适配 RNN | pack_padded_sequence() |
| RNN 输出解包 | pad_packed_sequence() |
pad_packed_sequence() 是用于 还原 RNN 变长输出的关键函数,它确保可以在后续模块中 使用统一形状的数据 进行处理。
更多推荐
所有评论(0)