AIGC 基本功|算子融合与 CUDA Graph-Fusion

人工智能炼丹君
2026-10-02 / 0 评论 / 1 阅读 / 正在检测是否收录...

算子融合与 CUDA Graph

所属方向:推理加速 | 难度:进阶 | 前置知识:性能建模与 Profiling(performance_profiling)、混合精度(mixed_precision)、自注意力机制(attention_basics)
关键词:算子融合、kernel fusion、CUDA Graph、发射开销、访存瓶颈、torch.compile、Inductor

关于本文的数字:作者手里这台机器是 Apple M1 Pro,没有 NVIDIA GPU,也没有安装 torch。所以凡是标「实测」的数字,都出自文末附录里那三个能在纯 numpy 上跑通的脚本;凡是 CUDA/HBM 相关的数字,都来自公开资料并显式标注出处,我没有拿 CPU 的数字去外推 GPU。第六节 6.5 专门交代了这条边界。


01. 为什么需要它

先摆三个数,都是本篇附录里真跑出来的。

第一个数:一步 decode 大约要发射 1093 个 kernel。 这不是实测,是按 Llama-3-8B 的公开结构一层层手数出来的:每层 34 个 kernel(输入 RMSNorm 拆 4 个、q/k/v 投影各 1 个、两组 RoPE 各 3 个、KV 写入 2 个、QK 转置乘 1 个、缩放加掩码 2 个、softmax 3 个、乘 V 1 个、o 投影 1 个、残差 1 个、后置 RMSNorm 4 个、gate/up 各 1 个、SiLU 1 个、逐元素乘 1 个、down 投影 1 个、残差 1 个),乘 32 层再加输出头约 5 个,得到 1093。公开资料给的量级是 300~1300,这个数落在里面,可以互为印证。

第二个数:这 1093 个 kernel 里,绝大多数只干几微秒的活。 每次 kernel 发射在 CPU 侧要花 1~5 微秒——NVIDIA 开发者论坛上 njuffa 给的口径是「空 kernel 约 5 微秒」,《CUDA Handbook》实测 NULL launch 约 4.9 微秒(老机器)、GeForce RTX 3060 上约 1.2 微秒,A100 上常见引用是 2~3 微秒。假设发射 3 微秒、执行 2 微秒,附录 launch_model.py 的时间线推演给出的结果是:eager 模式总时长 2102 微秒,其中 GPU 空转 702 微秒,也就是 33% 的时间卡在等 CPU 把下一个 kernel 发过来;换成一次提交的 CUDA Graph,总时长降到 1753 微秒,加速 1.20 倍,GPU 基本不空转了。

第三个数:8 个串起来的逐元素算子,其中 93.8% 的时间花在「发射」上而不是「算」。 这是 fusion_lab.py 在 512 个元素的张量上实测的。同样一条 8 段链,在 400 万元素的张量上反过来:加速的 100% 来自少搬字节,发射部分只占 0%。同一个优化在两个极端上赚的是完全不同的钱,这件事后面会用一整套账把它算清楚。

还有一个更刺眼的数在注意力上。按 $B=8$、$H=16$、$S=2048$、$d=4096$、bf16 算:输入激活 $[B,S,d]$ 是 128.0 MiB,而 score 矩阵 $[B,H,S,S]$ 是 1.000 GiB,正好是输入的 8 倍。放大倍数是 $H \cdot S / d$——$S$ 翻一倍,它翻一倍;$S=8192$ 时是 32 倍。softmax 如果不融合,max / 减 / exp / 求和 / 除五步各读写一遍 score,就是 8.000 GiB 的访存;融合后只要 2.000 GiB。

你以为模型卡在算力上,其实它经常卡在「把中间结果写出去再读回来」和「让 CPU 一次又一次地通知 GPU 开工」这两件事上。

02. 最小可用理解

三句话:

  1. 融合(fusion):一串相邻算子的中间结果,本来每个算子都要写回显存、下一个再读回来。融合就是把它们放进同一个 kernel,中间结果留在寄存器或共享内存里,只在最开头读一次输入、最末尾写一次输出。它不产生新的数学,只是让已经算出来的东西少走几趟路。
  2. CUDA Graph:那串 kernel 的发射动作(函数名、参数、依赖关系、显存地址)可以录下来,之后一次 cudaGraphLaunch 提交整张图。它一个字节都不省,省的是 CPU 侧那 1093 次提交。
  3. 两者的收益都能提前算出来,也都有限。融合的上限是字节比 $K$($K$ 段链最多省 $K$ 倍访存);CUDA Graph 的上限是「发射时间 / 执行时间」,而且只在单 kernel 执行时间接近或小于发射时间时才为正。

一句总结:融合省字节,Graph 省发射。把两者混为一谈,是这一块最常见的错误。

03. 数学推导

3.1 先回到那条判据:roofline

前置文章 performance_profiling 已经建过这条判据,这里只取结论。一个算子的耗时下界由两件事里更慢的那个决定:

$$T \ge \max\left(\frac{F}{P_{\text{peak}}},\ \frac{B}{B_{\text{off}}}\right)$$

其中 $F$ 是浮点运算次数(FLOPs),$P_{\text{peak}}$ 是峰值算力(FLOP/s),$B$ 是需要搬动的字节数(含读和写),$B_{\text{off}}$ 是片外内存带宽(字节/秒)。两者的比值就是算术强度:

$$I = \frac{F}{B}$$

$I$ 高的算子卡在算力上,$I$ 低的算子卡在带宽上。逐元素算子的 $I$ 是常数:算一个元素做 1 次操作、读 4 字节写 4 字节,$I = 1/8$。在 fp32 下算力峰值除以带宽的量级是每字节几十次操作,所以逐元素算子永远在带宽那一侧。

这就是融合的着力点:它不改变 $F$,只把 $B$ 变小。

3.2 字节账本:K 段链为什么最多省 K 倍

考虑 $K$ 个串起来的逐元素算子,作用在一个 $N$ 元素的张量上($b$ 为每元素字节数)。

不融合:每个算子单独成一个 kernel,各自读一遍输入、写一遍输出。第 $k$ 个 kernel 的访存是 $2Nb$ 字节(因为输入和输出都是 $N$ 个元素),$K$ 个加起来:

$$B_{\text{unfused}} = 2KNb$$

融合:一个 kernel 从头做到尾。读输入 $Nb$、写输出 $Nb$,中间那 $K-1$ 步全在片上:

$$B_{\text{fused}} = 2Nb$$

两式相除,字节比正好是 $K$:

$$\frac{B_{\text{unfused}}}{B_{\text{fused}}} = K$$

本文附录 fusion_lab.py 的 [A] 节把这个账算在了三处:8 段链 244.1 MiB 降到 30.5 MiB(8 倍);RMSNorm 从 4 遍降到 1 遍,122.1 MiB 降到 30.5 MiB(4 倍);attention softmax 从 8.000 GiB 降到 2.000 GiB(4 倍),另外还有一份被物化的 mask 值 1.000 GiB,融合后直接消失。

注意字节比 $K$ 只是「访存减少到 1/K」,不是「时间减少到 1/K」。 时间还取决于中间结果到底停在多快的存储上。设片上带宽为 $B_{\text{on}}$,定义片内片外带宽比:

$$\rho = \frac{B_{\text{on}}}{B_{\text{off}}}$$

那么融合后的时间是「片外读一次写一次」加上「$K-1$ 步在片上走」:

$$T_{\text{fused}} = \frac{2Nb}{B_{\text{off}}} + \frac{2(K-1)Nb}{B_{\text{on}}}$$

不融合的时间是:

$$T_{\text{unfused}} = \frac{2KNb}{B_{\text{off}}}$$

两者相除,$N$ 和 $b$ 全部约掉:

$$S = \frac{T_{\text{unfused}}}{T_{\text{fused}}} = \frac{K}{1 + (K-1)/\rho}$$

这个式子值得盯着看三秒。 它的两个极限都很干净:

  • $\rho \to 1$(片上不比片外快):$S \to 1$,融合白干。
  • $\rho \to \infty$(片上无限快):$S \to K$,也就是字节比。

所以「融合能加速几倍」这个问题,答案既不是 $K$,也不是玄学,而是 $K$ 和 $\rho$ 共同决定的。记住这一点,第六节会看到它把一台机器上的实测结果解释得干干净净。

3.3 单次调用的固定成本,和它的临界规模

一次算子调用的耗时,在规模足够小时和数据量无关。写成仿射形式:

$$t(n) = a + b_{\text{el}} \cdot n$$

$a$ 是固定开销(派发、参数检查、缓冲建立),$b_{\text{el}}$ 是每元素的边际成本。开销和数据各占一半的那个规模是:

$$n^{*} = \frac{a}{b_{\text{el}}}$$

fusion_lab.py 的 [D] 节实测这台机器:$a = 0.430$ 微秒/次,$b_{\text{el}} = 55.47$ 微秒/百万元素,$n^{*} = 7746$ 个元素,也就是 30.3 KiB。张量小于 30 KiB 时,大部分时间不是在算数据。

如果把这条链的 $2K$ 次调用全加起来,固定开销的占比是:

$$\text{share} = \frac{2K a}{2K(a + b_{\text{el}} n)} = \frac{1}{1 + n/n^{*}}$$

实测($K=8$,16 次调用):$n=512$ 时 93.8%,$n=4096$ 时 65.4%,$n=65536$ 时 10.6%,$n=1048576$ 时 0.7%。张量一小,你花在「组织计算」上的钱就超过「做计算」的钱。 这就是 GPU 上 kernel launch 开销在 CPU 上的同构物,也是为什么 GPU 上会出现同样性质的墙。

3.4 发射与执行的时间线:CUDA Graph 到底省什么

考虑单流、异步提交的 $N$ 个 kernel。CPU 发射第 $i$ 个要花 $t_{\text{launch}}$,发完就可以去发下一个,不等 GPU。GPU 执行第 $i$ 个要花 $t_{\text{exec}}$,它有两个前提:前一个跑完了,并且这一个已经发出去了。所以 GPU 开工时刻是:

$$\text{start}_i = \max\left(\text{end}_{i-1},\ (i+1) \cdot t_{\text{launch}}\right)$$

两个条件谁慢,谁决定节奏。分两种情形:

  • $t_{\text{launch}} \le t_{\text{exec}}$:CPU 能跑在 GPU 前面,GPU 一个接一个不停,总时长约 $N \cdot t_{\text{exec}} + t_{\text{launch}}$。
  • $t_{\text{launch}} > t_{\text{exec}}$:CPU 喂不上,GPU 每跑完一个就得等 $t_{\text{launch}} - t_{\text{exec}}$,总时长约 $N \cdot t_{\text{launch}}$。GPU 空转的比例是 $1 - t_{\text{exec}}/t_{\text{launch}}$。

CUDA Graph 把 $N$ 次提交压成 1 次,但 GPU 前端仍然要逐个节点过一遍(记 $t_{\text{disp}}$):

$$T_{\text{graph}} = t_{\text{launch}} + N \cdot (t_{\text{exec}} + t_{\text{disp}})$$

于是加速比是:

$$S_{\text{graph}} = \frac{\max(N t_{\text{launch}},\ N t_{\text{exec}})}{t_{\text{launch}} + N(t_{\text{exec}} + t_{\text{disp}})} \approx \frac{\max(t_{\text{launch}}, t_{\text{exec}})}{t_{\text{exec}} + t_{\text{disp}}}$$

