brintos

brintos / llvm-project-archived public Read only

0
0
Text · 20.0 KiB · f28a6cc Raw
519 lines · cpp
1//===- Pattern.cpp - 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/Pattern.h"10#include "mlir/Dialect/LLVMIR/FunctionCallUtils.h"11#include "mlir/Dialect/LLVMIR/LLVMDialect.h"12#include "mlir/Dialect/LLVMIR/LLVMTypes.h"13#include "mlir/IR/AffineMap.h"14#include "mlir/IR/BuiltinAttributes.h"15 16using namespace mlir;17 18//===----------------------------------------------------------------------===//19// ConvertToLLVMPattern20//===----------------------------------------------------------------------===//21 22ConvertToLLVMPattern::ConvertToLLVMPattern(23    StringRef rootOpName, MLIRContext *context,24    const LLVMTypeConverter &typeConverter, PatternBenefit benefit)25    : ConversionPattern(typeConverter, rootOpName, benefit, context) {}26 27const LLVMTypeConverter *ConvertToLLVMPattern::getTypeConverter() const {28  return static_cast<const LLVMTypeConverter *>(29      ConversionPattern::getTypeConverter());30}31 32LLVM::LLVMDialect &ConvertToLLVMPattern::getDialect() const {33  return *getTypeConverter()->getDialect();34}35 36Type ConvertToLLVMPattern::getIndexType() const {37  return getTypeConverter()->getIndexType();38}39 40Type ConvertToLLVMPattern::getIntPtrType(unsigned addressSpace) const {41  return IntegerType::get(&getTypeConverter()->getContext(),42                          getTypeConverter()->getPointerBitwidth(addressSpace));43}44 45Type ConvertToLLVMPattern::getVoidType() const {46  return LLVM::LLVMVoidType::get(&getTypeConverter()->getContext());47}48 49Type ConvertToLLVMPattern::getPtrType(unsigned addressSpace) const {50  return LLVM::LLVMPointerType::get(&getTypeConverter()->getContext(),51                                    addressSpace);52}53 54Type ConvertToLLVMPattern::getVoidPtrType() const { return getPtrType(); }55 56Value ConvertToLLVMPattern::createIndexAttrConstant(OpBuilder &builder,57                                                    Location loc,58                                                    Type resultType,59                                                    int64_t value) {60  return LLVM::ConstantOp::create(builder, loc, resultType,61                                  builder.getIndexAttr(value));62}63 64Value ConvertToLLVMPattern::getStridedElementPtr(65    ConversionPatternRewriter &rewriter, Location loc, MemRefType type,66    Value memRefDesc, ValueRange indices,67    LLVM::GEPNoWrapFlags noWrapFlags) const {68  return LLVM::getStridedElementPtr(rewriter, loc, *getTypeConverter(), type,69                                    memRefDesc, indices, noWrapFlags);70}71 72// Check if the MemRefType `type` is supported by the lowering. We currently73// only support memrefs with identity maps.74bool ConvertToLLVMPattern::isConvertibleAndHasIdentityMaps(75    MemRefType type) const {76  if (!type.getLayout().isIdentity())77    return false;78  return static_cast<bool>(typeConverter->convertType(type));79}80 81Type ConvertToLLVMPattern::getElementPtrType(MemRefType type) const {82  auto addressSpace = getTypeConverter()->getMemRefAddressSpace(type);83  if (failed(addressSpace))84    return {};85  return LLVM::LLVMPointerType::get(type.getContext(), *addressSpace);86}87 88void ConvertToLLVMPattern::getMemRefDescriptorSizes(89    Location loc, MemRefType memRefType, ValueRange dynamicSizes,90    ConversionPatternRewriter &rewriter, SmallVectorImpl<Value> &sizes,91    SmallVectorImpl<Value> &strides, Value &size, bool sizeInBytes) const {92  assert(isConvertibleAndHasIdentityMaps(memRefType) &&93         "layout maps must have been normalized away");94  assert(count(memRefType.getShape(), ShapedType::kDynamic) ==95             static_cast<ssize_t>(dynamicSizes.size()) &&96         "dynamicSizes size doesn't match dynamic sizes count in memref shape");97 98  sizes.reserve(memRefType.getRank());99  unsigned dynamicIndex = 0;100  Type indexType = getIndexType();101  for (int64_t size : memRefType.getShape()) {102    sizes.push_back(103        size == ShapedType::kDynamic104            ? dynamicSizes[dynamicIndex++]105            : createIndexAttrConstant(rewriter, loc, indexType, size));106  }107 108  // Strides: iterate sizes in reverse order and multiply.109  int64_t stride = 1;110  Value runningStride = createIndexAttrConstant(rewriter, loc, indexType, 1);111  strides.resize(memRefType.getRank());112  for (auto i = memRefType.getRank(); i-- > 0;) {113    strides[i] = runningStride;114 115    int64_t staticSize = memRefType.getShape()[i];116    bool useSizeAsStride = stride == 1;117    if (staticSize == ShapedType::kDynamic)118      stride = ShapedType::kDynamic;119    if (stride != ShapedType::kDynamic)120      stride *= staticSize;121 122    if (useSizeAsStride)123      runningStride = sizes[i];124    else if (stride == ShapedType::kDynamic)125      runningStride =126          LLVM::MulOp::create(rewriter, loc, runningStride, sizes[i]);127    else128      runningStride = createIndexAttrConstant(rewriter, loc, indexType, stride);129  }130  if (sizeInBytes) {131    // Buffer size in bytes.132    Type elementType = typeConverter->convertType(memRefType.getElementType());133    auto elementPtrType = LLVM::LLVMPointerType::get(rewriter.getContext());134    Value nullPtr = LLVM::ZeroOp::create(rewriter, loc, elementPtrType);135    Value gepPtr = LLVM::GEPOp::create(rewriter, loc, elementPtrType,136                                       elementType, nullPtr, runningStride);137    size = LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), gepPtr);138  } else {139    size = runningStride;140  }141}142 143Value ConvertToLLVMPattern::getSizeInBytes(144    Location loc, Type type, ConversionPatternRewriter &rewriter) const {145  // Compute the size of an individual element. This emits the MLIR equivalent146  // of the following sizeof(...) implementation in LLVM IR:147  //   %0 = getelementptr %elementType* null, %indexType 1148  //   %1 = ptrtoint %elementType* %0 to %indexType149  // which is a common pattern of getting the size of a type in bytes.150  Type llvmType = typeConverter->convertType(type);151  auto convertedPtrType = LLVM::LLVMPointerType::get(rewriter.getContext());152  auto nullPtr = LLVM::ZeroOp::create(rewriter, loc, convertedPtrType);153  auto gep = LLVM::GEPOp::create(rewriter, loc, convertedPtrType, llvmType,154                                 nullPtr, ArrayRef<LLVM::GEPArg>{1});155  return LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), gep);156}157 158Value ConvertToLLVMPattern::getNumElements(159    Location loc, MemRefType memRefType, ValueRange dynamicSizes,160    ConversionPatternRewriter &rewriter) const {161  assert(count(memRefType.getShape(), ShapedType::kDynamic) ==162             static_cast<ssize_t>(dynamicSizes.size()) &&163         "dynamicSizes size doesn't match dynamic sizes count in memref shape");164 165  Type indexType = getIndexType();166  Value numElements = memRefType.getRank() == 0167                          ? createIndexAttrConstant(rewriter, loc, indexType, 1)168                          : nullptr;169  unsigned dynamicIndex = 0;170 171  // Compute the total number of memref elements.172  for (int64_t staticSize : memRefType.getShape()) {173    if (numElements) {174      Value size =175          staticSize == ShapedType::kDynamic176              ? dynamicSizes[dynamicIndex++]177              : createIndexAttrConstant(rewriter, loc, indexType, staticSize);178      numElements = LLVM::MulOp::create(rewriter, loc, numElements, size);179    } else {180      numElements =181          staticSize == ShapedType::kDynamic182              ? dynamicSizes[dynamicIndex++]183              : createIndexAttrConstant(rewriter, loc, indexType, staticSize);184    }185  }186  return numElements;187}188 189/// Creates and populates the memref descriptor struct given all its fields.190MemRefDescriptor ConvertToLLVMPattern::createMemRefDescriptor(191    Location loc, MemRefType memRefType, Value allocatedPtr, Value alignedPtr,192    ArrayRef<Value> sizes, ArrayRef<Value> strides,193    ConversionPatternRewriter &rewriter) const {194  auto structType = typeConverter->convertType(memRefType);195  auto memRefDescriptor = MemRefDescriptor::poison(rewriter, loc, structType);196 197  // Field 1: Allocated pointer, used for malloc/free.198  memRefDescriptor.setAllocatedPtr(rewriter, loc, allocatedPtr);199 200  // Field 2: Actual aligned pointer to payload.201  memRefDescriptor.setAlignedPtr(rewriter, loc, alignedPtr);202 203  // Field 3: Offset in aligned pointer.204  Type indexType = getIndexType();205  memRefDescriptor.setOffset(206      rewriter, loc, createIndexAttrConstant(rewriter, loc, indexType, 0));207 208  // Fields 4: Sizes.209  for (const auto &en : llvm::enumerate(sizes))210    memRefDescriptor.setSize(rewriter, loc, en.index(), en.value());211 212  // Field 5: Strides.213  for (const auto &en : llvm::enumerate(strides))214    memRefDescriptor.setStride(rewriter, loc, en.index(), en.value());215 216  return memRefDescriptor;217}218 219Value ConvertToLLVMPattern::copyUnrankedDescriptor(220    OpBuilder &builder, Location loc, UnrankedMemRefType memRefType,221    Value operand, bool toDynamic) const {222  // Convert memory space.223  FailureOr<unsigned> addressSpace =224      getTypeConverter()->getMemRefAddressSpace(memRefType);225  if (failed(addressSpace))226    return {};227 228  // Get frequently used types.229  Type indexType = getTypeConverter()->getIndexType();230 231  // Find the malloc and free, or declare them if necessary.232  auto module = builder.getInsertionPoint()->getParentOfType<ModuleOp>();233  FailureOr<LLVM::LLVMFuncOp> freeFunc, mallocFunc;234  if (toDynamic) {235    mallocFunc = LLVM::lookupOrCreateMallocFn(builder, module, indexType);236    if (failed(mallocFunc))237      return {};238  }239  if (!toDynamic) {240    freeFunc = LLVM::lookupOrCreateFreeFn(builder, module);241    if (failed(freeFunc))242      return {};243  }244 245  UnrankedMemRefDescriptor desc(operand);246  Value allocationSize = UnrankedMemRefDescriptor::computeSize(247      builder, loc, *getTypeConverter(), desc, *addressSpace);248 249  // Allocate memory, copy, and free the source if necessary.250  Value memory = toDynamic251                     ? LLVM::CallOp::create(builder, loc, mallocFunc.value(),252                                            allocationSize)253                           .getResult()254                     : LLVM::AllocaOp::create(builder, loc, getPtrType(),255                                              IntegerType::get(getContext(), 8),256                                              allocationSize,257                                              /*alignment=*/0);258  Value source = desc.memRefDescPtr(builder, loc);259  LLVM::MemcpyOp::create(builder, loc, memory, source, allocationSize, false);260  if (!toDynamic)261    LLVM::CallOp::create(builder, loc, freeFunc.value(), source);262 263  // Create a new descriptor. The same descriptor can be returned multiple264  // times, attempting to modify its pointer can lead to memory leaks265  // (allocated twice and overwritten) or double frees (the caller does not266  // know if the descriptor points to the same memory).267  Type descriptorType = getTypeConverter()->convertType(memRefType);268  if (!descriptorType)269    return {};270  auto updatedDesc =271      UnrankedMemRefDescriptor::poison(builder, loc, descriptorType);272  Value rank = desc.rank(builder, loc);273  updatedDesc.setRank(builder, loc, rank);274  updatedDesc.setMemRefDescPtr(builder, loc, memory);275  return updatedDesc;276}277 278LogicalResult ConvertToLLVMPattern::copyUnrankedDescriptors(279    OpBuilder &builder, Location loc, TypeRange origTypes,280    SmallVectorImpl<Value> &operands, bool toDynamic) const {281  assert(origTypes.size() == operands.size() &&282         "expected as may original types as operands");283  for (unsigned i = 0, e = operands.size(); i < e; ++i) {284    if (auto memRefType = dyn_cast<UnrankedMemRefType>(origTypes[i])) {285      Value updatedDesc = copyUnrankedDescriptor(builder, loc, memRefType,286                                                 operands[i], toDynamic);287      if (!updatedDesc)288        return failure();289      operands[i] = updatedDesc;290    }291  }292  return success();293}294 295//===----------------------------------------------------------------------===//296// Detail methods297//===----------------------------------------------------------------------===//298 299/// Replaces the given operation "op" with a new operation of type "targetOp"300/// and given operands.301LogicalResult LLVM::detail::oneToOneRewrite(302    Operation *op, StringRef targetOp, ValueRange operands,303    ArrayRef<NamedAttribute> targetAttrs, Attribute propertiesAttr,304    const LLVMTypeConverter &typeConverter,305    ConversionPatternRewriter &rewriter) {306  unsigned numResults = op->getNumResults();307 308  SmallVector<Type> resultTypes;309  if (numResults != 0) {310    resultTypes.push_back(311        typeConverter.packOperationResults(op->getResultTypes()));312    if (!resultTypes.back())313      return failure();314  }315 316  // Create the operation through state since we don't know its C++ type.317  OperationState state(op->getLoc(), rewriter.getStringAttr(targetOp), operands,318                       resultTypes, targetAttrs);319  state.propertiesAttr = propertiesAttr;320  Operation *newOp = rewriter.create(state);321 322  // If the operation produced 0 or 1 result, return them immediately.323  if (numResults == 0)324    return rewriter.eraseOp(op), success();325  if (numResults == 1)326    return rewriter.replaceOp(op, newOp->getResult(0)), success();327 328  // Otherwise, it had been converted to an operation producing a structure.329  // Extract individual results from the structure and return them as list.330  SmallVector<Value, 4> results;331  results.reserve(numResults);332  for (unsigned i = 0; i < numResults; ++i) {333    results.push_back(LLVM::ExtractValueOp::create(rewriter, op->getLoc(),334                                                   newOp->getResult(0), i));335  }336  rewriter.replaceOp(op, results);337  return success();338}339 340LogicalResult LLVM::detail::intrinsicRewrite(341    Operation *op, StringRef intrinsic, ValueRange operands,342    const LLVMTypeConverter &typeConverter, RewriterBase &rewriter) {343  auto loc = op->getLoc();344 345  if (!llvm::all_of(operands, [](Value value) {346        return LLVM::isCompatibleType(value.getType());347      }))348    return failure();349 350  unsigned numResults = op->getNumResults();351  Type resType;352  if (numResults != 0)353    resType = typeConverter.packOperationResults(op->getResultTypes());354 355  auto callIntrOp = LLVM::CallIntrinsicOp::create(356      rewriter, loc, resType, rewriter.getStringAttr(intrinsic), operands);357  // Propagate attributes.358  callIntrOp->setAttrs(op->getAttrDictionary());359 360  if (numResults <= 1) {361    // Directly replace the original op.362    rewriter.replaceOp(op, callIntrOp);363    return success();364  }365 366  // Extract individual results from packed structure and use them as367  // replacements.368  SmallVector<Value, 4> results;369  results.reserve(numResults);370  Value intrRes = callIntrOp.getResults();371  for (unsigned i = 0; i < numResults; ++i)372    results.push_back(LLVM::ExtractValueOp::create(rewriter, loc, intrRes, i));373  rewriter.replaceOp(op, results);374 375  return success();376}377 378static unsigned getBitWidth(Type type) {379  if (type.isIntOrFloat())380    return type.getIntOrFloatBitWidth();381 382  auto vec = cast<VectorType>(type);383  assert(!vec.isScalable() && "scalable vectors are not supported");384  return vec.getNumElements() * getBitWidth(vec.getElementType());385}386 387static Value createI32Constant(OpBuilder &builder, Location loc,388                               int32_t value) {389  Type i32 = builder.getI32Type();390  return LLVM::ConstantOp::create(builder, loc, i32, value);391}392 393SmallVector<Value> mlir::LLVM::decomposeValue(OpBuilder &builder, Location loc,394                                              Value src, Type dstType) {395  Type srcType = src.getType();396  if (srcType == dstType)397    return {src};398 399  unsigned srcBitWidth = getBitWidth(srcType);400  unsigned dstBitWidth = getBitWidth(dstType);401  if (srcBitWidth == dstBitWidth) {402    Value cast = LLVM::BitcastOp::create(builder, loc, dstType, src);403    return {cast};404  }405 406  if (dstBitWidth > srcBitWidth) {407    auto smallerInt = builder.getIntegerType(srcBitWidth);408    if (srcType != smallerInt)409      src = LLVM::BitcastOp::create(builder, loc, smallerInt, src);410 411    auto largerInt = builder.getIntegerType(dstBitWidth);412    Value res = LLVM::ZExtOp::create(builder, loc, largerInt, src);413    return {res};414  }415  assert(srcBitWidth % dstBitWidth == 0 &&416         "src bit width must be a multiple of dst bit width");417  int64_t numElements = srcBitWidth / dstBitWidth;418  auto vecType = VectorType::get(numElements, dstType);419 420  src = LLVM::BitcastOp::create(builder, loc, vecType, src);421 422  SmallVector<Value> res;423  for (auto i : llvm::seq(numElements)) {424    Value idx = createI32Constant(builder, loc, i);425    Value elem = LLVM::ExtractElementOp::create(builder, loc, src, idx);426    res.emplace_back(elem);427  }428 429  return res;430}431 432Value mlir::LLVM::composeValue(OpBuilder &builder, Location loc, ValueRange src,433                               Type dstType) {434  assert(!src.empty() && "src range must not be empty");435  if (src.size() == 1) {436    Value res = src.front();437    if (res.getType() == dstType)438      return res;439 440    unsigned srcBitWidth = getBitWidth(res.getType());441    unsigned dstBitWidth = getBitWidth(dstType);442    if (dstBitWidth < srcBitWidth) {443      auto largerInt = builder.getIntegerType(srcBitWidth);444      if (res.getType() != largerInt)445        res = LLVM::BitcastOp::create(builder, loc, largerInt, res);446 447      auto smallerInt = builder.getIntegerType(dstBitWidth);448      res = LLVM::TruncOp::create(builder, loc, smallerInt, res);449    }450 451    if (res.getType() != dstType)452      res = LLVM::BitcastOp::create(builder, loc, dstType, res);453 454    return res;455  }456 457  int64_t numElements = src.size();458  auto srcType = VectorType::get(numElements, src.front().getType());459  Value res = LLVM::PoisonOp::create(builder, loc, srcType);460  for (auto &&[i, elem] : llvm::enumerate(src)) {461    Value idx = createI32Constant(builder, loc, i);462    res = LLVM::InsertElementOp::create(builder, loc, srcType, res, elem, idx);463  }464 465  if (res.getType() != dstType)466    res = LLVM::BitcastOp::create(builder, loc, dstType, res);467 468  return res;469}470 471Value mlir::LLVM::getStridedElementPtr(OpBuilder &builder, Location loc,472                                       const LLVMTypeConverter &converter,473                                       MemRefType type, Value memRefDesc,474                                       ValueRange indices,475                                       LLVM::GEPNoWrapFlags noWrapFlags) {476  auto [strides, offset] = type.getStridesAndOffset();477 478  MemRefDescriptor memRefDescriptor(memRefDesc);479  // Use a canonical representation of the start address so that later480  // optimizations have a longer sequence of instructions to CSE.481  // If we don't do that we would sprinkle the memref.offset in various482  // position of the different address computations.483  Value base = memRefDescriptor.bufferPtr(builder, loc, converter, type);484 485  LLVM::IntegerOverflowFlags intOverflowFlags =486      LLVM::IntegerOverflowFlags::none;487  if (LLVM::bitEnumContainsAny(noWrapFlags, LLVM::GEPNoWrapFlags::nusw)) {488    intOverflowFlags = intOverflowFlags | LLVM::IntegerOverflowFlags::nsw;489  }490  if (LLVM::bitEnumContainsAny(noWrapFlags, LLVM::GEPNoWrapFlags::nuw)) {491    intOverflowFlags = intOverflowFlags | LLVM::IntegerOverflowFlags::nuw;492  }493 494  Type indexType = converter.getIndexType();495  Value index;496  for (int i = 0, e = indices.size(); i < e; ++i) {497    Value increment = indices[i];498    if (strides[i] != 1) { // Skip if stride is 1.499      Value stride =500          ShapedType::isDynamic(strides[i])501              ? memRefDescriptor.stride(builder, loc, i)502              : LLVM::ConstantOp::create(builder, loc, indexType,503                                         builder.getIndexAttr(strides[i]));504      increment = LLVM::MulOp::create(builder, loc, increment, stride,505                                      intOverflowFlags);506    }507    index = index ? LLVM::AddOp::create(builder, loc, index, increment,508                                        intOverflowFlags)509                  : increment;510  }511 512  Type elementPtrType = memRefDescriptor.getElementPtrType();513  return index514             ? LLVM::GEPOp::create(builder, loc, elementPtrType,515                                   converter.convertType(type.getElementType()),516                                   base, index, noWrapFlags)517             : base;518}519