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