MLIR 24.0.0git
BufferizableOpInterfaceImpl.cpp
Go to the documentation of this file.
1//===- BufferizableOpInterfaceImpl.cpp - Impl. of BufferizableOpInterface -===//
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
14#include "mlir/IR/Dialect.h"
15#include "mlir/IR/Operation.h"
17
18using namespace mlir;
19using namespace linalg;
20using namespace mlir::bufferization;
21
22namespace {
23
24/// Generic conversion for any DestinationStyleOpInterface on tensors.
25static LogicalResult bufferizeDestinationStyleOpInterface(
26 RewriterBase &rewriter, DestinationStyleOpInterface op,
27 const BufferizationOptions &options, const BufferizationState &state) {
28 // Take a guard before anything else.
29 OpBuilder::InsertionGuard g(rewriter);
30 rewriter.setInsertionPoint(op);
31
32 // Nothing to do. This op is already bufferized.
33 if (op.hasPureBufferSemantics())
34 return success();
35
36 // Ensure op has only tensors. Allow mixed tensor-buffer mode on a per-need
37 // basis.
38 if (!op.hasPureTensorSemantics())
39 return op->emitError() << "op does not have pure tensor semantics";
40
41 // New input operands for the cloned op.
42 SmallVector<Value> newInputBuffers;
43 newInputBuffers.reserve(op.getNumDpsInputs());
44 for (OpOperand *opOperand : op.getDpsInputOperands()) {
45 if (op.isScalar(opOperand)) {
46 newInputBuffers.push_back(opOperand->get());
47 continue;
48 }
49 FailureOr<Value> buffer =
50 getBuffer(rewriter, opOperand->get(), options, state);
51 if (failed(buffer))
52 return failure();
53 newInputBuffers.push_back(*buffer);
54 }
55
56 // New output operands for the cloned op.
57 SmallVector<Value> newOutputBuffers;
58 for (OpResult opResult : op->getOpResults()) {
59 OpOperand *opOperand = op.getDpsInitOperand(opResult.getResultNumber());
60 FailureOr<Value> resultBuffer =
61 getBuffer(rewriter, opOperand->get(), options, state);
62 if (failed(resultBuffer))
63 return failure();
64 newOutputBuffers.push_back(*resultBuffer);
65 }
66
67 // Merge input/output operands.
68 SmallVector<Value> newOperands = newInputBuffers;
69 newOperands.append(newOutputBuffers.begin(), newOutputBuffers.end());
70
71 // Set insertion point now that potential alloc/dealloc are introduced.
72 rewriter.setInsertionPoint(op);
73 // Clone the op, but use the new operands. Move the existing block into the
74 // new op. Since the new op does not have any tensor results, it does not
75 // return anything.
76 assert(op->getNumRegions() == 1 && "expected that op has 1 region");
78 op->getLoc(), op->getName(), TypeRange{}, newOperands,
79 op->getDiscardableAttrDictionary(), op->getPropertiesStorage(),
80 /*successors=*/{}, /*numRegions=*/1);
81 newOp->getRegion(0).getBlocks().splice(newOp->getRegion(0).begin(),
82 op->getRegion(0).getBlocks());
83
84 // We don't want the rewriter tracks an incomplete operation, so insert new
85 // operation after op was fully constructed.
86 rewriter.insert(newOp);
87
88 // Replace the results of the old op with the new output buffers.
89 replaceOpWithBufferizedValues(rewriter, op, newOutputBuffers);
90
91 return success();
92}
93
94/// Bufferization of linalg.generic. Replace with a new linalg.generic that
95/// operates entirely on memrefs.
96template <typename OpTy>
97struct LinalgOpInterface
98 : public DstBufferizableOpInterfaceExternalModel<LinalgOpInterface<OpTy>,
99 OpTy> {
100 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
101 const AnalysisState &state) const {
102 // Operand is read if it is used in the computation.
103 auto linalgOp = cast<linalg::LinalgOp>(op);
104 return linalgOp.payloadUsesValueFromOperand(&opOperand);
105 }
106
107 bool bufferizesToMemoryWrite(Operation *op, OpOperand &opOperand,
108 const AnalysisState &state) const {
109 // Operand is written to if it is not an input/init.
110 auto dpsOp = cast<DestinationStyleOpInterface>(op);
111 return dpsOp.isDpsInit(&opOperand);
112 }
113
114 bool bufferizesToElementwiseAccess(Operation *op, const AnalysisState &state,
115 ArrayRef<OpOperand *> opOperands) const {
116 auto linalgOp = cast<linalg::LinalgOp>(op);
117
118 // Accesses into sparse data structures are not necessarily elementwise.
120 return false;
121
122 // All loops must be parallel.
123 if (linalgOp.getNumLoops() != linalgOp.getNumParallelLoops())
124 return false;
125
126 // All indexing maps of participating tensors must be the same
127 // permutation.
128 SmallVector<AffineMap> indexingMaps = linalgOp.getIndexingMapsArray();
129 assert(linalgOp->getNumOperands() == indexingMaps.size() &&
130 "unexpected number of indexing maps");
131 AffineMap commonIndexingMap;
132 for (auto [operand, map] :
133 llvm::zip(linalgOp->getOpOperands(), indexingMaps)) {
134 // Non-tensors do not participate in bufferization, so they can be
135 // ignored.
136 if (!isa<RankedTensorType, MemRefType>(operand.get().getType()))
137 continue;
138 // Only consider operands in `opOperands`.
139 if (!llvm::is_contained(opOperands, &operand))
140 continue;
141 if (!map.isPermutation())
142 return false;
143 if (commonIndexingMap && commonIndexingMap != map)
144 return false;
145 commonIndexingMap = map;
146 }
147
148 return true;
149 }
150
151 LogicalResult bufferize(Operation *op, RewriterBase &rewriter,
152 const BufferizationOptions &options,
153 BufferizationState &state) const {
154 return bufferizeDestinationStyleOpInterface(
155 rewriter, cast<DestinationStyleOpInterface>(op), options, state);
156 }
157};
158
159/// Helper structure that iterates over all LinalgOps in `OpTys` and registers
160/// the `BufferizableOpInterface` with each of them.
161template <typename... Ops>
162struct LinalgOpInterfaceHelper {
163 static void registerOpInterface(MLIRContext *ctx) {
164 (Ops::template attachInterface<LinalgOpInterface<Ops>>(*ctx), ...);
165 }
166};
167
168struct SoftmaxOpInterface
169 : public DstBufferizableOpInterfaceExternalModel<SoftmaxOpInterface,
170 linalg::SoftmaxOp> {
171 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
172 const AnalysisState &state) const {
173 // Output operand is not read.
174 auto softmaxOp = cast<linalg::SoftmaxOp>(op);
175 return &opOperand == &softmaxOp.getInputMutable();
176 }
177
178 LogicalResult bufferize(Operation *op, RewriterBase &rewriter,
179 const BufferizationOptions &options,
180 BufferizationState &state) const {
181 auto softmaxOp = cast<linalg::SoftmaxOp>(op);
182 FailureOr<Value> inputBuffer =
183 getBuffer(rewriter, softmaxOp.getInput(), options, state);
184 if (failed(inputBuffer))
185 return failure();
186 FailureOr<Value> outputBuffer =
187 getBuffer(rewriter, softmaxOp.getOutput(), options, state);
188 if (failed(outputBuffer))
189 return failure();
190 linalg::SoftmaxOp::create(rewriter, softmaxOp.getLoc(),
191 /*result=*/TypeRange(), *inputBuffer,
192 *outputBuffer, softmaxOp.getDimension());
193 replaceOpWithBufferizedValues(rewriter, op, *outputBuffer);
194 return success();
195 }
196};
197
198struct PackOpInterface
199 : public DstBufferizableOpInterfaceExternalModel<PackOpInterface,
200 linalg::PackOp> {
201 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
202 const AnalysisState &state) const {
203 auto packOp = cast<linalg::PackOp>(op);
204 return !packOp.isDpsInit(&opOperand);
205 }
206
207 LogicalResult bufferize(Operation *op, RewriterBase &rewriter,
208 const BufferizationOptions &options,
209 BufferizationState &state) const {
210 auto packOp = cast<linalg::PackOp>(op);
211 assert(!packOp.hasPureBufferSemantics() && "expected op with tensors");
212 if (!packOp.hasPureTensorSemantics())
213 return packOp.emitError()
214 << "mixed tensor/buffer semantic op not supported yet";
215 FailureOr<Value> sourceBuffer =
216 getBuffer(rewriter, packOp.getSource(), options, state);
217 if (failed(sourceBuffer))
218 return failure();
219 FailureOr<Value> destBuffer =
220 getBuffer(rewriter, packOp.getDest(), options, state);
221 if (failed(destBuffer))
222 return failure();
223
224 SmallVector<Value> operands;
225 operands.push_back(*sourceBuffer);
226 operands.push_back(*destBuffer);
227 if (auto val = packOp.getPaddingValue())
228 operands.push_back(val);
229 llvm::append_range(operands, packOp.getInnerTiles());
230
231 linalg::PackOp::create(rewriter, packOp.getLoc(), TypeRange{}, operands,
232 packOp.getProperties(),
233 packOp->getDiscardableAttrDictionary().getValue());
234 replaceOpWithBufferizedValues(rewriter, op, *destBuffer);
235 return success();
236 }
237};
238
239struct UnPackOpInterface
240 : public DstBufferizableOpInterfaceExternalModel<UnPackOpInterface,
241 linalg::UnPackOp> {
242 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
243 const AnalysisState &state) const {
244 auto unPackOp = cast<linalg::UnPackOp>(op);
245 return !unPackOp.isDpsInit(&opOperand);
246 }
247
248 LogicalResult bufferize(Operation *op, RewriterBase &rewriter,
249 const BufferizationOptions &options,
250 BufferizationState &state) const {
251 auto unPackOp = cast<linalg::UnPackOp>(op);
252 assert(!unPackOp.hasPureBufferSemantics() && "expected op with tensors");
253 if (!unPackOp.hasPureTensorSemantics())
254 return unPackOp.emitError()
255 << "mixed tensor/buffer semantic op not supported yet";
256 FailureOr<Value> sourceBuffer =
257 getBuffer(rewriter, unPackOp.getSource(), options, state);
258 if (failed(sourceBuffer))
259 return failure();
260 FailureOr<Value> destBuffer =
261 getBuffer(rewriter, unPackOp.getDest(), options, state);
262 if (failed(destBuffer))
263 return failure();
264
265 SmallVector<Value> operands;
266 operands.push_back(*sourceBuffer);
267 operands.push_back(*destBuffer);
268 llvm::append_range(operands, unPackOp.getInnerTiles());
269
270 linalg::UnPackOp::create(
271 rewriter, unPackOp.getLoc(), TypeRange{}, operands,
272 unPackOp.getProperties(),
273 unPackOp->getDiscardableAttrDictionary().getValue());
274 replaceOpWithBufferizedValues(rewriter, op, *destBuffer);
275 return success();
276 }
277};
278} // namespace
279
281 DialectRegistry &registry) {
282 registry.addExtension(+[](MLIRContext *ctx, linalg::LinalgDialect *dialect) {
283 // Register all Linalg structured ops. `LinalgOp` is an interface and it is
284 // not possible to attach an external interface to an existing interface.
285 // Therefore, attach the `BufferizableOpInterface` to all ops one-by-one.
286 LinalgOpInterfaceHelper<
287#define GET_OP_LIST
288#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
289
290 >::registerOpInterface(ctx);
291
292 SoftmaxOp::attachInterface<SoftmaxOpInterface>(*ctx);
293 PackOp::attachInterface<PackOpInterface>(*ctx);
294 UnPackOp::attachInterface<UnPackOpInterface>(*ctx);
295 });
296}
return success()
static llvm::ManagedStatic< PassManagerOptions > options
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
IRValueT get() const
Return the current value being used by this operand.
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
Operation * insert(Operation *op)
Insert the given operation at the current insertion point and return it.
Definition Builders.cpp:430
This class represents an operand of an operation.
Definition Value.h:254
This is a value defined by a result of an operation.
Definition Value.h:454
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
static Operation * create(Location location, OperationName name, TypeRange resultTypes, ValueRange operands, NamedAttrList &&attributes, PropertyRef properties, BlockRange successors, unsigned numRegions)
Create a new Operation with the specific fields.
Definition Operation.cpp:65
iterator begin()
Definition Region.h:55
BlockListType & getBlocks()
Definition Region.h:45
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
void registerBufferizableOpInterfaceExternalModels(DialectRegistry &registry)
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
bool hasAnySparseOperand(Operation *op)
Returns true iff MLIR operand has any sparse operand.
Include the generated interface declarations.
Bufferizable ops that implement the DestinationStyleOpInterface can use this external model base clas...