AI

RL r3 的超高校级的实现

Posted by w@hidva.com on September 23, 2026

早在 qwen3.8-flash-next 设计之初, 一个任务就摆在了我们 RL infra 组面前: 如何高效地把 rollout 时 QSA indexer 选出的 top-k 送到 train 端, 让训练复用推理时的选择, 像 router replay 一样, 尽力减少训推不一致性. 两者要录制的东西不同. Router replay 记录的是一个 token 在 MoE 层选中了哪些 expert; QSA indexer replay 记录的是这个 token 做稀疏 attention 时选中了哪些位置. 训练端拿同一段 token 重新 forward, 两边的 kernel、并行方式和 batching 都可能不同, 浮点计算的微小差异又可能改变离散的 top-k 选择. Expert 换了, 或者 attention 看向的位置换了, 后面的 logprob 和梯度自然也会跟着变. Replay 要尽量固定住的, 就是这部分选择.

听起来无非是多带一个字段, 真正算一下数据量就没这么轻松了: 在当时按 token 索引录制的口径下, QSA top-k 的原始字节量是 router top-k 的 100+ 倍. 一条长轨迹附带的索引就能到 GB 级别, 一个 batch 再乘上去, 谁看了都得皱一下眉. 下文沿用这一录制表示来讲这次优化. 大家一开始都觉得这是个很艰巨的任务. 一位同事在 POC 中做了不少努力, 各种分片, 各种压缩, 结果整个 step 的耗时仍然额外增加了 8 倍. 录制和回放确实接上了, 但这样的性能, 很难进入日常训练.

后来我把整条数据流重新看了一遍, 发现优化方向从一开始就走偏了. 那些中间组件根本不需要理解 top-k, 却被要求带着它完成一整套样本处理和训练布局变换. 把数据压小一些, 只能让这趟搬运便宜一点; 我想做的是让它直接退出这些环节. 最后落地之后, 相对关闭 indexer R3 的完整 step, 开启后的额外耗时不足 7%. 这也就是我在 RL 其实很简单 里提到的那件事:

玩笑背后是个挺认真的判断: 真正的优化机会几乎都长在链路的接缝处, 而接缝恰好是没有人完整拥有的地方. 每个组件的 owner 都把自己那段做到了局部最优, 于是剩下的空间全在”上一跳交给下一跳的时候, 到底交了什么”上面. 靠着对这条链路的掌控, 我把 DSA R3 的数据流从 GB 粒度降到了 KB 级别, 让一个原本被数据通路压死的特性以近乎 zero overhead 跑了起来.

这里旧文里的 DSA R3, 指的就是本文讨论的 indexer replay 这条链路. 本来我还犹豫这事要不要公开讲. 直到后来在 DeepSeek-V4 技术报告里看到这一段:

For the inference and training phase, we decompose the rollout data format into lightweight metadata and heavy per-token fields. During data dispatching, the metadata for the entire rollout data can be loaded to perform global shuffling and packing layout computation. Heavy per-token fields are loaded via a shared-memory data loader to eliminate intra-node data redundancy and are released immediately upon consumption at the mini-batch granularity, substantially reducing both CPU and GPU memory pressure.

看到 lightweight metadata 和 heavy per-token fields 分开, 再用 metadata 做 shuffling 和 packing, 我就知道大家想到了同一个地方. 具体用什么存储、怎么编码引用、什么时候回收, 实现可以各不相同, 但这个拆分思路已经讲得很清楚了. 索性我这里也详细展开一下.

一份 top-k 的九站旅程

要讲清楚优化做了什么, 得先看优化前这份数据经历了什么. POC 代码已经找不到了, 好在 router r3 现在还能翻到, 走的是同一形状的链路, 就拿它演示. 再强调一遍别搞混: 这里说的是 router r3, 录的是 expert; topk r3 录的是 attention 位置, 借 router 看的只是链路的形状. 一个 token 的 router top-k, 从推理到训练, 一共九站:

  1. 逐层捕获. 推理引擎 (SGLang) 每个 MoE 层的 router 选完 expert, 顺手把选择写进一块 device buffer, 形状是 (max_tokens, 层数, topk).
  2. DtoH. forward 收尾, 把本轮有效的行切出来, 经 pinned buffer 拷回 host, 再按 KV cache slot 散射进一块 CPU buffer. 拿 slot 当下标有个隐藏福利: 投机解码被 reject 的重试行会落回同一个 slot, 覆盖自动完成, 不用写一行去重代码.
  3. 按请求收集. scheduler 每轮 decode 把属于这个请求的行 gather 出来, 拼到这份录制后面.
  4. 多段归并. 一个请求可能分好几段执行, 收尾时把各段拼成一份请求级记录.
  5. 随响应回家. 记录经 base64 编码塞进 HTTP 响应的 meta_info. 对, base64, 字节先膨胀三分之一再上路.
  6. 训练框架收口. RL 框架拿到响应, 多轮轨迹把各轮记录拼起来, 按训练流水线各 stage 负责的层把层维切开, 存进对象存储, 样本身上只带一份 blob 引用.
  7. micro-batch 加载. 训练侧读回本 stage 的层切片, 尾部补行, pad 到 tp*cp 的倍数, 再按 rank 连续切开.
  8. 布局机器. 这份记录跟着 token 走 fused unpad: 去 padding, 重排, all_reduce, 顺手再提升成 int64. 这一站的存在意义, 是让每个 token 的录制和它的 hidden states 始终逐 token 对齐.
  9. replay. 记录经 ContextVar 注入 Megatron, router 前向时用录制的 expert 选择, 概率拿当前 logits 重新算.

