跳到主要内容

数据截至 (上游 commit cfacd76a0bdd)

04 · 权重同步与三种训练模式

这一章讲什么: RL 训练里最脏的那块活——每训一步,就得把刚更新的权重搬到推理引擎里;训练和推理还要抢同一批 GPU 的显存。以及 verl 怎么用三个 trainer 子类,把 on-policy 到 off-policy 的整个谱系表达出来。


1. 它要解决的小问题

训练侧和推理侧对同一个模型的布局完全不同

训练侧(FSDP/Megatron)推理侧(vLLM/SGLang)
分片方式FSDP 按参数展平切;Megatron 按 TP/PP 切按推理 TP 切
数值精度常是 fp32 主权重bf16/fp16
参数名可能带 _fsdp_wrapped_module. 前缀、LoRA 包装标准 HF 命名
进程数训练 world_size推理 world_size(通常更小)

所以「同步权重」不是一次 copy_,而是一次重分片 + 改名 + 转精度 + 跨进程传输。而且这件事每个训练步都要做一次,慢一点整体吞吐就废了。


2. 思路:把「产出权重」和「搬运权重」拆开

verl 的切法是两个接口:

模型引擎 (ME) 检查点引擎 (CE)
──────────── ──────────────
get_per_tensor_param() send_weights(生成器)
→ 一个 (name, tensor) 生成器 receive_weights()
已经改好名、转好精度、聚合成完整张量 负责怎么传(NCCL/NIXL/共享内存/进程内)

接口是一个生成器而不是一个 dict——这点很关键。671B 模型的完整权重放不进一张卡的显存,逐张量 yield 才能边产边传边释放。

2.1 训练侧怎么产

以 FSDP 为例(verl/workers/engine/fsdp/transformer_impl.py:949get_per_tensor_param):

① 把参数从 CPU load 回 GPU(如果开了 offload)
② 处理 LoRA:
merge=True → merged_lora_context 里取合并后的 state_dict
merge=False → 只收 LoRA 增量参数,改名成 peft 约定
没 LoRA → 直接 state_dict()
③ convert_weight_keys —— 剥掉 FSDP 包装前缀
④ offload_fsdp_model_to_cpu —— 把 GPU 上的训练权重挪回 CPU(腾显存)
⑤ 返回生成器:逐个 DTensor.full_tensor().to(bfloat16)

第 ④ 步在第 ⑤ 步之前——先把 GPU 上的训练权重 offload 回 CPU 腾出显存(:851-854),再懒惰地一张一张聚合。第 ⑤ 步的 full_tensor() 是把 FSDP 的分片 DTensor 聚合成完整张量(:861),这是重分片真正发生的地方:聚合出来的完整张量需要显存,而这块显存正是第 ④ 步腾出来的。

2.2 搬运侧的拓扑

CheckpointEngineManager 的 docstring 里画了张图(verl/checkpoint_engine/base.py:390-402),它说明了两件事:

  • 训练侧:模型引擎和检查点引擎在同一进程里,直接拿到张量。
  • 推理侧:检查点引擎和推理 worker 在不同进程,通过 CUDA IPC 递显存句柄(不复制)。
  • 两侧之间走 NCCL / NIXL / Mooncake 等后端。
训练侧 (N 个进程) 推理侧 (M 个副本 × K 进程)
┌─────┬─────┬─────┐ ┌─────────────────┐
│ ME0 │ ME1 │ MEn │ │ Replica0 │
│ ↓ │ ↓ │ ↓ │ │ r0 r1 r2 r3 │ ← 推理 worker 进程
│ CE │ CE │ CE │ └─┬───┬───┬───┬───┘
└──┬──┴─────┴─────┘ ↑ ↑ ↑ ↑ cuda ipc(递显存句柄)
│ ┌─┴───┴───┴───┴───┐
└───── nccl / nixl / ... ────────►│ CE CE CE CE │
└─────────────────┘
(Replica1..M 同构)

2.3 分块传输

单个张量可能有几个 GB(比如 embedding),一次传会撑爆通信缓冲。split_weight_chunks / merge_weight_chunksverl/checkpoint_engine/base.py:561:546)把张量按 bucket_size 字节切块传、收端再拼回:

# 示意,非源码:切块的核心思路
buffer = weight.view(-1).view(torch.uint8) # 按字节看待,跟 dtype 解耦
while offset < weight.nbytes:
size = min(bucket_size, weight.nbytes - offset)
yield (TensorMeta(name, shape, dtype, offset, size), buffer[offset:offset+size])
offset += size

重点看 view(torch.uint8)——把所有 dtype 统一成字节流,切块逻辑就不用为 bf16/fp8/int8 各写一份。收端靠 TensorMeta 里的 shape/dtype 还原。


3. 一次完整的权重更新

CheckpointEngineManager.update_weightsverl/checkpoint_engine/base.py:506)分两条路。

3.1 naive 路径(共置同步训练)

训练和推理在同一进程,不需要跨进程传输:

if self.backend == "naive":
ray.get(self.actor_wg.update_weights(global_steps=global_steps, mode=self.backend))
return

