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