mlir-tutorial ex7-convert Pass 部分源码分析

mlir-tutorial ex7-convert Pass 部分源码分析

1. tablegen 文件

ex7-convert/include/toy/ToyPasses.td

#ifndef TOY_PASSES_TD
#define TOY_PASSES_TD

include "mlir/Pass/PassBase.td"

def ConvertToyToArith : Pass<"convert-toy-to-arith"> {
  let summary = "Convert Toy To Arith";
  let constructor = "toy::createConvertToyToArithPass()";
  let options = [
    Option<"name", "name", "std::string", "", "help">
  ];
}

def DCE : Pass<"toy-dce", "toy::FuncOp"> {
  let summary = "dce";
  let constructor = "toy::createDCEPass()";
}
#endif

这两个 pass 都继承自:

class Pass<string passArg, string operation = "">
  : PassBase<passArg, "::mlir::OperationPass<" # operation # ">">;

生成的 .cpp.inc 文件是 build/ex7-convert/include/toy/ToyPasses.h.inc

/* Autogenerated by mlir-tblgen; don't manually edit */

#ifdef GEN_PASS_DECL
// Generate declarations for all passes.
#define GEN_PASS_DECL_CONVERTTOYTOARITH
#define GEN_PASS_DECL_DCE
#undef GEN_PASS_DECL
#endif // GEN_PASS_DECL

//===----------------------------------------------------------------------===//
// ConvertToyToArith
//===----------------------------------------------------------------------===//
#ifdef GEN_PASS_DECL_CONVERTTOYTOARITH
struct ConvertToyToArithOptions {
  std::string name;
};
#undef GEN_PASS_DECL_CONVERTTOYTOARITH
#endif // GEN_PASS_DECL_CONVERTTOYTOARITH
#ifdef GEN_PASS_DEF_CONVERTTOYTOARITH
namespace impl {

template <typename DerivedT>
class ConvertToyToArithBase : public ::mlir::OperationPass<> {
public:
  using Base = ConvertToyToArithBase;

  ConvertToyToArithBase() : ::mlir::OperationPass<>(::mlir::TypeID::get<DerivedT>()) {}
  ConvertToyToArithBase(const ConvertToyToArithBase &other) : ::mlir::OperationPass<>(other) {}

  /// Returns the command-line argument attached to this pass.
  static constexpr ::llvm::StringLiteral getArgumentName() {
    return ::llvm::StringLiteral("convert-toy-to-arith");
  }
  ::llvm::StringRef getArgument() const override { return "convert-toy-to-arith"; }

  ::llvm::StringRef getDescription() const override { return "Convert Toy To Arith"; }

  /// Returns the derived pass name.
  static constexpr ::llvm::StringLiteral getPassName() {
    return ::llvm::StringLiteral("ConvertToyToArith");
  }
  ::llvm::StringRef getName() const override { return "ConvertToyToArith"; }

  /// Support isa/dyn_cast functionality for the derived pass class.
  static bool classof(const ::mlir::Pass *pass) {
    return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
  }

  /// A clone method to create a copy of this pass.
  std::unique_ptr<::mlir::Pass> clonePass() const override {
    return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
  }

  /// Return the dialect that must be loaded in the context before this pass.
  void getDependentDialects(::mlir::DialectRegistry &registry) const override {
    
  }

  /// Explicitly declare the TypeID for this class. We declare an explicit private
  /// instantiation because Pass classes should only be visible by the current
  /// library.
  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ConvertToyToArithBase<DerivedT>)

  ConvertToyToArithBase(const ConvertToyToArithOptions &options) : ConvertToyToArithBase() {
    name = options.name;
  }
protected:
  ::mlir::Pass::Option<std::string> name{*this, "name", ::llvm::cl::desc("help")};
private:
};
} // namespace impl
#undef GEN_PASS_DEF_CONVERTTOYTOARITH
#endif // GEN_PASS_DEF_CONVERTTOYTOARITH

//===----------------------------------------------------------------------===//
// DCE
//===----------------------------------------------------------------------===//
#ifdef GEN_PASS_DECL_DCE
#undef GEN_PASS_DECL_DCE
#endif // GEN_PASS_DECL_DCE
#ifdef GEN_PASS_DEF_DCE
namespace impl {

template <typename DerivedT>
class DCEBase : public ::mlir::OperationPass<toy::FuncOp> {
public:
  using Base = DCEBase;

  DCEBase() : ::mlir::OperationPass<toy::FuncOp>(::mlir::TypeID::get<DerivedT>()) {}
  DCEBase(const DCEBase &other) : ::mlir::OperationPass<toy::FuncOp>(other) {}

  /// Returns the command-line argument attached to this pass.
  static constexpr ::llvm::StringLiteral getArgumentName() {
    return ::llvm::StringLiteral("toy-dce");
  }
  ::llvm::StringRef getArgument() const override { return "toy-dce"; }

  ::llvm::StringRef getDescription() const override { return "dce"; }

  /// Returns the derived pass name.
  static constexpr ::llvm::StringLiteral getPassName() {
    return ::llvm::StringLiteral("DCE");
  }
  ::llvm::StringRef getName() const override { return "DCE"; }

  /// Support isa/dyn_cast functionality for the derived pass class.
  static bool classof(const ::mlir::Pass *pass) {
    return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
  }

  /// A clone method to create a copy of this pass.
  std::unique_ptr<::mlir::Pass> clonePass() const override {
    return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
  }

  /// Return the dialect that must be loaded in the context before this pass.
  void getDependentDialects(::mlir::DialectRegistry &registry) const override {
    
  }

  /// Explicitly declare the TypeID for this class. We declare an explicit private
  /// instantiation because Pass classes should only be visible by the current
  /// library.
  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(DCEBase<DerivedT>)

protected:
private:
};
} // namespace impl
#undef GEN_PASS_DEF_DCE
#endif // GEN_PASS_DEF_DCE
#ifdef GEN_PASS_REGISTRATION

