192、MLIR与Triton(GPU编程语言)的对比与集成
上周五晚上十一点,我盯着屏幕上一条Triton kernel的报错发呆。错误信息指向一个奇怪的“tensor shape mismatch”,但我在Triton代码里反复检查了所有维度,逻辑上完全正确。折腾了两个小时后,我决定把同一个计算逻辑用MLIR的GPU dialect重写一遍——结果发现,Triton在编译时悄悄对某些维度做了“对齐”优化,而MLIR的lowering pass没有做同样的假设。这个坑让我意识到,Triton和MLIR虽然都号称“面向GPU的中间表示”,但它们的抽象层次和设计哲学差异,远比文档里写的要深。
抽象层次的错位:Triton是“带约束的DSL”,MLIR是“可组合的IR”
Triton本质上是一个嵌入在Python里的领域特定语言(DSL)。你写的是看起来像Python的代码,但实际描述的是“线程块级别的计算模式”。Triton编译器会帮你做线程映射、共享内存分配、warp级别同步——这些在MLIR里需要你手动用gpu.launch_func、gpu.thread_id、gpu.shared_memory等操作显式表达。
举个例子,Triton里一个简单的向量加法:
@triton