找回密码
立即注册
搜索
热搜: Java Python Linux Go
发回帖 发新帖

4558

积分

0

好友

596

主题
发表于 6 小时前 | 查看: 9| 回复: 0

写在前面:本文是【CUDA编程】Flash Attention CUDA 算子全面超越 SDPA 的续篇。上篇我们把训练 / prefill 阶段(Q 长度 N 较大)的 Flash Attention v2 kernel 优化到了全面超越 SDPA。但上篇的并行度来自"按 Q 行切块"——在自回归 decode 阶段(N=1),这个假设塌了:Q 只有一行,Br=1,并行度直接塌缩成 B×H,且不能再直接使用 mma 加速。当 batch、head 较小时,GPU 的 SM 大面积空转。

本文要解决的问题正是 decode 阶段的这个困境:Flash Decoding 把并行度从"Q 行"搬到"K/V 行(chunk)",用"按 chunk 分块并行 + LSE 加权归约"让 SM 重新占满。完整记录两段式 kernel 的算法、实现、踩坑与 ncu 性能剖析。测试 GPU 与上篇一致:RTX 4060(Ada,sm_89)。所有代码与数据均来自实际工程。本文为云栈社区【CUDA编程】系列技术文章之一。


1 为什么 decode 阶段需要 Flash Decoding

1.1 标准 Flash Attention 的并行度来自 Q 行

Flash Attention v2 的核心是把 Q 切成 Br 行的块,每个 thread block 负责一段 [Br, M] 的注意力。block 之间相互独立,天然的并行度 = num_q_blocks = ceil(N / Br)。prefill 阶段 N 很大(几千几万),并行度天然充足。

1.2 decode:N=1,并行度塌缩

自回归生成时,每一步只产出 1 个新 token,于是 N=1。此时为了正确,每个 block 只能处理这 1 行 Q(Br 取 1,否则会破坏 causal / 正确性),于是:

并行度 = ceil(1 / 1) × B × H = B × H

不再是 N 驱动的,而是 B×H 驱动的。

1.3 后果:B、H 小时 SM 占不满

一套 decode 配置的典型值:B=2, H=8, M=65536。并行问题数 = 2×8 = 16。而 RTX 4060 有 24 个 SM。意味着:

  • 即使 16 个问题全并行,也只有 16/24 的 SM 在干活;
  • 每个 SM 上若只驻留 1 个 block,那每 SM 的 32 个 warp 里只有 1 个 block 的 warp 在跑,其余全空转等显存(KV 序列长,memory-bound 严重);
  • M 越大(KV 越长),单问题算得越久,但并行度没有随之增加——这正是标准 Flash Attention 在 decode 长序列下的尴尬。

1.4 数字说话

配置 并行问题数 (B×H) SM 数 占用率上限
B=2 H=8 M=65536 16 24 66%(且每 SM 仅 1 block)
B=1 H=8 M=512 8 24 33%

Flash Decoding 的洞察:KV 序列有 M 行,为什么不把并行度放在 K/V 行上?把 M 切成 num_chunks = M / chunk_size 个 chunk,每个 chunk 独立算一份"partial 注意力",再合并——并行度立刻从 B×H 变成 num_chunks × B×HM=65536, chunk=128 时 = 512 个 chunk block,SM 被彻底占满。


2 Flash Decoding 算法原理

2.1 核心思想

标准 Flash Attention 把 [N, M] 注意力沿 M(K/V) 切成 Bc 列块,逐块流式计算并就地维护 online softmax 状态(不物化 [N,M] 大矩阵)。Flash Decoding 在此基础上额外沿 Q 维度做一层分块

  • 因为 decode 时 N=1,Q 只有 1 行,所以沿 Q 切块没意义;
  • 改沿 K/V 行(chunk) 切块:每个 chunk 是一段 [chunk_size, d] 的 K/V;
  • 每个 chunk 独立算"partial O"和"该 chunk 的 LSE(log-sum-exp)";
  • 最后用一个轻量归约 kernel,把所有 chunk 的 partial O 按 LSE 权重加权合并成最终 O。

关键性质:partial 之间相互独立,可以完全并行,这正是把并行度铺满 SM 的来源。

2.2 数学:chunk partial + 全局 LSE 加权归约

对单个 query 行 q(已除 sqrt(d)),设第 c 个 chunk 覆盖 K/V 的行区间 [c·Cs, (c+1)·Cs)

