brintos

brintos / llvm-project-archived public Read only

0
0
Text · 19.1 KiB · c307fb4 Raw
473 lines · cpp
1//===- VectorUtils.cpp - MLIR Utilities for VectorOps   ------------------===//2//3// Part of the MLIR 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 utility methods for working with the Vector dialect.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Dialect/Vector/Utils/VectorUtils.h"14 15#include "mlir/Dialect/Affine/Analysis/LoopAnalysis.h"16#include "mlir/Dialect/Affine/IR/AffineOps.h"17#include "mlir/Dialect/Arith/IR/Arith.h"18#include "mlir/Dialect/Func/IR/FuncOps.h"19#include "mlir/Dialect/MemRef/IR/MemRef.h"20#include "mlir/Dialect/Tensor/IR/Tensor.h"21#include "mlir/Dialect/Utils/IndexingUtils.h"22#include "mlir/Dialect/Vector/IR/VectorOps.h"23#include "mlir/IR/Builders.h"24#include "mlir/IR/IntegerSet.h"25#include "mlir/IR/Operation.h"26#include "mlir/IR/TypeUtilities.h"27#include "mlir/Support/LLVM.h"28 29#include "llvm/ADT/DenseSet.h"30#include "llvm/Support/DebugLog.h"31#include "llvm/Support/InterleavedRange.h"32 33#define DEBUG_TYPE "vector-utils"34 35using namespace mlir;36 37/// Helper function that creates a memref::DimOp or tensor::DimOp depending on38/// the type of `source`.39Value mlir::vector::createOrFoldDimOp(OpBuilder &b, Location loc, Value source,40                                      int64_t dim) {41  if (isa<UnrankedMemRefType, MemRefType>(source.getType()))42    return b.createOrFold<memref::DimOp>(loc, source, dim);43  if (isa<UnrankedTensorType, RankedTensorType>(source.getType()))44    return b.createOrFold<tensor::DimOp>(loc, source, dim);45  llvm_unreachable("Expected MemRefType or TensorType");46}47 48/// Given the n-D transpose pattern 'transp', return true if 'dim0' and 'dim1'49/// should be transposed with each other within the context of their 2D50/// transposition slice.51///52/// Example 1: dim0 = 0, dim1 = 2, transp = [2, 1, 0]53///   Return true: dim0 and dim1 are transposed within the context of their 2D54///   transposition slice ([1, 0]).55///56/// Example 2: dim0 = 0, dim1 = 1, transp = [2, 1, 0]57///   Return true: dim0 and dim1 are transposed within the context of their 2D58///   transposition slice ([1, 0]). Paradoxically, note how dim1 (1) is *not*59///   transposed within the full context of the transposition.60///61/// Example 3: dim0 = 0, dim1 = 1, transp = [2, 0, 1]62///   Return false: dim0 and dim1 are *not* transposed within the context of63///   their 2D transposition slice ([0, 1]). Paradoxically, note how dim0 (0)64///   and dim1 (1) are transposed within the full context of the of the65///   transposition.66static bool areDimsTransposedIn2DSlice(int64_t dim0, int64_t dim1,67                                       ArrayRef<int64_t> transp) {68  // Perform a linear scan along the dimensions of the transposed pattern. If69  // dim0 is found first, dim0 and dim1 are not transposed within the context of70  // their 2D slice. Otherwise, 'dim1' is found first and they are transposed.71  for (int64_t permDim : transp) {72    if (permDim == dim0)73      return false;74    if (permDim == dim1)75      return true;76  }77 78  llvm_unreachable("Ill-formed transpose pattern");79}80 81FailureOr<std::pair<int, int>>82mlir::vector::isTranspose2DSlice(vector::TransposeOp op) {83  VectorType srcType = op.getSourceVectorType();84  SmallVector<int64_t> srcGtOneDims;85  for (auto [index, size] : llvm::enumerate(srcType.getShape()))86    if (size > 1)87      srcGtOneDims.push_back(index);88 89  if (srcGtOneDims.size() != 2)90    return failure();91 92  // Check whether the two source vector dimensions that are greater than one93  // must be transposed with each other so that we can apply one of the 2-D94  // transpose patterns. Otherwise, these patterns are not applicable.95  if (!areDimsTransposedIn2DSlice(srcGtOneDims[0], srcGtOneDims[1],96                                  op.getPermutation()))97    return failure();98 99  return std::pair<int, int>(srcGtOneDims[0], srcGtOneDims[1]);100}101 102/// Constructs a permutation map from memref indices to vector dimension.103///104/// The implementation uses the knowledge of the mapping of enclosing loop to105/// vector dimension. `enclosingLoopToVectorDim` carries this information as a106/// map with:107///   - keys representing "vectorized enclosing loops";108///   - values representing the corresponding vector dimension.109/// The algorithm traverses "vectorized enclosing loops" and extracts the110/// at-most-one MemRef index that is invariant along said loop. This index is111/// guaranteed to be at most one by construction: otherwise the MemRef is not112/// vectorizable.113/// If this invariant index is found, it is added to the permutation_map at the114/// proper vector dimension.115/// If no index is found to be invariant, 0 is added to the permutation_map and116/// corresponds to a vector broadcast along that dimension.117///118/// Returns an empty AffineMap if `enclosingLoopToVectorDim` is empty,119/// signalling that no permutation map can be constructed given120/// `enclosingLoopToVectorDim`.121///122/// Examples can be found in the documentation of `makePermutationMap`, in the123/// header file.124static AffineMap makePermutationMap(125    ArrayRef<Value> indices,126    const DenseMap<Operation *, unsigned> &enclosingLoopToVectorDim) {127  if (enclosingLoopToVectorDim.empty())128    return AffineMap();129  MLIRContext *context =130      enclosingLoopToVectorDim.begin()->getFirst()->getContext();131  SmallVector<AffineExpr> perm(enclosingLoopToVectorDim.size(),132                               getAffineConstantExpr(0, context));133 134  for (auto kvp : enclosingLoopToVectorDim) {135    assert(kvp.second < perm.size());136    auto invariants = affine::getInvariantAccesses(137        cast<affine::AffineForOp>(kvp.first).getInductionVar(), indices);138    unsigned numIndices = indices.size();139    unsigned countInvariantIndices = 0;140    for (unsigned dim = 0; dim < numIndices; ++dim) {141      if (!invariants.count(indices[dim])) {142        assert(perm[kvp.second] == getAffineConstantExpr(0, context) &&143               "permutationMap already has an entry along dim");144        perm[kvp.second] = getAffineDimExpr(dim, context);145      } else {146        ++countInvariantIndices;147      }148    }149    assert((countInvariantIndices == numIndices ||150            countInvariantIndices == numIndices - 1) &&151           "Vectorization prerequisite violated: at most 1 index may be "152           "invariant wrt a vectorized loop");153    (void)countInvariantIndices;154  }155  return AffineMap::get(indices.size(), 0, perm, context);156}157 158/// Implementation detail that walks up the parents and records the ones with159/// the specified type.160/// TODO: could also be implemented as a collect parents followed by a161/// filter and made available outside this file.162template <typename T>163static SetVector<Operation *> getParentsOfType(Block *block) {164  SetVector<Operation *> res;165  auto *current = block->getParentOp();166  while (current) {167    if ([[maybe_unused]] auto typedParent = dyn_cast<T>(current)) {168      assert(res.count(current) == 0 && "Already inserted");169      res.insert(current);170    }171    current = current->getParentOp();172  }173  return res;174}175 176/// Returns the enclosing AffineForOp, from closest to farthest.177static SetVector<Operation *> getEnclosingforOps(Block *block) {178  return getParentsOfType<affine::AffineForOp>(block);179}180 181AffineMap mlir::makePermutationMap(182    Block *insertPoint, ArrayRef<Value> indices,183    const DenseMap<Operation *, unsigned> &loopToVectorDim) {184  DenseMap<Operation *, unsigned> enclosingLoopToVectorDim;185  auto enclosingLoops = getEnclosingforOps(insertPoint);186  for (auto *forInst : enclosingLoops) {187    auto it = loopToVectorDim.find(forInst);188    if (it != loopToVectorDim.end()) {189      enclosingLoopToVectorDim.insert(*it);190    }191  }192  return ::makePermutationMap(indices, enclosingLoopToVectorDim);193}194 195AffineMap mlir::makePermutationMap(196    Operation *op, ArrayRef<Value> indices,197    const DenseMap<Operation *, unsigned> &loopToVectorDim) {198  return makePermutationMap(op->getBlock(), indices, loopToVectorDim);199}200 201bool matcher::operatesOnSuperVectorsOf(Operation &op,202                                       VectorType subVectorType) {203  // First, extract the vector type and distinguish between:204  //   a. ops that *must* lower a super-vector (i.e. vector.transfer_read,205  //      vector.transfer_write); and206  //   b. ops that *may* lower a super-vector (all other ops).207  // The ops that *may* lower a super-vector only do so if the super-vector to208  // sub-vector ratio exists. The ops that *must* lower a super-vector are209  // explicitly checked for this property.210  /// TODO: there should be a single function for all ops to do this so we211  /// do not have to special case. Maybe a trait, or just a method, unclear atm.212  bool mustDivide = false;213  (void)mustDivide;214  VectorType superVectorType;215  if (auto transfer = dyn_cast<VectorTransferOpInterface>(op)) {216    superVectorType = transfer.getVectorType();217    mustDivide = true;218  } else if (op.getNumResults() == 0) {219    if (!isa<func::ReturnOp>(op)) {220      op.emitError("NYI: assuming only return operations can have 0 "221                   " results at this point");222    }223    return false;224  } else if (op.getNumResults() == 1) {225    if (auto v = dyn_cast<VectorType>(op.getResult(0).getType())) {226      superVectorType = v;227    } else {228      // Not a vector type.229      return false;230    }231  } else {232    // Not a vector.transfer and has more than 1 result, fail hard for now to233    // wake us up when something changes.234    op.emitError("NYI: operation has more than 1 result");235    return false;236  }237 238  // Get the ratio.239  auto ratio =240      computeShapeRatio(superVectorType.getShape(), subVectorType.getShape());241 242  // Sanity check.243  assert((ratio || !mustDivide) &&244         "vector.transfer operation in which super-vector size is not an"245         " integer multiple of sub-vector size");246 247  // This catches cases that are not strictly necessary to have multiplicity but248  // still aren't divisible by the sub-vector shape.249  // This could be useful information if we wanted to reshape at the level of250  // the vector type (but we would have to look at the compute and distinguish251  // between parallel, reduction and possibly other cases.252  return ratio.has_value();253}254 255bool vector::isContiguousSlice(MemRefType memrefType, VectorType vectorType) {256  if (vectorType.isScalable())257    return false;258 259  // Ignore a leading sequence of adjacent unit dimensions in the vector.260  ArrayRef<int64_t> vectorShape =261      vectorType.getShape().drop_while([](auto v) { return v == 1; });262  auto vecRank = vectorShape.size();263 264  if (!memrefType.areTrailingDimsContiguous(vecRank))265    return false;266 267  // Extract the trailing dims of the input memref268  auto memrefShape = memrefType.getShape().take_back(vecRank);269 270  // Compare the dims of `vectorType` against `memrefType`.271  // All of the dimensions, except the first must match.272  return llvm::equal(vectorShape.drop_front(), memrefShape.drop_front());273}274 275std::optional<StaticTileOffsetRange>276vector::createUnrollIterator(VectorType vType, int64_t targetRank) {277  if (vType.getRank() <= targetRank)278    return {};279  // Attempt to unroll until targetRank or the first scalable dimension (which280  // cannot be unrolled).281  auto shapeToUnroll = vType.getShape().drop_back(targetRank);282  auto inputScalableVecDimsToUnroll =283      vType.getScalableDims().drop_back(targetRank);284  const auto *it = llvm::find(inputScalableVecDimsToUnroll, true);285  auto firstScalableDim = it - inputScalableVecDimsToUnroll.begin();286  if (firstScalableDim == 0)287    return {};288  // All scalable dimensions should be removed now.289  inputScalableVecDimsToUnroll =290      inputScalableVecDimsToUnroll.slice(0, firstScalableDim);291  assert(!llvm::is_contained(inputScalableVecDimsToUnroll, true) &&292         "unexpected leading scalable dimension");293  // Create an unroll iterator for leading dimensions.294  shapeToUnroll = shapeToUnroll.slice(0, firstScalableDim);295  return StaticTileOffsetRange(shapeToUnroll, /*unrollStep=*/1);296}297 298SmallVector<OpFoldResult> vector::getMixedSizesXfer(bool hasTensorSemantics,299                                                    Operation *xfer,300                                                    RewriterBase &rewriter) {301  auto loc = xfer->getLoc();302 303  Value base = TypeSwitch<Operation *, Value>(xfer)304                   .Case<vector::TransferReadOp>(305                       [&](auto readOp) { return readOp.getBase(); })306                   .Case<vector::TransferWriteOp>(307                       [&](auto writeOp) { return writeOp.getOperand(1); });308 309  SmallVector<OpFoldResult> mixedSourceDims =310      hasTensorSemantics ? tensor::getMixedSizes(rewriter, loc, base)311                         : memref::getMixedSizes(rewriter, loc, base);312  return mixedSourceDims;313}314 315bool vector::isLinearizableVector(VectorType type) {316  return (type.getRank() > 1) && (type.getNumScalableDims() <= 1);317}318 319Value vector::createReadOrMaskedRead(OpBuilder &builder, Location loc,320                                     Value source,321                                     ArrayRef<int64_t> inputVectorSizes,322                                     std::optional<Value> padValue,323                                     bool useInBoundsInsteadOfMasking,324                                     ArrayRef<bool> inputScalableVecDims) {325  VectorType vecToReadTy = VectorType::get(326      inputVectorSizes, cast<ShapedType>(source.getType()).getElementType(),327      inputScalableVecDims);328 329  return createReadOrMaskedRead(builder, loc, source, vecToReadTy, padValue,330                                useInBoundsInsteadOfMasking);331}332 333Value vector::createReadOrMaskedRead(OpBuilder &builder, Location loc,334                                     Value source,335                                     const VectorType &vecToReadTy,336                                     std::optional<Value> padValue,337                                     bool useInBoundsInsteadOfMasking) {338  assert(!llvm::is_contained(vecToReadTy.getScalableDims(),339                             ShapedType::kDynamic) &&340         "invalid input vector sizes");341  auto sourceShapedType = cast<ShapedType>(source.getType());342  auto sourceShape = sourceShapedType.getShape();343 344  int64_t vecToReadRank = vecToReadTy.getRank();345  auto vecToReadShape = vecToReadTy.getShape();346 347  assert(sourceShape.size() == static_cast<size_t>(vecToReadRank) &&348         "expected same ranks.");349  assert((!padValue.has_value() ||350          padValue.value().getType() == sourceShapedType.getElementType()) &&351         "expected same pad element type to match source element type");352 353  auto zero = arith::ConstantIndexOp::create(builder, loc, 0);354  SmallVector<bool> inBoundsVal(vecToReadRank, true);355 356  if (useInBoundsInsteadOfMasking) {357    // Update the inBounds attribute.358    // FIXME: This computation is too weak - it ignores the read indices.359    for (unsigned i = 0; i < vecToReadRank; i++)360      inBoundsVal[i] = (sourceShape[i] == vecToReadShape[i]) &&361                       ShapedType::isStatic(sourceShape[i]);362  }363  auto transferReadOp = vector::TransferReadOp::create(364      builder, loc,365      /*vectorType=*/vecToReadTy,366      /*source=*/source,367      /*indices=*/SmallVector<Value>(vecToReadRank, zero),368      /*padding=*/padValue,369      /*inBounds=*/inBoundsVal);370 371  if (llvm::equal(vecToReadTy.getShape(), sourceShape) ||372      useInBoundsInsteadOfMasking)373    return transferReadOp;374  SmallVector<OpFoldResult> mixedSourceDims =375      isa<MemRefType>(source.getType())376          ? memref::getMixedSizes(builder, loc, source)377          : tensor::getMixedSizes(builder, loc, source);378 379  auto maskType = vecToReadTy.cloneWith(/*shape=*/{}, builder.getI1Type());380  Value mask =381      vector::CreateMaskOp::create(builder, loc, maskType, mixedSourceDims);382  return mlir::vector::maskOperation(builder, transferReadOp, mask)383      ->getResult(0);384}385 386LogicalResult387vector::isValidMaskedInputVector(ArrayRef<int64_t> shape,388                                 ArrayRef<int64_t> inputVectorSizes) {389  LDBG() << "Iteration space static sizes:" << llvm::interleaved(shape);390 391  if (inputVectorSizes.size() != shape.size()) {392    LDBG() << "Input vector sizes don't match the number of loops";393    return failure();394  }395  if (ShapedType::isDynamicShape(inputVectorSizes)) {396    LDBG() << "Input vector sizes can't have dynamic dimensions";397    return failure();398  }399  if (!llvm::all_of(llvm::zip(shape, inputVectorSizes),400                    [](std::tuple<int64_t, int64_t> sizePair) {401                      int64_t staticSize = std::get<0>(sizePair);402                      int64_t inputSize = std::get<1>(sizePair);403                      return ShapedType::isDynamic(staticSize) ||404                             staticSize <= inputSize;405                    })) {406    LDBG() << "Input vector sizes must be greater than or equal to iteration "407              "space static sizes";408    return failure();409  }410  return success();411}412 413/// Takes a 2+ dimensional vector as an input414/// returns n vector values produced by n vector.extract operations.415/// I.e. calling unrollVectorValue([[%v]], rewriter) such that416///417///   %v : vector<nxaxb...>418///419/// will produce the following IR changes420///421///   %v0 = vector.extract %v[0] : vector<axbx...> from vector<nxaxb...>422///   %v1 = vector.extract %v[1] : vector<axbx...> from vector<nxaxb...>423///   ...424///   %vnminusone = vector.extract %v[n-1] : vector<axbx...> from ...425///426/// and returns SmallVector<Value> r = {[[%v0]], [[%v1]], ..., [[%vnminusone]]}427FailureOr<SmallVector<Value>>428vector::unrollVectorValue(TypedValue<VectorType> vector,429                          RewriterBase &rewriter) {430  SmallVector<Value> subvectors;431  VectorType ty = cast<VectorType>(vector.getType());432  Location loc = vector.getLoc();433  if (ty.getRank() < 2)434    return rewriter.notifyMatchFailure(loc, "already 1-D");435 436  // Unrolling doesn't take vscale into account. Pattern is disabled for437  // vectors with leading scalable dim(s).438  if (ty.getScalableDims().front())439    return rewriter.notifyMatchFailure(loc, "cannot unroll scalable dim");440 441  for (int64_t i = 0, e = ty.getShape().front(); i < e; ++i) {442    subvectors.push_back(vector::ExtractOp::create(rewriter, loc, vector, i));443  }444 445  return subvectors;446}447 448LogicalResult vector::unrollVectorOp(Operation *op, PatternRewriter &rewriter,449                                     vector::UnrollVectorOpFn unrollFn) {450  assert(op->getNumResults() == 1 && "expected single result");451  assert(isa<VectorType>(op->getResult(0).getType()) && "expected vector type");452  VectorType resultTy = cast<VectorType>(op->getResult(0).getType());453  if (resultTy.getRank() < 2)454    return rewriter.notifyMatchFailure(op, "already 1-D");455 456  // Unrolling doesn't take vscale into account. Pattern is disabled for457  // vectors with leading scalable dim(s).458  if (resultTy.getScalableDims().front())459    return rewriter.notifyMatchFailure(op, "cannot unroll scalable dim");460 461  Location loc = op->getLoc();462  Value result = ub::PoisonOp::create(rewriter, loc, resultTy);463  VectorType subTy = VectorType::Builder(resultTy).dropDim(0);464 465  for (int64_t i = 0, e = resultTy.getShape().front(); i < e; ++i) {466    Value subVector = unrollFn(rewriter, loc, subTy, i);467    result = vector::InsertOp::create(rewriter, loc, subVector, result, i);468  }469 470  rewriter.replaceOp(op, result);471  return success();472}473