MLIR 24.0.0git
BuiltinTypes.cpp
Go to the documentation of this file.
1//===- BuiltinTypes.cpp - MLIR Builtin Type 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
10#include "TypeDetail.h"
12#include "mlir/IR/AffineExpr.h"
13#include "mlir/IR/AffineMap.h"
17#include "mlir/IR/Diagnostics.h"
18#include "mlir/IR/Dialect.h"
21#include "llvm/ADT/APFloat.h"
22#include "llvm/ADT/APInt.h"
23#include "llvm/ADT/Sequence.h"
24#include "llvm/ADT/TypeSwitch.h"
25#include "llvm/Support/CheckedArithmetic.h"
26#include <cstring>
27
28using namespace mlir;
29using namespace mlir::detail;
30
31//===----------------------------------------------------------------------===//
32/// Tablegen Type Definitions
33//===----------------------------------------------------------------------===//
34
35#define GET_TYPEDEF_CLASSES
36#include "mlir/IR/BuiltinTypes.cpp.inc"
37
38namespace mlir {
39#include "mlir/IR/BuiltinTypeConstraints.cpp.inc"
40} // namespace mlir
41
42//===----------------------------------------------------------------------===//
43// BuiltinDialect
44//===----------------------------------------------------------------------===//
45
46void BuiltinDialect::registerTypes() {
47 addTypes<
48#define GET_TYPEDEF_LIST
49#include "mlir/IR/BuiltinTypes.cpp.inc"
50 >();
51}
52
53//===----------------------------------------------------------------------===//
54/// ComplexType
55//===----------------------------------------------------------------------===//
56
57/// Verify the construction of an integer type.
58LogicalResult ComplexType::verify(function_ref<InFlightDiagnostic()> emitError,
59 Type elementType) {
60 if (!elementType.isIntOrFloat())
61 return emitError() << "invalid element type for complex";
62 return success();
63}
64
65size_t ComplexType::getDenseElementBitSize() const {
66 auto elemTy = cast<DenseElementType>(getElementType());
67 return llvm::alignTo<8>(elemTy.getDenseElementBitSize()) * 2;
68}
69
70Attribute ComplexType::convertToAttribute(ArrayRef<char> rawData) const {
71 auto elemTy = cast<DenseElementType>(getElementType());
72 size_t singleElementBytes =
73 llvm::alignTo<8>(elemTy.getDenseElementBitSize()) / 8;
75 elemTy.convertToAttribute(rawData.take_front(singleElementBytes));
77 elemTy.convertToAttribute(rawData.take_back(singleElementBytes));
78 return ArrayAttr::get(getContext(), {real, imag});
79}
80
81LogicalResult
82ComplexType::convertFromAttribute(Attribute attr,
84 auto arrayAttr = dyn_cast<ArrayAttr>(attr);
85 if (!arrayAttr || arrayAttr.size() != 2)
86 return failure();
87 auto elemTy = cast<DenseElementType>(getElementType());
88 SmallVector<char> realData, imagData;
89 if (failed(elemTy.convertFromAttribute(arrayAttr[0], realData)))
90 return failure();
91 if (failed(elemTy.convertFromAttribute(arrayAttr[1], imagData)))
92 return failure();
93 result.append(realData);
94 result.append(imagData);
95 return success();
96}
97
98//===----------------------------------------------------------------------===//
99// Integer Type
100//===----------------------------------------------------------------------===//
101
102/// Verify the construction of an integer type.
103LogicalResult IntegerType::verify(function_ref<InFlightDiagnostic()> emitError,
104 unsigned width,
105 SignednessSemantics signedness) {
106 if (width > IntegerType::kMaxWidth) {
107 return emitError() << "integer bitwidth is limited to "
108 << IntegerType::kMaxWidth << " bits";
109 }
110 return success();
111}
112
113unsigned IntegerType::getWidth() const { return getImpl()->width; }
114
115IntegerType::SignednessSemantics IntegerType::getSignedness() const {
116 return getImpl()->signedness;
117}
118
119IntegerType IntegerType::scaleElementBitwidth(unsigned scale) {
120 if (!scale)
121 return IntegerType();
122 return IntegerType::get(getContext(), scale * getWidth(), getSignedness());
123}
124
125size_t IntegerType::getDenseElementBitSize() const {
126 // Return the actual bit width. Storage alignment is handled separately.
127 return getWidth();
128}
129
130Attribute IntegerType::convertToAttribute(ArrayRef<char> rawData) const {
131 APInt value = detail::readBits(rawData.data(), /*bitPos=*/0, getWidth());
132 return IntegerAttr::get(*this, value);
133}
134
136 size_t byteSize = llvm::divideCeil(apInt.getBitWidth(), CHAR_BIT);
137 size_t bitPos = result.size() * CHAR_BIT;
138 result.resize(result.size() + byteSize);
139 detail::writeBits(result.data(), bitPos, apInt);
140}
141
142LogicalResult
143IntegerType::convertFromAttribute(Attribute attr,
145 auto intAttr = dyn_cast<IntegerAttr>(attr);
146 if (!intAttr || intAttr.getType() != *this)
147 return failure();
148 writeAPIntToVector(intAttr.getValue(), result);
149 return success();
150}
151
152//===----------------------------------------------------------------------===//
153// Index Type
154//===----------------------------------------------------------------------===//
155
156size_t IndexType::getDenseElementBitSize() const {
157 return kInternalStorageBitWidth;
158}
159
160Attribute IndexType::convertToAttribute(ArrayRef<char> rawData) const {
161 APInt value =
162 detail::readBits(rawData.data(), /*bitPos=*/0, kInternalStorageBitWidth);
163 return IntegerAttr::get(*this, value);
164}
165
166LogicalResult
167IndexType::convertFromAttribute(Attribute attr,
169 auto intAttr = dyn_cast<IntegerAttr>(attr);
170 if (!intAttr || intAttr.getType() != *this)
171 return failure();
172 writeAPIntToVector(intAttr.getValue(), result);
173 return success();
174}
175
176//===----------------------------------------------------------------------===//
177// Float Types
178//===----------------------------------------------------------------------===//
179
180// Mapping from MLIR FloatType to APFloat semantics.
181#define FLOAT_TYPE_SEMANTICS(TYPE, SEM) \
182 const llvm::fltSemantics &TYPE::getFloatSemantics() const { \
183 return APFloat::SEM(); \
184 }
185FLOAT_TYPE_SEMANTICS(Float4E2M1FNType, Float4E2M1FN)
186FLOAT_TYPE_SEMANTICS(Float6E2M3FNType, Float6E2M3FN)
187FLOAT_TYPE_SEMANTICS(Float6E3M2FNType, Float6E3M2FN)
188FLOAT_TYPE_SEMANTICS(Float8E5M2Type, Float8E5M2)
189FLOAT_TYPE_SEMANTICS(Float8E4M3Type, Float8E4M3)
190FLOAT_TYPE_SEMANTICS(Float8E4M3FNType, Float8E4M3FN)
191FLOAT_TYPE_SEMANTICS(Float8E5M2FNUZType, Float8E5M2FNUZ)
192FLOAT_TYPE_SEMANTICS(Float8E4M3FNUZType, Float8E4M3FNUZ)
193FLOAT_TYPE_SEMANTICS(Float8E4M3B11FNUZType, Float8E4M3B11FNUZ)
194FLOAT_TYPE_SEMANTICS(Float8E3M4Type, Float8E3M4)
195FLOAT_TYPE_SEMANTICS(Float8E8M0FNUType, Float8E8M0FNU)
196FLOAT_TYPE_SEMANTICS(Float8E5M3FNUType, Float8E5M3FNU)
197FLOAT_TYPE_SEMANTICS(BFloat16Type, BFloat)
198FLOAT_TYPE_SEMANTICS(Float16Type, IEEEhalf)
199FLOAT_TYPE_SEMANTICS(FloatTF32Type, FloatTF32)
200FLOAT_TYPE_SEMANTICS(Float32Type, IEEEsingle)
201FLOAT_TYPE_SEMANTICS(Float64Type, IEEEdouble)
202FLOAT_TYPE_SEMANTICS(Float80Type, x87DoubleExtended)
203FLOAT_TYPE_SEMANTICS(Float128Type, IEEEquad)
204#undef FLOAT_TYPE_SEMANTICS
205
206FloatType Float16Type::scaleElementBitwidth(unsigned scale) const {
207 if (scale == 2)
208 return Float32Type::get(getContext());
209 if (scale == 4)
210 return Float64Type::get(getContext());
211 return FloatType();
212}
213
214FloatType BFloat16Type::scaleElementBitwidth(unsigned scale) const {
215 if (scale == 2)
216 return Float32Type::get(getContext());
217 if (scale == 4)
218 return Float64Type::get(getContext());
219 return FloatType();
220}
221
222FloatType Float32Type::scaleElementBitwidth(unsigned scale) const {
223 if (scale == 2)
224 return Float64Type::get(getContext());
225 return FloatType();
226}
227
228//===----------------------------------------------------------------------===//
229// FunctionType
230//===----------------------------------------------------------------------===//
231
232unsigned FunctionType::getNumInputs() const { return getImpl()->numInputs; }
233
234ArrayRef<Type> FunctionType::getInputs() const {
235 return getImpl()->getInputs();
236}
237
238unsigned FunctionType::getNumResults() const { return getImpl()->numResults; }
239
240ArrayRef<Type> FunctionType::getResults() const {
241 return getImpl()->getResults();
242}
243
244FunctionType FunctionType::clone(TypeRange inputs, TypeRange results) const {
245 return get(getContext(), inputs, results);
246}
247
248/// Returns a new function type with the specified arguments and results
249/// inserted.
250FunctionType FunctionType::getWithArgsAndResults(
251 ArrayRef<unsigned> argIndices, TypeRange argTypes,
252 ArrayRef<unsigned> resultIndices, TypeRange resultTypes) {
253 SmallVector<Type> argStorage, resultStorage;
254 TypeRange newArgTypes =
255 insertTypesInto(getInputs(), argIndices, argTypes, argStorage);
256 TypeRange newResultTypes =
257 insertTypesInto(getResults(), resultIndices, resultTypes, resultStorage);
258 return clone(newArgTypes, newResultTypes);
259}
260
261/// Returns a new function type without the specified arguments and results.
262FunctionType
263FunctionType::getWithoutArgsAndResults(const BitVector &argIndices,
264 const BitVector &resultIndices) {
265 SmallVector<Type> argStorage, resultStorage;
266 TypeRange newArgTypes = filterTypesOut(getInputs(), argIndices, argStorage);
267 TypeRange newResultTypes =
268 filterTypesOut(getResults(), resultIndices, resultStorage);
269 return clone(newArgTypes, newResultTypes);
270}
271
272//===----------------------------------------------------------------------===//
273// GraphType
274//===----------------------------------------------------------------------===//
275
276unsigned GraphType::getNumInputs() const { return getImpl()->numInputs; }
277
278ArrayRef<Type> GraphType::getInputs() const { return getImpl()->getInputs(); }
279
280unsigned GraphType::getNumResults() const { return getImpl()->numResults; }
281
282ArrayRef<Type> GraphType::getResults() const { return getImpl()->getResults(); }
283
284GraphType GraphType::clone(TypeRange inputs, TypeRange results) const {
285 return get(getContext(), inputs, results);
286}
287
288/// Returns a new function type with the specified arguments and results
289/// inserted.
290GraphType GraphType::getWithArgsAndResults(ArrayRef<unsigned> argIndices,
291 TypeRange argTypes,
292 ArrayRef<unsigned> resultIndices,
293 TypeRange resultTypes) {
294 SmallVector<Type> argStorage, resultStorage;
295 TypeRange newArgTypes =
296 insertTypesInto(getInputs(), argIndices, argTypes, argStorage);
297 TypeRange newResultTypes =
298 insertTypesInto(getResults(), resultIndices, resultTypes, resultStorage);
299 return clone(newArgTypes, newResultTypes);
300}
301
302/// Returns a new function type without the specified arguments and results.
303GraphType GraphType::getWithoutArgsAndResults(const BitVector &argIndices,
304 const BitVector &resultIndices) {
305 SmallVector<Type> argStorage, resultStorage;
306 TypeRange newArgTypes = filterTypesOut(getInputs(), argIndices, argStorage);
307 TypeRange newResultTypes =
308 filterTypesOut(getResults(), resultIndices, resultStorage);
309 return clone(newArgTypes, newResultTypes);
310}
311//===----------------------------------------------------------------------===//
312// OpaqueType
313//===----------------------------------------------------------------------===//
314
315/// Verify the construction of an opaque type.
316LogicalResult OpaqueType::verify(function_ref<InFlightDiagnostic()> emitError,
317 StringAttr dialect, StringRef typeData) {
318 if (!Dialect::isValidNamespace(dialect.strref()))
319 return emitError() << "invalid dialect namespace '" << dialect << "'";
320
321 // Check that the dialect is actually registered.
322 MLIRContext *context = dialect.getContext();
323 if (!context->allowsUnregisteredDialects() &&
324 !context->getLoadedDialect(dialect.strref())) {
325 return emitError()
326 << "`!" << dialect << "<\"" << typeData << "\">"
327 << "` type created with unregistered dialect. If this is "
328 "intended, please call allowUnregisteredDialects() on the "
329 "MLIRContext, or use -allow-unregistered-dialect with "
330 "the MLIR opt tool used";
331 }
332
333 return success();
334}
335
336//===----------------------------------------------------------------------===//
337// VectorType
338//===----------------------------------------------------------------------===//
339
340bool VectorType::isValidElementType(Type t) {
342}
343
344LogicalResult VectorType::verify(function_ref<InFlightDiagnostic()> emitError,
345 ArrayRef<int64_t> shape, Type elementType,
346 ArrayRef<bool> scalableDims) {
347 if (!isValidElementType(elementType))
348 return emitError()
349 << "vector elements must be int/index/float type but got "
350 << elementType;
351
352 if (any_of(shape, [](int64_t i) { return i <= 0; }))
353 return emitError()
354 << "vector types must have positive constant sizes but got "
355 << shape;
356
357 if (scalableDims.size() != shape.size())
358 return emitError() << "number of dims must match, got "
359 << scalableDims.size() << " and " << shape.size();
360
361 return success();
362}
363
364VectorType VectorType::scaleElementBitwidth(unsigned scale) {
365 if (!scale)
366 return VectorType();
367 if (auto et = llvm::dyn_cast<IntegerType>(getElementType()))
368 if (auto scaledEt = et.scaleElementBitwidth(scale))
369 return VectorType::get(getShape(), scaledEt, getScalableDims());
370 if (auto et = llvm::dyn_cast<FloatType>(getElementType()))
371 if (auto scaledEt = et.scaleElementBitwidth(scale))
372 return VectorType::get(getShape(), scaledEt, getScalableDims());
373 return VectorType();
374}
375
376VectorType VectorType::cloneWith(std::optional<ArrayRef<int64_t>> shape,
377 Type elementType) const {
378 return VectorType::get(shape.value_or(getShape()), elementType,
379 getScalableDims());
380}
381
382//===----------------------------------------------------------------------===//
383// TensorType
384//===----------------------------------------------------------------------===//
385
388 .Case<RankedTensorType, UnrankedTensorType>(
389 [](auto type) { return type.getElementType(); });
390}
391
393 return !llvm::isa<UnrankedTensorType>(*this);
394}
395
397 return llvm::cast<RankedTensorType>(*this).getShape();
398}
399
401 Type elementType) const {
402 if (llvm::dyn_cast<UnrankedTensorType>(*this)) {
403 if (shape)
404 return RankedTensorType::get(*shape, elementType);
405 return UnrankedTensorType::get(elementType);
406 }
407
408 auto rankedTy = llvm::cast<RankedTensorType>(*this);
409 if (!shape)
410 return RankedTensorType::get(rankedTy.getShape(), elementType,
411 rankedTy.getEncoding());
412 return RankedTensorType::get(shape.value_or(rankedTy.getShape()), elementType,
413 rankedTy.getEncoding());
414}
415
417 Type elementType) const {
418 return ::llvm::cast<RankedTensorType>(cloneWith(shape, elementType));
419}
420
421RankedTensorType TensorType::clone(::llvm::ArrayRef<int64_t> shape) const {
422 return ::llvm::cast<RankedTensorType>(cloneWith(shape, getElementType()));
423}
424
425// Check if "elementType" can be an element type of a tensor.
426static LogicalResult
428 Type elementType) {
429 if (!TensorType::isValidElementType(elementType))
430 return emitError() << "invalid tensor element type: " << elementType;
431 return success();
432}
433
434/// Return true if the specified element type is ok in a tensor.
436 // Note: Non standard/builtin types are allowed to exist within tensor
437 // types. Dialects are expected to verify that tensor types have a valid
438 // element type within that dialect.
439 return llvm::isa<ComplexType, FloatType, IntegerType, OpaqueType, VectorType,
440 IndexType>(type) ||
441 !llvm::isa<BuiltinDialect>(type.getDialect());
442}
443
444//===----------------------------------------------------------------------===//
445// RankedTensorType
446//===----------------------------------------------------------------------===//
447
448LogicalResult
449RankedTensorType::verify(function_ref<InFlightDiagnostic()> emitError,
450 ArrayRef<int64_t> shape, Type elementType,
451 Attribute encoding) {
452 for (int64_t s : shape)
453 if (s < 0 && ShapedType::isStatic(s))
454 return emitError() << "invalid tensor dimension size";
455 if (auto v = llvm::dyn_cast_or_null<VerifiableTensorEncoding>(encoding))
456 if (failed(v.verifyEncoding(shape, elementType, emitError)))
457 return failure();
458 return checkTensorElementType(emitError, elementType);
459}
460
461//===----------------------------------------------------------------------===//
462// UnrankedTensorType
463//===----------------------------------------------------------------------===//
464
465LogicalResult
466UnrankedTensorType::verify(function_ref<InFlightDiagnostic()> emitError,
467 Type elementType) {
468 return checkTensorElementType(emitError, elementType);
469}
470
471//===----------------------------------------------------------------------===//
472// BaseMemRefType
473//===----------------------------------------------------------------------===//
474
477 .Case<MemRefType, UnrankedMemRefType>(
478 [](auto type) { return type.getElementType(); });
479}
480
482 return !llvm::isa<UnrankedMemRefType>(*this);
483}
484
486 return llvm::cast<MemRefType>(*this).getShape();
487}
488
490 Type elementType) const {
491 if (llvm::dyn_cast<UnrankedMemRefType>(*this)) {
492 if (!shape)
493 return UnrankedMemRefType::get(elementType, getMemorySpace());
494 MemRefType::Builder builder(*shape, elementType);
496 return builder;
497 }
498
499 MemRefType::Builder builder(llvm::cast<MemRefType>(*this));
500 if (shape)
501 builder.setShape(*shape);
502 builder.setElementType(elementType);
503 return builder;
504}
505
506FailureOr<PtrLikeTypeInterface>
508 std::optional<Type> elementType) const {
509 Type eTy = elementType ? *elementType : getElementType();
510 if (llvm::dyn_cast<UnrankedMemRefType>(*this))
511 return cast<PtrLikeTypeInterface>(
512 UnrankedMemRefType::get(eTy, memorySpace));
513
514 MemRefType::Builder builder(llvm::cast<MemRefType>(*this));
515 builder.setElementType(eTy);
516 builder.setMemorySpace(memorySpace);
517 return cast<PtrLikeTypeInterface>(static_cast<MemRefType>(builder));
518}
519
521 Type elementType) const {
522 return ::llvm::cast<MemRefType>(cloneWith(shape, elementType));
523}
524
526 return ::llvm::cast<MemRefType>(cloneWith(shape, getElementType()));
527}
528
530 if (auto rankedMemRefTy = llvm::dyn_cast<MemRefType>(*this))
531 return rankedMemRefTy.getMemorySpace();
532 return llvm::cast<UnrankedMemRefType>(*this).getMemorySpace();
533}
534
536 if (auto rankedMemRefTy = llvm::dyn_cast<MemRefType>(*this))
537 return rankedMemRefTy.getMemorySpaceAsInt();
538 return llvm::cast<UnrankedMemRefType>(*this).getMemorySpaceAsInt();
539}
540
541//===----------------------------------------------------------------------===//
542// MemRefType
543//===----------------------------------------------------------------------===//
544
545std::optional<llvm::SmallDenseSet<unsigned>>
547 ArrayRef<int64_t> reducedShape,
548 bool matchDynamic) {
549 size_t originalRank = originalShape.size(), reducedRank = reducedShape.size();
550 llvm::SmallDenseSet<unsigned> unusedDims;
551 unsigned reducedIdx = 0;
552 for (unsigned originalIdx = 0; originalIdx < originalRank; ++originalIdx) {
553 // Greedily insert `originalIdx` if match.
554 int64_t origSize = originalShape[originalIdx];
555 // if `matchDynamic`, count dynamic dims as a match, unless `origSize` is 1.
556 if (matchDynamic && reducedIdx < reducedRank && origSize != 1 &&
557 (ShapedType::isDynamic(reducedShape[reducedIdx]) ||
558 ShapedType::isDynamic(origSize))) {
559 reducedIdx++;
560 continue;
561 }
562 if (reducedIdx < reducedRank && origSize == reducedShape[reducedIdx]) {
563 reducedIdx++;
564 continue;
565 }
566
567 unusedDims.insert(originalIdx);
568 // If no match on `originalIdx`, the `originalShape` at this dimension
569 // must be 1, otherwise we bail.
570 if (origSize != 1)
571 return std::nullopt;
572 }
573 // The whole reducedShape must be scanned, otherwise we bail.
574 if (reducedIdx != reducedRank)
575 return std::nullopt;
576 return unusedDims;
577}
578
580mlir::isRankReducedType(ShapedType originalType,
581 ShapedType candidateReducedType) {
582 if (originalType == candidateReducedType)
584
585 ShapedType originalShapedType = llvm::cast<ShapedType>(originalType);
586 ShapedType candidateReducedShapedType =
587 llvm::cast<ShapedType>(candidateReducedType);
588
589 // Rank and size logic is valid for all ShapedTypes.
590 ArrayRef<int64_t> originalShape = originalShapedType.getShape();
591 ArrayRef<int64_t> candidateReducedShape =
592 candidateReducedShapedType.getShape();
593 unsigned originalRank = originalShape.size(),
594 candidateReducedRank = candidateReducedShape.size();
595 if (candidateReducedRank > originalRank)
597
598 auto optionalUnusedDimsMask =
599 computeRankReductionMask(originalShape, candidateReducedShape);
600
601 // Sizes cannot be matched in case empty vector is returned.
602 if (!optionalUnusedDimsMask)
604
605 if (originalShapedType.getElementType() !=
606 candidateReducedShapedType.getElementType())
608
610}
611
613 MLIRContext *ctx) {
614 if (memorySpace == 0)
615 return nullptr;
616
617 return IntegerAttr::get(IntegerType::get(ctx, 64), memorySpace);
618}
619
621 IntegerAttr intMemorySpace = llvm::dyn_cast_or_null<IntegerAttr>(memorySpace);
622 if (intMemorySpace && intMemorySpace.getValue() == 0)
623 return nullptr;
624
625 return memorySpace;
626}
627
629 if (!memorySpace)
630 return 0;
631
632 assert(llvm::isa<IntegerAttr>(memorySpace) &&
633 "Using `getMemorySpaceInteger` with non-Integer attribute");
634
635 return static_cast<unsigned>(llvm::cast<IntegerAttr>(memorySpace).getInt());
636}
637
638unsigned MemRefType::getMemorySpaceAsInt() const {
639 return detail::getMemorySpaceAsInt(getMemorySpace());
640}
641
642MemRefType MemRefType::get(ArrayRef<int64_t> shape, Type elementType,
643 MemRefLayoutAttrInterface layout,
644 Attribute memorySpace) {
645 // Use default layout for empty attribute.
646 if (!layout)
647 layout = AffineMapAttr::get(AffineMap::getMultiDimIdentityMap(
648 shape.size(), elementType.getContext()));
649
650 // Drop default memory space value and replace it with empty attribute.
651 memorySpace = skipDefaultMemorySpace(memorySpace);
652
653 return Base::get(elementType.getContext(), shape, elementType, layout,
654 memorySpace);
655}
656
657MemRefType MemRefType::getChecked(
659 Type elementType, MemRefLayoutAttrInterface layout, Attribute memorySpace) {
660
661 // Use default layout for empty attribute.
662 if (!layout)
663 layout = AffineMapAttr::get(AffineMap::getMultiDimIdentityMap(
664 shape.size(), elementType.getContext()));
665
666 // Drop default memory space value and replace it with empty attribute.
667 memorySpace = skipDefaultMemorySpace(memorySpace);
668
669 return Base::getChecked(emitErrorFn, elementType.getContext(), shape,
670 elementType, layout, memorySpace);
671}
672
673MemRefType MemRefType::get(ArrayRef<int64_t> shape, Type elementType,
674 AffineMap map, Attribute memorySpace) {
675
676 // Use default layout for empty map.
677 if (!map)
679 elementType.getContext());
680
681 // Wrap AffineMap into Attribute.
682 auto layout = AffineMapAttr::get(map);
683
684 // Drop default memory space value and replace it with empty attribute.
685 memorySpace = skipDefaultMemorySpace(memorySpace);
686
687 return Base::get(elementType.getContext(), shape, elementType, layout,
688 memorySpace);
689}
690
691MemRefType
692MemRefType::getChecked(function_ref<InFlightDiagnostic()> emitErrorFn,
693 ArrayRef<int64_t> shape, Type elementType, AffineMap map,
694 Attribute memorySpace) {
695
696 // Use default layout for empty map.
697 if (!map)
699 elementType.getContext());
700
701 // Wrap AffineMap into Attribute.
702 auto layout = AffineMapAttr::get(map);
703
704 // Drop default memory space value and replace it with empty attribute.
705 memorySpace = skipDefaultMemorySpace(memorySpace);
706
707 return Base::getChecked(emitErrorFn, elementType.getContext(), shape,
708 elementType, layout, memorySpace);
709}
710
711MemRefType MemRefType::get(ArrayRef<int64_t> shape, Type elementType,
712 AffineMap map, unsigned memorySpaceInd) {
713
714 // Use default layout for empty map.
715 if (!map)
717 elementType.getContext());
718
719 // Wrap AffineMap into Attribute.
720 auto layout = AffineMapAttr::get(map);
721
722 // Convert deprecated integer-like memory space to Attribute.
723 Attribute memorySpace =
724 wrapIntegerMemorySpace(memorySpaceInd, elementType.getContext());
725
726 return Base::get(elementType.getContext(), shape, elementType, layout,
727 memorySpace);
728}
729
730MemRefType
731MemRefType::getChecked(function_ref<InFlightDiagnostic()> emitErrorFn,
732 ArrayRef<int64_t> shape, Type elementType, AffineMap map,
733 unsigned memorySpaceInd) {
734
735 // Use default layout for empty map.
736 if (!map)
738 elementType.getContext());
739
740 // Wrap AffineMap into Attribute.
741 auto layout = AffineMapAttr::get(map);
742
743 // Convert deprecated integer-like memory space to Attribute.
744 Attribute memorySpace =
745 wrapIntegerMemorySpace(memorySpaceInd, elementType.getContext());
746
747 return Base::getChecked(emitErrorFn, elementType.getContext(), shape,
748 elementType, layout, memorySpace);
749}
750
751LogicalResult MemRefType::verify(function_ref<InFlightDiagnostic()> emitError,
752 ArrayRef<int64_t> shape, Type elementType,
753 MemRefLayoutAttrInterface layout,
754 Attribute memorySpace) {
755 if (!BaseMemRefType::isValidElementType(elementType))
756 return emitError() << "invalid memref element type";
757
758 // Negative sizes are not allowed except for `kDynamic`.
759 for (int64_t s : shape)
760 if (s < 0 && ShapedType::isStatic(s))
761 return emitError() << "invalid memref size";
762
763 assert(layout && "missing layout specification");
764 if (failed(layout.verifyLayout(shape, emitError)))
765 return failure();
766
767 return success();
768}
769
770bool MemRefType::areTrailingDimsContiguous(int64_t n) {
771 assert(n <= getRank() &&
772 "number of dimensions to check must not exceed rank");
773 return n <= getNumContiguousTrailingDims();
774}
775
776int64_t MemRefType::getNumContiguousTrailingDims() {
777 const int64_t n = getRank();
778
779 // memrefs with identity layout are entirely contiguous.
780 if (getLayout().isIdentity())
781 return n;
782
783 // Get the strides (if any). Failing to do that, conservatively assume a
784 // non-contiguous layout.
785 int64_t offset;
786 SmallVector<int64_t> strides;
787 if (!succeeded(getStridesAndOffset(strides, offset)))
788 return 0;
789
791
792 // A memref with dimensions `d0, d1, ..., dn-1` and strides
793 // `s0, s1, ..., sn-1` is contiguous up to dimension `k`
794 // if each stride `si` is the product of the dimensions `di+1, ..., dn-1`,
795 // for `i` in `[k, n-1]`.
796 // Ignore stride elements if the corresponding dimension is 1, as they are
797 // of no consequence.
798 int64_t dimProduct = 1;
799 for (int64_t i = n - 1; i >= 0; --i) {
800 if (shape[i] == 1)
801 continue;
802 if (strides[i] != dimProduct)
803 return n - i - 1;
804 if (shape[i] == ShapedType::kDynamic)
805 return n - i;
806 dimProduct *= shape[i];
807 }
808
809 return n;
810}
811
812MemRefType MemRefType::canonicalizeStridedLayout() {
813 AffineMap m = getLayout().getAffineMap();
814
815 // Already in canonical form.
816 if (m.isIdentity())
817 return *this;
818
819 // Can't reduce to canonical identity form, return in canonical form.
820 if (m.getNumResults() > 1)
821 return *this;
822
823 // Corner-case for 0-D affine maps.
824 if (m.getNumDims() == 0 && m.getNumSymbols() == 0) {
825 if (auto cst = llvm::dyn_cast<AffineConstantExpr>(m.getResult(0)))
826 if (cst.getValue() == 0)
827 return MemRefType::Builder(*this).setLayout({});
828 return *this;
829 }
830
831 // 0-D corner case for empty shape that still have an affine map. Example:
832 // `memref<f32, affine_map<()[s0] -> (s0)>>`. This is a 1 element memref whose
833 // offset needs to remain, just return t.
834 if (getShape().empty())
835 return *this;
836
837 // If the canonical strided layout for the sizes of `t` is equal to the
838 // simplified layout of `t` we can just return an empty layout. Otherwise,
839 // just simplify the existing layout.
841 auto simplifiedLayoutExpr =
843 if (expr != simplifiedLayoutExpr)
844 return MemRefType::Builder(*this).setLayout(
845 AffineMapAttr::get(AffineMap::get(m.getNumDims(), m.getNumSymbols(),
846 simplifiedLayoutExpr)));
847 return MemRefType::Builder(*this).setLayout({});
848}
849
850LogicalResult MemRefType::getStridesAndOffset(SmallVectorImpl<int64_t> &strides,
851 int64_t &offset) const {
852 return getLayout().getStridesAndOffset(getShape(), strides, offset);
853}
854
855std::pair<SmallVector<int64_t>, int64_t>
856MemRefType::getStridesAndOffset() const {
857 SmallVector<int64_t> strides;
858 int64_t offset;
859 LogicalResult status = getStridesAndOffset(strides, offset);
860 (void)status;
861 assert(succeeded(status) && "Invalid use of check-free getStridesAndOffset");
862 return {strides, offset};
863}
864
865bool MemRefType::isStrided() {
866 int64_t offset;
868 auto res = getStridesAndOffset(strides, offset);
869 return succeeded(res);
870}
871
872bool MemRefType::isLastDimUnitStride() {
873 int64_t offset;
874 SmallVector<int64_t> strides;
875 auto successStrides = getStridesAndOffset(strides, offset);
876 return succeeded(successStrides) && (strides.empty() || strides.back() == 1);
877}
878
879//===----------------------------------------------------------------------===//
880// UnrankedMemRefType
881//===----------------------------------------------------------------------===//
882
883unsigned UnrankedMemRefType::getMemorySpaceAsInt() const {
884 return detail::getMemorySpaceAsInt(getMemorySpace());
885}
886
887LogicalResult
888UnrankedMemRefType::verify(function_ref<InFlightDiagnostic()> emitError,
889 Type elementType, Attribute memorySpace) {
890 if (!BaseMemRefType::isValidElementType(elementType))
891 return emitError() << "invalid memref element type";
892
893 return success();
894}
895
896//===----------------------------------------------------------------------===//
897/// TupleType
898//===----------------------------------------------------------------------===//
899
900/// Return the elements types for this tuple.
901ArrayRef<Type> TupleType::getTypes() const { return getImpl()->getTypes(); }
902
903/// Accumulate the types contained in this tuple and tuples nested within it.
904/// Note that this only flattens nested tuples, not any other container type,
905/// e.g. a tuple<i32, tensor<i32>, tuple<f32, tuple<i64>>> is flattened to
906/// (i32, tensor<i32>, f32, i64)
907void TupleType::getFlattenedTypes(SmallVectorImpl<Type> &types) {
908 for (Type type : getTypes()) {
909 if (auto nestedTuple = llvm::dyn_cast<TupleType>(type))
910 nestedTuple.getFlattenedTypes(types);
911 else
912 types.push_back(type);
913 }
914}
915
916/// Return the number of element types.
917size_t TupleType::size() const { return getImpl()->size(); }
918
919//===----------------------------------------------------------------------===//
920// Type Utilities
921//===----------------------------------------------------------------------===//
922
925 MLIRContext *context) {
926 // Size 0 corner case is useful for canonicalizations.
927 if (sizes.empty())
928 return getAffineConstantExpr(0, context);
929
930 assert(!exprs.empty() && "expected exprs");
931 auto maps = AffineMap::inferFromExprList(exprs, context);
932 assert(!maps.empty() && "Expected one non-empty map");
933 unsigned numDims = maps[0].getNumDims(), nSymbols = maps[0].getNumSymbols();
934
935 AffineExpr expr;
936 bool dynamicPoisonBit = false;
937 int64_t runningSize = 1;
938 for (auto en : llvm::zip(llvm::reverse(exprs), llvm::reverse(sizes))) {
939 int64_t size = std::get<1>(en);
940 AffineExpr dimExpr = std::get<0>(en);
941 AffineExpr stride = dynamicPoisonBit
942 ? getAffineSymbolExpr(nSymbols++, context)
943 : getAffineConstantExpr(runningSize, context);
944 expr = expr ? expr + dimExpr * stride : dimExpr * stride;
945 if (size > 0) {
946 auto result = llvm::checkedMul(runningSize, size);
947 if (!result) {
948 // Overflow occurred, treat as dynamic
949 dynamicPoisonBit = true;
950 } else {
951 runningSize = *result;
952 }
953 } else {
954 dynamicPoisonBit = true;
955 }
956 }
957 return simplifyAffineExpr(expr, numDims, nSymbols);
958}
959
961 MLIRContext *context) {
963 exprs.reserve(sizes.size());
964 for (auto dim : llvm::seq<unsigned>(0, sizes.size()))
965 exprs.push_back(getAffineDimExpr(dim, context));
966 return makeCanonicalStridedLayoutExpr(sizes, exprs, context);
967}
return success()
static LogicalResult getStridesAndOffset(AffineMap m, ArrayRef< int64_t > shape, SmallVectorImpl< AffineExpr > &strides, AffineExpr &offset)
A stride specification is a list of integer values that are either static or dynamic (encoded with Sh...
static void writeAPIntToVector(APInt apInt, SmallVectorImpl< char > &result)
static LogicalResult checkTensorElementType(function_ref< InFlightDiagnostic()> emitError, Type elementType)
#define FLOAT_TYPE_SEMANTICS(TYPE, SEM)
b getContext())
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
Definition SPIRVOps.cpp:229
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Definition Traits.cpp:117
Base type for affine expression.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
static AffineMap getMultiDimIdentityMap(unsigned numDims, MLIRContext *context)
Returns an AffineMap with 'numDims' identity result dim exprs.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
unsigned getNumSymbols() const
unsigned getNumDims() const
unsigned getNumResults() const
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
AffineExpr getResult(unsigned idx) const
bool isIdentity() const
Returns true if this affine map is an identity affine map.
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class provides a shared interface for ranked and unranked memref types.
ArrayRef< int64_t > getShape() const
Returns the shape of this memref type.
static bool isValidElementType(Type type)
Return true if the specified element type is ok in a memref.
FailureOr< PtrLikeTypeInterface > clonePtrWith(Attribute memorySpace, std::optional< Type > elementType) const
Clone this type with the given memory space and element type.
constexpr Type()=default
Attribute getMemorySpace() const
Returns the memory space in which data referred to by this memref resides.
unsigned getMemorySpaceAsInt() const
[deprecated] Returns the memory space in old raw integer representation.
BaseMemRefType cloneWith(std::optional< ArrayRef< int64_t > > shape, Type elementType) const
Clone this type with the given shape and element type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
Type getElementType() const
Returns the element type of this memref type.
MemRefType clone(ArrayRef< int64_t > shape, Type elementType) const
Return a clone of this type with the given new shape and element type.
static bool isValidNamespace(StringRef str)
Utility function that returns if the given string is a valid dialect namespace.
Definition Dialect.cpp:95
This class represents a diagnostic that is inflight and set to be reported.
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
bool allowsUnregisteredDialects()
Return true if we allow to create operation for unregistered dialects.
This is a builder type that keeps local references to arguments.
Builder & setShape(ArrayRef< int64_t > newShape)
Builder & setMemorySpace(Attribute newMemorySpace)
Builder & setElementType(Type newElementType)
Builder & setLayout(MemRefLayoutAttrInterface newLayout)
Tensor types represent multi-dimensional arrays, and have two variants: RankedTensorType and Unranked...
TensorType cloneWith(std::optional< ArrayRef< int64_t > > shape, Type elementType) const
Clone this type with the given shape and element type.
constexpr Type()=default
static bool isValidElementType(Type type)
Return true if the specified element type is ok in a tensor.
ArrayRef< int64_t > getShape() const
Returns the shape of this tensor type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
RankedTensorType clone(ArrayRef< int64_t > shape, Type elementType) const
Return a clone of this type with the given new shape and element type.
Type getElementType() const
Returns the element type of this tensor type.
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
Dialect & getDialect() const
Get the dialect this type is registered to.
Definition Types.h:107
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
AttrTypeReplacer.
Attribute wrapIntegerMemorySpace(unsigned memorySpace, MLIRContext *ctx)
Wraps deprecated integer memory space to the new Attribute form.
unsigned getMemorySpaceAsInt(Attribute memorySpace)
[deprecated] Returns the memory space in old raw integer representation.
Attribute skipDefaultMemorySpace(Attribute memorySpace)
Replaces default memorySpace (integer == 0) with empty Attribute.
void writeBits(char *rawData, size_t bitPos, llvm::APInt value)
Write value to byte-aligned position bitPos in rawData.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool isValidVectorTypeElementType(::mlir::Type type)
SliceVerificationResult
Enum that captures information related to verifier error conditions on slice insert/extract type of o...
constexpr T real(const NonFloatComplex< T > &x)
Definition Complex.h:255
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
TypeRange filterTypesOut(TypeRange types, const BitVector &indices, SmallVectorImpl< Type > &storage)
Filters out any elements referenced by indices.
constexpr T imag(const NonFloatComplex< T > &x)
Definition Complex.h:260
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
AffineExpr makeCanonicalStridedLayoutExpr(ArrayRef< int64_t > sizes, ArrayRef< AffineExpr > exprs, MLIRContext *context)
Given MemRef sizes that are either static or dynamic, returns the canonical "contiguous" strides Affi...
std::optional< llvm::SmallDenseSet< unsigned > > computeRankReductionMask(ArrayRef< int64_t > originalShape, ArrayRef< int64_t > reducedShape, bool matchDynamic=false)
Given an originalShape and a reducedShape assumed to be a subset of originalShape with some 1 entries...
AffineExpr simplifyAffineExpr(AffineExpr expr, unsigned numDims, unsigned numSymbols)
Simplify an affine expression by flattening and some amount of simple analysis.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
SliceVerificationResult isRankReducedType(ShapedType originalType, ShapedType candidateReducedType)
Check if originalType can be rank reduced to candidateReducedType type by dropping some dimensions wi...
TypeRange insertTypesInto(TypeRange oldTypes, ArrayRef< unsigned > indices, TypeRange newTypes, SmallVectorImpl< Type > &storage)
Insert a set of newTypes into oldTypes at the given indices.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
AffineExpr getAffineSymbolExpr(unsigned position, MLIRContext *context)