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");
95 NVVM::CTAGroupKindAttr &groupAttr) {
99 std::optional<NVVM::CTAGroupKind> group =
100 NVVM::symbolizeCTAGroupKind(keyword);
103 groupAttr = NVVM::CTAGroupKindAttr::get(parser.
getContext(), *group);
108 NVVM::CTAGroupKindAttr groupAttr) {
109 printer << NVVM::stringifyCTAGroupKind(groupAttr.getValue());
112template <
typename AttrTy>
119 using EnumTy =
decltype(attr.getValue());
120 std::optional<EnumTy> value = NVVM::symbolizeEnum<EnumTy>(keyword);
122 return parser.
emitError(loc) <<
"unknown enum value '" << keyword <<
"'";
124 attr = AttrTy::get(parser.
getContext(), *value);
137 size_t numIm2ColOffsets,
139 if (tensorDims < 1 || tensorDims > 5)
140 return emitError(loc,
"expects coordinates between 1 to 5 dimension");
148 "to use im2col mode, the tensor has to be at least 3-dimensional");
150 if (numIm2ColOffsets && (tensorDims != (numIm2ColOffsets + 2)))
152 loc,
"im2col offsets must be 2 less than number of coordinates");
161 if (!tensorSize.empty() && coordinates.size() != tensorSize.size()) {
163 emitError(loc,
"Expected coordinates size to be equal to tensor size");
166 if (!lowerStride.empty() && tensorSize.empty()) {
169 "Expected tensor_size to be present when lower_stride is provided");
170 }
else if (!lowerStride.empty() &&
171 lowerStride.size() != tensorSize.size() - 1) {
174 "Expected lower_stride size to be equal to one less than tensor size");
177 if (!lowerStride.empty() !=
static_cast<bool>(upperStride)) {
179 "Expected lower_stride and upper_stride to be either both "
180 "present or both absent");
183 bool isDimStride = tensorSize.size() > 0;
184 if (!isTile && isDimStride) {
186 loc,
"Only tile mode supports override address with dim and stride");
192LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOp::verify() {
193 TMAStoreMode mode = getMode();
197 if (getPredicate()) {
198 if (mode != TMAStoreMode::TILE)
199 return emitError(
"Inline-ptx lowering supported only for Tile mode.");
200 if (getL2CacheHint())
201 return emitError(
"Inline-ptx lowering unsupported with L2 cache-hint.");
206 case TMAStoreMode::TILE:
208 case TMAStoreMode::IM2COL:
209 case TMAStoreMode::IM2COL_W:
211 case TMAStoreMode::TILE_SCATTER4:
213 return emitError(
"Scatter4 mode expects 5 coordinates");
218LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp::verify() {
219 TMAStoreMode mode = getMode();
221 mode == TMAStoreMode::IM2COL || mode == TMAStoreMode::IM2COL_W;
222 bool isTile = mode == TMAStoreMode::TILE;
228 getCoordinates(), getTensorSize(), getLowerStride(), getUpperStride(),
231 if (mode == TMAStoreMode::TILE_SCATTER4 &&
getCoordinates().size() != 5)
232 overrideAddrRes =
emitError(
"Mode tile scatter4 expects 5 coordinates");
237LogicalResult CpAsyncOp::verify() {
238 if (getModifier() != LoadCacheModifierKind::CG &&
239 getModifier() != LoadCacheModifierKind::CA)
240 return emitError(
"Only CG and CA cache modifiers are supported.");
241 if (getSize() != 4 && getSize() != 8 && getSize() != 16)
242 return emitError(
"expected byte size to be either 4, 8 or 16.");
243 if (getModifier() == LoadCacheModifierKind::CG && getSize() != 16)
244 return emitError(
"CG cache modifier is only support for 16 bytes copy.");
251 if (tensorDims < 1 || tensorDims > 5)
252 return emitError(loc,
"expects coordinates between 1 to 5 dimension");
254 auto checkTMALoadParams = [&](TMALoadMode mode,
bool isIm2col,
255 size_t expectedIm2colOff) -> LogicalResult {
256 if (isIm2col && (tensorDims < 3))
259 <<
" mode, the tensor has to be at least 3-dimensional";
261 if (numIm2colOff != expectedIm2colOff)
262 return emitError(loc) <<
" im2col offsets expected " << expectedIm2colOff
263 <<
" (provided " << numIm2colOff <<
")";
269 case TMALoadMode::TILE:
270 return checkTMALoadParams(mode,
false, 0);
271 case TMALoadMode::IM2COL:
272 return checkTMALoadParams(mode,
true, tensorDims - 2);
273 case TMALoadMode::IM2COL_W:
274 case TMALoadMode::IM2COL_W_128:
275 return checkTMALoadParams(mode,
true, 2);
276 case TMALoadMode::TILE_GATHER4:
277 return (tensorDims == 5)
278 ? checkTMALoadParams(mode,
false, 0)
279 :
emitError(loc,
"Gather4 mode expects 5 coordinates");
284LogicalResult CpAsyncBulkTensorPrefetchOp::verify() {
286 getMode(), getLoc());
289LogicalResult CpAsyncBulkTensorGlobalToSharedClusterOp::verify() {
290 TMALoadMode mode = getMode();
291 bool isCTAOnly = getIsCTAOnly();
292 if (getPredicate()) {
294 return emitError(
"Predicate is supported only for shared::cluster mode.");
295 if (mode != TMALoadMode::TILE && mode != TMALoadMode::IM2COL)
297 "Predicate is supported only for Tile and Im2col modes.");
299 NVVMMemorySpace expectedAS =
300 isCTAOnly ? NVVMMemorySpace::Shared : NVVMMemorySpace::SharedCluster;
301 unsigned AS = llvm::cast<LLVM::LLVMPointerType>(getDstMem().
getType())
303 if (AS != expectedAS)
306 ?
"Shared::cta destination requires address-space 3."
307 :
"Shared::cluster destination requires address-space 7.");
310 if (getMulticastMask())
311 return emitError(
"Multicast is not supported with shared::cta mode.");
313 return emitError(
"CTAGroup is not supported with shared::cta mode.");
318 getMode(), getLoc());
321LogicalResult CpAsyncBulkTensorReduceOp::verify() {
322 TMAStoreMode mode = getMode();
325 case TMAStoreMode::TILE:
327 case TMAStoreMode::IM2COL:
328 case TMAStoreMode::IM2COL_W:
330 case TMAStoreMode::TILE_SCATTER4:
331 return emitError(
"Scatter mode unsupported for CpAsyncBulkTensorReduceOp");
336LogicalResult CpAsyncBulkTensorReduceOverrideAddrOp::verify() {
338 getMode() == TMAStoreMode::IM2COL || getMode() == TMAStoreMode::IM2COL_W;
339 bool isTile = getMode() == TMAStoreMode::TILE;
345 getCoordinates(), getTensorSize(), getLowerStride(), getUpperStride(),
348 if (getMode() == TMAStoreMode::TILE_SCATTER4)
350 "Scatter mode unsupported for CpAsyncBulkTensorReduceOverrideAddrOp");
355LogicalResult CpAsyncBulkGlobalToSharedClusterOp::verify() {
357 if (isSharedCTA && getMulticastMask())
358 return emitError(
"Multicast is not supported with shared::cta mode.");
364 NVVM::MemScopeKind scope,
365 Value retVal =
nullptr) {
366 if (scope != NVVM::MemScopeKind::CTA && scope != NVVM::MemScopeKind::CLUSTER)
367 return op->
emitError(
"mbarrier scope must be either CTA or Cluster");
370 bool hasRetValue =
static_cast<bool>(retVal);
371 if (isSharedCluster && hasRetValue)
373 "mbarrier in shared_cluster space cannot return any value");
378LogicalResult MBarrierArriveOp::verify() {
383LogicalResult MBarrierArriveDropOp::verify() {
388LogicalResult MBarrierArriveExpectTxOp::verify() {
392 if (getPredicate()) {
393 if (getScope() != NVVM::MemScopeKind::CTA)
394 return emitError(
"mbarrier scope must be CTA when using predicate");
397 return emitError(
"mbarrier in shared_cluster space is not supported when "
401 return emitError(
"return-value is not supported when using predicate");
403 if (getRelaxed() ==
true)
404 return emitError(
"mbarrier with relaxed semantics is not supported when "
411LogicalResult MBarrierArriveDropExpectTxOp::verify() {
426 inferredReturnTypes.push_back(IntegerType::get(context, 64));
431MBarrierArriveOp::inferReturnTypes(
MLIRContext *context,
432 std::optional<Location> location,
433 MBarrierArriveOp::Adaptor adaptor,
436 inferredReturnTypes);
439LogicalResult MBarrierArriveDropOp::inferReturnTypes(
440 MLIRContext *context, std::optional<Location> location,
441 MBarrierArriveDropOp::Adaptor adaptor,
444 inferredReturnTypes);
447LogicalResult MBarrierArriveExpectTxOp::inferReturnTypes(
448 MLIRContext *context, std::optional<Location> location,
449 MBarrierArriveExpectTxOp::Adaptor adaptor,
453 if (adaptor.getPredicate())
456 inferredReturnTypes);
459LogicalResult MBarrierArriveDropExpectTxOp::inferReturnTypes(
460 MLIRContext *context, std::optional<Location> location,
461 MBarrierArriveDropExpectTxOp::Adaptor adaptor,
464 inferredReturnTypes);
474 return inferred == actual;
483bool MBarrierArriveExpectTxOp::isCompatibleReturnTypes(
TypeRange l,
487bool MBarrierArriveDropExpectTxOp::isCompatibleReturnTypes(
TypeRange l,
492LogicalResult MBarrierExpectTxOp::verify() {
496LogicalResult MBarrierCompleteTxOp::verify() {
500LogicalResult MBarrierTestWaitOp::verify() {
504LogicalResult MBarrierTryWaitOp::verify() {
508LogicalResult ConvertFloatToTF32Op::verify() {
509 using RndMode = NVVM::FPRoundingMode;
513 return emitError(
"Relu not supported with rna rounding mode.");
520 "Only {rn,rz,rna} rounding modes supported for ConvertFloatToTF32Op.");
525LogicalResult ConvertF32x2ToF6x2Op::verify() {
528 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy())) {
530 << mlir::Float6E2M3FNType::get(ctx) <<
" and "
531 << mlir::Float6E3M2FNType::get(ctx)
532 <<
" types are supported for conversions from f32x2 to f6x2.";
537LogicalResult ConvertF32x2ToF8x2Op::verify() {
538 using RndMode = NVVM::FPRoundingMode;
539 using SatMode = NVVM::SaturationMode;
541 bool isRoundingModeRN = getRnd() == RndMode::RN;
542 bool isRoundingModeRZ = getRnd() == RndMode::RZ;
543 bool isRoundingModeRP = getRnd() == RndMode::RP;
544 bool isSatFinite = getSat() == SatMode::SATFINITE;
546 bool hasRelu = getRelu();
551 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(
553 if (!isRoundingModeRN) {
554 return emitOpError(
"Only RN rounding mode is supported for "
555 "conversions from f32x2 to ")
556 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
557 << mlir::Float8E5M2Type::get(ctx) <<
" types";
560 return emitOpError(
"Only SATFINITE saturation mode is supported "
563 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
564 << mlir::Float8E5M2Type::get(ctx) <<
" types";
568 .Case<mlir::Float8E8M0FNUType>([&](
mlir::Type) -> LogicalResult {
569 if (!(isRoundingModeRZ || isRoundingModeRP)) {
570 return emitOpError(
"Only RZ and RP rounding modes are supported for "
571 "conversions from f32x2 to ")
572 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
575 return emitOpError(
"relu not supported for conversions to ")
576 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
582 << mlir::Float8E4M3FNType::get(ctx) <<
", "
583 << mlir::Float8E5M2Type::get(ctx) <<
", and "
584 << mlir::Float8E8M0FNUType::get(ctx)
586 "supported for conversions from f32x2 to f8x2";
590LogicalResult ConvertF16x2ToF8x2Op::verify() {
593 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy())) {
595 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
596 << mlir::Float8E5M2Type::get(ctx)
597 <<
" types are supported for conversions from f16x2 to f8x2.";
602LogicalResult ConvertBF16x2ToF8x2Op::verify() {
603 using RndMode = NVVM::FPRoundingMode;
604 using SatMode = NVVM::SaturationMode;
606 bool isRoundingModeRN = getRnd() == RndMode::RN;
607 bool isRoundingModeRZ = getRnd() == RndMode::RZ;
608 bool isRoundingModeRP = getRnd() == RndMode::RP;
609 bool isSatFinite = getSat() == SatMode::SATFINITE;
610 bool hasRelu = getRelu();
615 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(
617 if (!isRoundingModeRN)
618 return emitOpError(
"Only RN rounding mode is supported for "
619 "conversions from bf16x2 to ")
620 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
621 << mlir::Float8E5M2Type::get(ctx) <<
" types";
623 return emitOpError(
"Only SATFINITE saturation mode is supported "
624 "for conversions from bf16x2 to ")
625 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
626 << mlir::Float8E5M2Type::get(ctx) <<
" types";
629 .Case<mlir::Float8E8M0FNUType>([&](
mlir::Type) -> LogicalResult {
630 if (!(isRoundingModeRZ || isRoundingModeRP))
631 return emitOpError(
"Only RZ and RP rounding modes are supported for "
632 "conversions from bf16x2 to ")
633 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
635 return emitOpError(
"relu not supported for conversions to ")
636 << mlir::Float8E8M0FNUType::get(ctx) <<
" type";
640 llvm_unreachable(
"Invalid conversion in ConvertBF16x2ToF8x2Op");
645LogicalResult ConvertF32x2ToF4x2Op::verify() {
648 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))
650 << mlir::Float4E2M1FNType::get(ctx)
651 <<
" type is supported for conversions from f32x2 to f4x2.";
656LogicalResult ConvertF8x2ToBF16x2Op::verify() {
658 if (llvm::isa<Float8E8M0FNUType>(getSrcType())) {
659 if (getSat() != SaturationMode::NONE)
661 "Only NONE saturation mode is supported for conversions from ")
662 << Float8E8M0FNUType::get(ctx) <<
" type";
663 if (getScaleFactor())
664 return emitOpError(
"scaleFactor not supported for conversions from ")
665 << Float8E8M0FNUType::get(ctx) <<
" type";
667 return emitOpError(
"relu not supported for conversions from ")
668 << Float8E8M0FNUType::get(ctx) <<
" type";
674LogicalResult PermuteOp::verify() {
675 using Mode = NVVM::PermuteMode;
676 bool hasHi =
static_cast<bool>(getHi());
683 return emitError(
"mode '") << getMode() <<
"' requires 'hi' operand.";
691 << getMode() <<
"' does not accept 'hi' operand.";
706 static constexpr FPRoundingMode validRndModes[] = {
707 FPRoundingMode::RN, FPRoundingMode::RZ, FPRoundingMode::RS};
709 if (!llvm::is_contained(validRndModes, rnd)) {
711 "Only RN, RZ, and RS rounding modes are supported for "
712 "conversions from f32x2 to ")
716 if (rnd == FPRoundingMode::RS) {
717 if (!hasRandomBits) {
718 return op->
emitOpError(
"random_bits is required for RS rounding mode.");
723 "random_bits not supported for RN and RZ rounding modes.");
730LogicalResult ConvertF32x2ToF16x2Op::verify() {
732 getRandomBits() ?
true :
false, *
this);
735LogicalResult ConvertF32x2ToBF16x2Op::verify() {
737 getRandomBits() ?
true :
false, *
this);
740LogicalResult ConvertF32x4ToF8x4Op::verify() {
743 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy()))
745 << mlir::Float8E4M3FNType::get(ctx) <<
" and "
746 << mlir::Float8E5M2Type::get(ctx)
747 <<
" types are supported for conversions from f32x4 to f8x4.";
752LogicalResult ConvertF32x4ToF6x4Op::verify() {
755 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy()))
757 << mlir::Float6E2M3FNType::get(ctx) <<
" and "
758 << mlir::Float6E3M2FNType::get(ctx)
759 <<
" types are supported for conversions from f32x4 to f6x4.";
764LogicalResult ConvertF32x4ToF4x4Op::verify() {
767 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))
768 return emitOpError(
"Only ") << mlir::Float4E2M1FNType::get(ctx)
769 <<
" type is supported for conversions from "
775LogicalResult BulkStoreOp::verify() {
776 if (getInitVal() != 0)
777 return emitOpError(
"only 0 is supported for initVal, got ") << getInitVal();
781LogicalResult AsyncStoreGlobalOp::verify() {
782 NVVM::MemScopeKind scope = getScope();
783 bool isMmio = getMmio();
784 bool isMultimem = getMultimem();
786 if (scope != MemScopeKind::SYS && scope != MemScopeKind::GPU)
787 return emitOpError(
"scope must be either SYS or GPU");
789 if (isMmio && scope != MemScopeKind::SYS)
790 return emitOpError(
"mmio is only supported for SYS scope");
792 if (isMmio && isMultimem)
793 return emitOpError(
"multimem is not supported with mmio");
798LogicalResult PMEventOp::verify() {
799 auto eventId = getEventId();
800 auto maskedEventId = getMaskedEventId();
801 if (!maskedEventId && !eventId) {
802 return emitOpError() <<
"either `id` or `mask` must be set";
805 if (maskedEventId && eventId) {
806 return emitOpError() <<
"`id` and `mask` cannot be set at the same time";
810 if (eventId < 0 || eventId > 15) {
811 return emitOpError() <<
"`id` must be between 0 and 15";
815 return llvm::success();
821std::optional<mlir::NVVM::MMATypes>
822MmaOp::inferOperandMMAType(
Type operandElType,
bool isAccumulator) {
824 VectorType::get(2, Float16Type::get(operandElType.
getContext()));
825 if (operandElType.
isF64())
826 return NVVM::MMATypes::f64;
827 if (operandElType.
isF16() || operandElType == half2Type)
828 return NVVM::MMATypes::f16;
829 if (operandElType.
isF32() && isAccumulator)
830 return NVVM::MMATypes::f32;
831 if (operandElType.
isF32() && !isAccumulator)
832 return NVVM::MMATypes::tf32;
833 if (llvm::isa<IntegerType>(operandElType)) {
835 return NVVM::MMATypes::s32;
839 if (
auto structType = llvm::dyn_cast<LLVM::LLVMStructType>(operandElType)) {
840 if (structType.getBody().empty())
842 return inferOperandMMAType(structType.getBody()[0], isAccumulator);
849 return (type == MMATypes::u4 || type == MMATypes::s4);
853 return (type == MMATypes::u8 || type == MMATypes::s8);
858 type == MMATypes::s32;
861MMATypes MmaOp::accumPtxType() {
862 std::optional<mlir::NVVM::MMATypes> val = inferOperandMMAType(
863 getODSOperands(2).getTypes().front(),
true);
864 assert(val.has_value() &&
"accumulator PTX type should always be inferrable");
868MMATypes MmaOp::resultPtxType() {
869 std::optional<mlir::NVVM::MMATypes> val =
870 inferOperandMMAType(getResult().
getType(),
true);
871 assert(val.has_value() &&
"result PTX type should always be inferrable");
875template <
typename AttrTy>
877 StringRef keyword, AttrTy value) {
878 printer << (isFirst ?
" " :
", ") << keyword <<
" = ";
883template <
typename AttrTy>
885 StringRef keyword, AttrTy value) {
886 printer << (isFirst ?
" " :
", ") << keyword <<
" = "
887 << NVVM::stringifyEnum(value.getValue());
893 printer << (isFirst ?
" " :
", ") << keyword;
897template <
typename AttrTy>
901 if (attributes.
get(name))
903 "duplicate property '" + name +
"'");
907 attributes.
append(name, value);
911template <
typename AttrTy>
915 if (attributes.
get(name))
917 "duplicate property '" + name +
"'");
921 attributes.
append(name, value);
926 return llvm::is_contained(
928 "shape",
"b1Op",
"intOverflowBehavior",
"layoutA",
"layoutB",
929 "multiplicandAPtxType",
"multiplicandBPtxType",
"orderedMetadata",
930 "kind",
"scaleVecSize",
"blockScaleFormat",
"operandSegmentSizes"},
942 if (!llvm::is_contained(allowedKeywords, keyword))
944 "unknown MMA property '" + keyword +
"'");
946 ParseResult parseResult =
success();
947 if (keyword ==
"shape")
950 else if (keyword ==
"b1_op")
953 else if (keyword ==
"int_overflow")
955 parser, attributes,
"intOverflowBehavior");
956 else if (keyword ==
"layout_a")
959 else if (keyword ==
"layout_b")
962 else if (keyword ==
"multiplicand_a_ptx_type")
964 parser, attributes,
"multiplicandAPtxType");
965 else if (keyword ==
"multiplicand_b_ptx_type")
967 parser, attributes,
"multiplicandBPtxType");
968 else if (keyword ==
"kind") {
969 if (llvm::is_contained(allowedKeywords,
"block_scale_format"))
971 parser, attributes,
"kind");
975 }
else if (keyword ==
"scale_vec_size")
977 parser, attributes,
"scaleVecSize");
978 else if (keyword ==
"block_scale_format")
980 parser, attributes,
"blockScaleFormat");
981 else if (keyword ==
"ordered_metadata") {
982 if (attributes.
get(
"orderedMetadata"))
984 "duplicate property 'orderedMetadata'");
988 "unknown MMA property '" + keyword +
"'");
990 if (failed(parseResult))
996 for (StringRef property : requiredProperties) {
997 if (!attributes.
get(property))
999 "missing required property '" + property +
"'");
1009 "inherent property '" + attribute.getName().getValue() +
1010 "' must be spelled directly in the operation syntax");
1011 attributes.
append(attribute);
1018 struct MMAOperandFragment {
1019 StringRef operandName;
1020 StringRef ptxTypeAttr;
1021 SmallVector<Value, 4> regs;
1022 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1023 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1026 std::array<MMAOperandFragment, 3> frags{
1027 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
1028 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
1029 MMAOperandFragment(
"C",
"")};
1031 mlir::NVVM::MmaOp::getOperandSegmentSizeAttr()};
1033 for (
unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {
1034 auto &frag = frags[fragIdx];
1035 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);
1036 for (
auto operandIdx = varOperandSpec.first;
1037 operandIdx < varOperandSpec.first + varOperandSpec.second;
1039 frag.regs.push_back(this->getOperand(operandIdx));
1040 if (operandIdx == 0) {
1041 regTypes.push_back(this->getOperand(operandIdx).
getType());
1044 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(
1045 regTypes.back(), fragIdx >= 2);
1047 ignoreAttrNames.push_back(frag.ptxTypeAttr);
1050 auto printMmaOperand = [&](
const MMAOperandFragment &frag) ->
void {
1051 p <<
" " << frag.operandName;
1057 for (
const auto &frag : frags) {
1058 printMmaOperand(frag);
1061 bool isFirstProperty =
true;
1065 if (getIntOverflowBehaviorAttr())
1067 getIntOverflowBehaviorAttr());
1070 if (getMultiplicandAPtxTypeAttr() &&
1071 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
1073 getMultiplicandAPtxTypeAttr());
1074 if (getMultiplicandBPtxTypeAttr() &&
1075 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
1077 getMultiplicandBPtxTypeAttr());
1078 llvm::append_range(ignoreAttrNames,
1080 getIntOverflowBehaviorAttrName(),
1081 getLayoutAAttrName(),
1082 getLayoutBAttrName(),
1083 getMultiplicandAPtxTypeAttrName(),
1084 getMultiplicandBPtxTypeAttrName()});
1090 frags[1].regs[0].getType(),
1091 frags[2].regs[0].getType()},
1100 std::optional<MMAIntOverflow> intOverflow,
1101 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
1102 std::optional<std::array<MMALayout, 2>> multiplicandLayouts) {
1104 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
1109 result.addOperands(operandA);
1110 result.addOperands(operandB);
1111 result.addOperands(operandC);
1113 if (multiplicandPtxTypes) {
1114 result.addAttribute(
"multiplicandAPtxType",
1115 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
1116 result.addAttribute(
"multiplicandBPtxType",
1117 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
1119 if (
auto res = inferOperandMMAType(operandA[0].
getType(),
false))
1120 result.addAttribute(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
1121 if (
auto res = inferOperandMMAType(operandB[0].
getType(),
false))
1122 result.addAttribute(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
1125 if (multiplicandLayouts) {
1126 result.addAttribute(
"layoutA",
1127 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[0]));
1128 result.addAttribute(
"layoutB",
1129 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[1]));
1131 result.addAttribute(
"layoutA", MMALayoutAttr::get(ctx, MMALayout::row));
1132 result.addAttribute(
"layoutB", MMALayoutAttr::get(ctx, MMALayout::col));
1135 if (intOverflow.has_value())
1136 result.addAttribute(
"intOverflowBehavior",
1137 MMAIntOverflowAttr::get(ctx, *intOverflow));
1138 if (b1Op.has_value())
1139 result.addAttribute(
"b1Op", MMAB1OpAttr::get(ctx, *b1Op));
1141 result.addTypes(resultType);
1143 MmaOp::getOperandSegmentSizeAttr(),
1145 static_cast<int32_t>(operandB.size()),
1146 static_cast<int32_t>(operandC.size())}));
1155 struct MMAOperandFragment {
1156 std::optional<MMATypes> elemtype;
1157 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
1158 SmallVector<Type> regTypes;
1162 std::array<MMAOperandFragment, 4> frags;
1168 MMAOperandFragment &frag) -> LogicalResult {
1187 {
"shape",
"b1_op",
"int_overflow",
"layout_a",
1188 "layout_b",
"multiplicand_a_ptx_type",
1189 "multiplicand_b_ptx_type"},
1190 {
"shape",
"layoutA",
"layoutB"}))
1202 if (operandTypes.size() != 3)
1205 "expected one type for each operand segment but got " +
1206 Twine(operandTypes.size()) +
" types");
1207 for (
const auto &iter : llvm::enumerate(operandTypes)) {
1208 auto &frag = frags[iter.index()];
1209 frag.regTypes.resize(frag.regs.size(), iter.value());
1213 frag.elemtype = inferOperandMMAType(frag.regTypes[0],
1220 frags[3].elemtype = inferOperandMMAType(resultType,
true);
1222 std::array<StringRef, 2> names{
"multiplicandAPtxType",
1223 "multiplicandBPtxType"};
1224 for (
unsigned idx = 0; idx < names.size(); idx++) {
1225 const auto &frag = frags[idx];
1226 std::optional<NamedAttribute> attr = namedAttributes.
getNamed(names[idx]);
1227 if (!frag.elemtype.has_value() && !attr.has_value()) {
1230 "attribute " + names[idx] +
1231 " is not provided explicitly and cannot be inferred");
1233 if (!attr.has_value())
1235 names[idx], MMATypesAttr::get(parser.
getContext(), *frag.elemtype));
1238 result.addTypes(resultType);
1239 if (!namedAttributes.
empty())
1240 result.addAttributes(namedAttributes);
1241 result.addAttribute(MmaOp::getOperandSegmentSizeAttr(),
1243 static_cast<int32_t>(frags[0].regs.size()),
1244 static_cast<int32_t>(frags[1].regs.size()),
1245 static_cast<int32_t>(frags[2].regs.size()),
1250LogicalResult MmaOp::verify() {
1252 auto f16Ty = Float16Type::get(context);
1253 auto i32Ty = IntegerType::get(context, 32);
1254 auto f16x2Ty = VectorType::get(2, f16Ty);
1255 auto f32Ty = Float32Type::get(context);
1256 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(
1257 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});
1259 auto s32x4StructTy =
1260 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});
1261 auto f32x8StructTy =
1263 auto f16x2x2StructTy =
1264 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});
1265 auto f32x4StructTy =
1266 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});
1267 auto s32x2StructTy =
1268 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});
1270 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),
1271 getShapeAttr().getK()};
1277 AllowedShapes allowedShapes;
1278 AllowedTypes expectedA;
1279 AllowedTypes expectedB;
1280 AllowedTypes expectedC;
1285 if (mmaShape[0] == 16) {
1287 Type multiplicandFragType;
1288 switch (*getMultiplicandAPtxType()) {
1289 case MMATypes::tf32:
1291 multiplicandFragType = i32Ty;
1292 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1293 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1295 case MMATypes::bf16:
1297 multiplicandFragType = i32Ty;
1298 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1299 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1303 multiplicandFragType = f16x2Ty;
1304 expectedResult.push_back(f16x2x2StructTy);
1305 expectedResult.push_back(f32x4StructTy);
1307 case MMATypes::e4m3:
1308 case MMATypes::e5m2:
1312 multiplicandFragType = i32Ty;
1313 expectedResult.push_back(f16x2x2StructTy);
1314 expectedResult.push_back(f32x4StructTy);
1328 return emitError(
"invalid shape or multiplicand type: ")
1329 << getMultiplicandAPtxType().value();
1333 expectedResult.push_back(s32x4StructTy);
1334 expectedC.emplace_back(4, i32Ty);
1335 multiplicandFragType = i32Ty;
1337 expectedC.emplace_back(2, f16x2Ty);
1338 expectedC.emplace_back(4, f32Ty);
1341 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor);
1342 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);
1343 expectedA.emplace_back(unitA, multiplicandFragType);
1344 expectedB.emplace_back(unitB, multiplicandFragType);
1345 allowedShapes.push_back({16, 8, kFactor});
1346 allowedShapes.push_back({16, 8, kFactor * 2});
1348 if (resultPtxType() != accumPtxType())
1353 if (mmaShape[0] == 8) {
1354 if (*getMultiplicandAPtxType() == MMATypes::f16) {
1355 expectedA.emplace_back(2, f16x2Ty);
1356 expectedB.emplace_back(2, f16x2Ty);
1357 expectedResult.push_back(f16x2x4StructTy);
1358 expectedResult.push_back(f32x8StructTy);
1359 expectedC.emplace_back(4, f16x2Ty);
1360 expectedC.emplace_back(8, f32Ty);
1361 allowedShapes.push_back({8, 8, 4});
1363 if (*getMultiplicandAPtxType() == MMATypes::f64) {
1364 Type f64Ty = Float64Type::get(context);
1365 expectedA.emplace_back(1, f64Ty);
1366 expectedB.emplace_back(1, f64Ty);
1367 expectedC.emplace_back(2, f64Ty);
1368 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(
1370 allowedShapes.push_back({8, 8, 4});
1373 expectedA.push_back({i32Ty});
1374 expectedB.push_back({i32Ty});
1375 expectedC.push_back({i32Ty, i32Ty});
1376 expectedResult.push_back(s32x2StructTy);
1378 allowedShapes.push_back({8, 8, 32});
1380 allowedShapes.push_back({8, 8, 16});
1381 if (getMultiplicandAPtxType().value() == MMATypes::b1)
1382 allowedShapes.push_back({8, 8, 128});
1386 std::string errorMessage;
1387 llvm::raw_string_ostream errorStream(errorMessage);
1390 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||
1391 !llvm::is_contained(allowedShapes, mmaShape)) {
1392 errorStream <<
"unimplemented variant for MMA shape <";
1393 llvm::interleaveComma(mmaShape, errorStream);
1399 std::array<StringRef, 3> operandNames{
"A",
"B",
"C"};
1400 for (
const auto &iter : llvm::enumerate(
1402 auto spec = this->getODSOperandIndexAndLength(iter.index());
1404 operand_type_begin() + spec.first +
1406 bool match = llvm::is_contained(iter.value(), operandTySeg);
1409 errorStream <<
"Could not match types for the "
1410 << operandNames[iter.index()]
1411 <<
" operands; expected one of ";
1412 for (
const auto &x : iter.value()) {
1413 errorStream << x.size() <<
"x" << x[0] <<
" ";
1415 errorStream <<
"but got ";
1416 llvm::interleaveComma(operandTySeg, errorStream);
1422 if (!llvm::any_of(expectedResult, [&](
Type expectedResultType) {
1423 return expectedResultType == getResult().getType();
1426 <<
"Could not match allowed types for the result; expected one of ";
1427 llvm::interleaveComma(expectedResult, errorStream);
1428 errorStream <<
" but got " << getResult().getType();
1433 if (getMultiplicandAPtxType() == MMATypes::b1 && !getB1Op()) {
1434 return emitOpError(
"op requires " + getB1OpAttrName().strref() +
1442 if (!getIntOverflowBehavior())
1444 getIntOverflowBehaviorAttrName().strref() +
1452 (mmaShape[0] == 8 && mmaShape[1] == 8 && mmaShape[2] == 4 &&
1453 getMultiplicandAPtxType() == MMATypes::f16);
1455 if (!isM8N8K4_F16) {
1457 if (getLayoutA() != MMALayout::row || getLayoutB() != MMALayout::col) {
1458 return emitOpError(
"requires layoutA = #nvvm.mma_layout<row> and "
1459 "layoutB = #nvvm.mma_layout<col> for shape <")
1460 << mmaShape[0] <<
", " << mmaShape[1] <<
", " << mmaShape[2]
1461 <<
"> with element types " << *getMultiplicandAPtxType() <<
" and "
1462 << *getMultiplicandBPtxType()
1463 <<
". Only m8n8k4 with f16 supports other layouts.";
1470MMATypes MmaSpOp::accumPtxType() {
1471 std::optional<mlir::NVVM::MMATypes> val = MmaOp::inferOperandMMAType(
1472 getODSOperands(2).getTypes().front(),
true);
1473 assert(val.has_value() &&
"accumulator PTX type should always be inferrable");
1477MMATypes MmaSpOp::resultPtxType() {
1478 std::optional<mlir::NVVM::MMATypes> val =
1479 MmaOp::inferOperandMMAType(getResult().
getType(),
true);
1480 assert(val.has_value() &&
"result PTX type should always be inferrable");
1486 llvm::IRBuilderBase &builder) {
1487 auto thisOp = cast<NVVM::MmaSpOp>(op);
1495 auto intId = MmaSpOp::getIntrinsicID(
1496 thisOp.getShape().getM(), thisOp.getShape().getN(),
1497 thisOp.getShape().getK(), thisOp.getIntOverflowBehavior(),
1498 thisOp.getOrderedMetadata(), thisOp.getKind(),
1499 *thisOp.getMultiplicandAPtxType(), *thisOp.getMultiplicandBPtxType(),
1500 thisOp.accumPtxType(), thisOp.resultPtxType());
1502 return {intId, args};
1507 struct MMAOperandFragment {
1508 StringRef operandName;
1509 StringRef ptxTypeAttr;
1510 SmallVector<Value, 4> regs;
1511 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1512 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1515 std::array<MMAOperandFragment, 5> frags{
1516 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
1517 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
1518 MMAOperandFragment(
"C",
""), MMAOperandFragment(
"sparseMetadata",
""),
1519 MMAOperandFragment(
"selector",
"")};
1521 mlir::NVVM::MmaSpOp::getOperandSegmentSizeAttr()};
1524 for (
unsigned fragIdx = 0; fragIdx < 3; fragIdx++) {
1525 auto &frag = frags[fragIdx];
1526 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);
1527 for (
auto operandIdx = varOperandSpec.first;
1528 operandIdx < varOperandSpec.first + varOperandSpec.second;
1530 frag.regs.push_back(this->getOperand(operandIdx));
1531 if (operandIdx == varOperandSpec.first) {
1532 regTypes.push_back(this->getOperand(operandIdx).
getType());
1535 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(
1536 regTypes.back(), fragIdx >= 2);
1538 ignoreAttrNames.push_back(frag.ptxTypeAttr);
1542 frags[3].regs.push_back(getSparseMetadata());
1543 frags[4].regs.push_back(getSparsitySelector());
1545 auto printMmaSpOperand = [&](
const MMAOperandFragment &frag) ->
void {
1546 p <<
" " << frag.operandName;
1552 for (
const auto &frag : frags)
1553 printMmaSpOperand(frag);
1555 bool isFirstProperty =
true;
1557 if (getIntOverflowBehaviorAttr())
1559 getIntOverflowBehaviorAttr());
1560 if (getMultiplicandAPtxTypeAttr() &&
1561 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
1563 getMultiplicandAPtxTypeAttr());
1564 if (getMultiplicandBPtxTypeAttr() &&
1565 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
1567 getMultiplicandBPtxTypeAttr());
1568 if (getOrderedMetadata())
1575 getMultiplicandAPtxTypeAttrName(),
1576 getMultiplicandBPtxTypeAttrName(),
1577 getOrderedMetadataAttrName(), getKindAttrName()});
1581 for (
int i = 0; i < 3; ++i) {
1586 p <<
") -> " << getResult().getType();
1593 std::optional<MMAIntOverflow> intOverflow,
1594 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {
1596 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
1601 result.addOperands(operandA);
1602 result.addOperands(operandB);
1603 result.addOperands(operandC);
1604 result.addOperands(sparseMetadata);
1605 result.addOperands(sparsitySelector);
1607 if (multiplicandPtxTypes) {
1608 result.addAttribute(
"multiplicandAPtxType",
1609 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
1610 result.addAttribute(
"multiplicandBPtxType",
1611 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
1613 if (
auto res = MmaOp::inferOperandMMAType(operandA[0].
getType(),
false))
1614 result.addAttribute(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
1615 if (
auto res = MmaOp::inferOperandMMAType(operandB[0].
getType(),
false))
1616 result.addAttribute(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
1619 if (intOverflow.has_value())
1620 result.addAttribute(
"intOverflowBehavior",
1621 MMAIntOverflowAttr::get(ctx, *intOverflow));
1623 result.addTypes(resultType);
1625 MmaSpOp::getOperandSegmentSizeAttr(),
1627 static_cast<int32_t>(operandB.size()),
1628 static_cast<int32_t>(operandC.size()), 1,
1633 struct MMAOperandFragment {
1634 std::optional<MMATypes> elemtype;
1635 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
1636 SmallVector<Type> regTypes;
1640 std::array<MMAOperandFragment, 6> frags;
1645 auto parseMmaSpOperand = [&](StringRef operandName,
1646 MMAOperandFragment &frag) -> LogicalResult {
1657 if (parseMmaSpOperand(
"A", frags[0]).
failed())
1659 if (parseMmaSpOperand(
"B", frags[1]).
failed())
1661 if (parseMmaSpOperand(
"C", frags[2]).
failed())
1663 if (parseMmaSpOperand(
"sparseMetadata", frags[3]).
failed())
1665 if (parseMmaSpOperand(
"selector", frags[4]).
failed())
1669 {
"shape",
"int_overflow",
"multiplicand_a_ptx_type",
1670 "multiplicand_b_ptx_type",
"ordered_metadata",
1685 if (operandTypes.size() != 3)
1688 "expected one type for each operand segment but got " +
1689 Twine(operandTypes.size()) +
" types");
1690 for (
const auto &iter : llvm::enumerate(operandTypes)) {
1691 auto &frag = frags[iter.index()];
1692 frag.regTypes.resize(frag.regs.size(), iter.value());
1697 MmaOp::inferOperandMMAType(frag.regTypes[0],
1705 MmaOp::inferOperandMMAType(resultType,
true);
1720 std::array<StringRef, 2> names{
"multiplicandAPtxType",
1721 "multiplicandBPtxType"};
1722 for (
unsigned idx = 0; idx < names.size(); idx++) {
1723 const auto &frag = frags[idx];
1724 std::optional<NamedAttribute> attr = namedAttributes.
getNamed(names[idx]);
1725 if (!frag.elemtype.has_value() && !attr.has_value()) {
1728 "attribute " + names[idx] +
1729 " is not provided explicitly and cannot be inferred");
1731 if (!attr.has_value())
1733 names[idx], MMATypesAttr::get(parser.
getContext(), *frag.elemtype));
1736 result.addTypes(resultType);
1737 if (!namedAttributes.
empty())
1738 result.addAttributes(namedAttributes);
1739 result.addAttribute(MmaSpOp::getOperandSegmentSizeAttr(),
1741 static_cast<int32_t>(frags[0].regs.size()),
1742 static_cast<int32_t>(frags[1].regs.size()),
1743 static_cast<int32_t>(frags[2].regs.size()),
1750LogicalResult MmaSpOp::verify() {
1752 auto f16Ty = Float16Type::get(context);
1753 auto i32Ty = IntegerType::get(context, 32);
1754 auto f16x2Ty = VectorType::get(2, f16Ty);
1755 auto f32Ty = Float32Type::get(context);
1756 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(
1757 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});
1759 auto s32x4StructTy =
1760 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});
1761 auto f32x8StructTy =
1763 auto f16x2x2StructTy =
1764 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});
1765 auto f32x4StructTy =
1766 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});
1767 auto s32x2StructTy =
1768 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});
1770 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),
1771 getShapeAttr().getK()};
1777 AllowedShapes allowedShapes;
1778 AllowedTypes expectedA;
1779 AllowedTypes expectedB;
1780 AllowedTypes expectedC;
1785 if (mmaShape[0] == 16) {
1787 Type multiplicandFragType;
1788 switch (*getMultiplicandAPtxType()) {
1789 case MMATypes::tf32:
1791 multiplicandFragType = i32Ty;
1792 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1793 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1795 allowedShapes.push_back({16, 8, 8});
1796 allowedShapes.push_back({16, 8, 16});
1798 case MMATypes::bf16:
1800 multiplicandFragType = i32Ty;
1801 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1802 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1804 allowedShapes.push_back({16, 8, 16});
1805 allowedShapes.push_back({16, 8, 32});
1809 multiplicandFragType = f16x2Ty;
1810 expectedResult.push_back(f16x2x2StructTy);
1811 expectedResult.push_back(f32x4StructTy);
1813 allowedShapes.push_back({16, 8, 16});
1814 allowedShapes.push_back({16, 8, 32});
1820 allowedShapes.push_back({16, 8, 64});
1821 allowedShapes.push_back({16, 8, 128});
1827 allowedShapes.push_back({16, 8, 32});
1828 allowedShapes.push_back({16, 8, 64});
1830 case MMATypes::e4m3:
1831 case MMATypes::e5m2:
1832 case MMATypes::e3m2:
1833 case MMATypes::e2m3:
1834 case MMATypes::e2m1:
1836 multiplicandFragType = i32Ty;
1837 expectedResult.push_back(f16x2x2StructTy);
1838 expectedResult.push_back(f32x4StructTy);
1840 allowedShapes.push_back({16, 8, 64});
1843 return emitError(
"invalid shape or multiplicand type: ")
1844 << getMultiplicandAPtxType().value();
1848 expectedResult.push_back(s32x4StructTy);
1849 expectedC.emplace_back(4, i32Ty);
1850 multiplicandFragType = i32Ty;
1851 }
else if (*getMultiplicandAPtxType() >= MMATypes::e4m3 &&
1852 *getMultiplicandAPtxType() <= MMATypes::e2m1) {
1854 expectedC.emplace_back(2, f16x2Ty);
1855 expectedC.emplace_back(4, f32Ty);
1857 expectedC.emplace_back(2, f16x2Ty);
1858 expectedC.emplace_back(4, f32Ty);
1863 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor) / 2;
1864 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);
1865 expectedA.emplace_back(unitA, multiplicandFragType);
1866 expectedB.emplace_back(unitB, multiplicandFragType);
1868 if (resultPtxType() != accumPtxType())
1873 if (mmaShape[0] == 8) {
1874 if (*getMultiplicandAPtxType() == MMATypes::f16) {
1875 expectedA.emplace_back(2, f16x2Ty);
1876 expectedB.emplace_back(2, f16x2Ty);
1877 expectedResult.push_back(f16x2x4StructTy);
1878 expectedResult.push_back(f32x8StructTy);
1879 expectedC.emplace_back(4, f16x2Ty);
1880 expectedC.emplace_back(8, f32Ty);
1881 allowedShapes.push_back({8, 8, 4});
1883 if (*getMultiplicandAPtxType() == MMATypes::f64) {
1884 Type f64Ty = Float64Type::get(context);
1885 expectedA.emplace_back(1, f64Ty);
1886 expectedB.emplace_back(1, f64Ty);
1887 expectedC.emplace_back(2, f64Ty);
1888 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(
1890 allowedShapes.push_back({8, 8, 4});
1893 expectedA.push_back({i32Ty});
1894 expectedB.push_back({i32Ty});
1895 expectedC.push_back({i32Ty, i32Ty});
1896 expectedResult.push_back(s32x2StructTy);
1898 allowedShapes.push_back({8, 8, 32});
1900 allowedShapes.push_back({8, 8, 16});
1904 std::string errorMessage;
1905 llvm::raw_string_ostream errorStream(errorMessage);
1908 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||
1909 !llvm::is_contained(allowedShapes, mmaShape)) {
1910 errorStream <<
"unimplemented variant for MMA shape <";
1911 llvm::interleaveComma(mmaShape, errorStream);
1917 std::array<StringRef, 3> operandNames{
"A",
"B",
"C"};
1918 for (
const auto &iter : llvm::enumerate(
1920 auto spec = this->getODSOperandIndexAndLength(iter.index());
1922 operand_type_begin() + spec.first +
1924 bool match = llvm::is_contained(iter.value(), operandTySeg);
1927 errorStream <<
"Could not match types for the "
1928 << operandNames[iter.index()]
1929 <<
" operands; expected one of ";
1930 for (
const auto &x : iter.value()) {
1931 errorStream << x.size() <<
"x" << x[0] <<
" ";
1933 errorStream <<
"but got ";
1934 llvm::interleaveComma(operandTySeg, errorStream);
1940 if (!llvm::any_of(expectedResult, [&](
Type expectedResultType) {
1941 return expectedResultType == getResult().getType();
1944 <<
"Could not match allowed types for the result; expected one of ";
1945 llvm::interleaveComma(expectedResult, errorStream);
1946 errorStream <<
" but got " << getResult().getType();
1954 if (!getIntOverflowBehavior())
1956 getIntOverflowBehaviorAttrName().strref() +
1961 if (!getSparseMetadata().
getType().isInteger(32)) {
1962 return emitOpError() <<
"sparse metadata must be i32 type";
1966 if (!getSparsitySelector().
getType().isInteger(32)) {
1967 return emitOpError() <<
"sparsity selector must be i32 type";
1979struct MMAOperandFragment {
1980 StringRef operandName;
1981 StringRef ptxTypeAttr;
1982 SmallVector<Value, 4> regs;
1983 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1984 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1991 p <<
" " << name <<
"[";
2010template <
typename Op>
2015 for (
unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {
2016 auto &frag = frags[fragIdx];
2017 auto varOperandSpec = op.getODSOperandIndexAndLength(fragIdx);
2018 for (
auto operandIdx = varOperandSpec.first;
2019 operandIdx < varOperandSpec.first + varOperandSpec.second;
2021 frag.regs.push_back(op.getOperand(operandIdx));
2022 if (fragIdx == 0 && operandIdx == varOperandSpec.first) {
2023 regTypes.push_back(op.getOperand(operandIdx).getType());
2027 regTypes.push_back(frag.regs[0].getType());
2029 std::optional<MMATypes> inferredType =
2030 MmaOp::inferOperandMMAType(regTypes.back(),
2033 ignoreAttrNames.push_back(frag.ptxTypeAttr);
2044 auto typeParser = [&]() {
2048 operandTypes.push_back(ty);
2054 if (operandTypes.size() != 3)
2056 "expected exactly 3 types");
2065 if (!attrs.
get(
"multiplicandAPtxType")) {
2066 if (
auto inferredType =
2067 MmaOp::inferOperandMMAType(operandTypes[0],
false)) {
2068 attrs.
set(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *inferredType));
2071 if (!attrs.
get(
"multiplicandBPtxType")) {
2072 if (
auto inferredType =
2073 MmaOp::inferOperandMMAType(operandTypes[1],
false)) {
2074 attrs.
set(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *inferredType));
2080template <
typename OpType>
2083 ScaleVecSize scaleVecSize,
2084 BlockScaleFormat blockScaleFormat,
2085 MMABlockScaleKind kind) {
2087 auto &properties =
result.getOrAddProperties<
typename OpType::Properties>();
2088 properties.setShape(
2090 properties.setScaleVecSize(ScaleVecSizeAttr::get(ctx, scaleVecSize));
2091 properties.setBlockScaleFormat(
2092 BlockScaleFormatAttr::get(ctx, blockScaleFormat));
2093 properties.setKind(MMABlockScaleKindAttr::get(ctx, kind));
2100 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {
2101 if (multiplicandPtxTypes) {
2102 result.addAttribute(
"multiplicandAPtxType",
2103 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
2104 result.addAttribute(
"multiplicandBPtxType",
2105 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
2107 if (
auto res = MmaOp::inferOperandMMAType(operandA[0].
getType(),
false))
2108 result.addAttribute(
"multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
2109 if (
auto res = MmaOp::inferOperandMMAType(operandB[0].
getType(),
false))
2110 result.addAttribute(
"multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
2115template <
typename OpTy>
2117 return *MmaOp::inferOperandMMAType(
2118 cast<LLVM::LLVMStructType>(op.getRes().getType()).getBody()[0],
2128 std::array<MMAOperandFragment, 3> frags{
2129 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
2130 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
2131 MMAOperandFragment(
"C",
"")};
2133 mlir::NVVM::MmaBlockScaleOp::getOperandSegmentSizeAttr()};
2138 for (
const auto &frag : frags)
2143 {getScaleAData(), getByteIdA(), getThreadIdA()});
2145 {getScaleBData(), getByteIdB(), getThreadIdB()});
2147 bool isFirstProperty =
true;
2149 if (getMultiplicandAPtxTypeAttr() &&
2150 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
2152 getMultiplicandAPtxTypeAttr());
2153 if (getMultiplicandBPtxTypeAttr() &&
2154 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
2156 getMultiplicandBPtxTypeAttr());
2158 getScaleVecSizeAttr());
2160 getBlockScaleFormatAttr());
2165 getMultiplicandBPtxTypeAttrName(),
2166 getScaleVecSizeAttrName(),
2167 getBlockScaleFormatAttrName(), getKindAttrName()});
2173 frags[1].regs[0].getType(),
2174 frags[2].regs[0].getType()},
2180ParseResult MmaBlockScaleOp::parse(
OpAsmParser &parser,
2182 struct LocalOperandFragment {
2183 std::optional<MMATypes> elemtype;
2184 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
2188 std::array<LocalOperandFragment, 3> frags;
2204 {
"shape",
"multiplicand_a_ptx_type",
2205 "multiplicand_b_ptx_type",
"scale_vec_size",
2206 "block_scale_format",
"kind"},
2207 {
"shape",
"scaleVecSize",
"blockScaleFormat",
"kind"}))
2221 for (
const auto &[idx, frag] : llvm::enumerate(frags)) {
2222 frag.elemtype = MmaOp::inferOperandMMAType(operandTypes[idx],
2225 .resolveOperands(frag.regs, operandTypes[idx], parser.
getNameLoc(),
2235 .resolveOperands(scaleAOperands, scaleTypes, parser.
getNameLoc(),
2245 result.addAttributes(namedAttributes);
2249 result.addTypes(resultTypes);
2250 result.addAttribute(MmaBlockScaleOp::getOperandSegmentSizeAttr(),
2252 static_cast<int32_t>(frags[0].regs.size()),
2253 static_cast<int32_t>(frags[1].regs.size()),
2254 static_cast<int32_t>(frags[2].regs.size()),
2265void MmaBlockScaleOp::build(
2270 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
2271 ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat,
2272 MMABlockScaleKind kind) {
2273 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
2276 blockScaleFormat, kind);
2278 result.addOperands(operandA);
2279 result.addOperands(operandB);
2280 result.addOperands(operandC);
2282 {scaleAData, byteIdA, threadIdA, scaleBData, byteIdB, threadIdB});
2285 multiplicandPtxTypes);
2287 result.addTypes(resultType);
2288 result.addAttribute(MmaBlockScaleOp::getOperandSegmentSizeAttr(),
2290 static_cast<int32_t>(operandA.size()),
2291 static_cast<int32_t>(operandB.size()),
2292 static_cast<int32_t>(operandC.size()),
2304 auto curOp = cast<NVVM::MmaBlockScaleOp>(op);
2308 for (
Value operand : curOp.getOperandA())
2310 for (
Value operand : curOp.getOperandB())
2312 for (
Value operand : curOp.getOperandC())
2316 args.push_back(mt.
lookupValue(curOp.getScaleAData()));
2317 args.push_back(mt.
lookupValue(curOp.getByteIdA()));
2318 args.push_back(mt.
lookupValue(curOp.getThreadIdA()));
2319 args.push_back(mt.
lookupValue(curOp.getScaleBData()));
2320 args.push_back(mt.
lookupValue(curOp.getByteIdB()));
2321 args.push_back(mt.
lookupValue(curOp.getThreadIdB()));
2323 unsigned intId = MmaBlockScaleOp::getIntrinsicID(
2324 curOp.getShape().getM(), curOp.getShape().getN(), curOp.getShape().getK(),
2325 *curOp.getMultiplicandAPtxType(), *curOp.getMultiplicandBPtxType(),
2327 curOp.getBlockScaleFormat(), curOp.getKind());
2329 return {intId, args};
2332LogicalResult MmaBlockScaleOp::verify() {
2338 if (m == 16 && n == 8 && k == 64) {
2339 if (getMultiplicandAPtxType() != NVVM::MMATypes::e2m1 ||
2340 getMultiplicandBPtxType() != NVVM::MMATypes::e2m1)
2342 "unsupported MMATypes attribute for mma.m16n8k64.(mxf4nvf4|mxf4)");
2343 if (getKind() == NVVM::MMABlockScaleKind::MXF4) {
2344 if (getScaleVecSize() != NVVM::ScaleVecSize::X2)
2346 "unsupported ScaleVecSize attribute for mma.m16n8k64.mxf4");
2347 if (getBlockScaleFormat() != NVVM::BlockScaleFormat::UE8M0)
2349 "unsupported BlockScaleFormat attribute for mma.m16n8k64.mxf4");
2350 }
else if (getKind() == NVVM::MMABlockScaleKind::MXF4NVF4) {
2351 if (!((getScaleVecSize() == NVVM::ScaleVecSize::X2 &&
2352 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0) ||
2353 (getScaleVecSize() == NVVM::ScaleVecSize::X4 &&
2354 (getBlockScaleFormat() == NVVM::BlockScaleFormat::UE4M3 ||
2355 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))))
2357 "attributes for mma.m16n8k64.mxf4nvf4");
2361 }
else if (m == 16 && n == 8 && k == 32) {
2362 if (!(getKind() == NVVM::MMABlockScaleKind::MXF8F6F4 &&
2363 getScaleVecSize() == NVVM::ScaleVecSize::X1 &&
2364 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))
2366 emitOpError(
"unsupported Kind, ScaleVecSize and BlockScaleFormat "
2367 "attributes for mma.m16n8k32");
2380 std::array<MMAOperandFragment, 3> frags{
2381 MMAOperandFragment(
"A", getMultiplicandAPtxTypeAttrName()),
2382 MMAOperandFragment(
"B", getMultiplicandBPtxTypeAttrName()),
2383 MMAOperandFragment(
"C",
"")};
2385 mlir::NVVM::MmaSpBlockScaleOp::getOperandSegmentSizeAttr()};
2390 for (
const auto &frag : frags)
2399 {getScaleAData(), getByteIdA(), getThreadIdA()});
2401 {getScaleBData(), getByteIdB(), getThreadIdB()});
2403 bool isFirstProperty =
true;
2405 if (getMultiplicandAPtxTypeAttr() &&
2406 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
2408 getMultiplicandAPtxTypeAttr());
2409 if (getMultiplicandBPtxTypeAttr() &&
2410 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
2412 getMultiplicandBPtxTypeAttr());
2415 getScaleVecSizeAttr());
2417 getBlockScaleFormatAttr());
2422 getMultiplicandBPtxTypeAttrName(),
2423 getOrderedMetadataAttrName(),
2424 getScaleVecSizeAttrName(),
2425 getBlockScaleFormatAttrName(), getKindAttrName()});
2431 frags[1].regs[0].getType(),
2432 frags[2].regs[0].getType()},
2438ParseResult MmaSpBlockScaleOp::parse(
OpAsmParser &parser,
2440 struct LocalOperandFragment {
2441 std::optional<MMATypes> elemtype;
2442 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
2446 std::array<LocalOperandFragment, 3> frags;
2469 {
"shape",
"multiplicand_a_ptx_type",
2470 "multiplicand_b_ptx_type",
"ordered_metadata",
2471 "scale_vec_size",
"block_scale_format",
"kind"},
2472 {
"shape",
"scaleVecSize",
"blockScaleFormat",
"kind"}))
2486 for (
const auto &[idx, frag] : llvm::enumerate(frags)) {
2487 frag.elemtype = MmaOp::inferOperandMMAType(operandTypes[idx],
2490 .resolveOperands(frag.regs, operandTypes[idx], parser.
getNameLoc(),
2499 .resolveOperands(metadataOperands, i32Type, parser.
getNameLoc(),
2512 .resolveOperands(scaleAOperands, scaleTypes, parser.
getNameLoc(),
2522 result.addAttributes(namedAttributes);
2527 if (!
result.attributes.get(
"orderedMetadata"))
2530 result.addTypes(resultTypes);
2531 result.addAttribute(MmaSpBlockScaleOp::getOperandSegmentSizeAttr(),
2533 static_cast<int32_t>(frags[0].regs.size()),
2534 static_cast<int32_t>(frags[1].regs.size()),
2535 static_cast<int32_t>(frags[2].regs.size()),
2548void MmaSpBlockScaleOp::build(
2554 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
2555 ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat,
2556 MMABlockScaleKind kind) {
2557 assert(
shape.size() == 3 &&
"expected shape to have size 3 (m, n, k)");
2560 builder,
result,
shape, scaleVecSize, blockScaleFormat, kind);
2563 result.addOperands(operandA);
2564 result.addOperands(operandB);
2565 result.addOperands(operandC);
2566 result.addOperands({sparseMetadata, sparsitySelector, scaleAData, byteIdA,
2567 threadIdA, scaleBData, byteIdB, threadIdB});
2570 multiplicandPtxTypes);
2572 result.addTypes(resultType);
2573 result.addAttribute(MmaSpBlockScaleOp::getOperandSegmentSizeAttr(),
2575 static_cast<int32_t>(operandA.size()),
2576 static_cast<int32_t>(operandB.size()),
2577 static_cast<int32_t>(operandC.size()),
2591 auto curOp = cast<NVVM::MmaSpBlockScaleOp>(op);
2595 for (
Value operand : curOp.getOperandA())
2597 for (
Value operand : curOp.getOperandB())
2599 for (
Value operand : curOp.getOperandC())
2603 args.push_back(mt.
lookupValue(curOp.getSparseMetadata()));
2604 args.push_back(mt.
lookupValue(curOp.getSparsitySelector()));
2607 args.push_back(mt.
lookupValue(curOp.getScaleAData()));
2608 args.push_back(mt.
lookupValue(curOp.getByteIdA()));
2609 args.push_back(mt.
lookupValue(curOp.getThreadIdA()));
2610 args.push_back(mt.
lookupValue(curOp.getScaleBData()));
2611 args.push_back(mt.
lookupValue(curOp.getByteIdB()));
2612 args.push_back(mt.
lookupValue(curOp.getThreadIdB()));
2614 unsigned intId = MmaSpBlockScaleOp::getIntrinsicID(
2615 curOp.getShape().getM(), curOp.getShape().getN(), curOp.getShape().getK(),
2616 *curOp.getMultiplicandAPtxType(), *curOp.getMultiplicandBPtxType(),
2618 curOp.getBlockScaleFormat(), curOp.getKind());
2620 return {intId, args};
2623LogicalResult MmaSpBlockScaleOp::verify() {
2625 if (!getOrderedMetadata()) {
2626 return emitOpError(
"'orderedMetadata' attribute is mandatory");
2634 if (m == 16 && n == 8 && k == 128) {
2635 if (getMultiplicandAPtxType() != NVVM::MMATypes::e2m1 ||
2636 getMultiplicandBPtxType() != NVVM::MMATypes::e2m1)
2638 "unsupported MMATypes attribute for mma.m16n8k128.(mxf4nvf4|mxf4)");
2639 if (getKind() == NVVM::MMABlockScaleKind::MXF4) {
2640 if (getScaleVecSize() != NVVM::ScaleVecSize::X2)
2642 "unsupported ScaleVecSize attribute for mma.m16n8k128.mxf4");
2643 if (getBlockScaleFormat() != NVVM::BlockScaleFormat::UE8M0)
2645 "unsupported BlockScaleFormat attribute for mma.m16n8k128.mxf4");
2646 }
else if (getKind() == NVVM::MMABlockScaleKind::MXF4NVF4) {
2647 if (!((getScaleVecSize() == NVVM::ScaleVecSize::X2 &&
2648 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0) ||
2649 (getScaleVecSize() == NVVM::ScaleVecSize::X4 &&
2650 (getBlockScaleFormat() == NVVM::BlockScaleFormat::UE4M3 ||
2651 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))))
2653 "attributes for mma.m16n8k128.mxf4nvf4");
2657 }
else if (m == 16 && n == 8 && k == 64) {
2658 if (!(getKind() == NVVM::MMABlockScaleKind::MXF8F6F4 &&
2659 getScaleVecSize() == NVVM::ScaleVecSize::X1 &&
2660 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))
2662 emitOpError(
"unsupported Kind, ScaleVecSize and BlockScaleFormat "
2663 "attributes for mma.m16n8k64");
2670LogicalResult ShflOp::verify() {
2671 auto returnStructType = llvm::dyn_cast<LLVM::LLVMStructType>(
getType());
2673 auto verifyTypeError = [&](Twine desc,
Type expectedType,
2674 Type actualType) -> LogicalResult {
2675 return emitOpError(
"expected " + desc +
" to be of type ")
2676 << expectedType <<
" but got " << actualType <<
" instead";
2679 if (returnStructType) {
2680 if (!getReturnValueAndIsValid())
2681 return emitOpError(
"\"return_value_and_is_valid\" attribute must be "
2682 "specified when the return type is a struct type");
2684 if (returnStructType.getBody().size() != 2)
2685 return emitOpError(
"expected return type to be a two-element struct");
2688 auto resultType = returnStruct[0];
2689 if (resultType != getVal().
getType())
2690 return verifyTypeError(
"first element in the returned struct",
2691 getVal().
getType(), resultType);
2693 auto predicateType = returnStruct[1];
2694 if (!predicateType.isInteger(1))
2695 return verifyTypeError(
"second element in the returned struct",
2699 if (getReturnValueAndIsValid())
2700 return emitOpError(
"expected return type to be a two-element struct");
2703 return verifyTypeError(
"return type", getVal().
getType(),
getType());
2709ShflOp::inferReturnTypes(
MLIRContext *context, std::optional<Location> location,
2710 ShflOp::Adaptor adaptor,
2712 Type valType = adaptor.getVal().getType();
2713 if (adaptor.getReturnValueAndIsValid())
2714 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2715 context, {valType, IntegerType::get(context, 1)}));
2717 inferredReturnTypes.push_back(valType);
2722 NVVM::MMAFrag frag,
int nRow,
2725 unsigned numberElements = 0;
2728 Type f16x2 = VectorType::get(2, builder.getF16Type());
2729 if (type == NVVM::MMATypes::f16) {
2730 elementType = f16x2;
2731 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)
2735 }
else if (type == NVVM::MMATypes::f32) {
2736 elementType = builder.getF32Type();
2738 }
else if (type == NVVM::MMATypes::f64) {
2739 elementType = builder.getF64Type();
2740 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)
2744 }
else if (type == NVVM::MMATypes::tf32) {
2745 elementType = builder.getI32Type();
2747 }
else if (type == NVVM::MMATypes::s8 || type == NVVM::MMATypes::u8) {
2748 elementType = builder.getI32Type();
2749 int parallelSize = 0;
2750 if (frag == NVVM::MMAFrag::a)
2751 parallelSize = nRow;
2752 if (frag == NVVM::MMAFrag::b)
2753 parallelSize = nCol;
2756 if (parallelSize == 16)
2759 else if (parallelSize == 8)
2761 else if (parallelSize == 32)
2763 }
else if (type == NVVM::MMATypes::s32) {
2764 elementType = builder.getI32Type();
2767 assert(numberElements != 0 && elementType !=
nullptr);
2768 return std::make_pair(elementType, numberElements);
2771static std::pair<mlir::Type, unsigned>
2775 if (frag == NVVM::MMAFrag::a) {
2778 }
else if (frag == NVVM::MMAFrag::b) {
2785 assert(nRow && nCol);
2789LogicalResult NVVM::WMMALoadOp::verify() {
2790 unsigned addressSpace =
2791 llvm::cast<LLVM::LLVMPointerType>(getPtr().
getType()).getAddressSpace();
2792 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&
2793 addressSpace != NVVMMemorySpace::Shared)
2794 return emitOpError(
"expected source pointer in memory "
2797 if (NVVM::WMMALoadOp::getIntrinsicID(
getM(),
getN(), getK(), getLayout(),
2798 getEltype(), getFrag()) == 0)
2799 return emitOpError() <<
"invalid attribute combination";
2804 if (typeInfo.first == f64Ty && typeInfo.second == 1) {
2806 return emitOpError(
"expected destination type to be f64");
2810 Type dstType = LLVM::LLVMStructType::getLiteral(
2813 return emitOpError(
"expected destination type is a structure of ")
2814 << typeInfo.second <<
" elements of type " << typeInfo.first;
2818LogicalResult NVVM::WMMAStoreOp::verify() {
2819 unsigned addressSpace =
2820 llvm::cast<LLVM::LLVMPointerType>(getPtr().
getType()).getAddressSpace();
2821 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&
2822 addressSpace != NVVMMemorySpace::Shared)
2823 return emitOpError(
"expected operands to be a source pointer in memory "
2826 if (NVVM::WMMAStoreOp::getIntrinsicID(
getM(),
getN(), getK(), getLayout(),
2828 return emitOpError() <<
"invalid attribute combination";
2831 if (getArgs().size() != typeInfo.second)
2832 return emitOpError() <<
"expected " << typeInfo.second <<
" data operands";
2833 if (llvm::any_of(getArgs(), [&typeInfo](
Value operands) {
2834 return operands.
getType() != typeInfo.first;
2836 return emitOpError() <<
"expected data operands of type " << typeInfo.first;
2840LogicalResult NVVM::WMMAMmaOp::verify() {
2841 if (NVVM::WMMAMmaOp::getIntrinsicID(
getM(),
getN(), getK(), getLayoutA(),
2842 getLayoutB(), getEltypeA(),
2844 return emitOpError() <<
"invalid attribute combination";
2852 arguments.append(typeInfoA.second, typeInfoA.first);
2853 arguments.append(typeInfoB.second, typeInfoB.first);
2854 arguments.append(typeInfoC.second, typeInfoC.first);
2855 unsigned numArgs = arguments.size();
2856 if (getArgs().size() != numArgs)
2857 return emitOpError() <<
"expected " << numArgs <<
" arguments";
2858 for (
unsigned i = 0; i < numArgs; i++) {
2859 if (getArgs()[i].
getType() != arguments[i])
2860 return emitOpError() <<
"expected argument " << i <<
" to be of type "
2863 Type dstType = LLVM::LLVMStructType::getLiteral(
2866 return emitOpError(
"expected destination type is a structure of ")
2867 << typeInfoC.second <<
" elements of type " << typeInfoC.first;
2871LogicalResult NVVM::LdMatrixOp::verify() {
2873 if (m == 8 && n == 8) {
2874 if (num != 1 && num != 2 && num != 4) {
2875 return emitOpError(
"expected num attribute to be 1, 2 or 4 for 8x8 "
2878 if (getEltType() != LdStMatrixEltType::B16) {
2879 return emitOpError(
"expected element type to be b16 for 8x8 matrix");
2881 }
else if (m == 8 && n == 16) {
2882 if (num != 1 && num != 2 && num != 4) {
2883 return emitOpError(
"expected num attribute to be 1, 2 or 4 for 8x16 "
2886 if (getLayout() != MMALayout::row) {
2887 return emitOpError(
"expected layout to be row for 8x16 matrix");
2889 if (getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&
2890 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {
2891 return emitOpError(
"expected element type to be b8x16.b4x16_p64 or "
2892 "b8x16.b6x16_p32 for 8x16 matrix");
2894 }
else if (m == 16 && n == 16) {
2895 if (num != 1 && num != 2) {
2896 return emitOpError(
"expected num attribute to be 1 or 2 for 16x16 "
2899 if (getLayout() != MMALayout::col) {
2900 return emitOpError(
"expected layout to be col for 16x16 matrix");
2902 if (getEltType() != LdStMatrixEltType::B8 &&
2903 getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&
2904 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {
2905 return emitOpError(
"expected element type to be b8, b8x16.b4x16_p64 or "
2906 "b8x16.b6x16_p32 for 16x16 matrix");
2909 return emitOpError(
"expected shape to be 8x8, 8x16 or 16x16");
2913 uint32_t numElements = (m == 16 && n == 16 ? num * 2 : num);
2914 if (numElements == 1 &&
getType() != i32)
2915 return emitOpError(
"expected destination type is i32");
2916 if (numElements == 2 || numElements == 4) {
2917 Type dstType = LLVM::LLVMStructType::getLiteral(
2920 return emitOpError(
"expected destination type is a structure of ")
2921 << numElements <<
" elements of type i32";
2927LogicalResult LdMatrixOp::inferReturnTypes(
2928 MLIRContext *context, std::optional<Location> location,
2930 uint32_t num = adaptor.getNum();
2931 uint32_t m = adaptor.getShape().getM();
2932 uint32_t n = adaptor.getShape().getN();
2933 uint32_t numElements = (m == 16 && n == 16) ? num * 2 : num;
2935 Type i32 = IntegerType::get(context, 32);
2936 if (numElements == 1)
2937 inferredReturnTypes.push_back(i32);
2939 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2944LogicalResult NVVM::StMatrixOp::verify() {
2945 int numMatrix = getSources().size();
2946 if (numMatrix != 1 && numMatrix != 2 && numMatrix != 4)
2947 return emitOpError(
"expected num attribute to be 1, 2 or 4");
2950 if (m == 8 && n == 8) {
2951 if (getEltType() != NVVM::LdStMatrixEltType::B16) {
2952 return emitOpError(
"expected element type to be B16 for 8x8 matrix");
2954 }
else if (m == 16 && n == 8) {
2955 if (getEltType() != NVVM::LdStMatrixEltType::B8) {
2956 return emitOpError(
"expected element type to be B8 for 16x8 matrix");
2958 if (getLayout() != NVVM::MMALayout::col) {
2959 return emitOpError(
"expected layout to be col for 16x8 matrix");
2962 return emitOpError(
"expected shape to be 8x8 or 16x8");
2968LogicalResult NVVM::MovMatrixOp::verify() {
2970 if (m != 8 || n != 8)
2972 if (getLayout() != NVVM::MMALayout::col)
2974 if (getEltType() != NVVM::LdStMatrixEltType::B16)
2975 return emitOpError(
"expected element type to be b16");
2980 if (typeA == NVVM::WGMMATypes::tf32)
2982 if (typeA == NVVM::WGMMATypes::f16 || typeA == NVVM::WGMMATypes::bf16)
2984 if (typeA == NVVM::WGMMATypes::s8 || typeA == NVVM::WGMMATypes::u8)
2986 if (typeA == NVVM::WGMMATypes::e4m3 || typeA == NVVM::WGMMATypes::e5m2)
2988 if (typeA == NVVM::WGMMATypes::b1)
2994 NVVM::WGMMATypes typeA,
2995 NVVM::WGMMATypes typeB) {
2997 case NVVM::WGMMATypes::f16:
2998 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
2999 typeB == NVVM::WGMMATypes::f16)
3002 case NVVM::WGMMATypes::tf32:
3003 if (typeD == NVVM::WGMMATypes::f32 && typeB == NVVM::WGMMATypes::tf32)
3006 case NVVM::WGMMATypes::u8:
3007 case NVVM::WGMMATypes::s8:
3008 if (typeD == NVVM::WGMMATypes::s32 &&
3009 (typeB == NVVM::WGMMATypes::u8 || typeB == NVVM::WGMMATypes::s8))
3012 case NVVM::WGMMATypes::b1:
3013 if (typeD == NVVM::WGMMATypes::s32 && typeB == NVVM::WGMMATypes::b1)
3016 case NVVM::WGMMATypes::bf16:
3017 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
3018 typeB == NVVM::WGMMATypes::bf16)
3021 case NVVM::WGMMATypes::e4m3:
3022 case NVVM::WGMMATypes::e5m2:
3023 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
3024 (typeB == NVVM::WGMMATypes::e5m2 || typeB == NVVM::WGMMATypes::e4m3))
3027 case WGMMATypes::f32:
3028 case WGMMATypes::s32:
3029 llvm_unreachable(
"unsupported input types");
3037 72, 80, 88, 96, 104, 112, 120, 128,
3038 136, 144, 152, 160, 168, 176, 184, 192,
3039 200, 208, 216, 224, 232, 240, 248, 256};
3041 80, 96, 112, 128, 144, 160,
3042 176, 192, 208, 224, 240, 256};
3044 case WGMMATypes::f16:
3045 case WGMMATypes::tf32:
3046 case WGMMATypes::bf16:
3047 case WGMMATypes::e4m3:
3048 case WGMMATypes::e5m2:
3049 if (llvm::is_contained(allowedN, sizeN))
3052 case WGMMATypes::u8:
3053 case WGMMATypes::s8:
3054 case WGMMATypes::b1:
3055 if (llvm::is_contained(allowedNshort, sizeN))
3058 case WGMMATypes::f32:
3059 case WGMMATypes::s32:
3060 llvm_unreachable(
"unsupported input types");
3066LogicalResult NVVM::WgmmaMmaAsyncOp::verify() {
3067 Value outValue = getResults();
3068 auto stype = dyn_cast<LLVM::LLVMStructType>(outValue.
getType());
3070 return emitOpError() <<
"expected results to be struct";
3071 int outputSize = stype.getBody().size();
3072 WGMMATypes typeD = getTypeD();
3073 WGMMATypes typeA = getTypeA();
3074 WGMMATypes typeB = getTypeB();
3076 for (
Type t : stype.getBody()) {
3077 if (t != stype.getBody().front())
3079 <<
"all elements in struct must be same type but there is " << t;
3082 if (typeD != WGMMATypes::f32 && typeD != WGMMATypes::f16 &&
3083 typeD != WGMMATypes::s32) {
3084 return emitOpError() <<
"does not support the given output type " << typeD;
3086 if (typeD == WGMMATypes::s32 &&
3087 (getScaleA() == WGMMAScaleIn::neg || getScaleB() == WGMMAScaleIn::neg)) {
3088 return emitOpError() <<
"has s32 output, scaleA and scaleB cannot be neg";
3092 return emitOpError() << typeD <<
" += " << typeA <<
" * " << typeB
3093 <<
", it is not supported.";
3103 return emitOpError() <<
"shape 'k' must be " << allowedK.value()
3104 <<
" for input type " << typeA;
3108 return emitOpError() <<
"has input type " << typeA <<
" n is set to "
3109 <<
getShape().getN() <<
", it is not supported.";
3116 if ((typeA != WGMMATypes::f16 && typeA != WGMMATypes::bf16) &&
3117 (getLayoutA() == mlir::NVVM::MMALayout::col ||
3118 getLayoutB() == mlir::NVVM::MMALayout::row)) {
3120 <<
"given layouts layout_a = " << getLayoutA()
3121 <<
" and layout_b = " << getLayoutB() <<
" for input types " << typeA
3123 <<
" requires transpose. However, this is only supported for: "
3124 << MMATypes::f16 <<
" and " << MMATypes::bf16;
3128 int expectedOutput = 0;
3129 if (typeD == WGMMATypes::f32 || typeD == WGMMATypes::s32)
3130 expectedOutput =
getShape().getN() / 2;
3131 if (typeD == WGMMATypes::f16)
3132 expectedOutput =
getShape().getN() / 4;
3133 if (outputSize != expectedOutput) {
3134 return emitOpError() <<
"results " << expectedOutput
3135 <<
", however output struct has " << outputSize
3139 if (typeD != WGMMATypes::s32 &&
3140 getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==
3141 NVVM::MMAIntOverflow::satfinite) {
3143 <<
" `satfinite` can be only used with s32 accumulator, however "
3144 "the current accumulator is "
3151std::string NVVM::WgmmaMmaAsyncOp::getPtx() {
3154 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;
3156 StringRef outputTypeName = stringifyWGMMATypes(getTypeD());
3158 int expectedOutputRegisters = 0;
3159 if (getTypeD() == WGMMATypes::f16)
3160 expectedOutputRegisters =
getShape().getN() / 4;
3162 expectedOutputRegisters =
getShape().getN() / 2;
3165 llvm::raw_string_ostream ss(ptx);
3170 << ((expectedOutputRegisters * 2) + 2)
3172 "wgmma.mma_async.sync.aligned.m"
3173 << m <<
"n" << n <<
"k" << k <<
"." << outputTypeName <<
"." << getTypeA()
3174 <<
"." << getTypeB();
3175 if (getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==
3176 NVVM::MMAIntOverflow::satfinite)
3180 for (; regCnt < expectedOutputRegisters; ++regCnt) {
3181 ss <<
"$" << regCnt;
3182 if (regCnt != expectedOutputRegisters - 1)
3188 regCnt = (regCnt * 2);
3189 ss <<
" $" << (regCnt) <<
"," <<
" $" << (regCnt + 1) <<
"," <<
" p";
3190 if (getTypeD() != WGMMATypes::s32) {
3191 ss <<
", $" << (regCnt + 3) <<
", $" << (regCnt + 4);
3195 ss <<
", $" << (regCnt + 5) <<
", $" << (regCnt + 6);
3202bool NVVM::WgmmaMmaAsyncOp::getAsmValues(
3206 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;
3213 asmValues.push_back({makeConstantI32(rewriter,
static_cast<int>(getScaleD())),
3215 if (getTypeD() != WGMMATypes::s32) {
3216 asmValues.push_back(
3217 {makeConstantI32(rewriter,
3218 getScaleA() == NVVM::WGMMAScaleIn::neg ? -1 : 1),
3220 asmValues.push_back(
3221 {makeConstantI32(rewriter,
3222 getScaleB() == NVVM::WGMMAScaleIn::neg ? -1 : 1),
3226 asmValues.push_back(
3227 {makeConstantI32(rewriter,
static_cast<int>(getLayoutA())),
3229 asmValues.push_back(
3230 {makeConstantI32(rewriter, 1 -
static_cast<int>(getLayoutB())),
3236LogicalResult NVVM::FenceProxyOp::verify() {
3237 if (getKind() == NVVM::ProxyKind::async_shared && !getSpace().has_value()) {
3238 return emitOpError() <<
"async_shared fence requires space attribute";
3240 if (getKind() != NVVM::ProxyKind::async_shared && getSpace().has_value()) {
3241 return emitOpError() <<
"only async_shared fence can have space attribute";
3246LogicalResult NVVM::FenceProxyAcquireOp::verify() {
3247 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
3248 return emitOpError(
"uni-directional proxies only support generic for "
3249 "from_proxy attribute");
3251 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)
3252 return emitOpError(
"uni-directional proxies only support tensormap "
3253 "for to_proxy attribute");
3257LogicalResult NVVM::FenceProxyReleaseOp::verify() {
3258 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
3259 return emitOpError(
"uni-directional proxies only support generic for "
3260 "from_proxy attribute");
3262 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)
3263 return emitOpError(
"uni-directional proxies only support tensormap "
3264 "for to_proxy attribute");
3268LogicalResult NVVM::FenceProxySyncRestrictOp::verify() {
3269 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
3270 return emitOpError(
"only generic is support for from_proxy attribute");
3272 if (getToProxy() != NVVM::ProxyKind::async)
3273 return emitOpError(
"only async is supported for to_proxy attribute");
3277LogicalResult NVVM::SetMaxRegisterOp::verify() {
3278 if (getRegCount() % 8)
3279 return emitOpError(
"new register size must be multiple of 8");
3280 if (getRegCount() < 24 || getRegCount() > 256)
3281 return emitOpError(
"new register size must be in between 24 to 256");
3285LogicalResult NVVM::Tcgen05CpOp::verify() {
3286 auto mc = getMulticast();
3288 using SH = Tcgen05CpShape;
3289 using MC = Tcgen05CpMulticast;
3291 case SH::SHAPE_128x256b:
3292 case SH::SHAPE_128x128b:
3293 case SH::SHAPE_4x256b:
3295 return emitError(
"Invalid multicast type for tcgen05.cp Op");
3297 case SH::SHAPE_64x128b:
3298 if (mc != MC::WARPX2_01_23 && mc != MC::WARPX2_02_13)
3299 return emitError(
"Shape 64x128b requires multicast warpx2_01_23 or "
3300 "warpx2_02_13 for tcgen05.cp Op");
3302 case SH::SHAPE_32x128b:
3303 if (mc != MC::WARPX4)
3305 "Shape 32x128b requires multicast warpx4 for tcgen05.cp Op");
3311LogicalResult NVVM::MatchSyncOp::verify() {
3312 if (getKind() == NVVM::MatchSyncKind::all) {
3313 auto type = llvm::dyn_cast<LLVM::LLVMStructType>(
getType());
3314 if (!type || type.getBody().size() != 2 ||
3315 !type.getBody()[0].isInteger(32) || !type.getBody()[1].isInteger(1)) {
3316 return emitOpError(
"match.sync 'all' returns a two element struct with "
3317 "first element as i32 and second element as i1");
3320 if (!
getType().isInteger(32)) {
3321 return emitOpError(
"match.sync 'any' returns an i32");
3327LogicalResult MatchSyncOp::inferReturnTypes(
3328 MLIRContext *context, std::optional<Location> location,
3330 if (adaptor.getKind() == NVVM::MatchSyncKind::all)
3331 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
3333 {IntegerType::get(context, 32), IntegerType::get(context, 1)}));
3335 inferredReturnTypes.push_back(IntegerType::get(context, 32));
3339LogicalResult NVVM::VoteSyncOp::verify() {
3340 if (getKind() == NVVM::VoteSyncKind::ballot) {
3341 if (!
getType().isInteger(32)) {
3342 return emitOpError(
"vote.sync 'ballot' returns an i32");
3345 if (!
getType().isInteger(1)) {
3346 return emitOpError(
"vote.sync 'any', 'all' and 'uni' returns an i1");
3352LogicalResult VoteSyncOp::inferReturnTypes(
3353 MLIRContext *context, std::optional<Location> location,
3355 unsigned width = adaptor.getKind() == NVVM::VoteSyncKind::ballot ? 32 : 1;
3356 inferredReturnTypes.push_back(IntegerType::get(context, width));
3360LogicalResult NVVM::PrefetchOp::verify() {
3361 using MemSpace = NVVM::NVVMMemorySpace;
3362 using CacheLevel = NVVM::PrefetchCacheLevel;
3364 unsigned addressSpace =
3365 llvm::cast<LLVM::LLVMPointerType>(getAddr().
getType()).getAddressSpace();
3366 std::optional<NVVM::CacheEvictionPriority> evictPriority = getEvictPriority();
3367 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = getCacheLevel();
3369 if (getTensormap() && cacheLevel)
3370 return emitOpError(
"cannot specify both tensormap and cache level");
3372 if (getTensormap()) {
3373 if (addressSpace != MemSpace::Generic &&
3374 addressSpace != MemSpace::Constant) {
3376 "prefetch tensormap requires a generic or constant pointer");
3379 if (evictPriority) {
3381 "prefetch tensormap does not support eviction priority");
3384 if (getInParamSpace() && addressSpace != MemSpace::Generic) {
3386 "in_param_space can only be specified for a generic pointer");
3389 }
else if (cacheLevel) {
3390 if (addressSpace != MemSpace::Generic && addressSpace != MemSpace::Global &&
3391 addressSpace != MemSpace::Local) {
3392 return emitOpError(
"prefetch to cache level requires a generic, global, "
3393 "or local pointer");
3397 if (*cacheLevel != CacheLevel::L1) {
3399 "unsupported cache level, the only supported uniform "
3400 "cache level is L1");
3403 if (addressSpace != MemSpace::Generic) {
3405 "prefetch to uniform cache requires a generic pointer");
3409 if (evictPriority) {
3410 if (*cacheLevel != CacheLevel::L2)
3412 "cache eviction priority supported only for cache level L2");
3414 if (addressSpace != MemSpace::Global)
3415 return emitOpError(
"cache eviction priority requires a global pointer");
3417 if (*evictPriority != NVVM::CacheEvictionPriority::EvictNormal &&
3418 *evictPriority != NVVM::CacheEvictionPriority::EvictLast)
3420 "unsupported cache eviction priority, only evict_last and "
3421 "evict_normal are supported");
3425 return emitOpError(
"predicate supported only on prefetch tensormap");
3429 "requires specification of either cache level or tensormap");
3435LogicalResult NVVM::ClusterLaunchControlQueryCancelOp::verify() {
3436 switch (getQueryType()) {
3437 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:
3439 return emitOpError(
"is_canceled query type returns an i1");
3441 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:
3442 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:
3443 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:
3444 if (!
getType().isInteger(32)) {
3445 return emitOpError(
"get_first_cta_id_x, get_first_cta_id_y, "
3446 "get_first_cta_id_z query types return an i32");
3453LogicalResult ClusterLaunchControlQueryCancelOp::inferReturnTypes(
3454 MLIRContext *context, std::optional<Location> location,
3455 ClusterLaunchControlQueryCancelOp::Adaptor adaptor,
3458 adaptor.getQueryType() == NVVM::ClusterLaunchControlQueryType::IS_CANCELED
3461 inferredReturnTypes.push_back(IntegerType::get(context, width));
3465LogicalResult NVVM::ReduxOp::verify() {
3468 if (!reduxType.
isF32()) {
3470 return emitOpError(
"abs attribute is supported only for f32 type");
3472 return emitOpError(
"nan attribute is supported only for f32 type");
3475 NVVM::ReductionKind kind = getKind();
3477 case NVVM::ReductionKind::ADD:
3478 case NVVM::ReductionKind::AND:
3479 case NVVM::ReductionKind::OR:
3480 case NVVM::ReductionKind::XOR:
3481 case NVVM::ReductionKind::MAX:
3482 case NVVM::ReductionKind::MIN:
3483 case NVVM::ReductionKind::UMAX:
3484 case NVVM::ReductionKind::UMIN:
3487 << kind <<
"' reduction kind unsupported with " << reduxType
3488 <<
" type. Only supported type is 'i32'.";
3490 case NVVM::ReductionKind::FMIN:
3491 case NVVM::ReductionKind::FMAX:
3492 if (!reduxType.isF32())
3494 << kind <<
"' reduction kind unsupported with " << reduxType
3495 <<
" type. Only supported type is 'f32'.";
3502LogicalResult NVVM::TensormapReplaceOp::verify() {
3503 auto ord = getOrd();
3504 Value newVal = getNewValue();
3505 auto newValAttr = getNewValueAttr();
3506 auto fieldName = stringifyEnum(getField());
3508 if (ord && !llvm::is_contained({NVVM::TensormapField::BOX_DIM,
3509 NVVM::TensormapField::GLOBAL_DIM,
3510 NVVM::TensormapField::GLOBAL_STRIDE,
3511 NVVM::TensormapField::ELEMENT_STRIDE},
3513 return emitOpError(
"ordinal is not supported for ")
3514 << fieldName <<
" field";
3516 auto invalidNewVal = [&](llvm::Twine type) -> std::string {
3517 return llvm::Twine(
"new_value must be specified and must be an " + type +
3518 " for " + llvm::Twine(fieldName) +
" field")
3522 auto invalidNewValAttr = [&]() -> std::string {
3523 return (llvm::Twine(
3524 "new_value_attr must be specified and must be a valid ") +
3525 llvm::Twine(fieldName) +
" attribute for " + fieldName +
" field")
3529 switch (getField()) {
3530 case NVVM::TensormapField::GLOBAL_ADDRESS:
3534 case NVVM::TensormapField::RANK:
3538 case NVVM::TensormapField::GLOBAL_STRIDE:
3540 return emitOpError(
"ordinal is required for global_stride field");
3544 case NVVM::TensormapField::BOX_DIM:
3545 case NVVM::TensormapField::GLOBAL_DIM:
3546 case NVVM::TensormapField::ELEMENT_STRIDE:
3549 << stringifyEnum(getField()) <<
" field";
3553 case NVVM::TensormapField::ELEMTYPE:
3554 if (!(newValAttr && llvm::isa<TensormapElemtypeAttr>(*newValAttr)))
3557 case NVVM::TensormapField::INTERLEAVE_LAYOUT:
3558 if (!(newValAttr && llvm::isa<TensormapInterleaveLayoutAttr>(*newValAttr)))
3561 case NVVM::TensormapField::SWIZZLE_MODE:
3562 if (!(newValAttr && llvm::isa<TensormapSwizzleModeAttr>(*newValAttr)))
3565 case NVVM::TensormapField::SWIZZLE_ATOMICITY:
3566 if (!(newValAttr && llvm::isa<TensormapSwizzleAtomicityAttr>(*newValAttr)))
3569 case NVVM::TensormapField::FILL_MODE:
3570 if (!(newValAttr && llvm::isa<TensormapFillModeAttr>(*newValAttr)))
3578template <
typename OpType>
3580 mlir::NVVM::FPRoundingMode rndMode = op.getRnd();
3581 mlir::NVVM::SaturationMode satMode = op.getSat();
3582 bool isFTZ = op.getFtz();
3585 mlir::Type opBaseType = isa<VectorType>(opType)
3586 ? cast<VectorType>(opType).getElementType()
3589 if (opBaseType.
isF64() && (satMode != NVVM::SaturationMode::NONE || isFTZ))
3590 return op.emitOpError(
"FTZ and saturation are not supported for "
3591 "additions/subtractions involving f64 type");
3593 if (opBaseType.
isF16() && !(rndMode == NVVM::FPRoundingMode::RN ||
3594 rndMode == NVVM::FPRoundingMode::NONE))
3595 return op.emitOpError(
"only RN rounding mode is supported for f16 and "
3596 "vector<2xf16> additions/subtractions");
3598 if (opBaseType.
isBF16()) {
3599 if (rndMode != NVVM::FPRoundingMode::RN &&
3600 rndMode != NVVM::FPRoundingMode::NONE)
3601 return op.emitOpError(
"only RN rounding mode is supported for bf16 and "
3602 "vector<2xbf16> additions/subtractions");
3603 if (satMode != NVVM::SaturationMode::NONE || isFTZ)
3604 return op.emitOpError(
"FTZ and saturation are not supported for bf16 and "
3605 "vector<2xbf16> additions/subtractions");
3612 if (opBaseType.
isF16() && isFTZ && satMode == NVVM::SaturationMode::NONE)
3613 return op.emitOpError(
"FTZ with no saturation is not supported for f16 and "
3614 "vector<2xf16> additions/subtractions");
3623LogicalResult NVVM::FmaOp::verify() {
3624 auto opType = getRes().getType();
3625 mlir::NVVM::FPRoundingMode rndMode = getRnd();
3626 mlir::NVVM::SaturationMode satMode = getSat();
3627 bool isFTZ = getFtz();
3628 bool isRelu = getRelu();
3629 bool hasOOB = getOob();
3631 auto getBaseFType = [](
Type type) ->
Type {
3632 if (isa<VectorType>(type))
3633 return cast<VectorType>(type).getElementType();
3637 auto opBaseType = getBaseFType(opType);
3639 if (rndMode == NVVM::FPRoundingMode::NONE)
3640 return emitOpError(
"rounding mode must be specified");
3642 if (isRelu && satMode == NVVM::SaturationMode::SAT)
3643 return emitOpError(
"relu and saturation are not supported together");
3645 if (hasOOB && (satMode == NVVM::SaturationMode::SAT || isFTZ))
3646 return emitOpError(
"oob is not supported with saturation or FTZ");
3648 if (!(opBaseType.isF16() || opBaseType.isBF16()) && (isRelu || hasOOB))
3649 return emitOpError(
"relu and oob are only supported for f16 and bf16");
3651 if (opBaseType.isF64() && (satMode != NVVM::SaturationMode::NONE || isFTZ))
3652 return emitOpError(
"FTZ and saturation are not supported for f64 type");
3654 if (opBaseType.isF16() && rndMode != NVVM::FPRoundingMode::RN)
3656 "only RN rounding mode is supported for f16 and vector<2xf16>");
3658 if (opBaseType.isBF16()) {
3659 if (rndMode != NVVM::FPRoundingMode::RN)
3661 "only RN rounding mode is supported for bf16 and vector<2xbf16>");
3662 if (satMode != NVVM::SaturationMode::NONE || isFTZ)
3664 "FTZ and saturation are not supported for bf16 and vector<2xbf16>");
3670LogicalResult NVVM::SqrtOp::verify() {
3671 if (getRnd() == NVVM::FPRoundingMode::NONE)
3672 return emitOpError(
"rounding mode cannot be None");
3674 if (getRes().
getType().isF64() && getFtz())
3675 return emitOpError(
"FTZ is not supported for f64");
3680LogicalResult NVVM::DivFOp::verify() {
3681 bool isApprox = getApprox();
3682 bool isFull = getFull();
3683 bool isF64 = getRes().getType().isF64();
3684 bool isFtz = getFtz();
3685 NVVM::FPRoundingMode rndMode = getRnd();
3687 if (isApprox && isFull)
3688 return emitOpError(
"'approx' and 'full' are mutually exclusive");
3690 if (isApprox || isFull) {
3692 return emitOpError(
"'approx' and 'full' forms are f32-only");
3693 if (rndMode != NVVM::FPRoundingMode::NONE)
3695 "'approx' and 'full' forms do not accept a rounding mode");
3700 if (rndMode == NVVM::FPRoundingMode::NONE)
3701 return emitOpError(
"rounding mode cannot be None for the rounded divide");
3703 return emitOpError(
"FTZ is not supported for f64");
3714 unsigned sizeInBits,
3716 field = builder.CreateZExtOrBitCast(field, builder.getInt32Ty());
3718 unsigned mask = (sizeInBits < 32 ? ((1u << sizeInBits) - 1) : 0xffffffffu);
3719 if (mask != 0xffffffffu)
3720 field = builder.CreateAnd(field, builder.getInt32(mask));
3722 field = builder.CreateZExtOrBitCast(field, builder.getInt64Ty());
3723 field = builder.CreateShl(field, start);
3725 return builder.CreateOr(
result, field);
3728void Tcgen05MmaSmemDescOp::createSmemDescriptor(
Operation &op,
3730 llvm::IRBuilderBase &builder) {
3731 auto thisOp = cast<NVVM::Tcgen05MmaSmemDescOp>(op);
3732 llvm::Value *smemDesc = builder.getInt64(0);
3737 builder, smemDesc, mt.
lookupValue(thisOp.getLeadingDimOffset()), 14, 16);
3739 builder, smemDesc, mt.
lookupValue(thisOp.getStrideDimOffset()), 14, 32);
3745 builder, smemDesc, mt.
lookupValue(thisOp.getLeadingDimMode()), 1, 52);
3749 mt.
mapValue(thisOp.getRes()) = smemDesc;
3756std::string NVVM::MBarrierInitOp::getPtx() {
3758 return isShared ? std::string(
"mbarrier.init.shared.b64 [%0], %1;")
3759 : std::string(
"mbarrier.init.b64 [%0], %1;");
3762std::string NVVM::MBarrierArriveExpectTxOp::getPtx() {
3765 ? std::string(
"mbarrier.arrive.expect_tx.shared.b64 _, [%0], %1;")
3766 : std::string(
"mbarrier.arrive.expect_tx.b64 _, [%0], %1;");
3769std::string NVVM::MBarrierTryWaitParityOp::getPtx() {
3771 llvm::StringRef space = isShared ?
".shared" :
"";
3773 return llvm::formatv(
"{\n\t"
3774 ".reg .pred P1; \n\t"
3776 "mbarrier.try_wait.parity{0}.b64 P1, [%0], %1, %2; \n\t"
3777 "@P1 bra.uni DONE; \n\t"
3778 "bra.uni LAB_WAIT; \n\t"
3795 LLVM::FNegOp::create(rewriter, loc, op.getRhs().getType(), op.getRhs());
3798 op.getRnd(), op.getSat(), op.getFtz());
3817 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_count
3818 : llvm::Intrinsic::nvvm_barrier_cta_sync_count;
3820 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_all
3821 : llvm::Intrinsic::nvvm_barrier_cta_sync_all;
3826static llvm::Intrinsic::ID
3829 case NVVM::BarrierReduction::AND:
3830 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_and_aligned_all
3831 : llvm::Intrinsic::nvvm_barrier_cta_red_and_all;
3832 case NVVM::BarrierReduction::OR:
3833 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_or_aligned_all
3834 : llvm::Intrinsic::nvvm_barrier_cta_red_or_all;
3835 case NVVM::BarrierReduction::POPC:
3836 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_popc_aligned_all
3837 : llvm::Intrinsic::nvvm_barrier_cta_red_popc_all;
3839 llvm_unreachable(
"unknown BarrierReduction kind");
3844 auto thisOp = cast<NVVM::BarrierOp>(op);
3845 llvm::Value *barrierId = thisOp.getBarrierId()
3847 : builder.getInt32(0);
3848 bool hasCount =
static_cast<bool>(thisOp.getNumberOfThreads());
3849 llvm::Intrinsic::ID
id =
3853 args.push_back(mt.
lookupValue(thisOp.getNumberOfThreads()));
3854 return {id, std::move(args)};
3859 auto thisOp = cast<NVVM::BarrierArriveOp>(op);
3860 llvm::Value *barrierId = thisOp.getBarrierId()
3862 : builder.getInt32(0);
3863 llvm::Value *numThreads = mt.
lookupValue(thisOp.getNumberOfThreads());
3864 llvm::Intrinsic::ID
id =
3866 ? llvm::Intrinsic::nvvm_barrier_cta_arrive_aligned_count
3867 : llvm::Intrinsic::nvvm_barrier_cta_arrive_count;
3868 return {id, {barrierId, numThreads}};
3873 auto thisOp = cast<NVVM::BarrierReductionOp>(op);
3875 thisOp.getAligned(), thisOp.getReductionOp());
3876 llvm::Value *barrierId = thisOp.getBarrierId()
3878 : builder.getInt32(0);
3881 builder.CreateICmpNE(mt.
lookupValue(thisOp.getReductionPredicate()),
3882 builder.getInt32(0))};
3883 return {id, std::move(args)};
3888 llvm::IRBuilderBase &builder) {
3889 auto thisOp = cast<NVVM::CosOp>(op);
3890 llvm::Intrinsic::ID
id = thisOp.getFtz()
3891 ? llvm::Intrinsic::nvvm_cos_approx_ftz_f
3892 : llvm::Intrinsic::nvvm_cos_approx_f;
3898 llvm::IRBuilderBase &builder) {
3899 auto thisOp = cast<NVVM::SinOp>(op);
3900 llvm::Intrinsic::ID
id = thisOp.getFtz()
3901 ? llvm::Intrinsic::nvvm_sin_approx_ftz_f
3902 : llvm::Intrinsic::nvvm_sin_approx_f;
3908 llvm::IRBuilderBase &builder) {
3909 auto thisOp = cast<NVVM::Log2Op>(op);
3910 llvm::Intrinsic::ID
id = thisOp.getFtz()
3911 ? llvm::Intrinsic::nvvm_lg2_approx_ftz_f
3912 : llvm::Intrinsic::nvvm_lg2_approx_f;
3918 llvm::IRBuilderBase &builder) {
3919 auto thisOp = cast<NVVM::Ex2Op>(op);
3920 llvm::Intrinsic::ID
id = thisOp.getFtz()
3921 ? llvm::Intrinsic::nvvm_ex2_approx_ftz
3922 : llvm::Intrinsic::nvvm_ex2_approx;
3928 llvm::IRBuilderBase &builder) {
3929 auto thisOp = cast<NVVM::RsqrtOp>(op);
3930 Type t = thisOp.getRes().getType();
3931 bool isFtz = thisOp.getFtz();
3933 llvm::Intrinsic::ID
id = [&] {
3935 return isFtz ? llvm::Intrinsic::nvvm_rsqrt_approx_ftz_f
3936 : llvm::Intrinsic::nvvm_rsqrt_approx_f;
3939 return isFtz ? llvm::Intrinsic::nvvm_rsqrt_approx_ftz_d
3940 : llvm::Intrinsic::nvvm_rsqrt_approx_d;
3948 llvm::IRBuilderBase &builder) {
3949 auto thisOp = cast<NVVM::SqrtOp>(op);
3950 Type t = thisOp.getRes().getType();
3951 NVVM::FPRoundingMode rndMode = thisOp.getRnd();
3952 bool isFtz = thisOp.getFtz();
3956 unsigned rndIndex =
static_cast<unsigned>(rndMode) - 1;
3958 static constexpr llvm::Intrinsic::ID f32IDs[] = {
3959 llvm::Intrinsic::nvvm_sqrt_rn_f,
3960 llvm::Intrinsic::nvvm_sqrt_rm_f,
3961 llvm::Intrinsic::nvvm_sqrt_rp_f,
3962 llvm::Intrinsic::nvvm_sqrt_rz_f,
3964 static constexpr llvm::Intrinsic::ID f32FTZIDs[] = {
3965 llvm::Intrinsic::nvvm_sqrt_rn_ftz_f,
3966 llvm::Intrinsic::nvvm_sqrt_rm_ftz_f,
3967 llvm::Intrinsic::nvvm_sqrt_rp_ftz_f,
3968 llvm::Intrinsic::nvvm_sqrt_rz_ftz_f,
3970 static constexpr llvm::Intrinsic::ID f64IDs[] = {
3971 llvm::Intrinsic::nvvm_sqrt_rn_d,
3972 llvm::Intrinsic::nvvm_sqrt_rm_d,
3973 llvm::Intrinsic::nvvm_sqrt_rp_d,
3974 llvm::Intrinsic::nvvm_sqrt_rz_d,
3977 llvm::Intrinsic::ID
id =
3978 t.
isF32() ? (isFtz ? f32FTZIDs[rndIndex] : f32IDs[rndIndex])
3986 llvm::IRBuilderBase &builder) {
3987 auto thisOp = cast<NVVM::SqrtApproxOp>(op);
3988 llvm::Intrinsic::ID
id = thisOp.getFtz()
3989 ? llvm::Intrinsic::nvvm_sqrt_approx_ftz_f
3990 : llvm::Intrinsic::nvvm_sqrt_approx_f;
3996 llvm::IRBuilderBase &builder) {
3997 auto thisOp = cast<NVVM::DivFOp>(op);
3998 bool isFtz = thisOp.getFtz();
4000 llvm::Intrinsic::ID id;
4002 if (thisOp.getApprox()) {
4003 id = isFtz ? llvm::Intrinsic::nvvm_div_approx_ftz_f
4004 : llvm::Intrinsic::nvvm_div_approx_f;
4005 }
else if (thisOp.getFull()) {
4008 id = isFtz ? llvm::Intrinsic::nvvm_div_full_ftz
4009 : llvm::Intrinsic::nvvm_div_full;
4012 unsigned rndIndex =
static_cast<unsigned>(thisOp.getRnd()) - 1;
4014 static constexpr llvm::Intrinsic::ID f32IDs[] = {
4015 llvm::Intrinsic::nvvm_div_rn_f,
4016 llvm::Intrinsic::nvvm_div_rm_f,
4017 llvm::Intrinsic::nvvm_div_rp_f,
4018 llvm::Intrinsic::nvvm_div_rz_f,
4020 static constexpr llvm::Intrinsic::ID f32FTZIDs[] = {
4021 llvm::Intrinsic::nvvm_div_rn_ftz_f,
4022 llvm::Intrinsic::nvvm_div_rm_ftz_f,
4023 llvm::Intrinsic::nvvm_div_rp_ftz_f,
4024 llvm::Intrinsic::nvvm_div_rz_ftz_f,
4026 static constexpr llvm::Intrinsic::ID f64IDs[] = {
4027 llvm::Intrinsic::nvvm_div_rn_d,
4028 llvm::Intrinsic::nvvm_div_rm_d,
4029 llvm::Intrinsic::nvvm_div_rp_d,
4030 llvm::Intrinsic::nvvm_div_rz_d,
4032 Type t = thisOp.getRes().getType();
4033 id = t.
isF32() ? (isFtz ? f32FTZIDs[rndIndex] : f32IDs[rndIndex])
4044 auto thisOp = cast<NVVM::AsyncStoreGlobalOp>(op);
4045 mlir::NVVM::MemScopeKind scope = thisOp.getScope();
4046 bool isMmio = thisOp.getMmio();
4048 llvm::Value *addr = mt.
lookupValue(thisOp.getAddr());
4049 llvm::Value *value = mt.
lookupValue(thisOp.getValue());
4050 llvm::Value *isMultimem = builder.getInt1(thisOp.getMultimem());
4052 if (scope == MemScopeKind::SYS) {
4053 return isMmio ?
IDArgPair(llvm::Intrinsic::nvvm_st_async_mmio_sys,
4056 {addr, value, isMultimem});
4057 }
else if (scope == MemScopeKind::GPU) {
4058 return IDArgPair(llvm::Intrinsic::nvvm_st_async_gpu,
4059 {addr, value, isMultimem});
4061 llvm_unreachable(
"unsupported scope for AsyncStoreGlobalOp");
4066 llvm::IRBuilderBase &builder) {
4067 auto thisOp = cast<NVVM::PMEventOp>(op);
4071 llvm::Value *maskVal;
4072 if (
auto eventAttr = thisOp.getEventIdAttr()) {
4073 uint16_t mask =
static_cast<uint16_t
>(1u << eventAttr.getInt());
4074 maskVal = llvm::ConstantInt::get(i16Ty, mask);
4077 llvm::ConstantInt::get(i16Ty, thisOp.getMaskedEventIdAttr().getValue());
4080 return {llvm::Intrinsic::nvvm_pm_event_mask, {maskVal}};
4085 auto thisOp = cast<NVVM::MBarrierInitOp>(op);
4087 llvm::Intrinsic::ID
id = isShared ? llvm::Intrinsic::nvvm_mbarrier_init_shared
4088 : llvm::Intrinsic::nvvm_mbarrier_init;
4093 args.push_back(mt.
lookupValue(thisOp.getCount()));
4095 return {id, std::move(args)};
4100 auto thisOp = cast<NVVM::MBarrierInvalOp>(op);
4102 llvm::Intrinsic::ID
id = isShared
4103 ? llvm::Intrinsic::nvvm_mbarrier_inval_shared
4104 : llvm::Intrinsic::nvvm_mbarrier_inval;
4111 auto thisOp = cast<NVVM::MBarrierExpectTxOp>(op);
4114 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4117 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4119 static constexpr llvm::Intrinsic::ID IDs[] = {
4120 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cta,
4121 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cluster,
4122 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cta,
4123 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cluster};
4128 args.push_back(mt.
lookupValue(thisOp.getTxcount()));
4130 return {IDs[
index], std::move(args)};
4135 auto thisOp = cast<NVVM::MBarrierCompleteTxOp>(op);
4138 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4141 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4143 static constexpr llvm::Intrinsic::ID IDs[] = {
4144 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cta,
4145 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cluster,
4146 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cta,
4147 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cluster};
4152 args.push_back(mt.
lookupValue(thisOp.getTxcount()));
4154 return {IDs[
index], std::move(args)};
4159 auto thisOp = cast<NVVM::MBarrierArriveOp>(op);
4162 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4165 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4167 static constexpr llvm::Intrinsic::ID IDs[] = {
4168 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta,
4169 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cluster,
4170 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cta,
4171 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cluster};
4172 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4173 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cta,
4174 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cluster,
4175 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cta,
4177 nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cluster};
4178 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4182 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4189 bool hasCount =
static_cast<bool>(thisOp.getCount());
4191 (
id == llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta))
4192 return {llvm::Intrinsic::nvvm_mbarrier_arrive_shared, {mbar}};
4196 llvm::Value *count =
4198 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);
4199 return {id, {mbar, count}};
4204 auto thisOp = cast<NVVM::MBarrierArriveDropOp>(op);
4207 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4210 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4212 static constexpr llvm::Intrinsic::ID IDs[] = {
4213 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cta,
4214 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cluster,
4215 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cta,
4216 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cluster};
4217 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4218 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cta,
4220 nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cluster,
4222 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cta,
4224 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cluster};
4225 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4229 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4235 bool hasCount =
static_cast<bool>(thisOp.getCount());
4236 llvm::Value *count =
4238 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);
4240 return {id, {mbar, count}};
4243bool MBarrierArriveExpectTxOp::getAsmValues(
4250 for (
auto val : getOperands())
4258 auto thisOp = cast<NVVM::MBarrierArriveExpectTxOp>(op);
4261 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4264 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4267 static constexpr llvm::Intrinsic::ID IDs[] = {
4268 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cta,
4269 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cluster,
4270 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cta,
4271 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cluster};
4272 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4273 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cta,
4274 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cluster,
4275 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cta,
4276 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cluster};
4278 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4281 llvm::Value *txcount = mt.
lookupValue(thisOp.getTxcount());
4282 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4287 return {id, {mbar, txcount}};
4292 auto thisOp = cast<NVVM::MBarrierArriveDropExpectTxOp>(op);
4295 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4298 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4301 static constexpr llvm::Intrinsic::ID IDs[] = {
4302 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cta,
4303 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cluster,
4304 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cta,
4305 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cluster};
4306 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4307 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cta,
4308 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cluster,
4309 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cta,
4310 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cluster};
4312 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4315 llvm::Value *txcount = mt.
lookupValue(thisOp.getTxcount());
4316 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4321 return {id, {mbar, txcount}};
4326 auto thisOp = cast<NVVM::MBarrierArriveNocompleteOp>(op);
4328 llvm::Intrinsic::ID
id =
4329 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete_shared
4330 : llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete;
4334 args.push_back(mt.
lookupValue(thisOp.getCount()));
4336 return {id, std::move(args)};
4341 auto thisOp = cast<NVVM::MBarrierArriveDropNocompleteOp>(op);
4343 llvm::Intrinsic::ID
id =
4344 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete_shared
4345 : llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete;
4349 args.push_back(mt.
lookupValue(thisOp.getCount()));
4351 return {id, std::move(args)};
4356 auto thisOp = cast<NVVM::MBarrierTestWaitOp>(op);
4357 bool isPhaseParity = thisOp.getStateOrPhase().getType().isInteger(32);
4358 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4361 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isPhaseParity ? 1 : 0);
4364 static constexpr llvm::Intrinsic::ID IDs[] = {
4365 llvm::Intrinsic::nvvm_mbarrier_test_wait_scope_cta_space_cta,
4366 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_scope_cta_space_cta,
4367 llvm::Intrinsic::nvvm_mbarrier_test_wait_scope_cluster_space_cta,
4368 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_scope_cluster_space_cta};
4369 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4370 llvm::Intrinsic::nvvm_mbarrier_test_wait_relaxed_scope_cta_space_cta,
4371 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_relaxed_scope_cta_space_cta,
4372 llvm::Intrinsic::nvvm_mbarrier_test_wait_relaxed_scope_cluster_space_cta,
4373 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_relaxed_scope_cluster_space_cta};
4375 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4378 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4379 llvm::Value *input = mt.
lookupValue(thisOp.getStateOrPhase());
4384 return {id, {mbar, input}};
4389 auto thisOp = cast<NVVM::MBarrierTryWaitOp>(op);
4390 bool isPhaseParity = thisOp.getStateOrPhase().getType().isInteger(32);
4391 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4392 bool hasTicks =
static_cast<bool>(thisOp.getTicks());
4396 size_t index = ((hasTicks ? 1 : 0) << 2) | ((isClusterScope ? 1 : 0) << 1) |
4397 (isPhaseParity ? 1 : 0);
4400 static constexpr llvm::Intrinsic::ID IDs[] = {
4401 llvm::Intrinsic::nvvm_mbarrier_try_wait_scope_cta_space_cta,
4402 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_scope_cta_space_cta,
4403 llvm::Intrinsic::nvvm_mbarrier_try_wait_scope_cluster_space_cta,
4404 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_scope_cluster_space_cta,
4405 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_scope_cta_space_cta,
4406 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_scope_cta_space_cta,
4407 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_scope_cluster_space_cta,
4408 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_scope_cluster_space_cta};
4409 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4410 llvm::Intrinsic::nvvm_mbarrier_try_wait_relaxed_scope_cta_space_cta,
4411 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_relaxed_scope_cta_space_cta,
4412 llvm::Intrinsic::nvvm_mbarrier_try_wait_relaxed_scope_cluster_space_cta,
4413 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_relaxed_scope_cluster_space_cta,
4414 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_relaxed_scope_cta_space_cta,
4415 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_relaxed_scope_cta_space_cta,
4416 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_relaxed_scope_cluster_space_cta,
4417 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_relaxed_scope_cluster_space_cta};
4419 auto id = thisOp.getRelaxed() ? relaxedIDs[
index] : IDs[
index];
4422 llvm::Value *mbar = mt.
lookupValue(thisOp.getAddr());
4429 args.push_back(mbar);
4430 args.push_back(mt.
lookupValue(thisOp.getStateOrPhase()));
4432 args.push_back(mt.
lookupValue(thisOp.getTicks()));
4434 return {id, std::move(args)};
4439 auto thisOp = cast<NVVM::CpAsyncMBarrierArriveOp>(op);
4442 llvm::Intrinsic::ID id;
4443 if (thisOp.getNoinc()) {
4444 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc_shared
4445 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc;
4447 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_shared
4448 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive;
4456 llvm::IRBuilderBase &builder) {
4457 auto thisOp = cast<NVVM::MovMatrixOp>(op);
4458 return {llvm::Intrinsic::nvvm_movmatrix_sync_aligned_m8n8_trans_b16,
4462#define CP_ASYNC_ID_IMPL(mod, size, suffix) \
4463 llvm::Intrinsic::nvvm_cp_async_##mod##_shared_global_##size##suffix
4465#define GET_CP_ASYNC_ID(mod, size, has_cpsize) \
4466 has_cpsize ? CP_ASYNC_ID_IMPL(mod, size, _s) : CP_ASYNC_ID_IMPL(mod, size, )
4471 llvm::Intrinsic::ID id;
4473 auto cpAsyncOp = cast<NVVM::CpAsyncOp>(op);
4474 bool hasCpSize =
static_cast<bool>(cpAsyncOp.getCpSize());
4475 switch (cpAsyncOp.getSize()) {
4483 id = (cpAsyncOp.getModifier() == NVVM::LoadCacheModifierKind::CG)
4488 llvm_unreachable(
"Invalid copy size in CpAsyncOp.");
4492 args.push_back(mt.
lookupValue(cpAsyncOp.getDst()));
4493 args.push_back(mt.
lookupValue(cpAsyncOp.getSrc()));
4495 args.push_back(mt.
lookupValue(cpAsyncOp.getCpSize()));
4502 auto thisOp = cast<NVVM::CpAsyncBulkPrefetchOp>(op);
4504 llvm::Intrinsic::ID
id = llvm::Intrinsic::nvvm_cp_async_bulk_prefetch_L2;
4507 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4511 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4512 llvm::Value *i64Unused =
4513 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4514 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4515 args.push_back(builder.getInt1(hasCacheHint));
4517 return {id, std::move(args)};
4522 auto thisOp = cast<NVVM::CpAsyncBulkGlobalToSharedClusterOp>(op);
4526 args.push_back(mt.
lookupValue(thisOp.getDstMem()));
4528 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4532 mlir::Value multicastMask = thisOp.getMulticastMask();
4533 const bool hasMulticastMask =
static_cast<bool>(multicastMask);
4536 llvm::Value *i16Unused = llvm::ConstantInt::get(builder.getInt16Ty(), 0);
4537 args.push_back(hasMulticastMask ? mt.
lookupValue(multicastMask)
4543 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4544 llvm::Value *i64Unused = llvm::ConstantInt::get(builder.getInt64Ty(), 0);
4545 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4549 args.push_back(builder.getInt1(hasMulticastMask));
4550 args.push_back(builder.getInt1(hasCacheHint));
4552 llvm::Intrinsic::ID
id =
4554 ? llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cta
4555 : llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cluster;
4557 return {id, std::move(args)};
4562 auto thisOp = cast<NVVM::CpAsyncBulkSharedCTAToGlobalOp>(op);
4564 llvm::Intrinsic::ID
id =
4565 llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global;
4568 args.push_back(mt.
lookupValue(thisOp.getDstMem()));
4569 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4573 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4574 llvm::Value *i64Unused =
4575 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4576 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4577 args.push_back(builder.getInt1(hasCacheHint));
4580 if (
mlir::Value byteMask = thisOp.getByteMask()) {
4582 id = llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global_bytemask;
4585 return {id, std::move(args)};
4588bool CpAsyncBulkTensorGlobalToSharedClusterOp::getAsmValues(
4595 for (
auto val : getOperands())
4602CpAsyncBulkTensorGlobalToSharedClusterOp::getIntrinsicIDAndArgs(
4604 auto thisOp = cast<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(op);
4605 const bool isCTAOnly = thisOp.getIsCTAOnly();
4609 args.push_back(mt.
lookupValue(thisOp.getDstMem()));
4611 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4621 const bool hasMC =
static_cast<bool>(mcMask);
4622 llvm::Value *i16Zero =
4623 llvm::ConstantInt::get(llvm::Type::getInt16Ty(mt.
getLLVMContext()), 0);
4627 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4628 llvm::Value *i64Zero =
4629 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4635 thisOp.getGroup() ? (
static_cast<int32_t
>(*thisOp.getGroup()) + 1) : 0;
4637 llvm::ConstantInt::get(llvm::Type::getInt32Ty(mt.
getLLVMContext()), val);
4641 args.push_back(hasMC ? mt.
lookupValue(mcMask) : i16Zero);
4642 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Zero);
4643 args.push_back(builder.getInt1(hasMC));
4644 args.push_back(builder.getInt1(hasCacheHint));
4648 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Zero);
4649 args.push_back(builder.getInt1(hasCacheHint));
4652 constexpr size_t numDims = 5;
4653 constexpr size_t numModes = 5;
4654 using rowTy = std::array<llvm::Intrinsic::ID, numDims + 1>;
4655 using TableTy = std::array<rowTy, numModes>;
4656 static constexpr TableTy IDTable{
4657 {{
notIntrinsic, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_1d,
4658 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_2d,
4659 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_3d,
4660 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_4d,
4661 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_5d},
4663 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_3d,
4664 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_4d,
4665 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_5d},
4667 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_3d,
4668 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_4d,
4669 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_5d},
4671 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_3d,
4672 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_4d,
4673 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_5d},
4675 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_gather4_2d}}};
4677 static constexpr TableTy IDTableCTA{
4679 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_1d,
4680 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_2d,
4681 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_3d,
4682 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_4d,
4683 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_5d},
4685 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_3d,
4686 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_4d,
4687 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_5d},
4689 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_3d,
4690 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_4d,
4691 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_5d},
4693 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_3d,
4694 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_4d,
4695 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_5d},
4697 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_gather4_2d}}};
4700 (getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1) &&
4701 (getMaxEnumValForTMALoadMode() == std::size(IDTableCTA) - 1),
4702 "TMALoadModes must match number of rows in IDTable and IDTableCTA");
4703 size_t mode =
static_cast<size_t>(thisOp.getMode());
4704 size_t dim = thisOp.getCoordinates().size();
4705 auto id = isCTAOnly ? IDTableCTA[mode][dim] : IDTable[mode][dim];
4707 "Invalid intrinsic for CpAsyncBulkTensorGlobalToSharedClusterOp.");
4709 return {id, std::move(args)};
4714 auto thisOp = cast<NVVM::CpAsyncBulkTensorPrefetchOp>(op);
4718 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4720 for (
auto v : thisOp.getCoordinates())
4722 for (
auto v : thisOp.getIm2colOffsets())
4726 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4727 llvm::Value *i64Unused =
4728 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4729 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4730 args.push_back(builder.getInt1(hasCacheHint));
4732 const unsigned NI = llvm::Intrinsic::not_intrinsic;
4733 static constexpr llvm::Intrinsic::ID IDTable[][6] = {
4734 {NI, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_1d,
4735 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_2d,
4736 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_3d,
4737 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_4d,
4738 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_5d},
4740 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_3d,
4741 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_4d,
4742 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_5d},
4744 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_3d,
4745 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_4d,
4746 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_5d},
4748 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_3d,
4749 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_4d,
4750 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_5d},
4751 {NI, NI, NI, NI, NI,
4752 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_gather4_2d}};
4754 static_assert(getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1,
4755 "TMALoadModes must match number of rows in IDTable");
4756 size_t mode =
static_cast<size_t>(thisOp.getMode());
4757 size_t dim = thisOp.getCoordinates().size();
4758 llvm::Intrinsic::ID
id = IDTable[mode][dim];
4759 if (
id == llvm::Intrinsic::not_intrinsic)
4760 llvm_unreachable(
"Invalid intrinsic for CpAsyncBulkTensorPrefetchOp.");
4762 return {id, std::move(args)};
4766CpAsyncBulkTensorSharedCTAToGlobalOp::getIntrinsicIDAndArgs(
4768 auto thisOp = cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(op);
4772 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4773 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4775 for (
auto v : thisOp.getCoordinates())
4779 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4780 llvm::Value *i64Unused =
4781 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.
getLLVMContext()), 0);
4782 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint) : i64Unused);
4783 args.push_back(builder.getInt1(hasCacheHint));
4785 using namespace llvm::Intrinsic;
4786 const unsigned NI = not_intrinsic;
4787 static constexpr ID IDTable[][6] = {
4788 {NI, nvvm_cp_async_bulk_tensor_s2g_tile_1d,
4789 nvvm_cp_async_bulk_tensor_s2g_tile_2d,
4790 nvvm_cp_async_bulk_tensor_s2g_tile_3d,
4791 nvvm_cp_async_bulk_tensor_s2g_tile_4d,
4792 nvvm_cp_async_bulk_tensor_s2g_tile_5d},
4793 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_3d,
4794 nvvm_cp_async_bulk_tensor_s2g_im2col_4d,
4795 nvvm_cp_async_bulk_tensor_s2g_im2col_5d},
4796 {NI, NI, NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_2d},
4797 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_w_3d,
4798 nvvm_cp_async_bulk_tensor_s2g_im2col_w_4d,
4799 nvvm_cp_async_bulk_tensor_s2g_im2col_w_5d}};
4801 static_assert(getMaxEnumValForTMAStoreMode() == std::size(IDTable) - 1,
4802 "TMAStoreModes must match number of rows in IDTable");
4803 size_t mode =
static_cast<size_t>(thisOp.getMode());
4804 size_t dim = thisOp.getCoordinates().size();
4805 ID
id = IDTable[mode][dim];
4806 if (
id == llvm::Intrinsic::not_intrinsic)
4808 "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOp.");
4810 return {id, std::move(args)};
4814CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp::getIntrinsicIDAndArgs(
4817 cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp>(op);
4820 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4821 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4822 args.push_back(mt.
lookupValue(thisOp.getOverrideAddr()));
4823 for (
Value v : thisOp.getTensorSize())
4825 for (
Value v : thisOp.getLowerStride())
4827 if (thisOp.getUpperStride())
4828 args.push_back(mt.
lookupValue(thisOp.getUpperStride()));
4829 for (
Value v : thisOp.getCoordinates())
4833 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4834 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint)
4835 : builder.getInt64(0));
4836 args.push_back(builder.getInt1(hasCacheHint));
4838 using namespace llvm::Intrinsic;
4839 const unsigned NI = not_intrinsic;
4842 static constexpr ID IDTable[][6] = {
4843 {NI, nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_1d,
4844 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_2d,
4845 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_3d,
4846 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_4d,
4847 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_5d},
4848 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_3d,
4849 nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_4d,
4850 nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_5d},
4851 {NI, NI, NI, NI, NI,
4852 nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_override_addr_2d},
4853 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_3d,
4854 nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_4d,
4855 nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_5d}};
4859 static constexpr ID dimStrideIDTable[] = {
4860 NI, nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_1d,
4861 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_2d,
4862 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_3d,
4863 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_4d,
4864 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_5d};
4867 size_t mode =
static_cast<size_t>(thisOp.getMode());
4868 size_t dim = thisOp.getCoordinates().size();
4869 bool isDimStride = !thisOp.getTensorSize().empty();
4871 assert(mode < std::size(IDTable) &&
4872 "Invalid mode for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
4873 assert(dim < std::size(IDTable[mode]) && dim < std::size(dimStrideIDTable) &&
4874 "Invalid dim for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
4876 ID intrinsicID = isDimStride ? dimStrideIDTable[dim] : IDTable[mode][dim];
4878 intrinsicID != NI &&
4879 "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
4880 return {intrinsicID, std::move(args)};
4885 auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOp>(op);
4888 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4889 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4890 for (
Value v : thisOp.getCoordinates())
4894 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4895 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint)
4896 : builder.getInt64(0));
4897 args.push_back(builder.getInt32(
static_cast<uint32_t
>(thisOp.getRedKind())));
4898 args.push_back(builder.getInt1(hasCacheHint));
4900 using namespace llvm::Intrinsic;
4901 const unsigned NI = not_intrinsic;
4902 static constexpr ID IDTable[][6] = {
4903 {NI, nvvm_cp_async_bulk_tensor_reduce_tile_1d,
4904 nvvm_cp_async_bulk_tensor_reduce_tile_2d,
4905 nvvm_cp_async_bulk_tensor_reduce_tile_3d,
4906 nvvm_cp_async_bulk_tensor_reduce_tile_4d,
4907 nvvm_cp_async_bulk_tensor_reduce_tile_5d},
4908 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_3d,
4909 nvvm_cp_async_bulk_tensor_reduce_im2col_4d,
4910 nvvm_cp_async_bulk_tensor_reduce_im2col_5d},
4911 {NI, NI, NI, NI, NI, NI},
4912 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_w_3d,
4913 nvvm_cp_async_bulk_tensor_reduce_im2col_w_4d,
4914 nvvm_cp_async_bulk_tensor_reduce_im2col_w_5d}};
4916 size_t mode =
static_cast<size_t>(thisOp.getMode());
4917 size_t dim = thisOp.getCoordinates().size();
4918 assert(mode < std::size(IDTable) &&
4919 "Invalid mode for CpAsyncBulkTensorReduceOp");
4920 assert(dim < std::size(IDTable[mode]) &&
4921 "Invalid dim for CpAsyncBulkTensorReduceOp");
4923 ID intrinsicID = IDTable[mode][dim];
4924 assert(intrinsicID != NI &&
4925 "Invalid intrinsic for CpAsyncBulkTensorReduceOp");
4926 return {intrinsicID, std::move(args)};
4929NVVM::IDArgPair CpAsyncBulkTensorReduceOverrideAddrOp::getIntrinsicIDAndArgs(
4931 auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOverrideAddrOp>(op);
4934 args.push_back(mt.
lookupValue(thisOp.getSrcMem()));
4935 args.push_back(mt.
lookupValue(thisOp.getTmaDescriptor()));
4936 args.push_back(mt.
lookupValue(thisOp.getOverrideAddr()));
4938 for (
Value v : thisOp.getTensorSize())
4940 for (
Value v : thisOp.getLowerStride())
4942 if (thisOp.getUpperStride())
4943 args.push_back(mt.
lookupValue(thisOp.getUpperStride()));
4944 for (
Value v : thisOp.getCoordinates())
4948 const bool hasCacheHint =
static_cast<bool>(cacheHint);
4949 args.push_back(hasCacheHint ? mt.
lookupValue(cacheHint)
4950 : builder.getInt64(0));
4951 args.push_back(builder.getInt32(
static_cast<uint32_t
>(thisOp.getRedKind())));
4952 args.push_back(builder.getInt1(hasCacheHint));
4954 using namespace llvm::Intrinsic;
4955 const unsigned NI = not_intrinsic;
4958static constexpr ID IDTable[][6] = {
4959 {NI, nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_1d,
4960 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_2d,
4961 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_3d,
4962 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_4d,
4963 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_5d},
4964 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_3d,
4965 nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_4d,
4966 nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_5d},
4967 {NI, NI, NI, NI, NI, NI},
4968 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_3d,
4969 nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_4d,
4970 nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_5d}};
4974static constexpr ID dimStrideIDTable[] = {
4975 NI, nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_1d,
4976 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_2d,
4977 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_3d,
4978 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_4d,
4979 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_5d};
4982 size_t mode =
static_cast<size_t>(thisOp.getMode());
4983 size_t dim = thisOp.getCoordinates().size();
4984 bool isDimStride = !thisOp.getTensorSize().empty();
4986 assert(mode < std::size(IDTable) &&
4987 "Invalid mode for CpAsyncBulkTensorReduceOverrideAddrOp");
4988 assert(dim < std::size(IDTable[mode]) && dim < std::size(dimStrideIDTable) &&
4989 "Invalid dim for CpAsyncBulkTensorReduceOverrideAddrOp");
4991 ID intrinsicID = isDimStride ? dimStrideIDTable[dim] : IDTable[mode][dim];
4992 assert(intrinsicID != NI &&
4993 "Invalid intrinsic for CpAsyncBulkTensorReduceOverrideAddrOp");
4994 return {intrinsicID, std::move(args)};
4999#define CVT_F2TF32_ID_IMPL(rnd, relu, sf) \
5000 hasRelu ? llvm::Intrinsic::nvvm_f2tf32_##rnd##relu##sf \
5001 : llvm::Intrinsic::nvvm_f2tf32_##rnd##sf
5003#define GET_CVT_F2TF32_ID(rnd, relu, sf) \
5004 hasSatFinite ? CVT_F2TF32_ID_IMPL(rnd, relu, sf) \
5005 : CVT_F2TF32_ID_IMPL(rnd, relu, )
5008ConvertFloatToTF32Op::getIntrinsicID(NVVM::FPRoundingMode rnd,
5009 NVVM::SaturationMode sat,
bool hasRelu) {
5010 using RndMode = NVVM::FPRoundingMode;
5011 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
5020 llvm_unreachable(
"Invalid RoundingMode for CvtFloatToTF32Op");
5025ConvertF32x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF4x2Op op,
5027 llvm::IRBuilderBase &builder) {
5032 bool hasRelu = op.getRelu();
5034 llvm::Intrinsic::ID intId =
5035 hasRelu ? llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_relu_satfinite
5036 : llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_satfinite;
5038 return {intId, std::move(args)};
5041#define GET_F32x2_TO_F6x2_ID(type, has_relu) \
5042 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu_satfinite \
5043 : llvm::Intrinsic::nvvm_ff_to_##type##_rn_satfinite
5045llvm::Intrinsic::ID ConvertF32x2ToF6x2Op::getIntrinsicID(
mlir::Type dstTy,
5048 .Case([&](mlir::Float6E2M3FNType) {
5051 .Case([&](mlir::Float6E3M2FNType) {
5055 llvm_unreachable(
"Invalid conversion in ConvertF32x2ToF6x2Op");
5056 return llvm::Intrinsic::not_intrinsic;
5061ConvertF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF16x2ToF4x2Op &op,
5063 llvm::IRBuilderBase &builder) {
5065 bool hasRelu = op.getRelu();
5067 llvm::Intrinsic::ID intId = llvm::Intrinsic::not_intrinsic;
5069 if (llvm::isa<mlir::Float4E2M1FNType>(dstTy))
5070 intId = hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e2m1x2_rn_relu_satfinite
5071 : llvm::Intrinsic::nvvm_f16x2_to_e2m1x2_rn_satfinite;
5076 return {intId, std::move(args)};
5080ConvertBF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertBF16x2ToF4x2Op &op,
5082 llvm::IRBuilderBase &builder) {
5084 bool hasRelu = op.getRelu();
5086 llvm::Intrinsic::ID intId = llvm::Intrinsic::not_intrinsic;
5088 if (llvm::isa<mlir::Float4E2M1FNType>(dstTy))
5089 intId = hasRelu ? llvm::Intrinsic::nvvm_bf16x2_to_e2m1x2_rn_relu_satfinite
5090 : llvm::Intrinsic::nvvm_bf16x2_to_e2m1x2_rn_satfinite;
5095 return {intId, std::move(args)};
5098llvm::Intrinsic::ID ConvertF16x2ToF6x2Op::getIntrinsicID(
mlir::Type dstTy,
5101 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {
5102 return hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e2m3x2_rn_relu_satfinite
5103 : llvm::Intrinsic::nvvm_f16x2_to_e2m3x2_rn_satfinite;
5105 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {
5106 return hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e3m2x2_rn_relu_satfinite
5107 : llvm::Intrinsic::nvvm_f16x2_to_e3m2x2_rn_satfinite;
5110 llvm_unreachable(
"Invalid conversion in ConvertF16x2ToF6x2Op");
5111 return llvm::Intrinsic::not_intrinsic;
5115llvm::Intrinsic::ID ConvertBF16x2ToF6x2Op::getIntrinsicID(
mlir::Type dstTy,
5118 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {
5120 ? llvm::Intrinsic::nvvm_bf16x2_to_e2m3x2_rn_relu_satfinite
5121 : llvm::Intrinsic::nvvm_bf16x2_to_e2m3x2_rn_satfinite;
5123 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {
5125 ? llvm::Intrinsic::nvvm_bf16x2_to_e3m2x2_rn_relu_satfinite
5126 : llvm::Intrinsic::nvvm_bf16x2_to_e3m2x2_rn_satfinite;
5129 llvm_unreachable(
"Invalid conversion in ConvertBF16x2ToF6x2Op");
5130 return llvm::Intrinsic::not_intrinsic;
5134#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf) \
5135 has_satf ? llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd##_satfinite \
5136 : llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd
5138#define GET_F32x2_TO_F8X2_S_ID(type, has_relu) \
5139 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu \
5140 : llvm::Intrinsic::nvvm_ff_to_##type##_rn
5143ConvertF32x2ToF8x2Op::getIntrinsicID(
mlir::Type dstTy, NVVM::FPRoundingMode rnd,
5144 NVVM::SaturationMode sat,
bool hasRelu) {
5145 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
5146 bool hasRoundingModeRZ = (rnd == NVVM::FPRoundingMode::RZ);
5147 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);
5150 .Case([&](mlir::Float8E4M3FNType) {
5153 .Case([&](mlir::Float8E5M2Type) {
5156 .Case([&](mlir::Float8E8M0FNUType) {
5157 if (hasRoundingModeRZ)
5159 else if (hasRoundingModeRP)
5162 llvm_unreachable(
"Invalid conversion in ConvertF32x2ToF8x2Op");
5165 llvm_unreachable(
"Invalid conversion in ConvertF32x2ToF8x2Op");
5166 return llvm::Intrinsic::not_intrinsic;
5170#define GET_F16x2_TO_F8X2_ID(type, has_relu) \
5171 has_relu ? llvm::Intrinsic::nvvm_f16x2_to_##type##_rn_relu \
5172 : llvm::Intrinsic::nvvm_f16x2_to_##type##_rn
5174llvm::Intrinsic::ID ConvertF16x2ToF8x2Op::getIntrinsicID(
mlir::Type dstTy,
5177 .Case([&](mlir::Float8E4M3FNType) {
5180 .Case([&](mlir::Float8E5M2Type) {
5184 llvm_unreachable(
"Invalid conversion in ConvertF16x2ToF8x2Op");
5185 return llvm::Intrinsic::not_intrinsic;
5190ConvertBF16x2ToF8x2Op::getIntrinsicID(
mlir::Type dstTy,
5191 NVVM::FPRoundingMode rnd,
5192 NVVM::SaturationMode sat,
bool hasRelu) {
5193 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
5195 static constexpr llvm::Intrinsic::ID ue8m0x2IDs[] = {
5196 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rz,
5197 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rp,
5198 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rz_satfinite,
5199 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rp_satfinite,
5203 .Case<mlir::Float8E4M3FNType>([&](mlir::Float8E4M3FNType) {
5205 ? llvm::Intrinsic::nvvm_bf16x2_to_e4m3x2_rn_relu_satfinite
5206 : llvm::Intrinsic::nvvm_bf16x2_to_e4m3x2_rn_satfinite;
5208 .Case<mlir::Float8E5M2Type>([&](mlir::Float8E5M2Type) {
5210 ? llvm::Intrinsic::nvvm_bf16x2_to_e5m2x2_rn_relu_satfinite
5211 : llvm::Intrinsic::nvvm_bf16x2_to_e5m2x2_rn_satfinite;
5213 .Case<mlir::Float8E8M0FNUType>([&](mlir::Float8E8M0FNUType) {
5214 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);
5215 unsigned index = (hasSatFinite << 1) | hasRoundingModeRP;
5216 return ue8m0x2IDs[
index];
5219 llvm_unreachable(
"Invalid conversion in ConvertBF16x2ToF8x2Op");
5220 return llvm::Intrinsic::not_intrinsic;
5226 auto curOp = cast<NVVM::ConvertF8x2ToF16x2Op>(op);
5228 bool hasRelu = curOp.getRelu();
5230 llvm::Intrinsic::ID intId =
5232 .Case([&](Float8E4M3FNType type) {
5233 return hasRelu ? llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn_relu
5234 : llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn;
5236 .Case([&](Float8E5M2Type type) {
5237 return hasRelu ? llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn_relu
5238 : llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn;
5241 llvm_unreachable(
"Invalid type for ConvertF8x2ToF16x2Op");
5242 return llvm::Intrinsic::not_intrinsic;
5245 llvm::Value *packedI16 =
5246 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
5247 llvm::Type::getInt16Ty(builder.getContext()));
5249 return {intId, {packedI16}};
5254 auto curOp = cast<NVVM::ConvertF8x2ToBF16x2Op>(op);
5255 bool hasScale =
static_cast<bool>(curOp.getScaleFactor());
5256 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
5257 bool hasRelu = curOp.getRelu();
5259 static constexpr llvm::Intrinsic::ID E4M3Ids[] = {
5260 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_scale_n2_ue8m0,
5261 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5262 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5263 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5266 static constexpr llvm::Intrinsic::ID E5M2Ids[] = {
5267 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_scale_n2_ue8m0,
5268 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5269 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5270 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5273 llvm::Intrinsic::ID intId =
5275 .Case([&](Float8E8M0FNUType type) {
5276 return llvm::Intrinsic::nvvm_ue8m0x2_to_bf16x2;
5278 .Case([&](Float8E4M3FNType type) {
5279 return E4M3Ids[hasSatfinite << 1 | hasRelu];
5281 .Case([&](Float8E5M2Type type) {
5282 return E5M2Ids[hasSatfinite << 1 | hasRelu];
5285 llvm_unreachable(
"Invalid type for ConvertF8x2ToBF16x2Op");
5286 return llvm::Intrinsic::not_intrinsic;
5288 llvm::Value *packedI16 =
5289 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
5290 llvm::Type::getInt16Ty(builder.getContext()));
5293 args.push_back(packedI16);
5294 if (!isa<Float8E8M0FNUType>(curOp.getSrcType()))
5297 : builder.getInt16(0x7f7f));
5300 return {intId, std::move(args)};
5305 auto curOp = cast<NVVM::ConvertF6x2ToF16x2Op>(op);
5307 bool hasRelu = curOp.getRelu();
5309 llvm::Intrinsic::ID intId =
5311 .Case([&](Float6E2M3FNType type) {
5312 return hasRelu ? llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn_relu
5313 : llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn;
5315 .Case([&](Float6E3M2FNType type) {
5316 return hasRelu ? llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn_relu
5317 : llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn;
5320 llvm_unreachable(
"Invalid type for ConvertF6x2ToF16x2Op");
5321 return llvm::Intrinsic::not_intrinsic;
5324 llvm::Value *packedI16 =
5325 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
5326 llvm::Type::getInt16Ty(builder.getContext()));
5328 return {intId, {packedI16}};
5333 auto curOp = cast<NVVM::ConvertF6x2ToBF16x2Op>(op);
5334 bool hasScale =
static_cast<bool>(curOp.getScaleFactor());
5335 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
5336 bool hasRelu = curOp.getRelu();
5338 static constexpr llvm::Intrinsic::ID E2M3Ids[] = {
5339 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_scale_n2_ue8m0,
5340 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5341 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5342 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5345 static constexpr llvm::Intrinsic::ID E3M2Ids[] = {
5346 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_scale_n2_ue8m0,
5347 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5348 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5349 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5352 unsigned idx = (hasSatfinite << 1) | hasRelu;
5353 llvm::Intrinsic::ID intId =
5355 .Case([&](Float6E2M3FNType type) {
return E2M3Ids[idx]; })
5356 .Case([&](Float6E3M2FNType type) {
return E3M2Ids[idx]; })
5358 llvm_unreachable(
"Invalid type for ConvertF6x2ToBF16x2Op");
5359 return llvm::Intrinsic::not_intrinsic;
5362 llvm::Value *packedI16 =
5363 builder.CreateBitCast(mt.
lookupValue(curOp.getSrc()),
5364 llvm::Type::getInt16Ty(builder.getContext()));
5367 args.push_back(packedI16);
5374 return {intId, std::move(args)};
5379 auto curOp = cast<NVVM::ConvertF4x2ToF16x2Op>(op);
5381 bool hasRelu = curOp.getRelu();
5383 llvm::Intrinsic::ID intId =
5385 .Case([&](Float4E2M1FNType type) {
5386 return hasRelu ? llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn_relu
5387 : llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn;
5390 llvm_unreachable(
"Invalid type for ConvertF4x2ToF16x2Op");
5391 return llvm::Intrinsic::not_intrinsic;
5394 llvm::Value *extendedI16 =
5395 builder.CreateZExt(mt.
lookupValue(curOp.getSrc()),
5396 llvm::Type::getInt16Ty(builder.getContext()));
5398 return {intId, {extendedI16}};
5403 auto curOp = cast<NVVM::ConvertF4x2ToBF16x2Op>(op);
5404 bool hasScale =
static_cast<bool>(curOp.getScaleFactor());
5405 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
5406 bool hasRelu = curOp.getRelu();
5408 static constexpr llvm::Intrinsic::ID E2M1Ids[] = {
5409 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_scale_n2_ue8m0,
5410 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5411 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5412 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5415 unsigned idx = (hasSatfinite << 1) | hasRelu;
5416 llvm::Intrinsic::ID intId =
5418 .Case([&](Float4E2M1FNType type) {
return E2M1Ids[idx]; })
5420 llvm_unreachable(
"Invalid type for ConvertF4x2ToBF16x2Op");
5421 return llvm::Intrinsic::not_intrinsic;
5424 llvm::Value *extendedI16 =
5425 builder.CreateZExt(mt.
lookupValue(curOp.getSrc()),
5426 llvm::Type::getInt16Ty(builder.getContext()));
5429 args.push_back(extendedI16);
5436 return {intId, std::move(args)};
5441 auto thisOp = cast<NVVM::ConvertF32x2ToS2F6x2Op>(op);
5442 bool hasRelu = thisOp.getRelu();
5443 bool hasScale =
static_cast<bool>(thisOp.getScaleFactor());
5445 llvm::Intrinsic::ID
id =
5447 ? llvm::Intrinsic::nvvm_ff_to_s2f6x2_rn_relu_satfinite_scale_n2_ue8m0
5448 : llvm::Intrinsic::nvvm_ff_to_s2f6x2_rn_satfinite_scale_n2_ue8m0;
5454 args.push_back(hasScale ? mt.
lookupValue(thisOp.getScaleFactor())
5455 : builder.getInt16(0x7f7f));
5456 return {id, std::move(args)};
5461 auto thisOp = cast<NVVM::ConvertBF16x2ToS2F6x2Op>(op);
5462 bool hasRelu = thisOp.getRelu();
5463 bool hasScale =
static_cast<bool>(thisOp.getScaleFactor());
5465 llvm::Intrinsic::ID
id =
5468 nvvm_bf16x2_to_s2f6x2_rn_relu_satfinite_scale_n2_ue8m0
5469 : llvm::Intrinsic::nvvm_bf16x2_to_s2f6x2_rn_satfinite_scale_n2_ue8m0;
5474 args.push_back(hasScale ? mt.
lookupValue(thisOp.getScaleFactor())
5475 : builder.getInt16(0x7f7f));
5476 return {id, std::move(args)};
5481 auto thisOp = cast<NVVM::ConvertS2F6x2ToBF16x2Op>(op);
5482 bool hasRelu = thisOp.getRelu();
5483 bool hasScale =
static_cast<bool>(thisOp.getScaleFactor());
5484 bool hasSat = thisOp.getSat() == NVVM::SaturationMode::SATFINITE;
5486 static constexpr llvm::Intrinsic::ID ids[] = {
5487 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_scale_n2_ue8m0,
5488 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5489 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5490 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5493 unsigned idx = (hasSat << 1) | hasRelu;
5497 llvm::Value *packedI16 =
5498 builder.CreateBitCast(mt.
lookupValue(thisOp.getSrc()),
5499 llvm::Type::getInt16Ty(builder.getContext()));
5500 args.push_back(packedI16);
5501 args.push_back(hasScale ? mt.
lookupValue(thisOp.getScaleFactor())
5502 : builder.getInt16(0x7f7f));
5504 return {ids[idx], std::move(args)};
5508Tcgen05AllocOp::getIntrinsicIDAndArgs(
Operation &op,
5511 auto curOp = cast<NVVM::Tcgen05AllocOp>(op);
5512 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
5514 llvm::Intrinsic::ID
id = is2CTAMode ? llvm::Intrinsic::nvvm_tcgen05_alloc_cg2
5515 : llvm::Intrinsic::nvvm_tcgen05_alloc_cg1;
5520 args.push_back(llvm::ConstantInt::getFalse(mt.
getLLVMContext()));
5525llvm::Intrinsic::ID Tcgen05DeallocOp::getIntrinsicIDAndArgs(
5528 auto curOp = cast<NVVM::Tcgen05DeallocOp>(op);
5529 auto id = (curOp.getGroup() == CTAGroupKind::CTA_1)
5530 ? llvm::Intrinsic::nvvm_tcgen05_dealloc_cg1
5531 : llvm::Intrinsic::nvvm_tcgen05_dealloc_cg2;
5536 args.push_back(llvm::ConstantInt::getFalse(mt.
getLLVMContext()));
5542Tcgen05CommitOp::getIntrinsicIDAndArgs(
Operation &op,
5545 auto curOp = cast<NVVM::Tcgen05CommitOp>(op);
5546 bool hasMulticast =
static_cast<bool>(curOp.getMulticastMask());
5547 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
5548 bool hasSmemARead = curOp.getSmemARead();
5549 unsigned index = (
static_cast<unsigned>(hasSmemARead) << 1) |
5550 static_cast<unsigned>(is2CTAMode);
5552 using namespace llvm::Intrinsic;
5553 static constexpr ID IDs[] = {
5554 nvvm_tcgen05_commit_cg1,
5555 nvvm_tcgen05_commit_cg2,
5556 nvvm_tcgen05_commit_smem_a_read_cg1,
5557 nvvm_tcgen05_commit_smem_a_read_cg2,
5560 static constexpr ID multicastIDs[] = {
5561 nvvm_tcgen05_commit_mc_cg1,
5562 nvvm_tcgen05_commit_mc_cg2,
5563 nvvm_tcgen05_commit_smem_a_read_mc_cg1,
5564 nvvm_tcgen05_commit_smem_a_read_mc_cg2,
5567 ID
id = hasMulticast ? multicastIDs[
index] : IDs[
index];
5571 args.push_back(mt.
lookupValue(curOp.getMulticastMask()));
5576#define TCGEN05_CP_IMPL(shape_mc, src_fmt, cg) \
5577 llvm::Intrinsic::nvvm_tcgen05_cp##shape_mc##src_fmt##cg
5579#define TCGEN05_CP_2CTA(shape_mc, src_fmt, is_2cta) \
5580 is_2cta ? TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg2) \
5581 : TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg1)
5583#define GET_TCGEN05_CP_ID(shape_mc, src_fmt, is_2cta) \
5585 if ((src_fmt) == Tcgen05CpSrcFormat::B6x16_P32) \
5586 return TCGEN05_CP_2CTA(shape_mc, _b6x16_p32, is_2cta); \
5587 if ((src_fmt) == Tcgen05CpSrcFormat::B4x16_P64) \
5588 return TCGEN05_CP_2CTA(shape_mc, _b4x16_p64, is_2cta); \
5589 return TCGEN05_CP_2CTA(shape_mc, , is_2cta); \
5593ConvertF32x2ToF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF16x2Op &op,
5595 llvm::IRBuilderBase &builder) {
5596 static constexpr llvm::Intrinsic::ID rndRNIds[] = {
5597 llvm::Intrinsic::nvvm_ff2f16x2_rn,
5598 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu,
5599 llvm::Intrinsic::nvvm_ff2f16x2_rn_satfinite,
5600 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu_satfinite,
5602 static constexpr llvm::Intrinsic::ID rndRZIds[] = {
5603 llvm::Intrinsic::nvvm_ff2f16x2_rz,
5604 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu,
5605 llvm::Intrinsic::nvvm_ff2f16x2_rz_satfinite,
5606 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu_satfinite,
5608 static constexpr llvm::Intrinsic::ID rndRSIds[] = {
5609 llvm::Intrinsic::nvvm_ff2f16x2_rs,
5610 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu,
5611 llvm::Intrinsic::nvvm_ff2f16x2_rs_satfinite,
5612 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu_satfinite,
5615 unsigned hasRelu = op.getRelu() ? 1 : 0;
5616 unsigned hasSatFinite =
5617 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;
5620 unsigned idx = (hasSatFinite << 1) | hasRelu;
5625 if (op.getRandomBits())
5626 args.push_back(mt.
lookupValue(op.getRandomBits()));
5629 args.push_back(builder.getInt1(
false));
5631 switch (op.getRnd()) {
5632 case FPRoundingMode::RN:
5633 return {rndRNIds[idx], std::move(args)};
5634 case FPRoundingMode::RZ:
5635 return {rndRZIds[idx], std::move(args)};
5636 case FPRoundingMode::RS:
5637 return {rndRSIds[idx], std::move(args)};
5639 llvm_unreachable(
"Invalid rounding mode for ConvertF32x2ToF16x2Op");
5644ConvertF32x2ToBF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToBF16x2Op &op,
5646 llvm::IRBuilderBase &builder) {
5647 static constexpr llvm::Intrinsic::ID rndRNIds[] = {
5648 llvm::Intrinsic::nvvm_ff2bf16x2_rn,
5649 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu,
5650 llvm::Intrinsic::nvvm_ff2bf16x2_rn_satfinite,
5651 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu_satfinite,
5653 static constexpr llvm::Intrinsic::ID rndRZIds[] = {
5654 llvm::Intrinsic::nvvm_ff2bf16x2_rz,
5655 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu,
5656 llvm::Intrinsic::nvvm_ff2bf16x2_rz_satfinite,
5657 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu_satfinite,
5659 static constexpr llvm::Intrinsic::ID rndRSIds[] = {
5660 llvm::Intrinsic::nvvm_ff2bf16x2_rs,
5661 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu,
5662 llvm::Intrinsic::nvvm_ff2bf16x2_rs_satfinite,
5663 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu_satfinite,
5666 unsigned hasRelu = op.getRelu() ? 1 : 0;
5667 unsigned hasSatFinite =
5668 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;
5671 unsigned idx = (hasSatFinite << 1) | hasRelu;
5676 if (op.getRandomBits())
5677 args.push_back(mt.
lookupValue(op.getRandomBits()));
5680 args.push_back(builder.getInt1(
false));
5682 switch (op.getRnd()) {
5683 case FPRoundingMode::RN:
5684 return {rndRNIds[idx], std::move(args)};
5685 case FPRoundingMode::RZ:
5686 return {rndRZIds[idx], std::move(args)};
5687 case FPRoundingMode::RS:
5688 return {rndRSIds[idx], std::move(args)};
5690 llvm_unreachable(
"Invalid rounding mode for ConvertF32x2ToBF16x2Op");
5694llvm::Intrinsic::ID ConvertF32x4ToF8x4Op::getIntrinsicID() {
5696 bool hasRelu = getRelu();
5699 .Case([&](mlir::Float8E4M3FNType) {
5700 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_relu_satfinite
5701 : llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_satfinite;
5703 .Case([&](mlir::Float8E5M2Type) {
5704 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_relu_satfinite
5705 : llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_satfinite;
5708 llvm_unreachable(
"Invalid F8 type in ConvertF32x4ToF8x4Op");
5709 return llvm::Intrinsic::not_intrinsic;
5713llvm::Intrinsic::ID ConvertF32x4ToF6x4Op::getIntrinsicID() {
5715 bool hasRelu = getRelu();
5718 .Case([&](mlir::Float6E2M3FNType) {
5719 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_relu_satfinite
5720 : llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_satfinite;
5722 .Case([&](mlir::Float6E3M2FNType) {
5723 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_relu_satfinite
5724 : llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_satfinite;
5727 llvm_unreachable(
"Invalid F6 type in ConvertF32x4ToF6x4Op");
5728 return llvm::Intrinsic::not_intrinsic;
5732llvm::Intrinsic::ID ConvertF32x4ToF4x4Op::getIntrinsicID() {
5734 bool hasRelu = getRelu();
5737 .Case([&](mlir::Float4E2M1FNType) {
5738 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_relu_satfinite
5739 : llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_satfinite;
5742 llvm_unreachable(
"Invalid F4 type in ConvertF32x4ToF4x4Op");
5743 return llvm::Intrinsic::not_intrinsic;
5747llvm::Intrinsic::ID Tcgen05CpOp::getIntrinsicID(
Operation &op) {
5748 auto curOp = cast<NVVM::Tcgen05CpOp>(op);
5749 bool is2CTA = curOp.getGroup() == CTAGroupKind::CTA_2;
5750 auto srcFmt = curOp.getSrcFormat();
5751 auto mc = curOp.getMulticast();
5753 switch (curOp.getShape()) {
5754 case Tcgen05CpShape::SHAPE_128x256b:
5756 case Tcgen05CpShape::SHAPE_128x128b:
5758 case Tcgen05CpShape::SHAPE_4x256b:
5760 case Tcgen05CpShape::SHAPE_32x128b:
5762 case Tcgen05CpShape::SHAPE_64x128b:
5763 return (mc == Tcgen05CpMulticast::WARPX2_01_23)
5767 llvm_unreachable(
"Invalid shape in tcgen05 cp Op");
5774 if (
shape == NVVM::Tcgen05LdStShape::SHAPE_16X128B)
5776 if (
shape == NVVM::Tcgen05LdStShape::SHAPE_16X256B)
5781LogicalResult Tcgen05LdOp::verify() {
5783 if (
getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())
5786 if (
getShape() != NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && getOffset())
5787 result =
emitError(
"offset argument is only supported for shape 16x32bx2");
5789 auto resTy = getRes().getType();
5790 unsigned resLen = isa<VectorType>(resTy)
5791 ? llvm::cast<VectorType>(resTy).getNumElements()
5794 result =
emitError(llvm::formatv(
"invalid result type length {0} for shape "
5795 "{1} in tcgen05.ld Op",
5796 resLen, stringifyEnum(
getShape())));
5801LogicalResult Tcgen05StOp::verify() {
5803 if (
getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())
5806 auto valTy = getVal().getType();
5807 unsigned valLen = isa<VectorType>(valTy)
5808 ? llvm::cast<VectorType>(valTy).getNumElements()
5811 result =
emitError(llvm::formatv(
"invalid input length {0} for shape "
5812 "{1} in tcgen05.st Op",
5813 valLen, stringifyEnum(
getShape())));
5823 if (
auto rangeAttr = op->
getAttrOfType<LLVM::ConstantRangeAttr>(
"range")) {
5824 setResultRanges(
result, {rangeAttr.getLower(), rangeAttr.getUpper(),
5825 rangeAttr.getLower(), rangeAttr.getUpper()});
5835 std::optional<LLVM::ConstantRangeAttr> rangeAttr) {
5839 const llvm::APInt &lower = rangeAttr->getLower();
5840 const llvm::APInt &upper = rangeAttr->getUpper();
5843 if (lower == upper && !lower.isMaxValue() && !lower.isMinValue()) {
5844 unsigned bitWidth = lower.getBitWidth();
5845 llvm::APInt minVal = llvm::APInt::getMinValue(bitWidth);
5846 llvm::APInt maxVal = llvm::APInt::getMaxValue(bitWidth);
5848 "invalid range attribute: Lower == Upper, but they aren't min (")
5849 << llvm::toString(minVal, 10,
false) <<
") or max ("
5850 << llvm::toString(maxVal, 10,
false)
5851 <<
") value! This is an invalid constant range.";
5858 llvm::IRBuilderBase &builder) {
5859 return builder.CreateBitCast(arg,
5860 llvm::Type::getInt32Ty(builder.getContext()));
5865 auto curOp = cast<NVVM::DotAccumulate4WayOp>(op);
5872 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;
5873 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;
5874 unsigned type = (isASigned << 1) | isBSigned;
5875 const llvm::Intrinsic::ID ids[] = {
5876 llvm::Intrinsic::nvvm_idp4a_u_u,
5877 llvm::Intrinsic::nvvm_idp4a_u_s,
5878 llvm::Intrinsic::nvvm_idp4a_s_u,
5879 llvm::Intrinsic::nvvm_idp4a_s_s,
5881 return {ids[type], args};
5886 auto curOp = cast<NVVM::DotAccumulate2WayOp>(op);
5891 args.push_back(builder.getInt1(curOp.getBHi()));
5894 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;
5895 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;
5896 unsigned type = (isASigned << 1) | isBSigned;
5897 const llvm::Intrinsic::ID ids[] = {
5898 llvm::Intrinsic::nvvm_idp2a_u_u,
5899 llvm::Intrinsic::nvvm_idp2a_u_s,
5900 llvm::Intrinsic::nvvm_idp2a_s_u,
5901 llvm::Intrinsic::nvvm_idp2a_s_s,
5903 return {ids[type], args};
5907 llvm::IRBuilderBase &builder) {
5908 return builder.CreateAddrSpaceCast(
5909 addr, builder.getPtrTy(llvm::NVPTXAS::ADDRESS_SPACE_ENTRY_PARAM));
5913PrefetchOp::getIntrinsicIDAndArgs(NVVM::PrefetchOp &op,
5915 llvm::IRBuilderBase &builder) {
5916 using MemSpace = NVVM::NVVMMemorySpace;
5917 using CacheLevel = NVVM::PrefetchCacheLevel;
5919 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = op.getCacheLevel();
5920 std::optional<NVVM::CacheEvictionPriority> evictPriority =
5921 op.getEvictPriority();
5922 unsigned addressSpace =
5923 llvm::cast<LLVM::LLVMPointerType>(op.getAddr().getType())
5931 if (op.getTensormap())
5932 return {llvm::Intrinsic::nvvm_prefetch_tensormap, args};
5934 assert(cacheLevel &&
"expected cache level for non-tensormap prefetch");
5936 if (op.getUniform() && *cacheLevel == CacheLevel::L1)
5937 return {llvm::Intrinsic::nvvm_prefetchu_L1, args};
5939 if (evictPriority && *cacheLevel == CacheLevel::L2) {
5940 switch (*evictPriority) {
5941 case NVVM::CacheEvictionPriority::EvictLast:
5942 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_last, args};
5943 case NVVM::CacheEvictionPriority::EvictNormal:
5944 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_normal, args};
5946 llvm_unreachable(
"Invalid cache eviction priority");
5950 switch (
static_cast<MemSpace
>(addressSpace)) {
5951 case MemSpace::Generic:
5952 return *cacheLevel == CacheLevel::L1
5954 :
NVVM::
IDArgPair({llvm::Intrinsic::nvvm_prefetch_L2, args});
5955 case MemSpace::Global:
5956 return *cacheLevel == CacheLevel::L1
5958 {llvm::Intrinsic::nvvm_prefetch_global_L1, args})
5960 {llvm::Intrinsic::nvvm_prefetch_global_L2, args});
5961 case MemSpace::Local:
5962 return *cacheLevel == CacheLevel::L1
5964 {llvm::Intrinsic::nvvm_prefetch_local_L1, args})
5966 {llvm::Intrinsic::nvvm_prefetch_local_L2, args});
5968 llvm_unreachable(
"Invalid pointer address space");
5972bool NVVM::InlinePtxOp::getAsmValues(
5976 for (
auto arg : getReadWriteArgs())
5978 for (
auto arg : getResults())
5980 for (
auto arg : getReadOnlyArgs())
5987NVVM::IDArgPair ClusterLaunchControlTryCancelOp::getIntrinsicIDAndArgs(
5989 auto curOp = cast<NVVM::ClusterLaunchControlTryCancelOp>(op);
5991 args.push_back(mt.
lookupValue(curOp.getSmemAddress()));
5992 args.push_back(mt.
lookupValue(curOp.getMbarrier()));
5994 llvm::Intrinsic::ID intrinsicID =
5995 curOp.getMulticast()
5997 nvvm_clusterlaunchcontrol_try_cancel_async_multicast_shared
5998 : llvm::Intrinsic::nvvm_clusterlaunchcontrol_try_cancel_async_shared;
6000 return {intrinsicID, args};
6003NVVM::IDArgPair ClusterLaunchControlQueryCancelOp::getIntrinsicIDAndArgs(
6005 auto curOp = cast<NVVM::ClusterLaunchControlQueryCancelOp>(op);
6007 args.push_back(mt.
lookupValue(curOp.getTryCancelResponse()));
6009 llvm::Intrinsic::ID intrinsicID;
6011 switch (curOp.getQueryType()) {
6012 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:
6014 llvm::Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_is_canceled;
6016 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:
6017 intrinsicID = llvm::Intrinsic::
6018 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_x;
6020 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:
6021 intrinsicID = llvm::Intrinsic::
6022 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_y;
6024 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:
6025 intrinsicID = llvm::Intrinsic::
6026 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_z;
6029 return {intrinsicID, args};
6034 llvm::IRBuilderBase &builder) {
6035 auto thisOp = cast<NVVM::PermuteOp>(op);
6036 NVVM::PermuteMode mode = thisOp.getMode();
6038 static constexpr llvm::Intrinsic::ID IDs[] = {
6039 llvm::Intrinsic::nvvm_prmt, llvm::Intrinsic::nvvm_prmt_f4e,
6040 llvm::Intrinsic::nvvm_prmt_b4e, llvm::Intrinsic::nvvm_prmt_rc8,
6041 llvm::Intrinsic::nvvm_prmt_ecl, llvm::Intrinsic::nvvm_prmt_ecr,
6042 llvm::Intrinsic::nvvm_prmt_rc16};
6044 unsigned modeIndex =
static_cast<unsigned>(mode);
6052 args.push_back(mt.
lookupValue(thisOp.getSelector()));
6054 return {IDs[modeIndex], args};
6059 auto thisOp = cast<NVVM::TensormapReplaceOp>(op);
6063 if (thisOp.getOrd())
6064 args.push_back(builder.getInt32(thisOp.getOrd().value()));
6065 if (thisOp.getNewValue())
6066 args.push_back(mt.
lookupValue(thisOp.getNewValue()));
6067 if (
auto attr = thisOp.getNewValueAttr()) {
6070 .Case<TensormapElemtypeAttr, TensormapInterleaveLayoutAttr,
6071 TensormapSwizzleModeAttr, TensormapSwizzleAtomicityAttr,
6072 TensormapFillModeAttr>([](
auto attr) {
6073 return static_cast<unsigned>(attr.getValue());
6075 .Default([](
auto attr) {
6076 llvm_unreachable(
"Invalid attribute type");
6079 args.push_back(builder.getInt32(val));
6082 static constexpr llvm::Intrinsic::ID IDs[] = {
6083 llvm::Intrinsic::nvvm_tensormap_replace_global_address,
6084 llvm::Intrinsic::nvvm_tensormap_replace_rank,
6085 llvm::Intrinsic::nvvm_tensormap_replace_box_dim,
6086 llvm::Intrinsic::nvvm_tensormap_replace_global_dim,
6087 llvm::Intrinsic::nvvm_tensormap_replace_global_stride,
6088 llvm::Intrinsic::nvvm_tensormap_replace_element_stride,
6089 llvm::Intrinsic::nvvm_tensormap_replace_elemtype,
6090 llvm::Intrinsic::nvvm_tensormap_replace_interleave_layout,
6091 llvm::Intrinsic::nvvm_tensormap_replace_swizzle_mode,
6092 llvm::Intrinsic::nvvm_tensormap_replace_swizzle_atomicity,
6093 llvm::Intrinsic::nvvm_tensormap_replace_fill_mode,
6096 unsigned fieldIndex =
static_cast<unsigned>(thisOp.getField());
6098 return {IDs[fieldIndex], args};
6105static llvm::nvvm::Tcgen05MMAKind
6108 case NVVM::Tcgen05MMAKind::F16:
6109 return llvm::nvvm::Tcgen05MMAKind::F16;
6110 case NVVM::Tcgen05MMAKind::TF32:
6111 return llvm::nvvm::Tcgen05MMAKind::TF32;
6112 case NVVM::Tcgen05MMAKind::F8F6F4:
6113 return llvm::nvvm::Tcgen05MMAKind::F8F6F4;
6114 case NVVM::Tcgen05MMAKind::I8:
6115 return llvm::nvvm::Tcgen05MMAKind::I8;
6116 case NVVM::Tcgen05MMAKind::TI16:
6117 return llvm::nvvm::Tcgen05MMAKind::TI16;
6118 case NVVM::Tcgen05MMAKind::MXF8F6F4:
6119 case NVVM::Tcgen05MMAKind::MXF4:
6120 case NVVM::Tcgen05MMAKind::MXF4NVF4:
6123 llvm_unreachable(
"Unsupported tcgen05.mma kind");
6129 llvm::IRBuilderBase &builder) {
6131 auto thisOp = cast<NVVM::Tcgen05MMAOp>(op);
6134 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6137 const bool isATensor = isa<llvm::PointerType>(
A->getType());
6140 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6141 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6142 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6144 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;
6145 using CtaGroupArray = std::array<EnableAShiftArray, 2>;
6146 using IsATensorArray = std::array<CtaGroupArray, 2>;
6147 using HasScaleInputDArray = std::array<IsATensorArray, 2>;
6148 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;
6151 static constexpr HasDisableOutputLaneArray tcgen05MMAIDs = {
6157 {llvm::Intrinsic::nvvm_tcgen05_mma_shared,
notIntrinsic},
6159 {llvm::Intrinsic::nvvm_tcgen05_mma_shared,
notIntrinsic}}},
6163 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,
6164 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,
6168 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,
6169 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,
6175 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d,
notIntrinsic},
6177 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d,
notIntrinsic}}},
6181 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,
6182 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,
6186 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,
6187 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,
6193 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1,
6196 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2,
6201 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1,
6203 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift,
6208 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2,
6210 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift,
6216 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1,
6220 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2,
6225 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1,
6227 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift},
6231 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2,
6233 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift,
6236 llvm::Value *ScaleInputD = mt.
lookupValue(thisOp.getScaleInputD());
6237 bool hasScaleInputD = ScaleInputD !=
nullptr;
6239 llvm::Value *DisableOutputLane =
6241 bool hasDisableOutputLane = DisableOutputLane !=
nullptr;
6243 const unsigned ctaGroup =
6246 llvm::Intrinsic::ID ID =
6247 tcgen05MMAIDs[hasDisableOutputLane][hasScaleInputD][isATensor]
6248 [ctaGroup - 1][thisOp.getAShift()];
6250 assert(ID !=
notIntrinsic &&
"Invalid intrinsic for Tcgen05MMAOp.");
6253 args.push_back(ScaleInputD);
6255 if (hasDisableOutputLane)
6256 args.push_back(DisableOutputLane);
6258 args.push_back(builder.getInt32(
6261 if (!hasDisableOutputLane)
6262 args.push_back(builder.getInt32(ctaGroup));
6265 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6268 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6275 NVVM::CTAGroupKind ctaGroup,
bool hasAShift,
6276 NVVM::Tcgen05MMACollectorOp collectorOp,
Location loc) {
6278 if (disableOutputLane) {
6279 mlir::VectorType disableOutputLaneType =
6280 cast<mlir::VectorType>(disableOutputLane.
getType());
6281 if ((ctaGroup == NVVM::CTAGroupKind::CTA_1 &&
6282 disableOutputLaneType.getNumElements() != 4) ||
6283 (ctaGroup == NVVM::CTAGroupKind::CTA_2 &&
6284 disableOutputLaneType.getNumElements() != 8))
6285 return emitError(loc) <<
"Disable Output Lane of length "
6286 << disableOutputLaneType.getNumElements()
6287 <<
" is incompatible with CtaGroupAttr";
6290 if (hasAShift && !isATensor)
6292 loc,
"A-shift can be applied only when matrix A is in tensor memory");
6294 if (hasAShift ==
true && (collectorOp == Tcgen05MMACollectorOp::FILL ||
6295 collectorOp == Tcgen05MMACollectorOp::USE))
6297 loc,
"Cannot use collector buffer operation fill or use with ashift");
6302LogicalResult Tcgen05MMAOp::verify() {
6304 getDisableOutputLane(), getCtaGroup(), getAShift(),
6305 getCollectorOp(), getLoc());
6315 auto thisOp = cast<NVVM::Tcgen05MMASparseOp>(op);
6318 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6321 bool isATensor = isa<llvm::PointerType>(
A->getType());
6324 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6325 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6326 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6327 args.push_back(mt.
lookupValue(thisOp.getSparseMetadata()));
6329 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;
6330 using CtaGroupArray = std::array<EnableAShiftArray, 2>;
6331 using IsATensorArray = std::array<CtaGroupArray, 2>;
6332 using HasScaleInputDArray = std::array<IsATensorArray, 2>;
6333 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;
6336 static constexpr HasDisableOutputLaneArray tcgen05MMASparseIDs = {
6342 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared,
notIntrinsic},
6344 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared,
notIntrinsic}}},
6348 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,
6349 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,
6353 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,
6354 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,
6360 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,
6363 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,
6368 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,
6369 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,
6373 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,
6374 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,
6381 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1,
6385 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2,
6390 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1,
6392 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift,
6397 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2,
6399 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift,
6405 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1,
6409 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2,
6414 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1,
6416 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift},
6420 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2,
6422 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift,
6425 llvm::Value *ScaleInputD = mt.
lookupValue(thisOp.getScaleInputD());
6426 bool hasScaleInputD = ScaleInputD !=
nullptr;
6428 llvm::Value *DisableOutputLane =
6430 bool hasDisableOutputLane = DisableOutputLane !=
nullptr;
6435 llvm::Intrinsic::ID ID =
6436 tcgen05MMASparseIDs[hasDisableOutputLane][hasScaleInputD][isATensor]
6437 [ctaGroup - 1][thisOp.getAShift()];
6439 assert(ID !=
notIntrinsic &&
"Invalid intrinsic for Tcgen05MMASparseOp.");
6442 args.push_back(ScaleInputD);
6444 if (hasDisableOutputLane)
6445 args.push_back(DisableOutputLane);
6447 args.push_back(builder.getInt32(
6450 if (!hasDisableOutputLane)
6451 args.push_back(builder.getInt32(ctaGroup));
6454 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6457 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6462LogicalResult Tcgen05MMASparseOp::verify() {
6464 getDisableOutputLane(), getCtaGroup(), getAShift(),
6465 getCollectorOp(), getLoc());
6475 auto thisOp = cast<NVVM::Tcgen05MMABlockScaleOp>(op);
6478 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6481 bool isATensor = isa<llvm::PointerType>(
A->getType());
6484 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6485 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6486 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6487 args.push_back(mt.
lookupValue(thisOp.getScaleA()));
6488 args.push_back(mt.
lookupValue(thisOp.getScaleB()));
6489 args.push_back(builder.getInt32(
6492 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6494 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6496 auto kind = thisOp.getKind();
6497 auto blockScale = thisOp.getBlockScale();
6498 llvm::Intrinsic::ID ID = [&]() {
6499 if (kind == NVVM::Tcgen05MMAKind::MXF8F6F4) {
6500 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6501 return isATensor ? llvm::Intrinsic::
6502 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale
6504 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale;
6505 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6508 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale_block32
6510 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale_block32;
6512 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4) {
6513 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6515 ? llvm::Intrinsic::nvvm_tcgen05_mma_tensor_mxf4_block_scale
6516 : llvm::Intrinsic::nvvm_tcgen05_mma_shared_mxf4_block_scale;
6517 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6518 return isATensor ? llvm::Intrinsic::
6519 nvvm_tcgen05_mma_tensor_mxf4_block_scale_block32
6521 nvvm_tcgen05_mma_shared_mxf4_block_scale_block32;
6523 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4NVF4) {
6524 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6527 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block32
6529 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block32;
6531 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {
6534 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block16
6536 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block16;
6539 llvm_unreachable(
"Invalid tcgen05.mma.block_scale attributes");
6546 NVVM::Tcgen05MMACollectorOp collectorOp, NVVM::Tcgen05MMAKind kind,
6547 NVVM::Tcgen05MMABlockScale blockScale,
Location loc) {
6548 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT &&
6549 kind == NVVM::Tcgen05MMAKind::MXF4NVF4)
6550 return emitError(loc,
"mxf4nvf4 requires block scale attribute");
6552 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16 &&
6553 kind != NVVM::Tcgen05MMAKind::MXF4NVF4)
6555 llvm::formatv(
"{} kind does not support block16 attribute",
6556 stringifyEnum(kind)));
6561LogicalResult Tcgen05MMABlockScaleOp::verify() {
6563 getBlockScale(), getLoc());
6573 auto thisOp = cast<NVVM::Tcgen05MMASparseBlockScaleOp>(op);
6576 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6579 bool isATensor = isa<llvm::PointerType>(
A->getType());
6582 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6583 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6584 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6585 args.push_back(mt.
lookupValue(thisOp.getSparseMetadata()));
6586 args.push_back(mt.
lookupValue(thisOp.getScaleA()));
6587 args.push_back(mt.
lookupValue(thisOp.getScaleB()));
6588 args.push_back(builder.getInt32(
6591 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6593 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6595 auto kind = thisOp.getKind();
6596 auto blockScale = thisOp.getBlockScale();
6597 llvm::Intrinsic::ID ID = [&]() {
6598 if (kind == NVVM::Tcgen05MMAKind::MXF8F6F4) {
6599 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6600 return isATensor ? llvm::Intrinsic::
6601 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale
6603 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale;
6604 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6607 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale_block32
6609 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale_block32;
6611 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4) {
6612 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6613 return isATensor ? llvm::Intrinsic::
6614 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale
6616 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale;
6617 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6620 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale_block32
6622 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale_block32;
6624 }
else if (kind == NVVM::Tcgen05MMAKind::MXF4NVF4) {
6625 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6628 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block32
6630 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block32;
6632 }
else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {
6635 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block16
6637 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block16;
6640 llvm_unreachable(
"Invalid tcgen05.mma.sp.block_scale attributes");
6646LogicalResult Tcgen05MMASparseBlockScaleOp::verify() {
6648 getBlockScale(), getLoc());
6658 auto thisOp = cast<NVVM::Tcgen05MMAWsOp>(op);
6661 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6664 bool isATensor = isa<llvm::PointerType>(
A->getType());
6667 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6668 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6669 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6671 mlir::Value ZeroColMask = thisOp.getZeroColMask();
6675 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor_zero_col_mask
6676 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared_zero_col_mask;
6678 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor
6679 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared;
6681 args.push_back(builder.getInt32(
6684 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorBBuffer())));
6686 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6698 auto thisOp = cast<NVVM::Tcgen05MMAWsSparseOp>(op);
6701 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6704 bool isATensor = isa<llvm::PointerType>(
A->getType());
6707 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6708 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6709 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6710 args.push_back(mt.
lookupValue(thisOp.getSparseMetadata()));
6712 mlir::Value ZeroColMask = thisOp.getZeroColMask();
6717 ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor_zero_col_mask
6718 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared_zero_col_mask;
6720 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor
6721 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared;
6723 args.push_back(builder.getInt32(
6726 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorBBuffer())));
6728 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOp())));
6739 auto thisOp = cast<Tcgen05MMADecompressBOp>(op);
6742 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6745 const bool isATensor = isa<llvm::PointerType>(
A->getType());
6748 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6749 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6750 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6751 args.push_back(mt.
lookupValue(thisOp.getDecompressBMetadata()));
6753 llvm::Value *DisableOutputLane =
6755 bool hasDisableOutputLane = DisableOutputLane !=
nullptr;
6757 NVVM::CTAGroupKind ctaGroup = thisOp.getCtaGroup();
6759 using namespace llvm::Intrinsic;
6760 ID intrinsicID = not_intrinsic;
6762 if (hasDisableOutputLane) {
6763 if (ctaGroup == NVVM::CTAGroupKind::CTA_1) {
6766 ? nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg1_decompress_b
6767 : nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg1_decompress_b;
6768 }
else if (ctaGroup == NVVM::CTAGroupKind::CTA_2) {
6771 ? nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg2_decompress_b
6772 : nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg2_decompress_b;
6774 llvm_unreachable(
"Unknown ctaGroup for tcgen05.mma.decompress_b");
6777 intrinsicID = isATensor ? nvvm_tcgen05_mma_tensor_f8f6f4_decompress_b
6778 : nvvm_tcgen05_mma_shared_f8f6f4_decompress_b;
6781 assert(intrinsicID != not_intrinsic &&
6782 "Invalid intrinsic for Tcgen05MMADecompressBOp.");
6784 if (hasDisableOutputLane)
6785 args.push_back(DisableOutputLane);
6791 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpA())));
6793 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6795 return {intrinsicID, args};
6798LogicalResult Tcgen05MMADecompressBOp::verify() {
6799 mlir::Value disableOutputLane = getDisableOutputLane();
6801 if (disableOutputLane) {
6802 NVVM::CTAGroupKind ctaGroup = getCtaGroup();
6804 mlir::VectorType disableOutputLaneType =
6805 cast<mlir::VectorType>(disableOutputLane.
getType());
6806 if ((ctaGroup == NVVM::CTAGroupKind::CTA_1 &&
6807 disableOutputLaneType.getNumElements() != 4) ||
6808 (ctaGroup == NVVM::CTAGroupKind::CTA_2 &&
6809 disableOutputLaneType.getNumElements() != 8))
6810 return emitOpError() <<
"Disable Output Lane of length "
6811 << disableOutputLaneType.getNumElements()
6812 <<
" is incompatible with CtaGroupAttr";
6824 auto thisOp = cast<Tcgen05MMABlockScaleDecompressBOp>(op);
6827 args.push_back(mt.
lookupValue(thisOp.getMatrixD()));
6830 const bool isATensor = isa<llvm::PointerType>(
A->getType());
6833 args.push_back(mt.
lookupValue(thisOp.getMatrixB()));
6834 args.push_back(mt.
lookupValue(thisOp.getIdesc()));
6835 args.push_back(mt.
lookupValue(thisOp.getEnableInputD()));
6836 args.push_back(mt.
lookupValue(thisOp.getScaleA()));
6837 args.push_back(mt.
lookupValue(thisOp.getScaleB()));
6838 args.push_back(mt.
lookupValue(thisOp.getDecompressBMetadata()));
6839 args.push_back(builder.getInt32(
6842 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpA())));
6844 builder.getInt32(
static_cast<unsigned>(thisOp.getCollectorOpB())));
6846 llvm::Intrinsic::ID intrinsicID =
6849 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale_block32_decompress_b
6851 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale_block32_decompress_b;
6853 return {intrinsicID, args};
6860#define TCGEN05LDRED(SHAPE, NUM, TYPE) \
6861 llvm::Intrinsic::nvvm_tcgen05_ld_red_##SHAPE##_##NUM##_##TYPE
6865 auto thisOp = cast<NVVM::Tcgen05LdRedOp>(op);
6868 mlir::VectorType VecResTy =
6869 cast<mlir::VectorType>(thisOp.getData().getType());
6870 unsigned Num = VecResTy.getNumElements();
6871 bool IsFloat = thisOp.getRedVal().getType().isF32();
6873 llvm::Intrinsic::ID Shape32x32b[][2] = {
6884 llvm::Intrinsic::ID Shape16x32bx2[][2] = {
6895 NVVM::Tcgen05LdStShape
shape = thisOp.getShape();
6896 unsigned ID = [&]() {
6899 unsigned idx = std::log2(Num);
6901 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
6902 return Shape32x32b[idx][IsFloat];
6903 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
6904 return Shape16x32bx2[idx][IsFloat];
6906 llvm_unreachable(
"unhandled tcgen05.ld lowering");
6912 if (
shape == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2)
6913 args.push_back(mt.
lookupValue(thisOp.getOffset()));
6916 builder.getInt32(thisOp.getOp() == NVVM::ReductionKind::MIN ? 0 : 1));
6919 args.push_back(builder.getInt1(
static_cast<unsigned>(thisOp.getAbs())));
6920 args.push_back(builder.getInt1(
static_cast<unsigned>(thisOp.getNan())));
6925LogicalResult Tcgen05LdRedOp::verify() {
6926 VectorType data = cast<VectorType>(getData().
getType());
6927 Type redVal = getRedVal().getType();
6929 if (data.getElementType() != redVal)
6931 "type of reduction value and element type of vector data should match");
6933 if (getOp() != NVVM::ReductionKind::MIN &&
6934 getOp() != NVVM::ReductionKind::MAX)
6935 return emitError(
"only min and max reduction kinds are supported");
6937 if (redVal.
isInteger() && (getAbs() || getNan())) {
6938 return emitError(
"abs or nan is only applicable for f32 type");
6948struct NVVMInlinerInterface final : DialectInlinerInterface {
6949 using DialectInlinerInterface::DialectInlinerInterface;
6950 bool isLegalToInline(Operation *, Region *,
bool, IRMapping &)
const final {
6957void NVVMDialect::initialize() {
6960#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"
6963#define GET_ATTRDEF_LIST
6964#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"
6969 allowUnknownOperations();
6970 addInterfaces<NVVMInlinerInterface>();
6971 declarePromisedInterface<ConvertToLLVMPatternInterface, NVVMDialect>();
6972 declarePromisedInterface<gpu::TargetAttrInterface, NVVMTargetAttr>();
6975LogicalResult NVVMDialect::verifyOperationAttribute(
Operation *op,
6977 StringAttr attrName = attr.
getName();
6979 if (attrName == NVVMDialect::getKernelFuncAttrName()) {
6980 if (!isa<LLVM::LLVMFuncOp>(op)) {
6981 return op->
emitError() <<
"'" << NVVMDialect::getKernelFuncAttrName()
6982 <<
"' attribute attached to unexpected op";
6987 if (attrName == NVVMDialect::getMaxntidAttrName() ||
6988 attrName == NVVMDialect::getReqntidAttrName() ||
6989 attrName == NVVMDialect::getClusterDimAttrName()) {
6990 auto values = llvm::dyn_cast<DenseI32ArrayAttr>(attr.
getValue());
6991 if (!values || values.empty() || values.size() > 3) {
6994 <<
"' attribute must be integer array with maximum 3 index";
6999 if (attrName == NVVMDialect::getMinctasmAttrName() ||
7000 attrName == NVVMDialect::getMaxnregAttrName() ||
7001 attrName == NVVMDialect::getClusterMaxBlocksAttrName()) {
7002 if (!llvm::dyn_cast<IntegerAttr>(attr.
getValue())) {
7004 <<
"'" << attrName <<
"' attribute must be integer constant";
7008 if (attrName == NVVMDialect::getBlocksAreClustersAttrName()) {
7009 if (!op->
hasAttr(NVVMDialect::getReqntidAttrName()) ||
7010 !op->
hasAttr(NVVMDialect::getClusterDimAttrName())) {
7012 <<
"'" << attrName <<
"' attribute must be used along with " <<
"'"
7013 << NVVMDialect::getReqntidAttrName() <<
"' and " <<
"'"
7014 << NVVMDialect::getClusterDimAttrName() <<
"'";
7021LogicalResult NVVMDialect::verifyRegionArgAttribute(
Operation *op,
7022 unsigned regionIndex,
7025 auto funcOp = dyn_cast<FunctionOpInterface>(op);
7029 bool isKernel = op->
hasAttr(NVVMDialect::getKernelFuncAttrName());
7030 StringAttr attrName = argAttr.
getName();
7031 if (attrName == NVVM::NVVMDialect::getGridConstantAttrName()) {
7035 <<
"' attribute must be present only on kernel arguments";
7037 if (!isa<UnitAttr>(argAttr.
getValue()))
7038 return op->
emitError() <<
"'" << attrName <<
"' must be a unit attribute";
7039 if (!funcOp.getArgAttr(argIndex, LLVM::LLVMDialect::getByValAttrName())) {
7042 <<
"' attribute requires the argument to also have attribute '"
7043 << LLVM::LLVMDialect::getByValAttrName() <<
"'";
7054unsigned NVVMMemorySpaceAttr::getAddressSpace()
const {
7055 return static_cast<unsigned>(getValue());
7058bool NVVMMemorySpaceAttr::isValidLoad(
7059 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,
7060 const ::mlir::DataLayout *dataLayout,
7066bool NVVMMemorySpaceAttr::isValidStore(
7067 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,
7068 const ::mlir::DataLayout *dataLayout,
7074bool NVVMMemorySpaceAttr::isValidAtomicOp(
7075 ptr::AtomicBinOp op,
Type type, ptr::AtomicOrdering ordering,
7076 std::optional<int64_t> alignment, const ::mlir::DataLayout *dataLayout,
7079 assert(
false &&
"unimplemented, see TODO in the source.");
7083bool NVVMMemorySpaceAttr::isValidAtomicXchg(
7084 Type type, ptr::AtomicOrdering successOrdering,
7085 ptr::AtomicOrdering failureOrdering, std::optional<int64_t> alignment,
7086 const ::mlir::DataLayout *dataLayout,
7089 assert(
false &&
"unimplemented, see TODO in the source.");
7093bool NVVMMemorySpaceAttr::isValidAddrSpaceCast(
7097 assert(
false &&
"unimplemented, see TODO in the source.");
7101bool NVVMMemorySpaceAttr::isValidPtrIntCast(
7106 assert(
false &&
"unimplemented, see TODO in the source.");
7115 int optLevel, StringRef triple, StringRef chip,
7116 StringRef features, DictionaryAttr flags,
7118 if (optLevel < 0 || optLevel > 3) {
7119 emitError() <<
"The optimization level must be a number between 0 and 3.";
7122 if (triple.empty()) {
7123 emitError() <<
"The target triple cannot be empty.";
7127 emitError() <<
"The target chip cannot be empty.";
7130 if (files && !llvm::all_of(files, [](::mlir::Attribute attr) {
7131 return mlir::isa_and_nonnull<StringAttr>(attr);
7133 emitError() <<
"All the elements in the `link` array must be strings.";
7139LogicalResult NVVMTargetAttr::verifyTarget(
Operation *gpuModule) {
7140 if (!getVerifyTarget())
7143 auto gpuModuleOp = llvm::dyn_cast<gpu::GPUModuleOp>(gpuModule);
7146 "NVVM target attribute must be attached to a GPU module");
7149 const unsigned targetFullSmVersion =
7153 "Minimum NVVM target SM version is sm_20");
7157 ->
walk([&](Operation *op) {
7158 if (
auto reqOp = llvm::dyn_cast<NVVM::RequiresSMInterface>(op)) {
7159 const NVVMCheckSMVersion requirement =
7160 reqOp.getRequiredMinSMVersion();
7162 op->
emitOpError() <<
"is not supported on " << getChip();
7174#define GET_OP_CLASSES
7175#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"
7177#define GET_ATTRDEF_CLASSES
7178#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 printCTAGroup(OpAsmPrinter &printer, Operation *, NVVM::CTAGroupKindAttr groupAttr)
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)
static ParseResult parseCTAGroup(OpAsmParser &parser, NVVM::CTAGroupKindAttr &groupAttr)
static ParseResult parseEnumKeyword(OpAsmParser &parser, AttrTy &attr)
#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf)
static llvm::nvvm::Tcgen05MMAKind getNVVMTcgen05MMAKind(NVVM::Tcgen05MMAKind kind)
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 void printMmaUnitProperty(OpAsmPrinter &printer, bool &isFirst, StringRef keyword)
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)
static void printMmaProperty(OpAsmPrinter &printer, bool &isFirst, StringRef keyword, AttrTy value)
static bool isMmaPropertyName(StringRef name)
#define GET_F32x2_TO_F6x2_ID(type, has_relu)
static llvm::Value * getAsPackedI32(llvm::Value *arg, llvm::IRBuilderBase &builder)
static void printMmaEnumProperty(OpAsmPrinter &printer, bool &isFirst, StringRef keyword, AttrTy value)
#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 ParseResult parseMmaProperties(OpAsmParser &parser, NamedAttrList &attributes, ArrayRef< StringRef > allowedKeywords, ArrayRef< StringRef > requiredProperties)
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)
LogicalResult CpAsyncBulkTensorOverrideAddrCommonVerifier(OperandRange coordinates, OperandRange tensorSize, OperandRange lowerStride, Value upperStride, bool isTile, Location loc)
static ParseResult parseMmaEnumPropertyValue(OpAsmParser &parser, NamedAttrList &attributes, StringRef name)
#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 ParseResult parseMmaPropertyValue(OpAsmParser &parser, NamedAttrList &attributes, StringRef name)
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.
ParseResult parseKeywordOrString(std::string *result)
Parse a keyword or a quoted string.
virtual ParseResult parseCustomAttributeWithFallback(Attribute &result, Type type, function_ref< ParseResult(Attribute &result, Type type)> parseAttribute)=0
Parse a custom attribute with the provided callback, unless the next token is #, in which case the ge...
virtual ParseResult parseEqual()=0
Parse a = token.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
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)
void printStrippedAttrOrType(AttrOrType attrOrType)
Print the provided attribute in the context of an operation custom printer/parser: this will invoke d...
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.
void append(StringRef name, Attribute attr)
Add an attribute with the specified name.
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.
This class implements the operand iterators for the Operation class.
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()
The OpAsmOpInterface, see OpAsmInterface.td for more details.
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.