26#include "llvm/ADT/STLExtras.h"
27#include "llvm/ADT/SmallBitVector.h"
28#include "llvm/ADT/SmallVectorExtras.h"
38 return arith::ConstantOp::materialize(builder, value, type, loc);
51 auto cast = operand.get().getDefiningOp<CastOp>();
52 if (cast && operand.get() != inner &&
53 !llvm::isa<UnrankedMemRefType>(cast.getOperand().getType())) {
54 operand.set(cast.getOperand());
64 if (
auto memref = llvm::dyn_cast<MemRefType>(type))
65 return RankedTensorType::get(
memref.getShape(),
memref.getElementType());
66 if (
auto memref = llvm::dyn_cast<UnrankedMemRefType>(type))
67 return UnrankedTensorType::get(
memref.getElementType());
73 auto memrefType = llvm::cast<MemRefType>(value.
getType());
74 if (memrefType.isDynamicDim(dim))
75 return builder.
createOrFold<memref::DimOp>(loc, value, dim);
82 auto memrefType = llvm::cast<MemRefType>(value.
getType());
84 for (
int64_t i = 0; i < memrefType.getRank(); ++i)
101 assert(constValues.size() == values.size() &&
102 "incorrect number of const values");
103 for (
auto [i, cstVal] : llvm::enumerate(constValues)) {
105 if (ShapedType::isStatic(cstVal)) {
119static std::tuple<MemorySpaceCastOpInterface, PtrLikeTypeInterface, Type>
121 MemorySpaceCastOpInterface castOp =
122 MemorySpaceCastOpInterface::getIfPromotableCast(src);
130 FailureOr<PtrLikeTypeInterface> srcTy = resultTy.
clonePtrWith(
131 castOp.getSourcePtr().getType().getMemorySpace(), std::nullopt);
135 FailureOr<PtrLikeTypeInterface> tgtTy = resultTy.
clonePtrWith(
136 castOp.getTargetPtr().getType().getMemorySpace(), std::nullopt);
141 if (!castOp.isValidMemorySpaceCast(*tgtTy, *srcTy))
144 return std::make_tuple(castOp, *tgtTy, *srcTy);
149template <
typename ConcreteOpTy>
150static FailureOr<std::optional<SmallVector<Value>>>
160 llvm::append_range(operands, op->getOperands());
164 auto newOp = ConcreteOpTy::create(
165 builder, op.getLoc(),
TypeRange(resTy), operands, op.getProperties(),
166 op->getDiscardableAttrDictionary().getValue());
169 MemorySpaceCastOpInterface
result = castOp.cloneMemorySpaceCastOp(
172 return std::optional<SmallVector<Value>>(
180void AllocOp::getAsmResultNames(
182 setNameFn(getResult(),
"alloc");
185void AllocaOp::getAsmResultNames(
187 setNameFn(getResult(),
"alloca");
190template <
typename AllocLikeOp>
192 static_assert(llvm::is_one_of<AllocLikeOp, AllocOp, AllocaOp>::value,
193 "applies to only alloc or alloca");
194 auto memRefType = llvm::dyn_cast<MemRefType>(op.getResult().getType());
196 return op.emitOpError(
"result must be a memref");
201 unsigned numSymbols = 0;
202 if (!memRefType.getLayout().isIdentity())
203 numSymbols = memRefType.getLayout().getAffineMap().getNumSymbols();
204 if (op.getSymbolOperands().size() != numSymbols)
205 return op.emitOpError(
"symbol operand count does not equal memref symbol "
207 << numSymbols <<
", got " << op.getSymbolOperands().size();
214LogicalResult AllocaOp::verify() {
218 "requires an ancestor op with AutomaticAllocationScope trait");
225template <
typename AllocLikeOp>
227 using OpRewritePattern<AllocLikeOp>::OpRewritePattern;
229 LogicalResult matchAndRewrite(AllocLikeOp alloc,
230 PatternRewriter &rewriter)
const override {
233 if (llvm::none_of(alloc.getDynamicSizes(), [](Value operand) {
235 if (!matchPattern(operand, m_ConstantInt(&constSizeArg)))
237 return constSizeArg.isNonNegative();
241 auto memrefType = alloc.getType();
245 SmallVector<int64_t, 4> newShapeConstants;
246 newShapeConstants.reserve(memrefType.getRank());
247 SmallVector<Value, 4> dynamicSizes;
249 unsigned dynamicDimPos = 0;
250 for (
unsigned dim = 0, e = memrefType.getRank(); dim < e; ++dim) {
251 int64_t dimSize = memrefType.getDimSize(dim);
253 if (ShapedType::isStatic(dimSize)) {
254 newShapeConstants.push_back(dimSize);
257 auto dynamicSize = alloc.getDynamicSizes()[dynamicDimPos];
260 constSizeArg.isNonNegative()) {
262 newShapeConstants.push_back(constSizeArg.getZExtValue());
265 newShapeConstants.push_back(ShapedType::kDynamic);
266 dynamicSizes.push_back(dynamicSize);
272 MemRefType newMemRefType =
273 MemRefType::Builder(memrefType).setShape(newShapeConstants);
274 assert(dynamicSizes.size() == newMemRefType.getNumDynamicDims());
277 auto newAlloc = AllocLikeOp::create(rewriter, alloc.getLoc(), newMemRefType,
278 dynamicSizes, alloc.getSymbolOperands(),
279 alloc.getAlignmentAttr());
289 using OpRewritePattern<T>::OpRewritePattern;
291 LogicalResult matchAndRewrite(T alloc,
292 PatternRewriter &rewriter)
const override {
293 if (llvm::any_of(alloc->getUsers(), [&](Operation *op) {
294 if (auto storeOp = dyn_cast<StoreOp>(op))
295 return storeOp.getValue() == alloc;
296 return !isa<DeallocOp>(op);
300 for (Operation *user : llvm::make_early_inc_range(alloc->getUsers()))
311 results.
add<SimplifyAllocConst<AllocOp>, SimplifyDeadAlloc<AllocOp>>(context);
316 results.
add<SimplifyAllocConst<AllocaOp>, SimplifyDeadAlloc<AllocaOp>>(
324LogicalResult ReallocOp::verify() {
325 auto sourceType = llvm::cast<MemRefType>(getOperand(0).
getType());
326 MemRefType resultType =
getType();
329 if (!sourceType.getLayout().isIdentity())
330 return emitError(
"unsupported layout for source memref type ")
334 if (!resultType.getLayout().isIdentity())
335 return emitError(
"unsupported layout for result memref type ")
339 if (sourceType.getMemorySpace() != resultType.getMemorySpace())
340 return emitError(
"different memory spaces specified for source memref "
342 << sourceType <<
" and result memref type " << resultType;
350 if (resultType.getNumDynamicDims() && !getDynamicResultSize())
351 return emitError(
"missing dimension operand for result type ")
353 if (!resultType.getNumDynamicDims() && getDynamicResultSize())
354 return emitError(
"unnecessary dimension operand for result type ")
362 results.
add<SimplifyDeadAlloc<ReallocOp>>(context);
370 bool printBlockTerminators =
false;
373 if (!getResults().empty()) {
374 p <<
" -> (" << getResultTypes() <<
")";
375 printBlockTerminators =
true;
380 printBlockTerminators);
386 result.regions.reserve(1);
396 AllocaScopeOp::ensureTerminator(*bodyRegion, parser.
getBuilder(),
406void AllocaScopeOp::getSuccessorRegions(
423 MemoryEffectOpInterface
interface = dyn_cast<MemoryEffectOpInterface>(op);
428 interface.getEffectOnValue<MemoryEffects::Allocate>(res)) {
429 if (isa<SideEffects::AutomaticAllocationScopeResource>(
430 effect->getResource()))
446 MemoryEffectOpInterface
interface = dyn_cast<MemoryEffectOpInterface>(op);
451 interface.getEffectOnValue<MemoryEffects::Allocate>(res)) {
452 if (isa<SideEffects::AutomaticAllocationScopeResource>(
453 effect->getResource()))
477 bool hasPotentialAlloca =
490 if (hasPotentialAlloca) {
523 if (!lastParentWithoutScope ||
536 lastParentWithoutScope = lastParentWithoutScope->
getParentOp();
537 if (!lastParentWithoutScope ||
544 Region *containingRegion =
nullptr;
545 for (
auto &r : lastParentWithoutScope->
getRegions()) {
546 if (r.isAncestor(op->getParentRegion())) {
547 assert(containingRegion ==
nullptr &&
548 "only one region can contain the op");
549 containingRegion = &r;
552 assert(containingRegion &&
"op must be contained in a region");
562 return containingRegion->isAncestor(v.getParentRegion());
565 toHoist.push_back(alloc);
572 for (
auto *op : toHoist) {
573 auto *cloned = rewriter.
clone(*op);
574 rewriter.
replaceOp(op, cloned->getResults());
589LogicalResult AssumeAlignmentOp::verify() {
590 if (!llvm::isPowerOf2_32(getAlignment()))
591 return emitOpError(
"alignment must be power of 2");
595void AssumeAlignmentOp::getAsmResultNames(
597 setNameFn(getResult(),
"assume_align");
600OpFoldResult AssumeAlignmentOp::fold(FoldAdaptor adaptor) {
601 auto source = getMemref().getDefiningOp<AssumeAlignmentOp>();
604 if (source.getAlignment() != getAlignment())
609FailureOr<std::optional<SmallVector<Value>>>
610AssumeAlignmentOp::bubbleDownCasts(
OpBuilder &builder) {
614FailureOr<OpFoldResult> AssumeAlignmentOp::reifyDimOfResult(
OpBuilder &builder,
617 assert(resultIndex == 0 &&
"AssumeAlignmentOp has a single result");
618 return getMixedSize(builder, getLoc(), getMemref(), dim);
625LogicalResult DistinctObjectsOp::verify() {
626 if (getOperandTypes() != getResultTypes())
627 return emitOpError(
"operand types and result types must match");
629 if (getOperandTypes().empty())
630 return emitOpError(
"expected at least one operand");
635LogicalResult DistinctObjectsOp::inferReturnTypes(
640 llvm::copy(operands.
getTypes(), std::back_inserter(inferredReturnTypes));
649 setNameFn(getResult(),
"cast");
689bool CastOp::canFoldIntoConsumerOp(CastOp castOp) {
690 MemRefType sourceType =
691 llvm::dyn_cast<MemRefType>(castOp.getSource().getType());
692 MemRefType resultType = llvm::dyn_cast<MemRefType>(castOp.getType());
695 if (!sourceType || !resultType)
699 if (sourceType.getElementType() != resultType.getElementType())
703 if (sourceType.getRank() != resultType.getRank())
707 int64_t sourceOffset, resultOffset;
709 if (
failed(sourceType.getStridesAndOffset(sourceStrides, sourceOffset)) ||
710 failed(resultType.getStridesAndOffset(resultStrides, resultOffset)))
714 for (
auto it : llvm::zip(sourceType.getShape(), resultType.getShape())) {
715 auto ss = std::get<0>(it), st = std::get<1>(it);
717 if (ShapedType::isDynamic(ss) && ShapedType::isStatic(st))
722 if (sourceOffset != resultOffset)
723 if (ShapedType::isDynamic(sourceOffset) &&
724 ShapedType::isStatic(resultOffset))
728 for (
auto it : llvm::zip(sourceStrides, resultStrides)) {
729 auto ss = std::get<0>(it), st = std::get<1>(it);
731 if (ShapedType::isDynamic(ss) && ShapedType::isStatic(st))
739 if (inputs.size() != 1 || outputs.size() != 1)
741 if (inputs == outputs)
743 Type a = inputs.front(),
b = outputs.front();
744 auto aT = llvm::dyn_cast<MemRefType>(a);
745 auto bT = llvm::dyn_cast<MemRefType>(
b);
747 auto uaT = llvm::dyn_cast<UnrankedMemRefType>(a);
748 auto ubT = llvm::dyn_cast<UnrankedMemRefType>(
b);
751 if (aT.getElementType() != bT.getElementType())
753 if (aT.getLayout() != bT.getLayout()) {
756 if (
failed(aT.getStridesAndOffset(aStrides, aOffset)) ||
757 failed(bT.getStridesAndOffset(bStrides, bOffset)) ||
758 aStrides.size() != bStrides.size())
767 return (ShapedType::isDynamic(a) || ShapedType::isDynamic(
b) || a ==
b);
769 if (!checkCompatible(aOffset, bOffset))
772 if (aT.getDimSize(
index) == 1 || bT.getDimSize(
index) == 1)
774 if (!checkCompatible(aStride, bStrides[
index]))
778 if (aT.getMemorySpace() != bT.getMemorySpace())
782 if (aT.getRank() != bT.getRank())
785 for (
unsigned i = 0, e = aT.getRank(); i != e; ++i) {
786 int64_t aDim = aT.getDimSize(i), bDim = bT.getDimSize(i);
787 if (ShapedType::isStatic(aDim) && ShapedType::isStatic(bDim) &&
801 auto aEltType = (aT) ? aT.getElementType() : uaT.getElementType();
802 auto bEltType = (bT) ? bT.getElementType() : ubT.getElementType();
803 if (aEltType != bEltType)
806 auto aMemSpace = (aT) ? aT.getMemorySpace() : uaT.getMemorySpace();
807 auto bMemSpace = (bT) ? bT.getMemorySpace() : ubT.getMemorySpace();
808 return aMemSpace == bMemSpace;
818FailureOr<std::optional<SmallVector<Value>>>
819CastOp::bubbleDownCasts(
OpBuilder &builder) {
831 using OpRewritePattern<CopyOp>::OpRewritePattern;
833 LogicalResult matchAndRewrite(CopyOp copyOp,
834 PatternRewriter &rewriter)
const override {
835 if (copyOp.getSource() != copyOp.getTarget())
844 using OpRewritePattern<CopyOp>::OpRewritePattern;
846 static bool isEmptyMemRef(BaseMemRefType type) {
850 LogicalResult matchAndRewrite(CopyOp copyOp,
851 PatternRewriter &rewriter)
const override {
852 if (isEmptyMemRef(copyOp.getSource().getType()) ||
853 isEmptyMemRef(copyOp.getTarget().getType())) {
865 results.
add<FoldEmptyCopy, FoldSelfCopy>(context);
872 for (
OpOperand &operand : op->getOpOperands()) {
874 if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) {
875 operand.set(castOp.getOperand());
882LogicalResult CopyOp::fold(FoldAdaptor adaptor,
883 SmallVectorImpl<OpFoldResult> &results) {
893LogicalResult DeallocOp::fold(FoldAdaptor adaptor,
894 SmallVectorImpl<OpFoldResult> &results) {
903void DimOp::getAsmResultNames(
function_ref<
void(Value, StringRef)> setNameFn) {
904 setNameFn(getResult(),
"dim");
907void DimOp::build(OpBuilder &builder, OperationState &
result, Value source,
909 auto loc =
result.location;
911 build(builder,
result, source, indexValue);
914std::optional<int64_t> DimOp::getConstantIndex() {
923 auto rankedSourceType = dyn_cast<MemRefType>(getSource().
getType());
924 if (!rankedSourceType)
927 if (rankedSourceType.getRank() <= constantIndex)
933void DimOp::inferResultRangesFromOptional(ArrayRef<IntegerValueRange> argRanges,
935 setResultRange(getResult(),
944 std::map<int64_t, unsigned> numOccurences;
945 for (
auto val : vals)
946 numOccurences[val]++;
947 return numOccurences;
957static FailureOr<llvm::SmallBitVector>
959 MemRefType reducedType,
961 int64_t rankReduction = originalType.getRank() - reducedType.getRank();
962 if (rankReduction <= 0)
963 return llvm::SmallBitVector(originalType.getRank());
967 for (
const auto &it : llvm::enumerate(sizes)) {
969 sourceSizes[it.index()] = *cst;
971 sourceSizes[it.index()] = ShapedType::kDynamic;
975 llvm::SmallBitVector usedSourceDims(originalType.getRank());
977 for (
int64_t resultSize : resultSizes) {
978 bool matched =
false;
979 for (
int64_t j = startJ;
j < originalType.getRank(); ++
j) {
980 if (sourceSizes[
j] == resultSize) {
981 usedSourceDims.set(
j);
991 llvm::SmallBitVector unusedDims(originalType.getRank());
992 for (
int64_t i = 0; i < originalType.getRank(); ++i)
993 if (!usedSourceDims.test(i))
1006 MemRefType originalType, MemRefType reducedType,
1008 llvm::SmallBitVector unusedDims) {
1016 std::map<int64_t, unsigned> currUnaccountedStrides =
1018 std::map<int64_t, unsigned> candidateStridesNumOccurences =
1020 for (
size_t dim = 0, e = unusedDims.size(); dim != e; ++dim) {
1021 if (!unusedDims.test(dim))
1023 int64_t originalStride = originalStrides[dim];
1024 if (currUnaccountedStrides[originalStride] >
1025 candidateStridesNumOccurences[originalStride]) {
1027 currUnaccountedStrides[originalStride]--;
1030 if (currUnaccountedStrides[originalStride] ==
1031 candidateStridesNumOccurences[originalStride]) {
1033 unusedDims.reset(dim);
1036 if (currUnaccountedStrides[originalStride] <
1037 candidateStridesNumOccurences[originalStride]) {
1043 if (
static_cast<int64_t>(unusedDims.count()) + reducedType.getRank() !=
1044 originalType.getRank())
1056static FailureOr<llvm::SmallBitVector>
1059 llvm::SmallBitVector unusedDims(originalType.getRank());
1060 if (originalType.getRank() == reducedType.getRank())
1063 for (
const auto &dim : llvm::enumerate(sizes))
1064 if (
auto attr = llvm::dyn_cast_if_present<Attribute>(dim.value()))
1065 if (llvm::cast<IntegerAttr>(attr).getInt() == 1)
1066 unusedDims.set(dim.index());
1070 if (
static_cast<int64_t>(unusedDims.count()) + reducedType.getRank() ==
1071 originalType.getRank())
1075 int64_t originalOffset, candidateOffset;
1077 originalType.getStridesAndOffset(originalStrides, originalOffset)) ||
1079 reducedType.getStridesAndOffset(candidateStrides, candidateOffset)))
1087 if (strides.size() <= 1)
1089 return llvm::any_of(strides.drop_back(),
1090 [](
int64_t s) { return !ShapedType::isDynamic(s); });
1092 if (hasNonTrivialStaticStride(originalStrides) ||
1093 hasNonTrivialStaticStride(candidateStrides)) {
1094 FailureOr<llvm::SmallBitVector> strideBased =
1097 candidateStrides, unusedDims);
1098 if (succeeded(strideBased))
1099 return *strideBased;
1105llvm::SmallBitVector SubViewOp::getDroppedDims() {
1106 MemRefType sourceType = getSourceType();
1107 MemRefType resultType =
getType();
1108 FailureOr<llvm::SmallBitVector> unusedDims =
1110 assert(succeeded(unusedDims) &&
"unable to find unused dims of subview");
1114OpFoldResult DimOp::fold(FoldAdaptor adaptor) {
1116 std::optional<int64_t> index = getConstantIndex();
1121 auto memrefType = llvm::dyn_cast<MemRefType>(getSource().
getType());
1127 int64_t indexVal = index.value();
1128 if (indexVal < 0 || indexVal >= memrefType.getRank())
1132 if (!memrefType.isDynamicDim(indexVal)) {
1134 return builder.
getIndexAttr(memrefType.getShape()[indexVal]);
1139 Operation *definingOp = getSource().getDefiningOp();
1141 if (
auto alloc = dyn_cast_or_null<AllocOp>(definingOp))
1142 return *(alloc.getDynamicSizes().begin() +
1143 memrefType.getDynamicDimIndex(indexVal));
1145 if (
auto alloca = dyn_cast_or_null<AllocaOp>(definingOp))
1146 return *(alloca.getDynamicSizes().begin() +
1147 memrefType.getDynamicDimIndex(indexVal));
1149 if (
auto view = dyn_cast_or_null<ViewOp>(definingOp))
1150 return *(view.getDynamicSizes().begin() +
1151 memrefType.getDynamicDimIndex(indexVal));
1153 if (
auto subview = dyn_cast_or_null<SubViewOp>(definingOp)) {
1158 unsigned dynamicResultDimIdx = memrefType.getDynamicDimIndex(indexVal);
1159 unsigned dynamicIdx = 0;
1160 for (OpFoldResult size : subview.getMixedSizes()) {
1161 if (llvm::isa<Attribute>(size))
1163 if (dynamicIdx == dynamicResultDimIdx)
1180struct DimOfMemRefReshape :
public OpRewritePattern<DimOp> {
1181 using OpRewritePattern<DimOp>::OpRewritePattern;
1183 LogicalResult matchAndRewrite(DimOp dim,
1184 PatternRewriter &rewriter)
const override {
1185 auto reshape = dim.getSource().getDefiningOp<ReshapeOp>();
1189 dim,
"Dim op is not defined by a reshape op.");
1200 if (dim.getIndex().getParentBlock() == reshape->getBlock()) {
1201 if (
auto *definingOp = dim.getIndex().getDefiningOp()) {
1202 if (reshape->isBeforeInBlock(definingOp)) {
1205 "dim.getIndex is not defined before reshape in the same block.");
1210 else if (dim->getBlock() != reshape->getBlock() &&
1211 !dim.getIndex().getParentRegion()->isProperAncestor(
1212 reshape->getParentRegion())) {
1217 dim,
"dim.getIndex does not dominate reshape.");
1223 Location loc = dim.getLoc();
1225 LoadOp::create(rewriter, loc, reshape.getShape(), dim.getIndex());
1226 if (
load.getType() != dim.getType())
1227 load = arith::IndexCastOp::create(rewriter, loc, dim.getType(),
load);
1235void DimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1236 MLIRContext *context) {
1237 results.
add<DimOfMemRefReshape>(context);
1244void DmaStartOp::build(OpBuilder &builder, OperationState &
result,
1245 Value srcMemRef,
ValueRange srcIndices, Value destMemRef,
1247 Value tagMemRef,
ValueRange tagIndices, Value stride,
1248 Value elementsPerStride) {
1249 result.addOperands(srcMemRef);
1250 result.addOperands(srcIndices);
1251 result.addOperands(destMemRef);
1252 result.addOperands(destIndices);
1253 result.addOperands({numElements, tagMemRef});
1254 result.addOperands(tagIndices);
1256 result.addOperands({stride, elementsPerStride});
1259void DmaStartOp::print(OpAsmPrinter &p) {
1260 p <<
" " << getSrcMemRef() <<
'[' << getSrcIndices() <<
"], "
1261 << getDstMemRef() <<
'[' << getDstIndices() <<
"], " <<
getNumElements()
1262 <<
", " << getTagMemRef() <<
'[' << getTagIndices() <<
']';
1264 p <<
", " << getStride() <<
", " << getNumElementsPerStride();
1267 p <<
" : " << getSrcMemRef().getType() <<
", " << getDstMemRef().getType()
1268 <<
", " << getTagMemRef().getType();
1279ParseResult DmaStartOp::parse(OpAsmParser &parser, OperationState &
result) {
1280 OpAsmParser::UnresolvedOperand srcMemRefInfo;
1281 SmallVector<OpAsmParser::UnresolvedOperand, 4> srcIndexInfos;
1282 OpAsmParser::UnresolvedOperand dstMemRefInfo;
1283 SmallVector<OpAsmParser::UnresolvedOperand, 4> dstIndexInfos;
1284 OpAsmParser::UnresolvedOperand numElementsInfo;
1285 OpAsmParser::UnresolvedOperand tagMemrefInfo;
1286 SmallVector<OpAsmParser::UnresolvedOperand, 4> tagIndexInfos;
1287 SmallVector<OpAsmParser::UnresolvedOperand, 2> strideInfo;
1289 SmallVector<Type, 3> types;
1309 bool isStrided = strideInfo.size() == 2;
1310 if (!strideInfo.empty() && !isStrided) {
1312 "expected two stride related operands");
1317 if (types.size() != 3)
1339LogicalResult DmaStartOp::verify() {
1344 if (numOperands < 4)
1345 return emitOpError(
"expected at least 4 operands");
1350 if (!llvm::isa<MemRefType>(getSrcMemRef().
getType()))
1351 return emitOpError(
"expected source to be of memref type");
1352 if (numOperands < getSrcMemRefRank() + 4)
1353 return emitOpError() <<
"expected at least " << getSrcMemRefRank() + 4
1355 if (!getSrcIndices().empty() &&
1356 !llvm::all_of(getSrcIndices().getTypes(),
1357 [](Type t) {
return t.
isIndex(); }))
1358 return emitOpError(
"expected source indices to be of index type");
1361 if (!llvm::isa<MemRefType>(getDstMemRef().
getType()))
1362 return emitOpError(
"expected destination to be of memref type");
1363 unsigned numExpectedOperands = getSrcMemRefRank() + getDstMemRefRank() + 4;
1364 if (numOperands < numExpectedOperands)
1365 return emitOpError() <<
"expected at least " << numExpectedOperands
1367 if (!getDstIndices().empty() &&
1368 !llvm::all_of(getDstIndices().getTypes(),
1369 [](Type t) {
return t.
isIndex(); }))
1370 return emitOpError(
"expected destination indices to be of index type");
1374 return emitOpError(
"expected num elements to be of index type");
1377 if (!llvm::isa<MemRefType>(getTagMemRef().
getType()))
1378 return emitOpError(
"expected tag to be of memref type");
1379 numExpectedOperands += getTagMemRefRank();
1380 if (numOperands < numExpectedOperands)
1381 return emitOpError() <<
"expected at least " << numExpectedOperands
1383 if (!getTagIndices().empty() &&
1384 !llvm::all_of(getTagIndices().getTypes(),
1385 [](Type t) {
return t.
isIndex(); }))
1386 return emitOpError(
"expected tag indices to be of index type");
1390 if (numOperands != numExpectedOperands &&
1391 numOperands != numExpectedOperands + 2)
1392 return emitOpError(
"incorrect number of operands");
1396 if (!getStride().
getType().isIndex() ||
1397 !getNumElementsPerStride().
getType().isIndex())
1399 "expected stride and num elements per stride to be of type index");
1405LogicalResult DmaStartOp::fold(FoldAdaptor adaptor,
1406 SmallVectorImpl<OpFoldResult> &results) {
1411void DmaStartOp::setMemrefsAndIndices(RewriterBase &rewriter, Value newSrc,
1415 SmallVector<Value> newOperands;
1416 newOperands.push_back(newSrc);
1417 llvm::append_range(newOperands, newSrcIndices);
1418 newOperands.push_back(newDst);
1419 llvm::append_range(newOperands, newDstIndices);
1421 newOperands.push_back(getTagMemRef());
1422 llvm::append_range(newOperands, getTagIndices());
1424 newOperands.push_back(getStride());
1425 newOperands.push_back(getNumElementsPerStride());
1428 rewriter.
modifyOpInPlace(*
this, [&]() { (*this)->setOperands(newOperands); });
1435LogicalResult DmaWaitOp::fold(FoldAdaptor adaptor,
1436 SmallVectorImpl<OpFoldResult> &results) {
1441LogicalResult DmaWaitOp::verify() {
1443 unsigned numTagIndices = getTagIndices().size();
1444 unsigned tagMemRefRank = getTagMemRefRank();
1445 if (numTagIndices != tagMemRefRank)
1446 return emitOpError() <<
"expected tagIndices to have the same number of "
1447 "elements as the tagMemRef rank, expected "
1448 << tagMemRefRank <<
", but got " << numTagIndices;
1456void ExtractAlignedPointerAsIndexOp::getAsmResultNames(
1458 setNameFn(getResult(),
"intptr");
1467LogicalResult ExtractStridedMetadataOp::inferReturnTypes(
1468 MLIRContext *context, std::optional<Location> location,
1469 ExtractStridedMetadataOp::Adaptor adaptor,
1470 SmallVectorImpl<Type> &inferredReturnTypes) {
1471 auto sourceType = llvm::dyn_cast<MemRefType>(adaptor.getSource().getType());
1475 unsigned sourceRank = sourceType.getRank();
1476 IndexType indexType = IndexType::get(context);
1478 MemRefType::get({}, sourceType.getElementType(),
1479 MemRefLayoutAttrInterface{}, sourceType.getMemorySpace());
1481 inferredReturnTypes.push_back(memrefType);
1483 inferredReturnTypes.push_back(indexType);
1485 for (
unsigned i = 0; i < sourceRank * 2; ++i)
1486 inferredReturnTypes.push_back(indexType);
1490void ExtractStridedMetadataOp::getAsmResultNames(
1492 setNameFn(getBaseBuffer(),
"base_buffer");
1493 setNameFn(getOffset(),
"offset");
1496 if (!getSizes().empty()) {
1497 setNameFn(getSizes().front(),
"sizes");
1498 setNameFn(getStrides().front(),
"strides");
1505template <
typename Container>
1509 assert(values.size() == maybeConstants.size() &&
1510 " expected values and maybeConstants of the same size");
1511 bool atLeastOneReplacement =
false;
1512 for (
auto [maybeConstant,
result] : llvm::zip(maybeConstants, values)) {
1517 assert(isa<Attribute>(maybeConstant) &&
1518 "The constified value should be either unchanged (i.e., == result) "
1522 llvm::cast<IntegerAttr>(cast<Attribute>(maybeConstant)).getInt());
1527 atLeastOneReplacement =
true;
1530 return atLeastOneReplacement;
1534ExtractStridedMetadataOp::fold(FoldAdaptor adaptor,
1535 SmallVectorImpl<OpFoldResult> &results) {
1536 OpBuilder builder(*
this);
1540 getConstifiedMixedOffset());
1542 getConstifiedMixedSizes());
1544 builder, getLoc(), getStrides(), getConstifiedMixedStrides());
1547 if (
auto prev = getSource().getDefiningOp<CastOp>())
1548 if (isa<MemRefType>(prev.getSource().getType())) {
1549 getSourceMutable().assign(prev.getSource());
1550 atLeastOneReplacement =
true;
1553 return success(atLeastOneReplacement);
1556SmallVector<OpFoldResult> ExtractStridedMetadataOp::getConstifiedMixedSizes() {
1562SmallVector<OpFoldResult>
1563ExtractStridedMetadataOp::getConstifiedMixedStrides() {
1565 SmallVector<int64_t> staticValues;
1567 LogicalResult status =
1568 getSource().getType().getStridesAndOffset(staticValues, unused);
1570 assert(succeeded(status) &&
"could not get strides from type");
1575OpFoldResult ExtractStridedMetadataOp::getConstifiedMixedOffset() {
1577 SmallVector<OpFoldResult> values(1, offsetOfr);
1578 SmallVector<int64_t> staticValues, unused;
1580 LogicalResult status =
1581 getSource().getType().getStridesAndOffset(unused, offset);
1583 assert(succeeded(status) &&
"could not get offset from type");
1584 staticValues.push_back(offset);
1593void GenericAtomicRMWOp::build(OpBuilder &builder, OperationState &
result,
1595 OpBuilder::InsertionGuard g(builder);
1596 result.addOperands(memref);
1599 if (
auto memrefType = llvm::dyn_cast<MemRefType>(memref.
getType())) {
1600 Type elementType = memrefType.getElementType();
1601 result.addTypes(elementType);
1603 Region *bodyRegion =
result.addRegion();
1609LogicalResult GenericAtomicRMWOp::verify() {
1610 auto &body = getRegion();
1611 if (body.getNumArguments() != 1)
1612 return emitOpError(
"expected single number of entry block arguments");
1614 if (getResult().
getType() != body.getArgument(0).getType())
1615 return emitOpError(
"expected block argument of the same type result type");
1618 body.walk([&](Operation *nestedOp) {
1622 "body of 'memref.generic_atomic_rmw' should contain "
1623 "only operations with no side effects");
1630ParseResult GenericAtomicRMWOp::parse(OpAsmParser &parser,
1631 OperationState &
result) {
1632 OpAsmParser::UnresolvedOperand memref;
1634 SmallVector<OpAsmParser::UnresolvedOperand, 4> ivs;
1644 Region *body =
result.addRegion();
1652void GenericAtomicRMWOp::print(OpAsmPrinter &p) {
1653 p <<
' ' << getMemref() <<
"[" <<
getIndices()
1654 <<
"] : " << getMemref().
getType() <<
' ';
1663std::optional<SmallVector<Value>> GenericAtomicRMWOp::updateMemrefAndIndices(
1664 RewriterBase &rewriter, Value newMemref,
ValueRange newIndices) {
1666 getMemrefMutable().assign(newMemref);
1667 getIndicesMutable().assign(newIndices);
1669 return std::nullopt;
1676LogicalResult AtomicYieldOp::verify() {
1677 Type parentType = (*this)->getParentOp()->getResultTypes().front();
1678 Type resultType = getResult().getType();
1679 if (parentType != resultType)
1680 return emitOpError() <<
"types mismatch between yield op: " << resultType
1681 <<
" and its parent: " << parentType;
1693 if (!op.isExternal()) {
1695 if (op.isUninitialized())
1696 p <<
"uninitialized";
1709 auto memrefType = llvm::dyn_cast<MemRefType>(type);
1710 if (!memrefType || !memrefType.hasStaticShape())
1712 <<
"type should be static shaped memref, but got " << type;
1713 typeAttr = TypeAttr::get(type);
1719 initialValue = UnitAttr::get(parser.
getContext());
1726 if (!llvm::isa<ElementsAttr>(initialValue))
1728 <<
"initial value should be a unit or elements attribute";
1732LogicalResult GlobalOp::verify() {
1733 auto memrefType = llvm::dyn_cast<MemRefType>(
getType());
1734 if (!memrefType || !memrefType.hasStaticShape())
1735 return emitOpError(
"type should be static shaped memref, but got ")
1740 if (getInitialValue().has_value()) {
1741 Attribute initValue = getInitialValue().value();
1742 if (!llvm::isa<UnitAttr>(initValue) && !llvm::isa<ElementsAttr>(initValue))
1743 return emitOpError(
"initial value should be a unit or elements "
1744 "attribute, but got ")
1749 if (
auto elementsAttr = llvm::dyn_cast<ElementsAttr>(initValue)) {
1751 auto initElementType =
1752 cast<TensorType>(elementsAttr.getType()).getElementType();
1753 auto memrefElementType = memrefType.getElementType();
1755 if (initElementType != memrefElementType)
1756 return emitOpError(
"initial value element expected to be of type ")
1757 << memrefElementType <<
", but was of type " << initElementType;
1762 auto initShape = elementsAttr.getShapedType().getShape();
1763 auto memrefShape = memrefType.getShape();
1764 if (initShape != memrefShape)
1765 return emitOpError(
"initial value shape expected to be ")
1766 << memrefShape <<
" but was " << initShape;
1774ElementsAttr GlobalOp::getConstantInitValue() {
1775 auto initVal = getInitialValue();
1776 if (getConstant() && initVal.has_value())
1777 return llvm::cast<ElementsAttr>(initVal.value());
1786GetGlobalOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
1792 return emitOpError(
"'")
1793 << getName() <<
"' does not reference a valid global memref";
1795 Type resultType = getResult().getType();
1796 if (global.getType() != resultType)
1797 return emitOpError(
"result type ")
1798 << resultType <<
" does not match type " << global.getType()
1799 <<
" of the global memref @" << getName();
1811 result = dyn_cast<BoolAttr>(attr);
1814 "expected boolean attribute");
1822OpFoldResult LoadOp::fold(FoldAdaptor adaptor) {
1828 auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>();
1834 getGlobalOp, getGlobalOp.getNameAttr());
1839 dyn_cast_or_null<SplatElementsAttr>(global.getConstantInitValue());
1843 return splatAttr.getSplatValue<Attribute>();
1848std::optional<SmallVector<Value>>
1849LoadOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
1852 getMemrefMutable().assign(newMemref);
1853 getIndicesMutable().assign(newIndices);
1855 return std::nullopt;
1858FailureOr<std::optional<SmallVector<Value>>>
1859LoadOp::bubbleDownCasts(OpBuilder &builder) {
1868void MemorySpaceCastOp::getAsmResultNames(
1870 setNameFn(getResult(),
"memspacecast");
1874 if (inputs.size() != 1 || outputs.size() != 1)
1876 Type a = inputs.front(),
b = outputs.front();
1877 auto aT = llvm::dyn_cast<MemRefType>(a);
1878 auto bT = llvm::dyn_cast<MemRefType>(
b);
1880 auto uaT = llvm::dyn_cast<UnrankedMemRefType>(a);
1881 auto ubT = llvm::dyn_cast<UnrankedMemRefType>(
b);
1884 if (aT.getElementType() != bT.getElementType())
1886 if (aT.getLayout() != bT.getLayout())
1888 if (aT.getShape() != bT.getShape())
1893 return uaT.getElementType() == ubT.getElementType();
1898OpFoldResult MemorySpaceCastOp::fold(FoldAdaptor adaptor) {
1901 if (
auto parentCast = getSource().getDefiningOp<MemorySpaceCastOp>()) {
1902 getSourceMutable().assign(parentCast.getSource());
1916bool MemorySpaceCastOp::isValidMemorySpaceCast(PtrLikeTypeInterface tgt,
1917 PtrLikeTypeInterface src) {
1918 return isa<BaseMemRefType>(tgt) &&
1919 tgt.clonePtrWith(src.getMemorySpace(), std::nullopt) == src;
1922MemorySpaceCastOpInterface MemorySpaceCastOp::cloneMemorySpaceCastOp(
1923 OpBuilder &
b, PtrLikeTypeInterface tgt,
1925 assert(isValidMemorySpaceCast(tgt, src.getType()) &&
"invalid arguments");
1926 return MemorySpaceCastOp::create(
b, getLoc(), tgt, src);
1930bool MemorySpaceCastOp::isSourcePromotable() {
1931 return getDest().getType().getMemorySpace() ==
nullptr;
1938void PrefetchOp::print(OpAsmPrinter &p) {
1939 p <<
" " << getMemref() <<
'[';
1941 p <<
']' <<
", " << (getIsWrite() ?
"write" :
"read");
1942 p <<
", locality<" << getLocalityHint();
1943 p <<
">, " << (getIsDataCache() ?
"data" :
"instr");
1945 (*this)->getDiscardableAttrDictionary(),
1946 {
"localityHint",
"isWrite",
"isDataCache"});
1950ParseResult PrefetchOp::parse(OpAsmParser &parser, OperationState &
result) {
1951 OpAsmParser::UnresolvedOperand memrefInfo;
1952 SmallVector<OpAsmParser::UnresolvedOperand, 4> indexInfo;
1953 IntegerAttr localityHint;
1955 StringRef readOrWrite, cacheType;
1972 if (readOrWrite !=
"read" && readOrWrite !=
"write")
1974 "rw specifier has to be 'read' or 'write'");
1975 result.addAttribute(PrefetchOp::getIsWriteAttrStrName(),
1978 if (cacheType !=
"data" && cacheType !=
"instr")
1980 "cache type has to be 'data' or 'instr'");
1982 result.addAttribute(PrefetchOp::getIsDataCacheAttrStrName(),
1988LogicalResult PrefetchOp::verify() {
1990 return emitOpError(
"too few indices");
1995LogicalResult PrefetchOp::fold(FoldAdaptor adaptor,
1996 SmallVectorImpl<OpFoldResult> &results) {
2003std::optional<SmallVector<Value>>
2004PrefetchOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
2007 getMemrefMutable().assign(newMemref);
2008 getIndicesMutable().assign(newIndices);
2010 return std::nullopt;
2017OpFoldResult RankOp::fold(FoldAdaptor adaptor) {
2019 auto type = getOperand().getType();
2020 auto shapedType = llvm::dyn_cast<ShapedType>(type);
2021 if (shapedType && shapedType.hasRank())
2022 return IntegerAttr::get(IndexType::get(
getContext()), shapedType.getRank());
2023 return IntegerAttr();
2032struct PrintDynamicOrValue {
2036Diagnostic &
operator<<(Diagnostic &diag, PrintDynamicOrValue printed) {
2037 if (ShapedType::isDynamic(printed.value))
2038 return diag <<
"dynamic";
2039 return diag << printed.value;
2044void ReinterpretCastOp::getAsmResultNames(
2046 setNameFn(getResult(),
"reinterpret_cast");
2052void ReinterpretCastOp::build(OpBuilder &
b, OperationState &
result,
2053 MemRefType resultType, Value source,
2054 OpFoldResult offset, ArrayRef<OpFoldResult> sizes,
2055 ArrayRef<OpFoldResult> strides,
2056 ArrayRef<NamedAttribute> attrs) {
2057 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
2058 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
2062 result.addAttributes(attrs);
2063 build(
b,
result, resultType, source, dynamicOffsets, dynamicSizes,
2064 dynamicStrides,
b.getDenseI64ArrayAttr(staticOffsets),
2065 b.getDenseI64ArrayAttr(staticSizes),
2066 b.getDenseI64ArrayAttr(staticStrides));
2069void ReinterpretCastOp::build(OpBuilder &
b, OperationState &
result,
2070 Value source, OpFoldResult offset,
2071 ArrayRef<OpFoldResult> sizes,
2072 ArrayRef<OpFoldResult> strides,
2073 ArrayRef<NamedAttribute> attrs) {
2074 auto sourceType = cast<BaseMemRefType>(source.
getType());
2075 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
2076 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
2080 auto stridedLayout = StridedLayoutAttr::get(
2081 b.getContext(), staticOffsets.front(), staticStrides);
2082 auto resultType = MemRefType::get(staticSizes, sourceType.getElementType(),
2083 stridedLayout, sourceType.getMemorySpace());
2084 build(
b,
result, resultType, source, offset, sizes, strides, attrs);
2087void ReinterpretCastOp::build(OpBuilder &
b, OperationState &
result,
2088 MemRefType resultType, Value source,
2089 int64_t offset, ArrayRef<int64_t> sizes,
2090 ArrayRef<int64_t> strides,
2091 ArrayRef<NamedAttribute> attrs) {
2092 SmallVector<OpFoldResult> sizeValues = llvm::map_to_vector<4>(
2093 sizes, [&](int64_t v) -> OpFoldResult {
return b.getI64IntegerAttr(v); });
2094 SmallVector<OpFoldResult> strideValues =
2095 llvm::map_to_vector<4>(strides, [&](int64_t v) -> OpFoldResult {
2096 return b.getI64IntegerAttr(v);
2098 build(
b,
result, resultType, source,
b.getI64IntegerAttr(offset), sizeValues,
2099 strideValues, attrs);
2102void ReinterpretCastOp::build(OpBuilder &
b, OperationState &
result,
2103 MemRefType resultType, Value source, Value offset,
2105 ArrayRef<NamedAttribute> attrs) {
2106 SmallVector<OpFoldResult> sizeValues =
2107 llvm::map_to_vector<4>(sizes, [](Value v) -> OpFoldResult {
return v; });
2108 SmallVector<OpFoldResult> strideValues = llvm::map_to_vector<4>(
2109 strides, [](Value v) -> OpFoldResult {
return v; });
2110 build(
b,
result, resultType, source, offset, sizeValues, strideValues, attrs);
2115LogicalResult ReinterpretCastOp::verify() {
2117 auto srcType = llvm::cast<BaseMemRefType>(getSource().
getType());
2118 auto resultType = llvm::cast<MemRefType>(
getType());
2119 if (srcType.getMemorySpace() != resultType.getMemorySpace())
2120 return emitError(
"different memory spaces specified for source type ")
2121 << srcType <<
" and result memref type " << resultType;
2127 for (
auto [idx, resultSize, expectedSize] :
2128 llvm::enumerate(resultType.getShape(), getStaticSizes())) {
2129 if (resultSize != expectedSize)
2130 return emitError(
"expected result type with size = ")
2131 << PrintDynamicOrValue{expectedSize} <<
" instead of "
2132 << PrintDynamicOrValue{resultSize} <<
" in dim = " << idx;
2138 int64_t resultOffset;
2139 SmallVector<int64_t, 4> resultStrides;
2140 if (
failed(resultType.getStridesAndOffset(resultStrides, resultOffset)))
2141 return emitError(
"expected result type to have strided layout but found ")
2145 int64_t expectedOffset = getStaticOffsets().front();
2146 if (resultOffset != expectedOffset)
2147 return emitError(
"expected result type with offset = ")
2148 << PrintDynamicOrValue{expectedOffset} <<
" instead of "
2149 << PrintDynamicOrValue{resultOffset};
2152 for (
auto [idx, resultStride, expectedStride] :
2153 llvm::enumerate(resultStrides, getStaticStrides())) {
2154 if (resultStride != expectedStride)
2155 return emitError(
"expected result type with stride = ")
2156 << PrintDynamicOrValue{expectedStride} <<
" instead of "
2157 << PrintDynamicOrValue{resultStride} <<
" in dim = " << idx;
2163OpFoldResult ReinterpretCastOp::fold(FoldAdaptor ) {
2164 Value src = getSource();
2165 auto getPrevSrc = [&]() -> Value {
2168 return prev.getSource();
2172 return prev.getSource();
2178 return prev.getSource();
2183 if (
auto prevSrc = getPrevSrc()) {
2184 getSourceMutable().assign(prevSrc);
2197SmallVector<OpFoldResult> ReinterpretCastOp::getConstifiedMixedSizes() {
2203SmallVector<OpFoldResult> ReinterpretCastOp::getConstifiedMixedStrides() {
2204 SmallVector<OpFoldResult> values = getMixedStrides();
2205 SmallVector<int64_t> staticValues;
2207 LogicalResult status =
getType().getStridesAndOffset(staticValues, unused);
2209 assert(succeeded(status) &&
"could not get strides from type");
2214OpFoldResult ReinterpretCastOp::getConstifiedMixedOffset() {
2215 SmallVector<OpFoldResult> values = getMixedOffsets();
2216 assert(values.size() == 1 &&
2217 "reinterpret_cast must have one and only one offset");
2218 SmallVector<int64_t> staticValues, unused;
2220 LogicalResult status =
getType().getStridesAndOffset(unused, offset);
2222 assert(succeeded(status) &&
"could not get offset from type");
2223 staticValues.push_back(offset);
2271struct ReinterpretCastOpExtractStridedMetadataFolder
2272 :
public OpRewritePattern<ReinterpretCastOp> {
2274 using OpRewritePattern<ReinterpretCastOp>::OpRewritePattern;
2276 LogicalResult matchAndRewrite(ReinterpretCastOp op,
2277 PatternRewriter &rewriter)
const override {
2278 auto extractStridedMetadata =
2279 op.getSource().getDefiningOp<ExtractStridedMetadataOp>();
2280 if (!extractStridedMetadata)
2285 auto isReinterpretCastNoop = [&]() ->
bool {
2287 if (!llvm::equal(extractStridedMetadata.getConstifiedMixedStrides(),
2288 op.getConstifiedMixedStrides()))
2292 if (!llvm::equal(extractStridedMetadata.getConstifiedMixedSizes(),
2293 op.getConstifiedMixedSizes()))
2297 assert(op.getMixedOffsets().size() == 1 &&
2298 "reinterpret_cast with more than one offset should have been "
2299 "rejected by the verifier");
2300 return extractStridedMetadata.getConstifiedMixedOffset() ==
2301 op.getConstifiedMixedOffset();
2304 if (!isReinterpretCastNoop()) {
2321 op.getSourceMutable().assign(extractStridedMetadata.getSource());
2331 Type srcTy = extractStridedMetadata.getSource().getType();
2332 if (srcTy == op.getResult().getType())
2333 rewriter.
replaceOp(op, extractStridedMetadata.getSource());
2336 extractStridedMetadata.getSource());
2342struct ReinterpretCastOpConstantFolder
2343 :
public OpRewritePattern<ReinterpretCastOp> {
2345 using OpRewritePattern<ReinterpretCastOp>::OpRewritePattern;
2347 LogicalResult matchAndRewrite(ReinterpretCastOp op,
2348 PatternRewriter &rewriter)
const override {
2349 unsigned srcStaticCount = llvm::count_if(
2350 llvm::concat<OpFoldResult>(op.getMixedOffsets(), op.getMixedSizes(),
2351 op.getMixedStrides()),
2352 [](OpFoldResult ofr) { return isa<Attribute>(ofr); });
2354 SmallVector<OpFoldResult> offsets = {op.getConstifiedMixedOffset()};
2355 SmallVector<OpFoldResult> sizes = op.getConstifiedMixedSizes();
2356 SmallVector<OpFoldResult> strides = op.getConstifiedMixedStrides();
2363 offsets[0] = op.getMixedOffsets()[0];
2368 for (
auto it : llvm::zip(op.getMixedSizes(), sizes)) {
2369 auto &srcSizeOfr = std::get<0>(it);
2370 auto &sizeOfr = std::get<1>(it);
2373 sizeOfr = srcSizeOfr;
2380 if (srcStaticCount ==
2381 llvm::count_if(llvm::concat<OpFoldResult>(offsets, sizes, strides),
2382 [](OpFoldResult ofr) {
return isa<Attribute>(ofr); }))
2385 auto newReinterpretCast = ReinterpretCastOp::create(
2386 rewriter, op->getLoc(), op.getSource(), offsets[0], sizes, strides);
2394void ReinterpretCastOp::getCanonicalizationPatterns(RewritePatternSet &results,
2395 MLIRContext *context) {
2396 results.
add<ReinterpretCastOpExtractStridedMetadataFolder,
2397 ReinterpretCastOpConstantFolder>(context);
2400FailureOr<std::optional<SmallVector<Value>>>
2401ReinterpretCastOp::bubbleDownCasts(OpBuilder &builder) {
2409void CollapseShapeOp::getAsmResultNames(
2411 setNameFn(getResult(),
"collapse_shape");
2414void ExpandShapeOp::getAsmResultNames(
2416 setNameFn(getResult(),
"expand_shape");
2419LogicalResult ExpandShapeOp::reifyResultShapes(
2421 reifiedResultShapes = {
2422 getMixedValues(getStaticOutputShape(), getOutputShape(), builder)};
2435 bool allowMultipleDynamicDimsPerGroup) {
2437 if (collapsedShape.size() != reassociation.size())
2438 return op->
emitOpError(
"invalid number of reassociation groups: found ")
2439 << reassociation.size() <<
", expected " << collapsedShape.size();
2444 for (
const auto &it : llvm::enumerate(reassociation)) {
2446 int64_t collapsedDim = it.index();
2448 bool foundDynamic =
false;
2449 for (
int64_t expandedDim : group) {
2450 if (expandedDim != nextDim++)
2451 return op->
emitOpError(
"reassociation indices must be contiguous");
2453 if (expandedDim >=
static_cast<int64_t>(expandedShape.size()))
2455 << expandedDim <<
" is out of bounds";
2458 if (ShapedType::isDynamic(expandedShape[expandedDim])) {
2459 if (foundDynamic && !allowMultipleDynamicDimsPerGroup)
2461 "at most one dimension in a reassociation group may be dynamic");
2462 foundDynamic =
true;
2467 if (ShapedType::isDynamic(collapsedShape[collapsedDim]) != foundDynamic)
2470 <<
") must be dynamic if and only if reassociation group is "
2475 if (!foundDynamic) {
2477 for (
int64_t expandedDim : group)
2478 groupSize *= expandedShape[expandedDim];
2479 if (groupSize != collapsedShape[collapsedDim])
2481 << collapsedShape[collapsedDim]
2482 <<
") must equal reassociation group size (" << groupSize <<
")";
2486 if (collapsedShape.empty()) {
2488 for (
int64_t d : expandedShape)
2491 "rank 0 memrefs can only be extended/collapsed with/from ones");
2492 }
else if (nextDim !=
static_cast<int64_t>(expandedShape.size())) {
2496 << expandedShape.size()
2497 <<
") inconsistent with number of reassociation indices (" << nextDim
2504SmallVector<AffineMap, 4> CollapseShapeOp::getReassociationMaps() {
2508SmallVector<ReassociationExprs, 4> CollapseShapeOp::getReassociationExprs() {
2510 getReassociationIndices());
2513SmallVector<AffineMap, 4> ExpandShapeOp::getReassociationMaps() {
2517SmallVector<ReassociationExprs, 4> ExpandShapeOp::getReassociationExprs() {
2519 getReassociationIndices());
2524static FailureOr<StridedLayoutAttr>
2529 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
2531 assert(srcStrides.size() == reassociation.size() &&
"invalid reassociation");
2546 reverseResultStrides.reserve(resultShape.size());
2547 unsigned shapeIndex = resultShape.size() - 1;
2548 for (
auto it : llvm::reverse(llvm::zip(reassociation, srcStrides))) {
2550 int64_t currentStrideToExpand = std::get<1>(it);
2551 for (
unsigned idx = 0, e = reassoc.size(); idx < e; ++idx) {
2552 reverseResultStrides.push_back(currentStrideToExpand);
2553 currentStrideToExpand =
2559 auto resultStrides = llvm::to_vector<8>(llvm::reverse(reverseResultStrides));
2560 resultStrides.resize(resultShape.size(), 1);
2561 return StridedLayoutAttr::get(srcType.getContext(), srcOffset, resultStrides);
2564FailureOr<MemRefType> ExpandShapeOp::computeExpandedType(
2565 MemRefType srcType, ArrayRef<int64_t> resultShape,
2566 ArrayRef<ReassociationIndices> reassociation) {
2567 if (srcType.getLayout().isIdentity()) {
2570 MemRefLayoutAttrInterface layout;
2571 return MemRefType::get(resultShape, srcType.getElementType(), layout,
2572 srcType.getMemorySpace());
2576 FailureOr<StridedLayoutAttr> computedLayout =
2578 if (
failed(computedLayout))
2580 return MemRefType::get(resultShape, srcType.getElementType(), *computedLayout,
2581 srcType.getMemorySpace());
2584FailureOr<SmallVector<OpFoldResult>>
2585ExpandShapeOp::inferOutputShape(OpBuilder &
b, Location loc,
2586 MemRefType expandedType,
2587 ArrayRef<ReassociationIndices> reassociation,
2588 ArrayRef<OpFoldResult> inputShape) {
2589 std::optional<SmallVector<OpFoldResult>> outputShape =
2594 return *outputShape;
2597void ExpandShapeOp::build(OpBuilder &builder, OperationState &
result,
2598 Type resultType, Value src,
2599 ArrayRef<ReassociationIndices> reassociation,
2600 ArrayRef<OpFoldResult> outputShape) {
2601 auto [staticOutputShape, dynamicOutputShape] =
2603 build(builder,
result, llvm::cast<MemRefType>(resultType), src,
2605 dynamicOutputShape, staticOutputShape);
2608void ExpandShapeOp::build(OpBuilder &builder, OperationState &
result,
2609 Type resultType, Value src,
2610 ArrayRef<ReassociationIndices> reassociation) {
2611 SmallVector<OpFoldResult> inputShape =
2613 MemRefType memrefResultTy = llvm::cast<MemRefType>(resultType);
2614 FailureOr<SmallVector<OpFoldResult>> outputShape = inferOutputShape(
2615 builder,
result.location, memrefResultTy, reassociation, inputShape);
2618 assert(succeeded(outputShape) &&
"unable to infer output shape");
2619 build(builder,
result, memrefResultTy, src, reassociation, *outputShape);
2622void ExpandShapeOp::build(OpBuilder &builder, OperationState &
result,
2623 ArrayRef<int64_t> resultShape, Value src,
2624 ArrayRef<ReassociationIndices> reassociation) {
2626 auto srcType = llvm::cast<MemRefType>(src.
getType());
2627 FailureOr<MemRefType> resultType =
2628 ExpandShapeOp::computeExpandedType(srcType, resultShape, reassociation);
2631 assert(succeeded(resultType) &&
"could not compute layout");
2632 build(builder,
result, *resultType, src, reassociation);
2635void ExpandShapeOp::build(OpBuilder &builder, OperationState &
result,
2636 ArrayRef<int64_t> resultShape, Value src,
2637 ArrayRef<ReassociationIndices> reassociation,
2638 ArrayRef<OpFoldResult> outputShape) {
2640 auto srcType = llvm::cast<MemRefType>(src.
getType());
2641 FailureOr<MemRefType> resultType =
2642 ExpandShapeOp::computeExpandedType(srcType, resultShape, reassociation);
2645 assert(succeeded(resultType) &&
"could not compute layout");
2646 build(builder,
result, *resultType, src, reassociation, outputShape);
2649LogicalResult ExpandShapeOp::verify() {
2653 MemRefType srcType = getSrcType();
2654 MemRefType resultType = getResultType();
2656 if (srcType.getRank() > resultType.getRank()) {
2657 auto r0 = srcType.getRank();
2658 auto r1 = resultType.getRank();
2659 return emitOpError(
"has source rank ")
2660 << r0 <<
" and result rank " << r1 <<
". This is not an expansion ("
2661 << r0 <<
" > " << r1 <<
").";
2666 resultType.getShape(),
2667 getReassociationIndices(),
2672 FailureOr<MemRefType> expectedResultType = ExpandShapeOp::computeExpandedType(
2673 srcType, resultType.getShape(), getReassociationIndices());
2674 if (
failed(expectedResultType))
2675 return emitOpError(
"invalid source layout map");
2678 if (*expectedResultType != resultType)
2679 return emitOpError(
"expected expanded type to be ")
2680 << *expectedResultType <<
" but found " << resultType;
2682 if ((int64_t)getStaticOutputShape().size() != resultType.getRank())
2683 return emitOpError(
"expected number of static shape bounds to be equal to "
2684 "the output rank (")
2685 << resultType.getRank() <<
") but found "
2686 << getStaticOutputShape().size() <<
" inputs instead";
2688 if ((int64_t)getOutputShape().size() !=
2689 llvm::count(getStaticOutputShape(), ShapedType::kDynamic))
2690 return emitOpError(
"mismatch in dynamic dims in output_shape and "
2691 "static_output_shape: static_output_shape has ")
2692 << llvm::count(getStaticOutputShape(), ShapedType::kDynamic)
2693 <<
" dynamic dims while output_shape has " << getOutputShape().size()
2704 ArrayRef<int64_t> resShape = getResult().getType().getShape();
2705 for (
auto [pos, shape] : llvm::enumerate(resShape)) {
2706 if (ShapedType::isStatic(shape) && shape != staticOutputShapes[pos]) {
2707 return emitOpError(
"invalid output shape provided at pos ") << pos;
2720 auto cast = op.getSrc().getDefiningOp<CastOp>();
2724 if (!CastOp::canFoldIntoConsumerOp(cast))
2732 for (
auto [dimIdx, dimSize] : enumerate(originalOutputShape)) {
2734 if (!sizeOpt.has_value()) {
2735 newOutputShapeSizes.push_back(ShapedType::kDynamic);
2739 newOutputShapeSizes.push_back(sizeOpt.value());
2740 newOutputShape[dimIdx] = rewriter.
getIndexAttr(sizeOpt.value());
2743 Value castSource = cast.getSource();
2744 auto castSourceType = llvm::cast<MemRefType>(castSource.
getType());
2746 op.getReassociationIndices();
2747 for (
auto [idx, group] : llvm::enumerate(reassociationIndices)) {
2748 auto newOutputShapeSizesSlice =
2749 ArrayRef(newOutputShapeSizes).slice(group.front(), group.size());
2750 bool newOutputDynamic =
2751 llvm::is_contained(newOutputShapeSizesSlice, ShapedType::kDynamic);
2752 if (castSourceType.isDynamicDim(idx) != newOutputDynamic)
2754 op,
"folding cast will result in changing dynamicity in "
2755 "reassociation group");
2758 FailureOr<MemRefType> newResultTypeOrFailure =
2759 ExpandShapeOp::computeExpandedType(castSourceType, newOutputShapeSizes,
2760 reassociationIndices);
2762 if (failed(newResultTypeOrFailure))
2764 op,
"could not compute new expanded type after folding cast");
2766 if (*newResultTypeOrFailure == op.getResultType()) {
2768 op, [&]() { op.getSrcMutable().assign(castSource); });
2770 Value newOp = ExpandShapeOp::create(rewriter, op->getLoc(),
2771 *newResultTypeOrFailure, castSource,
2772 reassociationIndices, newOutputShape);
2779void ExpandShapeOp::getCanonicalizationPatterns(RewritePatternSet &results,
2780 MLIRContext *context) {
2782 ComposeReassociativeReshapeOps<ExpandShapeOp, ReshapeOpKind::kExpand>,
2783 ComposeExpandOfCollapseOp<ExpandShapeOp, CollapseShapeOp, CastOp>,
2784 ExpandShapeOpMemRefCastFolder>(context);
2787FailureOr<std::optional<SmallVector<Value>>>
2788ExpandShapeOp::bubbleDownCasts(OpBuilder &builder) {
2799static FailureOr<StridedLayoutAttr>
2802 bool strict =
false) {
2805 auto srcShape = srcType.getShape();
2806 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
2815 resultStrides.reserve(reassociation.size());
2818 while (srcShape[ref.back()] == 1 && ref.size() > 1)
2819 ref = ref.drop_back();
2820 if (ShapedType::isStatic(srcShape[ref.back()]) || ref.size() == 1) {
2821 resultStrides.push_back(srcStrides[ref.back()]);
2827 resultStrides.push_back(ShapedType::kDynamic);
2832 unsigned resultStrideIndex = resultStrides.size() - 1;
2836 for (
int64_t idx : llvm::reverse(trailingReassocs)) {
2841 if (srcShape[idx - 1] == 1)
2853 if (strict && (stride.saturated || srcStride.saturated))
2856 if (!stride.saturated && !srcStride.saturated && stride != srcStride)
2860 return StridedLayoutAttr::get(srcType.getContext(), srcOffset, resultStrides);
2863bool CollapseShapeOp::isGuaranteedCollapsible(
2864 MemRefType srcType, ArrayRef<ReassociationIndices> reassociation) {
2866 if (srcType.getLayout().isIdentity())
2873MemRefType CollapseShapeOp::computeCollapsedType(
2874 MemRefType srcType, ArrayRef<ReassociationIndices> reassociation) {
2875 SmallVector<int64_t> resultShape;
2876 resultShape.reserve(reassociation.size());
2879 for (int64_t srcDim : group)
2882 resultShape.push_back(groupSize.asInteger());
2885 if (srcType.getLayout().isIdentity()) {
2888 MemRefLayoutAttrInterface layout;
2889 return MemRefType::get(resultShape, srcType.getElementType(), layout,
2890 srcType.getMemorySpace());
2896 FailureOr<StridedLayoutAttr> computedLayout =
2898 assert(succeeded(computedLayout) &&
2899 "invalid source layout map or collapsing non-contiguous dims");
2900 return MemRefType::get(resultShape, srcType.getElementType(), *computedLayout,
2901 srcType.getMemorySpace());
2904void CollapseShapeOp::build(OpBuilder &
b, OperationState &
result, Value src,
2905 ArrayRef<ReassociationIndices> reassociation,
2906 ArrayRef<NamedAttribute> attrs) {
2907 auto srcType = llvm::cast<MemRefType>(src.
getType());
2908 MemRefType resultType =
2909 CollapseShapeOp::computeCollapsedType(srcType, reassociation);
2910 buildPropertiesAndDiscardableAttributes(
result, attrs);
2911 result.getOrAddProperties<Properties>().reassociation =
2914 result.addTypes(resultType);
2917LogicalResult CollapseShapeOp::verify() {
2921 MemRefType srcType = getSrcType();
2922 MemRefType resultType = getResultType();
2924 if (srcType.getRank() < resultType.getRank()) {
2925 auto r0 = srcType.getRank();
2926 auto r1 = resultType.getRank();
2927 return emitOpError(
"has source rank ")
2928 << r0 <<
" and result rank " << r1 <<
". This is not a collapse ("
2929 << r0 <<
" < " << r1 <<
").";
2934 srcType.getShape(), getReassociationIndices(),
2939 MemRefType expectedResultType;
2940 if (srcType.getLayout().isIdentity()) {
2943 MemRefLayoutAttrInterface layout;
2944 expectedResultType =
2945 MemRefType::get(resultType.getShape(), srcType.getElementType(), layout,
2946 srcType.getMemorySpace());
2951 FailureOr<StridedLayoutAttr> computedLayout =
2953 if (
failed(computedLayout))
2955 "invalid source layout map or collapsing non-contiguous dims");
2956 expectedResultType =
2957 MemRefType::get(resultType.getShape(), srcType.getElementType(),
2958 *computedLayout, srcType.getMemorySpace());
2961 if (expectedResultType != resultType)
2962 return emitOpError(
"expected collapsed type to be ")
2963 << expectedResultType <<
" but found " << resultType;
2975 auto cast = op.getOperand().getDefiningOp<CastOp>();
2979 if (!CastOp::canFoldIntoConsumerOp(cast))
2982 Type newResultType = CollapseShapeOp::computeCollapsedType(
2983 llvm::cast<MemRefType>(cast.getOperand().getType()),
2984 op.getReassociationIndices());
2986 if (newResultType == op.getResultType()) {
2988 op, [&]() { op.getSrcMutable().assign(cast.getSource()); });
2991 CollapseShapeOp::create(rewriter, op->getLoc(), cast.getSource(),
2992 op.getReassociationIndices());
2999void CollapseShapeOp::getCanonicalizationPatterns(RewritePatternSet &results,
3000 MLIRContext *context) {
3002 ComposeReassociativeReshapeOps<CollapseShapeOp, ReshapeOpKind::kCollapse>,
3003 ComposeCollapseOfExpandOp<CollapseShapeOp, ExpandShapeOp, CastOp,
3004 memref::DimOp, MemRefType>,
3005 CollapseShapeOpMemRefCastFolder>(context);
3008OpFoldResult ExpandShapeOp::fold(FoldAdaptor adaptor) {
3010 adaptor.getOperands());
3013OpFoldResult CollapseShapeOp::fold(FoldAdaptor adaptor) {
3015 adaptor.getOperands());
3018FailureOr<std::optional<SmallVector<Value>>>
3019CollapseShapeOp::bubbleDownCasts(OpBuilder &builder) {
3027void ReshapeOp::getAsmResultNames(
3029 setNameFn(getResult(),
"reshape");
3032LogicalResult ReshapeOp::verify() {
3033 Type operandType = getSource().getType();
3034 Type resultType = getResult().getType();
3036 Type operandElementType =
3037 llvm::cast<ShapedType>(operandType).getElementType();
3038 Type resultElementType = llvm::cast<ShapedType>(resultType).getElementType();
3039 if (operandElementType != resultElementType)
3040 return emitOpError(
"element types of source and destination memref "
3041 "types should be the same");
3043 if (
auto operandMemRefType = llvm::dyn_cast<MemRefType>(operandType))
3044 if (!operandMemRefType.getLayout().isIdentity())
3045 return emitOpError(
"source memref type should have identity affine map");
3049 auto resultMemRefType = llvm::dyn_cast<MemRefType>(resultType);
3050 if (resultMemRefType) {
3051 if (!resultMemRefType.getLayout().isIdentity())
3052 return emitOpError(
"result memref type should have identity affine map");
3053 if (shapeSize == ShapedType::kDynamic)
3054 return emitOpError(
"cannot use shape operand with dynamic length to "
3055 "reshape to statically-ranked memref type");
3056 if (shapeSize != resultMemRefType.getRank())
3058 "length of shape operand differs from the result's memref rank");
3063FailureOr<std::optional<SmallVector<Value>>>
3064ReshapeOp::bubbleDownCasts(OpBuilder &builder) {
3072LogicalResult StoreOp::fold(FoldAdaptor adaptor,
3073 SmallVectorImpl<OpFoldResult> &results) {
3080std::optional<SmallVector<Value>>
3081StoreOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
3084 getMemrefMutable().assign(newMemref);
3085 getIndicesMutable().assign(newIndices);
3087 return std::nullopt;
3090FailureOr<std::optional<SmallVector<Value>>>
3091StoreOp::bubbleDownCasts(OpBuilder &builder) {
3100void SubViewOp::getAsmResultNames(
3102 setNameFn(getResult(),
"subview");
3108MemRefType SubViewOp::inferResultType(MemRefType sourceMemRefType,
3109 ArrayRef<int64_t> staticOffsets,
3110 ArrayRef<int64_t> staticSizes,
3111 ArrayRef<int64_t> staticStrides) {
3112 unsigned rank = sourceMemRefType.getRank();
3114 assert(staticOffsets.size() == rank &&
"staticOffsets length mismatch");
3115 assert(staticSizes.size() == rank &&
"staticSizes length mismatch");
3116 assert(staticStrides.size() == rank &&
"staticStrides length mismatch");
3119 auto [sourceStrides, sourceOffset] = sourceMemRefType.getStridesAndOffset();
3123 int64_t targetOffset = sourceOffset;
3124 for (
auto it : llvm::zip(staticOffsets, sourceStrides)) {
3125 auto staticOffset = std::get<0>(it), sourceStride = std::get<1>(it);
3134 SmallVector<int64_t, 4> targetStrides;
3135 targetStrides.reserve(staticOffsets.size());
3136 for (
auto it : llvm::zip(sourceStrides, staticStrides)) {
3137 auto sourceStride = std::get<0>(it), staticStride = std::get<1>(it);
3144 return MemRefType::get(staticSizes, sourceMemRefType.getElementType(),
3145 StridedLayoutAttr::get(sourceMemRefType.getContext(),
3146 targetOffset, targetStrides),
3147 sourceMemRefType.getMemorySpace());
3150MemRefType SubViewOp::inferResultType(MemRefType sourceMemRefType,
3151 ArrayRef<OpFoldResult> offsets,
3152 ArrayRef<OpFoldResult> sizes,
3153 ArrayRef<OpFoldResult> strides) {
3154 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
3155 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
3165 return SubViewOp::inferResultType(sourceMemRefType, staticOffsets,
3166 staticSizes, staticStrides);
3169MemRefType SubViewOp::inferRankReducedResultType(
3170 ArrayRef<int64_t> resultShape, MemRefType sourceRankedTensorType,
3171 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes,
3172 ArrayRef<int64_t> strides) {
3173 MemRefType inferredType =
3174 inferResultType(sourceRankedTensorType, offsets, sizes, strides);
3175 assert(inferredType.getRank() >=
static_cast<int64_t
>(resultShape.size()) &&
3177 if (inferredType.getRank() ==
static_cast<int64_t
>(resultShape.size()))
3178 return inferredType;
3181 std::optional<llvm::SmallDenseSet<unsigned>> dimsToProject =
3183 assert(dimsToProject.has_value() &&
"invalid rank reduction");
3186 auto inferredLayout = llvm::cast<StridedLayoutAttr>(inferredType.getLayout());
3187 SmallVector<int64_t> rankReducedStrides;
3188 rankReducedStrides.reserve(resultShape.size());
3189 for (
auto [idx, value] : llvm::enumerate(inferredLayout.getStrides())) {
3190 if (!dimsToProject->contains(idx))
3191 rankReducedStrides.push_back(value);
3193 return MemRefType::get(resultShape, inferredType.getElementType(),
3194 StridedLayoutAttr::get(inferredLayout.getContext(),
3195 inferredLayout.getOffset(),
3196 rankReducedStrides),
3197 inferredType.getMemorySpace());
3200MemRefType SubViewOp::inferRankReducedResultType(
3201 ArrayRef<int64_t> resultShape, MemRefType sourceRankedTensorType,
3202 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
3203 ArrayRef<OpFoldResult> strides) {
3204 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
3205 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
3209 return SubViewOp::inferRankReducedResultType(
3210 resultShape, sourceRankedTensorType, staticOffsets, staticSizes,
3216void SubViewOp::build(OpBuilder &
b, OperationState &
result,
3217 MemRefType resultType, Value source,
3218 ArrayRef<OpFoldResult> offsets,
3219 ArrayRef<OpFoldResult> sizes,
3220 ArrayRef<OpFoldResult> strides,
3221 ArrayRef<NamedAttribute> attrs) {
3222 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
3223 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
3227 auto sourceMemRefType = llvm::cast<MemRefType>(source.
getType());
3230 resultType = SubViewOp::inferResultType(sourceMemRefType, staticOffsets,
3231 staticSizes, staticStrides);
3233 result.addAttributes(attrs);
3234 build(
b,
result, resultType, source, dynamicOffsets, dynamicSizes,
3235 dynamicStrides,
b.getDenseI64ArrayAttr(staticOffsets),
3236 b.getDenseI64ArrayAttr(staticSizes),
3237 b.getDenseI64ArrayAttr(staticStrides));
3242void SubViewOp::build(OpBuilder &
b, OperationState &
result, Value source,
3243 ArrayRef<OpFoldResult> offsets,
3244 ArrayRef<OpFoldResult> sizes,
3245 ArrayRef<OpFoldResult> strides,
3246 ArrayRef<NamedAttribute> attrs) {
3247 build(
b,
result, MemRefType(), source, offsets, sizes, strides, attrs);
3251void SubViewOp::build(OpBuilder &
b, OperationState &
result, Value source,
3252 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes,
3253 ArrayRef<int64_t> strides,
3254 ArrayRef<NamedAttribute> attrs) {
3255 SmallVector<OpFoldResult> offsetValues =
3256 llvm::map_to_vector<4>(offsets, [&](int64_t v) -> OpFoldResult {
3257 return b.getI64IntegerAttr(v);
3259 SmallVector<OpFoldResult> sizeValues = llvm::map_to_vector<4>(
3260 sizes, [&](int64_t v) -> OpFoldResult {
return b.getI64IntegerAttr(v); });
3261 SmallVector<OpFoldResult> strideValues =
3262 llvm::map_to_vector<4>(strides, [&](int64_t v) -> OpFoldResult {
3263 return b.getI64IntegerAttr(v);
3265 build(
b,
result, source, offsetValues, sizeValues, strideValues, attrs);
3270void SubViewOp::build(OpBuilder &
b, OperationState &
result,
3271 MemRefType resultType, Value source,
3272 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes,
3273 ArrayRef<int64_t> strides,
3274 ArrayRef<NamedAttribute> attrs) {
3275 SmallVector<OpFoldResult> offsetValues =
3276 llvm::map_to_vector<4>(offsets, [&](int64_t v) -> OpFoldResult {
3277 return b.getI64IntegerAttr(v);
3279 SmallVector<OpFoldResult> sizeValues = llvm::map_to_vector<4>(
3280 sizes, [&](int64_t v) -> OpFoldResult {
return b.getI64IntegerAttr(v); });
3281 SmallVector<OpFoldResult> strideValues =
3282 llvm::map_to_vector<4>(strides, [&](int64_t v) -> OpFoldResult {
3283 return b.getI64IntegerAttr(v);
3285 build(
b,
result, resultType, source, offsetValues, sizeValues, strideValues,
3291void SubViewOp::build(OpBuilder &
b, OperationState &
result,
3292 MemRefType resultType, Value source,
ValueRange offsets,
3294 ArrayRef<NamedAttribute> attrs) {
3295 SmallVector<OpFoldResult> offsetValues = llvm::map_to_vector<4>(
3296 offsets, [](Value v) -> OpFoldResult {
return v; });
3297 SmallVector<OpFoldResult> sizeValues =
3298 llvm::map_to_vector<4>(sizes, [](Value v) -> OpFoldResult {
return v; });
3299 SmallVector<OpFoldResult> strideValues = llvm::map_to_vector<4>(
3300 strides, [](Value v) -> OpFoldResult {
return v; });
3301 build(
b,
result, resultType, source, offsetValues, sizeValues, strideValues);
3305void SubViewOp::build(OpBuilder &
b, OperationState &
result, Value source,
3307 ArrayRef<NamedAttribute> attrs) {
3308 build(
b,
result, MemRefType(), source, offsets, sizes, strides, attrs);
3312Value SubViewOp::getViewSource() {
return getSource(); }
3319 auto res1 = t1.getStridesAndOffset(t1Strides, t1Offset);
3320 auto res2 = t2.getStridesAndOffset(t2Strides, t2Offset);
3321 return succeeded(res1) && succeeded(res2) && t1Offset == t2Offset;
3328 const llvm::SmallBitVector &droppedDims) {
3329 assert(
size_t(t1.getRank()) == droppedDims.size() &&
3330 "incorrect number of bits");
3331 assert(
size_t(t1.getRank() - t2.getRank()) == droppedDims.count() &&
3332 "incorrect number of dropped dims");
3335 auto res1 = t1.getStridesAndOffset(t1Strides, t1Offset);
3336 auto res2 = t2.getStridesAndOffset(t2Strides, t2Offset);
3337 if (failed(res1) || failed(res2))
3339 for (
int64_t i = 0,
j = 0, e = t1.getRank(); i < e; ++i) {
3342 if (t1Strides[i] != t2Strides[
j])
3350 SubViewOp op,
Type expectedType) {
3351 auto memrefType = llvm::cast<ShapedType>(expectedType);
3356 return op->emitError(
"expected result rank to be smaller or equal to ")
3357 <<
"the source rank, but got " << op.getType();
3359 return op->emitError(
"expected result type to be ")
3361 <<
" or a rank-reduced version. (mismatch of result sizes), but got "
3364 return op->emitError(
"expected result element type to be ")
3365 << memrefType.getElementType() <<
", but got " << op.getType();
3367 return op->emitError(
3368 "expected result and source memory spaces to match, but got ")
3371 return op->emitError(
"expected result type to be ")
3373 <<
" or a rank-reduced version. (mismatch of result layout), but "
3377 llvm_unreachable(
"unexpected subview verification result");
3381LogicalResult SubViewOp::verify() {
3382 MemRefType baseType = getSourceType();
3383 MemRefType subViewType =
getType();
3384 ArrayRef<int64_t> staticOffsets = getStaticOffsets();
3385 ArrayRef<int64_t> staticSizes = getStaticSizes();
3386 ArrayRef<int64_t> staticStrides = getStaticStrides();
3389 if (baseType.getMemorySpace() != subViewType.getMemorySpace())
3390 return emitError(
"different memory spaces specified for base memref "
3392 << baseType <<
" and subview memref type " << subViewType;
3395 if (!baseType.isStrided())
3396 return emitError(
"base type ") << baseType <<
" is not strided";
3400 MemRefType expectedType = SubViewOp::inferResultType(
3401 baseType, staticOffsets, staticSizes, staticStrides);
3406 expectedType, subViewType);
3411 if (expectedType.getMemorySpace() != subViewType.getMemorySpace())
3413 *
this, expectedType);
3418 *
this, expectedType);
3428 *
this, expectedType);
3433 *
this, expectedType);
3437 SliceBoundsVerificationResult boundsResult =
3439 staticStrides,
true);
3441 return getOperation()->emitError(boundsResult.
errorMessage);
3447 return os <<
"range " << range.
offset <<
":" << range.
size <<
":"
3456 std::array<unsigned, 3> ranks = op.getArrayAttrMaxRanks();
3457 assert(ranks[0] == ranks[1] &&
"expected offset and sizes of equal ranks");
3458 assert(ranks[1] == ranks[2] &&
"expected sizes and strides of equal ranks");
3460 unsigned rank = ranks[0];
3462 for (
unsigned idx = 0; idx < rank; ++idx) {
3464 op.isDynamicOffset(idx)
3465 ? op.getDynamicOffset(idx)
3468 op.isDynamicSize(idx)
3469 ? op.getDynamicSize(idx)
3472 op.isDynamicStride(idx)
3473 ? op.getDynamicStride(idx)
3475 res.emplace_back(
Range{offset, size, stride});
3488 MemRefType currentResultType, MemRefType currentSourceType,
3491 MemRefType nonRankReducedType = SubViewOp::inferResultType(
3492 sourceType, mixedOffsets, mixedSizes, mixedStrides);
3494 currentSourceType, currentResultType, mixedSizes);
3495 if (failed(unusedDims))
3498 auto layout = llvm::cast<StridedLayoutAttr>(nonRankReducedType.getLayout());
3500 unsigned numDimsAfterReduction =
3501 nonRankReducedType.getRank() - unusedDims->count();
3502 shape.reserve(numDimsAfterReduction);
3503 strides.reserve(numDimsAfterReduction);
3504 for (
const auto &[idx, size, stride] :
3505 llvm::zip(llvm::seq<unsigned>(0, nonRankReducedType.getRank()),
3506 nonRankReducedType.getShape(), layout.getStrides())) {
3507 if (unusedDims->test(idx))
3509 shape.push_back(size);
3510 strides.push_back(stride);
3513 return MemRefType::get(
shape, nonRankReducedType.getElementType(),
3514 StridedLayoutAttr::get(sourceType.getContext(),
3515 layout.getOffset(), strides),
3516 nonRankReducedType.getMemorySpace());
3521 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3522 unsigned rank = memrefType.getRank();
3526 MemRefType targetType = SubViewOp::inferRankReducedResultType(
3527 targetShape, memrefType, offsets, sizes, strides);
3528 return b.createOrFold<memref::SubViewOp>(loc, targetType,
memref, offsets,
3535 auto sourceMemrefType = llvm::dyn_cast<MemRefType>(value.
getType());
3536 assert(sourceMemrefType &&
"not a ranked memref type");
3537 auto sourceShape = sourceMemrefType.getShape();
3538 if (sourceShape.equals(desiredShape))
3540 auto maybeRankReductionMask =
3542 if (!maybeRankReductionMask)
3552 if (subViewOp.getSourceType().getRank() != subViewOp.getType().getRank())
3555 auto mixedOffsets = subViewOp.getMixedOffsets();
3556 auto mixedSizes = subViewOp.getMixedSizes();
3557 auto mixedStrides = subViewOp.getMixedStrides();
3562 return !intValue || intValue.value() != 0;
3569 return !intValue || intValue.value() != 1;
3575 for (
const auto &size : llvm::enumerate(mixedSizes)) {
3577 if (!intValue || *intValue != sourceShape[size.index()])
3601class SubViewOpMemRefCastFolder final :
public OpRewritePattern<SubViewOp> {
3603 using OpRewritePattern<SubViewOp>::OpRewritePattern;
3605 LogicalResult matchAndRewrite(SubViewOp subViewOp,
3606 PatternRewriter &rewriter)
const override {
3609 if (llvm::any_of(subViewOp.getOperands(), [](Value operand) {
3610 return matchPattern(operand, matchConstantIndex());
3614 auto castOp = subViewOp.getSource().getDefiningOp<CastOp>();
3618 if (!CastOp::canFoldIntoConsumerOp(castOp))
3626 subViewOp.getType(), subViewOp.getSourceType(),
3627 llvm::cast<MemRefType>(castOp.getSource().getType()),
3628 subViewOp.getMixedOffsets(), subViewOp.getMixedSizes(),
3629 subViewOp.getMixedStrides());
3633 Value newSubView = SubViewOp::create(
3634 rewriter, subViewOp.getLoc(), resultType, castOp.getSource(),
3635 subViewOp.getOffsets(), subViewOp.getSizes(), subViewOp.getStrides(),
3636 subViewOp.getStaticOffsets(), subViewOp.getStaticSizes(),
3637 subViewOp.getStaticStrides());
3646class TrivialSubViewOpFolder final :
public OpRewritePattern<SubViewOp> {
3648 using OpRewritePattern<SubViewOp>::OpRewritePattern;
3650 LogicalResult matchAndRewrite(SubViewOp subViewOp,
3651 PatternRewriter &rewriter)
const override {
3654 if (subViewOp.getSourceType() == subViewOp.getType()) {
3655 rewriter.
replaceOp(subViewOp, subViewOp.getSource());
3659 subViewOp.getSource());
3671 MemRefType resTy = SubViewOp::inferResultType(
3672 op.getSourceType(), mixedOffsets, mixedSizes, mixedStrides);
3675 MemRefType nonReducedType = resTy;
3678 llvm::SmallBitVector droppedDims = op.getDroppedDims();
3679 if (droppedDims.none())
3680 return nonReducedType;
3683 auto [nonReducedStrides, offset] = nonReducedType.getStridesAndOffset();
3688 for (
int64_t i = 0; i < static_cast<int64_t>(mixedSizes.size()); ++i) {
3689 if (droppedDims.test(i))
3691 targetStrides.push_back(nonReducedStrides[i]);
3692 targetShape.push_back(nonReducedType.getDimSize(i));
3695 return MemRefType::get(targetShape, nonReducedType.getElementType(),
3696 StridedLayoutAttr::get(nonReducedType.getContext(),
3697 offset, targetStrides),
3698 nonReducedType.getMemorySpace());
3709void SubViewOp::getCanonicalizationPatterns(RewritePatternSet &results,
3710 MLIRContext *context) {
3712 .
add<OpWithOffsetSizesAndStridesConstantArgumentFolder<
3713 SubViewOp, SubViewReturnTypeCanonicalizer, SubViewCanonicalizer>,
3714 SubViewOpMemRefCastFolder, TrivialSubViewOpFolder>(context);
3717OpFoldResult SubViewOp::fold(FoldAdaptor adaptor) {
3718 MemRefType sourceMemrefType = getSource().getType();
3719 MemRefType resultMemrefType = getResult().getType();
3721 dyn_cast_if_present<StridedLayoutAttr>(resultMemrefType.getLayout());
3723 if (resultMemrefType == sourceMemrefType &&
3724 resultMemrefType.hasStaticShape() &&
3725 (!resultLayout || resultLayout.hasStaticLayout())) {
3726 return getViewSource();
3732 if (
auto srcSubview = getViewSource().getDefiningOp<SubViewOp>()) {
3733 auto srcSizes = srcSubview.getMixedSizes();
3735 auto offsets = getMixedOffsets();
3737 auto strides = getMixedStrides();
3738 bool allStridesOne = llvm::all_of(strides,
isOneInteger);
3739 bool allSizesSame = llvm::equal(sizes, srcSizes);
3740 if (allOffsetsZero && allStridesOne && allSizesSame &&
3741 resultMemrefType == sourceMemrefType)
3742 return getViewSource();
3748FailureOr<std::optional<SmallVector<Value>>>
3749SubViewOp::bubbleDownCasts(OpBuilder &builder) {
3753void SubViewOp::inferStridedMetadataRanges(
3754 ArrayRef<StridedMetadataRange> ranges,
GetIntRangeFn getIntRange,
3756 auto isUninitialized =
3757 +[](IntegerValueRange range) {
return range.isUninitialized(); };
3760 SmallVector<IntegerValueRange> offsetOperands =
3762 if (llvm::any_of(offsetOperands, isUninitialized))
3765 SmallVector<IntegerValueRange> sizeOperands =
3767 if (llvm::any_of(sizeOperands, isUninitialized))
3770 SmallVector<IntegerValueRange> stridesOperands =
3772 if (llvm::any_of(stridesOperands, isUninitialized))
3775 StridedMetadataRange sourceRange =
3776 ranges[getSourceMutable().getOperandNumber()];
3780 ArrayRef<ConstantIntRanges> srcStrides = sourceRange.
getStrides();
3786 ConstantIntRanges offset = sourceRange.
getOffsets()[0];
3787 SmallVector<ConstantIntRanges> strides, sizes;
3789 for (
size_t i = 0, e = droppedDims.size(); i < e; ++i) {
3790 bool dropped = droppedDims.test(i);
3792 ConstantIntRanges off =
3803 sizes.push_back(sizeOperands[i].getValue());
3806 setMetadata(getResult(),
3808 SmallVector<ConstantIntRanges>({std::move(offset)}),
3809 std::move(sizes), std::move(strides)));
3816void TransposeOp::getAsmResultNames(
3818 setNameFn(getResult(),
"transpose");
3824 auto originalSizes = memRefType.getShape();
3825 auto [originalStrides, offset] = memRefType.getStridesAndOffset();
3826 assert(originalStrides.size() ==
static_cast<unsigned>(memRefType.getRank()));
3835 StridedLayoutAttr::get(memRefType.getContext(), offset, strides));
3838Value TransposeOp::getViewSource() {
return getIn(); }
3840void TransposeOp::build(OpBuilder &
b, OperationState &
result, Value in,
3841 AffineMapAttr permutation,
3842 ArrayRef<NamedAttribute> attrs) {
3843 auto permutationMap = permutation.getValue();
3844 assert(permutationMap);
3846 auto memRefType = llvm::cast<MemRefType>(in.
getType());
3850 buildPropertiesAndDiscardableAttributes(
result, attrs);
3851 result.getOrAddProperties<Properties>().permutation = permutation;
3853 result.addTypes(resultType);
3857void TransposeOp::print(OpAsmPrinter &p) {
3858 p <<
" " << getIn() <<
" " << getPermutation();
3860 {getPermutationAttrStrName()});
3861 p <<
" : " << getIn().getType() <<
" to " <<
getType();
3864ParseResult TransposeOp::parse(OpAsmParser &parser, OperationState &
result) {
3865 OpAsmParser::UnresolvedOperand in;
3866 AffineMap permutation;
3867 MemRefType srcType, dstType;
3876 result.addAttribute(TransposeOp::getPermutationAttrStrName(),
3877 AffineMapAttr::get(permutation));
3881LogicalResult TransposeOp::verify() {
3883 return emitOpError(
"expected a permutation map");
3884 if (getPermutation().getNumDims() != getIn().
getType().getRank())
3885 return emitOpError(
"expected a permutation map of same rank as the input");
3887 auto srcType = llvm::cast<MemRefType>(getIn().
getType());
3888 auto resultType = llvm::cast<MemRefType>(
getType());
3890 .canonicalizeStridedLayout();
3892 if (resultType.canonicalizeStridedLayout() != canonicalResultType)
3893 return emitOpError(
"result type ")
3895 <<
" is not equivalent to the canonical transposed input type "
3896 << canonicalResultType;
3900OpFoldResult TransposeOp::fold(FoldAdaptor) {
3903 if (getPermutation().isIdentity() &&
getType() == getIn().
getType())
3907 if (
auto otherTransposeOp = getIn().getDefiningOp<memref::TransposeOp>()) {
3908 AffineMap composedPermutation =
3909 getPermutation().compose(otherTransposeOp.getPermutation());
3910 getInMutable().assign(otherTransposeOp.getIn());
3911 setPermutation(composedPermutation);
3917FailureOr<std::optional<SmallVector<Value>>>
3918TransposeOp::bubbleDownCasts(OpBuilder &builder) {
3926void ViewOp::getAsmResultNames(
function_ref<
void(Value, StringRef)> setNameFn) {
3927 setNameFn(getResult(),
"view");
3930LogicalResult ViewOp::verify() {
3931 auto baseType = llvm::cast<MemRefType>(getOperand(0).
getType());
3935 if (!baseType.getLayout().isIdentity())
3936 return emitError(
"unsupported map for base memref type ") << baseType;
3939 if (!viewType.getLayout().isIdentity())
3940 return emitError(
"unsupported map for result memref type ") << viewType;
3943 if (baseType.getMemorySpace() != viewType.getMemorySpace())
3944 return emitError(
"different memory spaces specified for base memref "
3946 << baseType <<
" and view memref type " << viewType;
3955Value ViewOp::getViewSource() {
return getSource(); }
3957OpFoldResult ViewOp::fold(FoldAdaptor adaptor) {
3958 MemRefType sourceMemrefType = getSource().getType();
3959 MemRefType resultMemrefType = getResult().getType();
3961 if (resultMemrefType == sourceMemrefType &&
3962 resultMemrefType.hasStaticShape() &&
isZeroInteger(getByteShift()))
3963 return getViewSource();
3968SmallVector<OpFoldResult> ViewOp::getMixedSizes() {
3969 SmallVector<OpFoldResult>
result;
3973 if (ShapedType::isDynamic(dim)) {
3974 result.push_back(getSizes()[ctr++]);
3976 result.push_back(
b.getIndexAttr(dim));
3988 SmallVectorImpl<Value> &foldedDynamicSizes) {
3989 SmallVector<int64_t> staticShape(type.getShape());
3990 assert(type.getNumDynamicDims() == dynamicSizes.size() &&
3991 "incorrect number of dynamic sizes");
3995 for (
auto [dim, dimSize] : llvm::enumerate(type.getShape())) {
3996 if (ShapedType::isStatic(dimSize))
3999 Value dynamicSize = dynamicSizes[ctr++];
4002 if (cst.value() < 0) {
4003 foldedDynamicSizes.push_back(dynamicSize);
4006 staticShape[dim] = cst.value();
4008 foldedDynamicSizes.push_back(dynamicSize);
4012 return MemRefType::Builder(type).setShape(staticShape);
4026struct ViewOpShapeFolder :
public OpRewritePattern<ViewOp> {
4029 LogicalResult matchAndRewrite(ViewOp viewOp,
4030 PatternRewriter &rewriter)
const override {
4031 SmallVector<Value> foldedDynamicSizes;
4032 MemRefType resultType = viewOp.getType();
4034 resultType, viewOp.getSizes(), foldedDynamicSizes);
4037 if (foldedMemRefType == resultType)
4041 auto newViewOp = ViewOp::create(rewriter, viewOp.getLoc(), foldedMemRefType,
4042 viewOp.getSource(), viewOp.getByteShift(),
4043 foldedDynamicSizes);
4051struct ViewOpMemrefCastFolder :
public OpRewritePattern<ViewOp> {
4054 LogicalResult matchAndRewrite(ViewOp viewOp,
4055 PatternRewriter &rewriter)
const override {
4056 auto memrefCastOp = viewOp.getSource().getDefiningOp<CastOp>();
4061 viewOp, viewOp.getType(), memrefCastOp.getSource(),
4062 viewOp.getByteShift(), viewOp.getSizes());
4068void ViewOp::getCanonicalizationPatterns(RewritePatternSet &results,
4069 MLIRContext *context) {
4070 results.
add<ViewOpShapeFolder, ViewOpMemrefCastFolder>(context);
4073FailureOr<std::optional<SmallVector<Value>>>
4074ViewOp::bubbleDownCasts(OpBuilder &builder) {
4082LogicalResult AtomicRMWOp::verify() {
4083 switch (getKind()) {
4084 case arith::AtomicRMWKind::addf:
4085 case arith::AtomicRMWKind::maximumf:
4086 case arith::AtomicRMWKind::minimumf:
4087 case arith::AtomicRMWKind::mulf:
4088 if (!llvm::isa<FloatType>(getValue().
getType()))
4089 return emitOpError() <<
"with kind '"
4090 << arith::stringifyAtomicRMWKind(getKind())
4091 <<
"' expects a floating-point type";
4093 case arith::AtomicRMWKind::addi:
4094 case arith::AtomicRMWKind::maxs:
4095 case arith::AtomicRMWKind::maxu:
4096 case arith::AtomicRMWKind::mins:
4097 case arith::AtomicRMWKind::minu:
4098 case arith::AtomicRMWKind::muli:
4099 case arith::AtomicRMWKind::ori:
4100 case arith::AtomicRMWKind::xori:
4101 case arith::AtomicRMWKind::andi:
4102 if (!llvm::isa<IntegerType>(getValue().
getType()))
4103 return emitOpError() <<
"with kind '"
4104 << arith::stringifyAtomicRMWKind(getKind())
4105 <<
"' expects an integer type";
4113OpFoldResult AtomicRMWOp::fold(FoldAdaptor adaptor) {
4117 return OpFoldResult();
4120FailureOr<std::optional<SmallVector<Value>>>
4121AtomicRMWOp::bubbleDownCasts(OpBuilder &builder) {
4128std::optional<SmallVector<Value>>
4129AtomicRMWOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
4132 getMemrefMutable().assign(newMemref);
4133 getIndicesMutable().assign(newIndices);
4135 return std::nullopt;
4142#define GET_OP_CLASSES
4143#include "mlir/Dialect/MemRef/IR/MemRefOps.cpp.inc"
getNumOperands() - 1))) return failure()
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static bool hasSideEffects(Operation *op)
static bool isPermutation(const std::vector< PermutationTy > &permutation)
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
static LogicalResult foldCopyOfCast(CopyOp op)
If the source/target of a CopyOp is a CastOp that does not modify the shape and element type,...
static void constifyIndexValues(SmallVectorImpl< OpFoldResult > &values, ArrayRef< int64_t > constValues)
Helper function that sets values[i] to constValues[i] if the latter is a static value,...
static void printGlobalMemrefOpTypeAndInitialValue(OpAsmPrinter &p, GlobalOp op, TypeAttr type, Attribute initialValue)
static LogicalResult verifyCollapsedShape(Operation *op, ArrayRef< int64_t > collapsedShape, ArrayRef< int64_t > expandedShape, ArrayRef< ReassociationIndices > reassociation, bool allowMultipleDynamicDimsPerGroup)
Helper function for verifying the shape of ExpandShapeOp and ResultShapeOp result and operand.
static bool isOpItselfPotentialAutomaticAllocation(Operation *op)
Given an operation, return whether this op itself could allocate an AutomaticAllocationScopeResource.
static MemRefType inferTransposeResultType(MemRefType memRefType, AffineMap permutationMap)
Build a strided memref type by applying permutationMap to memRefType.
static ParseResult parseBoolAttr(OpAsmParser &parser, BoolAttr &result)
static bool isGuaranteedAutomaticAllocation(Operation *op)
Given an operation, return whether this op is guaranteed to allocate an AutomaticAllocationScopeResou...
static FailureOr< llvm::SmallBitVector > computeMemRefRankReductionMaskByStrides(MemRefType originalType, MemRefType reducedType, ArrayRef< int64_t > originalStrides, ArrayRef< int64_t > candidateStrides, llvm::SmallBitVector unusedDims)
Returns the set of source dimensions that are dropped in a rank reduction.
static FailureOr< StridedLayoutAttr > computeExpandedLayoutMap(MemRefType srcType, ArrayRef< int64_t > resultShape, ArrayRef< ReassociationIndices > reassociation)
Compute the layout map after expanding a given source MemRef type with the specified reassociation in...
static bool haveCompatibleOffsets(MemRefType t1, MemRefType t2)
Return true if t1 and t2 have equal offsets (both dynamic or of same static value).
static void printBoolAttr(OpAsmPrinter &printer, Operation *, BoolAttr attr)
static bool replaceConstantUsesOf(OpBuilder &rewriter, Location loc, Container values, ArrayRef< OpFoldResult > maybeConstants)
Helper function to perform the replacement of all constant uses of values by a materialized constant ...
static LogicalResult produceSubViewErrorMsg(SliceVerificationResult result, SubViewOp op, Type expectedType)
static MemRefType getCanonicalSubViewResultType(MemRefType currentResultType, MemRefType currentSourceType, MemRefType sourceType, ArrayRef< OpFoldResult > mixedOffsets, ArrayRef< OpFoldResult > mixedSizes, ArrayRef< OpFoldResult > mixedStrides)
Compute the canonical result type of a SubViewOp.
static ParseResult parseGlobalMemrefOpTypeAndInitialValue(OpAsmParser &parser, TypeAttr &typeAttr, Attribute &initialValue)
static std::tuple< MemorySpaceCastOpInterface, PtrLikeTypeInterface, Type > getMemorySpaceCastInfo(BaseMemRefType resultTy, Value src)
Helper function to retrieve a lossless memory-space cast, and the corresponding new result memref typ...
static FailureOr< llvm::SmallBitVector > computeMemRefRankReductionMask(MemRefType originalType, MemRefType reducedType, ArrayRef< OpFoldResult > sizes)
Given the originalType and a candidateReducedType whose shape is assumed to be a subset of originalTy...
static bool isTrivialSubViewOp(SubViewOp subViewOp)
Helper method to check if a subview operation is trivially a no-op.
static bool lastNonTerminatorInRegion(Operation *op)
Return whether this op is the last non terminating op in a region.
static std::map< int64_t, unsigned > getNumOccurences(ArrayRef< int64_t > vals)
Return a map with key being elements in vals and data being number of occurences of it.
static bool haveCompatibleStrides(MemRefType t1, MemRefType t2, const llvm::SmallBitVector &droppedDims)
Return true if t1 and t2 have equal strides (both dynamic or of same static value).
static FailureOr< StridedLayoutAttr > computeCollapsedLayoutMap(MemRefType srcType, ArrayRef< ReassociationIndices > reassociation, bool strict=false)
Compute the layout map after collapsing a given source MemRef type with the specified reassociation i...
static FailureOr< std::optional< SmallVector< Value > > > bubbleDownCastsPassthroughOpImpl(ConcreteOpTy op, OpBuilder &builder, OpOperand &src)
Implementation of bubbleDownCasts method for memref operations that return a single memref result.
static FailureOr< llvm::SmallBitVector > computeMemRefRankReductionMaskByPosition(MemRefType originalType, MemRefType reducedType, ArrayRef< OpFoldResult > sizes)
Returns the set of source dimensions that are dropped in a rank reduction.
static LogicalResult verifyAllocLikeOp(AllocLikeOp op)
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 RankedTensorType foldDynamicToStaticDimSizes(RankedTensorType type, ValueRange dynamicSizes, SmallVector< Value > &foldedDynamicSizes)
Given a ranked tensor type and a range of values that defines its dynamic dimension sizes,...
static llvm::SmallBitVector getDroppedDims(ArrayRef< int64_t > reducedShape, ArrayRef< OpFoldResult > mixedSizes)
Compute the dropped dimensions of a rank-reducing tensor.extract_slice op or rank-extending tensor....
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
@ Square
Square brackets surrounding zero or more operands.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalEqual()=0
Parse a = token if present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseAffineMap(AffineMap &map)=0
Parse an affine map instance into 'map'.
ParseResult addTypeToList(Type type, SmallVectorImpl< Type > &result)
Add the specified type to the end of the specified type list and return success.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
ParseResult parseKeywordType(const char *keyword, Type &result)
Parse a keyword followed by a type.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
virtual void printAttributeWithoutType(Attribute attr)
Print the given attribute without its type.
virtual void printAttribute(Attribute attr)
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.
FailureOr< PtrLikeTypeInterface > clonePtrWith(Attribute memorySpace, std::optional< Type > elementType) const
Clone this type with the given memory space and element type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
Block represents an ordered list of Operations.
Operation * getTerminator()
Get the terminator operation of this block.
bool mightHaveTerminator()
Return "true" if this block might have a terminator.
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
This class is a general helper class for creating context-global objects like types,...
IntegerAttr getIndexAttr(int64_t value)
IntegerType getIntegerType(unsigned width)
BoolAttr getBoolAttr(bool value)
IRValueT get() const
Return the current value being used by this operand.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This is a builder type that keeps local references to arguments.
Builder & setShape(ArrayRef< int64_t > newShape)
Builder & setLayout(MemRefLayoutAttrInterface newLayout)
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
ParseResult parseTrailingOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None)
Parse zero or more trailing SSA comma-separated trailing operand references with a specified surround...
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
void printOperands(OperandRange operands)
Print a comma separated range of operation operands out of line to avoid instantiating the range iter...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
This class helps build Operations.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
This class represents a single result from folding an operation.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
A trait of region holding operations that define a new scope for automatic allocations,...
This trait indicates that the memory effects of an operation includes the effects of operations neste...
type_range getType() const
Operation is the basic unit of execution within MLIR.
void replaceUsesOfWith(Value from, Value to)
Replace any uses of 'from' with 'to' within this operation.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Block * getBlock()
Returns the operation block that contains this operation.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
MutableArrayRef< OpOperand > getOpOperands()
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
Region * getParentRegion()
Returns the region to which the instruction belongs.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
Type-safe wrapper around a void* for passing properties, including the properties structs of operatio...
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
bool isParent() const
Returns true if branching from the parent op.
This class provides an abstraction over the different types of ranges over Regions.
This class represents a successor of a region.
bool isOperation() const
Return true if the successor is an operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
bool hasOneBlock()
Return true if this region has exactly one block.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
virtual void inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
virtual Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
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...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
This class provides an abstraction over the different types of ranges over Values.
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static WalkResult advance()
static WalkResult interrupt()
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
constexpr void enumerate(std::tuple< Tys... > &tuple, CallbackT &&callback)
FailureOr< std::optional< SmallVector< Value > > > bubbleDownInPlaceMemorySpaceCastImpl(OpOperand &operand, ValueRange results)
Tries to bubble-down inplace a MemorySpaceCastOpInterface operation referenced by operand.
ConstantIntRanges inferAdd(ArrayRef< ConstantIntRanges > argRanges, OverflowFlags ovfFlags=OverflowFlags::None)
ConstantIntRanges inferMul(ArrayRef< ConstantIntRanges > argRanges, OverflowFlags ovfFlags=OverflowFlags::None)
ConstantIntRanges inferShapedDimOpInterface(ShapedDimOpInterface op, const IntegerValueRange &maybeDim)
Returns the integer range for the result of a ShapedDimOpInterface given the optional inferred ranges...
Type getTensorTypeFromMemRefType(Type type)
Return an unranked/ranked tensor type for the given unranked/ranked memref type.
OpFoldResult getMixedSize(OpBuilder &builder, Location loc, Value value, int64_t dim)
Return the dimension of the given memref value.
LogicalResult foldMemRefCast(Operation *op, Value inner=nullptr)
This is a common utility used for patterns of the form "someop(memref.cast) -> someop".
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given memref value.
Value createCanonicalRankReducingSubViewOp(OpBuilder &b, Location loc, Value memref, ArrayRef< int64_t > targetShape)
Create a rank-reducing SubViewOp @[0 .
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
DynamicAPInt getIndex(const ConeV &cone)
Get the index of a cone, i.e., the volume of the parallelepiped spanned by its generators,...
Value constantIndex(OpBuilder &builder, Location loc, int64_t i)
Generates a constant of index type.
MemRefType getMemRefType(T &&t)
Convenience method to abbreviate casting getType().
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
detail::constant_int_value_binder m_ConstantInt(IntegerAttr::ValueType *bind_value)
Matches a constant holding a scalar/vector/tensor integer (splat) and writes the integer value to bin...
SliceVerificationResult
Enum that captures information related to verifier error conditions on slice insert/extract type of o...
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
raw_ostream & operator<<(raw_ostream &os, const AliasResult &result)
llvm::function_ref< void(Value, const IntegerValueRange &)> SetIntLatticeFn
Similar to SetIntRangeFn, but operating on IntegerValueRange lattice values.
SliceBoundsVerificationResult verifyInBoundsSlice(ArrayRef< int64_t > shape, ArrayRef< int64_t > staticOffsets, ArrayRef< int64_t > staticSizes, ArrayRef< int64_t > staticStrides, bool generateErrorMessage=false)
Verify that the offsets/sizes/strides-style access into the given shape is in-bounds.
LogicalResult verifyDynamicDimensionCount(Operation *op, ShapedType type, ValueRange dynamicSizes)
Verify that the number of dynamic size operands matches the number of dynamic dimensions in the shape...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
SmallVector< Range, 8 > getOrCreateRanges(OffsetSizeAndStrideOpInterface op, OpBuilder &b, Location loc)
Return the list of Range (i.e.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
SmallVector< AffineMap, 4 > getSymbolLessAffineMaps(ArrayRef< ReassociationExprs > reassociation)
Constructs affine maps out of Array<Array<AffineExpr>>.
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
OpFoldResult foldReshapeOp(ReshapeOpTy reshapeOp, ArrayRef< Attribute > operands)
bool hasValidSizesOffsets(SmallVector< int64_t > sizesOrOffsets)
Helper function to check whether the passed in sizes or offsets are valid.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
SmallVector< IntegerValueRange > getIntValueRanges(ArrayRef< OpFoldResult > values, GetIntRangeFn getIntRange, int32_t indexBitwidth)
Helper function to collect the integer range values of an array of op fold results.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
bool hasValidStrides(SmallVector< int64_t > strides)
Helper function to check whether the passed in strides are valid.
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
SmallVector< SmallVector< AffineExpr, 2 >, 2 > convertReassociationIndicesToExprs(MLIRContext *context, ArrayRef< ReassociationIndices > reassociationIndices)
Convert reassociation indices to affine expressions.
std::optional< SmallVector< OpFoldResult > > inferExpandShapeOutputShape(OpBuilder &b, Location loc, ShapedType expandedType, ArrayRef< ReassociationIndices > reassociation, ArrayRef< OpFoldResult > inputShape)
Infer the output shape for a {memref|tensor}.expand_shape when it is possible to do so.
LogicalResult verifyElementTypesMatch(Operation *op, ShapedType lhs, ShapedType rhs, StringRef lhsName, StringRef rhsName)
Verify that two shaped types have matching element types.
SmallVector< T > applyPermutationMap(AffineMap map, llvm::ArrayRef< T > source)
Apply a permutation from map to source and return the result.
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
function_ref< void(Value, const StridedMetadataRange &)> SetStridedMetadataRangeFn
Callback function type for setting the strided metadata of a value.
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...
SmallVector< int64_t, 2 > ReassociationIndices
SliceVerificationResult isRankReducedType(ShapedType originalType, ShapedType candidateReducedType)
Check if originalType can be rank reduced to candidateReducedType type by dropping some dimensions wi...
ArrayAttr getReassociationIndicesAttribute(Builder &b, ArrayRef< ReassociationIndices > reassociation)
Wraps a list of reassociations in an ArrayAttr.
llvm::function_ref< Fn > function_ref
bool isOneInteger(OpFoldResult v)
Return true if v is an IntegerAttr with value 1.
std::pair< SmallVector< int64_t >, SmallVector< Value > > decomposeMixedValues(ArrayRef< OpFoldResult > mixedValues)
Decompose a vector of mixed static or dynamic values into the corresponding pair of arrays.
LogicalResult verifyReassociationIndicesNotEmpty(ReshapeOpTy op)
Verify that none of the reassociation groups is empty.
function_ref< IntegerValueRange(Value)> GetIntRangeFn
Helper callback type to get the integer range of a value.
Move allocations into an allocation scope, if it is legal to move them (e.g.
LogicalResult matchAndRewrite(AllocaScopeOp op, PatternRewriter &rewriter) const override
Inline an AllocaScopeOp if either the direct parent is an allocation scope or it contains no allocati...
LogicalResult matchAndRewrite(AllocaScopeOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(CollapseShapeOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(ExpandShapeOp op, PatternRewriter &rewriter) const override
A canonicalizer wrapper to replace SubViewOps.
void operator()(PatternRewriter &rewriter, SubViewOp op, SubViewOp newOp)
Return the canonical type of the result of a subview.
MemRefType operator()(SubViewOp op, ArrayRef< OpFoldResult > mixedOffsets, ArrayRef< OpFoldResult > mixedSizes, ArrayRef< OpFoldResult > mixedStrides)
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.
Represents a range (offset, size, and stride) where each element of the triple may be dynamic or stat...
static SaturatedInteger wrap(int64_t v)
bool isValid
If set to "true", the slice bounds verification was successful.
std::string errorMessage
An error message that can be printed during op verification.
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.