Skip to content
Gains Summary
Main Navigation 首页 / Home
C++ 编程 / C++ Programming
系统与高性能 / Systems & Performance
Web 开发 / Web Development
人工智能 / Artificial Intelligence
工业软件 / Industrial Software
其他内容 / Other Topics
C++ 编程 / C++系统与性能 / SystemsWeb 开发 / Web人工智能 / AI工业软件 / Industrial

外观

Sidebar Navigation

← 人工智能 / Artificial Intelligence

AI 编译器 / AI Compilers

1. AI 编译器全景——为什么模型需要编译器 / The AI Compiler Landscape and Why Models Need Compilers

2. 编译原理速通——面向 ML 工程师的核心概念 / Compiler Fundamentals for Machine Learning Engineers

3. 中间表示基础——理解 IR 层级与 lowering 链路 / Intermediate Representation Levels and Lowering Pipelines

4. 计算图的构建与表示 / Building and Representing Computational Graphs

5. MLIR 架构、方言与渐进式降级 / MLIR Architecture, Dialects, and Progressive Lowering

6. 算子语义、广播、归约与形状推导 / Operator Semantics, Broadcasting, Reduction, and Shape Inference

7. 模型前端格式:ONNX、TFLite、HLO 与 SavedModel / Model Frontend Formats: ONNX, TFLite, HLO, and SavedModel

8. 图优化 Pass——经典优化在 ML 中的应用 / Graph Optimization Passes for Machine Learning

9. 算子融合——编译器最重要的性能优化 / Operator Fusion as a Core Compiler Optimization

10. 内存规划——Buffer 分配与生命周期管理 / Memory Planning, Buffer Allocation, and Lifetime Management

11. Layout 优化——数据排布转换与内存效率 / Layout Optimization for Data Movement and Memory Efficiency

12. 动态 Shape——符号分析与形状处理 / Dynamic Shapes, Symbolic Analysis, and Shape Processing

13. 硬件约束下的操作调度 / Operation Scheduling Under Hardware Constraints

14. 从模板、DSL 到 IR 降级的代码生成架构 / Code Generation Architectures from Templates and DSLs to IR Lowering

15. CPU 后端:SIMD、分块与多线程 / CPU Backends with SIMD, Tiling, and Multithreading

16. CUDA 后端:合并访存与 Tensor Core / CUDA Backends, Memory Coalescing, and Tensor Cores

17. NPU 后端:脉动阵列与端侧 AI 生态 / NPU Backends, Systolic Arrays, and Edge AI Ecosystems

18. Kernel 性能基础:Roofline 与 Occupancy / Kernel Performance Fundamentals with Roofline and Occupancy

19. CUTLASS 与分层 GEMM 模板 / CUTLASS and Hierarchical GEMM Templates

20. TVM Tensor Expression 与计算调度分离 / TVM Tensor Expressions and Compute-Schedule Separation

21. 使用 Triton 编写高性能 GPU Kernel / Triton for High-Performance GPU Kernels in Python

22. 基于成本模型与实测搜索的自动调度 / Automatic Scheduling with Cost Models and Measurement-Based Search

23. XLA 内部机制:HLO、融合与 SPMD / XLA Internals, HLO, Fusion, and SPMD

24. Torch-MLIR:从 PyTorch 算子到 MLIR 方言 / Torch-MLIR from PyTorch Operators to MLIR Dialects

25. torch.compile:Dynamo、AOTAutograd、Inductor 与 Triton / Torch Compile with Dynamo, AOTAutograd, Inductor, and Triton

26. 从 MLIR 经 LLVM 降级到机器码 / Lowering from MLIR Through LLVM to Machine Code

27. 量化——低精度推理的工程实践 / Engineering Low-Precision Inference with Quantization

28. 分布式编译与训练——多设备编排的编译器支持 / Compiler Support for Distributed Training and Multi-Device Orchestration

29. 生产调试——真实问题的编译器视角排查 / Production Debugging from the Compiler Perspective

30. 未来方向——AI 编译器的新挑战与机遇 / Future Challenges and Opportunities for AI Compilers

本页目录

📅 创建时间:2026-06-03 🏷️ 标签:#IR #多级IR #lowering #HLO #Linalg #LLVM-IR #PTX 📚 前置知识:[[01-compiler-primer]](编译原理速通) 📚 相关知识:[[03-graph-representation]](计算图表示) [[04-mlir-architecture]](MLIR 架构)


┌─────────────────────────────────────────────────────────────────────────────┐
│  📍 场景:同一个 ResNet-50 模型在不同框架里格式完全不同                      │
├─────────────────────────────────────────────────────────────────────────────┤
│  你看到的现状:                                                             │
│  • PyTorch: TorchScript, FX Graph, ExportedProgram                         │
│  • TensorFlow: SavedModel, GraphDef, XLA HLO                              │
│  • ONNX: Protocol Buffer with opset version                                │
│  • TFLite: FlatBuffer with flex buffer                                    │
│  • TVM: Relay IR, TE, TIR                                                │
│                                                                             │
│  你的困惑:                                                                 │
│  这些格式之间能互转吗?哪个是"底层"的?                                      │
│  为什么需要这么多不同的 IR?它们之间的关系是什么?                           │
└─────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14

中间表示基础——理解 IR 层级与 lowering 链路 / Intermediate Representation Levels and Lowering Pipelines ​

第1节 IR 的抽象层级 ​

1.1 为什么需要多级 IR ​

IR 不是单一的东西,而是多个抽象层级:

┌─────────────────────────────────────────────────────────────────────────────────┐
│                           IR 的抽象层级金字塔                                   │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│                              ▲                                                  │
│                             ╱ ╲                                                 │
│                            ╱   ╲                                                │
│                           ╱  L4 ╲        ← 最高层:高层语义 IR                   │
│                          ╱IR(L4)╲           (最接近用户语义)                    │
│                         ╱─────────╲                                                │
│                        ╱    L3     ╲       ← 操作级 IR                         │
│                       ╱  IR(L3)   ╲          (细粒度算子表示)                  │
│                      ╱─────────────╲                                             │
│                     ╱      L2       ╲      ← 指令级 IR                        │
│                    ╱   IR(L2)       ╲         (接近硬件)                       │
│                   ╱───────────────────╲                                         │
│                  ╱         L1          ╲     ← 机器码级 IR                    │
│                 ╱      IR(L1)          ╲        (接近机器指令)                  │
│                ╱─────────────────────────╲                                      │
│               ╱           L0             ╲    ← 硬件描述                       │
│              ╱         IR(L0)             ╲      (硬件行为模型)                 │
│             ╱────────────────────────────────╲                                  │
│                                                                                  │
│  ┌─────────────────────────────────────────────────────────────────────────┐  │
│  │  层级越高 → 抽象程度越高 → 语义丰富 → 优化空间大                          │  │
│  │  层级越低 → 抽象程度越低 → 接近硬件 → 优化粒度细                          │  │
│  └─────────────────────────────────────────────────────────────────────────┘  │
│                                                                                  │
│  每降低一层 (lowering),编译器做一次"翻译"                                      │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31

