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 ®istry) 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 ®istry) {786 registry.addExtension(+[](MLIRContext *ctx, NVVMDialect *dialect) {787 NVVMTargetAttr::attachInterface<NVVMTargetConvertToLLVMAttrInterface>(*ctx);788 });789}790