1 //===- AsmPrinter.cpp - MLIR Assembly Printer Implementation --------------===//
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 // This file implements the MLIR AsmPrinter class, which is used to implement
10 // the various print() methods on the core IR objects.
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "mlir/IR/AffineExpr.h"
15 #include "mlir/IR/AffineMap.h"
16 #include "mlir/IR/AsmState.h"
17 #include "mlir/IR/Attributes.h"
18 #include "mlir/IR/BuiltinTypes.h"
19 #include "mlir/IR/Dialect.h"
20 #include "mlir/IR/DialectImplementation.h"
21 #include "mlir/IR/IntegerSet.h"
22 #include "mlir/IR/MLIRContext.h"
23 #include "mlir/IR/OpImplementation.h"
24 #include "mlir/IR/Operation.h"
25 #include "llvm/ADT/APFloat.h"
26 #include "llvm/ADT/DenseMap.h"
27 #include "llvm/ADT/MapVector.h"
28 #include "llvm/ADT/STLExtras.h"
29 #include "llvm/ADT/ScopedHashTable.h"
30 #include "llvm/ADT/SetVector.h"
31 #include "llvm/ADT/SmallString.h"
32 #include "llvm/ADT/StringExtras.h"
33 #include "llvm/ADT/StringSet.h"
34 #include "llvm/ADT/TypeSwitch.h"
35 #include "llvm/Support/CommandLine.h"
36 #include "llvm/Support/Endian.h"
37 #include "llvm/Support/Regex.h"
38 #include "llvm/Support/SaveAndRestore.h"
39 using namespace mlir;
40 using namespace mlir::detail;
41 
42 void Identifier::print(raw_ostream &os) const { os << str(); }
43 
44 void Identifier::dump() const { print(llvm::errs()); }
45 
46 void OperationName::print(raw_ostream &os) const { os << getStringRef(); }
47 
48 void OperationName::dump() const { print(llvm::errs()); }
49 
50 DialectAsmPrinter::~DialectAsmPrinter() {}
51 
52 //===--------------------------------------------------------------------===//
53 // OpAsmPrinter
54 //===--------------------------------------------------------------------===//
55 
56 OpAsmPrinter::~OpAsmPrinter() {}
57 
58 void OpAsmPrinter::printFunctionalType(Operation *op) {
59   auto &os = getStream();
60   os << '(';
61   llvm::interleaveComma(op->getOperands(), os, [&](Value operand) {
62     // Print the types of null values as <<NULL TYPE>>.
63     *this << (operand ? operand.getType() : Type());
64   });
65   os << ") -> ";
66 
67   // Print the result list.  We don't parenthesize single result types unless
68   // it is a function (avoiding a grammar ambiguity).
69   bool wrapped = op->getNumResults() != 1;
70   if (!wrapped && op->getResult(0).getType() &&
71       op->getResult(0).getType().isa<FunctionType>())
72     wrapped = true;
73 
74   if (wrapped)
75     os << '(';
76 
77   llvm::interleaveComma(op->getResults(), os, [&](const OpResult &result) {
78     // Print the types of null values as <<NULL TYPE>>.
79     *this << (result ? result.getType() : Type());
80   });
81 
82   if (wrapped)
83     os << ')';
84 }
85 
86 //===--------------------------------------------------------------------===//
87 // Operation OpAsm interface.
88 //===--------------------------------------------------------------------===//
89 
90 /// The OpAsmOpInterface, see OpAsmInterface.td for more details.
91 #include "mlir/IR/OpAsmInterface.cpp.inc"
92 
93 //===----------------------------------------------------------------------===//
94 // OpPrintingFlags
95 //===----------------------------------------------------------------------===//
96 
97 namespace {
98 /// This struct contains command line options that can be used to initialize
99 /// various bits of the AsmPrinter. This uses a struct wrapper to avoid the need
100 /// for global command line options.
101 struct AsmPrinterOptions {
102   llvm::cl::opt<int64_t> printElementsAttrWithHexIfLarger{
103       "mlir-print-elementsattrs-with-hex-if-larger",
104       llvm::cl::desc(
105           "Print DenseElementsAttrs with a hex string that have "
106           "more elements than the given upper limit (use -1 to disable)")};
107 
108   llvm::cl::opt<unsigned> elideElementsAttrIfLarger{
109       "mlir-elide-elementsattrs-if-larger",
110       llvm::cl::desc("Elide ElementsAttrs with \"...\" that have "
111                      "more elements than the given upper limit")};
112 
113   llvm::cl::opt<bool> printDebugInfoOpt{
114       "mlir-print-debuginfo", llvm::cl::init(false),
115       llvm::cl::desc("Print debug info in MLIR output")};
116 
117   llvm::cl::opt<bool> printPrettyDebugInfoOpt{
118       "mlir-pretty-debuginfo", llvm::cl::init(false),
119       llvm::cl::desc("Print pretty debug info in MLIR output")};
120 
121   // Use the generic op output form in the operation printer even if the custom
122   // form is defined.
123   llvm::cl::opt<bool> printGenericOpFormOpt{
124       "mlir-print-op-generic", llvm::cl::init(false),
125       llvm::cl::desc("Print the generic op form"), llvm::cl::Hidden};
126 
127   llvm::cl::opt<bool> printLocalScopeOpt{
128       "mlir-print-local-scope", llvm::cl::init(false),
129       llvm::cl::desc("Print assuming in local scope by default"),
130       llvm::cl::Hidden};
131 };
132 } // end anonymous namespace
133 
134 static llvm::ManagedStatic<AsmPrinterOptions> clOptions;
135 
136 /// Register a set of useful command-line options that can be used to configure
137 /// various flags within the AsmPrinter.
138 void mlir::registerAsmPrinterCLOptions() {
139   // Make sure that the options struct has been initialized.
140   *clOptions;
141 }
142 
143 /// Initialize the printing flags with default supplied by the cl::opts above.
144 OpPrintingFlags::OpPrintingFlags()
145     : printDebugInfoFlag(false), printDebugInfoPrettyFormFlag(false),
146       printGenericOpFormFlag(false), printLocalScope(false) {
147   // Initialize based upon command line options, if they are available.
148   if (!clOptions.isConstructed())
149     return;
150   if (clOptions->elideElementsAttrIfLarger.getNumOccurrences())
151     elementsAttrElementLimit = clOptions->elideElementsAttrIfLarger;
152   printDebugInfoFlag = clOptions->printDebugInfoOpt;
153   printDebugInfoPrettyFormFlag = clOptions->printPrettyDebugInfoOpt;
154   printGenericOpFormFlag = clOptions->printGenericOpFormOpt;
155   printLocalScope = clOptions->printLocalScopeOpt;
156 }
157 
158 /// Enable the elision of large elements attributes, by printing a '...'
159 /// instead of the element data, when the number of elements is greater than
160 /// `largeElementLimit`. Note: The IR generated with this option is not
161 /// parsable.
162 OpPrintingFlags &
163 OpPrintingFlags::elideLargeElementsAttrs(int64_t largeElementLimit) {
164   elementsAttrElementLimit = largeElementLimit;
165   return *this;
166 }
167 
168 /// Enable printing of debug information. If 'prettyForm' is set to true,
169 /// debug information is printed in a more readable 'pretty' form.
170 OpPrintingFlags &OpPrintingFlags::enableDebugInfo(bool prettyForm) {
171   printDebugInfoFlag = true;
172   printDebugInfoPrettyFormFlag = prettyForm;
173   return *this;
174 }
175 
176 /// Always print operations in the generic form.
177 OpPrintingFlags &OpPrintingFlags::printGenericOpForm() {
178   printGenericOpFormFlag = true;
179   return *this;
180 }
181 
182 /// Use local scope when printing the operation. This allows for using the
183 /// printer in a more localized and thread-safe setting, but may not necessarily
184 /// be identical of what the IR will look like when dumping the full module.
185 OpPrintingFlags &OpPrintingFlags::useLocalScope() {
186   printLocalScope = true;
187   return *this;
188 }
189 
190 /// Return if the given ElementsAttr should be elided.
191 bool OpPrintingFlags::shouldElideElementsAttr(ElementsAttr attr) const {
192   return elementsAttrElementLimit.hasValue() &&
193          *elementsAttrElementLimit < int64_t(attr.getNumElements()) &&
194          !attr.isa<SplatElementsAttr>();
195 }
196 
197 /// Return the size limit for printing large ElementsAttr.
198 Optional<int64_t> OpPrintingFlags::getLargeElementsAttrLimit() const {
199   return elementsAttrElementLimit;
200 }
201 
202 /// Return if debug information should be printed.
203 bool OpPrintingFlags::shouldPrintDebugInfo() const {
204   return printDebugInfoFlag;
205 }
206 
207 /// Return if debug information should be printed in the pretty form.
208 bool OpPrintingFlags::shouldPrintDebugInfoPrettyForm() const {
209   return printDebugInfoPrettyFormFlag;
210 }
211 
212 /// Return if operations should be printed in the generic form.
213 bool OpPrintingFlags::shouldPrintGenericOpForm() const {
214   return printGenericOpFormFlag;
215 }
216 
217 /// Return if the printer should use local scope when dumping the IR.
218 bool OpPrintingFlags::shouldUseLocalScope() const { return printLocalScope; }
219 
220 /// Returns true if an ElementsAttr with the given number of elements should be
221 /// printed with hex.
222 static bool shouldPrintElementsAttrWithHex(int64_t numElements) {
223   // Check to see if a command line option was provided for the limit.
224   if (clOptions.isConstructed()) {
225     if (clOptions->printElementsAttrWithHexIfLarger.getNumOccurrences()) {
226       // -1 is used to disable hex printing.
227       if (clOptions->printElementsAttrWithHexIfLarger == -1)
228         return false;
229       return numElements > clOptions->printElementsAttrWithHexIfLarger;
230     }
231   }
232 
233   // Otherwise, default to printing with hex if the number of elements is >100.
234   return numElements > 100;
235 }
236 
237 //===----------------------------------------------------------------------===//
238 // NewLineCounter
239 //===----------------------------------------------------------------------===//
240 
241 namespace {
242 /// This class is a simple formatter that emits a new line when inputted into a
243 /// stream, that enables counting the number of newlines emitted. This class
244 /// should be used whenever emitting newlines in the printer.
245 struct NewLineCounter {
246   unsigned curLine = 1;
247 };
248 } // end anonymous namespace
249 
250 static raw_ostream &operator<<(raw_ostream &os, NewLineCounter &newLine) {
251   ++newLine.curLine;
252   return os << '\n';
253 }
254 
255 //===----------------------------------------------------------------------===//
256 // AliasInitializer
257 //===----------------------------------------------------------------------===//
258 
259 namespace {
260 /// This class represents a specific instance of a symbol Alias.
261 class SymbolAlias {
262 public:
263   SymbolAlias(StringRef name, bool isDeferrable)
264       : name(name), suffixIndex(0), hasSuffixIndex(false),
265         isDeferrable(isDeferrable) {}
266   SymbolAlias(StringRef name, uint32_t suffixIndex, bool isDeferrable)
267       : name(name), suffixIndex(suffixIndex), hasSuffixIndex(true),
268         isDeferrable(isDeferrable) {}
269 
270   /// Print this alias to the given stream.
271   void print(raw_ostream &os) const {
272     os << name;
273     if (hasSuffixIndex)
274       os << suffixIndex;
275   }
276 
277   /// Returns true if this alias supports deferred resolution when parsing.
278   bool canBeDeferred() const { return isDeferrable; }
279 
280 private:
281   /// The main name of the alias.
282   StringRef name;
283   /// The optional suffix index of the alias, if multiple aliases had the same
284   /// name.
285   uint32_t suffixIndex : 30;
286   /// A flag indicating whether this alias has a suffix or not.
287   bool hasSuffixIndex : 1;
288   /// A flag indicating whether this alias may be deferred or not.
289   bool isDeferrable : 1;
290 };
291 
292 /// This class represents a utility that initializes the set of attribute and
293 /// type aliases, without the need to store the extra information within the
294 /// main AliasState class or pass it around via function arguments.
295 class AliasInitializer {
296 public:
297   AliasInitializer(
298       DialectInterfaceCollection<OpAsmDialectInterface> &interfaces,
299       llvm::BumpPtrAllocator &aliasAllocator)
300       : interfaces(interfaces), aliasAllocator(aliasAllocator),
301         aliasOS(aliasBuffer) {}
302 
303   void initialize(Operation *op, const OpPrintingFlags &printerFlags,
304                   llvm::MapVector<Attribute, SymbolAlias> &attrToAlias,
305                   llvm::MapVector<Type, SymbolAlias> &typeToAlias);
306 
307   /// Visit the given attribute to see if it has an alias. `canBeDeferred` is
308   /// set to true if the originator of this attribute can resolve the alias
309   /// after parsing has completed (e.g. in the case of operation locations).
310   void visit(Attribute attr, bool canBeDeferred = false);
311 
312   /// Visit the given type to see if it has an alias.
313   void visit(Type type);
314 
315 private:
316   /// Try to generate an alias for the provided symbol. If an alias is
317   /// generated, the provided alias mapping and reverse mapping are updated.
318   /// Returns success if an alias was generated, failure otherwise.
319   template <typename T>
320   LogicalResult
321   generateAlias(T symbol,
322                 llvm::MapVector<StringRef, std::vector<T>> &aliasToSymbol);
323 
324   /// The set of asm interfaces within the context.
325   DialectInterfaceCollection<OpAsmDialectInterface> &interfaces;
326 
327   /// Mapping between an alias and the set of symbols mapped to it.
328   llvm::MapVector<StringRef, std::vector<Attribute>> aliasToAttr;
329   llvm::MapVector<StringRef, std::vector<Type>> aliasToType;
330 
331   /// An allocator used for alias names.
332   llvm::BumpPtrAllocator &aliasAllocator;
333 
334   /// The set of visited attributes.
335   DenseSet<Attribute> visitedAttributes;
336 
337   /// The set of attributes that have aliases *and* can be deferred.
338   DenseSet<Attribute> deferrableAttributes;
339 
340   /// The set of visited types.
341   DenseSet<Type> visitedTypes;
342 
343   /// Storage and stream used when generating an alias.
344   SmallString<32> aliasBuffer;
345   llvm::raw_svector_ostream aliasOS;
346 };
347 
348 /// This class implements a dummy OpAsmPrinter that doesn't print any output,
349 /// and merely collects the attributes and types that *would* be printed in a
350 /// normal print invocation so that we can generate proper aliases. This allows
351 /// for us to generate aliases only for the attributes and types that would be
352 /// in the output, and trims down unnecessary output.
353 class DummyAliasOperationPrinter : private OpAsmPrinter {
354 public:
355   explicit DummyAliasOperationPrinter(const OpPrintingFlags &flags,
356                                       AliasInitializer &initializer)
357       : printerFlags(flags), initializer(initializer) {}
358 
359   /// Print the given operation.
360   void print(Operation *op) {
361     // Visit the operation location.
362     if (printerFlags.shouldPrintDebugInfo())
363       initializer.visit(op->getLoc(), /*canBeDeferred=*/true);
364 
365     // If requested, always print the generic form.
366     if (!printerFlags.shouldPrintGenericOpForm()) {
367       // Check to see if this is a known operation.  If so, use the registered
368       // custom printer hook.
369       if (auto *opInfo = op->getAbstractOperation()) {
370         opInfo->printAssembly(op, *this);
371         return;
372       }
373     }
374 
375     // Otherwise print with the generic assembly form.
376     printGenericOp(op);
377   }
378 
379 private:
380   /// Print the given operation in the generic form.
381   void printGenericOp(Operation *op) override {
382     // Consider nested operations for aliases.
383     if (op->getNumRegions() != 0) {
384       for (Region &region : op->getRegions())
385         printRegion(region, /*printEntryBlockArgs=*/true,
386                     /*printBlockTerminators=*/true);
387     }
388 
389     // Visit all the types used in the operation.
390     for (Type type : op->getOperandTypes())
391       printType(type);
392     for (Type type : op->getResultTypes())
393       printType(type);
394 
395     // Consider the attributes of the operation for aliases.
396     for (const NamedAttribute &attr : op->getAttrs())
397       printAttribute(attr.second);
398   }
399 
400   /// Print the given block. If 'printBlockArgs' is false, the arguments of the
401   /// block are not printed. If 'printBlockTerminator' is false, the terminator
402   /// operation of the block is not printed.
403   void print(Block *block, bool printBlockArgs = true,
404              bool printBlockTerminator = true) {
405     // Consider the types of the block arguments for aliases if 'printBlockArgs'
406     // is set to true.
407     if (printBlockArgs) {
408       for (Type type : block->getArgumentTypes())
409         printType(type);
410     }
411 
412     // Consider the operations within this block, ignoring the terminator if
413     // requested.
414     auto range = llvm::make_range(
415         block->begin(), std::prev(block->end(), printBlockTerminator ? 0 : 1));
416     for (Operation &op : range)
417       print(&op);
418   }
419 
420   /// Print the given region.
421   void printRegion(Region &region, bool printEntryBlockArgs,
422                    bool printBlockTerminators,
423                    bool printEmptyBlock = false) override {
424     if (region.empty())
425       return;
426 
427     auto *entryBlock = &region.front();
428     print(entryBlock, printEntryBlockArgs, printBlockTerminators);
429     for (Block &b : llvm::drop_begin(region, 1))
430       print(&b);
431   }
432 
433   /// Consider the given type to be printed for an alias.
434   void printType(Type type) override { initializer.visit(type); }
435 
436   /// Consider the given attribute to be printed for an alias.
437   void printAttribute(Attribute attr) override { initializer.visit(attr); }
438   void printAttributeWithoutType(Attribute attr) override {
439     printAttribute(attr);
440   }
441 
442   /// Print the given set of attributes with names not included within
443   /// 'elidedAttrs'.
444   void printOptionalAttrDict(ArrayRef<NamedAttribute> attrs,
445                              ArrayRef<StringRef> elidedAttrs = {}) override {
446     if (attrs.empty())
447       return;
448     if (elidedAttrs.empty()) {
449       for (const NamedAttribute &attr : attrs)
450         printAttribute(attr.second);
451       return;
452     }
453     llvm::SmallDenseSet<StringRef> elidedAttrsSet(elidedAttrs.begin(),
454                                                   elidedAttrs.end());
455     for (const NamedAttribute &attr : attrs)
456       if (!elidedAttrsSet.contains(attr.first.strref()))
457         printAttribute(attr.second);
458   }
459   void printOptionalAttrDictWithKeyword(
460       ArrayRef<NamedAttribute> attrs,
461       ArrayRef<StringRef> elidedAttrs = {}) override {
462     printOptionalAttrDict(attrs, elidedAttrs);
463   }
464 
465   /// Return a null stream as the output stream, this will ignore any data fed
466   /// to it.
467   raw_ostream &getStream() const override { return os; }
468 
469   /// The following are hooks of `OpAsmPrinter` that are not necessary for
470   /// determining potential aliases.
471   void printAffineMapOfSSAIds(AffineMapAttr, ValueRange) override {}
472   void printNewline() override {}
473   void printOperand(Value) override {}
474   void printOperand(Value, raw_ostream &os) override {
475     // Users expect the output string to have at least the prefixed % to signal
476     // a value name. To maintain this invariant, emit a name even if it is
477     // guaranteed to go unused.
478     os << "%";
479   }
480   void printSymbolName(StringRef) override {}
481   void printSuccessor(Block *) override {}
482   void printSuccessorAndUseList(Block *, ValueRange) override {}
483   void shadowRegionArgs(Region &, ValueRange) override {}
484 
485   /// The printer flags to use when determining potential aliases.
486   const OpPrintingFlags &printerFlags;
487 
488   /// The initializer to use when identifying aliases.
489   AliasInitializer &initializer;
490 
491   /// A dummy output stream.
492   mutable llvm::raw_null_ostream os;
493 };
494 } // end anonymous namespace
495 
496 /// Sanitize the given name such that it can be used as a valid identifier. If
497 /// the string needs to be modified in any way, the provided buffer is used to
498 /// store the new copy,
499 static StringRef sanitizeIdentifier(StringRef name, SmallString<16> &buffer,
500                                     StringRef allowedPunctChars = "$._-",
501                                     bool allowTrailingDigit = true) {
502   assert(!name.empty() && "Shouldn't have an empty name here");
503 
504   auto copyNameToBuffer = [&] {
505     for (char ch : name) {
506       if (llvm::isAlnum(ch) || allowedPunctChars.contains(ch))
507         buffer.push_back(ch);
508       else if (ch == ' ')
509         buffer.push_back('_');
510       else
511         buffer.append(llvm::utohexstr((unsigned char)ch));
512     }
513   };
514 
515   // Check to see if this name is valid. If it starts with a digit, then it
516   // could conflict with the autogenerated numeric ID's, so add an underscore
517   // prefix to avoid problems.
518   if (isdigit(name[0])) {
519     buffer.push_back('_');
520     copyNameToBuffer();
521     return buffer;
522   }
523 
524   // If the name ends with a trailing digit, add a '_' to avoid potential
525   // conflicts with autogenerated ID's.
526   if (!allowTrailingDigit && isdigit(name.back())) {
527     copyNameToBuffer();
528     buffer.push_back('_');
529     return buffer;
530   }
531 
532   // Check to see that the name consists of only valid identifier characters.
533   for (char ch : name) {
534     if (!llvm::isAlnum(ch) && !allowedPunctChars.contains(ch)) {
535       copyNameToBuffer();
536       return buffer;
537     }
538   }
539 
540   // If there are no invalid characters, return the original name.
541   return name;
542 }
543 
544 /// Given a collection of aliases and symbols, initialize a mapping from a
545 /// symbol to a given alias.
546 template <typename T>
547 static void
548 initializeAliases(llvm::MapVector<StringRef, std::vector<T>> &aliasToSymbol,
549                   llvm::MapVector<T, SymbolAlias> &symbolToAlias,
550                   DenseSet<T> *deferrableAliases = nullptr) {
551   std::vector<std::pair<StringRef, std::vector<T>>> aliases =
552       aliasToSymbol.takeVector();
553   llvm::array_pod_sort(aliases.begin(), aliases.end(),
554                        [](const auto *lhs, const auto *rhs) {
555                          return lhs->first.compare(rhs->first);
556                        });
557 
558   for (auto &it : aliases) {
559     // If there is only one instance for this alias, use the name directly.
560     if (it.second.size() == 1) {
561       T symbol = it.second.front();
562       bool isDeferrable = deferrableAliases && deferrableAliases->count(symbol);
563       symbolToAlias.insert({symbol, SymbolAlias(it.first, isDeferrable)});
564       continue;
565     }
566     // Otherwise, add the index to the name.
567     for (int i = 0, e = it.second.size(); i < e; ++i) {
568       T symbol = it.second[i];
569       bool isDeferrable = deferrableAliases && deferrableAliases->count(symbol);
570       symbolToAlias.insert({symbol, SymbolAlias(it.first, i, isDeferrable)});
571     }
572   }
573 }
574 
575 void AliasInitializer::initialize(
576     Operation *op, const OpPrintingFlags &printerFlags,
577     llvm::MapVector<Attribute, SymbolAlias> &attrToAlias,
578     llvm::MapVector<Type, SymbolAlias> &typeToAlias) {
579   // Use a dummy printer when walking the IR so that we can collect the
580   // attributes/types that will actually be used during printing when
581   // considering aliases.
582   DummyAliasOperationPrinter aliasPrinter(printerFlags, *this);
583   aliasPrinter.print(op);
584 
585   // Initialize the aliases sorted by name.
586   initializeAliases(aliasToAttr, attrToAlias, &deferrableAttributes);
587   initializeAliases(aliasToType, typeToAlias);
588 }
589 
590 void AliasInitializer::visit(Attribute attr, bool canBeDeferred) {
591   if (!visitedAttributes.insert(attr).second) {
592     // If this attribute already has an alias and this instance can't be
593     // deferred, make sure that the alias isn't deferred.
594     if (!canBeDeferred)
595       deferrableAttributes.erase(attr);
596     return;
597   }
598 
599   // Try to generate an alias for this attribute.
600   if (succeeded(generateAlias(attr, aliasToAttr))) {
601     if (canBeDeferred)
602       deferrableAttributes.insert(attr);
603     return;
604   }
605 
606   if (auto arrayAttr = attr.dyn_cast<ArrayAttr>()) {
607     for (Attribute element : arrayAttr.getValue())
608       visit(element);
609   } else if (auto dictAttr = attr.dyn_cast<DictionaryAttr>()) {
610     for (const NamedAttribute &attr : dictAttr)
611       visit(attr.second);
612   } else if (auto typeAttr = attr.dyn_cast<TypeAttr>()) {
613     visit(typeAttr.getValue());
614   }
615 }
616 
617 void AliasInitializer::visit(Type type) {
618   if (!visitedTypes.insert(type).second)
619     return;
620 
621   // Try to generate an alias for this type.
622   if (succeeded(generateAlias(type, aliasToType)))
623     return;
624 
625   // Visit several subtypes that contain types or attributes.
626   if (auto funcType = type.dyn_cast<FunctionType>()) {
627     // Visit input and result types for functions.
628     for (auto input : funcType.getInputs())
629       visit(input);
630     for (auto result : funcType.getResults())
631       visit(result);
632   } else if (auto shapedType = type.dyn_cast<ShapedType>()) {
633     visit(shapedType.getElementType());
634 
635     // Visit affine maps in memref type.
636     if (auto memref = type.dyn_cast<MemRefType>())
637       for (auto map : memref.getAffineMaps())
638         visit(AffineMapAttr::get(map));
639   }
640 }
641 
642 template <typename T>
643 LogicalResult AliasInitializer::generateAlias(
644     T symbol, llvm::MapVector<StringRef, std::vector<T>> &aliasToSymbol) {
645   SmallString<16> tempBuffer;
646   for (const auto &interface : interfaces) {
647     if (failed(interface.getAlias(symbol, aliasOS)))
648       continue;
649     StringRef name = aliasOS.str();
650     assert(!name.empty() && "expected valid alias name");
651     name = sanitizeIdentifier(name, tempBuffer, /*allowedPunctChars=*/"$_-",
652                               /*allowTrailingDigit=*/false);
653     name = name.copy(aliasAllocator);
654 
655     aliasToSymbol[name].push_back(symbol);
656     aliasBuffer.clear();
657     return success();
658   }
659   return failure();
660 }
661 
662 //===----------------------------------------------------------------------===//
663 // AliasState
664 //===----------------------------------------------------------------------===//
665 
666 namespace {
667 /// This class manages the state for type and attribute aliases.
668 class AliasState {
669 public:
670   // Initialize the internal aliases.
671   void
672   initialize(Operation *op, const OpPrintingFlags &printerFlags,
673              DialectInterfaceCollection<OpAsmDialectInterface> &interfaces);
674 
675   /// Get an alias for the given attribute if it has one and print it in `os`.
676   /// Returns success if an alias was printed, failure otherwise.
677   LogicalResult getAlias(Attribute attr, raw_ostream &os) const;
678 
679   /// Get an alias for the given type if it has one and print it in `os`.
680   /// Returns success if an alias was printed, failure otherwise.
681   LogicalResult getAlias(Type ty, raw_ostream &os) const;
682 
683   /// Print all of the referenced aliases that can not be resolved in a deferred
684   /// manner.
685   void printNonDeferredAliases(raw_ostream &os, NewLineCounter &newLine) const {
686     printAliases(os, newLine, /*isDeferred=*/false);
687   }
688 
689   /// Print all of the referenced aliases that support deferred resolution.
690   void printDeferredAliases(raw_ostream &os, NewLineCounter &newLine) const {
691     printAliases(os, newLine, /*isDeferred=*/true);
692   }
693 
694 private:
695   /// Print all of the referenced aliases that support the provided resolution
696   /// behavior.
697   void printAliases(raw_ostream &os, NewLineCounter &newLine,
698                     bool isDeferred) const;
699 
700   /// Mapping between attribute and alias.
701   llvm::MapVector<Attribute, SymbolAlias> attrToAlias;
702   /// Mapping between type and alias.
703   llvm::MapVector<Type, SymbolAlias> typeToAlias;
704 
705   /// An allocator used for alias names.
706   llvm::BumpPtrAllocator aliasAllocator;
707 };
708 } // end anonymous namespace
709 
710 void AliasState::initialize(
711     Operation *op, const OpPrintingFlags &printerFlags,
712     DialectInterfaceCollection<OpAsmDialectInterface> &interfaces) {
713   AliasInitializer initializer(interfaces, aliasAllocator);
714   initializer.initialize(op, printerFlags, attrToAlias, typeToAlias);
715 }
716 
717 LogicalResult AliasState::getAlias(Attribute attr, raw_ostream &os) const {
718   auto it = attrToAlias.find(attr);
719   if (it == attrToAlias.end())
720     return failure();
721   it->second.print(os << '#');
722   return success();
723 }
724 
725 LogicalResult AliasState::getAlias(Type ty, raw_ostream &os) const {
726   auto it = typeToAlias.find(ty);
727   if (it == typeToAlias.end())
728     return failure();
729 
730   it->second.print(os << '!');
731   return success();
732 }
733 
734 void AliasState::printAliases(raw_ostream &os, NewLineCounter &newLine,
735                               bool isDeferred) const {
736   auto filterFn = [=](const auto &aliasIt) {
737     return aliasIt.second.canBeDeferred() == isDeferred;
738   };
739   for (const auto &it : llvm::make_filter_range(attrToAlias, filterFn)) {
740     it.second.print(os << '#');
741     os << " = " << it.first << newLine;
742   }
743   for (const auto &it : llvm::make_filter_range(typeToAlias, filterFn)) {
744     it.second.print(os << '!');
745     os << " = type " << it.first << newLine;
746   }
747 }
748 
749 //===----------------------------------------------------------------------===//
750 // SSANameState
751 //===----------------------------------------------------------------------===//
752 
753 namespace {
754 /// This class manages the state of SSA value names.
755 class SSANameState {
756 public:
757   /// A sentinel value used for values with names set.
758   enum : unsigned { NameSentinel = ~0U };
759 
760   SSANameState(Operation *op,
761                DialectInterfaceCollection<OpAsmDialectInterface> &interfaces);
762 
763   /// Print the SSA identifier for the given value to 'stream'. If
764   /// 'printResultNo' is true, it also presents the result number ('#' number)
765   /// of this value.
766   void printValueID(Value value, bool printResultNo, raw_ostream &stream) const;
767 
768   /// Return the result indices for each of the result groups registered by this
769   /// operation, or empty if none exist.
770   ArrayRef<int> getOpResultGroups(Operation *op);
771 
772   /// Get the ID for the given block.
773   unsigned getBlockID(Block *block);
774 
775   /// Renumber the arguments for the specified region to the same names as the
776   /// SSA values in namesToUse. See OperationPrinter::shadowRegionArgs for
777   /// details.
778   void shadowRegionArgs(Region &region, ValueRange namesToUse);
779 
780 private:
781   /// Number the SSA values within the given IR unit.
782   void numberValuesInRegion(
783       Region &region,
784       DialectInterfaceCollection<OpAsmDialectInterface> &interfaces);
785   void numberValuesInBlock(
786       Block &block,
787       DialectInterfaceCollection<OpAsmDialectInterface> &interfaces);
788   void numberValuesInOp(
789       Operation &op,
790       DialectInterfaceCollection<OpAsmDialectInterface> &interfaces);
791 
792   /// Given a result of an operation 'result', find the result group head
793   /// 'lookupValue' and the result of 'result' within that group in
794   /// 'lookupResultNo'. 'lookupResultNo' is only filled in if the result group
795   /// has more than 1 result.
796   void getResultIDAndNumber(OpResult result, Value &lookupValue,
797                             Optional<int> &lookupResultNo) const;
798 
799   /// Set a special value name for the given value.
800   void setValueName(Value value, StringRef name);
801 
802   /// Uniques the given value name within the printer. If the given name
803   /// conflicts, it is automatically renamed.
804   StringRef uniqueValueName(StringRef name);
805 
806   /// This is the value ID for each SSA value. If this returns NameSentinel,
807   /// then the valueID has an entry in valueNames.
808   DenseMap<Value, unsigned> valueIDs;
809   DenseMap<Value, StringRef> valueNames;
810 
811   /// This is a map of operations that contain multiple named result groups,
812   /// i.e. there may be multiple names for the results of the operation. The
813   /// value of this map are the result numbers that start a result group.
814   DenseMap<Operation *, SmallVector<int, 1>> opResultGroups;
815 
816   /// This is the block ID for each block in the current.
817   DenseMap<Block *, unsigned> blockIDs;
818 
819   /// This keeps track of all of the non-numeric names that are in flight,
820   /// allowing us to check for duplicates.
821   /// Note: the value of the map is unused.
822   llvm::ScopedHashTable<StringRef, char> usedNames;
823   llvm::BumpPtrAllocator usedNameAllocator;
824 
825   /// This is the next value ID to assign in numbering.
826   unsigned nextValueID = 0;
827   /// This is the next ID to assign to a region entry block argument.
828   unsigned nextArgumentID = 0;
829   /// This is the next ID to assign when a name conflict is detected.
830   unsigned nextConflictID = 0;
831 };
832 } // end anonymous namespace
833 
834 SSANameState::SSANameState(
835     Operation *op,
836     DialectInterfaceCollection<OpAsmDialectInterface> &interfaces) {
837   llvm::ScopedHashTable<StringRef, char>::ScopeTy usedNamesScope(usedNames);
838   numberValuesInOp(*op, interfaces);
839 
840   for (auto &region : op->getRegions())
841     numberValuesInRegion(region, interfaces);
842 }
843 
844 void SSANameState::printValueID(Value value, bool printResultNo,
845                                 raw_ostream &stream) const {
846   if (!value) {
847     stream << "<<NULL>>";
848     return;
849   }
850 
851   Optional<int> resultNo;
852   auto lookupValue = value;
853 
854   // If this is an operation result, collect the head lookup value of the result
855   // group and the result number of 'result' within that group.
856   if (OpResult result = value.dyn_cast<OpResult>())
857     getResultIDAndNumber(result, lookupValue, resultNo);
858 
859   auto it = valueIDs.find(lookupValue);
860   if (it == valueIDs.end()) {
861     stream << "<<UNKNOWN SSA VALUE>>";
862     return;
863   }
864 
865   stream << '%';
866   if (it->second != NameSentinel) {
867     stream << it->second;
868   } else {
869     auto nameIt = valueNames.find(lookupValue);
870     assert(nameIt != valueNames.end() && "Didn't have a name entry?");
871     stream << nameIt->second;
872   }
873 
874   if (resultNo.hasValue() && printResultNo)
875     stream << '#' << resultNo;
876 }
877 
878 ArrayRef<int> SSANameState::getOpResultGroups(Operation *op) {
879   auto it = opResultGroups.find(op);
880   return it == opResultGroups.end() ? ArrayRef<int>() : it->second;
881 }
882 
883 unsigned SSANameState::getBlockID(Block *block) {
884   auto it = blockIDs.find(block);
885   return it != blockIDs.end() ? it->second : NameSentinel;
886 }
887 
888 void SSANameState::shadowRegionArgs(Region &region, ValueRange namesToUse) {
889   assert(!region.empty() && "cannot shadow arguments of an empty region");
890   assert(region.getNumArguments() == namesToUse.size() &&
891          "incorrect number of names passed in");
892   assert(region.getParentOp()->hasTrait<OpTrait::IsIsolatedFromAbove>() &&
893          "only KnownIsolatedFromAbove ops can shadow names");
894 
895   SmallVector<char, 16> nameStr;
896   for (unsigned i = 0, e = namesToUse.size(); i != e; ++i) {
897     auto nameToUse = namesToUse[i];
898     if (nameToUse == nullptr)
899       continue;
900     auto nameToReplace = region.getArgument(i);
901 
902     nameStr.clear();
903     llvm::raw_svector_ostream nameStream(nameStr);
904     printValueID(nameToUse, /*printResultNo=*/true, nameStream);
905 
906     // Entry block arguments should already have a pretty "arg" name.
907     assert(valueIDs[nameToReplace] == NameSentinel);
908 
909     // Use the name without the leading %.
910     auto name = StringRef(nameStream.str()).drop_front();
911 
912     // Overwrite the name.
913     valueNames[nameToReplace] = name.copy(usedNameAllocator);
914   }
915 }
916 
917 void SSANameState::numberValuesInRegion(
918     Region &region,
919     DialectInterfaceCollection<OpAsmDialectInterface> &interfaces) {
920   // Save the current value ids to allow for numbering values in sibling regions
921   // the same.
922   llvm::SaveAndRestore<unsigned> valueIDSaver(nextValueID);
923   llvm::SaveAndRestore<unsigned> argumentIDSaver(nextArgumentID);
924   llvm::SaveAndRestore<unsigned> conflictIDSaver(nextConflictID);
925 
926   // Push a new used names scope.
927   llvm::ScopedHashTable<StringRef, char>::ScopeTy usedNamesScope(usedNames);
928 
929   // Number the values within this region in a breadth-first order.
930   unsigned nextBlockID = 0;
931   for (auto &block : region) {
932     // Each block gets a unique ID, and all of the operations within it get
933     // numbered as well.
934     blockIDs[&block] = nextBlockID++;
935     numberValuesInBlock(block, interfaces);
936   }
937 
938   // After that we traverse the nested regions.
939   // TODO: Rework this loop to not use recursion.
940   for (auto &block : region) {
941     for (auto &op : block)
942       for (auto &nestedRegion : op.getRegions())
943         numberValuesInRegion(nestedRegion, interfaces);
944   }
945 }
946 
947 void SSANameState::numberValuesInBlock(
948     Block &block,
949     DialectInterfaceCollection<OpAsmDialectInterface> &interfaces) {
950   auto setArgNameFn = [&](Value arg, StringRef name) {
951     assert(!valueIDs.count(arg) && "arg numbered multiple times");
952     assert(arg.cast<BlockArgument>().getOwner() == &block &&
953            "arg not defined in 'block'");
954     setValueName(arg, name);
955   };
956 
957   bool isEntryBlock = block.isEntryBlock();
958   if (isEntryBlock) {
959     if (auto *op = block.getParentOp()) {
960       if (auto asmInterface = interfaces.getInterfaceFor(op->getDialect()))
961         asmInterface->getAsmBlockArgumentNames(&block, setArgNameFn);
962     }
963   }
964 
965   // Number the block arguments. We give entry block arguments a special name
966   // 'arg'.
967   SmallString<32> specialNameBuffer(isEntryBlock ? "arg" : "");
968   llvm::raw_svector_ostream specialName(specialNameBuffer);
969   for (auto arg : block.getArguments()) {
970     if (valueIDs.count(arg))
971       continue;
972     if (isEntryBlock) {
973       specialNameBuffer.resize(strlen("arg"));
974       specialName << nextArgumentID++;
975     }
976     setValueName(arg, specialName.str());
977   }
978 
979   // Number the operations in this block.
980   for (auto &op : block)
981     numberValuesInOp(op, interfaces);
982 }
983 
984 void SSANameState::numberValuesInOp(
985     Operation &op,
986     DialectInterfaceCollection<OpAsmDialectInterface> &interfaces) {
987   unsigned numResults = op.getNumResults();
988   if (numResults == 0)
989     return;
990   Value resultBegin = op.getResult(0);
991 
992   // Function used to set the special result names for the operation.
993   SmallVector<int, 2> resultGroups(/*Size=*/1, /*Value=*/0);
994   auto setResultNameFn = [&](Value result, StringRef name) {
995     assert(!valueIDs.count(result) && "result numbered multiple times");
996     assert(result.getDefiningOp() == &op && "result not defined by 'op'");
997     setValueName(result, name);
998 
999     // Record the result number for groups not anchored at 0.
1000     if (int resultNo = result.cast<OpResult>().getResultNumber())
1001       resultGroups.push_back(resultNo);
1002   };
1003   if (OpAsmOpInterface asmInterface = dyn_cast<OpAsmOpInterface>(&op))
1004     asmInterface.getAsmResultNames(setResultNameFn);
1005   else if (auto *asmInterface = interfaces.getInterfaceFor(op.getDialect()))
1006     asmInterface->getAsmResultNames(&op, setResultNameFn);
1007 
1008   // If the first result wasn't numbered, give it a default number.
1009   if (valueIDs.try_emplace(resultBegin, nextValueID).second)
1010     ++nextValueID;
1011 
1012   // If this operation has multiple result groups, mark it.
1013   if (resultGroups.size() != 1) {
1014     llvm::array_pod_sort(resultGroups.begin(), resultGroups.end());
1015     opResultGroups.try_emplace(&op, std::move(resultGroups));
1016   }
1017 }
1018 
1019 void SSANameState::getResultIDAndNumber(OpResult result, Value &lookupValue,
1020                                         Optional<int> &lookupResultNo) const {
1021   Operation *owner = result.getOwner();
1022   if (owner->getNumResults() == 1)
1023     return;
1024   int resultNo = result.getResultNumber();
1025 
1026   // If this operation has multiple result groups, we will need to find the
1027   // one corresponding to this result.
1028   auto resultGroupIt = opResultGroups.find(owner);
1029   if (resultGroupIt == opResultGroups.end()) {
1030     // If not, just use the first result.
1031     lookupResultNo = resultNo;
1032     lookupValue = owner->getResult(0);
1033     return;
1034   }
1035 
1036   // Find the correct index using a binary search, as the groups are ordered.
1037   ArrayRef<int> resultGroups = resultGroupIt->second;
1038   auto it = llvm::upper_bound(resultGroups, resultNo);
1039   int groupResultNo = 0, groupSize = 0;
1040 
1041   // If there are no smaller elements, the last result group is the lookup.
1042   if (it == resultGroups.end()) {
1043     groupResultNo = resultGroups.back();
1044     groupSize = static_cast<int>(owner->getNumResults()) - resultGroups.back();
1045   } else {
1046     // Otherwise, the previous element is the lookup.
1047     groupResultNo = *std::prev(it);
1048     groupSize = *it - groupResultNo;
1049   }
1050 
1051   // We only record the result number for a group of size greater than 1.
1052   if (groupSize != 1)
1053     lookupResultNo = resultNo - groupResultNo;
1054   lookupValue = owner->getResult(groupResultNo);
1055 }
1056 
1057 void SSANameState::setValueName(Value value, StringRef name) {
1058   // If the name is empty, the value uses the default numbering.
1059   if (name.empty()) {
1060     valueIDs[value] = nextValueID++;
1061     return;
1062   }
1063 
1064   valueIDs[value] = NameSentinel;
1065   valueNames[value] = uniqueValueName(name);
1066 }
1067 
1068 StringRef SSANameState::uniqueValueName(StringRef name) {
1069   SmallString<16> tmpBuffer;
1070   name = sanitizeIdentifier(name, tmpBuffer);
1071 
1072   // Check to see if this name is already unique.
1073   if (!usedNames.count(name)) {
1074     name = name.copy(usedNameAllocator);
1075   } else {
1076     // Otherwise, we had a conflict - probe until we find a unique name. This
1077     // is guaranteed to terminate (and usually in a single iteration) because it
1078     // generates new names by incrementing nextConflictID.
1079     SmallString<64> probeName(name);
1080     probeName.push_back('_');
1081     while (true) {
1082       probeName += llvm::utostr(nextConflictID++);
1083       if (!usedNames.count(probeName)) {
1084         name = StringRef(probeName).copy(usedNameAllocator);
1085         break;
1086       }
1087       probeName.resize(name.size() + 1);
1088     }
1089   }
1090 
1091   usedNames.insert(name, char());
1092   return name;
1093 }
1094 
1095 //===----------------------------------------------------------------------===//
1096 // AsmState
1097 //===----------------------------------------------------------------------===//
1098 
1099 namespace mlir {
1100 namespace detail {
1101 class AsmStateImpl {
1102 public:
1103   explicit AsmStateImpl(Operation *op, AsmState::LocationMap *locationMap)
1104       : interfaces(op->getContext()), nameState(op, interfaces),
1105         locationMap(locationMap) {}
1106 
1107   /// Initialize the alias state to enable the printing of aliases.
1108   void initializeAliases(Operation *op, const OpPrintingFlags &printerFlags) {
1109     aliasState.initialize(op, printerFlags, interfaces);
1110   }
1111 
1112   /// Get an instance of the OpAsmDialectInterface for the given dialect, or
1113   /// null if one wasn't registered.
1114   const OpAsmDialectInterface *getOpAsmInterface(Dialect *dialect) {
1115     return interfaces.getInterfaceFor(dialect);
1116   }
1117 
1118   /// Get the state used for aliases.
1119   AliasState &getAliasState() { return aliasState; }
1120 
1121   /// Get the state used for SSA names.
1122   SSANameState &getSSANameState() { return nameState; }
1123 
1124   /// Register the location, line and column, within the buffer that the given
1125   /// operation was printed at.
1126   void registerOperationLocation(Operation *op, unsigned line, unsigned col) {
1127     if (locationMap)
1128       (*locationMap)[op] = std::make_pair(line, col);
1129   }
1130 
1131 private:
1132   /// Collection of OpAsm interfaces implemented in the context.
1133   DialectInterfaceCollection<OpAsmDialectInterface> interfaces;
1134 
1135   /// The state used for attribute and type aliases.
1136   AliasState aliasState;
1137 
1138   /// The state used for SSA value names.
1139   SSANameState nameState;
1140 
1141   /// An optional location map to be populated.
1142   AsmState::LocationMap *locationMap;
1143 };
1144 } // end namespace detail
1145 } // end namespace mlir
1146 
1147 AsmState::AsmState(Operation *op, LocationMap *locationMap)
1148     : impl(std::make_unique<AsmStateImpl>(op, locationMap)) {}
1149 AsmState::~AsmState() {}
1150 
1151 //===----------------------------------------------------------------------===//
1152 // ModulePrinter
1153 //===----------------------------------------------------------------------===//
1154 
1155 namespace {
1156 class ModulePrinter {
1157 public:
1158   ModulePrinter(raw_ostream &os, OpPrintingFlags flags = llvm::None,
1159                 AsmStateImpl *state = nullptr)
1160       : os(os), printerFlags(flags), state(state) {}
1161   explicit ModulePrinter(ModulePrinter &printer)
1162       : os(printer.os), printerFlags(printer.printerFlags),
1163         state(printer.state) {}
1164 
1165   /// Returns the output stream of the printer.
1166   raw_ostream &getStream() { return os; }
1167 
1168   template <typename Container, typename UnaryFunctor>
1169   inline void interleaveComma(const Container &c, UnaryFunctor each_fn) const {
1170     llvm::interleaveComma(c, os, each_fn);
1171   }
1172 
1173   /// This enum describes the different kinds of elision for the type of an
1174   /// attribute when printing it.
1175   enum class AttrTypeElision {
1176     /// The type must not be elided,
1177     Never,
1178     /// The type may be elided when it matches the default used in the parser
1179     /// (for example i64 is the default for integer attributes).
1180     May,
1181     /// The type must be elided.
1182     Must
1183   };
1184 
1185   /// Print the given attribute.
1186   void printAttribute(Attribute attr,
1187                       AttrTypeElision typeElision = AttrTypeElision::Never);
1188 
1189   void printType(Type type);
1190 
1191   /// Print the given location to the stream. If `allowAlias` is true, this
1192   /// allows for the internal location to use an attribute alias.
1193   void printLocation(LocationAttr loc, bool allowAlias = false);
1194 
1195   void printAffineMap(AffineMap map);
1196   void
1197   printAffineExpr(AffineExpr expr,
1198                   function_ref<void(unsigned, bool)> printValueName = nullptr);
1199   void printAffineConstraint(AffineExpr expr, bool isEq);
1200   void printIntegerSet(IntegerSet set);
1201 
1202 protected:
1203   void printOptionalAttrDict(ArrayRef<NamedAttribute> attrs,
1204                              ArrayRef<StringRef> elidedAttrs = {},
1205                              bool withKeyword = false);
1206   void printNamedAttribute(NamedAttribute attr);
1207   void printTrailingLocation(Location loc);
1208   void printLocationInternal(LocationAttr loc, bool pretty = false);
1209 
1210   /// Print a dense elements attribute. If 'allowHex' is true, a hex string is
1211   /// used instead of individual elements when the elements attr is large.
1212   void printDenseElementsAttr(DenseElementsAttr attr, bool allowHex);
1213 
1214   /// Print a dense string elements attribute.
1215   void printDenseStringElementsAttr(DenseStringElementsAttr attr);
1216 
1217   /// Print a dense elements attribute. If 'allowHex' is true, a hex string is
1218   /// used instead of individual elements when the elements attr is large.
1219   void printDenseIntOrFPElementsAttr(DenseIntOrFPElementsAttr attr,
1220                                      bool allowHex);
1221 
1222   void printDialectAttribute(Attribute attr);
1223   void printDialectType(Type type);
1224 
1225   /// This enum is used to represent the binding strength of the enclosing
1226   /// context that an AffineExprStorage is being printed in, so we can
1227   /// intelligently produce parens.
1228   enum class BindingStrength {
1229     Weak,   // + and -
1230     Strong, // All other binary operators.
1231   };
1232   void printAffineExprInternal(
1233       AffineExpr expr, BindingStrength enclosingTightness,
1234       function_ref<void(unsigned, bool)> printValueName = nullptr);
1235 
1236   /// The output stream for the printer.
1237   raw_ostream &os;
1238 
1239   /// A set of flags to control the printer's behavior.
1240   OpPrintingFlags printerFlags;
1241 
1242   /// An optional printer state for the module.
1243   AsmStateImpl *state;
1244 
1245   /// A tracker for the number of new lines emitted during printing.
1246   NewLineCounter newLine;
1247 };
1248 } // end anonymous namespace
1249 
1250 void ModulePrinter::printTrailingLocation(Location loc) {
1251   // Check to see if we are printing debug information.
1252   if (!printerFlags.shouldPrintDebugInfo())
1253     return;
1254 
1255   os << " ";
1256   printLocation(loc, /*allowAlias=*/true);
1257 }
1258 
1259 void ModulePrinter::printLocationInternal(LocationAttr loc, bool pretty) {
1260   TypeSwitch<LocationAttr>(loc)
1261       .Case<OpaqueLoc>([&](OpaqueLoc loc) {
1262         printLocationInternal(loc.getFallbackLocation(), pretty);
1263       })
1264       .Case<UnknownLoc>([&](UnknownLoc loc) {
1265         if (pretty)
1266           os << "[unknown]";
1267         else
1268           os << "unknown";
1269       })
1270       .Case<FileLineColLoc>([&](FileLineColLoc loc) {
1271         if (pretty) {
1272           os << loc.getFilename();
1273         } else {
1274           os << "\"";
1275           printEscapedString(loc.getFilename(), os);
1276           os << "\"";
1277         }
1278         os << ':' << loc.getLine() << ':' << loc.getColumn();
1279       })
1280       .Case<NameLoc>([&](NameLoc loc) {
1281         os << '\"';
1282         printEscapedString(loc.getName(), os);
1283         os << '\"';
1284 
1285         // Print the child if it isn't unknown.
1286         auto childLoc = loc.getChildLoc();
1287         if (!childLoc.isa<UnknownLoc>()) {
1288           os << '(';
1289           printLocationInternal(childLoc, pretty);
1290           os << ')';
1291         }
1292       })
1293       .Case<CallSiteLoc>([&](CallSiteLoc loc) {
1294         Location caller = loc.getCaller();
1295         Location callee = loc.getCallee();
1296         if (!pretty)
1297           os << "callsite(";
1298         printLocationInternal(callee, pretty);
1299         if (pretty) {
1300           if (callee.isa<NameLoc>()) {
1301             if (caller.isa<FileLineColLoc>()) {
1302               os << " at ";
1303             } else {
1304               os << newLine << " at ";
1305             }
1306           } else {
1307             os << newLine << " at ";
1308           }
1309         } else {
1310           os << " at ";
1311         }
1312         printLocationInternal(caller, pretty);
1313         if (!pretty)
1314           os << ")";
1315       })
1316       .Case<FusedLoc>([&](FusedLoc loc) {
1317         if (!pretty)
1318           os << "fused";
1319         if (Attribute metadata = loc.getMetadata())
1320           os << '<' << metadata << '>';
1321         os << '[';
1322         interleave(
1323             loc.getLocations(),
1324             [&](Location loc) { printLocationInternal(loc, pretty); },
1325             [&]() { os << ", "; });
1326         os << ']';
1327       });
1328 }
1329 
1330 /// Print a floating point value in a way that the parser will be able to
1331 /// round-trip losslessly.
1332 static void printFloatValue(const APFloat &apValue, raw_ostream &os) {
1333   // We would like to output the FP constant value in exponential notation,
1334   // but we cannot do this if doing so will lose precision.  Check here to
1335   // make sure that we only output it in exponential format if we can parse
1336   // the value back and get the same value.
1337   bool isInf = apValue.isInfinity();
1338   bool isNaN = apValue.isNaN();
1339   if (!isInf && !isNaN) {
1340     SmallString<128> strValue;
1341     apValue.toString(strValue, /*FormatPrecision=*/6, /*FormatMaxPadding=*/0,
1342                      /*TruncateZero=*/false);
1343 
1344     // Check to make sure that the stringized number is not some string like
1345     // "Inf" or NaN, that atof will accept, but the lexer will not.  Check
1346     // that the string matches the "[-+]?[0-9]" regex.
1347     assert(((strValue[0] >= '0' && strValue[0] <= '9') ||
1348             ((strValue[0] == '-' || strValue[0] == '+') &&
1349              (strValue[1] >= '0' && strValue[1] <= '9'))) &&
1350            "[-+]?[0-9] regex does not match!");
1351 
1352     // Parse back the stringized version and check that the value is equal
1353     // (i.e., there is no precision loss).
1354     if (APFloat(apValue.getSemantics(), strValue).bitwiseIsEqual(apValue)) {
1355       os << strValue;
1356       return;
1357     }
1358 
1359     // If it is not, use the default format of APFloat instead of the
1360     // exponential notation.
1361     strValue.clear();
1362     apValue.toString(strValue);
1363 
1364     // Make sure that we can parse the default form as a float.
1365     if (StringRef(strValue).contains('.')) {
1366       os << strValue;
1367       return;
1368     }
1369   }
1370 
1371   // Print special values in hexadecimal format. The sign bit should be included
1372   // in the literal.
1373   SmallVector<char, 16> str;
1374   APInt apInt = apValue.bitcastToAPInt();
1375   apInt.toString(str, /*Radix=*/16, /*Signed=*/false,
1376                  /*formatAsCLiteral=*/true);
1377   os << str;
1378 }
1379 
1380 void ModulePrinter::printLocation(LocationAttr loc, bool allowAlias) {
1381   if (printerFlags.shouldPrintDebugInfoPrettyForm())
1382     return printLocationInternal(loc, /*pretty=*/true);
1383 
1384   os << "loc(";
1385   if (!allowAlias || !state || failed(state->getAliasState().getAlias(loc, os)))
1386     printLocationInternal(loc);
1387   os << ')';
1388 }
1389 
1390 /// Returns true if the given dialect symbol data is simple enough to print in
1391 /// the pretty form, i.e. without the enclosing "".
1392 static bool isDialectSymbolSimpleEnoughForPrettyForm(StringRef symName) {
1393   // The name must start with an identifier.
1394   if (symName.empty() || !isalpha(symName.front()))
1395     return false;
1396 
1397   // Ignore all the characters that are valid in an identifier in the symbol
1398   // name.
1399   symName = symName.drop_while(
1400       [](char c) { return llvm::isAlnum(c) || c == '.' || c == '_'; });
1401   if (symName.empty())
1402     return true;
1403 
1404   // If we got to an unexpected character, then it must be a <>.  Check those
1405   // recursively.
1406   if (symName.front() != '<' || symName.back() != '>')
1407     return false;
1408 
1409   SmallVector<char, 8> nestedPunctuation;
1410   do {
1411     // If we ran out of characters, then we had a punctuation mismatch.
1412     if (symName.empty())
1413       return false;
1414 
1415     auto c = symName.front();
1416     symName = symName.drop_front();
1417 
1418     switch (c) {
1419     // We never allow null characters. This is an EOF indicator for the lexer
1420     // which we could handle, but isn't important for any known dialect.
1421     case '\0':
1422       return false;
1423     case '<':
1424     case '[':
1425     case '(':
1426     case '{':
1427       nestedPunctuation.push_back(c);
1428       continue;
1429     case '-':
1430       // Treat `->` as a special token.
1431       if (!symName.empty() && symName.front() == '>') {
1432         symName = symName.drop_front();
1433         continue;
1434       }
1435       break;
1436     // Reject types with mismatched brackets.
1437     case '>':
1438       if (nestedPunctuation.pop_back_val() != '<')
1439         return false;
1440       break;
1441     case ']':
1442       if (nestedPunctuation.pop_back_val() != '[')
1443         return false;
1444       break;
1445     case ')':
1446       if (nestedPunctuation.pop_back_val() != '(')
1447         return false;
1448       break;
1449     case '}':
1450       if (nestedPunctuation.pop_back_val() != '{')
1451         return false;
1452       break;
1453     default:
1454       continue;
1455     }
1456 
1457     // We're done when the punctuation is fully matched.
1458   } while (!nestedPunctuation.empty());
1459 
1460   // If there were extra characters, then we failed.
1461   return symName.empty();
1462 }
1463 
1464 /// Print the given dialect symbol to the stream.
1465 static void printDialectSymbol(raw_ostream &os, StringRef symPrefix,
1466                                StringRef dialectName, StringRef symString) {
1467   os << symPrefix << dialectName;
1468 
1469   // If this symbol name is simple enough, print it directly in pretty form,
1470   // otherwise, we print it as an escaped string.
1471   if (isDialectSymbolSimpleEnoughForPrettyForm(symString)) {
1472     os << '.' << symString;
1473     return;
1474   }
1475 
1476   // TODO: escape the symbol name, it could contain " characters.
1477   os << "<\"" << symString << "\">";
1478 }
1479 
1480 /// Returns true if the given string can be represented as a bare identifier.
1481 static bool isBareIdentifier(StringRef name) {
1482   assert(!name.empty() && "invalid name");
1483 
1484   // By making this unsigned, the value passed in to isalnum will always be
1485   // in the range 0-255. This is important when building with MSVC because
1486   // its implementation will assert. This situation can arise when dealing
1487   // with UTF-8 multibyte characters.
1488   unsigned char firstChar = static_cast<unsigned char>(name[0]);
1489   if (!isalpha(firstChar) && firstChar != '_')
1490     return false;
1491   return llvm::all_of(name.drop_front(), [](unsigned char c) {
1492     return isalnum(c) || c == '_' || c == '$' || c == '.';
1493   });
1494 }
1495 
1496 /// Print the given string as a symbol reference. A symbol reference is
1497 /// represented as a string prefixed with '@'. The reference is surrounded with
1498 /// ""'s and escaped if it has any special or non-printable characters in it.
1499 static void printSymbolReference(StringRef symbolRef, raw_ostream &os) {
1500   assert(!symbolRef.empty() && "expected valid symbol reference");
1501 
1502   // If the symbol can be represented as a bare identifier, write it directly.
1503   if (isBareIdentifier(symbolRef)) {
1504     os << '@' << symbolRef;
1505     return;
1506   }
1507 
1508   // Otherwise, output the reference wrapped in quotes with proper escaping.
1509   os << "@\"";
1510   printEscapedString(symbolRef, os);
1511   os << '"';
1512 }
1513 
1514 // Print out a valid ElementsAttr that is succinct and can represent any
1515 // potential shape/type, for use when eliding a large ElementsAttr.
1516 //
1517 // We choose to use an opaque ElementsAttr literal with conspicuous content to
1518 // hopefully alert readers to the fact that this has been elided.
1519 //
1520 // Unfortunately, neither of the strings of an opaque ElementsAttr literal will
1521 // accept the string "elided". The first string must be a registered dialect
1522 // name and the latter must be a hex constant.
1523 static void printElidedElementsAttr(raw_ostream &os) {
1524   os << R"(opaque<"_", "0xDEADBEEF">)";
1525 }
1526 
1527 void ModulePrinter::printAttribute(Attribute attr,
1528                                    AttrTypeElision typeElision) {
1529   if (!attr) {
1530     os << "<<NULL ATTRIBUTE>>";
1531     return;
1532   }
1533 
1534   // Try to print an alias for this attribute.
1535   if (state && succeeded(state->getAliasState().getAlias(attr, os)))
1536     return;
1537 
1538   auto attrType = attr.getType();
1539   if (auto opaqueAttr = attr.dyn_cast<OpaqueAttr>()) {
1540     printDialectSymbol(os, "#", opaqueAttr.getDialectNamespace(),
1541                        opaqueAttr.getAttrData());
1542   } else if (attr.isa<UnitAttr>()) {
1543     os << "unit";
1544     return;
1545   } else if (auto dictAttr = attr.dyn_cast<DictionaryAttr>()) {
1546     os << '{';
1547     interleaveComma(dictAttr.getValue(),
1548                     [&](NamedAttribute attr) { printNamedAttribute(attr); });
1549     os << '}';
1550 
1551   } else if (auto intAttr = attr.dyn_cast<IntegerAttr>()) {
1552     if (attrType.isSignlessInteger(1)) {
1553       os << (intAttr.getValue().getBoolValue() ? "true" : "false");
1554 
1555       // Boolean integer attributes always elides the type.
1556       return;
1557     }
1558 
1559     // Only print attributes as unsigned if they are explicitly unsigned or are
1560     // signless 1-bit values.  Indexes, signed values, and multi-bit signless
1561     // values print as signed.
1562     bool isUnsigned =
1563         attrType.isUnsignedInteger() || attrType.isSignlessInteger(1);
1564     intAttr.getValue().print(os, !isUnsigned);
1565 
1566     // IntegerAttr elides the type if I64.
1567     if (typeElision == AttrTypeElision::May && attrType.isSignlessInteger(64))
1568       return;
1569 
1570   } else if (auto floatAttr = attr.dyn_cast<FloatAttr>()) {
1571     printFloatValue(floatAttr.getValue(), os);
1572 
1573     // FloatAttr elides the type if F64.
1574     if (typeElision == AttrTypeElision::May && attrType.isF64())
1575       return;
1576 
1577   } else if (auto strAttr = attr.dyn_cast<StringAttr>()) {
1578     os << '"';
1579     printEscapedString(strAttr.getValue(), os);
1580     os << '"';
1581 
1582   } else if (auto arrayAttr = attr.dyn_cast<ArrayAttr>()) {
1583     os << '[';
1584     interleaveComma(arrayAttr.getValue(), [&](Attribute attr) {
1585       printAttribute(attr, AttrTypeElision::May);
1586     });
1587     os << ']';
1588 
1589   } else if (auto affineMapAttr = attr.dyn_cast<AffineMapAttr>()) {
1590     os << "affine_map<";
1591     affineMapAttr.getValue().print(os);
1592     os << '>';
1593 
1594     // AffineMap always elides the type.
1595     return;
1596 
1597   } else if (auto integerSetAttr = attr.dyn_cast<IntegerSetAttr>()) {
1598     os << "affine_set<";
1599     integerSetAttr.getValue().print(os);
1600     os << '>';
1601 
1602     // IntegerSet always elides the type.
1603     return;
1604 
1605   } else if (auto typeAttr = attr.dyn_cast<TypeAttr>()) {
1606     printType(typeAttr.getValue());
1607 
1608   } else if (auto refAttr = attr.dyn_cast<SymbolRefAttr>()) {
1609     printSymbolReference(refAttr.getRootReference(), os);
1610     for (FlatSymbolRefAttr nestedRef : refAttr.getNestedReferences()) {
1611       os << "::";
1612       printSymbolReference(nestedRef.getValue(), os);
1613     }
1614 
1615   } else if (auto opaqueAttr = attr.dyn_cast<OpaqueElementsAttr>()) {
1616     if (printerFlags.shouldElideElementsAttr(opaqueAttr)) {
1617       printElidedElementsAttr(os);
1618     } else {
1619       os << "opaque<\"" << opaqueAttr.getDialect() << "\", \"0x"
1620          << llvm::toHex(opaqueAttr.getValue()) << "\">";
1621     }
1622 
1623   } else if (auto intOrFpEltAttr = attr.dyn_cast<DenseIntOrFPElementsAttr>()) {
1624     if (printerFlags.shouldElideElementsAttr(intOrFpEltAttr)) {
1625       printElidedElementsAttr(os);
1626     } else {
1627       os << "dense<";
1628       printDenseIntOrFPElementsAttr(intOrFpEltAttr, /*allowHex=*/true);
1629       os << '>';
1630     }
1631 
1632   } else if (auto strEltAttr = attr.dyn_cast<DenseStringElementsAttr>()) {
1633     if (printerFlags.shouldElideElementsAttr(strEltAttr)) {
1634       printElidedElementsAttr(os);
1635     } else {
1636       os << "dense<";
1637       printDenseStringElementsAttr(strEltAttr);
1638       os << '>';
1639     }
1640 
1641   } else if (auto sparseEltAttr = attr.dyn_cast<SparseElementsAttr>()) {
1642     if (printerFlags.shouldElideElementsAttr(sparseEltAttr.getIndices()) ||
1643         printerFlags.shouldElideElementsAttr(sparseEltAttr.getValues())) {
1644       printElidedElementsAttr(os);
1645     } else {
1646       os << "sparse<";
1647       DenseIntElementsAttr indices = sparseEltAttr.getIndices();
1648       if (indices.getNumElements() != 0) {
1649         printDenseIntOrFPElementsAttr(indices, /*allowHex=*/false);
1650         os << ", ";
1651         printDenseElementsAttr(sparseEltAttr.getValues(), /*allowHex=*/true);
1652       }
1653       os << '>';
1654     }
1655 
1656   } else if (auto locAttr = attr.dyn_cast<LocationAttr>()) {
1657     printLocation(locAttr);
1658 
1659   } else {
1660     return printDialectAttribute(attr);
1661   }
1662 
1663   // Don't print the type if we must elide it, or if it is a None type.
1664   if (typeElision != AttrTypeElision::Must && !attrType.isa<NoneType>()) {
1665     os << " : ";
1666     printType(attrType);
1667   }
1668 }
1669 
1670 /// Print the integer element of a DenseElementsAttr.
1671 static void printDenseIntElement(const APInt &value, raw_ostream &os,
1672                                  bool isSigned) {
1673   if (value.getBitWidth() == 1)
1674     os << (value.getBoolValue() ? "true" : "false");
1675   else
1676     value.print(os, isSigned);
1677 }
1678 
1679 static void
1680 printDenseElementsAttrImpl(bool isSplat, ShapedType type, raw_ostream &os,
1681                            function_ref<void(unsigned)> printEltFn) {
1682   // Special case for 0-d and splat tensors.
1683   if (isSplat)
1684     return printEltFn(0);
1685 
1686   // Special case for degenerate tensors.
1687   auto numElements = type.getNumElements();
1688   if (numElements == 0)
1689     return;
1690 
1691   // We use a mixed-radix counter to iterate through the shape. When we bump a
1692   // non-least-significant digit, we emit a close bracket. When we next emit an
1693   // element we re-open all closed brackets.
1694 
1695   // The mixed-radix counter, with radices in 'shape'.
1696   int64_t rank = type.getRank();
1697   SmallVector<unsigned, 4> counter(rank, 0);
1698   // The number of brackets that have been opened and not closed.
1699   unsigned openBrackets = 0;
1700 
1701   auto shape = type.getShape();
1702   auto bumpCounter = [&] {
1703     // Bump the least significant digit.
1704     ++counter[rank - 1];
1705     // Iterate backwards bubbling back the increment.
1706     for (unsigned i = rank - 1; i > 0; --i)
1707       if (counter[i] >= shape[i]) {
1708         // Index 'i' is rolled over. Bump (i-1) and close a bracket.
1709         counter[i] = 0;
1710         ++counter[i - 1];
1711         --openBrackets;
1712         os << ']';
1713       }
1714   };
1715 
1716   for (unsigned idx = 0, e = numElements; idx != e; ++idx) {
1717     if (idx != 0)
1718       os << ", ";
1719     while (openBrackets++ < rank)
1720       os << '[';
1721     openBrackets = rank;
1722     printEltFn(idx);
1723     bumpCounter();
1724   }
1725   while (openBrackets-- > 0)
1726     os << ']';
1727 }
1728 
1729 void ModulePrinter::printDenseElementsAttr(DenseElementsAttr attr,
1730                                            bool allowHex) {
1731   if (auto stringAttr = attr.dyn_cast<DenseStringElementsAttr>())
1732     return printDenseStringElementsAttr(stringAttr);
1733 
1734   printDenseIntOrFPElementsAttr(attr.cast<DenseIntOrFPElementsAttr>(),
1735                                 allowHex);
1736 }
1737 
1738 void ModulePrinter::printDenseIntOrFPElementsAttr(DenseIntOrFPElementsAttr attr,
1739                                                   bool allowHex) {
1740   auto type = attr.getType();
1741   auto elementType = type.getElementType();
1742 
1743   // Check to see if we should format this attribute as a hex string.
1744   auto numElements = type.getNumElements();
1745   if (!attr.isSplat() && allowHex &&
1746       shouldPrintElementsAttrWithHex(numElements)) {
1747     ArrayRef<char> rawData = attr.getRawData();
1748     if (llvm::support::endian::system_endianness() ==
1749         llvm::support::endianness::big) {
1750       // Convert endianess in big-endian(BE) machines. `rawData` is BE in BE
1751       // machines. It is converted here to print in LE format.
1752       SmallVector<char, 64> outDataVec(rawData.size());
1753       MutableArrayRef<char> convRawData(outDataVec);
1754       DenseIntOrFPElementsAttr::convertEndianOfArrayRefForBEmachine(
1755           rawData, convRawData, type);
1756       os << '"' << "0x"
1757          << llvm::toHex(StringRef(convRawData.data(), convRawData.size()))
1758          << "\"";
1759     } else {
1760       os << '"' << "0x"
1761          << llvm::toHex(StringRef(rawData.data(), rawData.size())) << "\"";
1762     }
1763 
1764     return;
1765   }
1766 
1767   if (ComplexType complexTy = elementType.dyn_cast<ComplexType>()) {
1768     Type complexElementType = complexTy.getElementType();
1769     // Note: The if and else below had a common lambda function which invoked
1770     // printDenseElementsAttrImpl. This lambda was hitting a bug in gcc 9.1,9.2
1771     // and hence was replaced.
1772     if (complexElementType.isa<IntegerType>()) {
1773       bool isSigned = !complexElementType.isUnsignedInteger();
1774       printDenseElementsAttrImpl(attr.isSplat(), type, os, [&](unsigned index) {
1775         auto complexValue = *(attr.getComplexIntValues().begin() + index);
1776         os << "(";
1777         printDenseIntElement(complexValue.real(), os, isSigned);
1778         os << ",";
1779         printDenseIntElement(complexValue.imag(), os, isSigned);
1780         os << ")";
1781       });
1782     } else {
1783       printDenseElementsAttrImpl(attr.isSplat(), type, os, [&](unsigned index) {
1784         auto complexValue = *(attr.getComplexFloatValues().begin() + index);
1785         os << "(";
1786         printFloatValue(complexValue.real(), os);
1787         os << ",";
1788         printFloatValue(complexValue.imag(), os);
1789         os << ")";
1790       });
1791     }
1792   } else if (elementType.isIntOrIndex()) {
1793     bool isSigned = !elementType.isUnsignedInteger();
1794     auto intValues = attr.getIntValues();
1795     printDenseElementsAttrImpl(attr.isSplat(), type, os, [&](unsigned index) {
1796       printDenseIntElement(*(intValues.begin() + index), os, isSigned);
1797     });
1798   } else {
1799     assert(elementType.isa<FloatType>() && "unexpected element type");
1800     auto floatValues = attr.getFloatValues();
1801     printDenseElementsAttrImpl(attr.isSplat(), type, os, [&](unsigned index) {
1802       printFloatValue(*(floatValues.begin() + index), os);
1803     });
1804   }
1805 }
1806 
1807 void ModulePrinter::printDenseStringElementsAttr(DenseStringElementsAttr attr) {
1808   ArrayRef<StringRef> data = attr.getRawStringData();
1809   auto printFn = [&](unsigned index) {
1810     os << "\"";
1811     printEscapedString(data[index], os);
1812     os << "\"";
1813   };
1814   printDenseElementsAttrImpl(attr.isSplat(), attr.getType(), os, printFn);
1815 }
1816 
1817 void ModulePrinter::printType(Type type) {
1818   if (!type) {
1819     os << "<<NULL TYPE>>";
1820     return;
1821   }
1822 
1823   // Try to print an alias for this type.
1824   if (state && succeeded(state->getAliasState().getAlias(type, os)))
1825     return;
1826 
1827   TypeSwitch<Type>(type)
1828       .Case<OpaqueType>([&](OpaqueType opaqueTy) {
1829         printDialectSymbol(os, "!", opaqueTy.getDialectNamespace(),
1830                            opaqueTy.getTypeData());
1831       })
1832       .Case<IndexType>([&](Type) { os << "index"; })
1833       .Case<BFloat16Type>([&](Type) { os << "bf16"; })
1834       .Case<Float16Type>([&](Type) { os << "f16"; })
1835       .Case<Float32Type>([&](Type) { os << "f32"; })
1836       .Case<Float64Type>([&](Type) { os << "f64"; })
1837       .Case<Float80Type>([&](Type) { os << "f80"; })
1838       .Case<Float128Type>([&](Type) { os << "f128"; })
1839       .Case<IntegerType>([&](IntegerType integerTy) {
1840         if (integerTy.isSigned())
1841           os << 's';
1842         else if (integerTy.isUnsigned())
1843           os << 'u';
1844         os << 'i' << integerTy.getWidth();
1845       })
1846       .Case<FunctionType>([&](FunctionType funcTy) {
1847         os << '(';
1848         interleaveComma(funcTy.getInputs(), [&](Type ty) { printType(ty); });
1849         os << ") -> ";
1850         ArrayRef<Type> results = funcTy.getResults();
1851         if (results.size() == 1 && !results[0].isa<FunctionType>()) {
1852           os << results[0];
1853         } else {
1854           os << '(';
1855           interleaveComma(results, [&](Type ty) { printType(ty); });
1856           os << ')';
1857         }
1858       })
1859       .Case<VectorType>([&](VectorType vectorTy) {
1860         os << "vector<";
1861         for (int64_t dim : vectorTy.getShape())
1862           os << dim << 'x';
1863         os << vectorTy.getElementType() << '>';
1864       })
1865       .Case<RankedTensorType>([&](RankedTensorType tensorTy) {
1866         os << "tensor<";
1867         for (int64_t dim : tensorTy.getShape()) {
1868           if (ShapedType::isDynamic(dim))
1869             os << '?';
1870           else
1871             os << dim;
1872           os << 'x';
1873         }
1874         os << tensorTy.getElementType();
1875         // Only print the encoding attribute value if set.
1876         if (tensorTy.getEncoding()) {
1877           os << ", ";
1878           printAttribute(tensorTy.getEncoding());
1879         }
1880         os << '>';
1881       })
1882       .Case<UnrankedTensorType>([&](UnrankedTensorType tensorTy) {
1883         os << "tensor<*x";
1884         printType(tensorTy.getElementType());
1885         os << '>';
1886       })
1887       .Case<MemRefType>([&](MemRefType memrefTy) {
1888         os << "memref<";
1889         for (int64_t dim : memrefTy.getShape()) {
1890           if (ShapedType::isDynamic(dim))
1891             os << '?';
1892           else
1893             os << dim;
1894           os << 'x';
1895         }
1896         printType(memrefTy.getElementType());
1897         for (auto map : memrefTy.getAffineMaps()) {
1898           os << ", ";
1899           printAttribute(AffineMapAttr::get(map));
1900         }
1901         // Only print the memory space if it is the non-default one.
1902         if (memrefTy.getMemorySpace()) {
1903           os << ", ";
1904           printAttribute(memrefTy.getMemorySpace(), AttrTypeElision::May);
1905         }
1906         os << '>';
1907       })
1908       .Case<UnrankedMemRefType>([&](UnrankedMemRefType memrefTy) {
1909         os << "memref<*x";
1910         printType(memrefTy.getElementType());
1911         // Only print the memory space if it is the non-default one.
1912         if (memrefTy.getMemorySpace()) {
1913           os << ", ";
1914           printAttribute(memrefTy.getMemorySpace(), AttrTypeElision::May);
1915         }
1916         os << '>';
1917       })
1918       .Case<ComplexType>([&](ComplexType complexTy) {
1919         os << "complex<";
1920         printType(complexTy.getElementType());
1921         os << '>';
1922       })
1923       .Case<TupleType>([&](TupleType tupleTy) {
1924         os << "tuple<";
1925         interleaveComma(tupleTy.getTypes(),
1926                         [&](Type type) { printType(type); });
1927         os << '>';
1928       })
1929       .Case<NoneType>([&](Type) { os << "none"; })
1930       .Default([&](Type type) { return printDialectType(type); });
1931 }
1932 
1933 void ModulePrinter::printOptionalAttrDict(ArrayRef<NamedAttribute> attrs,
1934                                           ArrayRef<StringRef> elidedAttrs,
1935                                           bool withKeyword) {
1936   // If there are no attributes, then there is nothing to be done.
1937   if (attrs.empty())
1938     return;
1939 
1940   // Functor used to print a filtered attribute list.
1941   auto printFilteredAttributesFn = [&](auto filteredAttrs) {
1942     // Print the 'attributes' keyword if necessary.
1943     if (withKeyword)
1944       os << " attributes";
1945 
1946     // Otherwise, print them all out in braces.
1947     os << " {";
1948     interleaveComma(filteredAttrs,
1949                     [&](NamedAttribute attr) { printNamedAttribute(attr); });
1950     os << '}';
1951   };
1952 
1953   // If no attributes are elided, we can directly print with no filtering.
1954   if (elidedAttrs.empty())
1955     return printFilteredAttributesFn(attrs);
1956 
1957   // Otherwise, filter out any attributes that shouldn't be included.
1958   llvm::SmallDenseSet<StringRef> elidedAttrsSet(elidedAttrs.begin(),
1959                                                 elidedAttrs.end());
1960   auto filteredAttrs = llvm::make_filter_range(attrs, [&](NamedAttribute attr) {
1961     return !elidedAttrsSet.contains(attr.first.strref());
1962   });
1963   if (!filteredAttrs.empty())
1964     printFilteredAttributesFn(filteredAttrs);
1965 }
1966 
1967 void ModulePrinter::printNamedAttribute(NamedAttribute attr) {
1968   if (isBareIdentifier(attr.first)) {
1969     os << attr.first;
1970   } else {
1971     os << '"';
1972     printEscapedString(attr.first.strref(), os);
1973     os << '"';
1974   }
1975 
1976   // Pretty printing elides the attribute value for unit attributes.
1977   if (attr.second.isa<UnitAttr>())
1978     return;
1979 
1980   os << " = ";
1981   printAttribute(attr.second);
1982 }
1983 
1984 //===----------------------------------------------------------------------===//
1985 // CustomDialectAsmPrinter
1986 //===----------------------------------------------------------------------===//
1987 
1988 namespace {
1989 /// This class provides the main specialization of the DialectAsmPrinter that is
1990 /// used to provide support for print attributes and types. This hooks allows
1991 /// for dialects to hook into the main ModulePrinter.
1992 struct CustomDialectAsmPrinter : public DialectAsmPrinter {
1993 public:
1994   CustomDialectAsmPrinter(ModulePrinter &printer) : printer(printer) {}
1995   ~CustomDialectAsmPrinter() override {}
1996 
1997   raw_ostream &getStream() const override { return printer.getStream(); }
1998 
1999   /// Print the given attribute to the stream.
2000   void printAttribute(Attribute attr) override { printer.printAttribute(attr); }
2001 
2002   /// Print the given floating point value in a stablized form.
2003   void printFloat(const APFloat &value) override {
2004     printFloatValue(value, getStream());
2005   }
2006 
2007   /// Print the given type to the stream.
2008   void printType(Type type) override { printer.printType(type); }
2009 
2010   /// The main module printer.
2011   ModulePrinter &printer;
2012 };
2013 } // end anonymous namespace
2014 
2015 void ModulePrinter::printDialectAttribute(Attribute attr) {
2016   auto &dialect = attr.getDialect();
2017 
2018   // Ask the dialect to serialize the attribute to a string.
2019   std::string attrName;
2020   {
2021     llvm::raw_string_ostream attrNameStr(attrName);
2022     ModulePrinter subPrinter(attrNameStr, printerFlags, state);
2023     CustomDialectAsmPrinter printer(subPrinter);
2024     dialect.printAttribute(attr, printer);
2025   }
2026   printDialectSymbol(os, "#", dialect.getNamespace(), attrName);
2027 }
2028 
2029 void ModulePrinter::printDialectType(Type type) {
2030   auto &dialect = type.getDialect();
2031 
2032   // Ask the dialect to serialize the type to a string.
2033   std::string typeName;
2034   {
2035     llvm::raw_string_ostream typeNameStr(typeName);
2036     ModulePrinter subPrinter(typeNameStr, printerFlags, state);
2037     CustomDialectAsmPrinter printer(subPrinter);
2038     dialect.printType(type, printer);
2039   }
2040   printDialectSymbol(os, "!", dialect.getNamespace(), typeName);
2041 }
2042 
2043 //===----------------------------------------------------------------------===//
2044 // Affine expressions and maps
2045 //===----------------------------------------------------------------------===//
2046 
2047 void ModulePrinter::printAffineExpr(
2048     AffineExpr expr, function_ref<void(unsigned, bool)> printValueName) {
2049   printAffineExprInternal(expr, BindingStrength::Weak, printValueName);
2050 }
2051 
2052 void ModulePrinter::printAffineExprInternal(
2053     AffineExpr expr, BindingStrength enclosingTightness,
2054     function_ref<void(unsigned, bool)> printValueName) {
2055   const char *binopSpelling = nullptr;
2056   switch (expr.getKind()) {
2057   case AffineExprKind::SymbolId: {
2058     unsigned pos = expr.cast<AffineSymbolExpr>().getPosition();
2059     if (printValueName)
2060       printValueName(pos, /*isSymbol=*/true);
2061     else
2062       os << 's' << pos;
2063     return;
2064   }
2065   case AffineExprKind::DimId: {
2066     unsigned pos = expr.cast<AffineDimExpr>().getPosition();
2067     if (printValueName)
2068       printValueName(pos, /*isSymbol=*/false);
2069     else
2070       os << 'd' << pos;
2071     return;
2072   }
2073   case AffineExprKind::Constant:
2074     os << expr.cast<AffineConstantExpr>().getValue();
2075     return;
2076   case AffineExprKind::Add:
2077     binopSpelling = " + ";
2078     break;
2079   case AffineExprKind::Mul:
2080     binopSpelling = " * ";
2081     break;
2082   case AffineExprKind::FloorDiv:
2083     binopSpelling = " floordiv ";
2084     break;
2085   case AffineExprKind::CeilDiv:
2086     binopSpelling = " ceildiv ";
2087     break;
2088   case AffineExprKind::Mod:
2089     binopSpelling = " mod ";
2090     break;
2091   }
2092 
2093   auto binOp = expr.cast<AffineBinaryOpExpr>();
2094   AffineExpr lhsExpr = binOp.getLHS();
2095   AffineExpr rhsExpr = binOp.getRHS();
2096 
2097   // Handle tightly binding binary operators.
2098   if (binOp.getKind() != AffineExprKind::Add) {
2099     if (enclosingTightness == BindingStrength::Strong)
2100       os << '(';
2101 
2102     // Pretty print multiplication with -1.
2103     auto rhsConst = rhsExpr.dyn_cast<AffineConstantExpr>();
2104     if (rhsConst && binOp.getKind() == AffineExprKind::Mul &&
2105         rhsConst.getValue() == -1) {
2106       os << "-";
2107       printAffineExprInternal(lhsExpr, BindingStrength::Strong, printValueName);
2108       if (enclosingTightness == BindingStrength::Strong)
2109         os << ')';
2110       return;
2111     }
2112 
2113     printAffineExprInternal(lhsExpr, BindingStrength::Strong, printValueName);
2114 
2115     os << binopSpelling;
2116     printAffineExprInternal(rhsExpr, BindingStrength::Strong, printValueName);
2117 
2118     if (enclosingTightness == BindingStrength::Strong)
2119       os << ')';
2120     return;
2121   }
2122 
2123   // Print out special "pretty" forms for add.
2124   if (enclosingTightness == BindingStrength::Strong)
2125     os << '(';
2126 
2127   // Pretty print addition to a product that has a negative operand as a
2128   // subtraction.
2129   if (auto rhs = rhsExpr.dyn_cast<AffineBinaryOpExpr>()) {
2130     if (rhs.getKind() == AffineExprKind::Mul) {
2131       AffineExpr rrhsExpr = rhs.getRHS();
2132       if (auto rrhs = rrhsExpr.dyn_cast<AffineConstantExpr>()) {
2133         if (rrhs.getValue() == -1) {
2134           printAffineExprInternal(lhsExpr, BindingStrength::Weak,
2135                                   printValueName);
2136           os << " - ";
2137           if (rhs.getLHS().getKind() == AffineExprKind::Add) {
2138             printAffineExprInternal(rhs.getLHS(), BindingStrength::Strong,
2139                                     printValueName);
2140           } else {
2141             printAffineExprInternal(rhs.getLHS(), BindingStrength::Weak,
2142                                     printValueName);
2143           }
2144 
2145           if (enclosingTightness == BindingStrength::Strong)
2146             os << ')';
2147           return;
2148         }
2149 
2150         if (rrhs.getValue() < -1) {
2151           printAffineExprInternal(lhsExpr, BindingStrength::Weak,
2152                                   printValueName);
2153           os << " - ";
2154           printAffineExprInternal(rhs.getLHS(), BindingStrength::Strong,
2155                                   printValueName);
2156           os << " * " << -rrhs.getValue();
2157           if (enclosingTightness == BindingStrength::Strong)
2158             os << ')';
2159           return;
2160         }
2161       }
2162     }
2163   }
2164 
2165   // Pretty print addition to a negative number as a subtraction.
2166   if (auto rhsConst = rhsExpr.dyn_cast<AffineConstantExpr>()) {
2167     if (rhsConst.getValue() < 0) {
2168       printAffineExprInternal(lhsExpr, BindingStrength::Weak, printValueName);
2169       os << " - " << -rhsConst.getValue();
2170       if (enclosingTightness == BindingStrength::Strong)
2171         os << ')';
2172       return;
2173     }
2174   }
2175 
2176   printAffineExprInternal(lhsExpr, BindingStrength::Weak, printValueName);
2177 
2178   os << " + ";
2179   printAffineExprInternal(rhsExpr, BindingStrength::Weak, printValueName);
2180 
2181   if (enclosingTightness == BindingStrength::Strong)
2182     os << ')';
2183 }
2184 
2185 void ModulePrinter::printAffineConstraint(AffineExpr expr, bool isEq) {
2186   printAffineExprInternal(expr, BindingStrength::Weak);
2187   isEq ? os << " == 0" : os << " >= 0";
2188 }
2189 
2190 void ModulePrinter::printAffineMap(AffineMap map) {
2191   // Dimension identifiers.
2192   os << '(';
2193   for (int i = 0; i < (int)map.getNumDims() - 1; ++i)
2194     os << 'd' << i << ", ";
2195   if (map.getNumDims() >= 1)
2196     os << 'd' << map.getNumDims() - 1;
2197   os << ')';
2198 
2199   // Symbolic identifiers.
2200   if (map.getNumSymbols() != 0) {
2201     os << '[';
2202     for (unsigned i = 0; i < map.getNumSymbols() - 1; ++i)
2203       os << 's' << i << ", ";
2204     if (map.getNumSymbols() >= 1)
2205       os << 's' << map.getNumSymbols() - 1;
2206     os << ']';
2207   }
2208 
2209   // Result affine expressions.
2210   os << " -> (";
2211   interleaveComma(map.getResults(),
2212                   [&](AffineExpr expr) { printAffineExpr(expr); });
2213   os << ')';
2214 }
2215 
2216 void ModulePrinter::printIntegerSet(IntegerSet set) {
2217   // Dimension identifiers.
2218   os << '(';
2219   for (unsigned i = 1; i < set.getNumDims(); ++i)
2220     os << 'd' << i - 1 << ", ";
2221   if (set.getNumDims() >= 1)
2222     os << 'd' << set.getNumDims() - 1;
2223   os << ')';
2224 
2225   // Symbolic identifiers.
2226   if (set.getNumSymbols() != 0) {
2227     os << '[';
2228     for (unsigned i = 0; i < set.getNumSymbols() - 1; ++i)
2229       os << 's' << i << ", ";
2230     if (set.getNumSymbols() >= 1)
2231       os << 's' << set.getNumSymbols() - 1;
2232     os << ']';
2233   }
2234 
2235   // Print constraints.
2236   os << " : (";
2237   int numConstraints = set.getNumConstraints();
2238   for (int i = 1; i < numConstraints; ++i) {
2239     printAffineConstraint(set.getConstraint(i - 1), set.isEq(i - 1));
2240     os << ", ";
2241   }
2242   if (numConstraints >= 1)
2243     printAffineConstraint(set.getConstraint(numConstraints - 1),
2244                           set.isEq(numConstraints - 1));
2245   os << ')';
2246 }
2247 
2248 //===----------------------------------------------------------------------===//
2249 // OperationPrinter
2250 //===----------------------------------------------------------------------===//
2251 
2252 namespace {
2253 /// This class contains the logic for printing operations, regions, and blocks.
2254 class OperationPrinter : public ModulePrinter, private OpAsmPrinter {
2255 public:
2256   explicit OperationPrinter(raw_ostream &os, OpPrintingFlags flags,
2257                             AsmStateImpl &state)
2258       : ModulePrinter(os, flags, &state) {}
2259 
2260   /// Print the given top-level operation.
2261   void printTopLevelOperation(Operation *op);
2262 
2263   /// Print the given operation with its indent and location.
2264   void print(Operation *op);
2265   /// Print the bare location, not including indentation/location/etc.
2266   void printOperation(Operation *op);
2267   /// Print the given operation in the generic form.
2268   void printGenericOp(Operation *op) override;
2269 
2270   /// Print the name of the given block.
2271   void printBlockName(Block *block);
2272 
2273   /// Print the given block. If 'printBlockArgs' is false, the arguments of the
2274   /// block are not printed. If 'printBlockTerminator' is false, the terminator
2275   /// operation of the block is not printed.
2276   void print(Block *block, bool printBlockArgs = true,
2277              bool printBlockTerminator = true);
2278 
2279   /// Print the ID of the given value, optionally with its result number.
2280   void printValueID(Value value, bool printResultNo = true,
2281                     raw_ostream *streamOverride = nullptr) const;
2282 
2283   //===--------------------------------------------------------------------===//
2284   // OpAsmPrinter methods
2285   //===--------------------------------------------------------------------===//
2286 
2287   /// Return the current stream of the printer.
2288   raw_ostream &getStream() const override { return os; }
2289 
2290   /// Print a newline and indent the printer to the start of the current
2291   /// operation.
2292   void printNewline() override {
2293     os << newLine;
2294     os.indent(currentIndent);
2295   }
2296 
2297   /// Print the given type.
2298   void printType(Type type) override { ModulePrinter::printType(type); }
2299 
2300   /// Print the given attribute.
2301   void printAttribute(Attribute attr) override {
2302     ModulePrinter::printAttribute(attr);
2303   }
2304 
2305   /// Print the given attribute without its type. The corresponding parser must
2306   /// provide a valid type for the attribute.
2307   void printAttributeWithoutType(Attribute attr) override {
2308     ModulePrinter::printAttribute(attr, AttrTypeElision::Must);
2309   }
2310 
2311   /// Print the ID for the given value.
2312   void printOperand(Value value) override { printValueID(value); }
2313   void printOperand(Value value, raw_ostream &os) override {
2314     printValueID(value, /*printResultNo=*/true, &os);
2315   }
2316 
2317   /// Print an optional attribute dictionary with a given set of elided values.
2318   void printOptionalAttrDict(ArrayRef<NamedAttribute> attrs,
2319                              ArrayRef<StringRef> elidedAttrs = {}) override {
2320     ModulePrinter::printOptionalAttrDict(attrs, elidedAttrs);
2321   }
2322   void printOptionalAttrDictWithKeyword(
2323       ArrayRef<NamedAttribute> attrs,
2324       ArrayRef<StringRef> elidedAttrs = {}) override {
2325     ModulePrinter::printOptionalAttrDict(attrs, elidedAttrs,
2326                                          /*withKeyword=*/true);
2327   }
2328 
2329   /// Print the given successor.
2330   void printSuccessor(Block *successor) override;
2331 
2332   /// Print an operation successor with the operands used for the block
2333   /// arguments.
2334   void printSuccessorAndUseList(Block *successor,
2335                                 ValueRange succOperands) override;
2336 
2337   /// Print the given region.
2338   void printRegion(Region &region, bool printEntryBlockArgs,
2339                    bool printBlockTerminators, bool printEmptyBlock) override;
2340 
2341   /// Renumber the arguments for the specified region to the same names as the
2342   /// SSA values in namesToUse. This may only be used for IsolatedFromAbove
2343   /// operations. If any entry in namesToUse is null, the corresponding
2344   /// argument name is left alone.
2345   void shadowRegionArgs(Region &region, ValueRange namesToUse) override {
2346     state->getSSANameState().shadowRegionArgs(region, namesToUse);
2347   }
2348 
2349   /// Print the given affine map with the symbol and dimension operands printed
2350   /// inline with the map.
2351   void printAffineMapOfSSAIds(AffineMapAttr mapAttr,
2352                               ValueRange operands) override;
2353 
2354   /// Print the given string as a symbol reference.
2355   void printSymbolName(StringRef symbolRef) override {
2356     ::printSymbolReference(symbolRef, os);
2357   }
2358 
2359 private:
2360   /// The number of spaces used for indenting nested operations.
2361   const static unsigned indentWidth = 2;
2362 
2363   // This is the current indentation level for nested structures.
2364   unsigned currentIndent = 0;
2365 };
2366 } // end anonymous namespace
2367 
2368 void OperationPrinter::printTopLevelOperation(Operation *op) {
2369   // Output the aliases at the top level that can't be deferred.
2370   state->getAliasState().printNonDeferredAliases(os, newLine);
2371 
2372   // Print the module.
2373   print(op);
2374   os << newLine;
2375 
2376   // Output the aliases at the top level that can be deferred.
2377   state->getAliasState().printDeferredAliases(os, newLine);
2378 }
2379 
2380 void OperationPrinter::print(Operation *op) {
2381   // Track the location of this operation.
2382   state->registerOperationLocation(op, newLine.curLine, currentIndent);
2383 
2384   os.indent(currentIndent);
2385   printOperation(op);
2386   printTrailingLocation(op->getLoc());
2387 }
2388 
2389 void OperationPrinter::printOperation(Operation *op) {
2390   if (size_t numResults = op->getNumResults()) {
2391     auto printResultGroup = [&](size_t resultNo, size_t resultCount) {
2392       printValueID(op->getResult(resultNo), /*printResultNo=*/false);
2393       if (resultCount > 1)
2394         os << ':' << resultCount;
2395     };
2396 
2397     // Check to see if this operation has multiple result groups.
2398     ArrayRef<int> resultGroups = state->getSSANameState().getOpResultGroups(op);
2399     if (!resultGroups.empty()) {
2400       // Interleave the groups excluding the last one, this one will be handled
2401       // separately.
2402       interleaveComma(llvm::seq<int>(0, resultGroups.size() - 1), [&](int i) {
2403         printResultGroup(resultGroups[i],
2404                          resultGroups[i + 1] - resultGroups[i]);
2405       });
2406       os << ", ";
2407       printResultGroup(resultGroups.back(), numResults - resultGroups.back());
2408 
2409     } else {
2410       printResultGroup(/*resultNo=*/0, /*resultCount=*/numResults);
2411     }
2412 
2413     os << " = ";
2414   }
2415 
2416   // If requested, always print the generic form.
2417   if (!printerFlags.shouldPrintGenericOpForm()) {
2418     // Check to see if this is a known operation.  If so, use the registered
2419     // custom printer hook.
2420     if (auto *opInfo = op->getAbstractOperation()) {
2421       opInfo->printAssembly(op, *this);
2422       return;
2423     }
2424     // Otherwise try to dispatch to the dialect, if available.
2425     if (Dialect *dialect = op->getDialect()) {
2426       if (succeeded(dialect->printOperation(op, *this)))
2427         return;
2428     }
2429   }
2430 
2431   // Otherwise print with the generic assembly form.
2432   printGenericOp(op);
2433 }
2434 
2435 void OperationPrinter::printGenericOp(Operation *op) {
2436   os << '"';
2437   printEscapedString(op->getName().getStringRef(), os);
2438   os << "\"(";
2439   interleaveComma(op->getOperands(), [&](Value value) { printValueID(value); });
2440   os << ')';
2441 
2442   // For terminators, print the list of successors and their operands.
2443   if (op->getNumSuccessors() != 0) {
2444     os << '[';
2445     interleaveComma(op->getSuccessors(),
2446                     [&](Block *successor) { printBlockName(successor); });
2447     os << ']';
2448   }
2449 
2450   // Print regions.
2451   if (op->getNumRegions() != 0) {
2452     os << " (";
2453     interleaveComma(op->getRegions(), [&](Region &region) {
2454       printRegion(region, /*printEntryBlockArgs=*/true,
2455                   /*printBlockTerminators=*/true, /*printEmptyBlock=*/true);
2456     });
2457     os << ')';
2458   }
2459 
2460   auto attrs = op->getAttrs();
2461   printOptionalAttrDict(attrs);
2462 
2463   // Print the type signature of the operation.
2464   os << " : ";
2465   printFunctionalType(op);
2466 }
2467 
2468 void OperationPrinter::printBlockName(Block *block) {
2469   auto id = state->getSSANameState().getBlockID(block);
2470   if (id != SSANameState::NameSentinel)
2471     os << "^bb" << id;
2472   else
2473     os << "^INVALIDBLOCK";
2474 }
2475 
2476 void OperationPrinter::print(Block *block, bool printBlockArgs,
2477                              bool printBlockTerminator) {
2478   // Print the block label and argument list if requested.
2479   if (printBlockArgs) {
2480     os.indent(currentIndent);
2481     printBlockName(block);
2482 
2483     // Print the argument list if non-empty.
2484     if (!block->args_empty()) {
2485       os << '(';
2486       interleaveComma(block->getArguments(), [&](BlockArgument arg) {
2487         printValueID(arg);
2488         os << ": ";
2489         printType(arg.getType());
2490       });
2491       os << ')';
2492     }
2493     os << ':';
2494 
2495     // Print out some context information about the predecessors of this block.
2496     if (!block->getParent()) {
2497       os << "  // block is not in a region!";
2498     } else if (block->hasNoPredecessors()) {
2499       os << "  // no predecessors";
2500     } else if (auto *pred = block->getSinglePredecessor()) {
2501       os << "  // pred: ";
2502       printBlockName(pred);
2503     } else {
2504       // We want to print the predecessors in increasing numeric order, not in
2505       // whatever order the use-list is in, so gather and sort them.
2506       SmallVector<std::pair<unsigned, Block *>, 4> predIDs;
2507       for (auto *pred : block->getPredecessors())
2508         predIDs.push_back({state->getSSANameState().getBlockID(pred), pred});
2509       llvm::array_pod_sort(predIDs.begin(), predIDs.end());
2510 
2511       os << "  // " << predIDs.size() << " preds: ";
2512 
2513       interleaveComma(predIDs, [&](std::pair<unsigned, Block *> pred) {
2514         printBlockName(pred.second);
2515       });
2516     }
2517     os << newLine;
2518   }
2519 
2520   currentIndent += indentWidth;
2521   auto range = llvm::make_range(
2522       block->begin(), std::prev(block->end(), printBlockTerminator ? 0 : 1));
2523   for (auto &op : range) {
2524     print(&op);
2525     os << newLine;
2526   }
2527   currentIndent -= indentWidth;
2528 }
2529 
2530 void OperationPrinter::printValueID(Value value, bool printResultNo,
2531                                     raw_ostream *streamOverride) const {
2532   state->getSSANameState().printValueID(value, printResultNo,
2533                                         streamOverride ? *streamOverride : os);
2534 }
2535 
2536 void OperationPrinter::printSuccessor(Block *successor) {
2537   printBlockName(successor);
2538 }
2539 
2540 void OperationPrinter::printSuccessorAndUseList(Block *successor,
2541                                                 ValueRange succOperands) {
2542   printBlockName(successor);
2543   if (succOperands.empty())
2544     return;
2545 
2546   os << '(';
2547   interleaveComma(succOperands,
2548                   [this](Value operand) { printValueID(operand); });
2549   os << " : ";
2550   interleaveComma(succOperands,
2551                   [this](Value operand) { printType(operand.getType()); });
2552   os << ')';
2553 }
2554 
2555 void OperationPrinter::printRegion(Region &region, bool printEntryBlockArgs,
2556                                    bool printBlockTerminators,
2557                                    bool printEmptyBlock) {
2558   os << " {" << newLine;
2559   if (!region.empty()) {
2560     auto *entryBlock = &region.front();
2561     // Force printing the block header if printEmptyBlock is set and the block
2562     // is empty or if printEntryBlockArgs is set and there are arguments to
2563     // print.
2564     bool shouldAlwaysPrintBlockHeader =
2565         (printEmptyBlock && entryBlock->empty()) ||
2566         (printEntryBlockArgs && entryBlock->getNumArguments() != 0);
2567     print(entryBlock, shouldAlwaysPrintBlockHeader, printBlockTerminators);
2568     for (auto &b : llvm::drop_begin(region.getBlocks(), 1))
2569       print(&b);
2570   }
2571   os.indent(currentIndent) << "}";
2572 }
2573 
2574 void OperationPrinter::printAffineMapOfSSAIds(AffineMapAttr mapAttr,
2575                                               ValueRange operands) {
2576   AffineMap map = mapAttr.getValue();
2577   unsigned numDims = map.getNumDims();
2578   auto printValueName = [&](unsigned pos, bool isSymbol) {
2579     unsigned index = isSymbol ? numDims + pos : pos;
2580     assert(index < operands.size());
2581     if (isSymbol)
2582       os << "symbol(";
2583     printValueID(operands[index]);
2584     if (isSymbol)
2585       os << ')';
2586   };
2587 
2588   interleaveComma(map.getResults(), [&](AffineExpr expr) {
2589     printAffineExpr(expr, printValueName);
2590   });
2591 }
2592 
2593 //===----------------------------------------------------------------------===//
2594 // print and dump methods
2595 //===----------------------------------------------------------------------===//
2596 
2597 void Attribute::print(raw_ostream &os) const {
2598   ModulePrinter(os).printAttribute(*this);
2599 }
2600 
2601 void Attribute::dump() const {
2602   print(llvm::errs());
2603   llvm::errs() << "\n";
2604 }
2605 
2606 void Type::print(raw_ostream &os) { ModulePrinter(os).printType(*this); }
2607 
2608 void Type::dump() { print(llvm::errs()); }
2609 
2610 void AffineMap::dump() const {
2611   print(llvm::errs());
2612   llvm::errs() << "\n";
2613 }
2614 
2615 void IntegerSet::dump() const {
2616   print(llvm::errs());
2617   llvm::errs() << "\n";
2618 }
2619 
2620 void AffineExpr::print(raw_ostream &os) const {
2621   if (!expr) {
2622     os << "<<NULL AFFINE EXPR>>";
2623     return;
2624   }
2625   ModulePrinter(os).printAffineExpr(*this);
2626 }
2627 
2628 void AffineExpr::dump() const {
2629   print(llvm::errs());
2630   llvm::errs() << "\n";
2631 }
2632 
2633 void AffineMap::print(raw_ostream &os) const {
2634   if (!map) {
2635     os << "<<NULL AFFINE MAP>>";
2636     return;
2637   }
2638   ModulePrinter(os).printAffineMap(*this);
2639 }
2640 
2641 void IntegerSet::print(raw_ostream &os) const {
2642   ModulePrinter(os).printIntegerSet(*this);
2643 }
2644 
2645 void Value::print(raw_ostream &os) {
2646   if (auto *op = getDefiningOp())
2647     return op->print(os);
2648   // TODO: Improve this.
2649   BlockArgument arg = this->cast<BlockArgument>();
2650   os << "<block argument> of type '" << arg.getType()
2651      << "' at index: " << arg.getArgNumber() << '\n';
2652 }
2653 void Value::print(raw_ostream &os, AsmState &state) {
2654   if (auto *op = getDefiningOp())
2655     return op->print(os, state);
2656 
2657   // TODO: Improve this.
2658   BlockArgument arg = this->cast<BlockArgument>();
2659   os << "<block argument> of type '" << arg.getType()
2660      << "' at index: " << arg.getArgNumber() << '\n';
2661 }
2662 
2663 void Value::dump() {
2664   print(llvm::errs());
2665   llvm::errs() << "\n";
2666 }
2667 
2668 void Value::printAsOperand(raw_ostream &os, AsmState &state) {
2669   // TODO: This doesn't necessarily capture all potential cases.
2670   // Currently, region arguments can be shadowed when printing the main
2671   // operation. If the IR hasn't been printed, this will produce the old SSA
2672   // name and not the shadowed name.
2673   state.getImpl().getSSANameState().printValueID(*this, /*printResultNo=*/true,
2674                                                  os);
2675 }
2676 
2677 void Operation::print(raw_ostream &os, OpPrintingFlags flags) {
2678   // If this is a top level operation, we also print aliases.
2679   if (!getParent() && !flags.shouldUseLocalScope()) {
2680     AsmState state(this);
2681     state.getImpl().initializeAliases(this, flags);
2682     print(os, state, flags);
2683     return;
2684   }
2685 
2686   // Find the operation to number from based upon the provided flags.
2687   Operation *op = this;
2688   bool shouldUseLocalScope = flags.shouldUseLocalScope();
2689   do {
2690     // If we are printing local scope, stop at the first operation that is
2691     // isolated from above.
2692     if (shouldUseLocalScope && op->hasTrait<OpTrait::IsIsolatedFromAbove>())
2693       break;
2694 
2695     // Otherwise, traverse up to the next parent.
2696     Operation *parentOp = op->getParentOp();
2697     if (!parentOp)
2698       break;
2699     op = parentOp;
2700   } while (true);
2701 
2702   AsmState state(op);
2703   print(os, state, flags);
2704 }
2705 void Operation::print(raw_ostream &os, AsmState &state, OpPrintingFlags flags) {
2706   OperationPrinter printer(os, flags, state.getImpl());
2707   if (!getParent() && !flags.shouldUseLocalScope())
2708     printer.printTopLevelOperation(this);
2709   else
2710     printer.print(this);
2711 }
2712 
2713 void Operation::dump() {
2714   print(llvm::errs(), OpPrintingFlags().useLocalScope());
2715   llvm::errs() << "\n";
2716 }
2717 
2718 void Block::print(raw_ostream &os) {
2719   Operation *parentOp = getParentOp();
2720   if (!parentOp) {
2721     os << "<<UNLINKED BLOCK>>\n";
2722     return;
2723   }
2724   // Get the top-level op.
2725   while (auto *nextOp = parentOp->getParentOp())
2726     parentOp = nextOp;
2727 
2728   AsmState state(parentOp);
2729   print(os, state);
2730 }
2731 void Block::print(raw_ostream &os, AsmState &state) {
2732   OperationPrinter(os, /*flags=*/llvm::None, state.getImpl()).print(this);
2733 }
2734 
2735 void Block::dump() { print(llvm::errs()); }
2736 
2737 /// Print out the name of the block without printing its body.
2738 void Block::printAsOperand(raw_ostream &os, bool printType) {
2739   Operation *parentOp = getParentOp();
2740   if (!parentOp) {
2741     os << "<<UNLINKED BLOCK>>\n";
2742     return;
2743   }
2744   AsmState state(parentOp);
2745   printAsOperand(os, state);
2746 }
2747 void Block::printAsOperand(raw_ostream &os, AsmState &state) {
2748   OperationPrinter printer(os, /*flags=*/llvm::None, state.getImpl());
2749   printer.printBlockName(this);
2750 }
2751