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