所属方向:推理加速 | 难度:进阶 | 前置知识:性能建模与 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 专门交代了这条边界。
先摆三个数,都是本篇附录里真跑出来的。
第一个数:一步 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 开工」这两件事上。
三句话:
cudaGraphLaunch 提交整张图。它一个字节都不省,省的是 CPU 侧那 1093 次提交。一句总结:融合省字节,Graph 省发射。把两者混为一谈,是这一块最常见的错误。
前置文章 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$ 变小。
考虑 $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}$$
这个式子值得盯着看三秒。 它的两个极限都很干净:
所以「融合能加速几倍」这个问题,答案既不是 $K$,也不是玄学,而是 $K$ 和 $\rho$ 共同决定的。记住这一点,第六节会看到它把一台机器上的实测结果解释得干干净净。
一次算子调用的耗时,在规模足够小时和数据量无关。写成仿射形式:
$$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 上会出现同样性质的墙。
考虑单流、异步提交的 $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)$$
两个条件谁慢,谁决定节奏。分两种情形:
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 是纯负担。
因果注意力里,第 $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)$ 衰减,所以块越小越接近理论上界——但块越小,片上数据被切得越碎,调度开销越大。这是块大小必须折中的原因,不是调参玄学。

这张图要看什么:左图是 score 矩阵相对输入激活的放大倍数随 $S$ 的变化,灰虚线是「若按 $S$ 线性增长」的参考——实际曲线比线性还陡($H \cdot S/d$ 是 $S$ 的一次式,但在对数轴上叠加了 $H/d$ 的常数放大,$S=8192$ 时已到 32 倍)。右图是块级跳过的实际工作量占比,绿虚线是理论上界 50.0%:块边长 256 时要多算 6.2 个百分点,块边长 64 时只多算 1.6 个——这就是「块越小越准、代价越碎」的量化版本。
三个脚本,全部纯 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 张图我把「融合」拆成三种可测量的形态,这是本篇的核心实验设计:
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
这三种形态的物理含义不同,必须分清:
_v2_tiled):一块读进来,K 段都在这一块上算完,只写回一次。tile 就是片上缓冲。这是 GPU 融合 kernel 在 CPU 上能做到的最好近似——numpy 做不到寄存器级融合,每个 ufunc 调用仍然要把结果写回 tile。实测结果(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 差着一个数量级。
$\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)$ 会给出完全不同的答案。
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/百万元素。

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

这张图要看什么:左图是单次调用耗时随张量规模的变化,红色虚线是固定开销地板,紫色竖线是「开销和数据各占一半」的临界规模(7746 个元素);右图是 8 段链里固定开销占总时间的比例。曲线在最左边几乎是水平的——那一段里,你增加 500 倍的数据量,耗时只涨 14%。
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 后,边际成本实际上比拟合值高。这是线性模型的能力边界,写在这里免得读者拿它当精确预测器。