Step 1 — 每个 chunk 内算 partial(与标准 Flash Attention 的 single-block 计算完全一致):

S_c = q · K_c^T                 # [1, Cs]
m_c = rowmax(S_c)               # chunk 内最大 logit
P_c = exp(S_c - m_c)            # [1, Cs],chunk 内未归一化
l_c = rowsum(P_c)               # chunk 内概率和
O_c = P_c · V_c / l_c           # chunk 内归一化的 partial O,[1, d]
lse_c = m_c + log(l_c)          # 该 chunk 的 log-sum-exp,标量

Step 2 — 全局归约(整个序列的 softmax):

设全局最大值 m = max_c m_c。全局分母与分子:

L = Σ_c exp(m_c - m) · l_c = Σ_c exp(lse_c - m)      # 因 lse_c = m_c + log(l_c)
O = Σ_c exp(m_c - m) · l_c · O_c / L
  = Σ_c [exp(lse_c - m) / L] · O_c

global_lse = m + log(L) = log(Σ_c exp(lse_c))(就是把所有 chunk 的 lse_c 做 logsumexp),再记权重 w_c = exp(lse_c - global_lse) = exp(lse_c) / Σ_c exp(lse_c)(注意它已含 1/L),则:

O = Σ_c w_c · O_c

这就是我们 reduce kernel 里逐字实现的公式global_lse = global_max + log(denom)w = exp(lse_c - global_lse)O = Σ w_c·O_c

2.3 为什么加权和能等价

上面推导的核心是把 exp(S_cj - m) = exp(S_cj - m_c) · exp(m_c - m) = P_cj · exp(m_c - m) 拆开——P_cj(已含 exp(-m_c))与 exp(m_c - m)(chunk 级标量)分离后,chunk 内归一化的 O_c 可以原样参与全局加权和,只需再乘 chunk 级权重 w_c数值等价性由 logsumexp 的标准合并公式保证,不损失精度(实测 fp16 下 max_abs ≈ 1e-4)。

2.4 计算量没变,并行度翻了 ~64 倍

Flash Decoding 没有改变总计算量(注意力该算的矩阵乘一个没少),以 chunk_size = 128, M = 65536 为例,它只是重新组织了并行方式:把"串行的 512 个 chunk 依赖关系"解耦成"并行的 512 个独立 chunk + 1 次廉价归约"。这正是它能解决 SM 占不满问题的本质。


3 实现架构:两段式 kernel

3.1 整体数据流

输入:  Q[B,H,1,d]   K/V[B,H,M,d]
                    │
   ┌────────────────┴─────────────────┐
   │  chunk kernel  (grid = num_chunks × B×H) │
   │  每个 block 处理 1 个 chunk 的 128 行 K/V  │
   └────────────────┬─────────────────┘
                    ▼
   O_chunk[B,H,num_chunks,d]   +   lse[B,H,num_chunks]
                    │
   ┌────────────────┴─────────────────┐
   │  reduce kernel (grid = B×H)        │
   │  每个 block 把 num_chunks 个 partial │
   │  按 LSE 权重加权合并出 O[B,H,1,d]    │
   └────────────────┬─────────────────┘
                    ▼
   O[B,H,1,d]

两段式之间需要一个临时 buffer 存 O_chunklse,由 launch 函数内部 cudaMallocAsync 分配、用完 cudaFreeAsync 释放,调用方无感。

3.2 chunk kernel

dim3 grid_dim(num_chunks, num_head * batch_size);  // x: chunk 下标, y: (b,h)
dim3 block_dim(128);                               // 128 线程覆盖 128 行 K/V
flash_decoding_d64_chunk128_kernel<<<grid_dim, block_dim, 0, stream>>>(
    Q, K, V, O_chunk, lse, M, d, softmax_scale);

blockIdx.y 区分不同 (b,h),blockIdx.x 区分不同 chunk,chunk_start = chunk_id * 128 定位 K/V 行。

3.3 reduce kernel

dim3 grid_reduce(num_head * batch_size);            // 每个 (b,h) 一个 block
flash_decoding_reduce_chunk128_kernel<<<grid_reduce, 64,
    (num_chunks > 128 ? num_chunks : 128) * sizeof(float), stream>>>(
    O_chunk, lse, O, N, M, d);

