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