//===----------------------------------------------------------------------===//
// ConvertToyToArith Registration
//===----------------------------------------------------------------------===//

inline void registerConvertToyToArith() {
  ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
    return toy::createConvertToyToArithPass();
  });
}

// Old registration code, kept for temporary backwards compatibility.
inline void registerConvertToyToArithPass() {
  ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
    return toy::createConvertToyToArithPass();
  });
}

//===----------------------------------------------------------------------===//
// DCE Registration
//===----------------------------------------------------------------------===//

inline void registerDCE() {
  ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
    return toy::createDCEPass();
  });
}

// Old registration code, kept for temporary backwards compatibility.
inline void registerDCEPass() {
  ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
    return toy::createDCEPass();
  });
}

//===----------------------------------------------------------------------===//
//  Registration
//===----------------------------------------------------------------------===//

inline void registerPasses() {
  registerConvertToyToArith();
  registerDCE();
}
#undef GEN_PASS_REGISTRATION
#endif // GEN_PASS_REGISTRATION
// Deprecated. Please use the new per-pass macros.
#ifdef GEN_PASS_CLASSES

template <typename DerivedT>
class ConvertToyToArithBase : public ::mlir::OperationPass<> {
public:
  using Base = ConvertToyToArithBase;

  ConvertToyToArithBase() : ::mlir::OperationPass<>(::mlir::TypeID::get<DerivedT>()) {}
  ConvertToyToArithBase(const ConvertToyToArithBase &other) : ::mlir::OperationPass<>(other) {}

  /// Returns the command-line argument attached to this pass.
  static constexpr ::llvm::StringLiteral getArgumentName() {
    return ::llvm::StringLiteral("convert-toy-to-arith");
  }
  ::llvm::StringRef getArgument() const override { return "convert-toy-to-arith"; }

  ::llvm::StringRef getDescription() const override { return "Convert Toy To Arith"; }

  /// Returns the derived pass name.
  static constexpr ::llvm::StringLiteral getPassName() {
    return ::llvm::StringLiteral("ConvertToyToArith");
  }
  ::llvm::StringRef getName() const override { return "ConvertToyToArith"; }

  /// Support isa/dyn_cast functionality for the derived pass class.
  static bool classof(const ::mlir::Pass *pass) {
    return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
  }

  /// A clone method to create a copy of this pass.
  std::unique_ptr<::mlir::Pass> clonePass() const override {
    return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
  }

  /// Register the dialects that must be loaded in the context before this pass.
  void getDependentDialects(::mlir::DialectRegistry &registry) const override {
    
  }

  /// Explicitly declare the TypeID for this class. We declare an explicit private
  /// instantiation because Pass classes should only be visible by the current
  /// library.
  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ConvertToyToArithBase<DerivedT>)

protected:
  ::mlir::Pass::Option<std::string> name{*this, "name", ::llvm::cl::desc("help")};
};

template <typename DerivedT>
class DCEBase : public ::mlir::OperationPass<toy::FuncOp> {
public:
  using Base = DCEBase;

  DCEBase() : ::mlir::OperationPass<toy::FuncOp>(::mlir::TypeID::get<DerivedT>()) {}
  DCEBase(const DCEBase &other) : ::mlir::OperationPass<toy::FuncOp>(other) {}

  /// Returns the command-line argument attached to this pass.
  static constexpr ::llvm::StringLiteral getArgumentName() {
    return ::llvm::StringLiteral("toy-dce");
  }
  ::llvm::StringRef getArgument() const override { return "toy-dce"; }

  ::llvm::StringRef getDescription() const override { return "dce"; }

  /// Returns the derived pass name.
  static constexpr ::llvm::StringLiteral getPassName() {
    return ::llvm::StringLiteral("DCE");
  }
  ::llvm::StringRef getName() const override { return "DCE"; }

  /// Support isa/dyn_cast functionality for the derived pass class.
  static bool classof(const ::mlir::Pass *pass) {
    return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
  }

  /// A clone method to create a copy of this pass.
  std::unique_ptr<::mlir::Pass> clonePass() const override {
    return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
  }

  /// Register the dialects that must be loaded in the context before this pass.
  void getDependentDialects(::mlir::DialectRegistry &registry) const override {
    
  }

  /// Explicitly declare the TypeID for this class. We declare an explicit private
  /// instantiation because Pass classes should only be visible by the current
  /// library.
  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(DCEBase<DerivedT>)

protected:
};
#undef GEN_PASS_CLASSES
#endif // GEN_PASS_CLASSES

