RL / OPD 工程 — 后训练的发动机
Ch18 把"OPD 为什么可行"讲完了;这一章只回答"它怎么跑得起"。trillion 级 teacher × 全词表 logit × 1M context × on-policy rollout 这四件事任一单独都能压垮集群,V4 用四个工程支柱(FP4 / Teacher Scheduling / WAL Rollout / Million-Token RL)把它们同时跑起。
四个支柱 = (1) FP4 在 rollout/teacher/reference forward 全量启用;(2) Teacher Scheduling 把 N 个 trillion teacher 拆到中央存储 + 中央 hidden 缓冲 + GPU 上动态装载的 prediction head;(3) Token 粒度 WAL 让 rollout 抢占可恢复且无 length bias;(4) Million-Token RL 把 1M trajectory 切轻量元数据 + 重 per-token 字段
一句话:把"原理上可行"翻译成"工程上能跑"。 OPD 的物理障碍是显存(trillion teacher)+ 带宽(全词表 logit)+ 容错(on-policy rollout 抢占)+ 长度(1M context) 四个独立瓶颈。每个支柱单独解决一个,缺一不可。 这一章介绍全词表、多 teacher OPD 所需的系统实现,包括低精度 forward、按需 teacher 调度、可抢占 rollout 与长序列数据加载。
- FP4 全量启用(在 rollout / teacher / reference forward)
- Ch10 的 FP4 QAT 让训练时的 forward能用 FP4 权重 + 无损 dequant。OPD 阶段把这件事推到极致:所有 inference-only 路径(rollout、N 个 teacher 前向、reference 模型前向)全部走原生 FP4,只在 backward 路径维持 FP8。动机:teacher 与 reference 永远不更新,没必要 FP8 精度,FP4 砍内存 + 砍带宽。
- Teacher Scheduling(教师调度)
- OPD 框架要支持十余个大规模 teacher。把所有 teacher 同时常驻 GPU 会带来很高的显存需求。V4 将 teacher 状态分到中央存储、中央 hidden-state buffer 与按需加载的 prediction head,并按 teacher index 组织 mini-batch,使每个 head 在该 mini-batch 中只加载一次。
- Centralized Weight Offload(权重托管)
- 所有 teacher 权重放在中央分布式存储中,以 ZeRO-like sharding 按需加载,降低 I/O 与 DRAM 压力。报告没有给出每张卡约 100GB 的固定数字。
- 不实例化 logits(never materialize)
- 关键工程 trick。$|V| > 100\text{K}$ 全词表 logits × batch × seq 体积比模型本身还大(见 Ch18 数值演练 2.6 TB/step)。V4 在 teacher forward 期间只缓存最后一层 hidden state到中央 buffer;训练时再把它过对应 teacher 的 prediction head 现场重建 logits + 直接算 KL,logits 张量从来不在显存上完整出现。
- 按 Teacher Index 排序的 batching
- 训练样本在 dispatch 阶段按 teacher index 排序,使每个不同 teacher head 在一个 mini-batch 中只需装载一次,并保证任一时刻设备中至多驻留一个 teacher head。报告没有说每个 GPU 整个 step 只处理一个 teacher。
- Async I/O(异步 I/O)
- 权重 / hidden state 的 load/unload 全部走后台流,不阻塞前向计算。等价于把 I/O 时间藏到计算下面(与 Ch7 MegaMoE 同思路)。需要双 buffer + 预取,但工程上是经典模式。
- TileLang KL Kernel
- student / teacher logits 的精确 KL 由专门 TileLang kernel(Ch8)计算,避免 PyTorch 动态显存分配带来的碎片。动机:每 step 算上万次 KL,PyTorch 的 op-by-op 调度会让显存 fragment 在几小时内累积爆炸;TileLang 的静态 layout 让 KL 张量始终在固定地址,0 fragmentation。
- WAL(Write-Ahead Log,写前日志)
- 数据库经典技术。每次状态变更先写到顺序日志,再改主数据。崩溃恢复时凭日志重放。V4 RL 借这个思路:rollout 每生成 1 个 token 就 append 到 trajectory log,preempt 时凭 log + 落盘 KV 直接续 decode,不必从头重生。
- Length Bias(长度偏差)
- 关键概念。短 trajectory 容易在 preempt 窗口内"完成",长 trajectory 容易被中断。如果中断后整体丢弃重跑,训练数据中"成功生成"的样本被系统性偏向短回答。policy 学到"短答案更可能成功"的伪信号 → distribution shift。Token 粒度 WAL 是为了让长 trajectory 也能"续上",从根上消掉这个 bias。
- Preemptible & Fault-Tolerant Rollout
- 大规模集群里 rollout 服务常被高优任务抢占(GPU 资源争用)+ 常发硬件故障。RL 训练的 rollout 必须能在两种事件下保持训练数据无偏。WAL 同时解决这两个:preempt 时凭 WAL + 落盘 KV 续 decode;硬件错误时凭 WAL 已生成 token 重做 prefill 重建 KV。
- 1M Context RL 的 trajectory 切两类
- 1M-token rollout 的 per-token fields 体积很大。V4 将数据拆为 lightweight metadata 与 heavy per-token fields:dispatch 只加载元数据做 shuffle/packing,重字段通过共享内存 data loader 按需读取。报告没有公布单条 trajectory 的固定 MB 数。
1. FP4 在 OPD 的全量启用
Ch10 已经讲过 FP4 QAT 的核心 trick:FP32 master → FP4 → FP8 dequant 无损,让训练 forward 也能用 FP4 权重。OPD 阶段把这件事推到极致 —— 所有不需要 backward 的 forward 路径(rollout 推理、N 个 teacher 算 logit、reference 模型算 KL 基线)全部走原生 FP4。
- rollout 与所有 inference-only forward(teacher、reference)使用原生 FP4 (MXFP4);
- 训练 step 的 backward 仍走 FP8 主路径,FP4 → FP8 dequant 无损(Ch10 sub-block scale 嵌套),与现有 mixed-precision pipeline 无缝对接;
- 不修改既有 backward pipeline,并降低 inference-only 路径的权重读取量;整体显存降幅取决于 student、teacher 与缓存布局。
- 报告没有对未来硬件给出固定的单 token FLOPs 降幅。
- FP8 inference weights:每个被量化权重使用 8 bits,另有 scale 与运行时状态;
- FP4 inference weights:对应权重使用 4 bits,可降低权重内存与读取流量;
- 不能将这项局部减半直接换算为完整 step 54% 或 GPU 数减半。
2. Teacher Scheduling:让 N 个 trillion 教师同框
OPD 的核心瓶颈不是算 KL(那是廉价 GEMM),是把 N 个 trillion teacher 同时供给 KL 计算。论文给的工程套路是把 teacher 状态拆到三处,让任意时刻 GPU 上至多挂一个 teacher head:
- 中央分布式存储装 teacher 权重(按 ZeRO-like 切片),按需加载;
- 中央 buffer装最后一层 hidden state(teacher forward 的产物),训练时再过对应 prediction head 重建 logits;
- GPU 显存动态装载当前 batch 用到的 teacher prediction head(仅最后一层)。
配合按 teacher index 排序的 batching:同一 mini-batch 内按 teacher 组织样本,使每个不同 prediction head 只装载一次,并保证任一时刻至多一个 head 驻留设备。
这里的核心是把"全部 teacher 全部 logit"这个不可能装进 GPU 的对象,分时分空间地散到三处:
- 权重:从中央存储按需切片加载 → I/O 是瓶颈,但与计算 overlap 可吃掉;
- Hidden state(每 token 一个 $d$ 维向量,$d=7168$):体积是 logits 的 $1/|V| \approx 1/18$,能装进 GPU 中央 buffer;
- Logits:从来不在显存里完整出现,用一次算一次扔一次,靠 hidden + prediction head 现场重建。
合起来,teacher weights 与 hidden states 通过中央存储和 buffer 分时流动,设备侧只保留当前需要的 prediction head,并用专用 kernel 计算精确 KL。
3. WAL Rollout:可抢占可容错 + 消除 length bias
OPD 是 on-policy 的(Ch18 §4),需要从当前 student 采样 trajectory,再计算 KL。大规模集群还必须处理两类运行时事件:
- 抢占:rollout 服务被高优训练任务挤占(GPU 资源争用),随时可能被 evict;
- 故障:大规模训练中需要考虑硬件错误、网络异常与存储故障;报告没有披露每千卡每月的故障次数。
最朴素的容错方案是把未完成的 rollout 整体丢弃重跑。这个方案有一个系统性偏差,是整章最容易被忽视但也最重要的洞察:
若任务可能在生成期间被随机抢占,从头重生成会让长响应承受更多次被中断的机会,而短响应更容易完整保留下来。下面描述的是这一概率偏差,不是假设所有任务每隔固定 $T$ 秒必定中断。
- 长 trajectory(深度推理、长代码生成)大概率被丢;
- 短 trajectory 大概率"完成"被纳入训练数据;
- policy 看到的"成功完成"样本平均长度系统性低估;
- 反向梯度引导 policy 偏向短答案 —— 这是训练数据被 evict 机制污染。
报告明确指出,从头重生成未完成请求在数学上会引入 length bias,使模型更倾向生成短序列。token 粒度 WAL 让未完成请求从中断处继续,从而避免这项正确性问题,并省去重复 decode 的成本。
V4 的设计:
- 每个生成请求维护 token 粒度 Write-Ahead Log —— 每生成 1 个 token 立刻 append 到 trajectory log;
- preempt 时暂停 inference,把整个 KV cache 落盘到分布式存储;
- resume 时凭 WAL + 落盘 KV 直接续 decode,从中断的 token 接着生成;
- 致命硬件错误(KV cache 不可恢复)时,凭 WAL 中已生成的 token 重做 prefill,重建 KV。
读图法:上排 6 条 trajectory 走"丢弃重跑"路径 —— 每次抢占(红段)整条丢、下次从 0 重生。越长的 trajectory 越难穿越完整抢占窗口,所以长 trajectory 完成率低,整体数据偏向短样本(length bias)。
下排同样 6 条走"WAL 续"路径 —— 抢占时落盘 KV(黄段),恢复后从 WAL 记录的 token 位置接着 decode。即使 trajectory 跨多个抢占窗口,最终都能完成。
把 trajectory 长度从 5K 拉到 100K,可以观察这套示意模型中长回答更容易跨过抢占边界。实际完成率取决于抢占过程和恢复策略,不应把 demo 的 100% 当作实测。
4. Million-Token RL 框架:把 trajectory 拆轻重
1M token 的单条 trajectory 一旦完整 load 进显存就秒爆。V4 的优化分两步:
- 切两类字段:每条 trajectory 拆成
- lightweight metadata:prompt、reward、长度、teacher 归属、状态码 等;
- heavy per-token field:每 token 的 hidden state、KV、logit;
- 分层加载:
- dispatch 阶段只 load metadata,做 global shuffle 与 packing layout 计算;
- per-token 重字段通过共享内存 data loader 按需加载,节点内多 GPU 共享一份,消除冗余;
- mini-batch 消费完立即释放,CPU/GPU 内存压力随时间稳定;
- on-device mini-batch 数量按 workload 动态决定,在计算吞吐与I/O overlap之间取最优解。
- metadata:prompt、reward、长度、teacher index 与状态等轻量字段;
- heavy per-token fields:随序列长度增长的 hidden states、训练字段和其它逐 token 数据;
- 具体体积取决于保存字段、精度和切分方式,报告没有给出 10KB、7GB 或 60GB 的统一数字;
- logits(中央 buffer 缓存):根本不实例化(Ch19 §2);
5. 与 DSec Sandbox 的接缝
本章四节解决的是 RL/OPD 训练 + rollout 推理的工程问题。但 agentic rollout 还要执行外部命令(bash、文件、网页、单元测试),这部分的承载者是下一章的 DSec —— 整个 post-training 流水线由 "训练框架 + rollout 引擎 + sandbox 平台"三件套共同承担。
报告将十余个 teacher、全词表 KL、可抢占 rollout 与长序列训练组合在同一框架中。其工程复杂度很高,但论文没有比较社区实现,也没有宣称 24/7 无故障。
Ch19 + Ch20 合起来的工程深度,是 V4 区别于其它"做开源蒸馏"工作的真正分水岭 —— 不是论文里的公式新,是公式底下那层能让公式跑起来的水管新。
6. 一句话总结
四项设计分别处理不同瓶颈:FP4 降低 inference-only forward 的权重流量;Teacher Scheduling 分时加载权重、hidden states 与 prediction heads;token-level WAL 支持抢占和故障恢复并避免从头重生成引入长度偏差;Million-Token RL 将 metadata 与 heavy per-token fields 分层加载。