【TVM教程】转换
内容提要
TVM 0.21.0 更新后,本文介绍 Relax 程序转换。通过内置 Pass(如 LegalizeOps)将高层操作符转为低层,再经算子融合优化合并内核。同时展示自定义 Pass,用 Mutator 将 ReLU 替换为 GELU,实现灵活的程序变换。
延伸解读
Pass 是 Relax 程序转换的核心机制
在 Relax 中,Pass 是应用程序转换的主要方式。内置 Pass 如 LegalizeOps 可将高层操作符(如 matmul、add)转换为低层 call_tir 形式,便于后续优化。算子融合则通过 AnnotateTIROpPattern、FuseOps、FuseTIR 等 Pass 协作完成,将多个操作符合并为一个内核,减少内核启动开销和内存访问。理解 Pass 的流水线组织方式,是掌握 Relax 编译流程的关键。
自定义 Pass 实现灵活的程序变换
除了内置 Pass,用户还可以通过编写 Relax IR Mutator 自定义 Pass。示例中,通过继承 PyExprMutator 并重写 visit_call_ 方法,将 relu 操作符替换为 gelu,实现了对程序语义的修改。这种机制允许开发者针对特定模型或硬件进行定制优化,体现了 Relax 的可扩展性。
转换过程的可视化与验证
教程通过 mod.show() 展示了转换前后的 IRModule 结构,便于开发者直观地观察 Pass 的效果。从高层操作符到低层 call_tir,再到融合后的内核,每一步都清晰可见。这种可视化能力有助于调试和验证转换的正确性,是学习 Relax 编译流程的重要辅助手段。
Q&A
TVM 0.21.0 中 Relax 程序转换的主要方式是什么?
在 Relax 中,Pass 是应用程序转换的主要方式。通过应用内置 Pass(如 LegalizeOps)或自定义 Pass,可以改变程序的形式,实现优化和对接硬件后端。
LegalizeOps Pass 在 TVM Relax 中有什么作用?
LegalizeOps 是 Relax 中的一个内置 Pass,它可以将高层操作符(如 relax.op 中的操作)转换为低层操作符(如 relax.call_tir),从而将程序降低到更接近硬件后端的表示。
在 TVM Relax 中,算子融合是如何实现的?
在 Relax 中,算子融合是由一系列 Pass 协作完成的,通常按顺序应用 AnnotateTIROpPattern、FuseOps 和 FuseTIR 这三个 Pass。它们会分析操作符模式并融合成单个内核(call_tir),以减少内核启动开销和内存访问。
如何自定义一个 Pass 将 ReLU 替换为 GELU?
首先,需要编写一个继承自 PyExprMutator 的 Mutator(如 ReluRewriter),在 visit_call_ 方法中检测到 relax.nn.relu 时返回 relax.op.nn.gelu(call.args[0])。然后,使用 @tvm.transform.module_pass 装饰器定义一个 Pass 类(如 ReluToGelu),在 transform_module 方法中遍历模块中的函数,应用 Mutator 并更新函数。最后,将自定义 Pass 应用到模块上即可。
在 TVM Relax 中,LegalizeOps 转换后程序中的高层操作符变成了什么?
LegalizeOps 转换后,程序中的高层操作符(如 relax.op 中的 matmul、add、relu 等)被替换为对应的低层操作符 relax.call_tir,这些 call_tir 调用指向具体的 TIR 函数(如 matmul、add、relu 等)。
TVM Relax 中算子融合后,原来的多个操作符变成了什么?
算子融合后,原来的多个操作符(如 matmul、add、relu)被融合成一个内核,即一个 call_tir 调用,对应一个融合后的 TIR 函数(如 fused_matmul_add_relu),从而减少了内核启动次数和中间数据的读写。