19#include "llvm/ADT/StringExtras.h"
20#include "llvm/ADT/iterator_range.h"
21#include "llvm/IR/IRBuilder.h"
22#include "llvm/IR/IntrinsicsNVPTX.h"
23#include "llvm/Support/FormatVariadic.h"
24#include "llvm/Support/NVVMAttributes.h"
31#define REDUX_F32_ID_IMPL(op, abs, hasNaN) \
32 hasNaN ? llvm::Intrinsic::nvvm_redux_sync_f##op##abs##_NaN \
33 : llvm::Intrinsic::nvvm_redux_sync_f##op##abs
35#define GET_REDUX_F32_ID(op, hasAbs, hasNaN) \
36 hasAbs ? REDUX_F32_ID_IMPL(op, _abs, hasNaN) : REDUX_F32_ID_IMPL(op, , hasNaN)
39 NVVM::ReductionKind kind,
40 bool hasAbs,
bool hasNaN) {
42 case NVVM::ReductionKind::ADD:
43 return llvm::Intrinsic::nvvm_redux_sync_add;
44 case NVVM::ReductionKind::UMAX:
45 return llvm::Intrinsic::nvvm_redux_sync_umax;
46 case NVVM::ReductionKind::UMIN:
47 return llvm::Intrinsic::nvvm_redux_sync_umin;
48 case NVVM::ReductionKind::AND:
49 return llvm::Intrinsic::nvvm_redux_sync_and;
50 case NVVM::ReductionKind::OR:
51 return llvm::Intrinsic::nvvm_redux_sync_or;
52 case NVVM::ReductionKind::XOR:
53 return llvm::Intrinsic::nvvm_redux_sync_xor;
54 case NVVM::ReductionKind::MAX:
55 return llvm::Intrinsic::nvvm_redux_sync_max;
56 case NVVM::ReductionKind::MIN:
57 return llvm::Intrinsic::nvvm_redux_sync_min;
58 case NVVM::ReductionKind::FMIN:
60 case NVVM::ReductionKind::FMAX:
63 llvm_unreachable(
"unknown reduction kind");
71 resultType = cast<llvm::StructType>(resultType)->getElementType(0);
73 case NVVM::ShflKind::bfly:
74 return resultType->isFloatTy()
75 ? llvm::Intrinsic::nvvm_shfl_sync_bfly_f32p
76 : llvm::Intrinsic::nvvm_shfl_sync_bfly_i32p;
77 case NVVM::ShflKind::up:
78 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_up_f32p
79 : llvm::Intrinsic::nvvm_shfl_sync_up_i32p;
80 case NVVM::ShflKind::down:
81 return resultType->isFloatTy()
82 ? llvm::Intrinsic::nvvm_shfl_sync_down_f32p
83 : llvm::Intrinsic::nvvm_shfl_sync_down_i32p;
84 case NVVM::ShflKind::idx:
85 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_idx_f32p
86 : llvm::Intrinsic::nvvm_shfl_sync_idx_i32p;
90 case NVVM::ShflKind::bfly:
91 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_bfly_f32
92 : llvm::Intrinsic::nvvm_shfl_sync_bfly_i32;
93 case NVVM::ShflKind::up:
94 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_up_f32
95 : llvm::Intrinsic::nvvm_shfl_sync_up_i32;
96 case NVVM::ShflKind::down:
97 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_down_f32
98 : llvm::Intrinsic::nvvm_shfl_sync_down_i32;
99 case NVVM::ShflKind::idx:
100 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_idx_f32
101 : llvm::Intrinsic::nvvm_shfl_sync_idx_i32;
104 llvm_unreachable(
"unknown shuffle kind");
108 NVVM::MatchSyncKind kind) {
110 case NVVM::MatchSyncKind::any:
111 return valType.
isInteger(32) ? llvm::Intrinsic::nvvm_match_any_sync_i32
112 : llvm::Intrinsic::nvvm_match_any_sync_i64;
113 case NVVM::MatchSyncKind::all:
117 return valType.
isInteger(32) ? llvm::Intrinsic::nvvm_match_all_sync_i32p
118 : llvm::Intrinsic::nvvm_match_all_sync_i64p;
120 llvm_unreachable(
"unsupported match sync kind");
125 case NVVM::VoteSyncKind::any:
126 return llvm::Intrinsic::nvvm_vote_any_sync;
127 case NVVM::VoteSyncKind::all:
128 return llvm::Intrinsic::nvvm_vote_all_sync;
129 case NVVM::VoteSyncKind::ballot:
130 return llvm::Intrinsic::nvvm_vote_ballot_sync;
131 case NVVM::VoteSyncKind::uni:
132 return llvm::Intrinsic::nvvm_vote_uni_sync;
134 llvm_unreachable(
"unsupported vote kind");
137static llvm::Intrinsic::ID
139 NVVM::LdStMatrixShapeAttr
shape,
140 NVVM::LdStMatrixEltType eltType) {
144 return (layout == NVVM::MMALayout::row)
145 ? llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x1_b16
147 nvvm_ldmatrix_sync_aligned_m8n8_x1_trans_b16;
149 return (layout == NVVM::MMALayout::row)
150 ? llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x2_b16
152 nvvm_ldmatrix_sync_aligned_m8n8_x2_trans_b16;
154 return (layout == NVVM::MMALayout::row)
155 ? llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x4_b16
157 nvvm_ldmatrix_sync_aligned_m8n8_x4_trans_b16;
159 }
else if (
shape.getM() == 8 &&
shape.getN() == 16) {
160 if (eltType == NVVM::LdStMatrixEltType::B8X16_B6X16_P32) {
163 return llvm::Intrinsic::
164 nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b6x16_p32;
166 return llvm::Intrinsic::
167 nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b6x16_p32;
169 return llvm::Intrinsic::
170 nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b6x16_p32;
172 }
else if (eltType == NVVM::LdStMatrixEltType::B8X16_B4X16_P64) {
175 return llvm::Intrinsic::
176 nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b4x16_p64;
178 return llvm::Intrinsic::
179 nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b4x16_p64;
181 return llvm::Intrinsic::
182 nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b4x16_p64;
185 }
else if (
shape.getM() == 16 &&
shape.getN() == 16) {
186 if (eltType == NVVM::LdStMatrixEltType::B8) {
189 return llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8;
191 return llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8;
193 }
else if (eltType == NVVM::LdStMatrixEltType::B8X16_B6X16_P32) {
196 return llvm::Intrinsic::
197 nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b6x16_p32;
199 return llvm::Intrinsic::
200 nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b6x16_p32;
202 }
else if (eltType == NVVM::LdStMatrixEltType::B8X16_B4X16_P64) {
205 return llvm::Intrinsic::
206 nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b4x16_p64;
208 return llvm::Intrinsic::
209 nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b4x16_p64;
213 llvm_unreachable(
"unknown ldmatrix kind");
217static llvm::Intrinsic::ID
219 NVVM::LdStMatrixShapeAttr
shape,
220 NVVM::LdStMatrixEltType eltType) {
224 return (layout == NVVM::MMALayout::row)
225 ? llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x1_b16
227 nvvm_stmatrix_sync_aligned_m8n8_x1_trans_b16;
229 return (layout == NVVM::MMALayout::row)
230 ? llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x2_b16
232 nvvm_stmatrix_sync_aligned_m8n8_x2_trans_b16;
234 return (layout == NVVM::MMALayout::row)
235 ? llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x4_b16
237 nvvm_stmatrix_sync_aligned_m8n8_x4_trans_b16;
239 }
else if (
shape.getM() == 16 &&
shape.getN() == 8) {
242 return llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x1_trans_b8;
244 return llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x2_trans_b8;
246 return llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x4_trans_b8;
249 llvm_unreachable(
"unknown stmatrix kind");
253 NVVM::ProxyKind toProxy,
254 NVVM::MemScopeKind scope,
256 if (fromProxy == NVVM::ProxyKind::GENERIC &&
257 toProxy == NVVM::ProxyKind::TENSORMAP) {
259 case NVVM::MemScopeKind::CTA: {
261 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_release_cta;
262 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_acquire_cta;
264 case NVVM::MemScopeKind::CLUSTER: {
266 return llvm::Intrinsic::
267 nvvm_fence_proxy_tensormap_generic_release_cluster;
268 return llvm::Intrinsic::
269 nvvm_fence_proxy_tensormap_generic_acquire_cluster;
271 case NVVM::MemScopeKind::GPU: {
273 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_release_gpu;
274 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_acquire_gpu;
276 case NVVM::MemScopeKind::SYS: {
278 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_release_sys;
279 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_acquire_sys;
282 llvm_unreachable(
"Unknown scope for uni-directional fence.proxy operation");
284 llvm_unreachable(
"Unsupported proxy kinds");
289 case NVVM::MemScopeKind::CTA:
290 return llvm::Intrinsic::nvvm_membar_cta;
291 case NVVM::MemScopeKind::CLUSTER:
292 return llvm::Intrinsic::nvvm_fence_sc_cluster;
293 case NVVM::MemScopeKind::GPU:
294 return llvm::Intrinsic::nvvm_membar_gl;
295 case NVVM::MemScopeKind::SYS:
296 return llvm::Intrinsic::nvvm_membar_sys;
298 llvm_unreachable(
"Unknown scope for memory barrier");
301#define TCGEN05LD(SHAPE, NUM) llvm::Intrinsic::nvvm_tcgen05_ld_##SHAPE##_##NUM
303static llvm::Intrinsic::ID
305 llvm::Intrinsic::ID Shape16x64b[] = {
311 llvm::Intrinsic::ID Shape16x128b[] = {
317 llvm::Intrinsic::ID Shape16x256b[] = {
322 llvm::Intrinsic::ID Shape16x32bx2[] = {
329 llvm::Intrinsic::ID Shape32x32b[] = {
337 unsigned Idx = std::log2(num);
340 case NVVM::Tcgen05LdStShape::SHAPE_16X64B:
341 return Shape16x64b[Idx];
342 case NVVM::Tcgen05LdStShape::SHAPE_16X128B:
343 return Shape16x128b[Idx - 1];
344 case NVVM::Tcgen05LdStShape::SHAPE_16X256B:
345 return Shape16x256b[Idx - 2];
346 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
347 return Shape32x32b[Idx];
348 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
349 return Shape16x32bx2[Idx];
351 llvm_unreachable(
"unhandled tcgen05.ld lowering");
354#define TCGEN05ST(SHAPE, NUM) llvm::Intrinsic::nvvm_tcgen05_st_##SHAPE##_##NUM
356static llvm::Intrinsic::ID
358 llvm::Intrinsic::ID Shape16x64b[] = {
364 llvm::Intrinsic::ID Shape16x128b[] = {
370 llvm::Intrinsic::ID Shape16x256b[] = {
375 llvm::Intrinsic::ID Shape16x32bx2[] = {
382 llvm::Intrinsic::ID Shape32x32b[] = {
390 unsigned Idx = std::log2(num);
393 case NVVM::Tcgen05LdStShape::SHAPE_16X64B:
394 return Shape16x64b[Idx];
395 case NVVM::Tcgen05LdStShape::SHAPE_16X128B:
396 return Shape16x128b[Idx - 1];
397 case NVVM::Tcgen05LdStShape::SHAPE_16X256B:
398 return Shape16x256b[Idx - 2];
399 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
400 return Shape32x32b[Idx];
401 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
402 return Shape16x32bx2[Idx];
404 llvm_unreachable(
"unhandled tcgen05.st lowering");
408 return order == NVVM::MemOrderKind::ACQUIRE
410 nvvm_fence_acquire_sync_restrict_space_cluster_scope_cluster
412 nvvm_fence_release_sync_restrict_space_cta_scope_cluster;
415static llvm::Intrinsic::ID
418 case NVVM::ProxyKind::alias:
419 return llvm::Intrinsic::nvvm_fence_proxy_alias;
420 case NVVM::ProxyKind::async:
421 return llvm::Intrinsic::nvvm_fence_proxy_async;
422 case NVVM::ProxyKind::async_global:
423 return llvm::Intrinsic::nvvm_fence_proxy_async_global;
424 case NVVM::ProxyKind::async_shared:
425 return *space == NVVM::SharedSpace::shared_cta
426 ? llvm::Intrinsic::nvvm_fence_proxy_async_shared_cta
427 : llvm::Intrinsic::nvvm_fence_proxy_async_shared_cluster;
429 llvm_unreachable(
"unsupported proxy kind");
433static llvm::Intrinsic::ID
435 return order == NVVM::MemOrderKind::ACQUIRE
437 nvvm_fence_proxy_async_generic_acquire_sync_restrict_space_cluster_scope_cluster
439 nvvm_fence_proxy_async_generic_release_sync_restrict_space_cta_scope_cluster;
442static llvm::RoundingMode
445 case NVVM::FPRoundingMode::RN:
446 return llvm::RoundingMode::NearestTiesToEven;
447 case NVVM::FPRoundingMode::RM:
448 return llvm::RoundingMode::TowardNegative;
449 case NVVM::FPRoundingMode::RP:
450 return llvm::RoundingMode::TowardPositive;
451 case NVVM::FPRoundingMode::RZ:
452 return llvm::RoundingMode::TowardZero;
455 assert(rndMode == NVVM::FPRoundingMode::NONE &&
456 "unsupported rounding mode for nvvm fp arithmetic");
457 return llvm::RoundingMode::NearestTiesToEven;
467 llvm::Intrinsic::ID IID, llvm::Type *opTypeLLVM,
469 llvm::Type *retType) {
470 if (opTypeLLVM->isVectorTy() && (opTypeLLVM->getScalarType()->isFloatTy() ||
471 opTypeLLVM->getScalarType()->isDoubleTy())) {
472 llvm::Value *
result = llvm::PoisonValue::get(
473 llvm::FixedVectorType::get(opTypeLLVM->getScalarType(), 2));
474 for (
int64_t i = 0; i < 2; ++i) {
476 for (llvm::Value *op : operands)
477 scalarArgs.push_back(
478 op->getType()->isVectorTy()
479 ? builder.CreateExtractElement(op, builder.getInt32(i))
482 result = builder.CreateInsertElement(
result, res, builder.getInt32(i));
490void NVVM::AddFOp::lowerAddFToLLVMIR(llvm::Value *argLHS, llvm::Value *argRHS,
491 Value res, NVVM::FPRoundingMode rndMode,
492 NVVM::SaturationMode satMode,
bool isFTZ,
494 llvm::IRBuilderBase &builder) {
495 llvm::Type *opTypeLLVM = argLHS->getType();
496 bool isSat = satMode != NVVM::SaturationMode::NONE;
498 static constexpr llvm::Intrinsic::ID addIDs[2][2] = {
499 {llvm::Intrinsic::nvvm_fadd, llvm::Intrinsic::nvvm_fadd_sat},
500 {llvm::Intrinsic::nvvm_fadd_ftz, llvm::Intrinsic::nvvm_fadd_ftz_sat}};
502 llvm::Intrinsic::ID
id = addIDs[isFTZ][isSat];
503 llvm::Value *rnd = builder.getInt32(
508 llvm::Type *scalarTypeLLVM = opTypeLLVM->getScalarType();
509 if (opTypeLLVM->isVectorTy() && (scalarTypeLLVM->isDoubleTy() ||
510 (isSat && scalarTypeLLVM->isFloatTy()))) {
512 {argLHS, argRHS, rnd},
522 llvm::IRBuilderBase &builder) {
523 auto thisOp = cast<NVVM::FmaOp>(op);
524 mlir::NVVM::FPRoundingMode rndMode = thisOp.getRnd();
525 unsigned rndIndex =
static_cast<unsigned>(rndMode) - 1;
526 mlir::NVVM::SaturationMode satMode = thisOp.getSat();
527 bool isFTZ = thisOp.getFtz();
528 bool isRelu = thisOp.getRelu();
529 bool isSat = satMode == NVVM::SaturationMode::SAT;
530 bool isOOB = thisOp.getOob();
532 mlir::Type opType = thisOp.getRes().getType();
534 bool isVectorFma = opTypeLLVM->isVectorTy();
540 static constexpr llvm::Intrinsic::ID f16IDs[] = {
541 llvm::Intrinsic::nvvm_fma_rn_f16,
542 llvm::Intrinsic::nvvm_fma_rn_f16x2,
543 llvm::Intrinsic::nvvm_fma_rn_ftz_f16,
544 llvm::Intrinsic::nvvm_fma_rn_ftz_f16x2,
545 llvm::Intrinsic::nvvm_fma_rn_sat_f16,
546 llvm::Intrinsic::nvvm_fma_rn_sat_f16x2,
547 llvm::Intrinsic::nvvm_fma_rn_ftz_sat_f16,
548 llvm::Intrinsic::nvvm_fma_rn_ftz_sat_f16x2,
549 llvm::Intrinsic::nvvm_fma_rn_relu_f16,
550 llvm::Intrinsic::nvvm_fma_rn_relu_f16x2,
551 llvm::Intrinsic::nvvm_fma_rn_ftz_relu_f16,
552 llvm::Intrinsic::nvvm_fma_rn_ftz_relu_f16x2};
554 static constexpr llvm::Intrinsic::ID bf16IDs[] = {
555 llvm::Intrinsic::nvvm_fma_rn_bf16, llvm::Intrinsic::nvvm_fma_rn_bf16x2,
556 llvm::Intrinsic::nvvm_fma_rn_relu_bf16,
557 llvm::Intrinsic::nvvm_fma_rn_relu_bf16x2};
559 static constexpr llvm::Intrinsic::ID f32IDs[] = {
560 llvm::Intrinsic::nvvm_fma_rn_f,
561 llvm::Intrinsic::nvvm_fma_rm_f,
562 llvm::Intrinsic::nvvm_fma_rp_f,
563 llvm::Intrinsic::nvvm_fma_rz_f,
564 llvm::Intrinsic::nvvm_fma_rn_sat_f,
565 llvm::Intrinsic::nvvm_fma_rm_sat_f,
566 llvm::Intrinsic::nvvm_fma_rp_sat_f,
567 llvm::Intrinsic::nvvm_fma_rz_sat_f,
568 llvm::Intrinsic::nvvm_fma_rn_ftz_f,
569 llvm::Intrinsic::nvvm_fma_rm_ftz_f,
570 llvm::Intrinsic::nvvm_fma_rp_ftz_f,
571 llvm::Intrinsic::nvvm_fma_rz_ftz_f,
572 llvm::Intrinsic::nvvm_fma_rn_ftz_sat_f,
573 llvm::Intrinsic::nvvm_fma_rm_ftz_sat_f,
574 llvm::Intrinsic::nvvm_fma_rp_ftz_sat_f,
575 llvm::Intrinsic::nvvm_fma_rz_ftz_sat_f,
578 static constexpr llvm::Intrinsic::ID f64IDs[] = {
579 llvm::Intrinsic::nvvm_fma_rn_d, llvm::Intrinsic::nvvm_fma_rm_d,
580 llvm::Intrinsic::nvvm_fma_rp_d, llvm::Intrinsic::nvvm_fma_rz_d};
582 auto fmaIntrinsic = [&](llvm::Intrinsic::ID IID,
583 llvm::Type *retType) -> llvm::Value * {
585 builder, IID, opTypeLLVM, {argA, argB, argC}, retType);
589 if (opTypeLLVM->getScalarType()->isHalfTy()) {
592 result = fmaIntrinsic(isRelu ? llvm::Intrinsic::nvvm_fma_rn_oob_relu
593 : llvm::Intrinsic::nvvm_fma_rn_oob,
597 (isRelu << 3) | (isSat << 2) | (isFTZ << 1) |
606 if (opTypeLLVM->getScalarType()->isBFloatTy()) {
609 result = fmaIntrinsic(isRelu ? llvm::Intrinsic::nvvm_fma_rn_oob_relu
610 : llvm::Intrinsic::nvvm_fma_rn_oob,
613 unsigned index = (isRelu << 1) | isVectorFma;
621 if (opTypeLLVM->getScalarType()->isDoubleTy()) {
623 fmaIntrinsic(f64IDs[rndIndex], opTypeLLVM->getScalarType()));
628 const unsigned numRndModes = 4;
629 if (opTypeLLVM->getScalarType()->isFloatTy()) {
630 unsigned index = ((isFTZ << 1) | isSat) * numRndModes + rndIndex;
632 fmaIntrinsic(f32IDs[
index], opTypeLLVM->getScalarType()));
640class NVVMDialectLLVMIRTranslationInterface
641 :
public LLVMTranslationDialectInterface {
643 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
648 convertOperation(Operation *op, llvm::IRBuilderBase &builder,
649 LLVM::ModuleTranslation &moduleTranslation)
const final {
653 if (!builder.GetInsertBlock())
655 "cannot be translated to LLVM IR without an active insertion "
656 "point; make sure the op is inside a function");
657 Operation &opInst = *op;
658#include "mlir/Dialect/LLVMIR/NVVMConversions.inc"
666 amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,
667 NamedAttribute attribute,
668 LLVM::ModuleTranslation &moduleTranslation)
const final {
669 if (
auto globalOp = dyn_cast<LLVM::GlobalOp>(op)) {
670 if (attribute.getName() == NVVM::NVVMDialect::getManagedAttrName()) {
671 auto *gv = cast<llvm::GlobalVariable>(
672 moduleTranslation.lookupGlobal(globalOp));
673 llvm::Module *m = gv->getParent();
674 llvm::LLVMContext &ctx = m->getContext();
675 llvm::NamedMDNode *md = m->getOrInsertNamedMetadata(
"nvvm.annotations");
676 md->addOperand(llvm::MDNode::get(
677 ctx, {llvm::ConstantAsMetadata::get(gv),
678 llvm::MDString::get(ctx,
"managed"),
679 llvm::ConstantAsMetadata::get(llvm::ConstantInt::get(
680 llvm::Type::getInt32Ty(ctx), 1))}));
685 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
688 llvm::Function *llvmFunc = moduleTranslation.lookupFunction(func.getName());
690 if (attribute.getName() == NVVM::NVVMDialect::getMaxntidAttrName()) {
691 if (!isa<DenseI32ArrayAttr>(attribute.getValue()))
693 auto values = cast<DenseI32ArrayAttr>(attribute.getValue());
694 const std::string attr = llvm::formatv(
695 "{0:$[,]}", llvm::make_range(values.asArrayRef().begin(),
696 values.asArrayRef().end()));
697 llvmFunc->addFnAttr(llvm::NVVMAttr::MaxNTID, attr);
698 }
else if (attribute.getName() == NVVM::NVVMDialect::getReqntidAttrName()) {
699 if (!isa<DenseI32ArrayAttr>(attribute.getValue()))
701 auto values = cast<DenseI32ArrayAttr>(attribute.getValue());
702 const std::string attr = llvm::formatv(
703 "{0:$[,]}", llvm::make_range(values.asArrayRef().begin(),
704 values.asArrayRef().end()));
705 llvmFunc->addFnAttr(llvm::NVVMAttr::ReqNTID, attr);
706 }
else if (attribute.getName() ==
707 NVVM::NVVMDialect::getClusterDimAttrName()) {
708 if (!isa<DenseI32ArrayAttr>(attribute.getValue()))
710 auto values = cast<DenseI32ArrayAttr>(attribute.getValue());
711 const std::string attr = llvm::formatv(
712 "{0:$[,]}", llvm::make_range(values.asArrayRef().begin(),
713 values.asArrayRef().end()));
714 llvmFunc->addFnAttr(llvm::NVVMAttr::ClusterDim, attr);
715 }
else if (attribute.getName() ==
716 NVVM::NVVMDialect::getClusterMaxBlocksAttrName()) {
717 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
718 llvmFunc->addFnAttr(llvm::NVVMAttr::MaxClusterRank,
719 llvm::utostr(value.getInt()));
720 }
else if (attribute.getName() ==
721 NVVM::NVVMDialect::getMinctasmAttrName()) {
722 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
723 llvmFunc->addFnAttr(llvm::NVVMAttr::MinCTASm,
724 llvm::utostr(value.getInt()));
725 }
else if (attribute.getName() == NVVM::NVVMDialect::getMaxnregAttrName()) {
726 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
727 llvmFunc->addFnAttr(llvm::NVVMAttr::MaxNReg,
728 llvm::utostr(value.getInt()));
729 }
else if (attribute.getName() ==
730 NVVM::NVVMDialect::getKernelFuncAttrName()) {
731 llvmFunc->setCallingConv(llvm::CallingConv::PTX_Kernel);
732 }
else if (attribute.getName() ==
733 NVVM::NVVMDialect::getBlocksAreClustersAttrName()) {
734 llvmFunc->addFnAttr(llvm::NVVMAttr::BlocksAreClusters);
742 LLVM::ModuleTranslation &moduleTranslation)
const final {
744 llvm::LLVMContext &llvmContext = moduleTranslation.getLLVMContext();
745 llvm::Function *llvmFunc =
746 moduleTranslation.lookupFunction(funcOp.getName());
748 if (attribute.getName() == NVVM::NVVMDialect::getGridConstantAttrName()) {
749 llvmFunc->addParamAttr(
751 llvm::Attribute::get(llvmContext, llvm::NVVMAttr::GridConstant));
759 registry.
insert<NVVM::NVVMDialect>();
761 dialect->addInterfaces<NVVMDialectLLVMIRTranslationInterface>();
static LogicalResult convertParameterAttr(llvm::AttrBuilder &attrBuilder, llvm::Attribute::AttrKind llvmKind, NamedAttribute namedAttr, ModuleTranslation &moduleTranslation, Location loc)
static llvm::Intrinsic::ID getLdMatrixIntrinsicId(NVVM::MMALayout layout, int32_t num, NVVM::LdStMatrixShapeAttr shape, NVVM::LdStMatrixEltType eltType)
static llvm::Intrinsic::ID getFenceProxyID(NVVM::ProxyKind kind, std::optional< NVVM::SharedSpace > space)
#define GET_REDUX_F32_ID(op, hasAbs, hasNaN)
static llvm::Intrinsic::ID getStMatrixIntrinsicId(NVVM::MMALayout layout, int32_t num, NVVM::LdStMatrixShapeAttr shape, NVVM::LdStMatrixEltType eltType)
Return the intrinsic ID associated with stmatrix for the given paramters.
static llvm::Intrinsic::ID getTcgen05StIntrinsicID(mlir::NVVM::Tcgen05LdStShape shape, uint32_t num)
static llvm::Intrinsic::ID getTcgen05LdIntrinsicID(mlir::NVVM::Tcgen05LdStShape shape, uint32_t num)
static unsigned getMembarIntrinsicID(NVVM::MemScopeKind scope)
static unsigned getUnidirectionalFenceProxyID(NVVM::ProxyKind fromProxy, NVVM::ProxyKind toProxy, NVVM::MemScopeKind scope, bool isRelease)
llvm::CallInst * createIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID intrinsic, ArrayRef< llvm::Value * > args={}, ArrayRef< llvm::Type * > tys={})
Creates a call to an LLVM IR intrinsic function with the given arguments.
static llvm::Intrinsic::ID getFenceProxySyncRestrictID(NVVM::MemOrderKind order)
#define TCGEN05ST(SHAPE, NUM)
static llvm::Intrinsic::ID getReduxIntrinsicId(llvm::Type *resultType, NVVM::ReductionKind kind, bool hasAbs, bool hasNaN)
static llvm::RoundingMode getLLVMRoundingModeForFPArith(NVVM::FPRoundingMode rndMode)
static llvm::Value * createScalarizedIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID IID, llvm::Type *opTypeLLVM, ArrayRef< llvm::Value * > operands, llvm::Type *retType)
#define TCGEN05LD(SHAPE, NUM)
static llvm::Intrinsic::ID getFenceSyncRestrictID(NVVM::MemOrderKind order)
static llvm::Intrinsic::ID getShflIntrinsicId(llvm::Type *resultType, NVVM::ShflKind kind, bool withPredicate)
static llvm::Intrinsic::ID getVoteSyncIntrinsicId(NVVM::VoteSyncKind kind)
static llvm::Intrinsic::ID getMatchSyncIntrinsicId(Type valType, NVVM::MatchSyncKind kind)
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Implementation class for module translation.
llvm::Value * lookupValue(Value value) const
Finds an LLVM IR value corresponding to the given MLIR value.
llvm::Type * convertType(Type type)
Converts the type from MLIR LLVM dialect to LLVM.
void mapValue(Value mlir, llvm::Value *llvm)
Stores the mapping between an MLIR value and its LLVM IR counterpart.
MLIRContext is the top-level object for a collection of MLIR operations.
void appendDialectRegistry(const DialectRegistry ®istry)
Append the contents of the given dialect registry to the registry associated with this context.
Operation is the basic unit of execution within MLIR.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isInteger() const
Return true if this is an integer type (with the specified width).
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
llvm::CallInst * createIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID intrinsic, ArrayRef< llvm::Value * > args={}, ArrayRef< llvm::Type * > tys={})
Creates a call to an LLVM IR intrinsic function with the given arguments.
Include the generated interface declarations.
void registerNVVMDialectTranslation(DialectRegistry ®istry)
Register the NVVM dialect and the translation from it to the LLVM IR in the given registry;.