PyTorch中如何对多维度变长序列进行填充构建批次?
PyTorch 双层变长3D序列批次构建方案
针对两层序列长度均不固定、特征维度D固定的3D序列批次填充需求,不需要引入额外第三方库,基于PyTorch原生API即可快速实现,核心逻辑是先统计两个变长维度的全局最大长度,再初始化填充张量后逐位置拷贝真实数据,比嵌套调用pad_sequence更直观高效。
实现步骤
- 统计维度参数:遍历批次内所有样本,计算一级序列的最大长度
max_l1、所有二级序列的最大长度max_l2,结合固定的特征维度D、批次大小batch_size,确定最终输出张量的形状为(batch_size, max_l1, max_l2, D)。 - 初始化填充基底:创建形状匹配上述尺寸、填充值为指定值(默认0)的空张量。
- 拷贝真实数据:双层遍历每个样本的一级、二级序列,把真实存在的序列值写入填充张量的对应位置,未写入的位置自然保留填充值,即可得到对齐后的批次张量。
代码实现
对应给出的输入样例,可直接运行的代码如下:
import torch # 原始输入序列 input1 = [ torch.tensor([[1, 1, 1], [2, 2, 2], [3, 3, 3]]), torch.tensor([[4, 4, 4], [5, 5, 5]]) ] input2 = [ torch.tensor([[1, 1, 1], [2, 2, 2], [3, 3, 3]]), torch.tensor([[6, 6, 6]]), torch.tensor([[4, 4, 4], [5, 5, 5]]) ] batch_inputs = [input1, input2] pad_val = 0 # 计算全局维度参数 feat_dim = batch_inputs[0][0].shape[-1] max_len_l1 = max(len(seq) for seq in batch_inputs) max_len_l2 = max(subseq.shape[0] for seq in batch_inputs for subseq in seq) batch_size = len(batch_inputs) # 初始化全填充值的输出张量 padded_output = torch.full( (batch_size, max_len_l1, max_len_l2, feat_dim), fill_value=pad_val, dtype=batch_inputs[0][0].dtype ) # 逐位置写入真实数据 for sample_idx, l1_seq in enumerate(batch_inputs): for l1_pos, l2_seq in enumerate(l1_seq): cur_l2_len = l2_seq.shape[0] padded_output[sample_idx, l1_pos, :cur_l2_len, :] = l2_seq
运行后得到的padded_output形状为(2, 3, 3, 3),数值和期望输出完全一致。
如果后续需要配合PyTorch的pack_padded_sequence做变长序列训练,只需要在遍历过程中顺便记录每个样本的一级序列长度、每个一级序列对应的二级序列长度即可,不需要额外调整填充逻辑。
内容的提问来源于stack exchange,提问作者realmarv
相关产品推荐
相关产品推荐