真正的活在 ActorRolloutRefWorker.update_weightsverl/workers/engine_workers.py:719)里,顺序很讲究:

① 唤醒推理引擎的【权重】显存(KV cache 还没恢复)
② 从训练引擎取 per_tensor_param 生成器
③ rollout.update_weights(...) —— 进程内直接灌
④ 训练模型 offload 到 CPU + aggressive_empty_cache
⑤ 才唤醒推理引擎的【KV cache】显存

为什么要把 weights 和 kv_cache 分两次唤醒? 因为第 ③ 步进行时,显存里同时有训练权重和推理权重,是峰值时刻。KV cache 通常是最大的一块,把它推迟到训练权重释放之后再分配,峰值就低一截。这是很典型的显存峰值削减技巧。

代码里还包了 set_expandable_segments(False/True):706:750)—— 权重同步期间关掉 PyTorch 的可扩展段分配器,避免碎片。

3.2 非 naive 路径(跨进程/跨机)

① abort_replicas() 打断所有在途请求(partial rollout 会记住进度)
② 临时组一个包含所有副本 worker 的 RayWorkerGroup
③ release_kv_cache_replicas() 只放 KV cache,权重缓冲原地留着
④ build_process_group() prepare → build_topology → init_process_group
⑤ 两侧同时 update_weights 训练侧 send,推理侧 receive
⑥ finalize() 拆通信组
⑦ resume_kv_cache_replicas()
⑧ resume_generation_replicas() 被打断的请求继续跑

第 ③ 步的注释解释了它和 sleep_replicas() 的区别(verl/checkpoint_engine/base.py:489-493):只释放 KV cache、保留权重 buffer,这样 NCCL 可以直接写进现有的权重显存,省掉一次分配

第 ④ 步的 build_topologyverl/checkpoint_engine/base.py:153)是各后端自己实现的:给定训练 world_size 和推理 world_size,算出谁跟谁配对、各自的 rank 是多少。NCCL、NIXL、Mooncake、HCCL、KIMI 各有一份实现。


4. 三种训练模式

这是 verl 组织能力的高光:基类 PPOTrainer 把整个 step 编排定死,三个子类各自只覆写几个生命周期钩子

4.1 钩子设计

基类留了 9 个钩子(verl/trainer/ppo/v1/trainer_base.py:590-633):

钩子触发时机行号是否强制覆写
on_init_end初始化结束后:468
on_train_begin训练循环开始前:472
on_train_end训练循环结束后:476
on_validate_begin验证循环开始前:480
on_validate_end验证循环结束后:484
on_step_begin每个训练步开始时:488
on_step_end每个训练步结束时:492@abstractmethod
on_sample_begin从 replay buffer 采样前:497
on_sample_end采样一批之后:501@abstractmethod

九个里只有 on_step_endon_sample_end@abstractmethod——强制子类回答「一步结束时怎么办」和「采完一批样怎么办」这两个问题,因为这正是三种模式唯一的差别。其余七个有默认空实现,子类按需覆写。

4.2 sync:最朴素的 on-policy

class PPOTrainerSync(PPOTrainer):
def on_step_end(self):
self.checkpoint_manager.update_weights(self.global_steps) # 醒过来换新权重
def on_sample_end(self):
self.checkpoint_manager.sleep_replicas() # 睡下去让出显存

verl/trainer/ppo/v1/trainer_sync.py:24-42,全类只有 19 行)

时间线:

生成 ████████████ 生成 ████████████
训练 ██████████ ██████████
└─ sleep ──┘ └─ update ─┘

任一时刻只有一边在用 GPU ← 简单、严格 on-policy、但有气泡

4.3 colocate_async:同一批卡,但生成不停

class PPOTrainerColocateAsync(PPOTrainer):
def get_llm_client(self):
return self.llm_server_manager.get_client(client_cls=FullyAsyncLLMServerClient) # 换成会续跑的客户端
def on_train_begin(self):
for _ in range(num_warmup_batches):
self._add_batch_to_generate() # 预热:先灌几批进去
def on_step_end(self):
self.checkpoint_manager.update_weights(...)
self.checkpoint_manager.resume_generation_replicas() # 恢复被打断的
def on_sample_end(self):
self.checkpoint_manager.abort_replicas() # 打断在途请求
self.checkpoint_manager.sleep_replicas()

verl/trainer/ppo/v1/trainer_colocate_async.py:25-59

与 sync 的三处差别:

  1. 预热批次——训练开始前先塞 N 批 prompt 进生成队列,让 buffer 里始终有存货。
  2. abort 而非等待——要训练了就打断,靠 partial rollout 保住已生成部分。
  3. FullyAsync 客户端——让 abort 对 agent loop 透明。

代价是数据变成轻度 off-policy(一条轨迹可能跨几个版本),需要第 6 章的 rollout correction 来补偿。

4.4 separate_async:训推分离 + 空闲时反串

最复杂的一档(verl/trainer/ppo/v1/trainer_separate_async.py:40)。它同时维护两套推理资源

名字是什么谁在用
standalone_server_manager独立 GPU 上的推理副本,一直在生成agent loop 的主力
llm_server_manager(继承来的)和训练共置的 hybrid 引擎训练侧空闲时临时加进池子帮忙