1.2 各层级 IR 的用途 ​

层级IR 名称抽象程度主要用途代表系统
L4Graph IR极高用户级表示,接近框架 APITF GraphDef, PyTorch FX, JAX HLO
L3Op IR高算子融合/代数优化XLA HLO, TVM Relay, Linalg
L2Loop IR中循环变换/内存布局TVM TE, LLVM IR
L1Instruction IR低指令选择/寄存器分配LLVM MC, NVPTX (PTX)
L0硬件描述极低资源建模/调度NVVM, ISPC

第2节 计算图 IR(L4 层) ​

2.1 计算图的基本结构 ​

计算图是 ML 框架的标准表示:

┌─────────────────────────────────────────────────────────────────────────────────┐
│                            计算图的基本结构                                      │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│  节点 (Node) = 操作 (Operation/Op)                                             │
│  边 (Edge) = Tensor 数据                                                        │
│                                                                                  │
│  示例: y = (W * x + b).relu()                                                  │
│                                                                                  │
│       W                    b                                                     │
│        \                  /                                                      │
│         \                /                                                       │
│          \              /                                                        │
│           \            /                                                         │
│            \          /                                                          │
│             ┌────────────┐                                                      │
│             │   MatMul   │                                                      │
│             └────────────┘                                                      │
│                  │                                                             │
│                  ▼                                                             │
│             ┌────────────┐                                                      │
│             │    Add     │  ← BiasAdd                                          │
│             └────────────┘                                                      │
│                  │                                                             │
│                  ▼                                                             │
│             ┌────────────┐                                                      │
│             │   ReLU    │                                                      │
│             └────────────┘                                                      │
│                  │                                                             │
│                  ▼                                                             │
│                  y                                                             │
│                                                                                  │
│  计算图的数学表示:                                                               │
│  ════════════════════════════════════════════════                              │
│  y = ReLU(MatMul(x, W) + b)                                                   │
│                                                                                  │
│  DAG (有向无环图) 特性:                                                         │
│  • 无环:不会有循环依赖                                                         │
│  • 有向:数据从输入流向输出                                                     │
│  • 可拓扑排序:容易确定执行顺序                                                 │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42

2.2 PyTorch FX Graph ​

PyTorch FX (Framework eXchange) 是 PyTorch 2.0 的图表示:

python
import torch
import torch.fx

class SimpleModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = torch.nn.Linear(10, 10)
    
    def forward(self, x):
        x = self.linear(x)
        x = torch.relu(x)
        return x

model = SimpleModel()

# 追踪生成 FX Graph
traced = torch.fx.symbolic_trace(model)

# 打印 Graph
print(traced.graph)
"""
graph():
    %x : [#users=1] = placeholder[target=x]
    %linear_weight : [#users=1] = get_attr[target=linear_weight]
    %linear_bias : [#users=1] = get_attr[target=linear_bias]
    %linear : [#users=1] = call_module[target=linear](%x)
    %relu : [#users=1] = call_function[target=relu](%linear)
    return %relu
"""

# 获取节点的详细信息
for node in traced.graph.nodes:
    print(f"Op: {node.op}, Target: {node.target}")
    print(f"  Args: {node.args}")
    print(f"  Kwargs: {node.kwargs}")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35

FX Graph 的节点类型:

Op 类型说明示例
placeholder输入张量%x = placeholder[target=x]
get_attr获取模块参数%w = get_attr[target=weight]
call_module调用子模块%out = call_module[target=linear](%x)
call_function调用函数%out = call_function[target=relu](%x)
call_method调用方法%out = call_method[target=view](%x)
output输出节点return %out

2.3 TensorFlow GraphDef ​

TensorFlow 使用 GraphDef 作为核心图表示:

python
import tensorflow as tf

# 构建 TF 图
@tf.function
def model(x):
    x = tf.keras.layers.Dense(10)(x)
    x = tf.nn.relu(x)
    return x

# 获取 concrete function 的图
concrete = model.get_concrete_function(
    tf.TensorSpec(shape=[None, 10], dtype=tf.float32)
)

# 获取 GraphDef
graph_def = concrete.graph.as_graph_def()

# 查看节点
for node in graph_def.node[:10]:  # 只看前10个
    print(f"Node: {node.name}, Op: {node.op}")
    if node.input:
        print(f"  Inputs: {node.input}")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22

GraphDef 节点示例:

node {
  name: "dense/kernel"
  op: "VarHandleOp"
  attr {
    key: "dtype"
    value { type: DT_FLOAT }
  }
}
node {
  name: "dense/BiasAdd"
  op: "BiasAdd"
  input: "dense/MatMul"
  input: "dense/bias"
  attr {
    key: "T"
    value { type: DT_FLOAT }
  }
}
node {
  name: "dense/MatMul"
  op: "MatMul"
  input: "input_1"
  input: "dense/kernel"
  attr {
    key: "transpose_a"
    value { b: false }
  }
  attr {
    key: "transpose_b"
    value { b: false }
  }
}
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32

2.4 JAX HLO ​

JAX 使用 HLO (High-Level Operations) 作为计算表示:

python
import jax
import jax.numpy as jnp
from jax import jit

def model(x, W, b):
    return jnp.maximum(jnp.dot(x, W) + b, 0)

# JIT 编译后可以看到 HLO
from jax import make_jaxpr

jaxpr = make_jaxpr(model)(
    jnp.zeros((1, 10)),
    jnp.zeros((10, 10)),
    jnp.zeros((10,))
)
print(jaxpr)

"""
{ lambda ; a:f32[1,10] b:f32[10,10] c:f32[10] .
  let d:f32[1,10] = dot_general[
      dimension_numbers=(([1], [0]), ([], []))
  ] a b
      e:f32[1,10] = broadcast_in_dim[
          broadcast_dimensions=()
          shape=[1, 10]
      ] c
      f:f32[1,10] = add d e
      g:f32[1,10] = reduce_max[axes=(1,)] f
      h:f32[1,10] = broadcast_in_dim[
          broadcast_dimensions=[1]
          shape=[1, 10]
      ] g
      i:f32[1,10] = sub f h
      j:f32[1,10] = exp i
      k:f32[1,10] = reduce_sum[axes=(1,)] j
      l:f32[1,10] = broadcast_in_dim[
          broadcast_dimensions=[1]
          shape=[1, 10]
      ] k
      m:f32[1,10] = div j l
  in m }
"""
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42

第3节 Operation-level IR(L3 层) ​

3.1 为什么需要 Op-level IR ​

Graph-level IR 过于高层,无法进行细粒度优化:

问题:Graph-level IR 无法做算子融合
════════════════════════════════════════════════════════════

Graph-level 表示:
┌─────────┐     ┌─────────┐     ┌─────────┐
│  Conv   │────▶│   Add   │────▶│  ReLU   │
└─────────┘     └─────────┘     └─────────┘
   ↓ 3 个 kernel    ↓ 3 个 kernel    ↓ 3 个 kernel
 显存访问       显存访问       显存访问