reduce 每个 block 只处理 1 个 (b,h),但要把 num_chunks(可达 512)个 partial 合并,所以内部用 2 个 warp 协作(详见第 6 节)。

3.4 临时 buffer 与 launch 封装

void launch_flash_decoding_kernel(const half* Q, const half* K, const half* V, half* O,
    const int batch_size, const int num_head, const int N, const int M,
    const int d, cudaStream_t stream)
{
    assert(N == 1);
    assert(d == 64 && "flash_decoding_d64 only supports d=64");
    const float softmax_scale = 1.0f / sqrtf((float)d);
    constexpr int chunk_size = 128;
    const int num_chunks = div_ceil(M, chunk_size);

    half * lse;
    half * O_chunk;
    size_t buf_size = (batch_size * num_head * num_chunks          // lse
                     + batch_size * num_head * num_chunks * d) * sizeof(half);  // O_chunk
    CHECK_CUDA_ERROR(cudaMallocAsync((void **)&lse, buf_size, stream));
    O_chunk = lse + batch_size * num_head * num_chunks;

    // ... 启动 chunk + reduce(见 3.2 / 3.3)...
    CHECK_CUDA_ERROR(cudaFreeAsync(lse, stream));
}

调用方传入的 O 形状是 [B, H, 1, D](只有 1 个 query 行),但 chunk kernel 会写出 num_chunks 份 partial——这正是早期 OOB 坑的来源(见第 7.1 节),最终用 launch 内部分配的 O_chunk 中转解决。


4 chunk kernel 实现(D=64)

4.1 线程映射

blockDim=128,4 个 warp。Q 一行 64 列 = 32 个 half2,因此每 warp 的每 lane 持 2 个 half2(覆盖整行);4 个 warp 协同处理 chunk 内的 128 行 K/V(每 warp 负责 32 行)。

const uint32_t warp_id = threadIdx.x >> 5;
const uint32_t lane_id = threadIdx.x & 31;
const uint32_t num_chunks = gridDim.x;
const uint32_t chunk_id = blockIdx.x;
const uint32_t chunk_start = chunk_id * 128;
const uint32_t eff_chunk = min(128, M - chunk_start);  // 末 chunk 可能不足 128

const half* Q_ptr = Q + blockIdx.y * 1 * d;             // 该 (b,h) 的唯一 query 行
const half* K_ptr = K + blockIdx.y * M * d + chunk_start * d;
const half* V_ptr = V + blockIdx.y * M * d + chunk_start * d;

half2 RQ = reinterpret_cast<const half2 *>(Q_ptr)[lane_id];  // Q 整行常驻寄存器

4.2 Q@K^T:warp shuffle 归约

每个 warp 处理 eff_chunk 中的 32 行(按 warp_id 跨步)。每行 K 也是 64 列 = 32 个 half2,与 RQ 做 half2 点积后,warp 内 5 次 shfl_xor 把 32 个 lane 的部分积归约成标量

__shared__ float qk[128];  // 暂存 128 个 score

for (int i = warp_id; i < eff_chunk; i += (blockDim.x >> 5)) {
    half2 RK = reinterpret_cast<const half2 *>(K_ptr + i * d)[lane_id];
    half2 qk_mul = __hmul2_rn(RQ, RK);          // 2 个 half 点积,half2 向量化
    __syncwarp();
    qk_mul = __hadd2(qk_mul, __shfl_xor_sync(0xffffffff, qk_mul, 16));
    qk_mul = __hadd2(qk_mul, __shfl_xor_sync(0xffffffff, qk_mul, 8));
    qk_mul = __hadd2(qk_mul, __shfl_xor_sync(0xffffffff, qk_mul, 4));
    qk_mul = __hadd2(qk_mul, __shfl_xor_sync(0xffffffff, qk_mul, 2));
    qk_mul = __hadd2(qk_mul, __shfl_xor_sync(0xffffffff, qk_mul, 1));
    // 此时 qk_mul 已是 64 个部分积之和(half2 形式)
    if (lane_id == 0) {
        __half dot_val = __hadd(__low2half(qk_mul), __high2half(qk_mul));
        qk[i] = __half2float(dot_val) * softmax_scale;  // 写回 shared,i 号 score
    }
    __syncwarp();
}
__syncthreads();

