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