这个 inc 文件可以按照宏拆解成若干部分:

  • GEN_PASS_DECL:定义 GEN_PASS_DECL_CONVERTTOYTOARITHGEN_PASS_DECL_DCE,是该 inc 文件作为对外头文件使用

    #ifdef GEN_PASS_DECL
    // Generate declarations for all passes.
    #define GEN_PASS_DECL_CONVERTTOYTOARITH
    #define GEN_PASS_DECL_DCE
    #undef GEN_PASS_DECL
    #endif // GEN_PASS_DECL
    
    • GEN_PASS_DECL_CONVERTTOYTOARITH:ConvertToyToArith pass 对应的头文件

      #ifdef GEN_PASS_DECL_CONVERTTOYTOARITH
      struct ConvertToyToArithOptions {
        std::string name;
      };
      #undef GEN_PASS_DECL_CONVERTTOYTOARITH
      #endif // GEN_PASS_DECL_CONVERTTOYTOARITH
      
    • GEN_PASS_DECL_DCE:DCE 没有对外的头文件,所以这里没有内容

      #ifdef GEN_PASS_DECL_DCE
      #undef GEN_PASS_DECL_DCE
      #endif // GEN_PASS_DECL_DCE
      
  • GEN_PASS_DEF_CONVERTTOYTOARITH:ConvertToyToArith pass 对应的具体实现

    #ifdef GEN_PASS_DEF_CONVERTTOYTOARITH
    namespace impl {
    
    template <typename DerivedT>
    class ConvertToyToArithBase : public ::mlir::OperationPass<> {
    public:
      using Base = ConvertToyToArithBase;
    
      ConvertToyToArithBase() : ::mlir::OperationPass<>(::mlir::TypeID::get<DerivedT>()) {}
      ConvertToyToArithBase(const ConvertToyToArithBase &other) : ::mlir::OperationPass<>(other) {}
    
      /// Returns the command-line argument attached to this pass.
      static constexpr ::llvm::StringLiteral getArgumentName() {
        return ::llvm::StringLiteral("convert-toy-to-arith");
      }
      ::llvm::StringRef getArgument() const override { return "convert-toy-to-arith"; }
    
      ::llvm::StringRef getDescription() const override { return "Convert Toy To Arith"; }
    
      /// Returns the derived pass name.
      static constexpr ::llvm::StringLiteral getPassName() {
        return ::llvm::StringLiteral("ConvertToyToArith");
      }
      ::llvm::StringRef getName() const override { return "ConvertToyToArith"; }
    
      /// Support isa/dyn_cast functionality for the derived pass class.
      static bool classof(const ::mlir::Pass *pass) {
        return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
      }
    
      /// A clone method to create a copy of this pass.
      std::unique_ptr<::mlir::Pass> clonePass() const override {
        return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
      }
    
      /// Return the dialect that must be loaded in the context before this pass.
      void getDependentDialects(::mlir::DialectRegistry &registry) const override {
        
      }
    
      /// Explicitly declare the TypeID for this class. We declare an explicit private
      /// instantiation because Pass classes should only be visible by the current
      /// library.
      MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ConvertToyToArithBase<DerivedT>)
    
      ConvertToyToArithBase(const ConvertToyToArithOptions &options) : ConvertToyToArithBase() {
        name = options.name;
      }
    protected:
      ::mlir::Pass::Option<std::string> name{*this, "name", ::llvm::cl::desc("help")};
    private:
    };
    } // namespace impl
    #undef GEN_PASS_DEF_CONVERTTOYTOARITH
    #endif // GEN_PASS_DEF_CONVERTTOYTOARITH
    
  • GEN_PASS_DEF_DCE:DCE pass 对应的具体实现

    #ifdef GEN_PASS_DEF_DCE
    namespace impl {
    
    template <typename DerivedT>
    class DCEBase : public ::mlir::OperationPass<toy::FuncOp> {
    public:
      using Base = DCEBase;
    
      DCEBase() : ::mlir::OperationPass<toy::FuncOp>(::mlir::TypeID::get<DerivedT>()) {}
      DCEBase(const DCEBase &other) : ::mlir::OperationPass<toy::FuncOp>(other) {}
    
      /// Returns the command-line argument attached to this pass.
      static constexpr ::llvm::StringLiteral getArgumentName() {
        return ::llvm::StringLiteral("toy-dce");
      }
      ::llvm::StringRef getArgument() const override { return "toy-dce"; }
    
      ::llvm::StringRef getDescription() const override { return "dce"; }
    
      /// Returns the derived pass name.
      static constexpr ::llvm::StringLiteral getPassName() {
        return ::llvm::StringLiteral("DCE");
      }
      ::llvm::StringRef getName() const override { return "DCE"; }
    
      /// Support isa/dyn_cast functionality for the derived pass class.
      static bool classof(const ::mlir::Pass *pass) {
        return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
      }
    
      /// A clone method to create a copy of this pass.
      std::unique_ptr<::mlir::Pass> clonePass() const override {
        return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
      }
    
      /// Return the dialect that must be loaded in the context before this pass.
      void getDependentDialects(::mlir::DialectRegistry &registry) const override {
        
      }
    
      /// Explicitly declare the TypeID for this class. We declare an explicit private
      /// instantiation because Pass classes should only be visible by the current
      /// library.
      MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(DCEBase<DerivedT>)
    
    protected:
    private:
    };
    } // namespace impl
    #undef GEN_PASS_DEF_DCE
    #endif // GEN_PASS_DEF_DCE
    
  • GEN_PASS_REGISTRATION:两个 pass 的注册函数

    #ifdef GEN_PASS_REGISTRATION
    
    //===----------------------------------------------------------------------===//
    // ConvertToyToArith Registration
    //===----------------------------------------------------------------------===//
    
    inline void registerConvertToyToArith() {
      ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
        return toy::createConvertToyToArithPass();
      });
    }
    
    // Old registration code, kept for temporary backwards compatibility.
    inline void registerConvertToyToArithPass() {
      ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
        return toy::createConvertToyToArithPass();
      });
    }
    
    //===----------------------------------------------------------------------===//
    // DCE Registration
    //===----------------------------------------------------------------------===//
    
    inline void registerDCE() {
      ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
        return toy::createDCEPass();
      });
    }
    
    // Old registration code, kept for temporary backwards compatibility.
    inline void registerDCEPass() {
      ::mlir::registerPass([]() -> std::unique_ptr<::mlir::Pass> {
        return toy::createDCEPass();
      });
    }
    
    //===----------------------------------------------------------------------===//
    //  Registration
    //===----------------------------------------------------------------------===//
    
    inline void registerPasses() {
      registerConvertToyToArith();
      registerDCE();
    }
    #undef GEN_PASS_REGISTRATION
    #endif // GEN_PASS_REGISTRATION
    
  • GEN_PASS_CLASSES:Deprecated,兼容旧的使用方法,生成所有 pass 的实现

    #ifdef GEN_PASS_CLASSES
    
    template <typename DerivedT>
    class ConvertToyToArithBase : public ::mlir::OperationPass<> {
    public:
      using Base = ConvertToyToArithBase;
    
      ConvertToyToArithBase() : ::mlir::OperationPass<>(::mlir::TypeID::get<DerivedT>()) {}
      ConvertToyToArithBase(const ConvertToyToArithBase &other) : ::mlir::OperationPass<>(other) {}
    
      /// Returns the command-line argument attached to this pass.
      static constexpr ::llvm::StringLiteral getArgumentName() {
        return ::llvm::StringLiteral("convert-toy-to-arith");
      }
      ::llvm::StringRef getArgument() const override { return "convert-toy-to-arith"; }
    
      ::llvm::StringRef getDescription() const override { return "Convert Toy To Arith"; }
    
      /// Returns the derived pass name.
      static constexpr ::llvm::StringLiteral getPassName() {
        return ::llvm::StringLiteral("ConvertToyToArith");
      }
      ::llvm::StringRef getName() const override { return "ConvertToyToArith"; }
    
      /// Support isa/dyn_cast functionality for the derived pass class.
      static bool classof(const ::mlir::Pass *pass) {
        return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
      }
    
      /// A clone method to create a copy of this pass.
      std::unique_ptr<::mlir::Pass> clonePass() const override {
        return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
      }
    
      /// Register the dialects that must be loaded in the context before this pass.
      void getDependentDialects(::mlir::DialectRegistry &registry) const override {
        
      }
    
      /// Explicitly declare the TypeID for this class. We declare an explicit private
      /// instantiation because Pass classes should only be visible by the current
      /// library.
      MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ConvertToyToArithBase<DerivedT>)
    
    protected:
      ::mlir::Pass::Option<std::string> name{*this, "name", ::llvm::cl::desc("help")};
    };
    
    template <typename DerivedT>
    class DCEBase : public ::mlir::OperationPass<toy::FuncOp> {
    public:
      using Base = DCEBase;
    
      DCEBase() : ::mlir::OperationPass<toy::FuncOp>(::mlir::TypeID::get<DerivedT>()) {}
      DCEBase(const DCEBase &other) : ::mlir::OperationPass<toy::FuncOp>(other) {}
    
      /// Returns the command-line argument attached to this pass.
      static constexpr ::llvm::StringLiteral getArgumentName() {
        return ::llvm::StringLiteral("toy-dce");
      }
      ::llvm::StringRef getArgument() const override { return "toy-dce"; }
    
      ::llvm::StringRef getDescription() const override { return "dce"; }
    
      /// Returns the derived pass name.
      static constexpr ::llvm::StringLiteral getPassName() {
        return ::llvm::StringLiteral("DCE");
      }
      ::llvm::StringRef getName() const override { return "DCE"; }
    
      /// Support isa/dyn_cast functionality for the derived pass class.
      static bool classof(const ::mlir::Pass *pass) {
        return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
      }
    
      /// A clone method to create a copy of this pass.
      std::unique_ptr<::mlir::Pass> clonePass() const override {
        return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
      }
    
      /// Register the dialects that must be loaded in the context before this pass.
      void getDependentDialects(::mlir::DialectRegistry &registry) const override {
        
      }
    
      /// Explicitly declare the TypeID for this class. We declare an explicit private
      /// instantiation because Pass classes should only be visible by the current
      /// library.
      MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(DCEBase<DerivedT>)
    
    protected:
    };
    #undef GEN_PASS_CLASSES
    #endif // GEN_PASS_CLASSES
    

