You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.03 09:12:29