中文

DualKV:用于高效 RL 训练的共享提示 Flash Attention

机器学习 2026-05-29 v2

摘要

现代 RL 后训练方法,如 GRPO 和 DAPO,在 NN 个来自 PP 个 token 长度提示序列的 RR 个 token 响应序列上进行训练,但标准 FlashAttention 在前向和反向传递中复制所有 PP 个提示 token NN 次——在相同的隐藏状态上复制计算和内存。在大规模 rollout、长上下文 RL 训练(N16N{\geq}16P8KP{\geq}8\text{K})中,这种冗余占据了策略更新成本的主导地位。我们观察到在解码器模型中,因果掩码使提示表示在每一层对所有序列不变,因此所有 per-token 操作(规范化、投影、MLP)和注意力都可以只处理一次提示——这一属性尚未在训练时被用于内核层面。我们提出 \textbf{DualKV},这是一种首次消除 RL 训练中共享提示复制的 FlashAttention 内核变体,通过 (1) 对用 CUDA 融合前向和反向内核,遍历两个不相交的 KV 区域——共享上下文和 per-sequence 响应——在单个内核启动中,以及 (2) veRL 中的数据管道重新打包,将 N(P+R)N(P{+}R) 个 token 重新组织为 P+NRP{+}NR 个 token 每个 micro-batch,将注意力之外的整个模型的 token 减少率 ρ=N(P+R)/(P+NR)\rho = N(P{+}R)/(P{+}NR)。DualKV 在数学上等价于标准注意力,未引入任何近似。在 Qwen3-8B GRPO 训练中使用 8\times H100 GPU(N=32N{=}32,8K 上下文),DualKV 实现 1.631.63--2.09×2.09\times 策略更新加速,使 micro-batch 大小提升 2×2\times,MFU 从 36%36\% 提升至 76%76\%。DAPO 类似地获得 2.47×2.47\times 加速,77%77\% MFU。在 30B MoE 规模使用 16\times H100 时,DualKV 实现 3.82×3.82\times 策略更新和 3.38×3.38\times 端到端 step 加速(FlashAttention 需要 4 路 Ulysses 序列并行以避免 OOM)。

关键词

引用

@article{arxiv.2605.15422,
  title  = {DualKV: Shared-Prompt Flash Attention for Efficient RL Training with Large Rollouts and Long Contexts},
  author = {Jiading Gai and Shuai Zhang and Xiang Song and Bernie Wang and George Karypis},
  journal= {arXiv preprint arXiv:2605.15422},
  year   = {2026}
}