算子融合与 CUDA Graph
小 batch 时 GPU 是在等 CPU。融合算子、回放已捕获的 graph,能省掉每个 token 数百次的 kernel 启动。
- 启动开销
- 算子融合
- graph 捕获
- 静态 shape
- 分桶补齐
为什么重要
问题
在现代 GPU 上剖析一次 batch 为 1 的 decode,你会看到荒唐的一幕:GPU 有将近一半的时间在闲着。原因不是显存带宽,而是 CPU 派发工作的速度跟不上。
一个 32 层的模型每层大约跑十一个 kernel:两个 norm、三个投影、RoPE、注意力、输出投影,再加 FFN 的三个。算下来每个 token 约 350 次 kernel 启动,每次启动的构造与提交要占用 CPU 5–10 µs。每个 token 两毫秒的 CPU 工作,而每个 kernel 在 GPU 上只跑几微秒。
你有一台每秒能做一千万亿次运算的机器,正在等一个 Python 循环。
核心思路
解法
两种彼此独立的修法,通常一起用。
- 算子融合 —— 把相邻操作合并成一个 kernel。启动次数更少,中间值留在寄存器里而不必往返 HBM。RMSNorm 与残差加法融合;Q、K、V 打包成一次 matmul;RoPE 折进注意力的前导;SwiGLU 的门控、乘法和激活一趟做完。
- CUDA Graph —— 把整串启动序列录制一次,之后用一次 CPU 调用回放整张录好的图。GPU 驱动已经知道依赖结构,因此不再需要逐 kernel 提交。
class GraphRunner:
def __init__(self, model, batch_size):
# Static buffers: the graph records ADDRESSES, not values.
self.input_ids = torch.zeros(batch_size, dtype=torch.long, device="cuda")
self.positions = torch.zeros(batch_size, dtype=torch.long, device="cuda")
model(self.input_ids, self.positions) # warm up, then capture
self.graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(self.graph):
self.output = model(self.input_ids, self.positions)
def run(self, input_ids, positions):
# Copy into the SAME buffers the graph was captured against.
self.input_ids[: len(input_ids)].copy_(input_ids)
self.positions[: len(positions)].copy_(positions)
self.graph.replay() # ONE cpu call
return self.output.clone()graph 记录的是地址,不是数值
工作原理
工作原理
静态 shape,以及分桶这个变通办法
一张捕获好的 graph 只对捕获时的那组 shape 有效。而连续批处理每一步都在改变 batch 大小;让它变化,正是连续批处理的意义所在。
标准解法是分桶:为 batch 大小 1、2、4、8、16、24、32… 各捕获一张 graph,运行时把真实 batch 补齐到下一个已捕获的尺寸。补齐会浪费一点算力;在小 batch 下,省下的启动开销远比它值钱。
prefill 则完全不用 graph。它的序列长度实际上无界,不存在一小组可枚举的桶。这没有代价,因为 prefill 受算力限制,启动开销本来就被算术藏住了。
什么时候它不再重要
启动开销是每步的固定成本,因此随着每步 GPU 工作量上升,它的重要性下降。batch 为 1 时它可能占步时的大头;batch 为 256 时 kernel 已经足够长,启动完全被藏在后面。
于是 CUDA Graph 最重要的场景,恰好是批处理帮助最小的场景:低延迟、低并发的服务。两项技术互相兜底。
动手观察
动手试试
- 每步 kernel 数
- 288
- 每步耗时
- 1.73 ms
- 瓶颈在
- CPU 启动
- 相对未优化
- 1.00×
- CPU:构造并提交一次启动
- GPU:执行 kernel
- 空闲
GPU 那一行有明显的空隙:它在 CPU 提交下一个 kernel 之前就已经算完了。加速器在等一个 Python 速度的循环。这时候堆更多 GPU 完全没用。
两者是重叠的,所以每步耗时取的是最大值,而不是相加。只有缩短更大的那一项才有意义,先 profile 再动手。
52%
batch 为 1、32 层、且没做任何优化时,通常受启动开销限制:约 352 个 kernel × 每个 6 µs = 每 token 2.1 ms 的纯 CPU 开销,可能完全超过 GPU 的计算时间。把两个开关都打开,看它如何反转。
先用 batch 1、两个开关都关:GPU 那一行明显有空隙,读数显示受 CPU 限制。打开融合,再打开 graph,看瓶颈翻转。然后把两个开关重新关掉,把 batch 拖到 128,空隙会自己合上,因为 kernel 已经长到足以把启动藏住。
亲手实现
实现
分桶逻辑很短。真正容易写错的是补齐那一部分。
BUCKETS = [1, 2, 4, 8, 16, 24, 32, 48, 64, 96, 128, 192, 256]
class BucketedRunner:
def __init__(self, model):
self.graphs = {b: capture(model, b) for b in BUCKETS}
def forward(self, batch):
n = len(batch)
bucket = next(b for b in BUCKETS if b >= n)
padded = pad_to(batch, bucket) # real requests + dummies
out = self.graphs[bucket].run(padded)
return out[:n] # discard the dummy rows
# The dummies MUST be harmless: point them at a scratch KV block,
# or they will write garbage into a real request's cache.补齐行绝不能碰到真实显存
$ python code/s13_cuda_graphs.py只需要 NumPy — 查看环境准备.
生产实践
在生产环境中
- vLLM —— 按 batch 桶捕获 decode graph;
enforce_eager=True可以关掉它。排查 shape 或显存 bug 时,先试这一招。 - TensorRT-LLM —— 把这条路走到极致:整个模型被提前编译成一张固定的图。它的速度来自于此,改 shape 就要重新构建也来自于此。
- torch.compile —— 通过 Inductor 自动做融合,并可以捕获 graph(
mode="reduce-overhead"),按 shape 重新编译。 - llama.cpp / picoLM —— CPU 上没有启动开销,但融合依然重要,只是理由不同:融合的 dequant+dot 一趟做完,把内存流量砍了一半。
练习
- 1计算你的桶布局占用的 graph 显存,注意每张捕获的 graph 都会钉住自己的静态缓冲区。在固定显存预算下,找出补齐浪费最小的布局。
- 2在 S03 的实现里把 RMSNorm 与残差加法融合,确认输出不变。数一数消除了多少个中间张量。
- 3刻意实现一遍补齐污染 bug,观察输出只在某些 batch 大小下细微出错,然后写出能抓住它的断言。
继续学习
接下来
每一次前向现在已经尽可能便宜。剩下的一招是不再一个 token 跑一次前向。S14 讲投机解码:它打破「一次前向一个 token」的规则,同时完全不改变输出分布。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
CUDA Graph 记录的是什么?这为什么会限制你使用它的方式?