brintos

brintos / llvm-project-archived public Read only

0
0
Text · 5.7 KiB · f6bc225 Raw
146 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/Shape/Transforms/BufferizableOpInterfaceImpl.h"10 11#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h"12#include "mlir/Dialect/Bufferization/IR/Bufferization.h"13#include "mlir/Dialect/Shape/IR/Shape.h"14#include "mlir/IR/Operation.h"15#include "mlir/IR/PatternMatch.h"16 17using namespace mlir;18using namespace mlir::bufferization;19using namespace mlir::shape;20 21namespace mlir {22namespace shape {23namespace {24 25/// Bufferization of shape.assuming.26struct AssumingOpInterface27    : public BufferizableOpInterface::ExternalModel<AssumingOpInterface,28                                                    shape::AssumingOp> {29  AliasingOpOperandList30  getAliasingOpOperands(Operation *op, Value value,31                        const AnalysisState &state) const {32    // AssumingOps do not have tensor OpOperands. The yielded value can be any33    // SSA value that is in scope. To allow for use-def chain traversal through34    // AssumingOps in the analysis, the corresponding yield value is considered35    // to be aliasing with the result.36    auto assumingOp = cast<shape::AssumingOp>(op);37    size_t resultNum = std::distance(op->getOpResults().begin(),38                                     llvm::find(op->getOpResults(), value));39    // TODO: Support multiple blocks.40    assert(assumingOp.getDoRegion().hasOneBlock() &&41           "expected exactly 1 block");42    auto yieldOp = dyn_cast<shape::AssumingYieldOp>(43        assumingOp.getDoRegion().front().getTerminator());44    assert(yieldOp && "expected shape.assuming_yield terminator");45    return {{&yieldOp->getOpOperand(resultNum), BufferRelation::Equivalent}};46  }47 48  LogicalResult bufferize(Operation *op, RewriterBase &rewriter,49                          const BufferizationOptions &options,50                          BufferizationState &state) const {51    auto assumingOp = cast<shape::AssumingOp>(op);52    assert(assumingOp.getDoRegion().hasOneBlock() && "only 1 block supported");53    auto yieldOp = cast<shape::AssumingYieldOp>(54        assumingOp.getDoRegion().front().getTerminator());55 56    // Create new op and move over region.57    TypeRange newResultTypes(yieldOp.getOperands());58    auto newOp = shape::AssumingOp::create(59        rewriter, op->getLoc(), newResultTypes, assumingOp.getWitness());60    newOp.getDoRegion().takeBody(assumingOp.getRegion());61 62    // Update all uses of the old op.63    rewriter.setInsertionPointAfter(newOp);64    SmallVector<Value> newResults;65    for (const auto &it : llvm::enumerate(assumingOp->getResultTypes())) {66      if (isa<TensorType>(it.value())) {67        newResults.push_back(bufferization::ToTensorOp::create(68            rewriter, assumingOp.getLoc(), it.value(),69            newOp->getResult(it.index())));70      } else {71        newResults.push_back(newOp->getResult(it.index()));72      }73    }74 75    // Replace old op.76    rewriter.replaceOp(assumingOp, newResults);77 78    return success();79  }80};81 82/// Bufferization of shape.assuming_yield. Bufferized as part of their enclosing83/// ops, so this is for analysis only.84struct AssumingYieldOpInterface85    : public BufferizableOpInterface::ExternalModel<AssumingYieldOpInterface,86                                                    shape::AssumingYieldOp> {87  bool bufferizesToMemoryRead(Operation *op, OpOperand &opOperand,88                              const AnalysisState &state) const {89    return true;90  }91 92  bool bufferizesToMemoryWrite(Operation *op, OpOperand &opOperand,93                               const AnalysisState &state) const {94    return false;95  }96 97  AliasingValueList getAliasingValues(Operation *op, OpOperand &opOperand,98                                      const AnalysisState &state) const {99    assert(isa<shape::AssumingOp>(op->getParentOp()) &&100           "expected that parent is an AssumingOp");101    OpResult opResult =102        op->getParentOp()->getResult(opOperand.getOperandNumber());103    return {{opResult, BufferRelation::Equivalent}};104  }105 106  bool mustBufferizeInPlace(Operation *op, OpOperand &opOperand,107                            const AnalysisState &state) const {108    // Yield operands always bufferize inplace. Otherwise, an alloc + copy109    // may be generated inside the block. We should not return/yield allocations110    // when possible.111    return true;112  }113 114  LogicalResult bufferize(Operation *op, RewriterBase &rewriter,115                          const BufferizationOptions &options,116                          BufferizationState &state) const {117    auto yieldOp = cast<shape::AssumingYieldOp>(op);118    SmallVector<Value> newResults;119    for (Value value : yieldOp.getOperands()) {120      if (isa<TensorType>(value.getType())) {121        FailureOr<Value> buffer = getBuffer(rewriter, value, options, state);122        if (failed(buffer))123          return failure();124        newResults.push_back(*buffer);125      } else {126        newResults.push_back(value);127      }128    }129    replaceOpWithNewBufferizedOp<shape::AssumingYieldOp>(rewriter, op,130                                                         newResults);131    return success();132  }133};134 135} // namespace136} // namespace shape137} // namespace mlir138 139void mlir::shape::registerBufferizableOpInterfaceExternalModels(140    DialectRegistry &registry) {141  registry.addExtension(+[](MLIRContext *ctx, shape::ShapeDialect *dialect) {142    shape::AssumingOp::attachInterface<AssumingOpInterface>(*ctx);143    shape::AssumingYieldOp::attachInterface<AssumingYieldOpInterface>(*ctx);144  });145}146