053、跨Dialect优化:融合、消除死代码、常量折叠
从一次半夜的调试说起
凌晨两点,盯着终端里那个“Segmentation fault”已经快一个小时了。问题出在一个跨Dialect的优化pass上——我试图把Linalg的矩阵乘法和Arith的加法融合成一个自定义的算子,结果IR在转换过程中直接崩了。更诡异的是,同样的IR在单Dialect下跑得好好的,一跨Dialect就炸。后来发现,问题出在常量折叠阶段:一个被Arith::ConstantOp定义的常量,在Linalg的region里被当作动态值处理了,导致后续的消除死代码pass误判了依赖关系。
这种跨Dialect的坑,MLIR新手踩一次就长记性。今天这篇笔记,就聊聊跨Dialect优化里最常用的三个手段:融合、消除死代码、常量折叠。不讲教科书理论,只讲我踩过的坑和总结的经验。
融合:不是简单的“拼在一起”
很多人以为融合就是把两个op拼成一个op,比如把linalg.matmul和arith.addi合并成一个custom.matmul_add。但实际做的时候会发现,MLIR的Dialect之间往往有严格的类型系统和区域约束。
我犯过的第一个错误:直接修改op的operands列表。别这样写:
// 错误示范:直接修改operands %0 = linalg.matmul ins(%A, %B) outs(%C) %1