这张图要看什么:左图是时间线甘特图,上排是 eager。蓝条是 CPU 发射,绿条是 GPU 执行,斜纹是 GPU 空转——空转的宽度正好等于发射和执行的时间差。下排是 CUDA Graph,一次发射之后 GPU 一路不停。右图是加速比随单 kernel 执行时间的变化,红线是发射成本 3 μs:红线左边 Graph 赚,红线右边 Graph 亏。
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 大小。
最小实现里的融合是「一串逐元素算子合成一个 kernel」。生产实现里最有价值的融合形态是 epilogue fusion:把接在 GEMM 后面的 bias、激活、缩放并进 GEMM 的 kernel 里,让 GEMM 的输出直接以最终形态写出,不经过显存。
用 3.2 节的账本算一下就明白为什么值钱:一个 $[M,N]$ 的 fp32 输出,不融合时要写 GEMM 输出($4MN$ 字节写)、bias 加法读+写($8MN$)、激活读+写($8MN$),合计 $20MN$;融合后 GEMM 一趟写 $4MN$,加上读输入的开销,量级降到 $1/5$。表达式没变一行,访存少了五分之四。
torch/cuda/graphs.py 里的 CUDAGraph(约 289 行)暴露出和 CUDA 一一对应的四个动作:capture_begin / capture_end / instantiate / replay:
https://github.com/pytorch/pytorch/blob/main/torch/cuda/graphs.py
三个工程细节值得单独说:
graph_pool_handle()(约 96 行)必须存在。graph 在录制时把显存地址写死在节点里,回放时不会重新分配。所以你必须让图里所有张量都来自一个固定的内存池。这不是 API 的怪癖,是「录制-回放」这个机制的必然代价。instantiate 要在第一次 replay 之前显式调用,否则第一次回放的延迟会变高(图要在那时才编译)。这个问题在实时推理里就是一次 P99 抖动。torch/_inductor/cudagraph_trees.py 里是 CUDAGraphNode(约 982 行)和 TreeManagerContainer(约 220 行),还有一个 CUDAWarmupNode。为什么是「树」而不是「一张图」:训练时反向图的形状依赖前向的实际形状,一批数据一个形状;而且每次回放都会产生新的输出张量,需要知道哪些显存可以复用。Inductor 用 mode="reduce-overhead" 打开这套机制。| 差异 | 来源 |
|---|---|
| 要判断「能不能融合」,甚至要算清代价 | 融合不是永远划算,见第六节 6.1 与 6.3 |
| 要处理动态 shape | tile 大小不能预知,必须切分或退化成不融合 |
| 要维护固定的显存池 | CUDA Graph 把地址写死了 |
| 要为每个形状/profile 各录一份图 | 图是静态的;变长输入需要分段或重录 |
| 要处理随机流、原地写、别名 | 融合会改变这些语义 |
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$ 那条极限路径。但代数合并只在可化简的模式上成立(本链是仿射复合,可以闭式求解),一般算子是合不掉的。它展示的是融合的第二层价值:让编译器看见全貌,从而有机会化简。
实测相对误差约 5e-7(fp32 舍入量级)。这不是 bug,是代数重排的必然结果。任何声称「融合不改变数值」的说法都需要限定条件。在需要严格复现的训练里,这会让「同一份代码换个后端跑出不同 loss」——通常无害,但你必须知道它从哪来。
这是本篇最反直觉的一条,也是附录 [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 矩阵」的原因——它同时省了访存和一部分计算。
cudaMalloc。一句话:CUDA Graph 是给「小 kernel 洪水」用的药,不是通用加速开关。 一个理想的判断顺序是:先用 profiler 确认 GPU 有空转;再确认单个 kernel 确实小于发射成本;然后才考虑上 Graph。而如果那些 kernel 本来就该被融合掉,融合是更根本的解法——融合让 kernel 数从 1093 降到 355,Graph 只是把这 1093 次提交合并成 1 次;两者正交,可以叠加。
launch_model.py 文件头。launch_model.py 的 [F1]/[F2]/[F3]/[F4] 是在上述输入上做的算术,脚本可跑、输入可改。请把这些当作「给定这些输入会得出什么」,而不是「GPU 上就是这么快」。四篇,按「融合是怎么从手工技巧变成系统能力,又怎么被反过来审视」串起来。
TVM: An Automated End-to-End Optimizing Compiler for Deep Learning(arXiv:1802.04799,2018)。第一次把「算符融合」从工程师的手工活变成编译器的自动决策:给定计算图,搜索融合方案和调度模板。本篇 3.2 节那个「融合省多少字节」的账,在这篇里是被当作搜索目标函数的一部分来算的。
The Deep Learning Compiler: A Comprehensive Survey(arXiv:2002.03794,2020,本文知识树的锚点)。给出一张完整坐标系:图级优化(融合、常量折叠、CSE)、算子级优化、内存分配、后端代码生成。它把融合明确归到「图级优化」——融合对单个算子的数学一无所知,它只改写算子之间的边界。这句话是理解本篇全部内容的前提。
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 节说的那件事——省的不只是访存,还有被掩码白算的一半。
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 一节。把它和融合并列讨论,是因为它经常被当成融合的替代品——而它其实解决另一个问题。
can_fuse_* 那一组函数之所以是「判断」而不是「尽量合」,就是因为融合有代价。而本文 [E] 节的实测给了一个更直接的极端例子:分块融合的 softmax 比不融合慢 21%。exp 一个都没少([E] 节:实测是访存下界的 5.70 倍,说明瓶颈在计算)。都是纯 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 区间内。
09 节用到的脚本全文如下(launch_model.py、fusion_lab.py、make_figures.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()
#!/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()
#!/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)