问题:
1. Conv + Add + ReLU 分开执行,需要 3 次 kernel 启动
2. Conv 输出要写回显存,再读出来给 Add
3. Add 输出再写回显存,再读出来给 ReLU
4. 大量显存读写成为瓶颈


Op-level 表示(融合后):
┌─────────────────────────────┐
│    fused_conv_add_relu      │  ← 一个融合 kernel
└─────────────────────────────┘
   ↓ 1 个 kernel
  显存访问

好处:
1. 减少 2 次 kernel 启动
2. 中间结果无需写回显存
3. 充分利用 GPU 寄存器
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28

3.2 XLA HLO 详解 ​

HLO (High-Level Operations) 是 XLA 的核心 IR:

hlo
// XLA HLO 示例:简单的卷积层
HloModule convolution_module

ENTRY %convolution.10 (input: f32[1,224,224,3], filter: f32[7,7,3,64]) -> f32[1,112,112,64] {
  %input = f32[1,224,224,3] parameter(0)
  %filter = f32[7,7,3,64] parameter(1)
  
  // Convolution operation
  %convolution.6 = f32[1,112,112,64] convolution(
      f32[1,224,224,3] %input,
      f32[7,7,3,64] %filter),
    window={size=7x7 stride=2x2 pad=3_3_3_3},
    dim_labels=b01o_01io->b01o,
    backend_config="convolution_dimension_numbers"
  
  // Bias addition
  %bias = f32[64] parameter(2)  // 假设有 bias
  %broadcast = f32[1,112,112,64] broadcast(f32[64] %bias), dimensions={3}
  %result = f32[1,112,112,64] add(%convolution.6, %broadcast)
  
  // ReLU activation
  %zero = f32[] constant(0)
  %broadcast_zero = f32[1,112,112,64] broadcast(f32[] %zero)
  ROOT %activation = f32[1,112,112,64] maximum(%result, %broadcast_zero)
}
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25

HLO 操作分类:

类别操作说明
Element-wiseadd, mul, maximum, exp, log按元素操作
Shapereshape, broadcast, transpose改变 tensor 形状
Linear Algebradot, convolution矩阵乘法、卷积
Reductionreduce_sum, reduce_max归约操作
Control Flowconditional, while条件/循环
Communicationcollective-permute, all-reduce多设备通信

3.3 TVM Relay IR ​

Relay 是 TVM 的高层 IR,融合了函数式和命令式特性:

python
# TVM Relay IR 示例
from tvm import relay

# 定义一个简单的模型
x = relay.var("x", shape=(1, 3, 224, 224))
weight = relay.var("weight", shape=(64, 3, 7, 7))
bias = relay.var("bias", shape=(64,))

# 构建 Relay 表达式
y = relay.nn.conv2d(x, weight, kernel_size=(7, 7), padding=(3, 3), strides=(2, 2))
y = relay.add(y, relay.reshape(bias, (1, 64, 1, 1)))
y = relay.nn.relu(y)

# 完整的 ResNet block
func = relay.Function([x, weight, bias], y)
module = relay.Module.from_expr(func)

# Relay IR 可以打印出来
print(module.astext())
"""
v0.0.4
func @main(%x: Tensor[(1, 3, 224, 224), float32], 
           %weight: Tensor[(64, 3, 7, 7), float32], 
           %bias: Tensor[(64,), float32]) 
           -> Tensor[(1, 64, 112, 112), float32] {
  %0 = nn.conv2d(%x, %weight, kernel_size=[7, 7], padding=[3, 3, 3, 3], strides=[2, 2])
  %1 = reshape(%bias, newshape=[1, 64, 1, 1])
  %2 = add(%0, %1)
  %3 = nn.relu(%2)
  %4 = reshape(%bias, newshape=[1, 64, 1, 1])  // 冗余操作
  %5 = add(%3, %4)  // 编译器会发现并消除
  %6 = multiply(%3, %5)  // 其他融合机会
  %7 = nn.relu(%5)
  %8 = add(%7, %3)
  nn.relu(%8)
}
"""
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37

第4节 Instruction-level IR(L2/L1 层) ​

4.1 LLVM IR 详解 ​

LLVM IR 是最成熟的 instruction-level IR:

llvm
; LLVM IR 示例:矩阵乘法片段
define dso_local void @matmul_kernel(
    float* %C,      ; 输出矩阵 C
    float* %A,      ; 输入矩阵 A
    float* %B,      ; 输入矩阵 B
    i64 %M,         ; A 的行数
    i64 %N,         ; B 的列数
    i64 %K          ; A的列数/B的行数
) #0 {
entry:
  %i = call i64 @llvm.nvvm.read.ptx.sreg.ctaid.x()
  %j = call i64 @llvm.nvvm.read.ptx.sreg.tid.x()
  
  ; 计算 C[i][j] = sum_k A[i][k] * B[k][j]
  br label %outer_loop

outer_loop:
  %k = phi i64 [ 0, %entry ], [ %k.next, %outer_loop ]
  %accum = phi float [ 0.0, %entry ], [ %sum, %outer_loop ]
  
  ; C[i*N + j] += A[i*K + k] * B[k*N + j]
  %a_idx = add i64 %i, %k
  %a_val = load float, float* %A
  
  %b_idx = add i64 %k, %j
  %b_val = load float, float* %B
  
  %prod = fmul float %a_val, %b_val
  %sum = fadd float %accum, %prod
  
  %k.next = add i64 %k, 1
  %cond = icmp slt i64 %k.next, %K
  br i1 %cond, label %outer_loop, label %end

end:
  %C_idx = add i64 %i, %j
  store float %sum, float* %C
  ret void
}
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39

LLVM IR 核心概念:

概念说明示例
Function函数定义define void @foo(...)
BasicBlock基本块顺序执行的代码段
Instruction指令add i32 %a, %b
Value值(SSA)每个操作产生一个值
Type类型系统i32, float, ptr
Metadata元数据调试信息、优化提示

4.2 PTX (Parallel Thread Execution) ​

PTX 是 NVIDIA 的伪汇编语言:

asm
# PTX 示例:向量加法
//
// __global__ void vector_add(float *C, float *A, float *B, int N)
//
    .version 8.0
    .target sm_80
    .address_size 64

