1 //===----------------------------------------------------------------------===// 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/MemRef/IR/MemRef.h" 10 #include "mlir/Dialect/MemRef/Utils/MemRefUtils.h" 11 #include "mlir/Dialect/StandardOps/IR/Ops.h" 12 #include "mlir/Dialect/StandardOps/Utils/Utils.h" 13 #include "mlir/Dialect/Tensor/IR/Tensor.h" 14 #include "mlir/IR/AffineMap.h" 15 #include "mlir/IR/Builders.h" 16 #include "mlir/IR/BuiltinTypes.h" 17 #include "mlir/IR/Matchers.h" 18 #include "mlir/IR/PatternMatch.h" 19 #include "mlir/IR/TypeUtilities.h" 20 #include "mlir/Interfaces/InferTypeOpInterface.h" 21 #include "llvm/ADT/STLExtras.h" 22 23 using namespace mlir; 24 using namespace mlir::memref; 25 26 /// Materialize a single constant operation from a given attribute value with 27 /// the desired resultant type. 28 Operation *MemRefDialect::materializeConstant(OpBuilder &builder, 29 Attribute value, Type type, 30 Location loc) { 31 return builder.create<mlir::ConstantOp>(loc, type, value); 32 } 33 34 /// Extract int64_t values from the assumed ArrayAttr of IntegerAttr. 35 static SmallVector<int64_t, 4> extractFromI64ArrayAttr(Attribute attr) { 36 return llvm::to_vector<4>( 37 llvm::map_range(attr.cast<ArrayAttr>(), [](Attribute a) -> int64_t { 38 return a.cast<IntegerAttr>().getInt(); 39 })); 40 } 41 42 /// Helper function to dispatch an OpFoldResult into either the `dynamicVec` if 43 /// it is a Value or into `staticVec` if it is an IntegerAttr. 44 /// In the case of a Value, a copy of the `sentinel` value is also pushed to 45 /// `staticVec`. This is useful to extract mixed static and dynamic entries that 46 /// come from an AttrSizedOperandSegments trait. 47 static void dispatchIndexOpFoldResult(OpFoldResult ofr, 48 SmallVectorImpl<Value> &dynamicVec, 49 SmallVectorImpl<int64_t> &staticVec, 50 int64_t sentinel) { 51 if (auto v = ofr.dyn_cast<Value>()) { 52 dynamicVec.push_back(v); 53 staticVec.push_back(sentinel); 54 return; 55 } 56 APInt apInt = ofr.dyn_cast<Attribute>().cast<IntegerAttr>().getValue(); 57 staticVec.push_back(apInt.getSExtValue()); 58 } 59 60 static void dispatchIndexOpFoldResults(ArrayRef<OpFoldResult> ofrs, 61 SmallVectorImpl<Value> &dynamicVec, 62 SmallVectorImpl<int64_t> &staticVec, 63 int64_t sentinel) { 64 for (auto ofr : ofrs) 65 dispatchIndexOpFoldResult(ofr, dynamicVec, staticVec, sentinel); 66 } 67 68 //===----------------------------------------------------------------------===// 69 // Common canonicalization pattern support logic 70 //===----------------------------------------------------------------------===// 71 72 /// This is a common class used for patterns of the form 73 /// "someop(memrefcast) -> someop". It folds the source of any memref.cast 74 /// into the root operation directly. 75 static LogicalResult foldMemRefCast(Operation *op) { 76 bool folded = false; 77 for (OpOperand &operand : op->getOpOperands()) { 78 auto cast = operand.get().getDefiningOp<CastOp>(); 79 if (cast && !cast.getOperand().getType().isa<UnrankedMemRefType>()) { 80 operand.set(cast.getOperand()); 81 folded = true; 82 } 83 } 84 return success(folded); 85 } 86 87 //===----------------------------------------------------------------------===// 88 // Helpers for GlobalOp 89 //===----------------------------------------------------------------------===// 90 91 static Type getTensorTypeFromMemRefType(Type type) { 92 if (auto memref = type.dyn_cast<MemRefType>()) 93 return RankedTensorType::get(memref.getShape(), memref.getElementType()); 94 if (auto memref = type.dyn_cast<UnrankedMemRefType>()) 95 return UnrankedTensorType::get(memref.getElementType()); 96 return NoneType::get(type.getContext()); 97 } 98 99 //===----------------------------------------------------------------------===// 100 // AllocOp / AllocaOp 101 //===----------------------------------------------------------------------===// 102 103 template <typename AllocLikeOp> 104 static LogicalResult verifyAllocLikeOp(AllocLikeOp op) { 105 static_assert(llvm::is_one_of<AllocLikeOp, AllocOp, AllocaOp>::value, 106 "applies to only alloc or alloca"); 107 auto memRefType = op.getResult().getType().template dyn_cast<MemRefType>(); 108 if (!memRefType) 109 return op.emitOpError("result must be a memref"); 110 111 if (static_cast<int64_t>(op.dynamicSizes().size()) != 112 memRefType.getNumDynamicDims()) 113 return op.emitOpError("dimension operand count does not equal memref " 114 "dynamic dimension count"); 115 116 unsigned numSymbols = 0; 117 if (!memRefType.getAffineMaps().empty()) 118 numSymbols = memRefType.getAffineMaps().front().getNumSymbols(); 119 if (op.symbolOperands().size() != numSymbols) 120 return op.emitOpError( 121 "symbol operand count does not equal memref symbol count"); 122 123 return success(); 124 } 125 126 static LogicalResult verify(AllocOp op) { return verifyAllocLikeOp(op); } 127 128 static LogicalResult verify(AllocaOp op) { 129 // An alloca op needs to have an ancestor with an allocation scope trait. 130 if (!op->getParentWithTrait<OpTrait::AutomaticAllocationScope>()) 131 return op.emitOpError( 132 "requires an ancestor op with AutomaticAllocationScope trait"); 133 134 return verifyAllocLikeOp(op); 135 } 136 137 namespace { 138 /// Fold constant dimensions into an alloc like operation. 139 template <typename AllocLikeOp> 140 struct SimplifyAllocConst : public OpRewritePattern<AllocLikeOp> { 141 using OpRewritePattern<AllocLikeOp>::OpRewritePattern; 142 143 LogicalResult matchAndRewrite(AllocLikeOp alloc, 144 PatternRewriter &rewriter) const override { 145 // Check to see if any dimensions operands are constants. If so, we can 146 // substitute and drop them. 147 if (llvm::none_of(alloc.getOperands(), [](Value operand) { 148 return matchPattern(operand, matchConstantIndex()); 149 })) 150 return failure(); 151 152 auto memrefType = alloc.getType(); 153 154 // Ok, we have one or more constant operands. Collect the non-constant ones 155 // and keep track of the resultant memref type to build. 156 SmallVector<int64_t, 4> newShapeConstants; 157 newShapeConstants.reserve(memrefType.getRank()); 158 SmallVector<Value, 4> newOperands; 159 160 unsigned dynamicDimPos = 0; 161 for (unsigned dim = 0, e = memrefType.getRank(); dim < e; ++dim) { 162 int64_t dimSize = memrefType.getDimSize(dim); 163 // If this is already static dimension, keep it. 164 if (dimSize != -1) { 165 newShapeConstants.push_back(dimSize); 166 continue; 167 } 168 auto *defOp = alloc.getOperand(dynamicDimPos).getDefiningOp(); 169 if (auto constantIndexOp = dyn_cast_or_null<ConstantIndexOp>(defOp)) { 170 // Dynamic shape dimension will be folded. 171 newShapeConstants.push_back(constantIndexOp.getValue()); 172 } else { 173 // Dynamic shape dimension not folded; copy operand from old memref. 174 newShapeConstants.push_back(-1); 175 newOperands.push_back(alloc.getOperand(dynamicDimPos)); 176 } 177 dynamicDimPos++; 178 } 179 180 // Create new memref type (which will have fewer dynamic dimensions). 181 MemRefType newMemRefType = 182 MemRefType::Builder(memrefType).setShape(newShapeConstants); 183 assert(static_cast<int64_t>(newOperands.size()) == 184 newMemRefType.getNumDynamicDims()); 185 186 // Create and insert the alloc op for the new memref. 187 auto newAlloc = rewriter.create<AllocLikeOp>( 188 alloc.getLoc(), newMemRefType, newOperands, alloc.alignmentAttr()); 189 // Insert a cast so we have the same type as the old alloc. 190 auto resultCast = 191 rewriter.create<CastOp>(alloc.getLoc(), newAlloc, alloc.getType()); 192 193 rewriter.replaceOp(alloc, {resultCast}); 194 return success(); 195 } 196 }; 197 198 /// Fold alloc operations with no users or only store and dealloc uses. 199 template <typename T> 200 struct SimplifyDeadAlloc : public OpRewritePattern<T> { 201 using OpRewritePattern<T>::OpRewritePattern; 202 203 LogicalResult matchAndRewrite(T alloc, 204 PatternRewriter &rewriter) const override { 205 if (llvm::any_of(alloc->getUsers(), [](Operation *op) { 206 return !isa<StoreOp, DeallocOp>(op); 207 })) 208 return failure(); 209 210 for (Operation *user : llvm::make_early_inc_range(alloc->getUsers())) 211 rewriter.eraseOp(user); 212 213 rewriter.eraseOp(alloc); 214 return success(); 215 } 216 }; 217 } // end anonymous namespace. 218 219 void AllocOp::getCanonicalizationPatterns(RewritePatternSet &results, 220 MLIRContext *context) { 221 results.add<SimplifyAllocConst<AllocOp>, SimplifyDeadAlloc<AllocOp>>(context); 222 } 223 224 void AllocaOp::getCanonicalizationPatterns(RewritePatternSet &results, 225 MLIRContext *context) { 226 results.add<SimplifyAllocConst<AllocaOp>, SimplifyDeadAlloc<AllocaOp>>( 227 context); 228 } 229 230 //===----------------------------------------------------------------------===// 231 // AssumeAlignmentOp 232 //===----------------------------------------------------------------------===// 233 234 static LogicalResult verify(AssumeAlignmentOp op) { 235 unsigned alignment = op.alignment(); 236 if (!llvm::isPowerOf2_32(alignment)) 237 return op.emitOpError("alignment must be power of 2"); 238 return success(); 239 } 240 241 //===----------------------------------------------------------------------===// 242 // BufferCastOp 243 //===----------------------------------------------------------------------===// 244 245 OpFoldResult BufferCastOp::fold(ArrayRef<Attribute>) { 246 if (auto tensorLoad = tensor().getDefiningOp<TensorLoadOp>()) 247 if (tensorLoad.memref().getType() == getType()) 248 return tensorLoad.memref(); 249 return {}; 250 } 251 252 namespace { 253 /// Replace tensor_cast + buffer_cast by buffer_cast + memref_cast. 254 struct BufferCast : public OpRewritePattern<BufferCastOp> { 255 using OpRewritePattern<BufferCastOp>::OpRewritePattern; 256 257 LogicalResult matchAndRewrite(BufferCastOp bufferCast, 258 PatternRewriter &rewriter) const final { 259 auto tensorCastOperand = 260 bufferCast.getOperand().getDefiningOp<tensor::CastOp>(); 261 if (!tensorCastOperand) 262 return failure(); 263 auto srcTensorType = 264 tensorCastOperand.getOperand().getType().dyn_cast<RankedTensorType>(); 265 if (!srcTensorType) 266 return failure(); 267 auto memrefType = MemRefType::get(srcTensorType.getShape(), 268 srcTensorType.getElementType()); 269 Value memref = rewriter.create<BufferCastOp>( 270 bufferCast.getLoc(), memrefType, tensorCastOperand.getOperand()); 271 rewriter.replaceOpWithNewOp<CastOp>(bufferCast, bufferCast.getType(), 272 memref); 273 return success(); 274 } 275 }; 276 277 /// Canonicalize memref.tensor_load + memref.buffer_cast to memref.cast when 278 /// type mismatches prevent `BufferCastOp::fold` to kick in. 279 struct TensorLoadToMemRef : public OpRewritePattern<BufferCastOp> { 280 using OpRewritePattern<BufferCastOp>::OpRewritePattern; 281 282 LogicalResult matchAndRewrite(BufferCastOp bufferCast, 283 PatternRewriter &rewriter) const final { 284 auto tensorLoad = bufferCast.tensor().getDefiningOp<TensorLoadOp>(); 285 // Bail unless we have a tensor_load + memref.buffer_cast with different 286 // types. `BufferCastOp::fold` handles the same type case. 287 if (!tensorLoad || tensorLoad.memref().getType() == bufferCast.getType()) 288 return failure(); 289 // If types are not cast-compatible, bail. 290 if (!CastOp::areCastCompatible(tensorLoad.memref().getType(), 291 bufferCast.getType())) 292 return failure(); 293 rewriter.replaceOpWithNewOp<CastOp>(bufferCast, bufferCast.getType(), 294 tensorLoad.memref()); 295 return success(); 296 } 297 }; 298 299 } // namespace 300 301 void BufferCastOp::getCanonicalizationPatterns(RewritePatternSet &results, 302 MLIRContext *context) { 303 results.add<BufferCast, TensorLoadToMemRef>(context); 304 } 305 306 //===----------------------------------------------------------------------===// 307 // CastOp 308 //===----------------------------------------------------------------------===// 309 310 /// Determines whether MemRef_CastOp casts to a more dynamic version of the 311 /// source memref. This is useful to to fold a memref.cast into a consuming op 312 /// and implement canonicalization patterns for ops in different dialects that 313 /// may consume the results of memref.cast operations. Such foldable memref.cast 314 /// operations are typically inserted as `view` and `subview` ops are 315 /// canonicalized, to preserve the type compatibility of their uses. 316 /// 317 /// Returns true when all conditions are met: 318 /// 1. source and result are ranked memrefs with strided semantics and same 319 /// element type and rank. 320 /// 2. each of the source's size, offset or stride has more static information 321 /// than the corresponding result's size, offset or stride. 322 /// 323 /// Example 1: 324 /// ```mlir 325 /// %1 = memref.cast %0 : memref<8x16xf32> to memref<?x?xf32> 326 /// %2 = consumer %1 ... : memref<?x?xf32> ... 327 /// ``` 328 /// 329 /// may fold into: 330 /// 331 /// ```mlir 332 /// %2 = consumer %0 ... : memref<8x16xf32> ... 333 /// ``` 334 /// 335 /// Example 2: 336 /// ``` 337 /// %1 = memref.cast %0 : memref<?x16xf32, affine_map<(i, j)->(16 * i + j)>> 338 /// to memref<?x?xf32> 339 /// consumer %1 : memref<?x?xf32> ... 340 /// ``` 341 /// 342 /// may fold into: 343 /// 344 /// ``` 345 /// consumer %0 ... : memref<?x16xf32, affine_map<(i, j)->(16 * i + j)>> 346 /// ``` 347 bool CastOp::canFoldIntoConsumerOp(CastOp castOp) { 348 MemRefType sourceType = castOp.source().getType().dyn_cast<MemRefType>(); 349 MemRefType resultType = castOp.getType().dyn_cast<MemRefType>(); 350 351 // Requires ranked MemRefType. 352 if (!sourceType || !resultType) 353 return false; 354 355 // Requires same elemental type. 356 if (sourceType.getElementType() != resultType.getElementType()) 357 return false; 358 359 // Requires same rank. 360 if (sourceType.getRank() != resultType.getRank()) 361 return false; 362 363 // Only fold casts between strided memref forms. 364 int64_t sourceOffset, resultOffset; 365 SmallVector<int64_t, 4> sourceStrides, resultStrides; 366 if (failed(getStridesAndOffset(sourceType, sourceStrides, sourceOffset)) || 367 failed(getStridesAndOffset(resultType, resultStrides, resultOffset))) 368 return false; 369 370 // If cast is towards more static sizes along any dimension, don't fold. 371 for (auto it : llvm::zip(sourceType.getShape(), resultType.getShape())) { 372 auto ss = std::get<0>(it), st = std::get<1>(it); 373 if (ss != st) 374 if (MemRefType::isDynamic(ss) && !MemRefType::isDynamic(st)) 375 return false; 376 } 377 378 // If cast is towards more static offset along any dimension, don't fold. 379 if (sourceOffset != resultOffset) 380 if (MemRefType::isDynamicStrideOrOffset(sourceOffset) && 381 !MemRefType::isDynamicStrideOrOffset(resultOffset)) 382 return false; 383 384 // If cast is towards more static strides along any dimension, don't fold. 385 for (auto it : llvm::zip(sourceStrides, resultStrides)) { 386 auto ss = std::get<0>(it), st = std::get<1>(it); 387 if (ss != st) 388 if (MemRefType::isDynamicStrideOrOffset(ss) && 389 !MemRefType::isDynamicStrideOrOffset(st)) 390 return false; 391 } 392 393 return true; 394 } 395 396 bool CastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) { 397 if (inputs.size() != 1 || outputs.size() != 1) 398 return false; 399 Type a = inputs.front(), b = outputs.front(); 400 auto aT = a.dyn_cast<MemRefType>(); 401 auto bT = b.dyn_cast<MemRefType>(); 402 403 auto uaT = a.dyn_cast<UnrankedMemRefType>(); 404 auto ubT = b.dyn_cast<UnrankedMemRefType>(); 405 406 if (aT && bT) { 407 if (aT.getElementType() != bT.getElementType()) 408 return false; 409 if (aT.getAffineMaps() != bT.getAffineMaps()) { 410 int64_t aOffset, bOffset; 411 SmallVector<int64_t, 4> aStrides, bStrides; 412 if (failed(getStridesAndOffset(aT, aStrides, aOffset)) || 413 failed(getStridesAndOffset(bT, bStrides, bOffset)) || 414 aStrides.size() != bStrides.size()) 415 return false; 416 417 // Strides along a dimension/offset are compatible if the value in the 418 // source memref is static and the value in the target memref is the 419 // same. They are also compatible if either one is dynamic (see 420 // description of MemRefCastOp for details). 421 auto checkCompatible = [](int64_t a, int64_t b) { 422 return (a == MemRefType::getDynamicStrideOrOffset() || 423 b == MemRefType::getDynamicStrideOrOffset() || a == b); 424 }; 425 if (!checkCompatible(aOffset, bOffset)) 426 return false; 427 for (auto aStride : enumerate(aStrides)) 428 if (!checkCompatible(aStride.value(), bStrides[aStride.index()])) 429 return false; 430 } 431 if (aT.getMemorySpace() != bT.getMemorySpace()) 432 return false; 433 434 // They must have the same rank, and any specified dimensions must match. 435 if (aT.getRank() != bT.getRank()) 436 return false; 437 438 for (unsigned i = 0, e = aT.getRank(); i != e; ++i) { 439 int64_t aDim = aT.getDimSize(i), bDim = bT.getDimSize(i); 440 if (aDim != -1 && bDim != -1 && aDim != bDim) 441 return false; 442 } 443 return true; 444 } else { 445 if (!aT && !uaT) 446 return false; 447 if (!bT && !ubT) 448 return false; 449 // Unranked to unranked casting is unsupported 450 if (uaT && ubT) 451 return false; 452 453 auto aEltType = (aT) ? aT.getElementType() : uaT.getElementType(); 454 auto bEltType = (bT) ? bT.getElementType() : ubT.getElementType(); 455 if (aEltType != bEltType) 456 return false; 457 458 auto aMemSpace = (aT) ? aT.getMemorySpace() : uaT.getMemorySpace(); 459 auto bMemSpace = (bT) ? bT.getMemorySpace() : ubT.getMemorySpace(); 460 if (aMemSpace != bMemSpace) 461 return false; 462 463 return true; 464 } 465 466 return false; 467 } 468 469 OpFoldResult CastOp::fold(ArrayRef<Attribute> operands) { 470 return succeeded(foldMemRefCast(*this)) ? getResult() : Value(); 471 } 472 473 //===----------------------------------------------------------------------===// 474 // CloneOp 475 //===----------------------------------------------------------------------===// 476 477 void CloneOp::getEffects( 478 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 479 &effects) { 480 effects.emplace_back(MemoryEffects::Read::get(), input(), 481 SideEffects::DefaultResource::get()); 482 effects.emplace_back(MemoryEffects::Write::get(), output(), 483 SideEffects::DefaultResource::get()); 484 } 485 486 namespace { 487 /// Fold Dealloc operations that are deallocating an AllocOp that is only used 488 /// by other Dealloc operations. 489 struct SimplifyClones : public OpRewritePattern<CloneOp> { 490 using OpRewritePattern<CloneOp>::OpRewritePattern; 491 492 LogicalResult matchAndRewrite(CloneOp cloneOp, 493 PatternRewriter &rewriter) const override { 494 if (cloneOp.use_empty()) { 495 rewriter.eraseOp(cloneOp); 496 return success(); 497 } 498 499 Value source = cloneOp.input(); 500 501 // Removes the clone operation and the corresponding dealloc and alloc 502 // operation (if any). 503 auto tryRemoveClone = [&](Operation *sourceOp, Operation *dealloc, 504 Operation *alloc) { 505 if (!sourceOp || !dealloc || !alloc || 506 alloc->getBlock() != dealloc->getBlock()) 507 return false; 508 rewriter.replaceOp(cloneOp, source); 509 rewriter.eraseOp(dealloc); 510 return true; 511 }; 512 513 // Removes unnecessary clones that are derived from the result of the clone 514 // op. 515 Operation *deallocOp = findDealloc(cloneOp.output()); 516 Operation *sourceOp = source.getDefiningOp(); 517 if (tryRemoveClone(sourceOp, deallocOp, sourceOp)) 518 return success(); 519 520 // Removes unnecessary clones that are derived from the source of the clone 521 // op. 522 deallocOp = findDealloc(source); 523 if (tryRemoveClone(sourceOp, deallocOp, cloneOp)) 524 return success(); 525 526 return failure(); 527 } 528 }; 529 530 } // end anonymous namespace. 531 532 void CloneOp::getCanonicalizationPatterns(OwningRewritePatternList &results, 533 MLIRContext *context) { 534 results.insert<SimplifyClones>(context); 535 } 536 537 OpFoldResult CloneOp::fold(ArrayRef<Attribute> operands) { 538 return succeeded(foldMemRefCast(*this)) ? getResult() : Value(); 539 } 540 541 //===----------------------------------------------------------------------===// 542 // DeallocOp 543 //===----------------------------------------------------------------------===// 544 545 LogicalResult DeallocOp::fold(ArrayRef<Attribute> cstOperands, 546 SmallVectorImpl<OpFoldResult> &results) { 547 /// dealloc(memrefcast) -> dealloc 548 return foldMemRefCast(*this); 549 } 550 551 //===----------------------------------------------------------------------===// 552 // DimOp 553 //===----------------------------------------------------------------------===// 554 555 void DimOp::build(OpBuilder &builder, OperationState &result, Value memref, 556 int64_t index) { 557 auto loc = result.location; 558 Value indexValue = builder.create<ConstantIndexOp>(loc, index); 559 build(builder, result, memref, indexValue); 560 } 561 562 void DimOp::build(OpBuilder &builder, OperationState &result, Value memref, 563 Value index) { 564 auto indexTy = builder.getIndexType(); 565 build(builder, result, indexTy, memref, index); 566 } 567 568 Optional<int64_t> DimOp::getConstantIndex() { 569 if (auto constantOp = index().getDefiningOp<ConstantOp>()) 570 return constantOp.getValue().cast<IntegerAttr>().getInt(); 571 return {}; 572 } 573 574 static LogicalResult verify(DimOp op) { 575 // Assume unknown index to be in range. 576 Optional<int64_t> index = op.getConstantIndex(); 577 if (!index.hasValue()) 578 return success(); 579 580 // Check that constant index is not knowingly out of range. 581 auto type = op.memrefOrTensor().getType(); 582 if (auto memrefType = type.dyn_cast<MemRefType>()) { 583 if (index.getValue() >= memrefType.getRank()) 584 return op.emitOpError("index is out of range"); 585 } else if (auto tensorType = type.dyn_cast<RankedTensorType>()) { 586 if (index.getValue() >= tensorType.getRank()) 587 return op.emitOpError("index is out of range"); 588 } else if (type.isa<UnrankedMemRefType>() || type.isa<UnrankedTensorType>()) { 589 // Assume index to be in range. 590 } else { 591 llvm_unreachable("expected operand with memref type"); 592 } 593 return success(); 594 } 595 596 OpFoldResult DimOp::fold(ArrayRef<Attribute> operands) { 597 auto index = operands[1].dyn_cast_or_null<IntegerAttr>(); 598 599 // All forms of folding require a known index. 600 if (!index) 601 return {}; 602 603 auto argTy = memrefOrTensor().getType(); 604 // Fold if the shape extent along the given index is known. 605 if (auto shapedTy = argTy.dyn_cast<ShapedType>()) { 606 // Folding for unranked types (UnrankedMemRefType) is not supported. 607 if (!shapedTy.hasRank()) 608 return {}; 609 if (!shapedTy.isDynamicDim(index.getInt())) { 610 Builder builder(getContext()); 611 return builder.getIndexAttr(shapedTy.getShape()[index.getInt()]); 612 } 613 } 614 615 Operation *definingOp = memrefOrTensor().getDefiningOp(); 616 617 // dim(memref.tensor_load(memref)) -> dim(memref) 618 if (auto tensorLoadOp = dyn_cast_or_null<TensorLoadOp>(definingOp)) { 619 setOperand(0, tensorLoadOp.memref()); 620 return getResult(); 621 } 622 623 // Fold dim to the operand of tensor.generate. 624 if (auto fromElements = dyn_cast_or_null<tensor::GenerateOp>(definingOp)) { 625 auto resultType = 626 fromElements.getResult().getType().cast<RankedTensorType>(); 627 // The case where the type encodes the size of the dimension is handled 628 // above. 629 assert(resultType.getShape()[index.getInt()] == 630 RankedTensorType::kDynamicSize); 631 632 // Find the operand of the fromElements that corresponds to this index. 633 auto dynExtents = fromElements.dynamicExtents().begin(); 634 for (auto dim : resultType.getShape().take_front(index.getInt())) 635 if (dim == RankedTensorType::kDynamicSize) 636 dynExtents++; 637 638 return Value{*dynExtents}; 639 } 640 641 // The size at the given index is now known to be a dynamic size. 642 unsigned unsignedIndex = index.getValue().getZExtValue(); 643 644 if (auto subtensor = dyn_cast_or_null<mlir::SubTensorOp>(definingOp)) { 645 assert(subtensor.isDynamicSize(unsignedIndex) && 646 "Expected dynamic subtensor size"); 647 return subtensor.getDynamicSize(unsignedIndex); 648 } 649 650 // Fold dim to the size argument for an `AllocOp`, `ViewOp`, or `SubViewOp`. 651 auto memrefType = argTy.dyn_cast<MemRefType>(); 652 if (!memrefType) 653 return {}; 654 655 if (auto alloc = dyn_cast_or_null<AllocOp>(definingOp)) 656 return *(alloc.getDynamicSizes().begin() + 657 memrefType.getDynamicDimIndex(unsignedIndex)); 658 659 if (auto alloca = dyn_cast_or_null<AllocaOp>(definingOp)) 660 return *(alloca.getDynamicSizes().begin() + 661 memrefType.getDynamicDimIndex(unsignedIndex)); 662 663 if (auto view = dyn_cast_or_null<ViewOp>(definingOp)) 664 return *(view.getDynamicSizes().begin() + 665 memrefType.getDynamicDimIndex(unsignedIndex)); 666 667 if (auto subview = dyn_cast_or_null<SubViewOp>(definingOp)) { 668 assert(subview.isDynamicSize(unsignedIndex) && 669 "Expected dynamic subview size"); 670 return subview.getDynamicSize(unsignedIndex); 671 } 672 673 // dim(memrefcast) -> dim 674 if (succeeded(foldMemRefCast(*this))) 675 return getResult(); 676 677 return {}; 678 } 679 680 namespace { 681 /// Fold dim of a memref reshape operation to a load into the reshape's shape 682 /// operand. 683 struct DimOfMemRefReshape : public OpRewritePattern<DimOp> { 684 using OpRewritePattern<DimOp>::OpRewritePattern; 685 686 LogicalResult matchAndRewrite(DimOp dim, 687 PatternRewriter &rewriter) const override { 688 auto reshape = dim.memrefOrTensor().getDefiningOp<ReshapeOp>(); 689 690 if (!reshape) 691 return failure(); 692 693 // Place the load directly after the reshape to ensure that the shape memref 694 // was not mutated. 695 rewriter.setInsertionPointAfter(reshape); 696 rewriter.replaceOpWithNewOp<LoadOp>(dim, reshape.shape(), 697 llvm::makeArrayRef({dim.index()})); 698 return success(); 699 } 700 }; 701 702 /// Fold dim of a dim of a cast into the dim of the source of the tensor cast. 703 template <typename CastOpTy> 704 struct DimOfCastOp : public OpRewritePattern<DimOp> { 705 using OpRewritePattern<DimOp>::OpRewritePattern; 706 707 LogicalResult matchAndRewrite(DimOp dimOp, 708 PatternRewriter &rewriter) const override { 709 auto castOp = dimOp.memrefOrTensor().getDefiningOp<CastOpTy>(); 710 if (!castOp) 711 return failure(); 712 Value newSource = castOp.getOperand(); 713 rewriter.replaceOpWithNewOp<DimOp>(dimOp, newSource, dimOp.index()); 714 return success(); 715 } 716 }; 717 718 /// Helper method to get the `Value` that is the shape of the `resultIdx`-th 719 /// result at dimension `dimIndex` from the `ShapedTypeOpInterface`. 720 /// TODO(ravishankarm): This is better put as a interface utility method 721 /// somewhere, but that would imply the interface will depend on the `tensor` 722 /// dialect. Ideally maybe a utility method in the `tensor` dialect. 723 static Value getResultDimFromShapeInterface(OpBuilder &builder, OpResult result, 724 int64_t dimIndex) { 725 unsigned resultNumber = result.getResultNumber(); 726 auto shapedTypeOp = dyn_cast<InferShapedTypeOpInterface>(result.getOwner()); 727 Location loc = result.getOwner()->getLoc(); 728 if (!shapedTypeOp) 729 return nullptr; 730 731 // The interface exposes two methods, one that returns the shape of all the 732 // results as `Value` and other that returns the shape as a list of 733 // `SmallVector<Value>`. The former takes precedence over the latter. So first 734 // check if the op implements the first interface method or the second, and 735 // get the value to use appropriately. 736 SmallVector<Value> reifiedResultShapes; 737 if (succeeded(shapedTypeOp.reifyReturnTypeShapes( 738 builder, result.getOwner()->getOperands(), reifiedResultShapes))) { 739 if (reifiedResultShapes.size() <= resultNumber) 740 return nullptr; 741 Value resultShape = reifiedResultShapes[resultNumber]; 742 auto resultShapeType = resultShape.getType().dyn_cast<RankedTensorType>(); 743 if (!resultShapeType || !resultShapeType.getElementType().isa<IndexType>()) 744 return nullptr; 745 return builder.create<tensor::ExtractOp>( 746 loc, resultShape, builder.createOrFold<ConstantIndexOp>(loc, dimIndex)); 747 } 748 749 SmallVector<SmallVector<Value>> reifiedResultShapesPerDim; 750 if (failed(shapedTypeOp.reifyReturnTypeShapesPerResultDim( 751 builder, reifiedResultShapesPerDim))) 752 return nullptr; 753 if (reifiedResultShapesPerDim.size() <= resultNumber || 754 reifiedResultShapesPerDim[resultNumber].size() != 755 static_cast<size_t>(result.getType().cast<ShapedType>().getRank())) 756 return nullptr; 757 OpFoldResult valueOrAttr = reifiedResultShapesPerDim[resultNumber][dimIndex]; 758 if (auto attr = valueOrAttr.dyn_cast<Attribute>()) 759 return builder.createOrFold<ConstantIndexOp>( 760 loc, attr.cast<IntegerAttr>().getInt()); 761 return valueOrAttr.get<Value>(); 762 } 763 764 /// Fold dim of an operation that implements the InferShapedTypeOpInterface 765 struct DimOfShapedTypeOpInterface : public OpRewritePattern<DimOp> { 766 using OpRewritePattern<DimOp>::OpRewritePattern; 767 768 LogicalResult matchAndRewrite(DimOp dimOp, 769 PatternRewriter &rewriter) const override { 770 OpResult dimValue = dimOp.memrefOrTensor().dyn_cast<OpResult>(); 771 if (!dimValue) 772 return failure(); 773 auto shapedTypeOp = 774 dyn_cast<InferShapedTypeOpInterface>(dimValue.getOwner()); 775 if (!shapedTypeOp) 776 return failure(); 777 778 Optional<int64_t> dimIndex = dimOp.getConstantIndex(); 779 if (!dimIndex) 780 return failure(); 781 Value replacement = 782 getResultDimFromShapeInterface(rewriter, dimValue, *dimIndex); 783 if (!replacement) 784 return failure(); 785 rewriter.replaceOp(dimOp, replacement); 786 return success(); 787 } 788 }; 789 } // end anonymous namespace. 790 791 void DimOp::getCanonicalizationPatterns(RewritePatternSet &results, 792 MLIRContext *context) { 793 results.add<DimOfMemRefReshape, DimOfCastOp<BufferCastOp>, 794 DimOfCastOp<tensor::CastOp>, DimOfShapedTypeOpInterface>(context); 795 } 796 797 // --------------------------------------------------------------------------- 798 // DmaStartOp 799 // --------------------------------------------------------------------------- 800 801 void DmaStartOp::build(OpBuilder &builder, OperationState &result, 802 Value srcMemRef, ValueRange srcIndices, Value destMemRef, 803 ValueRange destIndices, Value numElements, 804 Value tagMemRef, ValueRange tagIndices, Value stride, 805 Value elementsPerStride) { 806 result.addOperands(srcMemRef); 807 result.addOperands(srcIndices); 808 result.addOperands(destMemRef); 809 result.addOperands(destIndices); 810 result.addOperands({numElements, tagMemRef}); 811 result.addOperands(tagIndices); 812 if (stride) 813 result.addOperands({stride, elementsPerStride}); 814 } 815 816 void DmaStartOp::print(OpAsmPrinter &p) { 817 p << getOperationName() << " " << getSrcMemRef() << '[' << getSrcIndices() 818 << "], " << getDstMemRef() << '[' << getDstIndices() << "], " 819 << getNumElements() << ", " << getTagMemRef() << '[' << getTagIndices() 820 << ']'; 821 if (isStrided()) 822 p << ", " << getStride() << ", " << getNumElementsPerStride(); 823 824 p.printOptionalAttrDict((*this)->getAttrs()); 825 p << " : " << getSrcMemRef().getType() << ", " << getDstMemRef().getType() 826 << ", " << getTagMemRef().getType(); 827 } 828 829 // Parse DmaStartOp. 830 // Ex: 831 // %dma_id = dma_start %src[%i, %j], %dst[%k, %l], %size, 832 // %tag[%index], %stride, %num_elt_per_stride : 833 // : memref<3076 x f32, 0>, 834 // memref<1024 x f32, 2>, 835 // memref<1 x i32> 836 // 837 ParseResult DmaStartOp::parse(OpAsmParser &parser, OperationState &result) { 838 OpAsmParser::OperandType srcMemRefInfo; 839 SmallVector<OpAsmParser::OperandType, 4> srcIndexInfos; 840 OpAsmParser::OperandType dstMemRefInfo; 841 SmallVector<OpAsmParser::OperandType, 4> dstIndexInfos; 842 OpAsmParser::OperandType numElementsInfo; 843 OpAsmParser::OperandType tagMemrefInfo; 844 SmallVector<OpAsmParser::OperandType, 4> tagIndexInfos; 845 SmallVector<OpAsmParser::OperandType, 2> strideInfo; 846 847 SmallVector<Type, 3> types; 848 auto indexType = parser.getBuilder().getIndexType(); 849 850 // Parse and resolve the following list of operands: 851 // *) source memref followed by its indices (in square brackets). 852 // *) destination memref followed by its indices (in square brackets). 853 // *) dma size in KiB. 854 if (parser.parseOperand(srcMemRefInfo) || 855 parser.parseOperandList(srcIndexInfos, OpAsmParser::Delimiter::Square) || 856 parser.parseComma() || parser.parseOperand(dstMemRefInfo) || 857 parser.parseOperandList(dstIndexInfos, OpAsmParser::Delimiter::Square) || 858 parser.parseComma() || parser.parseOperand(numElementsInfo) || 859 parser.parseComma() || parser.parseOperand(tagMemrefInfo) || 860 parser.parseOperandList(tagIndexInfos, OpAsmParser::Delimiter::Square)) 861 return failure(); 862 863 // Parse optional stride and elements per stride. 864 if (parser.parseTrailingOperandList(strideInfo)) 865 return failure(); 866 867 bool isStrided = strideInfo.size() == 2; 868 if (!strideInfo.empty() && !isStrided) { 869 return parser.emitError(parser.getNameLoc(), 870 "expected two stride related operands"); 871 } 872 873 if (parser.parseColonTypeList(types)) 874 return failure(); 875 if (types.size() != 3) 876 return parser.emitError(parser.getNameLoc(), "fewer/more types expected"); 877 878 if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) || 879 parser.resolveOperands(srcIndexInfos, indexType, result.operands) || 880 parser.resolveOperand(dstMemRefInfo, types[1], result.operands) || 881 parser.resolveOperands(dstIndexInfos, indexType, result.operands) || 882 // size should be an index. 883 parser.resolveOperand(numElementsInfo, indexType, result.operands) || 884 parser.resolveOperand(tagMemrefInfo, types[2], result.operands) || 885 // tag indices should be index. 886 parser.resolveOperands(tagIndexInfos, indexType, result.operands)) 887 return failure(); 888 889 if (isStrided) { 890 if (parser.resolveOperands(strideInfo, indexType, result.operands)) 891 return failure(); 892 } 893 894 return success(); 895 } 896 897 LogicalResult DmaStartOp::verify() { 898 unsigned numOperands = getNumOperands(); 899 900 // Mandatory non-variadic operands are: src memref, dst memref, tag memref and 901 // the number of elements. 902 if (numOperands < 4) 903 return emitOpError("expected at least 4 operands"); 904 905 // Check types of operands. The order of these calls is important: the later 906 // calls rely on some type properties to compute the operand position. 907 // 1. Source memref. 908 if (!getSrcMemRef().getType().isa<MemRefType>()) 909 return emitOpError("expected source to be of memref type"); 910 if (numOperands < getSrcMemRefRank() + 4) 911 return emitOpError() << "expected at least " << getSrcMemRefRank() + 4 912 << " operands"; 913 if (!getSrcIndices().empty() && 914 !llvm::all_of(getSrcIndices().getTypes(), 915 [](Type t) { return t.isIndex(); })) 916 return emitOpError("expected source indices to be of index type"); 917 918 // 2. Destination memref. 919 if (!getDstMemRef().getType().isa<MemRefType>()) 920 return emitOpError("expected destination to be of memref type"); 921 unsigned numExpectedOperands = getSrcMemRefRank() + getDstMemRefRank() + 4; 922 if (numOperands < numExpectedOperands) 923 return emitOpError() << "expected at least " << numExpectedOperands 924 << " operands"; 925 if (!getDstIndices().empty() && 926 !llvm::all_of(getDstIndices().getTypes(), 927 [](Type t) { return t.isIndex(); })) 928 return emitOpError("expected destination indices to be of index type"); 929 930 // 3. Number of elements. 931 if (!getNumElements().getType().isIndex()) 932 return emitOpError("expected num elements to be of index type"); 933 934 // 4. Tag memref. 935 if (!getTagMemRef().getType().isa<MemRefType>()) 936 return emitOpError("expected tag to be of memref type"); 937 numExpectedOperands += getTagMemRefRank(); 938 if (numOperands < numExpectedOperands) 939 return emitOpError() << "expected at least " << numExpectedOperands 940 << " operands"; 941 if (!getTagIndices().empty() && 942 !llvm::all_of(getTagIndices().getTypes(), 943 [](Type t) { return t.isIndex(); })) 944 return emitOpError("expected tag indices to be of index type"); 945 946 // Optional stride-related operands must be either both present or both 947 // absent. 948 if (numOperands != numExpectedOperands && 949 numOperands != numExpectedOperands + 2) 950 return emitOpError("incorrect number of operands"); 951 952 // 5. Strides. 953 if (isStrided()) { 954 if (!getStride().getType().isIndex() || 955 !getNumElementsPerStride().getType().isIndex()) 956 return emitOpError( 957 "expected stride and num elements per stride to be of type index"); 958 } 959 960 return success(); 961 } 962 963 LogicalResult DmaStartOp::fold(ArrayRef<Attribute> cstOperands, 964 SmallVectorImpl<OpFoldResult> &results) { 965 /// dma_start(memrefcast) -> dma_start 966 return foldMemRefCast(*this); 967 } 968 969 // --------------------------------------------------------------------------- 970 // DmaWaitOp 971 // --------------------------------------------------------------------------- 972 973 void DmaWaitOp::build(OpBuilder &builder, OperationState &result, 974 Value tagMemRef, ValueRange tagIndices, 975 Value numElements) { 976 result.addOperands(tagMemRef); 977 result.addOperands(tagIndices); 978 result.addOperands(numElements); 979 } 980 981 void DmaWaitOp::print(OpAsmPrinter &p) { 982 p << getOperationName() << " " << getTagMemRef() << '[' << getTagIndices() 983 << "], " << getNumElements(); 984 p.printOptionalAttrDict((*this)->getAttrs()); 985 p << " : " << getTagMemRef().getType(); 986 } 987 988 // Parse DmaWaitOp. 989 // Eg: 990 // dma_wait %tag[%index], %num_elements : memref<1 x i32, (d0) -> (d0), 4> 991 // 992 ParseResult DmaWaitOp::parse(OpAsmParser &parser, OperationState &result) { 993 OpAsmParser::OperandType tagMemrefInfo; 994 SmallVector<OpAsmParser::OperandType, 2> tagIndexInfos; 995 Type type; 996 auto indexType = parser.getBuilder().getIndexType(); 997 OpAsmParser::OperandType numElementsInfo; 998 999 // Parse tag memref, its indices, and dma size. 1000 if (parser.parseOperand(tagMemrefInfo) || 1001 parser.parseOperandList(tagIndexInfos, OpAsmParser::Delimiter::Square) || 1002 parser.parseComma() || parser.parseOperand(numElementsInfo) || 1003 parser.parseColonType(type) || 1004 parser.resolveOperand(tagMemrefInfo, type, result.operands) || 1005 parser.resolveOperands(tagIndexInfos, indexType, result.operands) || 1006 parser.resolveOperand(numElementsInfo, indexType, result.operands)) 1007 return failure(); 1008 1009 return success(); 1010 } 1011 1012 LogicalResult DmaWaitOp::fold(ArrayRef<Attribute> cstOperands, 1013 SmallVectorImpl<OpFoldResult> &results) { 1014 /// dma_wait(memrefcast) -> dma_wait 1015 return foldMemRefCast(*this); 1016 } 1017 1018 LogicalResult DmaWaitOp::verify() { 1019 // Mandatory non-variadic operands are tag and the number of elements. 1020 if (getNumOperands() < 2) 1021 return emitOpError() << "expected at least 2 operands"; 1022 1023 // Check types of operands. The order of these calls is important: the later 1024 // calls rely on some type properties to compute the operand position. 1025 if (!getTagMemRef().getType().isa<MemRefType>()) 1026 return emitOpError() << "expected tag to be of memref type"; 1027 1028 if (getNumOperands() != 2 + getTagMemRefRank()) 1029 return emitOpError() << "expected " << 2 + getTagMemRefRank() 1030 << " operands"; 1031 1032 if (!getTagIndices().empty() && 1033 !llvm::all_of(getTagIndices().getTypes(), 1034 [](Type t) { return t.isIndex(); })) 1035 return emitOpError() << "expected tag indices to be of index type"; 1036 1037 if (!getNumElements().getType().isIndex()) 1038 return emitOpError() 1039 << "expected the number of elements to be of index type"; 1040 1041 return success(); 1042 } 1043 1044 //===----------------------------------------------------------------------===// 1045 // GlobalOp 1046 //===----------------------------------------------------------------------===// 1047 1048 static void printGlobalMemrefOpTypeAndInitialValue(OpAsmPrinter &p, GlobalOp op, 1049 TypeAttr type, 1050 Attribute initialValue) { 1051 p << type; 1052 if (!op.isExternal()) { 1053 p << " = "; 1054 if (op.isUninitialized()) 1055 p << "uninitialized"; 1056 else 1057 p.printAttributeWithoutType(initialValue); 1058 } 1059 } 1060 1061 static ParseResult 1062 parseGlobalMemrefOpTypeAndInitialValue(OpAsmParser &parser, TypeAttr &typeAttr, 1063 Attribute &initialValue) { 1064 Type type; 1065 if (parser.parseType(type)) 1066 return failure(); 1067 1068 auto memrefType = type.dyn_cast<MemRefType>(); 1069 if (!memrefType || !memrefType.hasStaticShape()) 1070 return parser.emitError(parser.getNameLoc()) 1071 << "type should be static shaped memref, but got " << type; 1072 typeAttr = TypeAttr::get(type); 1073 1074 if (parser.parseOptionalEqual()) 1075 return success(); 1076 1077 if (succeeded(parser.parseOptionalKeyword("uninitialized"))) { 1078 initialValue = UnitAttr::get(parser.getBuilder().getContext()); 1079 return success(); 1080 } 1081 1082 Type tensorType = getTensorTypeFromMemRefType(memrefType); 1083 if (parser.parseAttribute(initialValue, tensorType)) 1084 return failure(); 1085 if (!initialValue.isa<ElementsAttr>()) 1086 return parser.emitError(parser.getNameLoc()) 1087 << "initial value should be a unit or elements attribute"; 1088 return success(); 1089 } 1090 1091 static LogicalResult verify(GlobalOp op) { 1092 auto memrefType = op.type().dyn_cast<MemRefType>(); 1093 if (!memrefType || !memrefType.hasStaticShape()) 1094 return op.emitOpError("type should be static shaped memref, but got ") 1095 << op.type(); 1096 1097 // Verify that the initial value, if present, is either a unit attribute or 1098 // an elements attribute. 1099 if (op.initial_value().hasValue()) { 1100 Attribute initValue = op.initial_value().getValue(); 1101 if (!initValue.isa<UnitAttr>() && !initValue.isa<ElementsAttr>()) 1102 return op.emitOpError("initial value should be a unit or elements " 1103 "attribute, but got ") 1104 << initValue; 1105 1106 // Check that the type of the initial value is compatible with the type of 1107 // the global variable. 1108 if (initValue.isa<ElementsAttr>()) { 1109 Type initType = initValue.getType(); 1110 Type tensorType = getTensorTypeFromMemRefType(memrefType); 1111 if (initType != tensorType) 1112 return op.emitOpError("initial value expected to be of type ") 1113 << tensorType << ", but was of type " << initType; 1114 } 1115 } 1116 1117 // TODO: verify visibility for declarations. 1118 return success(); 1119 } 1120 1121 //===----------------------------------------------------------------------===// 1122 // GetGlobalOp 1123 //===----------------------------------------------------------------------===// 1124 1125 LogicalResult 1126 GetGlobalOp::verifySymbolUses(SymbolTableCollection &symbolTable) { 1127 // Verify that the result type is same as the type of the referenced 1128 // memref.global op. 1129 auto global = 1130 symbolTable.lookupNearestSymbolFrom<GlobalOp>(*this, nameAttr()); 1131 if (!global) 1132 return emitOpError("'") 1133 << name() << "' does not reference a valid global memref"; 1134 1135 Type resultType = result().getType(); 1136 if (global.type() != resultType) 1137 return emitOpError("result type ") 1138 << resultType << " does not match type " << global.type() 1139 << " of the global memref @" << name(); 1140 return success(); 1141 } 1142 1143 //===----------------------------------------------------------------------===// 1144 // LoadOp 1145 //===----------------------------------------------------------------------===// 1146 1147 static LogicalResult verify(LoadOp op) { 1148 if (op.getNumOperands() != 1 + op.getMemRefType().getRank()) 1149 return op.emitOpError("incorrect number of indices for load"); 1150 return success(); 1151 } 1152 1153 OpFoldResult LoadOp::fold(ArrayRef<Attribute> cstOperands) { 1154 /// load(memrefcast) -> load 1155 if (succeeded(foldMemRefCast(*this))) 1156 return getResult(); 1157 return OpFoldResult(); 1158 } 1159 1160 namespace { 1161 /// Fold a load on a buffer_cast operation into an tensor.extract on the 1162 /// corresponding tensor. 1163 struct LoadOfBufferCast : public OpRewritePattern<LoadOp> { 1164 using OpRewritePattern<LoadOp>::OpRewritePattern; 1165 1166 LogicalResult matchAndRewrite(LoadOp load, 1167 PatternRewriter &rewriter) const override { 1168 auto buffercast = load.memref().getDefiningOp<BufferCastOp>(); 1169 if (!buffercast) 1170 return failure(); 1171 1172 rewriter.replaceOpWithNewOp<tensor::ExtractOp>(load, buffercast.tensor(), 1173 load.indices()); 1174 return success(); 1175 } 1176 }; 1177 } // end anonymous namespace. 1178 1179 void LoadOp::getCanonicalizationPatterns(RewritePatternSet &results, 1180 MLIRContext *context) { 1181 results.add<LoadOfBufferCast>(context); 1182 } 1183 1184 //===----------------------------------------------------------------------===// 1185 // PrefetchOp 1186 //===----------------------------------------------------------------------===// 1187 1188 static void print(OpAsmPrinter &p, PrefetchOp op) { 1189 p << PrefetchOp::getOperationName() << " " << op.memref() << '['; 1190 p.printOperands(op.indices()); 1191 p << ']' << ", " << (op.isWrite() ? "write" : "read"); 1192 p << ", locality<" << op.localityHint(); 1193 p << ">, " << (op.isDataCache() ? "data" : "instr"); 1194 p.printOptionalAttrDict( 1195 op->getAttrs(), 1196 /*elidedAttrs=*/{"localityHint", "isWrite", "isDataCache"}); 1197 p << " : " << op.getMemRefType(); 1198 } 1199 1200 static ParseResult parsePrefetchOp(OpAsmParser &parser, 1201 OperationState &result) { 1202 OpAsmParser::OperandType memrefInfo; 1203 SmallVector<OpAsmParser::OperandType, 4> indexInfo; 1204 IntegerAttr localityHint; 1205 MemRefType type; 1206 StringRef readOrWrite, cacheType; 1207 1208 auto indexTy = parser.getBuilder().getIndexType(); 1209 auto i32Type = parser.getBuilder().getIntegerType(32); 1210 if (parser.parseOperand(memrefInfo) || 1211 parser.parseOperandList(indexInfo, OpAsmParser::Delimiter::Square) || 1212 parser.parseComma() || parser.parseKeyword(&readOrWrite) || 1213 parser.parseComma() || parser.parseKeyword("locality") || 1214 parser.parseLess() || 1215 parser.parseAttribute(localityHint, i32Type, "localityHint", 1216 result.attributes) || 1217 parser.parseGreater() || parser.parseComma() || 1218 parser.parseKeyword(&cacheType) || parser.parseColonType(type) || 1219 parser.resolveOperand(memrefInfo, type, result.operands) || 1220 parser.resolveOperands(indexInfo, indexTy, result.operands)) 1221 return failure(); 1222 1223 if (!readOrWrite.equals("read") && !readOrWrite.equals("write")) 1224 return parser.emitError(parser.getNameLoc(), 1225 "rw specifier has to be 'read' or 'write'"); 1226 result.addAttribute( 1227 PrefetchOp::getIsWriteAttrName(), 1228 parser.getBuilder().getBoolAttr(readOrWrite.equals("write"))); 1229 1230 if (!cacheType.equals("data") && !cacheType.equals("instr")) 1231 return parser.emitError(parser.getNameLoc(), 1232 "cache type has to be 'data' or 'instr'"); 1233 1234 result.addAttribute( 1235 PrefetchOp::getIsDataCacheAttrName(), 1236 parser.getBuilder().getBoolAttr(cacheType.equals("data"))); 1237 1238 return success(); 1239 } 1240 1241 static LogicalResult verify(PrefetchOp op) { 1242 if (op.getNumOperands() != 1 + op.getMemRefType().getRank()) 1243 return op.emitOpError("too few indices"); 1244 1245 return success(); 1246 } 1247 1248 LogicalResult PrefetchOp::fold(ArrayRef<Attribute> cstOperands, 1249 SmallVectorImpl<OpFoldResult> &results) { 1250 // prefetch(memrefcast) -> prefetch 1251 return foldMemRefCast(*this); 1252 } 1253 1254 //===----------------------------------------------------------------------===// 1255 // ReinterpretCastOp 1256 //===----------------------------------------------------------------------===// 1257 1258 /// Build a ReinterpretCastOp with all dynamic entries: `staticOffsets`, 1259 /// `staticSizes` and `staticStrides` are automatically filled with 1260 /// source-memref-rank sentinel values that encode dynamic entries. 1261 void ReinterpretCastOp::build(OpBuilder &b, OperationState &result, 1262 MemRefType resultType, Value source, 1263 OpFoldResult offset, ArrayRef<OpFoldResult> sizes, 1264 ArrayRef<OpFoldResult> strides, 1265 ArrayRef<NamedAttribute> attrs) { 1266 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides; 1267 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides; 1268 dispatchIndexOpFoldResults(offset, dynamicOffsets, staticOffsets, 1269 ShapedType::kDynamicStrideOrOffset); 1270 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes, 1271 ShapedType::kDynamicSize); 1272 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides, 1273 ShapedType::kDynamicStrideOrOffset); 1274 build(b, result, resultType, source, dynamicOffsets, dynamicSizes, 1275 dynamicStrides, b.getI64ArrayAttr(staticOffsets), 1276 b.getI64ArrayAttr(staticSizes), b.getI64ArrayAttr(staticStrides)); 1277 result.addAttributes(attrs); 1278 } 1279 1280 void ReinterpretCastOp::build(OpBuilder &b, OperationState &result, 1281 MemRefType resultType, Value source, 1282 int64_t offset, ArrayRef<int64_t> sizes, 1283 ArrayRef<int64_t> strides, 1284 ArrayRef<NamedAttribute> attrs) { 1285 SmallVector<OpFoldResult> sizeValues = 1286 llvm::to_vector<4>(llvm::map_range(sizes, [&](int64_t v) -> OpFoldResult { 1287 return b.getI64IntegerAttr(v); 1288 })); 1289 SmallVector<OpFoldResult> strideValues = llvm::to_vector<4>( 1290 llvm::map_range(strides, [&](int64_t v) -> OpFoldResult { 1291 return b.getI64IntegerAttr(v); 1292 })); 1293 build(b, result, resultType, source, b.getI64IntegerAttr(offset), sizeValues, 1294 strideValues, attrs); 1295 } 1296 1297 void ReinterpretCastOp::build(OpBuilder &b, OperationState &result, 1298 MemRefType resultType, Value source, Value offset, 1299 ValueRange sizes, ValueRange strides, 1300 ArrayRef<NamedAttribute> attrs) { 1301 SmallVector<OpFoldResult> sizeValues = llvm::to_vector<4>( 1302 llvm::map_range(sizes, [](Value v) -> OpFoldResult { return v; })); 1303 SmallVector<OpFoldResult> strideValues = llvm::to_vector<4>( 1304 llvm::map_range(strides, [](Value v) -> OpFoldResult { return v; })); 1305 build(b, result, resultType, source, offset, sizeValues, strideValues, attrs); 1306 } 1307 1308 // TODO: ponder whether we want to allow missing trailing sizes/strides that are 1309 // completed automatically, like we have for subview and subtensor. 1310 static LogicalResult verify(ReinterpretCastOp op) { 1311 // The source and result memrefs should be in the same memory space. 1312 auto srcType = op.source().getType().cast<BaseMemRefType>(); 1313 auto resultType = op.getType().cast<MemRefType>(); 1314 if (srcType.getMemorySpace() != resultType.getMemorySpace()) 1315 return op.emitError("different memory spaces specified for source type ") 1316 << srcType << " and result memref type " << resultType; 1317 if (srcType.getElementType() != resultType.getElementType()) 1318 return op.emitError("different element types specified for source type ") 1319 << srcType << " and result memref type " << resultType; 1320 1321 // Match sizes in result memref type and in static_sizes attribute. 1322 for (auto &en : 1323 llvm::enumerate(llvm::zip(resultType.getShape(), 1324 extractFromI64ArrayAttr(op.static_sizes())))) { 1325 int64_t resultSize = std::get<0>(en.value()); 1326 int64_t expectedSize = std::get<1>(en.value()); 1327 if (resultSize != expectedSize) 1328 return op.emitError("expected result type with size = ") 1329 << expectedSize << " instead of " << resultSize 1330 << " in dim = " << en.index(); 1331 } 1332 1333 // Match offset and strides in static_offset and static_strides attributes if 1334 // result memref type has an affine map specified. 1335 if (!resultType.getAffineMaps().empty()) { 1336 int64_t resultOffset; 1337 SmallVector<int64_t, 4> resultStrides; 1338 if (failed(getStridesAndOffset(resultType, resultStrides, resultOffset))) 1339 return failure(); 1340 1341 // Match offset in result memref type and in static_offsets attribute. 1342 int64_t expectedOffset = 1343 extractFromI64ArrayAttr(op.static_offsets()).front(); 1344 if (resultOffset != expectedOffset) 1345 return op.emitError("expected result type with offset = ") 1346 << resultOffset << " instead of " << expectedOffset; 1347 1348 // Match strides in result memref type and in static_strides attribute. 1349 for (auto &en : llvm::enumerate(llvm::zip( 1350 resultStrides, extractFromI64ArrayAttr(op.static_strides())))) { 1351 int64_t resultStride = std::get<0>(en.value()); 1352 int64_t expectedStride = std::get<1>(en.value()); 1353 if (resultStride != expectedStride) 1354 return op.emitError("expected result type with stride = ") 1355 << expectedStride << " instead of " << resultStride 1356 << " in dim = " << en.index(); 1357 } 1358 } 1359 return success(); 1360 } 1361 1362 //===----------------------------------------------------------------------===// 1363 // ReshapeOp 1364 //===----------------------------------------------------------------------===// 1365 1366 static LogicalResult verify(ReshapeOp op) { 1367 Type operandType = op.source().getType(); 1368 Type resultType = op.result().getType(); 1369 1370 Type operandElementType = operandType.cast<ShapedType>().getElementType(); 1371 Type resultElementType = resultType.cast<ShapedType>().getElementType(); 1372 if (operandElementType != resultElementType) 1373 return op.emitOpError("element types of source and destination memref " 1374 "types should be the same"); 1375 1376 if (auto operandMemRefType = operandType.dyn_cast<MemRefType>()) 1377 if (!operandMemRefType.getAffineMaps().empty()) 1378 return op.emitOpError( 1379 "source memref type should have identity affine map"); 1380 1381 int64_t shapeSize = op.shape().getType().cast<MemRefType>().getDimSize(0); 1382 auto resultMemRefType = resultType.dyn_cast<MemRefType>(); 1383 if (resultMemRefType) { 1384 if (!resultMemRefType.getAffineMaps().empty()) 1385 return op.emitOpError( 1386 "result memref type should have identity affine map"); 1387 if (shapeSize == ShapedType::kDynamicSize) 1388 return op.emitOpError("cannot use shape operand with dynamic length to " 1389 "reshape to statically-ranked memref type"); 1390 if (shapeSize != resultMemRefType.getRank()) 1391 return op.emitOpError( 1392 "length of shape operand differs from the result's memref rank"); 1393 } 1394 return success(); 1395 } 1396 1397 //===----------------------------------------------------------------------===// 1398 // StoreOp 1399 //===----------------------------------------------------------------------===// 1400 1401 static LogicalResult verify(StoreOp op) { 1402 if (op.getNumOperands() != 2 + op.getMemRefType().getRank()) 1403 return op.emitOpError("store index operand count not equal to memref rank"); 1404 1405 return success(); 1406 } 1407 1408 LogicalResult StoreOp::fold(ArrayRef<Attribute> cstOperands, 1409 SmallVectorImpl<OpFoldResult> &results) { 1410 /// store(memrefcast) -> store 1411 return foldMemRefCast(*this); 1412 } 1413 1414 //===----------------------------------------------------------------------===// 1415 // SubViewOp 1416 //===----------------------------------------------------------------------===// 1417 1418 namespace { 1419 /// Helpers to write more idiomatic operations. 1420 namespace saturated_arith { 1421 struct Wrapper { 1422 explicit Wrapper(int64_t v) : v(v) {} 1423 operator int64_t() { return v; } 1424 int64_t v; 1425 }; 1426 Wrapper operator+(Wrapper a, int64_t b) { 1427 if (ShapedType::isDynamicStrideOrOffset(a) || 1428 ShapedType::isDynamicStrideOrOffset(b)) 1429 return Wrapper(ShapedType::kDynamicStrideOrOffset); 1430 return Wrapper(a.v + b); 1431 } 1432 Wrapper operator*(Wrapper a, int64_t b) { 1433 if (ShapedType::isDynamicStrideOrOffset(a) || 1434 ShapedType::isDynamicStrideOrOffset(b)) 1435 return Wrapper(ShapedType::kDynamicStrideOrOffset); 1436 return Wrapper(a.v * b); 1437 } 1438 } // end namespace saturated_arith 1439 } // end namespace 1440 1441 /// A subview result type can be fully inferred from the source type and the 1442 /// static representation of offsets, sizes and strides. Special sentinels 1443 /// encode the dynamic case. 1444 Type SubViewOp::inferResultType(MemRefType sourceMemRefType, 1445 ArrayRef<int64_t> leadingStaticOffsets, 1446 ArrayRef<int64_t> leadingStaticSizes, 1447 ArrayRef<int64_t> leadingStaticStrides) { 1448 // A subview may specify only a leading subset of offset/sizes/strides in 1449 // which case we complete with offset=0, sizes from memref type and strides=1. 1450 unsigned rank = sourceMemRefType.getRank(); 1451 assert(leadingStaticOffsets.size() <= rank && 1452 "unexpected leadingStaticOffsets overflow"); 1453 assert(leadingStaticSizes.size() <= rank && 1454 "unexpected leadingStaticSizes overflow"); 1455 assert(leadingStaticStrides.size() <= rank && 1456 "unexpected leadingStaticStrides overflow"); 1457 auto staticOffsets = llvm::to_vector<4>(leadingStaticOffsets); 1458 auto staticSizes = llvm::to_vector<4>(leadingStaticSizes); 1459 auto staticStrides = llvm::to_vector<4>(leadingStaticStrides); 1460 unsigned numTrailingOffsets = rank - staticOffsets.size(); 1461 unsigned numTrailingSizes = rank - staticSizes.size(); 1462 unsigned numTrailingStrides = rank - staticStrides.size(); 1463 staticOffsets.append(numTrailingOffsets, 0); 1464 llvm::append_range(staticSizes, 1465 sourceMemRefType.getShape().take_back(numTrailingSizes)); 1466 staticStrides.append(numTrailingStrides, 1); 1467 1468 // Extract source offset and strides. 1469 int64_t sourceOffset; 1470 SmallVector<int64_t, 4> sourceStrides; 1471 auto res = getStridesAndOffset(sourceMemRefType, sourceStrides, sourceOffset); 1472 assert(succeeded(res) && "SubViewOp expected strided memref type"); 1473 (void)res; 1474 1475 // Compute target offset whose value is: 1476 // `sourceOffset + sum_i(staticOffset_i * sourceStrides_i)`. 1477 int64_t targetOffset = sourceOffset; 1478 for (auto it : llvm::zip(staticOffsets, sourceStrides)) { 1479 auto staticOffset = std::get<0>(it), targetStride = std::get<1>(it); 1480 using namespace saturated_arith; 1481 targetOffset = Wrapper(targetOffset) + Wrapper(staticOffset) * targetStride; 1482 } 1483 1484 // Compute target stride whose value is: 1485 // `sourceStrides_i * staticStrides_i`. 1486 SmallVector<int64_t, 4> targetStrides; 1487 targetStrides.reserve(staticOffsets.size()); 1488 for (auto it : llvm::zip(sourceStrides, staticStrides)) { 1489 auto sourceStride = std::get<0>(it), staticStride = std::get<1>(it); 1490 using namespace saturated_arith; 1491 targetStrides.push_back(Wrapper(sourceStride) * staticStride); 1492 } 1493 1494 // The type is now known. 1495 return MemRefType::get( 1496 staticSizes, sourceMemRefType.getElementType(), 1497 makeStridedLinearLayoutMap(targetStrides, targetOffset, 1498 sourceMemRefType.getContext()), 1499 sourceMemRefType.getMemorySpace()); 1500 } 1501 1502 Type SubViewOp::inferResultType(MemRefType sourceMemRefType, 1503 ArrayRef<OpFoldResult> leadingStaticOffsets, 1504 ArrayRef<OpFoldResult> leadingStaticSizes, 1505 ArrayRef<OpFoldResult> leadingStaticStrides) { 1506 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides; 1507 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides; 1508 dispatchIndexOpFoldResults(leadingStaticOffsets, dynamicOffsets, 1509 staticOffsets, ShapedType::kDynamicStrideOrOffset); 1510 dispatchIndexOpFoldResults(leadingStaticSizes, dynamicSizes, staticSizes, 1511 ShapedType::kDynamicSize); 1512 dispatchIndexOpFoldResults(leadingStaticStrides, dynamicStrides, 1513 staticStrides, ShapedType::kDynamicStrideOrOffset); 1514 return SubViewOp::inferResultType(sourceMemRefType, staticOffsets, 1515 staticSizes, staticStrides) 1516 .cast<MemRefType>(); 1517 } 1518 1519 Type SubViewOp::inferRankReducedResultType( 1520 unsigned resultRank, MemRefType sourceRankedTensorType, 1521 ArrayRef<int64_t> leadingStaticOffsets, 1522 ArrayRef<int64_t> leadingStaticSizes, 1523 ArrayRef<int64_t> leadingStaticStrides) { 1524 auto inferredType = 1525 inferResultType(sourceRankedTensorType, leadingStaticOffsets, 1526 leadingStaticSizes, leadingStaticStrides) 1527 .cast<MemRefType>(); 1528 assert(inferredType.getRank() >= resultRank && "expected "); 1529 int rankDiff = inferredType.getRank() - resultRank; 1530 if (rankDiff > 0) { 1531 auto shape = inferredType.getShape(); 1532 llvm::SmallDenseSet<unsigned> dimsToProject; 1533 mlir::getPositionsOfShapeOne(rankDiff, shape, dimsToProject); 1534 SmallVector<int64_t> projectedShape; 1535 for (unsigned pos = 0, e = shape.size(); pos < e; ++pos) 1536 if (!dimsToProject.contains(pos)) 1537 projectedShape.push_back(shape[pos]); 1538 1539 AffineMap map; 1540 auto maps = inferredType.getAffineMaps(); 1541 if (!maps.empty() && maps.front()) 1542 map = getProjectedMap(maps.front(), dimsToProject); 1543 inferredType = 1544 MemRefType::get(projectedShape, inferredType.getElementType(), map, 1545 inferredType.getMemorySpace()); 1546 } 1547 return inferredType; 1548 } 1549 1550 Type SubViewOp::inferRankReducedResultType( 1551 unsigned resultRank, MemRefType sourceRankedTensorType, 1552 ArrayRef<OpFoldResult> leadingStaticOffsets, 1553 ArrayRef<OpFoldResult> leadingStaticSizes, 1554 ArrayRef<OpFoldResult> leadingStaticStrides) { 1555 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides; 1556 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides; 1557 dispatchIndexOpFoldResults(leadingStaticOffsets, dynamicOffsets, 1558 staticOffsets, ShapedType::kDynamicStrideOrOffset); 1559 dispatchIndexOpFoldResults(leadingStaticSizes, dynamicSizes, staticSizes, 1560 ShapedType::kDynamicSize); 1561 dispatchIndexOpFoldResults(leadingStaticStrides, dynamicStrides, 1562 staticStrides, ShapedType::kDynamicStrideOrOffset); 1563 return SubViewOp::inferRankReducedResultType( 1564 resultRank, sourceRankedTensorType, staticOffsets, staticSizes, 1565 staticStrides); 1566 } 1567 // Build a SubViewOp with mixed static and dynamic entries and custom result 1568 // type. If the type passed is nullptr, it is inferred. 1569 void SubViewOp::build(OpBuilder &b, OperationState &result, 1570 MemRefType resultType, Value source, 1571 ArrayRef<OpFoldResult> offsets, 1572 ArrayRef<OpFoldResult> sizes, 1573 ArrayRef<OpFoldResult> strides, 1574 ArrayRef<NamedAttribute> attrs) { 1575 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides; 1576 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides; 1577 dispatchIndexOpFoldResults(offsets, dynamicOffsets, staticOffsets, 1578 ShapedType::kDynamicStrideOrOffset); 1579 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes, 1580 ShapedType::kDynamicSize); 1581 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides, 1582 ShapedType::kDynamicStrideOrOffset); 1583 auto sourceMemRefType = source.getType().cast<MemRefType>(); 1584 // Structuring implementation this way avoids duplication between builders. 1585 if (!resultType) { 1586 resultType = SubViewOp::inferResultType(sourceMemRefType, staticOffsets, 1587 staticSizes, staticStrides) 1588 .cast<MemRefType>(); 1589 } 1590 build(b, result, resultType, source, dynamicOffsets, dynamicSizes, 1591 dynamicStrides, b.getI64ArrayAttr(staticOffsets), 1592 b.getI64ArrayAttr(staticSizes), b.getI64ArrayAttr(staticStrides)); 1593 result.addAttributes(attrs); 1594 } 1595 1596 // Build a SubViewOp with mixed static and dynamic entries and inferred result 1597 // type. 1598 void SubViewOp::build(OpBuilder &b, OperationState &result, Value source, 1599 ArrayRef<OpFoldResult> offsets, 1600 ArrayRef<OpFoldResult> sizes, 1601 ArrayRef<OpFoldResult> strides, 1602 ArrayRef<NamedAttribute> attrs) { 1603 build(b, result, MemRefType(), source, offsets, sizes, strides, attrs); 1604 } 1605 1606 // Build a SubViewOp with static entries and inferred result type. 1607 void SubViewOp::build(OpBuilder &b, OperationState &result, Value source, 1608 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes, 1609 ArrayRef<int64_t> strides, 1610 ArrayRef<NamedAttribute> attrs) { 1611 SmallVector<OpFoldResult> offsetValues = llvm::to_vector<4>( 1612 llvm::map_range(offsets, [&](int64_t v) -> OpFoldResult { 1613 return b.getI64IntegerAttr(v); 1614 })); 1615 SmallVector<OpFoldResult> sizeValues = 1616 llvm::to_vector<4>(llvm::map_range(sizes, [&](int64_t v) -> OpFoldResult { 1617 return b.getI64IntegerAttr(v); 1618 })); 1619 SmallVector<OpFoldResult> strideValues = llvm::to_vector<4>( 1620 llvm::map_range(strides, [&](int64_t v) -> OpFoldResult { 1621 return b.getI64IntegerAttr(v); 1622 })); 1623 build(b, result, source, offsetValues, sizeValues, strideValues, attrs); 1624 } 1625 1626 // Build a SubViewOp with dynamic entries and custom result type. If the 1627 // type passed is nullptr, it is inferred. 1628 void SubViewOp::build(OpBuilder &b, OperationState &result, 1629 MemRefType resultType, Value source, 1630 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes, 1631 ArrayRef<int64_t> strides, 1632 ArrayRef<NamedAttribute> attrs) { 1633 SmallVector<OpFoldResult> offsetValues = llvm::to_vector<4>( 1634 llvm::map_range(offsets, [&](int64_t v) -> OpFoldResult { 1635 return b.getI64IntegerAttr(v); 1636 })); 1637 SmallVector<OpFoldResult> sizeValues = 1638 llvm::to_vector<4>(llvm::map_range(sizes, [&](int64_t v) -> OpFoldResult { 1639 return b.getI64IntegerAttr(v); 1640 })); 1641 SmallVector<OpFoldResult> strideValues = llvm::to_vector<4>( 1642 llvm::map_range(strides, [&](int64_t v) -> OpFoldResult { 1643 return b.getI64IntegerAttr(v); 1644 })); 1645 build(b, result, resultType, source, offsetValues, sizeValues, strideValues, 1646 attrs); 1647 } 1648 1649 // Build a SubViewOp with dynamic entries and custom result type. If the type 1650 // passed is nullptr, it is inferred. 1651 void SubViewOp::build(OpBuilder &b, OperationState &result, 1652 MemRefType resultType, Value source, ValueRange offsets, 1653 ValueRange sizes, ValueRange strides, 1654 ArrayRef<NamedAttribute> attrs) { 1655 SmallVector<OpFoldResult> offsetValues = llvm::to_vector<4>( 1656 llvm::map_range(offsets, [](Value v) -> OpFoldResult { return v; })); 1657 SmallVector<OpFoldResult> sizeValues = llvm::to_vector<4>( 1658 llvm::map_range(sizes, [](Value v) -> OpFoldResult { return v; })); 1659 SmallVector<OpFoldResult> strideValues = llvm::to_vector<4>( 1660 llvm::map_range(strides, [](Value v) -> OpFoldResult { return v; })); 1661 build(b, result, resultType, source, offsetValues, sizeValues, strideValues); 1662 } 1663 1664 // Build a SubViewOp with dynamic entries and inferred result type. 1665 void SubViewOp::build(OpBuilder &b, OperationState &result, Value source, 1666 ValueRange offsets, ValueRange sizes, ValueRange strides, 1667 ArrayRef<NamedAttribute> attrs) { 1668 build(b, result, MemRefType(), source, offsets, sizes, strides, attrs); 1669 } 1670 1671 /// For ViewLikeOpInterface. 1672 Value SubViewOp::getViewSource() { return source(); } 1673 1674 enum SubViewVerificationResult { 1675 Success, 1676 RankTooLarge, 1677 SizeMismatch, 1678 ElemTypeMismatch, 1679 MemSpaceMismatch, 1680 AffineMapMismatch 1681 }; 1682 1683 /// Checks if `original` Type type can be rank reduced to `reduced` type. 1684 /// This function is slight variant of `is subsequence` algorithm where 1685 /// not matching dimension must be 1. 1686 static SubViewVerificationResult 1687 isRankReducedType(Type originalType, Type candidateReducedType, 1688 std::string *errMsg = nullptr) { 1689 if (originalType == candidateReducedType) 1690 return SubViewVerificationResult::Success; 1691 if (!originalType.isa<MemRefType>()) 1692 return SubViewVerificationResult::Success; 1693 if (originalType.isa<MemRefType>() && !candidateReducedType.isa<MemRefType>()) 1694 return SubViewVerificationResult::Success; 1695 1696 ShapedType originalShapedType = originalType.cast<ShapedType>(); 1697 ShapedType candidateReducedShapedType = 1698 candidateReducedType.cast<ShapedType>(); 1699 1700 // Rank and size logic is valid for all ShapedTypes. 1701 ArrayRef<int64_t> originalShape = originalShapedType.getShape(); 1702 ArrayRef<int64_t> candidateReducedShape = 1703 candidateReducedShapedType.getShape(); 1704 unsigned originalRank = originalShape.size(), 1705 candidateReducedRank = candidateReducedShape.size(); 1706 if (candidateReducedRank > originalRank) 1707 return SubViewVerificationResult::RankTooLarge; 1708 1709 auto optionalUnusedDimsMask = 1710 computeRankReductionMask(originalShape, candidateReducedShape); 1711 1712 // Sizes cannot be matched in case empty vector is returned. 1713 if (!optionalUnusedDimsMask.hasValue()) 1714 return SubViewVerificationResult::SizeMismatch; 1715 1716 if (originalShapedType.getElementType() != 1717 candidateReducedShapedType.getElementType()) 1718 return SubViewVerificationResult::ElemTypeMismatch; 1719 1720 // Strided layout logic is relevant for MemRefType only. 1721 MemRefType original = originalType.cast<MemRefType>(); 1722 MemRefType candidateReduced = candidateReducedType.cast<MemRefType>(); 1723 if (original.getMemorySpace() != candidateReduced.getMemorySpace()) 1724 return SubViewVerificationResult::MemSpaceMismatch; 1725 1726 llvm::SmallDenseSet<unsigned> unusedDims = optionalUnusedDimsMask.getValue(); 1727 auto inferredType = 1728 getProjectedMap(getStridedLinearLayoutMap(original), unusedDims); 1729 AffineMap candidateLayout; 1730 if (candidateReduced.getAffineMaps().empty()) 1731 candidateLayout = getStridedLinearLayoutMap(candidateReduced); 1732 else 1733 candidateLayout = candidateReduced.getAffineMaps().front(); 1734 assert(inferredType.getNumResults() == 1 && 1735 candidateLayout.getNumResults() == 1); 1736 if (inferredType.getNumSymbols() != candidateLayout.getNumSymbols() || 1737 inferredType.getNumDims() != candidateLayout.getNumDims()) { 1738 if (errMsg) { 1739 llvm::raw_string_ostream os(*errMsg); 1740 os << "inferred type: " << inferredType; 1741 } 1742 return SubViewVerificationResult::AffineMapMismatch; 1743 } 1744 // Check that the difference of the affine maps simplifies to 0. 1745 AffineExpr diffExpr = 1746 inferredType.getResult(0) - candidateLayout.getResult(0); 1747 diffExpr = simplifyAffineExpr(diffExpr, inferredType.getNumDims(), 1748 inferredType.getNumSymbols()); 1749 auto cst = diffExpr.dyn_cast<AffineConstantExpr>(); 1750 if (!(cst && cst.getValue() == 0)) { 1751 if (errMsg) { 1752 llvm::raw_string_ostream os(*errMsg); 1753 os << "inferred type: " << inferredType; 1754 } 1755 return SubViewVerificationResult::AffineMapMismatch; 1756 } 1757 return SubViewVerificationResult::Success; 1758 } 1759 1760 template <typename OpTy> 1761 static LogicalResult produceSubViewErrorMsg(SubViewVerificationResult result, 1762 OpTy op, Type expectedType, 1763 StringRef errMsg = "") { 1764 auto memrefType = expectedType.cast<ShapedType>(); 1765 switch (result) { 1766 case SubViewVerificationResult::Success: 1767 return success(); 1768 case SubViewVerificationResult::RankTooLarge: 1769 return op.emitError("expected result rank to be smaller or equal to ") 1770 << "the source rank. " << errMsg; 1771 case SubViewVerificationResult::SizeMismatch: 1772 return op.emitError("expected result type to be ") 1773 << expectedType 1774 << " or a rank-reduced version. (mismatch of result sizes) " 1775 << errMsg; 1776 case SubViewVerificationResult::ElemTypeMismatch: 1777 return op.emitError("expected result element type to be ") 1778 << memrefType.getElementType() << errMsg; 1779 case SubViewVerificationResult::MemSpaceMismatch: 1780 return op.emitError("expected result and source memory spaces to match.") 1781 << errMsg; 1782 case SubViewVerificationResult::AffineMapMismatch: 1783 return op.emitError("expected result type to be ") 1784 << expectedType 1785 << " or a rank-reduced version. (mismatch of result affine map) " 1786 << errMsg; 1787 } 1788 llvm_unreachable("unexpected subview verification result"); 1789 } 1790 1791 /// Verifier for SubViewOp. 1792 static LogicalResult verify(SubViewOp op) { 1793 MemRefType baseType = op.getSourceType(); 1794 MemRefType subViewType = op.getType(); 1795 1796 // The base memref and the view memref should be in the same memory space. 1797 if (baseType.getMemorySpace() != subViewType.getMemorySpace()) 1798 return op.emitError("different memory spaces specified for base memref " 1799 "type ") 1800 << baseType << " and subview memref type " << subViewType; 1801 1802 // Verify that the base memref type has a strided layout map. 1803 if (!isStrided(baseType)) 1804 return op.emitError("base type ") << baseType << " is not strided"; 1805 1806 // Verify result type against inferred type. 1807 auto expectedType = SubViewOp::inferResultType( 1808 baseType, extractFromI64ArrayAttr(op.static_offsets()), 1809 extractFromI64ArrayAttr(op.static_sizes()), 1810 extractFromI64ArrayAttr(op.static_strides())); 1811 1812 std::string errMsg; 1813 auto result = isRankReducedType(expectedType, subViewType, &errMsg); 1814 return produceSubViewErrorMsg(result, op, expectedType, errMsg); 1815 } 1816 1817 raw_ostream &mlir::operator<<(raw_ostream &os, Range &range) { 1818 return os << "range " << range.offset << ":" << range.size << ":" 1819 << range.stride; 1820 } 1821 1822 /// Return the list of Range (i.e. offset, size, stride). Each Range 1823 /// entry contains either the dynamic value or a ConstantIndexOp constructed 1824 /// with `b` at location `loc`. 1825 SmallVector<Range, 8> mlir::getOrCreateRanges(OffsetSizeAndStrideOpInterface op, 1826 OpBuilder &b, Location loc) { 1827 std::array<unsigned, 3> ranks = op.getArrayAttrMaxRanks(); 1828 assert(ranks[0] == ranks[1] && "expected offset and sizes of equal ranks"); 1829 assert(ranks[1] == ranks[2] && "expected sizes and strides of equal ranks"); 1830 SmallVector<Range, 8> res; 1831 unsigned rank = ranks[0]; 1832 res.reserve(rank); 1833 for (unsigned idx = 0; idx < rank; ++idx) { 1834 Value offset = 1835 op.isDynamicOffset(idx) 1836 ? op.getDynamicOffset(idx) 1837 : b.create<ConstantIndexOp>(loc, op.getStaticOffset(idx)); 1838 Value size = op.isDynamicSize(idx) 1839 ? op.getDynamicSize(idx) 1840 : b.create<ConstantIndexOp>(loc, op.getStaticSize(idx)); 1841 Value stride = 1842 op.isDynamicStride(idx) 1843 ? op.getDynamicStride(idx) 1844 : b.create<ConstantIndexOp>(loc, op.getStaticStride(idx)); 1845 res.emplace_back(Range{offset, size, stride}); 1846 } 1847 return res; 1848 } 1849 1850 /// Infer the canonical type of the result of a subview operation. Returns a 1851 /// type with rank `resultRank` that is either the rank of the rank-reduced 1852 /// type, or the non-rank-reduced type. 1853 static MemRefType 1854 getCanonicalSubViewResultType(unsigned resultRank, MemRefType sourceType, 1855 ArrayRef<OpFoldResult> mixedOffsets, 1856 ArrayRef<OpFoldResult> mixedSizes, 1857 ArrayRef<OpFoldResult> mixedStrides) { 1858 auto resultType = 1859 SubViewOp::inferRankReducedResultType( 1860 resultRank, sourceType, mixedOffsets, mixedSizes, mixedStrides) 1861 .cast<MemRefType>(); 1862 if (resultType.getRank() != resultRank) { 1863 resultType = SubViewOp::inferResultType(sourceType, mixedOffsets, 1864 mixedSizes, mixedStrides) 1865 .cast<MemRefType>(); 1866 } 1867 return resultType; 1868 } 1869 1870 namespace { 1871 /// Pattern to rewrite a subview op with MemRefCast arguments. 1872 /// This essentially pushes memref.cast past its consuming subview when 1873 /// `canFoldIntoConsumerOp` is true. 1874 /// 1875 /// Example: 1876 /// ``` 1877 /// %0 = memref.cast %V : memref<16x16xf32> to memref<?x?xf32> 1878 /// %1 = memref.subview %0[0, 0][3, 4][1, 1] : 1879 /// memref<?x?xf32> to memref<3x4xf32, offset:?, strides:[?, 1]> 1880 /// ``` 1881 /// is rewritten into: 1882 /// ``` 1883 /// %0 = memref.subview %V: memref<16x16xf32> to memref<3x4xf32, #[[map0]]> 1884 /// %1 = memref.cast %0: memref<3x4xf32, offset:0, strides:[16, 1]> to 1885 /// memref<3x4xf32, offset:?, strides:[?, 1]> 1886 /// ``` 1887 class SubViewOpMemRefCastFolder final : public OpRewritePattern<SubViewOp> { 1888 public: 1889 using OpRewritePattern<SubViewOp>::OpRewritePattern; 1890 1891 LogicalResult matchAndRewrite(SubViewOp subViewOp, 1892 PatternRewriter &rewriter) const override { 1893 // Any constant operand, just return to let SubViewOpConstantFolder kick in. 1894 if (llvm::any_of(subViewOp.getOperands(), [](Value operand) { 1895 return matchPattern(operand, matchConstantIndex()); 1896 })) 1897 return failure(); 1898 1899 auto castOp = subViewOp.source().getDefiningOp<CastOp>(); 1900 if (!castOp) 1901 return failure(); 1902 1903 if (!CastOp::canFoldIntoConsumerOp(castOp)) 1904 return failure(); 1905 1906 /// Deduce the resultType of the SubViewOp using `inferSubViewResultType` on 1907 /// the cast source operand type and the SubViewOp static information. This 1908 /// is the resulting type if the MemRefCastOp were folded. 1909 auto resultType = getCanonicalSubViewResultType( 1910 subViewOp.getType().getRank(), 1911 castOp.source().getType().cast<MemRefType>(), 1912 subViewOp.getMixedOffsets(), subViewOp.getMixedSizes(), 1913 subViewOp.getMixedStrides()); 1914 Value newSubView = rewriter.create<SubViewOp>( 1915 subViewOp.getLoc(), resultType, castOp.source(), subViewOp.offsets(), 1916 subViewOp.sizes(), subViewOp.strides(), subViewOp.static_offsets(), 1917 subViewOp.static_sizes(), subViewOp.static_strides()); 1918 rewriter.replaceOpWithNewOp<CastOp>(subViewOp, subViewOp.getType(), 1919 newSubView); 1920 return success(); 1921 } 1922 }; 1923 } // namespace 1924 1925 /// Return the canonical type of the result of a subview. 1926 struct SubViewReturnTypeCanonicalizer { 1927 MemRefType operator()(SubViewOp op, ArrayRef<OpFoldResult> mixedOffsets, 1928 ArrayRef<OpFoldResult> mixedSizes, 1929 ArrayRef<OpFoldResult> mixedStrides) { 1930 return getCanonicalSubViewResultType(op.getType().getRank(), 1931 op.getSourceType(), mixedOffsets, 1932 mixedSizes, mixedStrides); 1933 } 1934 }; 1935 1936 /// A canonicalizer wrapper to replace SubViewOps. 1937 struct SubViewCanonicalizer { 1938 void operator()(PatternRewriter &rewriter, SubViewOp op, SubViewOp newOp) { 1939 rewriter.replaceOpWithNewOp<CastOp>(op, newOp, op.getType()); 1940 } 1941 }; 1942 1943 void SubViewOp::getCanonicalizationPatterns(RewritePatternSet &results, 1944 MLIRContext *context) { 1945 results 1946 .add<OpWithOffsetSizesAndStridesConstantArgumentFolder< 1947 SubViewOp, SubViewReturnTypeCanonicalizer, SubViewCanonicalizer>, 1948 SubViewOpMemRefCastFolder>(context); 1949 } 1950 1951 OpFoldResult SubViewOp::fold(ArrayRef<Attribute> operands) { 1952 auto resultShapedType = getResult().getType().cast<ShapedType>(); 1953 auto sourceShapedType = source().getType().cast<ShapedType>(); 1954 1955 if (resultShapedType.hasStaticShape() && 1956 resultShapedType == sourceShapedType) { 1957 return getViewSource(); 1958 } 1959 1960 return {}; 1961 } 1962 1963 //===----------------------------------------------------------------------===// 1964 // TensorLoadOp 1965 //===----------------------------------------------------------------------===// 1966 1967 OpFoldResult TensorLoadOp::fold(ArrayRef<Attribute>) { 1968 if (auto bufferCast = memref().getDefiningOp<BufferCastOp>()) 1969 // Approximate alias analysis by conservatively folding only when no there 1970 // is no interleaved operation. 1971 if (bufferCast->getBlock() == this->getOperation()->getBlock() && 1972 bufferCast->getNextNode() == this->getOperation()) 1973 return bufferCast.tensor(); 1974 return {}; 1975 } 1976 1977 //===----------------------------------------------------------------------===// 1978 // TransposeOp 1979 //===----------------------------------------------------------------------===// 1980 1981 /// Build a strided memref type by applying `permutationMap` tp `memRefType`. 1982 static MemRefType inferTransposeResultType(MemRefType memRefType, 1983 AffineMap permutationMap) { 1984 auto rank = memRefType.getRank(); 1985 auto originalSizes = memRefType.getShape(); 1986 // Compute permuted sizes. 1987 SmallVector<int64_t, 4> sizes(rank, 0); 1988 for (auto en : llvm::enumerate(permutationMap.getResults())) 1989 sizes[en.index()] = 1990 originalSizes[en.value().cast<AffineDimExpr>().getPosition()]; 1991 1992 // Compute permuted strides. 1993 int64_t offset; 1994 SmallVector<int64_t, 4> strides; 1995 auto res = getStridesAndOffset(memRefType, strides, offset); 1996 assert(succeeded(res) && strides.size() == static_cast<unsigned>(rank)); 1997 (void)res; 1998 auto map = 1999 makeStridedLinearLayoutMap(strides, offset, memRefType.getContext()); 2000 map = permutationMap ? map.compose(permutationMap) : map; 2001 return MemRefType::Builder(memRefType).setShape(sizes).setAffineMaps(map); 2002 } 2003 2004 void TransposeOp::build(OpBuilder &b, OperationState &result, Value in, 2005 AffineMapAttr permutation, 2006 ArrayRef<NamedAttribute> attrs) { 2007 auto permutationMap = permutation.getValue(); 2008 assert(permutationMap); 2009 2010 auto memRefType = in.getType().cast<MemRefType>(); 2011 // Compute result type. 2012 MemRefType resultType = inferTransposeResultType(memRefType, permutationMap); 2013 2014 build(b, result, resultType, in, attrs); 2015 result.addAttribute(TransposeOp::getPermutationAttrName(), permutation); 2016 } 2017 2018 // transpose $in $permutation attr-dict : type($in) `to` type(results) 2019 static void print(OpAsmPrinter &p, TransposeOp op) { 2020 p << "memref.transpose " << op.in() << " " << op.permutation(); 2021 p.printOptionalAttrDict(op->getAttrs(), 2022 {TransposeOp::getPermutationAttrName()}); 2023 p << " : " << op.in().getType() << " to " << op.getType(); 2024 } 2025 2026 static ParseResult parseTransposeOp(OpAsmParser &parser, 2027 OperationState &result) { 2028 OpAsmParser::OperandType in; 2029 AffineMap permutation; 2030 MemRefType srcType, dstType; 2031 if (parser.parseOperand(in) || parser.parseAffineMap(permutation) || 2032 parser.parseOptionalAttrDict(result.attributes) || 2033 parser.parseColonType(srcType) || 2034 parser.resolveOperand(in, srcType, result.operands) || 2035 parser.parseKeywordType("to", dstType) || 2036 parser.addTypeToList(dstType, result.types)) 2037 return failure(); 2038 2039 result.addAttribute(TransposeOp::getPermutationAttrName(), 2040 AffineMapAttr::get(permutation)); 2041 return success(); 2042 } 2043 2044 static LogicalResult verify(TransposeOp op) { 2045 if (!op.permutation().isPermutation()) 2046 return op.emitOpError("expected a permutation map"); 2047 if (op.permutation().getNumDims() != op.getShapedType().getRank()) 2048 return op.emitOpError( 2049 "expected a permutation map of same rank as the input"); 2050 2051 auto srcType = op.in().getType().cast<MemRefType>(); 2052 auto dstType = op.getType().cast<MemRefType>(); 2053 auto transposedType = inferTransposeResultType(srcType, op.permutation()); 2054 if (dstType != transposedType) 2055 return op.emitOpError("output type ") 2056 << dstType << " does not match transposed input type " << srcType 2057 << ", " << transposedType; 2058 return success(); 2059 } 2060 2061 OpFoldResult TransposeOp::fold(ArrayRef<Attribute>) { 2062 if (succeeded(foldMemRefCast(*this))) 2063 return getResult(); 2064 return {}; 2065 } 2066 2067 //===----------------------------------------------------------------------===// 2068 // ViewOp 2069 //===----------------------------------------------------------------------===// 2070 2071 static ParseResult parseViewOp(OpAsmParser &parser, OperationState &result) { 2072 OpAsmParser::OperandType srcInfo; 2073 SmallVector<OpAsmParser::OperandType, 1> offsetInfo; 2074 SmallVector<OpAsmParser::OperandType, 4> sizesInfo; 2075 auto indexType = parser.getBuilder().getIndexType(); 2076 Type srcType, dstType; 2077 llvm::SMLoc offsetLoc; 2078 if (parser.parseOperand(srcInfo) || parser.getCurrentLocation(&offsetLoc) || 2079 parser.parseOperandList(offsetInfo, OpAsmParser::Delimiter::Square)) 2080 return failure(); 2081 2082 if (offsetInfo.size() != 1) 2083 return parser.emitError(offsetLoc) << "expects 1 offset operand"; 2084 2085 return failure( 2086 parser.parseOperandList(sizesInfo, OpAsmParser::Delimiter::Square) || 2087 parser.parseOptionalAttrDict(result.attributes) || 2088 parser.parseColonType(srcType) || 2089 parser.resolveOperand(srcInfo, srcType, result.operands) || 2090 parser.resolveOperands(offsetInfo, indexType, result.operands) || 2091 parser.resolveOperands(sizesInfo, indexType, result.operands) || 2092 parser.parseKeywordType("to", dstType) || 2093 parser.addTypeToList(dstType, result.types)); 2094 } 2095 2096 static void print(OpAsmPrinter &p, ViewOp op) { 2097 p << op.getOperationName() << ' ' << op.getOperand(0) << '['; 2098 p.printOperand(op.byte_shift()); 2099 p << "][" << op.sizes() << ']'; 2100 p.printOptionalAttrDict(op->getAttrs()); 2101 p << " : " << op.getOperand(0).getType() << " to " << op.getType(); 2102 } 2103 2104 static LogicalResult verify(ViewOp op) { 2105 auto baseType = op.getOperand(0).getType().cast<MemRefType>(); 2106 auto viewType = op.getType(); 2107 2108 // The base memref should have identity layout map (or none). 2109 if (baseType.getAffineMaps().size() > 1 || 2110 (baseType.getAffineMaps().size() == 1 && 2111 !baseType.getAffineMaps()[0].isIdentity())) 2112 return op.emitError("unsupported map for base memref type ") << baseType; 2113 2114 // The result memref should have identity layout map (or none). 2115 if (viewType.getAffineMaps().size() > 1 || 2116 (viewType.getAffineMaps().size() == 1 && 2117 !viewType.getAffineMaps()[0].isIdentity())) 2118 return op.emitError("unsupported map for result memref type ") << viewType; 2119 2120 // The base memref and the view memref should be in the same memory space. 2121 if (baseType.getMemorySpace() != viewType.getMemorySpace()) 2122 return op.emitError("different memory spaces specified for base memref " 2123 "type ") 2124 << baseType << " and view memref type " << viewType; 2125 2126 // Verify that we have the correct number of sizes for the result type. 2127 unsigned numDynamicDims = viewType.getNumDynamicDims(); 2128 if (op.sizes().size() != numDynamicDims) 2129 return op.emitError("incorrect number of size operands for type ") 2130 << viewType; 2131 2132 return success(); 2133 } 2134 2135 Value ViewOp::getViewSource() { return source(); } 2136 2137 namespace { 2138 2139 struct ViewOpShapeFolder : public OpRewritePattern<ViewOp> { 2140 using OpRewritePattern<ViewOp>::OpRewritePattern; 2141 2142 LogicalResult matchAndRewrite(ViewOp viewOp, 2143 PatternRewriter &rewriter) const override { 2144 // Return if none of the operands are constants. 2145 if (llvm::none_of(viewOp.getOperands(), [](Value operand) { 2146 return matchPattern(operand, matchConstantIndex()); 2147 })) 2148 return failure(); 2149 2150 // Get result memref type. 2151 auto memrefType = viewOp.getType(); 2152 2153 // Get offset from old memref view type 'memRefType'. 2154 int64_t oldOffset; 2155 SmallVector<int64_t, 4> oldStrides; 2156 if (failed(getStridesAndOffset(memrefType, oldStrides, oldOffset))) 2157 return failure(); 2158 assert(oldOffset == 0 && "Expected 0 offset"); 2159 2160 SmallVector<Value, 4> newOperands; 2161 2162 // Offset cannot be folded into result type. 2163 2164 // Fold any dynamic dim operands which are produced by a constant. 2165 SmallVector<int64_t, 4> newShapeConstants; 2166 newShapeConstants.reserve(memrefType.getRank()); 2167 2168 unsigned dynamicDimPos = 0; 2169 unsigned rank = memrefType.getRank(); 2170 for (unsigned dim = 0, e = rank; dim < e; ++dim) { 2171 int64_t dimSize = memrefType.getDimSize(dim); 2172 // If this is already static dimension, keep it. 2173 if (!ShapedType::isDynamic(dimSize)) { 2174 newShapeConstants.push_back(dimSize); 2175 continue; 2176 } 2177 auto *defOp = viewOp.sizes()[dynamicDimPos].getDefiningOp(); 2178 if (auto constantIndexOp = dyn_cast_or_null<ConstantIndexOp>(defOp)) { 2179 // Dynamic shape dimension will be folded. 2180 newShapeConstants.push_back(constantIndexOp.getValue()); 2181 } else { 2182 // Dynamic shape dimension not folded; copy operand from old memref. 2183 newShapeConstants.push_back(dimSize); 2184 newOperands.push_back(viewOp.sizes()[dynamicDimPos]); 2185 } 2186 dynamicDimPos++; 2187 } 2188 2189 // Create new memref type with constant folded dims. 2190 MemRefType newMemRefType = 2191 MemRefType::Builder(memrefType).setShape(newShapeConstants); 2192 // Nothing new, don't fold. 2193 if (newMemRefType == memrefType) 2194 return failure(); 2195 2196 // Create new ViewOp. 2197 auto newViewOp = rewriter.create<ViewOp>(viewOp.getLoc(), newMemRefType, 2198 viewOp.getOperand(0), 2199 viewOp.byte_shift(), newOperands); 2200 // Insert a cast so we have the same type as the old memref type. 2201 rewriter.replaceOpWithNewOp<CastOp>(viewOp, newViewOp, viewOp.getType()); 2202 return success(); 2203 } 2204 }; 2205 2206 struct ViewOpMemrefCastFolder : public OpRewritePattern<ViewOp> { 2207 using OpRewritePattern<ViewOp>::OpRewritePattern; 2208 2209 LogicalResult matchAndRewrite(ViewOp viewOp, 2210 PatternRewriter &rewriter) const override { 2211 Value memrefOperand = viewOp.getOperand(0); 2212 CastOp memrefCastOp = memrefOperand.getDefiningOp<CastOp>(); 2213 if (!memrefCastOp) 2214 return failure(); 2215 Value allocOperand = memrefCastOp.getOperand(); 2216 AllocOp allocOp = allocOperand.getDefiningOp<AllocOp>(); 2217 if (!allocOp) 2218 return failure(); 2219 rewriter.replaceOpWithNewOp<ViewOp>(viewOp, viewOp.getType(), allocOperand, 2220 viewOp.byte_shift(), viewOp.sizes()); 2221 return success(); 2222 } 2223 }; 2224 2225 } // end anonymous namespace 2226 2227 void ViewOp::getCanonicalizationPatterns(RewritePatternSet &results, 2228 MLIRContext *context) { 2229 results.add<ViewOpShapeFolder, ViewOpMemrefCastFolder>(context); 2230 } 2231 2232 //===----------------------------------------------------------------------===// 2233 // TableGen'd op method definitions 2234 //===----------------------------------------------------------------------===// 2235 2236 #define GET_OP_CLASSES 2237 #include "mlir/Dialect/MemRef/IR/MemRefOps.cpp.inc" 2238