这个式子给出了盈亏平衡点:只有当 $t_{\text{exec}} \lesssim t_{\text{launch}}$ 时 $S_{\text{graph}} > 1$。 launch_model.py 的 [F2] 节用 $N=700$、$t_{\text{launch}}=3$ 微秒、$t_{\text{disp}}=0.5$ 微秒扫了一遍:$t_{\text{exec}}=0.5$ 微秒时 2.99 倍,1 微秒时 2.00 倍,2 微秒时 1.20 倍,到 3 微秒正好跌破 1(0.86 倍),40 微秒时 0.99 倍——大 kernel 上 CUDA Graph 是纯负担。

3.5 因果掩码:白算的那一半

因果注意力里,第 $i$ 个 query 只能看见前 $i$ 个 key。如果老老实实算满 $S \times S$ 的 score 矩阵,被掩掉的部分是纯浪费。精确的保留比例是:

$$\text{kept} = \frac{S(S+1)/2}{S^2} = \frac{S+1}{2S}$$

$S=2048$ 时是 50.0%——一半的 score 元素算完就扔。逐元素地跳过在硬件上不现实,实际做法是按块跳:块边长 $B_r$ 时共有 $n_b = S/B_r$ 个 query 块,第 $i$ 个 query 块只算前 $i$ 个 key 块,算下来是 $n_b(n_b+1)/2$ 个块:

$$\text{kept}_{\text{block}} = \frac{n_b(n_b+1)/2}{n_b^2} = \frac{n_b+1}{2 n_b} = \frac{1}{2} + \frac{1}{2 n_b}$$

实测($S=2048$):块边长 64 时 51.6%,128 时 53.1%,256 时 56.2%。误差随 $n_b$ 按 $1/(2n_b)$ 衰减,所以块越小越接近理论上界——但块越小,片上数据被切得越碎,调度开销越大。这是块大小必须折中的原因,不是调参玄学。

attention 的两本账:中间产物多大、白算多少

这张图要看什么:左图是 score 矩阵相对输入激活的放大倍数随 $S$ 的变化,灰虚线是「若按 $S$ 线性增长」的参考——实际曲线比线性还陡($H \cdot S/d$ 是 $S$ 的一次式,但在对数轴上叠加了 $H/d$ 的常数放大,$S=8192$ 时已到 32 倍)。右图是块级跳过的实际工作量占比,绿虚线是理论上界 50.0%:块边长 256 时要多算 6.2 个百分点,块边长 64 时只多算 1.6 个——这就是「块越小越准、代价越碎」的量化版本。

04. 代码实现

三个脚本,全部纯 numpy,/usr/local/bin/python3 直接跑:

  • fusion_lab.py:[A] 字节账本、[B] 带宽与工作集、[C] 三种融合形态、[D] 固定开销、[E] attention 尾部
  • launch_model.py:[F1] 时间线、[F2] 扫描、[F3] kernel 计数、[F4] 收益分解
  • make_figures.py:把上面两个脚本落盘的 json 画成 5 张图

4.1 三种融合形态

我把「融合」拆成三种可测量的形态,这是本篇的核心实验设计:

K = 8                      # 8 段链,每段 h <- A*h + B
A_S, B_S = 1.01, 0.01


def _v1_unfused(x, buf, K=K):
    """V1 不融合:K 段,每段 2 次算子调用,每段结果都写回内存。"""
    np.multiply(x, A_S, out=buf)
    np.add(buf, B_S, out=buf)
    for _ in range(K - 1):
        np.multiply(buf, A_S, out=buf)
        np.add(buf, B_S, out=buf)
    return buf


def _v3_closed(x, out, K=K):
    """V3 代数合并:K 段仿射合成一个仿射,只剩 2 次调用。"""
    Ak = A_S ** K
    Bk = B_S * (Ak - 1.0) / (A_S - 1.0)
    np.multiply(x, Ak, out=out)
    np.add(out, Bk, out=out)
    return out

这三种形态的物理含义不同,必须分清:

  • V1 不融合:每个算子一个 kernel,中间结果每次都落回主存。
  • V2 分块融合(_v2_tiled):一块读进来,K 段都在这一块上算完,只写回一次。tile 就是片上缓冲。这是 GPU 融合 kernel 在 CPU 上能做到的最好近似——numpy 做不到寄存器级融合,每个 ufunc 调用仍然要把结果写回 tile。
  • V3 代数合并:把 $K$ 段仿射在数学上合成一段,中间结果根本不存在。它对应的是融合给编译器创造的二阶机会:串起来看得见全貌之后,可以化简。

实测结果(fusion_lab.py ALL 的 [C] 节):

张量元素数 V1 不融合 V3 代数合并 V1/V3 相对误差
512 7.75 μs 1.12 μs 6.89x 1.35e-07
4096 10.58 μs 1.46 μs 7.26x 1.97e-07
65536 84.63 μs 11.25 μs 7.52x 5.12e-07
262144 310.83 μs 41.62 μs 7.47x 5.80e-07
1048576 1240.04 μs 176.54 μs 7.02x 5.13e-07
4000000 5205.08 μs 792.83 μs 6.57x 6.16e-07

V1/V3 稳定在 6.57~7.52 倍,围着 $K=8$ 上下浮动——正如 3.2 节的推导,字节比就是上限。最后一列的相对误差约 5e-7,正好是 fp32 的舍入量级:代数合并改变了运算顺序,所以数值不会逐位相同,但量级完全在许可范围内。

同一节里 V2 的结果才是本篇最值得记住的一个数:

tile(元素) 耗时 相对 V1
16384 5.835 ms 0.89x
65536 5.626 ms 0.93x
262144 5.290 ms 0.98x
1048576 5.361 ms 0.97x

V2 一点都没变快,甚至略慢。 而 3.2 节的模型早就预测到了——代入 [B] 节测出的 $\rho = 1.124$:

$$S = \frac{8}{1 + 7/1.124} = 1.11$$

模型说 1.11 倍,实测 0.98 倍。同一量级,方向一致。为什么这台机器上 $\rho$ 这么小?因为单线程 numpy 逐元素循环的吞吐上限只有 91.5 GB/s,而主存带宽是 81.4 GB/s——两个数字离得太近,片上片下几乎没有差价可赚。

三种融合形态:耗时与加速比

这张图要看什么:左图是两种形态的耗时随规模的变化,两条虚线是各自的发射地板——$V1$ 是 $2K \cdot a = 6.9$ μs,$V3$ 是 $2a = 0.86$ μs。两条实线在小规模处几乎是水平的(贴着各自的地板走),到大规模才分开,这就是「小规模赚发射、大规模赚字节」的直接证据。灰色菱形是 V2 在最大规模上的结果,它几乎落在 V1 那条线上——$V2$ 省了字节但没省调用,所以拿不到 V3 的收益。右图把这件事翻译成加速比:橙线(V1/V3)贴着绿色上限 $K=8$ 走,而 V2 的实测点落在灰虚线(模型预测 1.11x)附近,离 8 差着一个数量级。

4.2 把带宽层级测出来

$\rho$ 不是查来的,是测来的:

for nbytes in [32 << 10, 128 << 10, 512 << 10, 2 << 20, 8 << 20,
               32 << 20, 128 << 20]:
    n = nbytes // 4
    x = rng.standard_normal(n).astype(np.float32)
    y = np.empty_like(x)
    bench(lambda: np.add(x, 1.0, out=y), reps=3)       # 预热
    t = bench(lambda: np.add(x, 1.0, out=y), reps=11)
    bw = 2 * nbytes / t / 1e9                          # 1 读 1 写

实测:32 KiB 62.9 GB/s、128 KiB 79.6、512 KiB 91.5、2 MiB 95.3、8 MiB 94.7、32 MiB 79.5、128 MiB 83.3。取 ≤8 MiB 的中位数 91.5 GB/s 作 $B_{\text{on}}$,≥32 MiB 的中位数 81.4 GB/s 作 $B_{\text{off}}$,得 $\rho = 1.124$。

注意这两个数字不是硬件规格里的峰值带宽,而是「单线程 numpy 逐元素循环能跑出来的吞吐」。它们的差距不代表 L2 和主存的差距,只代表在这条代码路径上「留片上」值多少钱。这个诚实的限定很重要——换一台有真 GPU 的机器,$\rho$ 会是另一个数,$K/(1+(K-1)/\rho)$ 会给出完全不同的答案。

4.3 固定开销

reps = 20000 if n <= 4096 else 500
t0 = time.perf_counter()
for _ in range(reps):
    np.add(a, 1.0, out=b)
t = (time.perf_counter() - t0) / reps

实测:$n=1$ 时 0.417 μs,$n=8$ 时 0.419,$n=64$ 时 0.425,$n=512$ 时 0.477,$n=4096$ 时 0.675,$n=16384$ 时 1.333,$n=65536$ 时 5.966,$n=262144$ 时 22.388,$n=1048576$ 时 89.331。

前四个点几乎是一条水平线——从 1 个元素到 512 个元素,元素数涨了 512 倍,耗时只从 0.417 涨到 0.477。用 $n \le 16384$ 的六个点做最小二乘,得 $a = 0.430$ μs、$b_{\text{el}}$ 对应 55.47 μs/百万元素。

4.4 图:先看两张

融合前后各搬多少字节

这张图要看什么:三组 pattern 的字节账本,横轴是对数刻度。橙红是不融合、绿是融合后,条右边的倍数就是字节比。注意最上面那根——attention softmax 的 8192.0 MiB 比最下面那根逐元素链的 244.1 MiB 大了三十多倍,真正需要融合的不是那串小算子,而是注意力里那张被反复读写的 score 矩阵。

每次调用都有一个和数据规模无关的地板

这张图要看什么:左图是单次调用耗时随张量规模的变化,红色虚线是固定开销地板,紫色竖线是「开销和数据各占一半」的临界规模(7746 个元素);右图是 8 段链里固定开销占总时间的比例。曲线在最左边几乎是水平的——那一段里,你增加 500 倍的数据量,耗时只涨 14%。

4.5 分解:两种收益各占多少

launch_model.py 的 [F4] 节把本机实测的加速拆成两部分:

张量元素数 V1 实测 V1 模型 V3 实测 V3 模型 发射贡献
512 7.75 μs 7.33 μs 1.12 μs 0.92 μs 94%
4096 10.58 μs 10.51 μs 1.46 μs 1.31 μs 65%
65536 84.63 μs 65.04 μs 11.25 μs 8.13 μs 11%
262144 310.83 μs 239.53 μs 41.62 μs 29.94 μs 3%
1048576 1240.04 μs 937.51 μs 176.54 μs 117.19 μs 1%
4000000 5205.08 μs 3556.95 μs 792.83 μs 444.62 μs 0%

$n=512$(2.0 KiB)时,加速的 94% 来自少发射,只有 6% 来自少搬字节;$n=4000000$(15.3 MiB)时反过来,100% 来自少搬字节。

模型在大规模那一端明显低估(5205 实测 vs 3557 模型)——因为 $b_{\text{el}}$ 是用缓存内的点拟合的,外推到 15 MiB 后,边际成本实际上比拟合值高。这是线性模型的能力边界,写在这里免得读者拿它当精确预测器。

CUDA Graph 省的是发射,不是字节

