MLIR 24.0.0git
ExpandStridedMetadata.cpp
Go to the documentation of this file.
1//===- ExpandStridedMetadata.cpp - Simplify this operation -------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9/// The pass expands memref operations that modify the metadata of a memref
10/// (sizes, offset, strides) into a sequence of easier to analyze constructs.
11/// In particular, this pass transforms operations into explicit sequence of
12/// operations that model the effect of this operation on the different
13/// metadata. This pass uses affine constructs to materialize these effects.
14//===----------------------------------------------------------------------===//
15
23#include "mlir/IR/AffineMap.h"
25#include "llvm/ADT/STLExtras.h"
26#include "llvm/ADT/SmallBitVector.h"
27#include <optional>
28
29namespace mlir {
30namespace memref {
31#define GEN_PASS_DEF_EXPANDSTRIDEDMETADATAPASS
32#include "mlir/Dialect/MemRef/Transforms/Passes.h.inc"
33} // namespace memref
34} // namespace mlir
35
36using namespace mlir;
37using namespace mlir::affine;
38
39namespace {
40
41struct StridedMetadata {
42 Value basePtr;
43 OpFoldResult offset;
44 SmallVector<OpFoldResult> sizes;
45 SmallVector<OpFoldResult> strides;
46};
47
48/// From `subview(memref, subOffset, subSizes, subStrides))` compute
49///
50/// \verbatim
51/// baseBuffer, baseOffset, baseSizes, baseStrides =
52/// extract_strided_metadata(memref)
53/// strides#i = baseStrides#i * subStrides#i
54/// offset = baseOffset + sum(subOffset#i * baseStrides#i)
55/// sizes = subSizes
56/// \endverbatim
57///
58/// and return {baseBuffer, offset, sizes, strides}
59static FailureOr<StridedMetadata>
60resolveSubviewStridedMetadata(RewriterBase &rewriter,
61 memref::SubViewOp subview) {
62 // Build a plain extract_strided_metadata(memref) from subview(memref).
63 Location origLoc = subview.getLoc();
64 Value source = subview.getSource();
65 auto sourceType = cast<MemRefType>(source.getType());
66 unsigned sourceRank = sourceType.getRank();
67
68 auto newExtractStridedMetadata =
69 memref::ExtractStridedMetadataOp::create(rewriter, origLoc, source);
70
71 auto [sourceStrides, sourceOffset] = sourceType.getStridesAndOffset();
72#ifndef NDEBUG
73 auto [resultStrides, resultOffset] = subview.getType().getStridesAndOffset();
74#endif // NDEBUG
75
76 // Compute the new strides and offset from the base strides and offset:
77 // newStride#i = baseStride#i * subStride#i
78 // offset = baseOffset + sum(subOffsets#i * newStrides#i)
80 SmallVector<OpFoldResult> subStrides = subview.getMixedStrides();
81 auto origStrides = newExtractStridedMetadata.getStrides();
82
83 // Hold the affine symbols and values for the computation of the offset.
84 SmallVector<OpFoldResult> values(2 * sourceRank + 1);
85 SmallVector<AffineExpr> symbols(2 * sourceRank + 1);
86
87 bindSymbolsList(rewriter.getContext(), MutableArrayRef{symbols});
88 AffineExpr expr = symbols.front();
89 values[0] = ShapedType::isDynamic(sourceOffset)
90 ? getAsOpFoldResult(newExtractStridedMetadata.getOffset())
91 : rewriter.getIndexAttr(sourceOffset);
92 SmallVector<OpFoldResult> subOffsets = subview.getMixedOffsets();
93
94 AffineExpr s0 = rewriter.getAffineSymbolExpr(0);
95 AffineExpr s1 = rewriter.getAffineSymbolExpr(1);
96 for (unsigned i = 0; i < sourceRank; ++i) {
97 // Compute the stride.
98 OpFoldResult origStride =
99 ShapedType::isDynamic(sourceStrides[i])
100 ? origStrides[i]
101 : OpFoldResult(rewriter.getIndexAttr(sourceStrides[i]));
102 strides.push_back(makeComposedFoldedAffineApply(
103 rewriter, origLoc, s0 * s1, {subStrides[i], origStride}));
104
105 // Build up the computation of the offset.
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;
112 }
113
114 // Compute the offset.
115 OpFoldResult finalOffset =
116 makeComposedFoldedAffineApply(rewriter, origLoc, expr, values);
117#ifndef NDEBUG
118 // Assert that the computed offset matches the offset of the result type of
119 // the subview op (if both are static).
120 std::optional<int64_t> computedOffset = getConstantIntValue(finalOffset);
121 if (computedOffset && ShapedType::isStatic(resultOffset))
122 assert(*computedOffset == resultOffset &&
123 "mismatch between computed offset and result type offset");
124#endif // NDEBUG
125
126 // The final result is <baseBuffer, offset, sizes, strides>.
127 // Thus we need 1 + 1 + subview.getRank() + subview.getRank(), to hold all
128 // the values.
129 auto subType = cast<MemRefType>(subview.getType());
130 unsigned subRank = subType.getRank();
131
132 // The sizes of the final type are defined directly by the input sizes of
133 // the subview.
134 // Moreover subviews can drop some dimensions, some strides and sizes may
135 // not end up in the final <base, offset, sizes, strides> value that we are
136 // replacing.
137 // Do the filtering here.
138 SmallVector<OpFoldResult> subSizes = subview.getMixedSizes();
139 llvm::SmallBitVector droppedDims = subview.getDroppedDims();
140
141 SmallVector<OpFoldResult> finalSizes;
142 finalSizes.reserve(subRank);
143
144 SmallVector<OpFoldResult> finalStrides;
145 finalStrides.reserve(subRank);
146
147#ifndef NDEBUG
148 // Iteration variable for result dimensions of the subview op.
149 int64_t j = 0;
150#endif // NDEBUG
151 for (unsigned i = 0; i < sourceRank; ++i) {
152 if (droppedDims.test(i))
153 continue;
154
155 finalSizes.push_back(subSizes[i]);
156 finalStrides.push_back(strides[i]);
157#ifndef NDEBUG
158 // Assert that the computed stride matches the stride of the result type of
159 // the subview op (if both are static).
160 std::optional<int64_t> computedStride = getConstantIntValue(strides[i]);
161 if (computedStride && ShapedType::isStatic(resultStrides[j]))
162 assert(*computedStride == resultStrides[j] &&
163 "mismatch between computed stride and result type stride");
164 ++j;
165#endif // NDEBUG
166 }
167 assert(finalSizes.size() == subRank &&
168 "Should have populated all the values at this point");
169 return StridedMetadata{newExtractStridedMetadata.getBaseBuffer(), finalOffset,
170 finalSizes, finalStrides};
171}
172
173/// Replace `dst = subview(memref, subOffset, subSizes, subStrides))`
174/// With
175///
176/// \verbatim
177/// baseBuffer, baseOffset, baseSizes, baseStrides =
178/// extract_strided_metadata(memref)
179/// strides#i = baseStrides#i * subSizes#i
180/// offset = baseOffset + sum(subOffset#i * baseStrides#i)
181/// sizes = subSizes
182/// dst = reinterpret_cast baseBuffer, offset, sizes, strides
183/// \endverbatim
184///
185/// In other words, get rid of the subview in that expression and canonicalize
186/// on its effects on the offset, the sizes, and the strides using affine.apply.
187struct SubviewFolder : public OpRewritePattern<memref::SubViewOp> {
188public:
189 using OpRewritePattern<memref::SubViewOp>::OpRewritePattern;
190
191 LogicalResult matchAndRewrite(memref::SubViewOp subview,
192 PatternRewriter &rewriter) const override {
193 FailureOr<StridedMetadata> stridedMetadata =
194 resolveSubviewStridedMetadata(rewriter, subview);
195 if (failed(stridedMetadata)) {
196 return rewriter.notifyMatchFailure(subview,
197 "failed to resolve subview metadata");
198 }
199
200 MemRefType resultType = updateTypeFromMetadata(
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);
209 return success();
210 }
211 // Preserve the original result type expected by existing users.
212 rewriter.replaceOpWithNewOp<memref::CastOp>(subview, subview.getType(),
213 foldedSubview);
214 return success();
215 }
216};
217
218/// Pattern to replace `extract_strided_metadata(subview)`
219/// With
220///
221/// \verbatim
222/// baseBuffer, baseOffset, baseSizes, baseStrides =
223/// extract_strided_metadata(memref)
224/// strides#i = baseStrides#i * subSizes#i
225/// offset = baseOffset + sum(subOffset#i * baseStrides#i)
226/// sizes = subSizes
227/// \verbatim
228///
229/// with `baseBuffer`, `offset`, `sizes` and `strides` being
230/// the replacements for the original `extract_strided_metadata`.
231struct ExtractStridedMetadataOpSubviewFolder
232 : OpRewritePattern<memref::ExtractStridedMetadataOp> {
234
235 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
236 PatternRewriter &rewriter) const override {
237 auto subviewOp = op.getSource().getDefiningOp<memref::SubViewOp>();
238 if (!subviewOp)
239 return failure();
240
241 FailureOr<StridedMetadata> stridedMetadata =
242 resolveSubviewStridedMetadata(rewriter, subviewOp);
243 if (failed(stridedMetadata)) {
244 return rewriter.notifyMatchFailure(
245 op, "failed to resolve metadata in terms of source subview op");
246 }
247 Location loc = subviewOp.getLoc();
248 SmallVector<Value> results;
249 results.reserve(subviewOp.getType().getRank() * 2 + 2);
250 results.push_back(stridedMetadata->basePtr);
251 results.push_back(getValueOrCreateConstantIndexOp(rewriter, loc,
252 stridedMetadata->offset));
253 results.append(
254 getValueOrCreateConstantIndexOp(rewriter, loc, stridedMetadata->sizes));
255 results.append(getValueOrCreateConstantIndexOp(rewriter, loc,
256 stridedMetadata->strides));
257 rewriter.replaceOp(op, results);
258
259 return success();
260 }
261};
262
263/// Compute the expanded sizes of the given \p expandShape for the
264/// \p groupId-th reassociation group.
265/// \p origSizes hold the sizes of the source shape as values.
266/// This is used to compute the new sizes in cases of dynamic shapes.
267///
268/// sizes#i = expandOutputShape#i
269///
270/// \post result.size() == expandShape.getReassociationIndices()[groupId].size()
271///
272/// TODO: Move this utility function directly within ExpandShapeOp. For now,
273/// this is not possible because this function uses the Affine dialect and the
274/// MemRef dialect cannot depend on the Affine dialect.
276getExpandedSizes(memref::ExpandShapeOp expandShape, OpBuilder &builder,
277 ArrayRef<OpFoldResult> origSizes, unsigned groupId) {
278 SmallVector<int64_t, 2> reassocGroup =
279 expandShape.getReassociationIndices()[groupId];
280 assert(!reassocGroup.empty() &&
281 "Reassociation group should have at least one dimension");
283 SmallVector<OpFoldResult> outputShape = expandShape.getMixedOutputShape();
285 for (auto index : reassocGroup)
286 expandedSizes.push_back(outputShape[index]);
288 return expandedSizes;
291/// Compute the expanded strides of the given \p expandShape for the
292/// \p groupId-th reassociation group.
293/// \p origStrides and \p origSizes hold respectively the strides and sizes
294/// of the source shape as values.
295/// This is used to compute the strides in cases of dynamic shapes and/or
296/// dynamic stride for this reassociation group.
297///
298/// strides#i =
299/// origStrides#reassDim * product(expandOutputShape#j, for j in
300/// reassIdx#i+1..reassIdx#i+group.size-1)
301///
302/// Where reassIdx#i is the reassociation index for at index i in \p groupId
303/// and expandOutputShape#j is taken directly from the mixed (static and
304/// dynamic) output shape
305///
306/// \post result.size() == expandShape.getReassociationIndices()[groupId].size()
307///
308/// TODO: Move this utility function directly within ExpandShapeOp. For now,
309/// this is not possible because this function uses the Affine dialect and the
310/// MemRef dialect cannot depend on the Affine dialect.
311SmallVector<OpFoldResult> getExpandedStrides(memref::ExpandShapeOp expandShape,
312 OpBuilder &builder,
314 ArrayRef<OpFoldResult> origStrides,
315 unsigned groupId) {
316 SmallVector<int64_t, 2> reassocGroup =
317 expandShape.getReassociationIndices()[groupId];
318 assert(!reassocGroup.empty() &&
319 "Reassociation group should have at least one dimension");
320
321 unsigned groupSize = reassocGroup.size();
322 Location loc = expandShape.getLoc();
323 AffineExpr s0, s1;
324 bindSymbols(builder.getContext(), s0, s1);
325 auto mul = [&](OpFoldResult v1, OpFoldResult v2) {
326 return affine::makeComposedFoldedAffineApply(builder, loc, s0 * s1,
327 {v1, v2});
328 };
329
330 // Collect the statically known information about the original stride.
331 Value source = expandShape.getSrc();
332 auto sourceType = cast<MemRefType>(source.getType());
333 auto [strides, offset] = sourceType.getStridesAndOffset();
334
335 OpFoldResult origStride = ShapedType::isDynamic(strides[groupId])
336 ? origStrides[groupId]
337 : builder.getIndexAttr(strides[groupId]);
338
339 // Fill up the expanded strides.
340 OpFoldResult currentStride = origStride;
341 SmallVector<OpFoldResult> outputShape = expandShape.getMixedOutputShape();
342 SmallVector<OpFoldResult> expandedStrides(groupSize);
343 for (int i = groupSize - 1; i >= 0; --i) {
344 expandedStrides[i] = currentStride;
345 currentStride = mul(currentStride, outputShape[reassocGroup[i]]);
346 }
347
348 return expandedStrides;
349}
350
351/// Produce an OpFoldResult object with \p builder at \p loc representing
352/// `prod(valueOrConstant#i, for i in {indices})`,
353/// where valueOrConstant#i is maybeConstant[i] when \p isDymamic is false,
354/// values[i] otherwise.
355///
356/// \pre for all index in indices: index < values.size()
357/// \pre for all index in indices: index < maybeConstants.size()
358static OpFoldResult
359getProductOfValues(ArrayRef<int64_t> indices, OpBuilder &builder, Location loc,
360 ArrayRef<int64_t> maybeConstants,
362 llvm::function_ref<bool(int64_t)> isDynamic) {
363 AffineExpr productOfValues = builder.getAffineConstantExpr(1);
364 SmallVector<OpFoldResult> inputValues;
365 unsigned numberOfSymbols = 0;
366 unsigned groupSize = indices.size();
367 for (unsigned i = 0; i < groupSize; ++i) {
368 productOfValues =
369 productOfValues * builder.getAffineSymbolExpr(numberOfSymbols++);
370 unsigned srcIdx = indices[i];
371 int64_t maybeConstant = maybeConstants[srcIdx];
372
373 inputValues.push_back(isDynamic(maybeConstant)
374 ? values[srcIdx]
375 : builder.getIndexAttr(maybeConstant));
376 }
377
378 return makeComposedFoldedAffineApply(builder, loc, productOfValues,
379 inputValues);
380}
381
382/// Compute the collapsed size of the given \p collpaseShape for the
383/// \p groupId-th reassociation group.
384/// \p origSizes hold the sizes of the source shape as values.
385/// This is used to compute the new sizes in cases of dynamic shapes.
386///
387/// Conceptually this helper function computes:
388/// `prod(origSizes#i, for i in {ressociationGroup[groupId]})`.
389///
390/// \post result.size() == 1, in other words, each group collapse to one
391/// dimension.
392///
393/// TODO: Move this utility function directly within CollapseShapeOp. For now,
394/// this is not possible because this function uses the Affine dialect and the
395/// MemRef dialect cannot depend on the Affine dialect.
397getCollapsedSize(memref::CollapseShapeOp collapseShape, OpBuilder &builder,
398 ArrayRef<OpFoldResult> origSizes, unsigned groupId) {
399 SmallVector<OpFoldResult> collapsedSize;
400
401 MemRefType collapseShapeType = collapseShape.getResultType();
402
403 uint64_t size = collapseShapeType.getDimSize(groupId);
404 if (ShapedType::isStatic(size)) {
405 collapsedSize.push_back(builder.getIndexAttr(size));
406 return collapsedSize;
407 }
408
409 // We are dealing with a dynamic size.
410 // Build the affine expr of the product of the original sizes involved in that
411 // group.
412 Value source = collapseShape.getSrc();
413 auto sourceType = cast<MemRefType>(source.getType());
414
415 SmallVector<int64_t, 2> reassocGroup =
416 collapseShape.getReassociationIndices()[groupId];
417
418 collapsedSize.push_back(getProductOfValues(
419 reassocGroup, builder, collapseShape.getLoc(), sourceType.getShape(),
420 origSizes, ShapedType::isDynamic));
421
422 return collapsedSize;
423}
424
425/// Compute the collapsed stride of the given \p collpaseShape for the
426/// \p groupId-th reassociation group.
427/// \p origStrides and \p origSizes hold respectively the strides and sizes
428/// of the source shape as values.
429/// This is used to compute the strides in cases of dynamic shapes and/or
430/// dynamic stride for this reassociation group.
431///
432/// Conceptually this helper function returns the stride of the inner most
433/// dimension of that group in the original shape.
434///
435/// \post result.size() == 1, in other words, each group collapse to one
436/// dimension.
438getCollapsedStride(memref::CollapseShapeOp collapseShape, OpBuilder &builder,
439 ArrayRef<OpFoldResult> origSizes,
440 ArrayRef<OpFoldResult> origStrides, unsigned groupId) {
441 SmallVector<int64_t, 2> reassocGroup =
442 collapseShape.getReassociationIndices()[groupId];
443 assert(!reassocGroup.empty() &&
444 "Reassociation group should have at least one dimension");
445
446 Value source = collapseShape.getSrc();
447 auto sourceType = cast<MemRefType>(source.getType());
448
449 auto [strides, offset] = sourceType.getStridesAndOffset();
450
451 ArrayRef<int64_t> srcShape = sourceType.getShape();
452
453 OpFoldResult lastValidStride = nullptr;
454 for (int64_t currentDim : reassocGroup) {
455 // Skip size-of-1 dimensions, since right now their strides may be
456 // meaningless.
457 // FIXME: size-of-1 dimensions shouldn't be used in collapse shape, unless
458 // they are truly contiguous. When they are truly contiguous, we shouldn't
459 // need to skip them.
460 if (srcShape[currentDim] == 1)
461 continue;
462
463 int64_t currentStride = strides[currentDim];
464 lastValidStride = ShapedType::isDynamic(currentStride)
465 ? origStrides[currentDim]
466 : builder.getIndexAttr(currentStride);
467 }
468 if (!lastValidStride) {
469 // We're dealing with a 1x1x...x1 shape. The stride is meaningless,
470 // but we still have to make the type system happy.
471 MemRefType collapsedType = collapseShape.getResultType();
472 auto [collapsedStrides, collapsedOffset] =
473 collapsedType.getStridesAndOffset();
474 int64_t finalStride = collapsedStrides[groupId];
475 if (ShapedType::isDynamic(finalStride)) {
476 // Look for a dynamic stride. At this point we don't know which one is
477 // desired, but they are all equally good/bad.
478 for (int64_t currentDim : reassocGroup) {
479 assert(srcShape[currentDim] == 1 &&
480 "We should be dealing with 1x1x...x1");
481
482 if (ShapedType::isDynamic(strides[currentDim]))
483 return {origStrides[currentDim]};
484 }
485 llvm_unreachable("We should have found a dynamic stride");
486 }
487 return {builder.getIndexAttr(finalStride)};
488 }
489
490 return {lastValidStride};
491}
492
493/// From `reshape_like(memref, subSizes, subStrides))` compute
494///
495/// \verbatim
496/// baseBuffer, baseOffset, baseSizes, baseStrides =
497/// extract_strided_metadata(memref)
498/// strides#i = baseStrides#i * subStrides#i
499/// sizes = subSizes
500/// \endverbatim
501///
502/// and return {baseBuffer, baseOffset, sizes, strides}
503template <typename ReassociativeReshapeLikeOp>
504static FailureOr<StridedMetadata> resolveReshapeStridedMetadata(
505 RewriterBase &rewriter, ReassociativeReshapeLikeOp reshape,
507 ReassociativeReshapeLikeOp, OpBuilder &,
508 ArrayRef<OpFoldResult> /*origSizes*/, unsigned /*groupId*/)>
509 getReshapedSizes,
511 ReassociativeReshapeLikeOp, OpBuilder &,
512 ArrayRef<OpFoldResult> /*origSizes*/,
513 ArrayRef<OpFoldResult> /*origStrides*/, unsigned /*groupId*/)>
514 getReshapedStrides) {
515 // Build a plain extract_strided_metadata(memref) from
516 // extract_strided_metadata(reassociative_reshape_like(memref)).
517 Location origLoc = reshape.getLoc();
518 Value source = reshape.getSrc();
519 auto sourceType = cast<MemRefType>(source.getType());
520 unsigned sourceRank = sourceType.getRank();
521
522 auto newExtractStridedMetadata =
523 memref::ExtractStridedMetadataOp::create(rewriter, origLoc, source);
524
525 // Collect statically known information.
526 auto [strides, offset] = sourceType.getStridesAndOffset();
527 MemRefType reshapeType = reshape.getResultType();
528 unsigned reshapeRank = reshapeType.getRank();
529
530 OpFoldResult offsetOfr =
531 ShapedType::isDynamic(offset)
532 ? getAsOpFoldResult(newExtractStridedMetadata.getOffset())
533 : rewriter.getIndexAttr(offset);
534
535 // Get the special case of 0-D out of the way.
536 if (sourceRank == 0) {
537 SmallVector<OpFoldResult> ones(reshapeRank, rewriter.getIndexAttr(1));
538 return StridedMetadata{newExtractStridedMetadata.getBaseBuffer(), offsetOfr,
539 /*sizes=*/ones, /*strides=*/ones};
540 }
541
542 SmallVector<OpFoldResult> finalSizes;
543 finalSizes.reserve(reshapeRank);
544 SmallVector<OpFoldResult> finalStrides;
545 finalStrides.reserve(reshapeRank);
546
547 // Compute the reshaped strides and sizes from the base strides and sizes.
548 SmallVector<OpFoldResult> origSizes =
549 getAsOpFoldResult(newExtractStridedMetadata.getSizes());
550 SmallVector<OpFoldResult> origStrides =
551 getAsOpFoldResult(newExtractStridedMetadata.getStrides());
552 unsigned idx = 0, endIdx = reshape.getReassociationIndices().size();
553 for (; idx != endIdx; ++idx) {
554 SmallVector<OpFoldResult> reshapedSizes =
555 getReshapedSizes(reshape, rewriter, origSizes, /*groupId=*/idx);
556 SmallVector<OpFoldResult> reshapedStrides = getReshapedStrides(
557 reshape, rewriter, origSizes, origStrides, /*groupId=*/idx);
558
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]);
563 }
564 }
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");
570
571 return StridedMetadata{newExtractStridedMetadata.getBaseBuffer(), offsetOfr,
572 finalSizes, finalStrides};
573}
574
575/// Replace `baseBuffer, offset, sizes, strides =
576/// extract_strided_metadata(reshapeLike(memref))`
577/// With
578///
579/// \verbatim
580/// baseBuffer, offset, baseSizes, baseStrides =
581/// extract_strided_metadata(memref)
582/// sizes = getReshapedSizes(reshapeLike)
583/// strides = getReshapedStrides(reshapeLike)
584/// \endverbatim
585///
586///
587/// Notice that `baseBuffer` and `offset` are unchanged.
588///
589/// In other words, get rid of the expand_shape in that expression and
590/// materialize its effects on the sizes and the strides using affine apply.
591template <typename ReassociativeReshapeLikeOp,
592 SmallVector<OpFoldResult> (*getReshapedSizes)(
593 ReassociativeReshapeLikeOp, OpBuilder &,
594 ArrayRef<OpFoldResult> /*origSizes*/, unsigned /*groupId*/),
595 SmallVector<OpFoldResult> (*getReshapedStrides)(
596 ReassociativeReshapeLikeOp, OpBuilder &,
597 ArrayRef<OpFoldResult> /*origSizes*/,
598 ArrayRef<OpFoldResult> /*origStrides*/, unsigned /*groupId*/)>
599struct ReshapeFolder : public OpRewritePattern<ReassociativeReshapeLikeOp> {
600public:
601 using OpRewritePattern<ReassociativeReshapeLikeOp>::OpRewritePattern;
602
603 LogicalResult matchAndRewrite(ReassociativeReshapeLikeOp reshape,
604 PatternRewriter &rewriter) const override {
605 FailureOr<StridedMetadata> stridedMetadata =
606 resolveReshapeStridedMetadata<ReassociativeReshapeLikeOp>(
607 rewriter, reshape, getReshapedSizes, getReshapedStrides);
608 if (failed(stridedMetadata)) {
609 return rewriter.notifyMatchFailure(reshape,
610 "failed to resolve reshape metadata");
611 }
612
613 MemRefType resultType = reshape.getResultType();
614 if (isa<memref::CollapseShapeOp>(reshape.getOperation()))
615 resultType = updateTypeFromMetadata(resultType, stridedMetadata->offset,
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);
624 return success();
625 }
626 // Preserve the original result type expected by existing users.
627 rewriter.replaceOpWithNewOp<memref::CastOp>(
628 reshape, reshape.getResultType(), foldedReshape);
629 return success();
630 }
631};
632
633/// Pattern to replace `extract_strided_metadata(collapse_shape)`
634/// With
635///
636/// \verbatim
637/// baseBuffer, baseOffset, baseSizes, baseStrides =
638/// extract_strided_metadata(memref)
639/// strides#i = baseStrides#i * subSizes#i
640/// offset = baseOffset + sum(subOffset#i * baseStrides#i)
641/// sizes = subSizes
642/// \verbatim
643///
644/// with `baseBuffer`, `offset`, `sizes` and `strides` being
645/// the replacements for the original `extract_strided_metadata`.
646struct ExtractStridedMetadataOpCollapseShapeFolder
649
650 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
651 PatternRewriter &rewriter) const override {
652 auto collapseShapeOp =
653 op.getSource().getDefiningOp<memref::CollapseShapeOp>();
654 if (!collapseShapeOp)
655 return failure();
656
657 FailureOr<StridedMetadata> stridedMetadata =
658 resolveReshapeStridedMetadata<memref::CollapseShapeOp>(
659 rewriter, collapseShapeOp, getCollapsedSize, getCollapsedStride);
660 if (failed(stridedMetadata)) {
661 return rewriter.notifyMatchFailure(
662 op,
663 "failed to resolve metadata in terms of source collapse_shape op");
664 }
665
666 Location loc = collapseShapeOp.getLoc();
667 SmallVector<Value> results;
668 results.push_back(stridedMetadata->basePtr);
669 results.push_back(getValueOrCreateConstantIndexOp(rewriter, loc,
670 stridedMetadata->offset));
671 results.append(
672 getValueOrCreateConstantIndexOp(rewriter, loc, stridedMetadata->sizes));
673 results.append(getValueOrCreateConstantIndexOp(rewriter, loc,
674 stridedMetadata->strides));
675 rewriter.replaceOp(op, results);
676 return success();
677 }
678};
679
680/// Pattern to replace `extract_strided_metadata(expand_shape)`
681/// with the results of computing the sizes and strides on the expanded shape
682/// and dividing up dimensions into static and dynamic parts as needed.
683struct ExtractStridedMetadataOpExpandShapeFolder
686
687 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
688 PatternRewriter &rewriter) const override {
689 auto expandShapeOp = op.getSource().getDefiningOp<memref::ExpandShapeOp>();
690 if (!expandShapeOp)
691 return failure();
692
693 FailureOr<StridedMetadata> stridedMetadata =
694 resolveReshapeStridedMetadata<memref::ExpandShapeOp>(
695 rewriter, expandShapeOp, getExpandedSizes, getExpandedStrides);
696 if (failed(stridedMetadata)) {
697 return rewriter.notifyMatchFailure(
698 op, "failed to resolve metadata in terms of source expand_shape op");
699 }
700
701 Location loc = expandShapeOp.getLoc();
702 SmallVector<Value> results;
703 results.push_back(stridedMetadata->basePtr);
704 results.push_back(getValueOrCreateConstantIndexOp(rewriter, loc,
705 stridedMetadata->offset));
706 results.append(
707 getValueOrCreateConstantIndexOp(rewriter, loc, stridedMetadata->sizes));
708 results.append(getValueOrCreateConstantIndexOp(rewriter, loc,
709 stridedMetadata->strides));
710 rewriter.replaceOp(op, results);
711 return success();
712 }
713};
714
715/// Replace `base, offset, sizes, strides =
716/// extract_strided_metadata(allocLikeOp)`
717///
718/// With
719///
720/// ```
721/// base = reinterpret_cast allocLikeOp(allocSizes) to a flat memref<eltTy>
722/// offset = 0
723/// sizes = allocSizes
724/// strides#i = prod(allocSizes#j, for j in {i+1..rank-1})
725/// ```
726///
727/// The transformation only applies if the allocLikeOp has been normalized.
728/// In other words, the affine_map must be an identity.
729template <typename AllocLikeOp>
730struct ExtractStridedMetadataOpAllocFolder
732public:
733 using OpRewritePattern<memref::ExtractStridedMetadataOp>::OpRewritePattern;
734
735 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
736 PatternRewriter &rewriter) const override {
737 auto allocLikeOp = op.getSource().getDefiningOp<AllocLikeOp>();
738 if (!allocLikeOp)
739 return failure();
740
741 auto memRefType = cast<MemRefType>(allocLikeOp.getResult().getType());
742 if (!memRefType.getLayout().isIdentity())
743 return rewriter.notifyMatchFailure(
744 allocLikeOp, "alloc-like operations should have been normalized");
745
746 Location loc = op.getLoc();
747 int rank = memRefType.getRank();
748
749 // Collect the sizes.
750 ValueRange dynamic = allocLikeOp.getDynamicSizes();
752 sizes.reserve(rank);
753 unsigned dynamicPos = 0;
754 for (int64_t size : memRefType.getShape()) {
755 if (ShapedType::isDynamic(size))
756 sizes.push_back(dynamic[dynamicPos++]);
757 else
758 sizes.push_back(rewriter.getIndexAttr(size));
759 }
760
761 // Strides (just creates identity strides).
762 SmallVector<OpFoldResult> strides(rank, rewriter.getIndexAttr(1));
763 AffineExpr expr = rewriter.getAffineConstantExpr(1);
764 unsigned symbolNumber = 0;
765 for (int i = rank - 2; i >= 0; --i) {
766 expr = expr * rewriter.getAffineSymbolExpr(symbolNumber++);
767 assert(i + 1 + symbolNumber == sizes.size() &&
768 "The ArrayRef should encompass the last #symbolNumber sizes");
769 ArrayRef<OpFoldResult> sizesInvolvedInStride(&sizes[i + 1], symbolNumber);
770 strides[i] = makeComposedFoldedAffineApply(rewriter, loc, expr,
771 sizesInvolvedInStride);
772 }
773
774 // Put all the values together to replace the results.
775 SmallVector<Value> results;
776 results.reserve(rank * 2 + 2);
777
778 auto baseBufferType = cast<MemRefType>(op.getBaseBuffer().getType());
779 int64_t offset = 0;
780 if (op.getBaseBuffer().use_empty()) {
781 results.push_back(nullptr);
782 } else {
783 if (allocLikeOp.getType() == baseBufferType)
784 results.push_back(allocLikeOp);
785 else
786 results.push_back(memref::ReinterpretCastOp::create(
787 rewriter, loc, baseBufferType, allocLikeOp, offset,
788 /*sizes=*/ArrayRef<int64_t>(),
789 /*strides=*/ArrayRef<int64_t>()));
790 }
791
792 // Offset.
793 results.push_back(arith::ConstantIndexOp::create(rewriter, loc, offset));
794
795 for (OpFoldResult size : sizes)
796 results.push_back(getValueOrCreateConstantIndexOp(rewriter, loc, size));
797
798 for (OpFoldResult stride : strides)
799 results.push_back(getValueOrCreateConstantIndexOp(rewriter, loc, stride));
800
801 rewriter.replaceOp(op, results);
802 return success();
803 }
804};
805
806/// Replace `base, offset, sizes, strides =
807/// extract_strided_metadata(get_global)`
808///
809/// With
810///
811/// ```
812/// base = reinterpret_cast get_global to a flat memref<eltTy>
813/// offset = 0
814/// sizes = allocSizes
815/// strides#i = prod(allocSizes#j, for j in {i+1..rank-1})
816/// ```
817///
818/// It is expected that the memref.get_global op has static shapes
819/// and identity affine_map for the layout.
820struct ExtractStridedMetadataOpGetGlobalFolder
821 : public OpRewritePattern<memref::ExtractStridedMetadataOp> {
822public:
823 using OpRewritePattern<memref::ExtractStridedMetadataOp>::OpRewritePattern;
824
825 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
826 PatternRewriter &rewriter) const override {
827 auto getGlobalOp = op.getSource().getDefiningOp<memref::GetGlobalOp>();
828 if (!getGlobalOp)
829 return failure();
830
831 auto memRefType = cast<MemRefType>(getGlobalOp.getResult().getType());
832 if (!memRefType.getLayout().isIdentity()) {
833 return rewriter.notifyMatchFailure(
834 getGlobalOp,
835 "get-global operation result should have been normalized");
836 }
837
838 Location loc = op.getLoc();
839 int rank = memRefType.getRank();
840
841 // Collect the sizes.
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");
845
846 // Strides (just creates identity strides).
847 SmallVector<int64_t> strides = computeSuffixProduct(sizes);
848
849 // Put all the values together to replace the results.
850 SmallVector<Value> results;
851 results.reserve(rank * 2 + 2);
852
853 auto baseBufferType = cast<MemRefType>(op.getBaseBuffer().getType());
854 int64_t offset = 0;
855 if (getGlobalOp.getType() == baseBufferType)
856 results.push_back(getGlobalOp);
857 else
858 results.push_back(memref::ReinterpretCastOp::create(
859 rewriter, loc, baseBufferType, getGlobalOp, offset,
860 /*sizes=*/ArrayRef<int64_t>(),
861 /*strides=*/ArrayRef<int64_t>()));
862
863 // Offset.
864 results.push_back(arith::ConstantIndexOp::create(rewriter, loc, offset));
865
866 for (auto size : sizes)
867 results.push_back(arith::ConstantIndexOp::create(rewriter, loc, size));
868
869 for (auto stride : strides)
870 results.push_back(arith::ConstantIndexOp::create(rewriter, loc, stride));
871
872 rewriter.replaceOp(op, results);
873 return success();
874 }
875};
876
877/// Pattern to replace `extract_strided_metadata(assume_alignment)`
878///
879/// With
880/// \verbatim
881/// extract_strided_metadata(memref)
882/// \endverbatim
883///
884/// Since `assume_alignment` is a view-like op that does not modify the
885/// underlying buffer, offset, sizes, or strides, extracting strided metadata
886/// from its result is equivalent to extracting it from its source. This
887/// canonicalization removes the unnecessary indirection.
888struct ExtractStridedMetadataOpAssumeAlignmentFolder
889 : public OpRewritePattern<memref::ExtractStridedMetadataOp> {
890public:
891 using OpRewritePattern<memref::ExtractStridedMetadataOp>::OpRewritePattern;
892
893 LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp op,
894 PatternRewriter &rewriter) const override {
895 auto assumeAlignmentOp =
896 op.getSource().getDefiningOp<memref::AssumeAlignmentOp>();
897 if (!assumeAlignmentOp)
898 return failure();
899
900 rewriter.replaceOpWithNewOp<memref::ExtractStridedMetadataOp>(
901 op, assumeAlignmentOp.getViewSource());
902 return success();
903 }
904};
905
906/// Rewrite memref.extract_aligned_pointer_as_index of a ViewLikeOp to the
907/// source of the ViewLikeOp.
908class RewriteExtractAlignedPointerAsIndexOfViewLikeOp
909 : public OpRewritePattern<memref::ExtractAlignedPointerAsIndexOp> {
911
912 LogicalResult
913 matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
914 PatternRewriter &rewriter) const override {
915 auto viewLikeOp =
916 extractOp.getSource().getDefiningOp<ViewLikeOpInterface>();
917 // ViewLikeOpInterface by itself doesn't guarantee to preserve the base
918 // pointer in general and `memref.view` is one such example, so just check
919 // for a few specific cases.
920 if (!viewLikeOp || extractOp.getSource() != viewLikeOp.getViewDest() ||
921 !isa<memref::SubViewOp, memref::ReinterpretCastOp>(viewLikeOp))
922 return rewriter.notifyMatchFailure(extractOp, "not a ViewLike source");
923 rewriter.modifyOpInPlace(extractOp, [&]() {
924 extractOp.getSourceMutable().assign(viewLikeOp.getViewSource());
925 });
926 return success();
927 }
928};
929
930/// Replace `base, offset, sizes, strides =
931/// extract_strided_metadata(
932/// reinterpret_cast(src, srcOffset, srcSizes, srcStrides))`
933/// With
934/// ```
935/// base, ... = extract_strided_metadata(src)
936/// offset = srcOffset
937/// sizes = srcSizes
938/// strides = srcStrides
939/// ```
940///
941/// In other words, consume the `reinterpret_cast` and apply its effects
942/// on the offset, sizes, and strides.
943class ExtractStridedMetadataOpReinterpretCastFolder
944 : public OpRewritePattern<memref::ExtractStridedMetadataOp> {
946
947 LogicalResult
948 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
949 PatternRewriter &rewriter) const override {
950 auto reinterpretCastOp = extractStridedMetadataOp.getSource()
951 .getDefiningOp<memref::ReinterpretCastOp>();
952 if (!reinterpretCastOp)
953 return failure();
954
955 Location loc = extractStridedMetadataOp.getLoc();
956 // Check if the source is suitable for extract_strided_metadata.
957 SmallVector<Type> inferredReturnTypes;
958 if (failed(extractStridedMetadataOp.inferReturnTypes(
959 rewriter.getContext(), loc, {reinterpretCastOp.getSource()},
960 /*attributes=*/{}, /*properties=*/{}, /*regions=*/{},
961 inferredReturnTypes)))
962 return rewriter.notifyMatchFailure(
963 reinterpretCastOp, "reinterpret_cast source's type is incompatible");
964
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);
969
970 auto newExtractStridedMetadata = memref::ExtractStridedMetadataOp::create(
971 rewriter, loc, reinterpretCastOp.getSource());
972
973 // Register the base_buffer.
974 results[0] = newExtractStridedMetadata.getBaseBuffer();
975
976 // Register the new offset.
978 rewriter, loc, reinterpretCastOp.getMixedOffsets()[0]);
979
980 const unsigned sizeStartIdx = 2;
981 const unsigned strideStartIdx = sizeStartIdx + rank;
982
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];
988 }
989 rewriter.replaceOp(extractStridedMetadataOp,
990 getValueOrCreateConstantIndexOp(rewriter, loc, results));
991 return success();
992 }
993};
994
995/// Replace `base, offset, sizes, strides = extract_strided_metadata(
996/// memory_space_cast(src) to dstTy)`
997/// with
998/// ```
999/// oldBase, offset, sizes, strides = extract_strided_metadata(src)
1000/// destBaseTy = type(oldBase) with memory space from destTy
1001/// base = memory_space_cast(oldBase) to destBaseTy
1002/// ```
1003///
1004/// In other words, propagate metadata extraction accross memory space casts.
1005class ExtractStridedMetadataOpMemorySpaceCastFolder
1006 : public OpRewritePattern<memref::ExtractStridedMetadataOp> {
1008
1009 LogicalResult
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)
1016 return failure();
1017 auto newExtractStridedMetadata = memref::ExtractStridedMetadataOp::create(
1018 rewriter, loc, memSpaceCastOp.getSource());
1019 SmallVector<Value> results(newExtractStridedMetadata.getResults());
1020 // As with most other strided metadata rewrite patterns, don't introduce
1021 // a use of the base pointer where non existed. This needs to happen here,
1022 // as opposed to in later dead-code elimination, because these patterns are
1023 // sometimes used during dialect conversion (see EmulateNarrowType, for
1024 // example), so adding spurious usages would cause a pre-legalization value
1025 // to be live that would be dead had this pattern not run.
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);
1034 } else {
1035 results[0] = nullptr;
1036 }
1037 rewriter.replaceOp(extractStridedMetadataOp, results);
1038 return success();
1039 }
1040};
1041
1042/// Replace `base, offset =
1043/// extract_strided_metadata(extract_strided_metadata(src)#0)`
1044/// With
1045/// ```
1046/// base, ... = extract_strided_metadata(src)
1047/// offset = 0
1048/// ```
1049class ExtractStridedMetadataOpExtractStridedMetadataFolder
1050 : public OpRewritePattern<memref::ExtractStridedMetadataOp> {
1052
1053 LogicalResult
1054 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
1055 PatternRewriter &rewriter) const override {
1056 auto sourceExtractStridedMetadataOp =
1057 extractStridedMetadataOp.getSource()
1058 .getDefiningOp<memref::ExtractStridedMetadataOp>();
1059 if (!sourceExtractStridedMetadataOp)
1060 return failure();
1061 Location loc = extractStridedMetadataOp.getLoc();
1062 rewriter.replaceOp(extractStridedMetadataOp,
1063 {sourceExtractStridedMetadataOp.getBaseBuffer(),
1065 rewriter, loc, rewriter.getIndexAttr(0))});
1066 return success();
1067 }
1068};
1069} // namespace
1070
1072 RewritePatternSet &patterns) {
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>(
1089 patterns.getContext());
1090}
1091
1093 RewritePatternSet &patterns) {
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>(
1105 patterns.getContext());
1106}
1107
1108//===----------------------------------------------------------------------===//
1109// Pass registration
1110//===----------------------------------------------------------------------===//
1111
1112namespace {
1113
1114struct ExpandStridedMetadataPass final
1116 ExpandStridedMetadataPass> {
1117 void runOnOperation() override;
1118};
1119
1120} // namespace
1121
1122void ExpandStridedMetadataPass::runOnOperation() {
1123 RewritePatternSet patterns(&getContext());
1125 (void)applyPatternsGreedily(getOperation(), std::move(patterns));
1126}
return success()
b getContext())
#define mul(a, b)
Base type for affine expression.
Definition AffineExpr.h:68
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
AffineExpr getAffineSymbolExpr(unsigned position)
Definition Builders.cpp:377
AffineExpr getAffineConstantExpr(int64_t constant)
Definition Builders.cpp:381
MLIRContext * getContext() const
Definition Builders.h:56
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class helps build Operations.
Definition Builders.h:210
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.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
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,...
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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 &region, 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 .
Definition AffineExpr.h:325
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.
Definition Utils.cpp:114
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
void bindSymbolsList(MLIRContext *ctx, MutableArrayRef< AffineExprTy > exprs)
Definition AffineExpr.h:330
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.