所属方向:注意力与位置编码 | 难度:进阶 | 前置知识:自注意力机制的计算与显存账本(本篇是它的直接后继)
关键词:FlashAttention、online softmax、tiling、IO感知、显存优化
上一篇《自注意力机制的计算与显存账本》结尾留了一个没解的结:N=4096 的视频 DiT,按保守教学分配模型计,单层注意力为 2.39 GiB 激活,无重算训练时按同一假设累加 32 层就是 76.5 GiB,尚未包含权重、梯度、优化器与 FFN。推理的临时激活则不能直接乘层数。当时摆出了两条出路:稀疏化(算得少,但要看清赔的是什么质量)和 FlashAttention。这一篇把后者讲透。
先把一个流传很广的说法钉死:FlashAttention 不是近似注意力。它算出来的就是标准的 softmax 注意力,在实数算术下等价,浮点下允许舍入差异(第 04 节有实测:float64 下最大误差 6.1e-16,纯浮点舍入)。它快的理由也不神秘——同目录 io_ledger.py 算过一笔账,单头 d=64、N=4096 时:
两种实现的主导矩阵乘 FLOPs 相同,但重算和归一化开销不同。在这组 A100 参数的 roofline 模型中,朴素实现算术强度为 31.5 FLOP/byte,低于约 200 的平衡点;这支持优先优化 IO 的方向,不证明所有序列、硬件或近似注意力方法都必然更慢。本文 130 MiB、17.8 MiB 是脚本估算值,不是 GPU 访存计数器实测。
另一个结论更值钱:显式保存概率矩阵 P 的朴素实现仍需二次存储,S 本身通常不必一并保留,显存永远是二次的;FlashAttention 把这两个矩阵整个从显存里删掉了,训练激活从 O(N²) 降到 O(N)——在统一采用六份线性张量、忽略小的 LSE 与工作区的教学预算中,上面那个 76.5 GiB 降为 4.5 GiB。N=32768 的长视频任务,本文教学模型的朴素实现要约 4.54 TiB 激活,物理上不存在能装下的卡;FlashAttention 只要 36 GiB。整个长序列时代(32K、128K 上下文)就是踩在这个技巧上站起来的。
三句话讲完核心思想:

这张图要看什么:左边朴素实现的 S、P 两个 N×N 是该前向示意中的主要中间量,它们必须写回 HBM 再读回来;右边是同一个 N×N 被切成 B_r×B_c 的小块,K/V 块进 SRAM 常驻、Q 块逐行流过,跨块只有 O_i、l_i、m_i 三个 O(N) 的量一直活着。
设一行分数为 $s_1, \dots, s_N$,softmax 的定义是
$$p_i = \frac{\exp(s_i)}{\sum_{j=1}^{N} \exp(s_j)}$$
分子分母都是 exp 的和。直接算会溢出:s 只要有 89 左右,$\exp(s)$ 在 fp32 就到 inf 了。工程上全部改用 safe softmax——先求这行的最大值 m,再算平移后的指数:
$$p_i = \frac{\exp(s_i - m)}{\sum_{j=1}^{N} \exp(s_j - m)}, \qquad m = \max_{1 \le j \le N} s_j$$
原始分子分母同乘 $\exp(-m)$,结果不变,但指数的输入全部落在 $(-\infty, 0]$,永不溢出。问题就出在这个 m 上:m 是对整行取的 max。你必须先把 N 个分数全部看过一遍才知道 m 是多少,然后才能开始算 exp——这解释了常见实现先保存分数再做 softmax 的流程;算法并不强制保存整个 N×N,也可重算分数或逐行处理,只是 IO 和效率不同。
而注意力输出对这一行还要再多两个量:分母 $l = \sum_j \exp(s_j - m)$,以及加权和 $O = \sum_j \exp(s_j - m)\, v_j$($v_j$ 是第 j 个 token 的 Value 向量)。最终输出就是 $O / l$。
所以真正要回答的问题是:如果分数是一块一块到来的(事先不知道后面块里有什么),这三个量还能算吗?
能。做法是把「以 m 为参考系」改成「以当前的 m 为参考系,m 变了就整体换算」。
设已经流过了前 t 块,维护三个量:参考系最大值 $m^{(t)}$、分母 $l^{(t)}$、未归一化加权和 $O^{(t)}$,它们满足不变式
$$O^{(t)} = \sum_{j \le t} \exp(s_j - m^{(t)})\, v_j, \qquad l^{(t)} = \sum_{j \le t} \exp(s_j - m^{(t)})$$
(这里 $j \le t$ 是「属于前 t 块的所有下标」的缩写。)现在第 $t+1$ 块到了,块内最大值是 $m_{\text{blk}}$。新的全局最大值是
$$m^{(t+1)} = \max(m^{(t)},\ m_{\text{blk}})$$
关键一步来了:旧累积量是按 $m^{(t)}$ 为参考系记的,而新的不变式要求参考系换成 $m^{(t+1)}$。把不变式里的 $\exp(s_j - m^{(t)})$ 拆成 $\exp(s_j - m^{(t+1)}) \cdot \exp(m^{(t+1)} - m^{(t)})$,旧量的换算系数就是 $\exp(m^{(t)} - m^{(t+1)})$:
$$l^{(t+1)} = \exp(m^{(t)} - m^{(t+1)})\, l^{(t)} + \sum_{j \in \text{blk}} \exp(s_j - m^{(t+1)})$$
$$O^{(t+1)} = \exp(m^{(t)} - m^{(t+1)})\, O^{(t)} + \sum_{j \in \text{blk}} \exp(s_j - m^{(t+1)})\, v_j$$
每一步只是把指数拆成两项相乘再重新合并,等价性是代入即可验证的恒等式;所有块流完后 $O^{(T)}/l^{(T)}$ 与朴素 softmax 在实数算术下相同(差在浮点舍入,第 04 节实测 1e-16 量级)。这个「换参考系」的系数在论文和代码里叫 rescale,跨块最大值、求和与 rescale 都会带来额外的标量操作;它们不改变主导矩阵乘次数。
严谨一点可以正向验证不变式:假设第 t 步的不变式成立,那么
$$O^{(t+1)} = e^{m^{(t)} - m^{(t+1)}} \sum_{j \le t} e^{s_j - m^{(t)}} v_j + \sum_{j \in \text{blk}} e^{s_j - m^{(t+1)}} v_j = \sum_{j \le t+1} e^{s_j - m^{(t+1)}} v_j$$
(第一个等号就是递推式,第二个等号把 $e^{m^{(t)} - m^{(t+1)}}$ 乘进求和号里、指数相加后正好变回 $e^{s_j - m^{(t+1)}}$。)旧块和新块在同一个参考系下合并,不变式保持。$l$ 的证明一字不差,把 $v_j$ 去掉就行。归纳基础是初始状态 $m^{(0)} = -\infty$、$l^{(0)} = 0$、$O^{(0)} = 0$:第一块到来时换算系数按 0 处理(对应代码里 np.where(np.isneginf(m_old), 0.0, ...) 那一行),三个量直接等于第一块的局部值。
三个细节值得停一下:
顺带说一句因果掩码:掩码就是把被遮位置的 $s$ 设成 $-\infty$,exp 之后是 0,对 m、l、O 都没有贡献;更进一步,如果一整块都被遮住(Q 块整体在 K 块之前),这块连算都不用算,直接跳过。第 04 节实测这个「整块跳过」在 N=4096、块 64 时省掉 49.2% 的块。
这里的 O 指最终归一化输出。令 $G=\partial L/\partial O$、$A=GV^\top=\partial L/\partial P$。softmax 沿每行归一化,其正确反向公式是:
$$\frac{\partial L}{\partial S}=P\odot\left(A-\operatorname{rowsum}(A\odot P)\right)$$
行和为 $N\times1$,沿 key 轴广播;它是上游梯度在概率权重下的行平均,不能写成 $1-P^\top\mathbf 1$。随后 $dQ=dS\,K/\sqrt d$、$dK=dS^\top Q/\sqrt d$、$dV=P^\top G$。朴素实现可保存 P 而无需同时保存 S;FlashAttention 则重算局部 P。
FlashAttention 的做法是:除 Q/K/V 外,前向额外存 O 和 LSE($N \times d$ 加 $N$ 个数,O(N));反向时把分块流程原样再走一遍,在每一块里用 $\exp(S_{ij} - \text{LSE})$ 把局部的那一小块 P 重新算出来,立刻用于梯度,算完就扔。整个反向里 P 从头到尾没有以 N×N 的形态存在过。
用重复计算换显存——这笔交易的换算率是:每层每个头多算一遍 QK^T 和一次 exp(具体比例取决于前后向统计口径与实现),对本文保守教学账本,这对应去掉 S、P、浮点 dropout 乘子三项;真实朴素反向通常不用同时保存 S,具体减少几份取决于实现。N=4096、本文简化模型上,单层激活 2.39 GiB → 0.141 GiB,17 倍。
设单头维度 d,片上存储预算为 M 个元素(不是字节);fp16 下 192 KiB 对应 M=98304。SRAM 里要同时放下 K 块、V 块(各 $B_c \times d$)和 Q 块、输出块(各 $B_r \times d$),论文取
$$B_c = \frac{M}{4d}, \qquad B_r = \min(B_c,\ d)$$
A100 每个 SM 有 192 KB SRAM,fp16 下 d=64 时 $B_c = 384$、$B_r = 64$——这是论文 IO 模型的粗略分块预算;真实 kernel 还受 score tile、累加器、寄存器、共享内存配额与 occupancy 约束,192 KiB 也不是每个 block 可独占的共享内存。
朴素实现的搬运量(单位:元素个数):QK^T 读 Q、K 各 Nd、写 S 一次 N²;softmax 读 S 写 P 各 N²;PV 读 P 一次 N²、读 V 写 O 各 Nd。合计 $4Nd + 4N^2$,主导项 4N²。
FlashAttention 的搬运量:K、V 各进 SRAM 一次($2Nd$);Q 块每换一个 K/V 块就要重读一遍,共 $T_c \cdot Nd$($T_c = \lceil N/B_c\rceil$ 是 K/V 块数);输出 O 同理要读出写回各一遍($2 T_c Nd$);l、m 两个 O(N) 的运行量共 $4T_c N$。合计约 $2Nd + 3T_cNd + 4T_cN$,代进去:
$$\text{HBM 搬运量} \;\approx\; 3 \cdot \frac{N}{B_c} \cdot Nd \;=\; \Theta\!\left(\frac{N^2 d^2}{M}\right)$$
两个量级一比:朴素是 $\Theta(N^2)$,Flash 是 $\Theta(N^2 d^2/M)$,比值 $d^2/M$ 在 d=64、fp16、M=98304 个元素(192 KiB) 时约等于 0.04——大 O 比例省略了常数,不能直接当作 25 倍的实际收益;脚本计入 Q/O 反复读写后的模型比值为 7.3 倍。这就是整篇论文的全部:不是新数学,是把「数据在哪」当成一等公民来优化。
按本文简化模型,朴素注意力的算术强度在 N 远大于 d 时约为 d/2 FLOP/byte(fp16),所以固定 d=64 时接近 32;它会随 d 改变,并非与头维度无关。A100 示例的约 200 FLOP/byte 是 dense FP16 Tensor Core 峰值与 HBM 带宽之比,真实 softmax 的非矩阵乘指令、缓存与调度仍会影响性能。

