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 🏷️ 标签:#Codegen #代码生成 #TemplateBased #DSLBased #JIT #AOT #TemplateGen #TVM #Triton 📚 前置知识:[[12-scheduling]](调度) 📚 相关知识:[[13-codegen-architecture]](代码生成架构) [[19-tvm-te]](TVM TE) [[20-triton]](Triton)


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

┌──────────────────────────────────────────────────────────────────────────────┐ │ 情 境 描 述 │ ├──────────────────────────────────────────────────────────────────────────────┤ │ 你在研究如何给新的 AI 芯片写编译器后端。查了资料发现有三种路线: │ │ 1. Template-based:手写 C++ 模板(cuBLAS 风格) │ │ 2. DSL-based:用 Tensor Expression DSL 描述计算(TVM 风格) │ │ 3. JIT from IR:直接 lower MLIR/LLVM IR(MLIR 风格) │ │ │ │ 每种路线各有什么优缺点?你的团队应该选哪个? │ └──────────────────────────────────────────────────────────────────────────────┘

第1节:代码生成的本质——从 IR 到机器码 ​

1.1 编译流程回顾 ​

代码生成(CodeGen)是编译器的最后阶段:

高级语言/IR
    ↓
┌─────────────────────────────────────────────────────────────────┐
│                     编译器后端流程                              │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│  1. High-Level IR (HLO)                                        │
│     - 操作级抽象:MatMul, Conv2d, ReLU                          │
│     - 目标:硬件无关优化                                         │
│                                                                 │
│  2. Mid-Level IR (MILO)                                        │
│     - 循环级抽象:for i, for j, load/store                     │
│     - 目标:内存访问优化                                         │
│                                                                 │
│  3. Low-Level IR (LLO)                                         │
│     - 指令级抽象:add, mul, load, store                         │
│     - 目标:寄存器分配、指令调度                                  │
│                                                                 │
│  4. Machine IR                                                 │
│     - 目标硬件特定:x86 addq, ARM fmla                          │
│                                                                 │
│  5. Binary / Assembly                                          │
│     - 可执行机器码                                               │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
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

1.2 代码生成器的核心职责 ​

python
class CodeGenerator:
    """
    代码生成器的核心职责
    """
    
    def __init__(self, target: HardwareTarget):
        self.target = target
        self.ir = None
    
    def lower_to_target(self, op_graph: OperationGraph) -> str:
        """
        核心入口:将操作图 lower 到目标硬件
        
        步骤:
        1. 操作拆分 (Op Decomposition)
        2. 循环变换 (Loop Transformation)  
        3. 指令选择 (Instruction Selection)
        4. 寄存器分配 (Register Allocation)
        5. 指令调度 (Instruction Scheduling)
        6. 代码发射 (Code Emission)
        """
        pass
    
    def op_decomposition(self, op):
        """
        1. 操作拆分:将复杂操作拆分为基本操作
        例如:BatchNorm → Mul + Add + Sub + Div
        """
        pass
    
    def loop_transform(self, loops):
        """
        2. 循环变换:tile, unroll, fuse, interchange
        """
        pass
    
    def instruction_selection(self, expr):
        """
        3. 指令选择:为每个 IR 节点选择目标机器指令
        使用树模式匹配或 DAG 覆盖算法
        """
        pass
    
    def register_allocation(self, instructions):
        """
        4. 寄存器分配:决定每个变量放在哪个寄存器
        图着色算法或线性扫描算法
        """
        pass
    
    def instruction_scheduling(self, instructions):
        """
        5. 指令调度:决定指令执行顺序
        隐藏延迟,暴露 ILP
        """
        pass
    
    def emit_code(self, scheduled_instructions):
        """
        6. 代码发射:生成最终的可执行代码
        """
        pass
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

1.3 三种代码生成范式概览 ​

范式核心理念性能可移植性开发成本
Template-based手写模板,编译期实例化⭐⭐⭐⭐⭐⭐极高
DSL-based高级 DSL 描述,编译器生成⭐⭐⭐⭐⭐⭐⭐⭐低
JIT from IR直接 JIT 编译 IR⭐⭐⭐⭐⭐⭐⭐⭐中

