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