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

第 5 章 Reduction:从 atomic 到 cluster reduce

作者 杨艺韬 · 3,824 字 · 发布于 · 更新于

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 停在三分之一的是另外三件事:

  1. 同步太密:blockDim=256 要跑 8 轮循环,每轮一次 __syncthreads(),8 次 block 级栅栏摊在只有 255 次加法的工作量上。
  2. 线程大面积空转:从 s=128 起每轮至少一半线程闲着,整个循环下来平均活跃线程数只有 blockDim 的 1/8 上下。
  3. 访存指令太少:每个线程只读 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;
    }
}

关键变化:

  1. warp 内 5 次 __shfl_xor_sync 替代了 5 轮"SMEM 写 + __syncthreads + SMEM 读"——赢的不是单条指令的延迟(SHFL 和 SMEM 是同一量级),而是指令条数少了三分之二、block 级栅栏一次都不用。
  2. 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 这一章给我们的工程哲学

读到这里读者应该感受到一种"层层压榨"的工程哲学:

  1. 永远先看 Roofline:reduce 是带宽 bound,目标是逼近 3.35 TB/s 峰值。
  2. 从 atomic 开始往下挖:全局 atomic → 分层 atomic → cluster atomic。
  3. 从访存开始往下挖:单 float → 多元素 unroll → vectorized。
  4. 从同步开始往下挖:每轮一次 __syncthreads → 整个 block 只同步一次 → cluster sync。
  5. 从指令带宽开始往下挖:每条指令处理更多数据,减少指令总数。

这五条线索贯穿 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 算法的灵魂。

本章动手练习:

  1. 在 H100 / A100 上把 v0..v7 都实现一遍,记录每个版本的实际带宽。
  2. 用 Nsight Compute 看 v3 vs v4 的 lts__t_sectors_op_read.sum.per_second 指标差异。
  3. 思考:如果 reduce 的不是 sum 而是 max,需要改哪些地方?为什么 float 的 max 没有像 atomicAdd 那样现成的 atomic(CUDA 只提供整数版 atomicMax),要怎么绕?