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