4.3 online softmax:blockAllReduceMax/Sum

整块 128 个 score 的 max / sum 跨所有 4 个 warp 归约——直接复用项目里现成的 blockAllReduceMax / blockAllReduceSum(用 shared memory 做跨 warp 合并,与上篇 Flash Attention 同款基础设施):

float score = threadIdx.x < eff_chunk ? qk[threadIdx.x] : -1e20f;
__syncthreads();
float row_max = blockAllReduceMax(score);

score = expf(score - row_max);
float local_sum = threadIdx.x < eff_chunk ? score : 0.0f;
__syncthreads();
float exp_sum = blockAllReduceSum(local_sum);
score /= exp_sum;
qk[threadIdx.x] = score;    // 此时 qk[] 存的是归一化后的 P(未乘 V)
__syncthreads();

4.4 PV:d=64 的跨组合并

d=64,128 线程只够"每线程 1 列"的一半(需 64 列)。于是让 128 线程两两配对覆盖 64 列:d_out = threadIdx.x & 63group = threadIdx.x / 64。两组各自沿 chunk 行累加 V 后,通过 shared memory 把同列的两组结果加总:

const uint32_t d_out = threadIdx.x & 63;
const uint32_t group = threadIdx.x / 64;
const int t_start = group * 64;
const int t_end   = min(t_start + 64, eff_chunk);
float acc = 0.0f;
for (int j = t_start; j < t_end; j++) {
    float v = __half2float(V_ptr[j * d + d_out]);
    acc += qk[j] * v;
}
__syncthreads();
qk[threadIdx.x] = acc;    // 复用 qk[] 暂存 partial
__syncthreads();

if (threadIdx.x < 64) {
    acc += qk[threadIdx.x + 64];  // 跨组合并同列
    const int out_offset = blockIdx.y * num_chunks * d + chunk_id * d + threadIdx.x;
    O[out_offset] = __float2half(acc);  // 写出 O_chunk[bh, chunk_id, d_out]
}

4.5 写出 LSE

LSE 是全局归约必需的标量,单独由一个线程写出(存 fp16,足够,因为 reduce 阶段会转回 float 做 logsumexp):

if (threadIdx.x == 0) {
    const int lse_offset = blockIdx.y * num_chunks + chunk_id;
    lse[lse_offset] = __float2half(row_max + logf(exp_sum));
}

这段 chunk kernel 从第一版就基本正确(正确性 bug 主要出在 OOB 与 reduce),是整条链路上最稳的一块。


5 chunk kernel 实现(D=128)

5.1 与 D=64 的异同

算法完全一致,差异全在"d=128 比 d=64 宽一倍"带来的线程映射变化。要点:

  • Q@K^T:d=128 = 64 个 half2,每 warp 每 lane 持 2 个 half2(即 4 个 half) 才能覆盖整行;每行产生 4 个 half2 部分积,转 float 后 warp-reduce 得标量 logit。
  • PV:d=128 恰好铺满 128 个线程——每个线程直接负责 1 个 d 列(threadIdx.x 即列号),沿全部 key 累加后写出。比 d=64 版那种"跨组 shared memory 合并"更干净,没有 qk[threadIdx.x]=acc+qk[threadIdx.x+64] 的合并步骤。
  • softmax 沿 key 维,与 d 无关,直接复用 blockAllReduceMax/Sum

5.2 Q@K^T 关键片段(每 lane 持 4 个 half)

// d=128:每行 64 个 half2;每 lane 持 2 个 half2(idx, idx+32),覆盖整行
const int h2 = lane_id * 2;                 // 该 lane 负责的 2 个 half2 起始
half2 RQ0 = reinterpret_cast<const half2*>(Q_ptr)[h2];
half2 RQ1 = reinterpret_cast<const half2*>(Q_ptr)[h2 + 1];

for (int i = warp_id; i < eff_chunk; i += (blockDim.x >> 5)) {
    const half2* Krow = reinterpret_cast<const half2*>(K_ptr + i * d);
    half2 p0 = __hmul2_rn(RQ0, Krow[h2]);
    half2 p1 = __hmul2_rn(RQ1, Krow[h2 + 1]);
    // 4 个 half 部分积 → float → warp reduce(同 d=64 的 shfl_xor 流程,略)
    // ... 得该行的标量 logit,写 qk[i] ...
}