这张图要看什么:左图是时间线甘特图,上排是 eager。蓝条是 CPU 发射,绿条是 GPU 执行,斜纹是 GPU 空转——空转的宽度正好等于发射和执行的时间差。下排是 CUDA Graph,一次发射之后 GPU 一路不停。右图是加速比随单 kernel 执行时间的变化,红线是发射成本 3 μs:红线左边 Graph 赚,红线右边 Graph 亏。

05. 工业级实现对照

5.1 torch.compile / Inductor:融合是调度器的决策

PyTorch 的主入口在 torch/_inductor/compile_fx.py 的 compile_fx(当前实现在该文件第 3122 行),它把 FX 图接给 Inductor,后者负责切分、调度、生成 Triton kernel:
https://github.com/pytorch/pytorch/blob/main/torch/_inductor/compile_fx.py

真正做融合决策的是调度器。torch/_inductor/scheduler.py 里:

  • Scheduler.fuse_nodes(nodes)(约 7048 行)是融合主循环;
  • can_fuse_vertical / can_fuse_horizontal / can_fuse_reduction_epilogue 决定两个节点能不能合;
  • score_fusion_memory(node1, node2, count_bytes=...)(约 11278 行)给一次融合打分——它的第一个形参就叫 count_bytes,因为融合的分数就是省下的字节数。
  • 融合决定返回 FusionResult(约 130 行),里面既可以是布尔值,也可以是一个待求值的 callable_fn,让代价模型并行算。

这就是本篇 3.2 节的账本在生产代码里的样子。 区别在于:编译器要在不能融合的时候正确地放弃。不能融合的情形包括跨块归纳(reduction 的中间结果必须先写完)、原地写与别名(写坏了输入后续还要用)、随机数算子(融合会改变随机流)、以及动态 shape 下无法预先确定 tile 大小。

5.2 一个必须在 05 节点名的差异:epilogue fusion

最小实现里的融合是「一串逐元素算子合成一个 kernel」。生产实现里最有价值的融合形态是 epilogue fusion:把接在 GEMM 后面的 bias、激活、缩放并进 GEMM 的 kernel 里,让 GEMM 的输出直接以最终形态写出,不经过显存。

用 3.2 节的账本算一下就明白为什么值钱:一个 $[M,N]$ 的 fp32 输出,不融合时要写 GEMM 输出($4MN$ 字节写)、bias 加法读+写($8MN$)、激活读+写($8MN$),合计 $20MN$;融合后 GEMM 一趟写 $4MN$,加上读输入的开销,量级降到 $1/5$。表达式没变一行,访存少了五分之四。

5.3 CUDA Graph 在 PyTorch 里的落地

torch/cuda/graphs.py 里的 CUDAGraph(约 289 行)暴露出和 CUDA 一一对应的四个动作:capture_begin / capture_end / instantiate / replay:
https://github.com/pytorch/pytorch/blob/main/torch/cuda/graphs.py

三个工程细节值得单独说:

  1. graph_pool_handle()(约 96 行)必须存在。graph 在录制时把显存地址写死在节点里,回放时不会重新分配。所以你必须让图里所有张量都来自一个固定的内存池。这不是 API 的怪癖,是「录制-回放」这个机制的必然代价。
  2. 上游警告 instantiate 要在第一次 replay 之前显式调用,否则第一次回放的延迟会变高(图要在那时才编译)。这个问题在实时推理里就是一次 P99 抖动。
  3. torch/_inductor/cudagraph_trees.py 里是 CUDAGraphNode(约 982 行)和 TreeManagerContainer(约 220 行),还有一个 CUDAWarmupNode。为什么是「树」而不是「一张图」:训练时反向图的形状依赖前向的实际形状,一批数据一个形状;而且每次回放都会产生新的输出张量,需要知道哪些显存可以复用。Inductor 用 mode="reduce-overhead" 打开这套机制。
    https://github.com/pytorch/pytorch/blob/main/torch/_inductor/cudagraph_trees.py

5.4 差异从哪来:为什么生产实现不能只有最小实现

差异 来源
要判断「能不能融合」,甚至要算清代价 融合不是永远划算,见第六节 6.1 与 6.3
要处理动态 shape tile 大小不能预知,必须切分或退化成不融合
要维护固定的显存池 CUDA Graph 把地址写死了
要为每个形状/profile 各录一份图 图是静态的;变长输入需要分段或重录
要处理随机流、原地写、别名 融合会改变这些语义

06. 代价与边界

6.1 收益的上限是字节比,不是魔法

3.2 节的 $S = K/(1+(K-1)/\rho)$ 是硬约束。在写本文这台机器上 $\rho = 1.124$,理论上限只有 1.11 倍,实测 V2 是 0.98 倍——收益在哪台机器上有、有多大,完全由 $\rho$ 决定。所以当有人告诉你「融合能快 6 倍」,第一个该问的问题不是「怎么融合」,而是「他的 $\rho$ 是多少、$K$ 是多少」。

顺带说,V3(代数合并)之所以能跑到 6.57~7.52 倍而不受 $\rho$ 限制,是因为它根本没产生中间结果——它走的是 $S \to K$ 那条极限路径。但代数合并只在可化简的模式上成立(本链是仿射复合,可以闭式求解),一般算子是合不掉的。它展示的是融合的第二层价值:让编译器看见全貌,从而有机会化简。

6.2 融合会改变数值

实测相对误差约 5e-7(fp32 舍入量级)。这不是 bug,是代数重排的必然结果。任何声称「融合不改变数值」的说法都需要限定条件。在需要严格复现的训练里,这会让「同一份代码换个后端跑出不同 loss」——通常无害,但你必须知道它从哪来。

6.3 融合省不了算错的东西

这是本篇最反直觉的一条,也是附录 [E] 节专门测的。softmax 五步在 $[8,1024,1024]$ 的 score 矩阵上:不融合 18.804 ms,按行分块融合 23.870 ms,加速比 0.79x——融合反而慢了 21%。

算一下就知道为什么。五步共读写 8 份 score 矩阵 = 256 MB,按 81.4 GB/s 的访存下界是 3.298 ms。实测 18.804 ms,是访存下界的 5.70 倍。也就是说这个 kernel 的时间几乎全在 exp 的计算上,不在访存上。 分块省了字节,但一个 exp 都没少,反而因为切碎了向量化还赔了一点。

所以:融合只对「本来就卡在访存或发射上」的算子有效。 一个计算受限的算子,融合不动它。这也是 FlashAttention 要在融合之外再做「不实例化 S×S 矩阵」的原因——它同时省了访存和一部分计算。

6.4 CUDA Graph 的硬约束

  • 静态 shape:图录下来的形状是死的。变长输入要么分段(prefill 一段、decode 一段各录一份),要么放弃。
  • 静态地址:所有张量必须来自固定内存池,中间不能有新的 cudaMalloc。
  • 不能有 host 同步和 D2H 拷贝:录制期间不允许任何把控制权交回 CPU 的操作。
  • 首次录制和实例化有成本,且出错时栈是「图里的第 N 个节点」,比直接调试难得多。
  • 大 kernel 上纯亏:3.4 节的模型里,$t_{\text{exec}} = 20$ μs 时加速比 0.98 倍。你付出了录制、显存和调试的成本,换来 2% 的倒退。

一句话:CUDA Graph 是给「小 kernel 洪水」用的药,不是通用加速开关。 一个理想的判断顺序是:先用 profiler 确认 GPU 有空转;再确认单个 kernel 确实小于发射成本;然后才考虑上 Graph。而如果那些 kernel 本来就该被融合掉,融合是更根本的解法——融合让 kernel 数从 1093 降到 355,Graph 只是把这 1093 次提交合并成 1 次;两者正交,可以叠加。

6.5 证据的范围

  • 本机实测:全部来自附录三个脚本在这台 Apple M1 Pro 上的真实运行输出([A]~[F4])。这台机器没有 NVIDIA GPU、没有 torch,所以没有任何一个 CUDA 数字是实测的。
  • 公开资料(非实测):kernel launch 的 1~5 μs 量级、一步 decode 的 300~1300 个 kernel、HBM 与片上带宽的量级差。出处见 launch_model.py 文件头。
  • 模型推演:launch_model.py 的 [F1]/[F2]/[F3]/[F4] 是在上述输入上做的算术,脚本可跑、输入可改。请把这些当作「给定这些输入会得出什么」,而不是「GPU 上就是这么快」。
  • 不确定的部分:[F3] 的 1093 是按架构手数的,真实值取决于编译器怎么 lowering,本文没有在真实 GPU 上验证过。Inductor 的融合策略也在持续变化,5.1 节引用的行号与函数名以 2026-10 时的 main 分支为准。

07. 经典论文脉络

四篇,按「融合是怎么从手工技巧变成系统能力,又怎么被反过来审视」串起来。

  1. TVM: An Automated End-to-End Optimizing Compiler for Deep Learning(arXiv:1802.04799,2018)。第一次把「算符融合」从工程师的手工活变成编译器的自动决策:给定计算图,搜索融合方案和调度模板。本篇 3.2 节那个「融合省多少字节」的账,在这篇里是被当作搜索目标函数的一部分来算的。

  2. The Deep Learning Compiler: A Comprehensive Survey(arXiv:2002.03794,2020,本文知识树的锚点)。给出一张完整坐标系:图级优化(融合、常量折叠、CSE)、算子级优化、内存分配、后端代码生成。它把融合明确归到「图级优化」——融合对单个算子的数学一无所知,它只改写算子之间的边界。这句话是理解本篇全部内容的前提。

  3. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(arXiv:2205.14135,2022)。把融合推到极致:不只把 softmax 融进注意力,还根本不实例化 $S \times S$ 的 score 矩阵,用在线 softmax 一边遍历一边归一化。这正是本篇 [A] 节那 1.000 GiB 中间产物的解法,而且它额外做到了 6.3 节说的那件事——省的不只是访存,还有被掩码白算的一半。

  4. Operator Fusion in XLA: Analysis and Evaluation(arXiv:2301.13062,2023)。反过来审视这件事:作者去读 XLA 的融合 pass 源码,实测在 Cartpole 上不同融合策略的效果,最好的实现拿到 10.56 倍。这篇的价值在于它把「融合是好事」这个默认假设变成了可度量的对象——和本篇的 V2 实测 0.98 倍是同一个姿势:先问一句「到底快了多少」,再决定要不要相信。

至于 CUDA Graph,它不是论文,是 CUDA 10 起提供的一项运行时机制(录制-实例化-回放),说明见 CUDA C++ Programming Guide 的 CUDA Graphs 一节。把它和融合并列讨论,是因为它经常被当成融合的替代品——而它其实解决另一个问题。

08. 常见误解

  1. 「融合越大越好」。 融得越狠,寄存器压力和共享内存占用越大,occupancy 掉下来,反而更慢。Inductor 的 can_fuse_* 那一组函数之所以是「判断」而不是「尽量合」,就是因为融合有代价。而本文 [E] 节的实测给了一个更直接的极端例子:分块融合的 softmax 比不融合慢 21%。
  2. 「CUDA Graph 是通用加速开关」。 它一个字节都不省,也不减少 kernel 数。3.4 节的公式说得很清楚:$t_{\text{exec}} \ge t_{\text{launch}}$ 时它只会让你付 $t_{\text{disp}}$ 的额外成本。先测 GPU 有没有空转,再决定要不要上。
  3. 「融合不改变数值」。 代数重排会改变舍入顺序。本文 V3 实测相对误差约 5e-7。无害,但要知道它存在。
  4. 「把中间结果留在片上就一定快」。 收益取决于 $\rho$。本文这台机器 $\rho = 1.124$,V2 实测 0.98 倍——留片上的动作做对了,收益是零。反面同样成立:GPU 上 $\rho$ 大得多,同一个动作就值几倍。
  5. 「融合省的是计算量」。 它不改变 FLOPs,只改变访存和发射。softmax 的 exp 一个都没少([E] 节:实测是访存下界的 5.70 倍,说明瓶颈在计算)。

