CUDA 算子工程:手写 FlashAttention v2 之路
第 5 章 Reduction:从 atomic 到 cluster reduce
Reduction 是 GPU 上最小的"完整题目":它同时考 SIMT、SMEM、warp shuffle、coalesced 访存、 atomic 和 occupancy。把它从"HBM 峰值的百分之一二"推到"逼近峰值",第 1-4 章的每一个概念都会被点名一次。
5.1 为什么 Reduction 是入门第一题
Reduction(归约)就是把一组数压成一个数——求和、求最大、求最小、求积。听起来简单,但在 GPU 上写好它需要把第 1-4 章的几乎所有概念都用上:
- SIMT 与 warp:reduce 的核心是"32 个值合并成 1 个值"。
- SMEM 与 bank conflict:跨 warp 的 partial sum 通过 SMEM 交换。
- Warp Shuffle:warp 内的 reduce 不应该走 SMEM。
- Coalesced 访存:HBM 读取需要 coalesced。
- Vectorized load:用
float4提升指令带宽。 - Cluster Reduce(Hopper 新):跨 block 协作。
- 算术强度:reduce 是极端 bandwidth-bound(AI ≈ 0.25),优化目标是逼近 HBM 峰值带宽。
Mark Harris 在 CUDA SDK 里附过一份经典材料《Optimizing Parallel Reduction in CUDA》(随 reduction sample 一起发布),把朴素 reduce 一步步优化了 7 个版本(reduce0–reduce6)。那份讲义的数字是在 G80(峰值带宽 86.4 GB/s)上对 4M 个 int 求和测的:从 2.083 GB/s 一路优化到 62.671 GB/s。这一章我们沿着他的思路,再加上 Kepler 引入、Volta 起改为 _sync 形式的 warp shuffle,以及 Hopper 的 cluster reduce,给读者一份"现代版 Mark Harris reduce"。
最终目标:对 1 亿个 float 求和,逼近 H100 SXM5 的 HBM3 官方峰值 3.35 TB/s(这是规格峰值;实际能达到几成随卡和实现而变,见 §5.10 表后的说明)。
5.2 V0:朴素 atomic 版
最直观的写法:每个线程读一个元素,atomic 加到全局结果上(*out 要在 launch 前清零,下面 v6、v7 同理)。
__global__ void reduce_v0(const float* arr, float* out, int N) {
int tid = blockIdx.x * blockDim.x + threadIdx.x;
if (tid < N) {
atomicAdd(out, arr[tid]);
}
}
带宽量级:HBM 峰值的百分之一二——惨不忍睹。
为什么这么慢?因为 N 个线程都在对同一个全局地址做 atomic add。原子操作并不是加锁:这一行编译出来是一条 REDG.E.ADD.F32(返回值没用到,编译器发的是不等结果的 RED 而不是 ATOM;nvcc 13.4、sm_90a 编译所得),读-改-写由 L2 的原子单元完成。但同一地址上的原子操作只能一个接一个地完成,N=1e8 时就是 1 亿次对同一个地址的串行读-改-写,整张卡的并行度在这里被压成了 1。
教训:永远不要在 reduce 的最内层用全局 atomic。atomic 应该是分层归约的最后一步,且 atomic 的次数要比元素总数小几个数量级——v6 里 1e8 个元素只发 1e5 次 atomic,就已经不是瓶颈了。
5.3 V1:Block 内 SMEM 归约
把 reduce 拆成两步:先在 block 内把一组元素归约成一个数(写到 partial sum 数组),再启动第二个 kernel 把 partial sum 归约成最终结果。
__global__ void reduce_v1(const float* arr, float* partial, int N) {
__shared__ float smem[256];
int tid = threadIdx.x;
int gid = blockIdx.x * blockDim.x + tid;
smem[tid] = (gid < N) ? arr[gid] : 0.0f;
__syncthreads();
// Block 内分层归约
for (int s = 1; s < blockDim.x; s *= 2) {
if (tid % (2 * s) == 0) {
smem[tid] += smem[tid + s];
}
__syncthreads();
}
if (tid == 0) partial[blockIdx.x] = smem[0];
}
带宽量级:跳到峰值的两成上下,比 v0 高一个数量级——主要是不再全局 atomic 了。
但两成还是太低。问题在哪?看这段代码的关键瓶颈:
if (tid % (2 * s) == 0) { ... }
这一行触发严重的 warp divergence:
- s=1 时:tid=0,2,4,6,...30 活跃(16 个),tid=1,3,5,...31 闲置 → 50% divergence。
- s=2 时:tid=0,4,8,...28 活跃(8 个),其余闲置 → 75% divergence。
- s=4 时:tid=0,8,16,24 活跃(4 个)。
- s=8 时:tid=0,16 活跃(2 个)。
- s=16 起:这个 warp 里只剩 tid=0 活跃(s ≥ 32 时有的 warp 整个闲置)。
虽然 warp 内"活跃 lane 数减少"不会直接降低带宽(带宽瓶颈在 HBM 读),但算术单元利用率下降会让整个 kernel 变长。
5.4 V2:避免 warp divergence
把分层方式改一下:让前 N/2 个线程做加法,避免奇偶交替。
__global__ void reduce_v2(const float* arr, float* partial, int N) {
__shared__ float smem[256];
int tid = threadIdx.x;
int gid = blockIdx.x * blockDim.x + tid;
smem[tid] = (gid < N) ? arr[gid] : 0.0f;
__syncthreads();
// 关键: 让前 s 个线程做加法
for (int s = blockDim.x / 2; s > 0; s >>= 1) {
if (tid < s) {
smem[tid] += smem[tid + s];
}
__syncthreads();
}
if (tid == 0) partial[blockIdx.x] = smem[0];
}
现在 warp divergence 大幅降低:
- s=128 时:tid=0..127 活跃(4 个完整 warp),tid=128..255 闲置(4 个完整 warp)。整个 warp 同进同退,无 warp 内 divergence。
- s=64 时:tid=0..63 活跃(2 个完整 warp)。
- s=32 时:tid=0..31 活跃(1 个完整 warp)。
- s=16 时:tid=0..15 活跃(半个 warp,开始有 divergence)。
只有最后几个迭代(s ≤ 16)有 divergence——这部分量级很小,影响不大。
带宽量级:约峰值的三分之一。
这里要先澄清一个流传很广的误解:这一句没有 bank conflict。
smem[tid] += smem[tid + s];
一个 warp 里 32 个连续的 tid,读 smem[tid + s] 落在 32 个连续的地址上,也就是 32 个不同的 bank;读 smem[tid] 同理。两条访存指令各自都是零冲突的,硬件不会因此多发射。真正让 v2 停在三分之一的是另外三件事:
- 同步太密:blockDim=256 要跑 8 轮循环,每轮一次
__syncthreads(),8 次 block 级栅栏摊在只有 255 次加法的工作量上。 - 线程大面积空转:从 s=128 起每轮至少一半线程闲着,整个循环下来平均活跃线程数只有 blockDim 的 1/8 上下。
- 访存指令太少:每个线程只读 1 个 float(4 字节),发一条 LDG 就没事干了——喂不饱 HBM 的是每线程在途的访存太少,不是带宽本身。
后面三版(v3 每线程多读、v4 float4、v5 shuffle 换掉 SMEM)分别对着这三点开刀。
5.5 V3:每线程读多个元素
reduce 的算术强度只有 0.25 FLOPs/byte——意味着每读 1 字节做 0.25 次加法。如果每线程读多个元素,可以摊薄"读元素到 SMEM"的开销。
__global__ void reduce_v3(const float* arr, float* partial, int N) {
__shared__ float smem[256];
int tid = threadIdx.x;
int gid = blockIdx.x * (blockDim.x * 2) + tid;
// 每线程读 2 个元素, 直接相加
float v = 0.0f;
if (gid < N) v += arr[gid];
if (gid + blockDim.x < N) v += arr[gid + blockDim.x];
smem[tid] = v;
__syncthreads();
for (int s = blockDim.x / 2; s > 0; s >>= 1) {
if (tid < s) {
smem[tid] += smem[tid + s];
}
__syncthreads();
}
if (tid == 0) partial[blockIdx.x] = smem[0];
}
每个 block 现在处理 2 * blockDim.x 个元素而不是 blockDim.x 个,grid_size 减半。带宽量级:接近峰值的一半。
如果每线程读 4 个元素:
// gid 相应改为 blockIdx.x * (blockDim.x * 4) + tid,每一路都要像上面那样做 < N 的边界检查
float v = 0.0f;
v += arr[gid];
v += arr[gid + blockDim.x];
v += arr[gid + blockDim.x * 2];
v += arr[gid + blockDim.x * 3];
带宽量级:过半。继续提升,但收益递减。
5.6 V4:Vectorized Load 用 float4
把"每线程读 4 个 float"换成"每线程读 1 个 float4"——这一行代码的改动让访存指令条数减为 1/4:
__global__ void reduce_v4(const float* arr, float* partial, int N) {
__shared__ float smem[256];
int tid = threadIdx.x;
int gid = blockIdx.x * blockDim.x + tid;
// Vectorized: 一次读 16 字节
float local = 0.0f;
if (gid * 4 + 4 <= N) {
float4 v = *reinterpret_cast<const float4*>(&arr[gid * 4]);
local = v.x + v.y + v.z + v.w;
} else {
for (int k = gid * 4; k < N; ++k) local += arr[k]; // 尾部不足 4 个时逐个读
}
smem[tid] = local;
__syncthreads();
for (int s = blockDim.x / 2; s > 0; s >>= 1) {
if (tid < s) {
smem[tid] += smem[tid + s];
}
__syncthreads();
}
if (tid == 0) partial[blockIdx.x] = smem[0];
}
这一版每线程处理 4 个元素,grid 取 ⌈N/1024⌉;float4 读取要求地址 16 字节对齐,cudaMalloc 返回的指针满足这一点。else 分支只在最后一个 block 里兜住 N 不是 4 的倍数时的尾巴——少了它,末尾最多 3 个元素会被悄悄丢掉。
带宽量级:约七成。一行代码(float4 替代 float)就把每条 LDG 搬运的字节数翻了两番。
5.7 V5:Warp Shuffle 替代 SMEM 归约
到现在为止 block 内归约还在用 SMEM。但 warp 内的归约其实可以用 warp shuffle,无需 SMEM、无需同步:
__inline__ __device__ float warp_reduce(float val) {
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
}
return val;
}
__global__ void reduce_v5(const float* arr, float* partial, int N) {
int tid = threadIdx.x;
int gid = blockIdx.x * blockDim.x + tid;
// 1. 每线程 vectorized load + 局部 sum
float local = 0.0f;
if (gid * 4 + 4 <= N) {
float4 v = *reinterpret_cast<const float4*>(&arr[gid * 4]);
local = v.x + v.y + v.z + v.w;
} else {
for (int k = gid * 4; k < N; ++k) local += arr[k]; // 尾部不足 4 个时逐个读
}
// 2. Warp 内归约 (无 SMEM!)
local = warp_reduce(local);
// 3. 每个 warp 的 lane 0 写到 SMEM
__shared__ float warp_sums[8]; // blockDim.x / 32 = 8 warps
int warp_id = tid / 32;
int lane_id = tid % 32;
if (lane_id == 0) warp_sums[warp_id] = local;
__syncthreads();
// 4. Warp 0 归约 8 个 warp_sum
if (warp_id == 0) {
local = (lane_id < 8) ? warp_sums[lane_id] : 0.0f;
local = warp_reduce(local); // 实际只需要 3 次 shuffle (8 -> 1)
if (lane_id == 0) partial[blockIdx.x] = local;
}
}
关键变化:
- warp 内 5 次
__shfl_xor_sync替代了 5 轮"SMEM 写 +__syncthreads+ SMEM 读"——赢的不是单条指令的延迟(SHFL和 SMEM 是同一量级),而是指令条数少了三分之二、block 级栅栏一次都不用。 - block 内只有一次
__syncthreads(写 warp_sums 之后)。
带宽量级:八成以上。
5.8 V6:移除 partial 数组,直接 atomic
V1-V5 都需要两个 kernel:第一个算 partial,第二个把 partial 数组归约成最终结果。两个 kernel 的开销和中间数组的 HBM 访问让性能受限。
如果每个 block 已经把自己的 partial sum 算到一个 float 了,这个数已经比元素数少了三个数量级(N=1e8 时约 1e5 个)——这时候用 atomic 就没问题了:
__global__ void reduce_v6(const float* arr, float* out, int N) {
int tid = threadIdx.x;
int gid = blockIdx.x * blockDim.x + tid;
float local = 0.0f;
if (gid * 4 + 4 <= N) {
float4 v = *reinterpret_cast<const float4*>(&arr[gid * 4]);
local = v.x + v.y + v.z + v.w;
} else {
for (int k = gid * 4; k < N; ++k) local += arr[k]; // 尾部不足 4 个时逐个读
}
local = warp_reduce(local);
__shared__ float warp_sums[8];
int warp_id = tid / 32;
int lane_id = tid % 32;
if (lane_id == 0) warp_sums[warp_id] = local;
__syncthreads();
if (warp_id == 0) {
local = (lane_id < 8) ? warp_sums[lane_id] : 0.0f;
local = warp_reduce(local);
// 关键: 不写中间数组, 直接 atomic
if (lane_id == 0) atomicAdd(out, local);
}
}
少了一次 kernel launch + 一次中间数组 HBM 访问。grid_size 大约 1e8 / 1024 ≈ 1e5,每个 block 一次 atomicAdd,atomic 数量从 v0 的 1e8 降到 1e5——降低 1000 倍,几乎不会成为瓶颈。
带宽量级:九成以上。已经非常接近 HBM 峰值。
5.9 V7:Cluster Reduce(Hopper)
Hopper 的 Cluster 让我们可以把 atomic 数量再除以 cluster 大小(可移植上限 8 即再降 8 倍,opt-in 到 16 为 16 倍)——把 cluster 内一组 block 的 partial sum 在分布式 SMEM 内汇总,再 atomic 出去。
下面用 __cluster_dims__(16,1,1) 是为了把效果放到最大;要注意 CUDA 可移植的 cluster 上限是 8,16 需要在 host 侧 opt-in cudaFuncAttributeNonPortableClusterSizeAllowed,否则 launch 会失败。另外 grid 必须是 cluster 大小的整数倍:N=1e8 时 ⌈N/1024⌉ = 97657 个 block,要向上补到 97664,多出来的 block 靠边界检查读到 0。
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
__global__ void __cluster_dims__(16, 1, 1)
reduce_v7(const float* arr, float* out, int N) {
auto cluster = cg::this_cluster();
auto block = cg::this_thread_block();
int tid = threadIdx.x;
// 注意 blockIdx.x 已经是**全网格**的 block 编号(cluster 不会改变它的含义,
// cluster.block_rank() 只是 blockIdx.x 对 cluster 大小取模),所以这里不能
// 再拿 block_rank 去叠一遍,否则索引会成倍跳过数据。
int gid = blockIdx.x * blockDim.x + tid;
// 1-3. 同 v6: vectorized load + warp reduce + block reduce
float local = 0.0f;
if (gid * 4 + 4 <= N) {
float4 v = *reinterpret_cast<const float4*>(&arr[gid * 4]);
local = v.x + v.y + v.z + v.w;
} else {
for (int k = gid * 4; k < N; ++k) local += arr[k]; // 尾部不足 4 个时逐个读
}
local = warp_reduce(local);
__shared__ float block_sum;
__shared__ float warp_sums[8];
int warp_id = tid / 32;
int lane_id = tid % 32;
if (lane_id == 0) warp_sums[warp_id] = local;
__syncthreads();
if (warp_id == 0) {
local = (lane_id < 8) ? warp_sums[lane_id] : 0.0f;
local = warp_reduce(local);
if (lane_id == 0) block_sum = local;
}
__syncthreads();
// 4. Cluster 内汇总: block 0 收集所有 block 的 block_sum
cluster.sync(); // 等所有 block 都写完 block_sum
if (cluster.block_rank() == 0 && warp_id == 0) {
float total = 0.0f;
for (int b = lane_id; b < cluster.num_blocks(); b += 32) {
float* peer = cluster.map_shared_rank(&block_sum, b);
total += *peer;
}
// warp 内归约这个 total
total = warp_reduce(total);
if (lane_id == 0) atomicAdd(out, total);
}
// 5. 退出前必须再同步一次:block 0 还在读别的 block 的 SMEM,
// 别的 block 一旦退出,那块 SMEM 就可能被下一个 block 复用。
// CUDA C++ Programming Guide 的 distributed shared memory 直方图示例
// 同样是"一头一尾各一次 cluster.sync()"。
cluster.sync();
}
Cluster=16 时,每 16 个 block 共享一次 atomic,atomic 总数从 v6 的 ~1e5 降到 ~6e3——再降 16 倍(cluster=8 则是 8 倍)。
带宽量级:与 v6 基本持平。v6 的 1e5 次 atomic 本来就不是瓶颈,再降 16 倍在纯求和里几乎看不出来;v7 的意义在于演示 DSMEM 这条跨 block 通路(见 §5.11)。
5.10 性能对比与小结
把 8 个版本的性能放在一张表里:
| 版本 | 优化点 | 达到的 HBM 峰值占比(量级) |
|---|---|---|
| v0 | 朴素 + 全局 atomic | ~1–2% |
| v1 | Block 内 SMEM 归约 | ~20% |
| v2 | 避免 warp divergence | ~33% |
| v3 | 每线程读 4 元素 | ~55% |
| v4 | float4 vectorized | ~70% |
| v5 | Warp shuffle | ~85% |
| v6 | 直接 atomic, 无中间数组 | ~91% |
| v7 | Cluster reduce (Hopper) | ~91%(≈ v6) |
关于这张表:本专栏没有条件在写作时对每一版重新跑 profile,上表是按"HBM 峰值 3.35 TB/s 的百分比"给出的量级示意, 用来表达各版之间的相对台阶(v0 的全局 atomic 慢在哪、v4 的
float4值多少、v7 的 cluster 为什么几乎抠不出带宽), 而不是某台机器上的实测报告。真实数字随卡、驱动、CUDA 版本、N 的取值而变——请照本章末的练习自己测一遍。
从百分之一二到逼近峰值——同一个算法,两个数量级的性能差距。这就是 GPU 编程的现实:正确性容易,性能难。一个看似"小改动"(比如 vectorized load)背后藏着对硬件的深刻理解。
flowchart LR V0[v0 ~1%] --> V1[v1 ~20%] V1 --> V2[v2 ~33%] V2 --> V3[v3 ~55%] V3 --> V4[v4 ~70%] V4 --> V5[v5 ~85%] V5 --> V6[v6 ~91%] V6 --> V7[v7 ≈ v6] style V0 fill:#fee2e2 style V7 fill:#bbf7d0
5.11 这一章给我们的工程哲学
读到这里读者应该感受到一种"层层压榨"的工程哲学:
- 永远先看 Roofline:reduce 是带宽 bound,目标是逼近 3.35 TB/s 峰值。
- 从 atomic 开始往下挖:全局 atomic → 分层 atomic → cluster atomic。
- 从访存开始往下挖:单 float → 多元素 unroll → vectorized。
- 从同步开始往下挖:每轮一次
__syncthreads→ 整个 block 只同步一次 → cluster sync。 - 从指令带宽开始往下挖:每条指令处理更多数据,减少指令总数。
这五条线索贯穿 LLM 算子的所有优化场景。GEMM、Softmax、LayerNorm、FA2——它们的优化思路都是这五条的某种组合。
特别值得记住的两个反直觉事实:
- warp shuffle 赢 SMEM,赢的不是延迟而是指令数与同步——一次交换 SMEM 要 "写 + 栅栏 + 读" 三步,shuffle 只要一条
SHFL,而且不占 SMEM 容量、不需要__syncthreads()。 - Cluster Reduce 不是为了 reduce 本身——v6 已经用一次 atomic 省掉了第二个 kernel,纯求和里 v7 能省的只剩那些本就不是瓶颈的 atomic。DSMEM 的真正价值在于让 cluster 内的 block 不经 HBM 交换中间结果:比如一行太长、要切给多个 block 的 Softmax/LayerNorm,归约结果还得广播回每个 block 继续算,这时 cluster 能省掉一次额外 kernel launch 或一轮 HBM 往返。
第 6 章我们把这一套手艺用到 Softmax 上。Softmax 的核心也是 reduction(找 max + 求 exp sum),但比纯 reduce 多一个数值稳定性的问题——这正好是 Online Softmax 要解决的,也是 FA2 算法的灵魂。
本章动手练习:
- 在 H100 / A100 上把 v0..v7 都实现一遍,记录每个版本的实际带宽。
- 用 Nsight Compute 看 v3 vs v4 的
lts__t_sectors_op_read.sum.per_second指标差异。- 思考:如果 reduce 的不是 sum 而是 max,需要改哪些地方?为什么 float 的 max 没有像
atomicAdd那样现成的 atomic(CUDA 只提供整数版atomicMax),要怎么绕?