.visible .entry vector_add(
    .param .u64 vector_add_param_0,  // C
    .param .u64 vector_add_param_1,  // A
    .param .u64 vector_add_param_2,  // B
    .param .u32 vector_add_param_3   // N
)
{
    .reg .pred %p;               // 谓词寄存器
    .reg .f32 %f;                // 32位浮点寄存器
    .reg .b32 %r;                // 32位通用寄存器
    .reg .b64 %rd;               // 64位通用寄存器

    ld.param.u64 %rd0, [vector_add_param_0];  // 加载 C 地址
    ld.param.u64 %rd1, [vector_add_param_1];  // 加载 A 地址
    ld.param.u64 %rd2, [vector_add_param_2];  // 加载 B 地址
    ld.param.u32 %r0, [vector_add_param_3];    // 加载 N

    cvta.to.global.u64 %rd0, %rd0;            // 转换为全局地址
    cvta.to.global.u64 %rd1, %rd1;
    cvta.to.global.u64 %rd2, %rd2;

    mov.u32 %r1, %tid.x;                       // 线程索引
    mul.wide.u32 %rd3, %r1, 4;                // 计算偏移
    add.s64 %rd4, %rd0, %rd3;                 // C[i] 地址
    add.s64 %rd5, %rd1, %rd3;                 // A[i] 地址
    add.s64 %rd6, %rd2, %rd3;                 // B[i] 地址

    ld.global.f32 %f0, [%rd5];                // 加载 A[i]
    ld.global.f32 %f1, [%rd6];                // 加载 B[i]
    add.f32 %f2, %f0, %f1;                   // C[i] = A[i] + B[i]
    st.global.f32 [%rd4], %f2;                // 存储结果

    ret;
}
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42

4.3 SPIR-V (Vulkan/OpenCL) ​

SPIR-V 是 Khronos 的 GPU 中间表示:

llvm
; SPIR-V (LLVM IR 格式)
; 用于 Vulkan/OpenCL 的 SPIR-V 生成
;
; 该 IR 会被工具链翻译成真正的 SPIR-V 二进制

; vector_add kernel
define spir_kernel void @vector_add(
    float* nocapture readonly %C,
    float* nocapture readonly %A,
    float* nocapture readonly %B,
    i32 %N
) #0 {
entry:
  %idx = call @_Z13get_global_idj(i32 0)  ; get_global_id(0)
  %cmp = icmp slt i32 %idx, %N
  br i1 %cmp, label %body, label %end

body:
  %a_ptr = getelementptr float, float* %A, i32 %idx
  %a_val = load float, float* %a_ptr
  
  %b_ptr = getelementptr float, float* %B, i32 %idx
  %b_val = load float, float* %b_ptr
  
  %sum = fadd float %a_val, %b_val
  
  %c_ptr = getelementptr float, float* %C, i32 %idx
  store float %sum, float* %c_ptr
  br label %end

end:
  ret void
}
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33

第5节 Lowering 链路 ​

5.1 什么是 Lowering ​

Lowering 是将高层 IR 逐步翻译到低层 IR 的过程:

┌─────────────────────────────────────────────────────────────────────────────────┐
│                              Lowering 概念                                      │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│  Lowering = 逐步降低抽象层级                                                    │
│  ════════════════════════════════════════════                                   │
│                                                                                  │
│      高层语义                                                                     │
│         │                                                                        │
│         ▼  lowering                                                              │
│      ┌────────┐                                                                 │
│      │ Conv2D │  ← 语义丰富:卷积操作                                         │
│      │ HLO    │                                                                 │
│      └────────┘                                                                 │
│         │                                                                        │
│         ▼  lowering                                                              │
│      ┌────────┐                                                                 │
│      │ Winograd│  ← 更具体:使用 Winograd 算法                                │
│      │ or GEMM │                                                                 │
│      └────────┘                                                                 │
│         │                                                                        │
│         ▼  lowering                                                              │
│      ┌────────┐                                                                 │
│      │ LLVM IR │  ← 指令级:具体的循环和内存访问                              │
│      └────────┘                                                                 │
│         │                                                                        │
│         ▼  lowering                                                              │
│      ┌────────┐                                                                 │
│      │   PTX   │  ← GPU 伪汇编                                                │
│      └────────┘                                                                 │
│         │                                                                        │
│         ▼  lowering                                                              │
│      ┌────────┐                                                                 │
│      │  SASS   │  ← GPU 机器码                                                │
│      └────────┘                                                                 │
│                                                                                  │
│  每一步 lowering:                                                                │
│  • 抽象层级降低                                                                  │
│  • 语义更具体                                                                   │
│  • 优化空间变小但优化粒度变细                                                    │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42

5.2 TensorFlow → GPU 完整 Lowering 链路 ​

