@Hass_Abdallah11: 我们有两篇新的 Colfax 文章,关于优化 FlashAttention-4,其中包括我(解码)和 Jack Carlisle(反向传播)的工作…

X AI KOLs Timeline 工具

摘要

文章详细介绍了 FlashAttention-4 的优化技术,包括用于解码的 S/P 乒乓方法以重叠操作,在 NVIDIA Blackwell GPU 上实现了高达 16% 的性能提升。

我们有两篇新的 Colfax 文章,关于优化 FlashAttention-4,其中包括我(解码)和 Jack Carlisle(反向传播)的工作:https://research.colfax-intl.com/optimization-diaries-s-p-ping-pong-for-flashattention-4-decode/… https://research.colfax-intl.com/optimization-diaries-improving-flashattention-4-backward-for-head-dimension-64/… 对于解码,使用 S/P 乒乓方法将 KV 块 i 的 softmax 与块 i+1 的 QK^T 重叠:
查看原文
查看缓存全文

缓存时间: 2026/09/09 03:44

优化日志:面向FlashAttention-4解码的S/P乒乓调度 - Colfax Research

来源:https://research.colfax-intl.com/optimization-diaries-s-p-ping-pong-for-flashattention-4-decode/

大语言模型推理分为预填充阶段和解码阶段。在预填充期间,模型处理大量输入令牌并填充键值缓存。在解码阶段,模型利用缓存的键和值自回归地一次生成一个或几个新令牌。

本文讨论针对NVIDIA Blackwell GPU上FlashAttention-4解码阶段的优化方案。当前FA4解码过程中,即使第i+1个KV块的QK^T矩阵乘法与第i个KV块的softmax运算之间不存在数学依赖关系,前者也需等待后者完成才能执行。为实现二者并行,我们可利用解码路径中存在的空闲张量存储器槽位。该空闲TMEM用于在两个槽位间以乒乓方式切换S/P计算——当softmax线程组将第i块的输出写入一个槽位时,矩阵乘法线程组可将第i+1块的QK^T计算任务写入另一个槽位。

此优化在支持的单令牌/多令牌解码配置(头维度64和128)中最高可带来16%的性能提升。相关代码已提交至FlashAttention仓库的PR #2817 (https://github.com/Dao-AILab/flash-attention/pull/2817)。

FA4前向传播回顾

设Q、K、V分别表示查询、键、值矩阵。注意力前向传播计算输出矩阵O的过程如下: S = (1/√d)QK^T, P = softmax(S), O = PV 其中softmax按行逐行计算。实际实现中,FA4不会生成完整的S或P矩阵,而是逐分块处理。因此softmax以“在线”方式计算,必要时对O进行重新缩放。

为实现前向传播,FA4采用跨越5种不同线程组角色(加载、矩阵乘法、softmax、校正和尾声)的流水线交叠架构。加载线程组将Q、K、V分块从全局存储器拷贝至共享存储器。矩阵乘法线程组使用加载的Q和K执行S=QK^T计算供softmax线程组使用。softmax线程组生成P并更新在线softmax统计量。校正线程组根据统计量按需重新缩放O。随后矩阵乘法线程组使用P和V执行PV计算。

后续讨论中将省略K的转置符号。预填充阶段每个CTA被分配两个128行Q分块:高地址Q块和低地址Q块,用于重叠softmax(S^H)与Q^L K计算。此方案中S^H存储于TMEM列0-127,S^L存储于列128-255,这两个槽位后续复用存储P^H和P^L。图1摘自FA4预印本论文,展示了该调度方案。

图1. 每个CTA分配两个Q分块,展示预填充期间两个矩阵乘法及其softmax的调度关系。 单令牌/多令牌解码通常只使用单个(通常带填充的)128行Q分块。此时图1的操作重叠方案不再适用。尽管如此,TMEM列128-255仍被分配但保持空闲。记QK(i)和S(i)分别为第i个KV块的QK矩阵乘法及其得分矩阵。仅存在单个Q分块时,QK(i)、softmax(S(i))和QK(i+1)将串行执行。因此softmax的延迟无法被隐藏,即便QK(i+1)并不依赖softmax(S(i))。我们实现了替代的并行策略来解决此问题。

S/P乒乓调度

为简化表述,将TMEM列0-127称为槽位0,列128-255称为槽位1。图2展示了基础路径初始阶段及主循环前几次迭代中各操作数在槽位中的分布情况。

图2. 基础路径中操作数在槽位0/1的分布状态。S分块为蓝色,P分块为黄色,灰色表示空闲。 乒乓调度通过利用槽位1的空闲TMEM实现QK(i+1)与softmax(S(i))的重叠。具体而言,矩阵乘法线程组写入S和softmax线程组写入P的目标槽位将如图3所示交替切换:

图3. S/P乒乓路径中操作数在槽位0/1的分布状态。 新的执行顺序如下:

图4. S/P乒乓路径的矩阵乘法/softmax调度时序。 注:图2中单元格宽度不代表实际执行时间。

