10#include "TypeDetail.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"
35#define GET_TYPEDEF_CLASSES
36#include "mlir/IR/BuiltinTypes.cpp.inc"
39#include "mlir/IR/BuiltinTypeConstraints.cpp.inc"
46void BuiltinDialect::registerTypes() {
48#define GET_TYPEDEF_LIST
49#include "mlir/IR/BuiltinTypes.cpp.inc"
61 return emitError() <<
"invalid element type for complex";
65size_t ComplexType::getDenseElementBitSize()
const {
67 return llvm::alignTo<8>(elemTy.getDenseElementBitSize()) * 2;
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));
82ComplexType::convertFromAttribute(
Attribute attr,
84 auto arrayAttr = dyn_cast<ArrayAttr>(attr);
85 if (!arrayAttr || arrayAttr.size() != 2)
89 if (
failed(elemTy.convertFromAttribute(arrayAttr[0], realData)))
91 if (
failed(elemTy.convertFromAttribute(arrayAttr[1], imagData)))
105 SignednessSemantics signedness) {
106 if (width > IntegerType::kMaxWidth) {
107 return emitError() <<
"integer bitwidth is limited to "
108 << IntegerType::kMaxWidth <<
" bits";
113unsigned IntegerType::getWidth()
const {
return getImpl()->width; }
115IntegerType::SignednessSemantics IntegerType::getSignedness()
const {
116 return getImpl()->signedness;
119IntegerType IntegerType::scaleElementBitwidth(
unsigned scale) {
121 return IntegerType();
122 return IntegerType::get(
getContext(), scale * getWidth(), getSignedness());
125size_t IntegerType::getDenseElementBitSize()
const {
131 APInt value = detail::readBits(rawData.data(), 0, getWidth());
132 return IntegerAttr::get(*
this, value);
136 size_t byteSize = llvm::divideCeil(apInt.getBitWidth(), CHAR_BIT);
137 size_t bitPos =
result.size() * CHAR_BIT;
143IntegerType::convertFromAttribute(
Attribute attr,
145 auto intAttr = dyn_cast<IntegerAttr>(attr);
146 if (!intAttr || intAttr.getType() != *
this)
156size_t IndexType::getDenseElementBitSize()
const {
157 return kInternalStorageBitWidth;
162 detail::readBits(rawData.data(), 0, kInternalStorageBitWidth);
163 return IntegerAttr::get(*
this, value);
167IndexType::convertFromAttribute(
Attribute attr,
169 auto intAttr = dyn_cast<IntegerAttr>(attr);
170 if (!intAttr || intAttr.getType() != *
this)
181#define FLOAT_TYPE_SEMANTICS(TYPE, SEM) \
182 const llvm::fltSemantics &TYPE::getFloatSemantics() const { \
183 return APFloat::SEM(); \
204#undef FLOAT_TYPE_SEMANTICS
206FloatType Float16Type::scaleElementBitwidth(
unsigned scale)
const {
214FloatType BFloat16Type::scaleElementBitwidth(
unsigned scale)
const {
222FloatType Float32Type::scaleElementBitwidth(
unsigned scale)
const {
232unsigned FunctionType::getNumInputs()
const {
return getImpl()->numInputs; }
235 return getImpl()->getInputs();
238unsigned FunctionType::getNumResults()
const {
return getImpl()->numResults; }
241 return getImpl()->getResults();
250FunctionType FunctionType::getWithArgsAndResults(
257 insertTypesInto(getResults(), resultIndices, resultTypes, resultStorage);
258 return clone(newArgTypes, newResultTypes);
263FunctionType::getWithoutArgsAndResults(
const BitVector &argIndices,
264 const BitVector &resultIndices) {
269 return clone(newArgTypes, newResultTypes);
276unsigned GraphType::getNumInputs()
const {
return getImpl()->numInputs; }
278ArrayRef<Type> GraphType::getInputs()
const {
return getImpl()->getInputs(); }
280unsigned GraphType::getNumResults()
const {
return getImpl()->numResults; }
282ArrayRef<Type> GraphType::getResults()
const {
return getImpl()->getResults(); }
298 insertTypesInto(getResults(), resultIndices, resultTypes, resultStorage);
299 return clone(newArgTypes, newResultTypes);
303GraphType GraphType::getWithoutArgsAndResults(
const BitVector &argIndices,
304 const BitVector &resultIndices) {
309 return clone(newArgTypes, newResultTypes);
317 StringAttr dialect, StringRef typeData) {
319 return emitError() <<
"invalid dialect namespace '" << dialect <<
"'";
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";
340bool VectorType::isValidElementType(
Type t) {
347 if (!isValidElementType(elementType))
349 <<
"vector elements must be int/index/float type but got "
352 if (any_of(shape, [](int64_t i) {
return i <= 0; }))
354 <<
"vector types must have positive constant sizes but got "
357 if (scalableDims.size() != shape.size())
358 return emitError() <<
"number of dims must match, got "
359 << scalableDims.size() <<
" and " << shape.size();
364VectorType VectorType::scaleElementBitwidth(
unsigned scale) {
368 if (
auto scaledEt = et.scaleElementBitwidth(scale))
369 return VectorType::get(
getShape(), scaledEt, getScalableDims());
371 if (
auto scaledEt = et.scaleElementBitwidth(scale))
372 return VectorType::get(
getShape(), scaledEt, getScalableDims());
377 Type elementType)
const {
378 return VectorType::get(shape.value_or(
getShape()), elementType,
388 .Case<RankedTensorType, UnrankedTensorType>(
389 [](
auto type) {
return type.getElementType(); });
393 return !llvm::isa<UnrankedTensorType>(*
this);
397 return llvm::cast<RankedTensorType>(*this).getShape();
401 Type elementType)
const {
402 if (llvm::dyn_cast<UnrankedTensorType>(*
this)) {
404 return RankedTensorType::get(*
shape, elementType);
405 return UnrankedTensorType::get(elementType);
408 auto rankedTy = llvm::cast<RankedTensorType>(*
this);
410 return RankedTensorType::get(rankedTy.getShape(), elementType,
411 rankedTy.getEncoding());
412 return RankedTensorType::get(
shape.value_or(rankedTy.getShape()), elementType,
413 rankedTy.getEncoding());
417 Type elementType)
const {
418 return ::llvm::cast<RankedTensorType>(
cloneWith(
shape, elementType));
430 return emitError() <<
"invalid tensor element type: " << elementType;
439 return llvm::isa<ComplexType, FloatType, IntegerType, OpaqueType, VectorType,
441 !llvm::isa<BuiltinDialect>(type.
getDialect());
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))
478 [](
auto type) {
return type.getElementType(); });
482 return !llvm::isa<UnrankedMemRefType>(*
this);
486 return llvm::cast<MemRefType>(*this).getShape();
490 Type elementType)
const {
491 if (llvm::dyn_cast<UnrankedMemRefType>(*
this)) {
506FailureOr<PtrLikeTypeInterface>
508 std::optional<Type> elementType)
const {
510 if (llvm::dyn_cast<UnrankedMemRefType>(*
this))
511 return cast<PtrLikeTypeInterface>(
512 UnrankedMemRefType::get(eTy, memorySpace));
517 return cast<PtrLikeTypeInterface>(
static_cast<MemRefType
>(builder));
521 Type elementType)
const {
530 if (
auto rankedMemRefTy = llvm::dyn_cast<MemRefType>(*
this))
531 return rankedMemRefTy.getMemorySpace();
532 return llvm::cast<UnrankedMemRefType>(*this).getMemorySpace();
536 if (
auto rankedMemRefTy = llvm::dyn_cast<MemRefType>(*
this))
537 return rankedMemRefTy.getMemorySpaceAsInt();
538 return llvm::cast<UnrankedMemRefType>(*this).getMemorySpaceAsInt();
545std::optional<llvm::SmallDenseSet<unsigned>>
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) {
554 int64_t origSize = originalShape[originalIdx];
556 if (matchDynamic && reducedIdx < reducedRank && origSize != 1 &&
557 (ShapedType::isDynamic(reducedShape[reducedIdx]) ||
558 ShapedType::isDynamic(origSize))) {
562 if (reducedIdx < reducedRank && origSize == reducedShape[reducedIdx]) {
567 unusedDims.insert(originalIdx);
574 if (reducedIdx != reducedRank)
581 ShapedType candidateReducedType) {
582 if (originalType == candidateReducedType)
585 ShapedType originalShapedType = llvm::cast<ShapedType>(originalType);
586 ShapedType candidateReducedShapedType =
587 llvm::cast<ShapedType>(candidateReducedType);
592 candidateReducedShapedType.getShape();
593 unsigned originalRank = originalShape.size(),
594 candidateReducedRank = candidateReducedShape.size();
595 if (candidateReducedRank > originalRank)
598 auto optionalUnusedDimsMask =
602 if (!optionalUnusedDimsMask)
605 if (originalShapedType.getElementType() !=
606 candidateReducedShapedType.getElementType())
614 if (memorySpace == 0)
617 return IntegerAttr::get(IntegerType::get(ctx, 64), memorySpace);
621 IntegerAttr intMemorySpace = llvm::dyn_cast_or_null<IntegerAttr>(memorySpace);
622 if (intMemorySpace && intMemorySpace.getValue() == 0)
632 assert(llvm::isa<IntegerAttr>(memorySpace) &&
633 "Using `getMemorySpaceInteger` with non-Integer attribute");
635 return static_cast<unsigned>(llvm::cast<IntegerAttr>(memorySpace).getInt());
638unsigned MemRefType::getMemorySpaceAsInt()
const {
643 MemRefLayoutAttrInterface layout,
657MemRefType MemRefType::getChecked(
659 Type elementType, MemRefLayoutAttrInterface layout,
Attribute memorySpace) {
670 elementType, layout, memorySpace);
682 auto layout = AffineMapAttr::get(map);
702 auto layout = AffineMapAttr::get(map);
708 elementType, layout, memorySpace);
712 AffineMap map,
unsigned memorySpaceInd) {
720 auto layout = AffineMapAttr::get(map);
733 unsigned memorySpaceInd) {
741 auto layout = AffineMapAttr::get(map);
748 elementType, layout, memorySpace);
753 MemRefLayoutAttrInterface layout,
756 return emitError() <<
"invalid memref element type";
759 for (int64_t s :
shape)
760 if (s < 0 && ShapedType::isStatic(s))
761 return emitError() <<
"invalid memref size";
763 assert(layout &&
"missing layout specification");
770bool MemRefType::areTrailingDimsContiguous(int64_t n) {
771 assert(n <= getRank() &&
772 "number of dimensions to check must not exceed rank");
773 return n <= getNumContiguousTrailingDims();
776int64_t MemRefType::getNumContiguousTrailingDims() {
777 const int64_t n = getRank();
780 if (getLayout().isIdentity())
798 int64_t dimProduct = 1;
799 for (int64_t i = n - 1; i >= 0; --i) {
802 if (strides[i] != dimProduct)
804 if (
shape[i] == ShapedType::kDynamic)
806 dimProduct *=
shape[i];
812MemRefType MemRefType::canonicalizeStridedLayout() {
813 AffineMap m = getLayout().getAffineMap();
825 if (
auto cst = llvm::dyn_cast<AffineConstantExpr>(m.
getResult(0)))
826 if (cst.getValue() == 0)
841 auto simplifiedLayoutExpr =
843 if (expr != simplifiedLayoutExpr)
846 simplifiedLayoutExpr)));
851 int64_t &offset)
const {
852 return getLayout().getStridesAndOffset(
getShape(), strides, offset);
855std::pair<SmallVector<int64_t>, int64_t>
856MemRefType::getStridesAndOffset()
const {
861 assert(succeeded(status) &&
"Invalid use of check-free getStridesAndOffset");
862 return {strides, offset};
865bool MemRefType::isStrided() {
869 return succeeded(res);
872bool MemRefType::isLastDimUnitStride() {
876 return succeeded(successStrides) && (strides.empty() || strides.back() == 1);
883unsigned UnrankedMemRefType::getMemorySpaceAsInt()
const {
891 return emitError() <<
"invalid memref element type";
901ArrayRef<Type> TupleType::getTypes()
const {
return getImpl()->getTypes(); }
908 for (
Type type : getTypes()) {
909 if (
auto nestedTuple = llvm::dyn_cast<TupleType>(type))
910 nestedTuple.getFlattenedTypes(types);
912 types.push_back(type);
917size_t TupleType::size()
const {
return getImpl()->size(); }
930 assert(!exprs.empty() &&
"expected exprs");
932 assert(!maps.empty() &&
"Expected one non-empty map");
933 unsigned numDims = maps[0].getNumDims(), nSymbols = maps[0].getNumSymbols();
936 bool dynamicPoisonBit =
false;
938 for (
auto en : llvm::zip(llvm::reverse(exprs), llvm::reverse(sizes))) {
939 int64_t size = std::get<1>(en);
944 expr = expr ? expr + dimExpr * stride : dimExpr * stride;
946 auto result = llvm::checkedMul(runningSize, size);
949 dynamicPoisonBit =
true;
954 dynamicPoisonBit =
true;
963 exprs.reserve(sizes.size());
964 for (
auto dim : llvm::seq<unsigned>(0, sizes.size()))
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)
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,...
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Base type for affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
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.
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.
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.
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.
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.
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.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Dialect & getDialect() const
Get the dialect this type is registered to.
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
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.
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)
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)
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
AffineExpr getAffineSymbolExpr(unsigned position, MLIRContext *context)