211 lines · cpp
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-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Dialect/Linalg/Transforms/BufferizableOpInterfaceImpl.h"10#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h"11#include "mlir/Dialect/Bufferization/IR/DstBufferizableOpInterfaceImpl.h"12#include "mlir/Dialect/Linalg/IR/Linalg.h"13#include "mlir/Dialect/SparseTensor/IR/SparseTensor.h"14#include "mlir/IR/Dialect.h"15#include "mlir/IR/Operation.h"16#include "mlir/Interfaces/DestinationStyleOpInterface.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-need37 // 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 the74 // new op. Since the new op does not have any tensor results, it does not75 // return anything.76 assert(op->getNumRegions() == 1 && "expected that op has 1 region");77 OperationState opState(op->getLoc(), op->getName(), newOperands, TypeRange{},78 op->getAttrs());79 opState.addRegion();80 Operation *newOp = Operation::create(opState);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 new85 // 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 that95/// operates entirely on memrefs.96template <typename OpTy>97struct LinalgOpInterface98 : 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.119 if (sparse_tensor::hasAnySparseOperand(linalgOp))120 return false;121 122 // All loops must be parallel.123 if (linalgOp.getNumLoops() != linalgOp.getNumParallelLoops())124 return false;125 126 // All index maps of tensors must be identity maps.127 SmallVector<AffineMap> indexingMaps = linalgOp.getIndexingMapsArray();128 assert(linalgOp->getNumOperands() == indexingMaps.size() &&129 "unexpected number of indexing maps");130 for (auto [operand, map] :131 llvm::zip(linalgOp->getOpOperands(), indexingMaps)) {132 // Non-tensors do not participate in bufferization, so they can be133 // ignored.134 if (!isa<RankedTensorType, MemRefType>(operand.get().getType()))135 continue;136 // Only consider operands in `opOperands`.137 if (!llvm::is_contained(opOperands, &operand))138 continue;139 // TODO: This could be generalized to other indexing maps. (All indexing140 // must be the same.)141 if (!map.isIdentity())142 return false;143 }144 145 return true;146 }147 148 LogicalResult bufferize(Operation *op, RewriterBase &rewriter,149 const BufferizationOptions &options,150 BufferizationState &state) const {151 return bufferizeDestinationStyleOpInterface(152 rewriter, cast<DestinationStyleOpInterface>(op), options, state);153 }154};155 156/// Helper structure that iterates over all LinalgOps in `OpTys` and registers157/// the `BufferizableOpInterface` with each of them.158template <typename... Ops>159struct LinalgOpInterfaceHelper {160 static void registerOpInterface(MLIRContext *ctx) {161 (Ops::template attachInterface<LinalgOpInterface<Ops>>(*ctx), ...);162 }163};164 165struct SoftmaxOpInterface166 : public DstBufferizableOpInterfaceExternalModel<SoftmaxOpInterface,167 linalg::SoftmaxOp> {168 bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,169 const AnalysisState &state) const {170 // Output operand is not read.171 auto softmaxOp = cast<linalg::SoftmaxOp>(op);172 return &opOperand == &softmaxOp.getInputMutable();173 }174 175 LogicalResult bufferize(Operation *op, RewriterBase &rewriter,176 const BufferizationOptions &options,177 BufferizationState &state) const {178 auto softmaxOp = cast<linalg::SoftmaxOp>(op);179 FailureOr<Value> inputBuffer =180 getBuffer(rewriter, softmaxOp.getInput(), options, state);181 if (failed(inputBuffer))182 return failure();183 FailureOr<Value> outputBuffer =184 getBuffer(rewriter, softmaxOp.getOutput(), options, state);185 if (failed(outputBuffer))186 return failure();187 linalg::SoftmaxOp::create(rewriter, softmaxOp.getLoc(),188 /*result=*/TypeRange(), *inputBuffer,189 *outputBuffer, softmaxOp.getDimension());190 replaceOpWithBufferizedValues(rewriter, op, *outputBuffer);191 return success();192 }193};194} // namespace195 196void mlir::linalg::registerBufferizableOpInterfaceExternalModels(197 DialectRegistry ®istry) {198 registry.addExtension(+[](MLIRContext *ctx, linalg::LinalgDialect *dialect) {199 // Register all Linalg structured ops. `LinalgOp` is an interface and it is200 // not possible to attach an external interface to an existing interface.201 // Therefore, attach the `BufferizableOpInterface` to all ops one-by-one.202 LinalgOpInterfaceHelper<203#define GET_OP_LIST204#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"205 206 >::registerOpInterface(ctx);207 208 SoftmaxOp::attachInterface<SoftmaxOpInterface>(*ctx);209 });210}211