图5总结了与乒乓调度相关的矩阵乘法、加载、softmax和校正线程组间的部分同步关系。虚线箭头表示TMEM操作数被消耗。O累加器使用单缓冲区,此处绘制两次仅为视觉清晰。虽然校正线程组每轮都显示执行“重缩放”,但实际是否执行取决于行最大值的变化程度。

图5. S/P乒乓调度的流水线与同步概览。虚线箭头表示TMEM操作数被消耗。两个O表示同一累加缓冲区。

实现细节

实现乒乓调度需要精细协调以确保线程组等待/消耗正确的槽位。原代码仅需单个相位位:单S槽位时存在单个屏障,每个块翻转一次相位。采用双S槽位后,每个槽位的屏障每两个块翻转一次。因此需跟踪屏障选择及该屏障已翻转次数。我们使用一对屏障和两个位:全局PVPV矩阵乘法计数(mma_pv_count)的位0和位1。位0选择屏障,位1记录(模2)给定槽位屏障的翻转次数。

初始化阶段执行流程如下:

  1. 加载线程组将Q0和K0拷贝至共享存储器。
  2. 矩阵乘法线程组发出QK(0)计算。
  3. 加载线程组将K1拷贝至共享存储器。

关键变化在于KV加载顺序。原流程按K0,V0,K1,V1,K2,V2…顺序生成K和V分块。乒乓路径按K0,K1,V0,K2,V1,K3,V2…顺序加载,最终追加最后一个V分块。加载第二个K分块是为了使矩阵乘法线程组进入主循环后可立即发出QK(1)。初始化阶段代码示例:

if const_expr(self.use_s_ping_pong):
    # 等待Q就绪
    pipeline_q.consumer_wait_w_index_phase(0, mma_q_consumer_phase)
    # 等待K(0)就绪
    pipeline_kv.consumer_wait(mma_kv_consumer_state)
    Ki_index, Ki_phase = mma_kv_consumer_state.index, mma_kv_consumer_state.phase
    sK_cur = sK[None, None, None, Ki_index]
    if const_expr(self.uneven_kv_smem):
        sK_cur = self.offset_kv_smem(sK_cur, Ki_index, Ki_phase)
    # 发出QK(0)计算
    if (mma_pv_count & 1) == 0:
        gemm_Si[0](smem_desc_start_b=sm100_desc.make_smem_desc_start_addr(sK_cur.iterator))
        pipeline_s_p_o.producer_commit_w_index(0)
    else:
        gemm_Si[1](smem_desc_start_b=sm100_desc.make_smem_desc_start_addr(sK_cur.iterator))
        pipeline_s_p_o.producer_commit_w_index(1)
    mma_q_consumer_phase ^= 1
    # 释放K(0)
    pipeline_kv.consumer_release(mma_kv_consumer_state)
    # 推进至K(1)
    mma_kv_consumer_state.advance() 
    O_should_accumulate = False

变量mma_pv_count是全局PVPV矩阵乘法计数。调用gemm_Si[0]指示QK(0)写入槽位0。pipeline_s_p_o.producer_commit_w_index(0)通知softmax线程组等待QK(0)完成后方可从槽位0消耗S(0)。类似地,gemm_Si[1]pipeline_s_p_o.producer_commit_w_index(1)对应槽位1。

乒乓主循环执行流程:

for i in cutlass.range(block_iter_count - 1, unroll=1):
    # 等待K(i+1)就绪
    pipeline_kv.consumer_wait(mma_kv_consumer_state)
    Ki_index, Ki_phase = mma_kv_consumer_state.index, mma_kv_consumer_state.phase
    sK_cur = sK[None, None, None, Ki_index]
    if const_expr(self.uneven_kv_smem):
        sK_cur = self.offset_kv_smem(sK_cur, Ki_index, Ki_phase)
    # 发出QK(i+1)。全局块计数为偶数时写入TMEM槽位0(0-127),奇数时写入槽位1(128-255)
    if ((mma_pv_count + 1) & 1) == 0:
        gemm_Si[0](smem_desc_start_b=sm100_desc.make_smem_desc_start_addr(sK_cur.iterator))
        pipeline_s_p_o.producer_commit_w_index(0)
    else:
        gemm_Si[1](smem_desc_start_b=sm100_desc.make_smem_desc_start_addr(sK_cur.iterator))
        pipeline_s_p_o.producer_commit_w_index(1)
    # 释放K(i+1)
    pipeline_kv.consumer_release(mma_kv_consumer_state)
    # 推进至V(i)
    mma_kv_consumer_state.advance()
    # 等待V(i)就绪
    pipeline_kv.consumer_wait(mma_kv_consumer_state)
    Vi_index, Vi_phase = mma_kv_consumer_state.index, mma_kv_consumer_state.phase
    tOrVi = tOrV[None, None, None, Vi_index]
    sV_cur = sV[None, None, None, Vi_index]
    if const_expr(self.uneven_kv_smem):
        sV_cur = self.offset_kv_smem(sV_cur, Vi_index, Vi_phase)
    # 当前块槽位的相位。每个槽位每两个块复用一次,复用时屏障相位翻转。
    pv_phase = (mma_pv_count >> 1) & 1
    # 发出PV(i)计算
    if (mma_pv_count & 1) == 0:
        pipeline_s_p_o.producer_acquire_w_index_phase(0, pv_phase)
        gemm_Pi[0](
            tCrB=tOrVi,
            sB=sV_cur,
            zero_init=not O_should_accumulate,
            mbar_ptr=pipeline_p_lastsplit.sync_object_full.get_barrier(0) if self.split_P_arrive > 0 else None,
            mbar_phase=pv_phase,
        )
    else:
        pipeline_s_p_o.producer_acquire_w_index_phase(1, pv_phase)
        gemm_Pi[1](
            tCrB=tOrVi,
            sB=sV_cur,
            zero_init=not O_should_accumulate,
            mbar_ptr=pipeline_p_lastsplit.sync_object_full.get_barrier(1) if self.split_P_arrive > 0 else None,
            mbar_phase=pv_phase,
        )
    pipeline_o_acc.producer_commit_w_index(mma_pv_count & 1)
    mma_pv_count += 1
    # 释放V(i)
    pipeline_kv.consumer_release(mma_kv_consumer_state)
    # 推进至K(i+2)
    mma_kv_consumer_state.advance() 
    O_should_accumulate = True

