25static LogicalResult bufferizeDestinationStyleOpInterface(
33 if (op.hasPureBufferSemantics())
38 if (!op.hasPureTensorSemantics())
39 return op->emitError() <<
"op does not have pure tensor semantics";
43 newInputBuffers.reserve(op.getNumDpsInputs());
44 for (
OpOperand *opOperand : op.getDpsInputOperands()) {
45 if (op.isScalar(opOperand)) {
46 newInputBuffers.push_back(opOperand->get());
49 FailureOr<Value> buffer =
50 getBuffer(rewriter, opOperand->get(),
options, state);
53 newInputBuffers.push_back(*buffer);
58 for (
OpResult opResult : op->getOpResults()) {
59 OpOperand *opOperand = op.getDpsInitOperand(opResult.getResultNumber());
60 FailureOr<Value> resultBuffer =
61 getBuffer(rewriter, opOperand->
get(),
options, state);
64 newOutputBuffers.push_back(*resultBuffer);
69 newOperands.append(newOutputBuffers.begin(), newOutputBuffers.end());
76 assert(op->getNumRegions() == 1 &&
"expected that op has 1 region");
78 op->getLoc(), op->getName(),
TypeRange{}, newOperands,
79 op->getDiscardableAttrDictionary(), op->getPropertiesStorage(),
82 op->getRegion(0).getBlocks());
89 replaceOpWithBufferizedValues(rewriter, op, newOutputBuffers);
96template <
typename OpTy>
97struct LinalgOpInterface
100 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
101 const AnalysisState &state)
const {
103 auto linalgOp = cast<linalg::LinalgOp>(op);
104 return linalgOp.payloadUsesValueFromOperand(&opOperand);
107 bool bufferizesToMemoryWrite(Operation *op, OpOperand &opOperand,
108 const AnalysisState &state)
const {
110 auto dpsOp = cast<DestinationStyleOpInterface>(op);
111 return dpsOp.isDpsInit(&opOperand);
114 bool bufferizesToElementwiseAccess(Operation *op,
const AnalysisState &state,
115 ArrayRef<OpOperand *> opOperands)
const {
116 auto linalgOp = cast<linalg::LinalgOp>(op);
123 if (linalgOp.getNumLoops() != linalgOp.getNumParallelLoops())
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)) {
136 if (!isa<RankedTensorType, MemRefType>(operand.get().getType()))
139 if (!llvm::is_contained(opOperands, &operand))
141 if (!map.isPermutation())
143 if (commonIndexingMap && commonIndexingMap != map)
145 commonIndexingMap = map;
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);
161template <
typename... Ops>
162struct LinalgOpInterfaceHelper {
163 static void registerOpInterface(MLIRContext *ctx) {
164 (Ops::template attachInterface<LinalgOpInterface<Ops>>(*ctx), ...);
168struct SoftmaxOpInterface
171 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
172 const AnalysisState &state)
const {
174 auto softmaxOp = cast<linalg::SoftmaxOp>(op);
175 return &opOperand == &softmaxOp.getInputMutable();
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);
186 FailureOr<Value> outputBuffer =
187 getBuffer(rewriter, softmaxOp.getOutput(),
options, state);
190 linalg::SoftmaxOp::create(rewriter, softmaxOp.getLoc(),
192 *outputBuffer, softmaxOp.getDimension());
193 replaceOpWithBufferizedValues(rewriter, op, *outputBuffer);
198struct PackOpInterface
201 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
202 const AnalysisState &state)
const {
203 auto packOp = cast<linalg::PackOp>(op);
204 return !packOp.isDpsInit(&opOperand);
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);
219 FailureOr<Value> destBuffer =
220 getBuffer(rewriter, packOp.getDest(),
options, state);
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());
231 linalg::PackOp::create(rewriter, packOp.getLoc(),
TypeRange{}, operands,
232 packOp.getProperties(),
233 packOp->getDiscardableAttrDictionary().getValue());
234 replaceOpWithBufferizedValues(rewriter, op, *destBuffer);
239struct UnPackOpInterface
242 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,
243 const AnalysisState &state)
const {
244 auto unPackOp = cast<linalg::UnPackOp>(op);
245 return !unPackOp.isDpsInit(&opOperand);
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);
260 FailureOr<Value> destBuffer =
261 getBuffer(rewriter, unPackOp.getDest(),
options, state);
265 SmallVector<Value> operands;
266 operands.push_back(*sourceBuffer);
267 operands.push_back(*destBuffer);
268 llvm::append_range(operands, unPackOp.getInnerTiles());
270 linalg::UnPackOp::create(
271 rewriter, unPackOp.getLoc(),
TypeRange{}, operands,
272 unPackOp.getProperties(),
273 unPackOp->getDiscardableAttrDictionary().getValue());
274 replaceOpWithBufferizedValues(rewriter, op, *destBuffer);
286 LinalgOpInterfaceHelper<
288#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
290 >::registerOpInterface(ctx);
292 SoftmaxOp::attachInterface<SoftmaxOpInterface>(*ctx);
293 PackOp::attachInterface<PackOpInterface>(*ctx);
294 UnPackOp::attachInterface<UnPackOpInterface>(*ctx);
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.
RAII guard to reset the insertion point of the builder when destroyed.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Operation * insert(Operation *op)
Insert the given operation at the current insertion point and return it.
This class represents an operand of an operation.
This is a value defined by a result of an operation.
Operation is the basic unit of execution within MLIR.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
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.
BlockListType & getBlocks()
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.
void registerBufferizableOpInterfaceExternalModels(DialectRegistry ®istry)
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...