第 6 站已经把数据放进对象存储了, 但问题并没有到此结束. 第 7 站取回本 stage 的序列记录之后, 第 8 站仍然拖着真正的 expert ID 张量去做布局变换. 所以, “RPC 里只传引用”解决了一段传输, 不代表后续每一跳都摆脱了大数据. 早期 indexer POC 则走了另一套具体实现: 推理侧先分片落文件, 样本收口时读回, 做多轮归并、padding、排序和压缩, 训练侧再加载、解压、切片, 最后参加布局变换和通信. 它与上表的传输格式不同, 但同事优化时围绕的也是这些交接环节: 这一段再切细一点, 那一段再压小一点, 想办法让大张量能够继续往下走.

不知道你有没有发现, 这些 top-k 数据一路随波逐流. 中间环节当然会 touch 它们的字节: 拷贝、编码、归并、压缩, 哪一个都少不了读写. 但从模型语义看, 它们根本不需要判断某个 expert 为什么被选中, 更不需要知道某个 attention 位置的分数是多少. 真正产生选择的是推理模型, 真正消费选择的是训练模型. 中间环节关心的始终是另一件事: 这个 token 现在属于哪个样本, 被排到哪里, 最后交给哪个 rank.

那么, 给它们一个能找到 top-k 的引用, 就够了.

超高校级的实现

优化后的链路一句话就能说完: payload (top-k 本体) 只写一次, 引用走全程, payload 在消费点物化一次.

[payload]  推理端定稿 ──→ 共享存储 ───────────────────────→ 训练消费点
[引用]     响应元数据 ──→ 样本 manifest ──→ 调度/packing ──→ 布局机器 ──┘

先讲 payload 怎么写出去. 推理侧录制先落在 device buffer 里. 投机解码下, 一次 forward 出现的行不一定作数: verify 之后尾部位置可能要重算, 新记录覆盖旧记录. 所以录制在内存里留一个尾窗, 窗口里的修改就地解决; 窗口之前已经稳了的行, 攒够一个 chunk 就写进共享存储, 我们直接复用了 Mooncake, (可参考我们 Qwen3.8: 从 0 到 1 用 Rust 重写 Mooncake: Suncake 了解 Mooncake 在我们 RL 系统栈中的应用). 落出去的对象 write-once, 存储侧不需要任何覆盖和删除语义. chunk 还按层拆成独立对象: 训练流水线一个 stage 只消费其中几层, 各取所需, 不用把全部层拉回来再扔掉大半. 写入丢给后台线程池, 但得有背压: 在途字节过了水位, 就让 scheduler 主循环等一等 —— 不然 rollout 快落盘慢, 内存迟早被后台任务吃光. 另外, 一堆并行 rank 里只有一个 owner rank 负责写: 对象按 uuid 寻址, 谁写出去的谁才拿着 key, 非 owner 也写的话, 那些对象没有任何人持有 key, 永远回收不掉.

然后, 真正随响应回家的, 只剩一串四元组: (对象 uuid, 起始行, 行数, 层数). 一条 GB 级轨迹的 payload, 用几十段这样的区间就描述完了, KB 级. RL 框架收口多轮轨迹时, 过去是读回 payload, 铺一张大画布互相覆盖; 现在覆盖只发生在区间归属上 —— 后一轮的段盖住前一轮重叠的区间, payload 一动不动. 收口的产物是一份 manifest: 哪些对象, 哪些区间, 归到哪一段. 样本调度, batch 重排, dynamic packing, 全都带着这份 KB 级的清单走.