┌─────────────────────────────────────────────────────────────────────────────────┐
│              TF GraphDef → XLA HLO → LLVM IR → PTX → SASS                      │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│  Level 1: TF GraphDef (高层)                                                    │
│  ══════════════════════════                                                      │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  node {                                                                  │    │
│  │    name: "conv2d"                                                        │    │
│  │    op: "Conv2D"                                                         │    │
│  │    input: "input:0"                                                     │    │
│  │    input: "kernel:0"                                                     │    │
│  │    attr { key: "padding" value { s: "SAME" } }                        │    │
│  │  }                                                                     │    │
│  │  node {                                                                  │    │
│  │    name: "bias_add"                                                     │    │
│  │    op: "BiasAdd"                                                       │    │
│  │    input: "conv2d:0"                                                     │    │
│  │    input: "bias:0"                                                      │    │
│  │  }                                                                     │    │
│  │  node {                                                                  │    │
│  │    name: "relu"                                                         │    │
│  │    op: "Relu"                                                           │    │
│  │    input: "bias_add:0"                                                  │    │
│  │  }                                                                     │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Level 2: Grappler 优化 (TF GraphDef → 优化后的 GraphDef)                      │
│  ══════════════════════════════════════════════════════════════               │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  • Constant folding: 折叠常量节点                                       │    │
│  │  • Layout optimizer: NHWC → NCHW (cuDNN 友好)                          │    │
│  │  • Arithmetic optimizer: 合并乘加运算                                   │    │
│  │  • Remapper: 将 Conv2D + BiasAdd + Relu 替换为 fused op                │    │
│  │                                                                          │    │
│  │  # 优化后的图                                                           │    │
│  │  node {                                                                  │    │
│  │    name: "fused_conv2d_bias_relu"                                      │    │
│  │    op: "_FusedConv2D"                                                  │    │
│  │    input: "input:0"                                                    │    │
│  │    input: "kernel:0"                                                   │    │
│  │    input: "bias:0"                                                      │    │
│  │    attr { key: "num_args" value { i: 1 } }                             │    │
│  │    attr { key: "fused_ops" value { list: { s: "BiasAdd" s: "Relu" }}} │    │
│  │  }                                                                     │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Level 3: XLA HLO (操作级 IR)                                                  │
│  ══════════════════════════════════════════════════════════════               │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  ENTRY %main.9 (input: f32[1,224,224,3], kernel: f32[7,7,3,64],       │    │
│  │                      bias: f32[64]) -> f32[1,112,112,64] {            │    │
│  │                                                                          │    │
│  │    %input = f32[1,224,224,3] parameter(0)                             │    │
│  │    %kernel = f32[7,7,3,64] parameter(1)                               │    │
│  │    %bias = f32[64] parameter(2)                                        │    │
│  │                                                                          │    │
│  │    # 融合后的 convolution                                               │    │
│  │    %conv = f32[1,112,112,64] convolution(%input, %kernel),           │    │
│  │      window={size=7x7 stride=2x2 pad=3_3_3_3},                         │    │
│  │      dim_labels=b01o_01io->b01o                                         │    │
│  │                                                                          │    │
│  │    # Broadcast bias                                                     │    │
│  │    %bias_bcast = f32[1,112,112,64] broadcast(%bias),                  │    │
│  │      dimensions={3}                                                     │    │
│  │                                                                          │    │
│  │    # Add bias                                                           │    │
│  │    %with_bias = f32[1,112,112,64] add(%conv, %bias_bcast)            │    │
│  │                                                                          │    │
│  │    # ReLU                                                               │    │
│  │    %zero = f32[] constant(0)                                           │    │
│  │    %zero_bcast = f32[1,112,112,64] broadcast(%zero),                   │    │
│  │      dimensions={}                                                      │    │
│  │    ROOT %output = f32[1,112,112,64] maximum(%with_bias, %zero_bcast) │    │
│  │  }                                                                     │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Level 4: XLA 优化 (HLO → 优化后的 HLO)                                        │
│  ══════════════════════════════════════════════════════════════               │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  # XLA HLO 优化器会对 HLO 做:                                           │    │
│  │                                                                          │    │
│  │  1. AlgebraicSimplifier: 代数化简                                       │    │
│  │     - 合并连续的 broadcast                                               │    │
│  │     - 消除冗余操作                                                       │    │
│  │                                                                          │    │
│  │  2. ConvolutionAlgorithmChooser:                                       │    │
│  │     - 选择最优卷积算法 (直接/FFT/Winograd)                              │    │
│  │     - 根据 shape 和硬件特性选择                                          │    │
│  │                                                                          │    │
│  │  3. LayoutAssignment:                                                 │    │
│  │     - 为每个 tensor 选择最优内存布局                                     │    │
│  │     - NCHW vs NHWC vs NCHWc                                            │    │
│  │                                                                          │    │
│  │  4. InstructionFusion:                                                 │    │
│  │     - 融合相邻的 element-wise 操作                                     │    │
│  │     - 融合 Conv + BiasAdd + Maximum → single fusion                   │    │
│  │                                                                          │    │
│  │  # 优化后的 HLO (融合了 Conv + Bias + ReLU)                           │    │
│  │  %fused_conv = f32[1,112,112,64] fusion(...),                         │    │
│  │    kind=kCustom, calls={                                                │    │
│  │      %conv = convolution(...)                                           │    │
│  │      %bias = broadcast(bias)                                            │    │
│  │      %add = add(%conv, %bias)                                           │    │
│  │      %relu = maximum(%add, zero)                                        │    │
│  │    }                                                                    │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Level 5: LLVM IR (指令级 IR)                                                   │
│  ══════════════════════════════════════════════════════════════               │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  define void @conv_kernel(                                              │    │
│  │      float* %C, float* %A, float* %B, float* %bias                    │    │
│  │  ) #0 {                                                                 │    │
│  │  entry:                                                                 │    │
│  │    %i = call i64 @llvm.nvvm.read.ptx.sreg.ctaid.x()                  │    │
│  │    %j = call i64 @llvm.nvvm.read.ptx.sreg.ctaid.y()                  │    │
│  │    %k = call i64 @llvm.nvvm.read.ptx.sreg.tid.x()                    │    │
│  │                                                                          │    │
│  │    ; 循环展开/向量化的卷积实现                                           │    │
│  │    %result = alloca float                                              │    │
│  │                                                                          │    │
│  │    br label %loop_start                                                │    │
│  │                                                                          │    │
│  │  loop_start:                                                           │    │
│  │    %k_val = phi i64 [0, %entry], [%k_next, %loop_start]              │    │
│  │    %acc = phi float [0.0, %entry], [%acc_next, %loop_start]         │    │
│  │                                                                          │    │
│  │    %A_val = call @llvm.nvvm.ldg.f32.f32(float* %A_ptr)              │    │
│  │    %B_val = call @llvm.nvvm.ldg.f32.f32(float* %B_ptr)              │    │
│  │    %prod = fmul float %A_val, %B_val                                  │    │
│  │    %acc_next = fadd float %acc, %prod                                 │    │
│  │                                                                          │    │
│  │    %k_next = add i64 %k_val, 1                                        │    │
│  │    %cond = icmp slt i64 %k_next, %K                                   │    │
│  │    br i1 %cond, label %loop_start, label %loop_end                    │    │
│  │                                                                          │    │
│  │  loop_end:                                                             │    │
│  │    %bias_val = load float, float* %bias_ptr                          │    │
│  │    %biased = fadd float %acc, %bias_val                               │    │
│  │    %relu_out = call @fmax.f32(%biased, 0.0)                          │    │
│  │    store float %relu_out, float* %C_ptr                               │    │
│  │    ret void                                                            │    │
│  │  }                                                                     │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Level 6: PTX (GPU 伪汇编)                                                     │
│  ══════════════════════════════════════════════════════════════               │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  .visible .entry _Z12conv_kernelPPfS_S_(                                │    │
│  │      .param .u64 _Z12conv_kernelPPfS_S__param_0,                       │    │
│  │      .param .u64 _Z12conv_kernelPPfS_S__param_1,                        │    │
│  │      .param .u64 _Z12conv_kernelPPfS_S__param_2,                        │    │
│  │      .param .u64 _Z12conv_kernelPPfS_S__param_3                         │    │
│  │  ) {                                                                    │    │
│  │      .reg .pred %p;                                                    │    │
│  │      .reg .f32 %f;                                                     │    │
│  │      .reg .b32 %r;                                                     │    │
│  │      .reg .b64 %rd;                                                    │    │
│  │                                                                          │    │
│  │      ld.param.u64 %rd0, [_Z12conv_kernelPPfS_S__param_0];             │    │
│  │      ...                                                                │    │
│  │      ld.global.f32 %f0, [%rd1];      // 加载 A[i]                      │    │
│  │      ld.global.f32 %f1, [%rd2];      // 加载 B[i]                      │    │
│  │      fma.rn.f32 %f2, %f0, %f1, %f3; // 乘加 + bias                    │    │
│  │      max.f32 %f4, %f2, 0;           // ReLU                           │    │
│  │      st.global.f32 [%rd0], %f4;     // 存储结果                        │    │
│  │      ret;                                                                │    │
│  │  }                                                                      │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Level 7: SASS (GPU 机器码)                                                    │
│  ══════════════════════════════════════════════════════════════               │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  // A100 SASS (通过 NVCC 或 PTXAS 编译得到)                              │    │
│  │                                                                          │    │
│  │  /*0058*/              MOV R4, R0;              // R4 = thread id      │    │
│  │  /*005c*/              SHR R4, R4, 0x2;          // R4 >>= 2           │    │
│  │  /*0060*/              ISETP.GE.AND P0, P1, R4, │                      │    │
│  │  /*0064*/              R0, P0, PT;              // Predication        │    │
│  │  /*0068*/              LDG.E.CONCAT  R0,         │                      │    │
│  │  /*0070*/              @P0  LDG.E  R6, [R6];    // Load with cache    │    │
│  │  /*0078*/              @P0  FMA R10, R6, R4, R5;│ // Fused multiply-add│    │
│  │  /*0080*/              @P0  MAX R10, R10, RZ;   │ // ReLU             │    │
│  │  /*0088*/              @P0  STG.E  [R0], R10;   │ // Store result     │    │
│  │  /*0090*/              EXIT;                     │ // Kernel exit     │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189

