CUDA 算子工程:手写 FlashAttention v2 之路

第 18 章 Persistent Kernel 与 Producer-Consumer

作者 杨艺韬 · 4,504 字 · 发布于 · 更新于

经典模式是"起足够多的 block 把 GPU 填满",persistent 模式是"只起 SM 数量那么多 block,让每个 block 连着做很多 tile"。 这一章讲这个反转解决的两个真问题:launch 开销,和最后一波只干一点活的波次量化。

18.1 经典 Kernel vs Persistent Kernel

到目前为止,本专栏所有 kernel 都是"经典"模式:

// 经典模式: grid_size = N_tiles, 每 block 一个 tile
dim3 grid(M_tiles, N_tiles);
my_kernel<<<grid, block>>>(...);

GPU 调度器把 M_tiles × N_tiles 个 block 自动分配到 SM 上。如果 block 数 >> SM 数,调度器自动队列、循环执行。

这个模式简单优雅,对大 tile 工作负载(比如训练时的大 GEMM)非常合适。但对小 tile 工作负载有几个隐藏问题:

  1. Block 启动开销:每个 block 上 SM 之前要走 GigaThread 引擎的派发、寄存器与 SMEM 的分配,还有 block 开头的初始化(算地址、建 barrier、预取描述符之类)。摊在一个"干活很久"的大 tile 上可以忽略;但 tile 又多又小(LLM decoding 时 batch_size=1、seq_len=1 就是这种形态)时,这笔开销在总时间里的占比就不能忽略了。
  2. L2 cache 不复用:每个 block 独立运行,相邻 block 的数据 L2 cache 命中无法保证。
  3. 调度延迟:block 数远超 SM 数时,调度器的全局调度有延迟。

Persistent Kernel 的思路:grid_size 固定为 SM 数(或其倍数),每个 block 通过 grid-stride loop 处理多个 tile:

// Persistent 模式
int n_sms = 132;  // H100
dim3 grid(n_sms);
persistent_kernel<<<grid, block>>>(...);

__global__ void persistent_kernel(int total_tiles, ...) {
    __shared__ int s_tile;
    while (true) {
        // 用 atomic counter 抢下一个 tile:只让 thread 0 抢,经 SMEM 广播给全 block
        // (若每个线程各自 atomicAdd,同一 block 的线程会拿到各不相同的 tile)
        if (threadIdx.x == 0) s_tile = atomicAdd(&global_counter, 1);
        __syncthreads();
        int tile_id = s_tile;
        if (tile_id >= total_tiles) break;

        process_tile(tile_id, ...);
        __syncthreads();  // 全员读完 s_tile 后 thread 0 才能写下一个
    }
}

每个 block 启动一次,活到所有 tile 都做完才退出。global_counter 每次 launch 前都要清零——FA3 在非 varlen 路径上就是每次调用前把它的 tile_count_semaphore 置零(flash-attn/hopper/flash_api.cpp:645)。

18.2 Persistent Kernel 的优势

18.2.1 摊薄启动开销

如果总 tile 数 = N,传统 kernel 要派发 N 个 block;persistent kernel 只派发 132 个(SM 数),剩下的工作由这 132 个 block 在 kernel 内部循环领取。派发次数降为原来的 132/N。

别把这笔账算成串行相加。 常见的错误估法是"(10000 − 132) × 每 block 开销",把它当成一段实打实的墙钟时间——但 GigaThread 引擎是并行派发的,一个 block 的派发开销与其他 block 的执行是重叠的。persistent 真正省下的不是这个乘积,而是尾部的调度抖动和反复的资源分配/回收。所以这一节该记住的是方向(tile 越小越碎,persistent 越划算),不是某个具体的毫秒数。

18.2.2 数据复用:在 L2,不在 SMEM