简单来说,对应关系如下:

  • GEN_PASS_DECL -> GEN_PASS_DECL_CONVERTTOYTOARITHGEN_PASS_DECL_DCE,用于放在头文件中使用
    • GEN_PASS_DECL_CONVERTTOYTOARITH -> 生成 ConvertToyToArithOptions 类作为 ConvertToyToArith pass 的对外接口
    • GEN_PASS_DECL_DCE -> 生成 DCE pass 的对外接口,本例为空
  • GEN_PASS_DEF_CONVERTTOYTOARITH -> 生成 ConvertToyToArithBase 类实现
  • GEN_PASS_DEF_DCE -> 生成 DCEBase 类实现
  • GEN_PASS_REGISTRATION -> 生成两个 pass 的注册代码
  • GEN_PASS_CLASSES -> 旧的使用方法,生成两个 pass 类的实现

这里有几点需要注意:

  • 为什么不直接生成类似 ConvertToyToArith/DCE 类,而是 ConvertToyToArithBase/DCEBase 类?
    TableGen 负责生成的是接入 MLIR 系统所需的样板代码,业务逻辑如 void runOnOperation() override { ... } 在继承自 TableGen 生成的 Base 类的子类中实现。

    • 什么叫接入 mlir 系统呢?
      首先是上述两个 pass 的继承链:::mlir::Pass -> ::mlir::OperationPass<> -> toy::impl::ConvertToyToArithBase<> -> ConvertToyToArithPass。MLIR 内部在管理使用 pass 的时候,所有 pass 作为 Pass * 使用,Pass * 基类同时提供其它必要的方法。OperationPass<> 基类为子类提供必要的运行在 operation 上的逻辑,如 classof() 方法返回该 operation 是否是某个类型,用于支持 llvm 中大量使用的 isa/dyn_cast 方法。getOperation() 方法直接返回对应的 operation(作为对比,比如 InterfacePass 类同名方法返回的是 interface)。TableGen 生成的 toy::impl::ConvertToyToArithBase<> 实现了 Pass 基类定义的一些虚方法如:getName()getArgument()getDescription(),这些方法的实现都依赖 td 文件中提供的信息。
  • CRTP 模式的使用和 impl namespace 架构约定

    • CRTP:基类保留模板参数,子类继承时传递自身的类型作为基类的模板参数。作用:

      • 用来取代虚函数,节约虚函数查表时间,消除重复逻辑

        static bool classof(const ::mlir::Pass *pass) {
            return pass->getTypeID() == ::mlir::TypeID::get<DerivedT>();
        }
        

        如果采用虚函数,则需要每个子类单独实现一下该方法,并且由于是虚函数,还有查表开销,llvm 中大量使用了 isa/dyn_cast 带来性能损失。

      • 防止对象切片问题

        std::unique_ptr<::mlir::Pass> clonePass() const override {
            return std::make_unique<DerivedT>(*static_cast<const DerivedT *>(this));
        }
        

        如果 cast 成自己的类型 ConvertToyToArithBase 会导致部分属性和方法丢失,并且 vptr 指向 Base 类的虚表,丢失多态性。

    • impl namespace:在 cpp 工程语境里 impl namespace 表达了一种类似 pimpl 设计模式的暗示,作为一个类的内部实现,不要去改动它,否则会造成错误。在 mlir 工程里就算你去改 Base 类也会由于重新构建导致改动丢失。

