跳到正文
LLM 推理
S13解码加速·185

算子融合与 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 提交。
图的捕获与重放python
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 记录的是地址,不是数值

捕获时写下的是「运行 kernel A,从指针 0x7f… 读、写到 0x7f…」,回放时原样重新执行。因此输入每次都必须拷贝进同一批缓冲区;在被捕获区域内新分配的任何张量,都会让这份录制失效。下面所有约束都由此推出。

工作原理

工作原理

示意图从启动受限到执行受限
EAGER 模式 —— 每个 KERNEL 一次 LAUNCHCPUGPU虚线区域是气泡:GPU 已经算完,正在等 CPU 提交下一个 kernel。一个 32 层的模型每个 token 要 launch 约 352 个 kernel。每个 6 µs,光 CPU 时间就是 2.1 ms。第 1 步 —— 融合rms_norm+ 残差q/k/v proj一个融合 kernel3 次 launch → 1 次,中间结果根本不进 HBM第 2 步 —— 把整个步捕获成一张图捕获一次记录这张 DAG每一步重放一次 CPU 调用CPUGPU没有气泡。GPU 一个接一个连着跑。约束条件图记录的是固定的形状和固定的地址。而 batch size 每一步都在变,所以要按桶(bucket)捕获。引擎会为 batch size 1、2、4、8…… 各捕获一张图,再把真实 batch 补齐到下一个桶。prefill 的形状太多,没法捕获,所以图只是 decode 阶段的优化。
融合减少启动次数;graph 去掉剩下那些启动的每次 CPU 开销。约束那一栏说明了这项优化为什么只适用于 decode。

静态 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 最重要的场景,恰好是批处理帮助最小的场景:低延迟、低并发的服务。两项技术互相兜底。

动手观察

动手试试

模拟器启动开销 vs kernel 耗时
每步 kernel 数
288
每步耗时
1.73 ms
瓶颈在
CPU 启动
相对未优化
1.00×
前 18 次 kernel 启动
CPU
GPU
  • CPU:构造并提交一次启动
  • GPU:执行 kernel
  • 空闲

GPU 那一行有明显的空隙:它在 CPU 提交下一个 kernel 之前就已经算完了。加速器在等一个 Python 速度的循环。这时候堆更多 GPU 完全没用。

时间花在哪里
CPU 启动开销1.73 ms
GPU 执行0.89 ms

两者是重叠的,所以每步耗时取的是最大值,而不是相加。只有缩短更大的那一项才有意义,先 profile 再动手。

GPU 利用率

52%

batch 为 1、32 层、且没做任何优化时,通常受启动开销限制:约 352 个 kernel × 每个 6 µs = 每 token 2.1 ms 的纯 CPU 开销,可能完全超过 GPU 的计算时间。把两个开关都打开,看它如何反转。

先用 batch 1、两个开关都关:GPU 那一行明显有空隙,读数显示受 CPU 限制。打开融合,再打开 graph,看瓶颈翻转。然后把两个开关重新关掉,把 batch 拖到 128,空隙会自己合上,因为 kernel 已经长到足以把启动藏住。

亲手实现

实现

分桶逻辑很短。真正容易写错的是补齐那一部分。

code/s13_cuda_graphs.py(节选)python
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.

补齐行绝不能碰到真实显存

补齐 batch 中的占位行会跑和真实请求一模一样的 kernel,其中包括把注意力结果写进 KV 缓存。如果某个占位行的 block table 指向了真实 block,它就会污染某个真实请求的缓存。给补齐行分配专属的临时 block,并且永不复用。这个 bug 间歇出现、依赖 batch 大小,输出还看起来合情合理,几乎是最糟糕的组合。
在本地运行
为一个可配置模型建模启动开销与 kernel 执行时间的关系,展示 CPU 受限与 GPU 受限的交叉点落在哪里。还包含一个为融合收益记账的模型,以及一个分桶模拟器,报告每种桶布局下的补齐浪费。
$ python code/s13_cuda_graphs.py
预期输出: 一次定位启动受限区间的 batch 大小扫描、一张显示每层 kernel 从 11 降到 4 的融合表,以及一次桶布局对比:补齐浪费与被 graph 钉住的显存之间的权衡。

只需要 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. 1
    计算你的桶布局占用的 graph 显存,注意每张捕获的 graph 都会钉住自己的静态缓冲区。在固定显存预算下,找出补齐浪费最小的布局。
  2. 2
    在 S03 的实现里把 RMSNorm 与残差加法融合,确认输出不变。数一数消除了多少个中间张量。
  3. 3
    刻意实现一遍补齐污染 bug,观察输出只在某些 batch 大小下细微出错,然后写出能抓住它的断言。

继续学习

接下来

每一次前向现在已经尽可能便宜。剩下的一招是不再一个 token 跑一次前向。S14 讲投机解码:它打破「一次前向一个 token」的规则,同时完全不改变输出分布。

自测

习题

先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。

自测 1 题 / 共 6

CUDA Graph 记录的是什么?这为什么会限制你使用它的方式?

得分 0/6