第2节:Template-based——极致性能的手写之路 ​

2.1 Template-based 原理 ​

Template-based 代码生成使用预写的 C++/CUDA 模板,在编译期通过模板参数实例化生成具体代码:

┌─────────────────────────────────────────────────────────────────┐
│                    Template-based Codegen 流程                   │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│  手写 C++/CUDA 模板库                                            │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │ template<typename T, int M, int N, int K>               │   │
│  │ class GemmKernel {                                      │   │
│  │     __device__ void run(T* C, const T* A, const T* B) { │   │
│  │         // 手写的 GEMM 实现                              │   │
│  │     }                                                    │   │
│  │ };                                                       │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  编译期实例化(模板参数化)                                      │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │ GemmKernel<float, 128, 256, 64> kernel_128_256_64;      │   │
│  │ GemmKernel<float, 64, 64, 64>   kernel_64_64_64;       │   │
│  │ // 编译器生成特化版本                                      │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  NVCC/编译器生成最终机器码                                       │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24

2.2 CUTLASS——NVIDIA 的 Template-based 典范 ​

CUTLASS 是 NVIDIA 官方的高性能 GEMM 模板库:

cpp
// CUTLASS 示例:使用模板实例化 GEMM kernel
#include <cutlass/gemm/device/gemm.h>

// 定义问题尺寸
using Gemm = cutlass::gemm::device::Gemm<
    cutlass::layout::RowMajor,                    // A, B, C 的内存布局
    cutlass::tensor::Coord<2>,                   // 索引类型
    float,                                       // A 元素类型
    cutlass::layout::RowMajor,                   // A 布局
    float,                                       // B 元素类型
    cutlass::layout::ColumnMajor,                // B 布局
    float,                                       // C 元素类型
    float,                                       // 累加器类型
    cutlass::arch::OpClassTensorOp,              // 使用 Tensor Core
    cutlass::arch::Sm75,                         // 针对 Turing 架构
    cutlass::gemm::GemmShape<128, 128, 32>,      // Thread Block 形状
    cutlass::gemm::GemmShape<64, 64, 32>,       // Warp 形状
    cutlass::gemm::GemmShape<16, 8, 8>,         // MMA 指令形状
    cutlass::epilogue::thread::LinearCombination<
        float,                                   // 输出类型
        1,                                       // 输出倍数
        float,                                   // 累加器类型
        float                                    // 偏置类型
    >
>;

// 创建 problem
Gemm gemm_op;
cutlass::Status status = gemm_op.initialize(
    {128, 256, 64},      // M, N, K
    {d_A, 256},          // A 指针和 LDA
    {d_B, 64},           // B 指针和 LDB
    {d_C, 256},          // C 指针和 LDC
    {d_C, 256},          // D 指针和 LDD
    {alpha, beta}        // 缩放因子
);

// 执行
status = gemm_op.run();
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

2.3 Template-based 的优缺点 ​

优点:

优点说明
极致性能专家手写的代码,经过大量优化,接近硬件理论峰值
硬件利用率高可以精确控制 Shared Memory、Register、Tensor Core 的使用
延迟可预测手动调度,可以精确控制数据流和流水线
调试方便代码可读性好,可以直接添加 profiler 和调试语句

缺点:

缺点说明
工程量大每个操作都需要手写模板,开发周期长
新硬件要重写硬件架构改变(如 H100)可能需要完全重写
维护成本高模板库需要持续更新以支持新操作和新优化
可移植性差不同硬件需要不同的模板库

2.4 Template-based 代表作品 ​

