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 上),又保证了损失计算的正确性。