5.3 PyTorch → GPU 完整 Lowering 链路 ​

┌─────────────────────────────────────────────────────────────────────────────────┐
│              PyTorch → FX → Inductor → Triton/CUDA → PTX → SASS                │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│  Step 1: PyTorch Eager (动态图)                                                │
│  ═══════════════════════════════════════                                         │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  import torch                                                            │    │
│  │  model = torch.nn.Sequential(                                            │    │
│  │      torch.nn.Conv2d(3, 64, 7, padding=3),                             │    │
│  │      torch.nn.BatchNorm2d(64),                                          │    │
│  │      torch.nn.ReLU()                                                     │    │
│  │  )                                                                       │    │
│  │  x = torch.randn(1, 3, 224, 224)                                        │    │
│  │  y = model(x)                                                           │    │
│  │                                                                          │    │
│  │  # 每次前向都重新追踪,重新调度                                           │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Step 2: TorchDynamo (图捕获)                                                   │
│  ════════════════════════════════════════════════                              │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  # torch.compile 启用 Dynamo                                            │    │
│  │  compiled = torch.compile(model, backend="inductor")                    │    │
│  │                                                                          │    │
│  │  # Dynamo 拦截 Python bytecode,生成 FX Graph:                            │    │
│  │  # graph():                                                              │    │
│  │  #   %x : [#users=1] = placeholder[target=x]                           │    │
│  │  #   %weight : [#users=1] = get_attr[target=weight]                   │    │
│  │  #   %bias : [#users=1] = get_attr[target=bias]                       │    │
│  │  #   %conv : [#users=1] = call_module[target=conv](%x)                │    │
│  │  #   %bn : [#users=1] = call_module[target=bn](%conv)                 │    │
│  │  #   %relu : [#users=1] = call_function[target=relu](%bn)            │    │
│  │  #   return %relu                                                       │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Step 3: AOTDispatch (提前追踪)                                                  │
│  ════════════════════════════════════════════════                              │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  # AOTAutograd 追踪 backward,生成 graph 和 gradient graph             │    │
│  │                                                                          │    │
│  │  # Forward graph:                                                       │    │
│  │  # fx_graph:                                                           │    │
│  │  #   %x = placeholder()                                                │    │
│  │  #   %w = placeholder()                                                │    │
│  │  #   %b = placeholder()                                                │    │
│  │  #   %conv = convolution(%x, %w)                                       │    │
│  │  #   %add = add(%conv, %b)                                             │    │
│  │  #   %relu = relu(%add)                                                 │    │
│  │  #   return (forward_graph, tuple(grad_outs))                          │    │
│  │                                                                          │    │
│  │  # Backward graph:                                                     │    │
│  │  #   记录梯度计算图,用于反向传播                                         │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Step 4: Inductor (核心编译)                                                     │
│  ════════════════════════════════════════════════                              │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  # Inductor 接收 FX Graph,进行 lowering 和 kernel 生成                  │    │
│  │                                                                          │    │
│  │  # 1. Scheduler: 调度计算                                              │    │
│  │  with Scheduler(scheduler) as scheduler:                                │    │
│  │      scheduler.enqueue("convolution_0")                                 │    │
│  │      scheduler.enqueue("add_1")                                         │    │
│  │      scheduler.enqueue("relu_2")                                        │    │
│  │                                                                          │    │
│  │  # 2. Fusion: 算子融合                                                 │    │
│  │  # Conv + Add + ReLU → single fused kernel                             │    │
│  │  scheduler.fuse_nodes(["convolution_0", "add_1", "relu_2"])            │    │
│  │                                                                          │    │
│  │  # 3. 选择 Compute Library                                             │    │
│  │  # 如果 cuDNN 可用,选择 cudnn_convolution; 否则用自己的 kernel         │    │
│  │  backend = "inductor"                                                   │    │
│  │  if has_cudnn():                                                       │    │
│  │      kernel = cudnn_convolution_plus_bias_relu()                        │    │
│  │  else:                                                                 │    │
│  │      kernel = inductor_fused_conv_add_relu()                           │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Step 5: Inductor Kernel 生成 (Triton/C++)                                      │
│  ══════════════════════════════════════════════════════                       │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  # Inductor 生成 Triton 代码                                            │    │
│  │                                                                          │    │
│  │  @triton.jit                                                             │    │
│  │  def triton_conv_fused(                                                  │    │
│  │      in_ptr, out_ptr, weight_ptr, bias_ptr,                             │    │
│  │      stride_h, stride_w, padding_h, padding_w,                          │    │
│  │      in_channel, out_channel, height, width                            │    │
│  │  ):                                                                      │    │
│  │      # Triton 自动管理线程和内存                                         │    │
│  │      pid_h = tl.program_id(0)                                           │    │
│  │      pid_w = tl.program_id(1)                                          │    │
│  │      pid_c = tl.program_id(2)                                          │    │
│  │                                                                          │    │
│  │      offs_h = pid_h * stride_h + tl.arange(0, 7) - padding_h           │    │
│  │      offs_w = pid_w * stride_w + tl.arange(0, 7) - padding_w           │    │
│  │      offs_in = offs_h[:, None] * width + offs_w[None, :]             │    │
│  │                                                                          │    │
│  │      # 向量化加载                                                        │    │
│  │      inp = tl.load(in_ptr + offs_in, mask=...)                         │    │
│  │      wgt = tl.load(weight_ptr + ...)                                   │    │
│  │                                                                          │    │
│  │      # 矩阵乘法 (Triton 自动向量化)                                      │    │
│  │      acc = tl.dot(inp, wgt)                                            │    │
│  │                                                                          │    │
│  │      # Bias + ReLU (融合)                                               │    │
│  │      acc = acc + tl.load(bias_ptr + pid_c)                             │    │
│  │      acc = tl.maximum(acc, 0)                                          │    │
│  │                                                                          │    │
│  │      tl.store(out_ptr + ..., acc)                                       │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    ↓                                            │
│  Step 6: PTX → SASS (GPU 机器码)                                                │
│  ══════════════════════════════════════════════════════                       │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  # Triton 编译器将 Triton 代码编译成 PTX                                  │    │
│  │  # PTXAS (NVIDIA PTX 汇编器) 将 PTX 编译成 SASS                         │    │
│  │                                                                          │    │
│  │  // 最终的 SASS 代码 (A100)                                              │    │
│  │  MOV R4, R0;                   // Thread ID                           │    │
│  │  I2F R4, R4;                  // Int to Float                         │    │
│  │  LDG.E.CONCAT R6, [R6];       // Load with L2 cache hint             │    │
│  │  FFMA R10, R6, R4, R5;       // FMA: a*b + bias                     │    │
│  │  MAX R10, R10, RZ;            // ReLU: max(result, 0)                │    │
│  │  STG.E [R0], R10;            // Store result                         │    │
│  │  EXIT;                        // Kernel exit                          │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130