5.3 PV:d 铺满线程,无需跨组合并

const int d_out = threadIdx.x;              // d=128 时 threadIdx.x ∈ [0,128) 即列号
float acc = 0.0f;
for (int j = 0; j < eff_chunk; ++j)
    acc += __half2float(V_ptr[j * d + d_out]) * qk[j];
const int out_offset = blockIdx.y * num_chunks * d + chunk_id * d + d_out;
O[out_offset] = __float2half(acc);          // 直接写出,无跨组合并

经验:d 维度能否被 block 线程数整除,直接决定 PV 要不要"跨组 shared memory 合并"。d=128 整除 128 线程,少了一层暂存与同步,是比 d=64 更优的映射。


6 reduce kernel 实现

reduce 把 num_chunks 个 partial O 加权合并成最终 O。这是整条链路上优化空间最大、也最容易写错的一块。

6.1 Phase 0:lse 整块进 shared(只从 global 读一次)

extern __shared__ float sbuf[];
float* lse_s = sbuf;

// lse 整块搬进 shared,全程只这 1 次 global 读
for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
    lse_s[c] = __half2float(lse[bh * num_chunks + c]);
__syncthreads();

6.2 Phase 1:跨 warp 规约 global_lse(复用 blockAllReduce*)

float local_max = -1e20f;
for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
    local_max = fmaxf(local_max, lse_s[c]);
float global_max = blockAllReduceMax(local_max);    // 跨 4 个 warp 合并

float local_den = 0.f;
for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
    local_den += expf(lse_s[c] - global_max);
float denom = blockAllReduceSum(local_den);
const float global_lse = global_max + logf(denom);  // = log(Σ exp(lse_c))

// 把 lse_s 原地改写成权重 w_c = exp(lse_c - global_lse),Phase2 零 global 读,还省一次 exp
for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
    lse_s[c] = expf(lse_s[c] - global_lse);
__syncthreads();

6.3 Phase 2:各 warp 处理不同 chunk,加权累加 + 跨 warp 求和

num_chunks 沿 warp 切分,warp0 处理前半 chunk、warp1 处理后半,两 warp 干不同的活。每个 warp 内部 32 个 lane 用 half2 覆盖完整 d=64:

const int num_warps = blockDim.x >> 5;
const int per_warp = (num_chunks + num_warps - 1) / num_warps;
const int c_start = warp_id * per_warp;
const int c_end = min(c_start + per_warp, num_chunks);

float2 acc = {0.f, 0.f};
for (int c = c_start; c < c_end; ++c) {
    const float w = lse_s[c];                                   // 权重已在 shared
    const float2 vf = __half22float2(reinterpret_cast<const half2*>(o_base + c * d)[lane]);
    acc.x += w * vf.x;
    acc.y += w * vf.y;
}
// 暂存部分和(64 线程 × float2 = 128 float;lse_s 此时已消费完,覆盖安全)
reinterpret_cast<float2*>(sbuf)[threadIdx.x] = acc;
__syncthreads();
if (warp_id == 0) {                                             // warp0 跨 warp 求和并写出
    float2 total = acc;
    for (int w = 1; w < num_warps; ++w) {
        const float2 oth = reinterpret_cast<float2*>(sbuf)[w * 32 + lane];
        total.x += oth.x; total.y += oth.y;
    }
    reinterpret_cast<__half2*>(O + bh * d)[lane] = __float22half2_rn(total);
}

设计要点:d=64 时 half2_idx = lane(0..31),两 warp 覆盖相同列,所以"按 chunk 拆 + 跨 warp 求和"合法且必要。d=128 版则更进一步——64 线程与 64 个 half2 构成双射(见 6.4),连跨 warp 求和都不需要,每个线程直接循环全部 chunk 累加自己的 2 列即可。

6.4 d=128 的 reduce:线程↔half2 双射,无需跨 warp 求和(最终形态)

d=128 = 64 个 half2;而 reduce 的 block 正好 64 线程。关键观察:threadIdx.x ∈ [0,63] 与 64 个输出 half2 构成一一映射(双射)。每个输出元素 O[o] = Σ_c w_c · O_chunk[c,o]列内独立的加权求和,最终写出根本不需要任何跨线程归约——于是每个线程自己循环全部 chunk、累加自己的 2 列、直接写出即可,既无跨 warp 求和、也无 acc 的 shared 暂存、更无 Phase 2 的 __syncthreads()