┌────────────────────────────────────────────────────────────────────┐
│                     Template-based 代表                             │
├────────────────────────────────────────────────────────────────────┤
│                                                                    │
│  cuBLAS (NVIDIA)                                                   │
│  ├─ 闭源,高性能,NVIDIA 官方                                       │
│  ├─ 覆盖:GEMM, Conv, FFT, RNG                                     │
│  └─ 不可修改,直接调用                                              │
│                                                                    │
│  cuDNN (NVIDIA)                                                    │
│  ├─ 神经网络算子:Conv, Pooling, BN, RNN                           │
│  └─ 业界标准 benchmark                                             │
│                                                                    │
│  CUTLASS (NVIDIA)                                                  │
│  ├─ 开源模板库,用户可定制                                          │
│  ├─ 学习价值高                                                      │
│  └─ https://github.com/NVIDIA/cutlass                             │
│                                                                    │
│  MKL-DNN / oneDNN (Intel)                                         │
│  ├─ x86 CPU 优化库                                                  │
│  ├─ 支持 AVX-512, AMX                                             │
│  └─ 跨 CPU 架构                                                     │
│                                                                    │
│  MIOpen (AMD)                                                      │
│  ├─ AMD GPU 深度学习优化库                                          │
│  └─ 类似 cuDNN 的定位                                               │
│                                                                    │
└────────────────────────────────────────────────────────────────────┘
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节:DSL-based——自动优化的声明式之路 ​

3.1 DSL-based 原理 ​

DSL-based 代码生成使用高级领域特定语言(Domain Specific Language)描述计算,然后由编译器自动生成优化代码:

┌─────────────────────────────────────────────────────────────────┐
│                    DSL-based Codegen 流程                        │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│  1. 开发者使用高级 DSL 描述计算                                  │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │  # TVM Tensor Expression 示例                             │   │
│  │  A = tvm.placeholder((M, K), name='A')                   │   │
│  │  B = tvm.placeholder((K, N), name='B')                  │   │
│  │  k = tvm.reduce_axis((0, K), name='k')                  │   │
│  │  C = tvm.compute((M, N), lambda i, j:                   │   │
│  │      tvm.sum(A[i, k] * B[k, j], axis=k), name='C')      │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  2. 编译器执行优化(Schedule 变换)                               │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │  # 手动指定 schedule(调度)                             │   │
│  │  s = tvm.create_schedule(C.op)                          │   │
│  │  block_x, block_y = s[C].split(factor=32)               │   │
│  │  thread_x, thread_y = s[C].split(factor=8)              │   │
│  │  s[C].bind(block_x, tb_x)                               │   │
│  │  s[C].bind(thread_x, tid_x)                              │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  3. 编译器代码生成                                               │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │  # 生成的 CUDA 代码片段                                   │   │
│  │  __global__ void kernel(float* C, float* A, float* B) { │   │
│  │      int i = blockIdx.x * 32 + threadIdx.x;             │   │
│  │      int j = blockIdx.y * 32 + threadIdx.y;             │   │
│  │      // ...                                             │   │
│  │  }                                                      │   │
│  └─────────────────────────────────────────────────────────┘   │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
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

3.2 TVM Tensor Expression ​

python
import tvm
from tvm import te

def tvm_gemm_example():
    """
    TVM Tensor Expression 实现 GEMM
    """
    # 定义输入占位符
    M, K, N = 1024, 512, 1024
    A = te.placeholder((M, K), name='A', dtype='float32')
    B = te.placeholder((K, N), name='B', dtype='float32')
    
    # 定义计算:C[i,j] = sum_k(A[i,k] * B[k,j])
    k = te.reduce_axis((0, K), name='k')
    C = te.compute(
        (M, N),
        lambda i, j: te.sum(A[i, k] * B[k, j], axis=k),
        name='C'
    )
    
    # 创建调度(Schedule)
    s = te.create_schedule(C.op)
    
    # 分块(Tiling)
    # 将 M 和 N 维度各分成 32x32 的块
    block_size = 32
    i_outer, i_inner = s[C].split(C.op.axis[0], factor=block_size)
    j_outer, j_inner = s[C].split(C.op.axis[1], factor=block_size)
    
    # 内部循环重排(改善数据局部性)
    s[C].reorder(i_outer, j_outer, k, i_inner, j_inner)
    
    # 绑定到 GPU thread/block
    block_x = te.thread_axis("blockIdx.x")
    block_y = te.thread_axis("blockIdx.y")
    thread_x = te.thread_axis("threadIdx.x")
    thread_y = te.thread_axis("threadIdx.y")
    
    s[C].bind(i_outer, block_x)
    s[C].bind(j_outer, block_y)
    s[C].bind(i_inner, thread_x)
    s[C].bind(j_inner, thread_y)
    
    # 生成目标代码
    target = 'cuda'
    with tvm.transform.PassContext(opt_level=3):
        func = tvm.build(s, [A, B, C], target=target)
    
    return func


