Online Draft Co-Training for Speculative Decoding in Large-Scale, Long-Context RL Post-Training

TL;DR

Online Draft Co-Training结合Zigzag Ring与TapChannel,实现最高1.88倍端到端加速。

cs.LG 🔴 高级 2026-09-07 41 次浏览
Zili Wang Zhaopeng Qiu Yuekai Zhang Shuang Yu Junjie Lai
推测解码 强化学习后训练 上下文并行 流水线并行 长上下文

核心发现

方法论

论文提出Online Draft Co-Training:策略模型参数θ通过RL目标更新,草稿模型参数ϕ利用同一批rollout token及策略中间特征Hθ训练,目标为L(θ,ϕ)=LRL(θ)+λLdraft(ϕ;x,sg(Hθ(x)))。系统在不改变策略TP/PP/CP拓扑的前提下,采用Zigzag Ring Attention处理主序列因果注意力,并将分支注意力保留在所属rank;再用TapChannel跨PP阶段传输目标特征。

关键结果

  • 在DAPOMath-17K、AIME 2024上的Qwen3-8B实验中,EAGLE-3、DFlash和DSpark均紧跟无草稿基线的reward、验证准确率及训练-推理KL轨迹;接受长度分别为2.28、3.45和3.63,端到端加速分别为1.50×、1.88×和1.83×。
  • 跨模型规模实验覆盖8B至122B:Qwen3.5-35B-A3B+DFlash端到端1.46×,Qwen3.5-122B-A10B+DFlash为1.35×,GPT-OSS-120B+DFlash为1.19×;rollout加速范围为1.19–2.23×。
  • 在最长256K上下文上,EAGLE-3 TTT的CP=1到8延迟由17.7秒降至2.35秒,达到7.5×扩展效率94%;内存由53.2GB降至7.5GB。相较USP,延迟最高降低2.9×、显存降低2.7×。

研究意义

该工作解决了RL后训练中“生成最慢、策略持续变化、长上下文成本高”的组合瓶颈。它表明推测解码不仅是推理服务优化,也可作为策略学习系统的在线协同模块。对学术界而言,论文把分支结构注意力、目标特征蒸馏与大规模并行训练统一起来;对工业界而言,它支持8B至122B模型及单轮、多轮工具调用任务,并将rollout加速转化为1.16–1.88倍训练加速。

技术贡献

核心工程贡献包括两点。其一,分支查询分别计算主序列因果部分和本地branch部分,并用在线softmax合并:ℓ=log(eℓm+eℓb),O=eℓm−ℓOm+eℓb−ℓOb;该设计兼容EAGLE-3、DFlash和DSpark。其二,TapChannel通过CUDA IPC或GPUDirect RDMA建立不进入PP调度的侧通道,以邮箱、序列戳和预分配槽位完成目标特征fan-in,避免改写流水线。

新颖性

新颖性不在于首次提出推测解码或在线草稿训练,而在于首次系统性地把多种分支草稿架构接入既有大模型CP/PP拓扑。相较SpecForge的顺序ring和USP填充式注意力,本文采用packed、load-balanced zigzag ring;相较标准相邻stage通信,TapChannel支持非相邻、单向、调度外特征传输。

局限性

  • 端到端收益依赖rollout占总步时的比例;在Workplace Assistant中,工具和环境延迟不可被解码加速覆盖,因此端到端收益仅约1.25–1.43×。
  • 强扩展会缩短每个rank的本地序列,使通信逐渐成为瓶颈;此外,MoE模型的验证前向仍需稀疏专家计算,故接受长度高不必然带来同等训练加速。
  • 实验主要基于NeMo-RL、指定模型和官方草稿checkpoint,尚未充分覆盖更广泛奖励函数、硬件及草稿训练目标。

未来方向

后续可研究通信与注意力的更细粒度重叠、动态选择草稿长度和架构,以及面向接受率的直接优化目标。还应评估更长多轮轨迹、异构集群、不同MoE路由策略和更大模型,并探索把TapChannel扩展为通用的跨阶段特征服务层。

AI 总览摘要

强化学习后训练的主要时间常被rollout生成占据。推测解码让小型草稿模型先提出多个token,再由目标策略并行验证;但随着策略不断更新,固定草稿会失配。更困难的是,EAGLE-3、DFlash和DSpark需要分支注意力及目标模型中间特征,而标准因果上下文并行和流水线并行并不支持这些数据流。

Wang等提出Online Draft Co-Training,在同一RL过程中联合更新策略与草稿模型,并以stop-gradient固定目标特征。其Zigzag Ring Attention将分支查询拆成主序列因果注意力和rank本地分支注意力,再以在线softmax合并;TapChannel则通过独立侧通道把远端PP阶段的特征直接送至草稿所在阶段,不改变原流水线调度。该设计覆盖EAGLE-3、DFlash和DSpark。