09. 动手验证

都是纯 numpy,不需要 GPU。

实验一:拿到你自己机器的地板。 跑 python fusion_lab.py D,读出 $a$ 和 $n^{*}$。在 M1 Pro 上预期 $a \approx 0.43$ μs、$n^{*} \approx 7746$ 元素(30.3 KiB)。换台机器这两个数会变,但「小张量上耗时和元素数无关」这段平台一定会出现。

实验二:把 K 调大。 把 fusion_lab.py 里的 K_CHAIN 从 8 改成 16,重跑 ALL。预期两件事:V1/V3 的加速比向 16 靠拢(而不是翻倍到 16 以上——字节比就是上限),以及 V1 的发射地板从 $2\times8\times0.43 \approx 6.9$ μs 涨到约 13.8 μs,曲线左端整体抬高。

实验三:看 $\rho$ 怎么被改坏。 把 [B] 节的 np.add(x, 1.0, out=y) 换成跨大步长的切片访问(例如每隔 32 个元素取一个),重测。预期 $B_{\text{on}}$ 明显下降、$\rho$ 趋近 1,随后 V2 的收益会进一步向 1 靠拢。这能让你亲眼看到「$\rho$ 小 ⇒ 融合白干」这条因果链。

实验四:自己数一遍 kernel。 跑 python launch_model.py,在 [F3] 的 per_layer 列表里按你自己熟悉的模型改,看总数落在哪。预期 Llama-3-8B 的 1093 落在公开资料给的 300~1300 区间内。

10. 延伸阅读

  • performance_profiling|性能建模与 Profiling:本篇 3.1 节的 roofline 判据来自那里,$\rho$ 和算术强度这两个量在那里第一次被建立起来。
  • attention_basics|自注意力机制:理解 3.5 节因果掩码和 [A] 节 score 矩阵形状的前提。
  • flash_attention|FlashAttention:本篇 [A] 节那 1.000 GiB 中间产物的正式解法,融合走到极致的样子。
  • mixed_precision|混合精度:数据类型决定 $b$(每元素字节数),直接进 3.2 节的账本。
  • kv_cache|KV Cache 与自回归视频生成:decode 场景为什么 kernel 又多又小,那篇给的 cache 账本是本篇 6.4 节「分段 capture」动机的来源。

附录:完整代码

