同一批短句为何被最长句拖慢?动态 Padding、长度分桶与 Token-based Batching
从变长样本的二次注意力浪费出发,手算长度分桶收益,实现当前 PyTorch 动态 padding 与 token 预算批处理,并解释 loss 归一化、分布式与编译形状陷阱。
上一篇把多篇短文装入同一行,并用文档边界阻止交叉污染。但监督微调、分类与线上请求常要求“一行对应一个独立样本”,不能随意拼接。此时固定补到全局最大长度仍然浪费,而每批只补到本批最长长度又可能被一个异常长样本拖慢。
本篇聚焦一个核心问题:怎样按长度组织样本,使每个 batch 的物理 token 接近有效 token,同时保持随机性、样本权重和梯度尺度可解释。
01 动态 Padding 为什么仍可能很浪费?#
对一个 batch 的长度 ,右侧动态补齐到 。token 利用率为:
Self-Attention 的粗略工作量却是 ,因为补出的 query/key 仍占据物理张量。mask 能阻止 PAD 影响语义,不保证内核完全跳过这些位置。
flowchart LR
A[样本与真实长度] --> B[打乱]
B --> C[长度桶]
C --> D[token 预算组 batch]
D --> E[本批动态 padding]
E --> F[input_ids B×T]
E --> G[attention_mask B×T]
E --> H[labels B×T]
F --> I[模型]
G --> I
I --> J[按有效 token 归一化 loss]
H --> Jmermaid02 四个样本怎样手算分桶收益?#
长度为 [8, 7, 2, 1],每批 2 个。随机配对 [8,1]、[7,2]:
按相近长度配对 [8,7]、[2,1]:
用 粗估 attention 单元,随机配对为 ,分桶为 。有效 token 完全相同,但物理计算减少约 40%。
03 分桶不是把数据永久排序#
若每个 epoch 都从短到长,模型会先连续看到短样本、后连续看到长样本;长度若与类别、语言或难度相关,训练顺序就变成隐式课程。
稳妥流程是:
全局样本索引
└─按 epoch seed 打乱
└─切成较大的局部池(例如 1,000 条)
└─池内按长度排序
└─组成 batch,再打乱 batch 顺序text局部池越大,padding 越少但全局随机性越弱、等待时间越长。应记录池大小与 seed,并按来源/标签检查每批分布,而不是只看吞吐。
04 固定样本数为何不是固定工作量?#
batch size 固定为 32 时,32 条 64-token 文本与 32 条 4,096-token 文本相差 64 倍 token,attention 粗略成本相差更多。Token-based Batching(按 token 预算组批)限制:
是每批物理 token 上限。长样本自动减少行数,短样本增加行数。更精细的调度器可约束 ,但 更简单,也更接近激活内存预算。
05 一个可检查的 Token Batch Sampler#
下面输入的是已按局部长度桶组织的 (index, length)。加入新样本后,若 新批大小 × 新最大长度 超过预算,就先产出旧批。
def batches_by_padded_tokens(index_lengths, max_tokens, max_examples=None):
batch, longest = [], 0
for index, length in index_lengths:
if length <= 0 or length > max_tokens:
raise ValueError(f"invalid length {length} for sample {index}")
next_longest = max(longest, length)
next_size = len(batch) + 1
over_tokens = next_size * next_longest > max_tokens
over_examples = max_examples is not None and next_size > max_examples
if batch and (over_tokens or over_examples):
yield batch
batch, longest = [], 0
batch.append(index)
longest = max(longest, length)
if batch:
yield batchpython输入长度 [8,7,2,1]、max_tokens=16 时,相近长度顺序会得到 [8,7] 与 [2,1]。应另设 max_examples,避免极短样本一次堆入数千行,耗尽 CPU 元数据或改变归一化层行为。
06 当前 PyTorch 怎样动态补齐?#
PyTorch 2.14 的 torch.nn.utils.rnn.pad_sequence 接收一组形如 [L_i,*] 的张量;batch_first=True 输出 [B,T,*],当前 API 还显式支持 padding_side='right'|'left'。
import torch
import torch.nn.functional as F
from torch.nn.utils.rnn import pad_sequence
PAD_ID = 0
IGNORE = -100
def collate_causal_lm(examples):
# examples: list[LongTensor[Li]],每条包含 EOS
input_ids = pad_sequence(
examples,
batch_first=True,
padding_value=PAD_ID,
padding_side="right",
) # [B,T]
lengths = torch.tensor([len(x) for x in examples]) # [B]
positions = torch.arange(input_ids.size(1))[None, :] # [1,T]
attention_mask = positions < lengths[:, None] # [B,T]
labels = input_ids[:, 1:].clone() # [B,T-1]
labels[~attention_mask[:, 1:]] = IGNORE
model_inputs = input_ids[:, :-1] # [B,T-1]
model_mask = attention_mask[:, :-1] # [B,T-1]
return model_inputs, model_mask, labels, lengthspython输入 token 是整数,因此 padding_value 虽接受浮点参数,也要传合法的词表 id。attention_mask 的 True/False 语义最终要按模型接口转换;不能假设所有库的布尔 attention mask 都同义。
07 右 Padding 与左 Padding 何时使用?#
训练因果 LM 时常用右 padding:每行真实 token 都从位置 0 开始,标签右移直观。批量自回归生成常用左 padding,使所有行的最后一个真实 token 对齐到同一列,便于取 logits[:, -1]。
right: [A B C EOS PAD] [D E EOS PAD PAD]
left: [A B C EOS] [PAD D E EOS]text左 padding 时不能直接把物理列号当位置 id。一个常见构造是:
position_ids = attention_mask.long().cumsum(dim=-1) - 1
position_ids.masked_fill_(~attention_mask, 0)python这样两行的第一个真实 token 都是位置 0。若模型用 RoPE 与 KV Cache,prefill 的 position id、cache 长度和后续 decode 位置必须遵守同一契约。
08 Loss 到底按 token 还是按样本平均?#
按 token 平均:
长样本贡献更多目标,适合把语料视为 token 流。按样本平均则先计算每行平均,再平均 行,使短样本与长样本权重相同。两者都合理,但目标不同。
token_loss = F.cross_entropy(
logits.transpose(1, 2), labels,
ignore_index=IGNORE, reduction="none",
) # [B,T-1]
valid = labels.ne(IGNORE)
token_mean = token_loss.sum() / valid.sum().clamp_min(1)
per_example = token_loss.sum(1) / valid.sum(1).clamp_min(1)
example_mean = per_example.mean()python当前 PyTorch 2.14 的 cross_entropy(ignore_index=...) 只在目标是类别索引时忽略该值。不要把 PAD id 同时当 ignore_index,因为 PAD id 可能是模型应预测的合法类别;使用词表外的 -100 更清楚。
09 变 batch size 后,梯度尺度怎样保持?#
若每步 loss.mean().backward(),长批与短批先各自变成一个均值,再做梯度累积,会让不同 step 获得相同权重,而不是每个 token 相同权重。
精确的 token 归一化应在一个优化窗口内累积 loss sum 与有效 token 总数。单进程可先对每个 micro-batch 的 loss_sum 反传,更新前把梯度除以窗口总 token 数;分布式还要 all-reduce 分母,并考虑 DDP 默认的梯度平均因子。
micro-batch 1: loss_sum=120, valid=80
micro-batch 2: loss_sum=30, valid=20
window loss = (120+30)/(80+20)=1.5text它不等于简单平均两个 micro-batch mean;当分母和平均 loss 不同时,差异会立刻出现。
10 分布式采样怎样避免重叠与失衡?#
错误做法是每个 rank 独立打乱全量索引再分桶:不同 rank 会抽到重复样本,step 数也可能不同而死锁。更稳妥的顺序是先由全局 epoch seed 产生确定索引,再按 rank 分片,各 rank 在自己的分片内建立长度桶;或者由统一 batch plan 分发各 rank 的 micro-batch。
需要验证:
- 同 epoch 的全局样本 id 是否恰好覆盖一次;
- 各 rank 是否产生相同步数;
- 每步最大 是否严重不均,导致快卡等待慢卡;
- 恢复训练后,epoch、seed、bucket cursor 与 batch plan 是否一致。
为保证相同步数而复制尾部样本时,必须记录重复并在统计权重中说明。
11 动态形状为何可能伤害编译性能?#
动态 padding 让每批 改变。GPU kernel、torch.compile 或图捕获可能为许多形状反复编译,节省的 FLOPs 被编译和调度开销抵消。
实用折中是把 向上取到少数边界,如 {128,256,512,1024},或取某个 tile 的倍数。此时略增 padding,却提高形状复用、内存规划稳定性和 Tensor Core 对齐。不要只比较 tokens/s;同时记录首次编译时间、steady-state 吞吐、峰值显存与重编译次数。
12 性能指标不能只报“每秒多少 batch”#
变长 batch 的行数不同,batches/s 会误导。至少同时报告:
| 指标 | 回答的问题 |
|---|---|
| 有效 tokens/s | 模型真正学习多少目标 |
| 物理 tokens/s | kernel 处理多少张量位置 |
| token 利用率 | padding 占比多大 |
| step latency 分位数 | 长尾 batch 是否卡顿 |
| 峰值显存 | token 上限是否安全 |
| 每来源/标签占比 | 分桶是否改变数据分布 |
端到端 profile 要包含 tokenizer、DataLoader、host-to-device copy 与模型;GPU 变快后,CPU 长度排序可能成为新瓶颈。
13 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| loss 低得异常 | PAD 标签未设 -100 | 数有效标签并查看末列 |
| 左 padding 生成错位 | position id 用物理列号 | 打印每行首个真实位置 |
| OOM 偶发 | 只限制样本数 | 记录每批 B,T,B*T |
| 吞吐未提升 | 形状过多导致重编译 | 统计唯一 与编译次数 |
| 指标偏向短文本 | 使用按样本平均 | 同时报 token/sample mean |
| 多卡偶发卡住 | rank step 数不同 | 启动前比较 batch-plan 长度 |
| 类别顺序成团 | 全局按长度永久排序 | 检查每批标签与来源直方图 |
最小测试集应包含长度 1、恰好等于上限、超过上限、全 PAD 非法输入和极端长尾;并固定 seed 比较断点恢复后的前 20 个 batch id。
14 失败场景与相近方法#
长度分桶只能减少同批长度差,无法消除每行尾部 padding;Sequence Packing 能继续提高利用率,但需要文档边界语义。梯度累积增加有效 batch token,不会减少单个 micro-batch padding。动态批处理改变 ,不等同于动态序列长度训练;后者可能特意改变上下文分布。
当绝大多数样本同长,分桶收益很小;当严格在线到达、延迟优先时,等待同长度请求会增加排队时间。推理服务必须在吞吐与尾延迟之间设最大等待时间,不能照搬离线训练策略。
15 今天真正需要记住什么?#
- 动态 padding 只补到本批最大长度;长度分桶进一步缩小批内差异。
- token-based batching 用 约束物理预算,使长样本自动减少行数。
- 可变 batch size 会暴露按 token/按样本 loss 的选择,也会影响梯度累积与多卡同步。
- 最佳形状不一定最紧凑;少量离散长度常能在 padding 与编译/kernel 复用间取得更好平衡。
16 思考题与小练习#
- 长度
[12,11,7,6,3,2]、max_tokens=24,分别按原顺序与降序运行 sampler,计算每批 和总体利用率。 - 实现按样本平均的因果 LM loss,并构造一条 2-token 与一条 8-token 样本,比较它和按 token 平均的权重。
- 设计四个离散边界
{128,256,512,1024}的 benchmark,说明怎样区分 padding 收益、重编译成本与 DataLoader 瓶颈。
相关工作#
- Kundu et al., Smart Batching: Fast Fine-Tuning of Transformer Language Models ↗,研究长度感知的 Transformer 批处理。
- Krell et al., Efficient Sequence Packing without Cross-contamination ↗,比较装箱与 padding 的效率边界。
- Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness ↗,从 IO 解释 attention 实际性能。
- You et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism ↗,展示大规模语言模型训练中的并行与批处理工程。
- Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM ↗,讨论吞吐、micro-batch 与并行调度。
17 下一篇预告#
数据已经去重、混合并高效组批,接下来要决定训练究竟持续多久。下一篇将研究 token 学习率日程、warmup、cosine decay 与按 step/按 token 计时的差异。