def tvm_schedule_exploration():
    """
    TVM 自动调度(AutoTVM / AutoScheduler)
    自动搜索最优 schedule
    """
    from tvm import auto_scheduler
    
    # 定义计算
    M, K, N = 1024, 512, 1024
    A = te.placeholder((M, K), name='A')
    B = te.placeholder((K, N), name='B')
    k = te.reduce_axis((0, K), name='k')
    C = te.compute((M, N), lambda i, j: te.sum(A[i, k] * B[k, j], axis=k))
    
    # 创建搜索任务
    task = auto_scheduler.SearchTask(
        func=lambda: C,
        args=(M, K, N),
        target='cuda'
    )
    
    # 调参器配置
    tune_option = auto_scheduler.TuningOptions(
        num_measure_trials=1000,  # 搜索 1000 个配置
        measure_callbacks=[auto_scheduler.RecordToFile('tune_log.json')],
        verbose=2
    )
    
    # 开始搜索
    task.tune(tune_option)
    
    # 应用最优 schedule
    sch, args = task.apply_best('tune_log.json')
    
    return sch, args
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

3.3 Tensor Comprehensions ​

Facebook (Meta) 开发的另一种 DSL-based 方案:

python
# Tensor Comprehensions 示例
from tensor_comprehensions import tc

# 使用声明式语言定义 GEMM
gemm = tc.define("""
def gemm(float(M,K) A, float(K,N) B) -> (C) {
    C(i, j) +=! A(i, kk) * B(kk, j)
}
""", name="gemm")

# 调用生成的 kernel
C = gemm(A, B, tiling=[32, 32, 8])
1
2
3
4
5
6
7
8
9
10
11
12

3.4 DSL-based 的优缺点 ​

优点:

优点说明
自动优化编译器自动应用各种优化(tile, unroll, fuse)
硬件可移植同一 DSL 可以 lower 到不同硬件后端
开发效率高开发者只需要描述"做什么",不需要描述"怎么做"
支持新硬件添加新后端只需实现 lower 规则

缺点:

缺点说明
性能上限低自动生成的代码难以达到手写模板的极致性能
优化难度大自动搜索空间巨大,搜索效率低
DSL 学习成本开发者需要学习新的 DSL
调试困难生成的代码难以理解和调试

第4节:JIT from IR——灵活高效的动态之路 ​

4.1 JIT from IR 原理 ​

JIT from IR 直接将 MLIR/LLVM IR 进行 JIT 编译:

┌─────────────────────────────────────────────────────────────────┐
│                    JIT from IR 流程                              │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│  输入:算子图 / 动态 shape                                        │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │  Graph:                                                    │   │
│  │    MatMul(1024, ?, 512)                                   │   │
│  │    Add(?, 256)                                            │   │
│  │    Softmax(?, ?)                                          │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  转换为 MLIR / LLVM IR                                          │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │  func @matmul(%a: tensor<?x512>, %b: tensor<?x?>) {      │   │
│  │    %c = linalg.matmul ins(%a, %b : ...) -> ...          │   │
│  │    return %c                                             │   │
│  │  }                                                       │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  JIT 编译(运行时)                                              │
│  ┌─────────────────────────────────────────────────────────┐   │
│  │  1. 类型特化:根据实际 shape 生成专门代码                   │   │
│  │  2. 优化:内联、向量化、窥孔优化                            │   │
│  │  3. 代码生成:LLVM → x86/ARM/CUDA                        │   │
│  │  4. 编译缓存:相同 shape 的调用复用编译结果                 │   │
│  └─────────────────────────────────────────────────────────┘   │
│                          ↓                                      │
│  返回编译后的函数指针                                             │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
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

