19#include "llvm/ADT/STLExtras.h"
25 if (!type.hasStaticShape())
34 if (failed(type.getStridesAndOffset(strides, offset)))
42 while (curDim >= 0 && strides[curDim] == runningStride) {
43 runningStride *= type.getDimSize(curDim);
48 while (curDim >= 0 && type.getDimSize(curDim) == 1) {
61 unsigned sourceRank = sizes.size();
62 assert(sizes.size() == strides.size() &&
63 "expected as many sizes as strides for a memref");
67 assert(indicesVec.size() == strides.size() &&
68 "expected as many indices as rank of memref");
77 for (
unsigned i = 0; i < sourceRank; ++i) {
78 unsigned offsetIdx = 2 * i;
79 addMulMap = addMulMap + symbols[offsetIdx] * symbols[offsetIdx + 1];
80 offsetValues[offsetIdx] = indicesVec[i];
81 offsetValues[offsetIdx + 1] = strides[i];
84 int64_t scaler = dstBits / srcBits;
86 builder, loc, addMulMap.
floorDiv(scaler), offsetValues);
88 size_t symbolIndex = 0;
91 for (
unsigned i = 0; i < sourceRank; ++i) {
92 AffineExpr strideExpr = symbols[symbolIndex++];
93 values.push_back(strides[i]);
95 values.push_back(sizes[i]);
103 0, symbolIndex, productExpressions,
112 builder, loc, s0.
floorDiv(scaler), {offset});
115 builder, loc, addMulMap % scaler, offsetValues);
117 return {{adjustBaseOffset, linearizedSize, intraVectorOffset},
127 if (!sizes.empty()) {
134 builder, loc, s0 * s1,
140 std::tie(linearizedMemRefInfo, std::ignore) =
144 return linearizedMemRefInfo;
152 std::vector<Operation *> opUses;
158 if (isa<memref::DeallocOp>(useOp) ||
162 opUses.push_back(useOp);
167 llvm::append_range(uses, opUses);
172 std::vector<Operation *> opToErase;
174 std::vector<Operation *> candidates;
175 if (isa<memref::AllocOp, memref::AllocaOp>(op) &&
177 llvm::append_range(opToErase, candidates);
178 opToErase.push_back(op);
194 for (
int64_t r =
static_cast<int64_t>(strides.size()) - 1; r > 0; --r) {
196 builder, loc, s0 * s1, {strides[r], sizes[r]});
209 while (
auto *op = source.getDefiningOp()) {
210 if (
auto subViewOp = dyn_cast<memref::SubViewOp>(op);
211 subViewOp && subViewOp.hasZeroOffset() && subViewOp.hasUnitStride()) {
214 source = cast<MemrefValue>(subViewOp.getSource());
215 }
else if (
auto castOp = dyn_cast<memref::CastOp>(op)) {
217 source = castOp.getSource();
226 while (
auto *op = source.getDefiningOp()) {
227 if (
auto viewLike = dyn_cast<ViewLikeOpInterface>(op)) {
228 if (source == viewLike.getViewDest()) {
229 source = cast<MemrefValue>(viewLike.getViewSource());
239 memref::ExpandShapeOp expandShapeOp,
242 bool startsInbounds) {
248 assert(!group.empty() &&
"association indices groups cannot be empty");
249 int64_t groupSize = group.size();
250 if (groupSize == 1) {
251 sourceIndices.push_back(
indices[group[0]]);
255 llvm::map_to_vector(group, [&](
int64_t d) {
return destShape[d]; });
257 llvm::map_to_vector(group, [&](
int64_t d) {
return indices[d]; });
258 Value collapsedIndex = affine::AffineLinearizeIndexOp::create(
259 rewriter, loc, groupIndices, groupBasis, startsInbounds);
260 sourceIndices.push_back(collapsedIndex);
265 memref::CollapseShapeOp collapseShapeOp,
268 bool startsInbounds) {
270 auto metadata = memref::ExtractStridedMetadataOp::create(
271 rewriter, loc, collapseShapeOp.getSrc());
273 for (
auto [
index, group] :
274 llvm::zip(
indices, collapseShapeOp.getReassociationIndices())) {
275 assert(!group.empty() &&
"association indices groups cannot be empty");
276 int64_t groupSize = group.size();
278 if (groupSize == 1) {
279 sourceIndices.push_back(
index);
289 trimmedGroup, [&](
int64_t d) {
return sourceSizes[d]; });
290 auto delinearize = affine::AffineDelinearizeIndexOp::create(
291 rewriter, loc,
index, basis, startsInbounds);
292 llvm::append_range(sourceIndices,
delinearize.getResults());
294 if (collapseShapeOp.getReassociationIndices().empty()) {
297 cast<MemRefType>(collapseShapeOp.getViewSource().getType()).getRank();
300 for (
int64_t i = 0; i < srcRank; i++) {
301 sourceIndices.push_back(
310 if (!subViewOp.hasZeroOffset() || !subViewOp.hasUnitStride())
313 MemRefType srcType = subViewOp.getSourceType();
314 MemRefType resType = subViewOp.getType();
315 unsigned srcRank = srcType.getRank();
316 unsigned resRank = resType.getRank();
317 if (srcRank <= resRank ||
indices.size() != resRank)
320 auto droppedDims = subViewOp.getDroppedDims();
321 if (droppedDims.none() || droppedDims.count() != srcRank - resRank)
324 auto mixedSizes = subViewOp.getMixedSizes();
325 if (mixedSizes.size() != srcRank)
328 unsigned resultDim = 0;
329 for (
unsigned sourceDim = 0; sourceDim < srcRank; ++sourceDim) {
330 if (droppedDims.test(sourceDim)) {
332 if (!sizeCst || *sizeCst != 1)
334 sourceIndices.push_back(
338 if (resultDim >=
indices.size())
340 sourceIndices.push_back(
indices[resultDim++]);
342 if (resultDim !=
indices.size())
349 auto [strides, offset] = memRefTy.getStridesAndOffset();
350 return llvm::any_of(strides, [](
int64_t stride) {
351 return ShapedType::isStatic(stride) && stride < 0;
368struct SubviewFootprint {
383 auto baseTy = dyn_cast<MemRefType>(base.
getType());
386 SubviewFootprint footprint;
389 footprint.offsets.assign(baseTy.getRank(), 0);
390 footprint.sizes.assign(baseTy.getShape().begin(), baseTy.getShape().end());
399 if (sv.getSourceType().getRank() != sv.getType().getRank() ||
400 llvm::any_of(staticOffsets, ShapedType::isDynamic) ||
401 llvm::any_of(staticStrides, [](
int64_t s) { return s != 1; }))
403 for (
unsigned d = 0, e = staticOffsets.size(); d < e; ++d)
404 footprint.offsets[d] += staticOffsets[d];
405 cur = sv.getSource();
411 if (isa<ViewLikeOpInterface>(def) && !isa<SubViewOp>(def))
413 footprint.root = cur;
420 const SubviewFootprint &
b) {
421 if (a.root !=
b.root)
423 for (
unsigned d = 0, e = a.offsets.size(); d < e; ++d) {
425 if (ShapedType::isDynamic(a.sizes[d]) || ShapedType::isDynamic(
b.sizes[d]))
427 int64_t aHi = a.offsets[d] + a.sizes[d];
429 if (aHi <=
b.offsets[d] || bHi <= a.offsets[d])
438 auto baseMemref = dyn_cast<MemrefValue>(base);
449 while (!worklist.empty()) {
450 Operation *user = worklist.pop_back_val();
451 if (!processed.insert(user).second)
454 if (
auto viewLike = dyn_cast<ViewLikeOpInterface>(user)) {
455 Value viewDest = viewLike.getViewDest();
467 auto slice = dyn_cast<MemrefValue>(operand);
471 std::optional<SubviewFootprint> sliceFp =
476 if (readsAreSafe && isa<MemoryEffectOpInterface>(user) &&
static int64_t product(ArrayRef< int64_t > vals)
Base type for affine expression.
AffineExpr floorDiv(uint64_t v) const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
IntegerAttr getIndexAttr(int64_t value)
AffineExpr getAffineConstantExpr(int64_t constant)
AffineMap getConstantAffineMap(int64_t val)
Returns a single constant result affine map with 0 dimensions and 0 symbols.
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.
This class represents an operand of an operation.
This class provides the API for ops that are known to be terminators.
Operation is the basic unit of execution within MLIR.
bool mightHaveTrait()
Returns true if the operation might have the provided trait.
unsigned getNumRegions()
Returns the number of regions held by this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
bool isAncestor(Operation *other)
Return true if this operation is an ancestor of the other operation.
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
use_range getUses()
Returns a range of all uses, which is useful for iterating over all uses.
unsigned getNumResults()
Return the number of results held by this operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
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.
user_range getUsers() const
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
OpFoldResult makeComposedFoldedAffineMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Constructs an AffineMinOp that computes a maximum across the results of applying map to operands,...
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...
std::pair< LinearizedMemRefInfo, OpFoldResult > getLinearizedMemRefOffsetAndSize(OpBuilder &builder, Location loc, int srcBits, int dstBits, OpFoldResult offset, ArrayRef< OpFoldResult > sizes, ArrayRef< OpFoldResult > strides, ArrayRef< OpFoldResult > indices={}, LinearizedDivKind sizeDivKind=LinearizedDivKind::Floor)
static bool resultIsNotRead(Operation *op, std::vector< Operation * > &uses)
Returns true if all the uses of op are not read/load.
bool hasNegativeStaticStride(MemRefType memRefTy)
Returns true if any stride of memRefTy is statically known to be negative.
MemrefValue skipFullyAliasingOperations(MemrefValue source)
Walk up the source chain until an operation that changes/defines the view of memory is found (i....
static std::optional< SubviewFootprint > resolveSubviewFootprint(Value base)
Resolve base through a memref.subview chain to a footprint, or std::nullopt if it cannot be modelled.
void eraseDeadAllocAndStores(RewriterBase &rewriter, Operation *parentOp)
Track temporary allocations that are never read from.
MemrefValue skipViewLikeOps(MemrefValue source)
Walk up the source chain until we find an operation that is not a view of the source memref (i....
bool isStaticShapeAndContiguousRowMajor(MemRefType type)
Returns true, if the memref type has static shapes and represents a contiguous chunk of memory.
bool hasNoAliasingAccessInScope(Value base, Operation *scope, ArrayRef< Operation * > excludedOps={}, bool readsAreSafe=false)
Return "true" when no other access nested in scope can alias base.
void resolveSourceIndicesCollapseShape(Location loc, PatternRewriter &rewriter, memref::CollapseShapeOp collapseShapeOp, ValueRange indices, SmallVectorImpl< Value > &sourceIndices, bool startsInbounds)
Given the 'indices' of a load/store operation where the memref is a result of a collapse_shape op,...
LinearizedDivKind
Controls how the per-dimension contribution to linearizedSize is divided by dstBits / srcBits when sc...
LogicalResult resolveSourceIndicesRankReducingSubview(Location loc, OpBuilder &b, memref::SubViewOp subViewOp, ValueRange indices, SmallVectorImpl< Value > &sourceIndices)
Given the 'indices' of a load/store operation where the memref is a result of a rank-reducing full su...
static SmallVector< OpFoldResult > computeSuffixProductIRBlockImpl(Location loc, OpBuilder &builder, ArrayRef< OpFoldResult > sizes, OpFoldResult unit)
SmallVector< OpFoldResult > computeSuffixProductIRBlock(Location loc, OpBuilder &builder, ArrayRef< OpFoldResult > sizes)
Given a set of sizes, return the suffix product.
void resolveSourceIndicesExpandShape(Location loc, PatternRewriter &rewriter, memref::ExpandShapeOp expandShapeOp, ValueRange indices, SmallVectorImpl< Value > &sourceIndices, bool startsInbounds)
Given the 'indices' of a load/store operation where the memref is a result of a expand_shape op,...
static bool areFootprintDisjoint(const SubviewFootprint &a, const SubviewFootprint &b)
True if two footprints provably cover disjoint memory: they must share the same root and be separated...
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
SmallVector< int64_t > delinearize(int64_t linearIndex, ArrayRef< int64_t > strides)
Given the strides together with a linear index in the dimension space, return the vector-space offset...
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
TypedValue< BaseMemRefType > MemrefValue
A value with a memref type.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
bool hasEffect(Operation *op)
Returns "true" if op has an effect of type EffectTy.
void bindSymbolsList(MLIRContext *ctx, MutableArrayRef< AffineExprTy > exprs)
For a memref with offset, sizes and strides, returns the offset, size, and potentially the size padde...