25#include "llvm/ADT/STLExtras.h"
26#include "llvm/ADT/SmallBitVector.h"
31#define GEN_PASS_DEF_EXPANDSTRIDEDMETADATAPASS
32#include "mlir/Dialect/MemRef/Transforms/Passes.h.inc"
41struct StridedMetadata {
44 SmallVector<OpFoldResult> sizes;
45 SmallVector<OpFoldResult> strides;
59static FailureOr<StridedMetadata>
61 memref::SubViewOp subview) {
64 Value source = subview.getSource();
65 auto sourceType = cast<MemRefType>(source.
getType());
66 unsigned sourceRank = sourceType.getRank();
68 auto newExtractStridedMetadata =
69 memref::ExtractStridedMetadataOp::create(rewriter, origLoc, source);
71 auto [sourceStrides, sourceOffset] = sourceType.getStridesAndOffset();
73 auto [resultStrides, resultOffset] = subview.getType().getStridesAndOffset();
81 auto origStrides = newExtractStridedMetadata.getStrides();
89 values[0] = ShapedType::isDynamic(sourceOffset)
91 : rewriter.getIndexAttr(sourceOffset);
96 for (
unsigned i = 0; i < sourceRank; ++i) {
99 ShapedType::isDynamic(sourceStrides[i])
103 rewriter, origLoc, s0 * s1, {subStrides[i], origStride}));
106 unsigned baseIdxForDim = 1 + 2 * i;
107 unsigned subOffsetForDim = baseIdxForDim;
108 unsigned origStrideForDim = baseIdxForDim + 1;
109 expr = expr + symbols[subOffsetForDim] * symbols[origStrideForDim];
110 values[subOffsetForDim] = subOffsets[i];
111 values[origStrideForDim] = origStride;
121 if (computedOffset && ShapedType::isStatic(resultOffset))
122 assert(*computedOffset == resultOffset &&
123 "mismatch between computed offset and result type offset");
129 auto subType = cast<MemRefType>(subview.getType());
130 unsigned subRank = subType.getRank();
139 llvm::SmallBitVector droppedDims = subview.getDroppedDims();
142 finalSizes.reserve(subRank);
145 finalStrides.reserve(subRank);
151 for (
unsigned i = 0; i < sourceRank; ++i) {
152 if (droppedDims.test(i))
155 finalSizes.push_back(subSizes[i]);
156 finalStrides.push_back(strides[i]);
161 if (computedStride && ShapedType::isStatic(resultStrides[
j]))
162 assert(*computedStride == resultStrides[
j] &&
163 "mismatch between computed stride and result type stride");
167 assert(finalSizes.size() == subRank &&
168 "Should have populated all the values at this point");
169 return StridedMetadata{newExtractStridedMetadata.getBaseBuffer(), finalOffset,
170 finalSizes, finalStrides};
189 using OpRewritePattern<memref::SubViewOp>::OpRewritePattern;
191 LogicalResult matchAndRewrite(memref::SubViewOp subview,
192 PatternRewriter &rewriter)
const override {
193 FailureOr<StridedMetadata> stridedMetadata =
194 resolveSubviewStridedMetadata(rewriter, subview);
195 if (
failed(stridedMetadata)) {
197 "failed to resolve subview metadata");
201 subview.getType(), stridedMetadata->offset, stridedMetadata->sizes,
202 stridedMetadata->strides);
203 auto foldedSubview = memref::ReinterpretCastOp::create(
204 rewriter, subview.getLoc(), resultType, stridedMetadata->basePtr,
205 stridedMetadata->offset, stridedMetadata->sizes,
206 stridedMetadata->strides);
207 if (resultType == subview.getType()) {
208 rewriter.
replaceOp(subview, foldedSubview);
231struct ExtractStridedMetadataOpSubviewFolder
236 PatternRewriter &rewriter)
const override {
237 auto subviewOp = op.getSource().getDefiningOp<memref::SubViewOp>();
241 FailureOr<StridedMetadata> stridedMetadata =
242 resolveSubviewStridedMetadata(rewriter, subviewOp);
243 if (
failed(stridedMetadata)) {
245 op,
"failed to resolve metadata in terms of source subview op");
247 Location loc = subviewOp.getLoc();
248 SmallVector<Value> results;
249 results.reserve(subviewOp.getType().getRank() * 2 + 2);
250 results.push_back(stridedMetadata->basePtr);
252 stridedMetadata->offset));
256 stridedMetadata->strides));
276getExpandedSizes(memref::ExpandShapeOp expandShape,
OpBuilder &builder,
279 expandShape.getReassociationIndices()[groupId];
280 assert(!reassocGroup.empty() &&
281 "Reassociation group should have at least one dimension");
285 for (
auto index : reassocGroup)
286 expandedSizes.push_back(outputShape[
index]);
288 return expandedSizes;
317 expandShape.getReassociationIndices()[groupId];
318 assert(!reassocGroup.empty() &&
319 "Reassociation group should have at least one dimension");
321 unsigned groupSize = reassocGroup.size();
322 Location loc = expandShape.getLoc();
331 Value source = expandShape.getSrc();
332 auto sourceType = cast<MemRefType>(source.
getType());
333 auto [strides, offset] = sourceType.getStridesAndOffset();
335 OpFoldResult origStride = ShapedType::isDynamic(strides[groupId])
336 ? origStrides[groupId]
343 for (
int i = groupSize - 1; i >= 0; --i) {
344 expandedStrides[i] = currentStride;
345 currentStride =
mul(currentStride, outputShape[reassocGroup[i]]);
348 return expandedStrides;
365 unsigned numberOfSymbols = 0;
366 unsigned groupSize =
indices.size();
367 for (
unsigned i = 0; i < groupSize; ++i) {
371 int64_t maybeConstant = maybeConstants[srcIdx];
373 inputValues.push_back(isDynamic(maybeConstant)
397getCollapsedSize(memref::CollapseShapeOp collapseShape,
OpBuilder &builder,
401 MemRefType collapseShapeType = collapseShape.getResultType();
403 uint64_t size = collapseShapeType.getDimSize(groupId);
404 if (ShapedType::isStatic(size)) {
406 return collapsedSize;
412 Value source = collapseShape.getSrc();
413 auto sourceType = cast<MemRefType>(source.
getType());
416 collapseShape.getReassociationIndices()[groupId];
418 collapsedSize.push_back(getProductOfValues(
419 reassocGroup, builder, collapseShape.getLoc(), sourceType.getShape(),
420 origSizes, ShapedType::isDynamic));
422 return collapsedSize;
438getCollapsedStride(memref::CollapseShapeOp collapseShape,
OpBuilder &builder,
442 collapseShape.getReassociationIndices()[groupId];
443 assert(!reassocGroup.empty() &&
444 "Reassociation group should have at least one dimension");
446 Value source = collapseShape.getSrc();
447 auto sourceType = cast<MemRefType>(source.
getType());
449 auto [strides, offset] = sourceType.getStridesAndOffset();
454 for (
int64_t currentDim : reassocGroup) {
460 if (srcShape[currentDim] == 1)
463 int64_t currentStride = strides[currentDim];
464 lastValidStride = ShapedType::isDynamic(currentStride)
465 ? origStrides[currentDim]
468 if (!lastValidStride) {
471 MemRefType collapsedType = collapseShape.getResultType();
472 auto [collapsedStrides, collapsedOffset] =
473 collapsedType.getStridesAndOffset();
474 int64_t finalStride = collapsedStrides[groupId];
475 if (ShapedType::isDynamic(finalStride)) {
478 for (
int64_t currentDim : reassocGroup) {
479 assert(srcShape[currentDim] == 1 &&
480 "We should be dealing with 1x1x...x1");
482 if (ShapedType::isDynamic(strides[currentDim]))
483 return {origStrides[currentDim]};
485 llvm_unreachable(
"We should have found a dynamic stride");
490 return {lastValidStride};
503template <
typename ReassociativeReshapeLikeOp>
504static FailureOr<StridedMetadata> resolveReshapeStridedMetadata(
505 RewriterBase &rewriter, ReassociativeReshapeLikeOp reshape,
514 getReshapedStrides) {
517 Location origLoc = reshape.getLoc();
518 Value source = reshape.getSrc();
519 auto sourceType = cast<MemRefType>(source.
getType());
520 unsigned sourceRank = sourceType.getRank();
522 auto newExtractStridedMetadata =
523 memref::ExtractStridedMetadataOp::create(rewriter, origLoc, source);
526 auto [strides, offset] = sourceType.getStridesAndOffset();
527 MemRefType reshapeType = reshape.getResultType();
528 unsigned reshapeRank = reshapeType.getRank();
531 ShapedType::isDynamic(offset)
533 : rewriter.getIndexAttr(offset);
536 if (sourceRank == 0) {
538 return StridedMetadata{newExtractStridedMetadata.getBaseBuffer(), offsetOfr,
543 finalSizes.reserve(reshapeRank);
545 finalStrides.reserve(reshapeRank);
552 unsigned idx = 0, endIdx = reshape.getReassociationIndices().size();
553 for (; idx != endIdx; ++idx) {
555 getReshapedSizes(reshape, rewriter, origSizes, idx);
557 reshape, rewriter, origSizes, origStrides, idx);
559 unsigned groupSize = reshapedSizes.size();
560 for (
unsigned i = 0; i < groupSize; ++i) {
561 finalSizes.push_back(reshapedSizes[i]);
562 finalStrides.push_back(reshapedStrides[i]);
565 assert(((isa<memref::ExpandShapeOp>(reshape) && idx == sourceRank) ||
566 (isa<memref::CollapseShapeOp>(reshape) && idx == reshapeRank)) &&
567 "We should have visited all the input dimensions");
568 assert(finalSizes.size() == reshapeRank &&
569 "We should have populated all the values");
571 return StridedMetadata{newExtractStridedMetadata.getBaseBuffer(), offsetOfr,
572 finalSizes, finalStrides};
591template <
typename ReassociativeReshapeLikeOp,
603 LogicalResult matchAndRewrite(ReassociativeReshapeLikeOp reshape,
605 FailureOr<StridedMetadata> stridedMetadata =
606 resolveReshapeStridedMetadata<ReassociativeReshapeLikeOp>(
607 rewriter, reshape, getReshapedSizes, getReshapedStrides);
608 if (
failed(stridedMetadata)) {
610 "failed to resolve reshape metadata");
613 MemRefType resultType = reshape.getResultType();
614 if (isa<memref::CollapseShapeOp>(reshape.getOperation()))
616 stridedMetadata->sizes,
617 stridedMetadata->strides);
618 auto foldedReshape = memref::ReinterpretCastOp::create(
619 rewriter, reshape.getLoc(), resultType, stridedMetadata->basePtr,
620 stridedMetadata->offset, stridedMetadata->sizes,
621 stridedMetadata->strides);
622 if (resultType == reshape.getResultType()) {
623 rewriter.
replaceOp(reshape, foldedReshape);
628 reshape, reshape.getResultType(), foldedReshape);
646struct ExtractStridedMetadataOpCollapseShapeFolder
650 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
652 auto collapseShapeOp =
653 op.getSource().getDefiningOp<memref::CollapseShapeOp>();
654 if (!collapseShapeOp)
657 FailureOr<StridedMetadata> stridedMetadata =
658 resolveReshapeStridedMetadata<memref::CollapseShapeOp>(
659 rewriter, collapseShapeOp, getCollapsedSize, getCollapsedStride);
660 if (
failed(stridedMetadata)) {
663 "failed to resolve metadata in terms of source collapse_shape op");
666 Location loc = collapseShapeOp.getLoc();
668 results.push_back(stridedMetadata->basePtr);
670 stridedMetadata->offset));
674 stridedMetadata->strides));
683struct ExtractStridedMetadataOpExpandShapeFolder
687 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
689 auto expandShapeOp = op.getSource().getDefiningOp<memref::ExpandShapeOp>();
693 FailureOr<StridedMetadata> stridedMetadata =
694 resolveReshapeStridedMetadata<memref::ExpandShapeOp>(
695 rewriter, expandShapeOp, getExpandedSizes, getExpandedStrides);
696 if (
failed(stridedMetadata)) {
698 op,
"failed to resolve metadata in terms of source expand_shape op");
701 Location loc = expandShapeOp.getLoc();
703 results.push_back(stridedMetadata->basePtr);
705 stridedMetadata->offset));
709 stridedMetadata->strides));
729template <
typename AllocLikeOp>
730struct ExtractStridedMetadataOpAllocFolder
735 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
737 auto allocLikeOp = op.getSource().getDefiningOp<AllocLikeOp>();
741 auto memRefType = cast<MemRefType>(allocLikeOp.getResult().getType());
742 if (!memRefType.getLayout().isIdentity())
744 allocLikeOp,
"alloc-like operations should have been normalized");
747 int rank = memRefType.getRank();
750 ValueRange dynamic = allocLikeOp.getDynamicSizes();
753 unsigned dynamicPos = 0;
754 for (
int64_t size : memRefType.getShape()) {
755 if (ShapedType::isDynamic(size))
756 sizes.push_back(dynamic[dynamicPos++]);
764 unsigned symbolNumber = 0;
765 for (
int i = rank - 2; i >= 0; --i) {
767 assert(i + 1 + symbolNumber == sizes.size() &&
768 "The ArrayRef should encompass the last #symbolNumber sizes");
771 sizesInvolvedInStride);
776 results.reserve(rank * 2 + 2);
778 auto baseBufferType = cast<MemRefType>(op.getBaseBuffer().getType());
780 if (op.getBaseBuffer().use_empty()) {
781 results.push_back(
nullptr);
783 if (allocLikeOp.getType() == baseBufferType)
784 results.push_back(allocLikeOp);
786 results.push_back(memref::ReinterpretCastOp::create(
787 rewriter, loc, baseBufferType, allocLikeOp, offset,
820struct ExtractStridedMetadataOpGetGlobalFolder
823 using OpRewritePattern<memref::ExtractStridedMetadataOp>::OpRewritePattern;
826 PatternRewriter &rewriter)
const override {
827 auto getGlobalOp = op.getSource().getDefiningOp<memref::GetGlobalOp>();
831 auto memRefType = cast<MemRefType>(getGlobalOp.getResult().getType());
832 if (!memRefType.getLayout().isIdentity()) {
835 "get-global operation result should have been normalized");
838 Location loc = op.getLoc();
839 int rank = memRefType.getRank();
842 ArrayRef<int64_t> sizes = memRefType.getShape();
843 assert(!llvm::any_of(sizes, ShapedType::isDynamic) &&
844 "unexpected dynamic shape for result of `memref.get_global` op");
850 SmallVector<Value> results;
851 results.reserve(rank * 2 + 2);
853 auto baseBufferType = cast<MemRefType>(op.getBaseBuffer().getType());
855 if (getGlobalOp.getType() == baseBufferType)
856 results.push_back(getGlobalOp);
858 results.push_back(memref::ReinterpretCastOp::create(
859 rewriter, loc, baseBufferType, getGlobalOp, offset,
861 ArrayRef<int64_t>()));
866 for (
auto size : sizes)
869 for (
auto stride : strides)
888struct ExtractStridedMetadataOpAssumeAlignmentFolder
891 using OpRewritePattern<memref::ExtractStridedMetadataOp>::OpRewritePattern;
893 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
894 PatternRewriter &rewriter)
const override {
895 auto assumeAlignmentOp =
896 op.getSource().getDefiningOp<memref::AssumeAlignmentOp>();
897 if (!assumeAlignmentOp)
901 op, assumeAlignmentOp.getViewSource());
908class RewriteExtractAlignedPointerAsIndexOfViewLikeOp
913 matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
914 PatternRewriter &rewriter)
const override {
916 extractOp.getSource().getDefiningOp<ViewLikeOpInterface>();
920 if (!viewLikeOp || extractOp.getSource() != viewLikeOp.getViewDest() ||
921 !isa<memref::SubViewOp, memref::ReinterpretCastOp>(viewLikeOp))
924 extractOp.getSourceMutable().assign(viewLikeOp.getViewSource());
943class ExtractStridedMetadataOpReinterpretCastFolder
948 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
949 PatternRewriter &rewriter)
const override {
950 auto reinterpretCastOp = extractStridedMetadataOp.getSource()
951 .getDefiningOp<memref::ReinterpretCastOp>();
952 if (!reinterpretCastOp)
955 Location loc = extractStridedMetadataOp.getLoc();
957 SmallVector<Type> inferredReturnTypes;
958 if (
failed(extractStridedMetadataOp.inferReturnTypes(
959 rewriter.
getContext(), loc, {reinterpretCastOp.getSource()},
961 inferredReturnTypes)))
963 reinterpretCastOp,
"reinterpret_cast source's type is incompatible");
965 auto memrefType = cast<MemRefType>(reinterpretCastOp.getResult().getType());
966 unsigned rank = memrefType.getRank();
967 SmallVector<OpFoldResult> results;
968 results.resize_for_overwrite(rank * 2 + 2);
970 auto newExtractStridedMetadata = memref::ExtractStridedMetadataOp::create(
971 rewriter, loc, reinterpretCastOp.getSource());
974 results[0] = newExtractStridedMetadata.getBaseBuffer();
978 rewriter, loc, reinterpretCastOp.getMixedOffsets()[0]);
980 const unsigned sizeStartIdx = 2;
981 const unsigned strideStartIdx = sizeStartIdx + rank;
983 SmallVector<OpFoldResult> sizes = reinterpretCastOp.getMixedSizes();
984 SmallVector<OpFoldResult> strides = reinterpretCastOp.getMixedStrides();
985 for (
unsigned i = 0; i < rank; ++i) {
986 results[sizeStartIdx + i] = sizes[i];
987 results[strideStartIdx + i] = strides[i];
989 rewriter.
replaceOp(extractStridedMetadataOp,
1005class ExtractStridedMetadataOpMemorySpaceCastFolder
1010 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
1011 PatternRewriter &rewriter)
const override {
1012 Location loc = extractStridedMetadataOp.getLoc();
1013 Value source = extractStridedMetadataOp.getSource();
1014 auto memSpaceCastOp = source.
getDefiningOp<memref::MemorySpaceCastOp>();
1015 if (!memSpaceCastOp)
1017 auto newExtractStridedMetadata = memref::ExtractStridedMetadataOp::create(
1018 rewriter, loc, memSpaceCastOp.getSource());
1019 SmallVector<Value> results(newExtractStridedMetadata.getResults());
1026 if (!extractStridedMetadataOp.getBaseBuffer().use_empty()) {
1027 auto baseBuffer = results[0];
1028 auto baseBufferType = cast<MemRefType>(baseBuffer.getType());
1029 MemRefType::Builder newTypeBuilder(baseBufferType);
1030 newTypeBuilder.setMemorySpace(
1031 memSpaceCastOp.getResult().getType().getMemorySpace());
1032 results[0] = memref::MemorySpaceCastOp::create(
1033 rewriter, loc, Type{newTypeBuilder}, baseBuffer);
1035 results[0] =
nullptr;
1037 rewriter.
replaceOp(extractStridedMetadataOp, results);
1049class ExtractStridedMetadataOpExtractStridedMetadataFolder
1054 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
1055 PatternRewriter &rewriter)
const override {
1056 auto sourceExtractStridedMetadataOp =
1057 extractStridedMetadataOp.getSource()
1058 .getDefiningOp<memref::ExtractStridedMetadataOp>();
1059 if (!sourceExtractStridedMetadataOp)
1061 Location loc = extractStridedMetadataOp.getLoc();
1062 rewriter.
replaceOp(extractStridedMetadataOp,
1063 {sourceExtractStridedMetadataOp.getBaseBuffer(),
1073 patterns.
add<SubviewFolder,
1074 ReshapeFolder<memref::ExpandShapeOp, getExpandedSizes,
1075 getExpandedStrides>,
1076 ReshapeFolder<memref::CollapseShapeOp, getCollapsedSize,
1077 getCollapsedStride>,
1078 ExtractStridedMetadataOpAllocFolder<memref::AllocOp>,
1079 ExtractStridedMetadataOpAllocFolder<memref::AllocaOp>,
1080 ExtractStridedMetadataOpCollapseShapeFolder,
1081 ExtractStridedMetadataOpExpandShapeFolder,
1082 ExtractStridedMetadataOpGetGlobalFolder,
1083 RewriteExtractAlignedPointerAsIndexOfViewLikeOp,
1084 ExtractStridedMetadataOpReinterpretCastFolder,
1085 ExtractStridedMetadataOpSubviewFolder,
1086 ExtractStridedMetadataOpMemorySpaceCastFolder,
1087 ExtractStridedMetadataOpAssumeAlignmentFolder,
1088 ExtractStridedMetadataOpExtractStridedMetadataFolder>(
1094 patterns.
add<ExtractStridedMetadataOpAllocFolder<memref::AllocOp>,
1095 ExtractStridedMetadataOpAllocFolder<memref::AllocaOp>,
1096 ExtractStridedMetadataOpCollapseShapeFolder,
1097 ExtractStridedMetadataOpExpandShapeFolder,
1098 ExtractStridedMetadataOpGetGlobalFolder,
1099 ExtractStridedMetadataOpSubviewFolder,
1100 RewriteExtractAlignedPointerAsIndexOfViewLikeOp,
1101 ExtractStridedMetadataOpReinterpretCastFolder,
1102 ExtractStridedMetadataOpMemorySpaceCastFolder,
1103 ExtractStridedMetadataOpAssumeAlignmentFolder,
1104 ExtractStridedMetadataOpExtractStridedMetadataFolder>(
1114struct ExpandStridedMetadataPass final
1116 ExpandStridedMetadataPass> {
1117 void runOnOperation()
override;
1122void ExpandStridedMetadataPass::runOnOperation() {
Base type for affine expression.
IntegerAttr getIndexAttr(int64_t value)
AffineExpr getAffineSymbolExpr(unsigned position)
AffineExpr getAffineConstantExpr(int64_t constant)
MLIRContext * getContext() const
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
This class represents a single result from folding an operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
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...
This class provides an abstraction over the different types of ranges over Values.
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.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
void populateResolveExtractStridedMetadataPatterns(RewritePatternSet &patterns)
Appends patterns for resolving memref.extract_strided_metadata into memref.extract_strided_metadata o...
void populateExpandStridedMetadataPatterns(RewritePatternSet &patterns)
Appends patterns for expanding memref operations that modify the metadata (sizes, offset,...
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
MemRefType updateTypeFromMetadata(MemRefType type, OpFoldResult offset, ArrayRef< OpFoldResult > sizes, ArrayRef< OpFoldResult > strides)
Returns a memref type matching the provided offset, size, and stride metadata.
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
SmallVector< int64_t > computeSuffixProduct(ArrayRef< int64_t > sizes)
Given a set of sizes, return the suffix product.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
llvm::function_ref< Fn > function_ref
void bindSymbolsList(MLIRContext *ctx, MutableArrayRef< AffineExprTy > exprs)
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const final
Wrapper around the RewritePattern method that passes the derived op type.
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.