4.2 MLIR 的 JIT 编译 ​

cpp
// MLIR JIT 编译示例
#include <mlir/InitAllDialects.h>
#include <mlir/ExecutionEngine/JitRunner.h>

int main() {
    mlir::DialectRegistry registry;
    registry.insert<mlir::arith::ArithDialect,
                    mlir::linalg::LinalgDialect,
                    mlir::func::FuncDialect>();
    
    mlir::MLIRContext context(registry);
    
    // 解析 MLIR 模块(可以从字符串或文件加载)
    mlir::OwningOpRef<mlir::ModuleOp> module = mlir::parseSourceString<mlir::ModuleOp>(
        R"(
            func @gemm(%A: tensor<1024x512xf32>, %B: tensor<512x1024xf32>) 
                    -> tensor<1024x1024xf32> {
                %C = linalg.matmul 
                    ins(%A, %B: tensor<1024x512xf32>, tensor<512x1024xf32>)
                    outs(%empty: tensor<1024x1024xf32>) -> tensor<1024x1024xf32>
                return %C : tensor<1024x1024xf32>
            }
        )",
        &context
    );
    
    // 创建 JIT 编译器
    mlir::JitRunnerConfig config;
    config.symbolInvoker = [](mlir::Operation* op) {
        // 符号解析回调
        return MLIR_EXPORT(__mlir_ciface_gemm);
    };
    
    // 编译并执行
    auto result = mlir::JitRunnerMain(
        std::cout,
        config,
        *module,
        /*llvmArgc=*/0,
        /*llvmArgv=*/nullptr
    );
    
    return result.succeeded() ? 0 : 1;
}
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

4.3 XLA 的 JIT 编译 ​

python
import jax
import jax.numpy as jnp

# XLA JIT 编译
@jax.jit
def gemm(A, B):
    return jnp.dot(A, B)

# 第一次调用:JIT 编译
# JAX 将 Python 函数 lower 到 HLO,然后 XLA JIT 编译
import time
start = time.time()
for _ in range(100):
    # 后续调用:直接执行编译后的代码
    C = gemm(jnp.ones((1024, 512)), jnp.ones((512, 1024)))
print(f"Time: {(time.time() - start) / 100 * 1000:.2f} ms")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16

4.4 JIT from IR 的优缺点 ​

优点:

优点说明
动态 shape 支持运行时根据实际 shape 生成专门代码
灵活性高可以与 Python/PyTorch 紧密集成
编译缓存相同 shape 的调用只需编译一次
调试友好可以保留高层语义用于调试

缺点:

缺点说明
JIT 开销首次运行需要编译,有冷启动延迟
峰值性能难以达到 Template-based 的极致性能
资源消耗JIT 编译器本身占用内存和 CPU
复杂优化难跨函数的复杂优化难以 JIT 化

第5节:三种范式的选择决策树 ​

┌─────────────────────────────────────────────────────────────────┐
│                    代码生成范式选择决策树                         │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│                          开始                                    │
│                            ↓                                     │
│                   性能是关键指标吗?                              │
│                   ↙              ↘                              │
│                 是                否                            │
│                  ↓                   ↓                           │
│          是固定 shape 吗?         使用动态 shape?                │
│           ↙        ↘               ↙         ↘                 │
│         是         否              是          否                │
│          ↓          ↓              ↓            ↓               │
│     Template     DSL-based    JIT from IR    DSL-based          │
│     based                          ↓                ↓            │
│          ↓                   首次延迟敏感?                       │
│    是否需要支持                        ↙           ↘            │
│    多种硬件?                      是       否                     │
│      ↙    ↘                      ↓        ↓                      │
│    是    否                  JIT缓存   Template                  │
│      ↓    ↓                          based +                    │
│   DSL    选择:                  JIT fallback                   │
│   based  Template +                                          
│   + DSL  DSL fallback                                   
│   lower                                                      
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
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

第6节:混合策略——取长补短 ​

6.1 实际系统的混合架构 ​

现代编译器通常采用混合策略:

