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