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