这里要先说清一个常见误解:K/V tile 并不会跨 tile 留在 SMEM 里复用。attention 中相邻 query block 确实共享同一批 K/V,但一个 query block 要扫过整条序列的 K/V,而 SMEM 只放得下流水线的几个 stage,是个环形缓冲——扫到下一个 query tile 时,前面的 K/V 早被覆盖了。FA3 的 producer 对每个 work tile 都重新沿 n_block 把 K/V 用 TMA 载入这几个 stage(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:749 的 load_K,每次 pipeline_k.producer_acquire 占一个 stage);CUTLASS 的 PersistentTileSchedulerSm90 也只负责给出下一个 (m, n) 坐标,A/B 同样按 K 维逐 stage 重新加载。persistent 在片上真正省下的是每个 tile 的"开场":barrier 与流水线状态只初始化一次、TMA 描述符只预取一次,并且 producer 可以在 consumer 做上一个 tile 的 epilogue 时就开始载入下一个 tile(FA3 的 DynamicPersistentTileScheduler 用 prefetch_next_work 提前一个 tile 抢号,flash-attn/hopper/tile_scheduler.hpp:336-340)。

K/V 的复用发生在 L2:persistent 模式下同时在跑的 block 数固定,tile 的执行顺序由 scheduler 决定,可以有意让同一时刻在跑的 tile 共享 K/V。FA3 的 DynamicPersistentTileScheduler 就按 L2 容量(代码里按 32 MB 给 K/V 预留)把 head × batch 分成若干 section,一个 section 的 K/V 放得进 L2,再在 section 内排 tile(flash-attn/hopper/tile_scheduler.hpp:226-229 的注释与 :254-260 的 swizzle 计算);载入 K/V 时还带 EVICT_LAST 的 cache hint(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:753、:767)。经典模式下 block 的上机顺序由硬件决定,做不到这种有意的编排。

18.2.3 更友好的负载均衡

如果 tiles 之间的工作量不均衡(比如 causal 或 sparse attention,有的 tile 几乎是空的),经典 kernel 其实也有一层动态均衡:某个 block 做完,硬件就把排队的下一个 block 派到这个 SM 上。它管不了的是顺序——重 tile 若恰好排在最后,就会拖出一条长尾。Persistent kernel 用 atomic counter 抢任务,既保留了"完成快的 block 多抢几个"的动态均衡,又能由 scheduler 决定先发哪些 tile。FA3 在 causal/local 时用的 DynamicPersistentTileScheduler 就是按"最长任务优先"(longest-processing-time-first)把最重的 m_block 排在前面(flash-attn/hopper/tile_scheduler.hpp:222-224、:309-310);非 causal、非 varlen 时各 tile 等重,用的是无 atomic 的 StaticPersistentTileScheduler(flash-attn/hopper/flash_fwd_launch_template.h:63-69)。

18.2.4 CUDA Graph 友好

LLM 推理的核心优化之一是 CUDA Graph——把多个 kernel launch 录成一个 graph,整体提交,省去重复 launch 开销。但录制时每个 kernel 的 grid 维度和参数都被固定下来,形状一变就得更新图里的节点参数或换一张图(附录 A.3)。Persistent kernel 的 grid 永远 = SM 数,不随问题规模变;只要再把 tile 总数之类的形状信息放在显存里由 kernel 读取,而不是作为值参数录进图里,同一张图就能覆盖不同的规模。

18.3 Tile Scheduler 设计

Persistent kernel 的核心是 tile scheduler——决定哪个 block 处理哪些 tile。最简单的 scheduler 是 atomic counter:

__device__ int next_tile() {
    __shared__ int s_tile;
    __syncthreads();  // 等上一轮全员读完 s_tile,再让 thread 0 覆盖
    if (threadIdx.x == 0) {
        s_tile = atomicAdd(&global_counter, 1);
    }
    __syncthreads();
    return s_tile;
}

但 atomic 有竞争开销。更高级的 scheduler:

18.3.1 Static Scheduler

在 host 端预先分配每个 block 处理哪些 tile:

// Host 端
int tiles_per_sm = (total_tiles + n_sms - 1) / n_sms;