接下来是最难的一段, 也是我最得意的一段: 一行新的并行代码都没写. 训练前排布 token 要去 padding, packing, 按并行策略切分, “每个 token 的录制和它的 hidden states 逐 token 对齐” 是这条链路最难保证的性质, 错位就是静默错误. 但这套对齐逻辑, router r3 已经验证过一遍了. 所以我们的做法是: 到 micro-batch 真要排布 token 的时候, 把 manifest 展开成每 token 一个 24 字节的 handle —— 128 bit 的对象 uuid, 加行内坐标, 再加两个哨兵值区分 gap 和 pad —— 然后把它喂进 router r3 那台 fused unpad, 切分, all_gather 机器. 机器不挑货, 它只关心哪一行去哪一行, 至于一行是 96 KiB 还是 24 B, 它无所谓. 以前它搬的是 (行数, 层数 × 2048) 的 int32, 现在搬 (行数, 6) 的 int32, all_reduce 通信量差了三个数量级.

token 的最终位置定了之后, 才轮到 payload 第二次出场. 训练侧把 handle 按 uuid 归组, 本 stage 只取自己那几层的对象, 把对应行抽出来组装成 record, 经 ContextVar 注入 Megatron. 这是 payload 全程唯一一次物化, 常驻范围只有当前 micro-batch. 读不到就死等, 每 10 秒叫一声: rollout 是异步落盘的, “暂时读不到” 多半只是写入还没落地, 等一等就兑现了; 当然目前线上还没有遇到第一次读不到的情况, 毕竟从写入到读取会耽搁很久=。=真读不到说明是 bug 或容量不够, 这种时候静默丢一部分 replay, 比直接崩 job 危险得多. 回收也是单一真源: manifest 枚举了这个 batch 的全部对象, 包括那些被后轮覆盖, 已经不在最终布局里的, 消费完统一删, 不依赖任何旁路记账. 把账摆在一起:

  • 控制链路 (线上传输, 样本收口, 调度, packing): 原来拖着 GB 级 payload, 读回, 归并, 压缩; 现在只带 KB 级 manifest. 一条轨迹的 payload 3.7 GB, 换成几十段区间, KB 级.
  • 布局机器 (unpad, 切分, all_reduce, all_gather): 原来每 token 96 KiB; 现在每 token 24 B. 一条一百 K token 的轨迹, 全部 handle 加起来 2 MB 出头, 搁在原来的口径里, 只够 25 个 token 用.
  • payload 本身: 该走的路一步没少 —— 从 GPU 拷出来, 写存储, 被训练读回去 —— 但它只走了生产者到消费者这条直路, 不再陪中间环节巡演.

端到端的结果前面说过了: 相对不开 r3 的 baseline, 额外开销不足 7%. 这是整条链路全开起来的总账, 录制, 存储, 读取, 回放全在里面. 而且前面说了, 这套拆法不挑数据: 任何要从 rollout 带去 train 的 token 粒度状态, 只要中间环节不需要看内容做决策, 都可以让布局信息先走, 本体留到消费点.

人在回路

回头看, 这件事最关键的判断, 在动手写实现之前就已经做完了: 哪些组件真正需要 top-k, 哪些组件只是因为接口长成这样, 被迫替下一跳搬运它. AI 可以很快帮我们写出切片、压缩、线程池、批量读取, 也可以继续沿着现有链路增加一轮又一轮优化. 你让它减少这一跳的数据量, 它就认真地减少这一跳的数据量. 但如果没有人先问一句“这份数据为什么要经过这一跳”, 代码写得越快, 也可能只是在错误的方向上走得越远. 这也是我现在越来越在意的事. AI 时代, 每个人 rush 代码都很快, 实现一个想法的门槛低了很多. 人反而更需要从局部实现里退出来, 把生产者、消费者、中间状态和生命周期完整看一遍, 对问题保留自己的判断.

AI 可以帮我免去逐行掌握所有细节的负担. 但数据为什么在这里, 谁需要它, 它应该活多久, 这些问题必须想明白. 否则只是更快地写完了一条本来可以不存在的数据通路.

后记

在去年 啊?你怎么知道我从104kg健到了70kg! 之后,笔者终于在节前又一次健到了 70kg 的水位附近(这次是从 80kg 开始)。这一次到多了许淡定与顺其自然。坚持3年的每天早上中午 2*400 卡最近 3 个月换成早晨 1 次 500 卡。吃也没有再斤斤计较而是发动了俺寻思之力(俺寻思这样吃可以健肥)。现在抬头望去,好像什么都想吃,但好像又什么都没有必要吃。(当然这并不包括我永远最爱的年糕与开心果大米冰淇淋!