📅 创建时间: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 代码生成器的核心职责
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. 代码发射:生成最终的可执行代码
"""
pass1.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/编译器生成最终机器码 │
│ │
└─────────────────────────────────────────────────────────────────┘2.2 CUTLASS——NVIDIA 的 Template-based 典范
CUTLASS 是 NVIDIA 官方的高性能 GEMM 模板库:
// 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();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 的定位 │
│ │
└────────────────────────────────────────────────────────────────────┘第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; │ │
│ │ // ... │ │
│ │ } │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────┘3.2 TVM Tensor Expression
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, args3.3 Tensor Comprehensions
Facebook (Meta) 开发的另一种 DSL-based 方案:
# 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])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 的调用复用编译结果 │ │
│ └─────────────────────────────────────────────────────────┘ │
│ ↓ │
│ 返回编译后的函数指针 │
│ │
└─────────────────────────────────────────────────────────────────┘4.2 MLIR 的 JIT 编译
// 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;
}4.3 XLA 的 JIT 编译
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")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
│ │
└─────────────────────────────────────────────────────────────────┘第6节:混合策略——取长补短
6.1 实际系统的混合架构
现代编译器通常采用混合策略:
┌─────────────────────────────────────────────────────────────────┐
│ 混合代码生成架构 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ ┌──────────────┐ │
│ │ 用户代码 │ │
│ │ (PyTorch) │ │
│ └──────┬───────┘ │
│ ↓ │
│ ┌──────────────┐ ┌──────────────┐ │
│ │ DSL 路径 │ │ Template 路径│ │
│ │ (AutoTVM) │ │ (CUTLASS) │ │
│ └──────┬───────┘ └──────┬───────┘ │
│ ↓ ↓ │
│ ┌──────────────────────────────────────────────────────┐ │
│ │ 统一 IR (TVM Relay / MHLO) │ │
│ └──────────────────────────┬─────────────────────────────┘ │
│ ↓ │
│ ┌──────────────────────────────────────────────────────────┐ │
│ │ 统一代码生成 │ │
│ │ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐ │ │
│ │ │ CUDA │ │ ROCm │ │ CPU │ │ NPU │ │ │
│ │ │ Codegen │ │ Codegen │ │ Codegen │ │ Codegen │ │ │
│ │ └──────────┘ └──────────┘ └──────────┘ └──────────┘ │ │
│ └──────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────┘6.2 实际案例:TensorRT 的分层优化
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. 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 规则——非常细节,需要时再查
学习状态:🟡 开始学习