MLIR 24.0.0git
BufferizableOpInterface.h
Go to the documentation of this file.
1//===- BufferizableOpInterface.h - Bufferizable Ops -------------*- C++ -*-===//
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#ifndef MLIR_DIALECT_BUFFERIZATION_IR_BUFFERIZABLEOPINTERFACE_H_
10#define MLIR_DIALECT_BUFFERIZATION_IR_BUFFERIZABLEOPINTERFACE_H_
11
12#include "mlir/IR/Operation.h"
14#include "mlir/Support/LLVM.h"
15#include "llvm/ADT/DenseMapInfoVariant.h"
16#include "llvm/ADT/SetVector.h"
17#include <optional>
18
19#include "mlir/Dialect/Bufferization/IR/BufferizationEnums.h.inc"
21
22namespace mlir {
23class OpBuilder;
24namespace func {
25class FuncOp;
26}
27
28namespace bufferization {
29
30class AnalysisState;
31class BufferizableOpInterface;
32
33/// Specifies a fine-grain relationship between buffers to enable more analysis.
34enum class BufferRelation {
35 Unknown,
36 // TODO: ResultContainsOperand,
37 // TODO: OperandContainsResult,
38 Equivalent
39};
40
41/// A maybe aliasing OpOperand. If `isDefinite` is `true`, the OpOperand is
42/// guaranteed to alias at runtime.
43struct AliasingOpOperand {
44 AliasingOpOperand(OpOperand *opOperand, BufferRelation relation,
45 bool isDefinite = true)
46 : opOperand(opOperand), relation(relation), isDefinite(isDefinite) {}
47
48 OpOperand *opOperand;
49 BufferRelation relation;
50 bool isDefinite;
51};
52
53/// A maybe aliasing Value. If `isDefinite` is `true`, the Value is guaranteed
54/// to alias at runtime.
55struct AliasingValue {
56 AliasingValue(Value value, BufferRelation relation, bool isDefinite = true)
57 : value(value), relation(relation), isDefinite(isDefinite) {}
58
59 Value value;
60 BufferRelation relation;
61 bool isDefinite;
62};
63
64template <typename T>
65class AliasList {
66public:
67 /// Create an empty list of aliases.
68 AliasList() = default;
69
70 /// Create a list of aliases.
71 AliasList(std::initializer_list<T> elems) {
72 for (T alias : elems)
73 addAlias(alias);
74 }
75
76 /// Create a list of aliases.
77 AliasList(SmallVector<T> &&aliases) : aliases(std::move(aliases)) {}
78
79 ArrayRef<T> getAliases() const { return aliases; }
80
81 size_t getNumAliases() const { return aliases.size(); }
82
83 void addAlias(T alias) { aliases.push_back(alias); }
84
85 auto begin() const { return aliases.begin(); }
86 auto end() const { return aliases.end(); }
87
88private:
89 /// The list of aliases.
90 SmallVector<T> aliases;
91};
92
93/// A list of possible aliasing OpOperands. This list models the runtime
94/// aliasing relationship for a Value.
95using AliasingOpOperandList = AliasList<AliasingOpOperand>;
96
97/// A list of possible aliasing Values. This list models the runtime aliasing
98/// relationship for an OpOperand.
99using AliasingValueList = AliasList<AliasingValue>;
100
101class OpFilter {
102public:
103 /// An op filter entry. Filters can be used to specify which ops should be
104 /// processed by the bufferization.
105 struct Entry {
106 /// If the filter function evaluates to `true`, the filter matches.
107 using FilterFn = std::function<bool(Operation *)>;
108
109 /// Filter type: A filter can either be a DENY filter or an ALLOW filter.
110 enum FilterType : int8_t { DENY = 0, ALLOW = 1 };
111
112 FilterFn fn;
113 FilterType type;
114 };
115
116 /// Return whether the op is allowed or not.
117 ///
118 /// If the filter does not have an ALLOW rule, ops are allowed by default,
119 /// unless they are explicitly marked as DENY. If the filter has at least one
120 /// ALLOW rule, ops are denied by default and only allowed if they match
121 /// an ALLOW rule and no DENY rule.
122 bool isOpAllowed(Operation *op) const;
123
124 /// Allow the given dialects.
125 ///
126 /// This function adds one or multiple ALLOW entries.
127 template <typename... DialectTs>
128 void allowDialect() {
129 // The following expands a call to allowDialectImpl for each dialect
130 // in 'DialectTs'.
131 (allowDialectImpl<DialectTs>(), ...);
132 }
133
134 /// Deny the given dialects.
135 ///
136 /// This function adds one or multiple DENY entries.
137 template <typename... DialectTs>
138 void denyDialect() {
139 (denyDialectImpl<DialectTs>(), ...);
140 }
141
142 /// Allow the given dialect.
143 ///
144 /// This function adds an ALLOW entry.
145 void allowDialect(StringRef dialectNamespace) {
146 Entry::FilterFn filterFn = [=](Operation *op) {
147 return op->getName().getDialectNamespace() == dialectNamespace;
148 };
149 entries.push_back(Entry{filterFn, Entry::FilterType::ALLOW});
150 }
151
152 /// Deny the given dialect.
153 ///
154 /// This function adds a DENY entry.
155 void denyDialect(StringRef dialectNamespace) {
156 Entry::FilterFn filterFn = [=](Operation *op) {
157 return op->getName().getDialectNamespace() == dialectNamespace;
158 };
159 entries.push_back(Entry{filterFn, Entry::FilterType::DENY});
160 }
161
162 /// Allow the given ops.
163 ///
164 /// This function adds one or multiple ALLOW entries.
165 template <typename... OpTys>
166 void allowOperation() {
167 (allowOperationImpl<OpTys>(), ...);
168 }
169
170 /// Deny the given ops.
171 ///
172 /// This function adds one or multiple DENY entries.
173 template <typename... OpTys>
174 void denyOperation() {
175 (denyOperationImpl<OpTys>(), ...);
176 }
177
178 /// Allow the given op.
179 ///
180 /// This function adds an ALLOW entry.
181 void allowOperation(StringRef opName) {
182 Entry::FilterFn filterFn = [=](Operation *op) {
183 return op->getName().getStringRef() == opName;
184 };
185 allowOperation(filterFn);
186 }
187
188 /// Deny the given op.
189 ///
190 /// This function adds a DENY entry.
191 void denyOperation(StringRef opName) {
192 Entry::FilterFn filterFn = [=](Operation *op) {
193 return op->getName().getStringRef() == opName;
194 };
195 denyOperation(filterFn);
196 }
197
198 /// Allow ops that are matched by `fn`.
199 ///
200 /// This function adds an ALLOW entry.
201 void allowOperation(Entry::FilterFn fn) {
202 entries.push_back(Entry{fn, Entry::FilterType::ALLOW});
203 }
204
205 /// Deny ops that are matched by `fn`.
206 ///
207 /// This function adds a DENY entry.
208 void denyOperation(Entry::FilterFn fn) {
209 entries.push_back(Entry{fn, Entry::FilterType::DENY});
210 }
211
212private:
213 /// Return `true` if the filter has at least one ALLOW rule.
214 bool hasAllowRule() const {
215 for (const Entry &e : entries)
216 if (e.type == Entry::FilterType::ALLOW)
217 return true;
218 return false;
219 }
220
221 /// Allow a dialect.
222 template <typename DialectT>
223 void allowDialectImpl() {
224 allowDialect(DialectT::getDialectNamespace());
225 }
226
227 /// Deny a dialect.
228 template <typename DialectT>
229 void denyDialectImpl() {
230 denyDialect(DialectT::getDialectNamespace());
231 }
232
233 /// Allow an op.
234 template <typename OpTy>
235 void allowOperationImpl() {
236 allowOperation(OpTy::getOperationName());
237 }
238
239 /// Deny an op.
240 template <typename OpTy>
241 void denyOperationImpl() {
242 denyOperation(OpTy::getOperationName());
243 }
244
245 /// A list of filter entries that determine whether an op should be allowed or
246 /// denied. If the filter has an ALLOW rule, only ops that are allowed and not
247 /// denied are allowed. If the filter does not have an ALLOW rule, only ops
248 /// that are not denied are allowed.
249 SmallVector<Entry> entries;
250};
251
252/// Options for BufferizableOpInterface-based bufferization.
253struct BufferizationOptions {
254 /// Allocator function: Generate a memref allocation with the given type,
255 /// dynamic extents and alignment.
256 using AllocationFn = std::function<FailureOr<Value>(
257 OpBuilder &, Location, MemRefType, ValueRange, unsigned int)>;
258 /// Memcpy function: Generate a memcpy between two buffers.
259 using MemCpyFn =
260 std::function<LogicalResult(OpBuilder &, Location, Value, Value)>;
261 /// Cast function: Convert a buffer value to a new value with the specified
262 /// type. This method is typically used when a simple cast-like operation is
263 /// sufficient to convert the buffer value, for example, when layout maps
264 /// between buffer value and resulting type do not match.
265 using CastFn =
266 std::function<FailureOr<Value>(OpBuilder &, Location, Type, Value)>;
267 /// Initializer function for analysis state.
268 using AnalysisStateInitFn = std::function<void(AnalysisState &)>;
269 /// Tensor-like -> Buffer-like type conversion.
270 /// Parameters: tensor-like type, memory space, func op, bufferization options
271 using FunctionArgTypeConverterFn =
272 std::function<BufferLikeType(TensorLikeType, Attribute memorySpace,
273 func::FuncOp, const BufferizationOptions &)>;
274 /// Tensor -> MemRef type conversion.
275 /// Parameters: tensor type, memory space, bufferization options
276 using UnknownTypeConverterFn = std::function<BufferLikeType(
277 TensorLikeType, Attribute memorySpace, const BufferizationOptions &)>;
278 // Produce a MemorySpace attribute from a tensor type
279 using DefaultMemorySpaceFn =
280 std::function<std::optional<Attribute>(TensorLikeType t)>;
281
282 /// Resolve a mismatch between buffer types that were independently inferred,
283 /// which results in a conflict at the "merge" point. Returns `failure()` to
284 /// signal bufferization failure; returns a buffer-like type when
285 /// reconciliation suceeded.
286 using ReconcileBufferTypeMismatchFn = std::function<FailureOr<BufferLikeType>(
287 BufferLikeType, BufferLikeType, const BufferizationOptions &)>;
288
289 BufferizationOptions();
290
291 /// Try to cast the given op to BufferizableOpInterface if the op is allow
292 /// listed.
293 BufferizableOpInterface dynCastBufferizableOp(Operation *op) const;
294
295 /// Try to cast the given value to BufferizableOpInterface if the op is allow
296 /// listed.
297 BufferizableOpInterface dynCastBufferizableOp(Value value) const;
298
299 /// A filter that specifies which ops should be bufferized and which ops
300 /// should be ignored.
301 OpFilter opFilter;
302
303 /// Return `true` if the given op should be bufferized.
304 bool isOpAllowed(Operation *op) const;
305
306 /// Specifies whether not bufferizable ops are allowed in the input. If so,
307 /// bufferization.to_buffer and bufferization.to_tensor ops are inserted at
308 /// the boundaries.
309 bool allowUnknownOps = false;
310
311 /// Specifies whether function boundaries (ops in the func dialect) should be
312 /// bufferized or not.
313 bool bufferizeFunctionBoundaries = false;
314
315 /// Whether the IR may contain a parallel region. When unset, the analysis
316 /// walks the IR to compute it.
317 /// Note: If the IR contains a parallel region, but this flag is set to
318 /// "false", the bufferization may produce incorrect IR.
319 std::optional<bool> mayHaveParallelRegions = std::nullopt;
320
321 /// This function controls buffer types on function signatures. Sets
322 /// `functionArgTypeConverterFn` and `inferFunctionResultLayout` accordingly.
323 ///
324 /// * InferLayoutMap: All function parameter types have a fully dynamic layout
325 /// map, but function result types are inferred from the body of the
326 /// function.
327 /// * FullyDynamicLayoutMap: All function parameter types and result types
328 /// have a fully dynamic layout map. This option is most efficient because
329 /// any layout map can be casted to a fully dynamic one.
330 /// * IdentityLayoutMap: All function parameter types and result types have a
331 /// static identity layout (i.e., no layout map). This option may introduce
332 /// additional buffer allocs and copies because layout maps cannot be casted
333 /// away.
334 ///
335 /// Note: Inferred layout maps may not be desireable when interacting with
336 /// external functions, because the generated function signatures will be less
337 /// predictable.
338 void setFunctionBoundaryTypeConversion(LayoutMapOption layoutMapOption);
339
340 /// Create a memref allocation with the given type and dynamic extents.
341 AllocationFn allocationFn = nullptr;
342
343 /// Creates a memcpy between two given buffers.
344 MemCpyFn memCpyFn = nullptr;
345
346 /// Creates a cast function from a buffer value to a new type.
347 CastFn castFn = nullptr;
348
349 /// Type conversion from tensors to buffers. This type conversion is used to
350 /// determine bufferized function argument and result types.
351 ///
352 /// By default, if tensor is a (builtin) tensor type, it is converted to a
353 /// memref type with a fully dynamic layout map; if tensor is a (generic)
354 /// tensor-like type, it is converted using unknownTypeConverterFn.
355 ///
356 /// If `bufferizeFunctionBoundaries` is not set, this function isn't used.
357 FunctionArgTypeConverterFn functionArgTypeConverterFn = nullptr;
358
359 /// If true, function result types are inferred from the body of the function.
360 /// Otherwise, function result type is determined by
361 /// `functionArgTypeConverterFn`.
362 ///
363 /// If `bufferizeFunctionBoundaries` is not set, this flag has no effect.
364 bool inferFunctionResultLayout = true;
365
366 /// Type conversion from tensors to memrefs. This type conversion is used if
367 /// no memref type could be inferred during bufferization. By default, returns
368 /// a memref type with a fully dynamic layout map.
369 UnknownTypeConverterFn unknownTypeConverterFn = nullptr;
370
371 // Use during type conversion to determine the memory space for memref based
372 // on the original tensor type if the memory space cannot be inferred.
373 // Returning std::nullopt will cause bufferization to fail (useful to indicate
374 // failure to determine memory space for a tensor type).
375 DefaultMemorySpaceFn defaultMemorySpaceFn =
376 [](TensorLikeType t) -> std::optional<Attribute> { return Attribute(); };
377
378 /// Hook to resolve a mismatch between conflicting buffer types that were
379 /// independently inferred and have to now "converge" to a common buffer type
380 /// (e.g. due to differences in iterations of a loop or branches of
381 /// if-statements). Depending on the situation and the types involved, this
382 /// may produce a "joined" type (e.g. a type combining properties of both), or
383 /// either one of the two types, etc. The default keeps the framework
384 /// behavior: promote to fully-dynamic layout on layout mismatch, fail on
385 /// memory-space mismatch.
386 ReconcileBufferTypeMismatchFn reconcileBufferTypeMismatchFn = nullptr;
387
388 /// If set to `true`, the analysis is skipped. A buffer is copied before every
389 /// write. This flag cannot be used together with `testAnalysisOnly = true`.
390 bool copyBeforeWrite = false;
391
392 /// If set to `true`, does not modify the IR apart from adding attributes (for
393 /// checking the results of the analysis) and post analysis steps.
394 bool testAnalysisOnly = false;
395
396 /// If set to `true`, the IR is annotated with details about RaW conflicts.
397 /// For debugging only. Should be used together with `testAnalysisOnly`.
398 bool printConflicts = false;
399
400 /// Buffer alignment for new memory allocations.
401 unsigned int bufferAlignment = 64;
402
403 /// Initializer functions for analysis state. These can be used to
404 /// initialize dialect-specific analysis state.
405 SmallVector<AnalysisStateInitFn> stateInitializers;
406};
407
408/// Traversal parameters for `findValueInReverseUseDefChain`.
409struct TraversalConfig {
410 /// Specifies if leaves (that do not have further OpOperands to follow)
411 /// should be returned even if they do not match the specified filter.
412 bool alwaysIncludeLeaves = true;
413
414 /// Specifies whether out-of-place/undecided OpOperands should be followed.
415 bool followInPlaceOnly = false;
416
417 /// Specifies whether non-equivalent OpOperands should be followed.
418 bool followEquivalentOnly = false;
419
420 /// Specifies whether unknown/non-bufferizable/ops not included in the
421 /// OpFilter of BufferizationOptions should be followed.
422 bool followUnknownOps = false;
423
424 /// Specifies whether OpOperands with a different type that are not the result
425 /// of a CastOpInterface op should be followed.
426 bool followSameTypeOrCastsOnly = false;
427
428 /// Specifies whether already visited values should be visited again.
429 /// (Note: This can result in infinite looping.)
430 bool revisitAlreadyVisitedValues = false;
431};
432
433/// AnalysisState provides a variety of helper functions for dealing with
434/// tensor values.
435class AnalysisState {
436public:
437 /// Determine which OpOperand* will alias with `value` if the op is
438 /// bufferized in place. Return all tensor OpOperand* if the op is not
439 /// bufferizable.
440 AliasingOpOperandList getAliasingOpOperands(Value value) const;
441
442 /// Determine which Value will alias with `opOperand` if the op is bufferized
443 /// in place. Return all tensor Values if the op is not bufferizable.
444 AliasingValueList getAliasingValues(OpOperand &opOperand) const;
445
446 /// Return true if `opOperand` bufferizes to a memory read. Return `true` if
447 /// the op is not bufferizable.
448 bool bufferizesToMemoryRead(OpOperand &opOperand) const;
449
450 /// Return true if `opOperand` bufferizes to a memory write. Return true` if
451 /// the op is not bufferizable.
452 bool bufferizesToMemoryWrite(OpOperand &opOperand) const;
453
454 /// Return true if the given `value` bufferizes to a memory write. Return
455 /// true if the value is a block argument. Return `true` if the defining op is
456 /// not bufferizable. Otherwise, consult the BufferizableOpInterface.
457 bool bufferizesToMemoryWrite(Value value) const;
458
459 /// Return true if `opOperand` does neither read nor write but bufferizes to
460 /// an alias. Return false if the op is not bufferizable.
461 bool bufferizesToAliasOnly(OpOperand &opOperand) const;
462
463 /// Return true if a copy can always be avoided when allocating a new tensor
464 /// for the given OpOperand.
465 bool canOmitTensorCopy(OpOperand &opOperand) const;
466
467 /// Return true if the given value is read by an op that bufferizes to a
468 /// memory read. Also takes into account ops that create an alias but do not
469 /// read by themselves (e.g., ExtractSliceOp).
470 bool isValueRead(Value value) const;
471
472 /// Starting from `opOperand`, follow the use-def chain in reverse, always
473 /// selecting the aliasing OpOperands. Find and return Values for which
474 /// `condition` evaluates to true. OpOperands of such matching Values are not
475 /// traversed any further, the visited aliasing opOperands will be preserved
476 /// through `visitedOpOperands`.
477 ///
478 /// When reaching the end of a chain, also return the last Value of that
479 /// chain if `config.alwaysIncludeLeaves` is set.
480 ///
481 /// Example:
482 ///
483 /// 8
484 /// |
485 /// 6* 7* +-----+----+
486 /// | | | |
487 /// 2* 3 4* 5
488 /// | | | |
489 /// +----------+----------+----------+
490 /// |
491 /// 1
492 ///
493 /// In the above example, Values with a star satisfy the condition. When
494 /// starting the traversal from Value 1, the resulting SetVector is:
495 /// { 2, 7, 8, 5 }
496 ///
497 /// Additional stopping conditions for the traversal can be specified in
498 /// `config`.
499 SetVector<Value> findValueInReverseUseDefChain(
500 OpOperand *opOperand, llvm::function_ref<bool(Value)> condition,
501 TraversalConfig config = TraversalConfig(),
502 llvm::DenseSet<OpOperand *> *visitedOpOperands = nullptr) const;
503
504 /// Find the values that may define the contents of the given value at
505 /// runtime. A block argument is always a definition. An OpResult is a
506 /// definition if it bufferizes to memory write. If it does not bufferize to
507 /// a memory write but has aliasing operands, we continue the lookup on these
508 /// values.
509 ///
510 /// Example: %r = tensor.insert %f into %t[%c0] : tensor<?xf32>
511 /// findDefinitions(%r) = {%r} because %r bufferizes to memory write.
512 ///
513 /// Example: %r = tensor.empty() : tensor<10xf32>
514 /// findDefinitions(%r) = {} because tensor.empty does not the define the
515 /// contents of its result (i.e., it does not bufferize to a memory write)
516 /// and it has no aliasing OpOperands.
517 ///
518 /// Example:
519 /// %a = arith.constant ... : tensor<10xf32>
520 /// %b1 = tensor.insert %f into %t : tensor<50xf32>
521 /// %b2 = tensor.extract_slice %b1[0][10][1] : tensor<50xf32> tensor<10xf32>
522 /// %r = arith.select %cond, %a, %b : tensor<10xf32>
523 /// findDefinitions(%r) = {%a, %b1}. %r and %b2 are skipped (lookup continues
524 /// in the operands) because their defining ops do not define the contents of
525 /// the tensor.
526 ///
527 /// Example:
528 /// %a = tensor.empty() : tensor<10xf32>
529 /// %b = arith.constant ... : tensor<10xf32>
530 /// %r = arith.select %cond, %a, %b : tensor<10xf32>
531 /// findDefinitions(%r) = {%b}. %a is excluded because it does not define the
532 /// contents of the tensor.
533 ///
534 /// Note: OpResults of unknown ops are handled conservatively and assumed to
535 /// be definitions.
536 SetVector<Value> findDefinitions(OpOperand *opOperand) const;
537
538 /// Return `true` if the given OpResult has been decided to bufferize inplace.
539 virtual bool isInPlace(OpOperand &opOperand) const;
540
541 /// Return true if `v1` and `v2` bufferize to equivalent buffers.
542 virtual bool areEquivalentBufferizedValues(Value v1, Value v2) const;
543
544 /// Return true if `v1` and `v2` may bufferize to aliasing buffers.
545 virtual bool areAliasingBufferizedValues(Value v1, Value v2) const;
546
547 /// Return `true` if the given tensor has undefined contents.
548 virtual bool hasUndefinedContents(OpOperand *opOperand) const;
549
550 /// Return a reference to the BufferizationOptions.
551 const BufferizationOptions &getOptions() const { return options; }
552
553 AnalysisState(const BufferizationOptions &options);
554
555 // AnalysisState should be passed as a reference.
556 AnalysisState(const AnalysisState &) = delete;
557
558 virtual ~AnalysisState() = default;
559
560 static bool classof(const AnalysisState *base) { return true; }
561
562 TypeID getType() const { return type; }
563
564 /// Return the closest enclosing repetitive region around the given op.
565 Region *getEnclosingRepetitiveRegion(Operation *op,
566 const BufferizationOptions &options);
567
568 /// Return the closest enclosing repetitive region around the place where the
569 /// given value is defined.
570 Region *getEnclosingRepetitiveRegion(Value value,
571 const BufferizationOptions &options);
572
573 /// Return the closest enclosing repetitive region around the given block.
574 Region *getEnclosingRepetitiveRegion(Block *block,
575 const BufferizationOptions &options);
576
577 virtual void resetCache();
578
579 /// Checks whether `op0` and `op1` are inside mutually exclusive regions.
580 /// The logic defers to `mlir::insideMutuallyExclusiveRegions`, but the
581 /// result is cached.
582 bool insideMutuallyExclusiveRegions(Operation *op0, Operation *op1);
583
584protected:
585 AnalysisState(const BufferizationOptions &options, TypeID type);
586
587private:
588 /// A reference to current bufferization options.
589 const BufferizationOptions &options;
590
591 /// The type of analysis.
592 TypeID type;
593
594 /// Cache containing closest ancestor repetitive Region.
595 DenseMap<std::variant<Operation *, Block *, Region *, Value>, Region *>
596 enclosingRepetitiveRegionCache;
597
598 /// Cache that specifies whether the two operations are in mutually exclusive
599 /// regions.
600 DenseMap<std::pair<Operation *, Operation *>, bool>
601 insideMutuallyExclusiveRegionsCache;
602};
603
604/// BufferizationState provides information about the state of the IR during the
605/// bufferization process.
606class BufferizationState {
607public:
608 /// Get a reference to the collection of cached symbol tables.
609 SymbolTableCollection &getSymbolTables();
610 /// Const overload so callers can reuse the cache from a const state.
611 SymbolTableCollection &getSymbolTables() const;
612
613private:
614 /// The cached symbol tables.
615 /// The user is expected to update / invalidate the cached symbol tables if
616 /// the bufferized operation has the Symbol or SymbolTable traits.
617 mutable SymbolTableCollection symbolTables;
618};
619
620/// Create an AllocTensorOp for the given shaped value (memref or tensor).
621/// If `copy` is set, the shaped value is copied. Otherwise, a tensor with
622/// undefined contents is allocated.
623FailureOr<Value>
624allocateTensorForShapedValue(OpBuilder &b, Location loc, Value shapedValue,
625 const BufferizationOptions &options,
626 const BufferizationState &state, bool copy = true);
627
628/// Lookup the buffer for the given value. If the value was not bufferized
629/// yet, wrap it in a ToBufferOp. Otherwise, it is the result of a ToTensorOp,
630/// from which the memref operand is returned.
631FailureOr<Value> getBuffer(RewriterBase &rewriter, Value value,
632 const BufferizationOptions &options,
633 const BufferizationState &state);
634
635/// Return the buffer type for a given Value (tensor) after bufferization
636/// without bufferizing any IR.
637///
638/// Note: It should be sufficient to call `getBuffer()->getType()` in most
639/// cases. However, when a buffer type should be predicted without modifying any
640/// IR, this function can be used.
641///
642/// This function is a wrapper around BufferizableOpInterface::getBufferType.
643FailureOr<BufferLikeType> getBufferType(Value value,
644 const BufferizationOptions &options,
645 const BufferizationState &state);
646
647/// Return the buffer type for a given Value (tensor) after bufferization
648/// without bufferizing any IR. This function (and not the other overload
649/// without `invocationStack`) can be used from `getBufferType` implementations
650/// of the `BufferizableOpInterface`.
651///
652/// Note: It should be sufficient to call `getBuffer()->getType()` in most
653/// cases. However, when a buffer type should be predicted without modifying any
654/// IR, this function can be used.
655///
656/// This function is a wrapper around `BufferizableOpInterface::getBufferType`.
657FailureOr<BufferLikeType> getBufferType(Value value,
658 const BufferizationOptions &options,
659 const BufferizationState &state,
660 SmallVector<Value> &invocationStack);
661
662/// Return "true" if the given op has tensor semantics and should be bufferized.
663/// If the op is bufferizable, the BufferizableOpInterface is queried.
664/// Otherwise, an op has tensor semantics if it has tensor operands, tensor
665/// op results and/or tensor block arguments.
666bool hasTensorSemantics(Operation *op);
667
668/// Replace an op with replacement values. The op is deleted. Tensor OpResults
669/// must be replaced with memref values.
670void replaceOpWithBufferizedValues(RewriterBase &rewriter, Operation *op,
671 ValueRange values);
672
673/// Replace an op with a new op. The new op must have the same number of
674/// results as the replaced op. The new op may not return any tensor values.
675template <typename OpTy, typename... Args>
676OpTy replaceOpWithNewBufferizedOp(RewriterBase &rewriter, Operation *op,
677 Args &&...args) {
678 auto newOp =
679 OpTy::create(rewriter, op->getLoc(), std::forward<Args>(args)...);
680 replaceOpWithBufferizedValues(rewriter, op, newOp->getResults());
681 return newOp;
682}
683
684/// Return a MemRef type with fully dynamic layout. If the given tensor type
685/// is unranked, return an unranked MemRef type.
686BaseMemRefType
687getMemRefTypeWithFullyDynamicLayout(TensorType tensorType,
688 Attribute memorySpace = nullptr);
689
690/// Return a MemRef type with a static identity layout (i.e., no layout map). If
691/// the given tensor type is unranked, return an unranked MemRef type.
692BaseMemRefType
693getMemRefTypeWithStaticIdentityLayout(TensorType tensorType,
694 Attribute memorySpace = nullptr);
695
696/// Return the owner of the given value. In case of a BlockArgument that is the
697/// owner of the block. In case of an OpResult that is the defining op.
698Operation *getOwnerOfValue(Value value);
699
700/// Assuming that the given region is repetitive, find the next enclosing
701/// repetitive region.
702Region *getNextEnclosingRepetitiveRegion(Region *region,
703 const BufferizationOptions &options);
704
705/// If `region` is a parallel region, return `region`. Otherwise, find the first
706/// enclosing parallel region of `region`. If there is no such region, return
707/// "nullptr".
708///
709/// Note: Whether a region is parallel or sequential is queried from the
710/// `BufferizableOpInterface`.
711Region *getParallelRegion(Region *region, const BufferizationOptions &options);
712
713namespace detail {
714/// This is the default implementation of
715/// BufferizableOpInterface::getAliasingOpOperands. Should not be called from
716/// other places.
717AliasingOpOperandList defaultGetAliasingOpOperands(Value value,
718 const AnalysisState &state);
719
720/// This is the default implementation of
721/// BufferizableOpInterface::getBufferType. Should not be called from other
722/// places.
723FailureOr<BufferLikeType>
724defaultGetBufferType(Value value, const BufferizationOptions &options,
725 const BufferizationState &state,
726 SmallVector<Value> &invocationStack);
727
728/// This is the default implementation of
729/// BufferizableOpInterface::resultBufferizesToMemoryWrite. Should not be called
730/// from other places.
731bool defaultResultBufferizesToMemoryWrite(OpResult opResult,
732 const AnalysisState &state);
733
734/// This is the default implementation of
735/// BufferizableOpInterface::isRepetitiveRegion. Should not be called from other
736/// places.
737bool defaultIsRepetitiveRegion(BufferizableOpInterface bufferizableOp,
738 unsigned index);
739
740/// This is the default implementation of getAliasingOpOperands in case the
741/// defining op does not implement the BufferizableOpInterface.
742AliasingOpOperandList unknownGetAliasingOpOperands(Value value);
743
744/// This is the default implementation of getAliasingValues in case the owner
745/// op does not implement the BufferizableOpInterface.
746AliasingValueList unknownGetAliasingValues(OpOperand &opOperand);
747
748/// This is the default implementation of
749/// BufferizableOpInterface::hasTensorSemantics
750bool defaultHasTensorSemantics(Operation *op);
751
752/// This is a helper function used when buffer type is guaranteed to be memref.
753/// It performs two actions: failure state checking and an explicit llvm::cast<>
754/// from the buffer-like type interface to a BaseMemRefType. This allows easier
755/// management of differences in C++ types at the API boundaries. Valid buffer
756/// type is casted to the memref type. Otherwise, the failure state is
757/// propagated i.e. asMemRefType(mlir::failure()) returns mlir::failure().
758FailureOr<BaseMemRefType> asMemRefType(FailureOr<BufferLikeType> bufferType);
759
760/// This function is a free-standing helper that relies on
761/// bufferization::TensorLikeTypeInterface to verify the types in tensor and
762/// buffer worlds match.
763bool typesMatchAfterBufferization(Operation &op, Value tensor, Value buffer);
764} // namespace detail
765
766} // namespace bufferization
767} // namespace mlir
768
769MLIR_DECLARE_EXPLICIT_TYPE_ID(mlir::bufferization::AnalysisState)
770
771//===----------------------------------------------------------------------===//
772// Bufferization Interfaces
773//===----------------------------------------------------------------------===//
774
775#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h.inc"
776
777#endif // MLIR_DIALECT_BUFFERIZATION_IR_BUFFERIZABLEOPINTERFACE_H_
bufferization::BufferResultsToOutParamsOpts::AllocationFn AllocationFn
bufferization::BufferResultsToOutParamsOpts::MemCpyFn MemCpyFn
static void copy(Location loc, Value dst, Value src, Value size, OpBuilder &builder)
Copies the given number of bytes from src to dst pointers.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static llvm::ManagedStatic< PassManagerOptions > options
static RankedTensorType getBufferType(const SparseTensorType &stt, bool needTmpCOO)
#define MLIR_DECLARE_EXPLICIT_TYPE_ID(CLASS_NAME)
Definition TypeID.h:321
static Operation * getOwnerOfValue(Value value)
This class helps build Operations.
Definition Builders.h:210
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
bool insideMutuallyExclusiveRegions(Operation *a, Operation *b)
Return true if a and b are in mutually exclusive regions as per RegionBranchOpInterface.
Region * getEnclosingRepetitiveRegion(Operation *op)
Return the first enclosing region of the given op that may be executed repetitively as per RegionBran...