模式切换靠 HybridEngineMode 枚举(:32):

ROLLOUT 模式 TRAINER 模式
┌──────────────────┐ ┌──────────────────┐
│ hybrid 引擎醒着 │ ──────────► │ hybrid 引擎睡了 │
│ 已加入负载均衡池 │ switch_to_ │ 已从池子里摘除 │
│ 帮着一起生成 │ trainer() │ 显存给训练用 │
└──────────────────┘ └──────────────────┘
▲ │
└────── switch_to_rollout() ───────┘

独立推理副本:全程不停,不受影响

switch_to_trainer 三步(:135):从均衡器摘服务器 → abort → sleep。switch_to_rollout 反过来(:128)。

构造函数里有一串很有信息量的构造期约束:47-68)。注意最后一行不是断言——配置里写错不会报错,会被静默改掉:

约束形式为什么
train_batch_size == ppo_mini_batch_sizeassert:47分离模式下一步只做一次梯度更新
rollout.nnodes > 0n_gpus_per_node > 0assert:51:52独立推理副本必须真的分到资源
checkpoint_engine.backend != "naive"assert:55跨机同步不能用进程内路径
用 RM 时必须 enable_resource_pool=Trueassert:59独立推理副本永不暂停,共置 RM 抢不到显存
rollout_correction.bypass_mode = True构造函数强制覆写,非断言:68你在 yaml 里写 False 不会报错,会被静默改成 True;见 §4.5

诚实标注: should_switch_to_rollout() 目前硬编码返回 False:153-155),TODO 写着「按 replay buffer 状态和切换开销实现策略」。也就是说训练侧空闲反串生成这条路,当前只在验证时会走(on_validate_begin:107)。这是一个已搭好骨架但策略未落地的功能。

4.5 三模式对比

维度synccolocate_asyncseparate_async
训练/推理 GPU共用共用分离 + 可反串
partial rollout
on/off policy严格 on-policy轻度 off明显 off
权重同步频率每步每步parameter_sync_step 步(默认 4)
old_log_prob重算(decoupled)重算强制 bypass(直接用 rollout 的 logprob)
预热批次014
适合稳定复现提吞吐大规模、生成远慢于训练

默认值在 verl/trainer/config/ppo_trainer.yaml:204-224


5. 巧妙之处

第一,把「模式」表达成钩子而不是 if。 如果用配置分支写,step() 里会散落十几个 if async_mode:。改成钩子后,step() 那 10 步编排(源码注释 # 1 ~ # 10verl/trainer/ppo/v1/trainer_base.py:509)对三种模式完全一致,差异全在钩子里、加起来不到 60 行。想加第四种模式(比如 one-step-off-policy),只需要再写一个子类。

第二,用 registry 而不是 if-else 选 trainer。 @register_trainer("sync")verl/trainer/ppo/v1/trainer_base.py:1838)+ get_trainer_cls(name):1597),配置里写字符串就行。

第三,检查点引擎和推理副本之间走 CUDA IPC。 同一节点上,显存句柄可以直接递过去而不做任何拷贝。这让「跨进程」这个架构选择的代价接近于零。

第四,权重同步与 KV cache 释放的顺序被反复调优过。 naive 路径里 weights/kv_cache 分两阶段唤醒、非 naive 路径里只放 KV cache 保留权重 buffer——这两个决定都是纯显存峰值考虑,注释也都写明了理由。


6. 代码地图

主题文件路径符号名
检查点引擎抽象verl/checkpoint_engine/base.pyCheckpointEngineCheckpointEngineRegistryTensorMeta
同步编排verl/checkpoint_engine/base.pyCheckpointEngineManager.update_weightsbuild_process_group
显存生命周期verl/checkpoint_engine/base.pysleep_replicaswake_up_replicasrelease_kv_cache_replicasabort_replicasresume_generation_replicas
分块传输verl/checkpoint_engine/base.pysplit_weight_chunksmerge_weight_chunks
推理侧承接进程verl/checkpoint_engine/base.pyCheckpointEngineWorkerColocatedCheckpointEngine
具体后端verl/checkpoint_engine/nccl_checkpoint_engine.pynixl_checkpoint_engine.pymooncake_checkpoint_engine.pyhccl_checkpoint_engine.pykimi_checkpoint_engine.py
naive 路径verl/workers/engine_workers.pyActorRolloutRefWorker.update_weights
训练侧产权重verl/workers/engine/fsdp/transformer_impl.pyFSDPEngine.get_per_tensor_param
trainer 基类与钩子verl/trainer/ppo/v1/trainer_base.pyPPOTrainer.stepon_step_endon_sample_endregister_trainerget_trainer_cls
同步模式verl/trainer/ppo/v1/trainer_sync.pyPPOTrainerSync
共置异步verl/trainer/ppo/v1/trainer_colocate_async.pyPPOTrainerColocateAsync
分离异步verl/trainer/ppo/v1/trainer_separate_async.pyPPOTrainerSeparateAsyncHybridEngineModeswitch_to_trainerswitch_to_rollout