这张图要看什么:固定 d 与 SRAM 大小时,左图两者的大 N 主导项均为 N²,比值渐近趋于常数;线性项和取整会让有限 N 的比值变化;右图斜率不同(N² 对 N),N 越大两条线离得越远,朴素教学激活曲线越过示例 80 GiB 预算线的位置,FlashAttention 还有几百倍余量。
完整脚本在文末附录(flash_online_softmax.py、io_ledger.py、make_figures.py),只依赖 numpy,全部用 /usr/local/bin/python3 实跑过,下面的数字都是真实输出。
先看一个能盯着看的例子(N=8、D=4、K/V 块大小 B_c=2,Q 整行一起处理):
def attention_flash(Q, K, V, B_r, B_c):
N, D = Q.shape
scale = 1.0 / np.sqrt(D)
O = np.zeros((N, D)) # 未归一化的输出累加器
l = np.zeros(N) # 归一化分母(exp 之和)
m = np.full(N, -np.inf) # 到目前为止见过的最大值
for j0 in range(0, N, B_c):
Kj = K[j0:j0 + B_c] # [B_c, D]
Vj = V[j0:j0 + B_c] # [B_c, D]
for i0 in range(0, N, B_r):
Qi = Q[i0:i0 + B_r] # [B_r, D]
Sij = (Qi @ Kj.T) * scale # 局部分数 [B_r, B_c];Pij 也是同尺寸临时量
m_blk = Sij.max(axis=-1) # [B_r]
m_old = m[i0:i0 + B_r]
m_new = np.maximum(m_old, m_blk)
Pij = np.exp(Sij - m_new[:, None]) # [B_r, B_c]
l_blk = Pij.sum(axis=-1)
corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new))
l[i0:i0 + B_r] = l[i0:i0 + B_r] * corr + l_blk
O[i0:i0 + B_r] = O[i0:i0 + B_r] * corr[:, None] + Pij @ Vj
m[i0:i0 + B_r] = m_new
return O / l[:, None], (m + np.log(l))
out, lse = attention_flash(Q, K, V, B_r=64, B_c=64)
print(out.shape, lse.shape) # (N, D) (N,) —— 输出和 LSE 都只有 O(N)
和第 3.2 节逐符号对上:m_new 是 $m^{(t+1)}$,corr 是换参考系的 $\exp(m^{(t)} - m^{(t+1)})$,l_blk 是块内指数和,Pij @ Vj 是块内加权和。corr 里的 np.where(np.isneginf(m_old), 0.0, ...) 处理的是第一块之前 $m = -\infty$ 的情况($-\infty - (-\infty)$ 会出 nan,直接规定换算系数为 0,旧累积量本来就是 0)。最后一行把 $m + \log l$ 作为 LSE 返回,留给反向。
实跑的逐块轨迹(第 0 号 query):
j= 0 m=+0.2968 l=1.7748 corr=0.0000
j= 2 m=+0.6071 l=3.0317 corr=0.7333
j= 4 m=+0.6071 l=3.9732 corr=1.0000
j= 6 m=+0.6071 l=5.3007 corr=1.0000
看两点:j=2 时新块里出现了更大的分数,m 被抬高、旧累积量被打了个 0.733 的折扣;j=4 之后 m 没再变,corr 恒为 1,rescale 白做——此时乘子在数学上为 1;具体 GPU kernel 是否跳过这些操作要核对实现,不能仅由轨迹推断。最大的单个临时矩阵只有 16 个元素([8, 2]),而 N×N 是 64;这里不等于同时存活临时数组的总元素数。
attention_naive 是上一篇的标准实现(S、P 都落地),两者对拍:
N D causal 最大绝对误差 相对误差
256 64 False 6.106e-16 1.311e-15
256 64 True 8.882e-16 3.520e-16
1024 64 False 6.106e-16 2.629e-15
1024 64 True 6.661e-16 2.332e-16
2048 64 False 9.437e-16 3.724e-15
2048 64 True 8.327e-16 3.570e-16
4096 64 False 7.702e-16 5.679e-15
4096 64 True 7.772e-16 2.426e-16
误差全是 1e-16 量级——float64 的舍入级别。这就是「精确注意力」四个字的实测含义:不是「误差很小」,是算法本身和朴素 softmax 完全等价。
用 tracemalloc 量函数内新分配的峰值内存(float64、D=64、块 64×64;Q/K/V 已提前分配,不计入本表):
N naive 实测 naive 理论 2N²·8 flash 实测 比值
1024 16.5 MiB 16.0 MiB 1.1 MiB 14.4x
2048 65.0 MiB 64.0 MiB 2.2 MiB 30.1x
4096 258.0 MiB 256.0 MiB 4.2 MiB 61.5x
8192 1028.0 MiB 1024.0 MiB 8.3 MiB 123.6x
N 从 4096 翻到 8192:naive 峰值 ×4.0,flash 峰值 ×2.0
naive 的实测和理论列($2N^2 \times 8$ 字节,S 和 P 两个 N×N)对得上,说明量的方法可信。最后那行是阶数的直接证据:N 翻倍,naive 峰值 ×4(二次),flash 峰值 ×2(线性)。
N naive flash(B=64) flash/naive
1024 7.2 ms 10.2 ms 1.4x
2048 26.3 ms 42.4 ms 1.6x
4096 98.7 ms 166.3 ms 1.7x
CPU + numpy 上 flash 慢约 1.6 倍。三个原因:乘加次数一样还多了 rescale 和逐块 exp;一次大矩阵乘被拆成 (N/64)² 个 64×64 小矩阵乘,BLAS 跑不满小块,Python 循环开销也进来了;CPU 也有 SRAM 缓存,但这里的 Python/NumPy 分块没有实现专门的缓存与线程优化,不能把它当作 GPU 内核速度的预测。第三张图展示 IO 模型为何支持在 GPU 上尝试这一优化。
朴素实现加因果掩码,N×N 还是得整块算完再往被遮的位置上写 $-\infty$,一个 FLOP 都省不下来。FlashAttention 按「Q 块整体在 K 块之前就整块跳过」处理,N=4096、块 64 时实测:
块总数(非因果) : 4096
块总数(因果跳过): 2080
跳过的块占比 : 49.2%
保留的是含对角线的下三角,块数是 $T(T+1)/2$,占比 $(T+1)/2T$,T 大时趋近一半。训练 GPT 类因果模型、以及视频 DiT 里的时序因果注意力,这半是免费的。
最小实现讲清了原理,但生产 kernel 和它有四处本质差异,每处都值得知道为什么:
第一,块大小不是从公式算的,是 autotune 出来的。 第 3.4 节的 $B_c = M/4d$ 是IO 分析采用的可行块预算;真实的 flash-attention kernel(flash_attn/flash_attn_interface.py 的 flash_attn_func,以 2025-09 的实现为准)里,块大小是按(头维度、数据类型、是否因果、显存架构)在若干组预编译配置里选的,还受 warp 数量、寄存器压力、shared memory bank conflict 的影响——公式只负责告诉你「必须小于某个数」,调优负责在约束内找最快的。
第二,减少非矩阵乘操作。 FlashAttention-2 使用未归一化输出累计,减少 rescale、除法等非矩阵乘工作。corr=1 时数学上无需改变旧值,但不能笼统声称所有 kernel 都按每行最大值是否变化来分支跳过。
第三,前向的结构是「外层 Q、内层 K/V」。 我们按论文 v1 的写法外层遍历 K/V 块;FlashAttention-2 把循环反过来(外层 Q 块),好处是输出 O 常驻寄存器不用反复读写、且不同 Q 块之间天然并行,能吃满更多 SM。论文 v1 的伪代码适合理解递推,v2 的循环结构才是现在 kernel 的样子。
第四,dropout 不存掩码,存随机数种子。 朴素实现要为反向留一个 B×H×N×N 的 dropout 掩码;kernel 里只存 Philox 计数器的 seed 和 offset(几十字节),反向时用同一个种子重新生成同样的掩码。这是「重算换显存」哲学最极致的一次应用——连随机数本身都可以重算。
另外两条工程事实:PyTorch 2.0 起 F.scaled_dot_product_attention 会自动按(头维度、掩码、数据类型、硬件)在 flash / memory-efficient / math 三个后端里挑,你不写一行 CUDA 也在用它;论文报告的端到端收益是 BERT-large(seq 512)比 MLPerf 1.1 训练记录快 15%、GPT-2(seq 1K)快 3 倍、Long Range Arena(seq 1K-4K)快 2.4 倍——注意 seq 512 时只有 15%,因为那时注意力在整层里占比还小,收益随序列长度涨,这正是 IO 复杂度模型的预测。
排查问题时你会想知道「此刻到底在用哪个后端」。PyTorch 留了一个官方口子:
from torch.nn.attention import sdpa_kernel, SDPBackend
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
# 强制走 flash;如果这个头维度/掩码组合它不支持,这里会直接报错,
# 而不是悄悄退回 math 后端——「悄悄降级」正是性能莫名掉一半时最该先查的事

