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 ®ion : 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