DualKV: 针对大规模生成和长上下文的共享提示Flash Attention,用于高效RL训练

arXiv cs.LG 论文

摘要

介绍DualKV,一种FlashAttention内核变体,可消除RL后训练(GRPO/DAPO)中冗余的提示词元计算,在30B MoE模型上实现高达3.82倍的加速。

arXiv:2605.15422v1 公告类型:new \n摘要:现代RL后训练方法(如GRPO和DAPO)基于一个包含$P$个词元的共享提示,采样$N$个响应序列(每个序列$R$个词元)进行训练。但标准FlashAttention会在前向和后向传播中将所有$P$个提示词元重复$N$次——在相同的隐藏状态上重复计算和内存开销。在大规模生成、长上下文的RL训练中($N{\geq}16$,$P{\geq}8\text{K}$),这种冗余主导了策略更新的成本。我们观察到,在仅解码器模型中,因果掩码使得提示表示在每一层的所有序列中保持不变,因此所有逐词元操作(归一化、投影、MLP)和注意力机制都可以只处理一次提示——这一特性此前未在训练的内核层面被利用。我们提出 \textbf{DualKV},这是首个消除RL训练中共享提示重复的FlashAttention内核变体,通过:(1)~融合的CUDA前向和后向内核,在单次内核启动中迭代两个不相交的KV区域——共享上下文和每个序列的响应;(2)~veRL中的数据管线重新设计,将每个微批中的$N(P{+}R)$个词元重新打包为$P{+}NR$个词元,将词元减少从注意力扩展到整个模型,因子为$\rho = N(P{+}R)/(P{+}NR)$。DualKV在数学上与标准注意力等价,且不引入任何近似。在基于8$\times$H100 GPU的Qwen3-8B GRPO训练中($N{=}32$,8K上下文),DualKV实现了$1.63$--$2.09\times$的策略更新加速,支持$2\times$更大的微批,并将MFU从$36\%$提升至$76\%$。相似的增益在DAPO上也有体现($2.47\times$加速,$77\%$ MFU)。在16$\times$H100上的30B MoE规模下,DualKV在策略更新上实现了$3.82\times$加速,端到端步骤加速达到$3.38\times$,优于FlashAttention(后者需要4路Ulysses序列并行以避免OOM)。
查看原文
查看缓存全文

缓存时间: 2026/05/18 06:41

# DualKV: 面向大规模rollout与长上下文的高效RL训练——共享提示的Flash注意力机制

来源: https://arxiv.org/html/2605.15422

Jiading Gai¹\*, Shuai Zhang¹\*, Xiang Song², Bernie Wang¹, George Karypis³

¹Amazon Web Services, ²Google, ³明尼苏达大学

\{jiadingg, shuaizs, yuyawang\}@amazon\.com, xiangsx@google\.com, karypis@umn\.edu  
*共同第一作者*

###### 摘要

现代RL后训练方法(如GRPO和DAPO)基于一个长度为 \(P\) 的共享提示(prompt),采样得到 \(N\) 个响应序列(response sequences)共 \(R\) 个词元(tokens)进行训练。标准的FlashAttention在前向和反向传播中将所有 \(P\) 个提示词元重复 \(N\) 次——在相同的隐藏状态上重复计算和存储。在大规模rollout、长上下文的RL训练(\(N \geq 16\), \(P \geq 8\text{K}\))中,这种冗余主导了策略更新的成本。我们观察到,在仅解码器(decoder-only)模型中,因果掩码(causal masking)使得提示表示(representations)在每一层对所有序列都保持不变,因此所有逐词元操作(norm、投影、MLP)和注意力机制都可以只处理提示一次——这一性质此前在训练层面未被内核级实现利用。我们提出 **DualKV**,这是首个在RL训练中消除共享提示重复的FlashAttention内核变体,其核心手段包括:(1) 融合的CUDA前向和反向内核,在单次内核启动中迭代两个不相交的KV区域——共享上下文和逐序列响应;(2) 对veRL的数据流水线进行重新设计,将 \(N(P+R)\) 个词元重新打包为每个微批次(micro-batch) \(P+NR\) 个词元,将词元减少因子 \(\rho = N(P+R)/(P+NR)\) 从注意力层扩展到整个模型。DualKV在数学上与标准注意力等价,且不引入任何近似。在与FlashAttention(需要4路Ulysses序列并行以避免OOM)对比时,在Qwen3-8B的GRPO训练(8×H100 GPU, \(N=32\), 8K上下文)中,DualKV实现了策略更新1.63–2.09倍加速,支持2倍更大的微批次,并将MFU从36%提升至76%。DAPO中也有类似增益(2.47倍加速,77% MFU)。在30B MoE规模、16×H100上,DualKV实现策略更新3.82倍、端到端step 3.38倍加速。

