跳到主要内容

数据截至 (上游 commit cfacd76a0bdd)

05 · 训练引擎与批次编排

这一章讲什么: 一个 batch 从 driver 发出去之后,在 GPU 上到底经历了什么——怎么被均衡地摊到各张卡、怎么被切成 micro-batch、loss 怎么归一化才能跨并行度一致。以及 verl 怎么用一个注册表容纳六种训练后端。


1. 后端抽象:EngineRegistry

1.1 它要解决的小问题

verl 要同时支持 FSDP、FSDP2、Megatron、VeOmni、TorchTitan、Automodel、MindSpeed(其中 FSDP 与 FSDP2 共用同一份实现,verl/workers/engine/ 下是六个后端目录),还要区分「语言模型(出 logits)」和「价值模型(出标量)」,还要区分 CUDA / NPU、区分 NVIDIA / 沐曦等厂商。

如果用 if-else 选类,会长成一棵谁也不想维护的树。

1.2 解法:四元组查表 + 三级回退

EngineRegistry.registerverl/workers/engine/base.py:351)的装饰器签名是:

@EngineRegistry.register(model_type="language_model", backend=["fsdp", "fsdp2"], device=["cuda", "npu"])
class FSDPEngineWithLMHead(FSDPEngine): ...

verl/workers/engine/fsdp/transformer_impl.py:1115

参数都支持传列表,一次注册多个组合。查表时(get_engine_cls:329)走三级回退:

想要 (model_type, backend),运行时探测到 device=cuda, vendor=metax

├─① 精确匹配 (cuda, metax) ── 命中就用

├─② 只按 device 匹配 "cuda" ── 命中就用(没写 vendor 的通用注册)

└─③ device 是 cuda 但 vendor 不是 nvidia?
试 (cuda, "nvidia") ── 大多数 CUDA 兼容卡能直接跑 NVIDIA 那份

还留了 VERL_ENGINE_DEVICE / VERL_ENGINE_VENDOR 两个环境变量强制覆盖(:335-338)——调试异构硬件时很有用。

1.3 BaseEngine 的接口面

BaseEngineverl/workers/engine/base.py:30)暴露的方法可以分五组:

方法说明
生命周期initializeto(device)train_mode()eval_mode()train_mode/eval_mode 返回上下文管理器
计算forward_backward_batchtrain_batchinfer_batch核心计算入口
优化器optimizer_zero_gradoptimizer_steplr_scheduler_step
拓扑自述get_data_parallel_size/rank/groupis_mp_src_rank_with_outputs供第 1 章的 mesh 分发用
权重导出get_per_tensor_param供第 4 章的权重同步用

train_mode() / eval_mode() 返回的 BaseEngineCtx:230)在 __enter__ 时把模型搬上 GPU、__exit__ 时搬回 CPU(_context_switch:244)——显存 offload 被藏在 with 语句里,调用方不用手写 load/offload 配对。

1.4 引擎导入也是懒的

verl/workers/engine/__init__.py 里每个可选后端都包在 try/except ImportError 里(:24-65),失败就把符号设成 None。所以只装了 FSDP 的环境不会因为没有 Megatron 而 import 失败。


2. 序列长度均衡:让各卡干一样多的活

2.1 问题

RL 生成出来的回答长度差异极大——同一批里可能有 50 token 的也有 4000 token 的。如果按顺序均分给 8 个 DP rank,某张卡可能拿到一堆长序列,其他卡等它。

注意力的计算量还不是线性的,是 O(L) + O(L²)

2.2 第一步:把长度换算成工作量

calculate_workloadverl/utils/seqlen_balancing.py:27):

return 24576 * seqlen_list + seqlen_list**2

docstring 说明了来历:transformer 前反向 FLOPs ≈ 12·h²·L + 2·h·L²,按 7B 模型 h=4096 代入约掉常数,得到 24576·L + L²

为什么不直接用长度: 一条 4000 token 的序列,二次项贡献占比远高于一条 500 token 的。按长度均分会系统性低估长序列的开销。

2.3 第二步:Karmarkar-Karp 数字划分

把 N 个工作量分成 K 组、让各组和尽量接近——这是经典的 NP-hard 多路数字划分问题。verl 用 Karmarkar-Karp 最大差分法(verl/utils/seqlen_balancing.py:49)近似求解。

直觉:每次取出「内部差距最大」的两个状态合并,合并时大配小——大的和小的凑一起,差距就被抵消了。

初始:把元素装进若干个 State,压进一个按 spread(组内最大-最小)排序的堆
—— spread 最大的先弹出(State.__lt__ 把比较反了过来,:121-127)
equal_size=False → 每个元素单独一个 State
equal_size=True → 先整体排序,每连续 k 个元素做一个 State

循环:弹出 spread 最大的两个 State,merge 它们
merge 时把 A 的第 i 大配 B 的第 (k-1-i) 大 ← 关键:反序配对
合并结果重新入堆

