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/MemRef/IR/MemRef.h" 12 #include "mlir/Dialect/Tensor/IR/Tensor.h" 13 #include "mlir/IR/AffineExprVisitor.h" 14 #include "mlir/IR/BlockAndValueMapping.h" 15 #include "mlir/IR/IntegerSet.h" 16 #include "mlir/IR/Matchers.h" 17 #include "mlir/IR/OpDefinition.h" 18 #include "mlir/IR/PatternMatch.h" 19 #include "mlir/Transforms/InliningUtils.h" 20 #include "llvm/ADT/SmallBitVector.h" 21 #include "llvm/ADT/TypeSwitch.h" 22 #include "llvm/Support/Debug.h" 23 24 using namespace mlir; 25 26 #define DEBUG_TYPE "affine-analysis" 27 28 #include "mlir/Dialect/Affine/IR/AffineOpsDialect.cpp.inc" 29 30 /// A utility function to check if a value is defined at the top level of 31 /// `region` or is an argument of `region`. A value of index type defined at the 32 /// top level of a `AffineScope` region is always a valid symbol for all 33 /// uses in that region. 34 bool mlir::isTopLevelValue(Value value, Region *region) { 35 if (auto arg = value.dyn_cast<BlockArgument>()) 36 return arg.getParentRegion() == region; 37 return value.getDefiningOp()->getParentRegion() == region; 38 } 39 40 /// Checks if `value` known to be a legal affine dimension or symbol in `src` 41 /// region remains legal if the operation that uses it is inlined into `dest` 42 /// with the given value mapping. `legalityCheck` is either `isValidDim` or 43 /// `isValidSymbol`, depending on the value being required to remain a valid 44 /// dimension or symbol. 45 static bool 46 remainsLegalAfterInline(Value value, Region *src, Region *dest, 47 const BlockAndValueMapping &mapping, 48 function_ref<bool(Value, Region *)> legalityCheck) { 49 // If the value is a valid dimension for any other reason than being 50 // a top-level value, it will remain valid: constants get inlined 51 // with the function, transitive affine applies also get inlined and 52 // will be checked themselves, etc. 53 if (!isTopLevelValue(value, src)) 54 return true; 55 56 // If it's a top-level value because it's a block operand, i.e. a 57 // function argument, check whether the value replacing it after 58 // inlining is a valid dimension in the new region. 59 if (value.isa<BlockArgument>()) 60 return legalityCheck(mapping.lookup(value), dest); 61 62 // If it's a top-level value because it's defined in the region, 63 // it can only be inlined if the defining op is a constant or a 64 // `dim`, which can appear anywhere and be valid, since the defining 65 // op won't be top-level anymore after inlining. 66 Attribute operandCst; 67 return matchPattern(value.getDefiningOp(), m_Constant(&operandCst)) || 68 value.getDefiningOp<memref::DimOp>() || 69 value.getDefiningOp<tensor::DimOp>(); 70 } 71 72 /// Checks if all values known to be legal affine dimensions or symbols in `src` 73 /// remain so if their respective users are inlined into `dest`. 74 static bool 75 remainsLegalAfterInline(ValueRange values, Region *src, Region *dest, 76 const BlockAndValueMapping &mapping, 77 function_ref<bool(Value, Region *)> legalityCheck) { 78 return llvm::all_of(values, [&](Value v) { 79 return remainsLegalAfterInline(v, src, dest, mapping, legalityCheck); 80 }); 81 } 82 83 /// Checks if an affine read or write operation remains legal after inlining 84 /// from `src` to `dest`. 85 template <typename OpTy> 86 static bool remainsLegalAfterInline(OpTy op, Region *src, Region *dest, 87 const BlockAndValueMapping &mapping) { 88 static_assert(llvm::is_one_of<OpTy, AffineReadOpInterface, 89 AffineWriteOpInterface>::value, 90 "only ops with affine read/write interface are supported"); 91 92 AffineMap map = op.getAffineMap(); 93 ValueRange dimOperands = op.getMapOperands().take_front(map.getNumDims()); 94 ValueRange symbolOperands = 95 op.getMapOperands().take_back(map.getNumSymbols()); 96 if (!remainsLegalAfterInline( 97 dimOperands, src, dest, mapping, 98 static_cast<bool (*)(Value, Region *)>(isValidDim))) 99 return false; 100 if (!remainsLegalAfterInline( 101 symbolOperands, src, dest, mapping, 102 static_cast<bool (*)(Value, Region *)>(isValidSymbol))) 103 return false; 104 return true; 105 } 106 107 /// Checks if an affine apply operation remains legal after inlining from `src` 108 /// to `dest`. 109 // Use "unused attribute" marker to silence clang-tidy warning stemming from 110 // the inability to see through "llvm::TypeSwitch". 111 template <> 112 bool LLVM_ATTRIBUTE_UNUSED 113 remainsLegalAfterInline(AffineApplyOp op, Region *src, Region *dest, 114 const BlockAndValueMapping &mapping) { 115 // If it's a valid dimension, we need to check that it remains so. 116 if (isValidDim(op.getResult(), src)) 117 return remainsLegalAfterInline( 118 op.getMapOperands(), src, dest, mapping, 119 static_cast<bool (*)(Value, Region *)>(isValidDim)); 120 121 // Otherwise it must be a valid symbol, check that it remains so. 122 return remainsLegalAfterInline( 123 op.getMapOperands(), src, dest, mapping, 124 static_cast<bool (*)(Value, Region *)>(isValidSymbol)); 125 } 126 127 //===----------------------------------------------------------------------===// 128 // AffineDialect Interfaces 129 //===----------------------------------------------------------------------===// 130 131 namespace { 132 /// This class defines the interface for handling inlining with affine 133 /// operations. 134 struct AffineInlinerInterface : public DialectInlinerInterface { 135 using DialectInlinerInterface::DialectInlinerInterface; 136 137 //===--------------------------------------------------------------------===// 138 // Analysis Hooks 139 //===--------------------------------------------------------------------===// 140 141 /// Returns true if the given region 'src' can be inlined into the region 142 /// 'dest' that is attached to an operation registered to the current dialect. 143 /// 'wouldBeCloned' is set if the region is cloned into its new location 144 /// rather than moved, indicating there may be other users. 145 bool isLegalToInline(Region *dest, Region *src, bool wouldBeCloned, 146 BlockAndValueMapping &valueMapping) const final { 147 // We can inline into affine loops and conditionals if this doesn't break 148 // affine value categorization rules. 149 Operation *destOp = dest->getParentOp(); 150 if (!isa<AffineParallelOp, AffineForOp, AffineIfOp>(destOp)) 151 return false; 152 153 // Multi-block regions cannot be inlined into affine constructs, all of 154 // which require single-block regions. 155 if (!llvm::hasSingleElement(*src)) 156 return false; 157 158 // Side-effecting operations that the affine dialect cannot understand 159 // should not be inlined. 160 Block &srcBlock = src->front(); 161 for (Operation &op : srcBlock) { 162 // Ops with no side effects are fine, 163 if (auto iface = dyn_cast<MemoryEffectOpInterface>(op)) { 164 if (iface.hasNoEffect()) 165 continue; 166 } 167 168 // Assuming the inlined region is valid, we only need to check if the 169 // inlining would change it. 170 bool remainsValid = 171 llvm::TypeSwitch<Operation *, bool>(&op) 172 .Case<AffineApplyOp, AffineReadOpInterface, 173 AffineWriteOpInterface>([&](auto op) { 174 return remainsLegalAfterInline(op, src, dest, valueMapping); 175 }) 176 .Default([](Operation *) { 177 // Conservatively disallow inlining ops we cannot reason about. 178 return false; 179 }); 180 181 if (!remainsValid) 182 return false; 183 } 184 185 return true; 186 } 187 188 /// Returns true if the given operation 'op', that is registered to this 189 /// dialect, can be inlined into the given region, false otherwise. 190 bool isLegalToInline(Operation *op, Region *region, bool wouldBeCloned, 191 BlockAndValueMapping &valueMapping) const final { 192 // Always allow inlining affine operations into a region that is marked as 193 // affine scope, or into affine loops and conditionals. There are some edge 194 // cases when inlining *into* affine structures, but that is handled in the 195 // other 'isLegalToInline' hook above. 196 Operation *parentOp = region->getParentOp(); 197 return parentOp->hasTrait<OpTrait::AffineScope>() || 198 isa<AffineForOp, AffineParallelOp, AffineIfOp>(parentOp); 199 } 200 201 /// Affine regions should be analyzed recursively. 202 bool shouldAnalyzeRecursively(Operation *op) const final { return true; } 203 }; 204 } // namespace 205 206 //===----------------------------------------------------------------------===// 207 // AffineDialect 208 //===----------------------------------------------------------------------===// 209 210 void AffineDialect::initialize() { 211 addOperations<AffineDmaStartOp, AffineDmaWaitOp, 212 #define GET_OP_LIST 213 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc" 214 >(); 215 addInterfaces<AffineInlinerInterface>(); 216 } 217 218 /// Materialize a single constant operation from a given attribute value with 219 /// the desired resultant type. 220 Operation *AffineDialect::materializeConstant(OpBuilder &builder, 221 Attribute value, Type type, 222 Location loc) { 223 return builder.create<arith::ConstantOp>(loc, type, value); 224 } 225 226 /// A utility function to check if a value is defined at the top level of an 227 /// op with trait `AffineScope`. If the value is defined in an unlinked region, 228 /// conservatively assume it is not top-level. A value of index type defined at 229 /// the top level is always a valid symbol. 230 bool mlir::isTopLevelValue(Value value) { 231 if (auto arg = value.dyn_cast<BlockArgument>()) { 232 // The block owning the argument may be unlinked, e.g. when the surrounding 233 // region has not yet been attached to an Op, at which point the parent Op 234 // is null. 235 Operation *parentOp = arg.getOwner()->getParentOp(); 236 return parentOp && parentOp->hasTrait<OpTrait::AffineScope>(); 237 } 238 // The defining Op may live in an unlinked block so its parent Op may be null. 239 Operation *parentOp = value.getDefiningOp()->getParentOp(); 240 return parentOp && parentOp->hasTrait<OpTrait::AffineScope>(); 241 } 242 243 /// Returns the closest region enclosing `op` that is held by an operation with 244 /// trait `AffineScope`; `nullptr` if there is no such region. 245 Region *mlir::getAffineScope(Operation *op) { 246 auto *curOp = op; 247 while (auto *parentOp = curOp->getParentOp()) { 248 if (parentOp->hasTrait<OpTrait::AffineScope>()) 249 return curOp->getParentRegion(); 250 curOp = parentOp; 251 } 252 return nullptr; 253 } 254 255 // A Value can be used as a dimension id iff it meets one of the following 256 // conditions: 257 // *) It is valid as a symbol. 258 // *) It is an induction variable. 259 // *) It is the result of affine apply operation with dimension id arguments. 260 bool mlir::isValidDim(Value value) { 261 // The value must be an index type. 262 if (!value.getType().isIndex()) 263 return false; 264 265 if (auto *defOp = value.getDefiningOp()) 266 return isValidDim(value, getAffineScope(defOp)); 267 268 // This value has to be a block argument for an op that has the 269 // `AffineScope` trait or for an affine.for or affine.parallel. 270 auto *parentOp = value.cast<BlockArgument>().getOwner()->getParentOp(); 271 return parentOp && (parentOp->hasTrait<OpTrait::AffineScope>() || 272 isa<AffineForOp, AffineParallelOp>(parentOp)); 273 } 274 275 // Value can be used as a dimension id iff it meets one of the following 276 // conditions: 277 // *) It is valid as a symbol. 278 // *) It is an induction variable. 279 // *) It is the result of an affine apply operation with dimension id operands. 280 bool mlir::isValidDim(Value value, Region *region) { 281 // The value must be an index type. 282 if (!value.getType().isIndex()) 283 return false; 284 285 // All valid symbols are okay. 286 if (isValidSymbol(value, region)) 287 return true; 288 289 auto *op = value.getDefiningOp(); 290 if (!op) { 291 // This value has to be a block argument for an affine.for or an 292 // affine.parallel. 293 auto *parentOp = value.cast<BlockArgument>().getOwner()->getParentOp(); 294 return isa<AffineForOp, AffineParallelOp>(parentOp); 295 } 296 297 // Affine apply operation is ok if all of its operands are ok. 298 if (auto applyOp = dyn_cast<AffineApplyOp>(op)) 299 return applyOp.isValidDim(region); 300 // The dim op is okay if its operand memref/tensor is defined at the top 301 // level. 302 if (auto dimOp = dyn_cast<memref::DimOp>(op)) 303 return isTopLevelValue(dimOp.getSource()); 304 if (auto dimOp = dyn_cast<tensor::DimOp>(op)) 305 return isTopLevelValue(dimOp.getSource()); 306 return false; 307 } 308 309 /// Returns true if the 'index' dimension of the `memref` defined by 310 /// `memrefDefOp` is a statically shaped one or defined using a valid symbol 311 /// for `region`. 312 template <typename AnyMemRefDefOp> 313 static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp, unsigned index, 314 Region *region) { 315 auto memRefType = memrefDefOp.getType(); 316 // Statically shaped. 317 if (!memRefType.isDynamicDim(index)) 318 return true; 319 // Get the position of the dimension among dynamic dimensions; 320 unsigned dynamicDimPos = memRefType.getDynamicDimIndex(index); 321 return isValidSymbol(*(memrefDefOp.getDynamicSizes().begin() + dynamicDimPos), 322 region); 323 } 324 325 /// Returns true if the result of the dim op is a valid symbol for `region`. 326 template <typename OpTy> 327 static bool isDimOpValidSymbol(OpTy dimOp, Region *region) { 328 // The dim op is okay if its source is defined at the top level. 329 if (isTopLevelValue(dimOp.getSource())) 330 return true; 331 332 // Conservatively handle remaining BlockArguments as non-valid symbols. 333 // E.g. scf.for iterArgs. 334 if (dimOp.getSource().template isa<BlockArgument>()) 335 return false; 336 337 // The dim op is also okay if its operand memref is a view/subview whose 338 // corresponding size is a valid symbol. 339 Optional<int64_t> index = dimOp.getConstantIndex(); 340 assert(index.has_value() && 341 "expect only `dim` operations with a constant index"); 342 int64_t i = index.getValue(); 343 return TypeSwitch<Operation *, bool>(dimOp.getSource().getDefiningOp()) 344 .Case<memref::ViewOp, memref::SubViewOp, memref::AllocOp>( 345 [&](auto op) { return isMemRefSizeValidSymbol(op, i, region); }) 346 .Default([](Operation *) { return false; }); 347 } 348 349 // A value can be used as a symbol (at all its use sites) iff it meets one of 350 // the following conditions: 351 // *) It is a constant. 352 // *) Its defining op or block arg appearance is immediately enclosed by an op 353 // with `AffineScope` trait. 354 // *) It is the result of an affine.apply operation with symbol operands. 355 // *) It is a result of the dim op on a memref whose corresponding size is a 356 // valid symbol. 357 bool mlir::isValidSymbol(Value value) { 358 if (!value) 359 return false; 360 361 // The value must be an index type. 362 if (!value.getType().isIndex()) 363 return false; 364 365 // Check that the value is a top level value. 366 if (isTopLevelValue(value)) 367 return true; 368 369 if (auto *defOp = value.getDefiningOp()) 370 return isValidSymbol(value, getAffineScope(defOp)); 371 372 return false; 373 } 374 375 /// A value can be used as a symbol for `region` iff it meets one of the 376 /// following conditions: 377 /// *) It is a constant. 378 /// *) It is the result of an affine apply operation with symbol arguments. 379 /// *) It is a result of the dim op on a memref whose corresponding size is 380 /// a valid symbol. 381 /// *) It is defined at the top level of 'region' or is its argument. 382 /// *) It dominates `region`'s parent op. 383 /// If `region` is null, conservatively assume the symbol definition scope does 384 /// not exist and only accept the values that would be symbols regardless of 385 /// the surrounding region structure, i.e. the first three cases above. 386 bool mlir::isValidSymbol(Value value, Region *region) { 387 // The value must be an index type. 388 if (!value.getType().isIndex()) 389 return false; 390 391 // A top-level value is a valid symbol. 392 if (region && ::isTopLevelValue(value, region)) 393 return true; 394 395 auto *defOp = value.getDefiningOp(); 396 if (!defOp) { 397 // A block argument that is not a top-level value is a valid symbol if it 398 // dominates region's parent op. 399 Operation *regionOp = region ? region->getParentOp() : nullptr; 400 if (regionOp && !regionOp->hasTrait<OpTrait::IsIsolatedFromAbove>()) 401 if (auto *parentOpRegion = region->getParentOp()->getParentRegion()) 402 return isValidSymbol(value, parentOpRegion); 403 return false; 404 } 405 406 // Constant operation is ok. 407 Attribute operandCst; 408 if (matchPattern(defOp, m_Constant(&operandCst))) 409 return true; 410 411 // Affine apply operation is ok if all of its operands are ok. 412 if (auto applyOp = dyn_cast<AffineApplyOp>(defOp)) 413 return applyOp.isValidSymbol(region); 414 415 // Dim op results could be valid symbols at any level. 416 if (auto dimOp = dyn_cast<memref::DimOp>(defOp)) 417 return isDimOpValidSymbol(dimOp, region); 418 if (auto dimOp = dyn_cast<tensor::DimOp>(defOp)) 419 return isDimOpValidSymbol(dimOp, region); 420 421 // Check for values dominating `region`'s parent op. 422 Operation *regionOp = region ? region->getParentOp() : nullptr; 423 if (regionOp && !regionOp->hasTrait<OpTrait::IsIsolatedFromAbove>()) 424 if (auto *parentRegion = region->getParentOp()->getParentRegion()) 425 return isValidSymbol(value, parentRegion); 426 427 return false; 428 } 429 430 // Returns true if 'value' is a valid index to an affine operation (e.g. 431 // affine.load, affine.store, affine.dma_start, affine.dma_wait) where 432 // `region` provides the polyhedral symbol scope. Returns false otherwise. 433 static bool isValidAffineIndexOperand(Value value, Region *region) { 434 return isValidDim(value, region) || isValidSymbol(value, region); 435 } 436 437 /// Prints dimension and symbol list. 438 static void printDimAndSymbolList(Operation::operand_iterator begin, 439 Operation::operand_iterator end, 440 unsigned numDims, OpAsmPrinter &printer) { 441 OperandRange operands(begin, end); 442 printer << '(' << operands.take_front(numDims) << ')'; 443 if (operands.size() > numDims) 444 printer << '[' << operands.drop_front(numDims) << ']'; 445 } 446 447 /// Parses dimension and symbol list and returns true if parsing failed. 448 ParseResult mlir::parseDimAndSymbolList(OpAsmParser &parser, 449 SmallVectorImpl<Value> &operands, 450 unsigned &numDims) { 451 SmallVector<OpAsmParser::UnresolvedOperand, 8> opInfos; 452 if (parser.parseOperandList(opInfos, OpAsmParser::Delimiter::Paren)) 453 return failure(); 454 // Store number of dimensions for validation by caller. 455 numDims = opInfos.size(); 456 457 // Parse the optional symbol operands. 458 auto indexTy = parser.getBuilder().getIndexType(); 459 return failure(parser.parseOperandList( 460 opInfos, OpAsmParser::Delimiter::OptionalSquare) || 461 parser.resolveOperands(opInfos, indexTy, operands)); 462 } 463 464 /// Utility function to verify that a set of operands are valid dimension and 465 /// symbol identifiers. The operands should be laid out such that the dimension 466 /// operands are before the symbol operands. This function returns failure if 467 /// there was an invalid operand. An operation is provided to emit any necessary 468 /// errors. 469 template <typename OpTy> 470 static LogicalResult 471 verifyDimAndSymbolIdentifiers(OpTy &op, Operation::operand_range operands, 472 unsigned numDims) { 473 unsigned opIt = 0; 474 for (auto operand : operands) { 475 if (opIt++ < numDims) { 476 if (!isValidDim(operand, getAffineScope(op))) 477 return op.emitOpError("operand cannot be used as a dimension id"); 478 } else if (!isValidSymbol(operand, getAffineScope(op))) { 479 return op.emitOpError("operand cannot be used as a symbol"); 480 } 481 } 482 return success(); 483 } 484 485 //===----------------------------------------------------------------------===// 486 // AffineApplyOp 487 //===----------------------------------------------------------------------===// 488 489 AffineValueMap AffineApplyOp::getAffineValueMap() { 490 return AffineValueMap(getAffineMap(), getOperands(), getResult()); 491 } 492 493 ParseResult AffineApplyOp::parse(OpAsmParser &parser, OperationState &result) { 494 auto &builder = parser.getBuilder(); 495 auto indexTy = builder.getIndexType(); 496 497 AffineMapAttr mapAttr; 498 unsigned numDims; 499 if (parser.parseAttribute(mapAttr, "map", result.attributes) || 500 parseDimAndSymbolList(parser, result.operands, numDims) || 501 parser.parseOptionalAttrDict(result.attributes)) 502 return failure(); 503 auto map = mapAttr.getValue(); 504 505 if (map.getNumDims() != numDims || 506 numDims + map.getNumSymbols() != result.operands.size()) { 507 return parser.emitError(parser.getNameLoc(), 508 "dimension or symbol index mismatch"); 509 } 510 511 result.types.append(map.getNumResults(), indexTy); 512 return success(); 513 } 514 515 void AffineApplyOp::print(OpAsmPrinter &p) { 516 p << " " << getMapAttr(); 517 printDimAndSymbolList(operand_begin(), operand_end(), 518 getAffineMap().getNumDims(), p); 519 p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{"map"}); 520 } 521 522 LogicalResult AffineApplyOp::verify() { 523 // Check input and output dimensions match. 524 AffineMap affineMap = getMap(); 525 526 // Verify that operand count matches affine map dimension and symbol count. 527 if (getNumOperands() != affineMap.getNumDims() + affineMap.getNumSymbols()) 528 return emitOpError( 529 "operand count and affine map dimension and symbol count must match"); 530 531 // Verify that the map only produces one result. 532 if (affineMap.getNumResults() != 1) 533 return emitOpError("mapping must produce one value"); 534 535 return success(); 536 } 537 538 // The result of the affine apply operation can be used as a dimension id if all 539 // its operands are valid dimension ids. 540 bool AffineApplyOp::isValidDim() { 541 return llvm::all_of(getOperands(), 542 [](Value op) { return mlir::isValidDim(op); }); 543 } 544 545 // The result of the affine apply operation can be used as a dimension id if all 546 // its operands are valid dimension ids with the parent operation of `region` 547 // defining the polyhedral scope for symbols. 548 bool AffineApplyOp::isValidDim(Region *region) { 549 return llvm::all_of(getOperands(), 550 [&](Value op) { return ::isValidDim(op, region); }); 551 } 552 553 // The result of the affine apply operation can be used as a symbol if all its 554 // operands are symbols. 555 bool AffineApplyOp::isValidSymbol() { 556 return llvm::all_of(getOperands(), 557 [](Value op) { return mlir::isValidSymbol(op); }); 558 } 559 560 // The result of the affine apply operation can be used as a symbol in `region` 561 // if all its operands are symbols in `region`. 562 bool AffineApplyOp::isValidSymbol(Region *region) { 563 return llvm::all_of(getOperands(), [&](Value operand) { 564 return mlir::isValidSymbol(operand, region); 565 }); 566 } 567 568 OpFoldResult AffineApplyOp::fold(ArrayRef<Attribute> operands) { 569 auto map = getAffineMap(); 570 571 // Fold dims and symbols to existing values. 572 auto expr = map.getResult(0); 573 if (auto dim = expr.dyn_cast<AffineDimExpr>()) 574 return getOperand(dim.getPosition()); 575 if (auto sym = expr.dyn_cast<AffineSymbolExpr>()) 576 return getOperand(map.getNumDims() + sym.getPosition()); 577 578 // Otherwise, default to folding the map. 579 SmallVector<Attribute, 1> result; 580 if (failed(map.constantFold(operands, result))) 581 return {}; 582 return result[0]; 583 } 584 585 /// Replace all occurrences of AffineExpr at position `pos` in `map` by the 586 /// defining AffineApplyOp expression and operands. 587 /// When `dimOrSymbolPosition < dims.size()`, AffineDimExpr@[pos] is replaced. 588 /// When `dimOrSymbolPosition >= dims.size()`, 589 /// AffineSymbolExpr@[pos - dims.size()] is replaced. 590 /// Mutate `map`,`dims` and `syms` in place as follows: 591 /// 1. `dims` and `syms` are only appended to. 592 /// 2. `map` dim and symbols are gradually shifted to higher positions. 593 /// 3. Old `dim` and `sym` entries are replaced by nullptr 594 /// This avoids the need for any bookkeeping. 595 static LogicalResult replaceDimOrSym(AffineMap *map, 596 unsigned dimOrSymbolPosition, 597 SmallVectorImpl<Value> &dims, 598 SmallVectorImpl<Value> &syms) { 599 bool isDimReplacement = (dimOrSymbolPosition < dims.size()); 600 unsigned pos = isDimReplacement ? dimOrSymbolPosition 601 : dimOrSymbolPosition - dims.size(); 602 Value &v = isDimReplacement ? dims[pos] : syms[pos]; 603 if (!v) 604 return failure(); 605 606 auto affineApply = v.getDefiningOp<AffineApplyOp>(); 607 if (!affineApply) 608 return failure(); 609 610 // At this point we will perform a replacement of `v`, set the entry in `dim` 611 // or `sym` to nullptr immediately. 612 v = nullptr; 613 614 // Compute the map, dims and symbols coming from the AffineApplyOp. 615 AffineMap composeMap = affineApply.getAffineMap(); 616 assert(composeMap.getNumResults() == 1 && "affine.apply with >1 results"); 617 AffineExpr composeExpr = 618 composeMap.shiftDims(dims.size()).shiftSymbols(syms.size()).getResult(0); 619 ValueRange composeDims = 620 affineApply.getMapOperands().take_front(composeMap.getNumDims()); 621 ValueRange composeSyms = 622 affineApply.getMapOperands().take_back(composeMap.getNumSymbols()); 623 624 // Append the dims and symbols where relevant and perform the replacement. 625 MLIRContext *ctx = map->getContext(); 626 AffineExpr toReplace = isDimReplacement ? getAffineDimExpr(pos, ctx) 627 : getAffineSymbolExpr(pos, ctx); 628 dims.append(composeDims.begin(), composeDims.end()); 629 syms.append(composeSyms.begin(), composeSyms.end()); 630 *map = map->replace(toReplace, composeExpr, dims.size(), syms.size()); 631 632 return success(); 633 } 634 635 /// Iterate over `operands` and fold away all those produced by an AffineApplyOp 636 /// iteratively. Perform canonicalization of map and operands as well as 637 /// AffineMap simplification. `map` and `operands` are mutated in place. 638 static void composeAffineMapAndOperands(AffineMap *map, 639 SmallVectorImpl<Value> *operands) { 640 if (map->getNumResults() == 0) { 641 canonicalizeMapAndOperands(map, operands); 642 *map = simplifyAffineMap(*map); 643 return; 644 } 645 646 MLIRContext *ctx = map->getContext(); 647 SmallVector<Value, 4> dims(operands->begin(), 648 operands->begin() + map->getNumDims()); 649 SmallVector<Value, 4> syms(operands->begin() + map->getNumDims(), 650 operands->end()); 651 652 // Iterate over dims and symbols coming from AffineApplyOp and replace until 653 // exhaustion. This iteratively mutates `map`, `dims` and `syms`. Both `dims` 654 // and `syms` can only increase by construction. 655 // The implementation uses a `while` loop to support the case of symbols 656 // that may be constructed from dims ;this may be overkill. 657 while (true) { 658 bool changed = false; 659 for (unsigned pos = 0; pos != dims.size() + syms.size(); ++pos) 660 if ((changed |= succeeded(replaceDimOrSym(map, pos, dims, syms)))) 661 break; 662 if (!changed) 663 break; 664 } 665 666 // Clear operands so we can fill them anew. 667 operands->clear(); 668 669 // At this point we may have introduced null operands, prune them out before 670 // canonicalizing map and operands. 671 unsigned nDims = 0, nSyms = 0; 672 SmallVector<AffineExpr, 4> dimReplacements, symReplacements; 673 dimReplacements.reserve(dims.size()); 674 symReplacements.reserve(syms.size()); 675 for (auto *container : {&dims, &syms}) { 676 bool isDim = (container == &dims); 677 auto &repls = isDim ? dimReplacements : symReplacements; 678 for (const auto &en : llvm::enumerate(*container)) { 679 Value v = en.value(); 680 if (!v) { 681 assert(isDim ? !map->isFunctionOfDim(en.index()) 682 : !map->isFunctionOfSymbol(en.index()) && 683 "map is function of unexpected expr@pos"); 684 repls.push_back(getAffineConstantExpr(0, ctx)); 685 continue; 686 } 687 repls.push_back(isDim ? getAffineDimExpr(nDims++, ctx) 688 : getAffineSymbolExpr(nSyms++, ctx)); 689 operands->push_back(v); 690 } 691 } 692 *map = map->replaceDimsAndSymbols(dimReplacements, symReplacements, nDims, 693 nSyms); 694 695 // Canonicalize and simplify before returning. 696 canonicalizeMapAndOperands(map, operands); 697 *map = simplifyAffineMap(*map); 698 } 699 700 void mlir::fullyComposeAffineMapAndOperands(AffineMap *map, 701 SmallVectorImpl<Value> *operands) { 702 while (llvm::any_of(*operands, [](Value v) { 703 return isa_and_nonnull<AffineApplyOp>(v.getDefiningOp()); 704 })) { 705 composeAffineMapAndOperands(map, operands); 706 } 707 } 708 709 /// Given a list of `OpFoldResult`, build the necessary operations to populate 710 /// `actualValues` with values produced by operations. In particular, for any 711 /// attribute-typed element in `values`, call the constant materializer 712 /// associated with the Affine dialect to produce an operation. 713 static void materializeConstants(OpBuilder &b, Location loc, 714 ArrayRef<OpFoldResult> values, 715 SmallVectorImpl<Operation *> &constants, 716 SmallVectorImpl<Value> &actualValues) { 717 actualValues.reserve(values.size()); 718 auto *dialect = b.getContext()->getLoadedDialect<AffineDialect>(); 719 for (OpFoldResult ofr : values) { 720 if (auto value = ofr.dyn_cast<Value>()) { 721 actualValues.push_back(value); 722 continue; 723 } 724 constants.push_back(dialect->materializeConstant(b, ofr.get<Attribute>(), 725 b.getIndexType(), loc)); 726 actualValues.push_back(constants.back()->getResult(0)); 727 } 728 } 729 730 /// Create an operation of the type provided as template argument and attempt to 731 /// fold it immediately. The operation is expected to have a builder taking 732 /// arbitrary `leadingArguments`, followed by a list of Value-typed `operands`. 733 /// The operation is also expected to always produce a single result. Return an 734 /// `OpFoldResult` containing the Attribute representing the folded constant if 735 /// complete folding was possible and a Value produced by the created operation 736 /// otherwise. 737 template <typename OpTy, typename... Args> 738 static std::enable_if_t<OpTy::template hasTrait<OpTrait::OneResult>(), 739 OpFoldResult> 740 createOrFold(RewriterBase &b, Location loc, ValueRange operands, 741 Args &&...leadingArguments) { 742 // Identify the constant operands and extract their values as attributes. 743 // Note that we cannot use the original values directly because the list of 744 // operands may have changed due to canonicalization and composition. 745 SmallVector<Attribute> constantOperands; 746 constantOperands.reserve(operands.size()); 747 for (Value operand : operands) { 748 IntegerAttr attr; 749 if (matchPattern(operand, m_Constant(&attr))) 750 constantOperands.push_back(attr); 751 else 752 constantOperands.push_back(nullptr); 753 } 754 755 // Create the operation and immediately attempt to fold it. On success, 756 // delete the operation and prepare the (unmaterialized) value for being 757 // returned. On failure, return the operation result value. 758 // TODO: arguably, the main folder (createOrFold) API should support this use 759 // case instead of indiscriminately materializing constants. 760 OpTy op = 761 b.create<OpTy>(loc, std::forward<Args>(leadingArguments)..., operands); 762 SmallVector<OpFoldResult, 1> foldResults; 763 if (succeeded(op->fold(constantOperands, foldResults)) && 764 !foldResults.empty()) { 765 b.eraseOp(op); 766 return foldResults.front(); 767 } 768 return op->getResult(0); 769 } 770 771 AffineApplyOp mlir::makeComposedAffineApply(OpBuilder &b, Location loc, 772 AffineMap map, 773 ValueRange operands) { 774 AffineMap normalizedMap = map; 775 SmallVector<Value, 8> normalizedOperands(operands.begin(), operands.end()); 776 composeAffineMapAndOperands(&normalizedMap, &normalizedOperands); 777 assert(normalizedMap); 778 return b.create<AffineApplyOp>(loc, normalizedMap, normalizedOperands); 779 } 780 781 AffineApplyOp mlir::makeComposedAffineApply(OpBuilder &b, Location loc, 782 AffineExpr e, ValueRange values) { 783 return makeComposedAffineApply( 784 b, loc, AffineMap::inferFromExprList(ArrayRef<AffineExpr>{e}).front(), 785 values); 786 } 787 788 OpFoldResult 789 mlir::makeComposedFoldedAffineApply(RewriterBase &b, Location loc, 790 AffineMap map, 791 ArrayRef<OpFoldResult> operands) { 792 assert(map.getNumResults() == 1 && "building affine.apply with !=1 result"); 793 794 SmallVector<Operation *> constants; 795 SmallVector<Value> actualValues; 796 materializeConstants(b, loc, operands, constants, actualValues); 797 composeAffineMapAndOperands(&map, &actualValues); 798 OpFoldResult result = createOrFold<AffineApplyOp>(b, loc, actualValues, map); 799 if (result.is<Attribute>()) { 800 for (Operation *op : constants) 801 b.eraseOp(op); 802 } 803 return result; 804 } 805 806 OpFoldResult 807 mlir::makeComposedFoldedAffineApply(RewriterBase &b, Location loc, 808 AffineExpr expr, 809 ArrayRef<OpFoldResult> operands) { 810 return makeComposedFoldedAffineApply( 811 b, loc, AffineMap::inferFromExprList(ArrayRef<AffineExpr>{expr}).front(), 812 operands); 813 } 814 815 /// Composes the given affine map with the given list of operands, pulling in 816 /// the maps from any affine.apply operations that supply the operands. 817 static void composeMultiResultAffineMap(AffineMap &map, 818 SmallVectorImpl<Value> &operands) { 819 // Compose and canonicalize each expression in the map individually because 820 // composition only applies to single-result maps, collecting potentially 821 // duplicate operands in a single list with shifted dimensions and symbols. 822 SmallVector<Value> dims, symbols; 823 SmallVector<AffineExpr> exprs; 824 for (unsigned i : llvm::seq<unsigned>(0, map.getNumResults())) { 825 SmallVector<Value> submapOperands(operands.begin(), operands.end()); 826 AffineMap submap = map.getSubMap({i}); 827 fullyComposeAffineMapAndOperands(&submap, &submapOperands); 828 canonicalizeMapAndOperands(&submap, &submapOperands); 829 unsigned numNewDims = submap.getNumDims(); 830 submap = submap.shiftDims(dims.size()).shiftSymbols(symbols.size()); 831 llvm::append_range(dims, 832 ArrayRef<Value>(submapOperands).take_front(numNewDims)); 833 llvm::append_range(symbols, 834 ArrayRef<Value>(submapOperands).drop_front(numNewDims)); 835 exprs.push_back(submap.getResult(0)); 836 } 837 838 // Canonicalize the map created from composed expressions to deduplicate the 839 // dimension and symbol operands. 840 operands = llvm::to_vector(llvm::concat<Value>(dims, symbols)); 841 map = AffineMap::get(dims.size(), symbols.size(), exprs, map.getContext()); 842 canonicalizeMapAndOperands(&map, &operands); 843 } 844 845 Value mlir::makeComposedAffineMin(OpBuilder &b, Location loc, AffineMap map, 846 ValueRange operands) { 847 SmallVector<Value> allOperands = llvm::to_vector(operands); 848 composeMultiResultAffineMap(map, allOperands); 849 return b.createOrFold<AffineMinOp>(loc, b.getIndexType(), map, allOperands); 850 } 851 852 OpFoldResult 853 mlir::makeComposedFoldedAffineMin(RewriterBase &b, Location loc, AffineMap map, 854 ArrayRef<OpFoldResult> operands) { 855 SmallVector<Operation *> constants; 856 SmallVector<Value> actualValues; 857 materializeConstants(b, loc, operands, constants, actualValues); 858 composeMultiResultAffineMap(map, actualValues); 859 OpFoldResult result = 860 createOrFold<AffineMinOp>(b, loc, actualValues, b.getIndexType(), map); 861 if (result.is<Attribute>()) { 862 for (Operation *op : constants) 863 b.eraseOp(op); 864 } 865 return result; 866 } 867 868 /// Fully compose map with operands and canonicalize the result. 869 /// Return the `createOrFold`'ed AffineApply op. 870 static Value createFoldedComposedAffineApply(OpBuilder &b, Location loc, 871 AffineMap map, 872 ValueRange operandsRef) { 873 SmallVector<Value, 4> operands(operandsRef.begin(), operandsRef.end()); 874 fullyComposeAffineMapAndOperands(&map, &operands); 875 canonicalizeMapAndOperands(&map, &operands); 876 return b.createOrFold<AffineApplyOp>(loc, map, operands); 877 } 878 879 SmallVector<Value, 4> mlir::applyMapToValues(OpBuilder &b, Location loc, 880 AffineMap map, ValueRange values) { 881 SmallVector<Value, 4> res; 882 res.reserve(map.getNumResults()); 883 unsigned numDims = map.getNumDims(), numSym = map.getNumSymbols(); 884 // For each `expr` in `map`, applies the `expr` to the values extracted from 885 // ranges. If the resulting application can be folded into a Value, the 886 // folding occurs eagerly. 887 for (auto expr : map.getResults()) { 888 AffineMap map = AffineMap::get(numDims, numSym, expr); 889 res.push_back(createFoldedComposedAffineApply(b, loc, map, values)); 890 } 891 return res; 892 } 893 894 SmallVector<OpFoldResult> 895 mlir::applyMapToValues(RewriterBase &b, Location loc, AffineMap map, 896 ArrayRef<OpFoldResult> values) { 897 // Materialize constants and keep track of produced operations so we can clean 898 // them up later. 899 SmallVector<Operation *> constants; 900 SmallVector<Value> actualValues; 901 materializeConstants(b, loc, values, constants, actualValues); 902 903 // Compose, fold and construct maps for each result independently because they 904 // may simplify more effectively. 905 SmallVector<OpFoldResult> results; 906 results.reserve(map.getNumResults()); 907 bool foldedAll = true; 908 for (auto i : llvm::seq<unsigned>(0, map.getNumResults())) { 909 AffineMap submap = map.getSubMap({i}); 910 SmallVector<Value> operands = actualValues; 911 fullyComposeAffineMapAndOperands(&submap, &operands); 912 canonicalizeMapAndOperands(&submap, &operands); 913 results.push_back(createOrFold<AffineApplyOp>(b, loc, operands, submap)); 914 if (!results.back().is<Attribute>()) 915 foldedAll = false; 916 } 917 918 // If the entire map could be folded, remove the constants that were used in 919 // the initial ops. 920 if (foldedAll) { 921 for (Operation *constant : constants) 922 b.eraseOp(constant); 923 } 924 925 return results; 926 } 927 928 // A symbol may appear as a dim in affine.apply operations. This function 929 // canonicalizes dims that are valid symbols into actual symbols. 930 template <class MapOrSet> 931 static void canonicalizePromotedSymbols(MapOrSet *mapOrSet, 932 SmallVectorImpl<Value> *operands) { 933 if (!mapOrSet || operands->empty()) 934 return; 935 936 assert(mapOrSet->getNumInputs() == operands->size() && 937 "map/set inputs must match number of operands"); 938 939 auto *context = mapOrSet->getContext(); 940 SmallVector<Value, 8> resultOperands; 941 resultOperands.reserve(operands->size()); 942 SmallVector<Value, 8> remappedSymbols; 943 remappedSymbols.reserve(operands->size()); 944 unsigned nextDim = 0; 945 unsigned nextSym = 0; 946 unsigned oldNumSyms = mapOrSet->getNumSymbols(); 947 SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims()); 948 for (unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) { 949 if (i < mapOrSet->getNumDims()) { 950 if (isValidSymbol((*operands)[i])) { 951 // This is a valid symbol that appears as a dim, canonicalize it. 952 dimRemapping[i] = getAffineSymbolExpr(oldNumSyms + nextSym++, context); 953 remappedSymbols.push_back((*operands)[i]); 954 } else { 955 dimRemapping[i] = getAffineDimExpr(nextDim++, context); 956 resultOperands.push_back((*operands)[i]); 957 } 958 } else { 959 resultOperands.push_back((*operands)[i]); 960 } 961 } 962 963 resultOperands.append(remappedSymbols.begin(), remappedSymbols.end()); 964 *operands = resultOperands; 965 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, {}, nextDim, 966 oldNumSyms + nextSym); 967 968 assert(mapOrSet->getNumInputs() == operands->size() && 969 "map/set inputs must match number of operands"); 970 } 971 972 // Works for either an affine map or an integer set. 973 template <class MapOrSet> 974 static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet, 975 SmallVectorImpl<Value> *operands) { 976 static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value, 977 "Argument must be either of AffineMap or IntegerSet type"); 978 979 if (!mapOrSet || operands->empty()) 980 return; 981 982 assert(mapOrSet->getNumInputs() == operands->size() && 983 "map/set inputs must match number of operands"); 984 985 canonicalizePromotedSymbols<MapOrSet>(mapOrSet, operands); 986 987 // Check to see what dims are used. 988 llvm::SmallBitVector usedDims(mapOrSet->getNumDims()); 989 llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols()); 990 mapOrSet->walkExprs([&](AffineExpr expr) { 991 if (auto dimExpr = expr.dyn_cast<AffineDimExpr>()) 992 usedDims[dimExpr.getPosition()] = true; 993 else if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) 994 usedSyms[symExpr.getPosition()] = true; 995 }); 996 997 auto *context = mapOrSet->getContext(); 998 999 SmallVector<Value, 8> resultOperands; 1000 resultOperands.reserve(operands->size()); 1001 1002 llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims; 1003 SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims()); 1004 unsigned nextDim = 0; 1005 for (unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) { 1006 if (usedDims[i]) { 1007 // Remap dim positions for duplicate operands. 1008 auto it = seenDims.find((*operands)[i]); 1009 if (it == seenDims.end()) { 1010 dimRemapping[i] = getAffineDimExpr(nextDim++, context); 1011 resultOperands.push_back((*operands)[i]); 1012 seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i])); 1013 } else { 1014 dimRemapping[i] = it->second; 1015 } 1016 } 1017 } 1018 llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols; 1019 SmallVector<AffineExpr, 8> symRemapping(mapOrSet->getNumSymbols()); 1020 unsigned nextSym = 0; 1021 for (unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) { 1022 if (!usedSyms[i]) 1023 continue; 1024 // Handle constant operands (only needed for symbolic operands since 1025 // constant operands in dimensional positions would have already been 1026 // promoted to symbolic positions above). 1027 IntegerAttr operandCst; 1028 if (matchPattern((*operands)[i + mapOrSet->getNumDims()], 1029 m_Constant(&operandCst))) { 1030 symRemapping[i] = 1031 getAffineConstantExpr(operandCst.getValue().getSExtValue(), context); 1032 continue; 1033 } 1034 // Remap symbol positions for duplicate operands. 1035 auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]); 1036 if (it == seenSymbols.end()) { 1037 symRemapping[i] = getAffineSymbolExpr(nextSym++, context); 1038 resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]); 1039 seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()], 1040 symRemapping[i])); 1041 } else { 1042 symRemapping[i] = it->second; 1043 } 1044 } 1045 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping, 1046 nextDim, nextSym); 1047 *operands = resultOperands; 1048 } 1049 1050 void mlir::canonicalizeMapAndOperands(AffineMap *map, 1051 SmallVectorImpl<Value> *operands) { 1052 canonicalizeMapOrSetAndOperands<AffineMap>(map, operands); 1053 } 1054 1055 void mlir::canonicalizeSetAndOperands(IntegerSet *set, 1056 SmallVectorImpl<Value> *operands) { 1057 canonicalizeMapOrSetAndOperands<IntegerSet>(set, operands); 1058 } 1059 1060 namespace { 1061 /// Simplify AffineApply, AffineLoad, and AffineStore operations by composing 1062 /// maps that supply results into them. 1063 /// 1064 template <typename AffineOpTy> 1065 struct SimplifyAffineOp : public OpRewritePattern<AffineOpTy> { 1066 using OpRewritePattern<AffineOpTy>::OpRewritePattern; 1067 1068 /// Replace the affine op with another instance of it with the supplied 1069 /// map and mapOperands. 1070 void replaceAffineOp(PatternRewriter &rewriter, AffineOpTy affineOp, 1071 AffineMap map, ArrayRef<Value> mapOperands) const; 1072 1073 LogicalResult matchAndRewrite(AffineOpTy affineOp, 1074 PatternRewriter &rewriter) const override { 1075 static_assert( 1076 llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp, 1077 AffineStoreOp, AffineApplyOp, AffineMinOp, AffineMaxOp, 1078 AffineVectorStoreOp, AffineVectorLoadOp>::value, 1079 "affine load/store/vectorstore/vectorload/apply/prefetch/min/max op " 1080 "expected"); 1081 auto map = affineOp.getAffineMap(); 1082 AffineMap oldMap = map; 1083 auto oldOperands = affineOp.getMapOperands(); 1084 SmallVector<Value, 8> resultOperands(oldOperands); 1085 composeAffineMapAndOperands(&map, &resultOperands); 1086 canonicalizeMapAndOperands(&map, &resultOperands); 1087 if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(), 1088 resultOperands.begin())) 1089 return failure(); 1090 1091 replaceAffineOp(rewriter, affineOp, map, resultOperands); 1092 return success(); 1093 } 1094 }; 1095 1096 // Specialize the template to account for the different build signatures for 1097 // affine load, store, and apply ops. 1098 template <> 1099 void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp( 1100 PatternRewriter &rewriter, AffineLoadOp load, AffineMap map, 1101 ArrayRef<Value> mapOperands) const { 1102 rewriter.replaceOpWithNewOp<AffineLoadOp>(load, load.getMemRef(), map, 1103 mapOperands); 1104 } 1105 template <> 1106 void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp( 1107 PatternRewriter &rewriter, AffinePrefetchOp prefetch, AffineMap map, 1108 ArrayRef<Value> mapOperands) const { 1109 rewriter.replaceOpWithNewOp<AffinePrefetchOp>( 1110 prefetch, prefetch.getMemref(), map, mapOperands, 1111 prefetch.getLocalityHint(), prefetch.getIsWrite(), 1112 prefetch.getIsDataCache()); 1113 } 1114 template <> 1115 void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp( 1116 PatternRewriter &rewriter, AffineStoreOp store, AffineMap map, 1117 ArrayRef<Value> mapOperands) const { 1118 rewriter.replaceOpWithNewOp<AffineStoreOp>( 1119 store, store.getValueToStore(), store.getMemRef(), map, mapOperands); 1120 } 1121 template <> 1122 void SimplifyAffineOp<AffineVectorLoadOp>::replaceAffineOp( 1123 PatternRewriter &rewriter, AffineVectorLoadOp vectorload, AffineMap map, 1124 ArrayRef<Value> mapOperands) const { 1125 rewriter.replaceOpWithNewOp<AffineVectorLoadOp>( 1126 vectorload, vectorload.getVectorType(), vectorload.getMemRef(), map, 1127 mapOperands); 1128 } 1129 template <> 1130 void SimplifyAffineOp<AffineVectorStoreOp>::replaceAffineOp( 1131 PatternRewriter &rewriter, AffineVectorStoreOp vectorstore, AffineMap map, 1132 ArrayRef<Value> mapOperands) const { 1133 rewriter.replaceOpWithNewOp<AffineVectorStoreOp>( 1134 vectorstore, vectorstore.getValueToStore(), vectorstore.getMemRef(), map, 1135 mapOperands); 1136 } 1137 1138 // Generic version for ops that don't have extra operands. 1139 template <typename AffineOpTy> 1140 void SimplifyAffineOp<AffineOpTy>::replaceAffineOp( 1141 PatternRewriter &rewriter, AffineOpTy op, AffineMap map, 1142 ArrayRef<Value> mapOperands) const { 1143 rewriter.replaceOpWithNewOp<AffineOpTy>(op, map, mapOperands); 1144 } 1145 } // namespace 1146 1147 void AffineApplyOp::getCanonicalizationPatterns(RewritePatternSet &results, 1148 MLIRContext *context) { 1149 results.add<SimplifyAffineOp<AffineApplyOp>>(context); 1150 } 1151 1152 //===----------------------------------------------------------------------===// 1153 // Common canonicalization pattern support logic 1154 //===----------------------------------------------------------------------===// 1155 1156 /// This is a common class used for patterns of the form 1157 /// "someop(memrefcast) -> someop". It folds the source of any memref.cast 1158 /// into the root operation directly. 1159 static LogicalResult foldMemRefCast(Operation *op, Value ignore = nullptr) { 1160 bool folded = false; 1161 for (OpOperand &operand : op->getOpOperands()) { 1162 auto cast = operand.get().getDefiningOp<memref::CastOp>(); 1163 if (cast && operand.get() != ignore && 1164 !cast.getOperand().getType().isa<UnrankedMemRefType>()) { 1165 operand.set(cast.getOperand()); 1166 folded = true; 1167 } 1168 } 1169 return success(folded); 1170 } 1171 1172 //===----------------------------------------------------------------------===// 1173 // AffineDmaStartOp 1174 //===----------------------------------------------------------------------===// 1175 1176 // TODO: Check that map operands are loop IVs or symbols. 1177 void AffineDmaStartOp::build(OpBuilder &builder, OperationState &result, 1178 Value srcMemRef, AffineMap srcMap, 1179 ValueRange srcIndices, Value destMemRef, 1180 AffineMap dstMap, ValueRange destIndices, 1181 Value tagMemRef, AffineMap tagMap, 1182 ValueRange tagIndices, Value numElements, 1183 Value stride, Value elementsPerStride) { 1184 result.addOperands(srcMemRef); 1185 result.addAttribute(getSrcMapAttrStrName(), AffineMapAttr::get(srcMap)); 1186 result.addOperands(srcIndices); 1187 result.addOperands(destMemRef); 1188 result.addAttribute(getDstMapAttrStrName(), AffineMapAttr::get(dstMap)); 1189 result.addOperands(destIndices); 1190 result.addOperands(tagMemRef); 1191 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap)); 1192 result.addOperands(tagIndices); 1193 result.addOperands(numElements); 1194 if (stride) { 1195 result.addOperands({stride, elementsPerStride}); 1196 } 1197 } 1198 1199 void AffineDmaStartOp::print(OpAsmPrinter &p) { 1200 p << " " << getSrcMemRef() << '['; 1201 p.printAffineMapOfSSAIds(getSrcMapAttr(), getSrcIndices()); 1202 p << "], " << getDstMemRef() << '['; 1203 p.printAffineMapOfSSAIds(getDstMapAttr(), getDstIndices()); 1204 p << "], " << getTagMemRef() << '['; 1205 p.printAffineMapOfSSAIds(getTagMapAttr(), getTagIndices()); 1206 p << "], " << getNumElements(); 1207 if (isStrided()) { 1208 p << ", " << getStride(); 1209 p << ", " << getNumElementsPerStride(); 1210 } 1211 p << " : " << getSrcMemRefType() << ", " << getDstMemRefType() << ", " 1212 << getTagMemRefType(); 1213 } 1214 1215 // Parse AffineDmaStartOp. 1216 // Ex: 1217 // affine.dma_start %src[%i, %j], %dst[%k, %l], %tag[%index], %size, 1218 // %stride, %num_elt_per_stride 1219 // : memref<3076 x f32, 0>, memref<1024 x f32, 2>, memref<1 x i32> 1220 // 1221 ParseResult AffineDmaStartOp::parse(OpAsmParser &parser, 1222 OperationState &result) { 1223 OpAsmParser::UnresolvedOperand srcMemRefInfo; 1224 AffineMapAttr srcMapAttr; 1225 SmallVector<OpAsmParser::UnresolvedOperand, 4> srcMapOperands; 1226 OpAsmParser::UnresolvedOperand dstMemRefInfo; 1227 AffineMapAttr dstMapAttr; 1228 SmallVector<OpAsmParser::UnresolvedOperand, 4> dstMapOperands; 1229 OpAsmParser::UnresolvedOperand tagMemRefInfo; 1230 AffineMapAttr tagMapAttr; 1231 SmallVector<OpAsmParser::UnresolvedOperand, 4> tagMapOperands; 1232 OpAsmParser::UnresolvedOperand numElementsInfo; 1233 SmallVector<OpAsmParser::UnresolvedOperand, 2> strideInfo; 1234 1235 SmallVector<Type, 3> types; 1236 auto indexType = parser.getBuilder().getIndexType(); 1237 1238 // Parse and resolve the following list of operands: 1239 // *) dst memref followed by its affine maps operands (in square brackets). 1240 // *) src memref followed by its affine map operands (in square brackets). 1241 // *) tag memref followed by its affine map operands (in square brackets). 1242 // *) number of elements transferred by DMA operation. 1243 if (parser.parseOperand(srcMemRefInfo) || 1244 parser.parseAffineMapOfSSAIds(srcMapOperands, srcMapAttr, 1245 getSrcMapAttrStrName(), 1246 result.attributes) || 1247 parser.parseComma() || parser.parseOperand(dstMemRefInfo) || 1248 parser.parseAffineMapOfSSAIds(dstMapOperands, dstMapAttr, 1249 getDstMapAttrStrName(), 1250 result.attributes) || 1251 parser.parseComma() || parser.parseOperand(tagMemRefInfo) || 1252 parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr, 1253 getTagMapAttrStrName(), 1254 result.attributes) || 1255 parser.parseComma() || parser.parseOperand(numElementsInfo)) 1256 return failure(); 1257 1258 // Parse optional stride and elements per stride. 1259 if (parser.parseTrailingOperandList(strideInfo)) 1260 return failure(); 1261 1262 if (!strideInfo.empty() && strideInfo.size() != 2) { 1263 return parser.emitError(parser.getNameLoc(), 1264 "expected two stride related operands"); 1265 } 1266 bool isStrided = strideInfo.size() == 2; 1267 1268 if (parser.parseColonTypeList(types)) 1269 return failure(); 1270 1271 if (types.size() != 3) 1272 return parser.emitError(parser.getNameLoc(), "expected three types"); 1273 1274 if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) || 1275 parser.resolveOperands(srcMapOperands, indexType, result.operands) || 1276 parser.resolveOperand(dstMemRefInfo, types[1], result.operands) || 1277 parser.resolveOperands(dstMapOperands, indexType, result.operands) || 1278 parser.resolveOperand(tagMemRefInfo, types[2], result.operands) || 1279 parser.resolveOperands(tagMapOperands, indexType, result.operands) || 1280 parser.resolveOperand(numElementsInfo, indexType, result.operands)) 1281 return failure(); 1282 1283 if (isStrided) { 1284 if (parser.resolveOperands(strideInfo, indexType, result.operands)) 1285 return failure(); 1286 } 1287 1288 // Check that src/dst/tag operand counts match their map.numInputs. 1289 if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() || 1290 dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() || 1291 tagMapOperands.size() != tagMapAttr.getValue().getNumInputs()) 1292 return parser.emitError(parser.getNameLoc(), 1293 "memref operand count not equal to map.numInputs"); 1294 return success(); 1295 } 1296 1297 LogicalResult AffineDmaStartOp::verifyInvariantsImpl() { 1298 if (!getOperand(getSrcMemRefOperandIndex()).getType().isa<MemRefType>()) 1299 return emitOpError("expected DMA source to be of memref type"); 1300 if (!getOperand(getDstMemRefOperandIndex()).getType().isa<MemRefType>()) 1301 return emitOpError("expected DMA destination to be of memref type"); 1302 if (!getOperand(getTagMemRefOperandIndex()).getType().isa<MemRefType>()) 1303 return emitOpError("expected DMA tag to be of memref type"); 1304 1305 unsigned numInputsAllMaps = getSrcMap().getNumInputs() + 1306 getDstMap().getNumInputs() + 1307 getTagMap().getNumInputs(); 1308 if (getNumOperands() != numInputsAllMaps + 3 + 1 && 1309 getNumOperands() != numInputsAllMaps + 3 + 1 + 2) { 1310 return emitOpError("incorrect number of operands"); 1311 } 1312 1313 Region *scope = getAffineScope(*this); 1314 for (auto idx : getSrcIndices()) { 1315 if (!idx.getType().isIndex()) 1316 return emitOpError("src index to dma_start must have 'index' type"); 1317 if (!isValidAffineIndexOperand(idx, scope)) 1318 return emitOpError("src index must be a dimension or symbol identifier"); 1319 } 1320 for (auto idx : getDstIndices()) { 1321 if (!idx.getType().isIndex()) 1322 return emitOpError("dst index to dma_start must have 'index' type"); 1323 if (!isValidAffineIndexOperand(idx, scope)) 1324 return emitOpError("dst index must be a dimension or symbol identifier"); 1325 } 1326 for (auto idx : getTagIndices()) { 1327 if (!idx.getType().isIndex()) 1328 return emitOpError("tag index to dma_start must have 'index' type"); 1329 if (!isValidAffineIndexOperand(idx, scope)) 1330 return emitOpError("tag index must be a dimension or symbol identifier"); 1331 } 1332 return success(); 1333 } 1334 1335 LogicalResult AffineDmaStartOp::fold(ArrayRef<Attribute> cstOperands, 1336 SmallVectorImpl<OpFoldResult> &results) { 1337 /// dma_start(memrefcast) -> dma_start 1338 return foldMemRefCast(*this); 1339 } 1340 1341 //===----------------------------------------------------------------------===// 1342 // AffineDmaWaitOp 1343 //===----------------------------------------------------------------------===// 1344 1345 // TODO: Check that map operands are loop IVs or symbols. 1346 void AffineDmaWaitOp::build(OpBuilder &builder, OperationState &result, 1347 Value tagMemRef, AffineMap tagMap, 1348 ValueRange tagIndices, Value numElements) { 1349 result.addOperands(tagMemRef); 1350 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap)); 1351 result.addOperands(tagIndices); 1352 result.addOperands(numElements); 1353 } 1354 1355 void AffineDmaWaitOp::print(OpAsmPrinter &p) { 1356 p << " " << getTagMemRef() << '['; 1357 SmallVector<Value, 2> operands(getTagIndices()); 1358 p.printAffineMapOfSSAIds(getTagMapAttr(), operands); 1359 p << "], "; 1360 p.printOperand(getNumElements()); 1361 p << " : " << getTagMemRef().getType(); 1362 } 1363 1364 // Parse AffineDmaWaitOp. 1365 // Eg: 1366 // affine.dma_wait %tag[%index], %num_elements 1367 // : memref<1 x i32, (d0) -> (d0), 4> 1368 // 1369 ParseResult AffineDmaWaitOp::parse(OpAsmParser &parser, 1370 OperationState &result) { 1371 OpAsmParser::UnresolvedOperand tagMemRefInfo; 1372 AffineMapAttr tagMapAttr; 1373 SmallVector<OpAsmParser::UnresolvedOperand, 2> tagMapOperands; 1374 Type type; 1375 auto indexType = parser.getBuilder().getIndexType(); 1376 OpAsmParser::UnresolvedOperand numElementsInfo; 1377 1378 // Parse tag memref, its map operands, and dma size. 1379 if (parser.parseOperand(tagMemRefInfo) || 1380 parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr, 1381 getTagMapAttrStrName(), 1382 result.attributes) || 1383 parser.parseComma() || parser.parseOperand(numElementsInfo) || 1384 parser.parseColonType(type) || 1385 parser.resolveOperand(tagMemRefInfo, type, result.operands) || 1386 parser.resolveOperands(tagMapOperands, indexType, result.operands) || 1387 parser.resolveOperand(numElementsInfo, indexType, result.operands)) 1388 return failure(); 1389 1390 if (!type.isa<MemRefType>()) 1391 return parser.emitError(parser.getNameLoc(), 1392 "expected tag to be of memref type"); 1393 1394 if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs()) 1395 return parser.emitError(parser.getNameLoc(), 1396 "tag memref operand count != to map.numInputs"); 1397 return success(); 1398 } 1399 1400 LogicalResult AffineDmaWaitOp::verifyInvariantsImpl() { 1401 if (!getOperand(0).getType().isa<MemRefType>()) 1402 return emitOpError("expected DMA tag to be of memref type"); 1403 Region *scope = getAffineScope(*this); 1404 for (auto idx : getTagIndices()) { 1405 if (!idx.getType().isIndex()) 1406 return emitOpError("index to dma_wait must have 'index' type"); 1407 if (!isValidAffineIndexOperand(idx, scope)) 1408 return emitOpError("index must be a dimension or symbol identifier"); 1409 } 1410 return success(); 1411 } 1412 1413 LogicalResult AffineDmaWaitOp::fold(ArrayRef<Attribute> cstOperands, 1414 SmallVectorImpl<OpFoldResult> &results) { 1415 /// dma_wait(memrefcast) -> dma_wait 1416 return foldMemRefCast(*this); 1417 } 1418 1419 //===----------------------------------------------------------------------===// 1420 // AffineForOp 1421 //===----------------------------------------------------------------------===// 1422 1423 /// 'bodyBuilder' is used to build the body of affine.for. If iterArgs and 1424 /// bodyBuilder are empty/null, we include default terminator op. 1425 void AffineForOp::build(OpBuilder &builder, OperationState &result, 1426 ValueRange lbOperands, AffineMap lbMap, 1427 ValueRange ubOperands, AffineMap ubMap, int64_t step, 1428 ValueRange iterArgs, BodyBuilderFn bodyBuilder) { 1429 assert(((!lbMap && lbOperands.empty()) || 1430 lbOperands.size() == lbMap.getNumInputs()) && 1431 "lower bound operand count does not match the affine map"); 1432 assert(((!ubMap && ubOperands.empty()) || 1433 ubOperands.size() == ubMap.getNumInputs()) && 1434 "upper bound operand count does not match the affine map"); 1435 assert(step > 0 && "step has to be a positive integer constant"); 1436 1437 for (Value val : iterArgs) 1438 result.addTypes(val.getType()); 1439 1440 // Add an attribute for the step. 1441 result.addAttribute(getStepAttrStrName(), 1442 builder.getIntegerAttr(builder.getIndexType(), step)); 1443 1444 // Add the lower bound. 1445 result.addAttribute(getLowerBoundAttrStrName(), AffineMapAttr::get(lbMap)); 1446 result.addOperands(lbOperands); 1447 1448 // Add the upper bound. 1449 result.addAttribute(getUpperBoundAttrStrName(), AffineMapAttr::get(ubMap)); 1450 result.addOperands(ubOperands); 1451 1452 result.addOperands(iterArgs); 1453 // Create a region and a block for the body. The argument of the region is 1454 // the loop induction variable. 1455 Region *bodyRegion = result.addRegion(); 1456 bodyRegion->push_back(new Block); 1457 Block &bodyBlock = bodyRegion->front(); 1458 Value inductionVar = 1459 bodyBlock.addArgument(builder.getIndexType(), result.location); 1460 for (Value val : iterArgs) 1461 bodyBlock.addArgument(val.getType(), val.getLoc()); 1462 1463 // Create the default terminator if the builder is not provided and if the 1464 // iteration arguments are not provided. Otherwise, leave this to the caller 1465 // because we don't know which values to return from the loop. 1466 if (iterArgs.empty() && !bodyBuilder) { 1467 ensureTerminator(*bodyRegion, builder, result.location); 1468 } else if (bodyBuilder) { 1469 OpBuilder::InsertionGuard guard(builder); 1470 builder.setInsertionPointToStart(&bodyBlock); 1471 bodyBuilder(builder, result.location, inductionVar, 1472 bodyBlock.getArguments().drop_front()); 1473 } 1474 } 1475 1476 void AffineForOp::build(OpBuilder &builder, OperationState &result, int64_t lb, 1477 int64_t ub, int64_t step, ValueRange iterArgs, 1478 BodyBuilderFn bodyBuilder) { 1479 auto lbMap = AffineMap::getConstantMap(lb, builder.getContext()); 1480 auto ubMap = AffineMap::getConstantMap(ub, builder.getContext()); 1481 return build(builder, result, {}, lbMap, {}, ubMap, step, iterArgs, 1482 bodyBuilder); 1483 } 1484 1485 LogicalResult AffineForOp::verifyRegions() { 1486 // Check that the body defines as single block argument for the induction 1487 // variable. 1488 auto *body = getBody(); 1489 if (body->getNumArguments() == 0 || !body->getArgument(0).getType().isIndex()) 1490 return emitOpError("expected body to have a single index argument for the " 1491 "induction variable"); 1492 1493 // Verify that the bound operands are valid dimension/symbols. 1494 /// Lower bound. 1495 if (getLowerBoundMap().getNumInputs() > 0) 1496 if (failed(verifyDimAndSymbolIdentifiers(*this, getLowerBoundOperands(), 1497 getLowerBoundMap().getNumDims()))) 1498 return failure(); 1499 /// Upper bound. 1500 if (getUpperBoundMap().getNumInputs() > 0) 1501 if (failed(verifyDimAndSymbolIdentifiers(*this, getUpperBoundOperands(), 1502 getUpperBoundMap().getNumDims()))) 1503 return failure(); 1504 1505 unsigned opNumResults = getNumResults(); 1506 if (opNumResults == 0) 1507 return success(); 1508 1509 // If ForOp defines values, check that the number and types of the defined 1510 // values match ForOp initial iter operands and backedge basic block 1511 // arguments. 1512 if (getNumIterOperands() != opNumResults) 1513 return emitOpError( 1514 "mismatch between the number of loop-carried values and results"); 1515 if (getNumRegionIterArgs() != opNumResults) 1516 return emitOpError( 1517 "mismatch between the number of basic block args and results"); 1518 1519 return success(); 1520 } 1521 1522 /// Parse a for operation loop bounds. 1523 static ParseResult parseBound(bool isLower, OperationState &result, 1524 OpAsmParser &p) { 1525 // 'min' / 'max' prefixes are generally syntactic sugar, but are required if 1526 // the map has multiple results. 1527 bool failedToParsedMinMax = 1528 failed(p.parseOptionalKeyword(isLower ? "max" : "min")); 1529 1530 auto &builder = p.getBuilder(); 1531 auto boundAttrStrName = isLower ? AffineForOp::getLowerBoundAttrStrName() 1532 : AffineForOp::getUpperBoundAttrStrName(); 1533 1534 // Parse ssa-id as identity map. 1535 SmallVector<OpAsmParser::UnresolvedOperand, 1> boundOpInfos; 1536 if (p.parseOperandList(boundOpInfos)) 1537 return failure(); 1538 1539 if (!boundOpInfos.empty()) { 1540 // Check that only one operand was parsed. 1541 if (boundOpInfos.size() > 1) 1542 return p.emitError(p.getNameLoc(), 1543 "expected only one loop bound operand"); 1544 1545 // TODO: improve error message when SSA value is not of index type. 1546 // Currently it is 'use of value ... expects different type than prior uses' 1547 if (p.resolveOperand(boundOpInfos.front(), builder.getIndexType(), 1548 result.operands)) 1549 return failure(); 1550 1551 // Create an identity map using symbol id. This representation is optimized 1552 // for storage. Analysis passes may expand it into a multi-dimensional map 1553 // if desired. 1554 AffineMap map = builder.getSymbolIdentityMap(); 1555 result.addAttribute(boundAttrStrName, AffineMapAttr::get(map)); 1556 return success(); 1557 } 1558 1559 // Get the attribute location. 1560 SMLoc attrLoc = p.getCurrentLocation(); 1561 1562 Attribute boundAttr; 1563 if (p.parseAttribute(boundAttr, builder.getIndexType(), boundAttrStrName, 1564 result.attributes)) 1565 return failure(); 1566 1567 // Parse full form - affine map followed by dim and symbol list. 1568 if (auto affineMapAttr = boundAttr.dyn_cast<AffineMapAttr>()) { 1569 unsigned currentNumOperands = result.operands.size(); 1570 unsigned numDims; 1571 if (parseDimAndSymbolList(p, result.operands, numDims)) 1572 return failure(); 1573 1574 auto map = affineMapAttr.getValue(); 1575 if (map.getNumDims() != numDims) 1576 return p.emitError( 1577 p.getNameLoc(), 1578 "dim operand count and affine map dim count must match"); 1579 1580 unsigned numDimAndSymbolOperands = 1581 result.operands.size() - currentNumOperands; 1582 if (numDims + map.getNumSymbols() != numDimAndSymbolOperands) 1583 return p.emitError( 1584 p.getNameLoc(), 1585 "symbol operand count and affine map symbol count must match"); 1586 1587 // If the map has multiple results, make sure that we parsed the min/max 1588 // prefix. 1589 if (map.getNumResults() > 1 && failedToParsedMinMax) { 1590 if (isLower) { 1591 return p.emitError(attrLoc, "lower loop bound affine map with " 1592 "multiple results requires 'max' prefix"); 1593 } 1594 return p.emitError(attrLoc, "upper loop bound affine map with multiple " 1595 "results requires 'min' prefix"); 1596 } 1597 return success(); 1598 } 1599 1600 // Parse custom assembly form. 1601 if (auto integerAttr = boundAttr.dyn_cast<IntegerAttr>()) { 1602 result.attributes.pop_back(); 1603 result.addAttribute( 1604 boundAttrStrName, 1605 AffineMapAttr::get(builder.getConstantAffineMap(integerAttr.getInt()))); 1606 return success(); 1607 } 1608 1609 return p.emitError( 1610 p.getNameLoc(), 1611 "expected valid affine map representation for loop bounds"); 1612 } 1613 1614 ParseResult AffineForOp::parse(OpAsmParser &parser, OperationState &result) { 1615 auto &builder = parser.getBuilder(); 1616 OpAsmParser::Argument inductionVariable; 1617 inductionVariable.type = builder.getIndexType(); 1618 // Parse the induction variable followed by '='. 1619 if (parser.parseArgument(inductionVariable) || parser.parseEqual()) 1620 return failure(); 1621 1622 // Parse loop bounds. 1623 if (parseBound(/*isLower=*/true, result, parser) || 1624 parser.parseKeyword("to", " between bounds") || 1625 parseBound(/*isLower=*/false, result, parser)) 1626 return failure(); 1627 1628 // Parse the optional loop step, we default to 1 if one is not present. 1629 if (parser.parseOptionalKeyword("step")) { 1630 result.addAttribute( 1631 AffineForOp::getStepAttrStrName(), 1632 builder.getIntegerAttr(builder.getIndexType(), /*value=*/1)); 1633 } else { 1634 SMLoc stepLoc = parser.getCurrentLocation(); 1635 IntegerAttr stepAttr; 1636 if (parser.parseAttribute(stepAttr, builder.getIndexType(), 1637 AffineForOp::getStepAttrStrName().data(), 1638 result.attributes)) 1639 return failure(); 1640 1641 if (stepAttr.getValue().getSExtValue() < 0) 1642 return parser.emitError( 1643 stepLoc, 1644 "expected step to be representable as a positive signed integer"); 1645 } 1646 1647 // Parse the optional initial iteration arguments. 1648 SmallVector<OpAsmParser::Argument, 4> regionArgs; 1649 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands; 1650 1651 // Induction variable. 1652 regionArgs.push_back(inductionVariable); 1653 1654 if (succeeded(parser.parseOptionalKeyword("iter_args"))) { 1655 // Parse assignment list and results type list. 1656 if (parser.parseAssignmentList(regionArgs, operands) || 1657 parser.parseArrowTypeList(result.types)) 1658 return failure(); 1659 // Resolve input operands. 1660 for (auto argOperandType : 1661 llvm::zip(llvm::drop_begin(regionArgs), operands, result.types)) { 1662 Type type = std::get<2>(argOperandType); 1663 std::get<0>(argOperandType).type = type; 1664 if (parser.resolveOperand(std::get<1>(argOperandType), type, 1665 result.operands)) 1666 return failure(); 1667 } 1668 } 1669 1670 // Parse the body region. 1671 Region *body = result.addRegion(); 1672 if (regionArgs.size() != result.types.size() + 1) 1673 return parser.emitError( 1674 parser.getNameLoc(), 1675 "mismatch between the number of loop-carried values and results"); 1676 if (parser.parseRegion(*body, regionArgs)) 1677 return failure(); 1678 1679 AffineForOp::ensureTerminator(*body, builder, result.location); 1680 1681 // Parse the optional attribute list. 1682 return parser.parseOptionalAttrDict(result.attributes); 1683 } 1684 1685 static void printBound(AffineMapAttr boundMap, 1686 Operation::operand_range boundOperands, 1687 const char *prefix, OpAsmPrinter &p) { 1688 AffineMap map = boundMap.getValue(); 1689 1690 // Check if this bound should be printed using custom assembly form. 1691 // The decision to restrict printing custom assembly form to trivial cases 1692 // comes from the will to roundtrip MLIR binary -> text -> binary in a 1693 // lossless way. 1694 // Therefore, custom assembly form parsing and printing is only supported for 1695 // zero-operand constant maps and single symbol operand identity maps. 1696 if (map.getNumResults() == 1) { 1697 AffineExpr expr = map.getResult(0); 1698 1699 // Print constant bound. 1700 if (map.getNumDims() == 0 && map.getNumSymbols() == 0) { 1701 if (auto constExpr = expr.dyn_cast<AffineConstantExpr>()) { 1702 p << constExpr.getValue(); 1703 return; 1704 } 1705 } 1706 1707 // Print bound that consists of a single SSA symbol if the map is over a 1708 // single symbol. 1709 if (map.getNumDims() == 0 && map.getNumSymbols() == 1) { 1710 if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) { 1711 p.printOperand(*boundOperands.begin()); 1712 return; 1713 } 1714 } 1715 } else { 1716 // Map has multiple results. Print 'min' or 'max' prefix. 1717 p << prefix << ' '; 1718 } 1719 1720 // Print the map and its operands. 1721 p << boundMap; 1722 printDimAndSymbolList(boundOperands.begin(), boundOperands.end(), 1723 map.getNumDims(), p); 1724 } 1725 1726 unsigned AffineForOp::getNumIterOperands() { 1727 AffineMap lbMap = getLowerBoundMapAttr().getValue(); 1728 AffineMap ubMap = getUpperBoundMapAttr().getValue(); 1729 1730 return getNumOperands() - lbMap.getNumInputs() - ubMap.getNumInputs(); 1731 } 1732 1733 void AffineForOp::print(OpAsmPrinter &p) { 1734 p << ' '; 1735 p.printRegionArgument(getBody()->getArgument(0), /*argAttrs=*/{}, 1736 /*omitType=*/true); 1737 p << " = "; 1738 printBound(getLowerBoundMapAttr(), getLowerBoundOperands(), "max", p); 1739 p << " to "; 1740 printBound(getUpperBoundMapAttr(), getUpperBoundOperands(), "min", p); 1741 1742 if (getStep() != 1) 1743 p << " step " << getStep(); 1744 1745 bool printBlockTerminators = false; 1746 if (getNumIterOperands() > 0) { 1747 p << " iter_args("; 1748 auto regionArgs = getRegionIterArgs(); 1749 auto operands = getIterOperands(); 1750 1751 llvm::interleaveComma(llvm::zip(regionArgs, operands), p, [&](auto it) { 1752 p << std::get<0>(it) << " = " << std::get<1>(it); 1753 }); 1754 p << ") -> (" << getResultTypes() << ")"; 1755 printBlockTerminators = true; 1756 } 1757 1758 p << ' '; 1759 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false, 1760 printBlockTerminators); 1761 p.printOptionalAttrDict((*this)->getAttrs(), 1762 /*elidedAttrs=*/{getLowerBoundAttrStrName(), 1763 getUpperBoundAttrStrName(), 1764 getStepAttrStrName()}); 1765 } 1766 1767 /// Fold the constant bounds of a loop. 1768 static LogicalResult foldLoopBounds(AffineForOp forOp) { 1769 auto foldLowerOrUpperBound = [&forOp](bool lower) { 1770 // Check to see if each of the operands is the result of a constant. If 1771 // so, get the value. If not, ignore it. 1772 SmallVector<Attribute, 8> operandConstants; 1773 auto boundOperands = 1774 lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands(); 1775 for (auto operand : boundOperands) { 1776 Attribute operandCst; 1777 matchPattern(operand, m_Constant(&operandCst)); 1778 operandConstants.push_back(operandCst); 1779 } 1780 1781 AffineMap boundMap = 1782 lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap(); 1783 assert(boundMap.getNumResults() >= 1 && 1784 "bound maps should have at least one result"); 1785 SmallVector<Attribute, 4> foldedResults; 1786 if (failed(boundMap.constantFold(operandConstants, foldedResults))) 1787 return failure(); 1788 1789 // Compute the max or min as applicable over the results. 1790 assert(!foldedResults.empty() && "bounds should have at least one result"); 1791 auto maxOrMin = foldedResults[0].cast<IntegerAttr>().getValue(); 1792 for (unsigned i = 1, e = foldedResults.size(); i < e; i++) { 1793 auto foldedResult = foldedResults[i].cast<IntegerAttr>().getValue(); 1794 maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult) 1795 : llvm::APIntOps::smin(maxOrMin, foldedResult); 1796 } 1797 lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue()) 1798 : forOp.setConstantUpperBound(maxOrMin.getSExtValue()); 1799 return success(); 1800 }; 1801 1802 // Try to fold the lower bound. 1803 bool folded = false; 1804 if (!forOp.hasConstantLowerBound()) 1805 folded |= succeeded(foldLowerOrUpperBound(/*lower=*/true)); 1806 1807 // Try to fold the upper bound. 1808 if (!forOp.hasConstantUpperBound()) 1809 folded |= succeeded(foldLowerOrUpperBound(/*lower=*/false)); 1810 return success(folded); 1811 } 1812 1813 /// Canonicalize the bounds of the given loop. 1814 static LogicalResult canonicalizeLoopBounds(AffineForOp forOp) { 1815 SmallVector<Value, 4> lbOperands(forOp.getLowerBoundOperands()); 1816 SmallVector<Value, 4> ubOperands(forOp.getUpperBoundOperands()); 1817 1818 auto lbMap = forOp.getLowerBoundMap(); 1819 auto ubMap = forOp.getUpperBoundMap(); 1820 auto prevLbMap = lbMap; 1821 auto prevUbMap = ubMap; 1822 1823 composeAffineMapAndOperands(&lbMap, &lbOperands); 1824 canonicalizeMapAndOperands(&lbMap, &lbOperands); 1825 lbMap = removeDuplicateExprs(lbMap); 1826 1827 composeAffineMapAndOperands(&ubMap, &ubOperands); 1828 canonicalizeMapAndOperands(&ubMap, &ubOperands); 1829 ubMap = removeDuplicateExprs(ubMap); 1830 1831 // Any canonicalization change always leads to updated map(s). 1832 if (lbMap == prevLbMap && ubMap == prevUbMap) 1833 return failure(); 1834 1835 if (lbMap != prevLbMap) 1836 forOp.setLowerBound(lbOperands, lbMap); 1837 if (ubMap != prevUbMap) 1838 forOp.setUpperBound(ubOperands, ubMap); 1839 return success(); 1840 } 1841 1842 namespace { 1843 /// Returns constant trip count in trivial cases. 1844 static Optional<uint64_t> getTrivialConstantTripCount(AffineForOp forOp) { 1845 int64_t step = forOp.getStep(); 1846 if (!forOp.hasConstantBounds() || step <= 0) 1847 return None; 1848 int64_t lb = forOp.getConstantLowerBound(); 1849 int64_t ub = forOp.getConstantUpperBound(); 1850 return ub - lb <= 0 ? 0 : (ub - lb + step - 1) / step; 1851 } 1852 1853 /// This is a pattern to fold trivially empty loop bodies. 1854 /// TODO: This should be moved into the folding hook. 1855 struct AffineForEmptyLoopFolder : public OpRewritePattern<AffineForOp> { 1856 using OpRewritePattern<AffineForOp>::OpRewritePattern; 1857 1858 LogicalResult matchAndRewrite(AffineForOp forOp, 1859 PatternRewriter &rewriter) const override { 1860 // Check that the body only contains a yield. 1861 if (!llvm::hasSingleElement(*forOp.getBody())) 1862 return failure(); 1863 if (forOp.getNumResults() == 0) 1864 return success(); 1865 Optional<uint64_t> tripCount = getTrivialConstantTripCount(forOp); 1866 if (tripCount && *tripCount == 0) { 1867 // The initial values of the iteration arguments would be the op's 1868 // results. 1869 rewriter.replaceOp(forOp, forOp.getIterOperands()); 1870 return success(); 1871 } 1872 SmallVector<Value, 4> replacements; 1873 auto yieldOp = cast<AffineYieldOp>(forOp.getBody()->getTerminator()); 1874 auto iterArgs = forOp.getRegionIterArgs(); 1875 bool hasValDefinedOutsideLoop = false; 1876 bool iterArgsNotInOrder = false; 1877 for (unsigned i = 0, e = yieldOp->getNumOperands(); i < e; ++i) { 1878 Value val = yieldOp.getOperand(i); 1879 auto *iterArgIt = llvm::find(iterArgs, val); 1880 if (iterArgIt == iterArgs.end()) { 1881 // `val` is defined outside of the loop. 1882 assert(forOp.isDefinedOutsideOfLoop(val) && 1883 "must be defined outside of the loop"); 1884 hasValDefinedOutsideLoop = true; 1885 replacements.push_back(val); 1886 } else { 1887 unsigned pos = std::distance(iterArgs.begin(), iterArgIt); 1888 if (pos != i) 1889 iterArgsNotInOrder = true; 1890 replacements.push_back(forOp.getIterOperands()[pos]); 1891 } 1892 } 1893 // Bail out when the trip count is unknown and the loop returns any value 1894 // defined outside of the loop or any iterArg out of order. 1895 if (!tripCount.has_value() && 1896 (hasValDefinedOutsideLoop || iterArgsNotInOrder)) 1897 return failure(); 1898 // Bail out when the loop iterates more than once and it returns any iterArg 1899 // out of order. 1900 if (tripCount.has_value() && tripCount.getValue() >= 2 && 1901 iterArgsNotInOrder) 1902 return failure(); 1903 rewriter.replaceOp(forOp, replacements); 1904 return success(); 1905 } 1906 }; 1907 } // namespace 1908 1909 void AffineForOp::getCanonicalizationPatterns(RewritePatternSet &results, 1910 MLIRContext *context) { 1911 results.add<AffineForEmptyLoopFolder>(context); 1912 } 1913 1914 /// Return operands used when entering the region at 'index'. These operands 1915 /// correspond to the loop iterator operands, i.e., those excluding the 1916 /// induction variable. AffineForOp only has one region, so zero is the only 1917 /// valid value for `index`. 1918 OperandRange AffineForOp::getSuccessorEntryOperands(Optional<unsigned> index) { 1919 assert((!index || *index == 0) && "invalid region index"); 1920 1921 // The initial operands map to the loop arguments after the induction 1922 // variable or are forwarded to the results when the trip count is zero. 1923 return getIterOperands(); 1924 } 1925 1926 /// Given the region at `index`, or the parent operation if `index` is None, 1927 /// return the successor regions. These are the regions that may be selected 1928 /// during the flow of control. `operands` is a set of optional attributes that 1929 /// correspond to a constant value for each operand, or null if that operand is 1930 /// not a constant. 1931 void AffineForOp::getSuccessorRegions( 1932 Optional<unsigned> index, ArrayRef<Attribute> operands, 1933 SmallVectorImpl<RegionSuccessor> ®ions) { 1934 assert((!index.has_value() || index.getValue() == 0) && 1935 "expected loop region"); 1936 // The loop may typically branch back to its body or to the parent operation. 1937 // If the predecessor is the parent op and the trip count is known to be at 1938 // least one, branch into the body using the iterator arguments. And in cases 1939 // we know the trip count is zero, it can only branch back to its parent. 1940 Optional<uint64_t> tripCount = getTrivialConstantTripCount(*this); 1941 if (!index.has_value() && tripCount.has_value()) { 1942 if (tripCount.getValue() > 0) { 1943 regions.push_back(RegionSuccessor(&getLoopBody(), getRegionIterArgs())); 1944 return; 1945 } 1946 if (tripCount.getValue() == 0) { 1947 regions.push_back(RegionSuccessor(getResults())); 1948 return; 1949 } 1950 } 1951 1952 // From the loop body, if the trip count is one, we can only branch back to 1953 // the parent. 1954 if (index && tripCount && *tripCount == 1) { 1955 regions.push_back(RegionSuccessor(getResults())); 1956 return; 1957 } 1958 1959 // In all other cases, the loop may branch back to itself or the parent 1960 // operation. 1961 regions.push_back(RegionSuccessor(&getLoopBody(), getRegionIterArgs())); 1962 regions.push_back(RegionSuccessor(getResults())); 1963 } 1964 1965 /// Returns true if the affine.for has zero iterations in trivial cases. 1966 static bool hasTrivialZeroTripCount(AffineForOp op) { 1967 Optional<uint64_t> tripCount = getTrivialConstantTripCount(op); 1968 return tripCount && *tripCount == 0; 1969 } 1970 1971 LogicalResult AffineForOp::fold(ArrayRef<Attribute> operands, 1972 SmallVectorImpl<OpFoldResult> &results) { 1973 bool folded = succeeded(foldLoopBounds(*this)); 1974 folded |= succeeded(canonicalizeLoopBounds(*this)); 1975 if (hasTrivialZeroTripCount(*this)) { 1976 // The initial values of the loop-carried variables (iter_args) are the 1977 // results of the op. 1978 results.assign(getIterOperands().begin(), getIterOperands().end()); 1979 folded = true; 1980 } 1981 return success(folded); 1982 } 1983 1984 AffineBound AffineForOp::getLowerBound() { 1985 auto lbMap = getLowerBoundMap(); 1986 return AffineBound(AffineForOp(*this), 0, lbMap.getNumInputs(), lbMap); 1987 } 1988 1989 AffineBound AffineForOp::getUpperBound() { 1990 auto lbMap = getLowerBoundMap(); 1991 auto ubMap = getUpperBoundMap(); 1992 return AffineBound(AffineForOp(*this), lbMap.getNumInputs(), 1993 lbMap.getNumInputs() + ubMap.getNumInputs(), ubMap); 1994 } 1995 1996 void AffineForOp::setLowerBound(ValueRange lbOperands, AffineMap map) { 1997 assert(lbOperands.size() == map.getNumInputs()); 1998 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 1999 2000 SmallVector<Value, 4> newOperands(lbOperands.begin(), lbOperands.end()); 2001 2002 auto ubOperands = getUpperBoundOperands(); 2003 newOperands.append(ubOperands.begin(), ubOperands.end()); 2004 auto iterOperands = getIterOperands(); 2005 newOperands.append(iterOperands.begin(), iterOperands.end()); 2006 (*this)->setOperands(newOperands); 2007 2008 (*this)->setAttr(getLowerBoundAttrStrName(), AffineMapAttr::get(map)); 2009 } 2010 2011 void AffineForOp::setUpperBound(ValueRange ubOperands, AffineMap map) { 2012 assert(ubOperands.size() == map.getNumInputs()); 2013 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 2014 2015 SmallVector<Value, 4> newOperands(getLowerBoundOperands()); 2016 newOperands.append(ubOperands.begin(), ubOperands.end()); 2017 auto iterOperands = getIterOperands(); 2018 newOperands.append(iterOperands.begin(), iterOperands.end()); 2019 (*this)->setOperands(newOperands); 2020 2021 (*this)->setAttr(getUpperBoundAttrStrName(), AffineMapAttr::get(map)); 2022 } 2023 2024 void AffineForOp::setLowerBoundMap(AffineMap map) { 2025 auto lbMap = getLowerBoundMap(); 2026 assert(lbMap.getNumDims() == map.getNumDims() && 2027 lbMap.getNumSymbols() == map.getNumSymbols()); 2028 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 2029 (void)lbMap; 2030 (*this)->setAttr(getLowerBoundAttrStrName(), AffineMapAttr::get(map)); 2031 } 2032 2033 void AffineForOp::setUpperBoundMap(AffineMap map) { 2034 auto ubMap = getUpperBoundMap(); 2035 assert(ubMap.getNumDims() == map.getNumDims() && 2036 ubMap.getNumSymbols() == map.getNumSymbols()); 2037 assert(map.getNumResults() >= 1 && "bound map has at least one result"); 2038 (void)ubMap; 2039 (*this)->setAttr(getUpperBoundAttrStrName(), AffineMapAttr::get(map)); 2040 } 2041 2042 bool AffineForOp::hasConstantLowerBound() { 2043 return getLowerBoundMap().isSingleConstant(); 2044 } 2045 2046 bool AffineForOp::hasConstantUpperBound() { 2047 return getUpperBoundMap().isSingleConstant(); 2048 } 2049 2050 int64_t AffineForOp::getConstantLowerBound() { 2051 return getLowerBoundMap().getSingleConstantResult(); 2052 } 2053 2054 int64_t AffineForOp::getConstantUpperBound() { 2055 return getUpperBoundMap().getSingleConstantResult(); 2056 } 2057 2058 void AffineForOp::setConstantLowerBound(int64_t value) { 2059 setLowerBound({}, AffineMap::getConstantMap(value, getContext())); 2060 } 2061 2062 void AffineForOp::setConstantUpperBound(int64_t value) { 2063 setUpperBound({}, AffineMap::getConstantMap(value, getContext())); 2064 } 2065 2066 AffineForOp::operand_range AffineForOp::getLowerBoundOperands() { 2067 return {operand_begin(), operand_begin() + getLowerBoundMap().getNumInputs()}; 2068 } 2069 2070 AffineForOp::operand_range AffineForOp::getUpperBoundOperands() { 2071 return {operand_begin() + getLowerBoundMap().getNumInputs(), 2072 operand_begin() + getLowerBoundMap().getNumInputs() + 2073 getUpperBoundMap().getNumInputs()}; 2074 } 2075 2076 AffineForOp::operand_range AffineForOp::getControlOperands() { 2077 return {operand_begin(), operand_begin() + getLowerBoundMap().getNumInputs() + 2078 getUpperBoundMap().getNumInputs()}; 2079 } 2080 2081 bool AffineForOp::matchingBoundOperandList() { 2082 auto lbMap = getLowerBoundMap(); 2083 auto ubMap = getUpperBoundMap(); 2084 if (lbMap.getNumDims() != ubMap.getNumDims() || 2085 lbMap.getNumSymbols() != ubMap.getNumSymbols()) 2086 return false; 2087 2088 unsigned numOperands = lbMap.getNumInputs(); 2089 for (unsigned i = 0, e = lbMap.getNumInputs(); i < e; i++) { 2090 // Compare Value 's. 2091 if (getOperand(i) != getOperand(numOperands + i)) 2092 return false; 2093 } 2094 return true; 2095 } 2096 2097 Region &AffineForOp::getLoopBody() { return getRegion(); } 2098 2099 Optional<Value> AffineForOp::getSingleInductionVar() { 2100 return getInductionVar(); 2101 } 2102 2103 Optional<OpFoldResult> AffineForOp::getSingleLowerBound() { 2104 if (!hasConstantLowerBound()) 2105 return llvm::None; 2106 OpBuilder b(getContext()); 2107 return OpFoldResult(b.getI64IntegerAttr(getConstantLowerBound())); 2108 } 2109 2110 Optional<OpFoldResult> AffineForOp::getSingleStep() { 2111 OpBuilder b(getContext()); 2112 return OpFoldResult(b.getI64IntegerAttr(getStep())); 2113 } 2114 2115 Optional<OpFoldResult> AffineForOp::getSingleUpperBound() { 2116 if (!hasConstantUpperBound()) 2117 return llvm::None; 2118 OpBuilder b(getContext()); 2119 return OpFoldResult(b.getI64IntegerAttr(getConstantUpperBound())); 2120 } 2121 2122 /// Returns true if the provided value is the induction variable of a 2123 /// AffineForOp. 2124 bool mlir::isForInductionVar(Value val) { 2125 return getForInductionVarOwner(val) != AffineForOp(); 2126 } 2127 2128 /// Returns the loop parent of an induction variable. If the provided value is 2129 /// not an induction variable, then return nullptr. 2130 AffineForOp mlir::getForInductionVarOwner(Value val) { 2131 auto ivArg = val.dyn_cast<BlockArgument>(); 2132 if (!ivArg || !ivArg.getOwner()) 2133 return AffineForOp(); 2134 auto *containingInst = ivArg.getOwner()->getParent()->getParentOp(); 2135 if (auto forOp = dyn_cast<AffineForOp>(containingInst)) 2136 // Check to make sure `val` is the induction variable, not an iter_arg. 2137 return forOp.getInductionVar() == val ? forOp : AffineForOp(); 2138 return AffineForOp(); 2139 } 2140 2141 /// Extracts the induction variables from a list of AffineForOps and returns 2142 /// them. 2143 void mlir::extractForInductionVars(ArrayRef<AffineForOp> forInsts, 2144 SmallVectorImpl<Value> *ivs) { 2145 ivs->reserve(forInsts.size()); 2146 for (auto forInst : forInsts) 2147 ivs->push_back(forInst.getInductionVar()); 2148 } 2149 2150 /// Builds an affine loop nest, using "loopCreatorFn" to create individual loop 2151 /// operations. 2152 template <typename BoundListTy, typename LoopCreatorTy> 2153 static void buildAffineLoopNestImpl( 2154 OpBuilder &builder, Location loc, BoundListTy lbs, BoundListTy ubs, 2155 ArrayRef<int64_t> steps, 2156 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn, 2157 LoopCreatorTy &&loopCreatorFn) { 2158 assert(lbs.size() == ubs.size() && "Mismatch in number of arguments"); 2159 assert(lbs.size() == steps.size() && "Mismatch in number of arguments"); 2160 2161 // If there are no loops to be constructed, construct the body anyway. 2162 OpBuilder::InsertionGuard guard(builder); 2163 if (lbs.empty()) { 2164 if (bodyBuilderFn) 2165 bodyBuilderFn(builder, loc, ValueRange()); 2166 return; 2167 } 2168 2169 // Create the loops iteratively and store the induction variables. 2170 SmallVector<Value, 4> ivs; 2171 ivs.reserve(lbs.size()); 2172 for (unsigned i = 0, e = lbs.size(); i < e; ++i) { 2173 // Callback for creating the loop body, always creates the terminator. 2174 auto loopBody = [&](OpBuilder &nestedBuilder, Location nestedLoc, Value iv, 2175 ValueRange iterArgs) { 2176 ivs.push_back(iv); 2177 // In the innermost loop, call the body builder. 2178 if (i == e - 1 && bodyBuilderFn) { 2179 OpBuilder::InsertionGuard nestedGuard(nestedBuilder); 2180 bodyBuilderFn(nestedBuilder, nestedLoc, ivs); 2181 } 2182 nestedBuilder.create<AffineYieldOp>(nestedLoc); 2183 }; 2184 2185 // Delegate actual loop creation to the callback in order to dispatch 2186 // between constant- and variable-bound loops. 2187 auto loop = loopCreatorFn(builder, loc, lbs[i], ubs[i], steps[i], loopBody); 2188 builder.setInsertionPointToStart(loop.getBody()); 2189 } 2190 } 2191 2192 /// Creates an affine loop from the bounds known to be constants. 2193 static AffineForOp 2194 buildAffineLoopFromConstants(OpBuilder &builder, Location loc, int64_t lb, 2195 int64_t ub, int64_t step, 2196 AffineForOp::BodyBuilderFn bodyBuilderFn) { 2197 return builder.create<AffineForOp>(loc, lb, ub, step, /*iterArgs=*/llvm::None, 2198 bodyBuilderFn); 2199 } 2200 2201 /// Creates an affine loop from the bounds that may or may not be constants. 2202 static AffineForOp 2203 buildAffineLoopFromValues(OpBuilder &builder, Location loc, Value lb, Value ub, 2204 int64_t step, 2205 AffineForOp::BodyBuilderFn bodyBuilderFn) { 2206 auto lbConst = lb.getDefiningOp<arith::ConstantIndexOp>(); 2207 auto ubConst = ub.getDefiningOp<arith::ConstantIndexOp>(); 2208 if (lbConst && ubConst) 2209 return buildAffineLoopFromConstants(builder, loc, lbConst.value(), 2210 ubConst.value(), step, bodyBuilderFn); 2211 return builder.create<AffineForOp>(loc, lb, builder.getDimIdentityMap(), ub, 2212 builder.getDimIdentityMap(), step, 2213 /*iterArgs=*/llvm::None, bodyBuilderFn); 2214 } 2215 2216 void mlir::buildAffineLoopNest( 2217 OpBuilder &builder, Location loc, ArrayRef<int64_t> lbs, 2218 ArrayRef<int64_t> ubs, ArrayRef<int64_t> steps, 2219 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn) { 2220 buildAffineLoopNestImpl(builder, loc, lbs, ubs, steps, bodyBuilderFn, 2221 buildAffineLoopFromConstants); 2222 } 2223 2224 void mlir::buildAffineLoopNest( 2225 OpBuilder &builder, Location loc, ValueRange lbs, ValueRange ubs, 2226 ArrayRef<int64_t> steps, 2227 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn) { 2228 buildAffineLoopNestImpl(builder, loc, lbs, ubs, steps, bodyBuilderFn, 2229 buildAffineLoopFromValues); 2230 } 2231 2232 AffineForOp mlir::replaceForOpWithNewYields(OpBuilder &b, AffineForOp loop, 2233 ValueRange newIterOperands, 2234 ValueRange newYieldedValues, 2235 ValueRange newIterArgs, 2236 bool replaceLoopResults) { 2237 assert(newIterOperands.size() == newYieldedValues.size() && 2238 "newIterOperands must be of the same size as newYieldedValues"); 2239 // Create a new loop before the existing one, with the extra operands. 2240 OpBuilder::InsertionGuard g(b); 2241 b.setInsertionPoint(loop); 2242 auto operands = llvm::to_vector<4>(loop.getIterOperands()); 2243 operands.append(newIterOperands.begin(), newIterOperands.end()); 2244 SmallVector<Value, 4> lbOperands(loop.getLowerBoundOperands()); 2245 SmallVector<Value, 4> ubOperands(loop.getUpperBoundOperands()); 2246 SmallVector<Value, 4> steps(loop.getStep()); 2247 auto lbMap = loop.getLowerBoundMap(); 2248 auto ubMap = loop.getUpperBoundMap(); 2249 AffineForOp newLoop = 2250 b.create<AffineForOp>(loop.getLoc(), lbOperands, lbMap, ubOperands, ubMap, 2251 loop.getStep(), operands); 2252 // Take the body of the original parent loop. 2253 newLoop.getLoopBody().takeBody(loop.getLoopBody()); 2254 for (Value val : newIterArgs) 2255 newLoop.getLoopBody().addArgument(val.getType(), val.getLoc()); 2256 2257 // Update yield operation with new values to be added. 2258 if (!newYieldedValues.empty()) { 2259 auto yield = cast<AffineYieldOp>(newLoop.getBody()->getTerminator()); 2260 b.setInsertionPoint(yield); 2261 auto yieldOperands = llvm::to_vector<4>(yield.getOperands()); 2262 yieldOperands.append(newYieldedValues.begin(), newYieldedValues.end()); 2263 b.create<AffineYieldOp>(yield.getLoc(), yieldOperands); 2264 yield.erase(); 2265 } 2266 if (replaceLoopResults) { 2267 for (auto it : llvm::zip(loop.getResults(), newLoop.getResults().take_front( 2268 loop.getNumResults()))) { 2269 std::get<0>(it).replaceAllUsesWith(std::get<1>(it)); 2270 } 2271 } 2272 return newLoop; 2273 } 2274 2275 //===----------------------------------------------------------------------===// 2276 // AffineIfOp 2277 //===----------------------------------------------------------------------===// 2278 2279 namespace { 2280 /// Remove else blocks that have nothing other than a zero value yield. 2281 struct SimplifyDeadElse : public OpRewritePattern<AffineIfOp> { 2282 using OpRewritePattern<AffineIfOp>::OpRewritePattern; 2283 2284 LogicalResult matchAndRewrite(AffineIfOp ifOp, 2285 PatternRewriter &rewriter) const override { 2286 if (ifOp.getElseRegion().empty() || 2287 !llvm::hasSingleElement(*ifOp.getElseBlock()) || ifOp.getNumResults()) 2288 return failure(); 2289 2290 rewriter.startRootUpdate(ifOp); 2291 rewriter.eraseBlock(ifOp.getElseBlock()); 2292 rewriter.finalizeRootUpdate(ifOp); 2293 return success(); 2294 } 2295 }; 2296 2297 /// Removes affine.if cond if the condition is always true or false in certain 2298 /// trivial cases. Promotes the then/else block in the parent operation block. 2299 struct AlwaysTrueOrFalseIf : public OpRewritePattern<AffineIfOp> { 2300 using OpRewritePattern<AffineIfOp>::OpRewritePattern; 2301 2302 LogicalResult matchAndRewrite(AffineIfOp op, 2303 PatternRewriter &rewriter) const override { 2304 2305 auto isTriviallyFalse = [](IntegerSet iSet) { 2306 return iSet.isEmptyIntegerSet(); 2307 }; 2308 2309 auto isTriviallyTrue = [](IntegerSet iSet) { 2310 return (iSet.getNumEqualities() == 1 && iSet.getNumInequalities() == 0 && 2311 iSet.getConstraint(0) == 0); 2312 }; 2313 2314 IntegerSet affineIfConditions = op.getIntegerSet(); 2315 Block *blockToMove; 2316 if (isTriviallyFalse(affineIfConditions)) { 2317 // The absence, or equivalently, the emptiness of the else region need not 2318 // be checked when affine.if is returning results because if an affine.if 2319 // operation is returning results, it always has a non-empty else region. 2320 if (op.getNumResults() == 0 && !op.hasElse()) { 2321 // If the else region is absent, or equivalently, empty, remove the 2322 // affine.if operation (which is not returning any results). 2323 rewriter.eraseOp(op); 2324 return success(); 2325 } 2326 blockToMove = op.getElseBlock(); 2327 } else if (isTriviallyTrue(affineIfConditions)) { 2328 blockToMove = op.getThenBlock(); 2329 } else { 2330 return failure(); 2331 } 2332 Operation *blockToMoveTerminator = blockToMove->getTerminator(); 2333 // Promote the "blockToMove" block to the parent operation block between the 2334 // prologue and epilogue of "op". 2335 rewriter.mergeBlockBefore(blockToMove, op); 2336 // Replace the "op" operation with the operands of the 2337 // "blockToMoveTerminator" operation. Note that "blockToMoveTerminator" is 2338 // the affine.yield operation present in the "blockToMove" block. It has no 2339 // operands when affine.if is not returning results and therefore, in that 2340 // case, replaceOp just erases "op". When affine.if is not returning 2341 // results, the affine.yield operation can be omitted. It gets inserted 2342 // implicitly. 2343 rewriter.replaceOp(op, blockToMoveTerminator->getOperands()); 2344 // Erase the "blockToMoveTerminator" operation since it is now in the parent 2345 // operation block, which already has its own terminator. 2346 rewriter.eraseOp(blockToMoveTerminator); 2347 return success(); 2348 } 2349 }; 2350 } // namespace 2351 2352 LogicalResult AffineIfOp::verify() { 2353 // Verify that we have a condition attribute. 2354 // FIXME: This should be specified in the arguments list in ODS. 2355 auto conditionAttr = 2356 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName()); 2357 if (!conditionAttr) 2358 return emitOpError("requires an integer set attribute named 'condition'"); 2359 2360 // Verify that there are enough operands for the condition. 2361 IntegerSet condition = conditionAttr.getValue(); 2362 if (getNumOperands() != condition.getNumInputs()) 2363 return emitOpError("operand count and condition integer set dimension and " 2364 "symbol count must match"); 2365 2366 // Verify that the operands are valid dimension/symbols. 2367 if (failed(verifyDimAndSymbolIdentifiers(*this, getOperands(), 2368 condition.getNumDims()))) 2369 return failure(); 2370 2371 return success(); 2372 } 2373 2374 ParseResult AffineIfOp::parse(OpAsmParser &parser, OperationState &result) { 2375 // Parse the condition attribute set. 2376 IntegerSetAttr conditionAttr; 2377 unsigned numDims; 2378 if (parser.parseAttribute(conditionAttr, 2379 AffineIfOp::getConditionAttrStrName(), 2380 result.attributes) || 2381 parseDimAndSymbolList(parser, result.operands, numDims)) 2382 return failure(); 2383 2384 // Verify the condition operands. 2385 auto set = conditionAttr.getValue(); 2386 if (set.getNumDims() != numDims) 2387 return parser.emitError( 2388 parser.getNameLoc(), 2389 "dim operand count and integer set dim count must match"); 2390 if (numDims + set.getNumSymbols() != result.operands.size()) 2391 return parser.emitError( 2392 parser.getNameLoc(), 2393 "symbol operand count and integer set symbol count must match"); 2394 2395 if (parser.parseOptionalArrowTypeList(result.types)) 2396 return failure(); 2397 2398 // Create the regions for 'then' and 'else'. The latter must be created even 2399 // if it remains empty for the validity of the operation. 2400 result.regions.reserve(2); 2401 Region *thenRegion = result.addRegion(); 2402 Region *elseRegion = result.addRegion(); 2403 2404 // Parse the 'then' region. 2405 if (parser.parseRegion(*thenRegion, {}, {})) 2406 return failure(); 2407 AffineIfOp::ensureTerminator(*thenRegion, parser.getBuilder(), 2408 result.location); 2409 2410 // If we find an 'else' keyword then parse the 'else' region. 2411 if (!parser.parseOptionalKeyword("else")) { 2412 if (parser.parseRegion(*elseRegion, {}, {})) 2413 return failure(); 2414 AffineIfOp::ensureTerminator(*elseRegion, parser.getBuilder(), 2415 result.location); 2416 } 2417 2418 // Parse the optional attribute list. 2419 if (parser.parseOptionalAttrDict(result.attributes)) 2420 return failure(); 2421 2422 return success(); 2423 } 2424 2425 void AffineIfOp::print(OpAsmPrinter &p) { 2426 auto conditionAttr = 2427 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName()); 2428 p << " " << conditionAttr; 2429 printDimAndSymbolList(operand_begin(), operand_end(), 2430 conditionAttr.getValue().getNumDims(), p); 2431 p.printOptionalArrowTypeList(getResultTypes()); 2432 p << ' '; 2433 p.printRegion(getThenRegion(), /*printEntryBlockArgs=*/false, 2434 /*printBlockTerminators=*/getNumResults()); 2435 2436 // Print the 'else' regions if it has any blocks. 2437 auto &elseRegion = this->getElseRegion(); 2438 if (!elseRegion.empty()) { 2439 p << " else "; 2440 p.printRegion(elseRegion, 2441 /*printEntryBlockArgs=*/false, 2442 /*printBlockTerminators=*/getNumResults()); 2443 } 2444 2445 // Print the attribute list. 2446 p.printOptionalAttrDict((*this)->getAttrs(), 2447 /*elidedAttrs=*/getConditionAttrStrName()); 2448 } 2449 2450 IntegerSet AffineIfOp::getIntegerSet() { 2451 return (*this) 2452 ->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName()) 2453 .getValue(); 2454 } 2455 2456 void AffineIfOp::setIntegerSet(IntegerSet newSet) { 2457 (*this)->setAttr(getConditionAttrStrName(), IntegerSetAttr::get(newSet)); 2458 } 2459 2460 void AffineIfOp::setConditional(IntegerSet set, ValueRange operands) { 2461 setIntegerSet(set); 2462 (*this)->setOperands(operands); 2463 } 2464 2465 void AffineIfOp::build(OpBuilder &builder, OperationState &result, 2466 TypeRange resultTypes, IntegerSet set, ValueRange args, 2467 bool withElseRegion) { 2468 assert(resultTypes.empty() || withElseRegion); 2469 result.addTypes(resultTypes); 2470 result.addOperands(args); 2471 result.addAttribute(getConditionAttrStrName(), IntegerSetAttr::get(set)); 2472 2473 Region *thenRegion = result.addRegion(); 2474 thenRegion->push_back(new Block()); 2475 if (resultTypes.empty()) 2476 AffineIfOp::ensureTerminator(*thenRegion, builder, result.location); 2477 2478 Region *elseRegion = result.addRegion(); 2479 if (withElseRegion) { 2480 elseRegion->push_back(new Block()); 2481 if (resultTypes.empty()) 2482 AffineIfOp::ensureTerminator(*elseRegion, builder, result.location); 2483 } 2484 } 2485 2486 void AffineIfOp::build(OpBuilder &builder, OperationState &result, 2487 IntegerSet set, ValueRange args, bool withElseRegion) { 2488 AffineIfOp::build(builder, result, /*resultTypes=*/{}, set, args, 2489 withElseRegion); 2490 } 2491 2492 /// Canonicalize an affine if op's conditional (integer set + operands). 2493 LogicalResult AffineIfOp::fold(ArrayRef<Attribute>, 2494 SmallVectorImpl<OpFoldResult> &) { 2495 auto set = getIntegerSet(); 2496 SmallVector<Value, 4> operands(getOperands()); 2497 canonicalizeSetAndOperands(&set, &operands); 2498 2499 // Any canonicalization change always leads to either a reduction in the 2500 // number of operands or a change in the number of symbolic operands 2501 // (promotion of dims to symbols). 2502 if (operands.size() < getIntegerSet().getNumInputs() || 2503 set.getNumSymbols() > getIntegerSet().getNumSymbols()) { 2504 setConditional(set, operands); 2505 return success(); 2506 } 2507 2508 return failure(); 2509 } 2510 2511 void AffineIfOp::getCanonicalizationPatterns(RewritePatternSet &results, 2512 MLIRContext *context) { 2513 results.add<SimplifyDeadElse, AlwaysTrueOrFalseIf>(context); 2514 } 2515 2516 //===----------------------------------------------------------------------===// 2517 // AffineLoadOp 2518 //===----------------------------------------------------------------------===// 2519 2520 void AffineLoadOp::build(OpBuilder &builder, OperationState &result, 2521 AffineMap map, ValueRange operands) { 2522 assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands"); 2523 result.addOperands(operands); 2524 if (map) 2525 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map)); 2526 auto memrefType = operands[0].getType().cast<MemRefType>(); 2527 result.types.push_back(memrefType.getElementType()); 2528 } 2529 2530 void AffineLoadOp::build(OpBuilder &builder, OperationState &result, 2531 Value memref, AffineMap map, ValueRange mapOperands) { 2532 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info"); 2533 result.addOperands(memref); 2534 result.addOperands(mapOperands); 2535 auto memrefType = memref.getType().cast<MemRefType>(); 2536 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map)); 2537 result.types.push_back(memrefType.getElementType()); 2538 } 2539 2540 void AffineLoadOp::build(OpBuilder &builder, OperationState &result, 2541 Value memref, ValueRange indices) { 2542 auto memrefType = memref.getType().cast<MemRefType>(); 2543 int64_t rank = memrefType.getRank(); 2544 // Create identity map for memrefs with at least one dimension or () -> () 2545 // for zero-dimensional memrefs. 2546 auto map = 2547 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap(); 2548 build(builder, result, memref, map, indices); 2549 } 2550 2551 ParseResult AffineLoadOp::parse(OpAsmParser &parser, OperationState &result) { 2552 auto &builder = parser.getBuilder(); 2553 auto indexTy = builder.getIndexType(); 2554 2555 MemRefType type; 2556 OpAsmParser::UnresolvedOperand memrefInfo; 2557 AffineMapAttr mapAttr; 2558 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands; 2559 return failure( 2560 parser.parseOperand(memrefInfo) || 2561 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, 2562 AffineLoadOp::getMapAttrStrName(), 2563 result.attributes) || 2564 parser.parseOptionalAttrDict(result.attributes) || 2565 parser.parseColonType(type) || 2566 parser.resolveOperand(memrefInfo, type, result.operands) || 2567 parser.resolveOperands(mapOperands, indexTy, result.operands) || 2568 parser.addTypeToList(type.getElementType(), result.types)); 2569 } 2570 2571 void AffineLoadOp::print(OpAsmPrinter &p) { 2572 p << " " << getMemRef() << '['; 2573 if (AffineMapAttr mapAttr = 2574 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName())) 2575 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 2576 p << ']'; 2577 p.printOptionalAttrDict((*this)->getAttrs(), 2578 /*elidedAttrs=*/{getMapAttrStrName()}); 2579 p << " : " << getMemRefType(); 2580 } 2581 2582 /// Verify common indexing invariants of affine.load, affine.store, 2583 /// affine.vector_load and affine.vector_store. 2584 static LogicalResult 2585 verifyMemoryOpIndexing(Operation *op, AffineMapAttr mapAttr, 2586 Operation::operand_range mapOperands, 2587 MemRefType memrefType, unsigned numIndexOperands) { 2588 if (mapAttr) { 2589 AffineMap map = mapAttr.getValue(); 2590 if (map.getNumResults() != memrefType.getRank()) 2591 return op->emitOpError("affine map num results must equal memref rank"); 2592 if (map.getNumInputs() != numIndexOperands) 2593 return op->emitOpError("expects as many subscripts as affine map inputs"); 2594 } else { 2595 if (memrefType.getRank() != numIndexOperands) 2596 return op->emitOpError( 2597 "expects the number of subscripts to be equal to memref rank"); 2598 } 2599 2600 Region *scope = getAffineScope(op); 2601 for (auto idx : mapOperands) { 2602 if (!idx.getType().isIndex()) 2603 return op->emitOpError("index to load must have 'index' type"); 2604 if (!isValidAffineIndexOperand(idx, scope)) 2605 return op->emitOpError("index must be a dimension or symbol identifier"); 2606 } 2607 2608 return success(); 2609 } 2610 2611 LogicalResult AffineLoadOp::verify() { 2612 auto memrefType = getMemRefType(); 2613 if (getType() != memrefType.getElementType()) 2614 return emitOpError("result type must match element type of memref"); 2615 2616 if (failed(verifyMemoryOpIndexing( 2617 getOperation(), 2618 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()), 2619 getMapOperands(), memrefType, 2620 /*numIndexOperands=*/getNumOperands() - 1))) 2621 return failure(); 2622 2623 return success(); 2624 } 2625 2626 void AffineLoadOp::getCanonicalizationPatterns(RewritePatternSet &results, 2627 MLIRContext *context) { 2628 results.add<SimplifyAffineOp<AffineLoadOp>>(context); 2629 } 2630 2631 OpFoldResult AffineLoadOp::fold(ArrayRef<Attribute> cstOperands) { 2632 /// load(memrefcast) -> load 2633 if (succeeded(foldMemRefCast(*this))) 2634 return getResult(); 2635 2636 // Fold load from a global constant memref. 2637 auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>(); 2638 if (!getGlobalOp) 2639 return {}; 2640 // Get to the memref.global defining the symbol. 2641 auto *symbolTableOp = getGlobalOp->getParentWithTrait<OpTrait::SymbolTable>(); 2642 if (!symbolTableOp) 2643 return {}; 2644 auto global = dyn_cast_or_null<memref::GlobalOp>( 2645 SymbolTable::lookupSymbolIn(symbolTableOp, getGlobalOp.getNameAttr())); 2646 if (!global) 2647 return {}; 2648 2649 // Check if the global memref is a constant. 2650 auto cstAttr = 2651 global.getConstantInitValue().dyn_cast_or_null<DenseElementsAttr>(); 2652 if (!cstAttr) 2653 return {}; 2654 // If it's a splat constant, we can fold irrespective of indices. 2655 if (auto splatAttr = cstAttr.dyn_cast<SplatElementsAttr>()) 2656 return splatAttr.getSplatValue<Attribute>(); 2657 // Otherwise, we can fold only if we know the indices. 2658 if (!getAffineMap().isConstant()) 2659 return {}; 2660 auto indices = llvm::to_vector<4>( 2661 llvm::map_range(getAffineMap().getConstantResults(), 2662 [](int64_t v) -> uint64_t { return v; })); 2663 return cstAttr.getValues<Attribute>()[indices]; 2664 } 2665 2666 //===----------------------------------------------------------------------===// 2667 // AffineStoreOp 2668 //===----------------------------------------------------------------------===// 2669 2670 void AffineStoreOp::build(OpBuilder &builder, OperationState &result, 2671 Value valueToStore, Value memref, AffineMap map, 2672 ValueRange mapOperands) { 2673 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info"); 2674 result.addOperands(valueToStore); 2675 result.addOperands(memref); 2676 result.addOperands(mapOperands); 2677 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map)); 2678 } 2679 2680 // Use identity map. 2681 void AffineStoreOp::build(OpBuilder &builder, OperationState &result, 2682 Value valueToStore, Value memref, 2683 ValueRange indices) { 2684 auto memrefType = memref.getType().cast<MemRefType>(); 2685 int64_t rank = memrefType.getRank(); 2686 // Create identity map for memrefs with at least one dimension or () -> () 2687 // for zero-dimensional memrefs. 2688 auto map = 2689 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap(); 2690 build(builder, result, valueToStore, memref, map, indices); 2691 } 2692 2693 ParseResult AffineStoreOp::parse(OpAsmParser &parser, OperationState &result) { 2694 auto indexTy = parser.getBuilder().getIndexType(); 2695 2696 MemRefType type; 2697 OpAsmParser::UnresolvedOperand storeValueInfo; 2698 OpAsmParser::UnresolvedOperand memrefInfo; 2699 AffineMapAttr mapAttr; 2700 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands; 2701 return failure(parser.parseOperand(storeValueInfo) || parser.parseComma() || 2702 parser.parseOperand(memrefInfo) || 2703 parser.parseAffineMapOfSSAIds( 2704 mapOperands, mapAttr, AffineStoreOp::getMapAttrStrName(), 2705 result.attributes) || 2706 parser.parseOptionalAttrDict(result.attributes) || 2707 parser.parseColonType(type) || 2708 parser.resolveOperand(storeValueInfo, type.getElementType(), 2709 result.operands) || 2710 parser.resolveOperand(memrefInfo, type, result.operands) || 2711 parser.resolveOperands(mapOperands, indexTy, result.operands)); 2712 } 2713 2714 void AffineStoreOp::print(OpAsmPrinter &p) { 2715 p << " " << getValueToStore(); 2716 p << ", " << getMemRef() << '['; 2717 if (AffineMapAttr mapAttr = 2718 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName())) 2719 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 2720 p << ']'; 2721 p.printOptionalAttrDict((*this)->getAttrs(), 2722 /*elidedAttrs=*/{getMapAttrStrName()}); 2723 p << " : " << getMemRefType(); 2724 } 2725 2726 LogicalResult AffineStoreOp::verify() { 2727 // The value to store must have the same type as memref element type. 2728 auto memrefType = getMemRefType(); 2729 if (getValueToStore().getType() != memrefType.getElementType()) 2730 return emitOpError( 2731 "value to store must have the same type as memref element type"); 2732 2733 if (failed(verifyMemoryOpIndexing( 2734 getOperation(), 2735 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()), 2736 getMapOperands(), memrefType, 2737 /*numIndexOperands=*/getNumOperands() - 2))) 2738 return failure(); 2739 2740 return success(); 2741 } 2742 2743 void AffineStoreOp::getCanonicalizationPatterns(RewritePatternSet &results, 2744 MLIRContext *context) { 2745 results.add<SimplifyAffineOp<AffineStoreOp>>(context); 2746 } 2747 2748 LogicalResult AffineStoreOp::fold(ArrayRef<Attribute> cstOperands, 2749 SmallVectorImpl<OpFoldResult> &results) { 2750 /// store(memrefcast) -> store 2751 return foldMemRefCast(*this, getValueToStore()); 2752 } 2753 2754 //===----------------------------------------------------------------------===// 2755 // AffineMinMaxOpBase 2756 //===----------------------------------------------------------------------===// 2757 2758 template <typename T> static LogicalResult verifyAffineMinMaxOp(T op) { 2759 // Verify that operand count matches affine map dimension and symbol count. 2760 if (op.getNumOperands() != 2761 op.getMap().getNumDims() + op.getMap().getNumSymbols()) 2762 return op.emitOpError( 2763 "operand count and affine map dimension and symbol count must match"); 2764 return success(); 2765 } 2766 2767 template <typename T> static void printAffineMinMaxOp(OpAsmPrinter &p, T op) { 2768 p << ' ' << op->getAttr(T::getMapAttrStrName()); 2769 auto operands = op.getOperands(); 2770 unsigned numDims = op.getMap().getNumDims(); 2771 p << '(' << operands.take_front(numDims) << ')'; 2772 2773 if (operands.size() != numDims) 2774 p << '[' << operands.drop_front(numDims) << ']'; 2775 p.printOptionalAttrDict(op->getAttrs(), 2776 /*elidedAttrs=*/{T::getMapAttrStrName()}); 2777 } 2778 2779 template <typename T> 2780 static ParseResult parseAffineMinMaxOp(OpAsmParser &parser, 2781 OperationState &result) { 2782 auto &builder = parser.getBuilder(); 2783 auto indexType = builder.getIndexType(); 2784 SmallVector<OpAsmParser::UnresolvedOperand, 8> dimInfos; 2785 SmallVector<OpAsmParser::UnresolvedOperand, 8> symInfos; 2786 AffineMapAttr mapAttr; 2787 return failure( 2788 parser.parseAttribute(mapAttr, T::getMapAttrStrName(), 2789 result.attributes) || 2790 parser.parseOperandList(dimInfos, OpAsmParser::Delimiter::Paren) || 2791 parser.parseOperandList(symInfos, 2792 OpAsmParser::Delimiter::OptionalSquare) || 2793 parser.parseOptionalAttrDict(result.attributes) || 2794 parser.resolveOperands(dimInfos, indexType, result.operands) || 2795 parser.resolveOperands(symInfos, indexType, result.operands) || 2796 parser.addTypeToList(indexType, result.types)); 2797 } 2798 2799 /// Fold an affine min or max operation with the given operands. The operand 2800 /// list may contain nulls, which are interpreted as the operand not being a 2801 /// constant. 2802 template <typename T> 2803 static OpFoldResult foldMinMaxOp(T op, ArrayRef<Attribute> operands) { 2804 static_assert(llvm::is_one_of<T, AffineMinOp, AffineMaxOp>::value, 2805 "expected affine min or max op"); 2806 2807 // Fold the affine map. 2808 // TODO: Fold more cases: 2809 // min(some_affine, some_affine + constant, ...), etc. 2810 SmallVector<int64_t, 2> results; 2811 auto foldedMap = op.getMap().partialConstantFold(operands, &results); 2812 2813 // If some of the map results are not constant, try changing the map in-place. 2814 if (results.empty()) { 2815 // If the map is the same, report that folding did not happen. 2816 if (foldedMap == op.getMap()) 2817 return {}; 2818 op->setAttr("map", AffineMapAttr::get(foldedMap)); 2819 return op.getResult(); 2820 } 2821 2822 // Otherwise, completely fold the op into a constant. 2823 auto resultIt = std::is_same<T, AffineMinOp>::value 2824 ? std::min_element(results.begin(), results.end()) 2825 : std::max_element(results.begin(), results.end()); 2826 if (resultIt == results.end()) 2827 return {}; 2828 return IntegerAttr::get(IndexType::get(op.getContext()), *resultIt); 2829 } 2830 2831 /// Remove duplicated expressions in affine min/max ops. 2832 template <typename T> 2833 struct DeduplicateAffineMinMaxExpressions : public OpRewritePattern<T> { 2834 using OpRewritePattern<T>::OpRewritePattern; 2835 2836 LogicalResult matchAndRewrite(T affineOp, 2837 PatternRewriter &rewriter) const override { 2838 AffineMap oldMap = affineOp.getAffineMap(); 2839 2840 SmallVector<AffineExpr, 4> newExprs; 2841 for (AffineExpr expr : oldMap.getResults()) { 2842 // This is a linear scan over newExprs, but it should be fine given that 2843 // we typically just have a few expressions per op. 2844 if (!llvm::is_contained(newExprs, expr)) 2845 newExprs.push_back(expr); 2846 } 2847 2848 if (newExprs.size() == oldMap.getNumResults()) 2849 return failure(); 2850 2851 auto newMap = AffineMap::get(oldMap.getNumDims(), oldMap.getNumSymbols(), 2852 newExprs, rewriter.getContext()); 2853 rewriter.replaceOpWithNewOp<T>(affineOp, newMap, affineOp.getMapOperands()); 2854 2855 return success(); 2856 } 2857 }; 2858 2859 /// Merge an affine min/max op to its consumers if its consumer is also an 2860 /// affine min/max op. 2861 /// 2862 /// This pattern requires the producer affine min/max op is bound to a 2863 /// dimension/symbol that is used as a standalone expression in the consumer 2864 /// affine op's map. 2865 /// 2866 /// For example, a pattern like the following: 2867 /// 2868 /// %0 = affine.min affine_map<()[s0] -> (s0 + 16, s0 * 8)> ()[%sym1] 2869 /// %1 = affine.min affine_map<(d0)[s0] -> (s0 + 4, d0)> (%0)[%sym2] 2870 /// 2871 /// Can be turned into: 2872 /// 2873 /// %1 = affine.min affine_map< 2874 /// ()[s0, s1] -> (s0 + 4, s1 + 16, s1 * 8)> ()[%sym2, %sym1] 2875 template <typename T> struct MergeAffineMinMaxOp : public OpRewritePattern<T> { 2876 using OpRewritePattern<T>::OpRewritePattern; 2877 2878 LogicalResult matchAndRewrite(T affineOp, 2879 PatternRewriter &rewriter) const override { 2880 AffineMap oldMap = affineOp.getAffineMap(); 2881 ValueRange dimOperands = 2882 affineOp.getMapOperands().take_front(oldMap.getNumDims()); 2883 ValueRange symOperands = 2884 affineOp.getMapOperands().take_back(oldMap.getNumSymbols()); 2885 2886 auto newDimOperands = llvm::to_vector<8>(dimOperands); 2887 auto newSymOperands = llvm::to_vector<8>(symOperands); 2888 SmallVector<AffineExpr, 4> newExprs; 2889 SmallVector<T, 4> producerOps; 2890 2891 // Go over each expression to see whether it's a single dimension/symbol 2892 // with the corresponding operand which is the result of another affine 2893 // min/max op. If So it can be merged into this affine op. 2894 for (AffineExpr expr : oldMap.getResults()) { 2895 if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) { 2896 Value symValue = symOperands[symExpr.getPosition()]; 2897 if (auto producerOp = symValue.getDefiningOp<T>()) { 2898 producerOps.push_back(producerOp); 2899 continue; 2900 } 2901 } else if (auto dimExpr = expr.dyn_cast<AffineDimExpr>()) { 2902 Value dimValue = dimOperands[dimExpr.getPosition()]; 2903 if (auto producerOp = dimValue.getDefiningOp<T>()) { 2904 producerOps.push_back(producerOp); 2905 continue; 2906 } 2907 } 2908 // For the above cases we will remove the expression by merging the 2909 // producer affine min/max's affine expressions. Otherwise we need to 2910 // keep the existing expression. 2911 newExprs.push_back(expr); 2912 } 2913 2914 if (producerOps.empty()) 2915 return failure(); 2916 2917 unsigned numUsedDims = oldMap.getNumDims(); 2918 unsigned numUsedSyms = oldMap.getNumSymbols(); 2919 2920 // Now go over all producer affine ops and merge their expressions. 2921 for (T producerOp : producerOps) { 2922 AffineMap producerMap = producerOp.getAffineMap(); 2923 unsigned numProducerDims = producerMap.getNumDims(); 2924 unsigned numProducerSyms = producerMap.getNumSymbols(); 2925 2926 // Collect all dimension/symbol values. 2927 ValueRange dimValues = 2928 producerOp.getMapOperands().take_front(numProducerDims); 2929 ValueRange symValues = 2930 producerOp.getMapOperands().take_back(numProducerSyms); 2931 newDimOperands.append(dimValues.begin(), dimValues.end()); 2932 newSymOperands.append(symValues.begin(), symValues.end()); 2933 2934 // For expressions we need to shift to avoid overlap. 2935 for (AffineExpr expr : producerMap.getResults()) { 2936 newExprs.push_back(expr.shiftDims(numProducerDims, numUsedDims) 2937 .shiftSymbols(numProducerSyms, numUsedSyms)); 2938 } 2939 2940 numUsedDims += numProducerDims; 2941 numUsedSyms += numProducerSyms; 2942 } 2943 2944 auto newMap = AffineMap::get(numUsedDims, numUsedSyms, newExprs, 2945 rewriter.getContext()); 2946 auto newOperands = 2947 llvm::to_vector<8>(llvm::concat<Value>(newDimOperands, newSymOperands)); 2948 rewriter.replaceOpWithNewOp<T>(affineOp, newMap, newOperands); 2949 2950 return success(); 2951 } 2952 }; 2953 2954 /// Canonicalize the result expression order of an affine map and return success 2955 /// if the order changed. 2956 /// 2957 /// The function flattens the map's affine expressions to coefficient arrays and 2958 /// sorts them in lexicographic order. A coefficient array contains a multiplier 2959 /// for every dimension/symbol and a constant term. The canonicalization fails 2960 /// if a result expression is not pure or if the flattening requires local 2961 /// variables that, unlike dimensions and symbols, have no global order. 2962 static LogicalResult canonicalizeMapExprAndTermOrder(AffineMap &map) { 2963 SmallVector<SmallVector<int64_t>> flattenedExprs; 2964 for (const AffineExpr &resultExpr : map.getResults()) { 2965 // Fail if the expression is not pure. 2966 if (!resultExpr.isPureAffine()) 2967 return failure(); 2968 2969 SimpleAffineExprFlattener flattener(map.getNumDims(), map.getNumSymbols()); 2970 flattener.walkPostOrder(resultExpr); 2971 2972 // Fail if the flattened expression has local variables. 2973 if (flattener.operandExprStack.back().size() != 2974 map.getNumDims() + map.getNumSymbols() + 1) 2975 return failure(); 2976 2977 flattenedExprs.emplace_back(flattener.operandExprStack.back().begin(), 2978 flattener.operandExprStack.back().end()); 2979 } 2980 2981 // Fail if sorting is not necessary. 2982 if (llvm::is_sorted(flattenedExprs)) 2983 return failure(); 2984 2985 // Reorder the result expressions according to their flattened form. 2986 SmallVector<unsigned> resultPermutation = 2987 llvm::to_vector(llvm::seq<unsigned>(0, map.getNumResults())); 2988 llvm::sort(resultPermutation, [&](unsigned lhs, unsigned rhs) { 2989 return flattenedExprs[lhs] < flattenedExprs[rhs]; 2990 }); 2991 SmallVector<AffineExpr> newExprs; 2992 for (unsigned idx : resultPermutation) 2993 newExprs.push_back(map.getResult(idx)); 2994 2995 map = AffineMap::get(map.getNumDims(), map.getNumSymbols(), newExprs, 2996 map.getContext()); 2997 return success(); 2998 } 2999 3000 /// Canonicalize the affine map result expression order of an affine min/max 3001 /// operation. 3002 /// 3003 /// The pattern calls `canonicalizeMapExprAndTermOrder` to order the result 3004 /// expressions and replaces the operation if the order changed. 3005 /// 3006 /// For example, the following operation: 3007 /// 3008 /// %0 = affine.min affine_map<(d0, d1) -> (d0 + d1, d1 + 16, 32)> (%i0, %i1) 3009 /// 3010 /// Turns into: 3011 /// 3012 /// %0 = affine.min affine_map<(d0, d1) -> (32, d1 + 16, d0 + d1)> (%i0, %i1) 3013 template <typename T> 3014 struct CanonicalizeAffineMinMaxOpExprAndTermOrder : public OpRewritePattern<T> { 3015 using OpRewritePattern<T>::OpRewritePattern; 3016 3017 LogicalResult matchAndRewrite(T affineOp, 3018 PatternRewriter &rewriter) const override { 3019 AffineMap map = affineOp.getAffineMap(); 3020 if (failed(canonicalizeMapExprAndTermOrder(map))) 3021 return failure(); 3022 3023 rewriter.replaceOpWithNewOp<T>(affineOp, map, affineOp.getMapOperands()); 3024 return success(); 3025 } 3026 }; 3027 3028 template <typename T> 3029 struct CanonicalizeSingleResultAffineMinMaxOp : public OpRewritePattern<T> { 3030 using OpRewritePattern<T>::OpRewritePattern; 3031 3032 LogicalResult matchAndRewrite(T affineOp, 3033 PatternRewriter &rewriter) const override { 3034 if (affineOp.getMap().getNumResults() != 1) 3035 return failure(); 3036 rewriter.replaceOpWithNewOp<AffineApplyOp>(affineOp, affineOp.getMap(), 3037 affineOp.getOperands()); 3038 return success(); 3039 } 3040 }; 3041 3042 //===----------------------------------------------------------------------===// 3043 // AffineMinOp 3044 //===----------------------------------------------------------------------===// 3045 // 3046 // %0 = affine.min (d0) -> (1000, d0 + 512) (%i0) 3047 // 3048 3049 OpFoldResult AffineMinOp::fold(ArrayRef<Attribute> operands) { 3050 return foldMinMaxOp(*this, operands); 3051 } 3052 3053 void AffineMinOp::getCanonicalizationPatterns(RewritePatternSet &patterns, 3054 MLIRContext *context) { 3055 patterns.add<CanonicalizeSingleResultAffineMinMaxOp<AffineMinOp>, 3056 DeduplicateAffineMinMaxExpressions<AffineMinOp>, 3057 MergeAffineMinMaxOp<AffineMinOp>, SimplifyAffineOp<AffineMinOp>, 3058 CanonicalizeAffineMinMaxOpExprAndTermOrder<AffineMinOp>>( 3059 context); 3060 } 3061 3062 LogicalResult AffineMinOp::verify() { return verifyAffineMinMaxOp(*this); } 3063 3064 ParseResult AffineMinOp::parse(OpAsmParser &parser, OperationState &result) { 3065 return parseAffineMinMaxOp<AffineMinOp>(parser, result); 3066 } 3067 3068 void AffineMinOp::print(OpAsmPrinter &p) { printAffineMinMaxOp(p, *this); } 3069 3070 //===----------------------------------------------------------------------===// 3071 // AffineMaxOp 3072 //===----------------------------------------------------------------------===// 3073 // 3074 // %0 = affine.max (d0) -> (1000, d0 + 512) (%i0) 3075 // 3076 3077 OpFoldResult AffineMaxOp::fold(ArrayRef<Attribute> operands) { 3078 return foldMinMaxOp(*this, operands); 3079 } 3080 3081 void AffineMaxOp::getCanonicalizationPatterns(RewritePatternSet &patterns, 3082 MLIRContext *context) { 3083 patterns.add<CanonicalizeSingleResultAffineMinMaxOp<AffineMaxOp>, 3084 DeduplicateAffineMinMaxExpressions<AffineMaxOp>, 3085 MergeAffineMinMaxOp<AffineMaxOp>, SimplifyAffineOp<AffineMaxOp>, 3086 CanonicalizeAffineMinMaxOpExprAndTermOrder<AffineMaxOp>>( 3087 context); 3088 } 3089 3090 LogicalResult AffineMaxOp::verify() { return verifyAffineMinMaxOp(*this); } 3091 3092 ParseResult AffineMaxOp::parse(OpAsmParser &parser, OperationState &result) { 3093 return parseAffineMinMaxOp<AffineMaxOp>(parser, result); 3094 } 3095 3096 void AffineMaxOp::print(OpAsmPrinter &p) { printAffineMinMaxOp(p, *this); } 3097 3098 //===----------------------------------------------------------------------===// 3099 // AffinePrefetchOp 3100 //===----------------------------------------------------------------------===// 3101 3102 // 3103 // affine.prefetch %0[%i, %j + 5], read, locality<3>, data : memref<400x400xi32> 3104 // 3105 ParseResult AffinePrefetchOp::parse(OpAsmParser &parser, 3106 OperationState &result) { 3107 auto &builder = parser.getBuilder(); 3108 auto indexTy = builder.getIndexType(); 3109 3110 MemRefType type; 3111 OpAsmParser::UnresolvedOperand memrefInfo; 3112 IntegerAttr hintInfo; 3113 auto i32Type = parser.getBuilder().getIntegerType(32); 3114 StringRef readOrWrite, cacheType; 3115 3116 AffineMapAttr mapAttr; 3117 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands; 3118 if (parser.parseOperand(memrefInfo) || 3119 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, 3120 AffinePrefetchOp::getMapAttrStrName(), 3121 result.attributes) || 3122 parser.parseComma() || parser.parseKeyword(&readOrWrite) || 3123 parser.parseComma() || parser.parseKeyword("locality") || 3124 parser.parseLess() || 3125 parser.parseAttribute(hintInfo, i32Type, 3126 AffinePrefetchOp::getLocalityHintAttrStrName(), 3127 result.attributes) || 3128 parser.parseGreater() || parser.parseComma() || 3129 parser.parseKeyword(&cacheType) || 3130 parser.parseOptionalAttrDict(result.attributes) || 3131 parser.parseColonType(type) || 3132 parser.resolveOperand(memrefInfo, type, result.operands) || 3133 parser.resolveOperands(mapOperands, indexTy, result.operands)) 3134 return failure(); 3135 3136 if (!readOrWrite.equals("read") && !readOrWrite.equals("write")) 3137 return parser.emitError(parser.getNameLoc(), 3138 "rw specifier has to be 'read' or 'write'"); 3139 result.addAttribute( 3140 AffinePrefetchOp::getIsWriteAttrStrName(), 3141 parser.getBuilder().getBoolAttr(readOrWrite.equals("write"))); 3142 3143 if (!cacheType.equals("data") && !cacheType.equals("instr")) 3144 return parser.emitError(parser.getNameLoc(), 3145 "cache type has to be 'data' or 'instr'"); 3146 3147 result.addAttribute( 3148 AffinePrefetchOp::getIsDataCacheAttrStrName(), 3149 parser.getBuilder().getBoolAttr(cacheType.equals("data"))); 3150 3151 return success(); 3152 } 3153 3154 void AffinePrefetchOp::print(OpAsmPrinter &p) { 3155 p << " " << getMemref() << '['; 3156 AffineMapAttr mapAttr = 3157 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()); 3158 if (mapAttr) 3159 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 3160 p << ']' << ", " << (getIsWrite() ? "write" : "read") << ", " 3161 << "locality<" << getLocalityHint() << ">, " 3162 << (getIsDataCache() ? "data" : "instr"); 3163 p.printOptionalAttrDict( 3164 (*this)->getAttrs(), 3165 /*elidedAttrs=*/{getMapAttrStrName(), getLocalityHintAttrStrName(), 3166 getIsDataCacheAttrStrName(), getIsWriteAttrStrName()}); 3167 p << " : " << getMemRefType(); 3168 } 3169 3170 LogicalResult AffinePrefetchOp::verify() { 3171 auto mapAttr = (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()); 3172 if (mapAttr) { 3173 AffineMap map = mapAttr.getValue(); 3174 if (map.getNumResults() != getMemRefType().getRank()) 3175 return emitOpError("affine.prefetch affine map num results must equal" 3176 " memref rank"); 3177 if (map.getNumInputs() + 1 != getNumOperands()) 3178 return emitOpError("too few operands"); 3179 } else { 3180 if (getNumOperands() != 1) 3181 return emitOpError("too few operands"); 3182 } 3183 3184 Region *scope = getAffineScope(*this); 3185 for (auto idx : getMapOperands()) { 3186 if (!isValidAffineIndexOperand(idx, scope)) 3187 return emitOpError("index must be a dimension or symbol identifier"); 3188 } 3189 return success(); 3190 } 3191 3192 void AffinePrefetchOp::getCanonicalizationPatterns(RewritePatternSet &results, 3193 MLIRContext *context) { 3194 // prefetch(memrefcast) -> prefetch 3195 results.add<SimplifyAffineOp<AffinePrefetchOp>>(context); 3196 } 3197 3198 LogicalResult AffinePrefetchOp::fold(ArrayRef<Attribute> cstOperands, 3199 SmallVectorImpl<OpFoldResult> &results) { 3200 /// prefetch(memrefcast) -> prefetch 3201 return foldMemRefCast(*this); 3202 } 3203 3204 //===----------------------------------------------------------------------===// 3205 // AffineParallelOp 3206 //===----------------------------------------------------------------------===// 3207 3208 void AffineParallelOp::build(OpBuilder &builder, OperationState &result, 3209 TypeRange resultTypes, 3210 ArrayRef<arith::AtomicRMWKind> reductions, 3211 ArrayRef<int64_t> ranges) { 3212 SmallVector<AffineMap> lbs(ranges.size(), builder.getConstantAffineMap(0)); 3213 auto ubs = llvm::to_vector<4>(llvm::map_range(ranges, [&](int64_t value) { 3214 return builder.getConstantAffineMap(value); 3215 })); 3216 SmallVector<int64_t> steps(ranges.size(), 1); 3217 build(builder, result, resultTypes, reductions, lbs, /*lbArgs=*/{}, ubs, 3218 /*ubArgs=*/{}, steps); 3219 } 3220 3221 void AffineParallelOp::build(OpBuilder &builder, OperationState &result, 3222 TypeRange resultTypes, 3223 ArrayRef<arith::AtomicRMWKind> reductions, 3224 ArrayRef<AffineMap> lbMaps, ValueRange lbArgs, 3225 ArrayRef<AffineMap> ubMaps, ValueRange ubArgs, 3226 ArrayRef<int64_t> steps) { 3227 assert(llvm::all_of(lbMaps, 3228 [lbMaps](AffineMap m) { 3229 return m.getNumDims() == lbMaps[0].getNumDims() && 3230 m.getNumSymbols() == lbMaps[0].getNumSymbols(); 3231 }) && 3232 "expected all lower bounds maps to have the same number of dimensions " 3233 "and symbols"); 3234 assert(llvm::all_of(ubMaps, 3235 [ubMaps](AffineMap m) { 3236 return m.getNumDims() == ubMaps[0].getNumDims() && 3237 m.getNumSymbols() == ubMaps[0].getNumSymbols(); 3238 }) && 3239 "expected all upper bounds maps to have the same number of dimensions " 3240 "and symbols"); 3241 assert((lbMaps.empty() || lbMaps[0].getNumInputs() == lbArgs.size()) && 3242 "expected lower bound maps to have as many inputs as lower bound " 3243 "operands"); 3244 assert((ubMaps.empty() || ubMaps[0].getNumInputs() == ubArgs.size()) && 3245 "expected upper bound maps to have as many inputs as upper bound " 3246 "operands"); 3247 3248 result.addTypes(resultTypes); 3249 3250 // Convert the reductions to integer attributes. 3251 SmallVector<Attribute, 4> reductionAttrs; 3252 for (arith::AtomicRMWKind reduction : reductions) 3253 reductionAttrs.push_back( 3254 builder.getI64IntegerAttr(static_cast<int64_t>(reduction))); 3255 result.addAttribute(getReductionsAttrStrName(), 3256 builder.getArrayAttr(reductionAttrs)); 3257 3258 // Concatenates maps defined in the same input space (same dimensions and 3259 // symbols), assumes there is at least one map. 3260 auto concatMapsSameInput = [&builder](ArrayRef<AffineMap> maps, 3261 SmallVectorImpl<int32_t> &groups) { 3262 if (maps.empty()) 3263 return AffineMap::get(builder.getContext()); 3264 SmallVector<AffineExpr> exprs; 3265 groups.reserve(groups.size() + maps.size()); 3266 exprs.reserve(maps.size()); 3267 for (AffineMap m : maps) { 3268 llvm::append_range(exprs, m.getResults()); 3269 groups.push_back(m.getNumResults()); 3270 } 3271 return AffineMap::get(maps[0].getNumDims(), maps[0].getNumSymbols(), exprs, 3272 maps[0].getContext()); 3273 }; 3274 3275 // Set up the bounds. 3276 SmallVector<int32_t> lbGroups, ubGroups; 3277 AffineMap lbMap = concatMapsSameInput(lbMaps, lbGroups); 3278 AffineMap ubMap = concatMapsSameInput(ubMaps, ubGroups); 3279 result.addAttribute(getLowerBoundsMapAttrStrName(), 3280 AffineMapAttr::get(lbMap)); 3281 result.addAttribute(getLowerBoundsGroupsAttrStrName(), 3282 builder.getI32TensorAttr(lbGroups)); 3283 result.addAttribute(getUpperBoundsMapAttrStrName(), 3284 AffineMapAttr::get(ubMap)); 3285 result.addAttribute(getUpperBoundsGroupsAttrStrName(), 3286 builder.getI32TensorAttr(ubGroups)); 3287 result.addAttribute(getStepsAttrStrName(), builder.getI64ArrayAttr(steps)); 3288 result.addOperands(lbArgs); 3289 result.addOperands(ubArgs); 3290 3291 // Create a region and a block for the body. 3292 auto *bodyRegion = result.addRegion(); 3293 auto *body = new Block(); 3294 // Add all the block arguments. 3295 for (unsigned i = 0, e = steps.size(); i < e; ++i) 3296 body->addArgument(IndexType::get(builder.getContext()), result.location); 3297 bodyRegion->push_back(body); 3298 if (resultTypes.empty()) 3299 ensureTerminator(*bodyRegion, builder, result.location); 3300 } 3301 3302 Region &AffineParallelOp::getLoopBody() { return getRegion(); } 3303 3304 unsigned AffineParallelOp::getNumDims() { return getSteps().size(); } 3305 3306 AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() { 3307 return getOperands().take_front(getLowerBoundsMap().getNumInputs()); 3308 } 3309 3310 AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() { 3311 return getOperands().drop_front(getLowerBoundsMap().getNumInputs()); 3312 } 3313 3314 AffineMap AffineParallelOp::getLowerBoundMap(unsigned pos) { 3315 auto values = getLowerBoundsGroups().getValues<int32_t>(); 3316 unsigned start = 0; 3317 for (unsigned i = 0; i < pos; ++i) 3318 start += values[i]; 3319 return getLowerBoundsMap().getSliceMap(start, values[pos]); 3320 } 3321 3322 AffineMap AffineParallelOp::getUpperBoundMap(unsigned pos) { 3323 auto values = getUpperBoundsGroups().getValues<int32_t>(); 3324 unsigned start = 0; 3325 for (unsigned i = 0; i < pos; ++i) 3326 start += values[i]; 3327 return getUpperBoundsMap().getSliceMap(start, values[pos]); 3328 } 3329 3330 AffineValueMap AffineParallelOp::getLowerBoundsValueMap() { 3331 return AffineValueMap(getLowerBoundsMap(), getLowerBoundsOperands()); 3332 } 3333 3334 AffineValueMap AffineParallelOp::getUpperBoundsValueMap() { 3335 return AffineValueMap(getUpperBoundsMap(), getUpperBoundsOperands()); 3336 } 3337 3338 Optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() { 3339 if (hasMinMaxBounds()) 3340 return llvm::None; 3341 3342 // Try to convert all the ranges to constant expressions. 3343 SmallVector<int64_t, 8> out; 3344 AffineValueMap rangesValueMap; 3345 AffineValueMap::difference(getUpperBoundsValueMap(), getLowerBoundsValueMap(), 3346 &rangesValueMap); 3347 out.reserve(rangesValueMap.getNumResults()); 3348 for (unsigned i = 0, e = rangesValueMap.getNumResults(); i < e; ++i) { 3349 auto expr = rangesValueMap.getResult(i); 3350 auto cst = expr.dyn_cast<AffineConstantExpr>(); 3351 if (!cst) 3352 return llvm::None; 3353 out.push_back(cst.getValue()); 3354 } 3355 return out; 3356 } 3357 3358 Block *AffineParallelOp::getBody() { return &getRegion().front(); } 3359 3360 OpBuilder AffineParallelOp::getBodyBuilder() { 3361 return OpBuilder(getBody(), std::prev(getBody()->end())); 3362 } 3363 3364 void AffineParallelOp::setLowerBounds(ValueRange lbOperands, AffineMap map) { 3365 assert(lbOperands.size() == map.getNumInputs() && 3366 "operands to map must match number of inputs"); 3367 3368 auto ubOperands = getUpperBoundsOperands(); 3369 3370 SmallVector<Value, 4> newOperands(lbOperands); 3371 newOperands.append(ubOperands.begin(), ubOperands.end()); 3372 (*this)->setOperands(newOperands); 3373 3374 setLowerBoundsMapAttr(AffineMapAttr::get(map)); 3375 } 3376 3377 void AffineParallelOp::setUpperBounds(ValueRange ubOperands, AffineMap map) { 3378 assert(ubOperands.size() == map.getNumInputs() && 3379 "operands to map must match number of inputs"); 3380 3381 SmallVector<Value, 4> newOperands(getLowerBoundsOperands()); 3382 newOperands.append(ubOperands.begin(), ubOperands.end()); 3383 (*this)->setOperands(newOperands); 3384 3385 setUpperBoundsMapAttr(AffineMapAttr::get(map)); 3386 } 3387 3388 void AffineParallelOp::setLowerBoundsMap(AffineMap map) { 3389 AffineMap lbMap = getLowerBoundsMap(); 3390 assert(lbMap.getNumDims() == map.getNumDims() && 3391 lbMap.getNumSymbols() == map.getNumSymbols()); 3392 (void)lbMap; 3393 setLowerBoundsMapAttr(AffineMapAttr::get(map)); 3394 } 3395 3396 void AffineParallelOp::setUpperBoundsMap(AffineMap map) { 3397 AffineMap ubMap = getUpperBoundsMap(); 3398 assert(ubMap.getNumDims() == map.getNumDims() && 3399 ubMap.getNumSymbols() == map.getNumSymbols()); 3400 (void)ubMap; 3401 setUpperBoundsMapAttr(AffineMapAttr::get(map)); 3402 } 3403 3404 void AffineParallelOp::setSteps(ArrayRef<int64_t> newSteps) { 3405 setStepsAttr(getBodyBuilder().getI64ArrayAttr(newSteps)); 3406 } 3407 3408 LogicalResult AffineParallelOp::verify() { 3409 auto numDims = getNumDims(); 3410 if (getLowerBoundsGroups().getNumElements() != numDims || 3411 getUpperBoundsGroups().getNumElements() != numDims || 3412 getSteps().size() != numDims || getBody()->getNumArguments() != numDims) { 3413 return emitOpError() << "the number of region arguments (" 3414 << getBody()->getNumArguments() 3415 << ") and the number of map groups for lower (" 3416 << getLowerBoundsGroups().getNumElements() 3417 << ") and upper bound (" 3418 << getUpperBoundsGroups().getNumElements() 3419 << "), and the number of steps (" << getSteps().size() 3420 << ") must all match"; 3421 } 3422 3423 unsigned expectedNumLBResults = 0; 3424 for (APInt v : getLowerBoundsGroups()) 3425 expectedNumLBResults += v.getZExtValue(); 3426 if (expectedNumLBResults != getLowerBoundsMap().getNumResults()) 3427 return emitOpError() << "expected lower bounds map to have " 3428 << expectedNumLBResults << " results"; 3429 unsigned expectedNumUBResults = 0; 3430 for (APInt v : getUpperBoundsGroups()) 3431 expectedNumUBResults += v.getZExtValue(); 3432 if (expectedNumUBResults != getUpperBoundsMap().getNumResults()) 3433 return emitOpError() << "expected upper bounds map to have " 3434 << expectedNumUBResults << " results"; 3435 3436 if (getReductions().size() != getNumResults()) 3437 return emitOpError("a reduction must be specified for each output"); 3438 3439 // Verify reduction ops are all valid 3440 for (Attribute attr : getReductions()) { 3441 auto intAttr = attr.dyn_cast<IntegerAttr>(); 3442 if (!intAttr || !arith::symbolizeAtomicRMWKind(intAttr.getInt())) 3443 return emitOpError("invalid reduction attribute"); 3444 } 3445 3446 // Verify that the bound operands are valid dimension/symbols. 3447 /// Lower bounds. 3448 if (failed(verifyDimAndSymbolIdentifiers(*this, getLowerBoundsOperands(), 3449 getLowerBoundsMap().getNumDims()))) 3450 return failure(); 3451 /// Upper bounds. 3452 if (failed(verifyDimAndSymbolIdentifiers(*this, getUpperBoundsOperands(), 3453 getUpperBoundsMap().getNumDims()))) 3454 return failure(); 3455 return success(); 3456 } 3457 3458 LogicalResult AffineValueMap::canonicalize() { 3459 SmallVector<Value, 4> newOperands{operands}; 3460 auto newMap = getAffineMap(); 3461 composeAffineMapAndOperands(&newMap, &newOperands); 3462 if (newMap == getAffineMap() && newOperands == operands) 3463 return failure(); 3464 reset(newMap, newOperands); 3465 return success(); 3466 } 3467 3468 /// Canonicalize the bounds of the given loop. 3469 static LogicalResult canonicalizeLoopBounds(AffineParallelOp op) { 3470 AffineValueMap lb = op.getLowerBoundsValueMap(); 3471 bool lbCanonicalized = succeeded(lb.canonicalize()); 3472 3473 AffineValueMap ub = op.getUpperBoundsValueMap(); 3474 bool ubCanonicalized = succeeded(ub.canonicalize()); 3475 3476 // Any canonicalization change always leads to updated map(s). 3477 if (!lbCanonicalized && !ubCanonicalized) 3478 return failure(); 3479 3480 if (lbCanonicalized) 3481 op.setLowerBounds(lb.getOperands(), lb.getAffineMap()); 3482 if (ubCanonicalized) 3483 op.setUpperBounds(ub.getOperands(), ub.getAffineMap()); 3484 3485 return success(); 3486 } 3487 3488 LogicalResult AffineParallelOp::fold(ArrayRef<Attribute> operands, 3489 SmallVectorImpl<OpFoldResult> &results) { 3490 return canonicalizeLoopBounds(*this); 3491 } 3492 3493 /// Prints a lower(upper) bound of an affine parallel loop with max(min) 3494 /// conditions in it. `mapAttr` is a flat list of affine expressions and `group` 3495 /// identifies which of the those expressions form max/min groups. `operands` 3496 /// are the SSA values of dimensions and symbols and `keyword` is either "min" 3497 /// or "max". 3498 static void printMinMaxBound(OpAsmPrinter &p, AffineMapAttr mapAttr, 3499 DenseIntElementsAttr group, ValueRange operands, 3500 StringRef keyword) { 3501 AffineMap map = mapAttr.getValue(); 3502 unsigned numDims = map.getNumDims(); 3503 ValueRange dimOperands = operands.take_front(numDims); 3504 ValueRange symOperands = operands.drop_front(numDims); 3505 unsigned start = 0; 3506 for (llvm::APInt groupSize : group) { 3507 if (start != 0) 3508 p << ", "; 3509 3510 unsigned size = groupSize.getZExtValue(); 3511 if (size == 1) { 3512 p.printAffineExprOfSSAIds(map.getResult(start), dimOperands, symOperands); 3513 ++start; 3514 } else { 3515 p << keyword << '('; 3516 AffineMap submap = map.getSliceMap(start, size); 3517 p.printAffineMapOfSSAIds(AffineMapAttr::get(submap), operands); 3518 p << ')'; 3519 start += size; 3520 } 3521 } 3522 } 3523 3524 void AffineParallelOp::print(OpAsmPrinter &p) { 3525 p << " (" << getBody()->getArguments() << ") = ("; 3526 printMinMaxBound(p, getLowerBoundsMapAttr(), getLowerBoundsGroupsAttr(), 3527 getLowerBoundsOperands(), "max"); 3528 p << ") to ("; 3529 printMinMaxBound(p, getUpperBoundsMapAttr(), getUpperBoundsGroupsAttr(), 3530 getUpperBoundsOperands(), "min"); 3531 p << ')'; 3532 SmallVector<int64_t, 8> steps = getSteps(); 3533 bool elideSteps = llvm::all_of(steps, [](int64_t step) { return step == 1; }); 3534 if (!elideSteps) { 3535 p << " step ("; 3536 llvm::interleaveComma(steps, p); 3537 p << ')'; 3538 } 3539 if (getNumResults()) { 3540 p << " reduce ("; 3541 llvm::interleaveComma(getReductions(), p, [&](auto &attr) { 3542 arith::AtomicRMWKind sym = *arith::symbolizeAtomicRMWKind( 3543 attr.template cast<IntegerAttr>().getInt()); 3544 p << "\"" << arith::stringifyAtomicRMWKind(sym) << "\""; 3545 }); 3546 p << ") -> (" << getResultTypes() << ")"; 3547 } 3548 3549 p << ' '; 3550 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false, 3551 /*printBlockTerminators=*/getNumResults()); 3552 p.printOptionalAttrDict( 3553 (*this)->getAttrs(), 3554 /*elidedAttrs=*/{AffineParallelOp::getReductionsAttrStrName(), 3555 AffineParallelOp::getLowerBoundsMapAttrStrName(), 3556 AffineParallelOp::getLowerBoundsGroupsAttrStrName(), 3557 AffineParallelOp::getUpperBoundsMapAttrStrName(), 3558 AffineParallelOp::getUpperBoundsGroupsAttrStrName(), 3559 AffineParallelOp::getStepsAttrStrName()}); 3560 } 3561 3562 /// Given a list of lists of parsed operands, populates `uniqueOperands` with 3563 /// unique operands. Also populates `replacements with affine expressions of 3564 /// `kind` that can be used to update affine maps previously accepting a 3565 /// `operands` to accept `uniqueOperands` instead. 3566 static ParseResult deduplicateAndResolveOperands( 3567 OpAsmParser &parser, 3568 ArrayRef<SmallVector<OpAsmParser::UnresolvedOperand>> operands, 3569 SmallVectorImpl<Value> &uniqueOperands, 3570 SmallVectorImpl<AffineExpr> &replacements, AffineExprKind kind) { 3571 assert((kind == AffineExprKind::DimId || kind == AffineExprKind::SymbolId) && 3572 "expected operands to be dim or symbol expression"); 3573 3574 Type indexType = parser.getBuilder().getIndexType(); 3575 for (const auto &list : operands) { 3576 SmallVector<Value> valueOperands; 3577 if (parser.resolveOperands(list, indexType, valueOperands)) 3578 return failure(); 3579 for (Value operand : valueOperands) { 3580 unsigned pos = std::distance(uniqueOperands.begin(), 3581 llvm::find(uniqueOperands, operand)); 3582 if (pos == uniqueOperands.size()) 3583 uniqueOperands.push_back(operand); 3584 replacements.push_back( 3585 kind == AffineExprKind::DimId 3586 ? getAffineDimExpr(pos, parser.getContext()) 3587 : getAffineSymbolExpr(pos, parser.getContext())); 3588 } 3589 } 3590 return success(); 3591 } 3592 3593 namespace { 3594 enum class MinMaxKind { Min, Max }; 3595 } // namespace 3596 3597 /// Parses an affine map that can contain a min/max for groups of its results, 3598 /// e.g., max(expr-1, expr-2), expr-3, max(expr-4, expr-5, expr-6). Populates 3599 /// `result` attributes with the map (flat list of expressions) and the grouping 3600 /// (list of integers that specify how many expressions to put into each 3601 /// min/max) attributes. Deduplicates repeated operands. 3602 /// 3603 /// parallel-bound ::= `(` parallel-group-list `)` 3604 /// parallel-group-list ::= parallel-group (`,` parallel-group-list)? 3605 /// parallel-group ::= simple-group | min-max-group 3606 /// simple-group ::= expr-of-ssa-ids 3607 /// min-max-group ::= ( `min` | `max` ) `(` expr-of-ssa-ids-list `)` 3608 /// expr-of-ssa-ids-list ::= expr-of-ssa-ids (`,` expr-of-ssa-id-list)? 3609 /// 3610 /// Examples: 3611 /// (%0, min(%1 + %2, %3), %4, min(%5 floordiv 32, %6)) 3612 /// (%0, max(%1 - 2 * %2)) 3613 static ParseResult parseAffineMapWithMinMax(OpAsmParser &parser, 3614 OperationState &result, 3615 MinMaxKind kind) { 3616 constexpr llvm::StringLiteral tmpAttrStrName = "__pseudo_bound_map"; 3617 3618 StringRef mapName = kind == MinMaxKind::Min 3619 ? AffineParallelOp::getUpperBoundsMapAttrStrName() 3620 : AffineParallelOp::getLowerBoundsMapAttrStrName(); 3621 StringRef groupsName = 3622 kind == MinMaxKind::Min 3623 ? AffineParallelOp::getUpperBoundsGroupsAttrStrName() 3624 : AffineParallelOp::getLowerBoundsGroupsAttrStrName(); 3625 3626 if (failed(parser.parseLParen())) 3627 return failure(); 3628 3629 if (succeeded(parser.parseOptionalRParen())) { 3630 result.addAttribute( 3631 mapName, AffineMapAttr::get(parser.getBuilder().getEmptyAffineMap())); 3632 result.addAttribute(groupsName, parser.getBuilder().getI32TensorAttr({})); 3633 return success(); 3634 } 3635 3636 SmallVector<AffineExpr> flatExprs; 3637 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatDimOperands; 3638 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatSymOperands; 3639 SmallVector<int32_t> numMapsPerGroup; 3640 SmallVector<OpAsmParser::UnresolvedOperand> mapOperands; 3641 auto parseOperands = [&]() { 3642 if (succeeded(parser.parseOptionalKeyword( 3643 kind == MinMaxKind::Min ? "min" : "max"))) { 3644 mapOperands.clear(); 3645 AffineMapAttr map; 3646 if (failed(parser.parseAffineMapOfSSAIds(mapOperands, map, tmpAttrStrName, 3647 result.attributes, 3648 OpAsmParser::Delimiter::Paren))) 3649 return failure(); 3650 result.attributes.erase(tmpAttrStrName); 3651 llvm::append_range(flatExprs, map.getValue().getResults()); 3652 auto operandsRef = llvm::makeArrayRef(mapOperands); 3653 auto dimsRef = operandsRef.take_front(map.getValue().getNumDims()); 3654 SmallVector<OpAsmParser::UnresolvedOperand> dims(dimsRef.begin(), 3655 dimsRef.end()); 3656 auto symsRef = operandsRef.drop_front(map.getValue().getNumDims()); 3657 SmallVector<OpAsmParser::UnresolvedOperand> syms(symsRef.begin(), 3658 symsRef.end()); 3659 flatDimOperands.append(map.getValue().getNumResults(), dims); 3660 flatSymOperands.append(map.getValue().getNumResults(), syms); 3661 numMapsPerGroup.push_back(map.getValue().getNumResults()); 3662 } else { 3663 if (failed(parser.parseAffineExprOfSSAIds(flatDimOperands.emplace_back(), 3664 flatSymOperands.emplace_back(), 3665 flatExprs.emplace_back()))) 3666 return failure(); 3667 numMapsPerGroup.push_back(1); 3668 } 3669 return success(); 3670 }; 3671 if (parser.parseCommaSeparatedList(parseOperands) || parser.parseRParen()) 3672 return failure(); 3673 3674 unsigned totalNumDims = 0; 3675 unsigned totalNumSyms = 0; 3676 for (unsigned i = 0, e = flatExprs.size(); i < e; ++i) { 3677 unsigned numDims = flatDimOperands[i].size(); 3678 unsigned numSyms = flatSymOperands[i].size(); 3679 flatExprs[i] = flatExprs[i] 3680 .shiftDims(numDims, totalNumDims) 3681 .shiftSymbols(numSyms, totalNumSyms); 3682 totalNumDims += numDims; 3683 totalNumSyms += numSyms; 3684 } 3685 3686 // Deduplicate map operands. 3687 SmallVector<Value> dimOperands, symOperands; 3688 SmallVector<AffineExpr> dimRplacements, symRepacements; 3689 if (deduplicateAndResolveOperands(parser, flatDimOperands, dimOperands, 3690 dimRplacements, AffineExprKind::DimId) || 3691 deduplicateAndResolveOperands(parser, flatSymOperands, symOperands, 3692 symRepacements, AffineExprKind::SymbolId)) 3693 return failure(); 3694 3695 result.operands.append(dimOperands.begin(), dimOperands.end()); 3696 result.operands.append(symOperands.begin(), symOperands.end()); 3697 3698 Builder &builder = parser.getBuilder(); 3699 auto flatMap = AffineMap::get(totalNumDims, totalNumSyms, flatExprs, 3700 parser.getContext()); 3701 flatMap = flatMap.replaceDimsAndSymbols( 3702 dimRplacements, symRepacements, dimOperands.size(), symOperands.size()); 3703 3704 result.addAttribute(mapName, AffineMapAttr::get(flatMap)); 3705 result.addAttribute(groupsName, builder.getI32TensorAttr(numMapsPerGroup)); 3706 return success(); 3707 } 3708 3709 // 3710 // operation ::= `affine.parallel` `(` ssa-ids `)` `=` parallel-bound 3711 // `to` parallel-bound steps? region attr-dict? 3712 // steps ::= `steps` `(` integer-literals `)` 3713 // 3714 ParseResult AffineParallelOp::parse(OpAsmParser &parser, 3715 OperationState &result) { 3716 auto &builder = parser.getBuilder(); 3717 auto indexType = builder.getIndexType(); 3718 SmallVector<OpAsmParser::Argument, 4> ivs; 3719 if (parser.parseArgumentList(ivs, OpAsmParser::Delimiter::Paren) || 3720 parser.parseEqual() || 3721 parseAffineMapWithMinMax(parser, result, MinMaxKind::Max) || 3722 parser.parseKeyword("to") || 3723 parseAffineMapWithMinMax(parser, result, MinMaxKind::Min)) 3724 return failure(); 3725 3726 AffineMapAttr stepsMapAttr; 3727 NamedAttrList stepsAttrs; 3728 SmallVector<OpAsmParser::UnresolvedOperand, 4> stepsMapOperands; 3729 if (failed(parser.parseOptionalKeyword("step"))) { 3730 SmallVector<int64_t, 4> steps(ivs.size(), 1); 3731 result.addAttribute(AffineParallelOp::getStepsAttrStrName(), 3732 builder.getI64ArrayAttr(steps)); 3733 } else { 3734 if (parser.parseAffineMapOfSSAIds(stepsMapOperands, stepsMapAttr, 3735 AffineParallelOp::getStepsAttrStrName(), 3736 stepsAttrs, 3737 OpAsmParser::Delimiter::Paren)) 3738 return failure(); 3739 3740 // Convert steps from an AffineMap into an I64ArrayAttr. 3741 SmallVector<int64_t, 4> steps; 3742 auto stepsMap = stepsMapAttr.getValue(); 3743 for (const auto &result : stepsMap.getResults()) { 3744 auto constExpr = result.dyn_cast<AffineConstantExpr>(); 3745 if (!constExpr) 3746 return parser.emitError(parser.getNameLoc(), 3747 "steps must be constant integers"); 3748 steps.push_back(constExpr.getValue()); 3749 } 3750 result.addAttribute(AffineParallelOp::getStepsAttrStrName(), 3751 builder.getI64ArrayAttr(steps)); 3752 } 3753 3754 // Parse optional clause of the form: `reduce ("addf", "maxf")`, where the 3755 // quoted strings are a member of the enum AtomicRMWKind. 3756 SmallVector<Attribute, 4> reductions; 3757 if (succeeded(parser.parseOptionalKeyword("reduce"))) { 3758 if (parser.parseLParen()) 3759 return failure(); 3760 auto parseAttributes = [&]() -> ParseResult { 3761 // Parse a single quoted string via the attribute parsing, and then 3762 // verify it is a member of the enum and convert to it's integer 3763 // representation. 3764 StringAttr attrVal; 3765 NamedAttrList attrStorage; 3766 auto loc = parser.getCurrentLocation(); 3767 if (parser.parseAttribute(attrVal, builder.getNoneType(), "reduce", 3768 attrStorage)) 3769 return failure(); 3770 llvm::Optional<arith::AtomicRMWKind> reduction = 3771 arith::symbolizeAtomicRMWKind(attrVal.getValue()); 3772 if (!reduction) 3773 return parser.emitError(loc, "invalid reduction value: ") << attrVal; 3774 reductions.push_back(builder.getI64IntegerAttr( 3775 static_cast<int64_t>(reduction.getValue()))); 3776 // While we keep getting commas, keep parsing. 3777 return success(); 3778 }; 3779 if (parser.parseCommaSeparatedList(parseAttributes) || parser.parseRParen()) 3780 return failure(); 3781 } 3782 result.addAttribute(AffineParallelOp::getReductionsAttrStrName(), 3783 builder.getArrayAttr(reductions)); 3784 3785 // Parse return types of reductions (if any) 3786 if (parser.parseOptionalArrowTypeList(result.types)) 3787 return failure(); 3788 3789 // Now parse the body. 3790 Region *body = result.addRegion(); 3791 for (auto &iv : ivs) 3792 iv.type = indexType; 3793 if (parser.parseRegion(*body, ivs) || 3794 parser.parseOptionalAttrDict(result.attributes)) 3795 return failure(); 3796 3797 // Add a terminator if none was parsed. 3798 AffineParallelOp::ensureTerminator(*body, builder, result.location); 3799 return success(); 3800 } 3801 3802 //===----------------------------------------------------------------------===// 3803 // AffineYieldOp 3804 //===----------------------------------------------------------------------===// 3805 3806 LogicalResult AffineYieldOp::verify() { 3807 auto *parentOp = (*this)->getParentOp(); 3808 auto results = parentOp->getResults(); 3809 auto operands = getOperands(); 3810 3811 if (!isa<AffineParallelOp, AffineIfOp, AffineForOp>(parentOp)) 3812 return emitOpError() << "only terminates affine.if/for/parallel regions"; 3813 if (parentOp->getNumResults() != getNumOperands()) 3814 return emitOpError() << "parent of yield must have same number of " 3815 "results as the yield operands"; 3816 for (auto it : llvm::zip(results, operands)) { 3817 if (std::get<0>(it).getType() != std::get<1>(it).getType()) 3818 return emitOpError() << "types mismatch between yield op and its parent"; 3819 } 3820 3821 return success(); 3822 } 3823 3824 //===----------------------------------------------------------------------===// 3825 // AffineVectorLoadOp 3826 //===----------------------------------------------------------------------===// 3827 3828 void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result, 3829 VectorType resultType, AffineMap map, 3830 ValueRange operands) { 3831 assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands"); 3832 result.addOperands(operands); 3833 if (map) 3834 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map)); 3835 result.types.push_back(resultType); 3836 } 3837 3838 void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result, 3839 VectorType resultType, Value memref, 3840 AffineMap map, ValueRange mapOperands) { 3841 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info"); 3842 result.addOperands(memref); 3843 result.addOperands(mapOperands); 3844 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map)); 3845 result.types.push_back(resultType); 3846 } 3847 3848 void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result, 3849 VectorType resultType, Value memref, 3850 ValueRange indices) { 3851 auto memrefType = memref.getType().cast<MemRefType>(); 3852 int64_t rank = memrefType.getRank(); 3853 // Create identity map for memrefs with at least one dimension or () -> () 3854 // for zero-dimensional memrefs. 3855 auto map = 3856 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap(); 3857 build(builder, result, resultType, memref, map, indices); 3858 } 3859 3860 void AffineVectorLoadOp::getCanonicalizationPatterns(RewritePatternSet &results, 3861 MLIRContext *context) { 3862 results.add<SimplifyAffineOp<AffineVectorLoadOp>>(context); 3863 } 3864 3865 ParseResult AffineVectorLoadOp::parse(OpAsmParser &parser, 3866 OperationState &result) { 3867 auto &builder = parser.getBuilder(); 3868 auto indexTy = builder.getIndexType(); 3869 3870 MemRefType memrefType; 3871 VectorType resultType; 3872 OpAsmParser::UnresolvedOperand memrefInfo; 3873 AffineMapAttr mapAttr; 3874 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands; 3875 return failure( 3876 parser.parseOperand(memrefInfo) || 3877 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, 3878 AffineVectorLoadOp::getMapAttrStrName(), 3879 result.attributes) || 3880 parser.parseOptionalAttrDict(result.attributes) || 3881 parser.parseColonType(memrefType) || parser.parseComma() || 3882 parser.parseType(resultType) || 3883 parser.resolveOperand(memrefInfo, memrefType, result.operands) || 3884 parser.resolveOperands(mapOperands, indexTy, result.operands) || 3885 parser.addTypeToList(resultType, result.types)); 3886 } 3887 3888 void AffineVectorLoadOp::print(OpAsmPrinter &p) { 3889 p << " " << getMemRef() << '['; 3890 if (AffineMapAttr mapAttr = 3891 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName())) 3892 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 3893 p << ']'; 3894 p.printOptionalAttrDict((*this)->getAttrs(), 3895 /*elidedAttrs=*/{getMapAttrStrName()}); 3896 p << " : " << getMemRefType() << ", " << getType(); 3897 } 3898 3899 /// Verify common invariants of affine.vector_load and affine.vector_store. 3900 static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, 3901 VectorType vectorType) { 3902 // Check that memref and vector element types match. 3903 if (memrefType.getElementType() != vectorType.getElementType()) 3904 return op->emitOpError( 3905 "requires memref and vector types of the same elemental type"); 3906 return success(); 3907 } 3908 3909 LogicalResult AffineVectorLoadOp::verify() { 3910 MemRefType memrefType = getMemRefType(); 3911 if (failed(verifyMemoryOpIndexing( 3912 getOperation(), 3913 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()), 3914 getMapOperands(), memrefType, 3915 /*numIndexOperands=*/getNumOperands() - 1))) 3916 return failure(); 3917 3918 if (failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) 3919 return failure(); 3920 3921 return success(); 3922 } 3923 3924 //===----------------------------------------------------------------------===// 3925 // AffineVectorStoreOp 3926 //===----------------------------------------------------------------------===// 3927 3928 void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &result, 3929 Value valueToStore, Value memref, AffineMap map, 3930 ValueRange mapOperands) { 3931 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info"); 3932 result.addOperands(valueToStore); 3933 result.addOperands(memref); 3934 result.addOperands(mapOperands); 3935 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map)); 3936 } 3937 3938 // Use identity map. 3939 void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &result, 3940 Value valueToStore, Value memref, 3941 ValueRange indices) { 3942 auto memrefType = memref.getType().cast<MemRefType>(); 3943 int64_t rank = memrefType.getRank(); 3944 // Create identity map for memrefs with at least one dimension or () -> () 3945 // for zero-dimensional memrefs. 3946 auto map = 3947 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap(); 3948 build(builder, result, valueToStore, memref, map, indices); 3949 } 3950 void AffineVectorStoreOp::getCanonicalizationPatterns( 3951 RewritePatternSet &results, MLIRContext *context) { 3952 results.add<SimplifyAffineOp<AffineVectorStoreOp>>(context); 3953 } 3954 3955 ParseResult AffineVectorStoreOp::parse(OpAsmParser &parser, 3956 OperationState &result) { 3957 auto indexTy = parser.getBuilder().getIndexType(); 3958 3959 MemRefType memrefType; 3960 VectorType resultType; 3961 OpAsmParser::UnresolvedOperand storeValueInfo; 3962 OpAsmParser::UnresolvedOperand memrefInfo; 3963 AffineMapAttr mapAttr; 3964 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands; 3965 return failure( 3966 parser.parseOperand(storeValueInfo) || parser.parseComma() || 3967 parser.parseOperand(memrefInfo) || 3968 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, 3969 AffineVectorStoreOp::getMapAttrStrName(), 3970 result.attributes) || 3971 parser.parseOptionalAttrDict(result.attributes) || 3972 parser.parseColonType(memrefType) || parser.parseComma() || 3973 parser.parseType(resultType) || 3974 parser.resolveOperand(storeValueInfo, resultType, result.operands) || 3975 parser.resolveOperand(memrefInfo, memrefType, result.operands) || 3976 parser.resolveOperands(mapOperands, indexTy, result.operands)); 3977 } 3978 3979 void AffineVectorStoreOp::print(OpAsmPrinter &p) { 3980 p << " " << getValueToStore(); 3981 p << ", " << getMemRef() << '['; 3982 if (AffineMapAttr mapAttr = 3983 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName())) 3984 p.printAffineMapOfSSAIds(mapAttr, getMapOperands()); 3985 p << ']'; 3986 p.printOptionalAttrDict((*this)->getAttrs(), 3987 /*elidedAttrs=*/{getMapAttrStrName()}); 3988 p << " : " << getMemRefType() << ", " << getValueToStore().getType(); 3989 } 3990 3991 LogicalResult AffineVectorStoreOp::verify() { 3992 MemRefType memrefType = getMemRefType(); 3993 if (failed(verifyMemoryOpIndexing( 3994 *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()), 3995 getMapOperands(), memrefType, 3996 /*numIndexOperands=*/getNumOperands() - 2))) 3997 return failure(); 3998 3999 if (failed(verifyVectorMemoryOp(*this, memrefType, getVectorType()))) 4000 return failure(); 4001 4002 return success(); 4003 } 4004 4005 //===----------------------------------------------------------------------===// 4006 // TableGen'd op method definitions 4007 //===----------------------------------------------------------------------===// 4008 4009 #define GET_OP_CLASSES 4010 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc" 4011