ChatPaper.aiChatPaper

대규모·장문맥 RL 후학습에서의 추측 디코딩을 위한 온라인 드래프트 공동 학습

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

September 7, 2026
저자: Zili Wang, Zhaopeng Qiu, Yuekai Zhang, Shuang Yu, Junjie Lai
cs.AI

초록

추측 디코딩은 강화 학습(RL) 사후 학습 비용을 지배하는 롤아웃 생성을 가속한다. 온라인 공동 학습은 드래프트의 정확도를 더욱 높여 더 큰 속도 향상을 낳을 수 있다. 그러나 이 접근법을 긴 컨텍스트를 가진 대규모 모델의 공동 학습으로 확장하는 데에는 두 가지 장애물이 있다: (1) 분기 어텐션은 표준 인과적 컨텍스트 병렬(CP) 구현에서 지원되지 않으며, (2) 타깃 특징이 파이프라인 병렬(PP) 스테이지에 걸쳐 있다. 우리는 대규모 온라인 드래프트 공동 학습을 위한 엔드투엔드 시스템으로 두 문제를 모두 해결한다. CP의 경우, 랭크-로컬 분기 어텐션을 인과적 메인 시퀀스 어텐션과 병합하여 패킹되고 로드 밸런싱된 지그재그 링 어텐션을 확장한다. PP의 경우, TapChannel은 별도 경로를 통해 중간 타깃 특징을 스테이지 간에 전송하며 파이프라인 스케줄에는 영향을 주지 않는다. 실험은 공동 학습된 드래프트가 정책 기준선을 밀접하게 추종하면서 최대 122B의 모델 규모 전반에서 상당한 롤아웃 및 엔드투엔드 속도 향상을 제공함을 보여준다. 우리의 CP 설계는 256K 토큰에서 강한 스케일링을 달성하고 이전 연구보다 상당한 메모리 절감을 제공하며, 우리의 PP 전송은 약간의 오버헤드를 발생시킨다. 코드는 https://github.com/NVIDIA-NeMo/RL/issues/3698에서 확인할 수 있다.
English
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.