## 1 引言

现代RL后训练方法,如GRPO(Shao et al., 2024 (https://arxiv.org/html/2605.15422#bib.bib20))和DAPO(Yu et al., 2025 (https://arxiv.org/html/2605.15422#bib.bib24)),为 \(N\) 个共享一个长度为 \(P\) 的提示的响应序列计算对数概率和梯度。使用标准的FlashAttention-2(FA2)(Dao, 2024 (https://arxiv.org/html/2605.15422#bib.bib3))时,训练微批次包含 \(N\) 个长度为 \(S_i = P + R_i\) 的序列,将长度为 \(P\) 的提示重复了 \(N\) 次。每步前向和反向传播对每层中 \((N-1) \times P\) 个冗余提示词元重新计算 \(K, V\) 激活值及其梯度。近期工作表明,随着rollout因子不断增大,模型精度持续提升。Lightman等人 (2023 (https://arxiv.org/html/2605.15422#bib.bib12)) 展示了在MATH数据集上,Best-of-N性能与 \(N\) 呈对数线性关系,直至 \(N=1860\)。GRPO及相关方法同样受益于较大的 \(N\) 以获得更精确的奖励信号估计(例如,DAPO使用 \(N=16\))。然而,标准 \(N\) 副本打包(\(N\)-copy packing)的计算开销为 \(O(N \cdot S^2)\):每个 \(N\) 个序列独立地重新计算完整 \(P\) 个提示词元上的注意力,在每层的计算和内存中都引入了 \((N-1) \times P\) 个提示词元的重复。正如我们在第4.2节(https://arxiv.org/html/2605.15422#S4.SS2)中所示,消除这种重复可将策略更新时间削减至原本的 \(2\times\)。

提示KV共享在推理方面已有多个方向的研究:分页注意力(paged attention)(Kwon et al., 2023 (https://arxiv.org/html/2605.15422#bib.bib11))和前缀缓存(prefix caching)通过写时复制(copy-on-write)块表避免冗余的提示KV存储;分叉注意力(bifurcated attention)(Athiwaratkun et al., 2024 (https://arxiv.org/html/2605.15422#bib.bib1))将注意力分解为共享提示和逐序列两个阶段以实现并行解码。然而,这些机制仅针对推理场景,其中共享提示KV是只读的;它们没有训练反向传播,没有兼容自动求导(autograd)的共享提示梯度累积,也无法直接嵌入RL策略更新流水线(附录I (https://arxiv.org/html/2605.15422#A9))。训练带来了三个新挑战:(1) 反向传播中有 \(N\) 个序列并发地将梯度累积到共享KV缓冲区——这种模式在推理中不存在,需要原子累加以保证无竞态写入,以及fp32累加器结合最终类型转换以匹配FA2的逐元素精度;(2) 自动求导必须正确聚合来自上下文自注意力(context self-attention)和解码注意力(decoded attention)两处调用的梯度,因为它们都涉及共享KV;(3) RL训练框架的数据流水线必须重组,以便在每个微批次内将同一提示下的多个响应归组(附录B (https://arxiv.org/html/2605.15422#A2))。DualKV在内核和系统层面同时解决了这三个挑战。

我们观察到,在仅解码器模型中,提示隐藏状态在每一层对所有 \(N\) 个序列都是**相同**的:因果掩码确保每个提示词元仅关注之前的提示词元,而不关注其后不同的响应词元,因此提示的表示与生成了哪个响应无关。Prefix Grouper(Liu et al., 2025 (https://arxiv.org/html/2605.15422#bib.bib17))在框架层面应用了这一性质,但未提供内核支持——它仍然将通过标准FA2传递完整的 \(N\) 副本KV,保留了 \(O(N \cdot P \cdot d)\) 的内存瓶颈,导致OOM并迫使使用序列并行。DualKV贡献了首个消除这种重复的FlashAttention内核变体,使用融合的前向和反向CUDA内核,从单个物理缓冲区读取共享KV,并通过fp32原子写累积来自 \(N\) 个并发序列的梯度。通过将微批次打包为单个提示副本,并将注意力分解为上下文自注意力(计算一次)和融合的DualKV内核(用于解码注意力),DualKV在整个模型中消除了冗余的提示计算——不仅限于注意力,还包括norm、投影和MLP层——从而在无需近似的情况下同时节省内存和计算。

## 2 消除RL训练中的提示重复:DualKV设计

每个GRPO训练步骤处理 \(N\) 个长度为 \(S_i = P + R_i\) 的响应序列,这些序列共享一个长度为 \(P\) 的提示。该步骤在优化器步骤之前运行三遍前向/反向:(1) `old_log_prob`:通过 \(\pi_{\theta_{\text{old}}}\) 前向计算重要性比例;(2) `ref_log_prob`:通过 \(\pi_{\text{ref}}\) 前向计算KL惩罚;(3) 策略更新:通过 \(\pi_{\theta}\) 前向+反向产生策略梯度。这三遍都看到相同的打包微批次,其中包含 \((N-1) \cdot P\) 个冗余提示词元,因此DualKV可加速全部三遍。该方案同样适用于任何将 \(N\) 个同一提示下的响应进行打包的RL方法,例如DAPO。

### 2.1 基线:使用副本提示的标准打包

标准实现(例如veRL配合FA2)将 \(N\) 个序列打包到一个 `flash_attn_varlen_func` 调用中,使用累积序列长度:
\[\underbrace{[P, R_1]}_{\text{seq 1}},\; \underbrace{[P, R_2]}_{\text{seq 2}},\; \ldots,\; \underbrace{[P, R_N]}_{\text{seq }N} \quad \text{--- 总词元数 } T_{\text{std}} = N(P+R) \tag{1}\]
提示词元被**重复**了 \(N\) 次:每个序列携带自己的提示隐藏状态、QKV投影和注意力计算副本。每个逐词元操作(Norm、QKV投影、RoPE、MLP、输出投影)处理全部 \(N(P+R)\) 个词元。对于注意力,FA2计算 \(N\) 个独立的因果自注意力,每个序列一个:
\[O^{(i)} = \text{softmax}\!\left(\frac{Q^{(i)}(K^{(i)})^\top}{\sqrt{d}} + M_{\text{causal}}\right) V^{(i)},\qquad i=1,\ldots,N, \tag{2}\]
其中 \(Q^{(i)}, K^{(i)}, V^{(i)} \in \mathbb{R}^{(P+R_i) \times H \times d}\) 是序列 \(i\) 的打包投影,\(M_{\text{causal}}\) 是逐序列的因果掩码。FA2通过分块在线softmax计算公式(2)。FLOPs: \(O(N \cdot S_i^2 \cdot H \cdot d)\)。

### 2.2 DualKV:单提示打包

**提示不变性与冗余性。** 在公式(2)中,将输出 \(O^{(i)}\) 的行分解为提示查询行 \(O_P^{(i)}\)(前 \(P\) 行)和响应查询行 \(O_{R_i}^{(i)}\)(剩余 \(R_i\) 行)。对于提示查询行,因果掩码阻止关注响应键,因此只有提示键参与:
\[O_P = \text{softmax}\!\left(\frac{Q_P (K_P)^\top}{\sqrt{d}} + M_{\text{causal},P}\right) V_P, \tag{3}\]
其中 \(Q_P, K_P, V_P\) 是提示行的投影。由于所有 \(N\) 个序列共享相同的提示词元,因此 \(Q_P, K_P, V_P\) 在所有 \(i\) 上相同(附录A.2 (https://arxiv.org/html/2605.15422#A1.SS2)),所以 \(O_P\) 对每个序列都是一样的。基线中 \(N-1\) 次提示-提示计算因此是冗余的。对于响应查询行,每个响应因果地关注**既有**共享提示键也有自身的响应键:
\[O_{R_i}^{(i)} = \text{softmax}\!\left(\frac{Q_{R_i}^{(i)} [K_P;\, K_{R_i}^{(i)}]^\top}{\sqrt{d}} + M_{\text{causal}}\right) [V_P;\, V_{R_i}^{(i)}], \tag{4}\]
其中 \([\cdot; \cdot]\) 表示沿词元轴拼接。这一块并不冗余——每个响应具有不同的查询——但所有 \(i\) 都重用相同的 \(K_P, V_P\)。

**DualKV打包。** DualKV将微批次打包为一个共享提示后接 \(N\) 个逐序列响应:
\[\underbrace{[P]}_{\text{prompt (一次)}},\; \underbrace{[R_1]}_{\text{resp 1}},\; \ldots,\; \underbrace{[R_N]}_{\text{resp }N} \quad \text{--- 总词元数 } T_{\text{dk}} = P + NR \text{ 词元} \tag{5}\]
所有逐词元操作(norm、投影、RoPE、MLP、输出投影)现在只处理 \(P+NR\) 个词元,而不是 \(N(P+R)\)——在整个模型中节省了 \((N-1)P\) 个词元。

**两调用分解。** 公式(3)和公式(4)可以分别计算以消除冗余,同时保留精确注意力:
- 调用1——上下文自注意力(`flash_attn_varlen_func`,标准FA2):对单份提示副本**一次**计算公式(3)。FLOPs: \(O(P^2 \cdot H \cdot d)\)。
- 调用2——解码注意力(`flash_attn_dualkv_varlen_func`,DualKV内核):对所有 \(N\) 个响应计算公式(4),关注共享的 \(K_P, V_P\)(来自调用1)加上逐序列的 \(K_{R_i}, V_{R_i}\)。FLOPs: \(O(N \cdot R \cdot S \cdot H \cdot d)\)。

**为什么需要新内核。** 公式(4)的注意力模式——响应查询从拼接 \([K_P; K_{R_i}]\) 中读取——无法被标准FA2直接消费,除非将 \(K_P, V_P\) 实例化为 \(N\) 个副本(这会重新引入我们刚刚消除的重复)。DualKV内核在一次启动中迭代两个物理上不相交的KV区域(\(K_P, V_P\) 在批次内共享;\(K_{R_i}, V_{R_i}\) 逐序列)。内核接口和算法详见第3节(https://arxiv.org/html/2605.15422#S3)。

**数据流水线。** 这种打包要求来自同一提示的所有 \(N\) 个响应位于同一GPU的同一微批次内。veRL的rollout引擎已经生成提示连续的序列;DualKV通过跳过veRL的 `balance_batch` 步骤(该步骤按序列长度重新排序)并禁用小批次迭代器中的epoch内混洗来保持这一顺序。生成的流水线保留了精确的小批次梯度估计(附录B (https://arxiv.org/html/2605.15422#A2))。当微批次包含多个提示组时,DualKV为每个组独立处理,拥有各自的上下文KV和逐组的 `cu_seqlens`,支持单次前向传递中的异构提示长度。

**梯度正确性。** 共享的 \(K_P, V_P\) 接收来自两次调用的梯度——自动求导将两个贡献求和,而DualKV内核在调用2内部累加 \(N\) 个逐序列项:
\[\frac{\partial \mathcal{L}}{\partial K_P} = \left(\frac{\partial \mathcal{L}}{\partial K_P}\right)_{\text{Call 1}} + \sum_{i=1}^N \left(\frac{\partial \mathcal{L}}{\partial K_P}\right)^{(i)}_{\text{Call 2}} \tag{6}\]
(\(V_P\) 同理)。第一项是调用1的提示自注意力梯度;\(\sum_{i=1}^N\) 收集调用2中每个响应的贡献,形成了对共享 \(K_P, V_P\) 的 \(N\) 路并发写入——这种模式在FA2的反向中没有对应物,在第3.2节(https://arxiv.org/html/2605.15422#S3.SS2)的内核层面解决。总和在数学上与FA2基线相同(附录A.5 (https://arxiv.org/html/2605.15422#A1.SS5))。

**DualKV何时有效。** 词元减少比率 \(\rho = N(P+R)/(P+NR)\) 决定了整个模型中逐词元操作(norm、投影、RoPE、MLP、输出投影)的加速倍数;注意力还有额外的结构性节省(附录C (https://arxiv.org/html/2605.15422#A3))。\(\rho\) 在两种情况下最大:**大规模rollout**(\(N \geq 16\),常见于GRPO和DAPO;当 \(N=32, P=16\text{K}\) 时,减少倍数达 \(7.2\times\))和**长提示**(\(P \geq 16\text{K}\),常见于智能体任务(Jimenez et al., 2024 (https://arxiv.org/html/2605.15422#bib.bib9))和仓库级代码生成(Liu et al., 2023 (https://arxiv.org/html/2605.15422#bib.bib16));当 \(P=64\text{K}, N=16\) 时,减少倍数达 \(14.3\times\))。实践中,\(\rho\) 受到每GPU微批次大小的限制(第4.2节 (https://arxiv.org/html/2605.15422#S4.SS2))。完整分析见附录C.3 (https://arxiv.org/html/2605.15422#A3.SS3)。

## 3 新计算原语:DualKV内核

本节描述执行两调用分解(第2.2节 (https://arxiv.org/html/2605.15422#S2.SS2))中调用2的DualKV CUDA内核。该内核作为FA2代码库的扩展实现,接受五个输入张量——解码查询 \(q\)、共享上下文 \(K_c, V_c\)(单副本,仅存储一次)、以及逐序列的变长打包解码 \(K_d, V_d\),外加一个参数 `context_seqlen = P`,用于移动因果掩码,使得解码查询关注所有 \(P\) 个上下文键及其自身之前的解码键。我们遵循FlashAttention的命名法,用 \(K_c, V_c\) 和 \(K_d, V_d\) 分别表示 \(K_P, V_P\) 和 \(K_{R_i}, V_{R_i}\)。内核的接口和算法在下文详述。

(注:原文第3节在此处中断,翻译已覆盖至该位置。后续内容(第3.1节等)若存在,用户可继续提供。)

相似文章

SparDA:用于高效长上下文 LLM 推理的稀疏解耦注意力

arXiv cs.CL

SparDA 提出了一种解耦稀疏注意力架构,通过添加轻量级"Forecast"投影来预测未来的 KV 缓存需求,从而实现从 CPU 到 GPU 的预取(lookahead prefetching),并降低选择开销。在基于稀疏预训练的 8B 模型上,其 prefill 速度最高可提升 1.25×,decode 速度最高可提升 1.7×,相比非 offload 基线,decode 吞吐量最高可提升 5.3×。