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 ®istry) 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 ®istry) 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 ®istry) 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 ®istry) 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_CONVERTTOYTOARITH和GEN_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 ®istry) 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 ®istry) 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 ®istry) 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 ®istry) 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_CONVERTTOYTOARITH,GEN_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 文件中提供的信息。
- 什么叫接入 mlir 系统呢?
-
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
- 如果不写 constructor 属性也会自动生成对应工厂函数,放在 inc 文件对应 pass 的头文件部分
-
引入 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 ®istry) 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 -> PatternPattern包含一个模式所需要的全部元数据,但是不包含具体的模式,子类继承该类描述具体的模式: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之上,提出了match,rewrite,matchAndRrite,create方法: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(); };match,rewrite,matchAndRrite比较容易理解,就是用来进行匹配以及重写。create主要用于注册一个 pattern 时,安全的实例化一个 pattern 对象,而非手动new或者make_unique,因为有的 pattern 还需要进行额外的 initialize。ConvertionPattern在RewritePattern的基础上提供了自动类型转换的功能,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这个中间类的目的。没有using将void 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 ®istry) 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); }
浙公网安备 33010602011771号