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