Back to writing

/ What Is Series

What Is a Compiler Pass?

An engineering-oriented introduction to compiler passes: how they analyze and rewrite graphs and IR, and how fusion, lowering, pattern rewriting, and kernels fit together.

3 minCompiler · Pass · IR · Graph Optimization

在学习算子、编译器或推理框架时,经常会听到一句话:

这里写个 Pass 就行了。

这里的 Pass 指的是 编译器 Pass 或图优化 Pass,不是 Python 里的 pass 空语句。

一句话概括:

Pass 是对计算图、IR 或低层程序执行的一轮分析或变换。它读取一种程序表示,收集信息,或者把它改写成更适合后续优化和硬件执行的形式。

例如,模型里原本有三步计算:

x -> MatMul -> Add Bias -> GELU -> y

一个算子融合 Pass 可以识别这段子图,检查 shape、dtype、layout 和用户关系,然后把它改写成:

x -> FusedMatMulBiasGELU -> y

Pass 本身通常不计算业务 Tensor。它修改的是“程序如何被描述和执行”。

1. 为什么叫 Pass#

pass 的原意是“经过一遍”。程序从源代码到机器代码,通常不会一步完成,而是依次经过多轮处理:

模型代码
  -> 计算图 / 高级 IR
  -> 常量折叠 Pass
  -> 死代码消除 Pass
  -> 算子融合 Pass
  -> Layout 优化 Pass
  -> Lowering Pass
  -> 低级 IR
  -> 机器代码

每一轮只负责一个相对明确的问题。例如:

Pass 1:推导类型和 shape
Pass 2:删除无效计算
Pass 3:融合可组合的算子
Pass 4:选择或转换数据布局
Pass 5:把高级操作降低到目标硬件指令

按顺序组织起来的一组 Pass 叫做 Pass Pipeline;负责注册、排序和执行它们的组件通常叫 PassManager

不过,“经过一遍”不必机械地理解成每个 Pass 都完整线性扫描整份程序。具体实现可能遍历一个函数、一个 region、一类 operation,也可能借助 worklist 反复应用 rewrite,直到不能继续改写。Pass 更准确的含义是编译 Pipeline 中一个边界清晰的处理阶段。

2. Pass 操作的对象是什么#

Pass 并不只存在于某一个抽象层次。随着编译流程逐步接近硬件,它操作的对象也会变化。

2.1 计算图 Pass#

计算图 Pass 直接处理模型中的节点和边:

Conv -> BatchNorm -> ReLU

可以被改写为:

FusedConvBnReLU

这类 Pass 常见于 PyTorch FX、ONNX、TensorFlow Graph、TensorRT 以及各类大模型推理框架。它通常关心:

  • 算子类型和连接关系;
  • Tensor 的 shape、dtype 和 layout;
  • 一个节点是否有多个消费者;
  • 后端是否有可用的融合算子或 Kernel。

2.2 IR Pass#

IR 是 Intermediate Representation,即中间表示。它是编译器用于描述程序的内部语言。

例如下面的 PyTorch 表达式:

y = torch.relu(x @ w + b)

可能先变成较高级的 IR:

%0 = matmul %x, %w
%1 = add %0, %b
%2 = relu %1
return %2

融合 Pass 可以把它改写为:

%0 = fused_matmul_bias_relu %x, %w, %b
return %0

之后,Lowering Pass 再把高级操作逐步转换为 Triton IR、LLVM IR、GPU/MLU 相关 IR,最终生成目标设备可以执行的代码。

2.3 Kernel 级 Pass#

进入更低层以后,Pass 操作的不再是 AttentionLayerNorm 这类模型算子,而是循环、内存访问和指令。

例如朴素矩阵乘法:

for (int i = 0; i < M; ++i)
    for (int j = 0; j < N; ++j)
        for (int k = 0; k < K; ++k)
            C[i][j] += A[i][k] * B[k][j];

低层 Pass 可能依次完成:

循环分块
  -> 向量化加载
  -> 共享内存缓存
  -> 软件流水
  -> 循环展开
  -> 映射到 Tensor Core / MLU 指令

同一个“Pass”概念因此贯穿了从模型结构到机器指令的完整编译链路。

