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