数据截至 (上游 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:949,get_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_chunks(verl/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 还原。