3. 两类最基本的 Pass#

按照是否修改程序,Pass 可以先分成两大类。

3.1 Analysis Pass:分析程序#

Analysis Pass 只收集信息,不直接改变程序语义或结构。例如:

  • 活跃变量和 Tensor 生命周期分析;
  • 数据依赖、支配关系和别名分析;
  • shape、dtype 和 layout 推导;
  • 常量传播分析;
  • 内存复用机会分析;
  • 融合收益和代价分析。

例如,生命周期分析得到:

Tensor A:从节点 1 活跃到节点 5
Tensor B:从节点 2 活跃到节点 3
Tensor C:从节点 6 活跃到节点 8

后续的内存规划 Pass 就可以利用这些结果,让生命周期不重叠的 Tensor 复用同一块内存。

分析结果通常会被缓存。某个 Transformation Pass 修改 IR 后,旧结果可能失效,因此成熟的 Pass 基础设施还要管理 analysis 的依赖、保留和失效关系。

3.2 Transformation Pass:改写程序#

Transformation Pass 会直接修改计算图或 IR。例如代数化简:

Add(x, 0) -> x
Mul(x, 1) -> x
Mul(x, 0) -> 0

或者算子融合:

MatMul -> Add -> GELU

改写成:

FusedMatMulBiasGELU

工程中说“写个 Pass”,大多数时候是在说这种变换型 Pass,但它往往仍然依赖 shape inference、alias analysis 等分析结果。

4. 常见的 Pass 在做什么#

4.1 常量折叠:在编译期完成计算#

原程序:

a = 2 + 3
y = x * a

优化后:

y = x * 5

只依赖常量的表达式无需留到运行时计算。

4.2 死代码消除:删除不可观察的计算#

原程序:

a = x + 1
b = x * 2
return b

a 不影响返回值,也没有其他副作用,可以删除:

b = x * 2
return b

4.3 公共子表达式消除:避免重复计算#

原程序:

a = x + y
b = x + y
z = a * b

优化后:

a = x + y
z = a * a

这个优化成立的前提是两次表达式的操作数和语义确实相同,并且没有破坏等价性的副作用。

4.4 算子融合:减少中间结果和调度开销#

原图:

MatMul -> BiasAdd -> GELU

优化后:

FusedMatMulBiasGELU

它可能减少 Kernel Launch、中间 Tensor 的显存读写和调度开销,并提高数据局部性。

4.5 Layout Transform:选择更合适的存储布局#

例如:

NCHW -> NHWC
连续布局 -> 硬件友好的分块布局

Layout Pass 会在单个算子的局部效率和整张图的转换开销之间做权衡。如果一次布局转换只能让一个小算子略微加速,额外的搬运成本可能反而更高。

4.6 Quantization:插入或融合量化计算#

量化 Pass 可能把:

FP16 MatMul

改写为带有 scale、zero point、quantize 和 dequantize 语义的 INT8 计算,也可能继续把这些节点融合进后端量化算子。

4.7 Lowering:逐步接近目标硬件#

Lowering 的核心是:

把抽象程度高的操作,转换成语义更具体、更接近目标硬件的操作。

例如一个高级 Attention 操作可以被分解为 QK MatMulSoftmaxPV MatMul,也可以在条件合适时直接降低为 FlashAttention 风格的融合 Kernel。Lowering 并不等于“拆开”:它的方向是降低抽象层级,最终形式既可能更细,也可能是一个后端专用融合操作。

5. Pass、Pattern Rewrite 和 Kernel 的关系#

这三个概念经常一起出现,但职责不同。

概念输入与输出主要职责
Operator / KernelTensor -> Tensor真正执行计算
IR程序结构描述要执行的计算
Pattern一段 IR 结构描述要寻找什么
Rewrite旧 IR -> 新 IR描述匹配后如何替换
PassIR -> 分析结果或新 IR在指定范围内组织和应用分析、Pattern 与 Rewrite

一个 Fusion Pass 里可以注册多条规则:

Conv + ReLU       -> FusedConvReLU
MatMul + GELU     -> FusedMatMulGELU
LayerNorm + Linear -> FusedLayerNormLinear

