MLIR中操作映射至外部函数的实现与自定义映射问询
MLIR Linalg Dialect 外部库调用映射详解
一、核心实现方式
Linalg的外部库映射核心依赖模式重写(Pattern Rewriting)与属性标记:
- 先给目标Linalg操作添加专属属性(比如
library_call),标记其需要映射到外部库函数 - 自定义
RewritePattern子类实现匹配逻辑:识别带标记的Linalg操作,将其替换为func.call操作,直接调用外部库函数 - 依托MLIR重写框架,通过
applyPatternsAndFoldGreedily等API触发重写流程
二、支持范围
- 原生覆盖常见线性代数操作:矩阵乘(matmul)、卷积(conv)、转置(transpose)等标准Linalg generic/named ops
- 自定义Linalg操作可通过手动添加重写模式支持,只要能在重写逻辑中匹配到操作的结构与属性
- 兼容主流数值计算库:BLAS、MKL、CuBLAS等,也支持用户私有自定义库
三、参数重映射与额外元数据添加
完全可以实现参数顺序调整、形状元数据添加的需求,核心是在重写模式中手动构造func.call的参数列表:
- 参数重映射:匹配到原Linalg操作后,直接调整操作数顺序传入
func.call,比如把原操作的u1, u2替换为u2, u1 - 添加元数据/形状参数:通过
getShape或getDimSize等API获取张量维度信息(比如numRow_of_u1就是u1的第0维大小),将这些常量值作为额外参数传入外部函数
简化版伪代码示例:
struct FooNodeToMyFooFcnPattern : public OpRewritePattern<FooNodeOp> { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(FooNodeOp op, PatternRewriter &rewriter) const override { // 获取原操作数 Value u1 = op.getOperand(0); Value u2 = op.getOperand(1); // 获取维度信息 Value numRowU1 = rewriter.create<arith.ConstantIndexOp>(op.getLoc(), u1.getType().getDimSize(0)); Value numRowU2 = rewriter.create<arith.ConstantIndexOp>(op.getLoc(), u2.getType().getDimSize(0)); // 构造新调用参数:调整顺序+添加维度 SmallVector<Value> callArgs = {u2, u1, numRowU1, numRowU2}; // 创建外部函数调用 rewriter.replaceOpWithNewOp<func.CallOp>(op, "myfooFcn", op.getResultTypes(), callArgs); return success(); } };
四、LLVM论坛提到的模式重写代码是否对应该功能?
是的,这类模式重写代码就是Linalg映射到外部库调用的核心实现。Linalg原生的库映射逻辑(比如到BLAS的映射)均通过类似的RewritePattern子类完成,论坛中的代码大概率是具体的映射案例或用户自定义的映射实现。
内容的提问来源于stack exchange,提问作者knightyangpku
相关产品推荐
相关产品推荐