09 节用到的脚本全文如下(launch_model.py、fusion_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。

launch_model.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
launch_model.py —— 发射开销、CUDA Graph,以及它们和融合的分工

融合省的是字节,CUDA Graph 省的是 CPU 侧的发射次数,一个字节都不省。
这个脚本把两件事分开算:

  [F1] 离散事件时间线:eager 逐个发射 vs CUDA Graph 一次发射
  [F2] 扫描:单个 kernel 执行多久时,CUDA Graph 才赚?
  [F3] 按 Llama-3-8B 架构手数一步 decode 有多少个 kernel
  [F4] 分解:本机实测的 K 段链加速里,多少来自「少发射」,多少来自「少搬字节」

公开量级的出处(都不是本机实测,本机没有 NVIDIA GPU):
  * kernel launch 的 CPU 侧成本:NVIDIA 开发者论坛 njuffa 给的是"空 kernel
    约 5 us";CUDA Handbook 实测 NULL launch 约 4.9 us(老机器)、RTX 3060
    约 1.2 us;A100 上常见引用是 2~3 us。所以本文取 1~5 us 这个区间,
    并且强调**比值比绝对值重要**。
  * 一个 decode step 的 kernel 数:公开资料给的量级是 300~1300。
    [F3] 按架构手数出来的数落在这个区间里,可互为印证。

运行:
    python launch_model.py
"""

from __future__ import annotations

import json
import os

import numpy as np

HERE = os.path.dirname(os.path.abspath(__file__))


# ══════════════════════════════════════════════════════════════
# [F1] 离散事件时间线
# ══════════════════════════════════════════════════════════════

def simulate_eager(n_kernel, t_launch, t_exec):
    """eager:CPU 逐个 cudaLaunchKernel,GPU 逐个执行,单流异步。

    CPU 发第 i 个要花 t_launch,发完就不用管了,可以继续发下一个;
    GPU 要等「前一个跑完」且「这个已经被发出去」才能开始。
    两个条件里慢的那个决定 GPU 什么时候开工,差出来的就是 GPU 空转。
    """
    gpu_free = 0.0
    idle = 0.0
    for i in range(n_kernel):
        cpu_done = (i + 1) * t_launch          # CPU 发完第 i 个的时刻
        start = max(gpu_free, cpu_done)        # GPU 能开工的时刻
        idle += start - gpu_free
        gpu_free = start + t_exec
    return {"total": gpu_free, "idle": idle,
            "idle_frac": idle / gpu_free if gpu_free > 0 else 0.0}


def simulate_graph(n_kernel, t_launch, t_exec, t_dispatch):
    """CUDA Graph:一次 cudaGraphLaunch 提交整张图。

    GPU 仍然要逐个节点过一遍前端,所以每个 kernel 还留一个 t_dispatch 的
    GPU 侧派发成本——只是不再需要 CPU 每次都插手。
    """
    total = t_launch + n_kernel * (t_exec + t_dispatch)
    return {"total": total, "idle": t_launch,
            "idle_frac": t_launch / total if total > 0 else 0.0}


def section_F1():
    print("\n" + "=" * 72)
    print("[F1] 时间线:eager 逐个发射 vs CUDA Graph 一次发射")
    print("=" * 72)

    N, tl, te, td = 700, 3.0, 2.0, 0.5
    e = simulate_eager(N, tl, te)
    g = simulate_graph(N, tl, te, td)
    print(f"\n  设定:N={N} 个 kernel,发射 {tl} us,执行 {te} us,"
          f"graph 内派发 {td} us")
    print(f"  eager      总时长 {e['total']:8.1f} us,GPU 空转 {e['idle']:7.1f} us"
          f"({e['idle_frac']*100:.0f}%)")
    print(f"  CUDA Graph 总时长 {g['total']:8.1f} us,GPU 空转 {g['idle']:7.1f} us"
          f"({g['idle_frac']*100:.0f}%)")
    print(f"  加速比 = {e['total']/g['total']:.2f}x")

    # 换一个大 kernel 场景:执行时间远大于发射
    e2 = simulate_eager(N, tl, 20.0)
    g2 = simulate_graph(N, tl, 20.0, td)
    print(f"\n  换成大 kernel(执行 20 us):")
    print(f"  eager      {e2['total']:8.1f} us,空转 {e2['idle_frac']*100:.0f}%")
    print(f"  CUDA Graph {g2['total']:8.1f} us,空转 {g2['idle_frac']*100:.0f}%")
    print(f"  加速比 = {e2['total']/g2['total']:.2f}x  -> 几乎没用")
    print(f"\n  结论:CUDA Graph 只在「单个 kernel 的执行时间接近或小于发射时间」")
    print(f"  时才赚钱。大 kernel 上它是纯负担(录制、显存、调试成本)。")

    return {"N": N, "t_launch": tl, "t_exec": te, "t_dispatch": td,
            "eager": e, "graph": g, "speedup": e["total"] / g["total"],
            "big": {"eager": e2, "graph": g2,
                    "speedup": e2["total"] / g2["total"]}}


# ══════════════════════════════════════════════════════════════
# [F2] 扫描:什么时候值得上 CUDA Graph
# ══════════════════════════════════════════════════════════════

def section_F2():
    print("\n" + "=" * 72)
    print("[F2] 扫描:单 kernel 执行时间 t_exec 对加速比的影响")
    print("=" * 72)

    N, tl, td = 700, 3.0, 0.5
    print(f"\n  N={N}, 发射 {tl} us, graph 内派发 {td} us")
    print(f"\n  {'t_exec(us)':>11} {'eager(us)':>11} {'graph(us)':>11} "
          f"{'加速':>7} {'eager空转':>9}")
    rows = []
    for te in [0.5, 1.0, 2.0, 3.0, 5.0, 8.0, 12.0, 20.0, 40.0]:
        e = simulate_eager(N, tl, te)
        g = simulate_graph(N, tl, te, td)
        rows.append({"t_exec": te, "eager_us": e["total"], "graph_us": g["total"],
                     "speedup": e["total"] / g["total"],
                     "idle_frac": e["idle_frac"]})
        print(f"  {te:>11.1f} {e['total']:>11.1f} {g['total']:>11.1f} "
              f"{e['total']/g['total']:6.2f}x {e['idle_frac']*100:8.0f}%")

    # 盈亏平衡点:eager 总时长 == graph 总时长
    print(f"\n  盈亏平衡:eager 靠 CPU 逐个发射,graph 每个 kernel 多付 {td} us 派发。")
    print(f"  当 t_exec > t_launch 时 eager 已经不让 GPU 空转,graph 只是白付 {td}。"
          f"")
    print(f"  本例盈亏点在 t_exec ≈ t_launch = {tl} us 附近。")
    return {"N": N, "t_launch": tl, "t_dispatch": td, "rows": rows}


# ══════════════════════════════════════════════════════════════
# [F3] 一步 decode 有多少个 kernel(按架构手数)
# ══════════════════════════════════════════════════════════════

# Llama-3-8B 的公开配置
LLAMA3_8B = dict(n_layer=32, d_model=4096, n_head=32, n_kv_head=8,
                 head_dim=128, d_ffn=14336, vocab=128256)


def section_F3():
    print("\n" + "=" * 72)
    print("[F3] 一步 decode 有多少个 kernel(按 Llama-3-8B 架构手数)")
    print("=" * 72)
    print("\n   这是**按架构推导**,不是实测。真实数字取决于编译器怎么 lowering,")
    print("    融合后能少一半以上。公开资料给的量级是 300~1300。")

    cfg = LLAMA3_8B
    # 每层:括号里是不融合时这个模块会拆成几个 kernel
    per_layer = [
        ("输入 RMSNorm", 4),        # 平方 / 归约 / 乘 / 缩放
        ("q_proj", 1),
        ("k_proj", 1),
        ("v_proj", 1),
        ("RoPE on q", 3),           # cos/sin 表取 + 旋转 + 拼接
        ("RoPE on k", 3),
        ("KV cache 写入", 2),       # k、v 各一次 scatter
        ("QK^T", 1),
        ("scale + mask", 2),
        ("softmax", 3),             # max / sub+exp / sum+div
        ("@V", 1),
        ("o_proj", 1),
        ("残差加", 1),
        ("后注意力 RMSNorm", 4),
        ("gate_proj", 1),
        ("up_proj", 1),
        ("SiLU", 1),
        ("逐元素乘", 1),
        ("down_proj", 1),
        ("残差加", 1),
    ]
    total_per_layer = sum(k for _, k in per_layer)
    n_layer = cfg["n_layer"]
    total = total_per_layer * n_layer + 5      # + 输出层 norm / lm_head / 采样等
    print(f"\n  每层 {total_per_layer} 个 kernel:")
    for name, k in per_layer:
        print(f"    {name:<22} {k}")
    print(f"\n  x {n_layer} 层 = {total_per_layer*n_layer},加输出头约 5 个")
    print(f"  合计 ≈ {total} 个 kernel / token")
    print(f"  (公开资料给的区间 300~1300,这个数落在里面)")

    # 融合后:把每层里能合的合掉
    fused_per_layer = [
        ("RMSNorm 融合", 1),
        ("QKV 一次 GEMM", 1),
        ("RoPE 融合", 1),
        ("KV cache 写入", 1),
        ("注意力融合(Flash)", 1),
        ("o_proj", 1),
        ("残差 + RMSNorm 融合", 1),
        ("gate/up 一次 GEMM", 1),
        ("SiLU + 乘 融合", 1),
        ("down_proj", 1),
        ("残差加", 1),
    ]
    fused_pl = sum(k for _, k in fused_per_layer)
    fused_total = fused_pl * n_layer + 3
    print(f"\n  融合后每层 {fused_pl} 个 -> 合计 ≈ {fused_total} 个/token")
    print(f"  kernel 数减少 {total/fused_total:.1f}x")
    print(f"\n  注意:融合减少的是 kernel 数,CUDA Graph 不减少 kernel 数,")
    print(f"  它只是把 {total} 次 CPU 发射合并成 1 次。两者正交,可以同时用。")

    return {"cfg": cfg, "per_layer": per_layer,
            "total_per_layer": total_per_layer, "total": total,
            "fused_per_layer": fused_per_layer, "fused_total": fused_total,
            "reduce": total / fused_total}


# ══════════════════════════════════════════════════════════════
# [F4] 分解:本机实测的加速里,发射和字节各占多少
# ══════════════════════════════════════════════════════════════

def section_F4(fus):
    print("\n" + "=" * 72)
    print("[F4] 分解:本机 K 段链的加速里,多少来自少发射、多少来自少搬字节")
    print("=" * 72)

    D = fus["D"]
    C = fus["C"]
    a = D["a_us"]                      # 每次调用的固定开销(实测)
    b = D["b_us_per_elem"]             # 每元素的边际成本(实测)
    K = C["K"]

    print(f"\n  实测固定开销 a = {a:.3f} us/次,边际 b = {b*1e6:.2f} us/百万元素")
    print(f"\n  {'n':>9} {'V1 实测':>10} {'V1 模型':>10} {'V3 实测':>10} "
          f"{'V3 模型':>10} {'发射贡献':>9}")

    rows = []
    for c in C["curves"]:
        n = c["n"]
        # V1: 2K 次调用,每次搬 n 个元素;V3: 2 次调用
        t1_model = 2 * K * (a + b * n)
        t3_model = 2 * (a + b * n)
        # 反事实:只把调用次数从 2K 降到 2,数据量不变(纯发射收益)
        t_launch_only = 2 * K * a + 2 * K * b * n - (2 * a + 2 * K * b * n)
        total_gain = t1_model - t3_model
        frac = t_launch_only / total_gain if total_gain > 0 else 1.0
        rows.append({"n": n, "t1_model": t1_model, "t3_model": t3_model,
                     "launch_frac": frac})
        print(f"  {n:>9} {c['t1_us']:9.2f} us {t1_model:9.2f} us "
              f"{c['t3_us']:9.2f} us {t3_model:9.2f} us {frac*100:8.0f}%")

    small = rows[0]
    big = rows[-1]
    print(f"\n  n={small['n']}({small['n']*4/1024:.1f} KiB):"
          f"模型说加速的 {small['launch_frac']*100:.0f}% 来自少发射,"
          f"只有 {(1-small['launch_frac'])*100:.0f}% 来自少搬字节。")
    print(f"  n={big['n']}({big['n']*4/1024/1024:.1f} MiB):"
          f"反过来,{big['launch_frac']*100:.0f}% 来自少发射,"
          f"{(1-big['launch_frac'])*100:.0f}% 来自少搬字节。")
    print(f"\n  这就是本文最想说的一句话:")
    print(f"    **小张量上,融合赚的几乎全是发射次数;大张量上,赚的才是字节。**")
    print(f"    CUDA Graph 只做前一件,而且做得更彻底(一次发射整张图)。")

    return {"a_us": a, "b_us_per_elem": b, "rows": rows}


# ══════════════════════════════════════════════════════════════

def main():
    res = {}
    res["F1"] = section_F1()
    res["F2"] = section_F2()
    res["F3"] = section_F3()

    p = os.path.join(HERE, "_fusion_results.json")
    if os.path.exists(p):
        with open(p, encoding="utf-8") as f:
            fus = json.load(f)
        if "C" in fus and "D" in fus:
            res["F4"] = section_F4(fus)
    else:
        print("\n[F4] 跳过:先跑 fusion_lab.py ALL 生成 _fusion_results.json")

    p = os.path.join(HERE, "_launch_results.json")
    with open(p, "w", encoding="utf-8") as f:
        json.dump(res, f, ensure_ascii=False, indent=2)
    print(f"\n结果已写入 {p}")


if __name__ == "__main__":
    main()

fusion_lab.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
fusion_lab.py —— 算子融合的「字节账本 + 实测校验」

写这篇文章的机器是一台 Apple M1 Pro,**没有 NVIDIA GPU,也没有 torch**,
所以这里不假装能测 CUDA kernel。能测的是两件在这台机器上真实存在的事:

  1. 内存有层级:数据留在片上和落回主存,代价差一个可测的倍数。
     融合做的事就是「少落回几次」,这个倍数能测,收益也能算。
  2. 每次调用算子都有一个与数据规模无关的固定成本(函数派发、缓冲检查、
     循环建立)。它能测出来,是 GPU 上 kernel launch 开销在 CPU 上的同构物。

GPU 上的具体数字(HBM 带宽、launch 微秒数)本文一律引公开资料并明确标注,
不拿这台机器外推。反过来,凡是标「实测」的数字,都出自本文件。

五组实验:
    [A] 字节账本    —— 几类常见 pattern 融合前后各搬多少字节(纯算术,精确)
    [B] 带宽–工作集 —— 实测片上/片外带宽比 rho
    [C] 融合三形态  —— 不融合 / 分块(留片上) / 代数合并,扫规模看各自值多少
    [D] 固定开销    —— 拟合出每次调用的固定成本,算它占总耗时的比例
    [E] attention尾 —— score 矩阵账本、本机 softmax 卡在哪

运行:
    python fusion_lab.py ALL
"""

from __future__ import annotations

import json
import os
import time

import numpy as np

HERE = os.path.dirname(os.path.abspath(__file__))
RNG_SEED = 20261002

MIB = float(1024 ** 2)
GIB = float(1024 ** 3)

# 逐元素链的形态:每段 h <- A*h + B,共 K 段
K_CHAIN = 8
A_S, B_S = 1.01, 0.01


def bench(fn, reps=7):
    """取多次里的最小值:最小值受系统抖动影响最小。"""
    ts = []
    for _ in range(reps):
        t0 = time.perf_counter()
        fn()
        ts.append(time.perf_counter() - t0)
    return min(ts)


def passes_to_bytes(n_elem, dtype_bytes, n_passes):
    """1 个 pass = 把整个张量读一遍再写一遍 = 2 * 字节数。"""
    return 2.0 * n_elem * dtype_bytes * n_passes


# ══════════════════════════════════════════════════════════════
# [A] 字节账本:融合前后各搬多少字节
# ══════════════════════════════════════════════════════════════

def section_A():
    print("\n" + "=" * 72)
    print("[A] 字节账本:融合前后各搬多少字节(纯算术,精确)")
    print("=" * 72)

    out = {}
    n = 4_000_000                            # 一个中等激活张量,fp32
    b = 4

    # ── A1. K 段逐元素链 ──────────────────────────────────
    K = K_CHAIN
    unf = passes_to_bytes(n, b, K)
    fus = passes_to_bytes(n, b, 1)
    print(f"\n  A1  {K} 段逐元素链({n/1e6:.0f}M 元素,fp32)")
    print(f"      不融合 {K} 个 kernel : {unf/MIB:8.1f} MiB")
    print(f"      融合成 1 个 kernel   : {fus/MIB:8.1f} MiB   -> 省 {unf/fus:.0f}x")
    out["chain"] = {"K": K, "unfused_MiB": unf / MIB, "fused_MiB": fus / MIB,
                    "ratio": unf / fus}

    # ── A2. RMSNorm ───────────────────────────────────────
    # 不融合:x*x(读写) / reduce mean(读) / x*r(读写) / *g(读写)
    unf = passes_to_bytes(n, b, 4)
    fus = passes_to_bytes(n, b, 1)
    print(f"\n  A2  RMSNorm({n/1e6:.0f}M 元素)")
    print(f"      不融合 4 遍 : {unf/MIB:8.1f} MiB")
    print(f"      融合  1 遍 : {fus/MIB:8.1f} MiB   -> 省 {unf/fus:.0f}x")
    out["rmsnorm"] = {"unfused_MiB": unf / MIB, "fused_MiB": fus / MIB,
                      "ratio": unf / fus}

    # ── A3. attention 的 score 矩阵 ───────────────────────
    Bs, H, S, D = 8, 16, 2048, 4096
    dt = 2                                    # bf16
    act_bytes = Bs * S * D * dt
    score_bytes = Bs * H * S * S * dt
    half = 1 + 2 + 2 + 1 + 2                  # softmax 五步各读写几份 score
    unf = half * score_bytes
    fus = 2 * score_bytes
    print(f"\n  A3  attention score(B={Bs}, H={H}, S={S}, d={D}, bf16)")
    print(f"      输入激活 [B,S,d]           : {act_bytes/MIB:8.1f} MiB")
    print(f"      score 矩阵 [B,H,S,S]       : {score_bytes/GIB:8.3f} GiB"
          f"   = 输入的 {score_bytes/act_bytes:.0f}x")
    print(f"      softmax 不融合 {half} 份读写 : {unf/GIB:8.3f} GiB")
    print(f"      softmax 融合             : {fus/GIB:8.3f} GiB   -> 省 {unf/fus:.1f}x")
    print(f"      额外物化一份 mask          : +{score_bytes/GIB:.3f} GiB(融合后消失)")
    print(f"\n      放大倍数 = H*S/d = {H}*{S}/{D} = {H*S/D:.0f}x")
    print(f"      S 翻一倍 -> 放大倍数翻一倍(score 按 S 的平方长,激活按 S 长)")
    for s in (1024, 2048, 4096, 8192):
        print(f"        S={s:>5}: score/激活 = {H*s/D:5.1f}x")
    out["softmax"] = {"B": Bs, "H": H, "S": S, "d": D,
                      "act_MiB": act_bytes / MIB, "score_GiB": score_bytes / GIB,
                      "blowup": score_bytes / act_bytes,
                      "unfused_GiB": unf / GIB, "fused_GiB": fus / GIB,
                      "ratio": unf / fus, "mask_GiB": score_bytes / GIB,
                      "blowup_curve": [{"S": s, "x": H * s / D}
                                       for s in (1024, 2048, 4096, 8192)]}

    # ── A4. 因果掩码:块级跳过能省多少 ────────────────────
    total = S * S
    kept_exact = S * (S + 1) // 2
    rows = []
    print(f"\n  A4  因果掩码 S={S}:块级跳过 vs 精确跳过")
    for br in (64, 128, 256):
        nb = S // br
        blocks_done = nb * (nb + 1) // 2      # query 块 i 只算 j<=i 的 key 块
        rows.append({"br": br, "nb": nb, "frac": blocks_done / (nb * nb)})
        print(f"      块边长 {br:>4}: 算 {blocks_done:>4}/{nb*nb:<4} 块"
              f" = {blocks_done/(nb*nb)*100:5.1f}% 工作量")
    print(f"      理论上界(逐元素精确): {kept_exact/total*100:5.1f}%")
    print(f"      完全不跳            : {100.0:5.1f}%  -> 白算 "
          f"{100-kept_exact/total*100:.1f}%")
    out["causal"] = {"S": S, "exact_frac": kept_exact / total, "blocks": rows}

    return out


# ══════════════════════════════════════════════════════════════
# [B] 带宽 vs 工作集:实测片上/片外带宽比
# ══════════════════════════════════════════════════════════════

def section_B():
    print("\n" + "=" * 72)
    print("[B] 带宽 vs 工作集(np.add(x, 1, out=y):1 读 1 写)")
    print("=" * 72)

    rng = np.random.default_rng(RNG_SEED)
    print(f"\n  {'工作集':>12}  {'耗时':>10}  {'带宽':>10}   归属")
    rows = []
    for nbytes in [32 << 10, 128 << 10, 512 << 10, 2 << 20, 8 << 20,
                   32 << 20, 128 << 20]:
        n = nbytes // 4
        x = rng.standard_normal(n).astype(np.float32)
        y = np.empty_like(x)
        bench(lambda: np.add(x, 1.0, out=y), reps=3)
        t = bench(lambda: np.add(x, 1.0, out=y), reps=11)
        bw = 2 * nbytes / t / 1e9
        tag = "onchip" if nbytes <= (8 << 20) else "dram"
        rows.append({"KiB": nbytes >> 10, "ms": t * 1e3, "GBs": bw, "tag": tag})
        print(f"  {nbytes>>10:>10} KiB  {t*1e3:9.3f} ms  {bw:9.1f} GB/s   {tag}")
        del x, y

    b_cache = float(np.median([r["GBs"] for r in rows if r["tag"] == "onchip"]))
    b_dram = float(np.median([r["GBs"] for r in rows if r["tag"] == "dram"]))
    rho = b_cache / b_dram
    print(f"\n  片上带宽中位数 B_cache = {b_cache:.1f} GB/s")
    print(f"  主存带宽中位数 B_dram  = {b_dram:.1f} GB/s")
    print(f"  比值 rho = {rho:.3f}   <- 决定融合能兑现多少的关键参数")
    print(f"\n  注意:这里的 B_cache 不是缓存的峰值带宽,而是单线程 numpy 逐元素")
    print(f"  循环的吞吐上限。它只比主存快 {rho:.2f} 倍,所以在这台机器上")
    print(f"  「把中间结果留在片上」这件事本身几乎不值钱(见 [C] 的 V2)。")
    return {"rows": rows, "B_cache_GBs": b_cache, "B_dram_GBs": b_dram, "rho": rho}


# ══════════════════════════════════════════════════════════════
# [C] 融合三形态 × 规模扫描
# ══════════════════════════════════════════════════════════════

def _v1_unfused(x, buf, K=K_CHAIN):
    """V1 不融合:K 段,每段 2 次算子调用,每段结果都写回内存。"""
    np.multiply(x, A_S, out=buf)
    np.add(buf, B_S, out=buf)
    for _ in range(K - 1):
        np.multiply(buf, A_S, out=buf)
        np.add(buf, B_S, out=buf)
    return buf


def _v2_tiled(x, out, tile, K=K_CHAIN):
    """V2 分块融合:一块读进来,K 段都在片上算完,只写回一次。

    这是 GPU 融合 kernel 的 CPU 近似:tile 就是寄存器/共享内存那块片上缓冲。
    numpy 做不到寄存器级融合(每个 ufunc 仍要写回 tile),所以 V2 是本机能
    做到的最好近似。
    """
    t = np.empty(tile, dtype=np.float32)
    for s in range(0, x.size, tile):
        m = min(tile, x.size - s)
        v = t[:m]
        np.multiply(x[s:s + m], A_S, out=v)
        np.add(v, B_S, out=v)
        for _ in range(K - 1):
            np.multiply(v, A_S, out=v)
            np.add(v, B_S, out=v)
        out[s:s + m] = v
    return out


def _v3_closed(x, out, K=K_CHAIN):
    """V3 代数合并:K 段仿射合成一个仿射,只剩 2 次调用。

    h_K = A^K * h_0 + B*(A^K - 1)/(A - 1)
    中间结果根本不存在,连「留在片上」都不需要。
    """
    Ak = A_S ** K
    Bk = B_S * (Ak - 1.0) / (A_S - 1.0)
    np.multiply(x, Ak, out=out)
    np.add(out, Bk, out=out)
    return out


def section_C(rho):
    print("\n" + "=" * 72)
    print(f"[C] 融合三形态 × 规模扫描(K={K_CHAIN} 段 h <- {A_S}*h + {B_S})")
    print("=" * 72)

    rng = np.random.default_rng(RNG_SEED)
    sizes = [512, 4096, 65536, 262144, 1 << 20, 4_000_000]
    curves = []
    print(f"\n  {'n':>9} {'V1 不融合':>11} {'V3 代数合并':>11} "
          f"{'V1/V3':>7} {'相对误差':>10}")
    for n in sizes:
        x = rng.standard_normal(n).astype(np.float32)
        buf, out = np.empty_like(x), np.empty_like(x)
        reps = 5000 if n <= 65536 else 11
        bench(lambda: _v1_unfused(x, buf), reps=3)
        t1 = bench(lambda: _v1_unfused(x, buf), reps=reps)
        bench(lambda: _v3_closed(x, out), reps=3)
        t3 = bench(lambda: _v3_closed(x, out), reps=reps)
        rel = float(np.max(np.abs(buf - out)) / np.max(np.abs(buf)))
        curves.append({"n": n, "t1_us": t1 * 1e6, "t3_us": t3 * 1e6,
                       "speedup": t1 / t3, "rel_err": rel})
        print(f"  {n:>9} {t1*1e6:9.2f} us {t3*1e6:9.2f} us "
              f"{t1/t3:6.2f}x {rel:10.2e}")

    # V2 只在最大规模上跑:它慢,且只有这里才谈得上「主存 vs 片上」
    print(f"\n  V2 分块融合(n=4,000,000,扫 tile):")
    x = rng.standard_normal(4_000_000).astype(np.float32)
    buf, out = np.empty_like(x), np.empty_like(x)
    bench(lambda: _v1_unfused(x, buf), reps=3)
    t1_big = bench(lambda: _v1_unfused(x, buf), reps=11)
    print(f"    V1 不融合基准: {t1_big*1e3:.3f} ms")
    v2 = []
    for tile in [16384, 65536, 262144, 1 << 20]:
        bench(lambda: _v2_tiled(x, out, tile), reps=3)
        t = bench(lambda: _v2_tiled(x, out, tile), reps=11)
        v2.append({"tile": tile, "ms": t * 1e3, "speedup": t1_big / t})
        print(f"    tile={tile:>8}: {t*1e3:8.3f} ms  {t1_big/t:5.2f}x")

    best2 = max(v2, key=lambda r: r["speedup"])
    K = K_CHAIN
    pred = K / (1 + (K - 1) / rho)
    print(f"\n  V2 最佳 {best2['speedup']:.2f}x(tile={best2['tile']})")
    print(f"  模型预测 K/(1+(K-1)/rho) = {K}/(1+{K-1}/{rho:.3f}) = {pred:.2f}x")
    print(f"  -> 模型说 V2 几乎没收益,实测确实几乎没有。两者一致。")
    print(f"\n  V3 的收益不受 rho 限制:中间结果根本不存在,")
    print(f"  直接就是字节比 K = {K}x(实测 "
          f"{min(c['speedup'] for c in curves):.2f}~"
          f"{max(c['speedup'] for c in curves):.2f}x,在 K 附近浮动)。")

    return {"K": K, "rho": rho, "curves": curves,
            "v2": v2, "v2_best": best2, "v2_pred": pred,
            "t1_big_ms": t1_big * 1e3}


# ══════════════════════════════════════════════════════════════
# [D] 固定开销:每次调用的地板成本
# ══════════════════════════════════════════════════════════════

def section_D():
    print("\n" + "=" * 72)
    print("[D] 单次调用的固定开销(np.add(x, 1, out=y))")
    print("=" * 72)

    rng = np.random.default_rng(RNG_SEED)
    sizes = [1, 8, 64, 512, 4096, 16384, 65536, 262144, 1 << 20]
    print(f"\n  {'n':>9}  {'单次耗时':>11}")
    rows = []
    for n in sizes:
        a = rng.standard_normal(n).astype(np.float32)
        b = np.empty_like(a)
        reps = 20000 if n <= 4096 else 500
        t0 = time.perf_counter()
        for _ in range(reps):
            np.add(a, 1.0, out=b)
        t = (time.perf_counter() - t0) / reps
        rows.append({"n": n, "us": t * 1e6})
        print(f"  {n:>9}  {t*1e6:9.3f} us")

    # 只用小端拟合:大端会被缓存/主存的拐点污染
    ns = np.array([r["n"] for r in rows if r["n"] <= 16384], dtype=float)
    ts = np.array([r["us"] for r in rows if r["n"] <= 16384], dtype=float)
    A = np.stack([np.ones_like(ns), ns], axis=1)
    coef, *_ = np.linalg.lstsq(A, ts, rcond=None)
    a_us, slope = float(coef[0]), float(coef[1])
    print(f"\n  拟合 t = a + b*n(n <= 16384):")
    print(f"    固定开销 a = {a_us:.3f} us / 次")
    print(f"    边际成本 b = {slope*1e6:.3f} us / 百万元素")

    n_half = a_us / slope if slope > 0 else float("inf")
    print(f"    开销与数据各占一半的临界规模 n* = a/b = {n_half:.0f} 元素"
          f"({n_half*4/1024:.1f} KiB)")

    print(f"\n  {K_CHAIN} 段链({2*K_CHAIN} 次调用)里固定开销的占比:")
    share = []
    for r in rows:
        n = r["n"]
        fixed = 2 * K_CHAIN * a_us
        tc = 2 * K_CHAIN * (a_us + slope * n)
        frac = fixed / tc if tc > 0 else 1.0
        share.append({"n": n, "fixed_frac": frac})
        print(f"    n={n:>9}: 固定部分占 {frac*100:5.1f}%")

    return {"rows": rows, "a_us": a_us, "b_us_per_elem": slope,
            "n_half": n_half, "share": share, "K": K_CHAIN}


# ══════════════════════════════════════════════════════════════
# [E] attention 尾部:本机 softmax 卡在哪
# ══════════════════════════════════════════════════════════════

def _softmax_unfused(x, out):
    m = np.max(x, axis=-1, keepdims=True)          # 读
    np.subtract(x, m, out=out)                     # 读 x + 写 out
    np.exp(out, out=out)                           # 读 + 写
    s = np.sum(out, axis=-1, keepdims=True)        # 读
    np.divide(out, s, out=out)                     # 读 + 写
    return out


def _softmax_tiled(x, out, rows):
    for i in range(0, x.shape[1], rows):
        sl = slice(i, i + rows)
        blk, dst = x[:, sl, :], out[:, sl, :]
        m = np.max(blk, axis=-1, keepdims=True)
        np.subtract(blk, m, out=dst)
        np.exp(dst, out=dst)
        s = np.sum(dst, axis=-1, keepdims=True)
        np.divide(dst, s, out=dst)
    return out


def section_E(b_dram):
    print("\n" + "=" * 72)
    print("[E] attention 尾部:本机 softmax 到底卡在哪")
    print("=" * 72)

    rng = np.random.default_rng(RNG_SEED)
    H, S = 8, 1024
    x = rng.standard_normal((H, S, S)).astype(np.float32)
    nbytes = x.nbytes
    print(f"\n  score 矩阵 [H={H}, S={S}, S={S}] = {nbytes/MIB:.1f} MB")

    out = np.empty_like(x)
    bench(lambda: _softmax_unfused(x, out), reps=3)
    t_unf = bench(lambda: _softmax_unfused(x, out), reps=11)
    ref = out.copy()

    o2 = np.empty_like(x)
    bench(lambda: _softmax_tiled(x, o2, 64), reps=3)
    t_til = bench(lambda: _softmax_tiled(x, o2, 64), reps=11)
    diff = float(np.max(np.abs(ref - o2)))

    traffic = 8 * nbytes                            # [A3] 里的 8 份读写
    roof = traffic / (b_dram * 1e9) * 1e3
    print(f"\n  不融合(5 个 kernel,共 {traffic/MIB:.0f} MB 读写): {t_unf*1e3:8.3f} ms")
    print(f"  分块融合(rows=64)                        : {t_til*1e3:8.3f} ms"
          f"   speedup={t_unf/t_til:.2f}x  diff={diff:.1e}")
    print(f"\n  纯访存下界 = {traffic/MIB:.0f} MB / {b_dram:.0f} GB/s = {roof:.3f} ms")
    print(f"  实测 / 下界 = {t_unf*1e3/roof:.2f}x")
    print(f"  -> 远高于 1:这里卡的是 exp 的计算本身,不是访存。")
    print(f"     分块省了字节但一个 exp 都没少,所以反而更慢。")
    print(f"     这正是「融合省不了算错的东西」的直接证据。")

    return {"H": H, "S": S, "score_MB": nbytes / MIB,
            "t_unfused_ms": t_unf * 1e3, "t_tiled_ms": t_til * 1e3,
            "speedup": t_unf / t_til, "diff": diff,
            "traffic_MB": traffic / MIB, "roof_ms": roof,
            "over_roof": t_unf * 1e3 / roof}


# ══════════════════════════════════════════════════════════════

def main():
    import sys
    which = sys.argv[1] if len(sys.argv) > 1 else "ALL"
    res = {}

    if which in ("ALL", "A"):
        res["A"] = section_A()
    if which in ("ALL", "B"):
        res["B"] = section_B()
    if which in ("ALL", "C"):
        rho = res.get("B", {}).get("rho")
        if rho is None:
            res["B"] = section_B()
            rho = res["B"]["rho"]
        res["C"] = section_C(rho)
    if which in ("ALL", "D"):
        res["D"] = section_D()
    if which in ("ALL", "E"):
        bd = res.get("B", {}).get("B_dram_GBs")
        if bd is None:
            res["B"] = section_B()
            bd = res["B"]["B_dram_GBs"]
        res["E"] = section_E(bd)

    p = os.path.join(HERE, "_fusion_results.json")
    with open(p, "w", encoding="utf-8") as f:
        json.dump(res, f, ensure_ascii=False, indent=2)
    print(f"\n结果已写入 {p}")


if __name__ == "__main__":
    main()

make_figures.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
make_figures.py —— 画本文的 5 张配图。

数据全部读已经跑完的实验(_fusion_results.json / _launch_results.json),
不在这里重新算,避免图上的数字和正文漂移。

运行:
    python make_figures.py
"""

from __future__ import annotations

import json
import os

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np

HERE = os.path.dirname(os.path.abspath(__file__))
FIGDIR = os.path.join(HERE, "..", "figures")
try:
    os.makedirs(FIGDIR, exist_ok=True)
except FileExistsError:
    pass

# 配色(正文写「这张图要看什么」时按这几个名字描述)
C_MAIN = "#1f4e79"     # 深蓝:主曲线
C_ALT = "#c1440e"      # 橙红:对照曲线
C_GREEN = "#2e7d32"    # 绿:第三组 / 融合后
C_PURPLE = "#7b1fa2"   # 紫:标注线
C_GRAY = "#8a8a8a"     # 灰:参考线
C_LIGHT = "#bcd7ee"    # 浅蓝:填充
C_RED = "#b71c1c"      # 红:越界 / 地板

plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"]
plt.rcParams["axes.unicode_minus"] = False
plt.rcParams["figure.dpi"] = 130
plt.rcParams["savefig.dpi"] = 130

MIB = float(1024 ** 2)


def _load(name):
    p = os.path.join(HERE, name)
    if not os.path.exists(p):
        raise SystemExit(f"缺少 {name},先跑 fusion_lab.py ALL / launch_model.py")
    with open(p, encoding="utf-8") as f:
        return json.load(f)


# ══════════════════════════════════════════════════════════════
# 图 1:字节账本
# ══════════════════════════════════════════════════════════════
def fig_traffic(fus):
    A = fus["A"]
    items = [
        ("8 段逐元素链\n(4M 元素 fp32)",
         A["chain"]["unfused_MiB"], A["chain"]["fused_MiB"]),
        ("RMSNorm\n(4M 元素 fp32)",
         A["rmsnorm"]["unfused_MiB"], A["rmsnorm"]["fused_MiB"]),
        ("attention softmax\n(B=8,H=16,S=2048,bf16)",
         A["softmax"]["unfused_GiB"] * 1024, A["softmax"]["fused_GiB"] * 1024),
    ]
    labels = [t[0] for t in items]
    unf = [t[1] for t in items]
    fs = [t[2] for t in items]

    fig, ax = plt.subplots(figsize=(9.2, 4.4))
    y = np.arange(len(items))
    h = 0.34
    ax.barh(y + h / 2, unf, height=h, color=C_ALT, label="不融合")
    ax.barh(y - h / 2, fs, height=h, color=C_GREEN, label="融合后")
    for i, (u, f) in enumerate(zip(unf, fs)):
        ax.text(u * 1.15, i + h / 2, f"{u:.1f} MiB", va="center",
                fontsize=9, color=C_ALT)
        ax.text(f * 1.15, i - h / 2, f"{f:.1f} MiB", va="center",
                fontsize=9, color=C_GREEN)
        ax.annotate(f"{u/f:.0f}x", xy=(u, i), xytext=(u * 1.15, i + 0.42),
                    fontsize=9, color=C_MAIN, fontweight="bold")
    ax.set_yticks(y)
    ax.set_yticklabels(labels, fontsize=9)
    ax.set_xscale("log")
    ax.set_xlim(8, 40000)
    ax.set_xlabel("搬动的字节数(对数刻度)", fontsize=10)
    ax.set_title("图 1:融合前后各搬多少字节", fontsize=12, color=C_MAIN)
    ax.legend(loc="lower right", fontsize=9)
    ax.grid(axis="x", alpha=0.3, linestyle=":")
    ax.set_axisbelow(True)
    fig.tight_layout()
    fig.savefig(os.path.join(FIGDIR, "fig_traffic.png"))
    plt.close(fig)


# ══════════════════════════════════════════════════════════════
# 图 2:score 矩阵的放大倍数 + 因果块级跳过
# ══════════════════════════════════════════════════════════════
def fig_blowup(fus):
    A = fus["A"]
    sm = A["softmax"]
    ca = A["causal"]

    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.4, 4.2))

    # 左:放大倍数随 S 增长
    ss = [d["S"] for d in sm["blowup_curve"]]
    xx = [d["x"] for d in sm["blowup_curve"]]
    ax1.plot(ss, xx, "o-", color=C_MAIN, linewidth=2, markersize=7)
    ax1.plot(ss, ss, "--", color=C_GRAY, linewidth=1.2,
             label="线性参考(若按 S 增长)")
    ax1.set_xscale("log", base=2)
    ax1.set_yscale("log", base=2)
    ax1.set_xticks(ss)
    ax1.set_xticklabels([str(s) for s in ss], fontsize=9)
    ax1.set_yticks(xx)
    ax1.set_yticklabels([f"{v:.0f}x" for v in xx], fontsize=9)
    ax1.minorticks_off()
    ax1.set_xlabel("序列长度 S", fontsize=10)
    ax1.set_ylabel("score 矩阵 / 输入激活", fontsize=10)
    ax1.set_title(f"中间产物被放大 H*S/d = {sm['H']}*S/{sm['d']} 倍",
                  fontsize=11, color=C_MAIN)
    ax1.grid(alpha=0.3, linestyle=":")
    ax1.legend(fontsize=8, loc="upper left")
    ax1.set_axisbelow(True)
    for s, v in zip(ss, xx):
        ax1.annotate(f"{v:.0f}x", (s, v), textcoords="offset points",
                     xytext=(6, -12), fontsize=8, color=C_MAIN)

    # 右:因果块级跳过
    brs = [str(b["br"]) for b in ca["blocks"]]
    fr = [b["frac"] * 100 for b in ca["blocks"]]
    bars = ax2.bar(brs, fr, color=C_LIGHT, edgecolor=C_MAIN, width=0.55)
    ax2.axhline(ca["exact_frac"] * 100, color=C_GREEN, linestyle="--",
                linewidth=1.6, label=f"理论上界 {ca['exact_frac']*100:.1f}%")
    ax2.axhline(100, color=C_GRAY, linestyle=":", linewidth=1.2,
                label="不跳过 100%")
    for b, v in zip(bars, fr):
        ax2.text(b.get_x() + b.get_width() / 2, v + 1.2, f"{v:.1f}%",
                 ha="center", fontsize=9, color=C_MAIN)
    ax2.set_ylim(0, 118)
    ax2.set_xlabel("块边长", fontsize=10)
    ax2.set_ylabel("实际算的工作量占比 (%)", fontsize=10)
    ax2.set_title(f"因果掩码块级跳过(S={ca['S']})", fontsize=11, color=C_MAIN)
    ax2.legend(fontsize=8, loc="upper left")
    ax2.grid(axis="y", alpha=0.3, linestyle=":")
    ax2.set_axisbelow(True)

    fig.suptitle("图 2:attention 的两本账——中间产物多大、白算多少",
                 fontsize=12, color=C_MAIN, y=1.0)
    fig.tight_layout()
    fig.savefig(os.path.join(FIGDIR, "fig_blowup.png"))
    plt.close(fig)


# ══════════════════════════════════════════════════════════════
# 图 3:融合三形态 × 规模扫描
# ══════════════════════════════════════════════════════════════
def fig_three_forms(fus):
    C = fus["C"]
    D = fus["D"]
    K = C["K"]
    a = D["a_us"]
    cur = C["curves"]
    ns = [c["n"] for c in cur]
    t1 = [c["t1_us"] for c in cur]
    t3 = [c["t3_us"] for c in cur]
    sp = [c["speedup"] for c in cur]

    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.6, 4.3))

    # 左:耗时 vs 规模
    ax1.loglog(ns, t1, "o-", color=C_ALT, linewidth=2, markersize=6,
               label="V1 不融合(2K 次调用)")
    ax1.loglog(ns, t3, "s-", color=C_GREEN, linewidth=2, markersize=6,
               label="V3 代数合并(2 次调用)")
    floor = 2 * K * a
    ax1.axhline(floor, color=C_RED, linestyle="--", linewidth=1.4,
                label=f"V1 的发射地板 2K*a = {floor:.1f} μs")
    ax1.axhline(2 * a, color=C_PURPLE, linestyle=":", linewidth=1.4,
                label=f"V3 的发射地板 2*a = {2*a:.1f} μs")
    # V2 在最大规模上的结果
    v2b = C["v2_best"]
    ax1.plot([ns[-1]], [v2b["ms"] * 1e3], "D", color=C_GRAY, markersize=9,
             label=f"V2 分块融合 {v2b['ms']*1e3:.0f} μs")
    ax1.set_xlabel("张量元素数", fontsize=10)
    ax1.set_ylabel("单次耗时 (μs)", fontsize=10)
    ax1.set_title(f"三种形态的耗时(K={K})", fontsize=11, color=C_MAIN)
    ax1.legend(fontsize=8, loc="upper left")
    ax1.grid(alpha=0.3, linestyle=":")
    ax1.set_axisbelow(True)

    # 右:加速比
    ax2.semilogx(ns, sp, "o-", color=C_MAIN, linewidth=2, markersize=6,
                 label="V1 / V3(代数合并)")
    ax2.axhline(K, color=C_GREEN, linestyle="--", linewidth=1.6,
                label=f"字节比上限 K = {K}x")
    ax2.axhline(C["v2_pred"], color=C_GRAY, linestyle=":", linewidth=1.6,
                label=f"V2 模型预测 {C['v2_pred']:.2f}x")
    ax2.plot([ns[-1]], [v2b["speedup"]], "D", color=C_ALT, markersize=9,
             label=f"V2 实测 {v2b['speedup']:.2f}x")
    ax2.set_ylim(0, K * 1.25)
    ax2.set_xlabel("张量元素数", fontsize=10)
    ax2.set_ylabel("加速比", fontsize=10)
    ax2.set_title("各自兑现了多少", fontsize=11, color=C_MAIN)
    ax2.legend(fontsize=8, loc="lower left")
    ax2.grid(alpha=0.3, linestyle=":")
    ax2.set_axisbelow(True)

    fig.suptitle("图 3:融合三形态 × 规模扫描(本机实测,M1 Pro / numpy)",
                 fontsize=12, color=C_MAIN, y=1.0)
    fig.tight_layout()
    fig.savefig(os.path.join(FIGDIR, "fig_three_forms.png"))
    plt.close(fig)


# ══════════════════════════════════════════════════════════════
# 图 4:单次调用的固定开销
# ══════════════════════════════════════════════════════════════
def fig_launch(fus):
    D = fus["D"]
    rows = D["rows"]
    ns = np.array([r["n"] for r in rows], dtype=float)
    us = np.array([r["us"] for r in rows], dtype=float)
    a, b = D["a_us"], D["b_us_per_elem"]
    n_half = D["n_half"]

    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.6, 4.2))

    # 左:延迟 vs 规模,含地板
    ax1.loglog(ns, us, "o", color=C_MAIN, markersize=7, label="实测")
    grid = np.logspace(0, np.log10(ns.max()), 60)
    ax1.loglog(grid, a + b * grid, "-", color=C_ALT, linewidth=1.8,
               label=f"拟合 a + b*n(a={a:.2f} μs)")
    ax1.axhline(a, color=C_RED, linestyle="--", linewidth=1.6,
                label=f"固定开销地板 a = {a:.2f} μs")
    ax1.axvline(n_half, color=C_PURPLE, linestyle=":", linewidth=1.6,
                label=f"各占一半 n* = {n_half:.0f} 元素")
    ax1.set_xlabel("张量元素数", fontsize=10)
    ax1.set_ylabel("单次调用耗时 (μs)", fontsize=10)
    ax1.set_title("小到一定程度,耗时就和数据无关了", fontsize=11, color=C_MAIN)
    ax1.legend(fontsize=8, loc="upper left")
    ax1.grid(alpha=0.3, linestyle=":")
    ax1.set_axisbelow(True)

    # 右:一条 K 段链里固定开销的占比
    share = D["share"]
    sn = np.array([s["n"] for s in share], dtype=float)
    sf = np.array([s["fixed_frac"] for s in share]) * 100
    ax2.semilogx(sn, sf, "o-", color=C_MAIN, linewidth=2, markersize=6)
    ax2.axhline(50, color=C_GRAY, linestyle=":", linewidth=1.3)
    ax2.fill_between(sn, 0, sf, color=C_LIGHT, alpha=0.55)
    for x, v in zip(sn, sf):
        if v > 55 or v < 12:
            ax2.annotate(f"{v:.0f}%", (x, v), textcoords="offset points",
                         xytext=(0, 8), fontsize=8, color=C_MAIN,
                         ha="center")
    ax2.set_ylim(0, 108)
    ax2.set_xlabel("张量元素数", fontsize=10)
    ax2.set_ylabel("固定开销占总耗时 (%)", fontsize=10)
    ax2.set_title(f"{D['K']} 段链({2*D['K']} 次调用)里,发射占多少",
                  fontsize=11, color=C_MAIN)
    ax2.grid(alpha=0.3, linestyle=":")
    ax2.set_axisbelow(True)

    fig.suptitle("图 4:每次调用都有一个和数据规模无关的地板(本机实测)",
                 fontsize=12, color=C_MAIN, y=1.0)
    fig.tight_layout()
    fig.savefig(os.path.join(FIGDIR, "fig_launch.png"))
    plt.close(fig)


# ══════════════════════════════════════════════════════════════
# 图 5:CUDA Graph 的时间线
# ══════════════════════════════════════════════════════════════
def fig_graph(lnc):
    F1 = lnc["F1"]
    F2 = lnc["F2"]
    tl, te, td = F1["t_launch"], F1["t_exec"], F1["t_dispatch"]

    fig, (ax1, ax2) = plt.subplots(
        1, 2, figsize=(11.2, 4.3), gridspec_kw={"width_ratios": [1.25, 1]})

    # ── 左:甘特图,只画前 8 个 kernel ──
    n_show = 8
    # eager:CPU 逐个发射,GPU 要等「前一个跑完」且「这个已发出」才能开工。
    # 时刻逐事件推算,和图里的条形严格同源。
    gpu_free, segs_cpu, segs_gpu, segs_idle = 0.0, [], [], []
    for i in range(n_show):
        cpu_done = (i + 1) * tl
        start = max(gpu_free, cpu_done)
        if start > gpu_free:
            segs_idle.append((gpu_free, start - gpu_free))
        segs_gpu.append((start, te))
        segs_cpu.append((i * tl, tl))
        gpu_free = start + te
    ax1.broken_barh(segs_cpu, (3.4, 0.8), facecolors=C_LIGHT,
                    edgecolor=C_MAIN, linewidth=0.6)
    ax1.broken_barh(segs_gpu, (2.2, 0.8), facecolors=C_GREEN,
                    edgecolor=C_GREEN, linewidth=0.6)
    ax1.broken_barh(segs_idle, (2.2, 0.8), facecolors="#f0c9c9",
                    edgecolor=C_RED, linewidth=0.6, hatch="//")

    # graph:一次发射,之后背靠背
    g_cpu = [(0, tl)]
    g_gpu = [(tl, n_show * (te + td))]
    ax1.broken_barh(g_cpu, (1.0, 0.8), facecolors=C_LIGHT,
                    edgecolor=C_MAIN, linewidth=0.6)
    ax1.broken_barh(g_gpu, (-0.2, 0.8), facecolors=C_GREEN,
                    edgecolor=C_GREEN, linewidth=0.6)

    ax1.set_yticks([3.8, 2.6, 1.4, 0.2])
    ax1.set_yticklabels(["CPU 发射", "GPU 执行", "CPU 发射", "GPU 执行"],
                        fontsize=9)
    ax1.set_xlabel("时间 (μs)", fontsize=10)
    ax1.set_xlim(-0.5, n_show * tl + te + 3)
    ax1.set_ylim(-0.6, 5.4)
    ax1.set_title(f"eager(上)vs CUDA Graph(下)  "
                  f"发射 {tl} μs / 执行 {te} μs", fontsize=10.5, color=C_MAIN)
    ax1.grid(axis="x", alpha=0.3, linestyle=":")
    ax1.set_axisbelow(True)
    ax1.text(0.4, 4.95, "eager:每次发射 GPU 都要等(斜纹 = 空转)",
             fontsize=8.5, color=C_ALT, va="center")
    ax1.text(0.4, 0.78, "graph:只发射一次,之后 GPU 背靠背(无空转)",
             fontsize=8.5, color=C_GREEN, va="center")

    # ── 右:加速比 vs 单 kernel 执行时间 ──
    rws = F2["rows"]
    tes = [r["t_exec"] for r in rws]
    sps = [r["speedup"] for r in rws]
    ax2.plot(tes, sps, "o-", color=C_MAIN, linewidth=2, markersize=6)
    ax2.axhline(1.0, color=C_GRAY, linestyle=":", linewidth=1.3,
                label="盈亏线 1.0x")
    ax2.axvline(F2["t_launch"], color=C_RED, linestyle="--", linewidth=1.5,
                label=f"发射成本 {F2['t_launch']} μs")
    ax2.fill_between(tes, 0, 1.0, where=np.array(sps) < 1.0,
                     color="#f0c9c9", alpha=0.6, label="graph 反而更慢")
    ax2.set_xscale("log")
    ax2.set_xlabel("单个 kernel 的执行时间 (μs)", fontsize=10)
    ax2.set_ylabel("CUDA Graph 加速比", fontsize=10)
    ax2.set_title(f"N={F2['N']} 个 kernel", fontsize=11, color=C_MAIN)
    ax2.legend(fontsize=8, loc="lower right")
    ax2.grid(alpha=0.3, linestyle=":")
    ax2.set_axisbelow(True)

    fig.suptitle("图 5:CUDA Graph 省的是发射,不是字节(模型推演)",
                 fontsize=12, color=C_MAIN, y=1.0)
    fig.tight_layout()
    fig.savefig(os.path.join(FIGDIR, "fig_graph.png"))
    plt.close(fig)


# ══════════════════════════════════════════════════════════════

def main():
    fus = _load("_fusion_results.json")
    lnc = _load("_launch_results.json")
    fig_traffic(fus)
    fig_blowup(fus)
    fig_three_forms(fus)
    fig_launch(fus)
    fig_graph(lnc)
    print("已生成 5 张图:")
    for f in sorted(os.listdir(FIGDIR)):
        if f.endswith(".png"):
            p = os.path.join(FIGDIR, f)
            print(f"  {f}  {os.path.getsize(p)/1024:.0f} KiB")


if __name__ == "__main__":
    main()
0

评论 (0)

取消
粤ICP备2021042327号