5.4 ONNX → 多后端 Lowering ​

ONNX 作为中间格式,支持多种 lowering 路径:

┌─────────────────────────────────────────────────────────────────────────────────┐
│                        ONNX → 多后端 Lowering                                   │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│  ONNX Model (Protobuf)                                                          │
│  ════════════════════════                                                         │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  ir_version: 8                                                         │    │
│  │  producer_name: "pytorch"                                             │    │
│  │  graph {                                                                │    │
│  │    node {                                                               │    │
│  │      input: "input"                                                    │    │
│  │      output: "conv_out"                                                │    │
│  │      op_type: "Conv"                                                  │    │
│  │      attribute { name: "kernel_shape" i: [7, 7] }                     │    │
│  │      attribute { name: "strides" i: [2, 2] }                         │    │
│  │      attribute { name: "pads" i: [3, 3, 3, 3] }                     │    │
│  │    }                                                                   │    │
│  │    node {                                                              │    │
│  │      input: "conv_out"                                                │    │
│  │      output: "relu_out"                                               │    │
│  │      op_type: "Relu"                                                  │    │
│  │    }                                                                   │    │
│  │  }                                                                     │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                    │                                             │
│           ┌────────────────────────┼────────────────────────┐                    │
│           │                        │                        │                     │
│           ▼                        ▼                        ▼                     │
│  ┌─────────────────┐    ┌─────────────────┐    ┌─────────────────┐              │
│  │   ONNX Runtime  │    │    TVM          │    │   TensorRT      │              │
│  │   (Direct)      │    │   (ONNX → TVM)  │    │   (ONNX → TRT)  │              │
│  └────────┬────────┘    └────────┬────────┘    └────────┬────────┘              │
│           │                       │                       │                       │
│           ▼                       ▼                       ▼                       │
│  ┌─────────────────┐    ┌─────────────────┐    ┌─────────────────┐              │
│  │ ONNX → MLAS    │    │ TVM Relay IR    │    │ TRT Network     │              │
│  │ (Microsoft BLAS)│    │ ↓               │    │ ↓               │              │
│  └────────┬────────┘    │ TVM TE/TIR     │    │ TRT Engine      │              │
│           │              │ ↓               │    │ ↓               │              │
│           ▼              │ LLVM IR         │    │ CUDA/CuDNN      │              │
│  ┌─────────────────┐    │ ↓               │    │ ↓               │              │
│  │  CPU BLAS       │    │ PTX/SASS        │    │ SASS            │              │
│  │  (x86/ARM)      │    │ ↓               │    │ ↓               │              │
│  └─────────────────┘    │ GPU Exec        │    │ GPU Exec        │              │
│                         └─────────────────┘    └─────────────────┘              │
│                                                                                  │
│  ┌─────────────────────────────────────────────────────────────────────────┐    │
│  │  其他后端路径:                                                           │    │
│  │                                                                          │    │
│  │  ONNX → ONNX-MLIR → Linalg → LLVM → SPIR-V → Vulkan/Vulkan Compute   │    │
│  │  ONNX → ONNX-TFlite → FlatBuffer → ARM NEON / Hexagon DSP            │    │
│  │  ONNX → ONNX-JS → WebGL / WebGPU / WASM                              │    │
│  │  ONNX → ONNX-TRT → TensorRT INT8 → Calibration → INT8 Engine         │    │
│  └─────────────────────────────────────────────────────────────────────────┘    │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57

第6节 为什么需要多级 IR ​

6.1 抽象能力 vs 优化粒度的权衡 ​

多级 IR 的核心原因:不同抽象层级适合做不同类型的优化:

┌─────────────────────────────────────────────────────────────────────────────────┐
│                      IR 层级的优化任务分配                                       │
├─────────────────────────────────────────────────────────────────────────────────┤
│                                                                                  │
│  Graph-level IR (L4):                                                          │
│  ═════════════════════                                                          │
│  ✅ 算子融合 (Op Fusion)                                                        │
│     Conv + BN + ReLU → fused_op                                               │
│  ✅ 子图替换 (Graph Substitution)                                               │
│     替换整个计算模式                                                             │
│  ✅ 内存共享 (Memory Sharing)                                                   │
│     识别可共享内存的 tensor                                                     │
│  ❌ 循环变换 (Loop Transformation) - 层级太高                                   │
│  ❌ 寄存器分配 (Register Allocation) - 层级太高                               │
│                                                                                  │
│  Op-level IR (L3):                                                             │
│  ═════════════════════                                                          │
│  ✅ 算法选择 (Algorithm Selection)                                               │
│     Conv: direct vs FFT vs Winograd                                           │
│  ✅ 布局转换 (Layout Transformation)                                           │
│     NHWC ↔ NCHW ↔ NCHWc                                                       │
│  ✅ 精度量化 (Precision Quantization)                                          │
│     FP32 → FP16 → INT8                                                        │
│  ❌ 循环展开 (Loop Unrolling) - 层级太高                                        │
│  ❌ 指令调度 (Instruction Scheduling) - 层级太高                               │
│                                                                                  │
│  Loop-level IR (L2):                                                           │
│  ═════════════════════                                                          │
│  ✅ 循环变换 (Loop Transformation)                                              │
│     tiling, unrolling, fusion                                                 │
│  ✅ 内存层级优化 (Memory Hierarchy)                                             │
│     register file, shared memory, global memory                               │
│  ✅ 数据局部性 (Data Locality)                                                  │
│     reuse, blocking                                                           │
│  ❌ 指令选择 (Instruction Selection) - 层级太低                               │
│  ❌ 寄存器分配 (Register Allocation) - 层级太低                               │
│                                                                                  │
│  Instruction-level IR (L1):                                                    │
│  ═════════════════════════                                                      │
│  ✅ 指令选择 (Instruction Selection)                                           │
│     选择最优指令                                                                │
│  ✅ 寄存器分配 (Register Allocation)                                           │
│     图着色、线性扫描                                                            │
│  ✅ 指令调度 (Instruction Scheduling)                                          │
│     隐藏延迟、最大化发射率                                                       │
│                                                                                  │
│  结论: 如果只有单一层级,要么丢失优化机会,要么实现复杂                          │
│                                                                                  │
└─────────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49

6.2 模块化与可复用性 ​

多级 IR 使编译器各部分可以独立演进和复用:

