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