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