// Kernel 端
__global__ void kernel(int tiles_per_sm, ...) {
    int my_first = blockIdx.x * tiles_per_sm;
    int my_last = min((blockIdx.x + 1) * tiles_per_sm, total_tiles);
    for (int t = my_first; t < my_last; ++t) {
        process_tile(t, ...);
    }
}

简单、无竞争,但负载不均衡时差。

18.3.2 Round-Robin Scheduler

__global__ void kernel(...) {
    for (int t = blockIdx.x; t < total_tiles; t += gridDim.x) {
        process_tile(t, ...);
    }
}

也是无竞争,且天然循环——所有 block 平均分担。CUTLASS 3.x 起的 Sm90 默认 scheduler(PersistentTileSchedulerSm90,TileScheduler 为 void 时即选它,cutlass-4.7.0/include/cutlass/gemm/kernel/tile_scheduler.hpp:107-121)就属于这一类:它不用 atomic,而是按一个可配置的 raster order(沿 M 还是沿 N 先走、以及 swizzle 宽度)把线性的 tile 序号映射成 (m, n) 坐标——顺序选得好,相邻时刻在跑的 tile 在 L2 里能共享更多 A/B 数据。

18.3.3 Dynamic Atomic Scheduler

__global__ void kernel(...) {
    while (true) {
        int t = atomic_get_next_tile();
        if (t >= total_tiles) break;
        process_tile(t, ...);
    }
}

最灵活,但 atomic 开销。tile 内工作量大时(GEMM 这种),atomic 开销可以忽略。

18.3.4 GPU-side Tile Scheduler with Optimization

CUTLASS 3.x 在 include/cutlass/gemm/kernel/ 下为 Sm90 GEMM 给了两个层次的 scheduler(另有 Grouped GEMM 专用的 PersistentTileSchedulerSm90Group):

  • PersistentTileSchedulerSm90(cutlass-4.7.0/include/cutlass/gemm/kernel/sm90_tile_scheduler.hpp):纯静态映射,无 atomic。grid 取 SM 数(有 cluster 时按能同时驻留的 cluster 数折算),且不超过 tile 总数(cutlass-4.7.0/include/cutlass/gemm/kernel/tile_scheduler_params.h:238-308);每个 block 从自己的线性序号出发、每次前进整个 grid 的大小(cutlass-4.7.0/include/cutlass/gemm/kernel/static_tile_scheduler.hpp:193),再按 raster order 映射成 tile 坐标。
  • PersistentTileSchedulerSm90StreamK(cutlass-4.7.0/include/cutlass/gemm/kernel/sm90_tile_scheduler_stream_k.hpp):解决波次量化(wave quantization)的长尾。

Stream-K 常被误解成"把大 tile 拆成小 tile"——它拆的其实是 K 维(归约维):当输出 tile 数不是 SM 数的整数倍时,最后一"波"只有部分 SM 有活干。Stream-K 把总工作量按 MAC 循环迭代次数均分给固定数量的 block,一个输出 tile 的 K 维可能被切给多个 block 各算一段部分和,再由其中一个 block 归约。代价是需要 workspace 存部分和、需要 barrier 协调归约顺序。这套方法出自 Osama、Merrill、Cecka、Garland、Owens 的论文 Stream-K: Work-centric Parallel Decomposition for Dense Matrix-Matrix Multiplication on the GPU(PPoPP 2023,arXiv:2301.03598)。CUTLASS 的实现是混合式的:只有出现波次量化时,才拿出两波的 tile 量交给 stream-K(源码注释写的是排在最前的两波),其余整波照常按 tile 数据并行;Heuristic 模式下尾波过半满还会直接退回纯数据并行(cutlass-4.7.0/include/cutlass/gemm/kernel/tile_scheduler_params.h:1082-1106)。

18.3.5 波次量化:persistent 想解决的正主

假设 GPU 能同时驻留 132 个 block,而 GEMM 一共有 140 个输出 tile:

