/ 什么是系列
什么是编译器 Pass?
从计算图和 IR 出发,理解编译器 Pass 如何分析、匹配并改写程序,以及算子融合、Lowering、Pattern Rewrite 和 Kernel 之间的关系。
在学习算子、编译器或推理框架时,经常会听到一句话:
这里写个 Pass 就行了。这里的 Pass 指的是 编译器 Pass 或图优化 Pass,不是 Python 里的 pass 空语句。
一句话概括:
Pass 是对计算图、IR 或低层程序执行的一轮分析或变换。它读取一种程序表示,收集信息,或者把它改写成更适合后续优化和硬件执行的形式。
例如,模型里原本有三步计算:
x -> MatMul -> Add Bias -> GELU -> y一个算子融合 Pass 可以识别这段子图,检查 shape、dtype、layout 和用户关系,然后把它改写成:
x -> FusedMatMulBiasGELU -> yPass 本身通常不计算业务 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 操作的不再是 Attention、LayerNorm 这类模型算子,而是循环、内存访问和指令。
例如朴素矩阵乘法:
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 ba 不影响返回值,也没有其他副作用,可以删除:
b = x * 2
return b4.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 MatMul、Softmax 和 PV MatMul,也可以在条件合适时直接降低为 FlashAttention 风格的融合 Kernel。Lowering 并不等于“拆开”:它的方向是降低抽象层级,最终形式既可能更细,也可能是一个后端专用融合操作。
5. Pass、Pattern Rewrite 和 Kernel 的关系#
这三个概念经常一起出现,但职责不同。
| 概念 | 输入与输出 | 主要职责 |
|---|---|---|
| Operator / Kernel | Tensor -> Tensor | 真正执行计算 |
| IR | 程序结构 | 描述要执行的计算 |
| Pattern | 一段 IR 结构 | 描述要寻找什么 |
| Rewrite | 旧 IR -> 新 IR | 描述匹配后如何替换 |
| Pass | IR -> 分析结果或新 IR | 在指定范围内组织和应用分析、Pattern 与 Rewrite |
一个 Fusion Pass 里可以注册多条规则:
Conv + ReLU -> FusedConvReLU
MatMul + GELU -> FusedMatMulGELU
LayerNorm + Linear -> FusedLayerNormLinear所以可以把它们的关系记成:
Pass
-> 遍历 IR
-> 用 Pattern 找候选结构
-> 检查合法性和收益
-> 用 Rewrite 修改 IR
-> 交给后续 Pass 或 BackendPattern 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 -> ximport 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.op和node.target用于 Match;rhs == 0与lhs == 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 融合后直接删除 Add,OtherOp 的输入就被破坏了。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,并验证改写后的程序。
真正困难的部分往往不在“找到那几个节点”,而在证明这次替换对所有合法输入都正确,并且在目标硬件上确实更好。
参考资料#
- LLVM, Writing an LLVM New PM Pass: https://llvm.org/docs/WritingAnLLVMNewPMPass.html
- LLVM, LLVM's Analysis and Transform Passes: https://llvm.org/docs/Passes.html
- MLIR, Pass Infrastructure: https://mlir.llvm.org/docs/PassManagement/
- MLIR, Pattern Rewriting: https://mlir.llvm.org/docs/PatternRewriter/
- PyTorch, FX: https://docs.pytorch.org/docs/stable/fx.html
- Apache TVM, Pass Infrastructure: https://tvm.apache.org/docs/arch/pass_infra.html