1 //===- RewriterGen.cpp - MLIR pattern rewriter 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 // RewriterGen uses pattern rewrite definitions to generate rewriter matchers. 10 // 11 //===----------------------------------------------------------------------===// 12 13 #include "mlir/Support/IndentedOstream.h" 14 #include "mlir/TableGen/Attribute.h" 15 #include "mlir/TableGen/Format.h" 16 #include "mlir/TableGen/GenInfo.h" 17 #include "mlir/TableGen/Operator.h" 18 #include "mlir/TableGen/Pattern.h" 19 #include "mlir/TableGen/Predicate.h" 20 #include "mlir/TableGen/Type.h" 21 #include "llvm/ADT/StringExtras.h" 22 #include "llvm/ADT/StringSet.h" 23 #include "llvm/Support/CommandLine.h" 24 #include "llvm/Support/Debug.h" 25 #include "llvm/Support/FormatAdapters.h" 26 #include "llvm/Support/PrettyStackTrace.h" 27 #include "llvm/Support/Signals.h" 28 #include "llvm/TableGen/Error.h" 29 #include "llvm/TableGen/Main.h" 30 #include "llvm/TableGen/Record.h" 31 #include "llvm/TableGen/TableGenBackend.h" 32 33 using namespace mlir; 34 using namespace mlir::tblgen; 35 36 using llvm::formatv; 37 using llvm::Record; 38 using llvm::RecordKeeper; 39 40 #define DEBUG_TYPE "mlir-tblgen-rewritergen" 41 42 namespace llvm { 43 template <> 44 struct format_provider<mlir::tblgen::Pattern::IdentifierLine> { 45 static void format(const mlir::tblgen::Pattern::IdentifierLine &v, 46 raw_ostream &os, StringRef style) { 47 os << v.first << ":" << v.second; 48 } 49 }; 50 } // end namespace llvm 51 52 //===----------------------------------------------------------------------===// 53 // PatternEmitter 54 //===----------------------------------------------------------------------===// 55 56 namespace { 57 class PatternEmitter { 58 public: 59 PatternEmitter(Record *pat, RecordOperatorMap *mapper, raw_ostream &os); 60 61 // Emits the mlir::RewritePattern struct named `rewriteName`. 62 void emit(StringRef rewriteName); 63 64 private: 65 // Emits the code for matching ops. 66 void emitMatchLogic(DagNode tree); 67 68 // Emits the code for rewriting ops. 69 void emitRewriteLogic(); 70 71 //===--------------------------------------------------------------------===// 72 // Match utilities 73 //===--------------------------------------------------------------------===// 74 75 // Emits C++ statements for matching the op constrained by the given DAG 76 // `tree`. 77 void emitOpMatch(DagNode tree, int depth); 78 79 // Emits C++ statements for matching the `argIndex`-th argument of the given 80 // DAG `tree` as an operand. 81 void emitOperandMatch(DagNode tree, int argIndex, int depth); 82 83 // Emits C++ statements for matching the `argIndex`-th argument of the given 84 // DAG `tree` as an attribute. 85 void emitAttributeMatch(DagNode tree, int argIndex, int depth); 86 87 // Emits C++ for checking a match with a corresponding match failure 88 // diagnostic. 89 void emitMatchCheck(int depth, const FmtObjectBase &matchFmt, 90 const llvm::formatv_object_base &failureFmt); 91 92 // Emits C++ for checking a match with a corresponding match failure 93 // diagnostics. 94 void emitMatchCheck(int depth, const std::string &matchStr, 95 const std::string &failureStr); 96 97 //===--------------------------------------------------------------------===// 98 // Rewrite utilities 99 //===--------------------------------------------------------------------===// 100 101 // The entry point for handling a result pattern rooted at `resultTree`. This 102 // method dispatches to concrete handlers according to `resultTree`'s kind and 103 // returns a symbol representing the whole value pack. Callers are expected to 104 // further resolve the symbol according to the specific use case. 105 // 106 // `depth` is the nesting level of `resultTree`; 0 means top-level result 107 // pattern. For top-level result pattern, `resultIndex` indicates which result 108 // of the matched root op this pattern is intended to replace, which can be 109 // used to deduce the result type of the op generated from this result 110 // pattern. 111 std::string handleResultPattern(DagNode resultTree, int resultIndex, 112 int depth); 113 114 // Emits the C++ statement to replace the matched DAG with a value built via 115 // calling native C++ code. 116 std::string handleReplaceWithNativeCodeCall(DagNode resultTree); 117 118 // Returns the symbol of the old value serving as the replacement. 119 StringRef handleReplaceWithValue(DagNode tree); 120 121 // Returns the location value to use. 122 std::pair<bool, std::string> getLocation(DagNode tree); 123 124 // Returns the location value to use. 125 std::string handleLocationDirective(DagNode tree); 126 127 // Emits the C++ statement to build a new op out of the given DAG `tree` and 128 // returns the variable name that this op is assigned to. If the root op in 129 // DAG `tree` has a specified name, the created op will be assigned to a 130 // variable of the given name. Otherwise, a unique name will be used as the 131 // result value name. 132 std::string handleOpCreation(DagNode tree, int resultIndex, int depth); 133 134 using ChildNodeIndexNameMap = DenseMap<unsigned, std::string>; 135 136 // Emits a local variable for each value and attribute to be used for creating 137 // an op. 138 void createSeparateLocalVarsForOpArgs(DagNode node, 139 ChildNodeIndexNameMap &childNodeNames); 140 141 // Emits the concrete arguments used to call an op's builder. 142 void supplyValuesForOpArgs(DagNode node, 143 const ChildNodeIndexNameMap &childNodeNames); 144 145 // Emits the local variables for holding all values as a whole and all named 146 // attributes as a whole to be used for creating an op. 147 void createAggregateLocalVarsForOpArgs( 148 DagNode node, const ChildNodeIndexNameMap &childNodeNames); 149 150 // Returns the C++ expression to construct a constant attribute of the given 151 // `value` for the given attribute kind `attr`. 152 std::string handleConstantAttr(Attribute attr, StringRef value); 153 154 // Returns the C++ expression to build an argument from the given DAG `leaf`. 155 // `patArgName` is used to bound the argument to the source pattern. 156 std::string handleOpArgument(DagLeaf leaf, StringRef patArgName); 157 158 //===--------------------------------------------------------------------===// 159 // General utilities 160 //===--------------------------------------------------------------------===// 161 162 // Collects all of the operations within the given dag tree. 163 void collectOps(DagNode tree, llvm::SmallPtrSetImpl<const Operator *> &ops); 164 165 // Returns a unique symbol for a local variable of the given `op`. 166 std::string getUniqueSymbol(const Operator *op); 167 168 //===--------------------------------------------------------------------===// 169 // Symbol utilities 170 //===--------------------------------------------------------------------===// 171 172 // Returns how many static values the given DAG `node` correspond to. 173 int getNodeValueCount(DagNode node); 174 175 private: 176 // Pattern instantiation location followed by the location of multiclass 177 // prototypes used. This is intended to be used as a whole to 178 // PrintFatalError() on errors. 179 ArrayRef<llvm::SMLoc> loc; 180 181 // Op's TableGen Record to wrapper object. 182 RecordOperatorMap *opMap; 183 184 // Handy wrapper for pattern being emitted. 185 Pattern pattern; 186 187 // Map for all bound symbols' info. 188 SymbolInfoMap symbolInfoMap; 189 190 // The next unused ID for newly created values. 191 unsigned nextValueId; 192 193 raw_indented_ostream os; 194 195 // Format contexts containing placeholder substitutions. 196 FmtContext fmtCtx; 197 198 // Number of op processed. 199 int opCounter = 0; 200 }; 201 } // end anonymous namespace 202 203 PatternEmitter::PatternEmitter(Record *pat, RecordOperatorMap *mapper, 204 raw_ostream &os) 205 : loc(pat->getLoc()), opMap(mapper), pattern(pat, mapper), 206 symbolInfoMap(pat->getLoc()), nextValueId(0), os(os) { 207 fmtCtx.withBuilder("rewriter"); 208 } 209 210 std::string PatternEmitter::handleConstantAttr(Attribute attr, 211 StringRef value) { 212 if (!attr.isConstBuildable()) 213 PrintFatalError(loc, "Attribute " + attr.getAttrDefName() + 214 " does not have the 'constBuilderCall' field"); 215 216 // TODO: Verify the constants here 217 return std::string(tgfmt(attr.getConstBuilderTemplate(), &fmtCtx, value)); 218 } 219 220 // Helper function to match patterns. 221 void PatternEmitter::emitOpMatch(DagNode tree, int depth) { 222 Operator &op = tree.getDialectOp(opMap); 223 LLVM_DEBUG(llvm::dbgs() << "start emitting match for op '" 224 << op.getOperationName() << "' at depth " << depth 225 << '\n'); 226 227 int indent = 4 + 2 * depth; 228 os.indent(indent) << formatv( 229 "auto castedOp{0} = ::llvm::dyn_cast_or_null<{1}>(op{0}); " 230 "(void)castedOp{0};\n", 231 depth, op.getQualCppClassName()); 232 // Skip the operand matching at depth 0 as the pattern rewriter already does. 233 if (depth != 0) { 234 // Skip if there is no defining operation (e.g., arguments to function). 235 os << formatv("if (!castedOp{0})\n return failure();\n", depth); 236 } 237 if (tree.getNumArgs() != op.getNumArgs()) { 238 PrintFatalError(loc, formatv("op '{0}' argument number mismatch: {1} in " 239 "pattern vs. {2} in definition", 240 op.getOperationName(), tree.getNumArgs(), 241 op.getNumArgs())); 242 } 243 244 // If the operand's name is set, set to that variable. 245 auto name = tree.getSymbol(); 246 if (!name.empty()) 247 os << formatv("{0} = castedOp{1};\n", name, depth); 248 249 for (int i = 0, e = tree.getNumArgs(); i != e; ++i) { 250 auto opArg = op.getArg(i); 251 252 // Handle nested DAG construct first 253 if (DagNode argTree = tree.getArgAsNestedDag(i)) { 254 if (auto *operand = opArg.dyn_cast<NamedTypeConstraint *>()) { 255 if (operand->isVariableLength()) { 256 auto error = formatv("use nested DAG construct to match op {0}'s " 257 "variadic operand #{1} unsupported now", 258 op.getOperationName(), i); 259 PrintFatalError(loc, error); 260 } 261 } 262 os << "{\n"; 263 264 os.indent() << formatv( 265 "auto *op{0} = " 266 "(*castedOp{1}.getODSOperands({2}).begin()).getDefiningOp();\n", 267 depth + 1, depth, i); 268 emitOpMatch(argTree, depth + 1); 269 os << formatv("tblgen_ops[{0}] = op{1};\n", ++opCounter, depth + 1); 270 os.unindent() << "}\n"; 271 continue; 272 } 273 274 // Next handle DAG leaf: operand or attribute 275 if (opArg.is<NamedTypeConstraint *>()) { 276 emitOperandMatch(tree, i, depth); 277 } else if (opArg.is<NamedAttribute *>()) { 278 emitAttributeMatch(tree, i, depth); 279 } else { 280 PrintFatalError(loc, "unhandled case when matching op"); 281 } 282 } 283 LLVM_DEBUG(llvm::dbgs() << "done emitting match for op '" 284 << op.getOperationName() << "' at depth " << depth 285 << '\n'); 286 } 287 288 void PatternEmitter::emitOperandMatch(DagNode tree, int argIndex, int depth) { 289 Operator &op = tree.getDialectOp(opMap); 290 auto *operand = op.getArg(argIndex).get<NamedTypeConstraint *>(); 291 auto matcher = tree.getArgAsLeaf(argIndex); 292 293 // If a constraint is specified, we need to generate C++ statements to 294 // check the constraint. 295 if (!matcher.isUnspecified()) { 296 if (!matcher.isOperandMatcher()) { 297 PrintFatalError( 298 loc, formatv("the {1}-th argument of op '{0}' should be an operand", 299 op.getOperationName(), argIndex + 1)); 300 } 301 302 // Only need to verify if the matcher's type is different from the one 303 // of op definition. 304 Constraint constraint = matcher.getAsConstraint(); 305 if (operand->constraint != constraint) { 306 if (operand->isVariableLength()) { 307 auto error = formatv( 308 "further constrain op {0}'s variadic operand #{1} unsupported now", 309 op.getOperationName(), argIndex); 310 PrintFatalError(loc, error); 311 } 312 auto self = 313 formatv("(*castedOp{0}.getODSOperands({1}).begin()).getType()", depth, 314 argIndex); 315 emitMatchCheck( 316 depth, 317 tgfmt(constraint.getConditionTemplate(), &fmtCtx.withSelf(self)), 318 formatv("\"operand {0} of op '{1}' failed to satisfy constraint: " 319 "'{2}'\"", 320 operand - op.operand_begin(), op.getOperationName(), 321 constraint.getDescription())); 322 } 323 } 324 325 // Capture the value 326 auto name = tree.getArgName(argIndex); 327 // `$_` is a special symbol to ignore op argument matching. 328 if (!name.empty() && name != "_") { 329 // We need to subtract the number of attributes before this operand to get 330 // the index in the operand list. 331 auto numPrevAttrs = std::count_if( 332 op.arg_begin(), op.arg_begin() + argIndex, 333 [](const Argument &arg) { return arg.is<NamedAttribute *>(); }); 334 335 auto res = symbolInfoMap.findBoundSymbol(name, op, argIndex); 336 os << formatv("{0} = castedOp{1}.getODSOperands({2});\n", 337 res->second.getVarName(name), depth, argIndex - numPrevAttrs); 338 } 339 } 340 341 void PatternEmitter::emitAttributeMatch(DagNode tree, int argIndex, int depth) { 342 Operator &op = tree.getDialectOp(opMap); 343 auto *namedAttr = op.getArg(argIndex).get<NamedAttribute *>(); 344 const auto &attr = namedAttr->attr; 345 346 os << "{\n"; 347 os.indent() << formatv( 348 "auto tblgen_attr = op{0}->getAttrOfType<{1}>(\"{2}\"); " 349 "(void)tblgen_attr;\n", 350 depth, attr.getStorageType(), namedAttr->name); 351 352 // TODO: This should use getter method to avoid duplication. 353 if (attr.hasDefaultValue()) { 354 os << "if (!tblgen_attr) tblgen_attr = " 355 << std::string(tgfmt(attr.getConstBuilderTemplate(), &fmtCtx, 356 attr.getDefaultValue())) 357 << ";\n"; 358 } else if (attr.isOptional()) { 359 // For a missing attribute that is optional according to definition, we 360 // should just capture a mlir::Attribute() to signal the missing state. 361 // That is precisely what getAttr() returns on missing attributes. 362 } else { 363 emitMatchCheck(depth, tgfmt("tblgen_attr", &fmtCtx), 364 formatv("\"expected op '{0}' to have attribute '{1}' " 365 "of type '{2}'\"", 366 op.getOperationName(), namedAttr->name, 367 attr.getStorageType())); 368 } 369 370 auto matcher = tree.getArgAsLeaf(argIndex); 371 if (!matcher.isUnspecified()) { 372 if (!matcher.isAttrMatcher()) { 373 PrintFatalError( 374 loc, formatv("the {1}-th argument of op '{0}' should be an attribute", 375 op.getOperationName(), argIndex + 1)); 376 } 377 378 // If a constraint is specified, we need to generate C++ statements to 379 // check the constraint. 380 emitMatchCheck( 381 depth, 382 tgfmt(matcher.getConditionTemplate(), &fmtCtx.withSelf("tblgen_attr")), 383 formatv("\"op '{0}' attribute '{1}' failed to satisfy constraint: " 384 "{2}\"", 385 op.getOperationName(), namedAttr->name, 386 matcher.getAsConstraint().getDescription())); 387 } 388 389 // Capture the value 390 auto name = tree.getArgName(argIndex); 391 // `$_` is a special symbol to ignore op argument matching. 392 if (!name.empty() && name != "_") { 393 os << formatv("{0} = tblgen_attr;\n", name); 394 } 395 396 os.unindent() << "}\n"; 397 } 398 399 void PatternEmitter::emitMatchCheck( 400 int depth, const FmtObjectBase &matchFmt, 401 const llvm::formatv_object_base &failureFmt) { 402 emitMatchCheck(depth, matchFmt.str(), failureFmt.str()); 403 } 404 405 void PatternEmitter::emitMatchCheck(int depth, const std::string &matchStr, 406 const std::string &failureStr) { 407 os << "if (!(" << matchStr << "))"; 408 os.scope("{\n", "\n}\n").os 409 << "return rewriter.notifyMatchFailure(op" << depth 410 << ", [&](::mlir::Diagnostic &diag) {\n diag << " << failureStr 411 << ";\n});"; 412 } 413 414 void PatternEmitter::emitMatchLogic(DagNode tree) { 415 LLVM_DEBUG(llvm::dbgs() << "--- start emitting match logic ---\n"); 416 int depth = 0; 417 emitOpMatch(tree, depth); 418 419 for (auto &appliedConstraint : pattern.getConstraints()) { 420 auto &constraint = appliedConstraint.constraint; 421 auto &entities = appliedConstraint.entities; 422 423 auto condition = constraint.getConditionTemplate(); 424 if (isa<TypeConstraint>(constraint)) { 425 auto self = formatv("({0}.getType())", 426 symbolInfoMap.getValueAndRangeUse(entities.front())); 427 emitMatchCheck( 428 depth, tgfmt(condition, &fmtCtx.withSelf(self.str())), 429 formatv("\"value entity '{0}' failed to satisfy constraint: {1}\"", 430 entities.front(), constraint.getDescription())); 431 432 } else if (isa<AttrConstraint>(constraint)) { 433 PrintFatalError( 434 loc, "cannot use AttrConstraint in Pattern multi-entity constraints"); 435 } else { 436 // TODO: replace formatv arguments with the exact specified 437 // args. 438 if (entities.size() > 4) { 439 PrintFatalError(loc, "only support up to 4-entity constraints now"); 440 } 441 SmallVector<std::string, 4> names; 442 int i = 0; 443 for (int e = entities.size(); i < e; ++i) 444 names.push_back(symbolInfoMap.getValueAndRangeUse(entities[i])); 445 std::string self = appliedConstraint.self; 446 if (!self.empty()) 447 self = symbolInfoMap.getValueAndRangeUse(self); 448 for (; i < 4; ++i) 449 names.push_back("<unused>"); 450 emitMatchCheck(depth, 451 tgfmt(condition, &fmtCtx.withSelf(self), names[0], 452 names[1], names[2], names[3]), 453 formatv("\"entities '{0}' failed to satisfy constraint: " 454 "{1}\"", 455 llvm::join(entities, ", "), 456 constraint.getDescription())); 457 } 458 } 459 460 // Some of the operands could be bound to the same symbol name, we need 461 // to enforce equality constraint on those. 462 // TODO: we should be able to emit equality checks early 463 // and short circuit unnecessary work if vars are not equal. 464 for (auto symbolInfoIt = symbolInfoMap.begin(); 465 symbolInfoIt != symbolInfoMap.end();) { 466 auto range = symbolInfoMap.getRangeOfEqualElements(symbolInfoIt->first); 467 auto startRange = range.first; 468 auto endRange = range.second; 469 470 auto firstOperand = symbolInfoIt->second.getVarName(symbolInfoIt->first); 471 for (++startRange; startRange != endRange; ++startRange) { 472 auto secondOperand = startRange->second.getVarName(symbolInfoIt->first); 473 emitMatchCheck( 474 depth, 475 formatv("*{0}.begin() == *{1}.begin()", firstOperand, secondOperand), 476 formatv("\"Operands '{0}' and '{1}' must be equal\"", firstOperand, 477 secondOperand)); 478 } 479 480 symbolInfoIt = endRange; 481 } 482 483 LLVM_DEBUG(llvm::dbgs() << "--- done emitting match logic ---\n"); 484 } 485 486 void PatternEmitter::collectOps(DagNode tree, 487 llvm::SmallPtrSetImpl<const Operator *> &ops) { 488 // Check if this tree is an operation. 489 if (tree.isOperation()) { 490 const Operator &op = tree.getDialectOp(opMap); 491 LLVM_DEBUG(llvm::dbgs() 492 << "found operation " << op.getOperationName() << '\n'); 493 ops.insert(&op); 494 } 495 496 // Recurse the arguments of the tree. 497 for (unsigned i = 0, e = tree.getNumArgs(); i != e; ++i) 498 if (auto child = tree.getArgAsNestedDag(i)) 499 collectOps(child, ops); 500 } 501 502 void PatternEmitter::emit(StringRef rewriteName) { 503 // Get the DAG tree for the source pattern. 504 DagNode sourceTree = pattern.getSourcePattern(); 505 506 const Operator &rootOp = pattern.getSourceRootOp(); 507 auto rootName = rootOp.getOperationName(); 508 509 // Collect the set of result operations. 510 llvm::SmallPtrSet<const Operator *, 4> resultOps; 511 LLVM_DEBUG(llvm::dbgs() << "start collecting ops used in result patterns\n"); 512 for (unsigned i = 0, e = pattern.getNumResultPatterns(); i != e; ++i) { 513 collectOps(pattern.getResultPattern(i), resultOps); 514 } 515 LLVM_DEBUG(llvm::dbgs() << "done collecting ops used in result patterns\n"); 516 517 // Emit RewritePattern for Pattern. 518 auto locs = pattern.getLocation(); 519 os << formatv("/* Generated from:\n {0:$[ instantiating\n ]}\n*/\n", 520 make_range(locs.rbegin(), locs.rend())); 521 os << formatv(R"(struct {0} : public ::mlir::RewritePattern { 522 {0}(::mlir::MLIRContext *context) 523 : ::mlir::RewritePattern("{1}", {{)", 524 rewriteName, rootName); 525 // Sort result operators by name. 526 llvm::SmallVector<const Operator *, 4> sortedResultOps(resultOps.begin(), 527 resultOps.end()); 528 llvm::sort(sortedResultOps, [&](const Operator *lhs, const Operator *rhs) { 529 return lhs->getOperationName() < rhs->getOperationName(); 530 }); 531 llvm::interleaveComma(sortedResultOps, os, [&](const Operator *op) { 532 os << '"' << op->getOperationName() << '"'; 533 }); 534 os << formatv(R"(}, {0}, context) {{})", pattern.getBenefit()) << "\n"; 535 536 // Emit matchAndRewrite() function. 537 { 538 auto classScope = os.scope(); 539 os.reindent(R"( 540 ::mlir::LogicalResult matchAndRewrite(::mlir::Operation *op0, 541 ::mlir::PatternRewriter &rewriter) const override {)") 542 << '\n'; 543 { 544 auto functionScope = os.scope(); 545 546 // Register all symbols bound in the source pattern. 547 pattern.collectSourcePatternBoundSymbols(symbolInfoMap); 548 549 LLVM_DEBUG(llvm::dbgs() 550 << "start creating local variables for capturing matches\n"); 551 os << "// Variables for capturing values and attributes used while " 552 "creating ops\n"; 553 // Create local variables for storing the arguments and results bound 554 // to symbols. 555 for (const auto &symbolInfoPair : symbolInfoMap) { 556 const auto &symbol = symbolInfoPair.first; 557 const auto &info = symbolInfoPair.second; 558 559 os << info.getVarDecl(symbol); 560 } 561 // TODO: capture ops with consistent numbering so that it can be 562 // reused for fused loc. 563 os << formatv("::mlir::Operation *tblgen_ops[{0}];\n\n", 564 pattern.getSourcePattern().getNumOps()); 565 LLVM_DEBUG(llvm::dbgs() 566 << "done creating local variables for capturing matches\n"); 567 568 os << "// Match\n"; 569 os << "tblgen_ops[0] = op0;\n"; 570 emitMatchLogic(sourceTree); 571 572 os << "\n// Rewrite\n"; 573 emitRewriteLogic(); 574 575 os << "return ::mlir::success();\n"; 576 } 577 os << "};\n"; 578 } 579 os << "};\n\n"; 580 } 581 582 void PatternEmitter::emitRewriteLogic() { 583 LLVM_DEBUG(llvm::dbgs() << "--- start emitting rewrite logic ---\n"); 584 const Operator &rootOp = pattern.getSourceRootOp(); 585 int numExpectedResults = rootOp.getNumResults(); 586 int numResultPatterns = pattern.getNumResultPatterns(); 587 588 // First register all symbols bound to ops generated in result patterns. 589 pattern.collectResultPatternBoundSymbols(symbolInfoMap); 590 591 // Only the last N static values generated are used to replace the matched 592 // root N-result op. We need to calculate the starting index (of the results 593 // of the matched op) each result pattern is to replace. 594 SmallVector<int, 4> offsets(numResultPatterns + 1, numExpectedResults); 595 // If we don't need to replace any value at all, set the replacement starting 596 // index as the number of result patterns so we skip all of them when trying 597 // to replace the matched op's results. 598 int replStartIndex = numExpectedResults == 0 ? numResultPatterns : -1; 599 for (int i = numResultPatterns - 1; i >= 0; --i) { 600 auto numValues = getNodeValueCount(pattern.getResultPattern(i)); 601 offsets[i] = offsets[i + 1] - numValues; 602 if (offsets[i] == 0) { 603 if (replStartIndex == -1) 604 replStartIndex = i; 605 } else if (offsets[i] < 0 && offsets[i + 1] > 0) { 606 auto error = formatv( 607 "cannot use the same multi-result op '{0}' to generate both " 608 "auxiliary values and values to be used for replacing the matched op", 609 pattern.getResultPattern(i).getSymbol()); 610 PrintFatalError(loc, error); 611 } 612 } 613 614 if (offsets.front() > 0) { 615 const char error[] = "no enough values generated to replace the matched op"; 616 PrintFatalError(loc, error); 617 } 618 619 os << "auto odsLoc = rewriter.getFusedLoc({"; 620 for (int i = 0, e = pattern.getSourcePattern().getNumOps(); i != e; ++i) { 621 os << (i ? ", " : "") << "tblgen_ops[" << i << "]->getLoc()"; 622 } 623 os << "}); (void)odsLoc;\n"; 624 625 // Process auxiliary result patterns. 626 for (int i = 0; i < replStartIndex; ++i) { 627 DagNode resultTree = pattern.getResultPattern(i); 628 auto val = handleResultPattern(resultTree, offsets[i], 0); 629 // Normal op creation will be streamed to `os` by the above call; but 630 // NativeCodeCall will only be materialized to `os` if it is used. Here 631 // we are handling auxiliary patterns so we want the side effect even if 632 // NativeCodeCall is not replacing matched root op's results. 633 if (resultTree.isNativeCodeCall()) 634 os << val << ";\n"; 635 } 636 637 if (numExpectedResults == 0) { 638 assert(replStartIndex >= numResultPatterns && 639 "invalid auxiliary vs. replacement pattern division!"); 640 // No result to replace. Just erase the op. 641 os << "rewriter.eraseOp(op0);\n"; 642 } else { 643 // Process replacement result patterns. 644 os << "::llvm::SmallVector<::mlir::Value, 4> tblgen_repl_values;\n"; 645 for (int i = replStartIndex; i < numResultPatterns; ++i) { 646 DagNode resultTree = pattern.getResultPattern(i); 647 auto val = handleResultPattern(resultTree, offsets[i], 0); 648 os << "\n"; 649 // Resolve each symbol for all range use so that we can loop over them. 650 // We need an explicit cast to `SmallVector` to capture the cases where 651 // `{0}` resolves to an `Operation::result_range` as well as cases that 652 // are not iterable (e.g. vector that gets wrapped in additional braces by 653 // RewriterGen). 654 // TODO: Revisit the need for materializing a vector. 655 os << symbolInfoMap.getAllRangeUse( 656 val, 657 "for (auto v: ::llvm::SmallVector<::mlir::Value, 4>{ {0} }) {{\n" 658 " tblgen_repl_values.push_back(v);\n}\n", 659 "\n"); 660 } 661 os << "\nrewriter.replaceOp(op0, tblgen_repl_values);\n"; 662 } 663 664 LLVM_DEBUG(llvm::dbgs() << "--- done emitting rewrite logic ---\n"); 665 } 666 667 std::string PatternEmitter::getUniqueSymbol(const Operator *op) { 668 return std::string( 669 formatv("tblgen_{0}_{1}", op->getCppClassName(), nextValueId++)); 670 } 671 672 std::string PatternEmitter::handleResultPattern(DagNode resultTree, 673 int resultIndex, int depth) { 674 LLVM_DEBUG(llvm::dbgs() << "handle result pattern: "); 675 LLVM_DEBUG(resultTree.print(llvm::dbgs())); 676 LLVM_DEBUG(llvm::dbgs() << '\n'); 677 678 if (resultTree.isLocationDirective()) { 679 PrintFatalError(loc, 680 "location directive can only be used with op creation"); 681 } 682 683 if (resultTree.isNativeCodeCall()) { 684 auto symbol = handleReplaceWithNativeCodeCall(resultTree); 685 symbolInfoMap.bindValue(symbol); 686 return symbol; 687 } 688 689 if (resultTree.isReplaceWithValue()) 690 return handleReplaceWithValue(resultTree).str(); 691 692 // Normal op creation. 693 auto symbol = handleOpCreation(resultTree, resultIndex, depth); 694 if (resultTree.getSymbol().empty()) { 695 // This is an op not explicitly bound to a symbol in the rewrite rule. 696 // Register the auto-generated symbol for it. 697 symbolInfoMap.bindOpResult(symbol, pattern.getDialectOp(resultTree)); 698 } 699 return symbol; 700 } 701 702 StringRef PatternEmitter::handleReplaceWithValue(DagNode tree) { 703 assert(tree.isReplaceWithValue()); 704 705 if (tree.getNumArgs() != 1) { 706 PrintFatalError( 707 loc, "replaceWithValue directive must take exactly one argument"); 708 } 709 710 if (!tree.getSymbol().empty()) { 711 PrintFatalError(loc, "cannot bind symbol to replaceWithValue"); 712 } 713 714 return tree.getArgName(0); 715 } 716 717 std::string PatternEmitter::handleLocationDirective(DagNode tree) { 718 assert(tree.isLocationDirective()); 719 auto lookUpArgLoc = [this, &tree](int idx) { 720 const auto *const lookupFmt = "(*{0}.begin()).getLoc()"; 721 return symbolInfoMap.getAllRangeUse(tree.getArgName(idx), lookupFmt); 722 }; 723 724 if (tree.getNumArgs() == 0) 725 llvm::PrintFatalError( 726 "At least one argument to location directive required"); 727 728 if (!tree.getSymbol().empty()) 729 PrintFatalError(loc, "cannot bind symbol to location"); 730 731 if (tree.getNumArgs() == 1) { 732 DagLeaf leaf = tree.getArgAsLeaf(0); 733 if (leaf.isStringAttr()) 734 return formatv("::mlir::NameLoc::get(rewriter.getIdentifier(\"{0}\"), " 735 "rewriter.getContext())", 736 leaf.getStringAttr()) 737 .str(); 738 return lookUpArgLoc(0); 739 } 740 741 std::string ret; 742 llvm::raw_string_ostream os(ret); 743 std::string strAttr; 744 os << "rewriter.getFusedLoc({"; 745 bool first = true; 746 for (int i = 0, e = tree.getNumArgs(); i != e; ++i) { 747 DagLeaf leaf = tree.getArgAsLeaf(i); 748 // Handle the optional string value. 749 if (leaf.isStringAttr()) { 750 if (!strAttr.empty()) 751 llvm::PrintFatalError("Only one string attribute may be specified"); 752 strAttr = leaf.getStringAttr(); 753 continue; 754 } 755 os << (first ? "" : ", ") << lookUpArgLoc(i); 756 first = false; 757 } 758 os << "}"; 759 if (!strAttr.empty()) { 760 os << ", rewriter.getStringAttr(\"" << strAttr << "\")"; 761 } 762 os << ")"; 763 return os.str(); 764 } 765 766 std::string PatternEmitter::handleOpArgument(DagLeaf leaf, 767 StringRef patArgName) { 768 if (leaf.isStringAttr()) 769 PrintFatalError(loc, "raw string not supported as argument"); 770 if (leaf.isConstantAttr()) { 771 auto constAttr = leaf.getAsConstantAttr(); 772 return handleConstantAttr(constAttr.getAttribute(), 773 constAttr.getConstantValue()); 774 } 775 if (leaf.isEnumAttrCase()) { 776 auto enumCase = leaf.getAsEnumAttrCase(); 777 if (enumCase.isStrCase()) 778 return handleConstantAttr(enumCase, enumCase.getSymbol()); 779 // This is an enum case backed by an IntegerAttr. We need to get its value 780 // to build the constant. 781 std::string val = std::to_string(enumCase.getValue()); 782 return handleConstantAttr(enumCase, val); 783 } 784 785 LLVM_DEBUG(llvm::dbgs() << "handle argument '" << patArgName << "'\n"); 786 auto argName = symbolInfoMap.getValueAndRangeUse(patArgName); 787 if (leaf.isUnspecified() || leaf.isOperandMatcher()) { 788 LLVM_DEBUG(llvm::dbgs() << "replace " << patArgName << " with '" << argName 789 << "' (via symbol ref)\n"); 790 return argName; 791 } 792 if (leaf.isNativeCodeCall()) { 793 auto repl = tgfmt(leaf.getNativeCodeTemplate(), &fmtCtx.withSelf(argName)); 794 LLVM_DEBUG(llvm::dbgs() << "replace " << patArgName << " with '" << repl 795 << "' (via NativeCodeCall)\n"); 796 return std::string(repl); 797 } 798 PrintFatalError(loc, "unhandled case when rewriting op"); 799 } 800 801 std::string PatternEmitter::handleReplaceWithNativeCodeCall(DagNode tree) { 802 LLVM_DEBUG(llvm::dbgs() << "handle NativeCodeCall pattern: "); 803 LLVM_DEBUG(tree.print(llvm::dbgs())); 804 LLVM_DEBUG(llvm::dbgs() << '\n'); 805 806 auto fmt = tree.getNativeCodeTemplate(); 807 // TODO: replace formatv arguments with the exact specified args. 808 SmallVector<std::string, 8> attrs(8); 809 if (tree.getNumArgs() > 8) { 810 PrintFatalError(loc, "unsupported NativeCodeCall argument numbers: " + 811 Twine(tree.getNumArgs())); 812 } 813 bool hasLocationDirective; 814 std::string locToUse; 815 std::tie(hasLocationDirective, locToUse) = getLocation(tree); 816 817 for (int i = 0, e = tree.getNumArgs() - hasLocationDirective; i != e; ++i) { 818 attrs[i] = handleOpArgument(tree.getArgAsLeaf(i), tree.getArgName(i)); 819 LLVM_DEBUG(llvm::dbgs() << "NativeCodeCall argument #" << i 820 << " replacement: " << attrs[i] << "\n"); 821 } 822 return std::string(tgfmt(fmt, &fmtCtx.addSubst("_loc", locToUse), attrs[0], 823 attrs[1], attrs[2], attrs[3], attrs[4], attrs[5], 824 attrs[6], attrs[7])); 825 } 826 827 int PatternEmitter::getNodeValueCount(DagNode node) { 828 if (node.isOperation()) { 829 // If the op is bound to a symbol in the rewrite rule, query its result 830 // count from the symbol info map. 831 auto symbol = node.getSymbol(); 832 if (!symbol.empty()) { 833 return symbolInfoMap.getStaticValueCount(symbol); 834 } 835 // Otherwise this is an unbound op; we will use all its results. 836 return pattern.getDialectOp(node).getNumResults(); 837 } 838 // TODO: This considers all NativeCodeCall as returning one 839 // value. Enhance if multi-value ones are needed. 840 return 1; 841 } 842 843 std::pair<bool, std::string> PatternEmitter::getLocation(DagNode tree) { 844 auto numPatArgs = tree.getNumArgs(); 845 846 if (numPatArgs != 0) { 847 if (auto lastArg = tree.getArgAsNestedDag(numPatArgs - 1)) 848 if (lastArg.isLocationDirective()) { 849 return std::make_pair(true, handleLocationDirective(lastArg)); 850 } 851 } 852 853 // If no explicit location is given, use the default, all fused, location. 854 return std::make_pair(false, "odsLoc"); 855 } 856 857 std::string PatternEmitter::handleOpCreation(DagNode tree, int resultIndex, 858 int depth) { 859 LLVM_DEBUG(llvm::dbgs() << "create op for pattern: "); 860 LLVM_DEBUG(tree.print(llvm::dbgs())); 861 LLVM_DEBUG(llvm::dbgs() << '\n'); 862 863 Operator &resultOp = tree.getDialectOp(opMap); 864 auto numOpArgs = resultOp.getNumArgs(); 865 auto numPatArgs = tree.getNumArgs(); 866 867 bool hasLocationDirective; 868 std::string locToUse; 869 std::tie(hasLocationDirective, locToUse) = getLocation(tree); 870 871 auto inPattern = numPatArgs - hasLocationDirective; 872 if (numOpArgs != inPattern) { 873 PrintFatalError(loc, 874 formatv("resultant op '{0}' argument number mismatch: " 875 "{1} in pattern vs. {2} in definition", 876 resultOp.getOperationName(), inPattern, numOpArgs)); 877 } 878 879 // A map to collect all nested DAG child nodes' names, with operand index as 880 // the key. This includes both bound and unbound child nodes. 881 ChildNodeIndexNameMap childNodeNames; 882 883 // First go through all the child nodes who are nested DAG constructs to 884 // create ops for them and remember the symbol names for them, so that we can 885 // use the results in the current node. This happens in a recursive manner. 886 for (int i = 0, e = resultOp.getNumOperands(); i != e; ++i) { 887 if (auto child = tree.getArgAsNestedDag(i)) 888 childNodeNames[i] = handleResultPattern(child, i, depth + 1); 889 } 890 891 // The name of the local variable holding this op. 892 std::string valuePackName; 893 // The symbol for holding the result of this pattern. Note that the result of 894 // this pattern is not necessarily the same as the variable created by this 895 // pattern because we can use `__N` suffix to refer only a specific result if 896 // the generated op is a multi-result op. 897 std::string resultValue; 898 if (tree.getSymbol().empty()) { 899 // No symbol is explicitly bound to this op in the pattern. Generate a 900 // unique name. 901 valuePackName = resultValue = getUniqueSymbol(&resultOp); 902 } else { 903 resultValue = std::string(tree.getSymbol()); 904 // Strip the index to get the name for the value pack and use it to name the 905 // local variable for the op. 906 valuePackName = std::string(SymbolInfoMap::getValuePackName(resultValue)); 907 } 908 909 // Create the local variable for this op. 910 os << formatv("{0} {1};\n{{\n", resultOp.getQualCppClassName(), 911 valuePackName); 912 913 // Right now ODS don't have general type inference support. Except a few 914 // special cases listed below, DRR needs to supply types for all results 915 // when building an op. 916 bool isSameOperandsAndResultType = 917 resultOp.getTrait("::mlir::OpTrait::SameOperandsAndResultType"); 918 bool useFirstAttr = 919 resultOp.getTrait("::mlir::OpTrait::FirstAttrDerivedResultType"); 920 921 if (isSameOperandsAndResultType || useFirstAttr) { 922 // We know how to deduce the result type for ops with these traits and we've 923 // generated builders taking aggregate parameters. Use those builders to 924 // create the ops. 925 926 // First prepare local variables for op arguments used in builder call. 927 createAggregateLocalVarsForOpArgs(tree, childNodeNames); 928 929 // Then create the op. 930 os.scope("", "\n}\n").os << formatv( 931 "{0} = rewriter.create<{1}>({2}, tblgen_values, tblgen_attrs);", 932 valuePackName, resultOp.getQualCppClassName(), locToUse); 933 return resultValue; 934 } 935 936 bool usePartialResults = valuePackName != resultValue; 937 938 if (usePartialResults || depth > 0 || resultIndex < 0) { 939 // For these cases (broadcastable ops, op results used both as auxiliary 940 // values and replacement values, ops in nested patterns, auxiliary ops), we 941 // still need to supply the result types when building the op. But because 942 // we don't generate a builder automatically with ODS for them, it's the 943 // developer's responsibility to make sure such a builder (with result type 944 // deduction ability) exists. We go through the separate-parameter builder 945 // here given that it's easier for developers to write compared to 946 // aggregate-parameter builders. 947 createSeparateLocalVarsForOpArgs(tree, childNodeNames); 948 949 os.scope().os << formatv("{0} = rewriter.create<{1}>({2}", valuePackName, 950 resultOp.getQualCppClassName(), locToUse); 951 supplyValuesForOpArgs(tree, childNodeNames); 952 os << "\n );\n}\n"; 953 return resultValue; 954 } 955 956 // If depth == 0 and resultIndex >= 0, it means we are replacing the values 957 // generated from the source pattern root op. Then we can use the source 958 // pattern's value types to determine the value type of the generated op 959 // here. 960 961 // First prepare local variables for op arguments used in builder call. 962 createAggregateLocalVarsForOpArgs(tree, childNodeNames); 963 964 // Then prepare the result types. We need to specify the types for all 965 // results. 966 os.indent() << formatv("::mlir::SmallVector<::mlir::Type, 4> tblgen_types; " 967 "(void)tblgen_types;\n"); 968 int numResults = resultOp.getNumResults(); 969 if (numResults != 0) { 970 for (int i = 0; i < numResults; ++i) 971 os << formatv("for (auto v: castedOp0.getODSResults({0})) {{\n" 972 " tblgen_types.push_back(v.getType());\n}\n", 973 resultIndex + i); 974 } 975 os << formatv("{0} = rewriter.create<{1}>({2}, tblgen_types, " 976 "tblgen_values, tblgen_attrs);\n", 977 valuePackName, resultOp.getQualCppClassName(), locToUse); 978 os.unindent() << "}\n"; 979 return resultValue; 980 } 981 982 void PatternEmitter::createSeparateLocalVarsForOpArgs( 983 DagNode node, ChildNodeIndexNameMap &childNodeNames) { 984 Operator &resultOp = node.getDialectOp(opMap); 985 986 // Now prepare operands used for building this op: 987 // * If the operand is non-variadic, we create a `Value` local variable. 988 // * If the operand is variadic, we create a `SmallVector<Value>` local 989 // variable. 990 991 int valueIndex = 0; // An index for uniquing local variable names. 992 for (int argIndex = 0, e = resultOp.getNumArgs(); argIndex < e; ++argIndex) { 993 const auto *operand = 994 resultOp.getArg(argIndex).dyn_cast<NamedTypeConstraint *>(); 995 // We do not need special handling for attributes. 996 if (!operand) 997 continue; 998 999 raw_indented_ostream::DelimitedScope scope(os); 1000 std::string varName; 1001 if (operand->isVariadic()) { 1002 varName = std::string(formatv("tblgen_values_{0}", valueIndex++)); 1003 os << formatv("::mlir::SmallVector<::mlir::Value, 4> {0};\n", varName); 1004 std::string range; 1005 if (node.isNestedDagArg(argIndex)) { 1006 range = childNodeNames[argIndex]; 1007 } else { 1008 range = std::string(node.getArgName(argIndex)); 1009 } 1010 // Resolve the symbol for all range use so that we have a uniform way of 1011 // capturing the values. 1012 range = symbolInfoMap.getValueAndRangeUse(range); 1013 os << formatv("for (auto v: {0}) {{\n {1}.push_back(v);\n}\n", range, 1014 varName); 1015 } else { 1016 varName = std::string(formatv("tblgen_value_{0}", valueIndex++)); 1017 os << formatv("::mlir::Value {0} = ", varName); 1018 if (node.isNestedDagArg(argIndex)) { 1019 os << symbolInfoMap.getValueAndRangeUse(childNodeNames[argIndex]); 1020 } else { 1021 DagLeaf leaf = node.getArgAsLeaf(argIndex); 1022 auto symbol = 1023 symbolInfoMap.getValueAndRangeUse(node.getArgName(argIndex)); 1024 if (leaf.isNativeCodeCall()) { 1025 os << std::string( 1026 tgfmt(leaf.getNativeCodeTemplate(), &fmtCtx.withSelf(symbol))); 1027 } else { 1028 os << symbol; 1029 } 1030 } 1031 os << ";\n"; 1032 } 1033 1034 // Update to use the newly created local variable for building the op later. 1035 childNodeNames[argIndex] = varName; 1036 } 1037 } 1038 1039 void PatternEmitter::supplyValuesForOpArgs( 1040 DagNode node, const ChildNodeIndexNameMap &childNodeNames) { 1041 Operator &resultOp = node.getDialectOp(opMap); 1042 for (int argIndex = 0, numOpArgs = resultOp.getNumArgs(); 1043 argIndex != numOpArgs; ++argIndex) { 1044 // Start each argument on its own line. 1045 os << ",\n "; 1046 1047 Argument opArg = resultOp.getArg(argIndex); 1048 // Handle the case of operand first. 1049 if (auto *operand = opArg.dyn_cast<NamedTypeConstraint *>()) { 1050 if (!operand->name.empty()) 1051 os << "/*" << operand->name << "=*/"; 1052 os << childNodeNames.lookup(argIndex); 1053 continue; 1054 } 1055 1056 // The argument in the op definition. 1057 auto opArgName = resultOp.getArgName(argIndex); 1058 if (auto subTree = node.getArgAsNestedDag(argIndex)) { 1059 if (!subTree.isNativeCodeCall()) 1060 PrintFatalError(loc, "only NativeCodeCall allowed in nested dag node " 1061 "for creating attribute"); 1062 os << formatv("/*{0}=*/{1}", opArgName, 1063 handleReplaceWithNativeCodeCall(subTree)); 1064 } else { 1065 auto leaf = node.getArgAsLeaf(argIndex); 1066 // The argument in the result DAG pattern. 1067 auto patArgName = node.getArgName(argIndex); 1068 if (leaf.isConstantAttr() || leaf.isEnumAttrCase()) { 1069 // TODO: Refactor out into map to avoid recomputing these. 1070 if (!opArg.is<NamedAttribute *>()) 1071 PrintFatalError(loc, Twine("expected attribute ") + Twine(argIndex)); 1072 if (!patArgName.empty()) 1073 os << "/*" << patArgName << "=*/"; 1074 } else { 1075 os << "/*" << opArgName << "=*/"; 1076 } 1077 os << handleOpArgument(leaf, patArgName); 1078 } 1079 } 1080 } 1081 1082 void PatternEmitter::createAggregateLocalVarsForOpArgs( 1083 DagNode node, const ChildNodeIndexNameMap &childNodeNames) { 1084 Operator &resultOp = node.getDialectOp(opMap); 1085 1086 auto scope = os.scope(); 1087 os << formatv("::mlir::SmallVector<::mlir::Value, 4> " 1088 "tblgen_values; (void)tblgen_values;\n"); 1089 os << formatv("::mlir::SmallVector<::mlir::NamedAttribute, 4> " 1090 "tblgen_attrs; (void)tblgen_attrs;\n"); 1091 1092 const char *addAttrCmd = 1093 "if (auto tmpAttr = {1}) {\n" 1094 " tblgen_attrs.emplace_back(rewriter.getIdentifier(\"{0}\"), " 1095 "tmpAttr);\n}\n"; 1096 for (int argIndex = 0, e = resultOp.getNumArgs(); argIndex < e; ++argIndex) { 1097 if (resultOp.getArg(argIndex).is<NamedAttribute *>()) { 1098 // The argument in the op definition. 1099 auto opArgName = resultOp.getArgName(argIndex); 1100 if (auto subTree = node.getArgAsNestedDag(argIndex)) { 1101 if (!subTree.isNativeCodeCall()) 1102 PrintFatalError(loc, "only NativeCodeCall allowed in nested dag node " 1103 "for creating attribute"); 1104 os << formatv(addAttrCmd, opArgName, 1105 handleReplaceWithNativeCodeCall(subTree)); 1106 } else { 1107 auto leaf = node.getArgAsLeaf(argIndex); 1108 // The argument in the result DAG pattern. 1109 auto patArgName = node.getArgName(argIndex); 1110 os << formatv(addAttrCmd, opArgName, 1111 handleOpArgument(leaf, patArgName)); 1112 } 1113 continue; 1114 } 1115 1116 const auto *operand = 1117 resultOp.getArg(argIndex).get<NamedTypeConstraint *>(); 1118 std::string varName; 1119 if (operand->isVariadic()) { 1120 std::string range; 1121 if (node.isNestedDagArg(argIndex)) { 1122 range = childNodeNames.lookup(argIndex); 1123 } else { 1124 range = std::string(node.getArgName(argIndex)); 1125 } 1126 // Resolve the symbol for all range use so that we have a uniform way of 1127 // capturing the values. 1128 range = symbolInfoMap.getValueAndRangeUse(range); 1129 os << formatv("for (auto v: {0}) {{\n tblgen_values.push_back(v);\n}\n", 1130 range); 1131 } else { 1132 os << formatv("tblgen_values.push_back("); 1133 if (node.isNestedDagArg(argIndex)) { 1134 os << symbolInfoMap.getValueAndRangeUse( 1135 childNodeNames.lookup(argIndex)); 1136 } else { 1137 DagLeaf leaf = node.getArgAsLeaf(argIndex); 1138 auto symbol = 1139 symbolInfoMap.getValueAndRangeUse(node.getArgName(argIndex)); 1140 if (leaf.isNativeCodeCall()) { 1141 os << std::string( 1142 tgfmt(leaf.getNativeCodeTemplate(), &fmtCtx.withSelf(symbol))); 1143 } else { 1144 os << symbol; 1145 } 1146 } 1147 os << ");\n"; 1148 } 1149 } 1150 } 1151 1152 static void emitRewriters(const RecordKeeper &recordKeeper, raw_ostream &os) { 1153 emitSourceFileHeader("Rewriters", os); 1154 1155 const auto &patterns = recordKeeper.getAllDerivedDefinitions("Pattern"); 1156 auto numPatterns = patterns.size(); 1157 1158 // We put the map here because it can be shared among multiple patterns. 1159 RecordOperatorMap recordOpMap; 1160 1161 std::vector<std::string> rewriterNames; 1162 rewriterNames.reserve(numPatterns); 1163 1164 std::string baseRewriterName = "GeneratedConvert"; 1165 int rewriterIndex = 0; 1166 1167 for (Record *p : patterns) { 1168 std::string name; 1169 if (p->isAnonymous()) { 1170 // If no name is provided, ensure unique rewriter names simply by 1171 // appending unique suffix. 1172 name = baseRewriterName + llvm::utostr(rewriterIndex++); 1173 } else { 1174 name = std::string(p->getName()); 1175 } 1176 LLVM_DEBUG(llvm::dbgs() 1177 << "=== start generating pattern '" << name << "' ===\n"); 1178 PatternEmitter(p, &recordOpMap, os).emit(name); 1179 LLVM_DEBUG(llvm::dbgs() 1180 << "=== done generating pattern '" << name << "' ===\n"); 1181 rewriterNames.push_back(std::move(name)); 1182 } 1183 1184 // Emit function to add the generated matchers to the pattern list. 1185 os << "void LLVM_ATTRIBUTE_UNUSED populateWithGenerated(::mlir::MLIRContext " 1186 "*context, ::mlir::OwningRewritePatternList &patterns) {\n"; 1187 for (const auto &name : rewriterNames) { 1188 os << " patterns.insert<" << name << ">(context);\n"; 1189 } 1190 os << "}\n"; 1191 } 1192 1193 static mlir::GenRegistration 1194 genRewriters("gen-rewriters", "Generate pattern rewriters", 1195 [](const RecordKeeper &records, raw_ostream &os) { 1196 emitRewriters(records, os); 1197 return false; 1198 }); 1199