brintos

brintos / llvm-project-archived public Read only

0
0
Text · 5.1 KiB · 9a16318 Raw
136 lines · cpp
1//===- BubbleUpExtractSlice.cpp - bubble up tensor.extract_slice ----------===//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// This file implements patterns that transforms linalg.<op> +10// tensor.extract_slice into tensor.extract_slice + linalg.<op> to reduce11// the computation for the linalg op.12//13//===----------------------------------------------------------------------===//14 15#include "mlir/Dialect/Affine/IR/AffineOps.h"16#include "mlir/Dialect/Linalg/IR/Linalg.h"17#include "mlir/Dialect/Linalg/Transforms/Transforms.h"18#include "mlir/Dialect/Linalg/Utils/Utils.h"19 20using namespace mlir;21using namespace mlir::linalg;22 23namespace {24/// Bubble up extract_slice above Linalg operation.25///26/// A sequence of operations27///28/// ```mlir29/// %0 = linalg.<op> ... arg0, arg1, ...30/// %1 = tensor.extract_slice %0 ...31/// ```32///33/// can be replaced with34///35/// ```mlir36/// %0 = tensor.extract_slice %arg037/// %1 = tensor.extract_slice %arg138/// %2 = linalg.<op> ... %0, %1, ...39/// ```40///41/// This results in the reduce computation of the linalg operation.42///43struct BubbleUpExtractSliceOpPattern44    : OpRewritePattern<tensor::ExtractSliceOp> {45  using OpRewritePattern<tensor::ExtractSliceOp>::OpRewritePattern;46 47  LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp,48                                PatternRewriter &rewriter) const final {49    Value source = sliceOp.getSource();50    auto linalgOp = source.getDefiningOp<LinalgOp>();51    if (!linalgOp) {52      return rewriter.notifyMatchFailure(sliceOp,53                                         "expected source to be linalg op");54    }55 56    // TODO: we might relax this if we want heuristics to detect that all uses57    // are small portion of the output.58    if (!linalgOp->hasOneUse()) {59      return rewriter.notifyMatchFailure(sliceOp,60                                         "expected single use of linalg op");61    }62 63    if (linalgOp.getNumDpsInits() != 1) {64      return rewriter.notifyMatchFailure(sliceOp,65                                         "expected single output of linalg op");66    }67 68    if (!linalgOp.hasPureTensorSemantics()) {69      return rewriter.notifyMatchFailure(sliceOp,70                                         "expected tensor of linalg op");71    }72 73    if (!sliceOp.hasUnitStride())74      return rewriter.notifyMatchFailure(sliceOp, "expected unit stride");75 76    if (sliceOp.getType().getRank() != sliceOp.getSourceType().getRank()) {77      return rewriter.notifyMatchFailure(sliceOp, "expected no rank reduction");78    }79 80    OpOperand *outOperand = linalgOp.getDpsInitOperand(0);81    AffineMap indexingMap = linalgOp.getMatchingIndexingMap(outOperand);82    if (!indexingMap.isProjectedPermutation()) {83      return rewriter.notifyMatchFailure(84          sliceOp, "expected a projected permutation for output");85    }86 87    auto linalgLoc = linalgOp.getLoc();88    SmallVector<OpFoldResult> allShapeSizes =89        linalgOp.createFlatListOfOperandDims(rewriter, linalgLoc);90    AffineMap shapeSizesToLoopsMap = linalgOp.getShapesToLoopsMap();91    if (!shapeSizesToLoopsMap) {92      return rewriter.notifyMatchFailure(93          linalgOp, "failed to get loops map from shape sizes");94    }95    SmallVector<OpFoldResult> sizeBounds =96        affine::makeComposedFoldedMultiResultAffineApply(97            rewriter, linalgLoc, shapeSizesToLoopsMap, allShapeSizes);98 99    // The offsets and sizes from the slice operation only give you the tile100    // size of the output. Use that compute the tile sizes and offsets of the101    // loops. For loops not used to access the output, set the tile sizes to102    // loop bounds and set the offset to 0.103    SmallVector<OpFoldResult> tileOffsets(sizeBounds.size(),104                                          rewriter.getIndexAttr(0));105    SmallVector<OpFoldResult> tileSizes = sizeBounds;106    for (auto const &result : enumerate(indexingMap.getResults())) {107      unsigned position = cast<AffineDimExpr>(result.value()).getPosition();108      tileOffsets[position] = sliceOp.getMixedOffsets()[result.index()];109      tileSizes[position] = sliceOp.getMixedSizes()[result.index()];110    }111 112    SmallVector<Value> valuesToTile = linalgOp->getOperands();113    SmallVector<Value> tiledOperands =114        makeTiledShapes(rewriter, linalgLoc, linalgOp, valuesToTile,115                        tileOffsets, tileSizes, sizeBounds,116                        /*omitPartialTileCheck=*/true);117 118    SmallVector<Type, 4> resultTensorTypes;119    for (OpOperand &opOperand : linalgOp.getDpsInitsMutable())120      resultTensorTypes.push_back(121          tiledOperands[opOperand.getOperandNumber()].getType());122 123    Operation *newOp =124        clone(rewriter, linalgOp, resultTensorTypes, tiledOperands);125    rewriter.replaceOp(sliceOp, newOp->getResults());126    return success();127  }128};129} // namespace130 131void mlir::linalg::populateBubbleUpExtractSliceOpPatterns(132    RewritePatternSet &patterns) {133  auto *context = patterns.getContext();134  patterns.add<BubbleUpExtractSliceOpPattern>(context);135}136