// d=128 的 reduce kernel(最终形态)。block=64(2 warp):
//  - Phase1:跨 warp 规约 lse(blockAllReduce*),与 d=64 版一致。
//  - Phase2:d=128 共 128 列 = 64 个 half2,block=64 线程与 64 个 half2 构成双射,
//            threadIdx.x 直接 1:1 映射到输出 half2;每个线程循环全部 chunk 累加自身 2 列,
//            直接写出,无需跨 warp 求和、无需 shared 暂存 acc。
__global__ void flash_decoding_reduce_d128_chunk128_new_kernel(
    const half* __restrict__ O_chunk, const half* __restrict__ lse,
    half* __restrict__ O, const int N, const int M, const int d)
{
    const int bh = blockIdx.x;
    const int num_chunks = div_ceil(M, 128);
    const half* o_base = O_chunk + bh * num_chunks * d;

    extern __shared__ float sbuf[];
    float* lse_s = sbuf;    // 仅缓存 lse_s[num_chunks],acc 全部在寄存器

    // Phase 0:lse 整块搬进 shared(只从 global 读这一次)
    for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
        lse_s[c] = __half2float(lse[bh * num_chunks + c]);
    __syncthreads();

    // Phase 1:跨 warp 协同规约 global_max / denom(复用 blockAllReduce*)
    float local_max = -1e20f;
    for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
        local_max = fmaxf(local_max, lse_s[c]);
    __syncthreads();
    float global_max = blockAllReduceMax(local_max);

    float local_den = 0.f;
    for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
        local_den += expf(lse_s[c] - global_max);
    __syncthreads();
    float denom = blockAllReduceSum(local_den);
    const float global_lse = global_max + logf(denom);

    // lse_s 原地改写成权重 w_c = exp(lse_c - global_lse),Phase2 零 global 读
    for (int c = threadIdx.x; c < num_chunks; c += blockDim.x)
        lse_s[c] = expf(lse_s[c] - global_lse);
    __syncthreads();

    // Phase 2:threadIdx.x ∈ [0,63] 与 64 个输出 half2 双射,每线程独占 2 列。
    //          O[o] = Σ_c w_c · O_chunk[c,o] 本就是列内独立加权求和,无需跨线程归约。
    float2 acc = {0.f, 0.f};
    for (int c = 0; c < num_chunks; ++c) {
        const float w = lse_s[c];
        const float2 vf = __half22float2(reinterpret_cast<const half2*>(o_base + c * d)[threadIdx.x]);
        acc.x += w * vf.x;
        acc.y += w * vf.y;
    }
    reinterpret_cast<half2*>(O + bh * d)[threadIdx.x] = __float22half2_rn(acc);
}

为什么 d=128 能比 d=64 更简洁:d=64 时 64 线程只对应 32 个 half2,两 warp 在列上重叠,必须跨 warp 求和(见 6.3);d=128 时 64 线程 ↔ 64 个 half2 恰好不重叠,"按 chunk 拆 + 跨 warp 求和"那套机制纯属多余。这一版彻底删掉了跨 warp merge,shared 也从"lse_s + acc 暂存"缩成只装 lse_s,代码更短、更不易错,且 ncu 实测还快了 28%(见 7.3)。数学上每列仍汇总了全部 chunk 的 partial,逐位等价于 d=64 版,实测 max_abs ≈ 1e-4


7 性能剖析(ncu + roofline)

正确性通过后,用 ncu --set full --nvtx 配合 nvtx range 抓取两个 decode kernel,聚焦长序列生产场景 B=2 H=8 N=1 M=65536

7.1 chunk kernel:贴死 DRAM 带宽天花板

指标 D=64 D=128
Duration 1.10 ms 2.17 ms
DRAM Throughput 96.37% (246 GB/s) 96.9% (248.4 GB/s)
Compute (SM) 56.44% 33.1%
Achieved Occupancy 98% ~100%
Registers / Thread 40 40
Eligible Warps / SM 0.46 0.31
L2 Hit Rate ~0.7% 0.62%