2. 头文件

ex7-convert/include/toy/ToyPasses.h

#pragma once

#include "mlir/Pass/Pass.h"
#include "mlir/Pass/PassRegistry.h"
#include "toy/ToyOps.h"
#include <memory>

namespace toy {

#define GEN_PASS_DECL
#include "toy/ToyPasses.h.inc"

std::unique_ptr<mlir::Pass> createConvertToyToArithPass(ConvertToyToArithOptions options={});
std::unique_ptr<mlir::Pass> createDCEPass();

#define GEN_PASS_REGISTRATION
#include "toy/ToyPasses.h.inc"

}

该头文件做了三件事:

  • 引入每个 pass 对应的头文件

  • 为 td 文件中定义的 constructor 工厂函数属性写个 cpp 的声明,这个函数被Pass 管理器或命令行工具在需要实例化 pass 时使用

    def ConvertToyToArith : Pass<"convert-toy-to-arith"> {
      let summary = "Convert Toy To Arith";
      let constructor = "toy::createConvertToyToArithPass()";
      let options = [
        Option<"name", "name", "std::string", "", "help">
      ];
    }
    
    def DCE : Pass<"toy-dce", "toy::FuncOp"> {
      let summary = "dce";
      let constructor = "toy::createDCEPass()";
    }
    
    • 如果不写 constructor 属性也会自动生成对应工厂函数,放在 inc 文件对应 pass 的头文件部分
      #ifdef GEN_PASS_DECL_CONVERTTOYTOARITH
      struct ConvertToyToArithOptions {
        std::string name;
      };
      std::unique_ptr<::mlir::Pass> createConvertToyToArith();
      std::unique_ptr<::mlir::Pass> createConvertToyToArith(const ConvertToyToArithOptions &options);
      #undef GEN_PASS_DECL_CONVERTTOYTOARITH
      #endif // GEN_PASS_DECL_CONVERTTOYTOARITH
      
  • 引入 pass 的注册函数·

3. 实现文件

ex7-convert/lib/Transforms/ConvertToyToArith.cpp

#include "mlir/IR/BuiltinDialect.h"
#include "mlir/IR/BuiltinOps.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/Support/LogicalResult.h"
#include "toy/ToyDialect.h"
#include "toy/ToyOps.h"
#include "toy/ToyTypes.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "llvm/ADT/STLExtras.h"
#include "llvm/ADT/SmallVector.h"
#include "llvm/Support/raw_ostream.h"
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
#define GEN_PASS_DEF_CONVERTTOYTOARITH
#include "toy/ToyPasses.h"

#include "mlir/Transforms/DialectConversion.h"

using namespace mlir;
using namespace llvm;
using namespace toy;

struct AddOpPat: OpConversionPattern<AddOp> {
  using OpConversionPattern<AddOp>::OpConversionPattern;
  LogicalResult matchAndRewrite(AddOp op, AddOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
    auto inputs = to_vector(adaptor.getInputs());
    auto result = inputs[0];
    for(size_t i = 1; i< inputs.size(); i++) {
      assert(inputs[i]);
      result = rewriter.create<arith::AddIOp>(op->getLoc(), result, inputs[i]);
    }
    rewriter.replaceOp(op, ValueRange(result));
    return success();
  }
};

struct SubOpPat: OpConversionPattern<SubOp> {
  using OpConversionPattern<SubOp>::OpConversionPattern;
  LogicalResult matchAndRewrite(SubOp op, SubOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
    rewriter.replaceOpWithNewOp<arith::SubIOp>(op, adaptor.getLhs(), adaptor.getRhs());
    return success();
  }
};

struct ConstantOpPat: OpConversionPattern<ConstantOp> {
  using OpConversionPattern<ConstantOp>::OpConversionPattern;
  LogicalResult matchAndRewrite(ConstantOp op, ConstantOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
    rewriter.replaceOpWithNewOp<arith::ConstantOp>(op, op.getValueAttr());
    return success();
  }
};

struct ReturnOpPat: OpConversionPattern<ReturnOp> {
  using OpConversionPattern<ReturnOp>::OpConversionPattern;
  LogicalResult matchAndRewrite(ReturnOp op, ReturnOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
    auto data = adaptor.getData();
    rewriter.startOpModification(op);
    op.getDataMutable().assign(data);
    rewriter.finalizeOpModification(op);
    return success();
  }
};