所以可以把它们的关系记成:

Pass
  -> 遍历 IR
  -> 用 Pattern 找候选结构
  -> 检查合法性和收益
  -> 用 Rewrite 修改 IR
  -> 交给后续 Pass 或 Backend

Pattern Rewrite 是实现 Pass 的常见机制,但不是所有 Pass 都必须基于 Pattern。数据流分析、循环优化和全局内存规划就可能使用完全不同的算法。

6. “写一个 Pass”具体在写什么#

一个典型的图变换 Pass 可以拆成四步:Match、Check、Rewrite 和 Verify。

6.1 Match:寻找候选模式#

先在图中寻找:

MatMul -> Add -> GELU

伪代码类似:

if node.op == "GELU":
    add = node.input
    if add.op == "Add":
        matmul = add.input[0]
        if matmul.op == "MatMul":
            matched = True

这一步只说明图的形状像目标模式,不代表它一定可以被替换。

6.2 Check:证明这次改写合法且值得做#

通常至少要检查:

dtype 是否被融合 Kernel 支持
shape 和广播语义是否兼容
stride 与 layout 是否满足要求
动态 shape 条件能否表达
中间节点是否还有其他消费者
操作是否有副作用或别名关系
目标硬件是否有对应 Kernel
融合后的成本是否真的更低

例如:

if matmul.dtype not in (torch.float16, torch.bfloat16):
    return
 
if len(add.users) != 1:
    return
 
if not bias.is_contiguous():
    return

很多 Pass 的 Bug 并不是 Pattern 写错,而是 Check 漏掉了某个语义条件。

6.3 Rewrite:创建新节点并维护 IR 合法性#

匹配和约束都通过后,创建融合节点:

fused = create_node(
    op="FusedMatMulBiasGELU",
    inputs=[x, weight, bias],
)

再替换旧输出:

replace_all_uses(gelu, fused)

最后删除已经没有用户的旧节点:

erase(gelu)
erase(add)
erase(matmul)

真实编译器还要同步维护 SSA use-def 关系、region、block、控制流、symbol table 等不变量。删除顺序错误或保留悬空引用,都可能让 IR 失去合法性。

6.4 Verify:验证结构和语义#

改写以后至少要做两类验证。

第一类是 IR 结构验证:

SSA use-def 是否合法
节点拓扑是否正确
shape 和 dtype 能否重新推导
是否存在悬空引用
操作约束是否满足

第二类是数值与行为验证:

torch.testing.assert_close(
    original_output,
    optimized_output,
    rtol=1e-3,
    atol=1e-3,
)

还应覆盖不同 shape、dtype、stride、动态维度、多消费者、空 Tensor、极值和数值溢出等边界情况。浮点优化不一定逐 bit 相同,因此容差应该根据算法和精度类型设定,而不是随意放宽。

7. 一个最小的 PyTorch FX Pass#

下面用 PyTorch FX 实现一个简单的代数化简:

x + 0 -> x
0 + x -> x
import operator
 
import torch
import torch.fx as fx
 
 
class Model(torch.nn.Module):
    def forward(self, x):
        y = x + 0
        return y * 2
 
 
def remove_add_zero_pass(
    graph_module: fx.GraphModule,
) -> fx.GraphModule:
    graph = graph_module.graph
 
    # list() creates a stable snapshot because the loop mutates the graph.
    for node in list(graph.nodes):
        if node.op != "call_function":
            continue
        if node.target not in (operator.add, torch.add):
            continue
 
        lhs, rhs = node.args
 
        if rhs == 0:
            node.replace_all_uses_with(lhs)
            graph.erase_node(node)
        elif lhs == 0:
            node.replace_all_uses_with(rhs)
            graph.erase_node(node)
 
    graph.lint()
    graph_module.recompile()
    return graph_module
 
 
model = Model()
traced = fx.symbolic_trace(model)
optimized = remove_add_zero_pass(traced)
 
print(optimized.graph)

优化前的 FX Graph 类似:

%x   = placeholder[target=x]
%add = call_function[target=operator.add](%x, 0)
%mul = call_function[target=operator.mul](%add, 2)
return %mul

优化后变成:

