1 //===- Attributes.cpp - MLIR Affine Expr Classes --------------------------===// 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 #include "mlir/IR/Attributes.h" 10 #include "AttributeDetail.h" 11 #include "mlir/IR/AffineMap.h" 12 #include "mlir/IR/Diagnostics.h" 13 #include "mlir/IR/Dialect.h" 14 #include "mlir/IR/Function.h" 15 #include "mlir/IR/IntegerSet.h" 16 #include "mlir/IR/Types.h" 17 #include "llvm/ADT/Sequence.h" 18 #include "llvm/ADT/Twine.h" 19 20 using namespace mlir; 21 using namespace mlir::detail; 22 23 //===----------------------------------------------------------------------===// 24 // AttributeStorage 25 //===----------------------------------------------------------------------===// 26 27 AttributeStorage::AttributeStorage(Type type) 28 : type(type.getAsOpaquePointer()) {} 29 AttributeStorage::AttributeStorage() : type(nullptr) {} 30 31 Type AttributeStorage::getType() const { 32 return Type::getFromOpaquePointer(type); 33 } 34 void AttributeStorage::setType(Type newType) { 35 type = newType.getAsOpaquePointer(); 36 } 37 38 //===----------------------------------------------------------------------===// 39 // Attribute 40 //===----------------------------------------------------------------------===// 41 42 /// Return the type of this attribute. 43 Type Attribute::getType() const { return impl->getType(); } 44 45 /// Return the context this attribute belongs to. 46 MLIRContext *Attribute::getContext() const { return getType().getContext(); } 47 48 /// Get the dialect this attribute is registered to. 49 Dialect &Attribute::getDialect() const { return impl->getDialect(); } 50 51 //===----------------------------------------------------------------------===// 52 // AffineMapAttr 53 //===----------------------------------------------------------------------===// 54 55 AffineMapAttr AffineMapAttr::get(AffineMap value) { 56 return Base::get(value.getContext(), StandardAttributes::AffineMap, value); 57 } 58 59 AffineMap AffineMapAttr::getValue() const { return getImpl()->value; } 60 61 //===----------------------------------------------------------------------===// 62 // ArrayAttr 63 //===----------------------------------------------------------------------===// 64 65 ArrayAttr ArrayAttr::get(ArrayRef<Attribute> value, MLIRContext *context) { 66 return Base::get(context, StandardAttributes::Array, value); 67 } 68 69 ArrayRef<Attribute> ArrayAttr::getValue() const { return getImpl()->value; } 70 71 Attribute ArrayAttr::operator[](unsigned idx) const { 72 assert(idx < size() && "index out of bounds"); 73 return getValue()[idx]; 74 } 75 76 //===----------------------------------------------------------------------===// 77 // BoolAttr 78 //===----------------------------------------------------------------------===// 79 80 bool BoolAttr::getValue() const { return getImpl()->value; } 81 82 //===----------------------------------------------------------------------===// 83 // DictionaryAttr 84 //===----------------------------------------------------------------------===// 85 86 /// Perform a three-way comparison between the names of the specified 87 /// NamedAttributes. 88 static int compareNamedAttributes(const NamedAttribute *lhs, 89 const NamedAttribute *rhs) { 90 return strcmp(lhs->first.data(), rhs->first.data()); 91 } 92 93 /// Returns if the name of the given attribute precedes that of 'name'. 94 static bool compareNamedAttributeWithName(const NamedAttribute &attr, 95 StringRef name) { 96 // This is correct even when attr.first.data()[name.size()] is not a zero 97 // string terminator, because we only care about a less than comparison. 98 // This can't use memcmp, because it doesn't guarantee that it will stop 99 // reading both buffers if one is shorter than the other, even if there is 100 // a difference. 101 return strncmp(attr.first.data(), name.data(), name.size()) < 0; 102 } 103 104 /// Helper function that does either an in place sort or sorts from source array 105 /// into destination. If inPlace then storage is both the source and the 106 /// destination, else value is the source and storage destination. Returns 107 /// whether source was sorted. 108 template <bool inPlace> 109 static bool dictionaryAttrSort(ArrayRef<NamedAttribute> value, 110 SmallVectorImpl<NamedAttribute> &storage) { 111 // Specialize for the common case. 112 switch (value.size()) { 113 case 0: 114 case 1: 115 // Zero or one elements are already sorted. 116 break; 117 case 2: 118 assert(value[0].first != value[1].first && 119 "DictionaryAttr element names must be unique"); 120 if (compareNamedAttributes(&value[0], &value[1]) > 0) { 121 if (inPlace) 122 std::swap(storage[0], storage[1]); 123 else 124 storage.append({value[1], value[0]}); 125 return true; 126 } 127 break; 128 default: 129 // Check to see they are sorted already. 130 bool isSorted = 131 llvm::is_sorted(value, [](NamedAttribute l, NamedAttribute r) { 132 return compareNamedAttributes(&l, &r) < 0; 133 }); 134 if (!isSorted) { 135 // If not, do a general sort. 136 if (!inPlace) 137 storage.append(value.begin(), value.end()); 138 llvm::array_pod_sort(storage.begin(), storage.end(), 139 compareNamedAttributes); 140 value = storage; 141 } 142 143 // Ensure that the attribute elements are unique. 144 assert(std::adjacent_find(value.begin(), value.end(), 145 [](NamedAttribute l, NamedAttribute r) { 146 return l.first == r.first; 147 }) == value.end() && 148 "DictionaryAttr element names must be unique"); 149 return !isSorted; 150 } 151 return false; 152 } 153 154 /// Sorts the NamedAttributes in the array ordered by name as expected by 155 /// getWithSorted. 156 /// Requires: uniquely named attributes. 157 void DictionaryAttr::sort(SmallVectorImpl<NamedAttribute> &array) { 158 dictionaryAttrSort</*inPlace=*/true>(array, array); 159 } 160 161 DictionaryAttr DictionaryAttr::get(ArrayRef<NamedAttribute> value, 162 MLIRContext *context) { 163 assert(llvm::all_of(value, 164 [](const NamedAttribute &attr) { return attr.second; }) && 165 "value cannot have null entries"); 166 167 // We need to sort the element list to canonicalize it. 168 SmallVector<NamedAttribute, 8> storage; 169 if (dictionaryAttrSort</*inPlace=*/false>(value, storage)) 170 value = storage; 171 172 return Base::get(context, StandardAttributes::Dictionary, value); 173 } 174 175 /// Construct a dictionary with an array of values that is known to already be 176 /// sorted by name and uniqued. 177 DictionaryAttr DictionaryAttr::getWithSorted(ArrayRef<NamedAttribute> value, 178 MLIRContext *context) { 179 // Ensure that the attribute elements are unique and sorted. 180 assert(llvm::is_sorted(value, 181 [](NamedAttribute l, NamedAttribute r) { 182 return l.first.strref() < r.first.strref(); 183 }) && 184 "expected attribute values to be sorted"); 185 assert(std::adjacent_find(value.begin(), value.end(), 186 [](NamedAttribute l, NamedAttribute r) { 187 return l.first == r.first; 188 }) == value.end() && 189 "DictionaryAttr element names must be unique"); 190 return Base::get(context, StandardAttributes::Dictionary, value); 191 } 192 193 ArrayRef<NamedAttribute> DictionaryAttr::getValue() const { 194 return getImpl()->getElements(); 195 } 196 197 /// Return the specified attribute if present, null otherwise. 198 Attribute DictionaryAttr::get(StringRef name) const { 199 Optional<NamedAttribute> attr = getNamed(name); 200 return attr ? attr->second : nullptr; 201 } 202 Attribute DictionaryAttr::get(Identifier name) const { 203 Optional<NamedAttribute> attr = getNamed(name); 204 return attr ? attr->second : nullptr; 205 } 206 207 /// Return the specified named attribute if present, None otherwise. 208 Optional<NamedAttribute> DictionaryAttr::getNamed(StringRef name) const { 209 ArrayRef<NamedAttribute> values = getValue(); 210 auto it = llvm::lower_bound(values, name, compareNamedAttributeWithName); 211 return it != values.end() && it->first == name ? *it 212 : Optional<NamedAttribute>(); 213 } 214 Optional<NamedAttribute> DictionaryAttr::getNamed(Identifier name) const { 215 for (auto elt : getValue()) 216 if (elt.first == name) 217 return elt; 218 return llvm::None; 219 } 220 221 DictionaryAttr::iterator DictionaryAttr::begin() const { 222 return getValue().begin(); 223 } 224 DictionaryAttr::iterator DictionaryAttr::end() const { 225 return getValue().end(); 226 } 227 size_t DictionaryAttr::size() const { return getValue().size(); } 228 229 //===----------------------------------------------------------------------===// 230 // FloatAttr 231 //===----------------------------------------------------------------------===// 232 233 FloatAttr FloatAttr::get(Type type, double value) { 234 return Base::get(type.getContext(), StandardAttributes::Float, type, value); 235 } 236 237 FloatAttr FloatAttr::getChecked(Type type, double value, Location loc) { 238 return Base::getChecked(loc, StandardAttributes::Float, type, value); 239 } 240 241 FloatAttr FloatAttr::get(Type type, const APFloat &value) { 242 return Base::get(type.getContext(), StandardAttributes::Float, type, value); 243 } 244 245 FloatAttr FloatAttr::getChecked(Type type, const APFloat &value, Location loc) { 246 return Base::getChecked(loc, StandardAttributes::Float, type, value); 247 } 248 249 APFloat FloatAttr::getValue() const { return getImpl()->getValue(); } 250 251 double FloatAttr::getValueAsDouble() const { 252 return getValueAsDouble(getValue()); 253 } 254 double FloatAttr::getValueAsDouble(APFloat value) { 255 if (&value.getSemantics() != &APFloat::IEEEdouble()) { 256 bool losesInfo = false; 257 value.convert(APFloat::IEEEdouble(), APFloat::rmNearestTiesToEven, 258 &losesInfo); 259 } 260 return value.convertToDouble(); 261 } 262 263 /// Verify construction invariants. 264 static LogicalResult verifyFloatTypeInvariants(Location loc, Type type) { 265 if (!type.isa<FloatType>()) 266 return emitError(loc, "expected floating point type"); 267 return success(); 268 } 269 270 LogicalResult FloatAttr::verifyConstructionInvariants(Location loc, Type type, 271 double value) { 272 return verifyFloatTypeInvariants(loc, type); 273 } 274 275 LogicalResult FloatAttr::verifyConstructionInvariants(Location loc, Type type, 276 const APFloat &value) { 277 // Verify that the type is correct. 278 if (failed(verifyFloatTypeInvariants(loc, type))) 279 return failure(); 280 281 // Verify that the type semantics match that of the value. 282 if (&type.cast<FloatType>().getFloatSemantics() != &value.getSemantics()) { 283 return emitError( 284 loc, "FloatAttr type doesn't match the type implied by its value"); 285 } 286 return success(); 287 } 288 289 //===----------------------------------------------------------------------===// 290 // SymbolRefAttr 291 //===----------------------------------------------------------------------===// 292 293 FlatSymbolRefAttr SymbolRefAttr::get(StringRef value, MLIRContext *ctx) { 294 return Base::get(ctx, StandardAttributes::SymbolRef, value, llvm::None) 295 .cast<FlatSymbolRefAttr>(); 296 } 297 298 SymbolRefAttr SymbolRefAttr::get(StringRef value, 299 ArrayRef<FlatSymbolRefAttr> nestedReferences, 300 MLIRContext *ctx) { 301 return Base::get(ctx, StandardAttributes::SymbolRef, value, nestedReferences); 302 } 303 304 StringRef SymbolRefAttr::getRootReference() const { return getImpl()->value; } 305 306 StringRef SymbolRefAttr::getLeafReference() const { 307 ArrayRef<FlatSymbolRefAttr> nestedRefs = getNestedReferences(); 308 return nestedRefs.empty() ? getRootReference() : nestedRefs.back().getValue(); 309 } 310 311 ArrayRef<FlatSymbolRefAttr> SymbolRefAttr::getNestedReferences() const { 312 return getImpl()->getNestedRefs(); 313 } 314 315 //===----------------------------------------------------------------------===// 316 // IntegerAttr 317 //===----------------------------------------------------------------------===// 318 319 IntegerAttr IntegerAttr::get(Type type, const APInt &value) { 320 return Base::get(type.getContext(), StandardAttributes::Integer, type, value); 321 } 322 323 IntegerAttr IntegerAttr::get(Type type, int64_t value) { 324 // This uses 64 bit APInts by default for index type. 325 if (type.isIndex()) 326 return get(type, APInt(IndexType::kInternalStorageBitWidth, value)); 327 328 auto intType = type.cast<IntegerType>(); 329 return get(type, APInt(intType.getWidth(), value, intType.isSignedInteger())); 330 } 331 332 APInt IntegerAttr::getValue() const { return getImpl()->getValue(); } 333 334 int64_t IntegerAttr::getInt() const { 335 assert((getImpl()->getType().isIndex() || 336 getImpl()->getType().isSignlessInteger()) && 337 "must be signless integer"); 338 return getValue().getSExtValue(); 339 } 340 341 int64_t IntegerAttr::getSInt() const { 342 assert(getImpl()->getType().isSignedInteger() && "must be signed integer"); 343 return getValue().getSExtValue(); 344 } 345 346 uint64_t IntegerAttr::getUInt() const { 347 assert(getImpl()->getType().isUnsignedInteger() && 348 "must be unsigned integer"); 349 return getValue().getZExtValue(); 350 } 351 352 static LogicalResult verifyIntegerTypeInvariants(Location loc, Type type) { 353 if (type.isa<IntegerType>() || type.isa<IndexType>()) 354 return success(); 355 return emitError(loc, "expected integer or index type"); 356 } 357 358 LogicalResult IntegerAttr::verifyConstructionInvariants(Location loc, Type type, 359 int64_t value) { 360 return verifyIntegerTypeInvariants(loc, type); 361 } 362 363 LogicalResult IntegerAttr::verifyConstructionInvariants(Location loc, Type type, 364 const APInt &value) { 365 if (failed(verifyIntegerTypeInvariants(loc, type))) 366 return failure(); 367 if (auto integerType = type.dyn_cast<IntegerType>()) 368 if (integerType.getWidth() != value.getBitWidth()) 369 return emitError(loc, "integer type bit width (") 370 << integerType.getWidth() << ") doesn't match value bit width (" 371 << value.getBitWidth() << ")"; 372 return success(); 373 } 374 375 //===----------------------------------------------------------------------===// 376 // IntegerSetAttr 377 //===----------------------------------------------------------------------===// 378 379 IntegerSetAttr IntegerSetAttr::get(IntegerSet value) { 380 return Base::get(value.getConstraint(0).getContext(), 381 StandardAttributes::IntegerSet, value); 382 } 383 384 IntegerSet IntegerSetAttr::getValue() const { return getImpl()->value; } 385 386 //===----------------------------------------------------------------------===// 387 // OpaqueAttr 388 //===----------------------------------------------------------------------===// 389 390 OpaqueAttr OpaqueAttr::get(Identifier dialect, StringRef attrData, Type type, 391 MLIRContext *context) { 392 return Base::get(context, StandardAttributes::Opaque, dialect, attrData, 393 type); 394 } 395 396 OpaqueAttr OpaqueAttr::getChecked(Identifier dialect, StringRef attrData, 397 Type type, Location location) { 398 return Base::getChecked(location, StandardAttributes::Opaque, dialect, 399 attrData, type); 400 } 401 402 /// Returns the dialect namespace of the opaque attribute. 403 Identifier OpaqueAttr::getDialectNamespace() const { 404 return getImpl()->dialectNamespace; 405 } 406 407 /// Returns the raw attribute data of the opaque attribute. 408 StringRef OpaqueAttr::getAttrData() const { return getImpl()->attrData; } 409 410 /// Verify the construction of an opaque attribute. 411 LogicalResult OpaqueAttr::verifyConstructionInvariants(Location loc, 412 Identifier dialect, 413 StringRef attrData, 414 Type type) { 415 if (!Dialect::isValidNamespace(dialect.strref())) 416 return emitError(loc, "invalid dialect namespace '") << dialect << "'"; 417 return success(); 418 } 419 420 //===----------------------------------------------------------------------===// 421 // StringAttr 422 //===----------------------------------------------------------------------===// 423 424 StringAttr StringAttr::get(StringRef bytes, MLIRContext *context) { 425 return get(bytes, NoneType::get(context)); 426 } 427 428 /// Get an instance of a StringAttr with the given string and Type. 429 StringAttr StringAttr::get(StringRef bytes, Type type) { 430 return Base::get(type.getContext(), StandardAttributes::String, bytes, type); 431 } 432 433 StringRef StringAttr::getValue() const { return getImpl()->value; } 434 435 //===----------------------------------------------------------------------===// 436 // TypeAttr 437 //===----------------------------------------------------------------------===// 438 439 TypeAttr TypeAttr::get(Type value) { 440 return Base::get(value.getContext(), StandardAttributes::Type, value); 441 } 442 443 Type TypeAttr::getValue() const { return getImpl()->value; } 444 445 //===----------------------------------------------------------------------===// 446 // ElementsAttr 447 //===----------------------------------------------------------------------===// 448 449 ShapedType ElementsAttr::getType() const { 450 return Attribute::getType().cast<ShapedType>(); 451 } 452 453 /// Returns the number of elements held by this attribute. 454 int64_t ElementsAttr::getNumElements() const { 455 return getType().getNumElements(); 456 } 457 458 /// Return the value at the given index. If index does not refer to a valid 459 /// element, then a null attribute is returned. 460 Attribute ElementsAttr::getValue(ArrayRef<uint64_t> index) const { 461 switch (getKind()) { 462 case StandardAttributes::DenseIntOrFPElements: 463 return cast<DenseElementsAttr>().getValue(index); 464 case StandardAttributes::OpaqueElements: 465 return cast<OpaqueElementsAttr>().getValue(index); 466 case StandardAttributes::SparseElements: 467 return cast<SparseElementsAttr>().getValue(index); 468 default: 469 llvm_unreachable("unknown ElementsAttr kind"); 470 } 471 } 472 473 /// Return if the given 'index' refers to a valid element in this attribute. 474 bool ElementsAttr::isValidIndex(ArrayRef<uint64_t> index) const { 475 auto type = getType(); 476 477 // Verify that the rank of the indices matches the held type. 478 auto rank = type.getRank(); 479 if (rank != static_cast<int64_t>(index.size())) 480 return false; 481 482 // Verify that all of the indices are within the shape dimensions. 483 auto shape = type.getShape(); 484 return llvm::all_of(llvm::seq<int>(0, rank), [&](int i) { 485 return static_cast<int64_t>(index[i]) < shape[i]; 486 }); 487 } 488 489 ElementsAttr 490 ElementsAttr::mapValues(Type newElementType, 491 function_ref<APInt(const APInt &)> mapping) const { 492 switch (getKind()) { 493 case StandardAttributes::DenseIntOrFPElements: 494 return cast<DenseElementsAttr>().mapValues(newElementType, mapping); 495 default: 496 llvm_unreachable("unsupported ElementsAttr subtype"); 497 } 498 } 499 500 ElementsAttr 501 ElementsAttr::mapValues(Type newElementType, 502 function_ref<APInt(const APFloat &)> mapping) const { 503 switch (getKind()) { 504 case StandardAttributes::DenseIntOrFPElements: 505 return cast<DenseElementsAttr>().mapValues(newElementType, mapping); 506 default: 507 llvm_unreachable("unsupported ElementsAttr subtype"); 508 } 509 } 510 511 /// Returns the 1 dimensional flattened row-major index from the given 512 /// multi-dimensional index. 513 uint64_t ElementsAttr::getFlattenedIndex(ArrayRef<uint64_t> index) const { 514 assert(isValidIndex(index) && "expected valid multi-dimensional index"); 515 auto type = getType(); 516 517 // Reduce the provided multidimensional index into a flattended 1D row-major 518 // index. 519 auto rank = type.getRank(); 520 auto shape = type.getShape(); 521 uint64_t valueIndex = 0; 522 uint64_t dimMultiplier = 1; 523 for (int i = rank - 1; i >= 0; --i) { 524 valueIndex += index[i] * dimMultiplier; 525 dimMultiplier *= shape[i]; 526 } 527 return valueIndex; 528 } 529 530 //===----------------------------------------------------------------------===// 531 // DenseElementAttr Utilities 532 //===----------------------------------------------------------------------===// 533 534 /// Get the bitwidth of a dense element type within the buffer. 535 /// DenseElementsAttr requires bitwidths greater than 1 to be aligned by 8. 536 static size_t getDenseElementStorageWidth(size_t origWidth) { 537 return origWidth == 1 ? origWidth : llvm::alignTo<8>(origWidth); 538 } 539 540 /// Set a bit to a specific value. 541 static void setBit(char *rawData, size_t bitPos, bool value) { 542 if (value) 543 rawData[bitPos / CHAR_BIT] |= (1 << (bitPos % CHAR_BIT)); 544 else 545 rawData[bitPos / CHAR_BIT] &= ~(1 << (bitPos % CHAR_BIT)); 546 } 547 548 /// Return the value of the specified bit. 549 static bool getBit(const char *rawData, size_t bitPos) { 550 return (rawData[bitPos / CHAR_BIT] & (1 << (bitPos % CHAR_BIT))) != 0; 551 } 552 553 /// Writes value to the bit position `bitPos` in array `rawData`. 554 static void writeBits(char *rawData, size_t bitPos, APInt value) { 555 size_t bitWidth = value.getBitWidth(); 556 557 // If the bitwidth is 1 we just toggle the specific bit. 558 if (bitWidth == 1) 559 return setBit(rawData, bitPos, value.isOneValue()); 560 561 // Otherwise, the bit position is guaranteed to be byte aligned. 562 assert((bitPos % CHAR_BIT) == 0 && "expected bitPos to be 8-bit aligned"); 563 std::copy_n(reinterpret_cast<const char *>(value.getRawData()), 564 llvm::divideCeil(bitWidth, CHAR_BIT), 565 rawData + (bitPos / CHAR_BIT)); 566 } 567 568 /// Reads the next `bitWidth` bits from the bit position `bitPos` in array 569 /// `rawData`. 570 static APInt readBits(const char *rawData, size_t bitPos, size_t bitWidth) { 571 // Handle a boolean bit position. 572 if (bitWidth == 1) 573 return APInt(1, getBit(rawData, bitPos) ? 1 : 0); 574 575 // Otherwise, the bit position must be 8-bit aligned. 576 assert((bitPos % CHAR_BIT) == 0 && "expected bitPos to be 8-bit aligned"); 577 APInt result(bitWidth, 0); 578 std::copy_n( 579 rawData + (bitPos / CHAR_BIT), llvm::divideCeil(bitWidth, CHAR_BIT), 580 const_cast<char *>(reinterpret_cast<const char *>(result.getRawData()))); 581 return result; 582 } 583 584 /// Returns if 'values' corresponds to a splat, i.e. one element, or has the 585 /// same element count as 'type'. 586 template <typename Values> 587 static bool hasSameElementsOrSplat(ShapedType type, const Values &values) { 588 return (values.size() == 1) || 589 (type.getNumElements() == static_cast<int64_t>(values.size())); 590 } 591 592 //===----------------------------------------------------------------------===// 593 // DenseElementAttr Iterators 594 //===----------------------------------------------------------------------===// 595 596 /// Constructs a new iterator. 597 DenseElementsAttr::AttributeElementIterator::AttributeElementIterator( 598 DenseElementsAttr attr, size_t index) 599 : llvm::indexed_accessor_iterator<AttributeElementIterator, const void *, 600 Attribute, Attribute, Attribute>( 601 attr.getAsOpaquePointer(), index) {} 602 603 /// Accesses the Attribute value at this iterator position. 604 Attribute DenseElementsAttr::AttributeElementIterator::operator*() const { 605 auto owner = getFromOpaquePointer(base).cast<DenseElementsAttr>(); 606 Type eltTy = owner.getType().getElementType(); 607 if (auto intEltTy = eltTy.dyn_cast<IntegerType>()) { 608 if (intEltTy.getWidth() == 1) 609 return BoolAttr::get((*IntElementIterator(owner, index)).isOneValue(), 610 owner.getContext()); 611 return IntegerAttr::get(eltTy, *IntElementIterator(owner, index)); 612 } 613 if (auto floatEltTy = eltTy.dyn_cast<FloatType>()) { 614 IntElementIterator intIt(owner, index); 615 FloatElementIterator floatIt(floatEltTy.getFloatSemantics(), intIt); 616 return FloatAttr::get(eltTy, *floatIt); 617 } 618 if (owner.isa<DenseStringElementsAttr>()) 619 return StringAttr::get(owner.getRawStringData()[index], eltTy); 620 llvm_unreachable("unexpected element type"); 621 } 622 623 /// Constructs a new iterator. 624 DenseElementsAttr::BoolElementIterator::BoolElementIterator( 625 DenseElementsAttr attr, size_t dataIndex) 626 : DenseElementIndexedIteratorImpl<BoolElementIterator, bool, bool, bool>( 627 attr.getRawData().data(), attr.isSplat(), dataIndex) {} 628 629 /// Accesses the bool value at this iterator position. 630 bool DenseElementsAttr::BoolElementIterator::operator*() const { 631 return getBit(getData(), getDataIndex()); 632 } 633 634 /// Constructs a new iterator. 635 DenseElementsAttr::IntElementIterator::IntElementIterator( 636 DenseElementsAttr attr, size_t dataIndex) 637 : DenseElementIndexedIteratorImpl<IntElementIterator, APInt, APInt, APInt>( 638 attr.getRawData().data(), attr.isSplat(), dataIndex), 639 bitWidth(getDenseElementBitWidth(attr.getType().getElementType())) {} 640 641 /// Accesses the raw APInt value at this iterator position. 642 APInt DenseElementsAttr::IntElementIterator::operator*() const { 643 return readBits(getData(), 644 getDataIndex() * getDenseElementStorageWidth(bitWidth), 645 bitWidth); 646 } 647 648 DenseElementsAttr::FloatElementIterator::FloatElementIterator( 649 const llvm::fltSemantics &smt, IntElementIterator it) 650 : llvm::mapped_iterator<IntElementIterator, 651 std::function<APFloat(const APInt &)>>( 652 it, [&](const APInt &val) { return APFloat(smt, val); }) {} 653 654 //===----------------------------------------------------------------------===// 655 // DenseElementsAttr 656 //===----------------------------------------------------------------------===// 657 658 DenseElementsAttr DenseElementsAttr::get(ShapedType type, 659 ArrayRef<Attribute> values) { 660 assert(hasSameElementsOrSplat(type, values)); 661 662 // If the element type is not based on int/float/index, assume it is a string 663 // type. 664 auto eltType = type.getElementType(); 665 if (!type.getElementType().isIntOrIndexOrFloat()) { 666 SmallVector<StringRef, 8> stringValues; 667 stringValues.reserve(values.size()); 668 for (Attribute attr : values) { 669 assert(attr.isa<StringAttr>() && 670 "expected string value for non integer/index/float element"); 671 stringValues.push_back(attr.cast<StringAttr>().getValue()); 672 } 673 return get(type, stringValues); 674 } 675 676 // Otherwise, get the raw storage width to use for the allocation. 677 size_t bitWidth = getDenseElementBitWidth(eltType); 678 size_t storageBitWidth = getDenseElementStorageWidth(bitWidth); 679 680 // Compress the attribute values into a character buffer. 681 SmallVector<char, 8> data(llvm::divideCeil(storageBitWidth, CHAR_BIT) * 682 values.size()); 683 APInt intVal; 684 for (unsigned i = 0, e = values.size(); i < e; ++i) { 685 assert(eltType == values[i].getType() && 686 "expected attribute value to have element type"); 687 688 switch (eltType.getKind()) { 689 case StandardTypes::BF16: 690 case StandardTypes::F16: 691 case StandardTypes::F32: 692 case StandardTypes::F64: 693 intVal = values[i].cast<FloatAttr>().getValue().bitcastToAPInt(); 694 break; 695 case StandardTypes::Integer: 696 case StandardTypes::Index: 697 intVal = values[i].isa<BoolAttr>() 698 ? APInt(1, values[i].cast<BoolAttr>().getValue() ? 1 : 0) 699 : values[i].cast<IntegerAttr>().getValue(); 700 break; 701 default: 702 llvm_unreachable("unexpected element type"); 703 } 704 assert(intVal.getBitWidth() == bitWidth && 705 "expected value to have same bitwidth as element type"); 706 writeBits(data.data(), i * storageBitWidth, intVal); 707 } 708 return DenseIntOrFPElementsAttr::getRaw(type, data, 709 /*isSplat=*/(values.size() == 1)); 710 } 711 712 DenseElementsAttr DenseElementsAttr::get(ShapedType type, 713 ArrayRef<bool> values) { 714 assert(hasSameElementsOrSplat(type, values)); 715 assert(type.getElementType().isInteger(1)); 716 717 std::vector<char> buff(llvm::divideCeil(values.size(), CHAR_BIT)); 718 for (int i = 0, e = values.size(); i != e; ++i) 719 setBit(buff.data(), i, values[i]); 720 return DenseIntOrFPElementsAttr::getRaw(type, buff, 721 /*isSplat=*/(values.size() == 1)); 722 } 723 724 DenseElementsAttr DenseElementsAttr::get(ShapedType type, 725 ArrayRef<StringRef> values) { 726 assert(!type.getElementType().isIntOrFloat()); 727 return DenseStringElementsAttr::get(type, values); 728 } 729 730 /// Constructs a dense integer elements attribute from an array of APInt 731 /// values. Each APInt value is expected to have the same bitwidth as the 732 /// element type of 'type'. 733 DenseElementsAttr DenseElementsAttr::get(ShapedType type, 734 ArrayRef<APInt> values) { 735 assert(type.getElementType().isIntOrIndex()); 736 return DenseIntOrFPElementsAttr::getRaw(type, values); 737 } 738 739 // Constructs a dense float elements attribute from an array of APFloat 740 // values. Each APFloat value is expected to have the same bitwidth as the 741 // element type of 'type'. 742 DenseElementsAttr DenseElementsAttr::get(ShapedType type, 743 ArrayRef<APFloat> values) { 744 assert(type.getElementType().isa<FloatType>()); 745 746 // Convert the APFloat values to APInt and create a dense elements attribute. 747 std::vector<APInt> intValues(values.size()); 748 for (unsigned i = 0, e = values.size(); i != e; ++i) 749 intValues[i] = values[i].bitcastToAPInt(); 750 return DenseIntOrFPElementsAttr::getRaw(type, intValues); 751 } 752 753 /// Construct a dense elements attribute from a raw buffer representing the 754 /// data for this attribute. Users should generally not use this methods as 755 /// the expected buffer format may not be a form the user expects. 756 DenseElementsAttr DenseElementsAttr::getFromRawBuffer(ShapedType type, 757 ArrayRef<char> rawBuffer, 758 bool isSplatBuffer) { 759 return DenseIntOrFPElementsAttr::getRaw(type, rawBuffer, isSplatBuffer); 760 } 761 762 /// Check the information for a C++ data type, check if this type is valid for 763 /// the current attribute. This method is used to verify specific type 764 /// invariants that the templatized 'getValues' method cannot. 765 static bool isValidIntOrFloat(ShapedType type, int64_t dataEltSize, bool isInt, 766 bool isSigned) { 767 // Make sure that the data element size is the same as the type element width. 768 if (getDenseElementBitWidth(type.getElementType()) != 769 static_cast<size_t>(dataEltSize * CHAR_BIT)) 770 return false; 771 772 // Check that the element type is either float or integer or index. 773 if (!isInt) 774 return type.getElementType().isa<FloatType>(); 775 776 if (type.getElementType().isIndex()) 777 return true; 778 779 auto intType = type.getElementType().dyn_cast<IntegerType>(); 780 if (!intType) 781 return false; 782 783 // Make sure signedness semantics is consistent. 784 if (intType.isSignless()) 785 return true; 786 return intType.isSigned() ? isSigned : !isSigned; 787 } 788 789 /// Defaults down the subclass implementation. 790 DenseElementsAttr DenseElementsAttr::getRawIntOrFloat(ShapedType type, 791 ArrayRef<char> data, 792 int64_t dataEltSize, 793 bool isInt, 794 bool isSigned) { 795 return DenseIntOrFPElementsAttr::getRawIntOrFloat(type, data, dataEltSize, 796 isInt, isSigned); 797 } 798 799 /// A method used to verify specific type invariants that the templatized 'get' 800 /// method cannot. 801 bool DenseElementsAttr::isValidIntOrFloat(int64_t dataEltSize, bool isInt, 802 bool isSigned) const { 803 return ::isValidIntOrFloat(getType(), dataEltSize, isInt, isSigned); 804 } 805 806 /// Returns if this attribute corresponds to a splat, i.e. if all element 807 /// values are the same. 808 bool DenseElementsAttr::isSplat() const { 809 return static_cast<DenseElementsAttributeStorage *>(impl)->isSplat; 810 } 811 812 /// Return the held element values as a range of Attributes. 813 auto DenseElementsAttr::getAttributeValues() const 814 -> llvm::iterator_range<AttributeElementIterator> { 815 return {attr_value_begin(), attr_value_end()}; 816 } 817 auto DenseElementsAttr::attr_value_begin() const -> AttributeElementIterator { 818 return AttributeElementIterator(*this, 0); 819 } 820 auto DenseElementsAttr::attr_value_end() const -> AttributeElementIterator { 821 return AttributeElementIterator(*this, getNumElements()); 822 } 823 824 /// Return the held element values as a range of bool. The element type of 825 /// this attribute must be of integer type of bitwidth 1. 826 auto DenseElementsAttr::getBoolValues() const 827 -> llvm::iterator_range<BoolElementIterator> { 828 auto eltType = getType().getElementType().dyn_cast<IntegerType>(); 829 assert(eltType && eltType.getWidth() == 1 && "expected i1 integer type"); 830 (void)eltType; 831 return {BoolElementIterator(*this, 0), 832 BoolElementIterator(*this, getNumElements())}; 833 } 834 835 /// Return the held element values as a range of APInts. The element type of 836 /// this attribute must be of integer type. 837 auto DenseElementsAttr::getIntValues() const 838 -> llvm::iterator_range<IntElementIterator> { 839 assert(getType().getElementType().isIntOrIndex() && "expected integral type"); 840 return {raw_int_begin(), raw_int_end()}; 841 } 842 auto DenseElementsAttr::int_value_begin() const -> IntElementIterator { 843 assert(getType().getElementType().isIntOrIndex() && "expected integral type"); 844 return raw_int_begin(); 845 } 846 auto DenseElementsAttr::int_value_end() const -> IntElementIterator { 847 assert(getType().getElementType().isIntOrIndex() && "expected integral type"); 848 return raw_int_end(); 849 } 850 851 /// Return the held element values as a range of APFloat. The element type of 852 /// this attribute must be of float type. 853 auto DenseElementsAttr::getFloatValues() const 854 -> llvm::iterator_range<FloatElementIterator> { 855 auto elementType = getType().getElementType().cast<FloatType>(); 856 assert(elementType.isa<FloatType>() && "expected float type"); 857 const auto &elementSemantics = elementType.getFloatSemantics(); 858 return {FloatElementIterator(elementSemantics, raw_int_begin()), 859 FloatElementIterator(elementSemantics, raw_int_end())}; 860 } 861 auto DenseElementsAttr::float_value_begin() const -> FloatElementIterator { 862 return getFloatValues().begin(); 863 } 864 auto DenseElementsAttr::float_value_end() const -> FloatElementIterator { 865 return getFloatValues().end(); 866 } 867 868 /// Return the raw storage data held by this attribute. 869 ArrayRef<char> DenseElementsAttr::getRawData() const { 870 return static_cast<DenseIntOrFPElementsAttributeStorage *>(impl)->data; 871 } 872 873 ArrayRef<StringRef> DenseElementsAttr::getRawStringData() const { 874 return static_cast<DenseStringElementsAttributeStorage *>(impl)->data; 875 } 876 877 /// Return a new DenseElementsAttr that has the same data as the current 878 /// attribute, but has been reshaped to 'newType'. The new type must have the 879 /// same total number of elements as well as element type. 880 DenseElementsAttr DenseElementsAttr::reshape(ShapedType newType) { 881 ShapedType curType = getType(); 882 if (curType == newType) 883 return *this; 884 885 (void)curType; 886 assert(newType.getElementType() == curType.getElementType() && 887 "expected the same element type"); 888 assert(newType.getNumElements() == curType.getNumElements() && 889 "expected the same number of elements"); 890 return DenseIntOrFPElementsAttr::getRaw(newType, getRawData(), isSplat()); 891 } 892 893 DenseElementsAttr 894 DenseElementsAttr::mapValues(Type newElementType, 895 function_ref<APInt(const APInt &)> mapping) const { 896 return cast<DenseIntElementsAttr>().mapValues(newElementType, mapping); 897 } 898 899 DenseElementsAttr DenseElementsAttr::mapValues( 900 Type newElementType, function_ref<APInt(const APFloat &)> mapping) const { 901 return cast<DenseFPElementsAttr>().mapValues(newElementType, mapping); 902 } 903 904 //===----------------------------------------------------------------------===// 905 // DenseStringElementsAttr 906 //===----------------------------------------------------------------------===// 907 908 DenseStringElementsAttr 909 DenseStringElementsAttr::get(ShapedType type, ArrayRef<StringRef> values) { 910 return Base::get(type.getContext(), StandardAttributes::DenseStringElements, 911 type, values, (values.size() == 1)); 912 } 913 914 //===----------------------------------------------------------------------===// 915 // DenseIntOrFPElementsAttr 916 //===----------------------------------------------------------------------===// 917 918 /// Constructs a dense elements attribute from an array of raw APInt values. 919 /// Each APInt value is expected to have the same bitwidth as the element type 920 /// of 'type'. 921 DenseElementsAttr DenseIntOrFPElementsAttr::getRaw(ShapedType type, 922 ArrayRef<APInt> values) { 923 assert(hasSameElementsOrSplat(type, values)); 924 925 size_t bitWidth = getDenseElementBitWidth(type.getElementType()); 926 size_t storageBitWidth = getDenseElementStorageWidth(bitWidth); 927 std::vector<char> elementData(llvm::divideCeil(storageBitWidth, CHAR_BIT) * 928 values.size()); 929 for (unsigned i = 0, e = values.size(); i != e; ++i) { 930 assert(values[i].getBitWidth() == bitWidth); 931 writeBits(elementData.data(), i * storageBitWidth, values[i]); 932 } 933 return DenseIntOrFPElementsAttr::getRaw(type, elementData, 934 /*isSplat=*/(values.size() == 1)); 935 } 936 937 DenseElementsAttr DenseIntOrFPElementsAttr::getRaw(ShapedType type, 938 ArrayRef<char> data, 939 bool isSplat) { 940 assert((type.isa<RankedTensorType>() || type.isa<VectorType>()) && 941 "type must be ranked tensor or vector"); 942 assert(type.hasStaticShape() && "type must have static shape"); 943 return Base::get(type.getContext(), StandardAttributes::DenseIntOrFPElements, 944 type, data, isSplat); 945 } 946 947 /// Overload of the 'getRaw' method that asserts that the given type is of 948 /// integer type. This method is used to verify type invariants that the 949 /// templatized 'get' method cannot. 950 DenseElementsAttr 951 DenseIntOrFPElementsAttr::getRawIntOrFloat(ShapedType type, ArrayRef<char> data, 952 int64_t dataEltSize, bool isInt, 953 bool isSigned) { 954 assert(::isValidIntOrFloat(type, dataEltSize, isInt, isSigned)); 955 956 int64_t numElements = data.size() / dataEltSize; 957 assert(numElements == 1 || numElements == type.getNumElements()); 958 return getRaw(type, data, /*isSplat=*/numElements == 1); 959 } 960 961 //===----------------------------------------------------------------------===// 962 // DenseFPElementsAttr 963 //===----------------------------------------------------------------------===// 964 965 template <typename Fn, typename Attr> 966 static ShapedType mappingHelper(Fn mapping, Attr &attr, ShapedType inType, 967 Type newElementType, 968 llvm::SmallVectorImpl<char> &data) { 969 size_t bitWidth = getDenseElementBitWidth(newElementType); 970 size_t storageBitWidth = getDenseElementStorageWidth(bitWidth); 971 972 ShapedType newArrayType; 973 if (inType.isa<RankedTensorType>()) 974 newArrayType = RankedTensorType::get(inType.getShape(), newElementType); 975 else if (inType.isa<UnrankedTensorType>()) 976 newArrayType = RankedTensorType::get(inType.getShape(), newElementType); 977 else if (inType.isa<VectorType>()) 978 newArrayType = VectorType::get(inType.getShape(), newElementType); 979 else 980 assert(newArrayType && "Unhandled tensor type"); 981 982 size_t numRawElements = attr.isSplat() ? 1 : newArrayType.getNumElements(); 983 data.resize(llvm::divideCeil(storageBitWidth, CHAR_BIT) * numRawElements); 984 985 // Functor used to process a single element value of the attribute. 986 auto processElt = [&](decltype(*attr.begin()) value, size_t index) { 987 auto newInt = mapping(value); 988 assert(newInt.getBitWidth() == bitWidth); 989 writeBits(data.data(), index * storageBitWidth, newInt); 990 }; 991 992 // Check for the splat case. 993 if (attr.isSplat()) { 994 processElt(*attr.begin(), /*index=*/0); 995 return newArrayType; 996 } 997 998 // Otherwise, process all of the element values. 999 uint64_t elementIdx = 0; 1000 for (auto value : attr) 1001 processElt(value, elementIdx++); 1002 return newArrayType; 1003 } 1004 1005 DenseElementsAttr DenseFPElementsAttr::mapValues( 1006 Type newElementType, function_ref<APInt(const APFloat &)> mapping) const { 1007 llvm::SmallVector<char, 8> elementData; 1008 auto newArrayType = 1009 mappingHelper(mapping, *this, getType(), newElementType, elementData); 1010 1011 return getRaw(newArrayType, elementData, isSplat()); 1012 } 1013 1014 /// Method for supporting type inquiry through isa, cast and dyn_cast. 1015 bool DenseFPElementsAttr::classof(Attribute attr) { 1016 return attr.isa<DenseElementsAttr>() && 1017 attr.getType().cast<ShapedType>().getElementType().isa<FloatType>(); 1018 } 1019 1020 //===----------------------------------------------------------------------===// 1021 // DenseIntElementsAttr 1022 //===----------------------------------------------------------------------===// 1023 1024 DenseElementsAttr DenseIntElementsAttr::mapValues( 1025 Type newElementType, function_ref<APInt(const APInt &)> mapping) const { 1026 llvm::SmallVector<char, 8> elementData; 1027 auto newArrayType = 1028 mappingHelper(mapping, *this, getType(), newElementType, elementData); 1029 1030 return getRaw(newArrayType, elementData, isSplat()); 1031 } 1032 1033 /// Method for supporting type inquiry through isa, cast and dyn_cast. 1034 bool DenseIntElementsAttr::classof(Attribute attr) { 1035 return attr.isa<DenseElementsAttr>() && 1036 attr.getType().cast<ShapedType>().getElementType().isIntOrIndex(); 1037 } 1038 1039 //===----------------------------------------------------------------------===// 1040 // OpaqueElementsAttr 1041 //===----------------------------------------------------------------------===// 1042 1043 OpaqueElementsAttr OpaqueElementsAttr::get(Dialect *dialect, ShapedType type, 1044 StringRef bytes) { 1045 assert(TensorType::isValidElementType(type.getElementType()) && 1046 "Input element type should be a valid tensor element type"); 1047 return Base::get(type.getContext(), StandardAttributes::OpaqueElements, type, 1048 dialect, bytes); 1049 } 1050 1051 StringRef OpaqueElementsAttr::getValue() const { return getImpl()->bytes; } 1052 1053 /// Return the value at the given index. If index does not refer to a valid 1054 /// element, then a null attribute is returned. 1055 Attribute OpaqueElementsAttr::getValue(ArrayRef<uint64_t> index) const { 1056 assert(isValidIndex(index) && "expected valid multi-dimensional index"); 1057 if (Dialect *dialect = getDialect()) 1058 return dialect->extractElementHook(*this, index); 1059 return Attribute(); 1060 } 1061 1062 Dialect *OpaqueElementsAttr::getDialect() const { return getImpl()->dialect; } 1063 1064 bool OpaqueElementsAttr::decode(ElementsAttr &result) { 1065 if (auto *d = getDialect()) 1066 return d->decodeHook(*this, result); 1067 return true; 1068 } 1069 1070 //===----------------------------------------------------------------------===// 1071 // SparseElementsAttr 1072 //===----------------------------------------------------------------------===// 1073 1074 SparseElementsAttr SparseElementsAttr::get(ShapedType type, 1075 DenseElementsAttr indices, 1076 DenseElementsAttr values) { 1077 assert(indices.getType().getElementType().isInteger(64) && 1078 "expected sparse indices to be 64-bit integer values"); 1079 assert((type.isa<RankedTensorType>() || type.isa<VectorType>()) && 1080 "type must be ranked tensor or vector"); 1081 assert(type.hasStaticShape() && "type must have static shape"); 1082 return Base::get(type.getContext(), StandardAttributes::SparseElements, type, 1083 indices.cast<DenseIntElementsAttr>(), values); 1084 } 1085 1086 DenseIntElementsAttr SparseElementsAttr::getIndices() const { 1087 return getImpl()->indices; 1088 } 1089 1090 DenseElementsAttr SparseElementsAttr::getValues() const { 1091 return getImpl()->values; 1092 } 1093 1094 /// Return the value of the element at the given index. 1095 Attribute SparseElementsAttr::getValue(ArrayRef<uint64_t> index) const { 1096 assert(isValidIndex(index) && "expected valid multi-dimensional index"); 1097 auto type = getType(); 1098 1099 // The sparse indices are 64-bit integers, so we can reinterpret the raw data 1100 // as a 1-D index array. 1101 auto sparseIndices = getIndices(); 1102 auto sparseIndexValues = sparseIndices.getValues<uint64_t>(); 1103 1104 // Check to see if the indices are a splat. 1105 if (sparseIndices.isSplat()) { 1106 // If the index is also not a splat of the index value, we know that the 1107 // value is zero. 1108 auto splatIndex = *sparseIndexValues.begin(); 1109 if (llvm::any_of(index, [=](uint64_t i) { return i != splatIndex; })) 1110 return getZeroAttr(); 1111 1112 // If the indices are a splat, we also expect the values to be a splat. 1113 assert(getValues().isSplat() && "expected splat values"); 1114 return getValues().getSplatValue(); 1115 } 1116 1117 // Build a mapping between known indices and the offset of the stored element. 1118 llvm::SmallDenseMap<llvm::ArrayRef<uint64_t>, size_t> mappedIndices; 1119 auto numSparseIndices = sparseIndices.getType().getDimSize(0); 1120 size_t rank = type.getRank(); 1121 for (size_t i = 0, e = numSparseIndices; i != e; ++i) 1122 mappedIndices.try_emplace( 1123 {&*std::next(sparseIndexValues.begin(), i * rank), rank}, i); 1124 1125 // Look for the provided index key within the mapped indices. If the provided 1126 // index is not found, then return a zero attribute. 1127 auto it = mappedIndices.find(index); 1128 if (it == mappedIndices.end()) 1129 return getZeroAttr(); 1130 1131 // Otherwise, return the held sparse value element. 1132 return getValues().getValue(it->second); 1133 } 1134 1135 /// Get a zero APFloat for the given sparse attribute. 1136 APFloat SparseElementsAttr::getZeroAPFloat() const { 1137 auto eltType = getType().getElementType().cast<FloatType>(); 1138 return APFloat(eltType.getFloatSemantics()); 1139 } 1140 1141 /// Get a zero APInt for the given sparse attribute. 1142 APInt SparseElementsAttr::getZeroAPInt() const { 1143 auto eltType = getType().getElementType().cast<IntegerType>(); 1144 return APInt::getNullValue(eltType.getWidth()); 1145 } 1146 1147 /// Get a zero attribute for the given attribute type. 1148 Attribute SparseElementsAttr::getZeroAttr() const { 1149 auto eltType = getType().getElementType(); 1150 1151 // Handle floating point elements. 1152 if (eltType.isa<FloatType>()) 1153 return FloatAttr::get(eltType, 0); 1154 1155 // Otherwise, this is an integer. 1156 auto intEltTy = eltType.cast<IntegerType>(); 1157 if (intEltTy.getWidth() == 1) 1158 return BoolAttr::get(false, eltType.getContext()); 1159 return IntegerAttr::get(eltType, 0); 1160 } 1161 1162 /// Flatten, and return, all of the sparse indices in this attribute in 1163 /// row-major order. 1164 std::vector<ptrdiff_t> SparseElementsAttr::getFlattenedSparseIndices() const { 1165 std::vector<ptrdiff_t> flatSparseIndices; 1166 1167 // The sparse indices are 64-bit integers, so we can reinterpret the raw data 1168 // as a 1-D index array. 1169 auto sparseIndices = getIndices(); 1170 auto sparseIndexValues = sparseIndices.getValues<uint64_t>(); 1171 if (sparseIndices.isSplat()) { 1172 SmallVector<uint64_t, 8> indices(getType().getRank(), 1173 *sparseIndexValues.begin()); 1174 flatSparseIndices.push_back(getFlattenedIndex(indices)); 1175 return flatSparseIndices; 1176 } 1177 1178 // Otherwise, reinterpret each index as an ArrayRef when flattening. 1179 auto numSparseIndices = sparseIndices.getType().getDimSize(0); 1180 size_t rank = getType().getRank(); 1181 for (size_t i = 0, e = numSparseIndices; i != e; ++i) 1182 flatSparseIndices.push_back(getFlattenedIndex( 1183 {&*std::next(sparseIndexValues.begin(), i * rank), rank})); 1184 return flatSparseIndices; 1185 } 1186 1187 //===----------------------------------------------------------------------===// 1188 // MutableDictionaryAttr 1189 //===----------------------------------------------------------------------===// 1190 1191 MutableDictionaryAttr::MutableDictionaryAttr( 1192 ArrayRef<NamedAttribute> attributes) { 1193 setAttrs(attributes); 1194 } 1195 1196 ArrayRef<NamedAttribute> MutableDictionaryAttr::getAttrs() const { 1197 return attrs ? attrs.getValue() : llvm::None; 1198 } 1199 1200 /// Replace the held attributes with ones provided in 'newAttrs'. 1201 void MutableDictionaryAttr::setAttrs(ArrayRef<NamedAttribute> attributes) { 1202 // Don't create an attribute list if there are no attributes. 1203 if (attributes.empty()) 1204 attrs = nullptr; 1205 else 1206 attrs = DictionaryAttr::get(attributes, attributes[0].second.getContext()); 1207 } 1208 1209 /// Return the specified attribute if present, null otherwise. 1210 Attribute MutableDictionaryAttr::get(StringRef name) const { 1211 return attrs ? attrs.get(name) : nullptr; 1212 } 1213 1214 /// Return the specified attribute if present, null otherwise. 1215 Attribute MutableDictionaryAttr::get(Identifier name) const { 1216 return attrs ? attrs.get(name) : nullptr; 1217 } 1218 1219 /// Return the specified named attribute if present, None otherwise. 1220 Optional<NamedAttribute> MutableDictionaryAttr::getNamed(StringRef name) const { 1221 return attrs ? attrs.getNamed(name) : Optional<NamedAttribute>(); 1222 } 1223 Optional<NamedAttribute> 1224 MutableDictionaryAttr::getNamed(Identifier name) const { 1225 return attrs ? attrs.getNamed(name) : Optional<NamedAttribute>(); 1226 } 1227 1228 /// If the an attribute exists with the specified name, change it to the new 1229 /// value. Otherwise, add a new attribute with the specified name/value. 1230 void MutableDictionaryAttr::set(Identifier name, Attribute value) { 1231 assert(value && "attributes may never be null"); 1232 1233 // Look for an existing value for the given name, and set it in-place. 1234 ArrayRef<NamedAttribute> values = getAttrs(); 1235 auto it = llvm::find_if( 1236 values, [name](NamedAttribute attr) { return attr.first == name; }); 1237 if (it != values.end()) { 1238 // Bail out early if the value is the same as what we already have. 1239 if (it->second == value) 1240 return; 1241 1242 SmallVector<NamedAttribute, 8> newAttrs(values.begin(), values.end()); 1243 newAttrs[it - values.begin()].second = value; 1244 attrs = DictionaryAttr::getWithSorted(newAttrs, value.getContext()); 1245 return; 1246 } 1247 1248 // Otherwise, insert the new attribute into its sorted position. 1249 it = llvm::lower_bound(values, name, compareNamedAttributeWithName); 1250 SmallVector<NamedAttribute, 8> newAttrs; 1251 newAttrs.reserve(values.size() + 1); 1252 newAttrs.append(values.begin(), it); 1253 newAttrs.push_back({name, value}); 1254 newAttrs.append(it, values.end()); 1255 attrs = DictionaryAttr::getWithSorted(newAttrs, value.getContext()); 1256 } 1257 1258 /// Remove the attribute with the specified name if it exists. The return 1259 /// value indicates whether the attribute was present or not. 1260 auto MutableDictionaryAttr::remove(Identifier name) -> RemoveResult { 1261 auto origAttrs = getAttrs(); 1262 for (unsigned i = 0, e = origAttrs.size(); i != e; ++i) { 1263 if (origAttrs[i].first == name) { 1264 // Handle the simple case of removing the only attribute in the list. 1265 if (e == 1) { 1266 attrs = nullptr; 1267 return RemoveResult::Removed; 1268 } 1269 1270 SmallVector<NamedAttribute, 8> newAttrs; 1271 newAttrs.reserve(origAttrs.size() - 1); 1272 newAttrs.append(origAttrs.begin(), origAttrs.begin() + i); 1273 newAttrs.append(origAttrs.begin() + i + 1, origAttrs.end()); 1274 attrs = DictionaryAttr::getWithSorted(newAttrs, 1275 newAttrs[0].second.getContext()); 1276 return RemoveResult::Removed; 1277 } 1278 } 1279 return RemoveResult::NotFound; 1280 } 1281