32#include "llvm/ADT/STLExtras.h"
33#include "llvm/ADT/TypeSwitch.h"
34#include "llvm/IR/IRBuilder.h"
35#include "llvm/IR/NVVMIntrinsicUtils.h"
36#include "llvm/Support/Casting.h"
37#include "llvm/Support/FormatVariadic.h"
38#include "llvm/Support/NVPTXAddrSpace.h"
39#include "llvm/Support/raw_ostream.h"
48#include "mlir/Dialect/LLVMIR/NVVMOpsDialect.cpp.inc"
49#include "mlir/Dialect/LLVMIR/NVVMOpsEnums.cpp.inc"
51static constexpr unsigned notIntrinsic = llvm::Intrinsic::not_intrinsic;
58 auto ptrTy = llvm::cast<LLVM::LLVMPointerType>(
ptr.getType());
59 return ptrTy.getAddressSpace() ==
static_cast<unsigned>(targetAS);
76 NVVMMemorySpace targetAS) {
77 unsigned AS =
static_cast<unsigned>(targetAS);
78 return builder.CreateAddrSpaceCast(
79 ptr, llvm::PointerType::get(builder.getContext(), AS));
83static llvm::nvvm::CTAGroupKind
86 case NVVM::CTAGroupKind::CTA_1:
87 return llvm::nvvm::CTAGroupKind::CG_1;
88 case NVVM::CTAGroupKind::CTA_2:
89 return llvm::nvvm::CTAGroupKind::CG_2;
91 llvm_unreachable(
"unsupported cta_group value");
103 size_t numIm2ColOffsets,
105 if (tensorDims < 1 || tensorDims > 5)
106 return emitError(loc,
"expects coordinates between 1 to 5 dimension");
114 "to use im2col mode, the tensor has to be at least 3-dimensional");
116 if (numIm2ColOffsets && (tensorDims != (numIm2ColOffsets + 2)))
118 loc,
"im2col offsets must be 2 less than number of coordinates");
123LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOp::verify() {
124 TMAStoreMode mode = getMode();
128 if (getPredicate()) {
129 if (mode != TMAStoreMode::TILE)
130 return emitError(
"Inline-ptx lowering supported only for Tile mode.");
131 if (getL2CacheHint())
132 return emitError(
"Inline-ptx lowering unsupported with L2 cache-hint.");
137 case TMAStoreMode::TILE:
139 case TMAStoreMode::IM2COL:
140 case TMAStoreMode::IM2COL_W:
142 case TMAStoreMode::TILE_SCATTER4:
144 return emitError(
"Scatter4 mode expects 5 coordinates");
149LogicalResult CpAsyncOp::verify() {
150 if (getModifier() != LoadCacheModifierKind::CG &&
151 getModifier() != LoadCacheModifierKind::CA)
152 return emitError(
"Only CG and CA cache modifiers are supported.");
153 if (getSize() != 4 && getSize() != 8 && getSize() != 16)
154 return emitError(
"expected byte size to be either 4, 8 or 16.");
155 if (getModifier() == LoadCacheModifierKind::CG && getSize() != 16)
156 return emitError(
"CG cache modifier is only support for 16 bytes copy.");
163 if (tensorDims < 1 || tensorDims > 5)
164 return emitError(loc,
"expects coordinates between 1 to 5 dimension");
166 auto checkTMALoadParams = [&](TMALoadMode mode,
bool isIm2col,
167 size_t expectedIm2colOff) -> LogicalResult {
168 if (isIm2col && (tensorDims < 3))
171 <<
" mode, the tensor has to be at least 3-dimensional";
173 if (numIm2colOff != expectedIm2colOff)
174 return emitError(loc) <<
" im2col offsets expected " << expectedIm2colOff
175 <<
" (provided " << numIm2colOff <<
")";
181 case TMALoadMode::TILE:
182 return checkTMALoadParams(mode,
false, 0);
183 case TMALoadMode::IM2COL:
184 return checkTMALoadParams(mode,
true, tensorDims - 2);
185 case TMALoadMode::IM2COL_W:
186 case TMALoadMode::IM2COL_W_128:
187 return checkTMALoadParams(mode,
true, 2);
188 case TMALoadMode::TILE_GATHER4:
189 return (tensorDims == 5)
190 ? checkTMALoadParams(mode,
false, 0)
191 :
emitError(loc,
"Gather4 mode expects 5 coordinates");
196LogicalResult CpAsyncBulkTensorPrefetchOp::verify() {
198 getMode(), getLoc());
201LogicalResult CpAsyncBulkTensorGlobalToSharedClusterOp::verify() {
202 TMALoadMode mode = getMode();
203 bool isCTAOnly = getIsCTAOnly();
204 if (getPredicate()) {
206 return emitError(
"Predicate is supported only for shared::cluster mode.");
207 if (mode != TMALoadMode::TILE && mode != TMALoadMode::IM2COL)
209 "Predicate is supported only for Tile and Im2col modes.");
211 NVVMMemorySpace expectedAS =
212 isCTAOnly ? NVVMMemorySpace::Shared : NVVMMemorySpace::SharedCluster;
213 unsigned AS = llvm::cast<LLVM::LLVMPointerType>(getDstMem().
getType())
215 if (AS != expectedAS)
218 ?
"Shared::cta destination requires address-space 3."
219 :
"Shared::cluster destination requires address-space 7.");
222 if (getMulticastMask())
223 return emitError(
"Multicast is not supported with shared::cta mode.");
225 return emitError(
"CTAGroup is not supported with shared::cta mode.");
230 getMode(), getLoc());
233LogicalResult CpAsyncBulkTensorReduceOp::verify() {
234 TMAStoreMode mode = getMode();
237 case TMAStoreMode::TILE:
239 case TMAStoreMode::IM2COL:
240 case TMAStoreMode::IM2COL_W:
242 case TMAStoreMode::TILE_SCATTER4:
243 return emitError(
"Scatter mode unsupported for CpAsyncBulkTensorReduceOp");
248LogicalResult CpAsyncBulkGlobalToSharedClusterOp::verify() {
250 if (isSharedCTA && getMulticastMask())
251 return emitError(
"Multicast is not supported with shared::cta mode.");
257 NVVM::MemScopeKind scope,
258 Value retVal =
nullptr) {
259 if (scope != NVVM::MemScopeKind::CTA && scope != NVVM::MemScopeKind::CLUSTER)
260 return op->
emitError(
"mbarrier scope must be either CTA or Cluster");
263 bool hasRetValue =
static_cast<bool>(retVal);
264 if (isSharedCluster && hasRetValue)
266 "mbarrier in shared_cluster space cannot return any value");
271LogicalResult MBarrierArriveOp::verify() {
276LogicalResult MBarrierArriveDropOp::verify() {
281LogicalResult MBarrierArriveExpectTxOp::verify() {
285 if (getPredicate()) {
286 if (getScope() != NVVM::MemScopeKind::CTA)
287 return emitError(
"mbarrier scope must be CTA when using predicate");
290 return emitError(
"mbarrier in shared_cluster space is not supported when "
294 return emitError(
"return-value is not supported when using predicate");
296 if (getRelaxed() ==
true)
297 return emitError(
"mbarrier with relaxed semantics is not supported when "
304LogicalResult MBarrierArriveDropExpectTxOp::verify() {
319 inferredReturnTypes.push_back(IntegerType::get(context, 64));
324MBarrierArriveOp::inferReturnTypes(
MLIRContext *context,
325 std::optional<Location> location,
326 MBarrierArriveOp::Adaptor adaptor,
329 inferredReturnTypes);
332LogicalResult MBarrierArriveDropOp::inferReturnTypes(
333 MLIRContext *context, std::optional<Location> location,
334 MBarrierArriveDropOp::Adaptor adaptor,
337 inferredReturnTypes);
340LogicalResult MBarrierArriveExpectTxOp::inferReturnTypes(
341 MLIRContext *context, std::optional<Location> location,
342 MBarrierArriveExpectTxOp::Adaptor adaptor,
346 if (adaptor.getPredicate())
349 inferredReturnTypes);
352LogicalResult MBarrierArriveDropExpectTxOp::inferReturnTypes(
353 MLIRContext *context, std::optional<Location> location,
354 MBarrierArriveDropExpectTxOp::Adaptor adaptor,
357 inferredReturnTypes);
367 return inferred == actual;
376bool MBarrierArriveExpectTxOp::isCompatibleReturnTypes(
TypeRange l,
380bool MBarrierArriveDropExpectTxOp::isCompatibleReturnTypes(
TypeRange l,
385LogicalResult MBarrierExpectTxOp::verify() {
389LogicalResult MBarrierCompleteTxOp::verify() {
393LogicalResult MBarrierTestWaitOp::verify() {
397LogicalResult MBarrierTryWaitOp::verify() {
401LogicalResult ConvertFloatToTF32Op::verify() {
402 using RndMode = NVVM::FPRoundingMode;
406 return emitError(
"Relu not supported with rna rounding mode.");
413 "Only {rn,rz,rna} rounding modes supported for ConvertFloatToTF32Op.");
418LogicalResult ConvertF32x2ToF6x2Op::verify() {
421 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy())) {
423 << mlir::Float6E2M3FNType::get(ctx) <<
" and "
424 << mlir::Float6E3M2FNType::get(ctx)
425 <<
" types are supported for conversions from f32x2 to f6x2.";
430LogicalResult ConvertF32x2ToF8x2Op::verify() {
431 using RndMode = NVVM::FPRoundingMode;
432 using SatMode = NVVM::SaturationMode;
434 bool isRoundingModeRN = getRnd() == RndMode::RN;
435 bool isRoundingModeRZ = getRnd() == RndMode::RZ;
436 bool isRoundingModeRP = getRnd() == RndMode::RP;
437 bool isSatFinite = getSat() == SatMode::SATFINITE;
439 bool hasRelu = getRelu();
444 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(
446 if (!isRoundingModeRN) {
447 return emitOpError(
"Only RN rounding mode is supported for "
448 "conversions from f32x2 to ")
449 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
450 << mlir::Float8E5M2Type::get(ctx) <<
" types";
453 return emitOpError(
"Only SATFINITE saturation mode is supported "
456 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
457 << mlir::Float8E5M2Type::get(ctx) <<
" types";
461 .Case<mlir::Float8E8M0FNUType>([&](
mlir::Type) -> LogicalResult {
462 if (!(isRoundingModeRZ || isRoundingModeRP)) {
463 return emitOpError(
"Only RZ and RP rounding modes are supported for "
464 "conversions from f32x2 to ")
465 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
468 return emitOpError(
"relu not supported for conversions to ")
469 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
475 << mlir::Float8E4M3FNType::get(ctx) <<
", "
476 << mlir::Float8E5M2Type::get(ctx) <<
", and "
477 << mlir::Float8E8M0FNUType::get(ctx)
479 "supported for conversions from f32x2 to f8x2";
483LogicalResult ConvertF16x2ToF8x2Op::verify() {
486 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy())) {
488 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
489 << mlir::Float8E5M2Type::get(ctx)
490 <<
" types are supported for conversions from f16x2 to f8x2.";
495LogicalResult ConvertBF16x2ToF8x2Op::verify() {
496 using RndMode = NVVM::FPRoundingMode;
497 using SatMode = NVVM::SaturationMode;
499 bool isRoundingModeRN = getRnd() == RndMode::RN;
500 bool isRoundingModeRZ = getRnd() == RndMode::RZ;
501 bool isRoundingModeRP = getRnd() == RndMode::RP;
502 bool isSatFinite = getSat() == SatMode::SATFINITE;
503 bool hasRelu = getRelu();
508 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(
510 if (!isRoundingModeRN)
511 return emitOpError(
"Only RN rounding mode is supported for "
512 "conversions from bf16x2 to ")
513 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
514 << mlir::Float8E5M2Type::get(ctx) <<
" types";
516 return emitOpError(
"Only SATFINITE saturation mode is supported "
517 "for conversions from bf16x2 to ")
518 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
519 << mlir::Float8E5M2Type::get(ctx) <<
" types";
522 .Case<mlir::Float8E8M0FNUType>([&](
mlir::Type) -> LogicalResult {
523 if (!(isRoundingModeRZ || isRoundingModeRP))
524 return emitOpError(
"Only RZ and RP rounding modes are supported for "
525 "conversions from bf16x2 to ")
526 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
528 return emitOpError(
"relu not supported for conversions to ")
529 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
533 llvm_unreachable(
"Invalid conversion in ConvertBF16x2ToF8x2Op");
538LogicalResult ConvertF32x2ToF4x2Op::verify() {
541 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))
543 << mlir::Float4E2M1FNType::get(ctx)
544 <<
" type is supported for conversions from f32x2 to f4x2.";
549LogicalResult ConvertF8x2ToBF16x2Op::verify() {
551 if (llvm::isa<Float8E8M0FNUType>(getSrcType())) {
552 if (getSat() != SaturationMode::NONE)
554 "Only NONE saturation mode is supported for conversions from ")
555 << Float8E8M0FNUType::get(ctx) <<
" type";
556 if (getScaleFactor())
557 return emitOpError(
"scaleFactor not supported for conversions from ")
558 << Float8E8M0FNUType::get(ctx) <<
" type";
560 return emitOpError(
"relu not supported for conversions from ")
561 << Float8E8M0FNUType::get(ctx) <<
" type";
567LogicalResult PermuteOp::verify() {
568 using Mode = NVVM::PermuteMode;
569 bool hasHi =
static_cast<bool>(getHi());
576 return emitError(
"mode '") << getMode() <<
"' requires 'hi' operand.";
584 << getMode() <<
"' does not accept 'hi' operand.";
599 static constexpr FPRoundingMode validRndModes[] = {
600 FPRoundingMode::RN, FPRoundingMode::RZ, FPRoundingMode::RS};
602 if (!llvm::is_contained(validRndModes, rnd)) {
604 "Only RN, RZ, and RS rounding modes are supported for "
605 "conversions from f32x2 to ")
609 if (rnd == FPRoundingMode::RS) {
610 if (!hasRandomBits) {
611 return op->
emitOpError(
"random_bits is required for RS rounding mode.");
616 "random_bits not supported for RN and RZ rounding modes.");
623LogicalResult ConvertF32x2ToF16x2Op::verify() {
625 getRandomBits() ?
true :
false, *
this);
628LogicalResult ConvertF32x2ToBF16x2Op::verify() {
630 getRandomBits() ?
true :
false, *
this);
633LogicalResult ConvertF32x4ToF8x4Op::verify() {
636 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy()))
638 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
639 << mlir::Float8E5M2Type::get(ctx)
640 <<
" types are supported for conversions from f32x4 to f8x4.";
645LogicalResult ConvertF32x4ToF6x4Op::verify() {
648 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy()))
650 << mlir::Float6E2M3FNType::get(ctx) <<
" and "
651 << mlir::Float6E3M2FNType::get(ctx)
652 <<
" types are supported for conversions from f32x4 to f6x4.";
657LogicalResult ConvertF32x4ToF4x4Op::verify() {
660 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))
661 return emitOpError(
"Only ") << mlir::Float4E2M1FNType::get(ctx)
662 <<
" type is supported for conversions from "
668LogicalResult BulkStoreOp::verify() {
669 if (getInitVal() != 0)
670 return emitOpError(
"only 0 is supported for initVal, got ") << getInitVal();
674LogicalResult PMEventOp::verify() {
675 auto eventId = getEventId();
676 auto maskedEventId = getMaskedEventId();
677 if (!maskedEventId && !eventId) {
678 return emitOpError() <<
"either `id` or `mask` must be set";
681 if (maskedEventId && eventId) {
682 return emitOpError() <<
"`id` and `mask` cannot be set at the same time";
686 if (eventId < 0 || eventId > 15) {
687 return emitOpError() <<
"`id` must be between 0 and 15";
691 return llvm::success();
697std::optional<mlir::NVVM::MMATypes>
698MmaOp::inferOperandMMAType(
Type operandElType,
bool isAccumulator) {
700 VectorType::get(2, Float16Type::get(operandElType.
getContext()));
701 if (operandElType.
isF64())
702 return NVVM::MMATypes::f64;
703 if (operandElType.
isF16() || operandElType == half2Type)
704 return NVVM::MMATypes::f16;
705 if (operandElType.
isF32() && isAccumulator)
706 return NVVM::MMATypes::f32;
707 if (operandElType.
isF32() && !isAccumulator)
708 return NVVM::MMATypes::tf32;
709 if (llvm::isa<IntegerType>(operandElType)) {
711 return NVVM::MMATypes::s32;
715 if (
auto structType = llvm::dyn_cast<LLVM::LLVMStructType>(operandElType)) {
716 if (structType.getBody().empty())
718 return inferOperandMMAType(structType.getBody()[0], isAccumulator);
725 return (type == MMATypes::u4 || type == MMATypes::s4);
729 return (type == MMATypes::u8 || type == MMATypes::s8);
734 type == MMATypes::s32;
737MMATypes MmaOp::accumPtxType() {
738 std::optional<mlir::NVVM::MMATypes> val = inferOperandMMAType(
739 getODSOperands(2).getTypes().front(),
true);
740 assert(val.has_value() &&
"accumulator PTX type should always be inferrable");
744MMATypes MmaOp::resultPtxType() {
745 std::optional<mlir::NVVM::MMATypes> val =
746 inferOperandMMAType(getResult().
getType(),
true);
747 assert(val.has_value() &&
"result PTX type should always be inferrable");
753 struct MMAOperandFragment {
754 StringRef operandName;
755 StringRef ptxTypeAttr;
756 SmallVector<Value, 4> regs;
757 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
758 : operandName(name), ptxTypeAttr(ptxTypeName) {}
761 std::array<MMAOperandFragment, 3> frags{
762 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
763 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
764 MMAOperandFragment(
"C",
"")};
766 mlir::NVVM::MmaOp::getOperandSegmentSizeAttr()};
768 for (
unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {
769 auto &frag = frags[fragIdx];
770 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);
771 for (
auto operandIdx = varOperandSpec.first;
772 operandIdx < varOperandSpec.first + varOperandSpec.second;
774 frag.regs.push_back(this->getOperand(operandIdx));
775 if (operandIdx == 0) {
776 regTypes.push_back(this->getOperand(operandIdx).
getType());
779 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(
780 regTypes.back(), fragIdx >= 2);
782 ignoreAttrNames.push_back(frag.ptxTypeAttr);
785 auto printMmaOperand = [&](
const MMAOperandFragment &frag) ->
void {
786 p <<
" " << frag.operandName;
792 for (
const auto &frag : frags) {
793 printMmaOperand(frag);
801 frags[1].regs[0].getType(),
802 frags[2].regs[0].getType()},
811 std::optional<MMAIntOverflow> intOverflow,
812 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
813 std::optional<std::array<MMALayout, 2>> multiplicandLayouts) {
815 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
820 result.addOperands(operandA);
821 result.addOperands(operandB);
822 result.addOperands(operandC);
824 if (multiplicandPtxTypes) {
825 result.addAttribute(
"multiplicandAPtxType",
826 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
827 result.addAttribute(
"multiplicandBPtxType",
828 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
830 if (
auto res = inferOperandMMAType(operandA[0].
getType(),
false))
831 result.addAttribute(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
832 if (
auto res = inferOperandMMAType(operandB[0].
getType(),
false))
833 result.addAttribute(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
836 if (multiplicandLayouts) {
837 result.addAttribute(
"layoutA",
838 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[0]));
839 result.addAttribute(
"layoutB",
840 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[1]));
842 result.addAttribute(
"layoutA", MMALayoutAttr::get(ctx, MMALayout::row));
843 result.addAttribute(
"layoutB", MMALayoutAttr::get(ctx, MMALayout::col));
846 if (intOverflow.has_value())
847 result.addAttribute(
"intOverflowBehavior",
848 MMAIntOverflowAttr::get(ctx, *intOverflow));
849 if (b1Op.has_value())
850 result.addAttribute(
"b1Op", MMAB1OpAttr::get(ctx, *b1Op));
852 result.addTypes(resultType);
854 MmaOp::getOperandSegmentSizeAttr(),
856 static_cast<int32_t>(operandB.size()),
857 static_cast<int32_t>(operandC.size())}));
865 struct MMAOperandFragment {
866 std::optional<MMATypes> elemtype;
867 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
868 SmallVector<Type> regTypes;
872 std::array<MMAOperandFragment, 4> frags;
878 MMAOperandFragment &frag) -> LogicalResult {
908 if (operandTypes.size() != 3)
911 "expected one type for each operand segment but got " +
912 Twine(operandTypes.size()) +
" types");
913 for (
const auto &iter : llvm::enumerate(operandTypes)) {
914 auto &frag = frags[iter.index()];
915 frag.regTypes.resize(frag.regs.size(), iter.value());
919 frag.elemtype = inferOperandMMAType(frag.regTypes[0],
926 frags[3].elemtype = inferOperandMMAType(resultType,
true);
928 std::array<StringRef, 2> names{
"multiplicandAPtxType",
929 "multiplicandBPtxType"};
930 for (
unsigned idx = 0; idx < names.size(); idx++) {
931 const auto &frag = frags[idx];
932 std::optional<NamedAttribute> attr = namedAttributes.
getNamed(names[idx]);
933 if (!frag.elemtype.has_value() && !attr.has_value()) {
936 "attribute " + names[idx] +
937 " is not provided explicitly and cannot be inferred");
939 if (!attr.has_value())
941 names[idx], MMATypesAttr::get(parser.
getContext(), *frag.elemtype));
944 result.addTypes(resultType);
945 if (!namedAttributes.
empty())
946 result.addAttributes(namedAttributes);
947 result.addAttribute(MmaOp::getOperandSegmentSizeAttr(),
949 static_cast<int32_t>(frags[0].regs.size()),
950 static_cast<int32_t>(frags[1].regs.size()),
951 static_cast<int32_t>(frags[2].regs.size()),
956LogicalResult MmaOp::verify() {
958 auto f16Ty = Float16Type::get(context);
959 auto i32Ty = IntegerType::get(context, 32);
960 auto f16x2Ty = VectorType::get(2, f16Ty);
961 auto f32Ty = Float32Type::get(context);
962 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(
963 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});
966 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});
969 auto f16x2x2StructTy =
970 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});
972 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});
974 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});
976 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),
977 getShapeAttr().getK()};
983 AllowedShapes allowedShapes;
984 AllowedTypes expectedA;
985 AllowedTypes expectedB;
986 AllowedTypes expectedC;
991 if (mmaShape[0] == 16) {
993 Type multiplicandFragType;
994 switch (*getMultiplicandAPtxType()) {
997 multiplicandFragType = i32Ty;
998 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
999 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1001 case MMATypes::bf16:
1003 multiplicandFragType = i32Ty;
1004 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1005 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1009 multiplicandFragType = f16x2Ty;
1010 expectedResult.push_back(f16x2x2StructTy);
1011 expectedResult.push_back(f32x4StructTy);
1013 case MMATypes::e4m3:
1014 case MMATypes::e5m2:
1018 multiplicandFragType = i32Ty;
1019 expectedResult.push_back(f16x2x2StructTy);
1020 expectedResult.push_back(f32x4StructTy);
1034 return emitError(
"invalid shape or multiplicand type: ")
1035 << getMultiplicandAPtxType().value();
1039 expectedResult.push_back(s32x4StructTy);
1040 expectedC.emplace_back(4, i32Ty);
1041 multiplicandFragType = i32Ty;
1043 expectedC.emplace_back(2, f16x2Ty);
1044 expectedC.emplace_back(4, f32Ty);
1047 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor);
1048 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);
1049 expectedA.emplace_back(unitA, multiplicandFragType);
1050 expectedB.emplace_back(unitB, multiplicandFragType);
1051 allowedShapes.push_back({16, 8, kFactor});
1052 allowedShapes.push_back({16, 8, kFactor * 2});
1054 if (resultPtxType() != accumPtxType())
1059 if (mmaShape[0] == 8) {
1060 if (*getMultiplicandAPtxType() == MMATypes::f16) {
1061 expectedA.emplace_back(2, f16x2Ty);
1062 expectedB.emplace_back(2, f16x2Ty);
1063 expectedResult.push_back(f16x2x4StructTy);
1064 expectedResult.push_back(f32x8StructTy);
1065 expectedC.emplace_back(4, f16x2Ty);
1066 expectedC.emplace_back(8, f32Ty);
1067 allowedShapes.push_back({8, 8, 4});
1069 if (*getMultiplicandAPtxType() == MMATypes::f64) {
1070 Type f64Ty = Float64Type::get(context);
1071 expectedA.emplace_back(1, f64Ty);
1072 expectedB.emplace_back(1, f64Ty);
1073 expectedC.emplace_back(2, f64Ty);
1074 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(
1076 allowedShapes.push_back({8, 8, 4});
1079 expectedA.push_back({i32Ty});
1080 expectedB.push_back({i32Ty});
1081 expectedC.push_back({i32Ty, i32Ty});
1082 expectedResult.push_back(s32x2StructTy);
1084 allowedShapes.push_back({8, 8, 32});
1086 allowedShapes.push_back({8, 8, 16});
1087 if (getMultiplicandAPtxType().value() == MMATypes::b1)
1088 allowedShapes.push_back({8, 8, 128});
1092 std::string errorMessage;
1093 llvm::raw_string_ostream errorStream(errorMessage);
1096 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||
1097 !llvm::is_contained(allowedShapes, mmaShape)) {
1098 errorStream <<
"unimplemented variant for MMA shape <";
1099 llvm::interleaveComma(mmaShape, errorStream);
1105 std::array<StringRef, 3> operandNames{
"A",
"B",
"C"};
1106 for (
const auto &iter : llvm::enumerate(
1108 auto spec = this->getODSOperandIndexAndLength(iter.index());
1110 operand_type_begin() + spec.first +
1112 bool match = llvm::is_contained(iter.value(), operandTySeg);
1115 errorStream <<
"Could not match types for the "
1116 << operandNames[iter.index()]
1117 <<
" operands; expected one of ";
1118 for (
const auto &x : iter.value()) {
1119 errorStream << x.size() <<
"x" << x[0] <<
" ";
1121 errorStream <<
"but got ";
1122 llvm::interleaveComma(operandTySeg, errorStream);
1128 if (!llvm::any_of(expectedResult, [&](
Type expectedResultType) {
1129 return expectedResultType == getResult().getType();
1132 <<
"Could not match allowed types for the result; expected one of ";
1133 llvm::interleaveComma(expectedResult, errorStream);
1134 errorStream <<
" but got " << getResult().getType();
1139 if (getMultiplicandAPtxType() == MMATypes::b1 && !getB1Op()) {
1140 return emitOpError(
"op requires " + getB1OpAttrName().strref() +
1148 if (!getIntOverflowBehavior())
1150 getIntOverflowBehaviorAttrName().strref() +
1158 (mmaShape[0] == 8 && mmaShape[1] == 8 && mmaShape[2] == 4 &&
1159 getMultiplicandAPtxType() == MMATypes::f16);
1161 if (!isM8N8K4_F16) {
1163 if (getLayoutA() != MMALayout::row || getLayoutB() != MMALayout::col) {
1164 return emitOpError(
"requires layoutA = #nvvm.mma_layout<row> and "
1165 "layoutB = #nvvm.mma_layout<col> for shape <")
1166 << mmaShape[0] <<
", " << mmaShape[1] <<
", " << mmaShape[2]
1167 <<
"> with element types " << *getMultiplicandAPtxType() <<
" and "
1168 << *getMultiplicandBPtxType()
1169 <<
". Only m8n8k4 with f16 supports other layouts.";
1176MMATypes MmaSpOp::accumPtxType() {
1177 std::optional<mlir::NVVM::MMATypes> val = MmaOp::inferOperandMMAType(
1178 getODSOperands(2).getTypes().front(),
true);
1179 assert(val.has_value() &&
"accumulator PTX type should always be inferrable");
1183MMATypes MmaSpOp::resultPtxType() {
1184 std::optional<mlir::NVVM::MMATypes> val =
1185 MmaOp::inferOperandMMAType(getResult().
getType(),
true);
1186 assert(val.has_value() &&
"result PTX type should always be inferrable");
1192 llvm::IRBuilderBase &builder) {
1193 auto thisOp = cast<NVVM::MmaSpOp>(op);
1201 auto intId = MmaSpOp::getIntrinsicID(
1202 thisOp.getShape().getM(), thisOp.getShape().getN(),
1203 thisOp.getShape().getK(), thisOp.getIntOverflowBehavior(),
1204 thisOp.getOrderedMetadata(), thisOp.getKind(),
1205 *thisOp.getMultiplicandAPtxType(), *thisOp.getMultiplicandBPtxType(),
1206 thisOp.accumPtxType(), thisOp.resultPtxType());
1208 return {intId, args};
1213 struct MMAOperandFragment {
1214 StringRef operandName;
1215 StringRef ptxTypeAttr;
1216 SmallVector<Value, 4> regs;
1217 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1218 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1221 std::array<MMAOperandFragment, 5> frags{
1222 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
1223 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
1224 MMAOperandFragment(
"C",
""), MMAOperandFragment(
"sparseMetadata",
""),
1225 MMAOperandFragment(
"selector",
"")};
1227 mlir::NVVM::MmaSpOp::getOperandSegmentSizeAttr()};
1230 for (
unsigned fragIdx = 0; fragIdx < 3; fragIdx++) {
1231 auto &frag = frags[fragIdx];
1232 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);
1233 for (
auto operandIdx = varOperandSpec.first;
1234 operandIdx < varOperandSpec.first + varOperandSpec.second;
1236 frag.regs.push_back(this->getOperand(operandIdx));
1237 if (operandIdx == varOperandSpec.first) {
1238 regTypes.push_back(this->getOperand(operandIdx).
getType());
1241 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(
1242 regTypes.back(), fragIdx >= 2);
1244 ignoreAttrNames.push_back(frag.ptxTypeAttr);
1248 frags[3].regs.push_back(getSparseMetadata());
1249 frags[4].regs.push_back(getSparsitySelector());
1251 auto printMmaSpOperand = [&](
const MMAOperandFragment &frag) ->
void {
1252 p <<
" " << frag.operandName;
1258 for (
const auto &frag : frags)
1259 printMmaSpOperand(frag);
1264 for (
int i = 0; i < 3; ++i) {
1269 p <<
") -> " << getResult().getType();
1276 std::optional<MMAIntOverflow> intOverflow,
1277 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {
1279 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
1284 result.addOperands(operandA);
1285 result.addOperands(operandB);
1286 result.addOperands(operandC);
1287 result.addOperands(sparseMetadata);
1288 result.addOperands(sparsitySelector);
1290 if (multiplicandPtxTypes) {
1291 result.addAttribute(
"multiplicandAPtxType",
1292 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
1293 result.addAttribute(
"multiplicandBPtxType",
1294 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
1296 if (
auto res = MmaOp::inferOperandMMAType(operandA[0].
getType(),
false))
1297 result.addAttribute(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
1298 if (
auto res = MmaOp::inferOperandMMAType(operandB[0].
getType(),
false))
1299 result.addAttribute(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
1302 if (intOverflow.has_value())
1303 result.addAttribute(
"intOverflowBehavior",
1304 MMAIntOverflowAttr::get(ctx, *intOverflow));
1306 result.addTypes(resultType);
1308 MmaSpOp::getOperandSegmentSizeAttr(),
1310 static_cast<int32_t>(operandB.size()),
1311 static_cast<int32_t>(operandC.size()), 1,
1316 struct MMAOperandFragment {
1317 std::optional<MMATypes> elemtype;
1318 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
1319 SmallVector<Type> regTypes;
1323 std::array<MMAOperandFragment, 6> frags;
1328 auto parseMmaSpOperand = [&](StringRef operandName,
1329 MMAOperandFragment &frag) -> LogicalResult {
1340 if (parseMmaSpOperand(
"A", frags[0]).
failed())
1342 if (parseMmaSpOperand(
"B", frags[1]).
failed())
1344 if (parseMmaSpOperand(
"C", frags[2]).
failed())
1346 if (parseMmaSpOperand(
"sparseMetadata", frags[3]).
failed())
1348 if (parseMmaSpOperand(
"selector", frags[4]).
failed())
1364 if (operandTypes.size() != 3)
1367 "expected one type for each operand segment but got " +
1368 Twine(operandTypes.size()) +
" types");
1369 for (
const auto &iter : llvm::enumerate(operandTypes)) {
1370 auto &frag = frags[iter.index()];
1371 frag.regTypes.resize(frag.regs.size(), iter.value());
1376 MmaOp::inferOperandMMAType(frag.regTypes[0],
1384 MmaOp::inferOperandMMAType(resultType,
true);
1399 std::array<StringRef, 2> names{
"multiplicandAPtxType",
1400 "multiplicandBPtxType"};
1401 for (
unsigned idx = 0; idx < names.size(); idx++) {
1402 const auto &frag = frags[idx];
1403 std::optional<NamedAttribute> attr = namedAttributes.
getNamed(names[idx]);
1404 if (!frag.elemtype.has_value() && !attr.has_value()) {
1407 "attribute " + names[idx] +
1408 " is not provided explicitly and cannot be inferred");
1410 if (!attr.has_value())
1412 names[idx], MMATypesAttr::get(parser.
getContext(), *frag.elemtype));
1415 result.addTypes(resultType);
1416 if (!namedAttributes.
empty())
1417 result.addAttributes(namedAttributes);
1418 result.addAttribute(MmaSpOp::getOperandSegmentSizeAttr(),
1420 static_cast<int32_t>(frags[0].regs.size()),
1421 static_cast<int32_t>(frags[1].regs.size()),
1422 static_cast<int32_t>(frags[2].regs.size()),
1429LogicalResult MmaSpOp::verify() {
1431 auto f16Ty = Float16Type::get(context);
1432 auto i32Ty = IntegerType::get(context, 32);
1433 auto f16x2Ty = VectorType::get(2, f16Ty);
1434 auto f32Ty = Float32Type::get(context);
1435 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(
1436 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});
1438 auto s32x4StructTy =
1439 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});
1440 auto f32x8StructTy =
1442 auto f16x2x2StructTy =
1443 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});
1444 auto f32x4StructTy =
1445 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});
1446 auto s32x2StructTy =
1447 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});
1449 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),
1450 getShapeAttr().getK()};
1456 AllowedShapes allowedShapes;
1457 AllowedTypes expectedA;
1458 AllowedTypes expectedB;
1459 AllowedTypes expectedC;
1464 if (mmaShape[0] == 16) {
1466 Type multiplicandFragType;
1467 switch (*getMultiplicandAPtxType()) {
1468 case MMATypes::tf32:
1470 multiplicandFragType = i32Ty;
1471 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1472 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1474 allowedShapes.push_back({16, 8, 8});
1475 allowedShapes.push_back({16, 8, 16});
1477 case MMATypes::bf16:
1479 multiplicandFragType = i32Ty;
1480 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1481 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1483 allowedShapes.push_back({16, 8, 16});
1484 allowedShapes.push_back({16, 8, 32});
1488 multiplicandFragType = f16x2Ty;
1489 expectedResult.push_back(f16x2x2StructTy);
1490 expectedResult.push_back(f32x4StructTy);
1492 allowedShapes.push_back({16, 8, 16});
1493 allowedShapes.push_back({16, 8, 32});
1499 allowedShapes.push_back({16, 8, 64});
1500 allowedShapes.push_back({16, 8, 128});
1506 allowedShapes.push_back({16, 8, 32});
1507 allowedShapes.push_back({16, 8, 64});
1509 case MMATypes::e4m3:
1510 case MMATypes::e5m2:
1511 case MMATypes::e3m2:
1512 case MMATypes::e2m3:
1513 case MMATypes::e2m1:
1515 multiplicandFragType = i32Ty;
1516 expectedResult.push_back(f16x2x2StructTy);
1517 expectedResult.push_back(f32x4StructTy);
1519 allowedShapes.push_back({16, 8, 64});
1522 return emitError(
"invalid shape or multiplicand type: ")
1523 << getMultiplicandAPtxType().value();
1527 expectedResult.push_back(s32x4StructTy);
1528 expectedC.emplace_back(4, i32Ty);
1529 multiplicandFragType = i32Ty;
1530 }
else if (*getMultiplicandAPtxType() >= MMATypes::e4m3 &&
1531 *getMultiplicandAPtxType() <= MMATypes::e2m1) {
1533 expectedC.emplace_back(2, f16x2Ty);
1534 expectedC.emplace_back(4, f32Ty);
1536 expectedC.emplace_back(2, f16x2Ty);
1537 expectedC.emplace_back(4, f32Ty);
1542 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor) / 2;
1543 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);
1544 expectedA.emplace_back(unitA, multiplicandFragType);
1545 expectedB.emplace_back(unitB, multiplicandFragType);
1547 if (resultPtxType() != accumPtxType())
1552 if (mmaShape[0] == 8) {
1553 if (*getMultiplicandAPtxType() == MMATypes::f16) {
1554 expectedA.emplace_back(2, f16x2Ty);
1555 expectedB.emplace_back(2, f16x2Ty);
1556 expectedResult.push_back(f16x2x4StructTy);
1557 expectedResult.push_back(f32x8StructTy);
1558 expectedC.emplace_back(4, f16x2Ty);
1559 expectedC.emplace_back(8, f32Ty);
1560 allowedShapes.push_back({8, 8, 4});
1562 if (*getMultiplicandAPtxType() == MMATypes::f64) {
1563 Type f64Ty = Float64Type::get(context);
1564 expectedA.emplace_back(1, f64Ty);
1565 expectedB.emplace_back(1, f64Ty);
1566 expectedC.emplace_back(2, f64Ty);
1567 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(
1569 allowedShapes.push_back({8, 8, 4});
1572 expectedA.push_back({i32Ty});
1573 expectedB.push_back({i32Ty});
1574 expectedC.push_back({i32Ty, i32Ty});
1575 expectedResult.push_back(s32x2StructTy);
1577 allowedShapes.push_back({8, 8, 32});
1579 allowedShapes.push_back({8, 8, 16});
1583 std::string errorMessage;
1584 llvm::raw_string_ostream errorStream(errorMessage);
1587 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||
1588 !llvm::is_contained(allowedShapes, mmaShape)) {
1589 errorStream <<
"unimplemented variant for MMA shape <";
1590 llvm::interleaveComma(mmaShape, errorStream);
1596 std::array<StringRef, 3> operandNames{
"A",
"B",
"C"};
1597 for (
const auto &iter : llvm::enumerate(
1599 auto spec = this->getODSOperandIndexAndLength(iter.index());
1601 operand_type_begin() + spec.first +
1603 bool match = llvm::is_contained(iter.value(), operandTySeg);
1606 errorStream <<
"Could not match types for the "
1607 << operandNames[iter.index()]
1608 <<
" operands; expected one of ";
1609 for (
const auto &x : iter.value()) {
1610 errorStream << x.size() <<
"x" << x[0] <<
" ";
1612 errorStream <<
"but got ";
1613 llvm::interleaveComma(operandTySeg, errorStream);
1619 if (!llvm::any_of(expectedResult, [&](
Type expectedResultType) {
1620 return expectedResultType == getResult().getType();
1623 <<
"Could not match allowed types for the result; expected one of ";
1624 llvm::interleaveComma(expectedResult, errorStream);
1625 errorStream <<
" but got " << getResult().getType();
1633 if (!getIntOverflowBehavior())
1635 getIntOverflowBehaviorAttrName().strref() +
1640 if (!getSparseMetadata().
getType().isInteger(32)) {
1641 return emitOpError() <<
"sparse metadata must be i32 type";
1645 if (!getSparsitySelector().
getType().isInteger(32)) {
1646 return emitOpError() <<
"sparsity selector must be i32 type";
1658struct MMAOperandFragment {
1659 StringRef operandName;
1660 StringRef ptxTypeAttr;
1661 SmallVector<Value, 4> regs;
1662 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1663 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1670 p <<
" " << name <<
"[";
1689template <
typename Op>
1694 for (
unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {
1695 auto &frag = frags[fragIdx];
1696 auto varOperandSpec = op.getODSOperandIndexAndLength(fragIdx);
1697 for (
auto operandIdx = varOperandSpec.first;
1698 operandIdx < varOperandSpec.first + varOperandSpec.second;
1700 frag.regs.push_back(op.getOperand(operandIdx));
1701 if (fragIdx == 0 && operandIdx == varOperandSpec.first) {
1702 regTypes.push_back(op.getOperand(operandIdx).getType());
1706 regTypes.push_back(frag.regs[0].getType());
1708 std::optional<MMATypes> inferredType =
1709 MmaOp::inferOperandMMAType(regTypes.back(),
1712 ignoreAttrNames.push_back(frag.ptxTypeAttr);
1723 auto typeParser = [&]() {
1727 operandTypes.push_back(ty);
1733 if (operandTypes.size() != 3)
1735 "expected exactly 3 types");
1744 if (!attrs.
get(
"multiplicandAPtxType")) {
1745 if (
auto inferredType =
1746 MmaOp::inferOperandMMAType(operandTypes[0],
false)) {
1747 attrs.
set(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *inferredType));
1750 if (!attrs.
get(
"multiplicandBPtxType")) {
1751 if (
auto inferredType =
1752 MmaOp::inferOperandMMAType(operandTypes[1],
false)) {
1753 attrs.
set(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *inferredType));
1759template <
typename OpType>
1762 ScaleVecSize scaleVecSize,
1763 BlockScaleFormat blockScaleFormat,
1764 MMABlockScaleKind kind) {
1766 auto &properties =
result.getOrAddProperties<
typename OpType::Properties>();
1767 properties.setShape(
1769 properties.setScaleVecSize(ScaleVecSizeAttr::get(ctx, scaleVecSize));
1770 properties.setBlockScaleFormat(
1771 BlockScaleFormatAttr::get(ctx, blockScaleFormat));
1772 properties.setKind(MMABlockScaleKindAttr::get(ctx, kind));
1779 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {
1780 if (multiplicandPtxTypes) {
1781 result.addAttribute(
"multiplicandAPtxType",
1782 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
1783 result.addAttribute(
"multiplicandBPtxType",
1784 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
1786 if (
auto res = MmaOp::inferOperandMMAType(operandA[0].
getType(),
false))
1787 result.addAttribute(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
1788 if (
auto res = MmaOp::inferOperandMMAType(operandB[0].
getType(),
false))
1789 result.addAttribute(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
1794template <
typename OpTy>
1796 return *MmaOp::inferOperandMMAType(
1797 cast<LLVM::LLVMStructType>(op.getRes().getType()).getBody()[0],
1807 std::array<MMAOperandFragment, 3> frags{
1808 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
1809 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
1810 MMAOperandFragment(
"C",
"")};
1812 mlir::NVVM::MmaBlockScaleOp::getOperandSegmentSizeAttr()};
1817 for (
const auto &frag : frags)
1822 {getScaleAData(), getByteIdA(), getThreadIdA()});
1824 {getScaleBData(), getByteIdB(), getThreadIdB()});
1831 frags[1].regs[0].getType(),
1832 frags[2].regs[0].getType()},
1838ParseResult MmaBlockScaleOp::parse(
OpAsmParser &parser,
1840 struct LocalOperandFragment {
1841 std::optional<MMATypes> elemtype;
1842 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
1846 std::array<LocalOperandFragment, 3> frags;
1875 for (
const auto &[idx, frag] : llvm::enumerate(frags)) {
1876 frag.elemtype = MmaOp::inferOperandMMAType(operandTypes[idx],
1879 .resolveOperands(frag.regs, operandTypes[idx], parser.
getNameLoc(),
1889 .resolveOperands(scaleAOperands, scaleTypes, parser.
getNameLoc(),
1899 result.addAttributes(namedAttributes);
1903 result.addTypes(resultTypes);
1904 result.addAttribute(MmaBlockScaleOp::getOperandSegmentSizeAttr(),
1906 static_cast<int32_t>(frags[0].regs.size()),
1907 static_cast<int32_t>(frags[1].regs.size()),
1908 static_cast<int32_t>(frags[2].regs.size()),
1919void MmaBlockScaleOp::build(
1924 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
1925 ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat,
1926 MMABlockScaleKind kind) {
1927 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
1930 blockScaleFormat, kind);
1932 result.addOperands(operandA);
1933 result.addOperands(operandB);
1934 result.addOperands(operandC);
1936 {scaleAData, byteIdA, threadIdA, scaleBData, byteIdB, threadIdB});
1939 multiplicandPtxTypes);
1941 result.addTypes(resultType);
1942 result.addAttribute(MmaBlockScaleOp::getOperandSegmentSizeAttr(),
1944 static_cast<int32_t>(operandA.size()),
1945 static_cast<int32_t>(operandB.size()),
1946 static_cast<int32_t>(operandC.size()),
1958 auto curOp = cast<NVVM::MmaBlockScaleOp>(op);
1962 for (
Value operand : curOp.getOperandA())
1964 for (
Value operand : curOp.getOperandB())
1966 for (
Value operand : curOp.getOperandC())
1970 args.push_back(mt.
lookupValue(curOp.getScaleAData()));
1971 args.push_back(mt.
lookupValue(curOp.getByteIdA()));
1972 args.push_back(mt.
lookupValue(curOp.getThreadIdA()));
1973 args.push_back(mt.
lookupValue(curOp.getScaleBData()));
1974 args.push_back(mt.
lookupValue(curOp.getByteIdB()));
1975 args.push_back(mt.
lookupValue(curOp.getThreadIdB()));
1977 unsigned intId = MmaBlockScaleOp::getIntrinsicID(
1978 curOp.getShape().getM(), curOp.getShape().getN(), curOp.getShape().getK(),
1979 *curOp.getMultiplicandAPtxType(), *curOp.getMultiplicandBPtxType(),
1981 curOp.getBlockScaleFormat(), curOp.getKind());
1983 return {intId, args};
1986LogicalResult MmaBlockScaleOp::verify() {
1992 if (m == 16 && n == 8 && k == 64) {
1993 if (getMultiplicandAPtxType() != NVVM::MMATypes::e2m1 ||
1994 getMultiplicandBPtxType() != NVVM::MMATypes::e2m1)
1996 "unsupported MMATypes attribute for mma.m16n8k64.(mxf4nvf4|mxf4)");
1997 if (getKind() == NVVM::MMABlockScaleKind::MXF4) {
1998 if (getScaleVecSize() != NVVM::ScaleVecSize::X2)
2000 "unsupported ScaleVecSize attribute for mma.m16n8k64.mxf4");
2001 if (getBlockScaleFormat() != NVVM::BlockScaleFormat::UE8M0)
2003 "unsupported BlockScaleFormat attribute for mma.m16n8k64.mxf4");
2004 }
else if (getKind() == NVVM::MMABlockScaleKind::MXF4NVF4) {
2005 if (!((getScaleVecSize() == NVVM::ScaleVecSize::X2 &&
2006 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0) ||
2007 (getScaleVecSize() == NVVM::ScaleVecSize::X4 &&
2008 (getBlockScaleFormat() == NVVM::BlockScaleFormat::UE4M3 ||
2009 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))))
2011 "attributes for mma.m16n8k64.mxf4nvf4");
2015 }
else if (m == 16 && n == 8 && k == 32) {
2016 if (!(getKind() == NVVM::MMABlockScaleKind::MXF8F6F4 &&
2017 getScaleVecSize() == NVVM::ScaleVecSize::X1 &&
2018 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))
2020 emitOpError(
"unsupported Kind, ScaleVecSize and BlockScaleFormat "
2021 "attributes for mma.m16n8k32");
2034 std::array<MMAOperandFragment, 3> frags{
2035 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
2036 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
2037 MMAOperandFragment(
"C",
"")};
2039 mlir::NVVM::MmaSpBlockScaleOp::getOperandSegmentSizeAttr()};
2044 for (
const auto &frag : frags)
2053 {getScaleAData(), getByteIdA(), getThreadIdA()});
2055 {getScaleBData(), getByteIdB(), getThreadIdB()});
2062 frags[1].regs[0].getType(),
2063 frags[2].regs[0].getType()},
2069ParseResult MmaSpBlockScaleOp::parse(
OpAsmParser &parser,
2071 struct LocalOperandFragment {
2072 std::optional<MMATypes> elemtype;
2073 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
2077 std::array<LocalOperandFragment, 3> frags;
2113 for (
const auto &[idx, frag] : llvm::enumerate(frags)) {
2114 frag.elemtype = MmaOp::inferOperandMMAType(operandTypes[idx],
2117 .resolveOperands(frag.regs, operandTypes[idx], parser.
getNameLoc(),
2126 .resolveOperands(metadataOperands, i32Type, parser.
getNameLoc(),
2139 .resolveOperands(scaleAOperands, scaleTypes, parser.
getNameLoc(),
2149 result.addAttributes(namedAttributes);
2154 if (!
result.attributes.get(
"orderedMetadata"))
2157 result.addTypes(resultTypes);
2158 result.addAttribute(MmaSpBlockScaleOp::getOperandSegmentSizeAttr(),
2160 static_cast<int32_t>(frags[0].regs.size()),
2161 static_cast<int32_t>(frags[1].regs.size()),
2162 static_cast<int32_t>(frags[2].regs.size()),
2175void MmaSpBlockScaleOp::build(
2181 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
2182 ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat,
2183 MMABlockScaleKind kind) {
2184 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
2187 builder,
result,
shape, scaleVecSize, blockScaleFormat, kind);
2190 result.addOperands(operandA);
2191 result.addOperands(operandB);
2192 result.addOperands(operandC);
2193 result.addOperands({sparseMetadata, sparsitySelector, scaleAData, byteIdA,
2194 threadIdA, scaleBData, byteIdB, threadIdB});
2197 multiplicandPtxTypes);
2199 result.addTypes(resultType);
2200 result.addAttribute(MmaSpBlockScaleOp::getOperandSegmentSizeAttr(),
2202 static_cast<int32_t>(operandA.size()),
2203 static_cast<int32_t>(operandB.size()),
2204 static_cast<int32_t>(operandC.size()),
2218 auto curOp = cast<NVVM::MmaSpBlockScaleOp>(op);
2222 for (
Value operand : curOp.getOperandA())
2224 for (
Value operand : curOp.getOperandB())
2226 for (
Value operand : curOp.getOperandC())
2230 args.push_back(mt.
lookupValue(curOp.getSparseMetadata()));
2231 args.push_back(mt.
lookupValue(curOp.getSparsitySelector()));
2234 args.push_back(mt.
lookupValue(curOp.getScaleAData()));
2235 args.push_back(mt.
lookupValue(curOp.getByteIdA()));
2236 args.push_back(mt.
lookupValue(curOp.getThreadIdA()));
2237 args.push_back(mt.
lookupValue(curOp.getScaleBData()));
2238 args.push_back(mt.
lookupValue(curOp.getByteIdB()));
2239 args.push_back(mt.
lookupValue(curOp.getThreadIdB()));
2241 unsigned intId = MmaSpBlockScaleOp::getIntrinsicID(
2242 curOp.getShape().getM(), curOp.getShape().getN(), curOp.getShape().getK(),
2243 *curOp.getMultiplicandAPtxType(), *curOp.getMultiplicandBPtxType(),
2245 curOp.getBlockScaleFormat(), curOp.getKind());
2247 return {intId, args};
2250LogicalResult MmaSpBlockScaleOp::verify() {
2252 if (!getOrderedMetadata()) {
2253 return emitOpError(
"'orderedMetadata' attribute is mandatory");
2261 if (m == 16 && n == 8 && k == 128) {
2262 if (getMultiplicandAPtxType() != NVVM::MMATypes::e2m1 ||
2263 getMultiplicandBPtxType() != NVVM::MMATypes::e2m1)
2265 "unsupported MMATypes attribute for mma.m16n8k128.(mxf4nvf4|mxf4)");
2266 if (getKind() == NVVM::MMABlockScaleKind::MXF4) {
2267 if (getScaleVecSize() != NVVM::ScaleVecSize::X2)
2269 "unsupported ScaleVecSize attribute for mma.m16n8k128.mxf4");
2270 if (getBlockScaleFormat() != NVVM::BlockScaleFormat::UE8M0)
2272 "unsupported BlockScaleFormat attribute for mma.m16n8k128.mxf4");
2273 }
else if (getKind() == NVVM::MMABlockScaleKind::MXF4NVF4) {
2274 if (!((getScaleVecSize() == NVVM::ScaleVecSize::X2 &&
2275 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0) ||
2276 (getScaleVecSize() == NVVM::ScaleVecSize::X4 &&
2277 (getBlockScaleFormat() == NVVM::BlockScaleFormat::UE4M3 ||
2278 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))))
2280 "attributes for mma.m16n8k128.mxf4nvf4");
2284 }
else if (m == 16 && n == 8 && k == 64) {
2285 if (!(getKind() == NVVM::MMABlockScaleKind::MXF8F6F4 &&
2286 getScaleVecSize() == NVVM::ScaleVecSize::X1 &&
2287 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))
2289 emitOpError(
"unsupported Kind, ScaleVecSize and BlockScaleFormat "
2290 "attributes for mma.m16n8k64");
2297LogicalResult ShflOp::verify() {
2298 auto returnStructType = llvm::dyn_cast<LLVM::LLVMStructType>(
getType());
2300 auto verifyTypeError = [&](Twine desc,
Type expectedType,
2301 Type actualType) -> LogicalResult {
2302 return emitOpError(
"expected " + desc +
" to be of type ")
2303 << expectedType <<
" but got " << actualType <<
" instead";
2306 if (returnStructType) {
2307 if (!getReturnValueAndIsValid())
2308 return emitOpError(
"\"return_value_and_is_valid\" attribute must be "
2309 "specified when the return type is a struct type");
2311 if (returnStructType.getBody().size() != 2)
2312 return emitOpError(
"expected return type to be a two-element struct");
2315 auto resultType = returnStruct[0];
2316 if (resultType != getVal().
getType())
2317 return verifyTypeError(
"first element in the returned struct",
2318 getVal().
getType(), resultType);
2320 auto predicateType = returnStruct[1];
2321 if (!predicateType.isInteger(1))
2322 return verifyTypeError(
"second element in the returned struct",
2326 if (getReturnValueAndIsValid())
2327 return emitOpError(
"expected return type to be a two-element struct");
2330 return verifyTypeError(
"return type", getVal().
getType(),
getType());
2336ShflOp::inferReturnTypes(
MLIRContext *context, std::optional<Location> location,
2337 ShflOp::Adaptor adaptor,
2339 Type valType = adaptor.getVal().getType();
2340 if (adaptor.getReturnValueAndIsValid())
2341 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2342 context, {valType, IntegerType::get(context, 1)}));
2344 inferredReturnTypes.push_back(valType);
2349 NVVM::MMAFrag frag,
int nRow,
2352 unsigned numberElements = 0;
2355 Type f16x2 = VectorType::get(2, builder.getF16Type());
2356 if (type == NVVM::MMATypes::f16) {
2357 elementType = f16x2;
2358 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)
2362 }
else if (type == NVVM::MMATypes::f32) {
2363 elementType = builder.getF32Type();
2365 }
else if (type == NVVM::MMATypes::f64) {
2366 elementType = builder.getF64Type();
2367 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)
2371 }
else if (type == NVVM::MMATypes::tf32) {
2372 elementType = builder.getI32Type();
2374 }
else if (type == NVVM::MMATypes::s8 || type == NVVM::MMATypes::u8) {
2375 elementType = builder.getI32Type();
2376 int parallelSize = 0;
2377 if (frag == NVVM::MMAFrag::a)
2378 parallelSize = nRow;
2379 if (frag == NVVM::MMAFrag::b)
2380 parallelSize = nCol;
2383 if (parallelSize == 16)
2386 else if (parallelSize == 8)
2388 else if (parallelSize == 32)
2390 }
else if (type == NVVM::MMATypes::s32) {
2391 elementType = builder.getI32Type();
2394 assert(numberElements != 0 && elementType !=
nullptr);
2395 return std::make_pair(elementType, numberElements);
2398static std::pair<mlir::Type, unsigned>
2402 if (frag == NVVM::MMAFrag::a) {
2405 }
else if (frag == NVVM::MMAFrag::b) {
2412 assert(nRow && nCol);
2416LogicalResult NVVM::WMMALoadOp::verify() {
2417 unsigned addressSpace =
2418 llvm::cast<LLVM::LLVMPointerType>(getPtr().
getType()).getAddressSpace();
2419 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&
2420 addressSpace != NVVMMemorySpace::Shared)
2421 return emitOpError(
"expected source pointer in memory "
2424 if (NVVM::WMMALoadOp::getIntrinsicID(
getM(),
getN(), getK(), getLayout(),
2425 getEltype(), getFrag()) == 0)
2426 return emitOpError() <<
"invalid attribute combination";
2431 if (typeInfo.first == f64Ty && typeInfo.second == 1) {
2433 return emitOpError(
"expected destination type to be f64");
2437 Type dstType = LLVM::LLVMStructType::getLiteral(
2440 return emitOpError(
"expected destination type is a structure of ")
2441 << typeInfo.second <<
" elements of type " << typeInfo.first;
2445LogicalResult NVVM::WMMAStoreOp::verify() {
2446 unsigned addressSpace =
2447 llvm::cast<LLVM::LLVMPointerType>(getPtr().
getType()).getAddressSpace();
2448 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&
2449 addressSpace != NVVMMemorySpace::Shared)
2450 return emitOpError(
"expected operands to be a source pointer in memory "
2453 if (NVVM::WMMAStoreOp::getIntrinsicID(
getM(),
getN(), getK(), getLayout(),
2455 return emitOpError() <<
"invalid attribute combination";
2458 if (getArgs().size() != typeInfo.second)
2459 return emitOpError() <<
"expected " << typeInfo.second <<
" data operands";
2460 if (llvm::any_of(getArgs(), [&typeInfo](
Value operands) {
2461 return operands.
getType() != typeInfo.first;
2463 return emitOpError() <<
"expected data operands of type " << typeInfo.first;
2467LogicalResult NVVM::WMMAMmaOp::verify() {
2468 if (NVVM::WMMAMmaOp::getIntrinsicID(
getM(),
getN(), getK(), getLayoutA(),
2469 getLayoutB(), getEltypeA(),
2471 return emitOpError() <<
"invalid attribute combination";
2479 arguments.append(typeInfoA.second, typeInfoA.first);
2480 arguments.append(typeInfoB.second, typeInfoB.first);
2481 arguments.append(typeInfoC.second, typeInfoC.first);
2482 unsigned numArgs = arguments.size();
2483 if (getArgs().size() != numArgs)
2484 return emitOpError() <<
"expected " << numArgs <<
" arguments";
2485 for (
unsigned i = 0; i < numArgs; i++) {
2486 if (getArgs()[i].
getType() != arguments[i])
2487 return emitOpError() <<
"expected argument " << i <<
" to be of type "
2490 Type dstType = LLVM::LLVMStructType::getLiteral(
2493 return emitOpError(
"expected destination type is a structure of ")
2494 << typeInfoC.second <<
" elements of type " << typeInfoC.first;
2498LogicalResult NVVM::LdMatrixOp::verify() {
2500 if (m == 8 && n == 8) {
2501 if (num != 1 && num != 2 && num != 4) {
2502 return emitOpError(
"expected num attribute to be 1, 2 or 4 for 8x8 "
2505 if (getEltType() != LdStMatrixEltType::B16) {
2506 return emitOpError(
"expected element type to be b16 for 8x8 matrix");
2508 }
else if (m == 8 && n == 16) {
2509 if (num != 1 && num != 2 && num != 4) {
2510 return emitOpError(
"expected num attribute to be 1, 2 or 4 for 8x16 "
2513 if (getLayout() != MMALayout::row) {
2514 return emitOpError(
"expected layout to be row for 8x16 matrix");
2516 if (getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&
2517 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {
2518 return emitOpError(
"expected element type to be b8x16.b4x16_p64 or "
2519 "b8x16.b6x16_p32 for 8x16 matrix");
2521 }
else if (m == 16 && n == 16) {
2522 if (num != 1 && num != 2) {
2523 return emitOpError(
"expected num attribute to be 1 or 2 for 16x16 "
2526 if (getLayout() != MMALayout::col) {
2527 return emitOpError(
"expected layout to be col for 16x16 matrix");
2529 if (getEltType() != LdStMatrixEltType::B8 &&
2530 getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&
2531 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {
2532 return emitOpError(
"expected element type to be b8, b8x16.b4x16_p64 or "
2533 "b8x16.b6x16_p32 for 16x16 matrix");
2536 return emitOpError(
"expected shape to be 8x8, 8x16 or 16x16");
2540 uint32_t numElements = (m == 16 && n == 16 ? num * 2 : num);
2541 if (numElements == 1 &&
getType() != i32)
2542 return emitOpError(
"expected destination type is i32");
2543 if (numElements == 2 || numElements == 4) {
2544 Type dstType = LLVM::LLVMStructType::getLiteral(
2547 return emitOpError(
"expected destination type is a structure of ")
2548 << numElements <<
" elements of type i32";
2554LogicalResult LdMatrixOp::inferReturnTypes(
2555 MLIRContext *context, std::optional<Location> location,
2557 uint32_t num = adaptor.getNum();
2558 uint32_t m = adaptor.getShape().getM();
2559 uint32_t n = adaptor.getShape().getN();
2560 uint32_t numElements = (m == 16 && n == 16) ? num * 2 : num;
2562 Type i32 = IntegerType::get(context, 32);
2563 if (numElements == 1)
2564 inferredReturnTypes.push_back(i32);
2566 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2571LogicalResult NVVM::StMatrixOp::verify() {
2572 int numMatrix = getSources().size();
2573 if (numMatrix != 1 && numMatrix != 2 && numMatrix != 4)
2574 return emitOpError(
"expected num attribute to be 1, 2 or 4");
2577 if (m == 8 && n == 8) {
2578 if (getEltType() != NVVM::LdStMatrixEltType::B16) {
2579 return emitOpError(
"expected element type to be B16 for 8x8 matrix");
2581 }
else if (m == 16 && n == 8) {
2582 if (getEltType() != NVVM::LdStMatrixEltType::B8) {
2583 return emitOpError(
"expected element type to be B8 for 16x8 matrix");
2585 if (getLayout() != NVVM::MMALayout::col) {
2586 return emitOpError(
"expected layout to be col for 16x8 matrix");
2589 return emitOpError(
"expected shape to be 8x8 or 16x8");
2595LogicalResult NVVM::MovMatrixOp::verify() {
2597 if (m != 8 || n != 8)
2599 if (getLayout() != NVVM::MMALayout::col)
2601 if (getEltType() != NVVM::LdStMatrixEltType::B16)
2602 return emitOpError(
"expected element type to be b16");
2607 if (typeA == NVVM::WGMMATypes::tf32)
2609 if (typeA == NVVM::WGMMATypes::f16 || typeA == NVVM::WGMMATypes::bf16)
2611 if (typeA == NVVM::WGMMATypes::s8 || typeA == NVVM::WGMMATypes::u8)
2613 if (typeA == NVVM::WGMMATypes::e4m3 || typeA == NVVM::WGMMATypes::e5m2)
2615 if (typeA == NVVM::WGMMATypes::b1)
2621 NVVM::WGMMATypes typeA,
2622 NVVM::WGMMATypes typeB) {
2624 case NVVM::WGMMATypes::f16:
2625 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
2626 typeB == NVVM::WGMMATypes::f16)
2629 case NVVM::WGMMATypes::tf32:
2630 if (typeD == NVVM::WGMMATypes::f32 && typeB == NVVM::WGMMATypes::tf32)
2633 case NVVM::WGMMATypes::u8:
2634 case NVVM::WGMMATypes::s8:
2635 if (typeD == NVVM::WGMMATypes::s32 &&
2636 (typeB == NVVM::WGMMATypes::u8 || typeB == NVVM::WGMMATypes::s8))
2639 case NVVM::WGMMATypes::b1:
2640 if (typeD == NVVM::WGMMATypes::s32 && typeB == NVVM::WGMMATypes::b1)
2643 case NVVM::WGMMATypes::bf16:
2644 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
2645 typeB == NVVM::WGMMATypes::bf16)
2648 case NVVM::WGMMATypes::e4m3:
2649 case NVVM::WGMMATypes::e5m2:
2650 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
2651 (typeB == NVVM::WGMMATypes::e5m2 || typeB == NVVM::WGMMATypes::e4m3))
2654 case WGMMATypes::f32:
2655 case WGMMATypes::s32:
2656 llvm_unreachable(
"unsupported input types");
2664 72, 80, 88, 96, 104, 112, 120, 128,
2665 136, 144, 152, 160, 168, 176, 184, 192,
2666 200, 208, 216, 224, 232, 240, 248, 256};
2668 80, 96, 112, 128, 144, 160,
2669 176, 192, 208, 224, 240, 256};
2671 case WGMMATypes::f16:
2672 case WGMMATypes::tf32:
2673 case WGMMATypes::bf16:
2674 case WGMMATypes::e4m3:
2675 case WGMMATypes::e5m2:
2676 if (llvm::is_contained(allowedN, sizeN))
2679 case WGMMATypes::u8:
2680 case WGMMATypes::s8:
2681 case WGMMATypes::b1:
2682 if (llvm::is_contained(allowedNshort, sizeN))
2685 case WGMMATypes::f32:
2686 case WGMMATypes::s32:
2687 llvm_unreachable(
"unsupported input types");
2693LogicalResult NVVM::WgmmaMmaAsyncOp::verify() {
2694 Value outValue = getResults();
2695 auto stype = dyn_cast<LLVM::LLVMStructType>(outValue.
getType());
2697 return emitOpError() <<
"expected results to be struct";
2698 int outputSize = stype.getBody().size();
2699 WGMMATypes typeD = getTypeD();
2700 WGMMATypes typeA = getTypeA();
2701 WGMMATypes typeB = getTypeB();
2703 for (
Type t : stype.getBody()) {
2704 if (t != stype.getBody().front())
2706 <<
"all elements in struct must be same type but there is " << t;
2709 if (typeD != WGMMATypes::f32 && typeD != WGMMATypes::f16 &&
2710 typeD != WGMMATypes::s32) {
2711 return emitOpError() <<
"does not support the given output type " << typeD;
2713 if (typeD == WGMMATypes::s32 &&
2714 (getScaleA() == WGMMAScaleIn::neg || getScaleB() == WGMMAScaleIn::neg)) {
2715 return emitOpError() <<
"has s32 output, scaleA and scaleB cannot be neg";
2719 return emitOpError() << typeD <<
" += " << typeA <<
" * " << typeB
2720 <<
", it is not supported.";
2730 return emitOpError() <<
"shape 'k' must be " << allowedK.value()
2731 <<
" for input type " << typeA;
2735 return emitOpError() <<
"has input type " << typeA <<
" n is set to "
2736 <<
getShape().getN() <<
", it is not supported.";
2743 if ((typeA != WGMMATypes::f16 && typeA != WGMMATypes::bf16) &&
2744 (getLayoutA() == mlir::NVVM::MMALayout::col ||
2745 getLayoutB() == mlir::NVVM::MMALayout::row)) {
2747 <<
"given layouts layout_a = " << getLayoutA()
2748 <<
" and layout_b = " << getLayoutB() <<
" for input types " << typeA
2750 <<
" requires transpose. However, this is only supported for: "
2751 << MMATypes::f16 <<
" and " << MMATypes::bf16;
2755 int expectedOutput = 0;
2756 if (typeD == WGMMATypes::f32 || typeD == WGMMATypes::s32)
2757 expectedOutput =
getShape().getN() / 2;
2758 if (typeD == WGMMATypes::f16)
2759 expectedOutput =
getShape().getN() / 4;
2760 if (outputSize != expectedOutput) {
2761 return emitOpError() <<
"results " << expectedOutput
2762 <<
", however output struct has " << outputSize
2766 if (typeD != WGMMATypes::s32 &&
2767 getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==
2768 NVVM::MMAIntOverflow::satfinite) {
2770 <<
" `satfinite` can be only used with s32 accumulator, however "
2771 "the current accumulator is "
2778std::string NVVM::WgmmaMmaAsyncOp::getPtx() {
2781 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;
2783 StringRef outputTypeName = stringifyWGMMATypes(getTypeD());
2785 int expectedOutputRegisters = 0;
2786 if (getTypeD() == WGMMATypes::f16)
2787 expectedOutputRegisters =
getShape().getN() / 4;
2789 expectedOutputRegisters =
getShape().getN() / 2;
2792 llvm::raw_string_ostream ss(ptx);
2797 << ((expectedOutputRegisters * 2) + 2)
2799 "wgmma.mma_async.sync.aligned.m"
2800 << m <<
"n" << n <<
"k" << k <<
"." << outputTypeName <<
"." << getTypeA()
2801 <<
"." << getTypeB();
2802 if (getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==
2803 NVVM::MMAIntOverflow::satfinite)
2807 for (; regCnt < expectedOutputRegisters; ++regCnt) {
2808 ss <<
"$" << regCnt;
2809 if (regCnt != expectedOutputRegisters - 1)
2815 regCnt = (regCnt * 2);
2816 ss <<
" $" << (regCnt) <<
"," <<
" $" << (regCnt + 1) <<
"," <<
" p";
2817 if (getTypeD() != WGMMATypes::s32) {
2818 ss <<
", $" << (regCnt + 3) <<
", $" << (regCnt + 4);
2822 ss <<
", $" << (regCnt + 5) <<
", $" << (regCnt + 6);
2829bool NVVM::WgmmaMmaAsyncOp::getAsmValues(
2833 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;
2840 asmValues.push_back({makeConstantI32(rewriter,
static_cast<int>(getScaleD())),
2842 if (getTypeD() != WGMMATypes::s32) {
2843 asmValues.push_back(
2844 {makeConstantI32(rewriter,
2845 getScaleA() == NVVM::WGMMAScaleIn::neg ? -1 : 1),
2847 asmValues.push_back(
2848 {makeConstantI32(rewriter,
2849 getScaleB() == NVVM::WGMMAScaleIn::neg ? -1 : 1),
2853 asmValues.push_back(
2854 {makeConstantI32(rewriter,
static_cast<int>(getLayoutA())),
2856 asmValues.push_back(
2857 {makeConstantI32(rewriter, 1 -
static_cast<int>(getLayoutB())),
2863LogicalResult NVVM::FenceProxyOp::verify() {
2864 if (getKind() == NVVM::ProxyKind::async_shared && !getSpace().has_value()) {
2865 return emitOpError() <<
"async_shared fence requires space attribute";
2867 if (getKind() != NVVM::ProxyKind::async_shared && getSpace().has_value()) {
2868 return emitOpError() <<
"only async_shared fence can have space attribute";
2873LogicalResult NVVM::FenceProxyAcquireOp::verify() {
2874 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
2875 return emitOpError(
"uni-directional proxies only support generic for "
2876 "from_proxy attribute");
2878 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)
2879 return emitOpError(
"uni-directional proxies only support tensormap "
2880 "for to_proxy attribute");
2884LogicalResult NVVM::FenceProxyReleaseOp::verify() {
2885 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
2886 return emitOpError(
"uni-directional proxies only support generic for "
2887 "from_proxy attribute");
2889 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)
2890 return emitOpError(
"uni-directional proxies only support tensormap "
2891 "for to_proxy attribute");
2895LogicalResult NVVM::FenceProxySyncRestrictOp::verify() {
2896 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
2897 return emitOpError(
"only generic is support for from_proxy attribute");
2899 if (getToProxy() != NVVM::ProxyKind::async)
2900 return emitOpError(
"only async is supported for to_proxy attribute");
2904LogicalResult NVVM::SetMaxRegisterOp::verify() {
2905 if (getRegCount() % 8)
2906 return emitOpError(
"new register size must be multiple of 8");
2907 if (getRegCount() < 24 || getRegCount() > 256)
2908 return emitOpError(
"new register size must be in between 24 to 256");
2912LogicalResult NVVM::Tcgen05CpOp::verify() {
2913 auto mc = getMulticast();
2915 using SH = Tcgen05CpShape;
2916 using MC = Tcgen05CpMulticast;
2918 case SH::SHAPE_128x256b:
2919 case SH::SHAPE_128x128b:
2920 case SH::SHAPE_4x256b:
2922 return emitError(
"Invalid multicast type for tcgen05.cp Op");
2924 case SH::SHAPE_64x128b:
2925 if (mc != MC::WARPX2_01_23 && mc != MC::WARPX2_02_13)
2926 return emitError(
"Shape 64x128b requires multicast warpx2_01_23 or "
2927 "warpx2_02_13 for tcgen05.cp Op");
2929 case SH::SHAPE_32x128b:
2930 if (mc != MC::WARPX4)
2932 "Shape 32x128b requires multicast warpx4 for tcgen05.cp Op");
2938LogicalResult NVVM::MatchSyncOp::verify() {
2939 if (getKind() == NVVM::MatchSyncKind::all) {
2940 auto type = llvm::dyn_cast<LLVM::LLVMStructType>(
getType());
2941 if (!type || type.getBody().size() != 2 ||
2942 !type.getBody()[0].isInteger(32) || !type.getBody()[1].isInteger(1)) {
2943 return emitOpError(
"match.sync 'all' returns a two element struct with "
2944 "first element as i32 and second element as i1");
2947 if (!
getType().isInteger(32)) {
2948 return emitOpError(
"match.sync 'any' returns an i32");
2954LogicalResult MatchSyncOp::inferReturnTypes(
2955 MLIRContext *context, std::optional<Location> location,
2957 if (adaptor.getKind() == NVVM::MatchSyncKind::all)
2958 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2960 {IntegerType::get(context, 32), IntegerType::get(context, 1)}));
2962 inferredReturnTypes.push_back(IntegerType::get(context, 32));
2966LogicalResult NVVM::VoteSyncOp::verify() {
2967 if (getKind() == NVVM::VoteSyncKind::ballot) {
2968 if (!
getType().isInteger(32)) {
2969 return emitOpError(
"vote.sync 'ballot' returns an i32");
2972 if (!
getType().isInteger(1)) {
2973 return emitOpError(
"vote.sync 'any', 'all' and 'uni' returns an i1");
2979LogicalResult VoteSyncOp::inferReturnTypes(
2980 MLIRContext *context, std::optional<Location> location,
2982 unsigned width = adaptor.getKind() == NVVM::VoteSyncKind::ballot ? 32 : 1;
2983 inferredReturnTypes.push_back(IntegerType::get(context, width));
2987LogicalResult NVVM::PrefetchOp::verify() {
2988 using MemSpace = NVVM::NVVMMemorySpace;
2989 using CacheLevel = NVVM::PrefetchCacheLevel;
2991 unsigned addressSpace =
2992 llvm::cast<LLVM::LLVMPointerType>(getAddr().
getType()).getAddressSpace();
2993 std::optional<NVVM::CacheEvictionPriority> evictPriority = getEvictPriority();
2994 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = getCacheLevel();
2996 if (getTensormap() && cacheLevel)
2997 return emitOpError(
"cannot specify both tensormap and cache level");
2999 if (getTensormap()) {
3000 if (addressSpace != MemSpace::Generic &&
3001 addressSpace != MemSpace::Constant) {
3003 "prefetch tensormap requires a generic or constant pointer");
3006 if (evictPriority) {
3008 "prefetch tensormap does not support eviction priority");
3011 if (getInParamSpace() && addressSpace != MemSpace::Generic) {
3013 "in_param_space can only be specified for a generic pointer");
3016 }
else if (cacheLevel) {
3017 if (addressSpace != MemSpace::Generic && addressSpace != MemSpace::Global &&
3018 addressSpace != MemSpace::Local) {
3019 return emitOpError(
"prefetch to cache level requires a generic, global, "
3020 "or local pointer");
3024 if (*cacheLevel != CacheLevel::L1) {
3026 "unsupported cache level, the only supported uniform "
3027 "cache level is L1");
3030 if (addressSpace != MemSpace::Generic) {
3032 "prefetch to uniform cache requires a generic pointer");
3036 if (evictPriority) {
3037 if (*cacheLevel != CacheLevel::L2)
3039 "cache eviction priority supported only for cache level L2");
3041 if (addressSpace != MemSpace::Global)
3042 return emitOpError(
"cache eviction priority requires a global pointer");
3044 if (*evictPriority != NVVM::CacheEvictionPriority::EvictNormal &&
3045 *evictPriority != NVVM::CacheEvictionPriority::EvictLast)
3047 "unsupported cache eviction priority, only evict_last and "
3048 "evict_normal are supported");
3052 return emitOpError(
"predicate supported only on prefetch tensormap");
3056 "requires specification of either cache level or tensormap");
3062LogicalResult NVVM::ClusterLaunchControlQueryCancelOp::verify() {
3063 switch (getQueryType()) {
3064 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:
3066 return emitOpError(
"is_canceled query type returns an i1");
3068 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:
3069 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:
3070 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:
3071 if (!
getType().isInteger(32)) {
3072 return emitOpError(
"get_first_cta_id_x, get_first_cta_id_y, "
3073 "get_first_cta_id_z query types return an i32");
3080LogicalResult ClusterLaunchControlQueryCancelOp::inferReturnTypes(
3081 MLIRContext *context, std::optional<Location> location,
3082 ClusterLaunchControlQueryCancelOp::Adaptor adaptor,
3085 adaptor.getQueryType() == NVVM::ClusterLaunchControlQueryType::IS_CANCELED
3088 inferredReturnTypes.push_back(IntegerType::get(context, width));
3092LogicalResult NVVM::ReduxOp::verify() {
3095 if (!reduxType.
isF32()) {
3097 return emitOpError(
"abs attribute is supported only for f32 type");
3099 return emitOpError(
"nan attribute is supported only for f32 type");
3102 NVVM::ReductionKind kind = getKind();
3104 case NVVM::ReductionKind::ADD:
3105 case NVVM::ReductionKind::AND:
3106 case NVVM::ReductionKind::OR:
3107 case NVVM::ReductionKind::XOR:
3108 case NVVM::ReductionKind::MAX:
3109 case NVVM::ReductionKind::MIN:
3110 case NVVM::ReductionKind::UMAX:
3111 case NVVM::ReductionKind::UMIN:
3114 << kind <<
"' reduction kind unsupported with " << reduxType
3115 <<
" type. Only supported type is 'i32'.";
3117 case NVVM::ReductionKind::FMIN:
3118 case NVVM::ReductionKind::FMAX:
3119 if (!reduxType.isF32())
3121 << kind <<
"' reduction kind unsupported with " << reduxType
3122 <<
" type. Only supported type is 'f32'.";
3129LogicalResult NVVM::TensormapReplaceOp::verify() {
3130 auto ord = getOrd();
3131 Value newVal = getNewValue();
3132 auto newValAttr = getNewValueAttr();
3133 auto fieldName = stringifyEnum(getField());
3135 if (ord && !llvm::is_contained({NVVM::TensormapField::BOX_DIM,
3136 NVVM::TensormapField::GLOBAL_DIM,
3137 NVVM::TensormapField::GLOBAL_STRIDE,
3138 NVVM::TensormapField::ELEMENT_STRIDE},
3140 return emitOpError(
"ordinal is not supported for ")
3141 << fieldName <<
" field";
3143 auto invalidNewVal = [&](llvm::Twine type) -> std::string {
3144 return llvm::Twine(
"new_value must be specified and must be an " + type +
3145 " for " + llvm::Twine(fieldName) +
" field")
3149 auto invalidNewValAttr = [&]() -> std::string {
3150 return (llvm::Twine(
3151 "new_value_attr must be specified and must be a valid ") +
3152 llvm::Twine(fieldName) +
" attribute for " + fieldName +
" field")
3156 switch (getField()) {
3157 case NVVM::TensormapField::GLOBAL_ADDRESS:
3161 case NVVM::TensormapField::RANK:
3165 case NVVM::TensormapField::GLOBAL_STRIDE:
3167 return emitOpError(
"ordinal is required for global_stride field");
3171 case NVVM::TensormapField::BOX_DIM:
3172 case NVVM::TensormapField::GLOBAL_DIM:
3173 case NVVM::TensormapField::ELEMENT_STRIDE:
3176 << stringifyEnum(getField()) <<
" field";
3180 case NVVM::TensormapField::ELEMTYPE:
3181 if (!(newValAttr && llvm::isa<TensormapElemtypeAttr>(*newValAttr)))
3184 case NVVM::TensormapField::INTERLEAVE_LAYOUT:
3185 if (!(newValAttr && llvm::isa<TensormapInterleaveLayoutAttr>(*newValAttr)))
3188 case NVVM::TensormapField::SWIZZLE_MODE:
3189 if (!(newValAttr && llvm::isa<TensormapSwizzleModeAttr>(*newValAttr)))
3192 case NVVM::TensormapField::SWIZZLE_ATOMICITY:
3193 if (!(newValAttr && llvm::isa<TensormapSwizzleAtomicityAttr>(*newValAttr)))
3196 case NVVM::TensormapField::FILL_MODE:
3197 if (!(newValAttr && llvm::isa<TensormapFillModeAttr>(*newValAttr)))
3205template <
typename OpType>
3207 mlir::NVVM::FPRoundingMode rndMode = op.getRnd();
3208 mlir::NVVM::SaturationMode satMode = op.getSat();
3209 bool isFTZ = op.getFtz();
3212 mlir::Type opBaseType = isa<VectorType>(opType)
3213 ? cast<VectorType>(opType).getElementType()
3216 if (opBaseType.
isF64() && (satMode != NVVM::SaturationMode::NONE || isFTZ))
3217 return op.emitOpError(
"FTZ and saturation are not supported for "
3218 "additions/subtractions involving f64 type");
3220 if (opBaseType.
isF16() && !(rndMode == NVVM::FPRoundingMode::RN ||
3221 rndMode == NVVM::FPRoundingMode::NONE))
3222 return op.emitOpError(
"only RN rounding mode is supported for f16 and "
3223 "vector<2xf16> additions/subtractions");
3225 if (opBaseType.
isBF16()) {
3226 if (rndMode != NVVM::FPRoundingMode::RN &&
3227 rndMode != NVVM::FPRoundingMode::NONE)
3228 return op.emitOpError(
"only RN rounding mode is supported for bf16 and "
3229 "vector<2xbf16> additions/subtractions");
3230 if (satMode != NVVM::SaturationMode::NONE || isFTZ)
3231 return op.emitOpError(
"FTZ and saturation are not supported for bf16 and "
3232 "vector<2xbf16> additions/subtractions");
3239 if (opBaseType.
isF16() && isFTZ && satMode == NVVM::SaturationMode::NONE)
3240 return op.emitOpError(
"FTZ with no saturation is not supported for f16 and "
3241 "vector<2xf16> additions/subtractions");
3250LogicalResult NVVM::FmaOp::verify() {
3251 auto opType = getRes().getType();
3252 mlir::NVVM::FPRoundingMode rndMode = getRnd();
3253 mlir::NVVM::SaturationMode satMode = getSat();
3254 bool isFTZ = getFtz();
3255 bool isRelu = getRelu();
3256 bool hasOOB = getOob();
3258 auto getBaseFType = [](
Type type) ->
Type {
3259 if (isa<VectorType>(type))
3260 return cast<VectorType>(type).getElementType();
3264 auto opBaseType = getBaseFType(opType);
3266 if (rndMode == NVVM::FPRoundingMode::NONE)
3267 return emitOpError(
"rounding mode must be specified");
3269 if (isRelu && satMode == NVVM::SaturationMode::SAT)
3270 return emitOpError(
"relu and saturation are not supported together");
3272 if (hasOOB && (satMode == NVVM::SaturationMode::SAT || isFTZ))
3273 return emitOpError(
"oob is not supported with saturation or FTZ");
3275 if (!(opBaseType.isF16() || opBaseType.isBF16()) && (isRelu || hasOOB))
3276 return emitOpError(
"relu and oob are only supported for f16 and bf16");
3278 if (opBaseType.isF64() && (satMode != NVVM::SaturationMode::NONE || isFTZ))
3279 return emitOpError(
"FTZ and saturation are not supported for f64 type");
3281 if (opBaseType.isF16() && rndMode != NVVM::FPRoundingMode::RN)
3283 "only RN rounding mode is supported for f16 and vector<2xf16>");
3285 if (opBaseType.isBF16()) {
3286 if (rndMode != NVVM::FPRoundingMode::RN)
3288 "only RN rounding mode is supported for bf16 and vector<2xbf16>");
3289 if (satMode != NVVM::SaturationMode::NONE || isFTZ)
3291 "FTZ and saturation are not supported for bf16 and vector<2xbf16>");
3297LogicalResult NVVM::SqrtOp::verify() {
3298 if (getRnd() == NVVM::FPRoundingMode::NONE)
3299 return emitOpError(
"rounding mode cannot be None");
3301 if (getRes().
getType().isF64() && getFtz())
3302 return emitOpError(
"FTZ is not supported for f64");
3307LogicalResult NVVM::DivFOp::verify() {
3308 bool isApprox = getApprox();
3309 bool isFull = getFull();
3310 bool isF64 = getRes().getType().isF64();
3311 bool isFtz = getFtz();
3312 NVVM::FPRoundingMode rndMode = getRnd();
3314 if (isApprox && isFull)
3315 return emitOpError(
"'approx' and 'full' are mutually exclusive");
3317 if (isApprox || isFull) {
3319 return emitOpError(
"'approx' and 'full' forms are f32-only");
3320 if (rndMode != NVVM::FPRoundingMode::NONE)
3322 "'approx' and 'full' forms do not accept a rounding mode");
3327 if (rndMode == NVVM::FPRoundingMode::NONE)
3328 return emitOpError(
"rounding mode cannot be None for the rounded divide");
3330 return emitOpError(
"FTZ is not supported for f64");
3341 unsigned sizeInBits,
3343 field = builder.CreateZExtOrBitCast(field, builder.getInt32Ty());
3345 unsigned mask = (sizeInBits < 32 ? ((1u << sizeInBits) - 1) : 0xffffffffu);
3346 if (mask != 0xffffffffu)
3347 field = builder.CreateAnd(field, builder.getInt32(mask));
3349 field = builder.CreateZExtOrBitCast(field, builder.getInt64Ty());
3350 field = builder.CreateShl(field, start);
3352 return builder.CreateOr(
result, field);
3355void Tcgen05MmaSmemDescOp::createSmemDescriptor(
Operation &op,
3357 llvm::IRBuilderBase &builder) {
3358 auto thisOp = cast<NVVM::Tcgen05MmaSmemDescOp>(op);
3359 llvm::Value *smemDesc = builder.getInt64(0);
3364 builder, smemDesc, mt.
lookupValue(thisOp.getLeadingDimOffset()), 14, 16);
3366 builder, smemDesc, mt.
lookupValue(thisOp.getStrideDimOffset()), 14, 32);
3372 builder, smemDesc, mt.
lookupValue(thisOp.getLeadingDimMode()), 1, 52);
3376 mt.
mapValue(thisOp.getRes()) = smemDesc;
3383std::string NVVM::MBarrierInitOp::getPtx() {
3385 return isShared ? std::string(
"mbarrier.init.shared.b64 [%0], %1;")
3386 : std::string(
"mbarrier.init.b64 [%0], %1;");
3389std::string NVVM::MBarrierArriveExpectTxOp::getPtx() {
3392 ? std::string(
"mbarrier.arrive.expect_tx.shared.b64 _, [%0], %1;")
3393 : std::string(
"mbarrier.arrive.expect_tx.b64 _, [%0], %1;");
3396std::string NVVM::MBarrierTryWaitParityOp::getPtx() {
3398 llvm::StringRef space = isShared ?
".shared" :
"";
3400 return llvm::formatv(
"{\n\t"
3401 ".reg .pred P1; \n\t"
3403 "mbarrier.try_wait.parity{0}.b64 P1, [%0], %1, %2; \n\t"
3404 "@P1 bra.uni DONE; \n\t"
3405 "bra.uni LAB_WAIT; \n\t"
3422 LLVM::FNegOp::create(rewriter, loc, op.getRhs().getType(), op.getRhs());
3425 op.getRnd(), op.getSat(), op.getFtz());
3444 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_count
3445 : llvm::Intrinsic::nvvm_barrier_cta_sync_count;
3447 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_all
3448 : llvm::Intrinsic::nvvm_barrier_cta_sync_all;
3453static llvm::Intrinsic::ID
3456 case NVVM::BarrierReduction::AND:
3457 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_and_aligned_all
3458 : llvm::Intrinsic::nvvm_barrier_cta_red_and_all;
3459 case NVVM::BarrierReduction::OR:
3460 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_or_aligned_all
3461 : llvm::Intrinsic::nvvm_barrier_cta_red_or_all;
3462 case NVVM::BarrierReduction::POPC:
3463 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_popc_aligned_all
3464 : llvm::Intrinsic::nvvm_barrier_cta_red_popc_all;
3466 llvm_unreachable(
"unknown BarrierReduction kind");
3471 auto thisOp = cast<NVVM::BarrierOp>(op);
3472 llvm::Value *barrierId = thisOp.getBarrierId()
3474 : builder.getInt32(0);
3475 bool hasCount =
static_cast<bool>(thisOp.getNumberOfThreads());
3476 llvm::Intrinsic::ID
id =
3480 args.push_back(mt.
lookupValue(thisOp.getNumberOfThreads()));
3481 return {id, std::move(args)};
3486 auto thisOp = cast<NVVM::BarrierArriveOp>(op);
3487 llvm::Value *barrierId = thisOp.getBarrierId()
3489 : builder.getInt32(0);
3490 llvm::Value *numThreads = mt.
lookupValue(thisOp.getNumberOfThreads());
3491 llvm::Intrinsic::ID
id =
3493 ? llvm::Intrinsic::nvvm_barrier_cta_arrive_aligned_count
3494 : llvm::Intrinsic::nvvm_barrier_cta_arrive_count;
3495 return {id, {barrierId, numThreads}};
3500 auto thisOp = cast<NVVM::BarrierReductionOp>(op);
3502 thisOp.getAligned(), thisOp.getReductionOp());
3503 llvm::Value *barrierId = thisOp.getBarrierId()
3505 : builder.getInt32(0);
3508 builder.CreateICmpNE(mt.
lookupValue(thisOp.getReductionPredicate()),
3509 builder.getInt32(0))};
3510 return {id, std::move(args)};
3515 llvm::IRBuilderBase &builder) {
3516 auto thisOp = cast<NVVM::CosOp>(op);
3517 llvm::Intrinsic::ID
id = thisOp.getFtz()
3518 ? llvm::Intrinsic::nvvm_cos_approx_ftz_f
3519 : llvm::Intrinsic::nvvm_cos_approx_f;
3525 llvm::IRBuilderBase &builder) {
3526 auto thisOp = cast<NVVM::SinOp>(op);
3527 llvm::Intrinsic::ID
id = thisOp.getFtz()
3528 ? llvm::Intrinsic::nvvm_sin_approx_ftz_f
3529 : llvm::Intrinsic::nvvm_sin_approx_f;
3535 llvm::IRBuilderBase &builder) {
3536 auto thisOp = cast<NVVM::Log2Op>(op);
3537 llvm::Intrinsic::ID
id = thisOp.getFtz()
3538 ? llvm::Intrinsic::nvvm_lg2_approx_ftz_f
3539 : llvm::Intrinsic::nvvm_lg2_approx_f;
3545 llvm::IRBuilderBase &builder) {
3546 auto thisOp = cast<NVVM::Ex2Op>(op);
3547 llvm::Intrinsic::ID
id = thisOp.getFtz()
3548 ? llvm::Intrinsic::nvvm_ex2_approx_ftz
3549 : llvm::Intrinsic::nvvm_ex2_approx;
3555 llvm::IRBuilderBase &builder) {
3556 auto thisOp = cast<NVVM::RsqrtOp>(op);
3557 Type t = thisOp.getRes().getType();
3558 bool isFtz = thisOp.getFtz();
3560 llvm::Intrinsic::ID
id = [&] {
3562 return isFtz ? llvm::Intrinsic::nvvm_rsqrt_approx_ftz_f
3563 : llvm::Intrinsic::nvvm_rsqrt_approx_f;
3566 return isFtz ? llvm::Intrinsic::nvvm_rsqrt_approx_ftz_d
3567 : llvm::Intrinsic::nvvm_rsqrt_approx_d;
3575 llvm::IRBuilderBase &builder) {
3576 auto thisOp = cast<NVVM::SqrtOp>(op);
3577 Type t = thisOp.getRes().getType();
3578 NVVM::FPRoundingMode rndMode = thisOp.getRnd();
3579 bool isFtz = thisOp.getFtz();
3583 unsigned rndIndex =
static_cast<unsigned>(rndMode) - 1;
3585 static constexpr llvm::Intrinsic::ID f32IDs[] = {
3586 llvm::Intrinsic::nvvm_sqrt_rn_f,
3587 llvm::Intrinsic::nvvm_sqrt_rm_f,
3588 llvm::Intrinsic::nvvm_sqrt_rp_f,
3589 llvm::Intrinsic::nvvm_sqrt_rz_f,
3591 static constexpr llvm::Intrinsic::ID f32FTZIDs[] = {
3592 llvm::Intrinsic::nvvm_sqrt_rn_ftz_f,
3593 llvm::Intrinsic::nvvm_sqrt_rm_ftz_f,
3594 llvm::Intrinsic::nvvm_sqrt_rp_ftz_f,
3595 llvm::Intrinsic::nvvm_sqrt_rz_ftz_f,
3597 static constexpr llvm::Intrinsic::ID f64IDs[] = {
3598 llvm::Intrinsic::nvvm_sqrt_rn_d,
3599 llvm::Intrinsic::nvvm_sqrt_rm_d,
3600 llvm::Intrinsic::nvvm_sqrt_rp_d,
3601 llvm::Intrinsic::nvvm_sqrt_rz_d,
3604 llvm::Intrinsic::ID
id =
3605 t.
isF32() ? (isFtz ? f32FTZIDs[rndIndex] : f32IDs[rndIndex])
3613 llvm::IRBuilderBase &builder) {
3614 auto thisOp = cast<NVVM::SqrtApproxOp>(op);
3615 llvm::Intrinsic::ID
id = thisOp.getFtz()
3616 ? llvm::Intrinsic::nvvm_sqrt_approx_ftz_f
3617 : llvm::Intrinsic::nvvm_sqrt_approx_f;
3623 llvm::IRBuilderBase &builder) {
3624 auto thisOp = cast<NVVM::DivFOp>(op);
3625 bool isFtz = thisOp.getFtz();
3627 llvm::Intrinsic::ID id;
3629 if (thisOp.getApprox()) {
3630 id = isFtz ? llvm::Intrinsic::nvvm_div_approx_ftz_f
3631 : llvm::Intrinsic::nvvm_div_approx_f;
3632 }
else if (thisOp.getFull()) {
3635 id = isFtz ? llvm::Intrinsic::nvvm_div_full_ftz
3636 : llvm::Intrinsic::nvvm_div_full;
3639 unsigned rndIndex =
static_cast<unsigned>(thisOp.getRnd()) - 1;
3641 static constexpr llvm::Intrinsic::ID f32IDs[] = {
3642 llvm::Intrinsic::nvvm_div_rn_f,
3643 llvm::Intrinsic::nvvm_div_rm_f,
3644 llvm::Intrinsic::nvvm_div_rp_f,
3645 llvm::Intrinsic::nvvm_div_rz_f,
3647 static constexpr llvm::Intrinsic::ID f32FTZIDs[] = {
3648 llvm::Intrinsic::nvvm_div_rn_ftz_f,
3649 llvm::Intrinsic::nvvm_div_rm_ftz_f,
3650 llvm::Intrinsic::nvvm_div_rp_ftz_f,
3651 llvm::Intrinsic::nvvm_div_rz_ftz_f,
3653 static constexpr llvm::Intrinsic::ID f64IDs[] = {
3654 llvm::Intrinsic::nvvm_div_rn_d,
3655 llvm::Intrinsic::nvvm_div_rm_d,
3656 llvm::Intrinsic::nvvm_div_rp_d,
3657 llvm::Intrinsic::nvvm_div_rz_d,
3659 Type t = thisOp.getRes().getType();
3660 id = t.
isF32() ? (isFtz ? f32FTZIDs[rndIndex] : f32IDs[rndIndex])
3670 llvm::IRBuilderBase &builder) {
3671 auto thisOp = cast<NVVM::PMEventOp>(op);
3675 llvm::Value *maskVal;
3676 if (
auto eventAttr = thisOp.getEventIdAttr()) {
3677 uint16_t mask =
static_cast<uint16_t
>(1u << eventAttr.getInt());
3678 maskVal = llvm::ConstantInt::get(i16Ty, mask);
3681 llvm::ConstantInt::get(i16Ty, thisOp.getMaskedEventIdAttr().getValue());
3684 return {llvm::Intrinsic::nvvm_pm_event_mask, {maskVal}};
3689 auto thisOp = cast<NVVM::MBarrierInitOp>(op);
3691 llvm::Intrinsic::ID
id = isShared ? llvm::Intrinsic::nvvm_mbarrier_init_shared
3692 : llvm::Intrinsic::nvvm_mbarrier_init;
3697 args.push_back(mt.
lookupValue(thisOp.getCount()));
3699 return {id, std::move(args)};
3704 auto thisOp = cast<NVVM::MBarrierInvalOp>(op);
3706 llvm::Intrinsic::ID
id = isShared
3707 ? llvm::Intrinsic::nvvm_mbarrier_inval_shared
3708 : llvm::Intrinsic::nvvm_mbarrier_inval;
3715 auto thisOp = cast<NVVM::MBarrierExpectTxOp>(op);
3718 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3721 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
3723 static constexpr llvm::Intrinsic::ID IDs[] = {
3724 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cta,
3725 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cluster,
3726 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cta,
3727 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cluster};
3732 args.push_back(mt.
lookupValue(thisOp.getTxcount()));
3734 return {IDs[
index], std::move(args)};
3739 auto thisOp = cast<NVVM::MBarrierCompleteTxOp>(op);
3742 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3745 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
3747 static constexpr llvm::Intrinsic::ID IDs[] = {
3748 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cta,
3749 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cluster,
3750 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cta,
3751 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cluster};
3756 args.push_back(mt.
lookupValue(thisOp.getTxcount()));
3758 return {IDs[
index], std::move(args)};
3763 auto thisOp = cast<NVVM::MBarrierArriveOp>(op);
3766 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3769 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
3771 static constexpr llvm::Intrinsic::ID IDs[] = {
3772 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta,
3773 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cluster,
3774 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cta,
3775 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cluster};
3776 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
3777 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cta,
3778 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cluster,
3779 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cta,
3781 nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cluster};
3782 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
3786 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
3793 bool hasCount =
static_cast<bool>(thisOp.getCount());
3795 (
id == llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta))
3796 return {llvm::Intrinsic::nvvm_mbarrier_arrive_shared, {mbar}};
3800 llvm::Value *count =
3802 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);
3803 return {id, {mbar, count}};
3808 auto thisOp = cast<NVVM::MBarrierArriveDropOp>(op);
3811 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3814 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
3816 static constexpr llvm::Intrinsic::ID IDs[] = {
3817 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cta,
3818 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cluster,
3819 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cta,
3820 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cluster};
3821 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
3822 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cta,
3824 nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cluster,
3826 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cta,
3828 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cluster};
3829 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
3833 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
3839 bool hasCount =
static_cast<bool>(thisOp.getCount());
3840 llvm::Value *count =
3842 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);
3844 return {id, {mbar, count}};
3847bool MBarrierArriveExpectTxOp::getAsmValues(
3854 for (
auto val : getOperands())
3862 auto thisOp = cast<NVVM::MBarrierArriveExpectTxOp>(op);
3865 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3868 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
3871 static constexpr llvm::Intrinsic::ID IDs[] = {
3872 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cta,
3873 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cluster,
3874 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cta,
3875 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cluster};
3876 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
3877 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cta,
3878 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cluster,
3879 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cta,
3880 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cluster};
3882 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
3885 llvm::Value *txcount = mt.
lookupValue(thisOp.getTxcount());
3886 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
3891 return {id, {mbar, txcount}};
3896 auto thisOp = cast<NVVM::MBarrierArriveDropExpectTxOp>(op);
3899 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3902 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
3905 static constexpr llvm::Intrinsic::ID IDs[] = {
3906 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cta,
3907 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cluster,
3908 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cta,
3909 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cluster};
3910 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
3911 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cta,
3912 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cluster,
3913 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cta,
3914 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cluster};
3916 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
3919 llvm::Value *txcount = mt.
lookupValue(thisOp.getTxcount());
3920 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
3925 return {id, {mbar, txcount}};
3930 auto thisOp = cast<NVVM::MBarrierArriveNocompleteOp>(op);
3932 llvm::Intrinsic::ID
id =
3933 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete_shared
3934 : llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete;
3938 args.push_back(mt.
lookupValue(thisOp.getCount()));
3940 return {id, std::move(args)};
3945 auto thisOp = cast<NVVM::MBarrierArriveDropNocompleteOp>(op);
3947 llvm::Intrinsic::ID
id =
3948 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete_shared
3949 : llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete;
3953 args.push_back(mt.
lookupValue(thisOp.getCount()));
3955 return {id, std::move(args)};
3960 auto thisOp = cast<NVVM::MBarrierTestWaitOp>(op);
3961 bool isPhaseParity = thisOp.getStateOrPhase().getType().isInteger(32);
3962 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3965 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isPhaseParity ? 1 : 0);
3968 static constexpr llvm::Intrinsic::ID IDs[] = {
3969 llvm::Intrinsic::nvvm_mbarrier_test_wait_scope_cta_space_cta,
3970 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_scope_cta_space_cta,
3971 llvm::Intrinsic::nvvm_mbarrier_test_wait_scope_cluster_space_cta,
3972 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_scope_cluster_space_cta};
3973 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
3974 llvm::Intrinsic::nvvm_mbarrier_test_wait_relaxed_scope_cta_space_cta,
3975 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_relaxed_scope_cta_space_cta,
3976 llvm::Intrinsic::nvvm_mbarrier_test_wait_relaxed_scope_cluster_space_cta,
3977 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_relaxed_scope_cluster_space_cta};
3979 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
3982 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
3983 llvm::Value *input = mt.
lookupValue(thisOp.getStateOrPhase());
3988 return {id, {mbar, input}};
3993 auto thisOp = cast<NVVM::MBarrierTryWaitOp>(op);
3994 bool isPhaseParity = thisOp.getStateOrPhase().getType().isInteger(32);
3995 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
3996 bool hasTicks =
static_cast<bool>(thisOp.getTicks());
4000 size_t index = ((hasTicks ? 1 : 0) << 2) | ((isClusterScope ? 1 : 0) << 1) |
4001 (isPhaseParity ? 1 : 0);
4004 static constexpr llvm::Intrinsic::ID IDs[] = {
4005 llvm::Intrinsic::nvvm_mbarrier_try_wait_scope_cta_space_cta,
4006 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_scope_cta_space_cta,
4007 llvm::Intrinsic::nvvm_mbarrier_try_wait_scope_cluster_space_cta,
4008 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_scope_cluster_space_cta,
4009 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_scope_cta_space_cta,
4010 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_scope_cta_space_cta,
4011 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_scope_cluster_space_cta,
4012 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_scope_cluster_space_cta};
4013 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4014 llvm::Intrinsic::nvvm_mbarrier_try_wait_relaxed_scope_cta_space_cta,
4015 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_relaxed_scope_cta_space_cta,
4016 llvm::Intrinsic::nvvm_mbarrier_try_wait_relaxed_scope_cluster_space_cta,
4017 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_relaxed_scope_cluster_space_cta,
4018 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_relaxed_scope_cta_space_cta,
4019 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_relaxed_scope_cta_space_cta,
4020 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_relaxed_scope_cluster_space_cta,
4021 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_relaxed_scope_cluster_space_cta};
4023 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4026 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4033 args.push_back(mbar);
4034 args.push_back(mt.
lookupValue(thisOp.getStateOrPhase()));
4036 args.push_back(mt.
lookupValue(thisOp.getTicks()));
4038 return {id, std::move(args)};
4043 auto thisOp = cast<NVVM::CpAsyncMBarrierArriveOp>(op);
4046 llvm::Intrinsic::ID id;
4047 if (thisOp.getNoinc()) {
4048 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc_shared
4049 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc;
4051 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_shared
4052 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive;
4060 llvm::IRBuilderBase &builder) {
4061 auto thisOp = cast<NVVM::MovMatrixOp>(op);
4062 return {llvm::Intrinsic::nvvm_movmatrix_sync_aligned_m8n8_trans_b16,
4066#define CP_ASYNC_ID_IMPL(mod, size, suffix) \
4067 llvm::Intrinsic::nvvm_cp_async_##mod##_shared_global_##size##suffix
4069#define GET_CP_ASYNC_ID(mod, size, has_cpsize) \
4070 has_cpsize ? CP_ASYNC_ID_IMPL(mod, size, _s) : CP_ASYNC_ID_IMPL(mod, size, )
4075 llvm::Intrinsic::ID id;
4077 auto cpAsyncOp = cast<NVVM::CpAsyncOp>(op);
4078 bool hasCpSize =
static_cast<bool>(cpAsyncOp.getCpSize());
4079 switch (cpAsyncOp.getSize()) {
4087 id = (cpAsyncOp.getModifier() == NVVM::LoadCacheModifierKind::CG)
4092 llvm_unreachable(
"Invalid copy size in CpAsyncOp.");
4096 args.push_back(mt.
lookupValue(cpAsyncOp.getDst()));
4097 args.push_back(mt.
lookupValue(cpAsyncOp.getSrc()));
4099 args.push_back(mt.
lookupValue(cpAsyncOp.getCpSize()));
4106 auto thisOp = cast<NVVM::CpAsyncBulkPrefetchOp>(op);
4108 llvm::Intrinsic::ID
id = llvm::Intrinsic::nvvm_cp_async_bulk_prefetch_L2;
4111 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4115 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4116 llvm::Value *i64Unused =
4117 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4118 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4119 args.push_back(builder.getInt1(hasCacheHint));
4121 return {id, std::move(args)};
4126 auto thisOp = cast<NVVM::CpAsyncBulkGlobalToSharedClusterOp>(op);
4130 args.push_back(mt.
lookupValue(thisOp.getDstMem()));
4132 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4136 mlir::Value multicastMask = thisOp.getMulticastMask();
4137 const bool hasMulticastMask =
static_cast<bool>(multicastMask);
4140 llvm::Value *i16Unused = llvm::ConstantInt::get(builder.getInt16Ty(), 0);
4141 args.push_back(hasMulticastMask ? mt.
lookupValue(multicastMask)
4147 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4148 llvm::Value *i64Unused = llvm::ConstantInt::get(builder.getInt64Ty(), 0);
4149 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4153 args.push_back(builder.getInt1(hasMulticastMask));
4154 args.push_back(builder.getInt1(hasCacheHint));
4156 llvm::Intrinsic::ID
id =
4158 ? llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cta
4159 : llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cluster;
4161 return {id, std::move(args)};
4166 auto thisOp = cast<NVVM::CpAsyncBulkSharedCTAToGlobalOp>(op);
4168 llvm::Intrinsic::ID
id =
4169 llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global;
4172 args.push_back(mt.
lookupValue(thisOp.getDstMem()));
4173 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4177 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4178 llvm::Value *i64Unused =
4179 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4180 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4181 args.push_back(builder.getInt1(hasCacheHint));
4184 if (
mlir::Value byteMask = thisOp.getByteMask()) {
4186 id = llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global_bytemask;
4189 return {id, std::move(args)};
4192bool CpAsyncBulkTensorGlobalToSharedClusterOp::getAsmValues(
4199 for (
auto val : getOperands())
4206CpAsyncBulkTensorGlobalToSharedClusterOp::getIntrinsicIDAndArgs(
4208 auto thisOp = cast<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(op);
4209 const bool isCTAOnly = thisOp.getIsCTAOnly();
4213 args.push_back(mt.
lookupValue(thisOp.getDstMem()));
4215 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4225 const bool hasMC =
static_cast<bool>(mcMask);
4226 llvm::Value *i16Zero =
4227 llvm::ConstantInt::get(llvm::Type::getInt16Ty(mt.
getLLVMContext()), 0);
4231 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4232 llvm::Value *i64Zero =
4233 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4239 thisOp.getGroup() ? (
static_cast<int32_t
>(*thisOp.getGroup()) + 1) : 0;
4241 llvm::ConstantInt::get(llvm::Type::getInt32Ty(mt.
getLLVMContext()), val);
4245 args.push_back(hasMC ? mt.
lookupValue(mcMask) : i16Zero);
4246 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Zero);
4247 args.push_back(builder.getInt1(hasMC));
4248 args.push_back(builder.getInt1(hasCacheHint));
4252 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Zero);
4253 args.push_back(builder.getInt1(hasCacheHint));
4256 constexpr size_t numDims = 5;
4257 constexpr size_t numModes = 5;
4258 using rowTy = std::array<llvm::Intrinsic::ID, numDims + 1>;
4259 using TableTy = std::array<rowTy, numModes>;
4260 static constexpr TableTy IDTable{
4261 {{
notIntrinsic, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_1d,
4262 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_2d,
4263 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_3d,
4264 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_4d,
4265 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_5d},
4267 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_3d,
4268 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_4d,
4269 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_5d},
4271 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_3d,
4272 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_4d,
4273 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_5d},
4275 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_3d,
4276 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_4d,
4277 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_5d},
4279 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_gather4_2d}}};
4281 static constexpr TableTy IDTableCTA{
4283 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_1d,
4284 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_2d,
4285 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_3d,
4286 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_4d,
4287 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_5d},
4289 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_3d,
4290 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_4d,
4291 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_5d},
4293 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_3d,
4294 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_4d,
4295 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_5d},
4297 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_3d,
4298 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_4d,
4299 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_5d},
4301 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_gather4_2d}}};
4304 (getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1) &&
4305 (getMaxEnumValForTMALoadMode() == std::size(IDTableCTA) - 1),
4306 "TMALoadModes must match number of rows in IDTable and IDTableCTA");
4307 size_t mode =
static_cast<size_t>(thisOp.getMode());
4308 size_t dim = thisOp.getCoordinates().size();
4309 auto id = isCTAOnly ? IDTableCTA[mode][dim] : IDTable[mode][dim];
4311 "Invalid intrinsic for CpAsyncBulkTensorGlobalToSharedClusterOp.");
4313 return {id, std::move(args)};
4318 auto thisOp = cast<NVVM::CpAsyncBulkTensorPrefetchOp>(op);
4322 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4324 for (
auto v : thisOp.getCoordinates())
4326 for (
auto v : thisOp.getIm2colOffsets())
4330 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4331 llvm::Value *i64Unused =
4332 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4333 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4334 args.push_back(builder.getInt1(hasCacheHint));
4336 const unsigned NI = llvm::Intrinsic::not_intrinsic;
4337 static constexpr llvm::Intrinsic::ID IDTable[][6] = {
4338 {NI, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_1d,
4339 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_2d,
4340 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_3d,
4341 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_4d,
4342 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_5d},
4344 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_3d,
4345 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_4d,
4346 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_5d},
4348 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_3d,
4349 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_4d,
4350 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_5d},
4352 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_3d,
4353 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_4d,
4354 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_5d},
4355 {NI, NI, NI, NI, NI,
4356 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_gather4_2d}};
4358 static_assert(getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1,
4359 "TMALoadModes must match number of rows in IDTable");
4360 size_t mode =
static_cast<size_t>(thisOp.getMode());
4361 size_t dim = thisOp.getCoordinates().size();
4362 llvm::Intrinsic::ID
id = IDTable[mode][dim];
4363 if (
id == llvm::Intrinsic::not_intrinsic)
4364 llvm_unreachable(
"Invalid intrinsic for CpAsyncBulkTensorPrefetchOp.");
4366 return {id, std::move(args)};
4370CpAsyncBulkTensorSharedCTAToGlobalOp::getIntrinsicIDAndArgs(
4372 auto thisOp = cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(op);
4376 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4377 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4379 for (
auto v : thisOp.getCoordinates())
4383 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4384 llvm::Value *i64Unused =
4385 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4386 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4387 args.push_back(builder.getInt1(hasCacheHint));
4389 using namespace llvm::Intrinsic;
4390 const unsigned NI = not_intrinsic;
4391 static constexpr ID IDTable[][6] = {
4392 {NI, nvvm_cp_async_bulk_tensor_s2g_tile_1d,
4393 nvvm_cp_async_bulk_tensor_s2g_tile_2d,
4394 nvvm_cp_async_bulk_tensor_s2g_tile_3d,
4395 nvvm_cp_async_bulk_tensor_s2g_tile_4d,
4396 nvvm_cp_async_bulk_tensor_s2g_tile_5d},
4397 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_3d,
4398 nvvm_cp_async_bulk_tensor_s2g_im2col_4d,
4399 nvvm_cp_async_bulk_tensor_s2g_im2col_5d},
4400 {NI, NI, NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_2d},
4401 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_w_3d,
4402 nvvm_cp_async_bulk_tensor_s2g_im2col_w_4d,
4403 nvvm_cp_async_bulk_tensor_s2g_im2col_w_5d}};
4405 static_assert(getMaxEnumValForTMAStoreMode() == std::size(IDTable) - 1,
4406 "TMAStoreModes must match number of rows in IDTable");
4407 size_t mode =
static_cast<size_t>(thisOp.getMode());
4408 size_t dim = thisOp.getCoordinates().size();
4409 ID
id = IDTable[mode][dim];
4410 if (
id == llvm::Intrinsic::not_intrinsic)
4412 "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOp.");
4414 return {id, std::move(args)};
4419 auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOp>(op);
4422 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4423 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4424 for (
Value v : thisOp.getCoordinates())
4428 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4429 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint)
4430 : builder.getInt64(0));
4431 args.push_back(builder.getInt32(
static_cast<uint32_t
>(thisOp.getRedKind())));
4432 args.push_back(builder.getInt1(hasCacheHint));
4434 using namespace llvm::Intrinsic;
4435 const unsigned NI = not_intrinsic;
4436 static constexpr ID IDTable[][6] = {
4437 {NI, nvvm_cp_async_bulk_tensor_reduce_tile_1d,
4438 nvvm_cp_async_bulk_tensor_reduce_tile_2d,
4439 nvvm_cp_async_bulk_tensor_reduce_tile_3d,
4440 nvvm_cp_async_bulk_tensor_reduce_tile_4d,
4441 nvvm_cp_async_bulk_tensor_reduce_tile_5d},
4442 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_3d,
4443 nvvm_cp_async_bulk_tensor_reduce_im2col_4d,
4444 nvvm_cp_async_bulk_tensor_reduce_im2col_5d},
4445 {NI, NI, NI, NI, NI, NI},
4446 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_w_3d,
4447 nvvm_cp_async_bulk_tensor_reduce_im2col_w_4d,
4448 nvvm_cp_async_bulk_tensor_reduce_im2col_w_5d}};
4450 size_t mode =
static_cast<size_t>(thisOp.getMode());
4451 size_t dim = thisOp.getCoordinates().size();
4452 assert(mode < std::size(IDTable) &&
4453 "Invalid mode for CpAsyncBulkTensorReduceOp");
4454 assert(dim < std::size(IDTable[mode]) &&
4455 "Invalid dim for CpAsyncBulkTensorReduceOp");
4457 ID intrinsicID = IDTable[mode][dim];
4458 assert(intrinsicID != NI &&
4459 "Invalid intrinsic for CpAsyncBulkTensorReduceOp");
4460 return {intrinsicID, std::move(args)};
4465#define CVT_F2TF32_ID_IMPL(rnd, relu, sf) \
4466 hasRelu ? llvm::Intrinsic::nvvm_f2tf32_##rnd##relu##sf \
4467 : llvm::Intrinsic::nvvm_f2tf32_##rnd##sf
4469#define GET_CVT_F2TF32_ID(rnd, relu, sf) \
4470 hasSatFinite ? CVT_F2TF32_ID_IMPL(rnd, relu, sf) \
4471 : CVT_F2TF32_ID_IMPL(rnd, relu, )
4474ConvertFloatToTF32Op::getIntrinsicID(NVVM::FPRoundingMode rnd,
4475 NVVM::SaturationMode sat,
bool hasRelu) {
4476 using RndMode = NVVM::FPRoundingMode;
4477 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
4486 llvm_unreachable(
"Invalid RoundingMode for CvtFloatToTF32Op");
4491ConvertF32x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF4x2Op op,
4493 llvm::IRBuilderBase &builder) {
4498 bool hasRelu = op.getRelu();
4500 llvm::Intrinsic::ID intId =
4501 hasRelu ? llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_relu_satfinite
4502 : llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_satfinite;
4504 return {intId, std::move(args)};
4507#define GET_F32x2_TO_F6x2_ID(type, has_relu) \
4508 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu_satfinite \
4509 : llvm::Intrinsic::nvvm_ff_to_##type##_rn_satfinite
4511llvm::Intrinsic::ID ConvertF32x2ToF6x2Op::getIntrinsicID(
mlir::Type dstTy,
4514 .Case([&](mlir::Float6E2M3FNType) {
4517 .Case([&](mlir::Float6E3M2FNType) {
4521 llvm_unreachable(
"Invalid conversion in ConvertF32x2ToF6x2Op");
4522 return llvm::Intrinsic::not_intrinsic;
4527ConvertF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF16x2ToF4x2Op &op,
4529 llvm::IRBuilderBase &builder) {
4531 bool hasRelu = op.getRelu();
4533 llvm::Intrinsic::ID intId = llvm::Intrinsic::not_intrinsic;
4535 if (llvm::isa<mlir::Float4E2M1FNType>(dstTy))
4536 intId = hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e2m1x2_rn_relu_satfinite
4537 : llvm::Intrinsic::nvvm_f16x2_to_e2m1x2_rn_satfinite;
4542 return {intId, std::move(args)};
4546ConvertBF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertBF16x2ToF4x2Op &op,
4548 llvm::IRBuilderBase &builder) {
4550 bool hasRelu = op.getRelu();
4552 llvm::Intrinsic::ID intId = llvm::Intrinsic::not_intrinsic;
4554 if (llvm::isa<mlir::Float4E2M1FNType>(dstTy))
4555 intId = hasRelu ? llvm::Intrinsic::nvvm_bf16x2_to_e2m1x2_rn_relu_satfinite
4556 : llvm::Intrinsic::nvvm_bf16x2_to_e2m1x2_rn_satfinite;
4561 return {intId, std::move(args)};
4564llvm::Intrinsic::ID ConvertF16x2ToF6x2Op::getIntrinsicID(
mlir::Type dstTy,
4567 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {
4568 return hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e2m3x2_rn_relu_satfinite
4569 : llvm::Intrinsic::nvvm_f16x2_to_e2m3x2_rn_satfinite;
4571 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {
4572 return hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e3m2x2_rn_relu_satfinite
4573 : llvm::Intrinsic::nvvm_f16x2_to_e3m2x2_rn_satfinite;
4576 llvm_unreachable(
"Invalid conversion in ConvertF16x2ToF6x2Op");
4577 return llvm::Intrinsic::not_intrinsic;
4581llvm::Intrinsic::ID ConvertBF16x2ToF6x2Op::getIntrinsicID(
mlir::Type dstTy,
4584 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {
4586 ? llvm::Intrinsic::nvvm_bf16x2_to_e2m3x2_rn_relu_satfinite
4587 : llvm::Intrinsic::nvvm_bf16x2_to_e2m3x2_rn_satfinite;
4589 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {
4591 ? llvm::Intrinsic::nvvm_bf16x2_to_e3m2x2_rn_relu_satfinite
4592 : llvm::Intrinsic::nvvm_bf16x2_to_e3m2x2_rn_satfinite;
4595 llvm_unreachable(
"Invalid conversion in ConvertBF16x2ToF6x2Op");
4596 return llvm::Intrinsic::not_intrinsic;
4600#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf) \
4601 has_satf ? llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd##_satfinite \
4602 : llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd
4604#define GET_F32x2_TO_F8X2_S_ID(type, has_relu) \
4605 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu \
4606 : llvm::Intrinsic::nvvm_ff_to_##type##_rn
4609ConvertF32x2ToF8x2Op::getIntrinsicID(
mlir::Type dstTy, NVVM::FPRoundingMode rnd,
4610 NVVM::SaturationMode sat,
bool hasRelu) {
4611 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
4612 bool hasRoundingModeRZ = (rnd == NVVM::FPRoundingMode::RZ);
4613 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);
4616 .Case([&](mlir::Float8E4M3FNType) {
4619 .Case([&](mlir::Float8E5M2Type) {
4622 .Case([&](mlir::Float8E8M0FNUType) {
4623 if (hasRoundingModeRZ)
4625 else if (hasRoundingModeRP)
4628 llvm_unreachable(
"Invalid conversion in ConvertF32x2ToF8x2Op");
4631 llvm_unreachable(
"Invalid conversion in ConvertF32x2ToF8x2Op");
4632 return llvm::Intrinsic::not_intrinsic;
4636#define GET_F16x2_TO_F8X2_ID(type, has_relu) \
4637 has_relu ? llvm::Intrinsic::nvvm_f16x2_to_##type##_rn_relu \
4638 : llvm::Intrinsic::nvvm_f16x2_to_##type##_rn
4640llvm::Intrinsic::ID ConvertF16x2ToF8x2Op::getIntrinsicID(
mlir::Type dstTy,
4643 .Case([&](mlir::Float8E4M3FNType) {
4646 .Case([&](mlir::Float8E5M2Type) {
4650 llvm_unreachable(
"Invalid conversion in ConvertF16x2ToF8x2Op");
4651 return llvm::Intrinsic::not_intrinsic;
4656ConvertBF16x2ToF8x2Op::getIntrinsicID(
mlir::Type dstTy,
4657 NVVM::FPRoundingMode rnd,
4658 NVVM::SaturationMode sat,
bool hasRelu) {
4659 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
4661 static constexpr llvm::Intrinsic::ID ue8m0x2IDs[] = {
4662 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rz,
4663 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rp,
4664 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rz_satfinite,
4665 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rp_satfinite,
4669 .Case<mlir::Float8E4M3FNType>([&](mlir::Float8E4M3FNType) {
4671 ? llvm::Intrinsic::nvvm_bf16x2_to_e4m3x2_rn_relu_satfinite
4672 : llvm::Intrinsic::nvvm_bf16x2_to_e4m3x2_rn_satfinite;
4674 .Case<mlir::Float8E5M2Type>([&](mlir::Float8E5M2Type) {
4676 ? llvm::Intrinsic::nvvm_bf16x2_to_e5m2x2_rn_relu_satfinite
4677 : llvm::Intrinsic::nvvm_bf16x2_to_e5m2x2_rn_satfinite;
4679 .Case<mlir::Float8E8M0FNUType>([&](mlir::Float8E8M0FNUType) {
4680 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);
4681 unsigned index = (hasSatFinite << 1) | hasRoundingModeRP;
4682 return ue8m0x2IDs[
index];
4685 llvm_unreachable(
"Invalid conversion in ConvertBF16x2ToF8x2Op");
4686 return llvm::Intrinsic::not_intrinsic;
4692 auto curOp = cast<NVVM::ConvertF8x2ToF16x2Op>(op);
4694 bool hasRelu = curOp.getRelu();
4696 llvm::Intrinsic::ID intId =
4698 .Case([&](Float8E4M3FNType type) {
4699 return hasRelu ? llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn_relu
4700 : llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn;
4702 .Case([&](Float8E5M2Type type) {
4703 return hasRelu ? llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn_relu
4704 : llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn;
4707 llvm_unreachable(
"Invalid type for ConvertF8x2ToF16x2Op");
4708 return llvm::Intrinsic::not_intrinsic;
4711 llvm::Value *packedI16 =
4712 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
4713 llvm::Type::getInt16Ty(builder.getContext()));
4715 return {intId, {packedI16}};
4720 auto curOp = cast<NVVM::ConvertF8x2ToBF16x2Op>(op);
4721 bool hasScale =
static_cast<bool>(curOp.getScaleFactor());
4722 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
4723 bool hasRelu = curOp.getRelu();
4725 static constexpr llvm::Intrinsic::ID E4M3Ids[] = {
4726 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_scale_n2_ue8m0,
4727 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
4728 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
4729 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
4732 static constexpr llvm::Intrinsic::ID E5M2Ids[] = {
4733 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_scale_n2_ue8m0,
4734 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
4735 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
4736 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
4739 llvm::Intrinsic::ID intId =
4741 .Case([&](Float8E8M0FNUType type) {
4742 return llvm::Intrinsic::nvvm_ue8m0x2_to_bf16x2;
4744 .Case([&](Float8E4M3FNType type) {
4745 return E4M3Ids[hasSatfinite << 1 | hasRelu];
4747 .Case([&](Float8E5M2Type type) {
4748 return E5M2Ids[hasSatfinite << 1 | hasRelu];
4751 llvm_unreachable(
"Invalid type for ConvertF8x2ToBF16x2Op");
4752 return llvm::Intrinsic::not_intrinsic;
4754 llvm::Value *packedI16 =
4755 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
4756 llvm::Type::getInt16Ty(builder.getContext()));
4759 args.push_back(packedI16);
4760 if (!isa<Float8E8M0FNUType>(curOp.getSrcType()))
4763 : builder.getInt16(0x7f7f));
4766 return {intId, std::move(args)};
4771 auto curOp = cast<NVVM::ConvertF6x2ToF16x2Op>(op);
4773 bool hasRelu = curOp.getRelu();
4775 llvm::Intrinsic::ID intId =
4777 .Case([&](Float6E2M3FNType type) {
4778 return hasRelu ? llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn_relu
4779 : llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn;
4781 .Case([&](Float6E3M2FNType type) {
4782 return hasRelu ? llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn_relu
4783 : llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn;
4786 llvm_unreachable(
"Invalid type for ConvertF6x2ToF16x2Op");
4787 return llvm::Intrinsic::not_intrinsic;
4790 llvm::Value *packedI16 =
4791 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
4792 llvm::Type::getInt16Ty(builder.getContext()));
4794 return {intId, {packedI16}};
4799 auto curOp = cast<NVVM::ConvertF6x2ToBF16x2Op>(op);
4800 bool hasScale =
static_cast<bool>(curOp.getScaleFactor());
4801 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
4802 bool hasRelu = curOp.getRelu();
4804 static constexpr llvm::Intrinsic::ID E2M3Ids[] = {
4805 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_scale_n2_ue8m0,
4806 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
4807 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
4808 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
4811 static constexpr llvm::Intrinsic::ID E3M2Ids[] = {
4812 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_scale_n2_ue8m0,
4813 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
4814 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
4815 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
4818 unsigned idx = (hasSatfinite << 1) | hasRelu;
4819 llvm::Intrinsic::ID intId =
4821 .Case([&](Float6E2M3FNType type) {
return E2M3Ids[idx]; })
4822 .Case([&](Float6E3M2FNType type) {
return E3M2Ids[idx]; })
4824 llvm_unreachable(
"Invalid type for ConvertF6x2ToBF16x2Op");
4825 return llvm::Intrinsic::not_intrinsic;
4828 llvm::Value *packedI16 =
4829 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
4830 llvm::Type::getInt16Ty(builder.getContext()));
4833 args.push_back(packedI16);
4840 return {intId, std::move(args)};
4845 auto curOp = cast<NVVM::ConvertF4x2ToF16x2Op>(op);
4847 bool hasRelu = curOp.getRelu();
4849 llvm::Intrinsic::ID intId =
4851 .Case([&](Float4E2M1FNType type) {
4852 return hasRelu ? llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn_relu
4853 : llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn;
4856 llvm_unreachable(
"Invalid type for ConvertF4x2ToF16x2Op");
4857 return llvm::Intrinsic::not_intrinsic;
4860 llvm::Value *extendedI16 =
4861 builder.CreateZExt(mt.
lookupValue(curOp.getSrc()),
4862 llvm::Type::getInt16Ty(builder.getContext()));
4864 return {intId, {extendedI16}};
4869 auto curOp = cast<NVVM::ConvertF4x2ToBF16x2Op>(op);
4870 bool hasScale =
static_cast<bool>(curOp.getScaleFactor());
4871 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
4872 bool hasRelu = curOp.getRelu();
4874 static constexpr llvm::Intrinsic::ID E2M1Ids[] = {
4875 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_scale_n2_ue8m0,
4876 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
4877 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
4878 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
4881 unsigned idx = (hasSatfinite << 1) | hasRelu;
4882 llvm::Intrinsic::ID intId =
4884 .Case([&](Float4E2M1FNType type) {
return E2M1Ids[idx]; })
4886 llvm_unreachable(
"Invalid type for ConvertF4x2ToBF16x2Op");
4887 return llvm::Intrinsic::not_intrinsic;
4890 llvm::Value *extendedI16 =
4891 builder.CreateZExt(mt.
lookupValue(curOp.getSrc()),
4892 llvm::Type::getInt16Ty(builder.getContext()));
4895 args.push_back(extendedI16);
4902 return {intId, std::move(args)};
4907 auto thisOp = cast<NVVM::ConvertF32x2ToS2F6x2Op>(op);
4908 bool hasRelu = thisOp.getRelu();
4909 bool hasScale =
static_cast<bool>(thisOp.getScaleFactor());
4911 llvm::Intrinsic::ID
id =
4913 ? llvm::Intrinsic::nvvm_ff_to_s2f6x2_rn_relu_satfinite_scale_n2_ue8m0
4914 : llvm::Intrinsic::nvvm_ff_to_s2f6x2_rn_satfinite_scale_n2_ue8m0;
4920 args.push_back(hasScale ? mt.
lookupValue(thisOp.getScaleFactor())
4921 : builder.getInt16(0x7f7f));
4922 return {id, std::move(args)};
4927 auto thisOp = cast<NVVM::ConvertBF16x2ToS2F6x2Op>(op);
4928 bool hasRelu = thisOp.getRelu();
4929 bool hasScale =
static_cast<bool>(thisOp.getScaleFactor());
4931 llvm::Intrinsic::ID
id =
4934 nvvm_bf16x2_to_s2f6x2_rn_relu_satfinite_scale_n2_ue8m0
4935 : llvm::Intrinsic::nvvm_bf16x2_to_s2f6x2_rn_satfinite_scale_n2_ue8m0;
4940 args.push_back(hasScale ? mt.
lookupValue(thisOp.getScaleFactor())
4941 : builder.getInt16(0x7f7f));
4942 return {id, std::move(args)};
4947 auto thisOp = cast<NVVM::ConvertS2F6x2ToBF16x2Op>(op);
4948 bool hasRelu = thisOp.getRelu();
4949 bool hasScale =
static_cast<bool>(thisOp.getScaleFactor());
4950 bool hasSat = thisOp.getSat() == NVVM::SaturationMode::SATFINITE;
4952 static constexpr llvm::Intrinsic::ID ids[] = {
4953 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_scale_n2_ue8m0,
4954 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
4955 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
4956 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
4959 unsigned idx = (hasSat << 1) | hasRelu;
4963 llvm::Value *packedI16 =
4964 builder.CreateBitCast(mt.
lookupValue(thisOp.getSrc()),
4965 llvm::Type::getInt16Ty(builder.getContext()));
4966 args.push_back(packedI16);
4967 args.push_back(hasScale ? mt.
lookupValue(thisOp.getScaleFactor())
4968 : builder.getInt16(0x7f7f));
4970 return {ids[idx], std::move(args)};
4974Tcgen05AllocOp::getIntrinsicIDAndArgs(
Operation &op,
4977 auto curOp = cast<NVVM::Tcgen05AllocOp>(op);
4978 unsigned as = llvm::cast<LLVM::LLVMPointerType>(curOp.getAddr().getType())
4980 bool isShared = as == NVVMMemorySpace::Shared;
4981 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
4983 llvm::Intrinsic::ID id;
4985 id = is2CTAMode ? llvm::Intrinsic::nvvm_tcgen05_alloc_shared_cg2
4986 : llvm::Intrinsic::nvvm_tcgen05_alloc_shared_cg1;
4988 id = is2CTAMode ? llvm::Intrinsic::nvvm_tcgen05_alloc_cg2
4989 : llvm::Intrinsic::nvvm_tcgen05_alloc_cg1;
4999llvm::Intrinsic::ID Tcgen05DeallocOp::getIntrinsicIDAndArgs(
5002 auto curOp = cast<NVVM::Tcgen05DeallocOp>(op);
5003 auto id = (curOp.getGroup() == CTAGroupKind::CTA_1)
5004 ? llvm::Intrinsic::nvvm_tcgen05_dealloc_cg1
5005 : llvm::Intrinsic::nvvm_tcgen05_dealloc_cg2;
5015Tcgen05CommitOp::getIntrinsicIDAndArgs(
Operation &op,
5018 auto curOp = cast<NVVM::Tcgen05CommitOp>(op);
5019 bool hasMulticast =
static_cast<bool>(curOp.getMulticastMask());
5020 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
5021 bool hasSmemARead = curOp.getSmemARead();
5022 unsigned index = (
static_cast<unsigned>(hasSmemARead) << 1) |
5023 static_cast<unsigned>(is2CTAMode);
5025 using namespace llvm::Intrinsic;
5026 static constexpr ID IDs[] = {
5027 nvvm_tcgen05_commit_cg1,
5028 nvvm_tcgen05_commit_cg2,
5029 nvvm_tcgen05_commit_smem_a_read_cg1,
5030 nvvm_tcgen05_commit_smem_a_read_cg2,
5033 static constexpr ID multicastIDs[] = {
5034 nvvm_tcgen05_commit_mc_cg1,
5035 nvvm_tcgen05_commit_mc_cg2,
5036 nvvm_tcgen05_commit_smem_a_read_mc_cg1,
5037 nvvm_tcgen05_commit_smem_a_read_mc_cg2,
5040 ID
id = hasMulticast ? multicastIDs[
index] : IDs[
index];
5044 args.push_back(mt.
lookupValue(curOp.getMulticastMask()));
5049#define TCGEN05_CP_IMPL(shape_mc, src_fmt, cg) \
5050 llvm::Intrinsic::nvvm_tcgen05_cp##shape_mc##src_fmt##cg
5052#define TCGEN05_CP_2CTA(shape_mc, src_fmt, is_2cta) \
5053 is_2cta ? TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg2) \
5054 : TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg1)
5056#define GET_TCGEN05_CP_ID(shape_mc, src_fmt, is_2cta) \
5058 if ((src_fmt) == Tcgen05CpSrcFormat::B6x16_P32) \
5059 return TCGEN05_CP_2CTA(shape_mc, _b6x16_p32, is_2cta); \
5060 if ((src_fmt) == Tcgen05CpSrcFormat::B4x16_P64) \
5061 return TCGEN05_CP_2CTA(shape_mc, _b4x16_p64, is_2cta); \
5062 return TCGEN05_CP_2CTA(shape_mc, , is_2cta); \
5066ConvertF32x2ToF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF16x2Op &op,
5068 llvm::IRBuilderBase &builder) {
5069 static constexpr llvm::Intrinsic::ID rndRNIds[] = {
5070 llvm::Intrinsic::nvvm_ff2f16x2_rn,
5071 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu,
5072 llvm::Intrinsic::nvvm_ff2f16x2_rn_satfinite,
5073 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu_satfinite,
5075 static constexpr llvm::Intrinsic::ID rndRZIds[] = {
5076 llvm::Intrinsic::nvvm_ff2f16x2_rz,
5077 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu,
5078 llvm::Intrinsic::nvvm_ff2f16x2_rz_satfinite,
5079 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu_satfinite,
5081 static constexpr llvm::Intrinsic::ID rndRSIds[] = {
5082 llvm::Intrinsic::nvvm_ff2f16x2_rs,
5083 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu,
5084 llvm::Intrinsic::nvvm_ff2f16x2_rs_satfinite,
5085 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu_satfinite,
5088 unsigned hasRelu = op.getRelu() ? 1 : 0;
5089 unsigned hasSatFinite =
5090 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;
5093 unsigned idx = (hasSatFinite << 1) | hasRelu;
5098 if (op.getRandomBits())
5099 args.push_back(mt.
lookupValue(op.getRandomBits()));
5101 switch (op.getRnd()) {
5102 case FPRoundingMode::RN:
5103 return {rndRNIds[idx], std::move(args)};
5104 case FPRoundingMode::RZ:
5105 return {rndRZIds[idx], std::move(args)};
5106 case FPRoundingMode::RS:
5107 return {rndRSIds[idx], std::move(args)};
5109 llvm_unreachable(
"Invalid rounding mode for ConvertF32x2ToF16x2Op");
5114ConvertF32x2ToBF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToBF16x2Op &op,
5116 llvm::IRBuilderBase &builder) {
5117 static constexpr llvm::Intrinsic::ID rndRNIds[] = {
5118 llvm::Intrinsic::nvvm_ff2bf16x2_rn,
5119 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu,
5120 llvm::Intrinsic::nvvm_ff2bf16x2_rn_satfinite,
5121 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu_satfinite,
5123 static constexpr llvm::Intrinsic::ID rndRZIds[] = {
5124 llvm::Intrinsic::nvvm_ff2bf16x2_rz,
5125 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu,
5126 llvm::Intrinsic::nvvm_ff2bf16x2_rz_satfinite,
5127 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu_satfinite,
5129 static constexpr llvm::Intrinsic::ID rndRSIds[] = {
5130 llvm::Intrinsic::nvvm_ff2bf16x2_rs,
5131 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu,
5132 llvm::Intrinsic::nvvm_ff2bf16x2_rs_satfinite,
5133 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu_satfinite,
5136 unsigned hasRelu = op.getRelu() ? 1 : 0;
5137 unsigned hasSatFinite =
5138 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;
5141 unsigned idx = (hasSatFinite << 1) | hasRelu;
5146 if (op.getRandomBits())
5147 args.push_back(mt.
lookupValue(op.getRandomBits()));
5149 switch (op.getRnd()) {
5150 case FPRoundingMode::RN:
5151 return {rndRNIds[idx], std::move(args)};
5152 case FPRoundingMode::RZ:
5153 return {rndRZIds[idx], std::move(args)};
5154 case FPRoundingMode::RS:
5155 return {rndRSIds[idx], std::move(args)};
5157 llvm_unreachable(
"Invalid rounding mode for ConvertF32x2ToBF16x2Op");
5161llvm::Intrinsic::ID ConvertF32x4ToF8x4Op::getIntrinsicID() {
5163 bool hasRelu = getRelu();
5166 .Case([&](mlir::Float8E4M3FNType) {
5167 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_relu_satfinite
5168 : llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_satfinite;
5170 .Case([&](mlir::Float8E5M2Type) {
5171 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_relu_satfinite
5172 : llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_satfinite;
5175 llvm_unreachable(
"Invalid F8 type in ConvertF32x4ToF8x4Op");
5176 return llvm::Intrinsic::not_intrinsic;
5180llvm::Intrinsic::ID ConvertF32x4ToF6x4Op::getIntrinsicID() {
5182 bool hasRelu = getRelu();
5185 .Case([&](mlir::Float6E2M3FNType) {
5186 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_relu_satfinite
5187 : llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_satfinite;
5189 .Case([&](mlir::Float6E3M2FNType) {
5190 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_relu_satfinite
5191 : llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_satfinite;
5194 llvm_unreachable(
"Invalid F6 type in ConvertF32x4ToF6x4Op");
5195 return llvm::Intrinsic::not_intrinsic;
5199llvm::Intrinsic::ID ConvertF32x4ToF4x4Op::getIntrinsicID() {
5201 bool hasRelu = getRelu();
5204 .Case([&](mlir::Float4E2M1FNType) {
5205 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_relu_satfinite
5206 : llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_satfinite;
5209 llvm_unreachable(
"Invalid F4 type in ConvertF32x4ToF4x4Op");
5210 return llvm::Intrinsic::not_intrinsic;
5214llvm::Intrinsic::ID Tcgen05CpOp::getIntrinsicID(
Operation &op) {
5215 auto curOp = cast<NVVM::Tcgen05CpOp>(op);
5216 bool is2CTA = curOp.getGroup() == CTAGroupKind::CTA_2;
5217 auto srcFmt = curOp.getSrcFormat();
5218 auto mc = curOp.getMulticast();
5220 switch (curOp.getShape()) {
5221 case Tcgen05CpShape::SHAPE_128x256b:
5223 case Tcgen05CpShape::SHAPE_128x128b:
5225 case Tcgen05CpShape::SHAPE_4x256b:
5227 case Tcgen05CpShape::SHAPE_32x128b:
5229 case Tcgen05CpShape::SHAPE_64x128b:
5230 return (mc == Tcgen05CpMulticast::WARPX2_01_23)
5234 llvm_unreachable(
"Invalid shape in tcgen05 cp Op");
5241 if (
shape == NVVM::Tcgen05LdStShape::SHAPE_16X128B)
5243 if (
shape == NVVM::Tcgen05LdStShape::SHAPE_16X256B)
5248LogicalResult Tcgen05LdOp::verify() {
5250 if (
getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())
5253 if (
getShape() != NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && getOffset())
5254 result =
emitError(
"offset argument is only supported for shape 16x32bx2");
5256 auto resTy = getRes().getType();
5257 unsigned resLen = isa<VectorType>(resTy)
5258 ? llvm::cast<VectorType>(resTy).getNumElements()
5261 result =
emitError(llvm::formatv(
"invalid result type length {0} for shape "
5262 "{1} in tcgen05.ld Op",
5263 resLen, stringifyEnum(
getShape())));
5268LogicalResult Tcgen05StOp::verify() {
5270 if (
getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())
5273 auto valTy = getVal().getType();
5274 unsigned valLen = isa<VectorType>(valTy)
5275 ? llvm::cast<VectorType>(valTy).getNumElements()
5278 result =
emitError(llvm::formatv(
"invalid input length {0} for shape "
5279 "{1} in tcgen05.st Op",
5280 valLen, stringifyEnum(
getShape())));
5290 if (
auto rangeAttr = op->
getAttrOfType<LLVM::ConstantRangeAttr>(
"range")) {
5291 setResultRanges(
result, {rangeAttr.getLower(), rangeAttr.getUpper(),
5292 rangeAttr.getLower(), rangeAttr.getUpper()});
5302 std::optional<LLVM::ConstantRangeAttr> rangeAttr) {
5306 const llvm::APInt &lower = rangeAttr->getLower();
5307 const llvm::APInt &upper = rangeAttr->getUpper();
5310 if (lower == upper && !lower.isMaxValue() && !lower.isMinValue()) {
5311 unsigned bitWidth = lower.getBitWidth();
5312 llvm::APInt minVal = llvm::APInt::getMinValue(bitWidth);
5313 llvm::APInt maxVal = llvm::APInt::getMaxValue(bitWidth);
5315 "invalid range attribute: Lower == Upper, but they aren't min (")
5316 << llvm::toString(minVal, 10,
false) <<
") or max ("
5317 << llvm::toString(maxVal, 10,
false)
5318 <<
") value! This is an invalid constant range.";
5325 llvm::IRBuilderBase &builder) {
5326 return builder.CreateBitCast(arg,
5327 llvm::Type::getInt32Ty(builder.getContext()));
5332 auto curOp = cast<NVVM::DotAccumulate4WayOp>(op);
5339 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;
5340 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;
5341 unsigned type = (isASigned << 1) | isBSigned;
5342 const llvm::Intrinsic::ID ids[] = {
5343 llvm::Intrinsic::nvvm_idp4a_u_u,
5344 llvm::Intrinsic::nvvm_idp4a_u_s,
5345 llvm::Intrinsic::nvvm_idp4a_s_u,
5346 llvm::Intrinsic::nvvm_idp4a_s_s,
5348 return {ids[type], args};
5353 auto curOp = cast<NVVM::DotAccumulate2WayOp>(op);
5358 args.push_back(builder.getInt1(curOp.getBHi()));
5361 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;
5362 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;
5363 unsigned type = (isASigned << 1) | isBSigned;
5364 const llvm::Intrinsic::ID ids[] = {
5365 llvm::Intrinsic::nvvm_idp2a_u_u,
5366 llvm::Intrinsic::nvvm_idp2a_u_s,
5367 llvm::Intrinsic::nvvm_idp2a_s_u,
5368 llvm::Intrinsic::nvvm_idp2a_s_s,
5370 return {ids[type], args};
5374 llvm::IRBuilderBase &builder) {
5375 return builder.CreateAddrSpaceCast(
5376 addr, builder.getPtrTy(llvm::NVPTXAS::ADDRESS_SPACE_ENTRY_PARAM));
5380PrefetchOp::getIntrinsicIDAndArgs(NVVM::PrefetchOp &op,
5382 llvm::IRBuilderBase &builder) {
5383 using MemSpace = NVVM::NVVMMemorySpace;
5384 using CacheLevel = NVVM::PrefetchCacheLevel;
5386 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = op.getCacheLevel();
5387 std::optional<NVVM::CacheEvictionPriority> evictPriority =
5388 op.getEvictPriority();
5389 unsigned addressSpace =
5390 llvm::cast<LLVM::LLVMPointerType>(op.getAddr().getType())
5398 if (op.getTensormap())
5399 return {llvm::Intrinsic::nvvm_prefetch_tensormap, args};
5401 assert(cacheLevel &&
"expected cache level for non-tensormap prefetch");
5403 if (op.getUniform() && *cacheLevel == CacheLevel::L1)
5404 return {llvm::Intrinsic::nvvm_prefetchu_L1, args};
5406 if (evictPriority && *cacheLevel == CacheLevel::L2) {
5407 switch (*evictPriority) {
5408 case NVVM::CacheEvictionPriority::EvictLast:
5409 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_last, args};
5410 case NVVM::CacheEvictionPriority::EvictNormal:
5411 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_normal, args};
5413 llvm_unreachable(
"Invalid cache eviction priority");
5417 switch (
static_cast<MemSpace
>(addressSpace)) {
5418 case MemSpace::Generic:
5419 return *cacheLevel == CacheLevel::L1
5421 :
NVVM::
IDArgPair({llvm::Intrinsic::nvvm_prefetch_L2, args});
5422 case MemSpace::Global:
5423 return *cacheLevel == CacheLevel::L1
5425 {llvm::Intrinsic::nvvm_prefetch_global_L1, args})
5427 {llvm::Intrinsic::nvvm_prefetch_global_L2, args});
5428 case MemSpace::Local:
5429 return *cacheLevel == CacheLevel::L1
5431 {llvm::Intrinsic::nvvm_prefetch_local_L1, args})
5433 {llvm::Intrinsic::nvvm_prefetch_local_L2, args});
5435 llvm_unreachable(
"Invalid pointer address space");
5439bool NVVM::InlinePtxOp::getAsmValues(
5443 for (
auto arg : getReadWriteArgs())
5445 for (
auto arg : getResults())
5447 for (
auto arg : getReadOnlyArgs())
5454NVVM::IDArgPair ClusterLaunchControlTryCancelOp::getIntrinsicIDAndArgs(
5456 auto curOp = cast<NVVM::ClusterLaunchControlTryCancelOp>(op);
5458 args.push_back(mt.
lookupValue(curOp.getSmemAddress()));
5459 args.push_back(mt.
lookupValue(curOp.getMbarrier()));
5461 llvm::Intrinsic::ID intrinsicID =
5462 curOp.getMulticast()
5464 nvvm_clusterlaunchcontrol_try_cancel_async_multicast_shared
5465 : llvm::Intrinsic::nvvm_clusterlaunchcontrol_try_cancel_async_shared;
5467 return {intrinsicID, args};
5470NVVM::IDArgPair ClusterLaunchControlQueryCancelOp::getIntrinsicIDAndArgs(
5472 auto curOp = cast<NVVM::ClusterLaunchControlQueryCancelOp>(op);
5474 args.push_back(mt.
lookupValue(curOp.getTryCancelResponse()));
5476 llvm::Intrinsic::ID intrinsicID;
5478 switch (curOp.getQueryType()) {
5479 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:
5481 llvm::Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_is_canceled;
5483 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:
5484 intrinsicID = llvm::Intrinsic::
5485 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_x;
5487 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:
5488 intrinsicID = llvm::Intrinsic::
5489 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_y;
5491 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:
5492 intrinsicID = llvm::Intrinsic::
5493 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_z;
5496 return {intrinsicID, args};
5501 llvm::IRBuilderBase &builder) {
5502 auto thisOp = cast<NVVM::PermuteOp>(op);
5503 NVVM::PermuteMode mode = thisOp.getMode();
5505 static constexpr llvm::Intrinsic::ID IDs[] = {
5506 llvm::Intrinsic::nvvm_prmt, llvm::Intrinsic::nvvm_prmt_f4e,
5507 llvm::Intrinsic::nvvm_prmt_b4e, llvm::Intrinsic::nvvm_prmt_rc8,
5508 llvm::Intrinsic::nvvm_prmt_ecl, llvm::Intrinsic::nvvm_prmt_ecr,
5509 llvm::Intrinsic::nvvm_prmt_rc16};
5511 unsigned modeIndex =
static_cast<unsigned>(mode);
5519 args.push_back(mt.
lookupValue(thisOp.getSelector()));
5521 return {IDs[modeIndex], args};
5526 auto thisOp = cast<NVVM::TensormapReplaceOp>(op);
5530 if (thisOp.getOrd())
5531 args.push_back(builder.getInt32(thisOp.getOrd().value()));
5532 if (thisOp.getNewValue())
5533 args.push_back(mt.
lookupValue(thisOp.getNewValue()));
5534 if (
auto attr = thisOp.getNewValueAttr()) {
5537 .Case<TensormapElemtypeAttr, TensormapInterleaveLayoutAttr,
5538 TensormapSwizzleModeAttr, TensormapSwizzleAtomicityAttr,
5539 TensormapFillModeAttr>([](
auto attr) {
5540 return static_cast<unsigned>(attr.getValue());
5542 .Default([](
auto attr) {
5543 llvm_unreachable(
"Invalid attribute type");
5546 args.push_back(builder.getInt32(val));
5549 static constexpr llvm::Intrinsic::ID IDs[] = {
5550 llvm::Intrinsic::nvvm_tensormap_replace_global_address,
5551 llvm::Intrinsic::nvvm_tensormap_replace_rank,
5552 llvm::Intrinsic::nvvm_tensormap_replace_box_dim,
5553 llvm::Intrinsic::nvvm_tensormap_replace_global_dim,
5554 llvm::Intrinsic::nvvm_tensormap_replace_global_stride,
5555 llvm::Intrinsic::nvvm_tensormap_replace_element_stride,
5556 llvm::Intrinsic::nvvm_tensormap_replace_elemtype,
5557 llvm::Intrinsic::nvvm_tensormap_replace_interleave_layout,
5558 llvm::Intrinsic::nvvm_tensormap_replace_swizzle_mode,
5559 llvm::Intrinsic::nvvm_tensormap_replace_swizzle_atomicity,
5560 llvm::Intrinsic::nvvm_tensormap_replace_fill_mode,
5563 unsigned fieldIndex =
static_cast<unsigned>(thisOp.getField());
5565 return {IDs[fieldIndex], args};
5574 llvm::IRBuilderBase &builder) {
5576 auto thisOp = cast<NVVM::Tcgen05MMAOp>(op);
5579 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
5582 const bool isATensor = isa<llvm::PointerType>(
A->getType());
5585 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
5586 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
5587 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
5589 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;
5590 using CtaGroupArray = std::array<EnableAShiftArray, 2>;
5591 using IsATensorArray = std::array<CtaGroupArray, 2>;
5592 using HasScaleInputDArray = std::array<IsATensorArray, 2>;
5593 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;
5596 static constexpr HasDisableOutputLaneArray tcgen05MMAIDs = {
5602 {llvm::Intrinsic::nvvm_tcgen05_mma_shared,
notIntrinsic},
5604 {llvm::Intrinsic::nvvm_tcgen05_mma_shared,
notIntrinsic}}},
5608 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,
5609 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,
5613 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,
5614 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,
5620 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d,
notIntrinsic},
5622 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d,
notIntrinsic}}},
5626 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,
5627 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,
5631 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,
5632 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,
5638 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1,
5641 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2,
5646 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1,
5648 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift,
5653 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2,
5655 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift,
5661 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1,
5665 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2,
5670 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1,
5672 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift},
5676 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2,
5678 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift,
5681 llvm::Value *ScaleInputD = mt.
lookupValue(thisOp.getScaleInputD());
5682 bool hasScaleInputD = ScaleInputD !=
nullptr;
5684 llvm::Value *DisableOutputLane =
5686 bool hasDisableOutputLane = DisableOutputLane !=
nullptr;
5688 const unsigned ctaGroup =
5691 llvm::Intrinsic::ID ID =
5692 tcgen05MMAIDs[hasDisableOutputLane][hasScaleInputD][isATensor]
5693 [ctaGroup - 1][thisOp.getAShift()];
5695 assert(ID !=
notIntrinsic &&
"Invalid intrinsic for Tcgen05MMAOp.");
5698 args.push_back(ScaleInputD);
5700 if (hasDisableOutputLane)
5701 args.push_back(DisableOutputLane);
5703 args.push_back(builder.getInt32(
static_cast<unsigned>(thisOp.getKind())));
5705 if (!hasDisableOutputLane)
5706 args.push_back(builder.getInt32(ctaGroup));
5709 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
5712 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
5719 NVVM::CTAGroupKind ctaGroup,
bool hasAShift,
5720 NVVM::Tcgen05MMACollectorOp collectorOp,
Location loc) {
5722 if (disableOutputLane) {
5723 mlir::VectorType disableOutputLaneType =
5724 cast<mlir::VectorType>(disableOutputLane.
getType());
5725 if ((ctaGroup == NVVM::CTAGroupKind::CTA_1 &&
5726 disableOutputLaneType.getNumElements() != 4) ||
5727 (ctaGroup == NVVM::CTAGroupKind::CTA_2 &&
5728 disableOutputLaneType.getNumElements() != 8))
5729 return emitError(loc) <<
"Disable Output Lane of length "
5730 << disableOutputLaneType.getNumElements()
5731 <<
" is incompatible with CtaGroupAttr";
5734 if (hasAShift && !isATensor)
5736 loc,
"A-shift can be applied only when matrix A is in tensor memory");
5738 if (hasAShift ==
true && (collectorOp == Tcgen05MMACollectorOp::FILL ||
5739 collectorOp == Tcgen05MMACollectorOp::USE))
5741 loc,
"Cannot use collector buffer operation fill or use with ashift");
5746LogicalResult Tcgen05MMAOp::verify() {
5748 getDisableOutputLane(), getCtaGroup(), getAShift(),
5749 getCollectorOp(), getLoc());
5759 auto thisOp = cast<NVVM::Tcgen05MMASparseOp>(op);
5762 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
5765 bool isATensor = isa<llvm::PointerType>(
A->getType());
5768 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
5769 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
5770 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
5771 args.push_back(mt.
lookupValue(thisOp.getSparseMetadata()));
5773 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;
5774 using CtaGroupArray = std::array<EnableAShiftArray, 2>;
5775 using IsATensorArray = std::array<CtaGroupArray, 2>;
5776 using HasScaleInputDArray = std::array<IsATensorArray, 2>;
5777 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;
5780 static constexpr HasDisableOutputLaneArray tcgen05MMASparseIDs = {
5786 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared,
notIntrinsic},
5788 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared,
notIntrinsic}}},
5792 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,
5793 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,
5797 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,
5798 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,
5804 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,
5807 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,
5812 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,
5813 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,
5817 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,
5818 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,
5825 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1,
5829 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2,
5834 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1,
5836 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift,
5841 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2,
5843 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift,
5849 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1,
5853 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2,
5858 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1,
5860 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift},
5864 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2,
5866 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift,
5869 llvm::Value *ScaleInputD = mt.
lookupValue(thisOp.getScaleInputD());
5870 bool hasScaleInputD = ScaleInputD !=
nullptr;
5872 llvm::Value *DisableOutputLane =
5874 bool hasDisableOutputLane = DisableOutputLane !=
nullptr;
5879 llvm::Intrinsic::ID ID =
5880 tcgen05MMASparseIDs[hasDisableOutputLane][hasScaleInputD][isATensor]
5881 [ctaGroup - 1][thisOp.getAShift()];
5883 assert(ID !=
notIntrinsic &&
"Invalid intrinsic for Tcgen05MMASparseOp.");
5886 args.push_back(ScaleInputD);
5888 if (hasDisableOutputLane)
5889 args.push_back(DisableOutputLane);
5891 args.push_back(builder.getInt32(
static_cast<unsigned>(thisOp.getKind())));
5893 if (!hasDisableOutputLane)
5894 args.push_back(builder.getInt32(ctaGroup));
5897 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
5900 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
5905LogicalResult Tcgen05MMASparseOp::verify() {
5907 getDisableOutputLane(), getCtaGroup(), getAShift(),
5908 getCollectorOp(), getLoc());
5918 auto thisOp = cast<NVVM::Tcgen05MMABlockScaleOp>(op);
5921 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
5924 bool isATensor = isa<llvm::PointerType>(
A->getType());
5927 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
5928 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
5929 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
5930 args.push_back(mt.
lookupValue(thisOp.getScaleA()));
5931 args.push_back(mt.
lookupValue(thisOp.getScaleB()));
5932 args.push_back(builder.getInt32(
5935 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
5937 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
5939 auto kind = thisOp.getKind();
5940 auto blockScale = thisOp.getBlockScale();
5941 llvm::Intrinsic::ID ID = [&]() {
5942 if (kind == NVVM::Tcgen05MMAKind::MXF8F6F4) {
5943 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
5944 return isATensor ? llvm::Intrinsic::
5945 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale
5947 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale;
5948 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
5951 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale_block32
5953 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale_block32;
5955 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4) {
5956 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
5958 ? llvm::Intrinsic::nvvm_tcgen05_mma_tensor_mxf4_block_scale
5959 : llvm::Intrinsic::nvvm_tcgen05_mma_shared_mxf4_block_scale;
5960 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
5961 return isATensor ? llvm::Intrinsic::
5962 nvvm_tcgen05_mma_tensor_mxf4_block_scale_block32
5964 nvvm_tcgen05_mma_shared_mxf4_block_scale_block32;
5966 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4NVF4) {
5967 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
5970 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block32
5972 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block32;
5974 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {
5977 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block16
5979 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block16;
5982 llvm_unreachable(
"Invalid tcgen05.mma.block_scale attributes");
5989 NVVM::Tcgen05MMACollectorOp collectorOp, NVVM::Tcgen05MMAKind kind,
5990 NVVM::Tcgen05MMABlockScale blockScale,
Location loc) {
5991 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT &&
5992 kind == NVVM::Tcgen05MMAKind::MXF4NVF4)
5993 return emitError(loc,
"mxf4nvf4 requires block scale attribute");
5995 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16 &&
5996 kind != NVVM::Tcgen05MMAKind::MXF4NVF4)
5998 llvm::formatv(
"{} kind does not support block16 attribute",
5999 stringifyEnum(kind)));
6004LogicalResult Tcgen05MMABlockScaleOp::verify() {
6006 getBlockScale(), getLoc());
6016 auto thisOp = cast<NVVM::Tcgen05MMASparseBlockScaleOp>(op);
6019 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6022 bool isATensor = isa<llvm::PointerType>(
A->getType());
6025 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6026 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6027 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6028 args.push_back(mt.
lookupValue(thisOp.getSparseMetadata()));
6029 args.push_back(mt.
lookupValue(thisOp.getScaleA()));
6030 args.push_back(mt.
lookupValue(thisOp.getScaleB()));
6031 args.push_back(builder.getInt32(
6034 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6036 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6038 auto kind = thisOp.getKind();
6039 auto blockScale = thisOp.getBlockScale();
6040 llvm::Intrinsic::ID ID = [&]() {
6041 if (kind == NVVM::Tcgen05MMAKind::MXF8F6F4) {
6042 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6043 return isATensor ? llvm::Intrinsic::
6044 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale
6046 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale;
6047 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6050 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale_block32
6052 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale_block32;
6054 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4) {
6055 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6056 return isATensor ? llvm::Intrinsic::
6057 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale
6059 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale;
6060 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6063 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale_block32
6065 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale_block32;
6067 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4NVF4) {
6068 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6071 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block32
6073 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block32;
6075 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {
6078 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block16
6080 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block16;
6083 llvm_unreachable(
"Invalid tcgen05.mma.sp.block_scale attributes");
6089LogicalResult Tcgen05MMASparseBlockScaleOp::verify() {
6091 getBlockScale(), getLoc());
6101 auto thisOp = cast<NVVM::Tcgen05MMAWsOp>(op);
6104 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6107 bool isATensor = isa<llvm::PointerType>(
A->getType());
6110 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6111 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6112 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6114 mlir::Value ZeroColMask = thisOp.getZeroColMask();
6118 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor_zero_col_mask
6119 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared_zero_col_mask;
6121 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor
6122 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared;
6124 args.push_back(builder.getInt32(
static_cast<unsigned>(thisOp.getKind())));
6126 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorBBuffer())));
6128 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6140 auto thisOp = cast<NVVM::Tcgen05MMAWsSparseOp>(op);
6143 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6146 bool isATensor = isa<llvm::PointerType>(
A->getType());
6149 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6150 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6151 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6152 args.push_back(mt.
lookupValue(thisOp.getSparseMetadata()));
6154 mlir::Value ZeroColMask = thisOp.getZeroColMask();
6159 ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor_zero_col_mask
6160 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared_zero_col_mask;
6162 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor
6163 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared;
6165 args.push_back(builder.getInt32(
static_cast<unsigned>(thisOp.getKind())));
6167 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorBBuffer())));
6169 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6178#define TCGEN05LDRED(SHAPE, NUM, TYPE) \
6179 llvm::Intrinsic::nvvm_tcgen05_ld_red_##SHAPE##_##NUM##_##TYPE
6183 auto thisOp = cast<NVVM::Tcgen05LdRedOp>(op);
6186 mlir::VectorType VecResTy =
6187 cast<mlir::VectorType>(thisOp.getData().getType());
6188 unsigned Num = VecResTy.getNumElements();
6189 bool IsFloat = thisOp.getRedVal().getType().isF32();
6191 llvm::Intrinsic::ID Shape32x32b[][2] = {
6202 llvm::Intrinsic::ID Shape16x32bx2[][2] = {
6213 NVVM::Tcgen05LdStShape
shape = thisOp.getShape();
6214 unsigned ID = [&]() {
6217 unsigned idx = std::log2(Num);
6219 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
6220 return Shape32x32b[idx][IsFloat];
6221 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
6222 return Shape16x32bx2[idx][IsFloat];
6224 llvm_unreachable(
"unhandled tcgen05.ld lowering");
6230 if (
shape == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2)
6231 args.push_back(mt.
lookupValue(thisOp.getOffset()));
6234 builder.getInt32(thisOp.getOp() == NVVM::ReductionKind::MIN ? 0 : 1));
6237 args.push_back(builder.getInt1(
static_cast<unsigned>(thisOp.getAbs())));
6238 args.push_back(builder.getInt1(
static_cast<unsigned>(thisOp.getNan())));
6243LogicalResult Tcgen05LdRedOp::verify() {
6244 VectorType data = cast<VectorType>(getData().
getType());
6245 Type redVal = getRedVal().getType();
6247 if (data.getElementType() != redVal)
6249 "type of reduction value and element type of vector data should match");
6251 if (getOp() != NVVM::ReductionKind::MIN &&
6252 getOp() != NVVM::ReductionKind::MAX)
6253 return emitError(
"only min and max reduction kinds are supported");
6255 if (redVal.
isInteger() && (getAbs() || getNan())) {
6256 return emitError(
"abs or nan is only applicable for f32 type");
6266struct NVVMInlinerInterface final : DialectInlinerInterface {
6267 using DialectInlinerInterface::DialectInlinerInterface;
6268 bool isLegalToInline(Operation *, Region *,
bool, IRMapping &)
const final {
6275void NVVMDialect::initialize() {
6278#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"
6281#define GET_ATTRDEF_LIST
6282#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"
6287 allowUnknownOperations();
6288 addInterfaces<NVVMInlinerInterface>();
6289 declarePromisedInterface<ConvertToLLVMPatternInterface, NVVMDialect>();
6290 declarePromisedInterface<gpu::TargetAttrInterface, NVVMTargetAttr>();
6293LogicalResult NVVMDialect::verifyOperationAttribute(
Operation *op,
6295 StringAttr attrName = attr.
getName();
6297 if (attrName == NVVMDialect::getKernelFuncAttrName()) {
6298 if (!isa<LLVM::LLVMFuncOp>(op)) {
6299 return op->
emitError() <<
"'" << NVVMDialect::getKernelFuncAttrName()
6300 <<
"' attribute attached to unexpected op";
6305 if (attrName == NVVMDialect::getMaxntidAttrName() ||
6306 attrName == NVVMDialect::getReqntidAttrName() ||
6307 attrName == NVVMDialect::getClusterDimAttrName()) {
6308 auto values = llvm::dyn_cast<DenseI32ArrayAttr>(attr.
getValue());
6309 if (!values || values.empty() || values.size() > 3) {
6312 <<
"' attribute must be integer array with maximum 3 index";
6317 if (attrName == NVVMDialect::getMinctasmAttrName() ||
6318 attrName == NVVMDialect::getMaxnregAttrName() ||
6319 attrName == NVVMDialect::getClusterMaxBlocksAttrName()) {
6320 if (!llvm::dyn_cast<IntegerAttr>(attr.
getValue())) {
6322 <<
"'" << attrName <<
"' attribute must be integer constant";
6326 if (attrName == NVVMDialect::getBlocksAreClustersAttrName()) {
6327 if (!op->
hasAttr(NVVMDialect::getReqntidAttrName()) ||
6328 !op->
hasAttr(NVVMDialect::getClusterDimAttrName())) {
6330 <<
"'" << attrName <<
"' attribute must be used along with " <<
"'"
6331 << NVVMDialect::getReqntidAttrName() <<
"' and " <<
"'"
6332 << NVVMDialect::getClusterDimAttrName() <<
"'";
6339LogicalResult NVVMDialect::verifyRegionArgAttribute(
Operation *op,
6340 unsigned regionIndex,
6343 auto funcOp = dyn_cast<FunctionOpInterface>(op);
6347 bool isKernel = op->
hasAttr(NVVMDialect::getKernelFuncAttrName());
6348 StringAttr attrName = argAttr.
getName();
6349 if (attrName == NVVM::NVVMDialect::getGridConstantAttrName()) {
6353 <<
"' attribute must be present only on kernel arguments";
6355 if (!isa<UnitAttr>(argAttr.
getValue()))
6356 return op->
emitError() <<
"'" << attrName <<
"' must be a unit attribute";
6357 if (!funcOp.getArgAttr(argIndex, LLVM::LLVMDialect::getByValAttrName())) {
6360 <<
"' attribute requires the argument to also have attribute '"
6361 << LLVM::LLVMDialect::getByValAttrName() <<
"'";
6372unsigned NVVMMemorySpaceAttr::getAddressSpace()
const {
6373 return static_cast<unsigned>(getValue());
6376bool NVVMMemorySpaceAttr::isValidLoad(
6377 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,
6378 const ::mlir::DataLayout *dataLayout,
6384bool NVVMMemorySpaceAttr::isValidStore(
6385 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,
6386 const ::mlir::DataLayout *dataLayout,
6392bool NVVMMemorySpaceAttr::isValidAtomicOp(
6393 ptr::AtomicBinOp op,
Type type, ptr::AtomicOrdering ordering,
6394 std::optional<int64_t> alignment, const ::mlir::DataLayout *dataLayout,
6397 assert(
false &&
"unimplemented, see TODO in the source.");
6401bool NVVMMemorySpaceAttr::isValidAtomicXchg(
6402 Type type, ptr::AtomicOrdering successOrdering,
6403 ptr::AtomicOrdering failureOrdering, std::optional<int64_t> alignment,
6404 const ::mlir::DataLayout *dataLayout,
6407 assert(
false &&
"unimplemented, see TODO in the source.");
6411bool NVVMMemorySpaceAttr::isValidAddrSpaceCast(
6415 assert(
false &&
"unimplemented, see TODO in the source.");
6419bool NVVMMemorySpaceAttr::isValidPtrIntCast(
6424 assert(
false &&
"unimplemented, see TODO in the source.");
6433 int optLevel, StringRef triple, StringRef chip,
6434 StringRef features, DictionaryAttr flags,
6436 if (optLevel < 0 || optLevel > 3) {
6437 emitError() <<
"The optimization level must be a number between 0 and 3.";
6440 if (triple.empty()) {
6441 emitError() <<
"The target triple cannot be empty.";
6445 emitError() <<
"The target chip cannot be empty.";
6448 if (files && !llvm::all_of(files, [](::mlir::Attribute attr) {
6449 return mlir::isa_and_nonnull<StringAttr>(attr);
6451 emitError() <<
"All the elements in the `link` array must be strings.";
6457LogicalResult NVVMTargetAttr::verifyTarget(
Operation *gpuModule) {
6458 if (!getVerifyTarget())
6461 auto gpuModuleOp = llvm::dyn_cast<gpu::GPUModuleOp>(gpuModule);
6464 "NVVM target attribute must be attached to a GPU module");
6467 const unsigned targetFullSmVersion =
6471 "Minimum NVVM target SM version is sm_20");
6475 ->
walk([&](Operation *op) {
6476 if (
auto reqOp = llvm::dyn_cast<NVVM::RequiresSMInterface>(op)) {
6477 const NVVMCheckSMVersion requirement =
6478 reqOp.getRequiredMinSMVersion();
6480 op->
emitOpError() <<
"is not supported on " << getChip();
6492#define GET_OP_CLASSES
6493#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"
6495#define GET_ATTRDEF_CLASSES
6496#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
#define GET_TCGEN05_CP_ID(shape_mc, src_fmt, is_2cta)
static LogicalResult verifyTMALoadParams(size_t tensorDims, size_t numIm2colOff, TMALoadMode mode, Location loc)
static LogicalResult verifyTcgen05MMAOp(bool isATensor, mlir::Value disableOutputLane, NVVM::CTAGroupKind ctaGroup, bool hasAShift, NVVM::Tcgen05MMACollectorOp collectorOp, Location loc)
static bool isPtrInAddrSpace(mlir::Value ptr, NVVMMemorySpace targetAS)
static bool isCompatibleReturnTypesOptionalResult(TypeRange inferred, TypeRange actual)
For ops with optional results, allow the user to omit the result even when inference would produce on...
static bool isPtrInSharedCTASpace(mlir::Value ptr)
static LogicalResult isAllowedSizeN(int sizeN, NVVM::WGMMATypes typeA)
static llvm::nvvm::CTAGroupKind getNVVMCtaGroupKind(NVVM::CTAGroupKind ctaGroup)
static void addInferredMultiplicandTypes(MLIRContext *ctx, OperationState &result, ValueRange operandA, ValueRange operandB, std::optional< std::array< MMATypes, 2 > > multiplicandPtxTypes)
#define GET_CVT_F2TF32_ID(rnd, relu, sf)
static void addBlockScaleProperties(OpBuilder &builder, OperationState &result, ArrayRef< int64_t > shape, ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat, MMABlockScaleKind kind)
#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf)
static llvm::Value * getParamCastedAddr(llvm::Value *addr, llvm::IRBuilderBase &builder)
static LogicalResult verifyAddSubFOp(OpType op)
static LogicalResult verifyTcgen05MMABlockScaleOp(NVVM::Tcgen05MMACollectorOp collectorOp, NVVM::Tcgen05MMAKind kind, NVVM::Tcgen05MMABlockScale blockScale, Location loc)
static llvm::Value * packValInto64Bits(llvm::IRBuilderBase &builder, llvm::Value *result, llvm::Value *field, unsigned sizeInBits, unsigned start)
Packs the given field into the result.
static void printOperandList(OpAsmPrinter &p, StringRef name, ArrayRef< Value > operands)
#define GET_F32x2_TO_F6x2_ID(type, has_relu)
static llvm::Value * getAsPackedI32(llvm::Value *arg, llvm::IRBuilderBase &builder)
#define GET_F16x2_TO_F8X2_ID(type, has_relu)
static LogicalResult verifyMBarrierArriveLikeOp(Operation *op, Value addr, NVVM::MemScopeKind scope, Value retVal=nullptr)
static llvm::Value * castPtrToAddrSpace(llvm::IRBuilderBase &builder, llvm::Value *ptr, NVVMMemorySpace targetAS)
static LogicalResult isAllowedWGMMADataType(NVVM::WGMMATypes typeD, NVVM::WGMMATypes typeA, NVVM::WGMMATypes typeB)
static llvm::Intrinsic::ID getBarrierReductionIntrinsic(bool aligned, NVVM::BarrierReduction kind)
Maps the (aligned, kind) pair to the @llvm.nvvm.barrier.cta.red.
static void inferAndSetMultiplicandTypes(MLIRContext *ctx, NamedAttrList &attrs, const SmallVectorImpl< Type > &operandTypes)
static LogicalResult parseMmaOperand(OpAsmParser &parser, StringRef operandName, SmallVectorImpl< OpAsmParser::UnresolvedOperand > ®s)
static std::pair< mlir::Type, unsigned > inferMMATypeFromMNK(NVVM::MMATypes type, NVVM::MMAFrag frag, int m, int n, int k, MLIRContext *context)
static bool isInt8PtxType(MMATypes type)
#define TCGEN05LDRED(SHAPE, NUM, TYPE)
static bool isInt4PtxType(MMATypes type)
static bool isIntegerPtxType(MMATypes type)
#define GET_F32x2_TO_F8X2_S_ID(type, has_relu)
static MMATypes inferPtxTypeFromResult(OpTy op)
static LogicalResult verifyConstantRangeAttr(Operation *op, std::optional< LLVM::ConstantRangeAttr > rangeAttr)
Verify the range attribute satisfies LLVM ConstantRange constructor requirements for NVVM SpecialRang...
static LogicalResult parseMmaTypeSignature(OpAsmParser &parser, SmallVectorImpl< Type > &operandTypes)
static FailureOr< int > getAllowedSizeK(NVVM::WGMMATypes typeA)
static bool isPtrInSharedClusterSpace(mlir::Value ptr)
#define GET_CP_ASYNC_ID(mod, size, has_cpsize)
static unsigned isValidVectorLength(NVVM::Tcgen05LdStShape shape, unsigned vecLen)
static LogicalResult verifyConvertF32x2ToFP16x2Op(Twine dstType, FPRoundingMode rnd, bool hasRandomBits, Operation *op)
static void nvvmInferResultRanges(Operation *op, Value result, ArrayRef<::mlir::ConstantIntRanges > argRanges, SetIntRangeFn setResultRanges)
Infer the result ranges for the NVVM SpecialRangeableRegisterOp that might have ConstantRangeAttr.
static LogicalResult cpAsyncBulkTensorCommonVerifier(size_t tensorDims, bool isIm2Col, size_t numIm2ColOffsets, Location loc)
static bool isPtrInGenericSpace(mlir::Value ptr)
static void processOperandFragments(Op &op, std::array< MMAOperandFragment, 3 > &frags, SmallVectorImpl< Type > ®Types, SmallVectorImpl< StringRef > &ignoreAttrNames)
static llvm::Intrinsic::ID getBarrierSyncIntrinsic(bool aligned, bool hasCount)
Maps the (aligned, hasCount) pair to the @llvm.nvvm.barrier.cta.sync.
static constexpr unsigned notIntrinsic
static LogicalResult inferMBarrierArriveResultTypes(MLIRContext *context, Value addr, SmallVectorImpl< Type > &inferredReturnTypes)
Only shared_cluster (ptr<7>) produces zero results; all other address spaces (including generic) retu...
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
@ OptionalSquare
Square brackets supporting zero or more ops, or nothing.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseArrow()=0
Parse a '->' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an arrow followed by a type list.
ParseResult parseTypeList(SmallVectorImpl< Type > &result)
Parse a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
void printArrowTypeList(TypeRange &&types)
This class is a general helper class for creating context-global objects like types,...
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
IntegerType getIntegerType(unsigned width)
MLIRContext * getContext() const
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
This class represents a diagnostic that is inflight and set to be reported.
static IntegerValueRange getMaxRange(Value value)
Create a maximal range ([0, uint_max(t)] / [int_min(t), int_max(t)]) range that is used to mark the v...
Implementation class for module translation.
llvm::Value * lookupValue(Value value) const
Finds an LLVM IR value corresponding to the given MLIR value.
void mapValue(Value mlir, llvm::Value *llvm)
Stores the mapping between an MLIR value and its LLVM IR counterpart.
llvm::LLVMContext & getLLVMContext() const
Returns the LLVM context in which the IR is being constructed.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
std::optional< NamedAttribute > getNamed(StringRef name) const
Return the specified named attribute if present, std::nullopt otherwise.
Attribute get(StringAttr name) const
Return the specified attribute if present, null otherwise.
Attribute set(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
NamedAttribute represents a combination of a name and an Attribute value.
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
void printOperands(const ContainerType &container)
Print a comma separated list of operands.
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
This class helps build Operations.
This provides public APIs that all operations should have.
Operation is the basic unit of execution within MLIR.
AttrClass getAttrOfType(StringAttr name)
bool hasAttr(StringAttr name)
Return true if the operation has an attribute with the provided name, false otherwise.
Location getLoc()
The source location the operation was defined or derived from.
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isInteger() const
Return true if this is an integer type (with the specified width).
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
static WalkResult advance()
static WalkResult interrupt()
bool isValidLoadStoreImpl(Type type, ptr::AtomicOrdering ordering, std::optional< int64_t > alignment, const ::mlir::DataLayout *dataLayout, function_ref< InFlightDiagnostic()> emitError)
Checks whether the given type is an LLVM type that can be loaded or stored.
SmallVector< int64_t, 4 > getCoordinates(ArrayRef< int64_t > basis, unsigned linearIndex)
@ Write
Write register with '=' modifier.
@ ReadWrite
ReadWrite register with '+' modifier.
@ Read
Read register with no modifier.
std::pair< mlir::Type, unsigned > inferMMAType(mlir::NVVM::MMATypes type, mlir::NVVM::MMAFrag frag, int nRow, int nCol, mlir::MLIRContext *context)
Return the element type and number of elements associated with a wmma matrix of given chracteristics.
std::pair< llvm::Intrinsic::ID, llvm::SmallVector< llvm::Value * > > IDArgPair
A pair type of LLVM's Intrinsic ID and args (which are llvm values).
void walk(Operation *op, function_ref< void(Region *)> callback, WalkOrder order)
Walk all of the regions, blocks, or operations nested under (and including) the given operation.
uint64_t getN(LevelType lt)
uint64_t getM(LevelType lt)
Include the generated interface declarations.
llvm::function_ref< void(Value, const ConstantIntRanges &)> SetIntRangeFn
The type of the setResultRanges callback provided to ops implementing InferIntRangeInterface.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
llvm::function_ref< Fn > function_ref
LogicalResult matchAndRewrite(SubFOp op, PatternRewriter &rewriter) const override
static bool isMinimumSMVersion(unsigned fullSmVersion)
static unsigned getTargetFullSmVersionFromStr(StringRef smVersionString)
bool isCompatibleWith(const unsigned &targetFullSmVersion) const
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.