struct CallOpPat: OpConversionPattern<CallOp> {
  using OpConversionPattern<CallOp>::OpConversionPattern;
  LogicalResult matchAndRewrite(CallOp op, CallOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
    SmallVector<Type> resTypes;
    assert(succeeded(getTypeConverter()->convertTypes(op->getResultTypes(), resTypes)));
    rewriter.replaceOpWithNewOp<CallOp>(op, resTypes, op.getCallee(), adaptor.getOperands());
    return success();
  }
};

struct ConvertToyToArithPass : toy::impl::ConvertToyToArithBase<ConvertToyToArithPass> {
  using toy::impl::ConvertToyToArithBase<ConvertToyToArithPass>::ConvertToyToArithBase;
  void getDependentDialects(DialectRegistry &registry) const final {
    registry.insert<arith::ArithDialect>();
  }
  void runOnOperation() final {
    ConversionTarget target(getContext());
    target.addLegalDialect<arith::ArithDialect>();
    // target.addDynamicallyLegalOp<FuncOp>([](FuncOp f) {
    //   return llvm::all_of(f.getArgumentTypes(), [](Type t) {return !isa<ToyIntegerType>(t);});
    // });
    auto checkValid = [](Operation* f) {
      return llvm::all_of(f->getOperandTypes(), [](Type t) {return !isa<ToyIntegerType>(t);});
    };
    target.addDynamicallyLegalOp<ReturnOp, CallOp>(checkValid);
    TypeConverter converter;
    converter.addConversion([&](ToyIntegerType t) -> std::optional<IntegerType> {
      return IntegerType::get(&getContext(), t.getWidth());
    });
    converter.addTargetMaterialization([](OpBuilder& builder, Type resultType, ValueRange inputs, Location loc) -> std::optional<Value> {
      return builder.create<UnrealizedConversionCastOp>(loc, resultType, inputs).getResult(0);
    });
    RewritePatternSet patterns(&getContext());
    patterns.add<AddOpPat, SubOpPat, ConstantOpPat, ReturnOpPat, CallOpPat>(converter, &getContext());
    populateFunctionOpInterfaceTypeConversionPattern<FuncOp>(patterns, converter);
    if(failed(applyPartialConversion(getOperation(), target, std::move(patterns))))
      signalPassFailure();
  }
};

std::unique_ptr<mlir::Pass> toy::createConvertToyToArithPass(ConvertToyToArithOptions options) {
  return std::make_unique<ConvertToyToArithPass>(options);
}

