brintos

brintos / llvm-project-archived public Read only

0
0
Text · 4.1 KiB · b6fdba3 Raw
106 lines · cpp
1//===- SubsetInsertionOpInterfaceImpl.cpp - Tensor subsets ----------------===//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/Tensor/Transforms/SubsetInsertionOpInterfaceImpl.h"10 11#include "mlir/Dialect/Tensor/IR/Tensor.h"12#include "mlir/Interfaces/SubsetOpInterface.h"13#include "mlir/Interfaces/ValueBoundsOpInterface.h"14 15using namespace mlir;16using namespace mlir::tensor;17 18namespace {19 20struct ExtractSliceOpSubsetOpInterface21    : public SubsetOpInterface::ExternalModel<ExtractSliceOpSubsetOpInterface,22                                              tensor::ExtractSliceOp> {23  FailureOr<HyperrectangularSlice>24  getAccessedHyperrectangularSlice(Operation *op) const {25    return HyperrectangularSlice(cast<OffsetSizeAndStrideOpInterface>(op));26  }27};28 29struct ExtractSliceOpSubsetExtractionOpInterface30    : public SubsetExtractionOpInterface::ExternalModel<31          ExtractSliceOpSubsetExtractionOpInterface, tensor::ExtractSliceOp> {32  OpOperand &getSourceOperand(Operation *op) const {33    return cast<tensor::ExtractSliceOp>(op).getSourceMutable();34  }35};36 37template <typename OpTy>38struct InsertSliceLikeOpSubsetOpInterface39    : public SubsetOpInterface::ExternalModel<40          InsertSliceLikeOpSubsetOpInterface<OpTy>, OpTy> {41  FailureOr<HyperrectangularSlice>42  getAccessedHyperrectangularSlice(Operation *op) const {43    return HyperrectangularSlice(cast<OffsetSizeAndStrideOpInterface>(op));44  }45};46 47template <typename OpTy>48struct InsertSliceLikeOpSubsetInsertionOpInterface49    : public SubsetInsertionOpInterface::ExternalModel<50          InsertSliceLikeOpSubsetInsertionOpInterface<OpTy>, OpTy> {51  OpOperand &getSourceOperand(Operation *op) const {52    return cast<OpTy>(op).getSourceMutable();53  }54 55  OpOperand &getDestinationOperand(Operation *op) const {56    return cast<OpTy>(op).getDestMutable();57  }58 59  Value buildSubsetExtraction(Operation *op, OpBuilder &builder,60                              Location loc) const {61    auto insertSliceOp = cast<OpTy>(op);62    auto extractOp = tensor::ExtractSliceOp::create(63        builder, loc, insertSliceOp.getSourceType(), insertSliceOp.getDest(),64        insertSliceOp.getMixedOffsets(), insertSliceOp.getMixedSizes(),65        insertSliceOp.getMixedStrides());66    return extractOp.getResult();67  }68 69  SmallVector<Value>70  getValuesNeededToBuildSubsetExtraction(Operation *op) const {71    auto insertSliceOp = cast<OpTy>(op);72    SmallVector<Value> neededValues;73    // Collect all values that are needed to construct the replacement op.74    neededValues.append(insertSliceOp.getOffsets().begin(),75                        insertSliceOp.getOffsets().end());76    neededValues.append(insertSliceOp.getSizes().begin(),77                        insertSliceOp.getSizes().end());78    neededValues.append(insertSliceOp.getStrides().begin(),79                        insertSliceOp.getStrides().end());80    neededValues.push_back(insertSliceOp.getDest());81    return neededValues;82  }83};84 85} // namespace86 87void mlir::tensor::registerSubsetOpInterfaceExternalModels(88    DialectRegistry &registry) {89  registry.addExtension(+[](MLIRContext *ctx, tensor::TensorDialect *dialect) {90    // Note: `SubsetExtractionOpInterface` and `SubsetInsertionOpInterface`91    // require `SubsetOpInterface`.92    ExtractSliceOp::attachInterface<ExtractSliceOpSubsetOpInterface>(*ctx);93    ExtractSliceOp::attachInterface<ExtractSliceOpSubsetExtractionOpInterface>(94        *ctx);95    InsertSliceOp::attachInterface<96        InsertSliceLikeOpSubsetOpInterface<InsertSliceOp>>(*ctx);97    InsertSliceOp::attachInterface<98        InsertSliceLikeOpSubsetInsertionOpInterface<InsertSliceOp>>(*ctx);99    ParallelInsertSliceOp::attachInterface<100        InsertSliceLikeOpSubsetOpInterface<ParallelInsertSliceOp>>(*ctx);101    ParallelInsertSliceOp::attachInterface<102        InsertSliceLikeOpSubsetInsertionOpInterface<ParallelInsertSliceOp>>(103        *ctx);104  });105}106