写在前面:本文是【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
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×H。M=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_chunk 和 lse,由 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 & 63,group = 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 经验教训
- decode 阶段的瓶颈是"并行度"不是"算力":标准 Flash Attention 沿 Q 切块,N=1 时并行度塌成 B×H;Flash Decoding 沿 K/V 切块,把并行度重建到 num_chunks×B×H,这正是它解决 SM 空转的本质。诊断 GPU 瓶颈时,先数"并行问题数 vs SM 数"。
- partial + LSE 加权归约是关键技巧:把"串行依赖的 softmax"解耦成"独立的 chunk partial + 一次 logsumexp 加权合并",是 Flash Decoding / 各类分块注意力(含 FlashDecoding++、PagedAttention)的公共范式。掌握
O = Σ w_c·O_c 这行公式,就掌握了这一类算子的合并逻辑。
- reduce 的 bug 几乎都在"列/块/warp 分工边界":写并行归约前,先静态推演"每个线程最终覆盖哪些列、哪些 chunk,边界是否需要跨 warp 合并"。
- lse 只从 global 读一次:归约类 kernel 里,被反复访问的小数组(如 lse)先整块搬 shared,能消掉每值多次 global 访存;顺手把 shared 原地改写成权重,还能省一次 exp。
- 优化必须基于 profile,且要看"端到端 vs kernel 内部"的差:flash 比 sdpa 慢那 ~2%,ncu 显示 kernel 已经 1.11ms,端到端却 1.13ms——差在 wrapper 固定开销。
- 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 算子优化感兴趣的朋友,欢迎在云栈社区交流讨论。