linalg.matmul 语义清晰但太“裸”,无法控制 tiling、融合等,生产代码需用 linalg.generic 手动定义 indexing map、iterator types 和 region 计算逻辑。

linalg.matmul 是最直接的方式,但实际项目里几乎没人只用它——因为太“裸”,没法控制 tiling、数据布局、融合或 bufferization。真正落地时,得从 linalg.generic 入手,再配合 lowering pass 链。
为什么不用 linalg.matmul 直接写生产代码
它语义清晰,适合快速验证,但参数固定、无法嵌入自定义计算逻辑,也不暴露迭代空间细节。比如你想把 bias 加进去、做 fused relu、或者适配非默认内存布局(如 column-major),linalg.matmul 就卡死了。
常见错误现象:linalg.matmul 生成的 IR 后续 pass(如 linalg.tiling)可能不生效,因为它的 op 没暴露 loop iterator 类型和 indexing map,优化器“看不清”怎么切块。
- 它内部其实是
linalg.generic的语法糖,但封装后隐藏了 region 和映射关系 - 不支持动态形状下的 affine constraint 推导(比如 batch 维度 runtime unknown)
- 在 IREE 或 TOSA 转换流程中,常被提前 lower 成
linalg.generic,你写的没意义
linalg.generic 怎么手动写矩阵乘法
核心是三件事:定义 indexing map(描述 A、B、C 各自的访存模式)、指定 iterator types(parallel/reduction)、写 region 里的 scalar 计算逻辑。
示例(A[4x8] × B[8x4] → C[4x4]):
%0 = linalg.generic {
indexing_maps = [
affine_map (d0, d1)>, // A[i, k]
affine_map (d1, d2)>, // B[k, j]
affine_map (d0, d2)> // C[i, j]
],
iterator_types = ["parallel", "parallel", "reduction"],
doc = "A * B"
} ins(%A, %B : tensor, tensor)
outs(%C : tensor) {
^bb0(%a: f32, %b: f32, %c: f32):
%mul = arith.mulf %a, %b : f32
%add = arith.addf %c, %mul : f32
linalg.yield %add : f32
} -> tensor
注意点:
- 三个 indexing map 必须维度对齐:输入两个 map 输出维度都是 3,第三个 map 是输出 shape 的索引投影
- iterator_types 顺序必须和 indexing_maps 中每个 map 的 domain 维度数一致(这里都是 3D)
- region 内必须用
linalg.yield,且类型要和 outs 的元素类型匹配
容易踩的坑:shape、layout 和 lowering 顺序
你以为写完 linalg.generic 就能跑了?大概率会卡在 bufferization 或 vectorization 阶段。
典型报错:failed to bufferize operation 'linalg.generic': cannot bufferize a generic op with non-identity output indexing map
- 问题出在第三个 indexing map 不是 identity —— 比如你写了
affine_map (d1, d0)>(转置输出),bufferization 默认不支持,得加--linalg-bufferize=allow-returning-new-buffers - tensor 默认是 row-major,但 GPU 上常用 col-major 布局,得先用
tensor.cast或linalg.transpose预处理 - lowering 顺序不能乱:必须先
linalg.tiling(带 tile size),再linalg.fuse(如果要融合 bias),最后linalg.bufferize→convert-linalg-to-loops→convert-scf-to-cf
实际调试建议:用 mlir-opt 分步看 IR 变化
别一上来就跑 end-to-end pipeline。用下面命令逐层观察:
-
mlir-opt --linalg-tile='tile-sizes=4,4,8' input.mlir看 tiling 后的嵌套 loop 结构 -
mlir-opt --linalg-bufferize --convert-linalg-to-loops input.mlir看是否成功转成 scf.for - 加
--verify-diagnostics能让报错定位到具体 op 行号,比 silent fail 强得多
最关键的细节往往藏在 indexing map 的括号里:少一个 d、多一个逗号、map 维度不一致,整个 lowering 就静默失败——它不会报错,只是跳过该 op。