第 1 波: 132 个 tile 同时算   ← 满载
第 2 波:   8 个 tile 同时算   ← 124 个 SM 闲置
总耗时 = 2 波 × 单 tile 耗时

利用率 = 140 / (2 × 132) ≈ 53%——近一半算力被"最后八个 tile"浪费掉了。tile 数越接近 SM 数的整数倍越好,越是"刚过一个整数倍一点点"越糟。这就是波次量化。它有三种解法:

  1. 调 tile 大小,让 tile 数凑近整数倍波——最简单,但 tile 大小还要同时满足 SMEM 和寄存器约束,自由度有限。
  2. Split-K:把 K 维切成若干段,人为把 tile 数乘上几倍,让最后一波的浪费占比变小。
  3. Stream-K:按迭代次数均分,从根上消掉波的概念。

18.4 Persistent + Producer/Consumer

把第 17 章的 Producer/Consumer 模式和 Persistent kernel 组合,给出 现代 attention kernel 的最终形态:

__global__ void persistent_fa3(...) {
    // 1. Persistent loop: 抢 tile
    //    注意广播必须走 SMEM 而不是 __shfl_sync —— 这个 block 有 384 线程(12 个 warp),
    //    shfl 只在 warp 内广播,warp 0 之外的线程拿不到 tile_id。
    __shared__ int s_tile_id;
    while (true) {
        if (threadIdx.x == 0) {
            s_tile_id = atomic_get_next_tile();
        }
        __syncthreads();
        int tile_id = s_tile_id;

        if (tile_id >= total_tiles) break;

        int q_tile_idx, head_idx, batch_idx;
        decode_tile_id(tile_id, &q_tile_idx, &head_idx, &batch_idx);

        // 2. 在这个 tile 上跑 FA3 producer/consumer 流水
        //    注意角色是按 warp-group (128 线程) 切的, 见第 17 章 17.2 节
        if (warp_group_idx == 0) {
            producer_main(q_tile_idx, head_idx, batch_idx, ...);
        } else {
            consumer_main(q_tile_idx, head_idx, batch_idx, ...);
        }

        // 3. 同步后再进下一轮: 不可省, 否则 thread 0 可能在别的线程读到 s_tile_id 之前就覆盖它
        __syncthreads();
    }
}

上面是为了看清结构的简化版。FA3 官方实现的骨架与之相同,但 producer 和 consumer 不共用一个循环、也不在 tile 之间 __syncthreads:flash-attn/hopper/tile_scheduler.hpp 提供 tile scheduler,flash-attn/hopper/flash_fwd_kernel_sm90.h 里 producer warp-group(:328-330)和 consumer warp-group(:376-378)各自跑一个 for (work_tile_info = scheduler.get_initial_work(...); work_tile_info.is_valid(...); work_tile_info = scheduler.get_next_work(...)) 循环。动态 scheduler 下由 producer warp 的 lane 0 做 atomicAdd,经 __shfl_sync 广播给 producer warp、经 SMEM 加两个 named barrier 交给 consumer(flash-attn/hopper/tile_scheduler.hpp:336-361)——这样 producer 不必等 consumer 做完当前 tile,就能开始载入下一个 tile。

18.5 LLM Decoding 的特殊优化:Flash-Decoding

LLM 推理的 decoding 阶段有一个特殊形态:Q 只有 1 个 token(当前生成的 token),但 K/V 有几千甚至几十万个(已生成的所有历史 tokens)。

这个形态下:

  • Q 只有 1 行,外层 Q 循环只有 1 次迭代。
  • K/V 维度极长,内层 K 循环可能跑几百次。
  • grid 只剩 (batch × head) 这一条并行轴——序列长度那条轴被压没了。
  • batch × head 小于 132 时(小 batch 长上下文正是这种情况),大量 SM 直接闲置,而那个在干活的 block 还要串行扫完几十万个 token 的 K/V。

Flash-Decoding 的解决方案:把 K 维度也切分给多个 SM,每个 SM 算一段 K 的 partial 结果,最后合并。

