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. 总结

步骤函数
对变长序列做 paddingpad_sequence()
打包以适配 RNNpack_padded_sequence()
RNN 输出解包pad_packed_sequence()

pad_packed_sequence() 是用于 还原 RNN 变长输出的关键函数,它确保可以在后续模块中 使用统一形状的数据 进行处理。

Logo

开源鸿蒙跨平台开发社区汇聚开发者与厂商,共建“一次开发,多端部署”的开源生态,致力于降低跨端开发门槛,推动万物智联创新。

更多推荐