DualKV:共享提示 Flash Attention,通过大规模部署和长上下文实现高效 RL 训练
DualKV: Shared-Prompt Flash Attention for Efficient RL Training with Large Rollouts and Long Contexts
摘要
现代 RL 后训练 方法(例如 GRPO 和 DAPO)在从 P 词元的共享提示中采样的 R 词元的 N 个响应序列上进行训练,但标准 FlashAttention 在前向和向后传递中复制所有 P 提示词元 N 次 - 在相同的隐藏状态上复制计算和内存。在大规模部署、长上下文 RL 训练中(N>=16,P>=8K),这种冗余在策略更新成本中占主导地位。我们观察到,在仅解码器模型中,因果掩蔽使得提示表示在每一层的序列之间保持不变,因此所有每个标记的操作(范数、投影、MLP)和注意力都可以处理一次提示——这一属性尚未在内核级别用于训练。我们提出 DualKV,第一个消除 RL 训练期间共享提示复制的 FlashAttention 内核变体,通过 (1) 融合 CUDA 前向和后向内核,在单个内核启动中迭代两个不相交的 KV 区域(共享上下文和每序列响应),以及 (2) veRL 中的数据管道重新设计,将 N(P+R) 词元重新打包为每个微批次的 P+NR 词元,将词元减少从注意力扩展到整个通过因子 rho = N(P+R)/(P+NR) 进行模型。 DualKV 在数学上等价于标准注意力,并且不引入近似值。在使用 8xH100 GPU(N=32,8K 上下文)的 Qwen3-8B GRPO 训练中,DualKV 实现了 1.63--2.09 倍的策略更新加速,支持 2 倍更大的微批次,并将 MFU 从 36% 提高到 76%。 DAPO 也有类似的收益(2.47 倍加速,77% MFU)。在 16xH100 上的 30B MoE 规模上,DualKV 比 FlashAttention(需要 4 路 Ulysses 序列并行性以避免 OOM)实现了 3.82 倍的策略更新和 3.38 倍的端到端步骤加速。 DualKV 还扩展到头部维度 512(FA2 不支持)的混合滑动/全局注意力,并与 Ulysses 序列并行性集成,在 64K 环境下的 Gemma-4-31B GRPO 上进行了演示。