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