使用MLIR实现矩阵乘法时linalg.fill编译错误求助
问题:MLIR实现矩阵乘法时linalg.fill编译失败
我尝试通过MLIR实现矩阵乘法,但IR编译失败,报错信息如下:
./matmult.mlir:24:16: error: custom op 'linalg.fill' [parseNamedStructuredOpRegion] ods-gen generated region expects 2 args, got 0
错误出现在该行:linalg.fill(%A, %cf1) : memref<2048x2048xf64>, f64。尽管已向linalg.fill传入两个参数,却提示未获取到参数。我使用的是LLVM16,编译命令为:mlir-opt -convert-linalg-to-loops ./matmult.mlir
MLIR文件内容:
// C += A * B. func.func @matmul(%A: memref<2048x2048xf64>, %B: memref<2048x2048xf64>, %C: memref<2048x2048xf64>) { affine.for %arg3 = 0 to 2048 { affine.for %arg4 = 0 to 2048 { affine.for %arg5 = 0 to 2048 { %a = affine.load %A[%arg3, %arg5] : memref<2048x2048xf64> %b = affine.load %B[%arg5, %arg4] : memref<2048x2048xf64> %ci = affine.load %C[%arg3, %arg4] : memref<2048x2048xf64> %p = arith.mulf %a, %b : f64 %co = arith.addf %ci, %p : f64 affine.store %co, %C[%arg3, %arg4] : memref<2048x2048xf64> } } } return } func.func @main() { %A = memref.alloc() : memref<2048x2048xf64> %B = memref.alloc() : memref<2048x2048xf64> %C = memref.alloc() : memref<2048x2048xf64> %cf1 = llvm.mlir.constant(1.00000e+00 : f64) : f64 linalg.fill(%A, %cf1) : memref<2048x2048xf64>, f64 linalg.fill(%B, %cf1) : memref<2048x2048xf64>, f64 linalg.fill(%C, %cf1) : memref<2048x2048xf64>, f64 call @matmul(%A, %B, %C) : (memref<2048x2048xf64>, memref<2048x2048xf64>, memref<2048x2048xf64>) -> () call @print_memref_2d_f64(%C): (memref<2048x2048xf64>) -> () return } func.func @print_memref_2d_f64(memref<2048x2048xf64>)
原因分析
LLVM 16中Linalg结构化操作的语法发生了变化,旧版本直接按顺序传递参数的写法已不再适用。linalg.fill需要显式用ins标记输入值,outs标记输出内存引用,以此区分操作的输入输出类型。
解决方法
将所有linalg.fill的调用语句修改为新语法:
linalg.fill ins(%cf1 : f64) outs(%A : memref<2048x2048xf64>) linalg.fill ins(%cf1 : f64) outs(%B : memref<2048x2048xf64>) linalg.fill ins(%cf1 : f64) outs(%C : memref<2048x2048xf64>)
修改后的完整MLIR代码:
// C += A * B. func.func @matmul(%A: memref<2048x2048xf64>, %B: memref<2048x2048xf64>, %C: memref<2048x2048xf64>) { affine.for %arg3 = 0 to 2048 { affine.for %arg4 = 0 to 2048 { affine.for %arg5 = 0 to 2048 { %a = affine.load %A[%arg3, %arg5] : memref<2048x2048xf64> %b = affine.load %B[%arg5, %arg4] : memref<2048x2048xf64> %ci = affine.load %C[%arg3, %arg4] : memref<2048x2048xf64> %p = arith.mulf %a, %b : f64 %co = arith.addf %ci, %p : f64 affine.store %co, %C[%arg3, %arg4] : memref<2048x2048xf64> } } } return } func.func @main() { %A = memref.alloc() : memref<2048x2048xf64> %B = memref.alloc() : memref<2048x2048xf64> %C = memref.alloc() : memref<2048x2048xf64> %cf1 = llvm.mlir.constant(1.00000e+00 : f64) : f64 linalg.fill ins(%cf1 : f64) outs(%A : memref<2048x2048xf64>) linalg.fill ins(%cf1 : f64) outs(%B : memref<2048x2048xf64>) linalg.fill ins(%cf1 : f64) outs(%C : memref<2048x2048xf64>) call @matmul(%A, %B, %C) : (memref<2048x2048xf64>, memref<2048x2048xf64>, memref<2048x2048xf64>) -> () call @print_memref_2d_f64(%C): (memref<2048x2048xf64>) -> () return } func.func @print_memref_2d_f64(memref<2048x2048xf64>)
重新运行编译命令即可正常通过。
内容的提问来源于stack exchange,提问作者Roy
相关产品推荐
相关产品推荐