┌─────────────────────────────────────────────────────────────────┐
│                    混合代码生成架构                              │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│  ┌──────────────┐                                               │
│  │  用户代码     │                                               │
│  │ (PyTorch)    │                                               │
│  └──────┬───────┘                                               │
│         ↓                                                       │
│  ┌──────────────┐     ┌──────────────┐                          │
│  │  DSL 路径    │     │  Template 路径│                          │
│  │  (AutoTVM)   │     │  (CUTLASS)    │                          │
│  └──────┬───────┘     └──────┬───────┘                          │
│         ↓                    ↓                                  │
│  ┌──────────────────────────────────────────────────────┐      │
│  │              统一 IR (TVM Relay / MHLO)                │      │
│  └──────────────────────────┬─────────────────────────────┘      │
│                             ↓                                    │
│  ┌──────────────────────────────────────────────────────────┐   │
│  │                    统一代码生成                            │   │
│  │  ┌──────────┐  ┌──────────┐  ┌──────────┐  ┌──────────┐   │   │
│  │  │  CUDA    │  │  ROCm    │  │  CPU     │  │  NPU     │   │   │
│  │  │  Codegen │  │  Codegen │  │  Codegen │  │  Codegen │   │   │
│  │  └──────────┘  └──────────┘  └──────────┘  └──────────┘   │   │
│  └──────────────────────────────────────────────────────────┘   │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
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

6.2 实际案例:TensorRT 的分层优化 ​

python
class TensorRTBuilder:
    """
    TensorRT 的混合代码生成策略
    """
    
    def build_engine(self, model):
        """
        TensorRT 的构建过程
        """
        # 1. 图优化(不需要生成代码)
        # 融合算子、消除冗余、shape 推理
        optimized_graph = self.fuse_and_optimize(model)
        
        # 2. 逐层 lower
        for layer in optimized_graph:
            if self.is_standard_op(layer):
                # 标准算子:使用预编译的 CUDA kernel
                layer.implementation = self.load_cublas_or_cudnn_kernel(layer)
            else:
                # 非标准算子:JIT 编译
                layer.implementation = self.jit_compile_layer(layer)
        
        # 3. 内存规划
        self.plan_memory(optimized_graph)
        
        # 4. 引擎序列化
        return self.serialize_engine(optimized_graph)
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

升华 ​

┌─────────────────────────────────────────────────────────────────────────────────┐ │ 代码生成的本质与选择 │ ├─────────────────────────────────────────────────────────────────────────────────┤ │ │ │ 1. 没有银弹:三种范式各有优劣,实际系统通常是混合架构 │ │ │ │ 2. Template-based 是性能极限:专家手写的模板是性能的天花板 │ │ │ │ 3. DSL-based 是可移植性的保障:一次描述,多硬件支持 │ │ │ │ 4. JIT 是动态性的钥匙:支持动态 shape 和运行时优化 │ │ │ └─────────────────────────────────────────────────────────────────────────────────┘

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

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

  • 🔴 三种代码生成范式的本质区别——Template 是"手把手教",DSL 是"告诉规则让机器学",JIT 是"运行时决定"
  • 🔴 CUTLASS 的模板参数体系——不知道就无法理解如何定制高性能 GEMM
  • 🔴 TVM 的 Schedule 机制——不知道就无法理解 DSL 如何控制代码生成
  • 🔴 混合架构的必要性——没有完美的单一方案,实际系统都是组合

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

  • ✅ 具体硬件的指令延迟数字——需要查官方文档
  • ✅ CUTLASS 的最新 API 变化——库更新频繁
  • ✅ TVM AutoScheduler 的搜索参数——每个版本略有不同
  • ✅ MLIR Dialect 的具体 lower 规则——非常细节,需要时再查

学习状态:🟡 开始学习

最后更新于:

Pager
上一篇13. 硬件约束下的操作调度 / Operation Scheduling Under Hardware Constraints
下一篇15. CPU 后端:SIMD、分块与多线程 / CPU Backends with SIMD, Tiling, and Multithreading

持续记录,持续成长

Copyright © Tidenflow