# Flash-Decoding 伪代码
# Q: [1, d] (单 token)
# K, V: [N, d] (N 可达 100K)

n_splits = 8  # 把 K 维度切 8 份
chunk_size = N // n_splits

# Phase 1: 每个 SM 算自己那一段 K 的 partial
for split_id in range(n_splits):  # 8 个 block 并行
    K_chunk = K[split_id * chunk_size : (split_id + 1) * chunk_size]
    V_chunk = V[split_id * chunk_size : (split_id + 1) * chunk_size]
    S_chunk = Q @ K_chunk.T
    P_chunk = softmax(S_chunk)  # 局部 softmax
    O_partial[split_id] = P_chunk @ V_chunk
    LSE_partial[split_id] = lse(S_chunk)

# Phase 2: 合并 partial
final_lse = combine_lse(LSE_partial)
final_O = sum(O_partial[s] * exp(LSE_partial[s] - final_lse) for s in range(n_splits))
final_O /= sum(exp(LSE_partial[s] - final_lse) for s in range(n_splits))  # final_lse 是 LSE 的 logsumexp 时此和恒为 1,可省

Flash-Decoding 把 K/V 维度变回一条并行轴,让原本闲置的 SM 参与进来。收益完全取决于原来闲了多少:batch × head 已经能填满 GPU 时它几乎没用,batch=1、长上下文时它是决定性的。

工程上这条路径到处都是:FlashAttention 仓库里是 flash_attn_with_kvcache(..., num_splits=...)(flash-attn/flash_attn/flash_attn_interface.py:1485)和 flash-attn/csrc/flash_attn/src/flash_fwd_split_*_sm80.cu 这一组 kernel;vLLM 里则是 vllm-0.8.5/csrc/attention/paged_attention_v2.cu——它的 grid 比 v1 多出一条 partition 轴,算完再由 reduce kernel 合并:

// vllm-0.8.5/csrc/attention/paged_attention_v1.cu:95
dim3 grid(num_heads, num_seqs, 1);
// vllm-0.8.5/csrc/attention/paged_attention_v2.cu:96 —— 多出的第三维就是 K/V 的切分
dim3 grid(num_heads, num_seqs, max_num_partitions);

18.6 Persistent Kernel 的代价

Persistent 模式不是免费的:

18.6.1 寄存器固定

Persistent kernel 的所有 tile 共享同一个寄存器布局。如果不同 tile 的最优寄存器需求不一样(比如 short context 的 attention 和 long context 的 attention),persistent 模式下只能取最大需求——可能浪费寄存器。

18.6.2 SMEM 复用要谨慎

不同 tile 之间复用 SMEM 听起来好,但需要小心 race condition——前一个 tile 还没写完 SMEM,后一个 tile 已经在读了。需要 __syncthreads(Hopper 上的 FA3/CUTLASS 则用按 stage 的 mbarrier 流水线和 named barrier)让线程真正互相等待;__threadfence_block 只是内存序栅栏,不会让任何线程停下来等别人,单靠它挡不住这种竞争。

18.6.3 调度复杂度

简单的 round-robin scheduler 可能导致负载不均;动态 atomic scheduler 又有竞争开销。CUTLASS 提供的 stream-K scheduler 是个不错的折中,但实现复杂。

18.6.4 不适合所有工作负载

Persistent 适合大量小 tile 或有复用机会的工作。如果 tile 都是大 GEMM(比如训练),传统模式同样高效,persistent 没有优势,反而增加复杂度。

18.7 一个反例:vLLM 的 PagedAttention 不是 persistent kernel

这里要修正一个流传很广的说法:vLLM 的 PagedAttention kernel 并不是 persistent kernel。看它的 launch 就清楚了(vllm-0.8.5/csrc/attention/paged_attention_v1.cu:95):

dim3 grid(num_heads, num_seqs, 1);
dim3 block(NUM_THREADS);

