跳转至

padding.py — 填充格式转换工具

文件概述

提供两个核心函数,在"有填充"(padded)和"无填充"(no-padding / remove-padding)两种数据格式之间进行转换。Remove Padding 是一种重要的性能优化,避免在 padding token 上浪费计算。

背景知识:两种数据格式

有填充格式(Padded):
  样本1: [tok1, tok2, tok3, PAD,  PAD ]  shape: (bsz, max_seq_len)
  样本2: [tok1, tok2, tok3, tok4, tok5]
  样本3: [tok1, tok2, PAD,  PAD,  PAD ]

无填充格式(No-Padding / Remove-Padding):
  展平: [tok1, tok2, tok3, tok1, tok2, tok3, tok4, tok5, tok1, tok2]
  长度: [3, 5, 2]                            shape: (total_nnz,)

  用 Nested Tensor 表示: nested_tensor([
    [tok1, tok2, tok3],
    [tok1, tok2, tok3, tok4, tok5],
    [tok1, tok2]
  ])

无填充格式的优势:不在 PAD token 上做无效计算,提升 GPU 利用率。

核心函数

1. left_right_2_no_padding - 有填充 -> 无填充

将标准的 padded TensorDict 转换为 nested tensor 格式,用于模型前向计算。

def left_right_2_no_padding(data: TensorDict) -> TensorDict:
    """
    输入: TensorDict 包含 padded 的 input_ids, attention_mask, response_mask, position_ids
    输出: TensorDict 包含 nested tensor 的 input_ids, position_ids, loss_mask
    """
    input_ids = data.pop("input_ids")
    attention_mask = data["attention_mask"]

    # 使用 unpad_input 去除 padding,获得展平的 token 和索引
    input_ids_rmpad, indices, cu_seqlens, *_ = unpad_input(
        input_ids.unsqueeze(-1), attention_mask
    )

    # 构建 nested tensor(PyTorch 的变长序列表示)
    input_ids_nested = torch.nested.nested_tensor_from_jagged(
        input_ids_rmpad.squeeze(-1), offsets=cu_seqlens
    )

    # 对 position_ids 做同样的处理
    position_ids_list = []
    for i in range(attention_mask.shape[0]):
        curr_mask = attention_mask[i].bool()
        curr_pos_ids = position_ids[i]
        if curr_pos_ids.dim() == 1:    # 标准: (seq_len,)
            valid_ids = curr_pos_ids[curr_mask]
        else:                           # 多头: (4, seq_len)
            valid_ids = curr_pos_ids[:, curr_mask]
        position_ids_list.append(valid_ids)
    position_ids_nested = torch.nested.as_nested_tensor(position_ids_list, layout=torch.jagged)

    data["input_ids"] = input_ids_nested
    data["position_ids"] = position_ids_nested
    data["loss_mask"] = data["response_mask"]

    # 处理 MoE Router Replay 的 routed_experts(如果存在)
    routed_experts = data.get("routed_experts", None)
    if routed_experts is not None and not routed_experts.is_nested:
        routed_experts_rmpad = index_first_axis(
            routed_experts.unsqueeze(-1).flatten(0, 1), indices
        )
        data["routed_experts"] = torch.nested.nested_tensor_from_jagged(
            routed_experts_rmpad.squeeze(-1), offsets=cu_seqlens
        )

    return data

关键步骤: 1. unpad_input: Flash Attention 的工具函数,去除 padding 并返回展平的 token 序列和累积长度 2. nested_tensor_from_jagged: 用累积长度信息构建 PyTorch nested tensor 3. 对 position_ids 逐样本过滤出有效位置(支持 1D 和多头 4D 两种格式) 4. 如果有 routed_experts(MoE 路由信息),也做相同的去 padding 处理

2. no_padding_2_padding - 无填充 -> 有填充(仅 response 部分)

从模型的无填充输出中提取 response 部分,恢复为标准的 padded 格式。

def no_padding_2_padding(tensor: torch.Tensor, data: TensorDict) -> torch.Tensor:
    """
    输入: nested tensor 或展平的 1D tensor (total_nnz,)
    输出: padded response tensor (bsz, max_response_len)
    """
    values = tensor.values() if tensor.is_nested else tensor
    prompt_ids = data["prompts"]
    response_ids = data["responses"]

    if prompt_ids.is_nested:
        prompt_lens = prompt_ids.offsets().diff()
        response_lens = response_ids.offsets().diff()
    else:
        prompt_lens = attention_mask[:, :prompt_ids.shape[1]].sum(dim=1)
        response_lens = attention_mask[:, prompt_ids.shape[1]:].sum(dim=1)

    sequence_lens = prompt_lens + response_lens
    sequence_offsets = sequence_lens.cumsum(dim=0)

    response_list = []
    for resp_len, seq_offset in zip(response_lens, sequence_offsets):
        pad_size = max_response_len - resp_len
        # 左移一个 token:log_prob[t] 对应预测 token[t+1]
        response_list.append(
            F.pad(values[seq_offset - resp_len - 1 : seq_offset - 1], (0, pad_size))
        )

    output = torch.stack(response_list, dim=0)
    return output

关键细节: - 只提取 response 部分(不需要 prompt 部分的 log_prob) - 左移一个 token:因为模型在位置 t 输出的 log_prob 是对位置 t+1 的预测 - 填充到 max_response_len 长度,使所有样本对齐

数据流示意

训练前向传播的数据流:

┌──────────────────────────────────────────────┐
│ 原始数据(Padded 格式)                        │
│ input_ids:      (bsz, max_seq_len)           │
│ attention_mask:  (bsz, max_seq_len)          │
│ position_ids:    (bsz, max_seq_len)          │
└────────────────┬─────────────────────────────┘
                 │
    left_right_2_no_padding()
                 │
                 ▼
┌──────────────────────────────────────────────┐
│ 无填充数据(Nested Tensor 格式)               │
│ input_ids:      nested (bsz, [j1])           │
│ position_ids:   nested (bsz, [j1])           │
└────────────────┬─────────────────────────────┘
                 │
          模型前向计算(高效,无 PAD 浪费)
                 │
                 ▼
┌──────────────────────────────────────────────┐
│ 模型输出(Nested Tensor 格式)                 │
│ log_probs:     nested (bsz, [j1])            │
│ values:        nested (bsz, [j1])            │
└────────────────┬─────────────────────────────┘
                 │
    no_padding_2_padding()
                 │
                 ▼
┌──────────────────────────────────────────────┐
│ Response 部分(Padded 格式)                   │
│ log_probs:     (bsz, max_response_len)       │
│ → 用于 PPO 损失计算                            │
└──────────────────────────────────────────────┘

与其他模块的关系

  • 被 losses.py 中的 ppo_loss 调用 no_padding_2_padding
  • 被各引擎(FSDP、Megatron、TorchTitan 等)的前向计算调用 left_right_2_no_padding
  • 依赖 verl/utils/attention_utils.py 中的 unpad_input 和 index_first_axis
  • 依赖 PyTorch 的 Nested Tensor API

小结

本文件实现了 verl 中 Remove Padding 优化的关键格式转换。通过在模型计算前去除 padding、计算后恢复 padding,既保证了计算效率(不浪费算力在 PAD token 上),又保证了损失计算的正确性。