组件复用方式
语言前端只需为每种框架写一个前端
优化器中层和高层优化可以跨框架复用
硬件后端只需为每种硬件写一个后端
复用示例:
════════════════════════════════════════════════════════════

如果没有多级 IR:
Caffe → x86
Caffe → ARM
TensorFlow → x86
TensorFlow → ARM
PyTorch → x86
PyTorch → ARM
... (每个组合都要单独实现)

如果有中间 IR:
Caffe ─┐
TF    ─┼──▶ Relay IR ─▶ LLVM IR ─▶ x86/ARM/CUDA/RISC-V
PyTorch┘                           (后端复用)

好处:
• 4 个前端 × 1 个优化器 × 1 个后端 = 4 个编译器
• 而不是 4 × 4 = 16 个编译器
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

第7节 IR 的属性 ​

7.1 类型系统 ​

IR 需要类型系统来描述值的种类:

类型说明示例
Scalar标量i32, f32, bool
Tensor张量f32[128, 512]
Tuple元组(f32, i32)
Function函数(i32) → f32
Pointer指针ptr<f32>
python
# TVM Relay 类型系统
from tvm import relay

# 类型注解
x = relay.var("x", shape=(1, 3, 224, 224), dtype="float32")
# x 的类型: Tensor[(1, 3, 224, 224), float32]

weight = relay.var("w", shape=(64, 3, 7, 7), dtype="float32")
# weight 的类型: Tensor[(64, 3, 7, 7), float32]

y = relay.nn.conv2d(x, weight, kernel_size=(7, 7), padding=(3, 3))
# y 的类型: Tensor[(1, 64, 112, 112), float32]
1
2
3
4
5
6
7
8
9
10
11
12

7.2 SSA 形式 ​

如前所述,现代 IR 使用 SSA 简化分析:

llvm
; LLVM IR 天然是 SSA 形式
%result1 = add i32 %a, %b
%result2 = mul i32 %result1, 2
%result3 = sub i32 %result2, %c
; 每个变量只赋值一次
1
2
3
4
5

7.3 副作用建模 ​

ML IR 需要建模副作用(Side Effects):

副作用类型示例IR 表示
无副作用纯函数pure 属性
读取全局状态RNGread_effect(state)
写入全局状态更新权重write_effect(state)
I/O打印、文件io_effect
异常错误处理nothrow 属性
python
# TVM Relay 副作用注解
from tvm.relay.op.annotation import compiler_assigned

# 标记为无副作用,编译器可以自由重排
no_side_effect = relay.expr.set_span(
    relu,
    relu_span.with_attr("Primitive", 1)
)

# 标记 RNG 操作
random_uniform = relay.random.uniform(shape=(10, 10), seed=42)
# 有副作用,必须按顺序执行
1
2
3
4
5
6
7
8
9
10
11
12

第8节 主要 IR 系统对比 ​

IR 系统层级抽象程度主要用户特点
XLA HLOL3高TensorFlow, JAX, TPU强于 TPU,自动融合
TVM RelayL3高TVM函数式 + 命令式混合
LinalgL3中高MLIR可组合的 op 库
LLVM IRL2中通用成熟的编译器基础设施
SPIR-VL2中Vulkan, OpenCL跨厂商 GPU
PTXL1低NVIDIA GPU接近硬件但可移植
ONNX OpsetL4高跨框架中间格式
PyTorch FXL4高PyTorchPythonic, 易用

升华 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                        IR 设计哲学                                           │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                              │
│  1. 分层是复杂系统的唯一出路                                                   │
│     ┌──────────────────────────────────────────────────────────────────┐     │
│     │  每层只解决一类问题                                                │     │
│     │  通过 lowering 连接各层                                            │     │
│     │  层与层之间可以独立演进                                            │     │
│     └──────────────────────────────────────────────────────────────────┘     │
│                                                                              │
│  2. 抽象层级决定优化类型                                                       │
│     ┌──────────────────────────────────────────────────────────────────┐     │
│     │  高层抽象 → 图级优化(融合、子图替换)                              │     │
│     │  中层抽象 → 算法选择、布局优化、量化                                │     │
│     │  低层抽象 → 循环变换、内存层次优化                                  │     │
│     │  指令级 → 指令选择、寄存器分配、调度                                │     │
│     └──────────────────────────────────────────────────────────────────┘     │
│                                                                              │
│  3. Lowering 是信息逐步具体化的过程                                            │
│     ┌──────────────────────────────────────────────────────────────────┐     │
│     │  高层 IR 丢失信息 → 低层 IR 更具体                                  │     │
│     │  编译器必须谨慎降低,避免丢失优化机会                                │     │
│     └──────────────────────────────────────────────────────────────────┘     │
│                                                                              │
│  4. 标准化 IR 使生态互联互通                                                  │
│     ┌──────────────────────────────────────────────────────────────────┐     │
│     │  ONNX 作为中间格式实现框架互通                                      │     │
│     │  MLIR 作为基础设施实现编译器模块化                                  │     │
│     │  LLVM IR 作为通用后端实现硬件无关                                   │     │
│     └──────────────────────────────────────────────────────────────────┘     │
│                                                                              │
└─────────────────────────────────────────────────────────────────────────────┘

核心一句话: IR 层级是编译器的架构语言,多级 lowering 是连接抽象与具体的桥梁,
            选择合适的 IR 层级决定了优化的可能性与实现的复杂度。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36

"AI 可查 vs 必须理解"清单 ​

必须理解(不理解就等于不会):

  • 🔴 IR 的四层抽象:Graph-level (L4) → Op-level (L3) → Loop-level (L2) → Instruction-level (L1)
  • 🔴 为什么需要多级 IR:不同层级适合不同类型的优化,图级做融合,指令级做调度
  • 🔴 Lowering 链路:从高层语义逐步翻译到机器码的过程,PyTorch → FX → Inductor → Triton → PTX
  • 🔴 TensorFlow → XLA HLO → LLVM IR → PTX → SASS 完整链路
  • 🔴 各层 IR 的代表系统:HLO (XLA)、Relay (TVM)、LLVM IR、Linalg (MLIR)、PTX
  • 🔴 为什么 Graph-level 无法做循环优化,Instruction-level 无法做算子融合
  • 🔴 ONNX 作为中间格式的作用:实现跨框架互通

AI 可查(知道去哪查就行):

  • ✅ XLA HLO 的具体操作列表和语义
  • ✅ TVM Relay IR 的完整语法
  • ✅ LLVM IR 指令的详细格式
  • ✅ PTX 汇编语言的具体语法
  • ✅ 某个特定 IR 的 passes 列表
  • ✅ 某个硬件的 SASS 指令格式
  • ✅ 某个 IR 的 C++ API 细节

学习状态:🟡 开始学习

最后更新于:

Pager
上一篇2. 编译原理速通——面向 ML 工程师的核心概念 / Compiler Fundamentals for Machine Learning Engineers
下一篇4. 计算图的构建与表示 / Building and Representing Computational Graphs

持续记录,持续成长

Copyright © Tidenflow