结果显示,Qwen3-8B在DAPOMath-17K上使用DFlash时接受长度3.45、rollout加速2.23×、端到端加速1.88×;跨8B至122B模型,端到端加速为1.16–1.88×。在256K上下文中,EAGLE-3 TTT达到94%并行效率,显存从53.2GB降至7.5GB。局限是工具调用、环境等待和MoE验证计算会削弱总收益,但该工作证明在线草稿训练可成为大规模RL基础设施的一部分。

深度分析

研究背景

推测解码由Leviathan等提出,利用草稿模型预测、目标模型并行验证,并以拒绝采样保持策略分布。RL系统如NeMo-RL已集成MTP和EAGLE-3;FastGRPO、ReSpec等进一步在线适配草稿。但EAGLE-3的TTT、DFlash的块扩散和DSpark的Markov head带来分支KV与目标中间特征,传统RingAttention只能处理因果主序列。

核心问题

论文要在大型、长上下文RL后训练中联合训练草稿,同时保留策略已有TP/PP/CP布局。CP侧,分支查询既要看主序列前缀又要看本地分支KV;PP侧,目标特征分散在非相邻阶段而草稿位于末阶段。若改变并行拓扑,会增加内存、调度和工程复杂度。

核心创新

第一,Zigzag Ring把变长主序列packed且负载均衡地分片,主序列K/V循环通信,分支K/V留在anchor rank。第二,在线softmax合并两类注意力结果,统一支持三种草稿。第三,TapChannel以预分配mailbox、sequence stamp、CUDA IPC和GPUDirect RDMA实现调度外feature fan-in。与USP相比避免2.25倍padding。

方法详解

  • �� 联合目标:用同一rollout token训练策略和草稿,草稿输入为sg(Hθ),避免辅助损失反向改变策略特征。
  • �� CP计算:每个branch query固定在本rank;主序列K/V进行C个ring step循环。主序列与分支分别得到(Om,ℓm)、(Ob,ℓb),按ℓ=log(eℓm+eℓb)和O=eℓm−ℓOm+eℓb−ℓOb合并。
  • �� 通信:前一attention step计算时预取下一K/V shard,暴露开销为max(tcomm,tattn);CP通信量为2(C−1)N/C·dkv·b。
  • �� PP传输:源stage写入末stage预分配槽位,消费者依据序列戳读取;跨节点使用NCCL+RDMA,同节点使用CUDA IPC。
  • �� 训练评估:在NeMo-RL中用GRPO联合更新,并测量reward、AIME准确率、KL、接受长度、吞吐和总step时间。

实验设计

实验使用DAPOMath-17K训练、AIME 2024验证,并在NeMo Gym Workplace Assistant上测试32,768长度多轮工具调用。目标包括Qwen3-8B、Qwen3.5-35B-A3B、Qwen3.5-122B-A10B、Nemotron-3.5-Lightning-30B-A3B和GPT-OSS-120B;草稿包括EAGLE-3、DFlash、DSpark。CP比较USP,序列长度扩展至256K;PP比较TapChannel与host staging。

结果分析

三类草稿均保持与baseline相近的RL轨迹。Qwen3-8B上DFlash端到端1.88×,DSpark为1.83×;最大规模122B上DFlash仍达1.35×。CP实验中,packed zigzag相对USP在CP=2、4、8分别快2.9×、2.3×、1.5×,显存统一少2.7×。TapChannel fan-in比host staging快4.5–8.5×,接收端HBM竞争最多1.6%。

应用场景

该系统适用于数学推理、代码生成、代理工具调用和长文档RL后训练。使用者需要支持NeMo-RL的分布式GPU集群、目标模型中间特征接口及可训练草稿模块。对在线服务而言,持续适配的草稿可减少策略更新后的失配;对训练平台而言,现有TP/PP/CP拓扑无需重构。

局限与展望

加速并非只由接受长度决定:多轮任务中工具执行和环境等待占据不可消除的时间,MoE验证也可能增加专家计算。CP在极强扩展时会暴露通信瓶颈,TapChannel还需管理跨节点缓冲和in-flight发送。论文尚未给出广泛消融、不同草稿损失的系统比较,也未证明在所有奖励任务和硬件上都能保持相同收益。

通俗解读 非专业人士也能看懂

把目标模型想成大型餐厅,完整做菜很慢;草稿模型像熟悉菜单的小助手,先把可能的几道菜准备好,大厨只需一次检查多道准备工作。若检查通过,顾客很快拿到餐品;若不通过,大厨仍按自己的标准修改,所以最终味道不变。

问题是餐厅有很多厨房楼层,而且长订单被分散到不同工作台。普通传菜系统只连接相邻楼层,也不懂“同一订单里临时分支的菜”。论文设计了两套新传输方式:Zigzag Ring让主订单资料在各工作台循环,而临时分支资料留在原地;TapChannel则像专用电梯,把远处厨房做好的关键材料直接送到最后的草稿工作台。

