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