grid 随 num_seqs 变——batch 里有多少个序列就起多少个 block,不是固定成 SM 数。每个 block 负责一个 (head, seq),在 kernel 内部沿这条序列的 KV page 列表循环。它真正的设计要点是:

  1. 一个 block 一个 (head, seq):block 内沿 block table 逐 page 遍历该序列的 KV。

  2. 每个 page 默认 16 个 token 的 KV:物理页不连续,靠 block table 间接寻址,这是"分页"的本体。

  3. v2 才是 split-KV:vllm-0.8.5/csrc/attention/paged_attention_v2.cu:96 的 grid 多一维 max_num_partitions,把长序列的 KV 切给多个 block,再用一个 reduce kernel 按 partial LSE 合并——就是 18.5 节讲的 Flash-Decoding 形态。选哪个由 vllm-0.8.5/vllm/attention/ops/paged_attn.py:128 决定,条件是两条与关系:

    use_v1 = (max_seq_len <= 8192
              and (max_num_partitions == 1 or num_seqs * num_heads > 512))

    也就是说不只看序列长度——就算序列不长,只要 num_seqs * num_heads 已经超过 512(GPU 本来就填满了,切 KV 没意义),也会留在 v1。分区粒度是同文件 :15 的 _PARTITION_SIZE = 512。源码注释写得很实在:context len > 8192 用 V2 是为了避开 SMEM 不够,不是为了并行度。

  4. CUDA Graph 的可捕获性来自别处:不是"grid 固定"。V0 引擎是整图录制,按一组预设的 batch size 分别捕获多张图,运行时把实际 batch padding 到不小于它的最小档位;v0.8.5 默认的 V1 引擎则是分段捕获,档位按本拍 token 总数计,attention(包括 PagedAttention)根本不进图(附录 A.3.4)。

那这一节讲的 persistent + producer/consumer 结构在哪儿?在 FA3 和 CUTLASS 的 GEMM 里——vLLM 是"经典 grid + 分页寻址"的代表,两种设计解决的是不同问题,别混。

详细设计请参考《vLLM 推理内核深度解析》第 4 章 PagedAttention。

18.8 第四篇收官:从理论到 SOTA

第四篇我们完成了 attention kernel 优化的完整旅程:

章节 主题 H100 FP16 上的水位
第 14 章 IO-Aware 思想 朴素分步 attention:算力利用率百分之几到一成
第 15 章 FA2 forward 骨架 官方 FA2 约 35% 利用率(FA3 论文口径);手写骨架更低
第 16 章 FA2 backward 耗时粗估约为前向的 3 倍(按 FA2 论文 A100 利用率估)
第 17 章 TMA + WGMMA + Warp Spec (FA3) ~740 TFLOPs / 75% 利用率;FP8 接近 1.2 PFLOPs
第 18 章 Persistent + Tile Scheduler 消掉波次量化的长尾;Flash-Decoding 救回长上下文 decoding

从"百分之几到一成"到"75%",差出近一个数量级。这就是 GPU 工程的分量——同一份算法、同一块卡,实现方式不同,差出的是近一个数量级。

到这里,读者已经具备了 LLM 推理 / 训练中所有核心 kernel(GEMM、Attention、LayerNorm、Softmax、量化)的优化能力。

第五篇(第 19-21 章)我们换个视角——讲性能工程的工具链:怎么用 Nsight Compute 找瓶颈、怎么读 PTX/SASS、常见性能反模式。这些工具是日常 kernel 调优中的瑞士军刀,没有它们,再好的优化思路也找不到落点。

本章动手练习:

  1. 把第 17 章的 FA3 forward 改成 persistent 模式,对比性能。
  2. 实现 Flash-Decoding(n_splits=8),在 N=64K 长上下文 decoding 上测试加速比。
  3. 对读 vLLM 的 vllm-0.8.5/csrc/attention/paged_attention_v1.cu:95 与 vllm-0.8.5/csrc/attention/paged_attention_v2.cu:96 的 grid 配置,说清楚 v2 多出来的那一维在算什么、为什么它需要一个额外的 reduce kernel。