← 返回首页
19_flash_attention_viz.html
手写 FlashAttention:拆开 SDPA 这个黑盒
从第 6 章起,注意力一直是一行 F.scaled_dot_product_attention。它和第 2 章手写的注意力算的是同一个数,
差别在 T×T 的分数矩阵放在哪。
FlashAttention
把 Q、K、V 切块,每块在片上算完就累加进输出,用
online softmax
保证结果精确。配套代码 phase5-fullstack/17_flash_attn.py 写了三层:朴素写法、纯 PyTorch 分块、Triton kernel,
在 RTX 4090 上比时间和显存,再换进第 6 章的 124M 核对 loss。
这一关
batch 4 测速用
heads 12 同 124M
head_dim 64 同 124M
T 512 … 16,384 序列长度
块 64 × 64 BLOCK_M × BLOCK_N
dtype bf16 kernel 内部累加用 fp32
heads 12 和 head_dim 64 与第 6 章相同;124M 的上下文是 1024,更长的 T 只用来测速
① T×T 的分数矩阵有多大
朴素写法 S = QKᵀ/√d → softmax → P·V:每个头都要一张 T×T 的 S,再来一张同样大的 P。
Q、K、V、输出只随 T 线性增长,S 随 T² 增长。拖 T 看两者在显存里各占多少(batch 4、12 个头、bf16):
↳ 代码:17_flash_attn.py 的 naive_attention(第 2 章的写法)
显存够的时候,朴素写法慢在哪?
慢在读写显存。GPU 的矩阵乘单元很快,片上 SRAM 也快,但显存(HBM)带宽有限。朴素写法把 S 写回显存、softmax 再整张读出来、写出 P、P·V 再读一遍,
T² 规模的数据来回搬了好几趟,算术单元大部分时间在等数据。FlashAttention 论文 §2.2 把注意力归为受显存带宽限制的操作,
§3.2 的定理 2 给出分块版的显存访问量是 Θ(T²d²/M)(M 是 SRAM 大小),朴素版是 Θ(Td + T²)。
第 5 章的 KV cache 也和 T 有关,是同一回事吗?
不是。KV cache 存的是每个 token 的 K 和 V,随 T 线性增长,生成时一直占着;第 12 章 STEP 5 算过它的账。
本页的 S 是一次前向里的临时中间结果,随 T 平方增长,算完就可以扔。FlashAttention 省的是后者,KV cache 一个字节也不省。
↳ 下一步:不存 S,就得分块算。可 softmax 要先知道整行的最大值和总和,分块时一行被切成好几段,怎么办?
② online softmax:分段读,结果不变
一行 8 个分数,每次只读 2 个(一块)。只记两个数:到目前为止的最大值 m,以及按 m 算的指数和 l。
新一块带来更大的最大值时,旧的 l 乘上 exp(m_旧 − m_新) 缩回去。点「读下一块」:
↳ 代码:17_flash_attn.py 的 online_softmax_trace;同一组 8 个数脚本里算出的最大误差 –
输出 P·V 也是分段累加的,也要跟着缩放吗?
要。累加器 acc 存的是 Σ exp(x − m)·v,用的是当时的 m;最大值变大后,acc 和 l 乘同一个系数 alpha = exp(m_旧 − m_新)。
读完所有块,acc / l 就是 softmax(x)·V。blocked_attention 里是这三行:
l = l * alpha + p.sum(-1)、acc = acc * alpha + p @ v、最后 acc / l。
FlashAttention-2(arXiv 2307.08691)§3.1 的改动之一就是把 ÷l 从每块挪到最后只做一次。
为什么要减去最大值,不直接求 Σ exp(x)?
防溢出。fp16 最大只到 65504,exp(12) 就超出;fp32 和 bf16 在 exp(89) 附近上溢。softmax 分子分母同乘 exp(−m) 结果不变,
减去最大值后每个指数都 ≤ 1。online softmax 在此基础上只多了一步:最大值中途变了,就把之前的和按新最大值重新缩放。
递推出自 Milakov & Gimelshein,Online normalizer calculation for softmax(arXiv 1805.02867)。
↳ 下一步:一行能分段,整张注意力就能切成块。看一个 Q 块怎样扫过 K/V 块,哪些数一直留在片上。
③ 分块:哪些数留在片上
把 T×T 的注意力切成 BLOCK_M × BLOCK_N 的格子。一个 Q 块(一行格子)由一个 program 负责:把这块 Q 读进片上,
依次读 K/V 块,每块在片上算出小块 S、更新 m / l / acc,不写回显存。因果掩码下右上方的格子全被遮住,直接跳过。
选一个 Q 块,点「下一格」看它的扫描顺序:
片上 SRAM(每个 program 自己的)
显存 HBM
正在算的格子
这一行已算完
对角线格(块内还要逐元素遮)
因果掩码全遮,跳过
↳ 代码:17_flash_attn.py 的 blocked_attention(外循环 Q 块、内循环只走到对角线)
块为什么取 64 × 64?
受片上存储限制。一个 program 同时要放 Q 块(64×64)、K 块、V 块、S 块(64×64,fp32)和 acc(64×64,fp32),
bf16 下合计几十 KB,落在一个 SM 的共享内存和寄存器容量内。块再大,单个 program 放不下或并发的 program 变少;块太小,矩阵乘单元吃不饱。
FlashAttention 论文 Algorithm 1 按 SRAM 大小 M 取 B_c = ⌈M/4d⌉。实际 kernel 一般对几组块大小做自动调优,本脚本固定 64 × 64,没有调。
不同的 Q 块之间要不要通信?
不用。每个 Q 块的输出只依赖它自己的 m、l、acc,和别的 Q 块无关。所以 Triton 里 grid = (Q 块数, batch × heads),
每个 program 各干各的,4 × 12 个头、T = 4096 时一次发出 64 × 48 = 3072 个 program。
按 Q 块并行是 FlashAttention-2 的做法;第一版的外循环是 K/V 块,需要把中间结果写回显存再合并。
↳ 下一步:同一个算法,用 Triton 写成 GPU kernel。
④ 写成 Triton kernel
Triton
以"块"为单位写 GPU 程序:tl.load 读一块到片上,tl.dot 做块间矩阵乘,线程怎么分由编译器处理。
下面是 _flash_fwd 的主体,点按钮高亮对应的几行:
对数四种实现和 fp32 SDPA 的最大绝对误差
输入是 bf16 的随机 q/k/v(batch 2、12 个头),参考答案用 fp32 算。T = 1000 不是 64 的整数倍,用来检查边界:
↳ 代码:17_flash_attn.py 的 _flash_fwd / flash_attn_triton
exp2 和 1.4427 是做什么的?
exp(x) = 2x · log₂e,log₂e ≈ 1.4427。GPU 上以 2 为底的指数有专门的快速指令,
所以 kernel 把 softmax 缩放系数 1/√d 和 log₂e 乘在一起,整个 kernel 里用 tl.math.exp2。
m 和 l 也就都是按 2 为底记的,最后 acc / l 时底数抵消,结果不变。Triton 官方 fused attention 教程也是这样写的。
这个 kernel 能拿来训练吗?
不能,只有前向。训练还要反向:需要 S 和 P 才能求梯度,而前向没存它们。
FlashAttention 论文 §3.1 的办法是前向只额外存每行的 m 和 l(logsumexp),反向时按块重算 S 和 P,多算一遍矩阵乘,换来不存 T² 的中间结果。
本页只做前向,换进 124M 也只做推理评测(第 5 步)。
↳ 下一步:三种实现在 4090 上比时间和显存,再把 kernel 换进第 6 章的 124M。
⑤ 实测:时间、显存,换进 124M
batch 4、12 个头、head_dim 64、bf16,因果注意力前向,每组取 10 次的中位数。切换看时间或显存,
「显存不够」表示朴素写法在这块卡上直接报 OOM(测试时卡上另有一个常驻进程占着约 10 GB,可用约 13 GB):
124M把第 6 章模型的注意力换成三种实现
同一份权重(10B token,step 19072)、同样 65,536 个 FineWeb-Edu val token,只换注意力函数:
自己写的 kernel 和 SDPA 差不多快,SDPA 还多做了什么?
PyTorch 的 SDPA 在 CUDA 上会按输入选后端(FlashAttention-2、memory-efficient、朴素数学实现),本页的 bf16、head_dim 64 走的是 FlashAttention-2 后端。
它另有反向 kernel、dropout、各种 head_dim 和掩码,对角线块单独处理(本脚本对每一块都逐元素判断掩码),块大小按硬件调过。
在这组规模上两者接近,说明性能主要来自"不落地 T×T"这个算法本身;更长的 T 或别的显卡上差距可能不同。
三种实现的 loss 为什么差在小数点后第四位?
bf16 下每种实现的累加顺序不同:朴素写法先整行 softmax 再乘 V,分块写法按块累加、中途缩放,舍入误差落在不同的位置。
第 4 步用随机输入测过,单个输出元素的误差在 10−2 量级(bf16 的尾数只有 7 位),经过 12 层、再平均到 65,536 个 token 上,loss 差在 10−4 量级。
三者算的是同一个数学式,没有近似。
达标本关实测
↳ 跑法:python 17_flash_attn.py --json runs/ch19_flash.json(Triton 需要 NVIDIA GPU;MPS / CPU 只跑朴素与分块的对数)
🎉 手写 FlashAttention · 通关
你把 SDPA 拆成了三层:online softmax 让一行可以分段算出精确结果,分块让 T×T 的 S 从不落地,Triton kernel 把同一个算法跑在 GPU 上,换进 124M 后 loss 与原来一致。下一关回到训练:用十几个小模型预测 124M 的 loss。