brintos

brintos / llvm-project-archived public Read only

0
0
Text · 34.4 KiB · 5848489 Raw
790 lines · cpp
1//===- LowerGpuOpsToNVVMOps.cpp - MLIR GPU to NVVM lowering passes --------===//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 a pass to generate NVVMIR operations for higher-level10// GPU operations.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Conversion/ConvertToLLVM/ToLLVMInterface.h"15#include "mlir/Conversion/ConvertToLLVM/ToLLVMPass.h"16#include "mlir/Conversion/GPUCommon/GPUCommonPass.h"17#include "mlir/Conversion/GPUToNVVM/GPUToNVVM.h"18#include "mlir/Conversion/GPUToNVVM/GPUToNVVMPass.h"19#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"20#include "mlir/Conversion/LLVMCommon/LoweringOptions.h"21#include "mlir/Conversion/LLVMCommon/TypeConverter.h"22#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"23#include "mlir/Dialect/Func/IR/FuncOps.h"24#include "mlir/Dialect/GPU/IR/GPUDialect.h"25#include "mlir/Dialect/GPU/Transforms/Passes.h"26#include "mlir/Dialect/LLVMIR/NVVMDialect.h"27#include "mlir/Dialect/Math/IR/Math.h"28#include "mlir/Dialect/MemRef/IR/MemRef.h"29#include "mlir/Dialect/NVGPU/IR/NVGPUDialect.h"30#include "mlir/Dialect/Vector/Transforms/LoweringPatterns.h"31#include "mlir/Transforms/DialectConversion.h"32#include "mlir/Transforms/GreedyPatternRewriteDriver.h"33 34#include "../GPUCommon/GPUOpsLowering.h"35#include "../GPUCommon/IndexIntrinsicsOpLowering.h"36#include "../GPUCommon/OpToFuncCallLowering.h"37#include <optional>38 39namespace mlir {40#define GEN_PASS_DEF_CONVERTGPUOPSTONVVMOPS41#include "mlir/Conversion/Passes.h.inc"42} // namespace mlir43 44using namespace mlir;45 46namespace {47 48/// Convert gpu dialect shfl mode enum to the equivalent nvvm one.49static NVVM::ShflKind convertShflKind(gpu::ShuffleMode mode) {50  switch (mode) {51  case gpu::ShuffleMode::XOR:52    return NVVM::ShflKind::bfly;53  case gpu::ShuffleMode::UP:54    return NVVM::ShflKind::up;55  case gpu::ShuffleMode::DOWN:56    return NVVM::ShflKind::down;57  case gpu::ShuffleMode::IDX:58    return NVVM::ShflKind::idx;59  }60  llvm_unreachable("unknown shuffle mode");61}62 63static std::optional<NVVM::ReduxKind>64convertReduxKind(gpu::AllReduceOperation mode) {65  switch (mode) {66  case gpu::AllReduceOperation::ADD:67    return NVVM::ReduxKind::ADD;68  case gpu::AllReduceOperation::MUL:69    return std::nullopt;70  case gpu::AllReduceOperation::MINSI:71    return NVVM::ReduxKind::MIN;72  case gpu::AllReduceOperation::MINUI:73    return std::nullopt;74  case gpu::AllReduceOperation::MINNUMF:75    return NVVM::ReduxKind::MIN;76  case gpu::AllReduceOperation::MAXSI:77    return NVVM::ReduxKind::MAX;78  case gpu::AllReduceOperation::MAXUI:79    return std::nullopt;80  case gpu::AllReduceOperation::MAXNUMF:81    return NVVM::ReduxKind::MAX;82  case gpu::AllReduceOperation::AND:83    return NVVM::ReduxKind::AND;84  case gpu::AllReduceOperation::OR:85    return NVVM::ReduxKind::OR;86  case gpu::AllReduceOperation::XOR:87    return NVVM::ReduxKind::XOR;88  case gpu::AllReduceOperation::MINIMUMF:89  case gpu::AllReduceOperation::MAXIMUMF:90    return std::nullopt;91  }92  return std::nullopt;93}94 95/// This pass lowers gpu.subgroup_reduce op into to the nvvm.redux op. The op96/// must be run by the entire subgroup, otherwise it is undefined behaviour.97struct GPUSubgroupReduceOpLowering98    : public ConvertOpToLLVMPattern<gpu::SubgroupReduceOp> {99  using ConvertOpToLLVMPattern<gpu::SubgroupReduceOp>::ConvertOpToLLVMPattern;100  LogicalResult101 102  matchAndRewrite(gpu::SubgroupReduceOp op, OpAdaptor adaptor,103                  ConversionPatternRewriter &rewriter) const override {104    if (op.getClusterSize())105      return rewriter.notifyMatchFailure(106          op, "lowering for clustered reduce not implemented");107 108    if (!op.getUniform())109      return rewriter.notifyMatchFailure(110          op, "cannot be lowered to redux as the op must be run "111              "uniformly (entire subgroup).");112    if (!op.getValue().getType().isInteger(32))113      return rewriter.notifyMatchFailure(op, "unsupported data type");114 115    std::optional<NVVM::ReduxKind> mode = convertReduxKind(op.getOp());116    if (!mode.has_value())117      return rewriter.notifyMatchFailure(118          op, "unsupported reduction mode for redux");119 120    Location loc = op->getLoc();121    auto int32Type = IntegerType::get(rewriter.getContext(), 32);122    Value offset = LLVM::ConstantOp::create(rewriter, loc, int32Type, -1);123 124    auto reduxOp = NVVM::ReduxOp::create(rewriter, loc, int32Type,125                                         op.getValue(), mode.value(), offset);126 127    rewriter.replaceOp(op, reduxOp->getResult(0));128    return success();129  }130};131 132struct GPUShuffleOpLowering : public ConvertOpToLLVMPattern<gpu::ShuffleOp> {133  using ConvertOpToLLVMPattern<gpu::ShuffleOp>::ConvertOpToLLVMPattern;134 135  /// Lowers a shuffle to the corresponding NVVM op.136  ///137  /// Convert the `width` argument into an activeMask (a bitmask which specifies138  /// which threads participate in the shuffle) and a maskAndClamp (specifying139  /// the highest lane which participates in the shuffle).140  ///141  ///     %one = llvm.constant(1 : i32) : i32142  ///     %minus_one = llvm.constant(-1 : i32) : i32143  ///     %thirty_two = llvm.constant(32 : i32) : i32144  ///     %num_lanes = llvm.sub %thirty_two, %width : i32145  ///     %active_mask = llvm.lshr %minus_one, %num_lanes : i32146  ///     %mask_and_clamp = llvm.sub %width, %one : i32147  ///     %shfl = nvvm.shfl.sync.bfly %active_mask, %value, %offset,148  ///         %mask_and_clamp : !llvm<"{ float, i1 }">149  ///     %shfl_value = llvm.extractvalue %shfl[0] :150  ///         !llvm<"{ float, i1 }">151  ///     %shfl_pred = llvm.extractvalue %shfl[1] :152  ///         !llvm<"{ float, i1 }">153  LogicalResult154  matchAndRewrite(gpu::ShuffleOp op, OpAdaptor adaptor,155                  ConversionPatternRewriter &rewriter) const override {156    Location loc = op->getLoc();157 158    auto valueTy = adaptor.getValue().getType();159    auto int32Type = IntegerType::get(rewriter.getContext(), 32);160    auto predTy = IntegerType::get(rewriter.getContext(), 1);161 162    Value one = LLVM::ConstantOp::create(rewriter, loc, int32Type, 1);163    Value minusOne = LLVM::ConstantOp::create(rewriter, loc, int32Type, -1);164    Value thirtyTwo = LLVM::ConstantOp::create(rewriter, loc, int32Type, 32);165    Value numLeadInactiveLane = LLVM::SubOp::create(166        rewriter, loc, int32Type, thirtyTwo, adaptor.getWidth());167    // Bit mask of active lanes: `(-1) >> (32 - activeWidth)`.168    Value activeMask = LLVM::LShrOp::create(rewriter, loc, int32Type, minusOne,169                                            numLeadInactiveLane);170    Value maskAndClamp;171    if (op.getMode() == gpu::ShuffleMode::UP) {172      // Clamp lane: `32 - activeWidth`173      maskAndClamp = numLeadInactiveLane;174    } else {175      // Clamp lane: `activeWidth - 1`176      maskAndClamp = LLVM::SubOp::create(rewriter, loc, int32Type,177                                         adaptor.getWidth(), one);178    }179 180    bool predIsUsed = !op->getResult(1).use_empty();181    UnitAttr returnValueAndIsValidAttr = nullptr;182    Type resultTy = valueTy;183    if (predIsUsed) {184      returnValueAndIsValidAttr = rewriter.getUnitAttr();185      resultTy = LLVM::LLVMStructType::getLiteral(rewriter.getContext(),186                                                  {valueTy, predTy});187    }188    Value shfl = NVVM::ShflOp::create(189        rewriter, loc, resultTy, activeMask, adaptor.getValue(),190        adaptor.getOffset(), maskAndClamp, convertShflKind(op.getMode()),191        returnValueAndIsValidAttr);192    if (predIsUsed) {193      Value shflValue = LLVM::ExtractValueOp::create(rewriter, loc, shfl, 0);194      Value isActiveSrcLane =195          LLVM::ExtractValueOp::create(rewriter, loc, shfl, 1);196      rewriter.replaceOp(op, {shflValue, isActiveSrcLane});197    } else {198      rewriter.replaceOp(op, {shfl, nullptr});199    }200    return success();201  }202};203 204struct GPULaneIdOpToNVVM : ConvertOpToLLVMPattern<gpu::LaneIdOp> {205  using ConvertOpToLLVMPattern<gpu::LaneIdOp>::ConvertOpToLLVMPattern;206 207  LogicalResult208  matchAndRewrite(gpu::LaneIdOp op, gpu::LaneIdOp::Adaptor adaptor,209                  ConversionPatternRewriter &rewriter) const override {210    auto loc = op->getLoc();211    MLIRContext *context = rewriter.getContext();212    LLVM::ConstantRangeAttr bounds = nullptr;213    if (std::optional<APInt> upperBound = op.getUpperBound())214      bounds = rewriter.getAttr<LLVM::ConstantRangeAttr>(215          /*bitWidth=*/32, /*lower=*/0, upperBound->getZExtValue());216    else217      bounds = rewriter.getAttr<LLVM::ConstantRangeAttr>(218          /*bitWidth=*/32, /*lower=*/0, /*upper=*/kWarpSize);219    Value newOp =220        NVVM::LaneIdOp::create(rewriter, loc, rewriter.getI32Type(), bounds);221    // Truncate or extend the result depending on the index bitwidth specified222    // by the LLVMTypeConverter options.223    const unsigned indexBitwidth = getTypeConverter()->getIndexTypeBitwidth();224    if (indexBitwidth > 32) {225      newOp = LLVM::SExtOp::create(226          rewriter, loc, IntegerType::get(context, indexBitwidth), newOp);227    } else if (indexBitwidth < 32) {228      newOp = LLVM::TruncOp::create(229          rewriter, loc, IntegerType::get(context, indexBitwidth), newOp);230    }231    rewriter.replaceOp(op, {newOp});232    return success();233  }234};235 236/// Lowering of cf.assert into a conditional __assertfail.237struct AssertOpToAssertfailLowering238    : public ConvertOpToLLVMPattern<cf::AssertOp> {239  using ConvertOpToLLVMPattern<cf::AssertOp>::ConvertOpToLLVMPattern;240 241  LogicalResult242  matchAndRewrite(cf::AssertOp assertOp, cf::AssertOpAdaptor adaptor,243                  ConversionPatternRewriter &rewriter) const override {244    MLIRContext *ctx = rewriter.getContext();245    Location loc = assertOp.getLoc();246    Type i8Type = typeConverter->convertType(rewriter.getIntegerType(8));247    Type i32Type = typeConverter->convertType(rewriter.getIntegerType(32));248    Type i64Type = typeConverter->convertType(rewriter.getIntegerType(64));249    Type ptrType = LLVM::LLVMPointerType::get(ctx);250    Type voidType = LLVM::LLVMVoidType::get(ctx);251 252    // Find or create __assertfail function declaration.253    auto moduleOp = assertOp->getParentOfType<gpu::GPUModuleOp>();254    auto assertfailType = LLVM::LLVMFunctionType::get(255        voidType, {ptrType, ptrType, i32Type, ptrType, i64Type});256    LLVM::LLVMFuncOp assertfailDecl = getOrDefineFunction(257        moduleOp, loc, rewriter, "__assertfail", assertfailType);258    assertfailDecl.setPassthroughAttr(259        ArrayAttr::get(ctx, StringAttr::get(ctx, "noreturn")));260 261    // Split blocks and insert conditional branch.262    // ^before:263    //   ...264    //   cf.cond_br %condition, ^after, ^assert265    // ^assert:266    //   cf.assert267    //   cf.br ^after268    // ^after:269    //   ...270    Block *beforeBlock = assertOp->getBlock();271    Block *assertBlock =272        rewriter.splitBlock(beforeBlock, assertOp->getIterator());273    Block *afterBlock =274        rewriter.splitBlock(assertBlock, ++assertOp->getIterator());275    rewriter.setInsertionPointToEnd(beforeBlock);276    cf::CondBranchOp::create(rewriter, loc, adaptor.getArg(), afterBlock,277                             assertBlock);278    rewriter.setInsertionPointToEnd(assertBlock);279    cf::BranchOp::create(rewriter, loc, afterBlock);280 281    // Continue cf.assert lowering.282    rewriter.setInsertionPoint(assertOp);283 284    // Populate file name, file number and function name from the location of285    // the AssertOp.286    StringRef fileName = "(unknown)";287    StringRef funcName = "(unknown)";288    int32_t fileLine = 0;289    while (auto callSiteLoc = dyn_cast<CallSiteLoc>(loc))290      loc = callSiteLoc.getCallee();291    if (auto fileLineColLoc = dyn_cast<FileLineColRange>(loc)) {292      fileName = fileLineColLoc.getFilename().strref();293      fileLine = fileLineColLoc.getStartLine();294    } else if (auto nameLoc = dyn_cast<NameLoc>(loc)) {295      funcName = nameLoc.getName().strref();296      if (auto fileLineColLoc =297              dyn_cast<FileLineColRange>(nameLoc.getChildLoc())) {298        fileName = fileLineColLoc.getFilename().strref();299        fileLine = fileLineColLoc.getStartLine();300      }301    }302 303    // Create constants.304    auto getGlobal = [&](LLVM::GlobalOp global) {305      // Get a pointer to the format string's first element.306      Value globalPtr = LLVM::AddressOfOp::create(307          rewriter, loc, LLVM::LLVMPointerType::get(ctx, global.getAddrSpace()),308          global.getSymNameAttr());309      Value start =310          LLVM::GEPOp::create(rewriter, loc, ptrType, global.getGlobalType(),311                              globalPtr, ArrayRef<LLVM::GEPArg>{0, 0});312      return start;313    };314    Value assertMessage = getGlobal(getOrCreateStringConstant(315        rewriter, loc, moduleOp, i8Type, "assert_message_", assertOp.getMsg()));316    Value assertFile = getGlobal(getOrCreateStringConstant(317        rewriter, loc, moduleOp, i8Type, "assert_file_", fileName));318    Value assertFunc = getGlobal(getOrCreateStringConstant(319        rewriter, loc, moduleOp, i8Type, "assert_func_", funcName));320    Value assertLine =321        LLVM::ConstantOp::create(rewriter, loc, i32Type, fileLine);322    Value c1 = LLVM::ConstantOp::create(rewriter, loc, i64Type, 1);323 324    // Insert function call to __assertfail.325    SmallVector<Value> arguments{assertMessage, assertFile, assertLine,326                                 assertFunc, c1};327    rewriter.replaceOpWithNewOp<LLVM::CallOp>(assertOp, assertfailDecl,328                                              arguments);329    return success();330  }331};332 333/// Import the GPU Ops to NVVM Patterns.334#include "GPUToNVVM.cpp.inc"335 336/// A pass that replaces all occurrences of GPU device operations with their337/// corresponding NVVM equivalent.338///339/// This pass only handles device code and is not meant to be run on GPU host340/// code.341struct LowerGpuOpsToNVVMOpsPass final342    : public impl::ConvertGpuOpsToNVVMOpsBase<LowerGpuOpsToNVVMOpsPass> {343  using Base::Base;344 345  void getDependentDialects(DialectRegistry &registry) const override {346    Base::getDependentDialects(registry);347    registerConvertToLLVMDependentDialectLoading(registry);348  }349 350  void runOnOperation() override {351    gpu::GPUModuleOp m = getOperation();352 353    // Request C wrapper emission.354    for (auto func : m.getOps<func::FuncOp>()) {355      func->setAttr(LLVM::LLVMDialect::getEmitCWrapperAttrName(),356                    UnitAttr::get(&getContext()));357    }358 359    // Customize the bitwidth used for the device side index computations.360    LowerToLLVMOptions options(361        m.getContext(),362        DataLayout(cast<DataLayoutOpInterface>(m.getOperation())));363    if (indexBitwidth != kDeriveIndexBitwidthFromDataLayout)364      options.overrideIndexBitwidth(indexBitwidth);365    options.useBarePtrCallConv = useBarePtrCallConv;366 367    // Apply in-dialect lowering. In-dialect lowering will replace368    // ops which need to be lowered further, which is not supported by a369    // single conversion pass.370    {371      RewritePatternSet patterns(m.getContext());372      populateGpuRewritePatterns(patterns);373      // Transform N-D vector.from_elements to 1-D vector.from_elements before374      // conversion.375      vector::populateVectorFromElementsUnrollPatterns(patterns);376      if (failed(applyPatternsGreedily(m, std::move(patterns))))377        return signalPassFailure();378    }379 380    LLVMTypeConverter converter(m.getContext(), options);381    configureGpuToNVVMTypeConverter(converter);382    RewritePatternSet llvmPatterns(m.getContext());383    LLVMConversionTarget target(getContext());384 385    // Set higher benefit, so patterns will run before generic LLVM lowering.386    populateGpuToNVVMConversionPatterns(converter, llvmPatterns,387                                        /*benefit=*/10);388 389    llvm::SmallDenseSet<StringRef> allowedDialectsSet(allowedDialects.begin(),390                                                      allowedDialects.end());391    for (Dialect *dialect : getContext().getLoadedDialects()) {392      // Skip math patterns as nvvm needs custom math lowering.393      if (isa<math::MathDialect>(dialect))394        continue;395 396      bool allowed = allowedDialectsSet.contains(dialect->getNamespace());397      // Empty `allowedDialectsSet` means all dialects are allowed.398      if (!allowedDialectsSet.empty() && !allowed)399        continue;400 401      auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);402      if (!iface) {403        // Error out if dialect was explicily specified but doesn't implement404        // conversion interface.405        if (allowed) {406          m.emitError()407              << "dialect does not implement ConvertToLLVMPatternInterface: "408              << dialect->getNamespace();409          return signalPassFailure();410        }411        continue;412      }413 414      iface->populateConvertToLLVMConversionPatterns(target, converter,415                                                     llvmPatterns);416    }417 418    populateGpuWMMAToNVVMConversionPatterns(converter, llvmPatterns);419    if (this->hasRedux)420      populateGpuSubgroupReduceOpLoweringPattern(converter, llvmPatterns);421    configureGpuToNVVMConversionLegality(target);422    ConversionConfig config;423    config.allowPatternRollback = allowPatternRollback;424    if (failed(425            applyPartialConversion(m, target, std::move(llvmPatterns), config)))426      signalPassFailure();427  }428};429 430} // namespace431 432void mlir::configureGpuToNVVMConversionLegality(ConversionTarget &target) {433  target.addIllegalOp<func::FuncOp>();434  target.addIllegalOp<cf::AssertOp>();435  target.addLegalDialect<::mlir::LLVM::LLVMDialect>();436  target.addLegalDialect<::mlir::NVVM::NVVMDialect>();437  target.addIllegalDialect<gpu::GPUDialect>();438  target.addIllegalOp<LLVM::CopySignOp, LLVM::CosOp, LLVM::ExpOp, LLVM::Exp2Op,439                      LLVM::FAbsOp, LLVM::FCeilOp, LLVM::FFloorOp, LLVM::FRemOp,440                      LLVM::LogOp, LLVM::Log10Op, LLVM::Log2Op, LLVM::PowOp,441                      LLVM::RoundEvenOp, LLVM::RoundOp, LLVM::SinOp,442                      LLVM::SincosOp, LLVM::SqrtOp>();443 444  // TODO: Remove once we support replacing non-root ops.445  target.addLegalOp<gpu::YieldOp, gpu::GPUModuleOp>();446}447 448void mlir::configureGpuToNVVMTypeConverter(LLVMTypeConverter &converter) {449  // NVVM uses alloca in the default address space to represent private450  // memory allocations, so drop private annotations. NVVM uses address451  // space 3 for shared memory. NVVM uses the default address space to452  // represent global memory.453  populateGpuMemorySpaceAttributeConversions(454      converter, [](gpu::AddressSpace space) -> unsigned {455        switch (space) {456        case gpu::AddressSpace::Global:457          return static_cast<unsigned>(NVVM::NVVMMemorySpace::Global);458        case gpu::AddressSpace::Workgroup:459          return static_cast<unsigned>(NVVM::NVVMMemorySpace::Shared);460        case gpu::AddressSpace::Private:461          return 0;462        }463        llvm_unreachable("unknown address space enum value");464        return static_cast<unsigned>(NVVM::NVVMMemorySpace::Generic);465      });466  // Lowering for MMAMatrixType.467  converter.addConversion([&](gpu::MMAMatrixType type) -> Type {468    return convertMMAToLLVMType(type);469  });470}471 472struct SincosOpLowering : public ConvertOpToLLVMPattern<math::SincosOp> {473  using ConvertOpToLLVMPattern<math::SincosOp>::ConvertOpToLLVMPattern;474 475  LogicalResult476  matchAndRewrite(math::SincosOp op, OpAdaptor adaptor,477                  ConversionPatternRewriter &rewriter) const override {478    Location loc = op.getLoc();479    Value input = adaptor.getOperand();480    Type inputType = input.getType();481    auto convertedInput = maybeExt(input, rewriter);482    auto computeType = convertedInput.getType();483 484    StringRef sincosFunc;485    if (isa<Float32Type>(computeType)) {486      const arith::FastMathFlags flag = op.getFastmath();487      const bool useApprox =488          mlir::arith::bitEnumContainsAny(flag, arith::FastMathFlags::afn);489      sincosFunc = useApprox ? "__nv_fast_sincosf" : "__nv_sincosf";490    } else if (isa<Float64Type>(computeType)) {491      sincosFunc = "__nv_sincos";492    } else {493      return rewriter.notifyMatchFailure(op,494                                         "unsupported operand type for sincos");495    }496 497    auto ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());498 499    Value sinPtr, cosPtr;500    {501      OpBuilder::InsertionGuard guard(rewriter);502      auto *scope =503          op->getParentWithTrait<mlir::OpTrait::AutomaticAllocationScope>();504      assert(scope && "Expected op to be inside automatic allocation scope");505      rewriter.setInsertionPointToStart(&scope->getRegion(0).front());506      auto one = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(),507                                          rewriter.getI32IntegerAttr(1));508      sinPtr =509          LLVM::AllocaOp::create(rewriter, loc, ptrType, computeType, one, 0);510      cosPtr =511          LLVM::AllocaOp::create(rewriter, loc, ptrType, computeType, one, 0);512    }513 514    createSincosCall(rewriter, loc, sincosFunc, convertedInput, sinPtr, cosPtr,515                     op);516 517    auto sinResult = LLVM::LoadOp::create(rewriter, loc, computeType, sinPtr);518    auto cosResult = LLVM::LoadOp::create(rewriter, loc, computeType, cosPtr);519 520    rewriter.replaceOp(op, {maybeTrunc(sinResult, inputType, rewriter),521                            maybeTrunc(cosResult, inputType, rewriter)});522    return success();523  }524 525private:526  Value maybeExt(Value operand, PatternRewriter &rewriter) const {527    if (isa<Float16Type, BFloat16Type>(operand.getType()))528      return LLVM::FPExtOp::create(rewriter, operand.getLoc(),529                                   Float32Type::get(rewriter.getContext()),530                                   operand);531    return operand;532  }533 534  Value maybeTrunc(Value operand, Type type, PatternRewriter &rewriter) const {535    if (operand.getType() != type)536      return LLVM::FPTruncOp::create(rewriter, operand.getLoc(), type, operand);537    return operand;538  }539 540  void createSincosCall(ConversionPatternRewriter &rewriter, Location loc,541                        StringRef funcName, Value input, Value sinPtr,542                        Value cosPtr, Operation *op) const {543    auto voidType = LLVM::LLVMVoidType::get(rewriter.getContext());544    auto ptrType = sinPtr.getType();545 546    SmallVector<Type> operandTypes = {input.getType(), ptrType, ptrType};547    auto funcType = LLVM::LLVMFunctionType::get(voidType, operandTypes);548 549    auto funcAttr = StringAttr::get(op->getContext(), funcName);550    auto funcOp =551        SymbolTable::lookupNearestSymbolFrom<LLVM::LLVMFuncOp>(op, funcAttr);552 553    if (!funcOp) {554      auto parentFunc = op->getParentOfType<FunctionOpInterface>();555      assert(parentFunc && "expected there to be a parent function");556      OpBuilder b(parentFunc);557 558      auto globalloc = loc->findInstanceOfOrUnknown<FileLineColLoc>();559      funcOp = LLVM::LLVMFuncOp::create(b, globalloc, funcName, funcType);560    }561 562    SmallVector<Value> callOperands = {input, sinPtr, cosPtr};563    LLVM::CallOp::create(rewriter, loc, funcOp, callOperands);564  }565};566 567template <typename OpTy>568static void populateOpPatterns(const LLVMTypeConverter &converter,569                               RewritePatternSet &patterns,570                               PatternBenefit benefit, StringRef f32Func,571                               StringRef f64Func, StringRef f32ApproxFunc = "",572                               StringRef f16Func = "") {573  patterns.add<ScalarizeVectorOpLowering<OpTy>>(converter, benefit);574  patterns.add<OpToFuncCallLowering<OpTy>>(converter, f32Func, f64Func,575                                           f32ApproxFunc, f16Func,576                                           /*i32Func=*/"", benefit);577}578 579template <typename OpTy>580static void populateIntOpPatterns(const LLVMTypeConverter &converter,581                                  RewritePatternSet &patterns,582                                  PatternBenefit benefit, StringRef i32Func) {583  patterns.add<ScalarizeVectorOpLowering<OpTy>>(converter, benefit);584  patterns.add<OpToFuncCallLowering<OpTy>>(converter, "", "", "", "", i32Func,585                                           benefit);586}587 588template <typename OpTy>589static void populateFloatIntOpPatterns(const LLVMTypeConverter &converter,590                                       RewritePatternSet &patterns,591                                       PatternBenefit benefit,592                                       StringRef f32Func, StringRef f64Func) {593  patterns.add<ScalarizeVectorOpLowering<OpTy>>(converter, benefit);594  patterns.add<OpToFuncCallLowering<OpTy>>(converter, f32Func, f64Func, "", "",595                                           /*i32Func=*/"", benefit);596}597 598void mlir::populateGpuSubgroupReduceOpLoweringPattern(599    const LLVMTypeConverter &converter, RewritePatternSet &patterns,600    PatternBenefit benefit) {601  patterns.add<GPUSubgroupReduceOpLowering>(converter, benefit);602}603 604void mlir::populateLibDeviceConversionPatterns(605    const LLVMTypeConverter &converter, RewritePatternSet &patterns,606    PatternBenefit benefit) {607  populateOpPatterns<arith::RemFOp>(converter, patterns, benefit, "__nv_fmodf",608                                    "__nv_fmod");609  populateOpPatterns<arith::MaxNumFOp>(converter, patterns, benefit,610                                       "__nv_fmaxf", "__nv_fmax");611  populateOpPatterns<arith::MinNumFOp>(converter, patterns, benefit,612                                       "__nv_fminf", "__nv_fmin");613 614  populateIntOpPatterns<math::AbsIOp>(converter, patterns, benefit, "__nv_abs");615  populateOpPatterns<math::AbsFOp>(converter, patterns, benefit, "__nv_fabsf",616                                   "__nv_fabs");617  populateOpPatterns<math::AcosOp>(converter, patterns, benefit, "__nv_acosf",618                                   "__nv_acos");619  populateOpPatterns<math::AcoshOp>(converter, patterns, benefit, "__nv_acoshf",620                                    "__nv_acosh");621  populateOpPatterns<math::AsinOp>(converter, patterns, benefit, "__nv_asinf",622                                   "__nv_asin");623  populateOpPatterns<math::AsinhOp>(converter, patterns, benefit, "__nv_asinhf",624                                    "__nv_asinh");625  populateOpPatterns<math::AtanOp>(converter, patterns, benefit, "__nv_atanf",626                                   "__nv_atan");627  populateOpPatterns<math::Atan2Op>(converter, patterns, benefit, "__nv_atan2f",628                                    "__nv_atan2");629  populateOpPatterns<math::AtanhOp>(converter, patterns, benefit, "__nv_atanhf",630                                    "__nv_atanh");631  populateOpPatterns<math::CbrtOp>(converter, patterns, benefit, "__nv_cbrtf",632                                   "__nv_cbrt");633  populateOpPatterns<math::CeilOp>(converter, patterns, benefit, "__nv_ceilf",634                                   "__nv_ceil");635  populateOpPatterns<math::CopySignOp>(converter, patterns, benefit,636                                       "__nv_copysignf", "__nv_copysign");637  populateOpPatterns<math::CosOp>(converter, patterns, benefit, "__nv_cosf",638                                  "__nv_cos", "__nv_fast_cosf");639  populateOpPatterns<math::CoshOp>(converter, patterns, benefit, "__nv_coshf",640                                   "__nv_cosh");641  populateOpPatterns<math::ErfOp>(converter, patterns, benefit, "__nv_erff",642                                  "__nv_erf");643  populateOpPatterns<math::ErfcOp>(converter, patterns, benefit, "__nv_erfcf",644                                   "__nv_erfc");645  populateOpPatterns<math::ExpOp>(converter, patterns, benefit, "__nv_expf",646                                  "__nv_exp", "__nv_fast_expf");647  populateOpPatterns<math::Exp2Op>(converter, patterns, benefit, "__nv_exp2f",648                                   "__nv_exp2");649  populateOpPatterns<math::ExpM1Op>(converter, patterns, benefit, "__nv_expm1f",650                                    "__nv_expm1");651  populateOpPatterns<math::FloorOp>(converter, patterns, benefit, "__nv_floorf",652                                    "__nv_floor");653  populateOpPatterns<math::FmaOp>(converter, patterns, benefit, "__nv_fmaf",654                                  "__nv_fma");655  // Note: libdevice uses a different name for 32-bit finite checking656  populateOpPatterns<math::IsFiniteOp>(converter, patterns, benefit,657                                       "__nv_finitef", "__nv_isfinited");658  populateOpPatterns<math::IsInfOp>(converter, patterns, benefit, "__nv_isinff",659                                    "__nv_isinfd");660  populateOpPatterns<math::IsNaNOp>(converter, patterns, benefit, "__nv_isnanf",661                                    "__nv_isnand");662  populateOpPatterns<math::LogOp>(converter, patterns, benefit, "__nv_logf",663                                  "__nv_log", "__nv_fast_logf");664  populateOpPatterns<math::Log10Op>(converter, patterns, benefit, "__nv_log10f",665                                    "__nv_log10", "__nv_fast_log10f");666  populateOpPatterns<math::Log1pOp>(converter, patterns, benefit, "__nv_log1pf",667                                    "__nv_log1p");668  populateOpPatterns<math::Log2Op>(converter, patterns, benefit, "__nv_log2f",669                                   "__nv_log2", "__nv_fast_log2f");670  populateOpPatterns<math::PowFOp>(converter, patterns, benefit, "__nv_powf",671                                   "__nv_pow", "__nv_fast_powf");672  populateFloatIntOpPatterns<math::FPowIOp>(converter, patterns, benefit,673                                            "__nv_powif", "__nv_powi");674  populateOpPatterns<math::RoundOp>(converter, patterns, benefit, "__nv_roundf",675                                    "__nv_round");676  populateOpPatterns<math::RoundEvenOp>(converter, patterns, benefit,677                                        "__nv_rintf", "__nv_rint");678  populateOpPatterns<math::RsqrtOp>(converter, patterns, benefit, "__nv_rsqrtf",679                                    "__nv_rsqrt");680  populateOpPatterns<math::SinOp>(converter, patterns, benefit, "__nv_sinf",681                                  "__nv_sin", "__nv_fast_sinf");682  populateOpPatterns<math::SinhOp>(converter, patterns, benefit, "__nv_sinhf",683                                   "__nv_sinh");684  populateOpPatterns<math::SqrtOp>(converter, patterns, benefit, "__nv_sqrtf",685                                   "__nv_sqrt");686  populateOpPatterns<math::TanOp>(converter, patterns, benefit, "__nv_tanf",687                                  "__nv_tan", "__nv_fast_tanf");688  populateOpPatterns<math::TanhOp>(converter, patterns, benefit, "__nv_tanhf",689                                   "__nv_tanh");690 691  // Custom pattern for sincos since it returns two values692  patterns.add<SincosOpLowering>(converter, benefit);693}694 695void mlir::populateGpuToNVVMConversionPatterns(696    const LLVMTypeConverter &converter, RewritePatternSet &patterns,697    PatternBenefit benefit) {698  using gpu::index_lowering::IndexKind;699  using gpu::index_lowering::IntrType;700 701  // TODO: Pass benefit to generated patterns.702  populateWithGenerated(patterns);703 704  patterns.add<GPUPrintfOpToVPrintfLowering, AssertOpToAssertfailLowering>(705      converter, benefit);706  patterns.add<707      gpu::index_lowering::OpLowering<gpu::ThreadIdOp, NVVM::ThreadIdXOp,708                                      NVVM::ThreadIdYOp, NVVM::ThreadIdZOp>>(709      converter, IndexKind::Block, IntrType::Id, benefit);710  patterns.add<711      gpu::index_lowering::OpLowering<gpu::BlockDimOp, NVVM::BlockDimXOp,712                                      NVVM::BlockDimYOp, NVVM::BlockDimZOp>>(713      converter, IndexKind::Block, IntrType::Dim, benefit);714  patterns.add<715      gpu::index_lowering::OpLowering<gpu::ClusterIdOp, NVVM::ClusterIdXOp,716                                      NVVM::ClusterIdYOp, NVVM::ClusterIdZOp>>(717      converter, IndexKind::Other, IntrType::Id, benefit);718  patterns.add<gpu::index_lowering::OpLowering<719      gpu::ClusterDimOp, NVVM::ClusterDimXOp, NVVM::ClusterDimYOp,720      NVVM::ClusterDimZOp>>(converter, IndexKind::Other, IntrType::Dim,721                            benefit);722  patterns.add<gpu::index_lowering::OpLowering<723      gpu::ClusterBlockIdOp, NVVM::BlockInClusterIdXOp,724      NVVM::BlockInClusterIdYOp, NVVM::BlockInClusterIdZOp>>(725      converter, IndexKind::Other, IntrType::Id, benefit);726  patterns.add<gpu::index_lowering::OpLowering<727      gpu::ClusterDimBlocksOp, NVVM::ClusterDimBlocksXOp,728      NVVM::ClusterDimBlocksYOp, NVVM::ClusterDimBlocksZOp>>(729      converter, IndexKind::Other, IntrType::Dim, benefit);730  patterns.add<gpu::index_lowering::OpLowering<731      gpu::BlockIdOp, NVVM::BlockIdXOp, NVVM::BlockIdYOp, NVVM::BlockIdZOp>>(732      converter, IndexKind::Grid, IntrType::Id, benefit);733  patterns.add<gpu::index_lowering::OpLowering<734      gpu::GridDimOp, NVVM::GridDimXOp, NVVM::GridDimYOp, NVVM::GridDimZOp>>(735      converter, IndexKind::Grid, IntrType::Dim, benefit);736  patterns.add<GPULaneIdOpToNVVM, GPUShuffleOpLowering, GPUReturnOpLowering>(737      converter, benefit);738 739  patterns.add<GPUDynamicSharedMemoryOpLowering>(740      converter, NVVM::kSharedMemoryAlignmentBit, benefit);741 742  // Explicitly drop memory space when lowering private memory743  // attributions since NVVM models it as `alloca`s in the default744  // memory space and does not support `alloca`s with addrspace(5).745  patterns.add<GPUFuncOpLowering>(746      converter,747      GPUFuncOpLoweringOptions{748          /*allocaAddrSpace=*/0,749          /*workgroupAddrSpace=*/750          static_cast<unsigned>(NVVM::NVVMMemorySpace::Shared),751          StringAttr::get(&converter.getContext(),752                          NVVM::NVVMDialect::getKernelFuncAttrName()),753          StringAttr::get(&converter.getContext(),754                          NVVM::NVVMDialect::getMaxntidAttrName())},755      benefit);756 757  populateLibDeviceConversionPatterns(converter, patterns, benefit);758}759 760//===----------------------------------------------------------------------===//761// NVVMTargetAttr convert to LLVM attr interface762//===----------------------------------------------------------------------===//763 764namespace {765struct NVVMTargetConvertToLLVMAttrInterface766    : public ConvertToLLVMAttrInterface::ExternalModel<767          NVVMTargetConvertToLLVMAttrInterface, NVVM::NVVMTargetAttr> {768  /// Configure GPU to NVVM.769  void populateConvertToLLVMConversionPatterns(770      Attribute attr, ConversionTarget &target,771      LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) const;772};773} // namespace774 775void NVVMTargetConvertToLLVMAttrInterface::776    populateConvertToLLVMConversionPatterns(Attribute attr,777                                            ConversionTarget &target,778                                            LLVMTypeConverter &typeConverter,779                                            RewritePatternSet &patterns) const {780  configureGpuToNVVMConversionLegality(target);781  configureGpuToNVVMTypeConverter(typeConverter);782  populateGpuToNVVMConversionPatterns(typeConverter, patterns);783}784 785void mlir::NVVM::registerConvertGpuToNVVMInterface(DialectRegistry &registry) {786  registry.addExtension(+[](MLIRContext *ctx, NVVMDialect *dialect) {787    NVVMTargetAttr::attachInterface<NVVMTargetConvertToLLVMAttrInterface>(*ctx);788  });789}790