1 //===- BufferizableOpInterface.cpp - Bufferizable Ops ---=----------------===// 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/Bufferization/IR/BufferizableOpInterface.h" 10 #include "mlir/Dialect/Bufferization/IR/Bufferization.h" 11 #include "mlir/Dialect/Func/IR/FuncOps.h" 12 #include "mlir/Dialect/MemRef/IR/MemRef.h" 13 #include "mlir/Dialect/Tensor/IR/Tensor.h" 14 #include "mlir/IR/AsmState.h" 15 #include "mlir/IR/BlockAndValueMapping.h" 16 #include "mlir/IR/BuiltinOps.h" 17 #include "mlir/IR/Operation.h" 18 #include "mlir/IR/TypeUtilities.h" 19 #include "mlir/IR/Value.h" 20 #include "llvm/Support/Debug.h" 21 22 //===----------------------------------------------------------------------===// 23 // BufferizableOpInterface 24 //===----------------------------------------------------------------------===// 25 26 namespace mlir { 27 namespace bufferization { 28 29 #include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.cpp.inc" 30 31 } // namespace bufferization 32 } // namespace mlir 33 34 #define DEBUG_TYPE "bufferizable-op-interface" 35 #define DBGS() (llvm::dbgs() << '[' << DEBUG_TYPE << "] ") 36 #define LDBG(X) LLVM_DEBUG(DBGS() << (X)) 37 38 using namespace mlir; 39 using namespace bufferization; 40 41 /// Attribute name used to mark region arguments that can be bufferized 42 /// in-place during linalg comprehensive bufferization. 43 constexpr const ::llvm::StringLiteral 44 bufferization::BufferizableOpInterface::kInplaceableAttrName; 45 46 /// Return the owner of the given value. 47 static Operation *getOwnerOfValue(Value value) { 48 if (auto opResult = value.dyn_cast<OpResult>()) 49 return opResult.getDefiningOp(); 50 return value.cast<BlockArgument>().getOwner()->getParentOp(); 51 } 52 53 bool bufferization::allocationDoesNotEscape(OpResult opResult) { 54 #ifndef NDEBUG 55 auto bufferizableOp = opResult.getDefiningOp<BufferizableOpInterface>(); 56 assert(bufferizableOp && bufferizableOp.bufferizesToAllocation(opResult) && 57 "expected op that bufferizes to an allocation"); 58 #endif // NDEBUG 59 60 Operation *op = opResult.getDefiningOp(); 61 // If there is no 'escape' attribute, we cannot say for sure. 62 if (!op->hasAttr(BufferizationDialect::kEscapeAttrName)) 63 return false; 64 auto attr = 65 op->getAttrOfType<ArrayAttr>(BufferizationDialect::kEscapeAttrName); 66 return !attr[opResult.getResultNumber()].cast<BoolAttr>().getValue(); 67 } 68 69 /// Create an AllocTensorOp for the given shaped value. If `copy` is set, the 70 /// shaped value is copied. Otherwise, a tensor with undefined contents is 71 /// allocated. 72 FailureOr<Value> bufferization::allocateTensorForShapedValue( 73 OpBuilder &b, Location loc, Value shapedValue, bool escape, 74 const BufferizationOptions &options, bool copy) { 75 Value tensor; 76 if (shapedValue.getType().isa<RankedTensorType>()) { 77 tensor = shapedValue; 78 } else if (shapedValue.getType().isa<MemRefType>()) { 79 tensor = b.create<ToTensorOp>(loc, shapedValue); 80 } else { 81 llvm_unreachable("expected RankedTensorType or MemRefType"); 82 } 83 RankedTensorType tensorType = tensor.getType().cast<RankedTensorType>(); 84 SmallVector<Value> dynamicSizes; 85 if (!copy) { 86 // Compute the dynamic part of the shape. 87 // First try to query the shape via ReifyRankedShapedTypeOpInterface. 88 bool reifiedShapes = false; 89 if (shapedValue.getType().isa<RankedTensorType>() && 90 shapedValue.isa<OpResult>()) { 91 if (auto rankedOp = dyn_cast_or_null<ReifyRankedShapedTypeOpInterface>( 92 shapedValue.getDefiningOp())) { 93 ReifiedRankedShapedTypeDims resultDims; 94 if (succeeded(rankedOp.reifyResultShapes(b, resultDims))) { 95 reifiedShapes = true; 96 auto &shape = 97 resultDims[shapedValue.cast<OpResult>().getResultNumber()]; 98 for (const auto &dim : enumerate(tensorType.getShape())) 99 if (ShapedType::isDynamic(dim.value())) 100 dynamicSizes.push_back(shape[dim.index()]); 101 } 102 } 103 } 104 105 // If the shape could not be reified, create DimOps. 106 if (!reifiedShapes) 107 populateDynamicDimSizes(b, loc, tensor, dynamicSizes); 108 } 109 110 // Create AllocTensorOp. 111 auto allocTensorOp = b.create<AllocTensorOp>(loc, tensorType, dynamicSizes, 112 copy ? tensor : Value()); 113 allocTensorOp->setAttr(BufferizationDialect::kEscapeAttrName, 114 b.getBoolArrayAttr({escape})); 115 116 // Add 'memory_space' attribute. Not needed if 'copy' operand is specified. 117 if (copy) 118 return allocTensorOp.getResult(); 119 FailureOr<BaseMemRefType> copyBufferType = getBufferType(tensor, options); 120 if (failed(copyBufferType)) 121 return failure(); 122 allocTensorOp.setMemorySpaceAttr( 123 b.getIntegerAttr(b.getIntegerType(64, /*isSigned=*/false), 124 copyBufferType->getMemorySpaceAsInt())); 125 return allocTensorOp.getResult(); 126 } 127 128 LogicalResult BufferizableOpInterface::resolveTensorOpOperandConflicts( 129 RewriterBase &rewriter, const AnalysisState &state) { 130 OpBuilder::InsertionGuard g(rewriter); 131 Operation *op = getOperation(); 132 SmallVector<OpOperand *> outOfPlaceOpOperands; 133 DenseSet<OpOperand *> copiedOpOperands; 134 DenseSet<OpOperand *> escapingOpOperandCopies; 135 SmallVector<OpResult> outOfPlaceOpResults; 136 DenseSet<OpResult> copiedOpResults; 137 DenseSet<OpResult> escapingOpResultCopies; 138 139 // Find all out-of-place OpOperands. 140 for (OpOperand &opOperand : op->getOpOperands()) { 141 Type operandType = opOperand.get().getType(); 142 if (!operandType.isa<TensorType>()) 143 continue; 144 if (state.isInPlace(opOperand)) 145 continue; 146 if (operandType.isa<UnrankedTensorType>()) 147 return op->emitError("copies of unranked tensors are not supported"); 148 149 SmallVector<OpResult> aliasingOpResults = 150 state.getAliasingOpResult(opOperand); 151 // Is the result yielded from a block? Or are deallocations turned off 152 // entirely? In either case, mark the allocation as "escaping", so that it 153 // will not be deallocated. 154 bool escape = !state.getOptions().createDeallocs || 155 llvm::any_of(aliasingOpResults, [&](Value v) { 156 return state.isTensorYielded(v); 157 }); 158 159 if (aliasingOpResults.size() == 1 && 160 !state.bufferizesToMemoryWrite(opOperand) && 161 state.getAliasingOpOperand(aliasingOpResults.front()).size() == 1) { 162 // The op itself does not write but may create exactly one alias. Instead 163 // of copying the OpOperand, copy the OpResult. The OpResult can sometimes 164 // be smaller than the OpOperand (e.g., in the case of an extract_slice, 165 // where the result is usually a smaller part of the source). 166 outOfPlaceOpResults.push_back(aliasingOpResults.front()); 167 if (!state.canOmitTensorCopy(opOperand)) 168 copiedOpResults.insert(aliasingOpResults.front()); 169 if (escape) 170 escapingOpResultCopies.insert(aliasingOpResults.front()); 171 } else { 172 // In all other cases, make a copy of the OpOperand. 173 outOfPlaceOpOperands.push_back(&opOperand); 174 if (!state.canOmitTensorCopy(opOperand)) 175 copiedOpOperands.insert(&opOperand); 176 if (escape) 177 escapingOpOperandCopies.insert(&opOperand); 178 } 179 } 180 181 // Insert copies of OpOperands. 182 rewriter.setInsertionPoint(op); 183 for (OpOperand *opOperand : outOfPlaceOpOperands) { 184 FailureOr<Value> copy = allocateTensorForShapedValue( 185 rewriter, op->getLoc(), opOperand->get(), 186 escapingOpOperandCopies.contains(opOperand), state.getOptions(), 187 copiedOpOperands.contains(opOperand)); 188 if (failed(copy)) 189 return failure(); 190 rewriter.updateRootInPlace(op, [&]() { opOperand->set(*copy); }); 191 } 192 193 // Insert copies of OpResults. 194 rewriter.setInsertionPointAfter(op); 195 for (OpResult opResult : outOfPlaceOpResults) { 196 FailureOr<Value> copy = allocateTensorForShapedValue( 197 rewriter, op->getLoc(), opResult, 198 escapingOpResultCopies.contains(opResult), state.getOptions(), 199 copiedOpResults.count(opResult)); 200 if (failed(copy)) 201 return failure(); 202 SmallVector<OpOperand *> uses = llvm::to_vector(llvm::map_range( 203 opResult.getUses(), [](OpOperand &use) { return &use; })); 204 for (OpOperand *use : uses) { 205 // Do not update the alloc_tensor op that we just created. 206 if (use->getOwner() != copy->getDefiningOp()) 207 rewriter.updateRootInPlace(use->getOwner(), [&]() { use->set(*copy); }); 208 } 209 } 210 211 return success(); 212 } 213 214 //===----------------------------------------------------------------------===// 215 // OpFilter 216 //===----------------------------------------------------------------------===// 217 218 bool OpFilter::isOpAllowed(Operation *op) const { 219 // All other ops: Allow/disallow according to filter. 220 bool isAllowed = !hasAllowRule(); 221 for (const Entry &entry : entries) { 222 bool filterResult = entry.fn(op); 223 switch (entry.type) { 224 case Entry::ALLOW: 225 isAllowed |= filterResult; 226 break; 227 case Entry::DENY: 228 if (filterResult) 229 // DENY filter matches. This op is no allowed. (Even if other ALLOW 230 // filters may match.) 231 return false; 232 }; 233 } 234 return isAllowed; 235 } 236 237 //===----------------------------------------------------------------------===// 238 // BufferizationOptions 239 //===----------------------------------------------------------------------===// 240 241 /// Default unknown type converter: Use a fully dynamic layout map. 242 static BaseMemRefType 243 defaultUnknownTypeConverter(Value value, unsigned memorySpace, 244 const BufferizationOptions &options) { 245 return getMemRefTypeWithFullyDynamicLayout(value.getType().cast<TensorType>(), 246 memorySpace); 247 } 248 249 // Default constructor for BufferizationOptions. 250 BufferizationOptions::BufferizationOptions() 251 : unknownTypeConverterFn(defaultUnknownTypeConverter) {} 252 253 bool BufferizationOptions::isOpAllowed(Operation *op) const { 254 // Special case: If function boundary bufferization is deactivated, do not 255 // allow ops that belong to the `func` dialect. 256 bool isFuncBoundaryOp = isa_and_nonnull<func::FuncDialect>(op->getDialect()); 257 if (!bufferizeFunctionBoundaries && isFuncBoundaryOp) 258 return false; 259 260 return opFilter.isOpAllowed(op); 261 } 262 263 BufferizableOpInterface 264 BufferizationOptions::dynCastBufferizableOp(Operation *op) const { 265 auto bufferizableOp = dyn_cast<BufferizableOpInterface>(op); 266 if (!bufferizableOp) 267 return nullptr; 268 if (!isOpAllowed(op)) 269 return nullptr; 270 return bufferizableOp; 271 } 272 273 BufferizableOpInterface 274 BufferizationOptions::dynCastBufferizableOp(Value value) const { 275 if (auto bufferizableOp = value.getDefiningOp<BufferizableOpInterface>()) 276 if (isOpAllowed(bufferizableOp.getOperation())) 277 return bufferizableOp; 278 return nullptr; 279 } 280 281 void BufferizationOptions::addDialectStateInitializer( 282 StringRef name, const DialectStateInitFn &fn) { 283 stateInitializers.push_back( 284 [=](AnalysisState &state) { state.insertDialectState(name, fn()); }); 285 } 286 287 //===----------------------------------------------------------------------===// 288 // Helper functions for BufferizableOpInterface 289 //===----------------------------------------------------------------------===// 290 291 static void setInsertionPointAfter(OpBuilder &b, Value value) { 292 if (auto bbArg = value.dyn_cast<BlockArgument>()) { 293 b.setInsertionPointToStart(bbArg.getOwner()); 294 } else { 295 b.setInsertionPointAfter(value.getDefiningOp()); 296 } 297 } 298 299 /// Determine which OpOperand* will alias with `result` if the op is bufferized 300 /// in place. Return an empty vector if the op is not bufferizable. 301 SmallVector<OpOperand *> 302 AnalysisState::getAliasingOpOperand(OpResult result) const { 303 if (Operation *op = result.getDefiningOp()) 304 if (auto bufferizableOp = getOptions().dynCastBufferizableOp(op)) 305 return bufferizableOp.getAliasingOpOperand(result, *this); 306 return {}; 307 } 308 309 /// Determine which OpResult will alias with `opOperand` if the op is bufferized 310 /// in place. Return an empty vector if the op is not bufferizable. 311 SmallVector<OpResult> 312 AnalysisState::getAliasingOpResult(OpOperand &opOperand) const { 313 if (auto bufferizableOp = 314 getOptions().dynCastBufferizableOp(opOperand.getOwner())) 315 return bufferizableOp.getAliasingOpResult(opOperand, *this); 316 return {}; 317 } 318 319 /// Return true if `opOperand` bufferizes to a memory read. Return `true` if the 320 /// op is not bufferizable. 321 bool AnalysisState::bufferizesToMemoryRead(OpOperand &opOperand) const { 322 if (auto bufferizableOp = 323 getOptions().dynCastBufferizableOp(opOperand.getOwner())) 324 return bufferizableOp.bufferizesToMemoryRead(opOperand, *this); 325 326 // Unknown op that returns a tensor. The inplace analysis does not support it. 327 // Conservatively return true. 328 return true; 329 } 330 331 /// Return true if `opOperand` bufferizes to a memory write. Return 332 /// `true` if the op is not bufferizable. 333 bool AnalysisState::bufferizesToMemoryWrite(OpOperand &opOperand) const { 334 if (auto bufferizableOp = 335 getOptions().dynCastBufferizableOp(opOperand.getOwner())) 336 return bufferizableOp.bufferizesToMemoryWrite(opOperand, *this); 337 338 // Unknown op that returns a tensor. The inplace analysis does not support it. 339 // Conservatively return true. 340 return true; 341 } 342 343 /// Return true if `opOperand` does neither read nor write but bufferizes to an 344 /// alias. Return false if the op is not bufferizable. 345 bool AnalysisState::bufferizesToAliasOnly(OpOperand &opOperand) const { 346 if (auto bufferizableOp = 347 getOptions().dynCastBufferizableOp(opOperand.getOwner())) 348 return bufferizableOp.bufferizesToAliasOnly(opOperand, *this); 349 350 // Unknown op that returns a tensor. The inplace analysis does not support it. 351 // Conservatively return false. 352 return false; 353 } 354 355 /// Return true if the given value is read by an op that bufferizes to a memory 356 /// read. Also takes into account ops that create an alias but do not read by 357 /// themselves (e.g., ExtractSliceOp). 358 bool AnalysisState::isValueRead(Value value) const { 359 assert(value.getType().isa<TensorType>() && "expected TensorType"); 360 SmallVector<OpOperand *> workingSet; 361 for (OpOperand &use : value.getUses()) 362 workingSet.push_back(&use); 363 364 while (!workingSet.empty()) { 365 OpOperand *uMaybeReading = workingSet.pop_back_val(); 366 // Skip over all ops that neither read nor write (but create an alias). 367 if (bufferizesToAliasOnly(*uMaybeReading)) 368 for (OpResult opResult : getAliasingOpResult(*uMaybeReading)) 369 for (OpOperand &use : opResult.getUses()) 370 workingSet.push_back(&use); 371 if (bufferizesToMemoryRead(*uMaybeReading)) 372 return true; 373 } 374 375 return false; 376 } 377 378 // Starting from `value`, follow the use-def chain in reverse, always selecting 379 // the aliasing OpOperands. Find and return Values for which `condition` 380 // evaluates to true. OpOperands of such matching Values are not traversed any 381 // further. 382 llvm::SetVector<Value> AnalysisState::findValueInReverseUseDefChain( 383 Value value, llvm::function_ref<bool(Value)> condition) const { 384 llvm::SetVector<Value> result, workingSet; 385 workingSet.insert(value); 386 387 while (!workingSet.empty()) { 388 Value value = workingSet.pop_back_val(); 389 if (condition(value) || value.isa<BlockArgument>()) { 390 result.insert(value); 391 continue; 392 } 393 394 OpResult opResult = value.cast<OpResult>(); 395 SmallVector<OpOperand *> opOperands = getAliasingOpOperand(opResult); 396 if (opOperands.empty() || !options.isOpAllowed(value.getDefiningOp())) { 397 result.insert(value); 398 continue; 399 } 400 401 for (OpOperand *o : opOperands) 402 workingSet.insert(o->get()); 403 } 404 405 return result; 406 } 407 408 // Find the Values of the last preceding write of a given Value. 409 llvm::SetVector<Value> 410 AnalysisState::findLastPrecedingWrite(Value value) const { 411 return findValueInReverseUseDefChain(value, [&](Value value) { 412 Operation *op = value.getDefiningOp(); 413 if (!op) 414 return true; 415 auto bufferizableOp = options.dynCastBufferizableOp(op); 416 if (!bufferizableOp) 417 return true; 418 return bufferizableOp.isMemoryWrite(value.cast<OpResult>(), *this); 419 }); 420 } 421 422 AnalysisState::AnalysisState(const BufferizationOptions &options) 423 : options(options) { 424 for (const BufferizationOptions::AnalysisStateInitFn &fn : 425 options.stateInitializers) 426 fn(*this); 427 } 428 429 bool AnalysisState::canOmitTensorCopy(OpOperand &opOperand) const { 430 // Do not copy if the tensor has undefined contents. 431 if (hasUndefinedContents(&opOperand)) 432 return true; 433 434 // Do not copy if the buffer of the tensor is entirely overwritten (with 435 // values that do not depend on the old tensor). 436 if (bufferizesToMemoryWrite(opOperand) && !bufferizesToMemoryRead(opOperand)) 437 return true; 438 439 // Do not copy if the tensor is never read. 440 SmallVector<OpResult> aliasingOpResults = getAliasingOpResult(opOperand); 441 if (!bufferizesToMemoryRead(opOperand) && 442 llvm::none_of(aliasingOpResults, 443 [&](OpResult opResult) { return isValueRead(opResult); })) 444 return true; 445 446 // Default: Cannot omit the copy. 447 return false; 448 } 449 450 bool AnalysisState::isInPlace(OpOperand &opOperand) const { 451 // ToMemrefOps are always in-place. 452 if (isa<ToMemrefOp>(opOperand.getOwner())) 453 return true; 454 455 // In the absence of analysis information, OpOperands that bufferize to a 456 // memory write are out-of-place, i.e., an alloc and copy is inserted. 457 return !bufferizesToMemoryWrite(opOperand); 458 } 459 460 bool AnalysisState::areEquivalentBufferizedValues(Value v1, Value v2) const { 461 // In the absence of analysis information, we do not know if the values are 462 // equivalent. The conservative answer is "false". 463 return false; 464 } 465 466 bool AnalysisState::areAliasingBufferizedValues(Value v1, Value v2) const { 467 // In the absence of analysis information, we do not know if the values may be 468 // aliasing. The conservative answer is "true". 469 return true; 470 } 471 472 bool AnalysisState::hasUndefinedContents(OpOperand *opOperand) const { 473 // In the absence of analysis information, the conservative answer is "false". 474 return false; 475 } 476 477 bool AnalysisState::isTensorYielded(Value tensor) const { 478 // In the absence of analysis information, the conservative answer is "true". 479 if (!tensor.getDefiningOp<AllocTensorOp>()) 480 return true; 481 482 // For AllocTensorOp results, we can do better: They do not alias with any 483 // preceding value, so we can follow SSA use-def chains and do a simple 484 // analysis. 485 SmallVector<OpOperand *> worklist; 486 for (OpOperand &use : tensor.getUses()) 487 worklist.push_back(&use); 488 489 while (!worklist.empty()) { 490 OpOperand *operand = worklist.pop_back_val(); 491 Operation *op = operand->getOwner(); 492 493 // If the op is not bufferizable, we can safely assume that the value is not 494 // yielded. (When bufferizing that op, it must handle such cases.) 495 if (!options.dynCastBufferizableOp(op)) 496 continue; 497 498 // We cannot analyze through ToMemrefOps, so we have to conservatively 499 // assume that the value is yielded. 500 if (isa<ToMemrefOp>(op)) 501 return true; 502 503 // Check if the op is returning/yielding. 504 if (isRegionReturnLike(op)) 505 return true; 506 507 // Add all aliasing OpResults to the worklist. 508 // Note: In the absence of detailed analysis information (e.g., there may be 509 // no function call analysis information), this `getAliasingOpResult` is 510 // conservative and may report additional OpResults as potentially aliasing. 511 for (OpResult opResult : getAliasingOpResult(*operand)) 512 for (OpOperand &use : opResult.getUses()) 513 worklist.push_back(&use); 514 } 515 516 // No ReturnLike op found: The value is not yielded. 517 return false; 518 } 519 520 // bufferization.to_memref is not allowed to change the rank. 521 static void ensureToMemrefOpIsValid(Value tensor, Type memrefType) { 522 #ifndef NDEBUG 523 auto rankedTensorType = tensor.getType().dyn_cast<RankedTensorType>(); 524 assert((!rankedTensorType || memrefType.cast<MemRefType>().getRank() == 525 rankedTensorType.getRank()) && 526 "to_memref would be invalid: mismatching ranks"); 527 #endif 528 } 529 530 FailureOr<Value> bufferization::getBuffer(RewriterBase &rewriter, Value value, 531 const BufferizationOptions &options) { 532 #ifndef NDEBUG 533 auto tensorType = value.getType().dyn_cast<TensorType>(); 534 assert(tensorType && "unexpected non-tensor type"); 535 #endif // NDEBUG 536 537 // Replace "%t = to_tensor %m" with %m. 538 if (auto toTensorOp = value.getDefiningOp<bufferization::ToTensorOp>()) 539 return toTensorOp.getMemref(); 540 541 // Insert to_memref op. 542 OpBuilder::InsertionGuard g(rewriter); 543 setInsertionPointAfter(rewriter, value); 544 FailureOr<BaseMemRefType> memrefType = getBufferType(value, options); 545 if (failed(memrefType)) 546 return failure(); 547 ensureToMemrefOpIsValid(value, *memrefType); 548 return rewriter 549 .create<bufferization::ToMemrefOp>(value.getLoc(), *memrefType, value) 550 .getResult(); 551 } 552 553 /// Return the buffer type for a given Value (tensor) after bufferization. 554 FailureOr<BaseMemRefType> 555 bufferization::getBufferType(Value value, const BufferizationOptions &options) { 556 assert(value.getType().isa<TensorType>() && "unexpected non-tensor type"); 557 Operation *op = getOwnerOfValue(value); 558 559 // ToTensorOp: Take buffer type directly from the op. 560 if (auto toTensorOp = value.getDefiningOp<bufferization::ToTensorOp>()) 561 return toTensorOp.getMemref().getType().cast<BaseMemRefType>(); 562 563 // If value is a bbArg of a bufferizable op: query op interface. 564 if (auto bbArg = value.dyn_cast<BlockArgument>()) 565 if (auto bufferizableOp = 566 options.dynCastBufferizableOp(bbArg.getOwner()->getParentOp())) 567 return bufferizableOp.getBufferType(bbArg, options); 568 569 // Check value is a new buffer allocation with a memory space attribute. In 570 // that case we can at least infer the memory space. 571 Optional<unsigned> memorySpace = None; 572 if (auto opResult = value.dyn_cast<OpResult>()) { 573 if (auto bufferizableOp = 574 options.dynCastBufferizableOp(opResult.getDefiningOp())) { 575 if (bufferizableOp.bufferizesToAllocation(opResult)) { 576 FailureOr<unsigned> queriedMemorySpace = 577 bufferizableOp.getMemorySpace(opResult); 578 if (!failed(queriedMemorySpace)) 579 memorySpace = *queriedMemorySpace; 580 } 581 } 582 } 583 584 // If we still do not know the memory space, use the default memory space (if 585 // any). 586 if (!memorySpace.has_value()) 587 memorySpace = options.defaultMemorySpace; 588 589 // If we still do not know the memory space, report a failure. 590 if (!memorySpace.has_value()) 591 return op->emitError("could not infer memory space"); 592 593 return getMemRefType(value, options, /*layout=*/{}, *memorySpace); 594 } 595 596 void bufferization::replaceOpWithBufferizedValues(RewriterBase &rewriter, 597 Operation *op, 598 ValueRange values) { 599 assert(values.size() == op->getNumResults() && 600 "expected one value per OpResult"); 601 OpBuilder::InsertionGuard g(rewriter); 602 603 // Replace all OpResults with the given values. 604 SmallVector<Value> replacements; 605 for (OpResult opResult : op->getOpResults()) { 606 Value replacement = values[opResult.getResultNumber()]; 607 if (opResult.getType().isa<TensorType>()) { 608 // The OpResult is a tensor. Such values are replaced with memrefs during 609 // bufferization. 610 assert((replacement.getType().isa<MemRefType>() || 611 replacement.getType().isa<UnrankedMemRefType>()) && 612 "tensor op result should be replaced with a memref value"); 613 // The existing uses of the OpResult still expect a tensor. Insert a 614 // ToTensorOp. Throughout bufferization, this ToTensorOp will gradually 615 // loose all of its users and eventually DCE away. 616 rewriter.setInsertionPointAfter(op); 617 replacement = rewriter.create<bufferization::ToTensorOp>( 618 replacement.getLoc(), replacement); 619 } 620 replacements.push_back(replacement); 621 } 622 623 rewriter.replaceOp(op, replacements); 624 } 625 626 //===----------------------------------------------------------------------===// 627 // Bufferization-specific scoped alloc/dealloc insertion support. 628 //===----------------------------------------------------------------------===// 629 630 /// Create a memref allocation with the given type and dynamic extents. 631 FailureOr<Value> BufferizationOptions::createAlloc(OpBuilder &b, Location loc, 632 MemRefType type, 633 ValueRange dynShape) const { 634 if (allocationFn) 635 return (*allocationFn)(b, loc, type, dynShape, bufferAlignment); 636 637 // Default bufferallocation via AllocOp. 638 if (bufferAlignment != 0) 639 return b 640 .create<memref::AllocOp>(loc, type, dynShape, 641 b.getI64IntegerAttr(bufferAlignment)) 642 .getResult(); 643 return b.create<memref::AllocOp>(loc, type, dynShape).getResult(); 644 } 645 646 /// Creates a memref deallocation. The given memref buffer must have been 647 /// allocated using `createAlloc`. 648 LogicalResult BufferizationOptions::createDealloc(OpBuilder &b, Location loc, 649 Value allocatedBuffer) const { 650 if (deallocationFn) 651 return (*deallocationFn)(b, loc, allocatedBuffer); 652 653 // Default buffer deallocation via DeallocOp. 654 b.create<memref::DeallocOp>(loc, allocatedBuffer); 655 return success(); 656 } 657 658 /// Create a memory copy between two memref buffers. 659 LogicalResult BufferizationOptions::createMemCpy(OpBuilder &b, Location loc, 660 Value from, Value to) const { 661 if (memCpyFn) 662 return (*memCpyFn)(b, loc, from, to); 663 664 b.create<memref::CopyOp>(loc, from, to); 665 return success(); 666 } 667 668 //===----------------------------------------------------------------------===// 669 // Bufferization-specific BlockAndValueMapping support with debugging. 670 //===----------------------------------------------------------------------===// 671 672 bool bufferization::isFunctionArgument(Value value) { 673 auto bbArg = value.dyn_cast<BlockArgument>(); 674 if (!bbArg) 675 return false; 676 return isa<func::FuncOp>(bbArg.getOwner()->getParentOp()); 677 } 678 679 BaseMemRefType bufferization::getMemRefType(Value value, 680 const BufferizationOptions &options, 681 MemRefLayoutAttrInterface layout, 682 unsigned memorySpace) { 683 auto tensorType = value.getType().cast<TensorType>(); 684 auto memorySpaceAttr = IntegerAttr::get( 685 IntegerType::get(tensorType.getContext(), 64), memorySpace); 686 687 // Case 1: Unranked memref type. 688 if (auto unrankedTensorType = tensorType.dyn_cast<UnrankedTensorType>()) { 689 assert(!layout && "UnrankedTensorType cannot have a layout map"); 690 return UnrankedMemRefType::get(unrankedTensorType.getElementType(), 691 memorySpaceAttr); 692 } 693 694 // Case 2: Ranked memref type with specified layout. 695 auto rankedTensorType = tensorType.cast<RankedTensorType>(); 696 if (layout) { 697 return MemRefType::get(rankedTensorType.getShape(), 698 rankedTensorType.getElementType(), layout, 699 memorySpaceAttr); 700 } 701 702 return options.unknownTypeConverterFn(value, memorySpace, options); 703 } 704 705 BaseMemRefType 706 bufferization::getMemRefTypeWithFullyDynamicLayout(TensorType tensorType, 707 unsigned memorySpace) { 708 // Case 1: Unranked memref type. 709 if (auto unrankedTensorType = tensorType.dyn_cast<UnrankedTensorType>()) { 710 return UnrankedMemRefType::get(unrankedTensorType.getElementType(), 711 memorySpace); 712 } 713 714 // Case 2: Ranked memref type. 715 auto memorySpaceAttr = IntegerAttr::get( 716 IntegerType::get(tensorType.getContext(), 64), memorySpace); 717 auto rankedTensorType = tensorType.cast<RankedTensorType>(); 718 int64_t dynamicOffset = ShapedType::kDynamicStrideOrOffset; 719 SmallVector<int64_t> dynamicStrides(rankedTensorType.getRank(), 720 ShapedType::kDynamicStrideOrOffset); 721 AffineMap stridedLayout = makeStridedLinearLayoutMap( 722 dynamicStrides, dynamicOffset, rankedTensorType.getContext()); 723 return MemRefType::get(rankedTensorType.getShape(), 724 rankedTensorType.getElementType(), stridedLayout, 725 memorySpaceAttr); 726 } 727 728 /// Return a MemRef type with a static identity layout (i.e., no layout map). If 729 /// the given tensor type is unranked, return an unranked MemRef type. 730 BaseMemRefType 731 bufferization::getMemRefTypeWithStaticIdentityLayout(TensorType tensorType, 732 unsigned memorySpace) { 733 // Case 1: Unranked memref type. 734 if (auto unrankedTensorType = tensorType.dyn_cast<UnrankedTensorType>()) { 735 return UnrankedMemRefType::get(unrankedTensorType.getElementType(), 736 memorySpace); 737 } 738 739 // Case 2: Ranked memref type. 740 auto rankedTensorType = tensorType.cast<RankedTensorType>(); 741 auto memorySpaceAttr = IntegerAttr::get( 742 IntegerType::get(tensorType.getContext(), 64), memorySpace); 743 MemRefLayoutAttrInterface layout = {}; 744 return MemRefType::get(rankedTensorType.getShape(), 745 rankedTensorType.getElementType(), layout, 746 memorySpaceAttr); 747 } 748