1 //===- AttrOrTypeDefGen.cpp - MLIR AttrOrType 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 #include "AttrOrTypeFormatGen.h" 10 #include "mlir/TableGen/AttrOrTypeDef.h" 11 #include "mlir/TableGen/Class.h" 12 #include "mlir/TableGen/CodeGenHelpers.h" 13 #include "mlir/TableGen/Format.h" 14 #include "mlir/TableGen/GenInfo.h" 15 #include "mlir/TableGen/Interfaces.h" 16 #include "llvm/ADT/StringSet.h" 17 #include "llvm/Support/CommandLine.h" 18 #include "llvm/TableGen/Error.h" 19 #include "llvm/TableGen/TableGenBackend.h" 20 21 #define DEBUG_TYPE "mlir-tblgen-attrortypedefgen" 22 23 using namespace mlir; 24 using namespace mlir::tblgen; 25 26 //===----------------------------------------------------------------------===// 27 // Utility Functions 28 //===----------------------------------------------------------------------===// 29 30 /// Find all the AttrOrTypeDef for the specified dialect. If no dialect 31 /// specified and can only find one dialect's defs, use that. 32 static void collectAllDefs(StringRef selectedDialect, 33 std::vector<llvm::Record *> records, 34 SmallVectorImpl<AttrOrTypeDef> &resultDefs) { 35 // Nothing to do if no defs were found. 36 if (records.empty()) 37 return; 38 39 auto defs = llvm::map_range( 40 records, [&](const llvm::Record *rec) { return AttrOrTypeDef(rec); }); 41 if (selectedDialect.empty()) { 42 // If a dialect was not specified, ensure that all found defs belong to the 43 // same dialect. 44 if (!llvm::is_splat(llvm::map_range( 45 defs, [](const auto &def) { return def.getDialect(); }))) { 46 llvm::PrintFatalError("defs belonging to more than one dialect. Must " 47 "select one via '--(attr|type)defs-dialect'"); 48 } 49 resultDefs.assign(defs.begin(), defs.end()); 50 } else { 51 // Otherwise, generate the defs that belong to the selected dialect. 52 auto dialectDefs = llvm::make_filter_range(defs, [&](const auto &def) { 53 return def.getDialect().getName().equals(selectedDialect); 54 }); 55 resultDefs.assign(dialectDefs.begin(), dialectDefs.end()); 56 } 57 } 58 59 //===----------------------------------------------------------------------===// 60 // DefGen 61 //===----------------------------------------------------------------------===// 62 63 namespace { 64 class DefGen { 65 public: 66 /// Create the attribute or type class. 67 DefGen(const AttrOrTypeDef &def); 68 69 void emitDecl(raw_ostream &os) const { 70 if (storageCls) { 71 NamespaceEmitter ns(os, def.getStorageNamespace()); 72 os << "struct " << def.getStorageClassName() << ";\n"; 73 } 74 defCls.writeDeclTo(os); 75 } 76 void emitDef(raw_ostream &os) const { 77 if (storageCls && def.genStorageClass()) { 78 NamespaceEmitter ns(os, def.getStorageNamespace()); 79 storageCls->writeDeclTo(os); // everything is inline 80 } 81 defCls.writeDefTo(os); 82 } 83 84 private: 85 /// Add traits from the TableGen definition to the class. 86 void createParentWithTraits(); 87 /// Emit top-level declarations: using declarations and any extra class 88 /// declarations. 89 void emitTopLevelDeclarations(); 90 /// Emit attribute or type builders. 91 void emitBuilders(); 92 /// Emit a verifier for the def. 93 void emitVerifier(); 94 /// Emit parsers and printers. 95 void emitParserPrinter(); 96 /// Emit parameter accessors, if required. 97 void emitAccessors(); 98 /// Emit interface methods. 99 void emitInterfaceMethods(); 100 101 //===--------------------------------------------------------------------===// 102 // Builder Emission 103 104 /// Emit the default builder `Attribute::get` 105 void emitDefaultBuilder(); 106 /// Emit the checked builder `Attribute::getChecked` 107 void emitCheckedBuilder(); 108 /// Emit a custom builder. 109 void emitCustomBuilder(const AttrOrTypeBuilder &builder); 110 /// Emit a checked custom builder. 111 void emitCheckedCustomBuilder(const AttrOrTypeBuilder &builder); 112 113 //===--------------------------------------------------------------------===// 114 // Interface Method Emission 115 116 /// Emit methods for a trait. 117 void emitTraitMethods(const InterfaceTrait &trait); 118 /// Emit a trait method. 119 void emitTraitMethod(const InterfaceMethod &method); 120 121 //===--------------------------------------------------------------------===// 122 // Storage Class Emission 123 void emitStorageClass(); 124 /// Generate the storage class constructor. 125 void emitStorageConstructor(); 126 /// Emit the key type `KeyTy`. 127 void emitKeyType(); 128 /// Emit the equality comparison operator. 129 void emitEquals(); 130 /// Emit the key hash function. 131 void emitHashKey(); 132 /// Emit the function to construct the storage class. 133 void emitConstruct(); 134 135 //===--------------------------------------------------------------------===// 136 // Utility Function Declarations 137 138 /// Get the method parameters for a def builder, where the first several 139 /// parameters may be different. 140 SmallVector<MethodParameter> 141 getBuilderParams(std::initializer_list<MethodParameter> prefix) const; 142 143 //===--------------------------------------------------------------------===// 144 // Class fields 145 146 /// The attribute or type definition. 147 const AttrOrTypeDef &def; 148 /// The list of attribute or type parameters. 149 ArrayRef<AttrOrTypeParameter> params; 150 /// The attribute or type class. 151 Class defCls; 152 /// An optional attribute or type storage class. The storage class will 153 /// exist if and only if the def has more than zero parameters. 154 Optional<Class> storageCls; 155 156 /// The C++ base value of the def, either "Attribute" or "Type". 157 StringRef valueType; 158 /// The prefix/suffix of the TableGen def name, either "Attr" or "Type". 159 StringRef defType; 160 }; 161 } // namespace 162 163 DefGen::DefGen(const AttrOrTypeDef &def) 164 : def(def), params(def.getParameters()), defCls(def.getCppClassName()), 165 valueType(isa<AttrDef>(def) ? "Attribute" : "Type"), 166 defType(isa<AttrDef>(def) ? "Attr" : "Type") { 167 // Check that all parameters have names. 168 for (const AttrOrTypeParameter ¶m : def.getParameters()) 169 if (param.isAnonymous()) 170 llvm::PrintFatalError("all parameters must have a name"); 171 172 // If a storage class is needed, create one. 173 if (def.getNumParameters() > 0) 174 storageCls.emplace(def.getStorageClassName(), /*isStruct=*/true); 175 176 // Create the parent class with any indicated traits. 177 createParentWithTraits(); 178 // Emit top-level declarations. 179 emitTopLevelDeclarations(); 180 // Emit builders for defs with parameters 181 if (storageCls) 182 emitBuilders(); 183 // Emit the verifier. 184 if (storageCls && def.genVerifyDecl()) 185 emitVerifier(); 186 // Emit the mnemonic, if there is one, and any associated parser and printer. 187 if (def.getMnemonic()) 188 emitParserPrinter(); 189 // Emit accessors 190 if (def.genAccessors()) 191 emitAccessors(); 192 // Emit trait interface methods 193 emitInterfaceMethods(); 194 defCls.finalize(); 195 // Emit a storage class if one is needed 196 if (storageCls && def.genStorageClass()) 197 emitStorageClass(); 198 } 199 200 void DefGen::createParentWithTraits() { 201 ParentClass defParent(strfmt("::mlir::{0}::{1}Base", valueType, defType)); 202 defParent.addTemplateParam(def.getCppClassName()); 203 defParent.addTemplateParam(def.getCppBaseClassName()); 204 defParent.addTemplateParam(storageCls 205 ? strfmt("{0}::{1}", def.getStorageNamespace(), 206 def.getStorageClassName()) 207 : strfmt("::mlir::{0}Storage", valueType)); 208 for (auto &trait : def.getTraits()) { 209 defParent.addTemplateParam( 210 isa<NativeTrait>(&trait) 211 ? cast<NativeTrait>(&trait)->getFullyQualifiedTraitName() 212 : cast<InterfaceTrait>(&trait)->getFullyQualifiedTraitName()); 213 } 214 defCls.addParent(std::move(defParent)); 215 } 216 217 /// Extra class definitions have a `$cppClass` substitution that is to be 218 /// replaced by the C++ class name. 219 static std::string formatExtraDefinitions(const AttrOrTypeDef &def) { 220 if (Optional<StringRef> extraDef = def.getExtraDefs()) { 221 FmtContext ctx = FmtContext().addSubst("cppClass", def.getCppClassName()); 222 return tgfmt(*extraDef, &ctx).str(); 223 } 224 return ""; 225 } 226 227 void DefGen::emitTopLevelDeclarations() { 228 // Inherit constructors from the attribute or type class. 229 defCls.declare<VisibilityDeclaration>(Visibility::Public); 230 defCls.declare<UsingDeclaration>("Base::Base"); 231 232 // Emit the extra declarations first in case there's a definition in there. 233 Optional<StringRef> extraDecl = def.getExtraDecls(); 234 std::string extraDef = formatExtraDefinitions(def); 235 defCls.declare<ExtraClassDeclaration>(extraDecl ? *extraDecl : "", 236 std::move(extraDef)); 237 } 238 239 void DefGen::emitBuilders() { 240 if (!def.skipDefaultBuilders()) { 241 emitDefaultBuilder(); 242 if (def.genVerifyDecl()) 243 emitCheckedBuilder(); 244 } 245 for (auto &builder : def.getBuilders()) { 246 emitCustomBuilder(builder); 247 if (def.genVerifyDecl()) 248 emitCheckedCustomBuilder(builder); 249 } 250 } 251 252 void DefGen::emitVerifier() { 253 defCls.declare<UsingDeclaration>("Base::getChecked"); 254 defCls.declareStaticMethod( 255 "::mlir::LogicalResult", "verify", 256 getBuilderParams({{"::llvm::function_ref<::mlir::InFlightDiagnostic()>", 257 "emitError"}})); 258 } 259 260 void DefGen::emitParserPrinter() { 261 auto *mnemonic = defCls.addStaticMethod<Method::Constexpr>( 262 "::llvm::StringLiteral", "getMnemonic"); 263 mnemonic->body().indent() << strfmt("return {\"{0}\"};", *def.getMnemonic()); 264 265 // Declare the parser and printer, if needed. 266 bool hasAssemblyFormat = def.getAssemblyFormat().has_value(); 267 if (!def.hasCustomAssemblyFormat() && !hasAssemblyFormat) 268 return; 269 270 // Declare the parser. 271 SmallVector<MethodParameter> parserParams; 272 parserParams.emplace_back("::mlir::AsmParser &", "odsParser"); 273 if (isa<AttrDef>(&def)) 274 parserParams.emplace_back("::mlir::Type", "odsType"); 275 auto *parser = defCls.addMethod(strfmt("::mlir::{0}", valueType), "parse", 276 hasAssemblyFormat ? Method::Static 277 : Method::StaticDeclaration, 278 std::move(parserParams)); 279 // Declare the printer. 280 auto props = hasAssemblyFormat ? Method::Const : Method::ConstDeclaration; 281 Method *printer = 282 defCls.addMethod("void", "print", props, 283 MethodParameter("::mlir::AsmPrinter &", "odsPrinter")); 284 // Emit the bodies if we are using the declarative format. 285 if (hasAssemblyFormat) 286 return generateAttrOrTypeFormat(def, parser->body(), printer->body()); 287 } 288 289 void DefGen::emitAccessors() { 290 for (auto ¶m : params) { 291 Method *m = defCls.addMethod( 292 param.getCppAccessorType(), param.getAccessorName(), 293 def.genStorageClass() ? Method::Const : Method::ConstDeclaration); 294 // Generate accessor definitions only if we also generate the storage 295 // class. Otherwise, let the user define the exact accessor definition. 296 if (!def.genStorageClass()) 297 continue; 298 auto scope = m->body().indent().scope("return getImpl()->", ";"); 299 if (isa<AttributeSelfTypeParameter>(param)) 300 m->body() << formatv("getType().cast<{0}>()", param.getCppType()); 301 else 302 m->body() << param.getName(); 303 } 304 } 305 306 void DefGen::emitInterfaceMethods() { 307 for (auto &traitDef : def.getTraits()) 308 if (auto *trait = dyn_cast<InterfaceTrait>(&traitDef)) 309 if (trait->shouldDeclareMethods()) 310 emitTraitMethods(*trait); 311 } 312 313 //===----------------------------------------------------------------------===// 314 // Builder Emission 315 316 SmallVector<MethodParameter> 317 DefGen::getBuilderParams(std::initializer_list<MethodParameter> prefix) const { 318 SmallVector<MethodParameter> builderParams; 319 builderParams.append(prefix.begin(), prefix.end()); 320 for (auto ¶m : params) 321 builderParams.emplace_back(param.getCppType(), param.getName()); 322 return builderParams; 323 } 324 325 void DefGen::emitDefaultBuilder() { 326 Method *m = defCls.addStaticMethod( 327 def.getCppClassName(), "get", 328 getBuilderParams({{"::mlir::MLIRContext *", "context"}})); 329 MethodBody &body = m->body().indent(); 330 auto scope = body.scope("return Base::get(context", ");"); 331 for (const auto ¶m : params) 332 body << ", " << param.getName(); 333 } 334 335 void DefGen::emitCheckedBuilder() { 336 Method *m = defCls.addStaticMethod( 337 def.getCppClassName(), "getChecked", 338 getBuilderParams( 339 {{"::llvm::function_ref<::mlir::InFlightDiagnostic()>", "emitError"}, 340 {"::mlir::MLIRContext *", "context"}})); 341 MethodBody &body = m->body().indent(); 342 auto scope = body.scope("return Base::getChecked(emitError, context", ");"); 343 for (const auto ¶m : params) 344 body << ", " << param.getName(); 345 } 346 347 static SmallVector<MethodParameter> 348 getCustomBuilderParams(std::initializer_list<MethodParameter> prefix, 349 const AttrOrTypeBuilder &builder) { 350 auto params = builder.getParameters(); 351 SmallVector<MethodParameter> builderParams; 352 builderParams.append(prefix.begin(), prefix.end()); 353 if (!builder.hasInferredContextParameter()) 354 builderParams.emplace_back("::mlir::MLIRContext *", "context"); 355 for (auto ¶m : params) { 356 builderParams.emplace_back(param.getCppType(), *param.getName(), 357 param.getDefaultValue()); 358 } 359 return builderParams; 360 } 361 362 void DefGen::emitCustomBuilder(const AttrOrTypeBuilder &builder) { 363 // Don't emit a body if there isn't one. 364 auto props = builder.getBody() ? Method::Static : Method::StaticDeclaration; 365 StringRef returnType = def.getCppClassName(); 366 if (Optional<StringRef> builderReturnType = builder.getReturnType()) 367 returnType = *builderReturnType; 368 Method *m = defCls.addMethod(returnType, "get", props, 369 getCustomBuilderParams({}, builder)); 370 if (!builder.getBody()) 371 return; 372 373 // Format the body and emit it. 374 FmtContext ctx; 375 ctx.addSubst("_get", "Base::get"); 376 if (!builder.hasInferredContextParameter()) 377 ctx.addSubst("_ctxt", "context"); 378 std::string bodyStr = tgfmt(*builder.getBody(), &ctx); 379 m->body().indent().getStream().printReindented(bodyStr); 380 } 381 382 /// Replace all instances of 'from' to 'to' in `str` and return the new string. 383 static std::string replaceInStr(std::string str, StringRef from, StringRef to) { 384 size_t pos = 0; 385 while ((pos = str.find(from.data(), pos, from.size())) != std::string::npos) 386 str.replace(pos, from.size(), to.data(), to.size()); 387 return str; 388 } 389 390 void DefGen::emitCheckedCustomBuilder(const AttrOrTypeBuilder &builder) { 391 // Don't emit a body if there isn't one. 392 auto props = builder.getBody() ? Method::Static : Method::StaticDeclaration; 393 StringRef returnType = def.getCppClassName(); 394 if (Optional<StringRef> builderReturnType = builder.getReturnType()) 395 returnType = *builderReturnType; 396 Method *m = defCls.addMethod( 397 returnType, "getChecked", props, 398 getCustomBuilderParams( 399 {{"::llvm::function_ref<::mlir::InFlightDiagnostic()>", "emitError"}}, 400 builder)); 401 if (!builder.getBody()) 402 return; 403 404 // Format the body and emit it. Replace $_get(...) with 405 // Base::getChecked(emitError, ...) 406 FmtContext ctx; 407 if (!builder.hasInferredContextParameter()) 408 ctx.addSubst("_ctxt", "context"); 409 std::string bodyStr = replaceInStr(builder.getBody()->str(), "$_get(", 410 "Base::getChecked(emitError, "); 411 bodyStr = tgfmt(bodyStr, &ctx); 412 m->body().indent().getStream().printReindented(bodyStr); 413 } 414 415 //===----------------------------------------------------------------------===// 416 // Interface Method Emission 417 418 void DefGen::emitTraitMethods(const InterfaceTrait &trait) { 419 // Get the set of methods that should always be declared. 420 auto alwaysDeclaredMethods = trait.getAlwaysDeclaredMethods(); 421 StringSet<> alwaysDeclared; 422 alwaysDeclared.insert(alwaysDeclaredMethods.begin(), 423 alwaysDeclaredMethods.end()); 424 425 Interface iface = trait.getInterface(); // causes strange bugs if elided 426 for (auto &method : iface.getMethods()) { 427 // Don't declare if the method has a body. Or if the method has a default 428 // implementation and the def didn't request that it always be declared. 429 if (method.getBody() || (method.getDefaultImplementation() && 430 !alwaysDeclared.count(method.getName()))) 431 continue; 432 emitTraitMethod(method); 433 } 434 } 435 436 void DefGen::emitTraitMethod(const InterfaceMethod &method) { 437 // All interface methods are declaration-only. 438 auto props = 439 method.isStatic() ? Method::StaticDeclaration : Method::ConstDeclaration; 440 SmallVector<MethodParameter> params; 441 for (auto ¶m : method.getArguments()) 442 params.emplace_back(param.type, param.name); 443 defCls.addMethod(method.getReturnType(), method.getName(), props, 444 std::move(params)); 445 } 446 447 //===----------------------------------------------------------------------===// 448 // Storage Class Emission 449 450 void DefGen::emitStorageConstructor() { 451 Constructor *ctor = 452 storageCls->addConstructor<Method::Inline>(getBuilderParams({})); 453 if (auto *attrDef = dyn_cast<AttrDef>(&def)) { 454 // For attributes, a parameter marked with AttributeSelfTypeParameter is 455 // the type initializer that must be passed to the parent constructor. 456 const auto isSelfType = [](const AttrOrTypeParameter ¶m) { 457 return isa<AttributeSelfTypeParameter>(param); 458 }; 459 auto *selfTypeParam = llvm::find_if(params, isSelfType); 460 if (std::count_if(selfTypeParam, params.end(), isSelfType) > 1) { 461 PrintFatalError(def.getLoc(), 462 "Only one attribute parameter can be marked as " 463 "AttributeSelfTypeParameter"); 464 } 465 // Alternatively, if a type builder was specified, use that instead. 466 std::string attrStorageInit = 467 selfTypeParam == params.end() ? "" : selfTypeParam->getName().str(); 468 if (attrDef->getTypeBuilder()) { 469 FmtContext ctx; 470 for (auto ¶m : params) 471 ctx.addSubst(strfmt("_{0}", param.getName()), param.getName()); 472 attrStorageInit = tgfmt(*attrDef->getTypeBuilder(), &ctx); 473 } 474 ctor->addMemberInitializer("::mlir::AttributeStorage", 475 std::move(attrStorageInit)); 476 // Initialize members that aren't the attribute's type. 477 for (auto ¶m : params) 478 if (selfTypeParam == params.end() || *selfTypeParam != param) 479 ctor->addMemberInitializer(param.getName(), param.getName()); 480 } else { 481 for (auto ¶m : params) 482 ctor->addMemberInitializer(param.getName(), param.getName()); 483 } 484 } 485 486 void DefGen::emitKeyType() { 487 std::string keyType("std::tuple<"); 488 llvm::raw_string_ostream os(keyType); 489 llvm::interleaveComma(params, os, 490 [&](auto ¶m) { os << param.getCppType(); }); 491 os << '>'; 492 storageCls->declare<UsingDeclaration>("KeyTy", std::move(os.str())); 493 } 494 495 void DefGen::emitEquals() { 496 Method *eq = storageCls->addConstMethod<Method::Inline>( 497 "bool", "operator==", MethodParameter("const KeyTy &", "tblgenKey")); 498 auto &body = eq->body().indent(); 499 auto scope = body.scope("return (", ");"); 500 const auto eachFn = [&](auto it) { 501 FmtContext ctx({{"_lhs", isa<AttributeSelfTypeParameter>(it.value()) 502 ? "getType()" 503 : it.value().getName()}, 504 {"_rhs", strfmt("std::get<{0}>(tblgenKey)", it.index())}}); 505 body << tgfmt(it.value().getComparator(), &ctx); 506 }; 507 llvm::interleave(llvm::enumerate(params), body, eachFn, ") && ("); 508 } 509 510 void DefGen::emitHashKey() { 511 Method *hash = storageCls->addStaticInlineMethod( 512 "::llvm::hash_code", "hashKey", 513 MethodParameter("const KeyTy &", "tblgenKey")); 514 auto &body = hash->body().indent(); 515 auto scope = body.scope("return ::llvm::hash_combine(", ");"); 516 llvm::interleaveComma(llvm::enumerate(params), body, [&](auto it) { 517 body << llvm::formatv("std::get<{0}>(tblgenKey)", it.index()); 518 }); 519 } 520 521 void DefGen::emitConstruct() { 522 Method *construct = storageCls->addMethod<Method::Inline>( 523 strfmt("{0} *", def.getStorageClassName()), "construct", 524 def.hasStorageCustomConstructor() ? Method::StaticDeclaration 525 : Method::Static, 526 MethodParameter(strfmt("::mlir::{0}StorageAllocator &", valueType), 527 "allocator"), 528 MethodParameter("const KeyTy &", "tblgenKey")); 529 if (!def.hasStorageCustomConstructor()) { 530 auto &body = construct->body().indent(); 531 for (const auto &it : llvm::enumerate(params)) { 532 body << formatv("auto {0} = std::get<{1}>(tblgenKey);\n", 533 it.value().getName(), it.index()); 534 } 535 // Use the parameters' custom allocator code, if provided. 536 FmtContext ctx = FmtContext().addSubst("_allocator", "allocator"); 537 for (auto ¶m : params) { 538 if (Optional<StringRef> allocCode = param.getAllocator()) { 539 ctx.withSelf(param.getName()).addSubst("_dst", param.getName()); 540 body << tgfmt(*allocCode, &ctx) << '\n'; 541 } 542 } 543 auto scope = 544 body.scope(strfmt("return new (allocator.allocate<{0}>()) {0}(", 545 def.getStorageClassName()), 546 ");"); 547 llvm::interleaveComma(params, body, 548 [&](auto ¶m) { body << param.getName(); }); 549 } 550 } 551 552 void DefGen::emitStorageClass() { 553 // Add the appropriate parent class. 554 storageCls->addParent(strfmt("::mlir::{0}Storage", valueType)); 555 // Add the constructor. 556 emitStorageConstructor(); 557 // Declare the key type. 558 emitKeyType(); 559 // Add the comparison method. 560 emitEquals(); 561 // Emit the key hash method. 562 emitHashKey(); 563 // Emit the storage constructor. Just declare it if the user wants to define 564 // it themself. 565 emitConstruct(); 566 // Emit the storage class members as public, at the very end of the struct. 567 storageCls->finalize(); 568 for (auto ¶m : params) 569 if (!isa<AttributeSelfTypeParameter>(param)) 570 storageCls->declare<Field>(param.getCppType(), param.getName()); 571 } 572 573 //===----------------------------------------------------------------------===// 574 // DefGenerator 575 //===----------------------------------------------------------------------===// 576 577 namespace { 578 /// This struct is the base generator used when processing tablegen interfaces. 579 class DefGenerator { 580 public: 581 bool emitDecls(StringRef selectedDialect); 582 bool emitDefs(StringRef selectedDialect); 583 584 protected: 585 DefGenerator(std::vector<llvm::Record *> &&defs, raw_ostream &os, 586 StringRef defType, StringRef valueType, bool isAttrGenerator) 587 : defRecords(std::move(defs)), os(os), defType(defType), 588 valueType(valueType), isAttrGenerator(isAttrGenerator) {} 589 590 /// Emit the list of def type names. 591 void emitTypeDefList(ArrayRef<AttrOrTypeDef> defs); 592 /// Emit the code to dispatch between different defs during parsing/printing. 593 void emitParsePrintDispatch(ArrayRef<AttrOrTypeDef> defs); 594 595 /// The set of def records to emit. 596 std::vector<llvm::Record *> defRecords; 597 /// The attribute or type class to emit. 598 /// The stream to emit to. 599 raw_ostream &os; 600 /// The prefix of the tablegen def name, e.g. Attr or Type. 601 StringRef defType; 602 /// The C++ base value type of the def, e.g. Attribute or Type. 603 StringRef valueType; 604 /// Flag indicating if this generator is for Attributes. False if the 605 /// generator is for types. 606 bool isAttrGenerator; 607 }; 608 609 /// A specialized generator for AttrDefs. 610 struct AttrDefGenerator : public DefGenerator { 611 AttrDefGenerator(const llvm::RecordKeeper &records, raw_ostream &os) 612 : DefGenerator(records.getAllDerivedDefinitionsIfDefined("AttrDef"), os, 613 "Attr", "Attribute", /*isAttrGenerator=*/true) {} 614 }; 615 /// A specialized generator for TypeDefs. 616 struct TypeDefGenerator : public DefGenerator { 617 TypeDefGenerator(const llvm::RecordKeeper &records, raw_ostream &os) 618 : DefGenerator(records.getAllDerivedDefinitionsIfDefined("TypeDef"), os, 619 "Type", "Type", /*isAttrGenerator=*/false) {} 620 }; 621 } // namespace 622 623 //===----------------------------------------------------------------------===// 624 // GEN: Declarations 625 //===----------------------------------------------------------------------===// 626 627 /// Print this above all the other declarations. Contains type declarations used 628 /// later on. 629 static const char *const typeDefDeclHeader = R"( 630 namespace mlir { 631 class AsmParser; 632 class AsmPrinter; 633 } // namespace mlir 634 )"; 635 636 bool DefGenerator::emitDecls(StringRef selectedDialect) { 637 emitSourceFileHeader((defType + "Def Declarations").str(), os); 638 IfDefScope scope("GET_" + defType.upper() + "DEF_CLASSES", os); 639 640 // Output the common "header". 641 os << typeDefDeclHeader; 642 643 SmallVector<AttrOrTypeDef, 16> defs; 644 collectAllDefs(selectedDialect, defRecords, defs); 645 if (defs.empty()) 646 return false; 647 { 648 NamespaceEmitter nsEmitter(os, defs.front().getDialect()); 649 650 // Declare all the def classes first (in case they reference each other). 651 for (const AttrOrTypeDef &def : defs) 652 os << "class " << def.getCppClassName() << ";\n"; 653 654 // Emit the declarations. 655 for (const AttrOrTypeDef &def : defs) 656 DefGen(def).emitDecl(os); 657 } 658 // Emit the TypeID explicit specializations to have a single definition for 659 // each of these. 660 for (const AttrOrTypeDef &def : defs) 661 if (!def.getDialect().getCppNamespace().empty()) 662 os << "MLIR_DECLARE_EXPLICIT_TYPE_ID(" 663 << def.getDialect().getCppNamespace() << "::" << def.getCppClassName() 664 << ")\n"; 665 666 return false; 667 } 668 669 //===----------------------------------------------------------------------===// 670 // GEN: Def List 671 //===----------------------------------------------------------------------===// 672 673 void DefGenerator::emitTypeDefList(ArrayRef<AttrOrTypeDef> defs) { 674 IfDefScope scope("GET_" + defType.upper() + "DEF_LIST", os); 675 auto interleaveFn = [&](const AttrOrTypeDef &def) { 676 os << def.getDialect().getCppNamespace() << "::" << def.getCppClassName(); 677 }; 678 llvm::interleave(defs, os, interleaveFn, ",\n"); 679 os << "\n"; 680 } 681 682 //===----------------------------------------------------------------------===// 683 // GEN: Definitions 684 //===----------------------------------------------------------------------===// 685 686 /// The code block for default attribute parser/printer dispatch boilerplate. 687 /// {0}: the dialect fully qualified class name. 688 /// {1}: the optional code for the dynamic attribute parser dispatch. 689 /// {2}: the optional code for the dynamic attribute printer dispatch. 690 static const char *const dialectDefaultAttrPrinterParserDispatch = R"( 691 /// Parse an attribute registered to this dialect. 692 ::mlir::Attribute {0}::parseAttribute(::mlir::DialectAsmParser &parser, 693 ::mlir::Type type) const {{ 694 ::llvm::SMLoc typeLoc = parser.getCurrentLocation(); 695 ::llvm::StringRef attrTag; 696 {{ 697 ::mlir::Attribute attr; 698 auto parseResult = generatedAttributeParser(parser, &attrTag, type, attr); 699 if (parseResult.hasValue()) 700 return attr; 701 } 702 {1} 703 parser.emitError(typeLoc) << "unknown attribute `" 704 << attrTag << "` in dialect `" << getNamespace() << "`"; 705 return {{}; 706 } 707 /// Print an attribute registered to this dialect. 708 void {0}::printAttribute(::mlir::Attribute attr, 709 ::mlir::DialectAsmPrinter &printer) const {{ 710 if (::mlir::succeeded(generatedAttributePrinter(attr, printer))) 711 return; 712 {2} 713 } 714 )"; 715 716 /// The code block for dynamic attribute parser dispatch boilerplate. 717 static const char *const dialectDynamicAttrParserDispatch = R"( 718 { 719 ::mlir::Attribute genAttr; 720 auto parseResult = parseOptionalDynamicAttr(attrTag, parser, genAttr); 721 if (parseResult.hasValue()) { 722 if (::mlir::succeeded(parseResult.getValue())) 723 return genAttr; 724 return Attribute(); 725 } 726 } 727 )"; 728 729 /// The code block for dynamic type printer dispatch boilerplate. 730 static const char *const dialectDynamicAttrPrinterDispatch = R"( 731 if (::mlir::succeeded(printIfDynamicAttr(attr, printer))) 732 return; 733 )"; 734 735 /// The code block for default type parser/printer dispatch boilerplate. 736 /// {0}: the dialect fully qualified class name. 737 /// {1}: the optional code for the dynamic type parser dispatch. 738 /// {2}: the optional code for the dynamic type printer dispatch. 739 static const char *const dialectDefaultTypePrinterParserDispatch = R"( 740 /// Parse a type registered to this dialect. 741 ::mlir::Type {0}::parseType(::mlir::DialectAsmParser &parser) const {{ 742 ::llvm::SMLoc typeLoc = parser.getCurrentLocation(); 743 ::llvm::StringRef mnemonic; 744 ::mlir::Type genType; 745 auto parseResult = generatedTypeParser(parser, &mnemonic, genType); 746 if (parseResult.hasValue()) 747 return genType; 748 {1} 749 parser.emitError(typeLoc) << "unknown type `" 750 << mnemonic << "` in dialect `" << getNamespace() << "`"; 751 return {{}; 752 } 753 /// Print a type registered to this dialect. 754 void {0}::printType(::mlir::Type type, 755 ::mlir::DialectAsmPrinter &printer) const {{ 756 if (::mlir::succeeded(generatedTypePrinter(type, printer))) 757 return; 758 {2} 759 } 760 )"; 761 762 /// The code block for dynamic type parser dispatch boilerplate. 763 static const char *const dialectDynamicTypeParserDispatch = R"( 764 { 765 auto parseResult = parseOptionalDynamicType(mnemonic, parser, genType); 766 if (parseResult.hasValue()) { 767 if (::mlir::succeeded(parseResult.getValue())) 768 return genType; 769 return Type(); 770 } 771 } 772 )"; 773 774 /// The code block for dynamic type printer dispatch boilerplate. 775 static const char *const dialectDynamicTypePrinterDispatch = R"( 776 if (::mlir::succeeded(printIfDynamicType(type, printer))) 777 return; 778 )"; 779 780 /// Emit the dialect printer/parser dispatcher. User's code should call these 781 /// functions from their dialect's print/parse methods. 782 void DefGenerator::emitParsePrintDispatch(ArrayRef<AttrOrTypeDef> defs) { 783 if (llvm::none_of(defs, [](const AttrOrTypeDef &def) { 784 return def.getMnemonic().has_value(); 785 })) { 786 return; 787 } 788 // Declare the parser. 789 SmallVector<MethodParameter> params = {{"::mlir::AsmParser &", "parser"}, 790 {"::llvm::StringRef *", "mnemonic"}}; 791 if (isAttrGenerator) 792 params.emplace_back("::mlir::Type", "type"); 793 params.emplace_back(strfmt("::mlir::{0} &", valueType), "value"); 794 Method parse("::mlir::OptionalParseResult", 795 strfmt("generated{0}Parser", valueType), Method::StaticInline, 796 std::move(params)); 797 // Declare the printer. 798 Method printer("::mlir::LogicalResult", 799 strfmt("generated{0}Printer", valueType), Method::StaticInline, 800 {{strfmt("::mlir::{0}", valueType), "def"}, 801 {"::mlir::AsmPrinter &", "printer"}}); 802 803 // The parser dispatch uses a KeywordSwitch, matching on the mnemonic and 804 // calling the def's parse function. 805 parse.body() << " return " 806 "::mlir::AsmParser::KeywordSwitch<::mlir::" 807 "OptionalParseResult>(parser)\n"; 808 const char *const getValueForMnemonic = 809 R"( .Case({0}::getMnemonic(), [&](llvm::StringRef, llvm::SMLoc) {{ 810 value = {0}::{1}; 811 return ::mlir::success(!!value); 812 }) 813 )"; 814 815 // The printer dispatch uses llvm::TypeSwitch to find and call the correct 816 // printer. 817 printer.body() << " return ::llvm::TypeSwitch<::mlir::" << valueType 818 << ", ::mlir::LogicalResult>(def)"; 819 const char *const printValue = R"( .Case<{0}>([&](auto t) {{ 820 printer << {0}::getMnemonic();{1} 821 return ::mlir::success(); 822 }) 823 )"; 824 for (auto &def : defs) { 825 if (!def.getMnemonic()) 826 continue; 827 bool hasParserPrinterDecl = 828 def.hasCustomAssemblyFormat() || def.getAssemblyFormat(); 829 std::string defClass = strfmt( 830 "{0}::{1}", def.getDialect().getCppNamespace(), def.getCppClassName()); 831 832 // If the def has no parameters or parser code, invoke a normal `get`. 833 std::string parseOrGet = 834 hasParserPrinterDecl 835 ? strfmt("parse(parser{0})", isAttrGenerator ? ", type" : "") 836 : "get(parser.getContext())"; 837 parse.body() << llvm::formatv(getValueForMnemonic, defClass, parseOrGet); 838 839 // If the def has no parameters and no printer, just print the mnemonic. 840 StringRef printDef = ""; 841 if (hasParserPrinterDecl) 842 printDef = "\nt.print(printer);"; 843 printer.body() << llvm::formatv(printValue, defClass, printDef); 844 } 845 parse.body() << " .Default([&](llvm::StringRef keyword, llvm::SMLoc) {\n" 846 " *mnemonic = keyword;\n" 847 " return llvm::None;\n" 848 " });"; 849 printer.body() << " .Default([](auto) { return ::mlir::failure(); });"; 850 851 raw_indented_ostream indentedOs(os); 852 parse.writeDeclTo(indentedOs); 853 printer.writeDeclTo(indentedOs); 854 } 855 856 bool DefGenerator::emitDefs(StringRef selectedDialect) { 857 emitSourceFileHeader((defType + "Def Definitions").str(), os); 858 859 SmallVector<AttrOrTypeDef, 16> defs; 860 collectAllDefs(selectedDialect, defRecords, defs); 861 if (defs.empty()) 862 return false; 863 emitTypeDefList(defs); 864 865 IfDefScope scope("GET_" + defType.upper() + "DEF_CLASSES", os); 866 emitParsePrintDispatch(defs); 867 for (const AttrOrTypeDef &def : defs) { 868 { 869 NamespaceEmitter ns(os, def.getDialect()); 870 DefGen gen(def); 871 gen.emitDef(os); 872 } 873 // Emit the TypeID explicit specializations to have a single symbol def. 874 if (!def.getDialect().getCppNamespace().empty()) 875 os << "MLIR_DEFINE_EXPLICIT_TYPE_ID(" 876 << def.getDialect().getCppNamespace() << "::" << def.getCppClassName() 877 << ")\n"; 878 } 879 880 Dialect firstDialect = defs.front().getDialect(); 881 882 // Emit the default parser/printer for Attributes if the dialect asked for it. 883 if (isAttrGenerator && firstDialect.useDefaultAttributePrinterParser()) { 884 NamespaceEmitter nsEmitter(os, firstDialect); 885 if (firstDialect.isExtensible()) { 886 os << llvm::formatv(dialectDefaultAttrPrinterParserDispatch, 887 firstDialect.getCppClassName(), 888 dialectDynamicAttrParserDispatch, 889 dialectDynamicAttrPrinterDispatch); 890 } else { 891 os << llvm::formatv(dialectDefaultAttrPrinterParserDispatch, 892 firstDialect.getCppClassName(), "", ""); 893 } 894 } 895 896 // Emit the default parser/printer for Types if the dialect asked for it. 897 if (!isAttrGenerator && firstDialect.useDefaultTypePrinterParser()) { 898 NamespaceEmitter nsEmitter(os, firstDialect); 899 if (firstDialect.isExtensible()) { 900 os << llvm::formatv(dialectDefaultTypePrinterParserDispatch, 901 firstDialect.getCppClassName(), 902 dialectDynamicTypeParserDispatch, 903 dialectDynamicTypePrinterDispatch); 904 } else { 905 os << llvm::formatv(dialectDefaultTypePrinterParserDispatch, 906 firstDialect.getCppClassName(), "", ""); 907 } 908 } 909 910 return false; 911 } 912 913 //===----------------------------------------------------------------------===// 914 // GEN: Registration hooks 915 //===----------------------------------------------------------------------===// 916 917 //===----------------------------------------------------------------------===// 918 // AttrDef 919 920 static llvm::cl::OptionCategory attrdefGenCat("Options for -gen-attrdef-*"); 921 static llvm::cl::opt<std::string> 922 attrDialect("attrdefs-dialect", 923 llvm::cl::desc("Generate attributes for this dialect"), 924 llvm::cl::cat(attrdefGenCat), llvm::cl::CommaSeparated); 925 926 static mlir::GenRegistration 927 genAttrDefs("gen-attrdef-defs", "Generate AttrDef definitions", 928 [](const llvm::RecordKeeper &records, raw_ostream &os) { 929 AttrDefGenerator generator(records, os); 930 return generator.emitDefs(attrDialect); 931 }); 932 static mlir::GenRegistration 933 genAttrDecls("gen-attrdef-decls", "Generate AttrDef declarations", 934 [](const llvm::RecordKeeper &records, raw_ostream &os) { 935 AttrDefGenerator generator(records, os); 936 return generator.emitDecls(attrDialect); 937 }); 938 939 //===----------------------------------------------------------------------===// 940 // TypeDef 941 942 static llvm::cl::OptionCategory typedefGenCat("Options for -gen-typedef-*"); 943 static llvm::cl::opt<std::string> 944 typeDialect("typedefs-dialect", 945 llvm::cl::desc("Generate types for this dialect"), 946 llvm::cl::cat(typedefGenCat), llvm::cl::CommaSeparated); 947 948 static mlir::GenRegistration 949 genTypeDefs("gen-typedef-defs", "Generate TypeDef definitions", 950 [](const llvm::RecordKeeper &records, raw_ostream &os) { 951 TypeDefGenerator generator(records, os); 952 return generator.emitDefs(typeDialect); 953 }); 954 static mlir::GenRegistration 955 genTypeDecls("gen-typedef-decls", "Generate TypeDef declarations", 956 [](const llvm::RecordKeeper &records, raw_ostream &os) { 957 TypeDefGenerator generator(records, os); 958 return generator.emitDecls(typeDialect); 959 }); 960