这张图要看什么:横轴是 N,纵轴是「每搬 1 字节做多少次运算」,灰色虚线是机器平衡点(201 FLOP/byte)。本模型中朴素实现的强度渐近接近 32;FlashAttention 的估计从 N=2048 起超过平衡点。这是按矩阵乘峰值做的模型分类,真实瓶颈还受非矩阵乘指令、缓存、并行度和调度影响。
FlashAttention 省下了 HBM 搬运和 N² 显存,赔进去的和没管住的也要说清楚。
代价:重计算和额外归一化操作。 反向要重算局部分数及概率,其中包含矩阵乘和逐元素运算,不能统称为固定 30% 的额外非矩阵乘开销。小序列的收益取决于 kernel、调度和硬件,没有统一的 seq<512 亏损阈值。
数值边界:实数算术等价不保证浮点逐位一致。 kernel 内部用 fp16/bf16 存储、fp32 累加,块内的归一化和朴素实现的一次性归一化在浮点上不同。对训练的影响需要结合 dtype、输入尺度与误差测试判断,但如果你在做数值敏感的分析(比如逐 token 概率对比),要知道它和参考实现差在舍入级别,不是 bug。
边界的核心一条:它没有改变复杂度,改变的常数。 显存从 $O(N^2)$ 降到 $O(N)$,但算力还是 $\Theta(N^2 d)$、HBM 搬运还是 $\Theta(N^2 d^2/M)$。N=32768 时 FlashAttention 的单头搬运是 1061 MiB——比朴素实现的 8.2 GiB 好得多,但随 N 继续平方增长这一点没变。上下文再往上涨(1M token),接力棒要交给稀疏注意力、线性注意力、状态空间模型这些真正改复杂度的方法。FlashAttention 的块结构恰恰是它们的底座:把注意力切成块之后,「整块跳过」才成为可能,第 4.5 节那个 49.2% 推广到任意稀疏模式就是块稀疏注意力。
不该用的场景:需要拿到完整注意力权重做分析或可视化的(P 从头到尾没存在过,想看它就得回到朴素实现);自定义的任意注意力偏置如果 kernel 不支持,绕过去的方法可能把优势吃掉;以及缺乏适配内核的环境;CPU 也有 SRAM 缓存,但本文 NumPy 循环没有实现专门的 CPU cache 优化(第 4.4 节的 CPU 实测就是例子)。
这条线的演进关系一句话各说清:
一条清晰的线:2018 年有技巧,2021 年有证明,2022 年才有产品——缺的从来不是数学,是「意识到瓶颈在 IO」这个视角。
以下几条都值得单独记住,前两条我当初也信过:
跑文末附录的 flash_online_softmax.py(只依赖 numpy),预期输出与正文一致:
再做一个一行代码的实验,直接看清递推式里 rescale 的分量:把 corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) 改成 corr = np.ones_like(m_old)(假装 m 永远不变,也就是退回「先见全家再算」之前的朴素流式假设),重跑等价性检验。我实测过:N=1024、D=64、随机高斯输入下,最大误差从 6.1e-16 恶化到 0.276,平均误差 0.0136——输出在量级上就是错的。这个对比说明:流式计算 softmax 时,「用新参考系换算旧累积量」这一步不是工程细节,是正确性本身。
最后改 io_ledger.py 开头的硬件常数(比如把 SRAM_PER_SM 调到 48 KB 模拟消费级卡),重跑看块大小和搬运量比值怎么变——你会看到 SRAM 越小,FlashAttention 相对朴素实现的搬运量优势越小,$d^2/M$ 里的 M 直接控制这一切。
按知识树的依赖关系,建议按这个顺序继续走:
延伸到知识树之外:想读 kernel 源码,从 Dao-AILab/flash-attention 的 flash_attn/flash_attn_interface.py 进,先读 forward 再读 backward;想读原始推导,Milakov & Gimelshein(1805.02867)给出了独立 softmax 在线归一化的直接推导。
09 节用到的脚本全文如下(io_ledger.py、flash_online_softmax.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""IO 账本:朴素注意力和 FlashAttention 在 HBM 上到底搬了多少字节。
这篇讲的是「为什么快」。答案不在 FLOPs 上——两者的乘加次数几乎一样——
而在于 GPU 有两级内存:
HBM(显存) 带宽约 1.5 TB/s,容量 40~80 GB
SRAM(片上) 带宽约 19 TB/s,但每个 SM 只有 192 KB
朴素实现把 N×N 的分数矩阵写回 HBM 再读出来,等于把数据在这条 13 倍带宽差的
通道上来回搬;FlashAttention 用 tiling 让这些中间结果根本不落 HBM。
下面的模型把每一笔读写都点清楚,参数是可调的,读者可以改硬件常数重算。
硬件数量级取自 FlashAttention 论文 Table 1(A100 40GB)。
只依赖 numpy。直接 `python io_ledger.py` 即可运行。
"""
import numpy as np
# ── 硬件常数(A100 40GB 量级,论文 Table 1)────────────────────
HBM_BW = 1.555e12 # HBM 带宽,字节/秒
SRAM_BW = 19.0e12 # 片上 SRAM 带宽,字节/秒
SRAM_PER_SM = 192 * 1024 # 每个 SM 的片上 SRAM,字节
FLOPS_PEAK = 312e12 # fp16 tensor core 峰值,FLOP/秒
GIB = 1024 ** 3
MIB = 1024 ** 2
# ────────────────────────────────────────────────────────────
# 块大小:SRAM 里能同时放下什么
# ────────────────────────────────────────────────────────────
def block_sizes(d, sram_bytes=SRAM_PER_SM, dtype_bytes=2):
"""返回 (B_r, B_c):Q 块行数与 K/V 块行数。
SRAM 里要同时放下 K_j、V_j 两个 [B_c, d] 和 Q_i、O_i 两个 [B_r, d]。
论文的取法是 B_c = M / (4d)、B_r = min(B_c, d),这里照抄:
先让 K_j+V_j 占掉一半 SRAM,Q 块则不超过 d 行(保证 softmax 按行算得下)。
"""
elems = sram_bytes / dtype_bytes
B_c = max(1, int(elems // (4 * d)))
B_r = min(B_c, d)
return B_r, B_c
# ────────────────────────────────────────────────────────────
# HBM 读写量(单位:元素个数,乘 dtype_bytes 得字节)
# ────────────────────────────────────────────────────────────
def hbm_elems_naive(N, d):
"""朴素实现:S 和 P 都要落地。
QK^T: 读 Q(Nd) + 读 K(Nd) + 写 S(N²)
softmax: 读 S(N²) + 写 P(N²)
PV: 读 P(N²) + 读 V(Nd) + 写 O(Nd)
"""
return 4 * N * d + 4 * N * N
def hbm_elems_flash(N, d, B_r, B_c):
"""FlashAttention:外层遍历 K/V 块,内层遍历 Q 块。
K、V 各读一遍(每个 j 块进 SRAM 后,内层 i 循环里一直复用)
Q 每个 j 都要重读一遍:T_c · N·d
O 每个 (j,i) 都要读出来再写回去(累加器跨 j 迭代):2 · T_c · N·d
l、m 两个 O(N) 的运行量同理:4 · T_c · N
"""
T_c = int(np.ceil(N / B_c))
return 2 * N * d + 3 * T_c * N * d + 4 * T_c * N
def flops_attention(N, d):
"""两个 N×N 矩阵乘,一次 [M,K]×[K,N] 算 2MKN 个浮点运算。"""
return 4 * N * N * d
def activation_bytes(N, d, H, B=1, dtype_bytes=2):
"""单层注意力的训练激活(反向要用的中间张量)。
保守教学模型:6 个 [B,N,d_model] 线性项 + S/P/浮点乘子 3 个 [B,H,N,N]。
此函数 d 是总宽度 d_model,不是其他 IO 函数中的头宽度。
融合侧保留同样六份线性预算,忽略小的 LSE 与工作区;不是框架峰值测量。
"""
lin = 6 * B * N * d * dtype_bytes
quad = 3 * B * H * N * N * dtype_bytes
return lin + quad, lin
def roofline(bytes_moved, flops):
"""算术强度(FLOP/byte)与两个上界时间。返回 (强度, 内存时间, 算力时间, 瓶颈)。"""
intensity = flops / bytes_moved
t_mem = bytes_moved / HBM_BW
t_comp = flops / FLOPS_PEAK
bound = "内存受限" if t_mem > t_comp else "算力受限"
return intensity, t_mem, t_comp, bound
# ────────────────────────────────────────────────────────────
# 报表
# ────────────────────────────────────────────────────────────
def report_io():
d, b = 64, 2 # 单头维度 64,fp16
B_r, B_c = block_sizes(d)
print("=" * 74)
print("1. SRAM 块大小与 HBM 读写量(单头,d=64,fp16)")
print("=" * 74)
print(f" SRAM {SRAM_PER_SM / 1024:.0f} KB / SM,fp16 下能放 "
f"{SRAM_PER_SM / b:.0f} 个元素")
print(f" → B_c = M/(4d) = {B_c},B_r = min(B_c, d) = {B_r}\n")
print(f"{'N':>7} {'naive HBM':>12} {'flash HBM':>12} {'比值':>8} "
f"{'naive 强度':>11} {'flash 强度':>11}")
rows = []
for N in (1024, 2048, 4096, 8192, 16384, 32768):
nb = hbm_elems_naive(N, d) * b
fb = hbm_elems_flash(N, d, B_r, B_c) * b
fl = flops_attention(N, d)
i_n, _, _, bound_n = roofline(nb, fl)
i_f, _, _, bound_f = roofline(fb, fl)
rows.append((N, nb, fb, i_n, i_f, bound_n, bound_f))
print(f"{N:>7} {nb / MIB:>10.1f} MiB {fb / MIB:>10.1f} MiB "
f"{nb / fb:>7.1f}x {i_n:>9.1f} {i_f:>9.1f} ")
print(f"\n 机器平衡点(峰值算力/带宽)= {FLOPS_PEAK / HBM_BW:.0f} FLOP/byte")
print(" 强度低于它 → 内存受限,加算力没用;高于它 → 才开始吃算力。\n")
print(f"{'N':>7} {'naive 瓶颈':>12} {'flash 瓶颈':>12}")
for N, nb, fb, i_n, i_f, bn, bf in rows:
print(f"{N:>7} {bn:>12} {bf:>12}")
print("\n 注意:flash 的强度在 N 大时越过平衡点,模型说它变成算力受限了。")
print(" 但真实 kernel 达不到峰值——softmax 的 exp 走的是特殊函数单元,")
print(" 不走 tensor core,这个「非矩阵乘开销」正是 FlashAttention-2 之后")
print(" 继续优化的地方。模型给出的是上界,不是承诺。")
def report_activation():
d, H, B, b = 3072, 24, 1, 2
print("\n" + "=" * 74)
print("2. 教学激活存储模型(简化 Transformer:d=3072, H=24, B=1, 32 层, bf16)")
print("=" * 74)
print(f"{'N':>7} {'朴素/层':>12} {'Flash/层':>12} {'比值':>8} "
f"{'朴素 32 层':>12} {'Flash 32 层':>13}")
for N in (1024, 4096, 8192, 32768):
naive, flash = activation_bytes(N, d, H, B, b)
print(f"{N:>7} {naive / GIB:>10.2f} GiB {flash / GIB:>10.3f} GiB "
f"{naive / flash:>7.0f}x {32 * naive / GIB:>10.1f} GiB "
f"{32 * flash / GIB:>11.2f} GiB")
n4096, f4096 = activation_bytes(4096, d, H, B, b)
print(f"\n N=4096 时,朴素实现单层 {n4096 / GIB:.2f} GiB,其中二次项占 "
f"{100 * (1 - f4096 / n4096):.1f}%——")
print(" 这正是 attention_basics 那篇里「三个 N×N 吃掉 94%」的那一格。")
print(" FlashAttention 删掉的就是这一格,剩下的是 O(N) 的线性项。")
def report_speed_limit():
d, b = 64, 2
B_r, B_c = block_sizes(d)
print("\n" + "=" * 74)
print("3. 理论上界:如果只受 HBM 带宽限制,两者各要多久(单头,d=64)")
print("=" * 74)
print(f"{'N':>7} {'FLOPs':>10} {'naive 内存时间':>16} {'flash 内存时间':>16} "
f"{'纯带宽模型比':>9}")
for N in (1024, 4096, 16384):
nb = hbm_elems_naive(N, d) * b
fb = hbm_elems_flash(N, d, B_r, B_c) * b
fl = flops_attention(N, d)
print(f"{N:>7} {fl / 1e9:>8.2f} G {nb / HBM_BW * 1e3:>14.3f} ms "
f"{fb / HBM_BW * 1e3:>14.3f} ms {nb / fb:>8.1f}x")
print("\n 此列只比较理想 HBM 时间,并非完整模型训练或实际 kernel 加速上界。")
print(" 论文的端到端训练加速与此处单头 IO 模型口径不同,不能直接相比。")
if __name__ == "__main__":
report_io()
report_activation()
report_speed_limit()
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""online softmax 的正确性与显存代价实测(FlashAttention 的核心那一步)。
对比两种实现:
naive —— 先算出完整的 N×N 分数矩阵 S,再逐行 softmax 得到 P,最后算 O = P V
flash —— 按块流过 K/V,只维护三个 O(N) 的运行量:输出累加器 O、分母 l、
运行最大值 m(第 03 节递推式的直接实现)
两者在数学上完全等价,唯一差别是中间有没有 N×N 的矩阵落到内存里。
这份脚本回答三个问题:
1. 递推式写出来的结果,和朴素 softmax 逐位一致吗?(1.1 节)
2. 峰值内存真的差一个 N 吗?(用 tracemalloc 量,不是估的)
3. 那算力呢?——在 CPU + numpy 上 flash 是**更慢**的,这一点必须诚实讲清楚
只依赖 numpy。直接 `python flash_online_softmax.py` 即可运行。
"""
import time
import tracemalloc
import numpy as np
NEG = -np.inf
# ────────────────────────────────────────────────────────────
# 两个被测实现
# ────────────────────────────────────────────────────────────
def attention_naive(Q, K, V, causal=False):
"""标准实现:S 和 P 都是完整的 N×N 常驻张量。"""
D = Q.shape[-1]
S = (Q @ K.T) / np.sqrt(D) # [N, N] ← 第一块 N×N
if causal:
S = np.where(np.triu(np.ones((Q.shape[0], K.shape[0])), 1) > 0, NEG, S)
S -= S.max(axis=-1, keepdims=True) # safe softmax,不改变结果
P = np.exp(S) # [N, N] ← 第二块 N×N
P /= P.sum(axis=-1, keepdims=True)
return P @ V
def attention_flash(Q, K, V, B_r, B_c, causal=False, trace=False):
"""按块流过 + online softmax。全程不出现 N×N 的张量。
B_r / B_c 分别是 Q 块和 K/V 块的行数,对应 SRAM 里各放得下多少行。
"""
N, D = Q.shape
scale = 1.0 / np.sqrt(D)
O = np.zeros((N, D)) # 未归一化的输出累加器
l = np.zeros(N) # 归一化分母(exp 之和)
m = np.full(N, NEG) # 到目前为止见过的最大值
peak_tmp = 0 # 记录出现过的最大临时矩阵(元素个数)
for j0 in range(0, N, B_c):
Kj = K[j0:j0 + B_c] # [B_c, D]
Vj = V[j0:j0 + B_c] # [B_c, D]
for i0 in range(0, N, B_r):
# 因果掩码下,若整个 Q 块都在 K 块之前(所有 query 下标 < 所有 key
# 下标),这一块全被遮掉,连算都不用算 —— 朴素实现做不到这一点。
# 条件:块内最大 query 下标 i0+B_r-1 < j0
if causal and i0 + B_r <= j0:
continue
Qi = Q[i0:i0 + B_r] # [B_r, D]
Sij = (Qi @ Kj.T) * scale # [B_r, B_c] 局部分数;Pij 也是临时块
if causal:
q_idx = i0 + np.arange(Sij.shape[0])[:, None]
k_idx = j0 + np.arange(Sij.shape[1])[None, :]
Sij = np.where(k_idx > q_idx, NEG, Sij)
peak_tmp = max(peak_tmp, Sij.size)
m_blk = Sij.max(axis=-1) # [B_r]
m_old = m[i0:i0 + B_r]
m_new = np.maximum(m_old, m_blk) # [B_r]
Pij = np.exp(Sij - m_new[:, None]) # [B_r, B_c]
l_blk = Pij.sum(axis=-1) # [B_r]
# 把之前累积的量从旧的参考最大值搬到新的(关键的 rescale 一步)
corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new))
l[i0:i0 + B_r] = l_old = l[i0:i0 + B_r] * corr + l_blk
O[i0:i0 + B_r] = O[i0:i0 + B_r] * corr[:, None] + Pij @ Vj
m[i0:i0 + B_r] = m_new
if trace:
print(f" j={j0:>2} i={i0:>2} m={m_new[0]:+.4f} "
f"l={l_old[0]:.4f} corr={corr[0]:.4f}")
return O / l[:, None], (m + np.log(l)), peak_tmp
def peak_bytes(fn, *args, **kwargs):
"""跑一次 fn,返回 (结果, 峰值字节数)。用 tracemalloc 实测。"""
tracemalloc.start()
tracemalloc.reset_peak()
out = fn(*args, **kwargs)
_, peak = tracemalloc.get_traced_memory()
tracemalloc.stop()
return out, peak
def timed(fn, *args, repeat=3, **kwargs):
best = float("inf")
out = None
for _ in range(repeat):
t0 = time.perf_counter()
out = fn(*args, **kwargs)
best = min(best, time.perf_counter() - t0)
return out, best
# ────────────────────────────────────────────────────────────
# 1. 递推式长什么样:一个能逐块打印的最小例子
# ────────────────────────────────────────────────────────────
def demo_recurrence():
rng = np.random.default_rng(0)
N, D, B_c = 8, 4, 2
Q = rng.normal(size=(N, D))
K = rng.normal(size=(N, D))
V = rng.normal(size=(N, D))
print("=" * 68)
print("1. online softmax 的递推过程(N=8, D=4, K/V 块大小 B_c=2)")
print("=" * 68)
print(" Q 块固定为整行(B_r=N),K/V 分成 4 块依次流过;")
print(" 每行打印第 0 号 query 的 m / l / corr,看它们怎么被逐次修正:\n")
_, _, peak = attention_flash(Q, K, V, B_r=N, B_c=B_c, trace=True)
print(f"\n 最大的单个临时矩阵只有 {peak} 个元素(并非临时内存总和) = [{N}, {B_c}],而 N×N = {N * N}")
# ────────────────────────────────────────────────────────────
# 2. 等价性:和朴素 softmax 逐位对得上吗
# ────────────────────────────────────────────────────────────
def demo_exactness():
print("\n" + "=" * 68)
print("2. 数值等价性(float64,非因果 / 因果两种掩码)")
print("=" * 68)
print(f"{'N':>6} {'D':>4} {'causal':>7} {'最大绝对误差':>14} {'相对误差':>12}")
for N in (256, 1024, 2048, 4096):
D = 64
rng = np.random.default_rng(N)
Q = rng.normal(size=(N, D))
K = rng.normal(size=(N, D))
V = rng.normal(size=(N, D))
for causal in (False, True):
ref = attention_naive(Q, K, V, causal)
got, _, _ = attention_flash(Q, K, V, B_r=64, B_c=64, causal=causal)
diff = np.abs(ref - got).max()
rel = diff / np.abs(ref).max()
print(f"{N:>6} {D:>4} {str(causal):>7} {diff:>14.3e} {rel:>12.3e}")
print("\n 误差量级是浮点舍入(1e-15),不是近似——FlashAttention 是精确算法。")
# ────────────────────────────────────────────────────────────
# 3. 峰值内存:是不是真的差一个 N
# ────────────────────────────────────────────────────────────
def demo_memory():
print("\n" + "=" * 68)
print("3. 峰值内存实测(tracemalloc,float64,D=64,B_r=B_c=64)")
print("=" * 68)
print(f"{'N':>6} {'naive 实测':>12} {'naive 理论 2N²·8':>18} "
f"{'flash 实测':>12} {'比值':>8}")
peaks = {}
for N in (1024, 2048, 4096, 8192):
D = 64
rng = np.random.default_rng(N)
Q = rng.normal(size=(N, D))
K = rng.normal(size=(N, D))
V = rng.normal(size=(N, D))
_, p_naive = peak_bytes(attention_naive, Q, K, V)
_, p_flash = peak_bytes(attention_flash, Q, K, V, 64, 64)
peaks[N] = (p_naive, p_flash)
print(f"{N:>6} {p_naive / 2**20:>10.1f} MiB "
f"{2 * N * N * 8 / 2**20:>16.1f} MiB "
f"{p_flash / 2**20:>10.1f} MiB {p_naive / p_flash:>7.1f}x")
g_n = peaks[8192][0] / peaks[4096][0]
g_f = peaks[8192][1] / peaks[4096][1]
print(f"\n N 从 4096 翻到 8192:naive 峰值 ×{g_n:.1f},flash 峰值 ×{g_f:.1f}")
print(" 这就是 O(N²) 和 O(N) 的区别:前者翻两倍(4×),后者跟着翻倍(2×)。")
print(" 朴素实现要同时留住 S 和 P 两个 N×N(理论列就是 2N²·8 字节,和实测对得上);")
print(" flash 只留 B_r×B_c 的块,剩下的是 O(N·D) 的输出累加器。")
# ────────────────────────────────────────────────────────────
# 4. 时间:CPU + numpy 上 flash 反而更慢,这才是重点
# ────────────────────────────────────────────────────────────
def demo_time():
print("\n" + "=" * 68)
print("4. 墙钟时间(CPU + numpy,BLAS 多线程)——反直觉的一项是这个")
print("=" * 68)
print(f"{'N':>6} {'naive':>10} {'flash(B=64)':>13} {'flash/naive':>12}")
for N in (1024, 2048, 4096):
D = 64
rng = np.random.default_rng(N)
Q = rng.normal(size=(N, D))
K = rng.normal(size=(N, D))
V = rng.normal(size=(N, D))
_, t_naive = timed(attention_naive, Q, K, V)
_, t_flash = timed(attention_flash, Q, K, V, 64, 64)
print(f"{N:>6} {t_naive * 1e3:>9.1f} ms {t_flash * 1e3:>12.1f} ms "
f"{t_flash / t_naive:>11.1f}x")
print("\n 以上比值是本机本次计时,不能据此预测 GPU;可能影响速度的因素包括:")
print(" 1. 乘加次数一模一样,还额外多了每块的 rescale 和逐元素 exp;")
print(" 2. 一次大矩阵乘被拆成 (N/B)² 次 64×64 的小矩阵乘,")
print(" BLAS 在小块上根本跑不满,Python 循环开销也进来了;")
print(" 3. 最关键的:它省的是**内存搬运**,不是 FLOPs,")
print(" CPU 同样有 SRAM 缓存,但本示例没有专门优化缓存和线程,")
print(" 因此分块可能省容量,却被 Python 循环与小 GEMM 开销抵消。")
print(" 第 05 节讲 GPU 上为什么结论会反过来。")
# ────────────────────────────────────────────────────────────
# 5. 因果掩码:flash 能顺手省掉一半算力,朴素实现不能
# ────────────────────────────────────────────────────────────
def demo_causal_blocks():
print("\n" + "=" * 68)
print("5. 因果掩码下实际算了多少块(N=4096, B_r=B_c=64)")
print("=" * 68)
N, B = 4096, 64
T_r = T_c = N // B
total = 0
for j in range(T_c):
for i in range(T_r):
if i * B + B <= j * B: # 整个 Q 块都在 K 块之前 → 全遮,跳过
continue
total += 1
print(f" 块总数(非因果) : {T_r * T_c}")
print(f" 块总数(因果跳过): {total}")
print(f" 跳过的块占比 : {100.0 * (1 - total / (T_r * T_c)):.1f}%")
print("\n 朴素实现就算加了因果掩码,N×N 的矩阵照样得整块算完再遮,")
print(" 省不了一点算力;flash 是整块跳过,这是免费的一半(严格说是")
print(" (T²+T)/2 / T² ≈ 一半多一点的对角块保留)。")
if __name__ == "__main__":
np.set_printoptions(precision=4, suppress=True)
demo_recurrence()
demo_exactness()
demo_memory()
demo_time()
demo_causal_blocks()
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""生成「FlashAttention」的三张解释图。
数值全部来自同目录下的 io_ledger.py(HBM 读写量、激活显存、算术强度),
改了那个脚本的话这里要跟着重跑,避免图与正文数字不一致。
三张图分别回答:
1. tiling 到底怎么切、什么留在 SRAM、什么留在 HBM
2. 两本账随 N 怎么长(HBM 读写量、教学激活存储模型)
3. 为什么说是「内存受限」变的「不那么受限」(算术强度对机器平衡点)
只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。
"""
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
from matplotlib.patches import FancyArrowPatch, FancyBboxPatch, Rectangle
from io_ledger import (
FLOPS_PEAK, HBM_BW,
activation_bytes, block_sizes, flops_attention,
hbm_elems_flash, hbm_elems_naive,
)
ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "figures"
OUT.mkdir(exist_ok=True)
plt.rcParams.update({
"font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"],
"axes.unicode_minus": False,
"figure.dpi": 160,
})
C_NAIVE = "#e05263" # 朴素实现:红
C_FLASH = "#2f9e6f" # FlashAttention:绿
C_SRAM = "#f0a03c" # SRAM 高亮:橙
C_GREY = "#94a3b8"
C_SKIP = "#e4e9f0"
INK = "#182238"
MIB = 1024 ** 2
GIB = 1024 ** 3
def _style(ax, title, xlabel=None, ylabel=None):
ax.set_title(title, fontsize=13.5, weight="bold", color=INK, pad=10)
if xlabel:
ax.set_xlabel(xlabel, fontsize=11, color="#475569")
if ylabel:
ax.set_ylabel(ylabel, fontsize=11, color="#475569")
ax.tick_params(colors="#475569", labelsize=10)
for s in ("top", "right"):
ax.spines[s].set_visible(False)
for s in ("left", "bottom"):
ax.spines[s].set_color("#cbd5e1")
ax.set_axisbelow(True)
def _box(ax, x, y, w, h, text, fc, fs=10.5, tc="white"):
ax.add_patch(FancyBboxPatch(
(x, y), w, h, boxstyle="round,pad=0.02,rounding_size=0.12",
linewidth=0, facecolor=fc, edgecolor="none"))
ax.text(x + w / 2, y + h / 2, text, ha="center", va="center",
fontsize=fs, color=tc, weight="bold", linespacing=1.5)
def _arrow(ax, x1, y1, x2, y2, label=None, color="#475569"):
ax.add_patch(FancyArrowPatch(
(x1, y1), (x2, y2), arrowstyle="-|>", mutation_scale=13,
linewidth=1.4, color=color, shrinkA=2, shrinkB=2))
if label:
ax.text((x1 + x2) / 2, max(y1, y2) + 0.22, label, ha="center",
fontsize=9, color=color)
def _caption(fig, text):
fig.text(0.5, 0.012, text, ha="center", fontsize=9.5, color="#94a3b8")
# ────────────────────────────────────────────────────────────
# 图 1:tiling 怎么切
# ────────────────────────────────────────────────────────────
def fig_tiling():
fig = plt.figure(figsize=(13.6, 6.0), facecolor="#fbfcfe")
# ── 左:朴素实现 ──
ax = fig.add_axes([0.03, 0.09, 0.46, 0.80])
ax.set_xlim(0, 10); ax.set_ylim(0, 7.4); ax.axis("off")
ax.text(0, 7.05, "本例朴素前向:物化 S 与 P", fontsize=13.5,
weight="bold", color=C_NAIVE)
_box(ax, 0.1, 5.1, 1.5, 1.1, "Q, K, V\n[N, d]", "#64748b", fs=10)
_box(ax, 2.4, 4.8, 2.0, 1.7, "S = QK^T / √d\n[N, N]", C_NAIVE, fs=11)
_box(ax, 5.3, 4.8, 2.0, 1.7, "P = softmax(S)\n[N, N]", C_NAIVE, fs=11)
_box(ax, 8.2, 5.1, 1.6, 1.1, "O = PV\n[N, d]", "#64748b", fs=10)
_arrow(ax, 1.65, 5.65, 2.35, 5.65)
_arrow(ax, 4.45, 5.65, 5.25, 5.65, "读 S 写 P")
_arrow(ax, 7.35, 5.65, 8.15, 5.65, "读 P")
ax.text(4.4, 4.25, "N=4096, d=64, fp16:S 和 P 各 32 MiB",
ha="center", fontsize=9.5, color="#475569")
ax.add_patch(Rectangle((0.1, 2.2), 9.7, 1.35, facecolor="#eef2f7",
edgecolor="#cbd5e1", linewidth=1.2))
ax.text(4.95, 3.2, "HBM 带宽 1.5 TB/s", ha="center", fontsize=10.5,
color="#475569", weight="bold")
ax.text(4.95, 2.5, "S 写一次读一次,P 写一次读一次 → 4N² 次元素搬运",
ha="center", fontsize=9.5, color="#64748b")
_arrow(ax, 3.4, 4.75, 3.4, 3.6, color=C_NAIVE)
_arrow(ax, 6.3, 4.75, 6.3, 3.6, color=C_NAIVE)
ax.text(0.1, 1.35, "代价:数据在这条通道上来回两趟,",
fontsize=10.5, color=C_NAIVE, weight="bold")
ax.text(0.1, 0.75, "而 HBM 比 SRAM 慢 12 倍", fontsize=10.5,
color=C_NAIVE, weight="bold")
# ── 右:FlashAttention ──
ax2 = fig.add_axes([0.53, 0.09, 0.44, 0.80])
ax2.set_xlim(0, 10); ax2.set_ylim(0, 7.4); ax2.axis("off")
ax2.text(0, 7.05, "FlashAttention:按块流过,只留 O(N)",
fontsize=13.5, weight="bold", color=C_FLASH)
T = 8
gx0, gy0, cell = 1.05, 3.1, 0.46
for i in range(T):
for j in range(T):
if i + 1 <= j: # 因果掩码下整块跳过(先判,优先级最高)
fc, ec = C_SKIP, "#c3ccd8"
elif j == 3: # 当前正在处理的 K/V 块列
fc, ec = "#cdeadb", C_SRAM
elif i == 5: # 当前 Q 块行
fc, ec = "#dfeaf6", "#94a3b8"
else:
fc, ec = "#f7fafc", "#dde3ea"
ax2.add_patch(Rectangle(
(gx0 + j * cell, gy0 + (T - 1 - i) * cell),
cell * 0.90, cell * 0.90,
facecolor=fc, edgecolor=ec, linewidth=1.1))
gtop = gy0 + T * cell
ax2.text(gx0 - 0.42, gy0 + T * cell / 2, "Q 块逐行流过",
fontsize=9.5, color="#475569", ha="center", va="center",
rotation=90)
ax2.add_patch(FancyBboxPatch(
(gx0 + 3 * cell - 0.06, gy0 - 0.06), cell * 1.02, T * cell + 0.12,
boxstyle="round,pad=0.03,rounding_size=0.1",
linewidth=1.6, edgecolor=C_SRAM, facecolor="none", linestyle="--"))
ax2.text(gx0 + 3.5 * cell, gy0 - 0.42, "K_j, V_j 常驻 SRAM",
ha="center", fontsize=9.5, color=C_SRAM, weight="bold")
ax2.text(gx0 + T * cell + 0.45, gtop - 0.35, "一块 = [B_r, B_c]\n= [64, 64]",
fontsize=9.5, color="#475569", va="top", linespacing=1.5)
ax2.text(gx0 + T * cell + 0.45, gtop - 1.55,
"灰色块:因果掩码下\n整块跳过,连算都不算",
fontsize=9, color="#94a3b8", va="top", linespacing=1.5)
ax2.text(gx0 + T * cell + 0.45, gtop - 2.95,
"片上空间受限:\n驻留 K/V 与 Q/O 块,\n还要容纳分数和统计量",
fontsize=9, color=C_SRAM, va="top", linespacing=1.5)
by = 1.85
ax2.text(0.15, by + 0.52, "跨 K/V 块一直复用的三个量:", fontsize=10.5,
color=INK, weight="bold")
for k, (label, color) in enumerate([
("O_i 输出累加器 [N, d]", C_FLASH),
("l_i 分母 [N]", "#7fb3d5"),
("m_i 运行最大值 [N]", "#c9a227")]):
ax2.add_patch(Rectangle((0.15, by - k * 0.46), 4.6, 0.28,
facecolor=color, edgecolor="none"))
ax2.text(4.9, by - k * 0.46 + 0.14, label, fontsize=9.5,
color="#475569", va="center")
ax2.text(0.15, 0.12, "每来一块就 rescale 一次(乘 e^{m旧−m新}),最后 O/l 收尾",
fontsize=9.5, color="#64748b")
_caption(fig, "图 1:切的是「K/V 块 × Q 块」这两层循环,不是把注意力切开。"
"左边突出两份二次中间量,右边只维护线性的累加量(固定 d)。")
fig.savefig(OUT / "tiling.png", facecolor="#fbfcfe")
plt.close(fig)
print(" figures/tiling.png")
# ────────────────────────────────────────────────────────────
# 图 2:两本账随 N 怎么长
# ────────────────────────────────────────────────────────────
def fig_curves():
d, b = 64, 2
B_r, B_c = block_sizes(d)
Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768])
naive_io = np.array([hbm_elems_naive(N, d) * b for N in Ns]) / MIB
flash_io = np.array([hbm_elems_flash(N, d, B_r, B_c) * b for N in Ns]) / MIB
fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.1), facecolor="#fbfcfe")
fig.subplots_adjust(bottom=0.20, top=0.86, wspace=0.25)
ax = axes[0]
ax.loglog(Ns, naive_io, "o-", color=C_NAIVE, lw=2.2, ms=6,
label="朴素实现(∝N²)")
ax.loglog(Ns, flash_io, "s-", color=C_FLASH, lw=2.2, ms=6,
label="FlashAttention(∝N²/M)")
ax.annotate(f"{naive_io[2] / flash_io[2]:.1f}×",
xy=(Ns[2], naive_io[2]), xytext=(Ns[2] * 1.6, naive_io[2] * 3.0),
fontsize=11, color=C_NAIVE, weight="bold",
arrowprops=dict(arrowstyle="->", color=C_NAIVE, lw=1.3))
ax.legend(fontsize=10, frameon=False, loc="upper left")
ax.grid(True, which="both", color="#eef2f7", lw=0.9)
_style(ax, "HBM 读写量(单头 d=64, fp16)", "序列长度 N", "MiB")
ax = axes[1]
d2, H, B = 3072, 24, 1
lay = 32
naive_act = np.array(
[activation_bytes(N, d2, H, B)[0] * lay for N in Ns]) / GIB
flash_act = np.array(
[activation_bytes(N, d2, H, B)[1] * lay for N in Ns]) / GIB
ax.loglog(Ns, naive_act, "o-", color=C_NAIVE, lw=2.2, ms=6,
label="朴素实现(∝N²)")
ax.loglog(Ns, flash_act, "s-", color=C_FLASH, lw=2.2, ms=6,
label="FlashAttention(∝N)")
ax.axhline(80, color="#64748b", ls="--", lw=1.4)
ax.text(Ns[0] * 1.15, 92, "示例预算:80 GiB(非硬件规格)", fontsize=9.5, color="#64748b")
ax.annotate("76.5 GiB", xy=(4096, naive_act[2]),
xytext=(4096 * 1.8, naive_act[2] * 2.2), fontsize=10.5,
color=C_NAIVE, weight="bold",
arrowprops=dict(arrowstyle="->", color=C_NAIVE, lw=1.3))
ax.annotate("4.5 GiB", xy=(4096, flash_act[2]),
xytext=(4096 * 0.55, flash_act[2] * 6.0), fontsize=10.5,
color=C_FLASH, weight="bold",
arrowprops=dict(arrowstyle="->", color=C_FLASH, lw=1.3))
ax.legend(fontsize=10, frameon=False, loc="upper left")
ax.grid(True, which="both", color="#eef2f7", lw=0.9)
_style(ax, "教学激活存储模型(简化视频 DiT,32 层)", "序列长度 N", "GiB")
_caption(fig, "图 2:两张都是对数轴,斜率就是复杂度阶数。"
"左图的大 N 主导项均为 N²,有限 N 时比例受线性项和取整影响;"
"右图斜率不同(N² 对 N),N 越大差距越离谱。")
fig.savefig(OUT / "io_curve.png", facecolor="#fbfcfe")
plt.close(fig)
print(" figures/io_curve.png")
# ────────────────────────────────────────────────────────────
# 图 3:算术强度 vs 机器平衡点
# ────────────────────────────────────────────────────────────
def fig_intensity():
d, b = 64, 2
B_r, B_c = block_sizes(d)
Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768])
naive_i = np.array([flops_attention(N, d) / (hbm_elems_naive(N, d) * b)
for N in Ns])
flash_i = np.array([flops_attention(N, d)
/ (hbm_elems_flash(N, d, B_r, B_c) * b) for N in Ns])
balance = FLOPS_PEAK / HBM_BW
fig, ax = plt.subplots(figsize=(11.4, 5.0), facecolor="#fbfcfe")
fig.subplots_adjust(bottom=0.22, top=0.86)
ax.axhspan(0, balance, color="#fdf1f2", zorder=0)
ax.axhline(balance, color="#64748b", ls="--", lw=1.6, zorder=3)
ax.text(Ns[-1] * 2.1, balance * 0.70,
f"机器平衡点 {balance:.0f} FLOP/byte\n(峰值算力 ÷ HBM 带宽)",
fontsize=10, color="#475569", va="center", linespacing=1.5)
ax.text(Ns[0] * 1.02, balance * 0.40,
"线下:本 roofline 模型由 HBM 项主导(实际瓶颈需测量)",
fontsize=10, color=C_NAIVE)
ax.text(Ns[0] * 1.02, balance * 1.32,
"线上:本模型计算项主导", fontsize=10, color=C_FLASH)
ax.semilogx(Ns, naive_i, "o-", color=C_NAIVE, lw=2.4, ms=7,
label="朴素实现", zorder=4)
ax.semilogx(Ns, flash_i, "s-", color=C_FLASH, lw=2.4, ms=7,
label="FlashAttention", zorder=4)
for x, v in zip(Ns, naive_i):
ax.annotate(f"{v:.0f}", (x, v), textcoords="offset points",
xytext=(0, -16), ha="center", fontsize=9, color=C_NAIVE)
for x, v in zip(Ns, flash_i):
ax.annotate(f"{v:.0f}", (x, v), textcoords="offset points",
xytext=(0, 9), ha="center", fontsize=9, color=C_FLASH)
ax.set_xticks(Ns)
ax.set_xticklabels([f"{n // 1024}K" for n in Ns])
ax.set_xlim(Ns[0] * 0.75, Ns[-1] * 4.5)
ax.set_ylim(0, balance * 1.85)
ax.legend(fontsize=10.5, frameon=False, loc="center right")
ax.grid(axis="y", color="#eef2f7", lw=0.9)
_style(ax, "算术强度:每搬 1 字节能做多少次运算(单头 d=64, fp16)",
"序列长度 N", "FLOP / byte")
_caption(fig, "图 3:朴素实现的强度几乎不随 N 变(一直贴在 32 附近),"
"在此模型中由 HBM 项主导;FlashAttention 的估计抬过平衡点,"
"并不单独证明实际 kernel 已充分利用计算单元。")
fig.savefig(OUT / "intensity.png", facecolor="#fbfcfe")
plt.close(fig)
print(" figures/intensity.png")
if __name__ == "__main__":
print("生成配图:")
fig_tiling()
fig_curves()
fig_intensity()
评论 (0)