%x   = placeholder[target=x]
%mul = call_function[target=operator.mul](%x, 2)
return %mul

这里:

  • node.opnode.target 用于 Match;
  • rhs == 0lhs == 0 是 Check;
  • replace_all_uses_with()erase_node() 完成 Rewrite;
  • graph.lint() 检查图结构,recompile() 根据新图重新生成可执行代码。

这个示例故意保持简单。生产级实现还要考虑 overload、标量和 Tensor 常量的区别、类型提升、复数和量化语义等问题。

8. 为什么算子融合能加速#

假设原始计算是:

A = MatMul(X, W)
B = A + Bias
C = GELU(B)

不融合时可能需要:

Kernel 1:读取 X、W,写回 A
Kernel 2:读取 A、Bias,写回 B
Kernel 3:读取 B,写回 C

融合后则可能是:

Kernel 1:
  计算 MatMul
  在寄存器或片上存储中加 Bias
  继续计算 GELU
  直接写回 C

它的潜在收益包括:

  • 减少 Kernel Launch;
  • 减少中间 Tensor 分配;
  • 减少显存读写;
  • 提高数据局部性;
  • 降低 Runtime 调度开销。

但融合并不总是更快。它也可能带来:

  • 寄存器压力上升;
  • Occupancy 下降;
  • Kernel 体积和编译时间增加;
  • 动态 shape 支持变差;
  • 中间结果无法被其他分支复用;
  • 为了融合而引入额外布局转换。

因此成熟的 Fusion Pass 不只是“看到模式就替换”,还需要合法性判断、启发式规则、成本模型或性能数据库。

9. 最容易写错的地方#

9.1 只看节点类型,不看用户数量#

考虑下面的分叉:

                 -> GELU
MatMul -> Add ---+
                 -> OtherOp

如果把 MatMul + Add + GELU 融合后直接删除 AddOtherOp 的输入就被破坏了。Pass 必须检查用户关系,或者保留原节点供另一条分支使用。

9.2 忽略 dtype 和数值语义#

融合 Kernel 可能只支持 FP16 和 BF16,不能拿它替换 FP32、INT8 或复杂的混合精度计算。即使 dtype 一样,累加精度、舍入顺序和 NaN 行为也可能不同。

9.3 忽略 shape 和广播规则#

同样是:

x + bias

它既可能是:

[M, N] + [N]

也可能是:

[M, N] + [M, 1]

两者都满足 Add 语义,但一个融合 Kernel 未必同时支持这两种广播方式。

9.4 忽略 stride、layout 和 alias#

两个 Tensor 的 shape 相同,不代表底层存储相同:

shape = [1024, 4096], stride = [4096, 1]
shape = [1024, 4096], stride = [1, 1024]

它们的访存模式完全不同。此外,view、原地更新和共享存储还会引入 alias。一次看似局部的改写,可能改变程序其他位置能够观察到的数据。

9.5 忽略副作用和操作顺序#

纯函数式节点通常更容易重排。但随机数、I/O、原地写、设备同步和通信操作具有副作用,不能仅凭输入输出形状相同就随意删除或移动。

9.6 改完图不重新验证#

图改写后往往需要重新执行:

IR verify / graph lint
shape 和 dtype inference
analysis invalidation
dead code elimination
canonicalization
重新生成或编译可执行代码

“Pass 运行成功”只代表代码没有当场报错,不代表改写后的程序仍然正确。

10. Pass 在 AI Infra 的哪些地方出现#

10.1 PyTorch 2.x#

一条简化后的编译路径是:

PyTorch Model
  -> TorchDynamo 捕获图
  -> AOTAutograd
  -> FX Graph Pass
  -> TorchInductor
  -> Triton / C++ Kernel

其中会发生算子分解、Pattern Match、Fusion、Layout 选择、内存规划和 Kernel 生成。用户看到的是 torch.compile(),背后执行的是一整条 Pass Pipeline。

10.2 Triton#

Triton Kernel 大致会经过:

Python DSL
  -> TTIR
  -> TTGIR
  -> LLVM IR
  -> PTX / 目标代码