这里拆解成三部分:

  • 定义要 convert 的 pattern:

    • AddOp:

      struct AddOpPat: OpConversionPattern<AddOp> {
        using OpConversionPattern<AddOp>::OpConversionPattern;
        LogicalResult matchAndRewrite(AddOp op, AddOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
          auto inputs = to_vector(adaptor.getInputs());
          auto result = inputs[0];
          for(size_t i = 1; i< inputs.size(); i++) {
            assert(inputs[i]);
            result = rewriter.create<arith::AddIOp>(op->getLoc(), result, inputs[i]);
          }
          rewriter.replaceOp(op, ValueRange(result));
          return success();
        }
      };
      

      这里的继承链是:

      AddOpPat -> OpConversionPattern<AddOp> -> ConversionPattern -> RewritePattern -> Pattern
      

      Pattern 包含一个模式所需要的全部元数据,但是不包含具体的模式,子类继承该类描述具体的模式:

      class Pattern {
        enum class RootKind {
          Any,
          OperationName,
          InterfaceID,
          TraitID
        };
      
      public:
        ArrayRef<OperationName> getGeneratedOps() const { return generatedOps; }
      
        std::optional<OperationName> getRootKind() const {
          if (rootKind == RootKind::OperationName)
            return OperationName::getFromOpaquePointer(rootValue);
          return std::nullopt;
        }
      
        std::optional<TypeID> getRootInterfaceID() const {
          if (rootKind == RootKind::InterfaceID)
            return TypeID::getFromOpaquePointer(rootValue);
          return std::nullopt;
        }
      
        std::optional<TypeID> getRootTraitID() const {
          if (rootKind == RootKind::TraitID)
            return TypeID::getFromOpaquePointer(rootValue);
          return std::nullopt;
        }
      
        PatternBenefit getBenefit() const { return benefit; }
      
        bool hasBoundedRewriteRecursion() const {
          return contextAndHasBoundedRecursion.getInt();
        }
      
        MLIRContext *getContext() const {
          return contextAndHasBoundedRecursion.getPointer();
        }
      
        StringRef getDebugName() const { return debugName; }
      
        void setDebugName(StringRef name) { debugName = name; }
      
        ArrayRef<StringRef> getDebugLabels() const { return debugLabels; }
      
        void addDebugLabels(ArrayRef<StringRef> labels) {
          debugLabels.append(labels.begin(), labels.end());
        }
        void addDebugLabels(StringRef label) { debugLabels.push_back(label); }
      
      protected:
        struct MatchAnyOpTypeTag {};
      
        struct MatchInterfaceOpTypeTag {};
      
        struct MatchTraitOpTypeTag {};
      
      
        Pattern(StringRef rootName, PatternBenefit benefit, MLIRContext *context,
                ArrayRef<StringRef> generatedNames = {});
      
        Pattern(MatchAnyOpTypeTag tag, PatternBenefit benefit, MLIRContext *context,
                ArrayRef<StringRef> generatedNames = {});
      
        Pattern(MatchInterfaceOpTypeTag tag, TypeID interfaceID,
                PatternBenefit benefit, MLIRContext *context,
                ArrayRef<StringRef> generatedNames = {});
      
        Pattern(MatchTraitOpTypeTag tag, TypeID traitID, PatternBenefit benefit,
                MLIRContext *context, ArrayRef<StringRef> generatedNames = {});
      
        void setHasBoundedRewriteRecursion(bool hasBoundedRecursionArg = true) {
          contextAndHasBoundedRecursion.setInt(hasBoundedRecursionArg);
        }
      
      private:
        Pattern(const void *rootValue, RootKind rootKind,
                ArrayRef<StringRef> generatedNames, PatternBenefit benefit,
                MLIRContext *context);
      
        const void *rootValue;
        RootKind rootKind;
        const PatternBenefit benefit;
        llvm::PointerIntPair<MLIRContext *, 1, bool> contextAndHasBoundedRecursion;
        SmallVector<OperationName, 2> generatedOps;
        StringRef debugName;
        SmallVector<StringRef, 0> debugLabels;
      };
      

      一些属性的解析:

      • RootKind rootKind:描述该模式是匹配任意类型/操作名/接口/特性
      • llvm::PointerIntPair<MLIRContext *, 1, bool> contextAndHasBoundedRecursion:这个属性描述的是该模式是否是有界的,引擎能反复应用它,不必担心无限循环。llvm::PointerIntPair<MLIRContext *, 1, bool> 逻辑上算是两个属性,这个数据结构实现使用 64 bit 存储一个指针 + 一个 bool 值,节约 64 bit 空间(因为 1 bit 的 bool 值会由于内存对齐扩展到 8 byte = 64 bit)
      • SmallVector<OperationName, 2> generatedOps:描述这个 pattern 可能会替换出来的 operator,用于合法性检查

      Pattern 的方法大都是 getter/setter,不做介绍。

      RewritePattern 建立在 Pattern 之上,提出了 matchrewritematchAndRritecreate 方法:

      class RewritePattern : public Pattern {
      public:
        virtual ~RewritePattern() = default;
      
        virtual void rewrite(Operation *op, PatternRewriter &rewriter) const;
      
        virtual LogicalResult match(Operation *op) const;
      
        virtual LogicalResult matchAndRewrite(Operation *op,
                                              PatternRewriter &rewriter) const {
          if (succeeded(match(op))) {
            rewrite(op, rewriter);
            return success();
          }
          return failure();
        }
      
        template <typename T, typename... Args>
        static std::unique_ptr<T> create(Args &&...args) {
          std::unique_ptr<T> pattern =
              std::make_unique<T>(std::forward<Args>(args)...);
          initializePattern<T>(*pattern);
      
          if (pattern->getDebugName().empty())
            pattern->setDebugName(llvm::getTypeName<T>());
          return pattern;
        }
      
      protected:
        using Pattern::Pattern;
      
      private:
        template <typename T, typename... Args>
        using has_initialize = decltype(std::declval<T>().initialize());
        template <typename T>
        using detect_has_initialize = llvm::is_detected<has_initialize, T>;
      
        template <typename T>
        static std::enable_if_t<detect_has_initialize<T>::value>
        initializePattern(T &pattern) {
          pattern.initialize();
        }
      
        template <typename T>
        static std::enable_if_t<!detect_has_initialize<T>::value>
        initializePattern(T &) {}
      
        virtual void anchor();
      };
      
      

      matchrewritematchAndRrite 比较容易理解,就是用来进行匹配以及重写。create 主要用于注册一个 pattern 时,安全的实例化一个 pattern 对象,而非手动 new 或者 make_unique,因为有的 pattern 还需要进行额外的 initialize。

      ConvertionPatternRewritePattern 的基础上提供了自动类型转换的功能,mlir 框架会利用 typeConverter 将需要转换类型的操作数转换成正确的类型,放入 ArrayRef<Value> operands 中,而无需在 rewrite 中手动为操作数转化类型:

      class ConversionPattern : public RewritePattern {
      public:
        virtual void rewrite(Operation *op, ArrayRef<Value> operands,
                             ConversionPatternRewriter &rewriter) const {
          llvm_unreachable("unimplemented rewrite");
        }
      
        virtual LogicalResult
        matchAndRewrite(Operation *op, ArrayRef<Value> operands,
                        ConversionPatternRewriter &rewriter) const {
          if (failed(match(op)))
            return failure();
          rewrite(op, operands, rewriter);
          return success();
      
        LogicalResult matchAndRewrite(Operation *op,
                                      PatternRewriter &rewriter) const final;
      
        const TypeConverter *getTypeConverter() const { return typeConverter; }
      
        template <typename ConverterTy>
        std::enable_if_t<std::is_base_of<TypeConverter, ConverterTy>::value,
                         const ConverterTy *>
        getTypeConverter() const {
          return static_cast<const ConverterTy *>(typeConverter);
        }
      
      protected:
        using RewritePattern::RewritePattern;
      
        template <typename... Args>
        ConversionPattern(const TypeConverter &typeConverter, Args &&...args)
            : RewritePattern(std::forward<Args>(args)...),
              typeConverter(&typeConverter) {}
      
      protected:
        const TypeConverter *typeConverter = nullptr;
      
      private:
        using RewritePattern::rewrite;
      };
      

      最后的 using 操作的作用是将父类没有 ArrayRef<Value> operands 参数的 rewrite 方法重新引入到 ConversionPattern 类的作用域,防止开发者写继承该类的子类时,错误又写一个和 ConversionPattern 的父类 RewritePattern 一样的 rewrite 方法,使得失去了 ConversionPattern 这个中间类的目的。没有 usingvoid rewrite(Operation *op, PatternRewriter &rewriter) 该函数签名引入作用域的话,子类在定义同签名方法时完全不会报错。

      最后 OpConversionPattern 通过使用 OpAdaptor adaptor 属性替代手动从 ArrayRef<Value> operands 列表中获取操作数,OpAdaptor adaptor 保留了操作数的语义,使用起来更自然:

      template <typename SourceOp>
      class OpConversionPattern : public ConversionPattern {
      public:
        using OpAdaptor = typename SourceOp::Adaptor;
      
        OpConversionPattern(MLIRContext *context, PatternBenefit benefit = 1)
            : ConversionPattern(SourceOp::getOperationName(), benefit, context) {}
        OpConversionPattern(const TypeConverter &typeConverter, MLIRContext *context,
                            PatternBenefit benefit = 1)
            : ConversionPattern(typeConverter, SourceOp::getOperationName(), benefit,
                                context) {}
      
        LogicalResult match(Operation *op) const final {
          return match(cast<SourceOp>(op));
        }
        void rewrite(Operation *op, ArrayRef<Value> operands,
                     ConversionPatternRewriter &rewriter) const final {
          auto sourceOp = cast<SourceOp>(op);
          rewrite(sourceOp, OpAdaptor(operands, sourceOp), rewriter);
        }
        LogicalResult
        matchAndRewrite(Operation *op, ArrayRef<Value> operands,
                        ConversionPatternRewriter &rewriter) const final {
          auto sourceOp = cast<SourceOp>(op);
          return matchAndRewrite(sourceOp, OpAdaptor(operands, sourceOp), rewriter);
        }
      
        virtual LogicalResult match(SourceOp op) const {
          llvm_unreachable("must override match or matchAndRewrite");
        }
        virtual void rewrite(SourceOp op, OpAdaptor adaptor,
                             ConversionPatternRewriter &rewriter) const {
          llvm_unreachable("must override matchAndRewrite or a rewrite method");
        }
        virtual LogicalResult
        matchAndRewrite(SourceOp op, OpAdaptor adaptor,
                        ConversionPatternRewriter &rewriter) const {
          if (failed(match(op)))
            return failure();
          rewrite(op, adaptor, rewriter);
          return success();
        }
      
      private:
        using ConversionPattern::matchAndRewrite;
      };
      

      OpAdaptor 是在定义 op 的时候 tablegen 自动生成的。

    • SubOpPat

      struct SubOpPat: OpConversionPattern<SubOp> {
        using OpConversionPattern<SubOp>::OpConversionPattern;
        LogicalResult matchAndRewrite(SubOp op, SubOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
          rewriter.replaceOpWithNewOp<arith::SubIOp>(op, adaptor.getLhs(), adaptor.getRhs());
          return success();
        }
      };
      

      这里 adaptor 的作用就体现出来了,直接使用 getLhs/getRhs 而非用下标访问 ArrayRef<Value> operands 的元素。

    • 下面雷同:

      struct ConstantOpPat: OpConversionPattern<ConstantOp> {
        using OpConversionPattern<ConstantOp>::OpConversionPattern;
        LogicalResult matchAndRewrite(ConstantOp op, ConstantOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
          rewriter.replaceOpWithNewOp<arith::ConstantOp>(op, op.getValueAttr());
          return success();
        }
      };
      
      struct ReturnOpPat: OpConversionPattern<ReturnOp> {
        using OpConversionPattern<ReturnOp>::OpConversionPattern;
        LogicalResult matchAndRewrite(ReturnOp op, ReturnOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
          auto data = adaptor.getData();
          rewriter.startOpModification(op);
          op.getDataMutable().assign(data);
          rewriter.finalizeOpModification(op);
          return success();
        }
      };
      
      struct CallOpPat: OpConversionPattern<CallOp> {
        using OpConversionPattern<CallOp>::OpConversionPattern;
        LogicalResult matchAndRewrite(CallOp op, CallOpAdaptor adaptor, ConversionPatternRewriter & rewriter) const {
          SmallVector<Type> resTypes;
          assert(succeeded(getTypeConverter()->convertTypes(op->getResultTypes(), resTypes)));
          rewriter.replaceOpWithNewOp<CallOp>(op, resTypes, op.getCallee(), adaptor.getOperands());
          return success();
        }
      };
      
  • 通过定义 ConvertToyToArithPass 类编写 convert pass 的业务逻辑:

    struct ConvertToyToArithPass : toy::impl::ConvertToyToArithBase<ConvertToyToArithPass> {
      using toy::impl::ConvertToyToArithBase<ConvertToyToArithPass>::ConvertToyToArithBase;
      void getDependentDialects(DialectRegistry &registry) const final {
        registry.insert<arith::ArithDialect>();
      }
      void runOnOperation() final {
        ConversionTarget target(getContext());
        target.addLegalDialect<arith::ArithDialect>();
        auto checkValid = [](Operation* f) {
          return llvm::all_of(f->getOperandTypes(), [](Type t) {return !isa<ToyIntegerType>(t);});
        };
        target.addDynamicallyLegalOp<ReturnOp, CallOp>(checkValid);
        TypeConverter converter;
        converter.addConversion([&](ToyIntegerType t) -> std::optional<IntegerType> {
          return IntegerType::get(&getContext(), t.getWidth());
        });
        converter.addTargetMaterialization([](OpBuilder& builder, Type resultType, ValueRange inputs, Location loc) -> std::optional<Value> {
          return builder.create<UnrealizedConversionCastOp>(loc, resultType, inputs).getResult(0);
        });
        RewritePatternSet patterns(&getContext());
        patterns.add<AddOpPat, SubOpPat, ConstantOpPat, ReturnOpPat, CallOpPat>(converter, &getContext());
        populateFunctionOpInterfaceTypeConversionPattern<FuncOp>(patterns, converter);
        if(failed(applyPartialConversion(getOperation(), target, std::move(patterns))))
          signalPassFailure();
      }
    };
    
  • 自定义的 pass constructor 的工厂函数实现:

    std::unique_ptr<mlir::Pass> toy::createConvertToyToArithPass(ConvertToyToArithOptions options) {
      return std::make_unique<ConvertToyToArithPass>(options);
    }
    
posted @ 2026-07-25 15:41  judesongd  阅读(6)  评论(0)    收藏  举报