1 //===- OpInterfacesGen.cpp - MLIR op interface utility 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 // OpInterfacesGen generates definitions for operation interfaces. 10 // 11 //===----------------------------------------------------------------------===// 12 13 #include "DocGenUtilities.h" 14 #include "mlir/TableGen/Format.h" 15 #include "mlir/TableGen/GenInfo.h" 16 #include "mlir/TableGen/OpInterfaces.h" 17 #include "llvm/ADT/SmallVector.h" 18 #include "llvm/ADT/StringExtras.h" 19 #include "llvm/Support/FormatVariadic.h" 20 #include "llvm/Support/raw_ostream.h" 21 #include "llvm/TableGen/Error.h" 22 #include "llvm/TableGen/Record.h" 23 #include "llvm/TableGen/TableGenBackend.h" 24 25 using namespace llvm; 26 using namespace mlir; 27 using mlir::tblgen::OpInterface; 28 using mlir::tblgen::OpInterfaceMethod; 29 30 // Emit the method name and argument list for the given method. If 31 // 'addOperationArg' is true, then an Operation* argument is added to the 32 // beginning of the argument list. 33 static void emitMethodNameAndArgs(const OpInterfaceMethod &method, 34 raw_ostream &os, bool addOperationArg) { 35 os << method.getName() << '('; 36 if (addOperationArg) 37 os << "Operation *tablegen_opaque_op" << (method.arg_empty() ? "" : ", "); 38 llvm::interleaveComma(method.getArguments(), os, 39 [&](const OpInterfaceMethod::Argument &arg) { 40 os << arg.type << " " << arg.name; 41 }); 42 os << ')'; 43 } 44 45 // Get an array of all OpInterface definitions but exclude those subclassing 46 // "DeclareOpInterfaceMethods". 47 static std::vector<Record *> 48 getAllOpInterfaceDefinitions(const RecordKeeper &recordKeeper) { 49 std::vector<Record *> defs = 50 recordKeeper.getAllDerivedDefinitions("OpInterface"); 51 52 llvm::erase_if(defs, [](const Record *def) { 53 return def->isSubClassOf("DeclareOpInterfaceMethods"); 54 }); 55 return defs; 56 } 57 58 //===----------------------------------------------------------------------===// 59 // GEN: Interface definitions 60 //===----------------------------------------------------------------------===// 61 62 static void emitInterfaceDef(OpInterface &interface, raw_ostream &os) { 63 StringRef interfaceName = interface.getName(); 64 65 // Insert the method definitions. 66 for (auto &method : interface.getMethods()) { 67 os << method.getReturnType() << " " << interfaceName << "::"; 68 emitMethodNameAndArgs(method, os, /*addOperationArg=*/false); 69 70 // Forward to the method on the concrete operation type. 71 os << " {\n return getImpl()->" << method.getName() << '('; 72 if (!method.isStatic()) 73 os << "getOperation()" << (method.arg_empty() ? "" : ", "); 74 llvm::interleaveComma( 75 method.getArguments(), os, 76 [&](const OpInterfaceMethod::Argument &arg) { os << arg.name; }); 77 os << ");\n }\n"; 78 } 79 } 80 81 static bool emitInterfaceDefs(const RecordKeeper &recordKeeper, 82 raw_ostream &os) { 83 llvm::emitSourceFileHeader("Operation Interface Definitions", os); 84 85 for (const auto *def : getAllOpInterfaceDefinitions(recordKeeper)) { 86 OpInterface interface(def); 87 emitInterfaceDef(interface, os); 88 } 89 return false; 90 } 91 92 //===----------------------------------------------------------------------===// 93 // GEN: Interface declarations 94 //===----------------------------------------------------------------------===// 95 96 static void emitConceptDecl(OpInterface &interface, raw_ostream &os) { 97 os << " class Concept {\n" 98 << " public:\n" 99 << " virtual ~Concept() = default;\n"; 100 101 // Insert each of the pure virtual concept methods. 102 for (auto &method : interface.getMethods()) { 103 os << " virtual " << method.getReturnType() << " "; 104 emitMethodNameAndArgs(method, os, /*addOperationArg=*/!method.isStatic()); 105 os << " = 0;\n"; 106 } 107 os << " };\n"; 108 } 109 110 static void emitModelDecl(OpInterface &interface, raw_ostream &os) { 111 os << " template<typename ConcreteOp>\n"; 112 os << " class Model : public Concept {\npublic:\n"; 113 114 // Insert each of the virtual method overrides. 115 for (auto &method : interface.getMethods()) { 116 os << " " << method.getReturnType() << " "; 117 emitMethodNameAndArgs(method, os, /*addOperationArg=*/!method.isStatic()); 118 os << " final {\n"; 119 120 // Provide a definition of the concrete op if this is non static. 121 if (!method.isStatic()) { 122 os << " auto op = llvm::cast<ConcreteOp>(tablegen_opaque_op);\n" 123 << " (void)op;\n"; 124 } 125 126 // Check for a provided body to the function. 127 if (auto body = method.getBody()) { 128 os << body << "\n }\n"; 129 continue; 130 } 131 132 // Forward to the method on the concrete operation type. 133 os << " return " << (method.isStatic() ? "ConcreteOp::" : "op."); 134 135 // Add the arguments to the call. 136 os << method.getName() << '('; 137 llvm::interleaveComma( 138 method.getArguments(), os, 139 [&](const OpInterfaceMethod::Argument &arg) { os << arg.name; }); 140 os << ");\n }\n"; 141 } 142 os << " };\n"; 143 } 144 145 static void emitTraitDecl(OpInterface &interface, raw_ostream &os, 146 StringRef interfaceName, 147 StringRef interfaceTraitsName) { 148 os << " template <typename ConcreteOp>\n " 149 << llvm::formatv("struct Trait : public OpInterface<{0}," 150 " detail::{1}>::Trait<ConcreteOp> {{\n", 151 interfaceName, interfaceTraitsName); 152 153 // Insert the default implementation for any methods. 154 for (auto &method : interface.getMethods()) { 155 // Flag interface methods named verifyTrait. 156 if (method.getName() == "verifyTrait") 157 PrintFatalError( 158 formatv("'verifyTrait' method cannot be specified as interface " 159 "method for '{0}'; set 'verify' on OpInterfaceTrait instead", 160 interfaceName)); 161 auto defaultImpl = method.getDefaultImplementation(); 162 if (!defaultImpl) 163 continue; 164 165 os << " " << (method.isStatic() ? "static " : "") << method.getReturnType() 166 << " "; 167 emitMethodNameAndArgs(method, os, /*addOperationArg=*/false); 168 os << " {\n" << defaultImpl.getValue() << " }\n"; 169 } 170 171 tblgen::FmtContext traitCtx; 172 traitCtx.withOp("op"); 173 if (auto verify = interface.getVerify()) { 174 os << " static LogicalResult verifyTrait(Operation* op) {\n" 175 << std::string(tblgen::tgfmt(*verify, &traitCtx)) << "\n }\n"; 176 } 177 178 os << " };\n"; 179 } 180 181 static void emitInterfaceDecl(OpInterface &interface, raw_ostream &os) { 182 StringRef interfaceName = interface.getName(); 183 auto interfaceTraitsName = (interfaceName + "InterfaceTraits").str(); 184 185 // Emit the traits struct containing the concept and model declarations. 186 os << "namespace detail {\n" 187 << "struct " << interfaceTraitsName << " {\n"; 188 emitConceptDecl(interface, os); 189 emitModelDecl(interface, os); 190 os << "};\n} // end namespace detail\n"; 191 192 // Emit the main interface class declaration. 193 os << llvm::formatv("class {0} : public OpInterface<{1}, detail::{2}> {\n" 194 "public:\n" 195 " using OpInterface<{1}, detail::{2}>::OpInterface;\n", 196 interfaceName, interfaceName, interfaceTraitsName); 197 198 // Emit the derived trait for the interface. 199 emitTraitDecl(interface, os, interfaceName, interfaceTraitsName); 200 201 // Insert the method declarations. 202 for (auto &method : interface.getMethods()) { 203 os << " " << method.getReturnType() << " "; 204 emitMethodNameAndArgs(method, os, /*addOperationArg=*/false); 205 os << ";\n"; 206 } 207 208 // Emit any extra declarations. 209 if (Optional<StringRef> extraDecls = interface.getExtraClassDeclaration()) 210 os << *extraDecls << "\n"; 211 212 os << "};\n"; 213 } 214 215 static bool emitInterfaceDecls(const RecordKeeper &recordKeeper, 216 raw_ostream &os) { 217 llvm::emitSourceFileHeader("Operation Interface Declarations", os); 218 219 for (const auto *def : getAllOpInterfaceDefinitions(recordKeeper)) { 220 OpInterface interface(def); 221 emitInterfaceDecl(interface, os); 222 } 223 return false; 224 } 225 226 //===----------------------------------------------------------------------===// 227 // GEN: Interface documentation 228 //===----------------------------------------------------------------------===// 229 230 /// Emit a string corresponding to a C++ type, followed by a space if necessary. 231 static raw_ostream &emitCPPType(StringRef type, raw_ostream &os) { 232 type = type.trim(); 233 os << type; 234 if (type.back() != '&' && type.back() != '*') 235 os << " "; 236 return os; 237 } 238 239 static void emitInterfaceDoc(const Record &interfaceDef, raw_ostream &os) { 240 OpInterface interface(&interfaceDef); 241 242 // Emit the interface name followed by the description. 243 os << "## " << interface.getName() << " (" << interfaceDef.getName() << ")"; 244 if (auto description = interface.getDescription()) 245 mlir::tblgen::emitDescription(*description, os); 246 247 // Emit the methods required by the interface. 248 os << "\n### Methods:\n"; 249 for (const auto &method : interface.getMethods()) { 250 // Emit the method name. 251 os << "#### `" << method.getName() << "`\n\n```c++\n"; 252 253 // Emit the method signature. 254 if (method.isStatic()) 255 os << "static "; 256 emitCPPType(method.getReturnType(), os) << method.getName() << '('; 257 llvm::interleaveComma(method.getArguments(), os, 258 [&](const OpInterfaceMethod::Argument &arg) { 259 emitCPPType(arg.type, os) << arg.name; 260 }); 261 os << ");\n```\n"; 262 263 // Emit the description. 264 if (auto description = method.getDescription()) 265 mlir::tblgen::emitDescription(*description, os); 266 267 // If the body is not provided, this method must be provided by the 268 // operation. 269 if (!method.getBody()) 270 os << "\nNOTE: This method *must* be implemented by the operation.\n\n"; 271 } 272 } 273 274 static bool emitInterfaceDocs(const RecordKeeper &recordKeeper, 275 raw_ostream &os) { 276 os << "<!-- Autogenerated by mlir-tblgen; don't manually edit -->\n"; 277 os << "# Operation Interface definition\n"; 278 279 for (const auto *def : getAllOpInterfaceDefinitions(recordKeeper)) 280 emitInterfaceDoc(*def, os); 281 return false; 282 } 283 284 //===----------------------------------------------------------------------===// 285 // GEN: Interface registration hooks 286 //===----------------------------------------------------------------------===// 287 288 // Registers the operation interface generator to mlir-tblgen. 289 static mlir::GenRegistration 290 genInterfaceDecls("gen-op-interface-decls", 291 "Generate op interface declarations", 292 [](const RecordKeeper &records, raw_ostream &os) { 293 return emitInterfaceDecls(records, os); 294 }); 295 296 // Registers the operation interface generator to mlir-tblgen. 297 static mlir::GenRegistration 298 genInterfaceDefs("gen-op-interface-defs", 299 "Generate op interface definitions", 300 [](const RecordKeeper &records, raw_ostream &os) { 301 return emitInterfaceDefs(records, os); 302 }); 303 304 // Registers the operation interface document generator to mlir-tblgen. 305 static mlir::GenRegistration 306 genInterfaceDocs("gen-op-interface-doc", 307 "Generate op interface documentation", 308 [](const RecordKeeper &records, raw_ostream &os) { 309 return emitInterfaceDocs(records, os); 310 }); 311