Ch 19 · RL / OPD 工程
第四部分 · 后训练 · 19

RL / OPD 工程 — 后训练的发动机

Ch18 把"OPD 为什么可行"讲完了;这一章只回答"它怎么跑得起"。trillion 级 teacher × 全词表 logit × 1M context × on-policy rollout 这四件事任一单独都能压垮集群,V4 用四个工程支柱(FP4 / Teacher Scheduling / WAL Rollout / Million-Token RL)把它们同时跑起。

名词速通 · 一分钟看懂"OPD 工程支柱"

四个支柱 = (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 数。
一句话定位:FP4 降低 inference-only forward 的内存流量;Teacher Scheduling 避免同时驻留全部 teacher 与 logits;token-level WAL 支持无长度偏差的抢占恢复;分层数据加载控制 1M-token rollout 的内存压力。

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 降幅。
数值演练 · FP4 在 OPD 阶段省的显存 示意:inference-only teacher、reference 与 rollout 使用 FP4 后,被量化权重的存储相对 FP8 约减半;但完整 OPD step 的显存还取决于 student training state、offload 与缓存。
  • FP8 inference weights:每个被量化权重使用 8 bits,另有 scale 与运行时状态;
  • FP4 inference weights:对应权重使用 4 bits,可降低权重内存与读取流量;
  • 不能将这项局部减半直接换算为完整 step 54% 或 GPU 数减半。
这项优化有助于容纳更多 inference-only forward,但报告没有量化为“不增加集群即可增加多少 teacher”。

2. Teacher Scheduling:让 N 个 trillion 教师同框

OPD 的核心瓶颈不是算 KL(那是廉价 GEMM),是把 N 个 trillion teacher 同时供给 KL 计算。论文给的工程套路是把 teacher 状态拆到三处,让任意时刻 GPU 上至多挂一个 teacher head:

  1. 中央分布式存储装 teacher 权重(按 ZeRO-like 切片),按需加载;
  2. 中央 buffer最后一层 hidden state(teacher forward 的产物),训练时再过对应 prediction head 重建 logits;
  3. GPU 显存动态装载当前 batch 用到的 teacher prediction head(仅最后一层)。

配合按 teacher index 排序的 batching:同一 mini-batch 内按 teacher 组织样本,使每个不同 prediction head 只装载一次,并保证任一时刻至多一个 head 驻留设备。

📖 三处状态拆分的物理账

这里的核心是把"全部 teacher 全部 logit"这个不可能装进 GPU 的对象,分时分空间地散到三处:

  1. 权重:从中央存储按需切片加载 → I/O 是瓶颈,但与计算 overlap 可吃掉;
  2. Hidden state(每 token 一个 $d$ 维向量,$d=7168$):体积是 logits 的 $1/|V| \approx 1/18$,能装进 GPU 中央 buffer
  3. 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 整体丢弃重跑。这个方案有一个系统性偏差,是整章最容易被忽视但也最重要的洞察:

为什么"丢弃重跑"会引入 length bias

若任务可能在生成期间被随机抢占,从头重生成会让长响应承受更多次被中断的机会,而短响应更容易完整保留下来。下面描述的是这一概率偏差,不是假设所有任务每隔固定 $T$ 秒必定中断。

  1. 长 trajectory(深度推理、长代码生成)大概率被丢;
  2. 短 trajectory 大概率"完成"被纳入训练数据;
  3. policy 看到的"成功完成"样本平均长度系统性低估
  4. 反向梯度引导 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。
Demo · 抢占下"丢弃重跑" vs "WAL 续 decode" 的时间线 + 成功完成比例(拖动 trajectory 长度)
交互
"丢弃重跑" 方案:抢占即整段丢,下次重头来过 "WAL 续 decode" 方案:抢占时落盘,恢复时从 WAL 接 完成比例 = 0% · length bias = 严重 完成比例 = 0% · length bias = 生成中 抢占丢弃 落盘 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 的优化分两步:

  1. 切两类字段:每条 trajectory 拆成
    • lightweight metadata:prompt、reward、长度、teacher 归属、状态码 等;
    • heavy per-token field:每 token 的 hidden state、KV、logit;
  2. 分层加载
    • dispatch 阶段只 load metadata,做 global shuffle 与 packing layout 计算;
    • per-token 重字段通过共享内存 data loader 按需加载,节点内多 GPU 共享一份,消除冗余;
    • mini-batch 消费完立即释放,CPU/GPU 内存压力随时间稳定;
    • on-device mini-batch 数量按 workload 动态决定,在计算吞吐I/O overlap之间取最优解。
数值演练 · 1M token trajectory 的体积分布 单条 1M token trajectory:
  • metadata:prompt、reward、长度、teacher index 与状态等轻量字段;
  • heavy per-token fields:随序列长度增长的 hidden states、训练字段和其它逐 token 数据;
  • 具体体积取决于保存字段、精度和切分方式,报告没有给出 10KB、7GB 或 60GB 的统一数字;
  • logits(中央 buffer 缓存):根本不实例化(Ch19 §2);
关键不在固定比例,而在访问模式:global shuffle 与 packing 只需要 metadata;heavy fields 到 mini-batch 消费时再加载并立即释放,可降低 CPU/GPU 内存峰值。

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 分层加载。