补充说明代码注释:QK矩阵乘法写入目标槽位取决于mma_pv_count + 1的奇偶性,因其执行比PVPV提前一步。pv_phase = (mma_pv_count >> 1) & 1提取mma_pv_count的位1,即槽位mma_pv_count & 1的复用次数奇偶性。图6总结了前几次迭代中PV写入槽位和pv_phasemma_pv_count的变化关系。

图6. 给定mma_pv_count对应的PV槽位及pv_phase值。 pipeline_s_p_o.producer_acquire_w_index_phase(0/1, pv_phase)调用具有双重作用:

  1. 等待softmax线程组完成P(i)生产并释放槽位。
  2. 等待校正线程组完成所有必要的O重缩放并释放槽位。

pipeline_o_acc.producer_commit_w_index(mma_pv_count & 1)调用表示矩阵乘法线程组已完成当前迭代PVPV在O累加缓冲区的累加,校正线程组可安全消耗。此操作在基础路径中非必需,因串行化保证了PV(i)在S(i+1)可能的校正重缩放前完成。

IKET性能分析

NVIDIA内核事件跟踪工具可观察线程组在内核生命周期中的活动。虽然我们的修改理论上允许矩阵乘法线程组在softmax线程组处理S(i)时发出QK(i+1),但IKET可验证此调度变更确实发生。我们提供基础路径与乒乓路径的跟踪对比。基础路径跟踪结果:

图7. FA4基础路径解码期间加载、矩阵乘法和softmax线程组的指令发射顺序。 相关序列为:

  1. QK矩阵乘法(mma_issue_QK
  2. softmax(S)(sm_compute
  3. PV矩阵乘法(mma_issue_PV
  4. QK矩阵乘法(mma_issue_QK) 这些操作的时间条几乎无重叠,实际呈串行执行。作为对比,乒乓路径跟踪结果:

图8. FA4 S/P乒乓路径解码期间加载、矩阵乘法和softmax线程组的指令发射顺序。 注意蓝色mma_issue_QK条与紫色sm_compute条存在显著重叠。这确凿证明乒乓路径实现了QK矩阵乘法与softmax的并行执行。

性能表现

解码是内存密集型任务,我们以实际达到的内存带宽作为性能指标。图9展示了在NVIDIA B200 Blackwell GPU上对部分配置(包括分组查询注意力比率H_q:H_kv=16:1和16:2)的基准测试结果。较短KV序列长度时收益较小,但随序列增长显著扩大。此现象与优化针对稳态场景的设计一致——较长序列暴露出更多可受益于QK与softmax重叠的迭代。头维度64时,16:2比率带宽提升达15.6%,16:1比率达16.0%。头维度128时,16:1比率提升9.3%,16:2比率提升14.9%。图10对比了固定序列长度128k下所有查询长度的单令牌与多令牌解码性能。

图9. NVIDIA B200 Blackwell GPU上基础路径(蓝色)与乒乓路径(橙色)的单令牌FA4解码基准测试结果。 图10. NVIDIA B200 Blackwell GPU上基础路径(蓝色)与乒乓路径(橙色)的单/多令牌FA4解码基准测试结果(固定序列长度128k)。

总结

本文讨论了FA4解码优化方案:通过S/P在两个TMEM槽位间乒乓调度(其中一个此前闲置),消除了QK与softmax间的非必要串行化。我们阐述了预填充并行策略如何不适用于解码场景,详解了乒乓调度实现细节,并通过IKET跟踪验证了预期的重叠效果。基准测试显示性能最高提升达16%。

相似文章