结束:堆里只剩一个 State,它的 k 个 Set 就是 k 个分区

反序配对那行是 self.sets[i].merge(other.sets[self.k - 1 - i]):114-115)——最重的配最轻的,这就是「差分法」三个字的含义。

equal_size=True 保证每组元素数相同,靠的不是比较函数,而是播种方式:初始就让每个 State 恰好装 k 个元素(:144-151),而 merge 是「k 个 Set 一一对应地合并」,所以各 Set 的元素数从头到尾同步增长;函数末尾还有一条断言复查这件事(:165-169)。DP 各 rank 必须拿到相同条数的样本,所以 driver 侧调用时固定传 equal_size=Trueverl/trainer/ppo/v1/trainer_base.py:1470)。

2.4 driver 上怎么用

PPOTrainer._balance_batchverl/trainer/ppo/v1/trainer_base.py:1453):

① 问 actor worker group:actor mesh 的 dp_size 是多少(第 1 章的懒查询)
② 算出 batch 必须是多少的倍数(见 §2.5),不够就 upsample 补齐
③ 从 tag 里取 seq_len(不读数据!),换算成 workload
④ Karmarkar-Karp 划成 dp_size 组
⑤ batch.reorder(打平后的索引) ← 只是重排 key 顺序

第 ③ 步是 TransferQueue 设计的红利:均衡决策所需的全部信息都在 tag 里,driver 一个 token 都不用读。

第 ⑤ 步只重排 key 列表——后续 dispatch 函数按顺序均分时,自然就分到了均衡的组。均衡这件事被彻底降维成「排序」

源码里有条重要警告(V0 版本 verl/trainer/ppo/ray_trainer.py:1521-1523):重排会改变数据顺序,这不影响优势计算(按 uid 分组),但会影响 mini-batch 的切分,进而影响 loss。

2.5 batch 大小的最小公倍数

_get_required_batch_multipleverl/trainer/ppo/v1/trainer_base.py:1434)算的是 batch 必须是谁的倍数:

required = dp_size
required = lcm(required, critic_mini_batch × rollout.n) # 若训 critic
required = lcm(required, actor_mini_batch × rollout.n) # 若过了 critic warmup

注释还点了句 lcm(a,b,c) == lcm(lcm(a,b),c),所以这是最优的。凑不齐就 upsample_batch_to_divisible_size 补占位样本(tag 里标 is_padding=True,算指标时会被 non_padding_mask 过滤掉,:1464)。


3. micro-batch:动态还是固定

prepare_micro_batchesverl/workers/engine/utils.py:92)两条路:

模式切法配置
固定条数每个 micro-batch 固定 N 条样本use_dynamic_bsz=False + micro_batch_size_per_gpu
动态 token 数每个 micro-batch 总 token 数 ≤ 上限,条数可变use_dynamic_bsz=True + max_token_len_per_gpu

动态模式为什么更好: 显存占用主要由 token 数决定,不由条数决定。固定条数时,为了兜住最坏情况(全是长序列)你必须把条数设得很小,短序列那批就浪费显存。动态模式让「全是短序列」的 micro-batch 塞进更多条。

注意 max_token_len = max_token_len_per_gpu * sp_size:76)—— 开了序列并行后,每卡实际承担的是 1/sp_size,所以上限要乘回去。

same_micro_num_in_dp=Trueverl/workers/engine/fsdp/transformer_impl.py:714)强制各 DP rank 的 micro-batch 数量一致。这不是为了均衡,是为了不死锁:FSDP 每个 micro-batch 都有集合通信,某个 rank 少跑一轮就会挂在 barrier 上。

动态切分会打乱顺序,所以 postprocess_batch_funcverl/workers/engine/utils.py:133)在收尾时用 restore_dynamic_batch(model_output[key], indices) 把结果还原回原顺序。


4. loss 归一化:跨并行度不变

4.1 问题

把同一个训练任务从 8 卡换到 16 卡,loss 数值应该完全一样——否则超参不可迁移、实验不可复现。

但天然不一样:每张卡只看到 1/dp_size 的数据,本地 mean 出来的值和全局 mean 不同。

4.2 解法:显式传全局分母

agg_lossverl/trainer/ppo/core_algos.py:1140)不做本地 mean,而是要求调用方传入全局统计量:

if loss_agg_mode == "token-mean":
loss = masked_sum(loss_mat, loss_mask) / batch_num_tokens * dp_size

拆开看:本地求 ÷ 全局 token 数 × dp_size。乘 dp_size 是因为 FSDP 的梯度 all-reduce 用的是 AVG,会再除一次 dp_size,两下抵消 —— 最终梯度恰好等于「全局 token 平均」。

而且它在 dp_size > 1 且没传全局分母时直接抛异常:1171),不给静默算错的机会。

全局 token 数哪来的?forward_backward_batch 开头做了一次 all_reduce(verl/workers/engine/fsdp/transformer_impl.py:706-710):

