brintos

brintos / llvm-project-archived public Read only

0
0
Text · 6.1 KiB · e5969c2 Raw
154 lines · cpp
1//===- VectorPattern.cpp - Vector conversion pattern to the LLVM dialect --===//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/Conversion/LLVMCommon/VectorPattern.h"10#include "mlir/Dialect/LLVMIR/LLVMDialect.h"11 12using namespace mlir;13 14// For >1-D vector types, extracts the necessary information to iterate over all15// 1-D subvectors in the underlying llrepresentation of the n-D vector16// Iterates on the llvm array type until we hit a non-array type (which is17// asserted to be an llvm vector type).18LLVM::detail::NDVectorTypeInfo19LLVM::detail::extractNDVectorTypeInfo(VectorType vectorType,20                                      const LLVMTypeConverter &converter) {21  assert(vectorType.getRank() > 1 && "expected >1D vector type");22  NDVectorTypeInfo info;23  info.llvmNDVectorTy = converter.convertType(vectorType);24  if (!info.llvmNDVectorTy || !LLVM::isCompatibleType(info.llvmNDVectorTy)) {25    info.llvmNDVectorTy = nullptr;26    return info;27  }28  info.arraySizes.reserve(vectorType.getRank() - 1);29  auto llvmTy = info.llvmNDVectorTy;30  while (isa<LLVM::LLVMArrayType>(llvmTy)) {31    info.arraySizes.push_back(32        cast<LLVM::LLVMArrayType>(llvmTy).getNumElements());33    llvmTy = cast<LLVM::LLVMArrayType>(llvmTy).getElementType();34  }35  if (!LLVM::isCompatibleVectorType(llvmTy))36    return info;37  info.llvm1DVectorTy = llvmTy;38  return info;39}40 41// Express `linearIndex` in terms of coordinates of `basis`.42// Returns the empty vector when linearIndex is out of the range [0, P] where43// P is the product of all the basis coordinates.44//45// Prerequisites:46//   Basis is an array of nonnegative integers (signed type inherited from47//   vector shape type).48SmallVector<int64_t, 4> LLVM::detail::getCoordinates(ArrayRef<int64_t> basis,49                                                     unsigned linearIndex) {50  SmallVector<int64_t, 4> res;51  res.reserve(basis.size());52  for (unsigned basisElement : llvm::reverse(basis)) {53    res.push_back(linearIndex % basisElement);54    linearIndex = linearIndex / basisElement;55  }56  if (linearIndex > 0)57    return {};58  std::reverse(res.begin(), res.end());59  return res;60}61 62// Iterate of linear index, convert to coords space and insert splatted 1-D63// vector in each position.64void LLVM::detail::nDVectorIterate(const LLVM::detail::NDVectorTypeInfo &info,65                                   OpBuilder &builder,66                                   function_ref<void(ArrayRef<int64_t>)> fun) {67  unsigned ub = 1;68  for (auto s : info.arraySizes)69    ub *= s;70  for (unsigned linearIndex = 0; linearIndex < ub; ++linearIndex) {71    auto coords = getCoordinates(info.arraySizes, linearIndex);72    // Linear index is out of bounds, we are done.73    if (coords.empty())74      break;75    assert(coords.size() == info.arraySizes.size());76    fun(coords);77  }78}79 80LogicalResult LLVM::detail::handleMultidimensionalVectors(81    Operation *op, ValueRange operands, const LLVMTypeConverter &typeConverter,82    std::function<Value(Type, ValueRange)> createOperand,83    ConversionPatternRewriter &rewriter) {84  auto resultNDVectorType = cast<VectorType>(op->getResult(0).getType());85  auto resultTypeInfo =86      extractNDVectorTypeInfo(resultNDVectorType, typeConverter);87  auto result1DVectorTy = resultTypeInfo.llvm1DVectorTy;88  auto resultNDVectoryTy = resultTypeInfo.llvmNDVectorTy;89  auto loc = op->getLoc();90  Value desc = LLVM::PoisonOp::create(rewriter, loc, resultNDVectoryTy);91  nDVectorIterate(resultTypeInfo, rewriter, [&](ArrayRef<int64_t> position) {92    // For this unrolled `position` corresponding to the `linearIndex`^th93    // element, extract operand vectors94    SmallVector<Value, 4> extractedOperands;95    for (const auto &operand : llvm::enumerate(operands)) {96      extractedOperands.push_back(LLVM::ExtractValueOp::create(97          rewriter, loc, operand.value(), position));98    }99    Value newVal = createOperand(result1DVectorTy, extractedOperands);100    desc = LLVM::InsertValueOp::create(rewriter, loc, desc, newVal, position);101  });102  rewriter.replaceOp(op, desc);103  return success();104}105 106LogicalResult LLVM::detail::vectorOneToOneRewrite(107    Operation *op, StringRef targetOp, ValueRange operands,108    ArrayRef<NamedAttribute> targetAttrs, Attribute propertiesAttr,109    const LLVMTypeConverter &typeConverter,110    ConversionPatternRewriter &rewriter) {111  assert(!operands.empty());112 113  // Cannot convert ops if their operands are not of LLVM type.114  if (!llvm::all_of(operands.getTypes(), isCompatibleType))115    return failure();116 117  auto llvmNDVectorTy = operands[0].getType();118  if (!isa<LLVM::LLVMArrayType>(llvmNDVectorTy))119    return oneToOneRewrite(op, targetOp, operands, targetAttrs, propertiesAttr,120                           typeConverter, rewriter);121  auto callback = [op, targetOp, targetAttrs, propertiesAttr,122                   &rewriter](Type llvm1DVectorTy, ValueRange operands) {123    OperationState state(op->getLoc(), rewriter.getStringAttr(targetOp),124                         operands, llvm1DVectorTy, targetAttrs);125    state.propertiesAttr = propertiesAttr;126    Operation *newOp = rewriter.create(state);127    return newOp->getResult(0);128  };129 130  return handleMultidimensionalVectors(op, operands, typeConverter, callback,131                                       rewriter);132}133 134/// Return the given type if it's a floating point type. If the given type is135/// a vector type, return its element type if it's a floating point type.136static FloatType getFloatingPointType(Type type) {137  if (auto floatType = dyn_cast<FloatType>(type))138    return floatType;139  if (auto vecType = dyn_cast<VectorType>(type))140    return dyn_cast<FloatType>(vecType.getElementType());141  return nullptr;142}143 144bool LLVM::detail::isUnsupportedFloatingPointType(145    const TypeConverter &typeConverter, Type type) {146  FloatType floatType = getFloatingPointType(type);147  if (!floatType)148    return false;149  Type convertedType = typeConverter.convertType(floatType);150  if (!convertedType)151    return true;152  return !isa<FloatType>(convertedType);153}154