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