batch_num_tokens = data["loss_mask"].sum().to(get_device_id())
torch.distributed.all_reduce(batch_num_tokens, op=ReduceOp.SUM, group=self.get_data_parallel_group())

4.3 四种聚合模式

模式公式(概念)效果
token-meanΣ所有token loss / 全局token数长序列权重更大(token 平权)
seq-mean-token-sumΣ(每条序列的token和) / 全局序列数序列平权,但长序列 loss 绝对值大
seq-mean-token-meanΣ(每条序列的token均值) / 全局序列数序列平权,长短同权
seq-mean-token-sum-norm同上再除固定 loss_scale_factorDr.GRPO 用,避免长度偏置

最后一种的 loss_scale_factor 参数有个细节(:1183-1186):不传的话会用 loss_mask.shape[-1],但那个值随 batch 变化。docstring 建议设成常数以保证整个训练过程归一化一致——这正是 Dr.GRPO 论文指出的长度偏置来源。


5. loss 是怎么被装配起来的

注意 verl 的一个设计:损失函数不写死在引擎里,而是从 driver 传进去的

driver: critic_wg.set_loss_fn(partial(value_loss, config=critic_cfg))
(verl/trainer/ppo/v1/trainer_base.py:227-228)


worker: self.loss_fn = loss_fn (engine_workers.py:161)


engine: forward_backward_batch(data, loss_function=self.loss_fn)


每个 micro-batch: loss, meta = forward_step(mb, loss_function, forward_only)
loss.backward()

ppo_lossverl/workers/utils/losses.py:57)就是被传进去的那个函数,它做三件事:

① 把全局统计量塞进 config.global_batch_info(dp_size / batch_num_tokens / global_batch_size / loss_scale_factor)
② 按 config.policy_loss.loss_mode 从注册表取策略损失函数,算 pg_loss
③ 叠加可选项:
- 熵正则:policy_loss -= entropy_coeff * entropy_loss
- KL 正则:policy_loss += kl_loss_coef * KL(π_θ ‖ π_ref)

还有个容易忽视但很关键的指标处理(:73-81):如果做了全局归一化,各 rank 的指标聚合方式必须是 SUM 而不是 MEAN——因为每个 rank 的值已经是「全局的一部分」,再取平均就少算了 dp_size 倍。这类细节写错了不会崩,只会让 wandb 上的曲线悄悄错一个常数因子。


6. 序列并行:Ulysses

超长序列单卡放不下时,verl 用 DeepSpeed-Ulysses 风格的序列并行(verl/utils/ulysses.py)。

核心是两次 all-to-all:

按序列切 按注意力头切
输入 ────────────────► 各卡拿 1/sp 的 token

│ gather_seq_scatter_heads (all-to-all)

各卡拿全部 token 的 1/sp 个头 ← 注意力在这里算

│ gather_heads_scatter_seq (all-to-all)

变回「1/sp token、全部头」 ← FFN 在这里算

为什么要来回换: 注意力需要看到完整序列(否则算不了跨位置的 attention),FFN 是逐 token 的(切序列最省)。两次 all-to-all 就是在两种最优切法之间转换。

SeqAllToAll:169)是自定义 autograd Function,反向就是正向的镜像操作。ulysses_pad_and_slice_inputs:302)负责把序列补到 sp_size 的倍数再切。


7. 代码地图

主题文件路径符号名
引擎接口verl/workers/engine/base.pyBaseEngineBaseEngineCtx
后端注册表与回退verl/workers/engine/base.pyEngineRegistry.registerget_engine_clsnew
FSDP 实现verl/workers/engine/fsdp/transformer_impl.pyFSDPEngineFSDPEngineWithLMHeadFSDPEngineWithValueHead
前反向主循环verl/workers/engine/fsdp/transformer_impl.pyFSDPEngine.forward_backward_batch
其他后端verl/workers/engine/megatron/veomni/torchtitan/automodel/mindspeed/
micro-batch 切分/还原verl/workers/engine/utils.pyprepare_micro_batchespostprocess_batch_func
工作量估计verl/utils/seqlen_balancing.pycalculate_workload
数字划分verl/utils/seqlen_balancing.pykarmarkar_karpget_seqlen_balanced_partitionslog_seqlen_unbalance
driver 侧均衡verl/trainer/ppo/v1/trainer_base.py_balance_batch_get_required_batch_multiple
占位补齐verl/trainer/ppo/padding_utils.pyupsample_batch_to_divisible_size
loss 归一化verl/trainer/ppo/core_algos.pyagg_loss
loss 装配verl/workers/utils/losses.pyppo_lossvalue_losssft_loss
worker 层 APIverl/workers/engine_workers.pyTrainingWorker.train_mini_batchtrain_batchinfer_batch_postprocess_output
序列并行verl/utils/ulysses.pySeqAllToAllgather_seq_scatter_headsulysses_pad_and_slice_inputs