结果很实用:Qwen3-8B配DFlash时,生成阶段快2.23倍,总训练快1.88倍;256K长度时仍能有效扩展。代价是电梯和检查本身仍要花时间,尤其是工具调用或专家模型计算很多时,总体收益会变小。

简单解释 像给14岁少年讲一样

想象你在玩一个超大型游戏,电脑里的“主玩家”很聪明,但每走一步都要花很久。旁边有个小助手,先猜接下来可能出现的几个动作。主玩家一次检查这些猜测,如果猜对,就不用一个动作一个动作慢慢算,游戏立刻快很多!

可是主玩家会不断学习,旧助手可能跟不上。论文让助手在主玩家训练时一起学习,所以它会越来越懂主玩家。这里有三种助手:EAGLE-3像边玩边预测,DFlash像一次猜一整组动作,DSpark再把这些动作按顺序稍微修正。

真正难点是游戏地图特别大,而且电脑把任务分给很多GPU。论文的Zigzag Ring像让资料在队伍里传递,同时把每个小分支留在自己的队员手上;TapChannel像专用快递通道,把远处GPU的资料送给助手,不打乱主任务。

实验中,Qwen3-8B配DFlash让rollout快2.23倍、整体训练快1.88倍;256K超长输入也能扩展。听起来很棒,但如果游戏中要等玩家操作、工具反馈,或者每次检查本来就很复杂,整体速度提升就不会同样大。

术语表

Speculative Decoding(推测解码)

小模型先提出多个token,大模型并行验证。验证和拒绝采样保证最终输出仍服从目标模型分布。

本文用它加速RL rollout生成。

Online Draft Co-Training(在线草稿协同训练)

草稿模型随不断变化的RL策略同步训练。它使用策略中间特征,但通过stop-gradient避免辅助损失改变策略。

论文的核心训练框架。

Context Parallelism, CP(上下文并行)

把长序列分布到多个GPU,并交换K/V以完成注意力。它降低单卡长上下文内存。

本文在CP中加入branch attention。

Pipeline Parallelism, PP(流水线并行)

把模型层切分到多个阶段,微批次依次通过这些阶段。非相邻阶段通常不能直接交换特征。

TapChannel专门解决该限制。

TapChannel

一种不进入PP调度的目标特征侧通道。它用mailbox、序列戳、CUDA IPC或GPUDirect RDMA完成传输。

将远端target features送到末stage草稿模块。

Acceptance Length(接受长度)

每次目标模型验证平均接受的草稿token数。数值越高,通常表示推测解码越有效。

本文报告范围为2.28–4.78。

开放问题 这项研究留下的未解疑问

  • 1 不同草稿训练损失、λ设置和奖励任务如何影响接受率与策略稳定性,论文尚未系统拆解。
  • 2 当上下文继续超过256K或集群网络更异构时,通信、缓存和调度是否会成为主要瓶颈仍待验证。
  • 3 工具延迟、MoE专家路由与草稿接受率之间的联合优化机制尚不清楚。

应用场景

近期应用

长上下文数学RL训练

研究团队可在NeMo-RL和GRPO中接入DFlash或DSpark,对DAPOMath类任务在线训练草稿。需要支持CP/PP的GPU集群;预期可获得约1.16–1.88倍端到端加速。

多轮代理与工具调用

代理系统可用EAGLE-3、DFlash或DSpark加速模型生成,同时保留目标策略的输出分布。应单独测量工具等待时间,因为它会限制总训练收益。

远期愿景

统一的分布式推测训练平台

未来可将TapChannel和分支注意力抽象为通用基础设施,使不同草稿架构自动适配TP、PP、CP及MoE模型,支持更长轨迹和更大参数规模。

原文摘要

Speculative decoding accelerates rollout generation, which dominates the cost of reinforcement learning (RL) post-training. Online co-training can further increase the draft's accuracy, yielding greater speedups. However, scaling this approach to co-training on large models with long contexts poses two obstacles: (1) branch attention is unsupported by standard causal context-parallel (CP) implementations, and (2) target features span across pipeline-parallel (PP) stages. We address both with an end-to-end system for large-scale online draft co-training. For CP, we extend packed, load-balanced zigzag ring attention by merging rank-local branch attention with causal main-sequence attention. For PP, TapChannel transports intermediate target features across stages via a separate path, leaving the pipeline schedule unaffected. Experiments demonstrate that co-trained drafts closely track the policy baseline while delivering substantial rollout and end-to-end speedups across model scales up to 122B. Our CP design achieves strong scaling at 256K tokens with significant memory savings over prior work, and our PP transport incurs modest overhead. Code can be found at https://github.com/NVIDIA-NeMo/RL/issues/3698.

cs.LG cs.DC