两个 head_dim 的 chunk kernel 都跑在 ~96% 的 DRAM 带宽上(本卡真实峰值 ≈ 255 GB/s)。这是典型的 memory-bound

  • 每个 chunk block 读 K(16KB) + V(16KB) = 32KB global,算 ~16K FMA;
  • 算术强度 AI ≈ 16K / 48KB ≈ 0.33 FLOP/byte,深度落在 roofline 的 memory-bound 区;
  • Eligible Warps/SM 仅 0.31~0.46,说明 SM 大多时间在等显存,不是算不过来。

结论:chunk kernel 的墙钟 ≈ KV字节数 / 带宽,算法层面绕不开(Flash Decoding 每个 chunk 必须独立读自己那段 K/V,跨 chunk 无 KV 复用)。

7.2 roofline 视角:斜线就是天花板

roofline 上贴着斜线 = 已撞 DRAM 带宽这堵墙。对 chunk kernel 而言:

时间 ≈ KV总字节 / 实测带宽
D=64 : 268 MB / 246 GB/s ≈ 1.09 ms  ≈ 实测 1.10 ms
D=128: 537 MB / 248 GB/s ≈ 2.16 ms  ≈ 实测 2.17 ms

两次实测都精确命中带宽地板——这是最好的证据:kernel 没有浪费,它已经把显存带宽吃满了。任何"加 Tensor Core / 提 occupancy / 寄存器调优"的尝试在这里都是无效投入(加 FLOPS 不动墙钟)。

7.3 flash vs sdpa:长序列 ~0.98x,差距在 wrapper 固定开销

多次 test_perf(取长序列场景):

Config flash (ms) sdpa (ms) flash vs sdpa
B=2 H=8 N=1 M=65536 D=64 1.12~1.13 1.10~1.11 0.97~0.98x
B=2 H=8 N=1 M=65536 D=128 2.22(三遍 2.223/2.220/2.221,无离群) 2.18~2.26 0.97~0.98x
B=2 H=8 N=1 M=8000 D=64 0.08~0.10 0.12~0.15 1.16~1.80x
  • 短序列(M=8000,D=64)flash 反超 sdpa:此时 KV 小、kernel 短,wrapper 固定开销占比低,两段式 kernel 的并行度优势显现。
  • 长序列(M=65536)flash 略慢 ~2%:但 ncu 里 chunk + reduce = 1.114 ms(D=64)/ 2.19 ms(D=128),端到端 flash 却多了 ~40µs。这多出来的时间不在 kernel 里,是 wrapper 每次调用的固定开销:

    • cudaMallocAsync + cudaFreeAsync(lse + O_chunk 缓冲)—— SDPA 零分配;
    • 两次 kernel launch + 跨 kernel 同步 —— SDPA 只一次 launch。

    另外 d=128 版把 reduce 的跨 warp 同步整个删掉后(见 6.4),长序列端到端的三遍 test_perf 标准差明显变小:旧版曾出现 2.639ms 离群,新版稳定 2.22ms——同步延迟的抖动也一并消除了,这是新写法除"更快"之外的额外红利。

关键认知:flash 与 sdpa 在长序列下本质并驾齐驱于带宽地板,残差纯属固定开销,不是 kernel 缺陷。要追平 / 反超,战场在 wrapper 而非 kernel(见第 8 节)。

7.4 D=128 vs D=64:纯带宽线性缩放,无算法退化

D=64 D=128 比值
KV 大小 268 MB 537 MB 2.00x
chunk Duration 1.10 ms 2.17 ms 1.97x
DRAM% 96.4% 96.9% 一致

chunk 时间正好 2x = KV 字节正好 2x,完美匹配带宽模型——证明 d=128 kernel 没有引入任何额外低效,只是搬了两倍字节。


8 优化空间

roofline 已贴斜线,计算侧优化全部无效。但还有几处真实杠杆:

杠杆 收益 代价 / 说明
① KV fp8 / int8 量化 唯一能翻倍的:537MB→269MB→~1.1ms 改数值精度,decode 对精度敏感需校验;生产标配
② cp.async + 异步搬运 DRAM 96%→99%(+3~5%) Eligible Warps/SM 仅 0.31~0.46,延迟没藏好;用 cp.async 让 warp 在搬运途中算 reduce,中等改动
③ 预分配 workspace 吃掉 ~2% 缺口 → flash ≈ sdpa 把 lse+O_chunk 缓冲改为一次分配、跨调用复用,去掉每次 cudaMallocAsync/cudaFreeAsync;低风险,建议优先
④ 融 reduce 进 chunk(last-block 归约) 去第二次 launch + 同步 grid.sync() 在 512 块会 deadlock,须用 last-block 模式(一个 block 等齐所有 partial 后合并);高风险
⑤ 提 chunk_size(128→512/1024) 小幅 KV 总字节不变,但块更胖→更好 MLP + 更少 reduce 迭代;d=128 下 reduce 比 d=64 重 2x,提 chunk 收益更大

