You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 12:30:55