1 //===- AffineOps.cpp - MLIR Affine Operations -----------------------------===// 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 #include "mlir/Dialect/Affine/IR/AffineOps.h" 10 #include "mlir/Dialect/Affine/IR/AffineValueMap.h" 11 #include "mlir/Dialect/StandardOps/IR/Ops.h" 12 #include "mlir/IR/Function.h" 13 #include "mlir/IR/IntegerSet.h" 14 #include "mlir/IR/Matchers.h" 15 #include "mlir/IR/OpImplementation.h" 16 #include "mlir/IR/PatternMatch.h" 17 #include "mlir/Transforms/InliningUtils.h" 18 #include "llvm/ADT/SetVector.h" 19 #include "llvm/ADT/SmallBitVector.h" 20 #include "llvm/Support/Debug.h" 21 22 using namespace mlir; 23 using llvm::dbgs; 24 25 #define DEBUG_TYPE "affine-analysis" 26 27 //===----------------------------------------------------------------------===// 28 // AffineDialect Interfaces 29 //===----------------------------------------------------------------------===// 30 31 namespace { 32 /// This class defines the interface for handling inlining with affine 33 /// operations. 34 struct AffineInlinerInterface : public DialectInlinerInterface { 35 using DialectInlinerInterface::DialectInlinerInterface; 36 37 //===--------------------------------------------------------------------===// 38 // Analysis Hooks 39 //===--------------------------------------------------------------------===// 40 41 /// Returns true if the given region 'src' can be inlined into the region 42 /// 'dest' that is attached to an operation registered to the current dialect. 43 bool isLegalToInline(Region *dest, Region *src, 44 BlockAndValueMapping &valueMapping) const final { 45 // Conservatively don't allow inlining into affine structures. 46 return false; 47 } 48 49 /// Returns true if the given operation 'op', that is registered to this 50 /// dialect, can be inlined into the given region, false otherwise. 51 bool isLegalToInline(Operation *op, Region *region, 52 BlockAndValueMapping &valueMapping) const final { 53 // Always allow inlining affine operations into the top-level region of a 54 // function. There are some edge cases when inlining *into* affine 55 // structures, but that is handled in the other 'isLegalToInline' hook 56 // above. 57 // TODO: We should be able to inline into other regions than functions. 58 return isa<FuncOp>(region->getParentOp()); 59 } 60 61 /// Affine regions should be analyzed recursively. 62 bool shouldAnalyzeRecursively(Operation *op) const final { return true; } 63 }; 64 } // end anonymous namespace 65 66 //===----------------------------------------------------------------------===// 67 // AffineDialect 68 //===----------------------------------------------------------------------===// 69 70 AffineDialect::AffineDialect(MLIRContext *context) 71 : Dialect(getDialectNamespace(), context) { 72 addOperations<AffineDmaStartOp, AffineDmaWaitOp, AffineLoadOp, AffineStoreOp, 73 #define GET_OP_LIST 74 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc" 75 >(); 76 addInterfaces<AffineInlinerInterface>(); 77 } 78 79 /// Materialize a single constant operation from a given attribute value with 80 /// the desired resultant type. 81 Operation *AffineDialect::materializeConstant(OpBuilder &builder, 82 Attribute value, Type type, 83 Location loc) { 84 return builder.create<ConstantOp>(loc, type, value); 85 } 86 87 /// A utility function to check if a given region is attached to a function. 88 static bool isFunctionRegion(Region *region) { 89 return llvm::isa<FuncOp>(region->getParentOp()); 90 } 91 92 /// A utility function to check if a value is defined at the top level of a 93 /// function. A value of index type defined at the top level is always a valid 94 /// symbol. 95 bool mlir::isTopLevelValue(Value value) { 96 if (auto arg = value.dyn_cast<BlockArgument>()) 97 return isFunctionRegion(arg.getOwner()->getParent()); 98 return isFunctionRegion(value.getDefiningOp()->getParentRegion()); 99 } 100 101 // Value can be used as a dimension id if it is valid as a symbol, or 102 // it is an induction variable, or it is a result of affine apply operation 103 // with dimension id arguments. 104 bool mlir::isValidDim(Value value) { 105 // The value must be an index type. 106 if (!value.getType().isIndex()) 107 return false; 108 109 if (auto *op = value.getDefiningOp()) { 110 // Top level operation or constant operation is ok. 111 if (isFunctionRegion(op->getParentRegion()) || isa<ConstantOp>(op)) 112 return true; 113 // Affine apply operation is ok if all of its operands are ok. 114 if (auto applyOp = dyn_cast<AffineApplyOp>(op)) 115 return applyOp.isValidDim(); 116 // The dim op is okay if its operand memref/tensor is defined at the top 117 // level. 118 if (auto dimOp = dyn_cast<DimOp>(op)) 119 return isTopLevelValue(dimOp.getOperand()); 120 return false; 121 } 122 // This value has to be a block argument of a FuncOp, an 'affine.for', or an 123 // 'affine.parallel'. 124 auto *parentOp = value.cast<BlockArgument>().getOwner()->getParentOp(); 125 return isa<FuncOp>(parentOp) || isa<AffineForOp>(parentOp) || 126 isa<AffineParallelOp>(parentOp); 127 } 128 129 /// Returns true if the 'index' dimension of the `memref` defined by 130 /// `memrefDefOp` is a statically shaped one or defined using a valid symbol. 131 template <typename AnyMemRefDefOp> 132 static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp, 133 unsigned index) { 134 auto memRefType = memrefDefOp.getType(); 135 // Statically shaped. 136 if (!memRefType.isDynamicDim(index)) 137 return true; 138 // Get the position of the dimension among dynamic dimensions; 139 unsigned dynamicDimPos = memRefType.getDynamicDimIndex(index); 140 return isValidSymbol( 141 *(memrefDefOp.getDynamicSizes().begin() + dynamicDimPos)); 142 } 143 144 /// Returns true if the result of the dim op is a valid symbol. 145 static bool isDimOpValidSymbol(DimOp dimOp) { 146 // The dim op is okay if its operand memref/tensor is defined at the top 147 // level. 148 if (isTopLevelValue(dimOp.getOperand())) 149 return true; 150 151 // The dim op is also okay if its operand memref/tensor is a view/subview 152 // whose corresponding size is a valid symbol. 153 unsigned index = dimOp.getIndex(); 154 if (auto viewOp = dyn_cast<ViewOp>(dimOp.getOperand().getDefiningOp())) 155 return isMemRefSizeValidSymbol<ViewOp>(viewOp, index); 156 if (auto subViewOp = dyn_cast<SubViewOp>(dimOp.getOperand().getDefiningOp())) 157 return isMemRefSizeValidSymbol<SubViewOp>(subViewOp, index); 158 if (auto allocOp = dyn_cast<AllocOp>(dimOp.getOperand().getDefiningOp())) 159 return isMemRefSizeValidSymbol<AllocOp>(allocOp, index); 160 return false; 161 } 162 163 // Value can be used as a symbol if it is a constant, or it is defined at 164 // the top level, or it is a result of affine apply operation with symbol 165 // arguments, or a result of the dim op on a memref satisfying certain 166 // constraints. 167 bool mlir::isValidSymbol(Value value) { 168 // The value must be an index type. 169 if (!value.getType().isIndex()) 170 return false; 171 172 if (auto *op = value.getDefiningOp()) { 173 // Top level operation or constant operation is ok. 174 if (isFunctionRegion(op->getParentRegion()) || isa<ConstantOp>(op)) 175 return true; 176 // Affine apply operation is ok if all of its operands are ok. 177 if (auto applyOp = dyn_cast<AffineApplyOp>(op)) 178 return applyOp.isValidSymbol(); 179 if (auto dimOp = dyn_cast<DimOp>(op)) { 180 return isDimOpValidSymbol(dimOp); 181 } 182 } 183 // Otherwise, check that the value is a top level value. 184 return isTopLevelValue(value); 185 } 186 187 // Returns true if 'value' is a valid index to an affine operation (e.g. 188 // affine.load, affine.store, affine.dma_start, affine.dma_wait). 189 // Returns false otherwise. 190 static bool isValidAffineIndexOperand(Value value) { 191 return isValidDim(value) || isValidSymbol(value); 192 } 193 194 /// Utility function to verify that a set of operands are valid dimension and 195 /// symbol identifiers. The operands should be laid out such that the dimension 196 /// operands are before the symbol operands. This function returns failure if 197 /// there was an invalid operand. An operation is provided to emit any necessary 198 /// errors. 199 template <typename OpTy> 200 static LogicalResult 201 verifyDimAndSymbolIdentifiers(OpTy &op, Operation::operand_range operands, 202 unsigned numDims) { 203 unsigned opIt = 0; 204 for (auto operand : operands) { 205 if (opIt++ < numDims) { 206 if (!isValidDim(operand)) 207 return op.emitOpError("operand cannot be used as a dimension id"); 208 } else if (!isValidSymbol(operand)) { 209 return op.emitOpError("operand cannot be used as a symbol"); 210 } 211 } 212 return success(); 213 } 214 215 //===----------------------------------------------------------------------===// 216 // AffineApplyOp 217 //===----------------------------------------------------------------------===// 218 219 AffineValueMap AffineApplyOp::getAffineValueMap() { 220 return AffineValueMap(getAffineMap(), getOperands(), getResult()); 221 } 222 223 static ParseResult parseAffineApplyOp(OpAsmParser &parser, 224 OperationState &result) { 225 auto &builder = parser.getBuilder(); 226 auto indexTy = builder.getIndexType(); 227 228 AffineMapAttr mapAttr; 229 unsigned numDims; 230 if (parser.parseAttribute(mapAttr, "map", result.attributes) || 231 parseDimAndSymbolList(parser, result.operands, numDims) || 232 parser.parseOptionalAttrDict(result.attributes)) 233 return failure(); 234 auto map = mapAttr.getValue(); 235 236 if (map.getNumDims() != numDims || 237 numDims + map.getNumSymbols() != result.operands.size()) { 238 return parser.emitError(parser.getNameLoc(), 239 "dimension or symbol index mismatch"); 240 } 241 242 result.types.append(map.getNumResults(), indexTy); 243 return success(); 244 } 245 246 static void print(OpAsmPrinter &p, AffineApplyOp op) { 247 p << AffineApplyOp::getOperationName() << " " << op.mapAttr(); 248 printDimAndSymbolList(op.operand_begin(), op.operand_end(), 249 op.getAffineMap().getNumDims(), p); 250 p.printOptionalAttrDict(op.getAttrs(), /*elidedAttrs=*/{"map"}); 251 } 252 253 static LogicalResult verify(AffineApplyOp op) { 254 // Check input and output dimensions match. 255 auto map = op.map(); 256 257 // Verify that operand count matches affine map dimension and symbol count. 258 if (op.getNumOperands() != map.getNumDims() + map.getNumSymbols()) 259 return op.emitOpError( 260 "operand count and affine map dimension and symbol count must match"); 261 262 // Verify that the map only produces one result. 263 if (map.getNumResults() != 1) 264 return op.emitOpError("mapping must produce one value"); 265 266 return success(); 267 } 268 269 // The result of the affine apply operation can be used as a dimension id if all 270 // its operands are valid dimension ids. 271 bool AffineApplyOp::isValidDim() { 272 return llvm::all_of(getOperands(), 273 [](Value op) { return mlir::isValidDim(op); }); 274 } 275 276 // The result of the affine apply operation can be used as a symbol if all its 277 // operands are symbols. 278 bool AffineApplyOp::isValidSymbol() { 279 return llvm::all_of(getOperands(), 280 [](Value op) { return mlir::isValidSymbol(op); }); 281 } 282 283 OpFoldResult AffineApplyOp::fold(ArrayRef<Attribute> operands) { 284 auto map = getAffineMap(); 285 286 // Fold dims and symbols to existing values. 287 auto expr = map.getResult(0); 288 if (auto dim = expr.dyn_cast<AffineDimExpr>()) 289 return getOperand(dim.getPosition()); 290 if (auto sym = expr.dyn_cast<AffineSymbolExpr>()) 291 return getOperand(map.getNumDims() + sym.getPosition()); 292 293 // Otherwise, default to folding the map. 294 SmallVector<Attribute, 1> result; 295 if (failed(map.constantFold(operands, result))) 296 return {}; 297 return result[0]; 298 } 299 300 AffineDimExpr AffineApplyNormalizer::renumberOneDim(Value v) { 301 DenseMap<Value, unsigned>::iterator iterPos; 302 bool inserted = false; 303 std::tie(iterPos, inserted) = 304 dimValueToPosition.insert(std::make_pair(v, dimValueToPosition.size())); 305 if (inserted) { 306 reorderedDims.push_back(v); 307 } 308 return getAffineDimExpr(iterPos->second, v.getContext()) 309 .cast<AffineDimExpr>(); 310 } 311 312 AffineMap AffineApplyNormalizer::renumber(const AffineApplyNormalizer &other) { 313 SmallVector<AffineExpr, 8> dimRemapping; 314 for (auto v : other.reorderedDims) { 315 auto kvp = other.dimValueToPosition.find(v); 316 if (dimRemapping.size() <= kvp->second) 317 dimRemapping.resize(kvp->second + 1); 318 dimRemapping[kvp->second] = renumberOneDim(kvp->first); 319 } 320 unsigned numSymbols = concatenatedSymbols.size(); 321 unsigned numOtherSymbols = other.concatenatedSymbols.size(); 322 SmallVector<AffineExpr, 8> symRemapping(numOtherSymbols); 323 for (unsigned idx = 0; idx < numOtherSymbols; ++idx) { 324 symRemapping[idx] = 325 getAffineSymbolExpr(idx + numSymbols, other.affineMap.getContext()); 326 } 327 concatenatedSymbols.insert(concatenatedSymbols.end(), 328 other.concatenatedSymbols.begin(), 329 other.concatenatedSymbols.end()); 330 auto map = other.affineMap; 331 return map.replaceDimsAndSymbols(dimRemapping, symRemapping, 332 reorderedDims.size(), 333 concatenatedSymbols.size()); 334 } 335 336 // Gather the positions of the operands that are produced by an AffineApplyOp. 337 static llvm::SetVector<unsigned> 338 indicesFromAffineApplyOp(ArrayRef<Value> operands) { 339 llvm::SetVector<unsigned> res; 340 for (auto en : llvm::enumerate(operands)) 341 if (isa_and_nonnull<AffineApplyOp>(en.value().getDefiningOp())) 342 res.insert(en.index()); 343 return res; 344 } 345 346 // Support the special case of a symbol coming from an AffineApplyOp that needs 347 // to be composed into the current AffineApplyOp. 348 // This case is handled by rewriting all such symbols into dims for the purpose 349 // of allowing mathematical AffineMap composition. 350 // Returns an AffineMap where symbols that come from an AffineApplyOp have been 351 // rewritten as dims and are ordered after the original dims. 352 // TODO(andydavis,ntv): This promotion makes AffineMap lose track of which 353 // symbols are represented as dims. This loss is static but can still be 354 // recovered dynamically (with `isValidSymbol`). Still this is annoying for the 355 // semi-affine map case. A dynamic canonicalization of all dims that are valid 356 // symbols (a.k.a `canonicalizePromotedSymbols`) into symbols helps and even 357 // results in better simplifications and foldings. But we should evaluate 358 // whether this behavior is what we really want after using more. 359 static AffineMap promoteComposedSymbolsAsDims(AffineMap map, 360 ArrayRef<Value> symbols) { 361 if (symbols.empty()) { 362 return map; 363 } 364 365 // Sanity check on symbols. 366 for (auto sym : symbols) { 367 assert(isValidSymbol(sym) && "Expected only valid symbols"); 368 (void)sym; 369 } 370 371 // Extract the symbol positions that come from an AffineApplyOp and 372 // needs to be rewritten as dims. 373 auto symPositions = indicesFromAffineApplyOp(symbols); 374 if (symPositions.empty()) { 375 return map; 376 } 377 378 // Create the new map by replacing each symbol at pos by the next new dim. 379 unsigned numDims = map.getNumDims(); 380 unsigned numSymbols = map.getNumSymbols(); 381 unsigned numNewDims = 0; 382 unsigned numNewSymbols = 0; 383 SmallVector<AffineExpr, 8> symReplacements(numSymbols); 384 for (unsigned i = 0; i < numSymbols; ++i) { 385 symReplacements[i] = 386 symPositions.count(i) > 0 387 ? getAffineDimExpr(numDims + numNewDims++, map.getContext()) 388 : getAffineSymbolExpr(numNewSymbols++, map.getContext()); 389 } 390 assert(numSymbols >= numNewDims); 391 AffineMap newMap = map.replaceDimsAndSymbols( 392 {}, symReplacements, numDims + numNewDims, numNewSymbols); 393 394 return newMap; 395 } 396 397 /// The AffineNormalizer composes AffineApplyOp recursively. Its purpose is to 398 /// keep a correspondence between the mathematical `map` and the `operands` of 399 /// a given AffineApplyOp. This correspondence is maintained by iterating over 400 /// the operands and forming an `auxiliaryMap` that can be composed 401 /// mathematically with `map`. To keep this correspondence in cases where 402 /// symbols are produced by affine.apply operations, we perform a local rewrite 403 /// of symbols as dims. 404 /// 405 /// Rationale for locally rewriting symbols as dims: 406 /// ================================================ 407 /// The mathematical composition of AffineMap must always concatenate symbols 408 /// because it does not have enough information to do otherwise. For example, 409 /// composing `(d0)[s0] -> (d0 + s0)` with itself must produce 410 /// `(d0)[s0, s1] -> (d0 + s0 + s1)`. 411 /// 412 /// The result is only equivalent to `(d0)[s0] -> (d0 + 2 * s0)` when 413 /// applied to the same mlir::Value for both s0 and s1. 414 /// As a consequence mathematical composition of AffineMap always concatenates 415 /// symbols. 416 /// 417 /// When AffineMaps are used in AffineApplyOp however, they may specify 418 /// composition via symbols, which is ambiguous mathematically. This corner case 419 /// is handled by locally rewriting such symbols that come from AffineApplyOp 420 /// into dims and composing through dims. 421 /// TODO(andydavis, ntv): Composition via symbols comes at a significant code 422 /// complexity. Alternatively we should investigate whether we want to 423 /// explicitly disallow symbols coming from affine.apply and instead force the 424 /// user to compose symbols beforehand. The annoyances may be small (i.e. 1 or 2 425 /// extra API calls for such uses, which haven't popped up until now) and the 426 /// benefit potentially big: simpler and more maintainable code for a 427 /// non-trivial, recursive, procedure. 428 AffineApplyNormalizer::AffineApplyNormalizer(AffineMap map, 429 ArrayRef<Value> operands) 430 : AffineApplyNormalizer() { 431 static_assert(kMaxAffineApplyDepth > 0, "kMaxAffineApplyDepth must be > 0"); 432 assert(map.getNumInputs() == operands.size() && 433 "number of operands does not match the number of map inputs"); 434 435 LLVM_DEBUG(map.print(dbgs() << "\nInput map: ")); 436 437 // Promote symbols that come from an AffineApplyOp to dims by rewriting the 438 // map to always refer to: 439 // (dims, symbols coming from AffineApplyOp, other symbols). 440 // The order of operands can remain unchanged. 441 // This is a simplification that relies on 2 ordering properties: 442 // 1. rewritten symbols always appear after the original dims in the map; 443 // 2. operands are traversed in order and either dispatched to: 444 // a. auxiliaryExprs (dims and symbols rewritten as dims); 445 // b. concatenatedSymbols (all other symbols) 446 // This allows operand order to remain unchanged. 447 unsigned numDimsBeforeRewrite = map.getNumDims(); 448 map = promoteComposedSymbolsAsDims(map, 449 operands.take_back(map.getNumSymbols())); 450 451 LLVM_DEBUG(map.print(dbgs() << "\nRewritten map: ")); 452 453 SmallVector<AffineExpr, 8> auxiliaryExprs; 454 bool furtherCompose = (affineApplyDepth() <= kMaxAffineApplyDepth); 455 // We fully spell out the 2 cases below. In this particular instance a little 456 // code duplication greatly improves readability. 457 // Note that the first branch would disappear if we only supported full 458 // composition (i.e. infinite kMaxAffineApplyDepth). 459 if (!furtherCompose) { 460 // 1. Only dispatch dims or symbols. 461 for (auto en : llvm::enumerate(operands)) { 462 auto t = en.value(); 463 assert(t.getType().isIndex()); 464 bool isDim = (en.index() < map.getNumDims()); 465 if (isDim) { 466 // a. The mathematical composition of AffineMap composes dims. 467 auxiliaryExprs.push_back(renumberOneDim(t)); 468 } else { 469 // b. The mathematical composition of AffineMap concatenates symbols. 470 // We do the same for symbol operands. 471 concatenatedSymbols.push_back(t); 472 } 473 } 474 } else { 475 assert(numDimsBeforeRewrite <= operands.size()); 476 // 2. Compose AffineApplyOps and dispatch dims or symbols. 477 for (unsigned i = 0, e = operands.size(); i < e; ++i) { 478 auto t = operands[i]; 479 auto affineApply = dyn_cast_or_null<AffineApplyOp>(t.getDefiningOp()); 480 if (affineApply) { 481 // a. Compose affine.apply operations. 482 LLVM_DEBUG(affineApply.getOperation()->print( 483 dbgs() << "\nCompose AffineApplyOp recursively: ")); 484 AffineMap affineApplyMap = affineApply.getAffineMap(); 485 SmallVector<Value, 8> affineApplyOperands( 486 affineApply.getOperands().begin(), affineApply.getOperands().end()); 487 AffineApplyNormalizer normalizer(affineApplyMap, affineApplyOperands); 488 489 LLVM_DEBUG(normalizer.affineMap.print( 490 dbgs() << "\nRenumber into current normalizer: ")); 491 492 auto renumberedMap = renumber(normalizer); 493 494 LLVM_DEBUG( 495 renumberedMap.print(dbgs() << "\nRecursive composition yields: ")); 496 497 auxiliaryExprs.push_back(renumberedMap.getResult(0)); 498 } else { 499 if (i < numDimsBeforeRewrite) { 500 // b. The mathematical composition of AffineMap composes dims. 501 auxiliaryExprs.push_back(renumberOneDim(t)); 502 } else { 503 // c. The mathematical composition of AffineMap concatenates symbols. 504 // Note that the map composition will put symbols already present 505 // in the map before any symbols coming from the auxiliary map, so 506 // we insert them before any symbols that are due to renumbering, 507 // and after the proper symbols we have seen already. 508 concatenatedSymbols.insert( 509 std::next(concatenatedSymbols.begin(), numProperSymbols++), t); 510 } 511 } 512 } 513 } 514 515 // Early exit if `map` is already composed. 516 if (auxiliaryExprs.empty()) { 517 affineMap = map; 518 return; 519 } 520 521 assert(concatenatedSymbols.size() >= map.getNumSymbols() && 522 "Unexpected number of concatenated symbols"); 523 auto numDims = dimValueToPosition.size(); 524 auto numSymbols = concatenatedSymbols.size() - map.getNumSymbols(); 525 auto auxiliaryMap = 526 AffineMap::get(numDims, numSymbols, auxiliaryExprs, map.getContext()); 527 528 LLVM_DEBUG(map.print(dbgs() << "\nCompose map: ")); 529 LLVM_DEBUG(auxiliaryMap.print(dbgs() << "\nWith map: ")); 530 LLVM_DEBUG(map.compose(auxiliaryMap).print(dbgs() << "\nResult: ")); 531 532 // TODO(andydavis,ntv): Disabling simplification results in major speed gains. 533 // Another option is to cache the results as it is expected a lot of redundant 534 // work is performed in practice. 535 affineMap = simplifyAffineMap(map.compose(auxiliaryMap)); 536 537 LLVM_DEBUG(affineMap.print(dbgs() << "\nSimplified result: ")); 538 LLVM_DEBUG(dbgs() << "\n"); 539 } 540 541 void AffineApplyNormalizer::normalize(AffineMap *otherMap, 542 SmallVectorImpl<Value> *otherOperands) { 543 AffineApplyNormalizer other(*otherMap, *otherOperands); 544 *otherMap = renumber(other); 545 546 otherOperands->reserve(reorderedDims.size() + concatenatedSymbols.size()); 547 otherOperands->assign(reorderedDims.begin(), reorderedDims.end()); 548 otherOperands->append(concatenatedSymbols.begin(), concatenatedSymbols.end()); 549 } 550 551 /// Implements `map` and `operands` composition and simplification to support 552 /// `makeComposedAffineApply`. This can be called to achieve the same effects 553 /// on `map` and `operands` without creating an AffineApplyOp that needs to be 554 /// immediately deleted. 555 static void composeAffineMapAndOperands(AffineMap *map, 556 SmallVectorImpl<Value> *operands) { 557 AffineApplyNormalizer normalizer(*map, *operands); 558 auto normalizedMap = normalizer.getAffineMap(); 559 auto normalizedOperands = normalizer.getOperands(); 560 canonicalizeMapAndOperands(&normalizedMap, &normalizedOperands); 561 *map = normalizedMap; 562 *operands = normalizedOperands; 563 assert(*map); 564 } 565 566 void mlir::fullyComposeAffineMapAndOperands(AffineMap *map, 567 SmallVectorImpl<Value> *operands) { 568 while (llvm::any_of(*operands, [](Value v) { 569 return isa_and_nonnull<AffineApplyOp>(v.getDefiningOp()); 570 })) { 571 composeAffineMapAndOperands(map, operands); 572 } 573 } 574 575 AffineApplyOp mlir::makeComposedAffineApply(OpBuilder &b, Location loc, 576 AffineMap map, 577 ArrayRef<Value> operands) { 578 AffineMap normalizedMap = map; 579 SmallVector<Value, 8> normalizedOperands(operands.begin(), operands.end()); 580 composeAffineMapAndOperands(&normalizedMap, &normalizedOperands); 581 assert(normalizedMap); 582 return b.create<AffineApplyOp>(loc, normalizedMap, normalizedOperands); 583 } 584 585 // A symbol may appear as a dim in affine.apply operations. This function 586 // canonicalizes dims that are valid symbols into actual symbols. 587 template <class MapOrSet> 588 static void canonicalizePromotedSymbols(MapOrSet *mapOrSet, 589 SmallVectorImpl<Value> *operands) { 590 if (!mapOrSet || operands->empty()) 591 return; 592 593 assert(mapOrSet->getNumInputs() == operands->size() && 594 "map/set inputs must match number of operands"); 595 596 auto *context = mapOrSet->getContext(); 597 SmallVector<Value, 8> resultOperands; 598 resultOperands.reserve(operands->size()); 599 SmallVector<Value, 8> remappedSymbols; 600 remappedSymbols.reserve(operands->size()); 601 unsigned nextDim = 0; 602 unsigned nextSym = 0; 603 unsigned oldNumSyms = mapOrSet->getNumSymbols(); 604 SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims()); 605 for (unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) { 606 if (i < mapOrSet->getNumDims()) { 607 if (isValidSymbol((*operands)[i])) { 608 // This is a valid symbol that appears as a dim, canonicalize it. 609 dimRemapping[i] = getAffineSymbolExpr(oldNumSyms + nextSym++, context); 610 remappedSymbols.push_back((*operands)[i]); 611 } else { 612 dimRemapping[i] = getAffineDimExpr(nextDim++, context); 613 resultOperands.push_back((*operands)[i]); 614 } 615 } else { 616 resultOperands.push_back((*operands)[i]); 617 } 618 } 619 620 resultOperands.append(remappedSymbols.begin(), remappedSymbols.end()); 621 *operands = resultOperands; 622 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, {}, nextDim, 623 oldNumSyms + nextSym); 624 625 assert(mapOrSet->getNumInputs() == operands->size() && 626 "map/set inputs must match number of operands"); 627 } 628 629 // Works for either an affine map or an integer set. 630 template <class MapOrSet> 631 static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet, 632 SmallVectorImpl<Value> *operands) { 633 static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value, 634 "Argument must be either of AffineMap or IntegerSet type"); 635 636 if (!mapOrSet || operands->empty()) 637 return; 638 639 assert(mapOrSet->getNumInputs() == operands->size() && 640 "map/set inputs must match number of operands"); 641 642 canonicalizePromotedSymbols<MapOrSet>(mapOrSet, operands); 643 644 // Check to see what dims are used. 645 llvm::SmallBitVector usedDims(mapOrSet->getNumDims()); 646 llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols()); 647 mapOrSet->walkExprs([&](AffineExpr expr) { 648 if (auto dimExpr = expr.dyn_cast<AffineDimExpr>()) 649 usedDims[dimExpr.getPosition()] = true; 650 else if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) 651 usedSyms[symExpr.getPosition()] = true; 652 }); 653 654 auto *context = mapOrSet->getContext(); 655 656 SmallVector<Value, 8> resultOperands; 657 resultOperands.reserve(operands->size()); 658 659 llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims; 660 SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims()); 661 unsigned nextDim = 0; 662 for (unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) { 663 if (usedDims[i]) { 664 // Remap dim positions for duplicate operands. 665 auto it = seenDims.find((*operands)[i]); 666 if (it == seenDims.end()) { 667 dimRemapping[i] = getAffineDimExpr(nextDim++, context); 668 resultOperands.push_back((*operands)[i]); 669 seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i])); 670 } else { 671 dimRemapping[i] = it->second; 672 } 673 } 674 } 675 llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols; 676 SmallVector<AffineExpr, 8> symRemapping(mapOrSet->getNumSymbols()); 677 unsigned nextSym = 0; 678 for (unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) { 679 if (!usedSyms[i]) 680 continue; 681 // Handle constant operands (only needed for symbolic operands since 682 // constant operands in dimensional positions would have already been 683 // promoted to symbolic positions above). 684 IntegerAttr operandCst; 685 if (matchPattern((*operands)[i + mapOrSet->getNumDims()], 686 m_Constant(&operandCst))) { 687 symRemapping[i] = 688 getAffineConstantExpr(operandCst.getValue().getSExtValue(), context); 689 continue; 690 } 691 // Remap symbol positions for duplicate operands. 692 auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]); 693 if (it == seenSymbols.end()) { 694 symRemapping[i] = getAffineSymbolExpr(nextSym++, context); 695 resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]); 696 seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()], 697 symRemapping[i])); 698 } else { 699 symRemapping[i] = it->second; 700 } 701 } 702 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping, 703 nextDim, nextSym); 704 *operands = resultOperands; 705 } 706 707 void mlir::canonicalizeMapAndOperands(AffineMap *map, 708 SmallVectorImpl<Value> *operands) { 709 canonicalizeMapOrSetAndOperands<AffineMap>(map, operands); 710 } 711 712 void mlir::canonicalizeSetAndOperands(IntegerSet *set, 713 SmallVectorImpl<Value> *operands) { 714 canonicalizeMapOrSetAndOperands<IntegerSet>(set, operands); 715 } 716 717 namespace { 718 /// Simplify AffineApply, AffineLoad, and AffineStore operations by composing 719 /// maps that supply results into them. 720 /// 721 template <typename AffineOpTy> 722 struct SimplifyAffineOp : public OpRewritePattern<AffineOpTy> { 723 using OpRewritePattern<AffineOpTy>::OpRewritePattern; 724 725 /// Replace the affine op with another instance of it with the supplied 726 /// map and mapOperands. 727 void replaceAffineOp(PatternRewriter &rewriter, AffineOpTy affineOp, 728 AffineMap map, ArrayRef<Value> mapOperands) const; 729 730 LogicalResult matchAndRewrite(AffineOpTy affineOp, 731 PatternRewriter &rewriter) const override { 732 static_assert(llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp, 733 AffineStoreOp, AffineApplyOp, AffineMinOp, 734 AffineMaxOp>::value, 735 "affine load/store/apply/prefetch/min/max op expected"); 736 auto map = affineOp.getAffineMap(); 737 AffineMap oldMap = map; 738 auto oldOperands = affineOp.getMapOperands(); 739 SmallVector<Value, 8> resultOperands(oldOperands); 740 composeAffineMapAndOperands(&map, &resultOperands); 741 if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(), 742 resultOperands.begin())) 743 return failure(); 744 745 replaceAffineOp(rewriter, affineOp, map, resultOperands); 746 return success(); 747 } 748 }; 749 750 // Specialize the template to account for the different build signatures for 751 // affine load, store, and apply ops. 752 template <> 753 void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp( 754 PatternRewriter &rewriter, AffineLoadOp load, AffineMap map, 755 ArrayRef<Value> mapOperands) const { 756 rewriter.replaceOpWithNewOp<AffineLoadOp>(load, load.getMemRef(), map, 757 mapOperands); 758 } 759 template <> 760 void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp( 761 PatternRewriter &rewriter, AffinePrefetchOp prefetch, AffineMap map, 762 ArrayRef<Value> mapOperands) const { 763 rewriter.replaceOpWithNewOp<AffinePrefetchOp>( 764 prefetch, prefetch.memref(), map, mapOperands, 765 prefetch.localityHint().getZExtValue(), prefetch.isWrite(), 766 prefetch.isDataCache()); 767 } 768 template <> 769 void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp( 770 PatternRewriter &rewriter, AffineStoreOp store, AffineMap map, 771 ArrayRef<Value> mapOperands) const { 772 rewriter.replaceOpWithNewOp<AffineStoreOp>( 773 store, store.getValueToStore(), store.getMemRef(), map, mapOperands); 774 } 775 776 // Generic version for ops that don't have extra operands. 777 template <typename AffineOpTy> 778 void SimplifyAffineOp<AffineOpTy>::replaceAffineOp( 779 PatternRewriter &rewriter, AffineOpTy op, AffineMap map, 780 ArrayRef<Value> mapOperands) const { 781 rewriter.replaceOpWithNewOp<AffineOpTy>(op, map, mapOperands); 782 } 783 } // end anonymous namespace. 784 785 void AffineApplyOp::getCanonicalizationPatterns( 786 OwningRewritePatternList &results, MLIRContext *context) { 787 results.insert<SimplifyAffineOp<AffineApplyOp>>(context); 788 } 789 790 //===----------------------------------------------------------------------===// 791 // Common canonicalization pattern support logic 792 //===----------------------------------------------------------------------===// 793 794 /// This is a common class used for patterns of the form 795 /// "someop(memrefcast) -> someop". It folds the source of any memref_cast 796 /// into the root operation directly. 797 static LogicalResult foldMemRefCast(Operation *op) { 798 bool folded = false; 799 for (OpOperand &operand : op->getOpOperands()) { 800 auto cast = dyn_cast_or_null<MemRefCastOp>(operand.get().getDefiningOp()); 801 if (cast && !cast.getOperand().getType().isa<UnrankedMemRefType>()) { 802 operand.set(cast.getOperand()); 803 folded = true; 804 } 805 } 806 return success(folded); 807 } 808 809 //===----------------------------------------------------------------------===// 810 // AffineDmaStartOp 811 //===----------------------------------------------------------------------===// 812 813 // TODO(b/133776335) Check that map operands are loop IVs or symbols. 814 void AffineDmaStartOp::build(Builder *builder, OperationState &result, 815 Value srcMemRef, AffineMap srcMap, 816 ValueRange srcIndices, Value destMemRef, 817 AffineMap dstMap, ValueRange destIndices, 818 Value tagMemRef, AffineMap tagMap, 819 ValueRange tagIndices, Value numElements, 820 Value stride, Value elementsPerStride) { 821 result.addOperands(srcMemRef); 822 result.addAttribute(getSrcMapAttrName(), AffineMapAttr::get(srcMap)); 823 result.addOperands(srcIndices); 824 result.addOperands(destMemRef); 825 result.addAttribute(getDstMapAttrName(), AffineMapAttr::get(dstMap)); 826 result.addOperands(destIndices); 827 result.addOperands(tagMemRef); 828 result.addAttribute(getTagMapAttrName(), AffineMapAttr::get(tagMap)); 829 result.addOperands(tagIndices); 830 result.addOperands(numElements); 831 if (stride) { 832 result.addOperands({stride, elementsPerStride}); 833 } 834 } 835 836 void AffineDmaStartOp::print(OpAsmPrinter &p) { 837 p << "affine.dma_start " << getSrcMemRef() << '['; 838 p.printAffineMapOfSSAIds(getSrcMapAttr(), getSrcIndices()); 839 p << "], " << getDstMemRef() << '['; 840 p.printAffineMapOfSSAIds(getDstMapAttr(), getDstIndices()); 841 p << "], " << getTagMemRef() << '['; 842 p.printAffineMapOfSSAIds(getTagMapAttr(), getTagIndices()); 843 p << "], " << getNumElements(); 844 if (isStrided()) { 845 p << ", " << getStride(); 846 p << ", " << getNumElementsPerStride(); 847 } 848 p << " : " << getSrcMemRefType() << ", " << getDstMemRefType() << ", " 849 << getTagMemRefType(); 850 } 851 852 // Parse AffineDmaStartOp. 853 // Ex: 854 // affine.dma_start %src[%i, %j], %dst[%k, %l], %tag[%index], %size, 855 // %stride, %num_elt_per_stride 856 // : memref<3076 x f32, 0>, memref<1024 x f32, 2>, memref<1 x i32> 857 // 858 ParseResult AffineDmaStartOp::parse(OpAsmParser &parser, 859 OperationState &result) { 860 OpAsmParser::OperandType srcMemRefInfo; 861 AffineMapAttr srcMapAttr; 862 SmallVector<OpAsmParser::OperandType, 4> srcMapOperands; 863 OpAsmParser::OperandType dstMemRefInfo; 864 AffineMapAttr dstMapAttr; 865 SmallVector<OpAsmParser::OperandType, 4> dstMapOperands; 866 OpAsmParser::OperandType tagMemRefInfo; 867 AffineMapAttr tagMapAttr; 868 SmallVector<OpAsmParser::OperandType, 4> tagMapOperands; 869 OpAsmParser::OperandType numElementsInfo; 870 SmallVector<OpAsmParser::OperandType, 2> strideInfo; 871 872 SmallVector<Type, 3> types; 873 auto indexType = parser.getBuilder().getIndexType(); 874 875 // Parse and resolve the following list of operands: 876 // *) dst memref followed by its affine maps operands (in square brackets). 877 // *) src memref followed by its affine map operands (in square brackets). 878 // *) tag memref followed by its affine map operands (in square brackets). 879 // *) number of elements transferred by DMA operation. 880 if (parser.parseOperand(srcMemRefInfo) || 881 parser.parseAffineMapOfSSAIds(srcMapOperands, srcMapAttr, 882 getSrcMapAttrName(), result.attributes) || 883 parser.parseComma() || parser.parseOperand(dstMemRefInfo) || 884 parser.parseAffineMapOfSSAIds(dstMapOperands, dstMapAttr, 885 getDstMapAttrName(), result.attributes) || 886 parser.parseComma() || parser.parseOperand(tagMemRefInfo) || 887 parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr, 888 getTagMapAttrName(), result.attributes) || 889 parser.parseComma() || parser.parseOperand(numElementsInfo)) 890 return failure(); 891 892 // Parse optional stride and elements per stride. 893 if (parser.parseTrailingOperandList(strideInfo)) { 894 return failure(); 895 } 896 if (!strideInfo.empty() && strideInfo.size() != 2) { 897 return parser.emitError(parser.getNameLoc(), 898 "expected two stride related operands"); 899 } 900 bool isStrided = strideInfo.size() == 2; 901 902 if (parser.parseColonTypeList(types)) 903 return failure(); 904 905 if (types.size() != 3) 906 return parser.emitError(parser.getNameLoc(), "expected three types"); 907 908 if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) || 909 parser.resolveOperands(srcMapOperands, indexType, result.operands) || 910 parser.resolveOperand(dstMemRefInfo, types[1], result.operands) || 911 parser.resolveOperands(dstMapOperands, indexType, result.operands) || 912 parser.resolveOperand(tagMemRefInfo, types[2], result.operands) || 913 parser.resolveOperands(tagMapOperands, indexType, result.operands) || 914 parser.resolveOperand(numElementsInfo, indexType, result.operands)) 915 return failure(); 916 917 if (isStrided) { 918 if (parser.resolveOperands(strideInfo, indexType, result.operands)) 919 return failure(); 920 } 921 922 // Check that src/dst/tag operand counts match their map.numInputs. 923 if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() || 924 dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() || 925 tagMapOperands.size() != tagMapAttr.getValue().getNumInputs()) 926 return parser.emitError(parser.getNameLoc(), 927 "memref operand count not equal to map.numInputs"); 928 return success(); 929 } 930 931 LogicalResult AffineDmaStartOp::verify() { 932 if (!getOperand(getSrcMemRefOperandIndex()).getType().isa<MemRefType>()) 933 return emitOpError("expected DMA source to be of memref type"); 934 if (!getOperand(getDstMemRefOperandIndex()).getType().isa<MemRefType>()) 935 return emitOpError("expected DMA destination to be of memref type"); 936 if (!getOperand(getTagMemRefOperandIndex()).getType().isa<MemRefType>()) 937 return emitOpError("expected DMA tag to be of memref type"); 938 939 // DMAs from different memory spaces supported. 940 if (getSrcMemorySpace() == getDstMemorySpace()) { 941 return emitOpError("DMA should be between different memory spaces"); 942 } 943 unsigned numInputsAllMaps = getSrcMap().getNumInputs() + 944 getDstMap().getNumInputs() + 945 getTagMap().getNumInputs(); 946 if (getNumOperands() != numInputsAllMaps + 3 + 1 && 947 getNumOperands() != numInputsAllMaps + 3 + 1 + 2) { 948 return emitOpError("incorrect number of operands"); 949 } 950 951 for (auto idx : getSrcIndices()) { 952 if (!idx.getType().isIndex()) 953 return emitOpError("src index to dma_start must have 'index' type"); 954 if (!isValidAffineIndexOperand(idx)) 955 return emitOpError("src index must be a dimension or symbol identifier"); 956 } 957 for (auto idx : getDstIndices()) { 958 if (!idx.getType().isIndex()) 959 return emitOpError("dst index to dma_start must have 'index' type"); 960 if (!isValidAffineIndexOperand(idx)) 961 return emitOpError("dst index must be a dimension or symbol identifier"); 962 } 963 for (auto idx : getTagIndices()) { 964 if (!idx.getType().isIndex()) 965 return emitOpError("tag index to dma_start must have 'index' type"); 966 if (!isValidAffineIndexOperand(idx)) 967 return emitOpError("tag index must be a dimension or symbol identifier"); 968 } 969 return success(); 970 } 971 972 LogicalResult AffineDmaStartOp::fold(ArrayRef<Attribute> cstOperands, 973 SmallVectorImpl<OpFoldResult> &results) { 974 /// dma_start(memrefcast) -> dma_start 975 return foldMemRefCast(*this); 976 } 977 978 //===----------------------------------------------------------------------===// 979 // AffineDmaWaitOp 980 //===----------------------------------------------------------------------===// 981 982 // TODO(b/133776335) Check that map operands are loop IVs or symbols. 983 void AffineDmaWaitOp::build(Builder *builder, OperationState &result, 984 Value tagMemRef, AffineMap tagMap, 985 ValueRange tagIndices, Value numElements) { 986 result.addOperands(tagMemRef); 987 result.addAttribute(getTagMapAttrName(), AffineMapAttr::get(tagMap)); 988 result.addOperands(tagIndices); 989 result.addOperands(numElements); 990 } 991 992 void AffineDmaWaitOp::print(OpAsmPrinter &p) { 993 p << "affine.dma_wait " << getTagMemRef() << '['; 994 SmallVector<Value, 2> operands(getTagIndices()); 995 p.printAffineMapOfSSAIds(getTagMapAttr(), operands); 996 p << "], "; 997 p.printOperand(getNumElements()); 998 p << " : " << getTagMemRef().getType(); 999 } 1000 1001 // Parse AffineDmaWaitOp. 1002 // Eg: 1003 // affine.dma_wait %tag[%index], %num_elements 1004 // : memref<1 x i32, (d0) -> (d0), 4> 1005 // 1006 ParseResult AffineDmaWaitOp::parse(OpAsmParser &parser, 1007 OperationState &result) { 1008 OpAsmParser::OperandType tagMemRefInfo; 1009 AffineMapAttr tagMapAttr; 1010 SmallVector<OpAsmParser::OperandType, 2> tagMapOperands; 1011 Type type; 1012 auto indexType = parser.getBuilder().getIndexType(); 1013 OpAsmParser::OperandType numElementsInfo; 1014 1015 // Parse tag memref, its map operands, and dma size. 1016 if (parser.parseOperand(tagMemRefInfo) || 1017 parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr, 1018 getTagMapAttrName(), result.attributes) || 1019 parser.parseComma() || parser.parseOperand(numElementsInfo) || 1020 parser.parseColonType(type) || 1021 parser.resolveOperand(tagMemRefInfo, type, result.operands) || 1022 parser.resolveOperands(tagMapOperands, indexType, result.operands) || 1023 parser.resolveOperand(numElementsInfo, indexType, result.operands)) 1024 return failure(); 1025 1026 if (!type.isa<MemRefType>()) 1027 return parser.emitError(parser.getNameLoc(), 1028 "expected tag to be of memref type"); 1029 1030 if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs()) 1031 return parser.emitError(parser.getNameLoc(), 1032 "tag memref operand count != to map.numInputs"); 1033 return success(); 1034 } 1035 1036 LogicalResult AffineDmaWaitOp::verify() { 1037 if (!getOperand(0).getType().isa<MemRefType>()) 1038 return emitOpError("expected DMA tag to be of memref type"); 1039 for (auto idx : getTagIndices()) { 1040 if (!idx.getType().isIndex()) 1041 return emitOpError("index to dma_wait must have 'index' type"); 1042 if (!isValidAffineIndexOperand(idx)) 1043 return emitOpError("index must be a dimension or symbol identifier"); 1044 } 1045 return success(); 1046 } 1047 1048 LogicalResult AffineDmaWaitOp::fold(ArrayRef<Attribute> cstOperands, 1049 SmallVectorImpl<OpFoldResult> &results) { 1050 /// dma_wait(memrefcast) -> dma_wait 1051 return foldMemRefCast(*this); 1052 } 1053 1054 //===----------------------------------------------------------------------===// 1055 // AffineForOp 1056 //===----------------------------------------------------------------------===// 1057 1058 void AffineForOp::build(Builder *builder, OperationState &result, 1059 ValueRange lbOperands, AffineMap lbMap, 1060 ValueRange ubOperands, AffineMap ubMap, int64_t step) { 1061 assert(((!lbMap && lbOperands.empty()) || 1062 lbOperands.size() == lbMap.getNumInputs()) && 1063 "lower bound operand count does not match the affine map"); 1064 assert(((!ubMap && ubOperands.empty()) || 1065 ubOperands.size() == ubMap.getNumInputs()) && 1066 "upper bound operand count does not match the affine map"); 1067 assert(step > 0 && "step has to be a positive integer constant"); 1068 1069 // Add an attribute for the step. 1070 result.addAttribute(getStepAttrName(), 1071 builder->getIntegerAttr(builder->getIndexType(), step)); 1072 1073 // Add the lower bound. 1074 result.addAttribute(getLowerBoundAttrName(), AffineMapAttr::get(lbMap)); 1075 result.addOperands(lbOperands); 1076 1077 // Add the upper bound. 1078 result.addAttribute(getUpperBoundAttrName(), AffineMapAttr::get(ubMap)); 1079 result.addOperands(ubOperands); 1080 1081 // Create a region and a block for the body. The argument of the region is 1082 // the loop induction variable. 1083 Region *bodyRegion = result.addRegion(); 1084 Block *body = new Block(); 1085 body->addArgument(IndexType::get(builder->getContext())); 1086 bodyRegion->push_back(body); 1087 ensureTerminator(*bodyRegion, *builder, result.location); 1088 1089 // Set the operands list as resizable so that we can freely modify the bounds. 1090 result.setOperandListToResizable(); 1091 } 1092 1093 void AffineForOp::build(Builder *builder, OperationState &result, int64_t lb, 1094 int64_t ub, int64_t step) { 1095 auto lbMap = AffineMap::getConstantMap(lb, builder->getContext()); 1096 auto ubMap = AffineMap::getConstantMap(ub, builder->getContext()); 1097 return build(builder, result, {}, lbMap, {}, ubMap, step); 1098 } 1099 1100 static LogicalResult verify(AffineForOp op) { 1101 // Check that the body defines as single block argument for the induction 1102 // variable. 1103 auto *body = op.getBody(); 1104 if (body->getNumArguments() != 1 || !body->getArgument(0).getType().isIndex()) 1105 return op.emitOpError( 1106 "expected body to have a single index argument for the " 1107 "induction variable"); 1108 1109 // Verify that there are enough operands for the bounds. 1110 AffineMap lowerBoundMap = op.getLowerBoundMap(), 1111 upperBoundMap = op.getUpperBoundMap(); 1112 if (op.getNumOperands() != 1113 (lowerBoundMap.getNumInputs() + upperBoundMap.getNumInputs())) 1114 return op.emitOpError( 1115 "operand count must match with affine map dimension and symbol count"); 1116 1117 // Verify that the bound operands are valid dimension/symbols. 1118 /// Lower bound. 1119 if (failed(verifyDimAndSymbolIdentifiers(op, op.getLowerBoundOperands(), 1120 op.getLowerBoundMap().getNumDims()))) 1121 return failure(); 1122 /// Upper bound. 1123 if (failed(verifyDimAndSymbolIdentifiers(op, op.getUpperBoundOperands(), 1124 op.getUpperBoundMap().getNumDims()))) 1125 return failure(); 1126 return success(); 1127 } 1128 1129 /// Parse a for operation loop bounds. 1130 static ParseResult parseBound(bool isLower, OperationState &result, 1131 OpAsmParser &p) { 1132 // 'min' / 'max' prefixes are generally syntactic sugar, but are required if 1133 // the map has multiple results. 1134 bool failedToParsedMinMax = 1135 failed(p.parseOptionalKeyword(isLower ? "max" : "min")); 1136 1137 auto &builder = p.getBuilder(); 1138 auto boundAttrName = isLower ? AffineForOp::getLowerBoundAttrName() 1139 : AffineForOp::getUpperBoundAttrName(); 1140 1141 // Parse ssa-id as identity map. 1142 SmallVector<OpAsmParser::OperandType, 1> boundOpInfos; 1143 if (p.parseOperandList(boundOpInfos)) 1144 return failure(); 1145 1146 if (!boundOpInfos.empty()) { 1147 // Check that only one operand was parsed. 1148 if (boundOpInfos.size() > 1) 1149 return p.emitError(p.getNameLoc(), 1150 "expected only one loop bound operand"); 1151 1152 // TODO: improve error message when SSA value is not of index type. 1153 // Currently it is 'use of value ... expects different type than prior uses' 1154 if (p.resolveOperand(boundOpInfos.front(), builder.getIndexType(), 1155 result.operands)) 1156 return failure(); 1157 1158 // Create an identity map using symbol id. This representation is optimized 1159 // for storage. Analysis passes may expand it into a multi-dimensional map 1160 // if desired. 1161 AffineMap map = builder.getSymbolIdentityMap(); 1162 result.addAttribute(boundAttrName, AffineMapAttr::get(map)); 1163 return success(); 1164 } 1165 1166 // Get the attribute location. 1167 llvm::SMLoc attrLoc = p.getCurrentLocation(); 1168 1169 Attribute boundAttr; 1170 if (p.parseAttribute(boundAttr, builder.getIndexType(), boundAttrName, 1171 result.attributes)) 1172 return failure(); 1173 1174 // Parse full form - affine map followed by dim and symbol list. 1175 if (auto affineMapAttr = boundAttr.dyn_cast<AffineMapAttr>()) { 1176 unsigned currentNumOperands = result.operands.size(); 1177 unsigned numDims; 1178 if (parseDimAndSymbolList(p, result.operands, numDims)) 1179 return failure(); 1180 1181 auto map = affineMapAttr.getValue(); 1182 if (map.getNumDims() != numDims) 1183 return p.emitError( 1184 p.getNameLoc(), 1185 "dim operand count and affine map dim count must match"); 1186 1187 unsigned numDimAndSymbolOperands = 1188 result.operands.size() - currentNumOperands; 1189 if (numDims + map.getNumSymbols() != numDimAndSymbolOperands) 1190 return p.emitError( 1191 p.getNameLoc(), 1192 "symbol operand count and affine map symbol count must match"); 1193 1194 // If the map has multiple results, make sure that we parsed the min/max 1195 // prefix. 1196 if (map.getNumResults() > 1 && failedToParsedMinMax) { 1197 if (isLower) { 1198 return p.emitError(attrLoc, "lower loop bound affine map with " 1199 "multiple results requires 'max' prefix"); 1200 } 1201 return p.emitError(attrLoc, "upper loop bound affine map with multiple " 1202 "results requires 'min' prefix"); 1203 } 1204 return success(); 1205 } 1206 1207 // Parse custom assembly form. 1208 if (auto integerAttr = boundAttr.dyn_cast<IntegerAttr>()) { 1209 result.attributes.pop_back(); 1210 result.addAttribute( 1211 boundAttrName, 1212 AffineMapAttr::get(builder.getConstantAffineMap(integerAttr.getInt()))); 1213 return success(); 1214 } 1215 1216 return p.emitError( 1217 p.getNameLoc(), 1218 "expected valid affine map representation for loop bounds"); 1219 } 1220 1221 static ParseResult parseAffineForOp(OpAsmParser &parser, 1222 OperationState &result) { 1223 auto &builder = parser.getBuilder(); 1224 OpAsmParser::OperandType inductionVariable; 1225 // Parse the induction variable followed by '='. 1226 if (parser.parseRegionArgument(inductionVariable) || parser.parseEqual()) 1227 return failure(); 1228 1229 // Parse loop bounds. 1230 if (parseBound(/*isLower=*/true, result, parser) || 1231 parser.parseKeyword("to", " between bounds") || 1232 parseBound(/*isLower=*/false, result, parser)) 1233 return failure(); 1234 1235 // Parse the optional loop step, we default to 1 if one is not present. 1236 if (parser.parseOptionalKeyword("step")) { 1237 result.addAttribute( 1238 AffineForOp::getStepAttrName(), 1239 builder.getIntegerAttr(builder.getIndexType(), /*value=*/1)); 1240 } else { 1241 llvm::SMLoc stepLoc = parser.getCurrentLocation(); 1242 IntegerAttr stepAttr; 1243 if (parser.parseAttribute(stepAttr, builder.getIndexType(), 1244 AffineForOp::getStepAttrName().data(), 1245 result.attributes)) 1246 return failure(); 1247 1248 if (stepAttr.getValue().getSExtValue() < 0) 1249 return parser.emitError( 1250 stepLoc, 1251 "expected step to be representable as a positive signed integer"); 1252 } 1253 1254 // Parse the body region. 1255 Region *body = result.addRegion(); 1256 if (parser.parseRegion(*body, inductionVariable, builder.getIndexType())) 1257 return failure(); 1258 1259 AffineForOp::ensureTerminator(*body, builder, result.location); 1260 1261 // Parse the optional attribute list. 1262 if (parser.parseOptionalAttrDict(result.attributes)) 1263 return failure(); 1264 1265 // Set the operands list as resizable so that we can freely modify the bounds. 1266 result.setOperandListToResizable(); 1267 return success(); 1268 } 1269 1270 static void printBound(AffineMapAttr boundMap, 1271 Operation::operand_range boundOperands, 1272 const char *prefix, OpAsmPrinter &p) { 1273 AffineMap map = boundMap.getValue(); 1274 1275 // Check if this bound should be printed using custom assembly form. 1276 // The decision to restrict printing custom assembly form to trivial cases 1277 // comes from the will to roundtrip MLIR binary -> text -> binary in a 1278 // lossless way. 1279 // Therefore, custom assembly form parsing and printing is only supported for 1280 // zero-operand constant maps and single symbol operand identity maps. 1281 if (map.getNumResults() == 1) { 1282 AffineExpr expr = map.getResult(0); 1283 1284 // Print constant bound. 1285 if (map.getNumDims() == 0 && map.getNumSymbols() == 0) { 1286 if (auto constExpr = expr.dyn_cast<AffineConstantExpr>()) { 1287 p << constExpr.getValue(); 1288 return; 1289 } 1290 } 1291 1292 // Print bound that consists of a single SSA symbol if the map is over a 1293 // single symbol. 1294 if (map.getNumDims() == 0 && map.getNumSymbols() == 1) { 1295 if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) { 1296 p.printOperand(*boundOperands.begin()); 1297 return; 1298 } 1299 } 1300 } else { 1301 // Map has multiple results. Print 'min' or 'max' prefix. 1302 p << prefix << ' '; 1303 } 1304 1305 // Print the map and its operands. 1306 p << boundMap; 1307 printDimAndSymbolList(boundOperands.begin(), boundOperands.end(), 1308 map.getNumDims(), p); 1309 } 1310 1311 static void print(OpAsmPrinter &p, AffineForOp op) { 1312 p << op.getOperationName() << ' '; 1313 p.printOperand(op.getBody()->getArgument(0)); 1314 p << " = "; 1315 printBound(op.getLowerBoundMapAttr(), op.getLowerBoundOperands(), "max", p); 1316 p << " to "; 1317 printBound(op.getUpperBoundMapAttr(), op.getUpperBoundOperands(), "min", p); 1318 1319 if (op.getStep() != 1) 1320 p << " step " << op.getStep(); 1321 p.printRegion(op.region(), 1322 /*printEntryBlockArgs=*/false, 1323 /*printBlockTerminators=*/false); 1324 p.printOptionalAttrDict(op.getAttrs(), 1325 /*elidedAttrs=*/{op.getLowerBoundAttrName(), 1326 op.getUpperBoundAttrName(), 1327 op.getStepAttrName()}); 1328 } 1329 1330 /// Fold the constant bounds of a loop. 1331 static LogicalResult foldLoopBounds(AffineForOp forOp) { 1332 auto foldLowerOrUpperBound = [&forOp](bool lower) { 1333 // Check to see if each of the operands is the result of a constant. If 1334 // so, get the value. If not, ignore it. 1335 SmallVector<Attribute, 8> operandConstants; 1336 auto boundOperands = 1337 lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands(); 1338 for (auto operand : boundOperands) { 1339 Attribute operandCst; 1340 matchPattern(operand, m_Constant(&operandCst)); 1341 operandConstants.push_back(operandCst); 1342 } 1343 1344 AffineMap boundMap = 1345 lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap(); 1346 assert(boundMap.getNumResults() >= 1 && 1347 "bound maps should have at least one result"); 1348 SmallVector<Attribute, 4> foldedResults; 1349 if (failed(boundMap.constantFold(operandConstants, foldedResults))) 1350 return failure(); 1351 1352 // Compute the max or min as applicable over the results. 1353 assert(!foldedResults.empty() && "bounds should have at least one result"); 1354 auto maxOrMin = foldedResults[0].cast<IntegerAttr>().getValue(); 1355 for (unsigned i = 1, e = foldedResults.size(); i < e; i++) { 1356 auto foldedResult = foldedResults[i].cast<IntegerAttr>().getValue(); 1357 maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult) 1358 : llvm::APIntOps::smin(maxOrMin, foldedResult); 1359 } 1360 lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue()) 1361 : forOp.setConstantUpperBound(maxOrMin.getSExtValue()); 1362 return success(); 1363 }; 1364 1365 // Try to fold the lower bound. 1366 bool folded = false; 1367 if (!forOp.hasConstantLowerBound()) 1368 folded |= succeeded(foldLowerOrUpperBound(/*lower=*/true)); 1369 1370 // Try to fold the upper bound. 1371 if (!forOp.hasConstantUpperBound()) 1372 folded |= succeeded(foldLowerOrUpperBound(/*lower=*/false)); 1373 return success(folded); 1374 } 1375 1376 /// Canonicalize the bounds of the given loop. 1377 static LogicalResult canonicalizeLoopBounds(AffineForOp forOp) { 1378 SmallVector<Value, 4> lbOperands(forOp.getLowerBoundOperands()); 1379 SmallVector<Value, 4> ubOperands(forOp.getUpperBoundOperands()); 1380 1381 auto lbMap = forOp.getLowerBoundMap(); 1382 auto ubMap = forOp.getUpperBoundMap(); 1383 auto prevLbMap = lbMap; 1384 auto prevUbMap = ubMap; 1385 1386 canonicalizeMapAndOperands(&lbMap, &lbOperands); 1387 lbMap = removeDuplicateExprs(lbMap); 1388 1389 canonicalizeMapAndOperands(&ubMap, &ubOperands); 1390 ubMap = removeDuplicateExprs(ubMap); 1391 1392 // Any canonicalization change always leads to updated map(s). 1393 if (lbMap == prevLbMap && ubMap == prevUbMap) 1394 return failure(); 1395 1396 if (lbMap != prevLbMap) 1397 forOp.setLowerBound(lbOperands, lbMap); 1398 if (ubMap != prevUbMap) 1399 forOp.setUpperBound(ubOperands, ubMap); 1400 return success(); 1401 } 1402 1403 namespace { 1404 /// This is a pattern to fold trivially empty loops. 1405 struct AffineForEmptyLoopFolder : public OpRewritePattern<AffineForOp> { 1406 using OpRewritePattern<AffineForOp>::OpRewritePattern; 1407 1408 LogicalResult matchAndRewrite(AffineForOp forOp, 1409 PatternRewriter &rewriter) const override { 1410 // Check that the body only contains a terminator. 1411 if (!llvm::hasSingleElement(*forOp.getBody())) 1412 return failure(); 1413 rewriter.eraseOp(forOp); 1414 return success(); 1415 } 1416 }; 1417 } // end anonymous namespace 1418 1419 void AffineForOp::getCanonicalizationPatterns(OwningRewritePatternList &results, 1420 MLIRContext *context) { 1421 results.insert<AffineForEmptyLoopFolder>(context); 1422 } 1423 1424 LogicalResult AffineForOp::fold(ArrayRef<Attribute> operands, 1425 SmallVectorImpl<OpFoldResult> &results) { 1426 bool folded = succeeded(foldLoopBounds(*this)); 1427 folded |= succeeded(canonicalizeLoopBounds(*this)); 1428 return success(folded); 1429 } 1430 1431 AffineBound AffineForOp::getLowerBound() { 1432 auto lbMap = getLowerBoundMap(); 1433 return AffineBound(AffineForOp(*this), 0, lbMap.getNumInputs(), lbMap); 1434 } 1435 1436 AffineBound AffineForOp::getUpperBound() { 1437 auto lbMap = getLowerBoundMap(); 1438 auto ubMap = getUpperBoundMap(); 1439 return AffineBound(AffineForOp(*this), lbMap.getNumInputs(), getNumOperands(), 1440 ubMap); 1441 } 1442 1443 void AffineForOp::setLowerBound(ValueRange lbOperands, AffineMap map) { 1444 assert(lbOperands.size() == map.getNumInputs()); 1445 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 1446 1447 SmallVector<Value, 4> newOperands(lbOperands.begin(), lbOperands.end()); 1448 1449 auto ubOperands = getUpperBoundOperands(); 1450 newOperands.append(ubOperands.begin(), ubOperands.end()); 1451 getOperation()->setOperands(newOperands); 1452 1453 setAttr(getLowerBoundAttrName(), AffineMapAttr::get(map)); 1454 } 1455 1456 void AffineForOp::setUpperBound(ValueRange ubOperands, AffineMap map) { 1457 assert(ubOperands.size() == map.getNumInputs()); 1458 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 1459 1460 SmallVector<Value, 4> newOperands(getLowerBoundOperands()); 1461 newOperands.append(ubOperands.begin(), ubOperands.end()); 1462 getOperation()->setOperands(newOperands); 1463 1464 setAttr(getUpperBoundAttrName(), AffineMapAttr::get(map)); 1465 } 1466 1467 void AffineForOp::setLowerBoundMap(AffineMap map) { 1468 auto lbMap = getLowerBoundMap(); 1469 assert(lbMap.getNumDims() == map.getNumDims() && 1470 lbMap.getNumSymbols() == map.getNumSymbols()); 1471 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 1472 (void)lbMap; 1473 setAttr(getLowerBoundAttrName(), AffineMapAttr::get(map)); 1474 } 1475 1476 void AffineForOp::setUpperBoundMap(AffineMap map) { 1477 auto ubMap = getUpperBoundMap(); 1478 assert(ubMap.getNumDims() == map.getNumDims() && 1479 ubMap.getNumSymbols() == map.getNumSymbols()); 1480 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 1481 (void)ubMap; 1482 setAttr(getUpperBoundAttrName(), AffineMapAttr::get(map)); 1483 } 1484 1485 bool AffineForOp::hasConstantLowerBound() { 1486 return getLowerBoundMap().isSingleConstant(); 1487 } 1488 1489 bool AffineForOp::hasConstantUpperBound() { 1490 return getUpperBoundMap().isSingleConstant(); 1491 } 1492 1493 int64_t AffineForOp::getConstantLowerBound() { 1494 return getLowerBoundMap().getSingleConstantResult(); 1495 } 1496 1497 int64_t AffineForOp::getConstantUpperBound() { 1498 return getUpperBoundMap().getSingleConstantResult(); 1499 } 1500 1501 void AffineForOp::setConstantLowerBound(int64_t value) { 1502 setLowerBound({}, AffineMap::getConstantMap(value, getContext())); 1503 } 1504 1505 void AffineForOp::setConstantUpperBound(int64_t value) { 1506 setUpperBound({}, AffineMap::getConstantMap(value, getContext())); 1507 } 1508 1509 AffineForOp::operand_range AffineForOp::getLowerBoundOperands() { 1510 return {operand_begin(), operand_begin() + getLowerBoundMap().getNumInputs()}; 1511 } 1512 1513 AffineForOp::operand_range AffineForOp::getUpperBoundOperands() { 1514 return {operand_begin() + getLowerBoundMap().getNumInputs(), operand_end()}; 1515 } 1516 1517 bool AffineForOp::matchingBoundOperandList() { 1518 auto lbMap = getLowerBoundMap(); 1519 auto ubMap = getUpperBoundMap(); 1520 if (lbMap.getNumDims() != ubMap.getNumDims() || 1521 lbMap.getNumSymbols() != ubMap.getNumSymbols()) 1522 return false; 1523 1524 unsigned numOperands = lbMap.getNumInputs(); 1525 for (unsigned i = 0, e = lbMap.getNumInputs(); i < e; i++) { 1526 // Compare Value 's. 1527 if (getOperand(i) != getOperand(numOperands + i)) 1528 return false; 1529 } 1530 return true; 1531 } 1532 1533 Region &AffineForOp::getLoopBody() { return region(); } 1534 1535 bool AffineForOp::isDefinedOutsideOfLoop(Value value) { 1536 return !region().isAncestor(value.getParentRegion()); 1537 } 1538 1539 LogicalResult AffineForOp::moveOutOfLoop(ArrayRef<Operation *> ops) { 1540 for (auto *op : ops) 1541 op->moveBefore(*this); 1542 return success(); 1543 } 1544 1545 /// Returns if the provided value is the induction variable of a AffineForOp. 1546 bool mlir::isForInductionVar(Value val) { 1547 return getForInductionVarOwner(val) != AffineForOp(); 1548 } 1549 1550 /// Returns the loop parent of an induction variable. If the provided value is 1551 /// not an induction variable, then return nullptr. 1552 AffineForOp mlir::getForInductionVarOwner(Value val) { 1553 auto ivArg = val.dyn_cast<BlockArgument>(); 1554 if (!ivArg || !ivArg.getOwner()) 1555 return AffineForOp(); 1556 auto *containingInst = ivArg.getOwner()->getParent()->getParentOp(); 1557 return dyn_cast<AffineForOp>(containingInst); 1558 } 1559 1560 /// Extracts the induction variables from a list of AffineForOps and returns 1561 /// them. 1562 void mlir::extractForInductionVars(ArrayRef<AffineForOp> forInsts, 1563 SmallVectorImpl<Value> *ivs) { 1564 ivs->reserve(forInsts.size()); 1565 for (auto forInst : forInsts) 1566 ivs->push_back(forInst.getInductionVar()); 1567 } 1568 1569 //===----------------------------------------------------------------------===// 1570 // AffineIfOp 1571 //===----------------------------------------------------------------------===// 1572 1573 namespace { 1574 /// Remove else blocks that have nothing other than the terminator. 1575 struct SimplifyDeadElse : public OpRewritePattern<AffineIfOp> { 1576 using OpRewritePattern<AffineIfOp>::OpRewritePattern; 1577 1578 LogicalResult matchAndRewrite(AffineIfOp ifOp, 1579 PatternRewriter &rewriter) const override { 1580 if (ifOp.elseRegion().empty() || 1581 !llvm::hasSingleElement(*ifOp.getElseBlock())) 1582 return failure(); 1583 1584 rewriter.startRootUpdate(ifOp); 1585 rewriter.eraseBlock(ifOp.getElseBlock()); 1586 rewriter.finalizeRootUpdate(ifOp); 1587 return success(); 1588 } 1589 }; 1590 } // end anonymous namespace. 1591 1592 static LogicalResult verify(AffineIfOp op) { 1593 // Verify that we have a condition attribute. 1594 auto conditionAttr = 1595 op.getAttrOfType<IntegerSetAttr>(op.getConditionAttrName()); 1596 if (!conditionAttr) 1597 return op.emitOpError( 1598 "requires an integer set attribute named 'condition'"); 1599 1600 // Verify that there are enough operands for the condition. 1601 IntegerSet condition = conditionAttr.getValue(); 1602 if (op.getNumOperands() != condition.getNumInputs()) 1603 return op.emitOpError( 1604 "operand count and condition integer set dimension and " 1605 "symbol count must match"); 1606 1607 // Verify that the operands are valid dimension/symbols. 1608 if (failed(verifyDimAndSymbolIdentifiers(op, op.getOperands(), 1609 condition.getNumDims()))) 1610 return failure(); 1611 1612 // Verify that the entry of each child region does not have arguments. 1613 for (auto ®ion : op.getOperation()->getRegions()) { 1614 for (auto &b : region) 1615 if (b.getNumArguments() != 0) 1616 return op.emitOpError( 1617 "requires that child entry blocks have no arguments"); 1618 } 1619 return success(); 1620 } 1621 1622 static ParseResult parseAffineIfOp(OpAsmParser &parser, 1623 OperationState &result) { 1624 // Parse the condition attribute set. 1625 IntegerSetAttr conditionAttr; 1626 unsigned numDims; 1627 if (parser.parseAttribute(conditionAttr, AffineIfOp::getConditionAttrName(), 1628 result.attributes) || 1629 parseDimAndSymbolList(parser, result.operands, numDims)) 1630 return failure(); 1631 1632 // Verify the condition operands. 1633 auto set = conditionAttr.getValue(); 1634 if (set.getNumDims() != numDims) 1635 return parser.emitError( 1636 parser.getNameLoc(), 1637 "dim operand count and integer set dim count must match"); 1638 if (numDims + set.getNumSymbols() != result.operands.size()) 1639 return parser.emitError( 1640 parser.getNameLoc(), 1641 "symbol operand count and integer set symbol count must match"); 1642 1643 // Create the regions for 'then' and 'else'. The latter must be created even 1644 // if it remains empty for the validity of the operation. 1645 result.regions.reserve(2); 1646 Region *thenRegion = result.addRegion(); 1647 Region *elseRegion = result.addRegion(); 1648 1649 // Parse the 'then' region. 1650 if (parser.parseRegion(*thenRegion, {}, {})) 1651 return failure(); 1652 AffineIfOp::ensureTerminator(*thenRegion, parser.getBuilder(), 1653 result.location); 1654 1655 // If we find an 'else' keyword then parse the 'else' region. 1656 if (!parser.parseOptionalKeyword("else")) { 1657 if (parser.parseRegion(*elseRegion, {}, {})) 1658 return failure(); 1659 AffineIfOp::ensureTerminator(*elseRegion, parser.getBuilder(), 1660 result.location); 1661 } 1662 1663 // Parse the optional attribute list. 1664 if (parser.parseOptionalAttrDict(result.attributes)) 1665 return failure(); 1666 1667 return success(); 1668 } 1669 1670 static void print(OpAsmPrinter &p, AffineIfOp op) { 1671 auto conditionAttr = 1672 op.getAttrOfType<IntegerSetAttr>(op.getConditionAttrName()); 1673 p << "affine.if " << conditionAttr; 1674 printDimAndSymbolList(op.operand_begin(), op.operand_end(), 1675 conditionAttr.getValue().getNumDims(), p); 1676 p.printRegion(op.thenRegion(), 1677 /*printEntryBlockArgs=*/false, 1678 /*printBlockTerminators=*/false); 1679 1680 // Print the 'else' regions if it has any blocks. 1681 auto &elseRegion = op.elseRegion(); 1682 if (!elseRegion.empty()) { 1683 p << " else"; 1684 p.printRegion(elseRegion, 1685 /*printEntryBlockArgs=*/false, 1686 /*printBlockTerminators=*/false); 1687 } 1688 1689 // Print the attribute list. 1690 p.printOptionalAttrDict(op.getAttrs(), 1691 /*elidedAttrs=*/op.getConditionAttrName()); 1692 } 1693 1694 IntegerSet AffineIfOp::getIntegerSet() { 1695 return getAttrOfType<IntegerSetAttr>(getConditionAttrName()).getValue(); 1696 } 1697 void AffineIfOp::setIntegerSet(IntegerSet newSet) { 1698 setAttr(getConditionAttrName(), IntegerSetAttr::get(newSet)); 1699 } 1700 1701 void AffineIfOp::setConditional(IntegerSet set, ValueRange operands) { 1702 setIntegerSet(set); 1703 getOperation()->setOperands(operands); 1704 } 1705 1706 void AffineIfOp::build(Builder *builder, OperationState &result, IntegerSet set, 1707 ValueRange args, bool withElseRegion) { 1708 result.addOperands(args); 1709 result.addAttribute(getConditionAttrName(), IntegerSetAttr::get(set)); 1710 Region *thenRegion = result.addRegion(); 1711 Region *elseRegion = result.addRegion(); 1712 AffineIfOp::ensureTerminator(*thenRegion, *builder, result.location); 1713 if (withElseRegion) 1714 AffineIfOp::ensureTerminator(*elseRegion, *builder, result.location); 1715 } 1716 1717 /// Canonicalize an affine if op's conditional (integer set + operands). 1718 LogicalResult AffineIfOp::fold(ArrayRef<Attribute>, 1719 SmallVectorImpl<OpFoldResult> &) { 1720 auto set = getIntegerSet(); 1721 SmallVector<Value, 4> operands(getOperands()); 1722 canonicalizeSetAndOperands(&set, &operands); 1723 1724 // Any canonicalization change always leads to either a reduction in the 1725 // number of operands or a change in the number of symbolic operands 1726 // (promotion of dims to symbols). 1727 if (operands.size() < getIntegerSet().getNumInputs() || 1728 set.getNumSymbols() > getIntegerSet().getNumSymbols()) { 1729 setConditional(set, operands); 1730 return success(); 1731 } 1732 1733 return failure(); 1734 } 1735 1736 void AffineIfOp::getCanonicalizationPatterns(OwningRewritePatternList &results, 1737 MLIRContext *context) { 1738 results.insert<SimplifyDeadElse>(context); 1739 } 1740 1741 //===----------------------------------------------------------------------===// 1742 // AffineLoadOp 1743 //===----------------------------------------------------------------------===// 1744 1745 void AffineLoadOp::build(Builder *builder, OperationState &result, 1746 AffineMap map, ValueRange operands) { 1747 assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands"); 1748 result.addOperands(operands); 1749 if (map) 1750 result.addAttribute(getMapAttrName(), AffineMapAttr::get(map)); 1751 auto memrefType = operands[0].getType().cast<MemRefType>(); 1752 result.types.push_back(memrefType.getElementType()); 1753 } 1754 1755 void AffineLoadOp::build(Builder *builder, OperationState &result, Value memref, 1756 AffineMap map, ValueRange mapOperands) { 1757 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info"); 1758 result.addOperands(memref); 1759 result.addOperands(mapOperands); 1760 auto memrefType = memref.getType().cast<MemRefType>(); 1761 result.addAttribute(getMapAttrName(), AffineMapAttr::get(map)); 1762 result.types.push_back(memrefType.getElementType()); 1763 } 1764 1765 void AffineLoadOp::build(Builder *builder, OperationState &result, Value memref, 1766 ValueRange indices) { 1767 auto memrefType = memref.getType().cast<MemRefType>(); 1768 auto rank = memrefType.getRank(); 1769 // Create identity map for memrefs with at least one dimension or () -> () 1770 // for zero-dimensional memrefs. 1771 auto map = rank ? builder->getMultiDimIdentityMap(rank) 1772 : builder->getEmptyAffineMap(); 1773 build(builder, result, memref, map, indices); 1774 } 1775 1776 ParseResult AffineLoadOp::parse(OpAsmParser &parser, OperationState &result) { 1777 auto &builder = parser.getBuilder(); 1778 auto indexTy = builder.getIndexType(); 1779 1780 MemRefType type; 1781 OpAsmParser::OperandType memrefInfo; 1782 AffineMapAttr mapAttr; 1783 SmallVector<OpAsmParser::OperandType, 1> mapOperands; 1784 return failure( 1785 parser.parseOperand(memrefInfo) || 1786 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, getMapAttrName(), 1787 result.attributes) || 1788 parser.parseOptionalAttrDict(result.attributes) || 1789 parser.parseColonType(type) || 1790 parser.resolveOperand(memrefInfo, type, result.operands) || 1791 parser.resolveOperands(mapOperands, indexTy, result.operands) || 1792 parser.addTypeToList(type.getElementType(), result.types)); 1793 } 1794 1795 void AffineLoadOp::print(OpAsmPrinter &p) { 1796 p << "affine.load " << getMemRef() << '['; 1797 if (AffineMapAttr mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName())) 1798 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 1799 p << ']'; 1800 p.printOptionalAttrDict(getAttrs(), /*elidedAttrs=*/{getMapAttrName()}); 1801 p << " : " << getMemRefType(); 1802 } 1803 1804 LogicalResult AffineLoadOp::verify() { 1805 if (getType() != getMemRefType().getElementType()) 1806 return emitOpError("result type must match element type of memref"); 1807 1808 auto mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName()); 1809 if (mapAttr) { 1810 AffineMap map = getAttrOfType<AffineMapAttr>(getMapAttrName()).getValue(); 1811 if (map.getNumResults() != getMemRefType().getRank()) 1812 return emitOpError("affine.load affine map num results must equal" 1813 " memref rank"); 1814 if (map.getNumInputs() != getNumOperands() - 1) 1815 return emitOpError("expects as many subscripts as affine map inputs"); 1816 } else { 1817 if (getMemRefType().getRank() != getNumOperands() - 1) 1818 return emitOpError( 1819 "expects the number of subscripts to be equal to memref rank"); 1820 } 1821 1822 for (auto idx : getMapOperands()) { 1823 if (!idx.getType().isIndex()) 1824 return emitOpError("index to load must have 'index' type"); 1825 if (!isValidAffineIndexOperand(idx)) 1826 return emitOpError("index must be a dimension or symbol identifier"); 1827 } 1828 return success(); 1829 } 1830 1831 void AffineLoadOp::getCanonicalizationPatterns( 1832 OwningRewritePatternList &results, MLIRContext *context) { 1833 results.insert<SimplifyAffineOp<AffineLoadOp>>(context); 1834 } 1835 1836 OpFoldResult AffineLoadOp::fold(ArrayRef<Attribute> cstOperands) { 1837 /// load(memrefcast) -> load 1838 if (succeeded(foldMemRefCast(*this))) 1839 return getResult(); 1840 return OpFoldResult(); 1841 } 1842 1843 //===----------------------------------------------------------------------===// 1844 // AffineStoreOp 1845 //===----------------------------------------------------------------------===// 1846 1847 void AffineStoreOp::build(Builder *builder, OperationState &result, 1848 Value valueToStore, Value memref, AffineMap map, 1849 ValueRange mapOperands) { 1850 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info"); 1851 result.addOperands(valueToStore); 1852 result.addOperands(memref); 1853 result.addOperands(mapOperands); 1854 result.addAttribute(getMapAttrName(), AffineMapAttr::get(map)); 1855 } 1856 1857 // Use identity map. 1858 void AffineStoreOp::build(Builder *builder, OperationState &result, 1859 Value valueToStore, Value memref, 1860 ValueRange indices) { 1861 auto memrefType = memref.getType().cast<MemRefType>(); 1862 auto rank = memrefType.getRank(); 1863 // Create identity map for memrefs with at least one dimension or () -> () 1864 // for zero-dimensional memrefs. 1865 auto map = rank ? builder->getMultiDimIdentityMap(rank) 1866 : builder->getEmptyAffineMap(); 1867 build(builder, result, valueToStore, memref, map, indices); 1868 } 1869 1870 ParseResult AffineStoreOp::parse(OpAsmParser &parser, OperationState &result) { 1871 auto indexTy = parser.getBuilder().getIndexType(); 1872 1873 MemRefType type; 1874 OpAsmParser::OperandType storeValueInfo; 1875 OpAsmParser::OperandType memrefInfo; 1876 AffineMapAttr mapAttr; 1877 SmallVector<OpAsmParser::OperandType, 1> mapOperands; 1878 return failure(parser.parseOperand(storeValueInfo) || parser.parseComma() || 1879 parser.parseOperand(memrefInfo) || 1880 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, 1881 getMapAttrName(), 1882 result.attributes) || 1883 parser.parseOptionalAttrDict(result.attributes) || 1884 parser.parseColonType(type) || 1885 parser.resolveOperand(storeValueInfo, type.getElementType(), 1886 result.operands) || 1887 parser.resolveOperand(memrefInfo, type, result.operands) || 1888 parser.resolveOperands(mapOperands, indexTy, result.operands)); 1889 } 1890 1891 void AffineStoreOp::print(OpAsmPrinter &p) { 1892 p << "affine.store " << getValueToStore(); 1893 p << ", " << getMemRef() << '['; 1894 if (AffineMapAttr mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName())) 1895 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 1896 p << ']'; 1897 p.printOptionalAttrDict(getAttrs(), /*elidedAttrs=*/{getMapAttrName()}); 1898 p << " : " << getMemRefType(); 1899 } 1900 1901 LogicalResult AffineStoreOp::verify() { 1902 // First operand must have same type as memref element type. 1903 if (getValueToStore().getType() != getMemRefType().getElementType()) 1904 return emitOpError("first operand must have same type memref element type"); 1905 1906 auto mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName()); 1907 if (mapAttr) { 1908 AffineMap map = mapAttr.getValue(); 1909 if (map.getNumResults() != getMemRefType().getRank()) 1910 return emitOpError("affine.store affine map num results must equal" 1911 " memref rank"); 1912 if (map.getNumInputs() != getNumOperands() - 2) 1913 return emitOpError("expects as many subscripts as affine map inputs"); 1914 } else { 1915 if (getMemRefType().getRank() != getNumOperands() - 2) 1916 return emitOpError( 1917 "expects the number of subscripts to be equal to memref rank"); 1918 } 1919 1920 for (auto idx : getMapOperands()) { 1921 if (!idx.getType().isIndex()) 1922 return emitOpError("index to store must have 'index' type"); 1923 if (!isValidAffineIndexOperand(idx)) 1924 return emitOpError("index must be a dimension or symbol identifier"); 1925 } 1926 return success(); 1927 } 1928 1929 void AffineStoreOp::getCanonicalizationPatterns( 1930 OwningRewritePatternList &results, MLIRContext *context) { 1931 results.insert<SimplifyAffineOp<AffineStoreOp>>(context); 1932 } 1933 1934 LogicalResult AffineStoreOp::fold(ArrayRef<Attribute> cstOperands, 1935 SmallVectorImpl<OpFoldResult> &results) { 1936 /// store(memrefcast) -> store 1937 return foldMemRefCast(*this); 1938 } 1939 1940 //===----------------------------------------------------------------------===// 1941 // AffineMinMaxOpBase 1942 //===----------------------------------------------------------------------===// 1943 1944 template <typename T> 1945 static LogicalResult verifyAffineMinMaxOp(T op) { 1946 // Verify that operand count matches affine map dimension and symbol count. 1947 if (op.getNumOperands() != op.map().getNumDims() + op.map().getNumSymbols()) 1948 return op.emitOpError( 1949 "operand count and affine map dimension and symbol count must match"); 1950 return success(); 1951 } 1952 1953 template <typename T> 1954 static void printAffineMinMaxOp(OpAsmPrinter &p, T op) { 1955 p << op.getOperationName() << ' ' << op.getAttr(T::getMapAttrName()); 1956 auto operands = op.getOperands(); 1957 unsigned numDims = op.map().getNumDims(); 1958 p << '(' << operands.take_front(numDims) << ')'; 1959 1960 if (operands.size() != numDims) 1961 p << '[' << operands.drop_front(numDims) << ']'; 1962 p.printOptionalAttrDict(op.getAttrs(), 1963 /*elidedAttrs=*/{T::getMapAttrName()}); 1964 } 1965 1966 template <typename T> 1967 static ParseResult parseAffineMinMaxOp(OpAsmParser &parser, 1968 OperationState &result) { 1969 auto &builder = parser.getBuilder(); 1970 auto indexType = builder.getIndexType(); 1971 SmallVector<OpAsmParser::OperandType, 8> dim_infos; 1972 SmallVector<OpAsmParser::OperandType, 8> sym_infos; 1973 AffineMapAttr mapAttr; 1974 return failure( 1975 parser.parseAttribute(mapAttr, T::getMapAttrName(), result.attributes) || 1976 parser.parseOperandList(dim_infos, OpAsmParser::Delimiter::Paren) || 1977 parser.parseOperandList(sym_infos, 1978 OpAsmParser::Delimiter::OptionalSquare) || 1979 parser.parseOptionalAttrDict(result.attributes) || 1980 parser.resolveOperands(dim_infos, indexType, result.operands) || 1981 parser.resolveOperands(sym_infos, indexType, result.operands) || 1982 parser.addTypeToList(indexType, result.types)); 1983 } 1984 1985 //===----------------------------------------------------------------------===// 1986 // AffineMinOp 1987 //===----------------------------------------------------------------------===// 1988 // 1989 // %0 = affine.min (d0) -> (1000, d0 + 512) (%i0) 1990 // 1991 1992 OpFoldResult AffineMinOp::fold(ArrayRef<Attribute> operands) { 1993 // Fold the affine map. 1994 // TODO(andydavis, ntv) Fold more cases: partial static information, 1995 // min(some_affine, some_affine + constant, ...). 1996 SmallVector<Attribute, 2> results; 1997 if (failed(map().constantFold(operands, results))) 1998 return {}; 1999 2000 // Compute and return min of folded map results. 2001 int64_t min = std::numeric_limits<int64_t>::max(); 2002 int minIndex = -1; 2003 for (unsigned i = 0, e = results.size(); i < e; ++i) { 2004 auto intAttr = results[i].cast<IntegerAttr>(); 2005 if (intAttr.getInt() < min) { 2006 min = intAttr.getInt(); 2007 minIndex = i; 2008 } 2009 } 2010 if (minIndex < 0) 2011 return {}; 2012 return results[minIndex]; 2013 } 2014 2015 void AffineMinOp::getCanonicalizationPatterns( 2016 OwningRewritePatternList &patterns, MLIRContext *context) { 2017 patterns.insert<SimplifyAffineOp<AffineMinOp>>(context); 2018 } 2019 2020 //===----------------------------------------------------------------------===// 2021 // AffineMaxOp 2022 //===----------------------------------------------------------------------===// 2023 // 2024 // %0 = affine.max (d0) -> (1000, d0 + 512) (%i0) 2025 // 2026 2027 OpFoldResult AffineMaxOp::fold(ArrayRef<Attribute> operands) { 2028 // Fold the affine map. 2029 // TODO(andydavis, ntv, ouhang) Fold more cases: partial static information, 2030 // max(some_affine, some_affine + constant, ...). 2031 SmallVector<Attribute, 2> results; 2032 if (failed(map().constantFold(operands, results))) 2033 return {}; 2034 2035 // Compute and return max of folded map results. 2036 int64_t max = std::numeric_limits<int64_t>::min(); 2037 int maxIndex = -1; 2038 for (unsigned i = 0, e = results.size(); i < e; ++i) { 2039 auto intAttr = results[i].cast<IntegerAttr>(); 2040 if (intAttr.getInt() > max) { 2041 max = intAttr.getInt(); 2042 maxIndex = i; 2043 } 2044 } 2045 if (maxIndex < 0) 2046 return {}; 2047 return results[maxIndex]; 2048 } 2049 2050 void AffineMaxOp::getCanonicalizationPatterns( 2051 OwningRewritePatternList &patterns, MLIRContext *context) { 2052 patterns.insert<SimplifyAffineOp<AffineMaxOp>>(context); 2053 } 2054 2055 //===----------------------------------------------------------------------===// 2056 // AffinePrefetchOp 2057 //===----------------------------------------------------------------------===// 2058 2059 // 2060 // affine.prefetch %0[%i, %j + 5], read, locality<3>, data : memref<400x400xi32> 2061 // 2062 static ParseResult parseAffinePrefetchOp(OpAsmParser &parser, 2063 OperationState &result) { 2064 auto &builder = parser.getBuilder(); 2065 auto indexTy = builder.getIndexType(); 2066 2067 MemRefType type; 2068 OpAsmParser::OperandType memrefInfo; 2069 IntegerAttr hintInfo; 2070 auto i32Type = parser.getBuilder().getIntegerType(32); 2071 StringRef readOrWrite, cacheType; 2072 2073 AffineMapAttr mapAttr; 2074 SmallVector<OpAsmParser::OperandType, 1> mapOperands; 2075 if (parser.parseOperand(memrefInfo) || 2076 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, 2077 AffinePrefetchOp::getMapAttrName(), 2078 result.attributes) || 2079 parser.parseComma() || parser.parseKeyword(&readOrWrite) || 2080 parser.parseComma() || parser.parseKeyword("locality") || 2081 parser.parseLess() || 2082 parser.parseAttribute(hintInfo, i32Type, 2083 AffinePrefetchOp::getLocalityHintAttrName(), 2084 result.attributes) || 2085 parser.parseGreater() || parser.parseComma() || 2086 parser.parseKeyword(&cacheType) || 2087 parser.parseOptionalAttrDict(result.attributes) || 2088 parser.parseColonType(type) || 2089 parser.resolveOperand(memrefInfo, type, result.operands) || 2090 parser.resolveOperands(mapOperands, indexTy, result.operands)) 2091 return failure(); 2092 2093 if (!readOrWrite.equals("read") && !readOrWrite.equals("write")) 2094 return parser.emitError(parser.getNameLoc(), 2095 "rw specifier has to be 'read' or 'write'"); 2096 result.addAttribute( 2097 AffinePrefetchOp::getIsWriteAttrName(), 2098 parser.getBuilder().getBoolAttr(readOrWrite.equals("write"))); 2099 2100 if (!cacheType.equals("data") && !cacheType.equals("instr")) 2101 return parser.emitError(parser.getNameLoc(), 2102 "cache type has to be 'data' or 'instr'"); 2103 2104 result.addAttribute( 2105 AffinePrefetchOp::getIsDataCacheAttrName(), 2106 parser.getBuilder().getBoolAttr(cacheType.equals("data"))); 2107 2108 return success(); 2109 } 2110 2111 static void print(OpAsmPrinter &p, AffinePrefetchOp op) { 2112 p << AffinePrefetchOp::getOperationName() << " " << op.memref() << '['; 2113 AffineMapAttr mapAttr = op.getAttrOfType<AffineMapAttr>(op.getMapAttrName()); 2114 if (mapAttr) { 2115 SmallVector<Value, 2> operands(op.getMapOperands()); 2116 p.printAffineMapOfSSAIds(mapAttr, operands); 2117 } 2118 p << ']' << ", " << (op.isWrite() ? "write" : "read") << ", " 2119 << "locality<" << op.localityHint() << ">, " 2120 << (op.isDataCache() ? "data" : "instr"); 2121 p.printOptionalAttrDict( 2122 op.getAttrs(), 2123 /*elidedAttrs=*/{op.getMapAttrName(), op.getLocalityHintAttrName(), 2124 op.getIsDataCacheAttrName(), op.getIsWriteAttrName()}); 2125 p << " : " << op.getMemRefType(); 2126 } 2127 2128 static LogicalResult verify(AffinePrefetchOp op) { 2129 auto mapAttr = op.getAttrOfType<AffineMapAttr>(op.getMapAttrName()); 2130 if (mapAttr) { 2131 AffineMap map = mapAttr.getValue(); 2132 if (map.getNumResults() != op.getMemRefType().getRank()) 2133 return op.emitOpError("affine.prefetch affine map num results must equal" 2134 " memref rank"); 2135 if (map.getNumInputs() + 1 != op.getNumOperands()) 2136 return op.emitOpError("too few operands"); 2137 } else { 2138 if (op.getNumOperands() != 1) 2139 return op.emitOpError("too few operands"); 2140 } 2141 2142 for (auto idx : op.getMapOperands()) { 2143 if (!isValidAffineIndexOperand(idx)) 2144 return op.emitOpError("index must be a dimension or symbol identifier"); 2145 } 2146 return success(); 2147 } 2148 2149 void AffinePrefetchOp::getCanonicalizationPatterns( 2150 OwningRewritePatternList &results, MLIRContext *context) { 2151 // prefetch(memrefcast) -> prefetch 2152 results.insert<SimplifyAffineOp<AffinePrefetchOp>>(context); 2153 } 2154 2155 LogicalResult AffinePrefetchOp::fold(ArrayRef<Attribute> cstOperands, 2156 SmallVectorImpl<OpFoldResult> &results) { 2157 /// prefetch(memrefcast) -> prefetch 2158 return foldMemRefCast(*this); 2159 } 2160 2161 //===----------------------------------------------------------------------===// 2162 // AffineParallelOp 2163 //===----------------------------------------------------------------------===// 2164 2165 void AffineParallelOp::build(Builder *builder, OperationState &result, 2166 ArrayRef<int64_t> ranges) { 2167 SmallVector<AffineExpr, 8> lbExprs(ranges.size(), 2168 builder->getAffineConstantExpr(0)); 2169 auto lbMap = AffineMap::get(0, 0, lbExprs, builder->getContext()); 2170 SmallVector<AffineExpr, 8> ubExprs; 2171 for (int64_t range : ranges) 2172 ubExprs.push_back(builder->getAffineConstantExpr(range)); 2173 auto ubMap = AffineMap::get(0, 0, ubExprs, builder->getContext()); 2174 build(builder, result, lbMap, {}, ubMap, {}); 2175 } 2176 2177 void AffineParallelOp::build(Builder *builder, OperationState &result, 2178 AffineMap lbMap, ValueRange lbArgs, 2179 AffineMap ubMap, ValueRange ubArgs) { 2180 auto numDims = lbMap.getNumResults(); 2181 // Verify that the dimensionality of both maps are the same. 2182 assert(numDims == ubMap.getNumResults() && 2183 "num dims and num results mismatch"); 2184 // Make default step sizes of 1. 2185 SmallVector<int64_t, 8> steps(numDims, 1); 2186 build(builder, result, lbMap, lbArgs, ubMap, ubArgs, steps); 2187 } 2188 2189 void AffineParallelOp::build(Builder *builder, OperationState &result, 2190 AffineMap lbMap, ValueRange lbArgs, 2191 AffineMap ubMap, ValueRange ubArgs, 2192 ArrayRef<int64_t> steps) { 2193 auto numDims = lbMap.getNumResults(); 2194 // Verify that the dimensionality of the maps matches the number of steps. 2195 assert(numDims == ubMap.getNumResults() && 2196 "num dims and num results mismatch"); 2197 assert(numDims == steps.size() && "num dims and num steps mismatch"); 2198 result.addAttribute(getLowerBoundsMapAttrName(), AffineMapAttr::get(lbMap)); 2199 result.addAttribute(getUpperBoundsMapAttrName(), AffineMapAttr::get(ubMap)); 2200 result.addAttribute(getStepsAttrName(), builder->getI64ArrayAttr(steps)); 2201 result.addOperands(lbArgs); 2202 result.addOperands(ubArgs); 2203 // Create a region and a block for the body. 2204 auto bodyRegion = result.addRegion(); 2205 auto body = new Block(); 2206 // Add all the block arguments. 2207 for (unsigned i = 0; i < numDims; ++i) 2208 body->addArgument(IndexType::get(builder->getContext())); 2209 bodyRegion->push_back(body); 2210 ensureTerminator(*bodyRegion, *builder, result.location); 2211 } 2212 2213 unsigned AffineParallelOp::getNumDims() { return steps().size(); } 2214 2215 AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() { 2216 return getOperands().take_front(lowerBoundsMap().getNumInputs()); 2217 } 2218 2219 AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() { 2220 return getOperands().drop_front(lowerBoundsMap().getNumInputs()); 2221 } 2222 2223 AffineValueMap AffineParallelOp::getLowerBoundsValueMap() { 2224 return AffineValueMap(lowerBoundsMap(), getLowerBoundsOperands()); 2225 } 2226 2227 AffineValueMap AffineParallelOp::getUpperBoundsValueMap() { 2228 return AffineValueMap(upperBoundsMap(), getUpperBoundsOperands()); 2229 } 2230 2231 AffineValueMap AffineParallelOp::getRangesValueMap() { 2232 AffineValueMap out; 2233 AffineValueMap::difference(getUpperBoundsValueMap(), getLowerBoundsValueMap(), 2234 &out); 2235 return out; 2236 } 2237 2238 Optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() { 2239 // Try to convert all the ranges to constant expressions. 2240 SmallVector<int64_t, 8> out; 2241 AffineValueMap rangesValueMap = getRangesValueMap(); 2242 out.reserve(rangesValueMap.getNumResults()); 2243 for (unsigned i = 0, e = rangesValueMap.getNumResults(); i < e; ++i) { 2244 auto expr = rangesValueMap.getResult(i); 2245 auto cst = expr.dyn_cast<AffineConstantExpr>(); 2246 if (!cst) 2247 return llvm::None; 2248 out.push_back(cst.getValue()); 2249 } 2250 return out; 2251 } 2252 2253 Block *AffineParallelOp::getBody() { return ®ion().front(); } 2254 2255 OpBuilder AffineParallelOp::getBodyBuilder() { 2256 return OpBuilder(getBody(), std::prev(getBody()->end())); 2257 } 2258 2259 void AffineParallelOp::setSteps(ArrayRef<int64_t> newSteps) { 2260 assert(newSteps.size() == getNumDims() && "steps & num dims mismatch"); 2261 setAttr(getStepsAttrName(), getBodyBuilder().getI64ArrayAttr(newSteps)); 2262 } 2263 2264 static LogicalResult verify(AffineParallelOp op) { 2265 auto numDims = op.getNumDims(); 2266 if (op.lowerBoundsMap().getNumResults() != numDims || 2267 op.upperBoundsMap().getNumResults() != numDims || 2268 op.steps().size() != numDims || 2269 op.getBody()->getNumArguments() != numDims) { 2270 return op.emitOpError("region argument count and num results of upper " 2271 "bounds, lower bounds, and steps must all match"); 2272 } 2273 // Verify that the bound operands are valid dimension/symbols. 2274 /// Lower bounds. 2275 if (failed(verifyDimAndSymbolIdentifiers(op, op.getLowerBoundsOperands(), 2276 op.lowerBoundsMap().getNumDims()))) 2277 return failure(); 2278 /// Upper bounds. 2279 if (failed(verifyDimAndSymbolIdentifiers(op, op.getUpperBoundsOperands(), 2280 op.upperBoundsMap().getNumDims()))) 2281 return failure(); 2282 return success(); 2283 } 2284 2285 static void print(OpAsmPrinter &p, AffineParallelOp op) { 2286 p << op.getOperationName() << " (" << op.getBody()->getArguments() << ") = ("; 2287 p.printAffineMapOfSSAIds(op.lowerBoundsMapAttr(), 2288 op.getLowerBoundsOperands()); 2289 p << ") to ("; 2290 p.printAffineMapOfSSAIds(op.upperBoundsMapAttr(), 2291 op.getUpperBoundsOperands()); 2292 p << ')'; 2293 SmallVector<int64_t, 4> steps; 2294 bool elideSteps = true; 2295 for (auto attr : op.steps()) { 2296 auto step = attr.cast<IntegerAttr>().getInt(); 2297 elideSteps &= (step == 1); 2298 steps.push_back(step); 2299 } 2300 if (!elideSteps) { 2301 p << " step ("; 2302 llvm::interleaveComma(steps, p); 2303 p << ')'; 2304 } 2305 p.printRegion(op.region(), /*printEntryBlockArgs=*/false, 2306 /*printBlockTerminators=*/false); 2307 p.printOptionalAttrDict( 2308 op.getAttrs(), 2309 /*elidedAttrs=*/{AffineParallelOp::getLowerBoundsMapAttrName(), 2310 AffineParallelOp::getUpperBoundsMapAttrName(), 2311 AffineParallelOp::getStepsAttrName()}); 2312 } 2313 2314 // 2315 // operation ::= `affine.parallel` `(` ssa-ids `)` `=` `(` map-of-ssa-ids `)` 2316 // `to` `(` map-of-ssa-ids `)` steps? region attr-dict? 2317 // steps ::= `steps` `(` integer-literals `)` 2318 // 2319 static ParseResult parseAffineParallelOp(OpAsmParser &parser, 2320 OperationState &result) { 2321 auto &builder = parser.getBuilder(); 2322 auto indexType = builder.getIndexType(); 2323 AffineMapAttr lowerBoundsAttr, upperBoundsAttr; 2324 SmallVector<OpAsmParser::OperandType, 4> ivs; 2325 SmallVector<OpAsmParser::OperandType, 4> lowerBoundsMapOperands; 2326 SmallVector<OpAsmParser::OperandType, 4> upperBoundsMapOperands; 2327 if (parser.parseRegionArgumentList(ivs, /*requiredOperandCount=*/-1, 2328 OpAsmParser::Delimiter::Paren) || 2329 parser.parseEqual() || 2330 parser.parseAffineMapOfSSAIds( 2331 lowerBoundsMapOperands, lowerBoundsAttr, 2332 AffineParallelOp::getLowerBoundsMapAttrName(), result.attributes, 2333 OpAsmParser::Delimiter::Paren) || 2334 parser.resolveOperands(lowerBoundsMapOperands, indexType, 2335 result.operands) || 2336 parser.parseKeyword("to") || 2337 parser.parseAffineMapOfSSAIds( 2338 upperBoundsMapOperands, upperBoundsAttr, 2339 AffineParallelOp::getUpperBoundsMapAttrName(), result.attributes, 2340 OpAsmParser::Delimiter::Paren) || 2341 parser.resolveOperands(upperBoundsMapOperands, indexType, 2342 result.operands)) 2343 return failure(); 2344 2345 AffineMapAttr stepsMapAttr; 2346 SmallVector<NamedAttribute, 1> stepsAttrs; 2347 SmallVector<OpAsmParser::OperandType, 4> stepsMapOperands; 2348 if (failed(parser.parseOptionalKeyword("step"))) { 2349 SmallVector<int64_t, 4> steps(ivs.size(), 1); 2350 result.addAttribute(AffineParallelOp::getStepsAttrName(), 2351 builder.getI64ArrayAttr(steps)); 2352 } else { 2353 if (parser.parseAffineMapOfSSAIds(stepsMapOperands, stepsMapAttr, 2354 AffineParallelOp::getStepsAttrName(), 2355 stepsAttrs, 2356 OpAsmParser::Delimiter::Paren)) 2357 return failure(); 2358 2359 // Convert steps from an AffineMap into an I64ArrayAttr. 2360 SmallVector<int64_t, 4> steps; 2361 auto stepsMap = stepsMapAttr.getValue(); 2362 for (const auto &result : stepsMap.getResults()) { 2363 auto constExpr = result.dyn_cast<AffineConstantExpr>(); 2364 if (!constExpr) 2365 return parser.emitError(parser.getNameLoc(), 2366 "steps must be constant integers"); 2367 steps.push_back(constExpr.getValue()); 2368 } 2369 result.addAttribute(AffineParallelOp::getStepsAttrName(), 2370 builder.getI64ArrayAttr(steps)); 2371 } 2372 2373 // Now parse the body. 2374 Region *body = result.addRegion(); 2375 SmallVector<Type, 4> types(ivs.size(), indexType); 2376 if (parser.parseRegion(*body, ivs, types) || 2377 parser.parseOptionalAttrDict(result.attributes)) 2378 return failure(); 2379 2380 // Add a terminator if none was parsed. 2381 AffineParallelOp::ensureTerminator(*body, builder, result.location); 2382 return success(); 2383 } 2384 2385 //===----------------------------------------------------------------------===// 2386 // TableGen'd op method definitions 2387 //===----------------------------------------------------------------------===// 2388 2389 #define GET_OP_CLASSES 2390 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc" 2391