MLIR/TVM/XLA深度学习编译器深度对比与实战
·
MLIR/TVM/XLA深度学习编译器深度对比与实战
一、引言
AI 芯片百花齐放:NVIDIA GPU、Apple M系列、Google TPU、华为昇腾、高通 Hexagon…每种芯片都有独特的指令集和内存模型。"写一次,到处优化"成为奢望。深度学习编译器正是破解这一困局的关键——将高层计算图自动编译为底层高效代码。
本文将深度对比三大编译器:Google XLA、Apache TVM、LLVM MLIR,并通过实战案例展示从模型定义到自动调优的完整流程。
二、编译器分层架构
┌─────────────────────────────────┐
│ 前端 (Frontend) │ PyTorch/TensorFlow/ONNX
├─────────────────────────────────┤
│ 高层IR (HLO/Relay) │ 算子融合、图优化
├─────────────────────────────────┤
│ 中层IR (Linalg/StableHLO) │ 循环变换、内存布局
├─────────────────────────────────┤
│ 底层IR (LLVM IR/SPIR-V/PTX) │ 向量化、指令选择
├─────────────────────────────────┤
│ 后端 (Backend) │ GPU/CPU/TPU/NPU
└─────────────────────────────────┘
三、XLA (Accelerated Linear Algebra)
XLA 是 Google 开发的 JIT 编译器,深度集成于 TensorFlow/JAX。
import torch
import torch_xla
import torch_xla.core.xla_model as xm
# 方式1: PyTorch/XLA (TPU/GPU)
device = xm.xla_device()
model = MyModel().to(device)
# JIT编译
@torch.jit.script
def compiled_forward(x):
return model(x)
# 训练循环
for data in dataloader:
data = data.to(device)
output = compiled_forward(data)
loss = criterion(output, target)
loss.backward()
xm.optimizer_step(optimizer)
# 方式2: JAX (原生XLA支持)
import jax
import jax.numpy as jnp
@jax.jit # 自动编译为XLA
def train_step(params, batch):
def loss_fn(params):
logits = model.apply(params, batch['x'])
return -jnp.mean(jax.nn.log_softmax(logits) * batch['y'])
grad = jax.grad(loss_fn)(params)
return grad
# 查看编译后的HLO
print(jax.xla_computation(train_step)(params, batch).as_hlo_text())
XLA核心优化
# XLA的算子融合示例
# 原始代码:
# y = matmul(W, x)
# y = y + b
# y = relu(y)
#
# XLA将其融合为单个kernel: FusedMatMulBiasRelu
# 查看JAX的HLO IR
import jax
computation = jax.xla_computation(my_function)(x)
print(computation.as_hlo_text())
# 输出:
# HloModule ...
# %fused_computation {
# %param_0 = f32[1024,512] parameter(0)
# %param_1 = f32[512,256] parameter(1)
# %dot = f32[1024,256] dot(%param_0, %param_1)
# %broadcast = f32[1024,256] broadcast(%bias)
# %add = f32[1024,256] add(%dot, %broadcast)
# ROOT %relu = f32[1024,256] maximum(%add, 0)
# }
四、TVM (Tensor Virtual Machine)
TVM 是 Apache 开源的端到端深度学习编译器。
4.1 从模型到部署
import tvm
from tvm import relay, auto_scheduler
import tvm.contrib.graph_executor as runtime
import onnx
# 1. 导入模型(支持ONNX/PyTorch/TF/Keras)
onnx_model = onnx.load("resnet18.onnx")
mod, params = relay.frontend.from_onnx(onnx_model)
# 2. 图级别优化(算子融合、常量折叠)
with tvm.transform.PassContext(opt_level=3):
mod = relay.transform.InferType()(mod)
mod = relay.transform.FuseOps(fuse_opt_level=3)(mod) # 算子融合
mod = relay.transform.FoldConstant()(mod) # 常量折叠
mod = relay.transform.AlterOpLayout()(mod) # 布局优化
# 3. 自动调优 (AutoTVM / AutoScheduler)
target = tvm.target.Target("cuda -arch=sm_80") # A100
tasks, task_weights = auto_scheduler.extract_tasks(
mod["main"], params, target
)
tuner = auto_scheduler.TaskScheduler(tasks, task_weights)
tune_option = auto_scheduler.TuningOptions(
num_measure_trials=200,
runner=auto_scheduler.LocalRunner(repeat=10, enable_cpu_cache_flush=True),
measure_callbacks=[auto_scheduler.RecordToFile("resnet18.json")],
)
tuner.tune(tune_option)
# 4. 应用最佳调优配置
with auto_scheduler.ApplyHistoryBest("resnet18.json"):
with tvm.transform.PassContext(
opt_level=3,
config={"relay.backend.use_auto_scheduler": True}
):
lib = relay.build(mod, target=target, params=params)
# 5. 部署运行
dev = tvm.cuda(0)
module = runtime.GraphModule(lib["default"](dev))
module.set_input("input", input_data)
module.run()
output = module.get_output(0)
4.2 手动调度示例
import tvm
from tvm import te
# 矩阵乘法的手动调度
M, N, K = 1024, 1024, 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))
# 创建调度
s = te.create_schedule(C.op)
# 分块(Tiling)
block_x, block_y = 32, 32
xo, yo, xi, yi = s[C].tile(C.op.axis[0], C.op.axis[1], block_x, block_y)
# 向量化
s[C].vectorize(yi)
# 缓存(共享内存)
AA = s.cache_read(A, "shared", [C])
BB = s.cache_read(B, "shared", [C])
# 绑定到GPU
s[AA].compute_at(s[C], xo)
s[BB].compute_at(s[C], xo)
# 编译
func = tvm.build(s, [A, B, C], target="cuda")
print(func.imported_modules[0].get_source())
4.3 性能对比
| 后端(TVM编译) | ResNet50 | MobileNetV2 | BERT |
|---|---|---|---|
| PyTorch Eager | 45ms | 12ms | 85ms |
| TVM AutoScheduler | 22ms | 5.5ms | 42ms |
| TVM + TensorRT | 15ms | 4.2ms | 30ms |
| 加速比 | 3x | 2.8x | 2.8x |
五、MLIR:多层中间表示
MLIR 是 LLVM 项目的子项目,提供可组合的编译器基础设施。
// MLIR方言示例:从高层到低层
// 1. StableHLO方言(XLA兼容)
func.func @main(%arg0: tensor<1x3x224x224xf32>) -> tensor<1x1000xf32> {
%0 = stablehlo.convolution(%arg0, %filter)
dim_numbers = [b, 0, 1, f]x[0, 1, i, o]->[b, 0, 1, f],
window = {stride = [2, 2], pad = [[1, 1], [1, 1]]}
: (tensor<1x3x224x224xf32>, tensor<64x3x7x7xf32>) -> tensor<1x64x112x112xf32>
%1 = stablehlo.batch_norm_inference %0, %scale, %offset, %mean, %variance
: tensor<1x64x112x112xf32>
return %1 : tensor<1x64x112x112xf32>
}
// 2. Linalg方言(线性代数操作)
func.func @matmul(%A: memref<1024x512xf32>, %B: memref<512x256xf32>,
%C: memref<1024x256xf32>) {
linalg.matmul ins(%A, %B : memref<1024x512xf32>, memref<512x256xf32>)
outs(%C : memref<1024x256xf32>)
return
}
// 3. SCF方言(结构化控制流)
scf.for %i = %c0 to %N step %c1 {
%val = memref.load %A[%i] : memref<1024xf32>
%squared = arith.mulf %val, %val : f32
memref.store %squared, %B[%i] : memref<1024xf32>
}
MLIR Python实战
from mlir.ir import *
from mlir.dialects import func, arith, scf, memref, linalg
def build_matmul():
"""用MLIR Python API构建矩阵乘法"""
with Context() as ctx, Location.unknown():
module = Module.create()
with InsertionPoint(module.body):
M, N, K = 1024, 256, 512
# 函数定义
ftype = FunctionType.get(
[MemRefType.get([M, K], F32Type.get()),
MemRefType.get([K, N], F32Type.get()),
MemRefType.get([M, N], F32Type.get())],
[]
)
func_op = func.FuncOp("matmul", ftype)
entry_block = func_op.add_entry_block()
with InsertionPoint(entry_block):
a, b, c = entry_block.arguments
# i循环
zero = arith.ConstantOp.create_index(0)
one = arith.ConstantOp.create_index(1)
i_loop = scf.ForOp(zero, arith.ConstantOp.create_index(M), one)
with InsertionPoint(i_loop.body):
j_loop = scf.ForOp(zero, arith.ConstantOp.create_index(N), one)
with InsertionPoint(j_loop.body):
# 初始化累加器
acc = memref.AllocaOp(MemRefType.get([1], F32Type.get()), [], [])
k_loop = scf.ForOp(zero, arith.ConstantOp.create_index(K), one)
with InsertionPoint(k_loop.body):
# C[i,j] += A[i,k] * B[k,j]
a_val = memref.LoadOp(a, [i_loop.induction_variable, k_loop.induction_variable])
b_val = memref.LoadOp(b, [k_loop.induction_variable, j_loop.induction_variable])
prod = arith.MulFOp(a_val, b_val)
old = memref.LoadOp(acc, [zero])
new_val = arith.AddFOp(old, prod)
memref.StoreOp(new_val, acc, [zero])
scf.YieldOp([])
final_val = memref.LoadOp(acc, [zero])
memref.StoreOp(final_val, c, [i_loop.induction_variable, j_loop.induction_variable])
scf.YieldOp([])
scf.YieldOp([])
func.ReturnOp([])
print(module)
return module
build_matmul()
六、三大编译器对比
| 特性 | XLA | TVM | MLIR |
|---|---|---|---|
| 所属 | Apache | LLVM | |
| 目标用户 | JAX/TF开发者 | 芯片/框架厂商 | 编译器开发者 |
| 输入格式 | HLO | Relay/ONNX | 自定义方言 |
| 自动调优 | ❌ | ✅ AutoTVM/Ansor | 基础passes |
| 硬件后端 | TPU/GPU/CPU | 全平台 | 全平台 |
| 学习曲线 | 低(JAX透明) | 中 | 高(DIY) |
| 生产成熟度 | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐ | ⭐⭐⭐ |
| 典型用户 | Google内部 | 华为/阿里/字节 | Apple/Google |
七、选择建议
| 场景 | 推荐 |
|---|---|
| JAX/TF用户,TPU部署 | XLA |
| 自定义芯片/极致性能 | TVM |
| 构建新编译器/DSL | MLIR |
| NVIDIA GPU通用优化 | TVM + TensorRT |
| 移动端部署 | TVM (ARM Mali/Adreno) |
| 浏览器推理 | TVM (WebGPU/WebAssembly) |
八、总结
深度学习编译器的核心价值:
- XLA — 零配置加速,JAX/TF用户首选
- TVM — 自动调优 + 全平台覆盖,极致性能
- MLIR — 构建下一代编译器的基础设施
三者关系:XLA专注TPU生态,TVM覆盖全硬件,MLIR提供编译器构建框架。实际项目中,TVM是通用性最好的选择。
更多推荐


所有评论(0)