9 经验教训

  1. decode 阶段的瓶颈是"并行度"不是"算力":标准 Flash Attention 沿 Q 切块,N=1 时并行度塌成 B×H;Flash Decoding 沿 K/V 切块,把并行度重建到 num_chunks×B×H,这正是它解决 SM 空转的本质。诊断 GPU 瓶颈时,先数"并行问题数 vs SM 数"。
  2. partial + LSE 加权归约是关键技巧:把"串行依赖的 softmax"解耦成"独立的 chunk partial + 一次 logsumexp 加权合并",是 Flash Decoding / 各类分块注意力(含 FlashDecoding++、PagedAttention)的公共范式。掌握 O = Σ w_c·O_c 这行公式,就掌握了这一类算子的合并逻辑。
  3. reduce 的 bug 几乎都在"列/块/warp 分工边界":写并行归约前,先静态推演"每个线程最终覆盖哪些列、哪些 chunk,边界是否需要跨 warp 合并"。
  4. lse 只从 global 读一次:归约类 kernel 里,被反复访问的小数组(如 lse)先整块搬 shared,能消掉每值多次 global 访存;顺手把 shared 原地改写成权重,还能省一次 exp。
  5. 优化必须基于 profile,且要看"端到端 vs kernel 内部"的差:flash 比 sdpa 慢那 ~2%,ncu 显示 kernel 已经 1.11ms,端到端却 1.13ms——差在 wrapper 固定开销。
  6. roofline 贴斜线 = 天花板,计算侧优化退出:一旦确认 chunk kernel 在 96% DRAM 带宽上,就该把精力从"算更快"转向"搬更少字节"(量化 KV)或"藏好延迟"(cp.async),而不是继续在算术上纠结。

10 小结

本文记录了 Flash Decoding 在 CUDA 上的完整实现:通过把并行度从 Q 行搬到 K/V 行(chunk),用 chunk kernel(算 partial O + LSE)加 reduce kernel(LSE 加权归约)的两段式结构,解决了标准 Flash Attention 在 decode 阶段(N=1)SM 占不满的问题。

核心结论:

  • 正确性:d=64 / d=128 实测 max_abs ≈ 1e-4,与 SDPA 数值等价;
  • 性能:chunk kernel 在长序列下已贴死 DRAM 带宽天花板(96%+、~247 GB/s),短序列 flash 反超 sdpa(最高 1.80x),长序列与 sdpa 并驾齐驱(0.97~0.98x),残差来自 wrapper 固定开销;
  • 优化杠杆:量化 KV(砍字节)、cp.async(吃 4% 余量)、预分配 workspace / 融 reduce(追平 sdpa)——计算侧因 memory-bound 已无空间。

与上篇 Flash Attention 形成对照:上篇是"寄存器化 + occupancy 提升"把 kernel 推过 SDPA;本篇是"分块重建并行度 + 接受带宽天花板",在 decode 这个不同场景下给出另一套答案。两者共同的哲学仍是那句——先看 profile 定瓶颈,再动手;数据比直觉可靠

完整 CUDA C++ 源码在 flash_attention 工程中(flash_decoding.cu,含 d=64 / d=128 两套 chunk + reduce kernel 及 launch 封装)。对 CUDA 算子优化感兴趣的朋友,欢迎在云栈社区交流讨论。




上一篇:三相电对地电压是多少?相电压/线电压220V、380V关系详解
下一篇:Silver Fox(银狐)虚假软件攻击分析:篡改Windows Defender实现持久化入侵
您需要登录后才可以回帖 登录 | 立即注册

手机版|小黑屋|网站地图|云栈社区 ( 苏ICP备2022046150号-2 )

GMT+8, 2026-9-6 08:41 , Processed in 0.792393 second(s), 42 queries , Gzip On.

Powered by Discuz! X3.5

© 2025-2026 云栈社区.

快速回复 返回顶部 返回列表