中间的 Pass 会完成 canonicalization、CSE、layout conversion、software pipelining、allocation 和 lowering 等工作。手写 Triton 代码决定了算法结构,而编译 Pass 会继续决定这段结构如何落到具体硬件上。

10.3 MLIR#

MLIR 把 Dialect 和 Pass 的边界表达得非常明确:

linalg.matmul
  -> Tiling Pass
分块循环
  -> Vectorization Pass
vector operations
  -> Lowering Pass
LLVM / GPU Dialect
  -> 目标代码

不同 Dialect 保留不同层次的语义,Pass 则负责在层次内部优化,或在不同 Dialect 之间转换。

10.4 TVM#

TVM 中常见:

with tvm.transform.PassContext(opt_level=3):
    module = relay.build(...)

PassContext 控制编译 Pipeline 的配置和优化级别。类型推导、常量折叠、公共子表达式消除、算子融合和布局改写等都可以作为 Pass 运行。

10.5 vLLM 和推理框架#

推理框架中的相关逻辑不一定都直接命名为 Pass,但思想一致。例如:

  • 把模型 Attention 替换为 PagedAttention 或设备自定义实现;
  • 把普通 Linear 替换为量化 Linear;
  • 改写 RoPE、MoE 或采样子图;
  • 划分可编译子图与 eager fallback;
  • 在 CUDA Graph capture 前稳定输入和执行结构;
  • 为 GPU、MLU 等不同设备选择对应 Custom Op。

这些工作都在回答同一个问题:如何把前端模型中的通用表达,可靠地改写成后端真正擅长执行的形式。

11. 一个算子如何真正接入框架#

学习算子时,可以把问题分成三层:

算法层:数学上算什么
Kernel 层:如何在 CUDA、Triton 或 MLU 上高效计算
Pass 层:如何让框架自动找到机会并调用这个 Kernel

例如你实现了一个高性能 Softmax Kernel,但模型图里可能表现为:

ReduceMax -> Sub -> Exp -> ReduceSum -> Div

要让框架自动使用新 Kernel,通常还需要:

识别 Softmax 子图
  -> 检查 axis、shape、dtype、stride 和数值语义
  -> 替换成 my_softmax
  -> 由 Backend Dispatch 到新 Kernel

因此,一个算子完整接入推理框架往往包含:

1. 算子语义定义
2. 前端注册或图捕获支持
3. Shape / dtype 推导
4. Kernel 实现
5. Backend Dispatch
6. Graph Pass / Rewrite
7. 正确性测试与 Benchmark

只写完 Kernel,解决的是“这个计算能不能高效执行”;补上 Pass,解决的才是“模型里的计算能不能自动落到这个 Kernel”。

12. 最核心的记忆方式#

可以记住四句话:

Operator / Kernel 负责计算。

IR 负责描述计算。

Pass 负责分析或修改 IR。

Backend 负责把 IR 变成目标硬件可执行的代码。

完整链路是:

模型代码
  -> 计算图 / IR
  -> 多个 Pass 分析和优化
  -> Lowering
  -> CUDA / Triton / MLU Kernel
  -> GPU / MLU 执行

所以,当别人说“这里写个 Pass 就行了”,更完整的意思通常是:

遍历指定范围内的计算图或 IR,找到目标算子或子图,证明改写在当前 shape、dtype、layout、用户关系和硬件条件下合法且有收益,然后替换成新的算子或 IR,并验证改写后的程序。

真正困难的部分往往不在“找到那几个节点”,而在证明这次替换对所有合法输入都正确,并且在目标硬件上确实更好。

参考资料#

  1. LLVM, Writing an LLVM New PM Pass: https://llvm.org/docs/WritingAnLLVMNewPMPass.html
  2. LLVM, LLVM's Analysis and Transform Passes: https://llvm.org/docs/Passes.html
  3. MLIR, Pass Infrastructure: https://mlir.llvm.org/docs/PassManagement/
  4. MLIR, Pattern Rewriting: https://mlir.llvm.org/docs/PatternRewriter/
  5. PyTorch, FX: https://docs.pytorch.org/docs/stable/fx.html
  6. Apache TVM, Pass Infrastructure: https://tvm.apache.org/docs/arch/pass_infra.html