📅 创建时间:2026-06-03 🏷️ 标签:#编译原理 #Pass #IR #SSA #优化遍 #传统编译器 📚 前置知识:[[00-compiler-overview]](AI 编译器全景) 📚 相关知识:[[02-ir-fundamentals]](IR 基础) [[03-graph-representation]](计算图表示)
┌─────────────────────────────────────────────────────────────────────────────┐
│ 📍 场景:运行 torch.compile(model)(input) 后报错 │
├─────────────────────────────────────────────────────────────────────────────┤
│ 报错信息: │
│ TorchScript error: │
│ 'Failed to compile subgraph with 12 operations. │
│ Unable to specialize type of tensor with dynamic shape.' │
│ │
│ 这是你第一次接触 AI 编译器。报错信息里的 │
│ "subgraph"、"specialize"、"dynamic shape" 都是什么意思? │
│ 编译器和 PyTorch 原生执行的本质区别是什么? │
└─────────────────────────────────────────────────────────────────────────────┘编译原理速通——面向 ML 工程师的核心概念 / Compiler Fundamentals for Machine Learning Engineers
第1节 传统编译器基础架构
1.1 编译器的三阶段架构
传统编译器(如 GCC、LLVM)遵循经典的三阶段架构:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ 传统编译器架构 (以 LLVM 为例) │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ 源代码 (Source Code) │ │
│ │ C / C++ / Rust │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ ↓ │
│ ═══════════════════════════════════ 前端 (Frontend) ══════════════════════ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ ┌──────────────┐ ┌──────────────┐ ┌──────────────┐ │ │
│ │ │ Lexer/ │───▶│ Parser/ │───▶│ Type │ │ │
│ │ │ Scanner │ │ Syntactic │ │ Checker │ │ │
│ │ │ (词法分析) │ │ (语法分析) │ │ (类型检查) │ │ │
│ │ └──────────────┘ └──────────────┘ └──────────────┘ │ │
│ │ │ │ │ │ │
│ │ ▼ ▼ ▼ │ │
│ │ 字符流 → Token流 Token流 → AST AST → Typed AST │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ ↓ │
│ ═══════════════════════════════════ 中间端 (Optimizer) ════════════════════ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ ┌────────────────────────────────────────────────────────────────┐ │ │
│ │ │ IR (中间表示) │ │ │
│ │ │ │ │ │
│ │ │ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐ │ │ │
│ │ │ │ Pass 1 │─▶│ Pass 2 │─▶│ Pass 3 │─▶│ Pass N │ │ │ │
│ │ │ │ 常量折叠 │ │ 代数化简 │ │ 死代码 │ │ 循环 │ │ │ │
│ │ │ │ │ │ │ │ 消除 │ │ 优化 │ │ │ │
│ │ │ └──────────┘ └──────────┘ └──────────┘ └──────────┘ │ │ │
│ │ │ │ │ │
│ │ └────────────────────────────────────────────────────────────────┘ │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ ↓ │
│ ═══════════════════════════════════ 后端 (Backend) ════════════════════════ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ ┌──────────────┐ ┌──────────────┐ ┌──────────────┐ │ │
│ │ │ Instruction │───▶│ Register │───▶│ Code │ │ │
│ │ │ Selection │ │ Allocation │ │ Emission │ │ │
│ │ │ (指令选择) │ │ (寄存器分配)│ │ (机器码生成)│ │ │
│ │ └──────────────┘ └──────────────┘ └──────────────┘ │ │
│ │ │ │ │ │ │
│ │ ▼ ▼ ▼ │ │
│ │ Selection DAG → Virtual Reg → Machine Code │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ 目标代码 (Object File / Executable) │ │
│ │ x86-64 / ARM / RISC-V │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘1.2 前端职责详解
词法分析 (Lexical Analysis):将源代码转换为 token 序列
# 示例:将代码分割成 token
source_code = "int result = a + b * 2;"
tokens = [
Token(type='keyword', value='int', line=1, col=1),
Token(type='identifier', value='result', line=1, col=5),
Token(type='operator', value='=', line=1, col=12),
Token(type='identifier', value='a', line=1, col=14),
Token(type='operator', value='+', line=1, col=16),
Token(type='identifier', value='b', line=1, col=18),
Token(type='operator', value='*', line=1, col=20),
Token(type='number', value='2', line=1, col=22),
Token(type='operator', value=';', line=1, col=23),
]语法分析 (Syntax Analysis):将 token 序列转换为抽象语法树 (AST)
源代码: int result = a + b * 2;
AST:
=
/ \
result +
/ \
a *
/ \
b 2语义分析 (Semantic Analysis):类型检查、作用域分析
# 类型检查示例
def type_check(node):
if node.type == 'binary_op':
left_type = type_check(node.left)
right_type = type_check(node.right)
if node.op in ['+', '-', '*', '/']:
if left_type == 'int' and right_type == 'int':
return 'int'
elif left_type == 'float' or right_type == 'float':
return 'float'
else:
raise TypeError(f"Cannot add {left_type} and {right_type}")1.3 优化器职责详解
优化器是编译器的核心,负责对 IR 进行各种优化:
# 优化类型示例
optimizations = {
# 1. 常量折叠 (Constant Folding)
# 在编译时计算出常量表达式的结果
"const_fold": {
"input": "x = 1 + 2 + 3",
"output": "x = 6", # 编译器直接算出结果
"impact": "减少 2 次运行时加法"
},
# 2. 代数化简 (Algebraic Simplification)
# 应用代数规则简化表达式
"algebraic_simplify": {
"input": "y = x * 1 + 0",
"output": "y = x", # x * 1 = x, x + 0 = x
"impact": "减少 2 次运行时运算"
},
# 3. 死代码消除 (Dead Code Elimination)
# 删除不会被使用的代码
"dce": {
"input": """
x = 5
y = x + 1 # 没人用 y
return x # 只用 x
""",
"output": "return 5",
"impact": "大幅减少代码量和执行时间"
},
# 4. 公共子表达式消除 (Common Subexpression Elimination)
# 避免重复计算相同的表达式
"cse": {
"input": "a = b + c; d = b + c + 1",
"output": "temp = b + c; a = temp; d = temp + 1",
"impact": "减少 1 次加法"
}
}1.4 后端职责详解
指令选择 (Instruction Selection):将 IR 映射到目标机器指令
IR: result = a + b
x86-64 选择:
mov eax, [a] ; 将 a 加载到 eax
add eax, [b] ; 加上 b
mov [result], eax ; 存储结果
ARM 选择:
ldr r0, [a] ; 将 a 加载到 r0
add r0, r0, [b] ; 加上 b
str r0, [result] ; 存储结果
RISC-V 选择:
ld a0, [a] ; 将 a 加载到 a0
add a0, a0, [b] ; 加上 b
sd a0, [result] ; 存储结果寄存器分配 (Register Allocation):将无限虚拟寄存器映射到有限物理寄存器
虚拟寄存器: v1, v2, v3, v4, v5, v6 ...
物理寄存器: rax, rbx, rcx, rdx, rsi, rdi, r8, r9, r10, r11 ...
分配策略:
- 图着色算法 (Graph Coloring)
- 线性扫描 (Linear Scan)
- 工作量证明 (Worklist)第2节 IR(中间表示)概念
2.1 为什么需要 IR
IR 是编译器的"通用语言",连接前后端:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ 为什么需要 IR │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ 问题: 如果没有 IR 会怎样? │
│ ═══════════════════════════════════════════════════ │
│ │
│ 源代码 → 机器码 │
│ │
│ 这意味着: │
│ • 每种语言 (C/C++/Rust/Go) 都需要为每种硬件 (x86/ARM/RISC-V) │
│ 写完整的编译器 │
│ • 这是一个 O(L×H) 的复杂度问题 │
│ (L = 语言数, H = 硬件数) │
│ │
│ 解决方案: IR 作为中间层 │
│ ═════════════════════════════════════════════════════ │
│ │
│ 源代码1 ─┐ │
│ 源代码2 ─┼──▶ IR ──▶ 机器码1 │
│ 源代码3 ─┤ ↑ 机器码2 │
│ ... ─┘ │ 机器码3 │
│ │ ... │
│ ┌─────┴─────┐ │
│ │ 语言前端 │ ← 只需为每种语言写一个前端 │
│ └───────────┘ │
│ │
│ 现在复杂度变成 O(L+H),大大减少! │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘2.2 IR 的特性
一个好的 IR 需要具备以下特性:
| 特性 | 说明 | 例子 |
|---|---|---|
| 足够抽象 | 脱离源语言细节 | 不关心是 C 还是 Rust 写的 |
| 足够具体 | 接近机器模型 | 有寄存器、内存、跳转概念 |
| 易于优化 | 便于进行分析和变换 | SSA 形式利于数据流分析 |
| 易于生成 | 前后端都容易生成和消费 | LLVM IR 是文本格式 |
2.3 LLVM IR 示例
; LLVM IR 示例:简单的函数
define i32 @add(i32 %a, i32 %b) {
entry:
%result = add i32 %a, %b ; a + b
ret i32 %result ; 返回结果
}
; 带有优化的 LLVM IR
define i32 @compute(i32 %x, i32 %y) {
entry:
%tmp1 = mul i32 %x, 2 ; x * 2
%tmp2 = add i32 %tmp1, %y ; (x * 2) + y
%tmp3 = mul i32 %tmp2, %x ; ((x * 2) + y) * x
ret i32 %tmp3
}
; 编译器会优化成: x * (2*x + y) 或 x * (2*x + y) = 2*x² + x*y第3节 Pass 遍(Pass)
3.1 什么是 Pass
Pass 是编译器中的基本优化单元。每个 Pass 对 IR 做一件事:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ Pass 工作方式 │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ IR 输入 │
│ │ │
│ ▼ │
│ ┌──────────┐ │
│ │ Pass 1 │ ──▶ 分析/变换 IR │
│ └──────────┘ │
│ │ │
│ ▼ │
│ ┌──────────┐ │
│ │ Pass 2 │ ──▶ 分析/变换 IR │
│ └──────────┘ │
│ │ │
│ ▼ │
│ ... │
│ │ │
│ ▼ │
│ IR 输出 │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘3.2 Pass 的两种类型
分析 Pass (Analysis Pass):收集 IR 信息,不修改 IR
| 分析 Pass | 收集的信息 |
|---|---|
| Dominator Tree | 控制流图的支配关系 |
| Loop Analysis | 循环结构 |
| Alias Analysis | 指针别名关系 |
| Data Flow Analysis | 数据流信息 |
| Type Analysis | 类型信息 |
# 分析 Pass 示例:数据流分析
def data_flow_analysis(ir):
"""
数据流分析追踪数据如何流经程序
"""
results = {
'reaching_definitions': {}, # 哪些定义能到达某个点
'available_expressions': [], # 哪些表达式已经计算过
'live_variables': [], # 哪些变量还被使用
}
for block in ir.blocks:
# 计算每个块的 IN/OUT 集合
block.in_set = gen(block) ∪ (block.out_set - kill(block))
block.out_set = union([pred.out_set for pred in block.predecessors])
return results变换 Pass (Transform Pass):修改 IR 以实现优化
| 变换 Pass | 效果 |
|---|---|
| Constant Folding | 常量折叠 |
| Loop Unrolling | 循环展开 |
| Function Inlining | 函数内联 |
| Loop Invariant Code Motion | 循环不变量外提 |
| Common Subexpression Elimination | 公共子表达式消除 |
| Dead Code Elimination | 死代码消除 |
# 变换 Pass 示例:常量折叠
def constant_folding(ir):
"""
将编译时可知的常量表达式折叠
"""
for instruction in ir.instructions:
if instruction.is_constant_expression():
# 在编译时计算结果
result = evaluate(instruction)
# 用常量替换原来的表达式
ir.replace(instruction, Constant(result))
print(f"Folded: {instruction} -> {result}")
# 示例
# 输入: %3 = add i32 1, 2
# 输出: %3 = i32 33.3 Pass 的执行顺序
Pass 之间的执行顺序很重要:
优化顺序示例(LLVM 的默认 pipeline):
════════════════════════════════════════════════════════════
1. Pass 依赖分析 (Pass dependency analysis)
└── 确定 pass 之间的依赖关系
2. 符号解析 (Symbol resolution)
└── 将符号链接到定义
3. 类型检查 (Type checking)
└── 验证类型安全
4. 基础优化 (Basic optimization)
├── 死代码删除
├── 常量折叠
└── 简单代数化简
5. 中级优化 (Mid-level optimization)
├── 内联
├── 循环不变代码外提
└── 重复GVN
6. 高级优化 (Advanced optimization)
├── 循环展开
├── 向量化
└── 尾部调用优化
7. 机器相关优化 (Machine-specific optimization)
├── 指令调度
└── 寄存器分配
8. 代码发射 (Code emission)
└── 生成目标机器码
════════════════════════════════════════════════════════════第4节 SSA(Static Single Assignment)
4.1 为什么需要 SSA
SSA 是现代编译器最重要的 IR 形式:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ SSA 形式 │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ 问题:同一个变量可能被多次赋值 │
│ ════════════════════════════════════════════════════════ │
│ │
│ 普通代码: │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ x = 5 # 第一次赋值 │ │
│ │ x = x + 1 # 第二次赋值 │ │
│ │ y = x * 2 # 使用 x │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
│ 问题:在 "y = x * 2" 处,x 的值是 5 还是 6? │
│ 编译器需要复杂的追踪才能知道 │
│ │
│ ──────────────────────────────────────────────────────────────────────────── │
│ │
│ SSA 形式: │
│ ══════════════════════════════════════════════════════ │
│ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ x1 = 5 # 第一次赋值 → x1 │ │
│ │ x2 = x1 + 1 # 第二次赋值 → x2 │ │
│ │ y1 = x2 * 2 # 使用 x2 │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
│ 优势:在 "y1 = x2 * 2" 处,x 的值明确是 x2 = 6 │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘4.2 SSA 的关键机制:Φ 函数
Φ 函数 (Phi Function) 用于合并不同控制流路径的值:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ Φ 函数 │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ 普通代码: │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ ┌─────────┐ │ │
│ │ │ start │ │ │
│ │ └────┬────┘ │ │
│ │ ┌────┴────┐ │ │
│ │ │ cond │ │ │
│ │ └────┬────┘ │ │
│ │ ┌────────────┼────────────┐ │ │
│ │ ▼ ▼ │ │
│ │ ┌─────────┐ ┌─────────┐ │ │
│ │ │ x=1 │ │ x=2 │ │ │
│ │ └────┬────┘ └────┬────┘ │ │
│ │ └────────────┬────────────┘ │ │
│ │ ┌────┴────┐ │ │
│ │ │ use │ │ │
│ │ │ x │ │ │
│ │ └─────────┘ │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
│ 在 "use x" 处,x 可能是 1 或 2(取决于条件分支) │
│ │
│ ──────────────────────────────────────────────────────────────────────────── │
│ │
│ SSA 形式: │
│ ══════════════════════════════════════════════════════ │
│ │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ ┌─────────┐ │ │
│ │ │ start │ │ │
│ │ └────┬────┘ │ │
│ │ ┌────┴────┐ │ │
│ │ │ cond │ │ │
│ │ └────┬────┘ │ │
│ │ ┌────────────┼────────────┐ │ │
│ │ ▼ ▼ │ │
│ │ ┌─────────┐ ┌─────────┐ │ │
│ │ │ x1 = 1 │ │ x2 = 2 │ │ │
│ │ └────┬────┘ └────┬────┘ │ │
│ │ └────────────┬────────────┘ │ │
│ │ ┌────┴────┐ │ │
│ │ │ x3 = φ │ ← Φ 函数:选择正确的版本 │ │
│ │ │ (x1, x2)│ │ │
│ │ └────┬────┘ │ │
│ │ ┌────┴────┐ │ │
│ │ │ use │ │ │
│ │ │ x3 │ │ │
│ │ └─────────┘ │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
│ Φ 函数在运行时根据控制流选择正确的值: │
│ • 如果从 x1 到达 → x3 = x1 │
│ • 如果从 x2 到达 → x3 = x2 │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘4.3 SSA 的优势
| 优势 | 说明 | 效果 |
|---|---|---|
| 简化数据流分析 | 每个变量只有一个定义点 | 分析时无需追踪多重赋值 |
| 精确的 Use-Def 链 | 明确的定义-使用关系 | 优化更容易、更安全 |
| 便于常量传播 | 常量赋值一目了然 | 折叠效果更好 |
| 便于死代码消除 | 未使用的定义清晰可见 | 消除更彻底 |
第5节 ML 编译器的特殊性
5.1 ML 编译器的独特挑战
ML 编译器与传统编译器有本质区别:
| 维度 | 传统编译器 | ML 编译器 |
|---|---|---|
| 计算模型 | 命令式,控制流复杂 | 数据流图,相对简单 |
| 操作类型 | 算术、逻辑、控制流 | 张量运算为主 |
| 优化目标 | 通用性能 | 吞吐量、显存、延迟 |
| 硬件 | CPU 为主 | GPU/TPU/NPU/FPGA |
| Shape | 静态 | 经常动态 |
| 精度 | 无特殊要求 | FP16/BF16/INT8/FP8 |
| 部署 | 编译一次,运行多次 | 需要优化权重 |
5.2 动态 Shape 问题
这是 ML 编译器最头疼的问题:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ 动态 Shape 问题 │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ 传统编译器: │
│ ════════════════ │
│ int arr[100]; // shape 固定,编译时可知 │
│ for (int i = 0; i < 100; i++) │
│ arr[i] = i; │
│ │
│ ML 编译器: │
│ ═════════════════ │
│ # Batch size 在运行时才知道 │
│ batch_size = get_batch_size() # 运行时获取 │
│ x = torch.randn(batch_size, 512) │
│ y = model(x) │
│ │
│ # Sequence length 在运行时才知道 │
│ seq_len = max(len(sentence) for sentence in batch) │
│ │
│ # 甚至维度数量都可能变化 │
│ ragged_tensor = ragged_from_padding(padded_tensor) │
│ │
│ 编译器的困境: │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ 为了最大化优化,编译器需要知道: │ │
│ │ • Tensor 的精确 shape │ │
│ │ • 内存布局 │ │
│ │ • 循环次数 │ │
│ │ │ │
│ │ 但 ML 场景这些经常是动态的! │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
│ 解决方案: │
│ ┌─────────────────────────────────────────────────────────────────────────┐ │
│ │ 1. 特殊化 (Specialization): │ │
│ │ 为常见 shape 编译专用 kernel │ │
│ │ │ │
│ │ 2. 动态调度 (Dynamic Dispatch): │ │
│ │ 运行时检查 shape,选择对应 kernel │ │
│ │ │ │
│ │ 3. Polytope 优化 (Polytope Optimization): │ │
│ │ 用数学方法处理任意 shape 的循环 │ │
│ │ │ │
│ │ 4. Symbolic Shape: │ │
│ │ 用符号表示 shape,参与优化决策 │ │
│ └─────────────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘5.3 Python 语义问题
PyTorch 的动态图特性带来额外的编译挑战:
# 问题 1: Python 控制流
def forward(x, condition):
if condition: # 运行时才知道条件
return model_a(x)
else:
return model_b(x)
# 编译器需要处理两种可能的执行路径
# 问题 2: Python 动态特性
def forward(x):
# hasattr, getattr, setattr 都在运行时决定
if hasattr(x, 'special_attr'):
x = x.special_attr
# exec, eval 更是无法静态分析
return torch.sum(x)
# 问题 3: Python 数据结构
def forward(data):
# 字典、元组、列表的嵌套
result = {'features': [], 'labels': []}
for item in data:
result['features'].append(extract(item))
return torch.stack(result['features'])5.4 数据依赖与别名分析
ML 中的数据依赖更复杂:
# 正常情况
x = torch.randn(10, 10)
y = x + 1 # y 依赖 x
# in-place 情况
x.add_(1) # x 自己被修改
# view/shallow copy
y = x.view(2, 5) # y 和 x 共享底层数据
# 编译器必须追踪这些关系,否则可能产生错误的结果第6节 图编译 vs 程序编译
6.1 ML 模型的本质是数据流图
深度学习模型本质上是计算图:
┌─────────────────────────────────────────────────────────────────────────────────┐
│ 计算图 vs 传统程序 │
├─────────────────────────────────────────────────────────────────────────────────┤
│ │
│ 传统程序: │
│ ════════════ │
│ def compute(a, b): │
│ x = a + b │
│ if x > 0: # 复杂控制流 │
│ y = foo(x) │
│ else: │
│ y = bar(x) # 条件分支 │
│ return y │
│ │
│ 特点: │
│ • 顺序执行 │
│ • 条件分支 │
│ • 循环 │
│ • 递归 │
│ │
│ ──────────────────────────────────────────────────────────────────────────── │
│ │
│ ML 计算图: │
│ ═════════════ │
│ x1 = input │
│ x2 = matmul(x1, W1) # 无条件执行 │
│ x3 = relu(x2) # 纯函数式,无副作用 │
│ x4 = matmul(x3, W2) # 前向无环 │
│ output = softmax(x4) │
│ │
│ 特点: │
│ • 主要是无副作用的数据流 │
│ • 控制流相对简单 │
│ • 可静态分析 │
│ • 适合并行执行 │
│ │
│ input │
│ │ │
│ ▼ │
│ ┌───────┐ │
│ │ MatMul │ │
│ └───────┘ │
│ │ │
│ ▼ │
│ ┌───────┐ │
│ │ ReLU │ │
│ └───────┘ │
│ │ │
│ ▼ │
│ ┌───────┐ │
│ │ MatMul │ │
│ └───────┘ │
│ │ │
│ ▼ │
│ ┌────────┐ │
│ │ Softmax│ │
│ └────────┘ │
│ │ │
│ ▼ │
│ output │
│ │
└─────────────────────────────────────────────────────────────────────────────────┘6.2 图编译的优势
| 优势 | 说明 |
|---|---|
| 全局优化 | 可以看到整个计算图,进行全局优化 |
| 算子融合 | 将相邻算子融合,减少 kernel 启动开销 |
| 内存规划 | 预先规划显存使用,实现 reuse |
| 并行调度 | 识别可并行的节点,最大化并行度 |
| 代码生成 | 为整个图生成高度优化的代码 |
第7节 三种编译策略
7.1 JIT(即时编译)
# PyTorch JIT 示例
import torch
class MyModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(10, 10)
def forward(self, x):
return torch.relu(self.linear(x))
model = MyModule()
# TorchScript JIT 编译
scripted = torch.jit.script(model) # 静态分析 + JIT
output = scripted(torch.randn(1, 10))
# torch.compile JIT (PyTorch 2.0)
compiled = torch.compile(model, mode="default")
output = compiled(torch.randn(1, 10))JIT 特点:
| 优点 | 缺点 |
|---|---|
| 无需重新编译 | 首次运行有编译开销 |
| 可利用运行时信息 | 不可预测的性能 |
| 支持动态 shape | 部署复杂 |
7.2 AOT(离线编译)
# AOT 编译示例
import torch
model = MyModel()
# 导出为静态图
exported = torch.export(model, args=(torch.randn(1, 10),))
# 保存
exported.save("model.pt2")
# 部署时加载
loaded = torch.export.load("model.pt2")AOT 特点:
| 优点 | 缺点 |
|---|---|
| 部署简单 | 需要指定输入 shape |
| 确定性性能 | 无法利用运行时信息 |
| 易于优化 | 重新编译耗时 |
7.3 Lazy Compilation(延迟编译)
# Lazy Compilation 策略
# 编译发生在真正需要时
@torch.jit.script
def lazy_func(x):
# 只有实际使用到的分支会被编译
if x.sum() > 0:
return x * 2
else:
return x / 2第8节 关键术语速查表
| 术语 | 英文全称 | 解释 | ML 场景例子 |
|---|---|---|---|
| IR | Intermediate Representation | 中间表示 | XLA HLO, TVM Relay, FX Graph |
| Pass | Optimization Pass | 优化遍 | 常量折叠 Pass, 融合 Pass |
| SSA | Static Single Assignment | 静态单赋值 | 每个变量只赋值一次 |
| Φ | Phi Function | φ 函数 | 合并不同分支的值 |
| AST | Abstract Syntax Tree | 抽象语法树 | 编译器前端输出 |
| CFG | Control Flow Graph | 控制流图 | 函数内控制流 |
| DAG | Directed Acyclic Graph | 有向无环图 | 计算图通常是无环的 |
| LICM | Loop Invariant Code Motion | 循环不变量外提 | 将不变计算移出循环 |
| DCE | Dead Code Elimination | 死代码消除 | 删除无用的代码 |
| CSE | Common Subexpression Elimination | 公共子表达式消除 | 避免重复计算 |
| CGVN | Conditional Global Value Numbering | 条件全局值编号 | 跨分支识别等价表达式 |
| Inline | Function Inlining | 函数内联 | 将函数调用展开为函数体 |
| Unroll | Loop Unrolling | 循环展开 | 减少循环控制开销 |
| Vectorize | Vectorization | 向量化 | SIMD 并行 |
| Fusion | Kernel Fusion | 核融合 | 合并多个 kernel |
| Lowering | Lowering | 降低 | 高层 IR → 低层 IR |
| Frontend | Compiler Frontend | 编译器前端 | 源码 → IR |
| Backend | Compiler Backend | 编译器后端 | IR → 机器码 |
| Opcode | Operation Code | 操作码 | add, mul, relu 等 |
| Operand | Operand | 操作数 | tensor, 参数 |
升华
┌─────────────────────────────────────────────────────────────────────────────┐
│ 编译原理核心要点 │
├─────────────────────────────────────────────────────────────────────────────┤
│ │
│ 1. 编译器 = 前端 + 中端 + 后端 │
│ ┌──────────────────────────────────────────────────────────────────┐ │
│ │ 前端:理解源语言 (词法/语法/语义分析) │ │
│ │ 中端:IR 优化 (Pass 遍,对 IR 进行分析和变换) │ │
│ │ 后端:生成代码 (指令选择/寄存器分配/机器码生成) │ │
│ └──────────────────────────────────────────────────────────────────┘ │
│ │
│ 2. IR 是编译器的灵魂 │
│ ┌──────────────────────────────────────────────────────────────────┐ │
│ │ IR 连接前后端,使编译器可复用 │ │
│ │ SSA 形式使分析和优化更简单、更精确 │ │
│ │ 多级 IR 平衡抽象能力和优化粒度 │ │
│ └──────────────────────────────────────────────────────────────────┘ │
│ │
│ 3. ML 编译器面临的独特挑战 │
│ ┌──────────────────────────────────────────────────────────────────┐ │
│ │ 动态 Shape、Python 语义、数据依赖复杂 │ │
│ │ 但 ML 模型本质是数据流图,适合图编译 │ │
│ │ 编译器通过特殊化/动态调度/Symbolic Shape 应对 │ │
│ └──────────────────────────────────────────────────────────────────┘ │
│ │
│ 4. 三种编译策略各有优劣 │
│ ┌──────────────────────────────────────────────────────────────────┐ │
│ │ JIT: 灵活但有首次开销 │ │
│ │ AOT: 确定性高但需指定 shape │ │
│ │ Lazy: 平衡前两者 │ │
│ └──────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────────┘
核心一句话: 编译器是将高层抽象逐步降低到机器可执行指令的艺术,
IR 是这个艺术的核心载体,Pass 是实现优化的基本手段。"AI 可查 vs 必须理解"清单
必须理解(不理解就等于不会):
- 🔴 编译器三阶段架构:Frontend (词法/语法/语义) → Optimizer (Pass 优化) → Backend (代码生成)
- 🔴 IR 的必要性:连接前后端,使多语言多硬件编译成为可能
- 🔴 SSA 形式:每个变量只赋值一次,简化数据流分析
- 🔴 Pass 分类:分析 Pass 收集信息,变换 Pass 修改 IR
- 🔴 ML 编译器的独特挑战:动态 Shape、Python 语义、数据依赖
- 🔴 图编译 vs 程序编译:ML 模型是数据流图,适合图级优化
- 🔴 三种编译策略:JIT(灵活但有开销)、AOT(确定但需指定 shape)、Lazy
AI 可查(知道去哪查就行):
- ✅ LLVM IR 的具体语法和 intrinsics
- ✅ 某个具体 Pass 的实现算法(如 Graph Coloring 寄存器分配)
- ✅ SSA 构造的具体算法(插入 φ 函数的位置)
- ✅ Polytope 模型的具体数学推导
- ✅ 某个特定硬件的指令选择规则
- ✅ 编译器优化的理论证明
学习状态:🟡 开始学习