1 //===- DialectGen.cpp - MLIR dialect definitions generator ----------------===// 2 // 3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. 4 // See https://llvm.org/LICENSE.txt for license information. 5 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception 6 // 7 //===----------------------------------------------------------------------===// 8 // 9 // DialectGen uses the description of dialects to generate C++ definitions. 10 // 11 //===----------------------------------------------------------------------===// 12 13 #include "mlir/TableGen/Format.h" 14 #include "mlir/TableGen/GenInfo.h" 15 #include "mlir/TableGen/OpClass.h" 16 #include "mlir/TableGen/OpInterfaces.h" 17 #include "mlir/TableGen/OpTrait.h" 18 #include "mlir/TableGen/Operator.h" 19 #include "llvm/ADT/Sequence.h" 20 #include "llvm/ADT/StringExtras.h" 21 #include "llvm/Support/CommandLine.h" 22 #include "llvm/Support/Signals.h" 23 #include "llvm/TableGen/Error.h" 24 #include "llvm/TableGen/Record.h" 25 #include "llvm/TableGen/TableGenBackend.h" 26 27 #define DEBUG_TYPE "mlir-tblgen-opdefgen" 28 29 using namespace mlir; 30 using namespace mlir::tblgen; 31 32 static llvm::cl::OptionCategory dialectGenCat("Options for -gen-dialect-*"); 33 static llvm::cl::opt<std::string> 34 selectedDialect("dialect", llvm::cl::desc("The dialect to gen for"), 35 llvm::cl::cat(dialectGenCat), llvm::cl::CommaSeparated); 36 37 /// Utility iterator used for filtering records for a specific dialect. 38 namespace { 39 using DialectFilterIterator = 40 llvm::filter_iterator<ArrayRef<llvm::Record *>::iterator, 41 std::function<bool(const llvm::Record *)>>; 42 } // end anonymous namespace 43 44 /// Given a set of records for a T, filter the ones that correspond to 45 /// the given dialect. 46 template <typename T> 47 static iterator_range<DialectFilterIterator> 48 filterForDialect(ArrayRef<llvm::Record *> records, Dialect &dialect) { 49 auto filterFn = [&](const llvm::Record *record) { 50 return T(record).getDialect() == dialect; 51 }; 52 return {DialectFilterIterator(records.begin(), records.end(), filterFn), 53 DialectFilterIterator(records.end(), records.end(), filterFn)}; 54 } 55 56 //===----------------------------------------------------------------------===// 57 // GEN: Dialect declarations 58 //===----------------------------------------------------------------------===// 59 60 /// The code block for the start of a dialect class declaration. 61 /// 62 /// {0}: The name of the dialect class. 63 /// {1}: The dialect namespace. 64 static const char *const dialectDeclBeginStr = R"( 65 class {0} : public ::mlir::Dialect { 66 public: 67 explicit {0}(::mlir::MLIRContext *context); 68 static ::llvm::StringRef getDialectNamespace() { return "{1}"; } 69 )"; 70 71 /// The code block for the attribute parser/printer hooks. 72 static const char *const attrParserDecl = R"( 73 /// Parse an attribute registered to this dialect. 74 ::mlir::Attribute parseAttribute(::mlir::DialectAsmParser &parser, 75 ::mlir::Type type) const override; 76 77 /// Print an attribute registered to this dialect. 78 void printAttribute(::mlir::Attribute attr, 79 ::mlir::DialectAsmPrinter &os) const override; 80 )"; 81 82 /// The code block for the type parser/printer hooks. 83 static const char *const typeParserDecl = R"( 84 /// Parse a type registered to this dialect. 85 ::mlir::Type parseType(::mlir::DialectAsmParser &parser) const override; 86 87 /// Print a type registered to this dialect. 88 void printType(::mlir::Type type, 89 ::mlir::DialectAsmPrinter &os) const override; 90 )"; 91 92 /// The code block for the constant materializer hook. 93 static const char *const constantMaterializerDecl = R"( 94 /// Materialize a single constant operation from a given attribute value with 95 /// the desired resultant type. 96 ::mlir::Operation *materializeConstant(::mlir::OpBuilder &builder, 97 ::mlir::Attribute value, 98 ::mlir::Type type, 99 ::mlir::Location loc) override; 100 )"; 101 102 /// The code block for the operation attribute verifier hook. 103 static const char *const opAttrVerifierDecl = R"( 104 /// Provides a hook for verifying dialect attributes attached to the given 105 /// op. 106 ::mlir::LogicalResult verifyOperationAttribute( 107 ::mlir::Operation *op, ::mlir::NamedAttribute attribute) override; 108 )"; 109 110 /// The code block for the region argument attribute verifier hook. 111 static const char *const regionArgAttrVerifierDecl = R"( 112 /// Provides a hook for verifying dialect attributes attached to the given 113 /// op's region argument. 114 ::mlir::LogicalResult verifyRegionArgAttribute( 115 ::mlir::Operation *op, unsigned regionIndex, unsigned argIndex, 116 ::mlir::NamedAttribute attribute) override; 117 )"; 118 119 /// The code block for the region result attribute verifier hook. 120 static const char *const regionResultAttrVerifierDecl = R"( 121 /// Provides a hook for verifying dialect attributes attached to the given 122 /// op's region result. 123 ::mlir::LogicalResult verifyRegionResultAttribute( 124 ::mlir::Operation *op, unsigned regionIndex, unsigned resultIndex, 125 ::mlir::NamedAttribute attribute) override; 126 )"; 127 128 /// Generate the declaration for the given dialect class. 129 static void emitDialectDecl(Dialect &dialect, 130 iterator_range<DialectFilterIterator> dialectAttrs, 131 iterator_range<DialectFilterIterator> dialectTypes, 132 raw_ostream &os) { 133 // Emit the start of the decl. 134 std::string cppName = dialect.getCppClassName(); 135 os << llvm::formatv(dialectDeclBeginStr, cppName, dialect.getName()); 136 137 // Check for any attributes/types registered to this dialect. If there are, 138 // add the hooks for parsing/printing. 139 if (!dialectAttrs.empty()) 140 os << attrParserDecl; 141 if (!dialectTypes.empty()) 142 os << typeParserDecl; 143 144 // Add the decls for the various features of the dialect. 145 if (dialect.hasConstantMaterializer()) 146 os << constantMaterializerDecl; 147 if (dialect.hasOperationAttrVerify()) 148 os << opAttrVerifierDecl; 149 if (dialect.hasRegionArgAttrVerify()) 150 os << regionArgAttrVerifierDecl; 151 if (dialect.hasRegionResultAttrVerify()) 152 os << regionResultAttrVerifierDecl; 153 if (llvm::Optional<StringRef> extraDecl = dialect.getExtraClassDeclaration()) 154 os << *extraDecl; 155 156 // End the dialect decl. 157 os << "};\n"; 158 } 159 160 static bool emitDialectDecls(const llvm::RecordKeeper &recordKeeper, 161 raw_ostream &os) { 162 emitSourceFileHeader("Dialect Declarations", os); 163 164 auto defs = recordKeeper.getAllDerivedDefinitions("Dialect"); 165 if (defs.empty()) 166 return false; 167 168 // Select the dialect to gen for. 169 const llvm::Record *dialectDef = nullptr; 170 if (defs.size() == 1 && selectedDialect.getNumOccurrences() == 0) { 171 dialectDef = defs.front(); 172 } else if (selectedDialect.getNumOccurrences() == 0) { 173 llvm::errs() << "when more than 1 dialect is present, one must be selected " 174 "via '-dialect'"; 175 return true; 176 } else { 177 auto dialectIt = llvm::find_if(defs, [](const llvm::Record *def) { 178 return Dialect(def).getName() == selectedDialect; 179 }); 180 if (dialectIt == defs.end()) { 181 llvm::errs() << "selected dialect with '-dialect' does not exist"; 182 return true; 183 } 184 dialectDef = *dialectIt; 185 } 186 187 auto attrDefs = recordKeeper.getAllDerivedDefinitions("DialectAttr"); 188 auto typeDefs = recordKeeper.getAllDerivedDefinitions("DialectType"); 189 Dialect dialect(dialectDef); 190 emitDialectDecl(dialect, filterForDialect<Attribute>(attrDefs, dialect), 191 filterForDialect<Type>(typeDefs, dialect), os); 192 return false; 193 } 194 195 //===----------------------------------------------------------------------===// 196 // GEN: Dialect registration hooks 197 //===----------------------------------------------------------------------===// 198 199 static mlir::GenRegistration 200 genDialectDecls("gen-dialect-decls", "Generate dialect declarations", 201 [](const llvm::RecordKeeper &records, raw_ostream &os) { 202 return emitDialectDecls(records, os); 203 }); 204