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