1 //===- Utils.cpp ---- Utilities for affine dialect transformation ---------===//
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 // This file implements miscellaneous transformation utilities for the Affine
10 // dialect.
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "mlir/Dialect/Affine/Utils.h"
15 
16 #include "mlir/Dialect/Affine/Analysis/Utils.h"
17 #include "mlir/Dialect/Affine/IR/AffineOps.h"
18 #include "mlir/Dialect/Affine/IR/AffineValueMap.h"
19 #include "mlir/Dialect/Affine/LoopUtils.h"
20 #include "mlir/Dialect/MemRef/IR/MemRef.h"
21 #include "mlir/IR/AffineExprVisitor.h"
22 #include "mlir/IR/BlockAndValueMapping.h"
23 #include "mlir/IR/Dominance.h"
24 #include "mlir/IR/IntegerSet.h"
25 #include "mlir/Transforms/GreedyPatternRewriteDriver.h"
26 
27 #define DEBUG_TYPE "affine-utils"
28 
29 using namespace mlir;
30 
31 namespace {
32 /// Visit affine expressions recursively and build the sequence of operations
33 /// that correspond to it.  Visitation functions return an Value of the
34 /// expression subtree they visited or `nullptr` on error.
35 class AffineApplyExpander
36     : public AffineExprVisitor<AffineApplyExpander, Value> {
37 public:
38   /// This internal class expects arguments to be non-null, checks must be
39   /// performed at the call site.
40   AffineApplyExpander(OpBuilder &builder, ValueRange dimValues,
41                       ValueRange symbolValues, Location loc)
42       : builder(builder), dimValues(dimValues), symbolValues(symbolValues),
43         loc(loc) {}
44 
45   template <typename OpTy>
46   Value buildBinaryExpr(AffineBinaryOpExpr expr) {
47     auto lhs = visit(expr.getLHS());
48     auto rhs = visit(expr.getRHS());
49     if (!lhs || !rhs)
50       return nullptr;
51     auto op = builder.create<OpTy>(loc, lhs, rhs);
52     return op.getResult();
53   }
54 
55   Value visitAddExpr(AffineBinaryOpExpr expr) {
56     return buildBinaryExpr<arith::AddIOp>(expr);
57   }
58 
59   Value visitMulExpr(AffineBinaryOpExpr expr) {
60     return buildBinaryExpr<arith::MulIOp>(expr);
61   }
62 
63   /// Euclidean modulo operation: negative RHS is not allowed.
64   /// Remainder of the euclidean integer division is always non-negative.
65   ///
66   /// Implemented as
67   ///
68   ///     a mod b =
69   ///         let remainder = srem a, b;
70   ///             negative = a < 0 in
71   ///         select negative, remainder + b, remainder.
72   Value visitModExpr(AffineBinaryOpExpr expr) {
73     auto rhsConst = expr.getRHS().dyn_cast<AffineConstantExpr>();
74     if (!rhsConst) {
75       emitError(
76           loc,
77           "semi-affine expressions (modulo by non-const) are not supported");
78       return nullptr;
79     }
80     if (rhsConst.getValue() <= 0) {
81       emitError(loc, "modulo by non-positive value is not supported");
82       return nullptr;
83     }
84 
85     auto lhs = visit(expr.getLHS());
86     auto rhs = visit(expr.getRHS());
87     assert(lhs && rhs && "unexpected affine expr lowering failure");
88 
89     Value remainder = builder.create<arith::RemSIOp>(loc, lhs, rhs);
90     Value zeroCst = builder.create<arith::ConstantIndexOp>(loc, 0);
91     Value isRemainderNegative = builder.create<arith::CmpIOp>(
92         loc, arith::CmpIPredicate::slt, remainder, zeroCst);
93     Value correctedRemainder =
94         builder.create<arith::AddIOp>(loc, remainder, rhs);
95     Value result = builder.create<arith::SelectOp>(
96         loc, isRemainderNegative, correctedRemainder, remainder);
97     return result;
98   }
99 
100   /// Floor division operation (rounds towards negative infinity).
101   ///
102   /// For positive divisors, it can be implemented without branching and with a
103   /// single division operation as
104   ///
105   ///        a floordiv b =
106   ///            let negative = a < 0 in
107   ///            let absolute = negative ? -a - 1 : a in
108   ///            let quotient = absolute / b in
109   ///                negative ? -quotient - 1 : quotient
110   Value visitFloorDivExpr(AffineBinaryOpExpr expr) {
111     auto rhsConst = expr.getRHS().dyn_cast<AffineConstantExpr>();
112     if (!rhsConst) {
113       emitError(
114           loc,
115           "semi-affine expressions (division by non-const) are not supported");
116       return nullptr;
117     }
118     if (rhsConst.getValue() <= 0) {
119       emitError(loc, "division by non-positive value is not supported");
120       return nullptr;
121     }
122 
123     auto lhs = visit(expr.getLHS());
124     auto rhs = visit(expr.getRHS());
125     assert(lhs && rhs && "unexpected affine expr lowering failure");
126 
127     Value zeroCst = builder.create<arith::ConstantIndexOp>(loc, 0);
128     Value noneCst = builder.create<arith::ConstantIndexOp>(loc, -1);
129     Value negative = builder.create<arith::CmpIOp>(
130         loc, arith::CmpIPredicate::slt, lhs, zeroCst);
131     Value negatedDecremented = builder.create<arith::SubIOp>(loc, noneCst, lhs);
132     Value dividend =
133         builder.create<arith::SelectOp>(loc, negative, negatedDecremented, lhs);
134     Value quotient = builder.create<arith::DivSIOp>(loc, dividend, rhs);
135     Value correctedQuotient =
136         builder.create<arith::SubIOp>(loc, noneCst, quotient);
137     Value result = builder.create<arith::SelectOp>(loc, negative,
138                                                    correctedQuotient, quotient);
139     return result;
140   }
141 
142   /// Ceiling division operation (rounds towards positive infinity).
143   ///
144   /// For positive divisors, it can be implemented without branching and with a
145   /// single division operation as
146   ///
147   ///     a ceildiv b =
148   ///         let negative = a <= 0 in
149   ///         let absolute = negative ? -a : a - 1 in
150   ///         let quotient = absolute / b in
151   ///             negative ? -quotient : quotient + 1
152   Value visitCeilDivExpr(AffineBinaryOpExpr expr) {
153     auto rhsConst = expr.getRHS().dyn_cast<AffineConstantExpr>();
154     if (!rhsConst) {
155       emitError(loc) << "semi-affine expressions (division by non-const) are "
156                         "not supported";
157       return nullptr;
158     }
159     if (rhsConst.getValue() <= 0) {
160       emitError(loc, "division by non-positive value is not supported");
161       return nullptr;
162     }
163     auto lhs = visit(expr.getLHS());
164     auto rhs = visit(expr.getRHS());
165     assert(lhs && rhs && "unexpected affine expr lowering failure");
166 
167     Value zeroCst = builder.create<arith::ConstantIndexOp>(loc, 0);
168     Value oneCst = builder.create<arith::ConstantIndexOp>(loc, 1);
169     Value nonPositive = builder.create<arith::CmpIOp>(
170         loc, arith::CmpIPredicate::sle, lhs, zeroCst);
171     Value negated = builder.create<arith::SubIOp>(loc, zeroCst, lhs);
172     Value decremented = builder.create<arith::SubIOp>(loc, lhs, oneCst);
173     Value dividend =
174         builder.create<arith::SelectOp>(loc, nonPositive, negated, decremented);
175     Value quotient = builder.create<arith::DivSIOp>(loc, dividend, rhs);
176     Value negatedQuotient =
177         builder.create<arith::SubIOp>(loc, zeroCst, quotient);
178     Value incrementedQuotient =
179         builder.create<arith::AddIOp>(loc, quotient, oneCst);
180     Value result = builder.create<arith::SelectOp>(
181         loc, nonPositive, negatedQuotient, incrementedQuotient);
182     return result;
183   }
184 
185   Value visitConstantExpr(AffineConstantExpr expr) {
186     auto op = builder.create<arith::ConstantIndexOp>(loc, expr.getValue());
187     return op.getResult();
188   }
189 
190   Value visitDimExpr(AffineDimExpr expr) {
191     assert(expr.getPosition() < dimValues.size() &&
192            "affine dim position out of range");
193     return dimValues[expr.getPosition()];
194   }
195 
196   Value visitSymbolExpr(AffineSymbolExpr expr) {
197     assert(expr.getPosition() < symbolValues.size() &&
198            "symbol dim position out of range");
199     return symbolValues[expr.getPosition()];
200   }
201 
202 private:
203   OpBuilder &builder;
204   ValueRange dimValues;
205   ValueRange symbolValues;
206 
207   Location loc;
208 };
209 } // namespace
210 
211 /// Create a sequence of operations that implement the `expr` applied to the
212 /// given dimension and symbol values.
213 mlir::Value mlir::expandAffineExpr(OpBuilder &builder, Location loc,
214                                    AffineExpr expr, ValueRange dimValues,
215                                    ValueRange symbolValues) {
216   return AffineApplyExpander(builder, dimValues, symbolValues, loc).visit(expr);
217 }
218 
219 /// Create a sequence of operations that implement the `affineMap` applied to
220 /// the given `operands` (as it it were an AffineApplyOp).
221 Optional<SmallVector<Value, 8>> mlir::expandAffineMap(OpBuilder &builder,
222                                                       Location loc,
223                                                       AffineMap affineMap,
224                                                       ValueRange operands) {
225   auto numDims = affineMap.getNumDims();
226   auto expanded = llvm::to_vector<8>(
227       llvm::map_range(affineMap.getResults(),
228                       [numDims, &builder, loc, operands](AffineExpr expr) {
229                         return expandAffineExpr(builder, loc, expr,
230                                                 operands.take_front(numDims),
231                                                 operands.drop_front(numDims));
232                       }));
233   if (llvm::all_of(expanded, [](Value v) { return v; }))
234     return expanded;
235   return None;
236 }
237 
238 /// Promotes the `then` or the `else` block of `ifOp` (depending on whether
239 /// `elseBlock` is false or true) into `ifOp`'s containing block, and discards
240 /// the rest of the op.
241 static void promoteIfBlock(AffineIfOp ifOp, bool elseBlock) {
242   if (elseBlock)
243     assert(ifOp.hasElse() && "else block expected");
244 
245   Block *destBlock = ifOp->getBlock();
246   Block *srcBlock = elseBlock ? ifOp.getElseBlock() : ifOp.getThenBlock();
247   destBlock->getOperations().splice(
248       Block::iterator(ifOp), srcBlock->getOperations(), srcBlock->begin(),
249       std::prev(srcBlock->end()));
250   ifOp.erase();
251 }
252 
253 /// Returns the outermost affine.for/parallel op that the `ifOp` is invariant
254 /// on. The `ifOp` could be hoisted and placed right before such an operation.
255 /// This method assumes that the ifOp has been canonicalized (to be correct and
256 /// effective).
257 static Operation *getOutermostInvariantForOp(AffineIfOp ifOp) {
258   // Walk up the parents past all for op that this conditional is invariant on.
259   auto ifOperands = ifOp.getOperands();
260   auto *res = ifOp.getOperation();
261   while (!isa<FuncOp>(res->getParentOp())) {
262     auto *parentOp = res->getParentOp();
263     if (auto forOp = dyn_cast<AffineForOp>(parentOp)) {
264       if (llvm::is_contained(ifOperands, forOp.getInductionVar()))
265         break;
266     } else if (auto parallelOp = dyn_cast<AffineParallelOp>(parentOp)) {
267       for (auto iv : parallelOp.getIVs())
268         if (llvm::is_contained(ifOperands, iv))
269           break;
270     } else if (!isa<AffineIfOp>(parentOp)) {
271       // Won't walk up past anything other than affine.for/if ops.
272       break;
273     }
274     // You can always hoist up past any affine.if ops.
275     res = parentOp;
276   }
277   return res;
278 }
279 
280 /// A helper for the mechanics of mlir::hoistAffineIfOp. Hoists `ifOp` just over
281 /// `hoistOverOp`. Returns the new hoisted op if any hoisting happened,
282 /// otherwise the same `ifOp`.
283 static AffineIfOp hoistAffineIfOp(AffineIfOp ifOp, Operation *hoistOverOp) {
284   // No hoisting to do.
285   if (hoistOverOp == ifOp)
286     return ifOp;
287 
288   // Create the hoisted 'if' first. Then, clone the op we are hoisting over for
289   // the else block. Then drop the else block of the original 'if' in the 'then'
290   // branch while promoting its then block, and analogously drop the 'then'
291   // block of the original 'if' from the 'else' branch while promoting its else
292   // block.
293   BlockAndValueMapping operandMap;
294   OpBuilder b(hoistOverOp);
295   auto hoistedIfOp = b.create<AffineIfOp>(ifOp.getLoc(), ifOp.getIntegerSet(),
296                                           ifOp.getOperands(),
297                                           /*elseBlock=*/true);
298 
299   // Create a clone of hoistOverOp to use for the else branch of the hoisted
300   // conditional. The else block may get optimized away if empty.
301   Operation *hoistOverOpClone = nullptr;
302   // We use this unique name to identify/find  `ifOp`'s clone in the else
303   // version.
304   StringAttr idForIfOp = b.getStringAttr("__mlir_if_hoisting");
305   operandMap.clear();
306   b.setInsertionPointAfter(hoistOverOp);
307   // We'll set an attribute to identify this op in a clone of this sub-tree.
308   ifOp->setAttr(idForIfOp, b.getBoolAttr(true));
309   hoistOverOpClone = b.clone(*hoistOverOp, operandMap);
310 
311   // Promote the 'then' block of the original affine.if in the then version.
312   promoteIfBlock(ifOp, /*elseBlock=*/false);
313 
314   // Move the then version to the hoisted if op's 'then' block.
315   auto *thenBlock = hoistedIfOp.getThenBlock();
316   thenBlock->getOperations().splice(thenBlock->begin(),
317                                     hoistOverOp->getBlock()->getOperations(),
318                                     Block::iterator(hoistOverOp));
319 
320   // Find the clone of the original affine.if op in the else version.
321   AffineIfOp ifCloneInElse;
322   hoistOverOpClone->walk([&](AffineIfOp ifClone) {
323     if (!ifClone->getAttr(idForIfOp))
324       return WalkResult::advance();
325     ifCloneInElse = ifClone;
326     return WalkResult::interrupt();
327   });
328   assert(ifCloneInElse && "if op clone should exist");
329   // For the else block, promote the else block of the original 'if' if it had
330   // one; otherwise, the op itself is to be erased.
331   if (!ifCloneInElse.hasElse())
332     ifCloneInElse.erase();
333   else
334     promoteIfBlock(ifCloneInElse, /*elseBlock=*/true);
335 
336   // Move the else version into the else block of the hoisted if op.
337   auto *elseBlock = hoistedIfOp.getElseBlock();
338   elseBlock->getOperations().splice(
339       elseBlock->begin(), hoistOverOpClone->getBlock()->getOperations(),
340       Block::iterator(hoistOverOpClone));
341 
342   return hoistedIfOp;
343 }
344 
345 LogicalResult
346 mlir::affineParallelize(AffineForOp forOp,
347                         ArrayRef<LoopReduction> parallelReductions) {
348   // Fail early if there are iter arguments that are not reductions.
349   unsigned numReductions = parallelReductions.size();
350   if (numReductions != forOp.getNumIterOperands())
351     return failure();
352 
353   Location loc = forOp.getLoc();
354   OpBuilder outsideBuilder(forOp);
355   AffineMap lowerBoundMap = forOp.getLowerBoundMap();
356   ValueRange lowerBoundOperands = forOp.getLowerBoundOperands();
357   AffineMap upperBoundMap = forOp.getUpperBoundMap();
358   ValueRange upperBoundOperands = forOp.getUpperBoundOperands();
359 
360   // Creating empty 1-D affine.parallel op.
361   auto reducedValues = llvm::to_vector<4>(llvm::map_range(
362       parallelReductions, [](const LoopReduction &red) { return red.value; }));
363   auto reductionKinds = llvm::to_vector<4>(llvm::map_range(
364       parallelReductions, [](const LoopReduction &red) { return red.kind; }));
365   AffineParallelOp newPloop = outsideBuilder.create<AffineParallelOp>(
366       loc, ValueRange(reducedValues).getTypes(), reductionKinds,
367       llvm::makeArrayRef(lowerBoundMap), lowerBoundOperands,
368       llvm::makeArrayRef(upperBoundMap), upperBoundOperands,
369       llvm::makeArrayRef(forOp.getStep()));
370   // Steal the body of the old affine for op.
371   newPloop.region().takeBody(forOp.region());
372   Operation *yieldOp = &newPloop.getBody()->back();
373 
374   // Handle the initial values of reductions because the parallel loop always
375   // starts from the neutral value.
376   SmallVector<Value> newResults;
377   newResults.reserve(numReductions);
378   for (unsigned i = 0; i < numReductions; ++i) {
379     Value init = forOp.getIterOperands()[i];
380     // This works because we are only handling single-op reductions at the
381     // moment. A switch on reduction kind or a mechanism to collect operations
382     // participating in the reduction will be necessary for multi-op reductions.
383     Operation *reductionOp = yieldOp->getOperand(i).getDefiningOp();
384     assert(reductionOp && "yielded value is expected to be produced by an op");
385     outsideBuilder.getInsertionBlock()->getOperations().splice(
386         outsideBuilder.getInsertionPoint(), newPloop.getBody()->getOperations(),
387         reductionOp);
388     reductionOp->setOperands({init, newPloop->getResult(i)});
389     forOp->getResult(i).replaceAllUsesWith(reductionOp->getResult(0));
390   }
391 
392   // Update the loop terminator to yield reduced values bypassing the reduction
393   // operation itself (now moved outside of the loop) and erase the block
394   // arguments that correspond to reductions. Note that the loop always has one
395   // "main" induction variable whenc coming from a non-parallel for.
396   unsigned numIVs = 1;
397   yieldOp->setOperands(reducedValues);
398   newPloop.getBody()->eraseArguments(
399       llvm::to_vector<4>(llvm::seq<unsigned>(numIVs, numReductions + numIVs)));
400 
401   forOp.erase();
402   return success();
403 }
404 
405 // Returns success if any hoisting happened.
406 LogicalResult mlir::hoistAffineIfOp(AffineIfOp ifOp, bool *folded) {
407   // Bail out early if the ifOp returns a result.  TODO: Consider how to
408   // properly support this case.
409   if (ifOp.getNumResults() != 0)
410     return failure();
411 
412   // Apply canonicalization patterns and folding - this is necessary for the
413   // hoisting check to be correct (operands should be composed), and to be more
414   // effective (no unused operands). Since the pattern rewriter's folding is
415   // entangled with application of patterns, we may fold/end up erasing the op,
416   // in which case we return with `folded` being set.
417   RewritePatternSet patterns(ifOp.getContext());
418   AffineIfOp::getCanonicalizationPatterns(patterns, ifOp.getContext());
419   bool erased;
420   FrozenRewritePatternSet frozenPatterns(std::move(patterns));
421   (void)applyOpPatternsAndFold(ifOp, frozenPatterns, &erased);
422   if (erased) {
423     if (folded)
424       *folded = true;
425     return failure();
426   }
427   if (folded)
428     *folded = false;
429 
430   // The folding above should have ensured this, but the affine.if's
431   // canonicalization is missing composition of affine.applys into it.
432   assert(llvm::all_of(ifOp.getOperands(),
433                       [](Value v) {
434                         return isTopLevelValue(v) || isForInductionVar(v);
435                       }) &&
436          "operands not composed");
437 
438   // We are going hoist as high as possible.
439   // TODO: this could be customized in the future.
440   auto *hoistOverOp = getOutermostInvariantForOp(ifOp);
441 
442   AffineIfOp hoistedIfOp = ::hoistAffineIfOp(ifOp, hoistOverOp);
443   // Nothing to hoist over.
444   if (hoistedIfOp == ifOp)
445     return failure();
446 
447   // Canonicalize to remove dead else blocks (happens whenever an 'if' moves up
448   // a sequence of affine.fors that are all perfectly nested).
449   (void)applyPatternsAndFoldGreedily(
450       hoistedIfOp->getParentWithTrait<OpTrait::IsIsolatedFromAbove>(),
451       frozenPatterns);
452 
453   return success();
454 }
455 
456 // Return the min expr after replacing the given dim.
457 AffineExpr mlir::substWithMin(AffineExpr e, AffineExpr dim, AffineExpr min,
458                               AffineExpr max, bool positivePath) {
459   if (e == dim)
460     return positivePath ? min : max;
461   if (auto bin = e.dyn_cast<AffineBinaryOpExpr>()) {
462     AffineExpr lhs = bin.getLHS();
463     AffineExpr rhs = bin.getRHS();
464     if (bin.getKind() == mlir::AffineExprKind::Add)
465       return substWithMin(lhs, dim, min, max, positivePath) +
466              substWithMin(rhs, dim, min, max, positivePath);
467 
468     auto c1 = bin.getLHS().dyn_cast<AffineConstantExpr>();
469     auto c2 = bin.getRHS().dyn_cast<AffineConstantExpr>();
470     if (c1 && c1.getValue() < 0)
471       return getAffineBinaryOpExpr(
472           bin.getKind(), c1, substWithMin(rhs, dim, min, max, !positivePath));
473     if (c2 && c2.getValue() < 0)
474       return getAffineBinaryOpExpr(
475           bin.getKind(), substWithMin(lhs, dim, min, max, !positivePath), c2);
476     return getAffineBinaryOpExpr(
477         bin.getKind(), substWithMin(lhs, dim, min, max, positivePath),
478         substWithMin(rhs, dim, min, max, positivePath));
479   }
480   return e;
481 }
482 
483 void mlir::normalizeAffineParallel(AffineParallelOp op) {
484   // Loops with min/max in bounds are not normalized at the moment.
485   if (op.hasMinMaxBounds())
486     return;
487 
488   AffineMap lbMap = op.lowerBoundsMap();
489   SmallVector<int64_t, 8> steps = op.getSteps();
490   // No need to do any work if the parallel op is already normalized.
491   bool isAlreadyNormalized =
492       llvm::all_of(llvm::zip(steps, lbMap.getResults()), [](auto tuple) {
493         int64_t step = std::get<0>(tuple);
494         auto lbExpr =
495             std::get<1>(tuple).template dyn_cast<AffineConstantExpr>();
496         return lbExpr && lbExpr.getValue() == 0 && step == 1;
497       });
498   if (isAlreadyNormalized)
499     return;
500 
501   AffineValueMap ranges;
502   AffineValueMap::difference(op.getUpperBoundsValueMap(),
503                              op.getLowerBoundsValueMap(), &ranges);
504   auto builder = OpBuilder::atBlockBegin(op.getBody());
505   auto zeroExpr = builder.getAffineConstantExpr(0);
506   SmallVector<AffineExpr, 8> lbExprs;
507   SmallVector<AffineExpr, 8> ubExprs;
508   for (unsigned i = 0, e = steps.size(); i < e; ++i) {
509     int64_t step = steps[i];
510 
511     // Adjust the lower bound to be 0.
512     lbExprs.push_back(zeroExpr);
513 
514     // Adjust the upper bound expression: 'range / step'.
515     AffineExpr ubExpr = ranges.getResult(i).ceilDiv(step);
516     ubExprs.push_back(ubExpr);
517 
518     // Adjust the corresponding IV: 'lb + i * step'.
519     BlockArgument iv = op.getBody()->getArgument(i);
520     AffineExpr lbExpr = lbMap.getResult(i);
521     unsigned nDims = lbMap.getNumDims();
522     auto expr = lbExpr + builder.getAffineDimExpr(nDims) * step;
523     auto map = AffineMap::get(/*dimCount=*/nDims + 1,
524                               /*symbolCount=*/lbMap.getNumSymbols(), expr);
525 
526     // Use an 'affine.apply' op that will be simplified later in subsequent
527     // canonicalizations.
528     OperandRange lbOperands = op.getLowerBoundsOperands();
529     OperandRange dimOperands = lbOperands.take_front(nDims);
530     OperandRange symbolOperands = lbOperands.drop_front(nDims);
531     SmallVector<Value, 8> applyOperands{dimOperands};
532     applyOperands.push_back(iv);
533     applyOperands.append(symbolOperands.begin(), symbolOperands.end());
534     auto apply = builder.create<AffineApplyOp>(op.getLoc(), map, applyOperands);
535     iv.replaceAllUsesExcept(apply, apply);
536   }
537 
538   SmallVector<int64_t, 8> newSteps(op.getNumDims(), 1);
539   op.setSteps(newSteps);
540   auto newLowerMap = AffineMap::get(
541       /*dimCount=*/0, /*symbolCount=*/0, lbExprs, op.getContext());
542   op.setLowerBounds({}, newLowerMap);
543   auto newUpperMap = AffineMap::get(ranges.getNumDims(), ranges.getNumSymbols(),
544                                     ubExprs, op.getContext());
545   op.setUpperBounds(ranges.getOperands(), newUpperMap);
546 }
547 
548 /// Normalizes affine.for ops. If the affine.for op has only a single iteration
549 /// only then it is simply promoted, else it is normalized in the traditional
550 /// way, by converting the lower bound to zero and loop step to one. The upper
551 /// bound is set to the trip count of the loop. For now, original loops must
552 /// have lower bound with a single result only. There is no such restriction on
553 /// upper bounds.
554 void mlir::normalizeAffineFor(AffineForOp op) {
555   if (succeeded(promoteIfSingleIteration(op)))
556     return;
557 
558   // Check if the forop is already normalized.
559   if (op.hasConstantLowerBound() && (op.getConstantLowerBound() == 0) &&
560       (op.getStep() == 1))
561     return;
562 
563   // Check if the lower bound has a single result only. Loops with a max lower
564   // bound can't be normalized without additional support like
565   // affine.execute_region's. If the lower bound does not have a single result
566   // then skip this op.
567   if (op.getLowerBoundMap().getNumResults() != 1)
568     return;
569 
570   Location loc = op.getLoc();
571   OpBuilder opBuilder(op);
572   int64_t origLoopStep = op.getStep();
573 
574   // Calculate upperBound for normalized loop.
575   SmallVector<Value, 4> ubOperands;
576   AffineBound lb = op.getLowerBound();
577   AffineBound ub = op.getUpperBound();
578   ubOperands.reserve(ub.getNumOperands() + lb.getNumOperands());
579   AffineMap origLbMap = lb.getMap();
580   AffineMap origUbMap = ub.getMap();
581 
582   // Add dimension operands from upper/lower bound.
583   for (unsigned j = 0, e = origUbMap.getNumDims(); j < e; ++j)
584     ubOperands.push_back(ub.getOperand(j));
585   for (unsigned j = 0, e = origLbMap.getNumDims(); j < e; ++j)
586     ubOperands.push_back(lb.getOperand(j));
587 
588   // Add symbol operands from upper/lower bound.
589   for (unsigned j = 0, e = origUbMap.getNumSymbols(); j < e; ++j)
590     ubOperands.push_back(ub.getOperand(origUbMap.getNumDims() + j));
591   for (unsigned j = 0, e = origLbMap.getNumSymbols(); j < e; ++j)
592     ubOperands.push_back(lb.getOperand(origLbMap.getNumDims() + j));
593 
594   // Add original result expressions from lower/upper bound map.
595   SmallVector<AffineExpr, 1> origLbExprs(origLbMap.getResults().begin(),
596                                          origLbMap.getResults().end());
597   SmallVector<AffineExpr, 2> origUbExprs(origUbMap.getResults().begin(),
598                                          origUbMap.getResults().end());
599   SmallVector<AffineExpr, 4> newUbExprs;
600 
601   // The original upperBound can have more than one result. For the new
602   // upperBound of this loop, take difference of all possible combinations of
603   // the ub results and lb result and ceildiv with the loop step. For e.g.,
604   //
605   //  affine.for %i1 = 0 to min affine_map<(d0)[] -> (d0 + 32, 1024)>(%i0)
606   //  will have an upperBound map as,
607   //  affine_map<(d0)[] -> (((d0 + 32) - 0) ceildiv 1, (1024 - 0) ceildiv
608   //  1)>(%i0)
609   //
610   // Insert all combinations of upper/lower bound results.
611   for (unsigned i = 0, e = origUbExprs.size(); i < e; ++i) {
612     newUbExprs.push_back(
613         (origUbExprs[i] - origLbExprs[0]).ceilDiv(origLoopStep));
614   }
615 
616   // Construct newUbMap.
617   AffineMap newUbMap =
618       AffineMap::get(origLbMap.getNumDims() + origUbMap.getNumDims(),
619                      origLbMap.getNumSymbols() + origUbMap.getNumSymbols(),
620                      newUbExprs, opBuilder.getContext());
621 
622   // Normalize the loop.
623   op.setUpperBound(ubOperands, newUbMap);
624   op.setLowerBound({}, opBuilder.getConstantAffineMap(0));
625   op.setStep(1);
626 
627   // Calculate the Value of new loopIV. Create affine.apply for the value of
628   // the loopIV in normalized loop.
629   opBuilder.setInsertionPointToStart(op.getBody());
630   SmallVector<Value, 4> lbOperands(lb.getOperands().begin(),
631                                    lb.getOperands().begin() +
632                                        lb.getMap().getNumDims());
633   // Add an extra dim operand for loopIV.
634   lbOperands.push_back(op.getInductionVar());
635   // Add symbol operands from lower bound.
636   for (unsigned j = 0, e = origLbMap.getNumSymbols(); j < e; ++j)
637     lbOperands.push_back(lb.getOperand(origLbMap.getNumDims() + j));
638 
639   AffineExpr origIVExpr = opBuilder.getAffineDimExpr(lb.getMap().getNumDims());
640   AffineExpr newIVExpr = origIVExpr * origLoopStep + origLbMap.getResult(0);
641   AffineMap ivMap = AffineMap::get(origLbMap.getNumDims() + 1,
642                                    origLbMap.getNumSymbols(), newIVExpr);
643   Operation *newIV = opBuilder.create<AffineApplyOp>(loc, ivMap, lbOperands);
644   op.getInductionVar().replaceAllUsesExcept(newIV->getResult(0), newIV);
645 }
646 
647 /// Ensure that all operations that could be executed after `start`
648 /// (noninclusive) and prior to `memOp` (e.g. on a control flow/op path
649 /// between the operations) do not have the potential memory effect
650 /// `EffectType` on `memOp`. `memOp`  is an operation that reads or writes to
651 /// a memref. For example, if `EffectType` is MemoryEffects::Write, this method
652 /// will check if there is no write to the memory between `start` and `memOp`
653 /// that would change the read within `memOp`.
654 template <typename EffectType, typename T>
655 static bool hasNoInterveningEffect(Operation *start, T memOp) {
656   Value memref = memOp.getMemRef();
657   bool isOriginalAllocation = memref.getDefiningOp<memref::AllocaOp>() ||
658                               memref.getDefiningOp<memref::AllocOp>();
659 
660   // A boolean representing whether an intervening operation could have impacted
661   // memOp.
662   bool hasSideEffect = false;
663 
664   // Check whether the effect on memOp can be caused by a given operation op.
665   std::function<void(Operation *)> checkOperation = [&](Operation *op) {
666     // If the effect has alreay been found, early exit,
667     if (hasSideEffect)
668       return;
669 
670     if (auto memEffect = dyn_cast<MemoryEffectOpInterface>(op)) {
671       SmallVector<MemoryEffects::EffectInstance, 1> effects;
672       memEffect.getEffects(effects);
673 
674       bool opMayHaveEffect = false;
675       for (auto effect : effects) {
676         // If op causes EffectType on a potentially aliasing location for
677         // memOp, mark as having the effect.
678         if (isa<EffectType>(effect.getEffect())) {
679           if (isOriginalAllocation && effect.getValue() &&
680               (effect.getValue().getDefiningOp<memref::AllocaOp>() ||
681                effect.getValue().getDefiningOp<memref::AllocOp>())) {
682             if (effect.getValue() != memref)
683               continue;
684           }
685           opMayHaveEffect = true;
686           break;
687         }
688       }
689 
690       if (!opMayHaveEffect)
691         return;
692 
693       // If the side effect comes from an affine read or write, try to
694       // prove the side effecting `op` cannot reach `memOp`.
695       if (isa<AffineReadOpInterface, AffineWriteOpInterface>(op)) {
696         MemRefAccess srcAccess(op);
697         MemRefAccess destAccess(memOp);
698         // Dependence analysis is only correct if both ops operate on the same
699         // memref.
700         if (srcAccess.memref == destAccess.memref) {
701           FlatAffineValueConstraints dependenceConstraints;
702 
703           // Number of loops containing the start op and the ending operation.
704           unsigned minSurroundingLoops =
705               getNumCommonSurroundingLoops(*start, *memOp);
706 
707           // Number of loops containing the operation `op` which has the
708           // potential memory side effect and can occur on a path between
709           // `start` and `memOp`.
710           unsigned nsLoops = getNumCommonSurroundingLoops(*op, *memOp);
711 
712           // For ease, let's consider the case that `op` is a store and we're
713           // looking for other potential stores (e.g `op`) that overwrite memory
714           // after `start`, and before being read in `memOp`. In this case, we
715           // only need to consider other potential stores with depth >
716           // minSurrounding loops since `start` would overwrite any store with a
717           // smaller number of surrounding loops before.
718           unsigned d;
719           for (d = nsLoops + 1; d > minSurroundingLoops; d--) {
720             DependenceResult result = checkMemrefAccessDependence(
721                 srcAccess, destAccess, d, &dependenceConstraints,
722                 /*dependenceComponents=*/nullptr);
723             if (hasDependence(result)) {
724               hasSideEffect = true;
725               return;
726             }
727           }
728 
729           // No side effect was seen, simply return.
730           return;
731         }
732       }
733       hasSideEffect = true;
734       return;
735     }
736 
737     if (op->hasTrait<OpTrait::HasRecursiveSideEffects>()) {
738       // Recurse into the regions for this op and check whether the internal
739       // operations may have the side effect `EffectType` on memOp.
740       for (Region &region : op->getRegions())
741         for (Block &block : region)
742           for (Operation &op : block)
743             checkOperation(&op);
744       return;
745     }
746 
747     // Otherwise, conservatively assume generic operations have the effect
748     // on the operation
749     hasSideEffect = true;
750   };
751 
752   // Check all paths from ancestor op `parent` to the operation `to` for the
753   // effect. It is known that `to` must be contained within `parent`.
754   auto until = [&](Operation *parent, Operation *to) {
755     // TODO check only the paths from `parent` to `to`.
756     // Currently we fallback and check the entire parent op, rather than
757     // just the paths from the parent path, stopping after reaching `to`.
758     // This is conservatively correct, but could be made more aggressive.
759     assert(parent->isAncestor(to));
760     checkOperation(parent);
761   };
762 
763   // Check for all paths from operation `from` to operation `untilOp` for the
764   // given memory effect.
765   std::function<void(Operation *, Operation *)> recur =
766       [&](Operation *from, Operation *untilOp) {
767         assert(
768             from->getParentRegion()->isAncestor(untilOp->getParentRegion()) &&
769             "Checking for side effect between two operations without a common "
770             "ancestor");
771 
772         // If the operations are in different regions, recursively consider all
773         // path from `from` to the parent of `to` and all paths from the parent
774         // of `to` to `to`.
775         if (from->getParentRegion() != untilOp->getParentRegion()) {
776           recur(from, untilOp->getParentOp());
777           until(untilOp->getParentOp(), untilOp);
778           return;
779         }
780 
781         // Now, assuming that `from` and `to` exist in the same region, perform
782         // a CFG traversal to check all the relevant operations.
783 
784         // Additional blocks to consider.
785         SmallVector<Block *, 2> todoBlocks;
786         {
787           // First consider the parent block of `from` an check all operations
788           // after `from`.
789           for (auto iter = ++from->getIterator(), end = from->getBlock()->end();
790                iter != end && &*iter != untilOp; ++iter) {
791             checkOperation(&*iter);
792           }
793 
794           // If the parent of `from` doesn't contain `to`, add the successors
795           // to the list of blocks to check.
796           if (untilOp->getBlock() != from->getBlock())
797             for (Block *succ : from->getBlock()->getSuccessors())
798               todoBlocks.push_back(succ);
799         }
800 
801         SmallPtrSet<Block *, 4> done;
802         // Traverse the CFG until hitting `to`.
803         while (!todoBlocks.empty()) {
804           Block *blk = todoBlocks.pop_back_val();
805           if (done.count(blk))
806             continue;
807           done.insert(blk);
808           for (auto &op : *blk) {
809             if (&op == untilOp)
810               break;
811             checkOperation(&op);
812             if (&op == blk->getTerminator())
813               for (Block *succ : blk->getSuccessors())
814                 todoBlocks.push_back(succ);
815           }
816         }
817       };
818   recur(start, memOp);
819   return !hasSideEffect;
820 }
821 
822 /// Attempt to eliminate loadOp by replacing it with a value stored into memory
823 /// which the load is guaranteed to retrieve. This check involves three
824 /// components: 1) The store and load must be on the same location 2) The store
825 /// must dominate (and therefore must always occur prior to) the load 3) No
826 /// other operations will overwrite the memory loaded between the given load
827 /// and store.  If such a value exists, the replaced `loadOp` will be added to
828 /// `loadOpsToErase` and its memref will be added to `memrefsToErase`.
829 static LogicalResult forwardStoreToLoad(
830     AffineReadOpInterface loadOp, SmallVectorImpl<Operation *> &loadOpsToErase,
831     SmallPtrSetImpl<Value> &memrefsToErase, DominanceInfo &domInfo) {
832 
833   // The store op candidate for forwarding that satisfies all conditions
834   // to replace the load, if any.
835   Operation *lastWriteStoreOp = nullptr;
836 
837   for (auto *user : loadOp.getMemRef().getUsers()) {
838     auto storeOp = dyn_cast<AffineWriteOpInterface>(user);
839     if (!storeOp)
840       continue;
841     MemRefAccess srcAccess(storeOp);
842     MemRefAccess destAccess(loadOp);
843 
844     // 1. Check if the store and the load have mathematically equivalent
845     // affine access functions; this implies that they statically refer to the
846     // same single memref element. As an example this filters out cases like:
847     //     store %A[%i0 + 1]
848     //     load %A[%i0]
849     //     store %A[%M]
850     //     load %A[%N]
851     // Use the AffineValueMap difference based memref access equality checking.
852     if (srcAccess != destAccess)
853       continue;
854 
855     // 2. The store has to dominate the load op to be candidate.
856     if (!domInfo.dominates(storeOp, loadOp))
857       continue;
858 
859     // 3. Ensure there is no intermediate operation which could replace the
860     // value in memory.
861     if (!hasNoInterveningEffect<MemoryEffects::Write>(storeOp, loadOp))
862       continue;
863 
864     // We now have a candidate for forwarding.
865     assert(lastWriteStoreOp == nullptr &&
866            "multiple simulataneous replacement stores");
867     lastWriteStoreOp = storeOp;
868   }
869 
870   if (!lastWriteStoreOp)
871     return failure();
872 
873   // Perform the actual store to load forwarding.
874   Value storeVal =
875       cast<AffineWriteOpInterface>(lastWriteStoreOp).getValueToStore();
876   // Check if 2 values have the same shape. This is needed for affine vector
877   // loads and stores.
878   if (storeVal.getType() != loadOp.getValue().getType())
879     return failure();
880   loadOp.getValue().replaceAllUsesWith(storeVal);
881   // Record the memref for a later sweep to optimize away.
882   memrefsToErase.insert(loadOp.getMemRef());
883   // Record this to erase later.
884   loadOpsToErase.push_back(loadOp);
885   return success();
886 }
887 
888 // This attempts to find stores which have no impact on the final result.
889 // A writing op writeA will be eliminated if there exists an op writeB if
890 // 1) writeA and writeB have mathematically equivalent affine access functions.
891 // 2) writeB postdominates writeA.
892 // 3) There is no potential read between writeA and writeB.
893 static void findUnusedStore(AffineWriteOpInterface writeA,
894                             SmallVectorImpl<Operation *> &opsToErase,
895                             PostDominanceInfo &postDominanceInfo) {
896 
897   for (Operation *user : writeA.getMemRef().getUsers()) {
898     // Only consider writing operations.
899     auto writeB = dyn_cast<AffineWriteOpInterface>(user);
900     if (!writeB)
901       continue;
902 
903     // The operations must be distinct.
904     if (writeB == writeA)
905       continue;
906 
907     // Both operations must lie in the same region.
908     if (writeB->getParentRegion() != writeA->getParentRegion())
909       continue;
910 
911     // Both operations must write to the same memory.
912     MemRefAccess srcAccess(writeB);
913     MemRefAccess destAccess(writeA);
914 
915     if (srcAccess != destAccess)
916       continue;
917 
918     // writeB must postdominate writeA.
919     if (!postDominanceInfo.postDominates(writeB, writeA))
920       continue;
921 
922     // There cannot be an operation which reads from memory between
923     // the two writes.
924     if (!hasNoInterveningEffect<MemoryEffects::Read>(writeA, writeB))
925       continue;
926 
927     opsToErase.push_back(writeA);
928     break;
929   }
930 }
931 
932 // The load to load forwarding / redundant load elimination is similar to the
933 // store to load forwarding.
934 // loadA will be be replaced with loadB if:
935 // 1) loadA and loadB have mathematically equivalent affine access functions.
936 // 2) loadB dominates loadA.
937 // 3) There is no write between loadA and loadB.
938 static void loadCSE(AffineReadOpInterface loadA,
939                     SmallVectorImpl<Operation *> &loadOpsToErase,
940                     DominanceInfo &domInfo) {
941   SmallVector<AffineReadOpInterface, 4> loadCandidates;
942   for (auto *user : loadA.getMemRef().getUsers()) {
943     auto loadB = dyn_cast<AffineReadOpInterface>(user);
944     if (!loadB || loadB == loadA)
945       continue;
946 
947     MemRefAccess srcAccess(loadB);
948     MemRefAccess destAccess(loadA);
949 
950     // 1. The accesses have to be to the same location.
951     if (srcAccess != destAccess) {
952       continue;
953     }
954 
955     // 2. The store has to dominate the load op to be candidate.
956     if (!domInfo.dominates(loadB, loadA))
957       continue;
958 
959     // 3. There is no write between loadA and loadB.
960     if (!hasNoInterveningEffect<MemoryEffects::Write>(loadB.getOperation(),
961                                                       loadA))
962       continue;
963 
964     // Check if two values have the same shape. This is needed for affine vector
965     // loads.
966     if (loadB.getValue().getType() != loadA.getValue().getType())
967       continue;
968 
969     loadCandidates.push_back(loadB);
970   }
971 
972   // Of the legal load candidates, use the one that dominates all others
973   // to minimize the subsequent need to loadCSE
974   Value loadB;
975   for (AffineReadOpInterface option : loadCandidates) {
976     if (llvm::all_of(loadCandidates, [&](AffineReadOpInterface depStore) {
977           return depStore == option ||
978                  domInfo.dominates(option.getOperation(),
979                                    depStore.getOperation());
980         })) {
981       loadB = option.getValue();
982       break;
983     }
984   }
985 
986   if (loadB) {
987     loadA.getValue().replaceAllUsesWith(loadB);
988     // Record this to erase later.
989     loadOpsToErase.push_back(loadA);
990   }
991 }
992 
993 // The store to load forwarding and load CSE rely on three conditions:
994 //
995 // 1) store/load providing a replacement value and load being replaced need to
996 // have mathematically equivalent affine access functions (checked after full
997 // composition of load/store operands); this implies that they access the same
998 // single memref element for all iterations of the common surrounding loop,
999 //
1000 // 2) the store/load op should dominate the load op,
1001 //
1002 // 3) no operation that may write to memory read by the load being replaced can
1003 // occur after executing the instruction (load or store) providing the
1004 // replacement value and before the load being replaced (thus potentially
1005 // allowing overwriting the memory read by the load).
1006 //
1007 // The above conditions are simple to check, sufficient, and powerful for most
1008 // cases in practice - they are sufficient, but not necessary --- since they
1009 // don't reason about loops that are guaranteed to execute at least once or
1010 // multiple sources to forward from.
1011 //
1012 // TODO: more forwarding can be done when support for
1013 // loop/conditional live-out SSA values is available.
1014 // TODO: do general dead store elimination for memref's. This pass
1015 // currently only eliminates the stores only if no other loads/uses (other
1016 // than dealloc) remain.
1017 //
1018 void mlir::affineScalarReplace(FuncOp f, DominanceInfo &domInfo,
1019                                PostDominanceInfo &postDomInfo) {
1020   // Load op's whose results were replaced by those forwarded from stores.
1021   SmallVector<Operation *, 8> opsToErase;
1022 
1023   // A list of memref's that are potentially dead / could be eliminated.
1024   SmallPtrSet<Value, 4> memrefsToErase;
1025 
1026   // Walk all load's and perform store to load forwarding.
1027   f.walk([&](AffineReadOpInterface loadOp) {
1028     if (failed(
1029             forwardStoreToLoad(loadOp, opsToErase, memrefsToErase, domInfo))) {
1030       loadCSE(loadOp, opsToErase, domInfo);
1031     }
1032   });
1033 
1034   // Erase all load op's whose results were replaced with store fwd'ed ones.
1035   for (auto *op : opsToErase)
1036     op->erase();
1037   opsToErase.clear();
1038 
1039   // Walk all store's and perform unused store elimination
1040   f.walk([&](AffineWriteOpInterface storeOp) {
1041     findUnusedStore(storeOp, opsToErase, postDomInfo);
1042   });
1043   // Erase all store op's which don't impact the program
1044   for (auto *op : opsToErase)
1045     op->erase();
1046 
1047   // Check if the store fwd'ed memrefs are now left with only stores and can
1048   // thus be completely deleted. Note: the canonicalize pass should be able
1049   // to do this as well, but we'll do it here since we collected these anyway.
1050   for (auto memref : memrefsToErase) {
1051     // If the memref hasn't been alloc'ed in this function, skip.
1052     Operation *defOp = memref.getDefiningOp();
1053     if (!defOp || !isa<memref::AllocOp>(defOp))
1054       // TODO: if the memref was returned by a 'call' operation, we
1055       // could still erase it if the call had no side-effects.
1056       continue;
1057     if (llvm::any_of(memref.getUsers(), [&](Operation *ownerOp) {
1058           return !isa<AffineWriteOpInterface, memref::DeallocOp>(ownerOp);
1059         }))
1060       continue;
1061 
1062     // Erase all stores, the dealloc, and the alloc on the memref.
1063     for (auto *user : llvm::make_early_inc_range(memref.getUsers()))
1064       user->erase();
1065     defOp->erase();
1066   }
1067 }
1068 
1069 // Perform the replacement in `op`.
1070 LogicalResult mlir::replaceAllMemRefUsesWith(Value oldMemRef, Value newMemRef,
1071                                              Operation *op,
1072                                              ArrayRef<Value> extraIndices,
1073                                              AffineMap indexRemap,
1074                                              ArrayRef<Value> extraOperands,
1075                                              ArrayRef<Value> symbolOperands,
1076                                              bool allowNonDereferencingOps) {
1077   unsigned newMemRefRank = newMemRef.getType().cast<MemRefType>().getRank();
1078   (void)newMemRefRank; // unused in opt mode
1079   unsigned oldMemRefRank = oldMemRef.getType().cast<MemRefType>().getRank();
1080   (void)oldMemRefRank; // unused in opt mode
1081   if (indexRemap) {
1082     assert(indexRemap.getNumSymbols() == symbolOperands.size() &&
1083            "symbolic operand count mismatch");
1084     assert(indexRemap.getNumInputs() ==
1085            extraOperands.size() + oldMemRefRank + symbolOperands.size());
1086     assert(indexRemap.getNumResults() + extraIndices.size() == newMemRefRank);
1087   } else {
1088     assert(oldMemRefRank + extraIndices.size() == newMemRefRank);
1089   }
1090 
1091   // Assert same elemental type.
1092   assert(oldMemRef.getType().cast<MemRefType>().getElementType() ==
1093          newMemRef.getType().cast<MemRefType>().getElementType());
1094 
1095   SmallVector<unsigned, 2> usePositions;
1096   for (const auto &opEntry : llvm::enumerate(op->getOperands())) {
1097     if (opEntry.value() == oldMemRef)
1098       usePositions.push_back(opEntry.index());
1099   }
1100 
1101   // If memref doesn't appear, nothing to do.
1102   if (usePositions.empty())
1103     return success();
1104 
1105   if (usePositions.size() > 1) {
1106     // TODO: extend it for this case when needed (rare).
1107     assert(false && "multiple dereferencing uses in a single op not supported");
1108     return failure();
1109   }
1110 
1111   unsigned memRefOperandPos = usePositions.front();
1112 
1113   OpBuilder builder(op);
1114   // The following checks if op is dereferencing memref and performs the access
1115   // index rewrites.
1116   auto affMapAccInterface = dyn_cast<AffineMapAccessInterface>(op);
1117   if (!affMapAccInterface) {
1118     if (!allowNonDereferencingOps) {
1119       // Failure: memref used in a non-dereferencing context (potentially
1120       // escapes); no replacement in these cases unless allowNonDereferencingOps
1121       // is set.
1122       return failure();
1123     }
1124     op->setOperand(memRefOperandPos, newMemRef);
1125     return success();
1126   }
1127   // Perform index rewrites for the dereferencing op and then replace the op
1128   NamedAttribute oldMapAttrPair =
1129       affMapAccInterface.getAffineMapAttrForMemRef(oldMemRef);
1130   AffineMap oldMap = oldMapAttrPair.getValue().cast<AffineMapAttr>().getValue();
1131   unsigned oldMapNumInputs = oldMap.getNumInputs();
1132   SmallVector<Value, 4> oldMapOperands(
1133       op->operand_begin() + memRefOperandPos + 1,
1134       op->operand_begin() + memRefOperandPos + 1 + oldMapNumInputs);
1135 
1136   // Apply 'oldMemRefOperands = oldMap(oldMapOperands)'.
1137   SmallVector<Value, 4> oldMemRefOperands;
1138   SmallVector<Value, 4> affineApplyOps;
1139   oldMemRefOperands.reserve(oldMemRefRank);
1140   if (oldMap != builder.getMultiDimIdentityMap(oldMap.getNumDims())) {
1141     for (auto resultExpr : oldMap.getResults()) {
1142       auto singleResMap = AffineMap::get(oldMap.getNumDims(),
1143                                          oldMap.getNumSymbols(), resultExpr);
1144       auto afOp = builder.create<AffineApplyOp>(op->getLoc(), singleResMap,
1145                                                 oldMapOperands);
1146       oldMemRefOperands.push_back(afOp);
1147       affineApplyOps.push_back(afOp);
1148     }
1149   } else {
1150     oldMemRefOperands.assign(oldMapOperands.begin(), oldMapOperands.end());
1151   }
1152 
1153   // Construct new indices as a remap of the old ones if a remapping has been
1154   // provided. The indices of a memref come right after it, i.e.,
1155   // at position memRefOperandPos + 1.
1156   SmallVector<Value, 4> remapOperands;
1157   remapOperands.reserve(extraOperands.size() + oldMemRefRank +
1158                         symbolOperands.size());
1159   remapOperands.append(extraOperands.begin(), extraOperands.end());
1160   remapOperands.append(oldMemRefOperands.begin(), oldMemRefOperands.end());
1161   remapOperands.append(symbolOperands.begin(), symbolOperands.end());
1162 
1163   SmallVector<Value, 4> remapOutputs;
1164   remapOutputs.reserve(oldMemRefRank);
1165 
1166   if (indexRemap &&
1167       indexRemap != builder.getMultiDimIdentityMap(indexRemap.getNumDims())) {
1168     // Remapped indices.
1169     for (auto resultExpr : indexRemap.getResults()) {
1170       auto singleResMap = AffineMap::get(
1171           indexRemap.getNumDims(), indexRemap.getNumSymbols(), resultExpr);
1172       auto afOp = builder.create<AffineApplyOp>(op->getLoc(), singleResMap,
1173                                                 remapOperands);
1174       remapOutputs.push_back(afOp);
1175       affineApplyOps.push_back(afOp);
1176     }
1177   } else {
1178     // No remapping specified.
1179     remapOutputs.assign(remapOperands.begin(), remapOperands.end());
1180   }
1181 
1182   SmallVector<Value, 4> newMapOperands;
1183   newMapOperands.reserve(newMemRefRank);
1184 
1185   // Prepend 'extraIndices' in 'newMapOperands'.
1186   for (Value extraIndex : extraIndices) {
1187     assert(extraIndex.getDefiningOp()->getNumResults() == 1 &&
1188            "single result op's expected to generate these indices");
1189     assert((isValidDim(extraIndex) || isValidSymbol(extraIndex)) &&
1190            "invalid memory op index");
1191     newMapOperands.push_back(extraIndex);
1192   }
1193 
1194   // Append 'remapOutputs' to 'newMapOperands'.
1195   newMapOperands.append(remapOutputs.begin(), remapOutputs.end());
1196 
1197   // Create new fully composed AffineMap for new op to be created.
1198   assert(newMapOperands.size() == newMemRefRank);
1199   auto newMap = builder.getMultiDimIdentityMap(newMemRefRank);
1200   // TODO: Avoid creating/deleting temporary AffineApplyOps here.
1201   fullyComposeAffineMapAndOperands(&newMap, &newMapOperands);
1202   newMap = simplifyAffineMap(newMap);
1203   canonicalizeMapAndOperands(&newMap, &newMapOperands);
1204   // Remove any affine.apply's that became dead as a result of composition.
1205   for (Value value : affineApplyOps)
1206     if (value.use_empty())
1207       value.getDefiningOp()->erase();
1208 
1209   OperationState state(op->getLoc(), op->getName());
1210   // Construct the new operation using this memref.
1211   state.operands.reserve(op->getNumOperands() + extraIndices.size());
1212   // Insert the non-memref operands.
1213   state.operands.append(op->operand_begin(),
1214                         op->operand_begin() + memRefOperandPos);
1215   // Insert the new memref value.
1216   state.operands.push_back(newMemRef);
1217 
1218   // Insert the new memref map operands.
1219   state.operands.append(newMapOperands.begin(), newMapOperands.end());
1220 
1221   // Insert the remaining operands unmodified.
1222   state.operands.append(op->operand_begin() + memRefOperandPos + 1 +
1223                             oldMapNumInputs,
1224                         op->operand_end());
1225 
1226   // Result types don't change. Both memref's are of the same elemental type.
1227   state.types.reserve(op->getNumResults());
1228   for (auto result : op->getResults())
1229     state.types.push_back(result.getType());
1230 
1231   // Add attribute for 'newMap', other Attributes do not change.
1232   auto newMapAttr = AffineMapAttr::get(newMap);
1233   for (auto namedAttr : op->getAttrs()) {
1234     if (namedAttr.getName() == oldMapAttrPair.getName())
1235       state.attributes.push_back({namedAttr.getName(), newMapAttr});
1236     else
1237       state.attributes.push_back(namedAttr);
1238   }
1239 
1240   // Create the new operation.
1241   auto *repOp = builder.createOperation(state);
1242   op->replaceAllUsesWith(repOp);
1243   op->erase();
1244 
1245   return success();
1246 }
1247 
1248 LogicalResult mlir::replaceAllMemRefUsesWith(
1249     Value oldMemRef, Value newMemRef, ArrayRef<Value> extraIndices,
1250     AffineMap indexRemap, ArrayRef<Value> extraOperands,
1251     ArrayRef<Value> symbolOperands, Operation *domOpFilter,
1252     Operation *postDomOpFilter, bool allowNonDereferencingOps,
1253     bool replaceInDeallocOp) {
1254   unsigned newMemRefRank = newMemRef.getType().cast<MemRefType>().getRank();
1255   (void)newMemRefRank; // unused in opt mode
1256   unsigned oldMemRefRank = oldMemRef.getType().cast<MemRefType>().getRank();
1257   (void)oldMemRefRank;
1258   if (indexRemap) {
1259     assert(indexRemap.getNumSymbols() == symbolOperands.size() &&
1260            "symbol operand count mismatch");
1261     assert(indexRemap.getNumInputs() ==
1262            extraOperands.size() + oldMemRefRank + symbolOperands.size());
1263     assert(indexRemap.getNumResults() + extraIndices.size() == newMemRefRank);
1264   } else {
1265     assert(oldMemRefRank + extraIndices.size() == newMemRefRank);
1266   }
1267 
1268   // Assert same elemental type.
1269   assert(oldMemRef.getType().cast<MemRefType>().getElementType() ==
1270          newMemRef.getType().cast<MemRefType>().getElementType());
1271 
1272   std::unique_ptr<DominanceInfo> domInfo;
1273   std::unique_ptr<PostDominanceInfo> postDomInfo;
1274   if (domOpFilter)
1275     domInfo =
1276         std::make_unique<DominanceInfo>(domOpFilter->getParentOfType<FuncOp>());
1277 
1278   if (postDomOpFilter)
1279     postDomInfo = std::make_unique<PostDominanceInfo>(
1280         postDomOpFilter->getParentOfType<FuncOp>());
1281 
1282   // Walk all uses of old memref; collect ops to perform replacement. We use a
1283   // DenseSet since an operation could potentially have multiple uses of a
1284   // memref (although rare), and the replacement later is going to erase ops.
1285   DenseSet<Operation *> opsToReplace;
1286   for (auto *op : oldMemRef.getUsers()) {
1287     // Skip this use if it's not dominated by domOpFilter.
1288     if (domOpFilter && !domInfo->dominates(domOpFilter, op))
1289       continue;
1290 
1291     // Skip this use if it's not post-dominated by postDomOpFilter.
1292     if (postDomOpFilter && !postDomInfo->postDominates(postDomOpFilter, op))
1293       continue;
1294 
1295     // Skip dealloc's - no replacement is necessary, and a memref replacement
1296     // at other uses doesn't hurt these dealloc's.
1297     if (isa<memref::DeallocOp>(op) && !replaceInDeallocOp)
1298       continue;
1299 
1300     // Check if the memref was used in a non-dereferencing context. It is fine
1301     // for the memref to be used in a non-dereferencing way outside of the
1302     // region where this replacement is happening.
1303     if (!isa<AffineMapAccessInterface>(*op)) {
1304       if (!allowNonDereferencingOps) {
1305         LLVM_DEBUG(llvm::dbgs()
1306                    << "Memref replacement failed: non-deferencing memref op: \n"
1307                    << *op << '\n');
1308         return failure();
1309       }
1310       // Non-dereferencing ops with the MemRefsNormalizable trait are
1311       // supported for replacement.
1312       if (!op->hasTrait<OpTrait::MemRefsNormalizable>()) {
1313         LLVM_DEBUG(llvm::dbgs() << "Memref replacement failed: use without a "
1314                                    "memrefs normalizable trait: \n"
1315                                 << *op << '\n');
1316         return failure();
1317       }
1318     }
1319 
1320     // We'll first collect and then replace --- since replacement erases the op
1321     // that has the use, and that op could be postDomFilter or domFilter itself!
1322     opsToReplace.insert(op);
1323   }
1324 
1325   for (auto *op : opsToReplace) {
1326     if (failed(replaceAllMemRefUsesWith(
1327             oldMemRef, newMemRef, op, extraIndices, indexRemap, extraOperands,
1328             symbolOperands, allowNonDereferencingOps)))
1329       llvm_unreachable("memref replacement guaranteed to succeed here");
1330   }
1331 
1332   return success();
1333 }
1334 
1335 /// Given an operation, inserts one or more single result affine
1336 /// apply operations, results of which are exclusively used by this operation
1337 /// operation. The operands of these newly created affine apply ops are
1338 /// guaranteed to be loop iterators or terminal symbols of a function.
1339 ///
1340 /// Before
1341 ///
1342 /// affine.for %i = 0 to #map(%N)
1343 ///   %idx = affine.apply (d0) -> (d0 mod 2) (%i)
1344 ///   "send"(%idx, %A, ...)
1345 ///   "compute"(%idx)
1346 ///
1347 /// After
1348 ///
1349 /// affine.for %i = 0 to #map(%N)
1350 ///   %idx = affine.apply (d0) -> (d0 mod 2) (%i)
1351 ///   "send"(%idx, %A, ...)
1352 ///   %idx_ = affine.apply (d0) -> (d0 mod 2) (%i)
1353 ///   "compute"(%idx_)
1354 ///
1355 /// This allows applying different transformations on send and compute (for eg.
1356 /// different shifts/delays).
1357 ///
1358 /// Returns nullptr either if none of opInst's operands were the result of an
1359 /// affine.apply and thus there was no affine computation slice to create, or if
1360 /// all the affine.apply op's supplying operands to this opInst did not have any
1361 /// uses besides this opInst; otherwise returns the list of affine.apply
1362 /// operations created in output argument `sliceOps`.
1363 void mlir::createAffineComputationSlice(
1364     Operation *opInst, SmallVectorImpl<AffineApplyOp> *sliceOps) {
1365   // Collect all operands that are results of affine apply ops.
1366   SmallVector<Value, 4> subOperands;
1367   subOperands.reserve(opInst->getNumOperands());
1368   for (auto operand : opInst->getOperands())
1369     if (isa_and_nonnull<AffineApplyOp>(operand.getDefiningOp()))
1370       subOperands.push_back(operand);
1371 
1372   // Gather sequence of AffineApplyOps reachable from 'subOperands'.
1373   SmallVector<Operation *, 4> affineApplyOps;
1374   getReachableAffineApplyOps(subOperands, affineApplyOps);
1375   // Skip transforming if there are no affine maps to compose.
1376   if (affineApplyOps.empty())
1377     return;
1378 
1379   // Check if all uses of the affine apply op's lie only in this op op, in
1380   // which case there would be nothing to do.
1381   bool localized = true;
1382   for (auto *op : affineApplyOps) {
1383     for (auto result : op->getResults()) {
1384       for (auto *user : result.getUsers()) {
1385         if (user != opInst) {
1386           localized = false;
1387           break;
1388         }
1389       }
1390     }
1391   }
1392   if (localized)
1393     return;
1394 
1395   OpBuilder builder(opInst);
1396   SmallVector<Value, 4> composedOpOperands(subOperands);
1397   auto composedMap = builder.getMultiDimIdentityMap(composedOpOperands.size());
1398   fullyComposeAffineMapAndOperands(&composedMap, &composedOpOperands);
1399 
1400   // Create an affine.apply for each of the map results.
1401   sliceOps->reserve(composedMap.getNumResults());
1402   for (auto resultExpr : composedMap.getResults()) {
1403     auto singleResMap = AffineMap::get(composedMap.getNumDims(),
1404                                        composedMap.getNumSymbols(), resultExpr);
1405     sliceOps->push_back(builder.create<AffineApplyOp>(
1406         opInst->getLoc(), singleResMap, composedOpOperands));
1407   }
1408 
1409   // Construct the new operands that include the results from the composed
1410   // affine apply op above instead of existing ones (subOperands). So, they
1411   // differ from opInst's operands only for those operands in 'subOperands', for
1412   // which they will be replaced by the corresponding one from 'sliceOps'.
1413   SmallVector<Value, 4> newOperands(opInst->getOperands());
1414   for (unsigned i = 0, e = newOperands.size(); i < e; i++) {
1415     // Replace the subOperands from among the new operands.
1416     unsigned j, f;
1417     for (j = 0, f = subOperands.size(); j < f; j++) {
1418       if (newOperands[i] == subOperands[j])
1419         break;
1420     }
1421     if (j < subOperands.size()) {
1422       newOperands[i] = (*sliceOps)[j];
1423     }
1424   }
1425   for (unsigned idx = 0, e = newOperands.size(); idx < e; idx++) {
1426     opInst->setOperand(idx, newOperands[idx]);
1427   }
1428 }
1429 
1430 /// Enum to set patterns of affine expr in tiled-layout map.
1431 /// TileFloorDiv: <dim expr> div <tile size>
1432 /// TileMod: <dim expr> mod <tile size>
1433 /// TileNone: None of the above
1434 /// Example:
1435 /// #tiled_2d_128x256 = affine_map<(d0, d1)
1436 ///            -> (d0 div 128, d1 div 256, d0 mod 128, d1 mod 256)>
1437 /// "d0 div 128" and "d1 div 256" ==> TileFloorDiv
1438 /// "d0 mod 128" and "d1 mod 256" ==> TileMod
1439 enum TileExprPattern { TileFloorDiv, TileMod, TileNone };
1440 
1441 /// Check if `map` is a tiled layout. In the tiled layout, specific k dimensions
1442 /// being floordiv'ed by respective tile sizes appeare in a mod with the same
1443 /// tile sizes, and no other expression involves those k dimensions. This
1444 /// function stores a vector of tuples (`tileSizePos`) including AffineExpr for
1445 /// tile size, positions of corresponding `floordiv` and `mod`. If it is not a
1446 /// tiled layout, an empty vector is returned.
1447 static LogicalResult getTileSizePos(
1448     AffineMap map,
1449     SmallVectorImpl<std::tuple<AffineExpr, unsigned, unsigned>> &tileSizePos) {
1450   // Create `floordivExprs` which is a vector of tuples including LHS and RHS of
1451   // `floordiv` and its position in `map` output.
1452   // Example: #tiled_2d_128x256 = affine_map<(d0, d1)
1453   //                -> (d0 div 128, d1 div 256, d0 mod 128, d1 mod 256)>
1454   // In this example, `floordivExprs` includes {d0, 128, 0} and {d1, 256, 1}.
1455   SmallVector<std::tuple<AffineExpr, AffineExpr, unsigned>, 4> floordivExprs;
1456   unsigned pos = 0;
1457   for (AffineExpr expr : map.getResults()) {
1458     if (expr.getKind() == AffineExprKind::FloorDiv) {
1459       AffineBinaryOpExpr binaryExpr = expr.cast<AffineBinaryOpExpr>();
1460       if (binaryExpr.getRHS().isa<AffineConstantExpr>())
1461         floordivExprs.emplace_back(
1462             std::make_tuple(binaryExpr.getLHS(), binaryExpr.getRHS(), pos));
1463     }
1464     pos++;
1465   }
1466   // Not tiled layout if `floordivExprs` is empty.
1467   if (floordivExprs.empty()) {
1468     tileSizePos = SmallVector<std::tuple<AffineExpr, unsigned, unsigned>>{};
1469     return success();
1470   }
1471 
1472   // Check if LHS of `floordiv` is used in LHS of `mod`. If not used, `map` is
1473   // not tiled layout.
1474   for (std::tuple<AffineExpr, AffineExpr, unsigned> fexpr : floordivExprs) {
1475     AffineExpr floordivExprLHS = std::get<0>(fexpr);
1476     AffineExpr floordivExprRHS = std::get<1>(fexpr);
1477     unsigned floordivPos = std::get<2>(fexpr);
1478 
1479     // Walk affinexpr of `map` output except `fexpr`, and check if LHS and RHS
1480     // of `fexpr` are used in LHS and RHS of `mod`. If LHS of `fexpr` is used
1481     // other expr, the map is not tiled layout. Example of non tiled layout:
1482     //   affine_map<(d0, d1, d2) -> (d0, d1, d2 floordiv 256, d2 floordiv 256)>
1483     //   affine_map<(d0, d1, d2) -> (d0, d1, d2 floordiv 256, d2 mod 128)>
1484     //   affine_map<(d0, d1, d2) -> (d0, d1, d2 floordiv 256, d2 mod 256, d2 mod
1485     //   256)>
1486     bool found = false;
1487     pos = 0;
1488     for (AffineExpr expr : map.getResults()) {
1489       bool notTiled = false;
1490       if (pos != floordivPos) {
1491         expr.walk([&](AffineExpr e) {
1492           if (e == floordivExprLHS) {
1493             if (expr.getKind() == AffineExprKind::Mod) {
1494               AffineBinaryOpExpr binaryExpr = expr.cast<AffineBinaryOpExpr>();
1495               // If LHS and RHS of `mod` are the same with those of floordiv.
1496               if (floordivExprLHS == binaryExpr.getLHS() &&
1497                   floordivExprRHS == binaryExpr.getRHS()) {
1498                 // Save tile size (RHS of `mod`), and position of `floordiv` and
1499                 // `mod` if same expr with `mod` is not found yet.
1500                 if (!found) {
1501                   tileSizePos.emplace_back(
1502                       std::make_tuple(binaryExpr.getRHS(), floordivPos, pos));
1503                   found = true;
1504                 } else {
1505                   // Non tiled layout: Have multilpe `mod` with the same LHS.
1506                   // eg. affine_map<(d0, d1, d2) -> (d0, d1, d2 floordiv 256, d2
1507                   // mod 256, d2 mod 256)>
1508                   notTiled = true;
1509                 }
1510               } else {
1511                 // Non tiled layout: RHS of `mod` is different from `floordiv`.
1512                 // eg. affine_map<(d0, d1, d2) -> (d0, d1, d2 floordiv 256, d2
1513                 // mod 128)>
1514                 notTiled = true;
1515               }
1516             } else {
1517               // Non tiled layout: LHS is the same, but not `mod`.
1518               // eg. affine_map<(d0, d1, d2) -> (d0, d1, d2 floordiv 256, d2
1519               // floordiv 256)>
1520               notTiled = true;
1521             }
1522           }
1523         });
1524       }
1525       if (notTiled) {
1526         tileSizePos = SmallVector<std::tuple<AffineExpr, unsigned, unsigned>>{};
1527         return success();
1528       }
1529       pos++;
1530     }
1531   }
1532   return success();
1533 }
1534 
1535 /// Check if `dim` dimension of memrefType with `layoutMap` becomes dynamic
1536 /// after normalization. Dimensions that include dynamic dimensions in the map
1537 /// output will become dynamic dimensions. Return true if `dim` is dynamic
1538 /// dimension.
1539 ///
1540 /// Example:
1541 /// #map0 = affine_map<(d0, d1) -> (d0, d1 floordiv 32, d1 mod 32)>
1542 ///
1543 /// If d1 is dynamic dimension, 2nd and 3rd dimension of map output are dynamic.
1544 /// memref<4x?xf32, #map0>  ==>  memref<4x?x?xf32>
1545 static bool
1546 isNormalizedMemRefDynamicDim(unsigned dim, AffineMap layoutMap,
1547                              SmallVectorImpl<unsigned> &inMemrefTypeDynDims,
1548                              MLIRContext *context) {
1549   bool isDynamicDim = false;
1550   AffineExpr expr = layoutMap.getResults()[dim];
1551   // Check if affine expr of the dimension includes dynamic dimension of input
1552   // memrefType.
1553   expr.walk([&inMemrefTypeDynDims, &isDynamicDim, &context](AffineExpr e) {
1554     if (e.isa<AffineDimExpr>()) {
1555       for (unsigned dm : inMemrefTypeDynDims) {
1556         if (e == getAffineDimExpr(dm, context)) {
1557           isDynamicDim = true;
1558         }
1559       }
1560     }
1561   });
1562   return isDynamicDim;
1563 }
1564 
1565 /// Create affine expr to calculate dimension size for a tiled-layout map.
1566 static AffineExpr createDimSizeExprForTiledLayout(AffineExpr oldMapOutput,
1567                                                   TileExprPattern pat) {
1568   // Create map output for the patterns.
1569   // "floordiv <tile size>" ==> "ceildiv <tile size>"
1570   // "mod <tile size>" ==> "<tile size>"
1571   AffineExpr newMapOutput;
1572   AffineBinaryOpExpr binaryExpr = nullptr;
1573   switch (pat) {
1574   case TileExprPattern::TileMod:
1575     binaryExpr = oldMapOutput.cast<AffineBinaryOpExpr>();
1576     newMapOutput = binaryExpr.getRHS();
1577     break;
1578   case TileExprPattern::TileFloorDiv:
1579     binaryExpr = oldMapOutput.cast<AffineBinaryOpExpr>();
1580     newMapOutput = getAffineBinaryOpExpr(
1581         AffineExprKind::CeilDiv, binaryExpr.getLHS(), binaryExpr.getRHS());
1582     break;
1583   default:
1584     newMapOutput = oldMapOutput;
1585   }
1586   return newMapOutput;
1587 }
1588 
1589 /// Create new maps to calculate each dimension size of `newMemRefType`, and
1590 /// create `newDynamicSizes` from them by using AffineApplyOp.
1591 ///
1592 /// Steps for normalizing dynamic memrefs for a tiled layout map
1593 /// Example:
1594 ///    #map0 = affine_map<(d0, d1) -> (d0, d1 floordiv 32, d1 mod 32)>
1595 ///    %0 = dim %arg0, %c1 :memref<4x?xf32>
1596 ///    %1 = alloc(%0) : memref<4x?xf32, #map0>
1597 ///
1598 /// (Before this function)
1599 /// 1. Check if `map`(#map0) is a tiled layout using `getTileSizePos()`. Only
1600 /// single layout map is supported.
1601 ///
1602 /// 2. Create normalized memrefType using `isNormalizedMemRefDynamicDim()`. It
1603 /// is memref<4x?x?xf32> in the above example.
1604 ///
1605 /// (In this function)
1606 /// 3. Create new maps to calculate each dimension of the normalized memrefType
1607 /// using `createDimSizeExprForTiledLayout()`. In the tiled layout, the
1608 /// dimension size can be calculated by replacing "floordiv <tile size>" with
1609 /// "ceildiv <tile size>" and "mod <tile size>" with "<tile size>".
1610 /// - New map in the above example
1611 ///   #map0 = affine_map<(d0, d1) -> (d0)>
1612 ///   #map1 = affine_map<(d0, d1) -> (d1 ceildiv 32)>
1613 ///   #map2 = affine_map<(d0, d1) -> (32)>
1614 ///
1615 /// 4. Create AffineApplyOp to apply the new maps. The output of AffineApplyOp
1616 /// is used in dynamicSizes of new AllocOp.
1617 ///   %0 = dim %arg0, %c1 : memref<4x?xf32>
1618 ///   %c4 = arith.constant 4 : index
1619 ///   %1 = affine.apply #map1(%c4, %0)
1620 ///   %2 = affine.apply #map2(%c4, %0)
1621 static void createNewDynamicSizes(MemRefType oldMemRefType,
1622                                   MemRefType newMemRefType, AffineMap map,
1623                                   memref::AllocOp *allocOp, OpBuilder b,
1624                                   SmallVectorImpl<Value> &newDynamicSizes) {
1625   // Create new input for AffineApplyOp.
1626   SmallVector<Value, 4> inAffineApply;
1627   ArrayRef<int64_t> oldMemRefShape = oldMemRefType.getShape();
1628   unsigned dynIdx = 0;
1629   for (unsigned d = 0; d < oldMemRefType.getRank(); ++d) {
1630     if (oldMemRefShape[d] < 0) {
1631       // Use dynamicSizes of allocOp for dynamic dimension.
1632       inAffineApply.emplace_back(allocOp->dynamicSizes()[dynIdx]);
1633       dynIdx++;
1634     } else {
1635       // Create ConstantOp for static dimension.
1636       Attribute constantAttr =
1637           b.getIntegerAttr(b.getIndexType(), oldMemRefShape[d]);
1638       inAffineApply.emplace_back(
1639           b.create<arith::ConstantOp>(allocOp->getLoc(), constantAttr));
1640     }
1641   }
1642 
1643   // Create new map to calculate each dimension size of new memref for each
1644   // original map output. Only for dynamic dimesion of `newMemRefType`.
1645   unsigned newDimIdx = 0;
1646   ArrayRef<int64_t> newMemRefShape = newMemRefType.getShape();
1647   SmallVector<std::tuple<AffineExpr, unsigned, unsigned>> tileSizePos;
1648   (void)getTileSizePos(map, tileSizePos);
1649   for (AffineExpr expr : map.getResults()) {
1650     if (newMemRefShape[newDimIdx] < 0) {
1651       // Create new maps to calculate each dimension size of new memref.
1652       enum TileExprPattern pat = TileExprPattern::TileNone;
1653       for (auto pos : tileSizePos) {
1654         if (newDimIdx == std::get<1>(pos))
1655           pat = TileExprPattern::TileFloorDiv;
1656         else if (newDimIdx == std::get<2>(pos))
1657           pat = TileExprPattern::TileMod;
1658       }
1659       AffineExpr newMapOutput = createDimSizeExprForTiledLayout(expr, pat);
1660       AffineMap newMap =
1661           AffineMap::get(map.getNumInputs(), map.getNumSymbols(), newMapOutput);
1662       Value affineApp =
1663           b.create<AffineApplyOp>(allocOp->getLoc(), newMap, inAffineApply);
1664       newDynamicSizes.emplace_back(affineApp);
1665     }
1666     newDimIdx++;
1667   }
1668 }
1669 
1670 // TODO: Currently works for static memrefs with a single layout map.
1671 LogicalResult mlir::normalizeMemRef(memref::AllocOp *allocOp) {
1672   MemRefType memrefType = allocOp->getType();
1673   OpBuilder b(*allocOp);
1674 
1675   // Fetch a new memref type after normalizing the old memref to have an
1676   // identity map layout.
1677   MemRefType newMemRefType =
1678       normalizeMemRefType(memrefType, b, allocOp->symbolOperands().size());
1679   if (newMemRefType == memrefType)
1680     // Either memrefType already had an identity map or the map couldn't be
1681     // transformed to an identity map.
1682     return failure();
1683 
1684   Value oldMemRef = allocOp->getResult();
1685 
1686   SmallVector<Value, 4> symbolOperands(allocOp->symbolOperands());
1687   AffineMap layoutMap = memrefType.getLayout().getAffineMap();
1688   memref::AllocOp newAlloc;
1689   // Check if `layoutMap` is a tiled layout. Only single layout map is
1690   // supported for normalizing dynamic memrefs.
1691   SmallVector<std::tuple<AffineExpr, unsigned, unsigned>> tileSizePos;
1692   (void)getTileSizePos(layoutMap, tileSizePos);
1693   if (newMemRefType.getNumDynamicDims() > 0 && !tileSizePos.empty()) {
1694     MemRefType oldMemRefType = oldMemRef.getType().cast<MemRefType>();
1695     SmallVector<Value, 4> newDynamicSizes;
1696     createNewDynamicSizes(oldMemRefType, newMemRefType, layoutMap, allocOp, b,
1697                           newDynamicSizes);
1698     // Add the new dynamic sizes in new AllocOp.
1699     newAlloc =
1700         b.create<memref::AllocOp>(allocOp->getLoc(), newMemRefType,
1701                                   newDynamicSizes, allocOp->alignmentAttr());
1702   } else {
1703     newAlloc = b.create<memref::AllocOp>(allocOp->getLoc(), newMemRefType,
1704                                          allocOp->alignmentAttr());
1705   }
1706   // Replace all uses of the old memref.
1707   if (failed(replaceAllMemRefUsesWith(oldMemRef, /*newMemRef=*/newAlloc,
1708                                       /*extraIndices=*/{},
1709                                       /*indexRemap=*/layoutMap,
1710                                       /*extraOperands=*/{},
1711                                       /*symbolOperands=*/symbolOperands,
1712                                       /*domOpFilter=*/nullptr,
1713                                       /*postDomOpFilter=*/nullptr,
1714                                       /*allowNonDereferencingOps=*/true))) {
1715     // If it failed (due to escapes for example), bail out.
1716     newAlloc.erase();
1717     return failure();
1718   }
1719   // Replace any uses of the original alloc op and erase it. All remaining uses
1720   // have to be dealloc's; RAMUW above would've failed otherwise.
1721   assert(llvm::all_of(oldMemRef.getUsers(), [](Operation *op) {
1722     return isa<memref::DeallocOp>(op);
1723   }));
1724   oldMemRef.replaceAllUsesWith(newAlloc);
1725   allocOp->erase();
1726   return success();
1727 }
1728 
1729 MemRefType mlir::normalizeMemRefType(MemRefType memrefType, OpBuilder b,
1730                                      unsigned numSymbolicOperands) {
1731   unsigned rank = memrefType.getRank();
1732   if (rank == 0)
1733     return memrefType;
1734 
1735   if (memrefType.getLayout().isIdentity()) {
1736     // Either no maps is associated with this memref or this memref has
1737     // a trivial (identity) map.
1738     return memrefType;
1739   }
1740   AffineMap layoutMap = memrefType.getLayout().getAffineMap();
1741 
1742   // We don't do any checks for one-to-one'ness; we assume that it is
1743   // one-to-one.
1744 
1745   // Normalize only static memrefs and dynamic memrefs with a tiled-layout map
1746   // for now.
1747   // TODO: Normalize the other types of dynamic memrefs.
1748   SmallVector<std::tuple<AffineExpr, unsigned, unsigned>> tileSizePos;
1749   (void)getTileSizePos(layoutMap, tileSizePos);
1750   if (memrefType.getNumDynamicDims() > 0 && tileSizePos.empty())
1751     return memrefType;
1752 
1753   // We have a single map that is not an identity map. Create a new memref
1754   // with the right shape and an identity layout map.
1755   ArrayRef<int64_t> shape = memrefType.getShape();
1756   // FlatAffineConstraint may later on use symbolicOperands.
1757   FlatAffineConstraints fac(rank, numSymbolicOperands);
1758   SmallVector<unsigned, 4> memrefTypeDynDims;
1759   for (unsigned d = 0; d < rank; ++d) {
1760     // Use constraint system only in static dimensions.
1761     if (shape[d] > 0) {
1762       fac.addBound(FlatAffineConstraints::LB, d, 0);
1763       fac.addBound(FlatAffineConstraints::UB, d, shape[d] - 1);
1764     } else {
1765       memrefTypeDynDims.emplace_back(d);
1766     }
1767   }
1768   // We compose this map with the original index (logical) space to derive
1769   // the upper bounds for the new index space.
1770   unsigned newRank = layoutMap.getNumResults();
1771   if (failed(fac.composeMatchingMap(layoutMap)))
1772     return memrefType;
1773   // TODO: Handle semi-affine maps.
1774   // Project out the old data dimensions.
1775   fac.projectOut(newRank, fac.getNumIds() - newRank - fac.getNumLocalIds());
1776   SmallVector<int64_t, 4> newShape(newRank);
1777   for (unsigned d = 0; d < newRank; ++d) {
1778     // Check if each dimension of normalized memrefType is dynamic.
1779     bool isDynDim = isNormalizedMemRefDynamicDim(
1780         d, layoutMap, memrefTypeDynDims, b.getContext());
1781     if (isDynDim) {
1782       newShape[d] = -1;
1783     } else {
1784       // The lower bound for the shape is always zero.
1785       auto ubConst = fac.getConstantBound(FlatAffineConstraints::UB, d);
1786       // For a static memref and an affine map with no symbols, this is
1787       // always bounded.
1788       assert(ubConst.hasValue() && "should always have an upper bound");
1789       if (ubConst.getValue() < 0)
1790         // This is due to an invalid map that maps to a negative space.
1791         return memrefType;
1792       // If dimension of new memrefType is dynamic, the value is -1.
1793       newShape[d] = ubConst.getValue() + 1;
1794     }
1795   }
1796 
1797   // Create the new memref type after trivializing the old layout map.
1798   MemRefType newMemRefType =
1799       MemRefType::Builder(memrefType)
1800           .setShape(newShape)
1801           .setLayout(AffineMapAttr::get(b.getMultiDimIdentityMap(newRank)));
1802 
1803   return newMemRefType;
1804 }
1805