4842 lines · cpp
1//===- NVVMDialect.cpp - NVVM IR Ops and Dialect registration -------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8//9// This file defines the types and operation details for the NVVM IR dialect in10// MLIR, and the LLVM IR dialect. It also registers the dialect.11//12// The NVVM dialect only contains GPU specific additions on top of the general13// LLVM dialect.14//15//===----------------------------------------------------------------------===//16 17#include "mlir/Dialect/LLVMIR/NVVMDialect.h"18 19#include "mlir/Conversion/ConvertToLLVM/ToLLVMInterface.h"20#include "mlir/Dialect/GPU/IR/CompilationInterfaces.h"21#include "mlir/Dialect/GPU/IR/GPUDialect.h"22#include "mlir/IR/Builders.h"23#include "mlir/IR/BuiltinAttributes.h"24#include "mlir/IR/BuiltinTypes.h"25#include "mlir/IR/Diagnostics.h"26#include "mlir/IR/DialectImplementation.h"27#include "mlir/IR/MLIRContext.h"28#include "mlir/IR/Operation.h"29#include "mlir/IR/OperationSupport.h"30#include "mlir/IR/Types.h"31#include "llvm/ADT/STLExtras.h"32#include "llvm/ADT/TypeSwitch.h"33#include "llvm/IR/IRBuilder.h"34#include "llvm/IR/NVVMIntrinsicUtils.h"35#include "llvm/Support/Casting.h"36#include "llvm/Support/FormatVariadic.h"37#include "llvm/Support/NVPTXAddrSpace.h"38#include "llvm/Support/raw_ostream.h"39#include <cassert>40#include <optional>41#include <string>42 43using namespace mlir;44using namespace NVVM;45 46#include "mlir/Dialect/LLVMIR/NVVMOpsDialect.cpp.inc"47#include "mlir/Dialect/LLVMIR/NVVMOpsEnums.cpp.inc"48 49static constexpr unsigned notIntrinsic = llvm::Intrinsic::not_intrinsic;50 51//===----------------------------------------------------------------------===//52// Helper/Utility methods53//===----------------------------------------------------------------------===//54 55static bool isPtrInAddrSpace(mlir::Value ptr, NVVMMemorySpace targetAS) {56 auto ptrTy = llvm::cast<LLVM::LLVMPointerType>(ptr.getType());57 return ptrTy.getAddressSpace() == static_cast<unsigned>(targetAS);58}59 60static bool isPtrInGenericSpace(mlir::Value ptr) {61 return isPtrInAddrSpace(ptr, NVVMMemorySpace::Generic);62}63 64static bool isPtrInSharedCTASpace(mlir::Value ptr) {65 return isPtrInAddrSpace(ptr, NVVMMemorySpace::Shared);66}67 68static bool isPtrInSharedClusterSpace(mlir::Value ptr) {69 return isPtrInAddrSpace(ptr, NVVMMemorySpace::SharedCluster);70}71 72static llvm::Value *castPtrToAddrSpace(llvm::IRBuilderBase &builder,73 llvm::Value *ptr,74 NVVMMemorySpace targetAS) {75 unsigned AS = static_cast<unsigned>(targetAS);76 return builder.CreateAddrSpaceCast(77 ptr, llvm::PointerType::get(builder.getContext(), AS));78}79 80// Helper method to convert CtaGroupKind in NVVM Dialect to CtaGroupKind in LLVM81static llvm::nvvm::CTAGroupKind82getNVVMCtaGroupKind(NVVM::CTAGroupKind ctaGroup) {83 switch (ctaGroup) {84 case NVVM::CTAGroupKind::CTA_1:85 return llvm::nvvm::CTAGroupKind::CG_1;86 case NVVM::CTAGroupKind::CTA_2:87 return llvm::nvvm::CTAGroupKind::CG_2;88 }89 llvm_unreachable("unsupported cta_group value");90}91 92//===----------------------------------------------------------------------===//93// Verifier methods94//===----------------------------------------------------------------------===//95 96// This verifier is shared among the following Ops:97// CpAsyncBulkTensorSharedCTAToGlobalOp (TMA Store)98// CpAsyncBulkTensorReduceOp (TMA Store-Reduce)99static LogicalResult cpAsyncBulkTensorCommonVerifier(size_t tensorDims,100 bool isIm2Col,101 size_t numIm2ColOffsets,102 Location loc) {103 if (tensorDims < 1 || tensorDims > 5)104 return emitError(loc, "expects coordinates between 1 to 5 dimension");105 106 // For Im2Col mode, there are two constraints:107 if (isIm2Col) {108 // 1. Tensor must always be at least 3-d.109 if (tensorDims < 3)110 return emitError(111 loc,112 "to use im2col mode, the tensor has to be at least 3-dimensional");113 // 2. When there are Im2ColOffsets, they must be (Dims - 2) in number.114 if (numIm2ColOffsets && (tensorDims != (numIm2ColOffsets + 2)))115 return emitError(116 loc, "im2col offsets must be 2 less than number of coordinates");117 }118 return success();119}120 121LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOp::verify() {122 TMAStoreMode mode = getMode();123 // We lower through inline-ptx when getPredicate() is true.124 // a) Only TILE mode is supported125 // b) Cache-hint is not supported126 if (getPredicate()) {127 if (mode != TMAStoreMode::TILE)128 return emitError("Inline-ptx lowering supported only for Tile mode.");129 if (getL2CacheHint())130 return emitError("Inline-ptx lowering unsupported with L2 cache-hint.");131 }132 133 size_t dims = getCoordinates().size();134 switch (mode) {135 case TMAStoreMode::TILE:136 return cpAsyncBulkTensorCommonVerifier(dims, false, 0, getLoc());137 case TMAStoreMode::IM2COL:138 return cpAsyncBulkTensorCommonVerifier(dims, true, 0, getLoc());139 case TMAStoreMode::TILE_SCATTER4:140 if (dims != 5)141 return emitError("Scatter4 mode expects 5 coordinates");142 }143 return success();144}145 146LogicalResult CpAsyncOp::verify() {147 if (getModifier() != LoadCacheModifierKind::CG &&148 getModifier() != LoadCacheModifierKind::CA)149 return emitError("Only CG and CA cache modifiers are supported.");150 if (getSize() != 4 && getSize() != 8 && getSize() != 16)151 return emitError("expected byte size to be either 4, 8 or 16.");152 if (getModifier() == LoadCacheModifierKind::CG && getSize() != 16)153 return emitError("CG cache modifier is only support for 16 bytes copy.");154 return success();155}156 157// This verify params can be shared across TMA Load and Prefetch Ops.158static LogicalResult verifyTMALoadParams(size_t tensorDims, size_t numIm2colOff,159 TMALoadMode mode, Location loc) {160 if (tensorDims < 1 || tensorDims > 5)161 return emitError(loc, "expects coordinates between 1 to 5 dimension");162 163 auto checkTMALoadParams = [&](TMALoadMode mode, bool isIm2col,164 size_t expectedIm2colOff) -> LogicalResult {165 if (isIm2col && (tensorDims < 3))166 return emitError(loc)167 << "to use " << stringifyEnum(mode)168 << " mode, the tensor has to be at least 3-dimensional";169 170 if (numIm2colOff != expectedIm2colOff)171 return emitError(loc) << " im2col offsets expected " << expectedIm2colOff172 << " (provided " << numIm2colOff << ")";173 174 return success();175 };176 177 switch (mode) {178 case TMALoadMode::TILE:179 return checkTMALoadParams(mode, false, 0);180 case TMALoadMode::IM2COL:181 return checkTMALoadParams(mode, true, tensorDims - 2);182 case TMALoadMode::IM2COL_W:183 case TMALoadMode::IM2COL_W_128:184 return checkTMALoadParams(mode, true, 2);185 case TMALoadMode::TILE_GATHER4:186 return (tensorDims == 5)187 ? checkTMALoadParams(mode, false, 0)188 : emitError(loc, "Gather4 mode expects 5 coordinates");189 }190 return success();191}192 193LogicalResult CpAsyncBulkTensorPrefetchOp::verify() {194 return verifyTMALoadParams(getCoordinates().size(), getIm2colOffsets().size(),195 getMode(), getLoc());196}197 198LogicalResult CpAsyncBulkTensorGlobalToSharedClusterOp::verify() {199 TMALoadMode mode = getMode();200 bool isCTAOnly = getIsCTAOnly();201 if (getPredicate()) { // Inline-asm based lowering202 if (isCTAOnly)203 return emitError("Predicate is supported only for shared::cluster mode.");204 if (mode != TMALoadMode::TILE && mode != TMALoadMode::IM2COL)205 return emitError(206 "Predicate is supported only for Tile and Im2col modes.");207 } else { // Intrinsics-based lowering208 NVVMMemorySpace expectedAS =209 isCTAOnly ? NVVMMemorySpace::Shared : NVVMMemorySpace::SharedCluster;210 unsigned AS = llvm::cast<LLVM::LLVMPointerType>(getDstMem().getType())211 .getAddressSpace();212 if (AS != expectedAS)213 return emitError()214 << (isCTAOnly215 ? "Shared::cta destination requires address-space 3."216 : "Shared::cluster destination requires address-space 7.");217 // Checks specific to shared::cta mode218 if (isCTAOnly) {219 if (getMulticastMask())220 return emitError("Multicast is not supported with shared::cta mode.");221 if (getGroup())222 return emitError("CTAGroup is not supported with shared::cta mode.");223 }224 }225 226 return verifyTMALoadParams(getCoordinates().size(), getIm2colOffsets().size(),227 getMode(), getLoc());228}229 230LogicalResult CpAsyncBulkTensorReduceOp::verify() {231 TMAStoreMode mode = getMode();232 size_t dims = getCoordinates().size();233 switch (mode) {234 case TMAStoreMode::TILE:235 return cpAsyncBulkTensorCommonVerifier(dims, false, 0, getLoc());236 case TMAStoreMode::IM2COL:237 return cpAsyncBulkTensorCommonVerifier(dims, true, 0, getLoc());238 case TMAStoreMode::TILE_SCATTER4:239 return emitError("Scatter mode unsupported for CpAsyncBulkTensorReduceOp");240 }241 return success();242}243 244LogicalResult CpAsyncBulkGlobalToSharedClusterOp::verify() {245 bool isSharedCTA = isPtrInSharedCTASpace(getDstMem());246 if (isSharedCTA && getMulticastMask())247 return emitError("Multicast is not supported with shared::cta mode.");248 249 return success();250}251 252static LogicalResult verifyMBarrierArriveLikeOp(Operation *op, Value addr,253 NVVM::MemScopeKind scope,254 Value retVal = nullptr) {255 bool isSharedCluster = isPtrInSharedClusterSpace(addr);256 if (scope != NVVM::MemScopeKind::CTA && scope != NVVM::MemScopeKind::CLUSTER)257 return op->emitError("mbarrier scope must be either CTA or Cluster");258 259 bool hasRetValue = static_cast<bool>(retVal);260 if (isSharedCluster && hasRetValue)261 return op->emitError(262 "mbarrier in shared_cluster space cannot return any value");263 264 return success();265}266 267LogicalResult MBarrierArriveOp::verify() {268 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),269 getRes());270}271 272LogicalResult MBarrierArriveDropOp::verify() {273 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),274 getRes());275}276 277LogicalResult MBarrierArriveExpectTxOp::verify() {278 // The inline-ptx version of this Op does not support all features.279 // With predicate, this Op lowers to inline-ptx. So, verify and280 // error-out if there are unsupported features.281 if (getPredicate()) {282 if (getScope() != NVVM::MemScopeKind::CTA)283 return emitError("mbarrier scope must be CTA when using predicate");284 285 if (isPtrInSharedClusterSpace(getAddr()))286 return emitError("mbarrier in shared_cluster space is not supported when "287 "using predicate");288 289 if (getRes())290 return emitError("return-value is not supported when using predicate");291 292 if (getRelaxed() == true)293 return emitError("mbarrier with relaxed semantics is not supported when "294 "using predicate");295 }296 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),297 getRes());298}299 300LogicalResult MBarrierArriveDropExpectTxOp::verify() {301 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),302 getRes());303}304 305LogicalResult MBarrierExpectTxOp::verify() {306 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope());307}308 309LogicalResult MBarrierCompleteTxOp::verify() {310 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope());311}312 313LogicalResult ConvertFloatToTF32Op::verify() {314 using RndMode = NVVM::FPRoundingMode;315 switch (getRnd()) {316 case RndMode::RNA:317 if (getRelu())318 return emitError("Relu not supported with rna rounding mode.");319 break;320 case RndMode::RN:321 case RndMode::RZ:322 break;323 default:324 return emitError(325 "Only {rn,rz,rna} rounding modes supported for ConvertFloatToTF32Op.");326 }327 return success();328}329 330LogicalResult ConvertF32x2ToF6x2Op::verify() {331 mlir::MLIRContext *ctx = getContext();332 333 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy())) {334 return emitOpError("Only ")335 << mlir::Float6E2M3FNType::get(ctx) << " and "336 << mlir::Float6E3M2FNType::get(ctx)337 << " types are supported for conversions from f32x2 to f6x2.";338 }339 return success();340}341 342LogicalResult ConvertF32x2ToF8x2Op::verify() {343 using RndMode = NVVM::FPRoundingMode;344 using SatMode = NVVM::SaturationMode;345 346 bool isRoundingModeRN = getRnd() == RndMode::RN;347 bool isRoundingModeRZ = getRnd() == RndMode::RZ;348 bool isRoundingModeRP = getRnd() == RndMode::RP;349 bool isSatFinite = getSat() == SatMode::SATFINITE;350 351 bool hasRelu = getRelu();352 353 mlir::MLIRContext *ctx = getContext();354 355 return llvm::TypeSwitch<mlir::Type, LogicalResult>(getDstTy())356 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(357 [&](mlir::Type) -> LogicalResult {358 if (!isRoundingModeRN) {359 return emitOpError("Only RN rounding mode is supported for "360 "conversions from f32x2 to ")361 << mlir::Float8E4M3FNType::get(ctx) << " and "362 << mlir::Float8E5M2Type::get(ctx) << " types";363 }364 if (!isSatFinite) {365 return emitOpError("Only SATFINITE saturation mode is supported "366 "for conversions "367 "from f32x2 to ")368 << mlir::Float8E4M3FNType::get(ctx) << " and "369 << mlir::Float8E5M2Type::get(ctx) << " types";370 }371 return success();372 })373 .Case<mlir::Float8E8M0FNUType>([&](mlir::Type) -> LogicalResult {374 if (!(isRoundingModeRZ || isRoundingModeRP)) {375 return emitOpError("Only RZ and RP rounding modes are supported for "376 "conversions from f32x2 to ")377 << mlir::Float8E8M0FNUType::get(ctx) << " type";378 }379 if (hasRelu) {380 return emitOpError("relu not supported for conversions to ")381 << mlir::Float8E8M0FNUType::get(ctx) << " type";382 }383 return success();384 })385 .Default([&](mlir::Type) {386 return emitOpError("Only ")387 << mlir::Float8E4M3FNType::get(ctx) << ", "388 << mlir::Float8E5M2Type::get(ctx) << ", and "389 << mlir::Float8E8M0FNUType::get(ctx)390 << " types are "391 "supported for conversions from f32x2 to f8x2";392 });393}394 395LogicalResult ConvertF16x2ToF8x2Op::verify() {396 mlir::MLIRContext *ctx = getContext();397 398 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy())) {399 return emitOpError("Only ")400 << mlir::Float8E4M3FNType::get(ctx) << " and "401 << mlir::Float8E5M2Type::get(ctx)402 << " types are supported for conversions from f16x2 to f8x2.";403 }404 return success();405}406 407LogicalResult ConvertBF16x2ToF8x2Op::verify() {408 using RndMode = NVVM::FPRoundingMode;409 410 if (!llvm::isa<mlir::Float8E8M0FNUType>(getDstTy()))411 return emitOpError("Only ") << mlir::Float8E8M0FNUType::get(getContext())412 << " type is supported for conversions from "413 "bf16x2 to f8x2.";414 415 auto rnd = getRnd();416 if (!(rnd == RndMode::RZ || rnd == RndMode::RP))417 return emitOpError("Only RZ and RP rounding modes are supported for "418 "conversions from bf16x2 to f8x2.");419 420 return success();421}422 423LogicalResult ConvertF32x2ToF4x2Op::verify() {424 mlir::MLIRContext *ctx = getContext();425 426 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))427 return emitOpError("Only ")428 << mlir::Float4E2M1FNType::get(ctx)429 << " type is supported for conversions from f32x2 to f4x2.";430 431 return success();432}433 434LogicalResult ConvertF8x2ToF16x2Op::verify() {435 mlir::MLIRContext *ctx = getContext();436 437 if (!llvm::isa<Float8E4M3FNType, Float8E5M2Type>(getSrcType()))438 return emitOpError("Only ")439 << mlir::Float8E4M3FNType::get(ctx) << " and "440 << mlir::Float8E5M2Type::get(ctx)441 << " types are supported for conversions from f8x2 to f16x2.";442 443 return success();444}445 446LogicalResult ConvertF8x2ToBF16x2Op::verify() {447 mlir::MLIRContext *ctx = getContext();448 if (!llvm::isa<Float8E8M0FNUType>(getSrcType()))449 return emitOpError("Only ")450 << mlir::Float8E8M0FNUType::get(ctx)451 << " type is supported for conversions from f8x2 to bf16x2.";452 453 return success();454}455 456LogicalResult ConvertF6x2ToF16x2Op::verify() {457 mlir::MLIRContext *ctx = getContext();458 459 if (!llvm::isa<Float6E2M3FNType, Float6E3M2FNType>(getSrcType()))460 return emitOpError("Only ")461 << mlir::Float6E2M3FNType::get(ctx) << " and "462 << mlir::Float6E3M2FNType::get(ctx)463 << " types are supported for conversions from f6x2 to f16x2.";464 465 return success();466}467 468LogicalResult ConvertF4x2ToF16x2Op::verify() {469 mlir::MLIRContext *ctx = getContext();470 471 if (!llvm::isa<Float4E2M1FNType>(getSrcType()))472 return emitOpError("Only ")473 << mlir::Float4E2M1FNType::get(ctx)474 << " type is supported for conversions from f4x2 to f16x2.";475 476 return success();477}478 479LogicalResult PermuteOp::verify() {480 using Mode = NVVM::PermuteMode;481 bool hasHi = static_cast<bool>(getHi());482 483 switch (getMode()) {484 case Mode::DEFAULT:485 case Mode::F4E:486 case Mode::B4E:487 if (!hasHi)488 return emitError("mode '")489 << stringifyPermuteMode(getMode()) << "' requires 'hi' operand.";490 break;491 case Mode::RC8:492 case Mode::ECL:493 case Mode::ECR:494 case Mode::RC16:495 if (hasHi)496 return emitError("mode '") << stringifyPermuteMode(getMode())497 << "' does not accept 'hi' operand.";498 break;499 }500 501 return success();502}503 504//===----------------------------------------------------------------------===//505// Stochastic Rounding Conversion Ops506//===----------------------------------------------------------------------===//507 508static LogicalResult verifyConvertF32x2ToFP16x2Op(Twine dstType,509 FPRoundingMode rnd,510 bool hasRandomBits,511 Operation *op) {512 static constexpr FPRoundingMode validRndModes[] = {513 FPRoundingMode::RN, FPRoundingMode::RZ, FPRoundingMode::RS};514 515 if (!llvm::is_contained(validRndModes, rnd)) {516 return op->emitOpError(517 "Only RN, RZ, and RS rounding modes are supported for "518 "conversions from f32x2 to ")519 << dstType << ".";520 }521 522 if (rnd == FPRoundingMode::RS) {523 if (!hasRandomBits) {524 return op->emitOpError("random_bits is required for RS rounding mode.");525 }526 } else {527 if (hasRandomBits) {528 return op->emitOpError(529 "random_bits not supported for RN and RZ rounding modes.");530 }531 }532 533 return success();534}535 536LogicalResult ConvertF32x2ToF16x2Op::verify() {537 return verifyConvertF32x2ToFP16x2Op("f16x2", getRnd(),538 getRandomBits() ? true : false, *this);539}540 541LogicalResult ConvertF32x2ToBF16x2Op::verify() {542 return verifyConvertF32x2ToFP16x2Op("bf16x2", getRnd(),543 getRandomBits() ? true : false, *this);544}545 546LogicalResult ConvertF32x4ToF8x4Op::verify() {547 mlir::MLIRContext *ctx = getContext();548 549 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy()))550 return emitOpError("Only ")551 << mlir::Float8E4M3FNType::get(ctx) << " and "552 << mlir::Float8E5M2Type::get(ctx)553 << " types are supported for conversions from f32x4 to f8x4.";554 555 return success();556}557 558LogicalResult ConvertF32x4ToF6x4Op::verify() {559 mlir::MLIRContext *ctx = getContext();560 561 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy()))562 return emitOpError("Only ")563 << mlir::Float6E2M3FNType::get(ctx) << " and "564 << mlir::Float6E3M2FNType::get(ctx)565 << " types are supported for conversions from f32x4 to f6x4.";566 567 return success();568}569 570LogicalResult ConvertF32x4ToF4x4Op::verify() {571 mlir::MLIRContext *ctx = getContext();572 573 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))574 return emitOpError("Only ") << mlir::Float4E2M1FNType::get(ctx)575 << " type is supported for conversions from "576 "f32x4 to f4x4.";577 578 return success();579}580 581LogicalResult BulkStoreOp::verify() {582 if (getInitVal() != 0)583 return emitOpError("only 0 is supported for initVal, got ") << getInitVal();584 return success();585}586 587LogicalResult PMEventOp::verify() {588 auto eventId = getEventId();589 auto maskedEventId = getMaskedEventId();590 if (!maskedEventId && !eventId) {591 return emitOpError() << "either `id` or `mask` must be set";592 }593 594 if (maskedEventId && eventId) {595 return emitOpError() << "`id` and `mask` cannot be set at the same time";596 }597 598 if (eventId) {599 if (eventId < 0 || eventId > 15) {600 return emitOpError() << "`id` must be between 0 and 15";601 }602 }603 604 return llvm::success();605}606 607// Given the element type of an operand and whether or not it is an accumulator,608// this function returns the PTX type (`NVVM::MMATypes`) that corresponds to the609// operand's element type.610std::optional<mlir::NVVM::MMATypes>611MmaOp::inferOperandMMAType(Type operandElType, bool isAccumulator) {612 auto half2Type =613 VectorType::get(2, Float16Type::get(operandElType.getContext()));614 if (operandElType.isF64())615 return NVVM::MMATypes::f64;616 if (operandElType.isF16() || operandElType == half2Type)617 return NVVM::MMATypes::f16;618 if (operandElType.isF32() && isAccumulator)619 return NVVM::MMATypes::f32;620 if (operandElType.isF32() && !isAccumulator)621 return NVVM::MMATypes::tf32;622 if (llvm::isa<IntegerType>(operandElType)) {623 if (isAccumulator)624 return NVVM::MMATypes::s32;625 return std::nullopt;626 }627 628 if (auto structType = llvm::dyn_cast<LLVM::LLVMStructType>(operandElType)) {629 if (structType.getBody().empty())630 return std::nullopt;631 return inferOperandMMAType(structType.getBody()[0], isAccumulator);632 }633 634 return std::nullopt;635}636 637static bool isInt4PtxType(MMATypes type) {638 return (type == MMATypes::u4 || type == MMATypes::s4);639}640 641static bool isInt8PtxType(MMATypes type) {642 return (type == MMATypes::u8 || type == MMATypes::s8);643}644 645static bool isIntegerPtxType(MMATypes type) {646 return isInt4PtxType(type) || isInt8PtxType(type) || type == MMATypes::b1 ||647 type == MMATypes::s32;648}649 650MMATypes MmaOp::accumPtxType() {651 std::optional<mlir::NVVM::MMATypes> val = inferOperandMMAType(652 getODSOperands(2).getTypes().front(), /*isAccumulator=*/true);653 assert(val.has_value() && "accumulator PTX type should always be inferrable");654 return val.value();655}656 657MMATypes MmaOp::resultPtxType() {658 std::optional<mlir::NVVM::MMATypes> val =659 inferOperandMMAType(getResult().getType(), /*isAccumulator=*/true);660 assert(val.has_value() && "result PTX type should always be inferrable");661 return val.value();662}663 664void MmaOp::print(OpAsmPrinter &p) {665 SmallVector<Type, 4> regTypes;666 struct OperandFragment {667 StringRef operandName;668 StringRef ptxTypeAttr;669 SmallVector<Value, 4> regs;670 explicit OperandFragment(StringRef name, StringRef ptxTypeName)671 : operandName(name), ptxTypeAttr(ptxTypeName) {}672 };673 674 std::array<OperandFragment, 3> frags{675 OperandFragment("A", getMultiplicandAPtxTypeAttrName()),676 OperandFragment("B", getMultiplicandBPtxTypeAttrName()),677 OperandFragment("C", "")};678 SmallVector<StringRef, 4> ignoreAttrNames{679 mlir::NVVM::MmaOp::getOperandSegmentSizeAttr()};680 681 for (unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {682 auto &frag = frags[fragIdx];683 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);684 for (auto operandIdx = varOperandSpec.first;685 operandIdx < varOperandSpec.first + varOperandSpec.second;686 operandIdx++) {687 frag.regs.push_back(this->getOperand(operandIdx));688 if (operandIdx == 0) {689 regTypes.push_back(this->getOperand(operandIdx).getType());690 }691 }692 std::optional<MMATypes> inferredType =693 inferOperandMMAType(regTypes.back(), /*isAccumulator=*/fragIdx >= 2);694 if (inferredType)695 ignoreAttrNames.push_back(frag.ptxTypeAttr);696 }697 698 auto printMmaOperand = [&](const OperandFragment &frag) -> void {699 p << " " << frag.operandName;700 p << "[";701 p.printOperands(frag.regs);702 p << "] ";703 };704 705 for (const auto &frag : frags) {706 printMmaOperand(frag);707 }708 709 p.printOptionalAttrDict(this->getOperation()->getAttrs(), ignoreAttrNames);710 711 // Print the types of the operands and result.712 p << " : "713 << "(";714 llvm::interleaveComma(SmallVector<Type, 3>{frags[0].regs[0].getType(),715 frags[1].regs[0].getType(),716 frags[2].regs[0].getType()},717 p);718 p << ")";719 p.printArrowTypeList(TypeRange{this->getRes().getType()});720}721 722void MmaOp::build(OpBuilder &builder, OperationState &result, Type resultType,723 ValueRange operandA, ValueRange operandB, ValueRange operandC,724 ArrayRef<int64_t> shape, std::optional<MMAB1Op> b1Op,725 std::optional<MMAIntOverflow> intOverflow,726 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,727 std::optional<std::array<MMALayout, 2>> multiplicandLayouts) {728 729 assert(shape.size() == 3 && "expected shape to have size 3 (m, n, k)");730 MLIRContext *ctx = builder.getContext();731 result.addAttribute(732 "shape", builder.getAttr<MMAShapeAttr>(shape[0], shape[1], shape[2]));733 734 result.addOperands(operandA);735 result.addOperands(operandB);736 result.addOperands(operandC);737 738 if (multiplicandPtxTypes) {739 result.addAttribute("multiplicandAPtxType",740 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));741 result.addAttribute("multiplicandBPtxType",742 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));743 } else {744 if (auto res = inferOperandMMAType(operandA[0].getType(), false))745 result.addAttribute("multiplicandAPtxType", MMATypesAttr::get(ctx, *res));746 if (auto res = inferOperandMMAType(operandB[0].getType(), false))747 result.addAttribute("multiplicandBPtxType", MMATypesAttr::get(ctx, *res));748 }749 750 if (multiplicandLayouts) {751 result.addAttribute("layoutA",752 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[0]));753 result.addAttribute("layoutB",754 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[1]));755 } else {756 result.addAttribute("layoutA", MMALayoutAttr::get(ctx, MMALayout::row));757 result.addAttribute("layoutB", MMALayoutAttr::get(ctx, MMALayout::col));758 }759 760 if (intOverflow.has_value())761 result.addAttribute("intOverflowBehavior",762 MMAIntOverflowAttr::get(ctx, *intOverflow));763 if (b1Op.has_value())764 result.addAttribute("b1Op", MMAB1OpAttr::get(ctx, *b1Op));765 766 result.addTypes(resultType);767 result.addAttribute(768 MmaOp::getOperandSegmentSizeAttr(),769 builder.getDenseI32ArrayAttr({static_cast<int32_t>(operandA.size()),770 static_cast<int32_t>(operandB.size()),771 static_cast<int32_t>(operandC.size())}));772}773 774// <operation> :=775// A `[` $operandA `]` B `[` $operandB `]` C `[` $operandC `]`776// attr-dict : (type($operandA[0]), type($operandB[0]), type($operandC[0]))777// `->` type($res)778ParseResult MmaOp::parse(OpAsmParser &parser, OperationState &result) {779 struct OperandFragment {780 std::optional<MMATypes> elemtype;781 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;782 SmallVector<Type> regTypes;783 };784 785 Builder &builder = parser.getBuilder();786 std::array<OperandFragment, 4> frags;787 788 NamedAttrList namedAttributes;789 790 // A helper to parse the operand segments.791 auto parseMmaOperand = [&](StringRef operandName,792 OperandFragment &frag) -> LogicalResult {793 if (parser.parseKeyword(operandName).failed())794 return failure();795 if (parser796 .parseOperandList(frag.regs, OpAsmParser::Delimiter::OptionalSquare)797 .failed())798 return failure();799 return success();800 };801 802 // Parse the operand segments.803 if (parseMmaOperand("A", frags[0]).failed())804 return failure();805 if (parseMmaOperand("B", frags[1]).failed())806 return failure();807 if (parseMmaOperand("C", frags[2]).failed())808 return failure();809 810 if (parser.parseOptionalAttrDict(namedAttributes).failed())811 return failure();812 813 // Parse the type specification and resolve operands.814 SmallVector<Type, 3> operandTypes;815 if (failed(parser.parseColon()))816 return failure();817 if (failed(parser.parseLParen()))818 return failure();819 if (failed(parser.parseTypeList(operandTypes)))820 return failure();821 if (failed(parser.parseRParen()))822 if (operandTypes.size() != 3)823 return parser.emitError(824 parser.getNameLoc(),825 "expected one type for each operand segment but got " +826 Twine(operandTypes.size()) + " types");827 for (const auto &iter : llvm::enumerate(operandTypes)) {828 auto &frag = frags[iter.index()];829 frag.regTypes.resize(frag.regs.size(), iter.value());830 if (failed(parser.resolveOperands(frag.regs, frag.regTypes,831 parser.getNameLoc(), result.operands)))832 return failure();833 frag.elemtype = inferOperandMMAType(frag.regTypes[0],834 /*isAccumulator*/ iter.index() < 2);835 }836 837 Type resultType;838 if (parser.parseArrow() || parser.parseType(resultType))839 return failure();840 frags[3].elemtype = inferOperandMMAType(resultType, /*isAccumulator*/ true);841 842 std::array<StringRef, 2> names{"multiplicandAPtxType",843 "multiplicandBPtxType"};844 for (unsigned idx = 0; idx < names.size(); idx++) {845 const auto &frag = frags[idx];846 std::optional<NamedAttribute> attr = namedAttributes.getNamed(names[idx]);847 if (!frag.elemtype.has_value() && !attr.has_value()) {848 return parser.emitError(849 parser.getNameLoc(),850 "attribute " + names[idx] +851 " is not provided explicitly and cannot be inferred");852 }853 if (!attr.has_value())854 result.addAttribute(855 names[idx], MMATypesAttr::get(parser.getContext(), *frag.elemtype));856 }857 858 result.addTypes(resultType);859 if (!namedAttributes.empty())860 result.addAttributes(namedAttributes);861 result.addAttribute(MmaOp::getOperandSegmentSizeAttr(),862 builder.getDenseI32ArrayAttr({863 static_cast<int32_t>(frags[0].regs.size()),864 static_cast<int32_t>(frags[1].regs.size()),865 static_cast<int32_t>(frags[2].regs.size()),866 }));867 return success();868}869 870LogicalResult MmaOp::verify() {871 MLIRContext *context = getContext();872 auto f16Ty = Float16Type::get(context);873 auto i32Ty = IntegerType::get(context, 32);874 auto f16x2Ty = VectorType::get(2, f16Ty);875 auto f32Ty = Float32Type::get(context);876 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(877 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});878 879 auto s32x4StructTy =880 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});881 auto f32x8StructTy =882 LLVM::LLVMStructType::getLiteral(context, SmallVector<Type>(8, f32Ty));883 auto f16x2x2StructTy =884 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});885 auto f32x4StructTy =886 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});887 auto s32x2StructTy =888 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});889 890 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),891 getShapeAttr().getK()};892 893 // These variables define the set of allowed data types for matrices A, B, C,894 // and result.895 using AllowedShapes = SmallVector<std::array<int64_t, 3>, 2>;896 using AllowedTypes = SmallVector<SmallVector<Type, 4>, 2>;897 AllowedShapes allowedShapes;898 AllowedTypes expectedA;899 AllowedTypes expectedB;900 AllowedTypes expectedC;901 SmallVector<Type> expectedResult;902 903 // When M = 16, we just need to calculate the number of 8xk tiles, where904 // k is a factor that depends on the data type.905 if (mmaShape[0] == 16) {906 int64_t kFactor;907 Type multiplicandFragType;908 switch (*getMultiplicandAPtxType()) {909 case MMATypes::tf32:910 kFactor = 4;911 multiplicandFragType = i32Ty;912 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(913 context, {f32Ty, f32Ty, f32Ty, f32Ty}));914 break;915 case MMATypes::bf16:916 kFactor = 8;917 multiplicandFragType = i32Ty;918 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(919 context, {f32Ty, f32Ty, f32Ty, f32Ty}));920 break;921 case MMATypes::f16:922 kFactor = 8;923 multiplicandFragType = f16x2Ty;924 expectedResult.push_back(f16x2x2StructTy);925 expectedResult.push_back(f32x4StructTy);926 break;927 case MMATypes::s4:928 case MMATypes::u4:929 kFactor = 32;930 break;931 case MMATypes::b1:932 kFactor = 128;933 break;934 case MMATypes::s8:935 case MMATypes::u8:936 kFactor = 16;937 break;938 default:939 return emitError("invalid shape or multiplicand type: " +940 stringifyEnum(getMultiplicandAPtxType().value()));941 }942 943 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {944 expectedResult.push_back(s32x4StructTy);945 expectedC.emplace_back(4, i32Ty);946 multiplicandFragType = i32Ty;947 } else {948 expectedC.emplace_back(2, f16x2Ty);949 expectedC.emplace_back(4, f32Ty);950 }951 952 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor);953 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);954 expectedA.emplace_back(unitA, multiplicandFragType);955 expectedB.emplace_back(unitB, multiplicandFragType);956 allowedShapes.push_back({16, 8, kFactor});957 allowedShapes.push_back({16, 8, kFactor * 2});958 959 if (resultPtxType() != accumPtxType())960 return emitOpError("ctype does not match dtype");961 }962 963 // In the M=8 case, there is only 1 possible case per data type.964 if (mmaShape[0] == 8) {965 if (*getMultiplicandAPtxType() == MMATypes::f16) {966 expectedA.emplace_back(2, f16x2Ty);967 expectedB.emplace_back(2, f16x2Ty);968 expectedResult.push_back(f16x2x4StructTy);969 expectedResult.push_back(f32x8StructTy);970 expectedC.emplace_back(4, f16x2Ty);971 expectedC.emplace_back(8, f32Ty);972 allowedShapes.push_back({8, 8, 4});973 }974 if (*getMultiplicandAPtxType() == MMATypes::f64) {975 Type f64Ty = Float64Type::get(context);976 expectedA.emplace_back(1, f64Ty);977 expectedB.emplace_back(1, f64Ty);978 expectedC.emplace_back(2, f64Ty);979 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(980 context, SmallVector<Type>(2, f64Ty)));981 allowedShapes.push_back({8, 8, 4});982 }983 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {984 expectedA.push_back({i32Ty});985 expectedB.push_back({i32Ty});986 expectedC.push_back({i32Ty, i32Ty});987 expectedResult.push_back(s32x2StructTy);988 if (isInt4PtxType(getMultiplicandAPtxType().value()))989 allowedShapes.push_back({8, 8, 32});990 if (isInt8PtxType(getMultiplicandAPtxType().value()))991 allowedShapes.push_back({8, 8, 16});992 if (getMultiplicandAPtxType().value() == MMATypes::b1)993 allowedShapes.push_back({8, 8, 128});994 }995 }996 997 std::string errorMessage;998 llvm::raw_string_ostream errorStream(errorMessage);999 1000 // Check that we matched an existing shape/dtype combination.1001 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||1002 !llvm::is_contained(allowedShapes, mmaShape)) {1003 errorStream << "unimplemented variant for MMA shape <";1004 llvm::interleaveComma(mmaShape, errorStream);1005 errorStream << ">";1006 return emitOpError(errorMessage);1007 }1008 1009 // Verify the operand types for segments of A, B, and C operands.1010 std::array<StringRef, 3> operandNames{"A", "B", "C"};1011 for (const auto &iter : llvm::enumerate(1012 SmallVector<AllowedTypes, 3>{expectedA, expectedB, expectedC})) {1013 auto spec = this->getODSOperandIndexAndLength(iter.index());1014 SmallVector<Type, 4> operandTySeg(operand_type_begin() + spec.first,1015 operand_type_begin() + spec.first +1016 spec.second);1017 bool match = llvm::is_contained(iter.value(), operandTySeg);1018 1019 if (!match) {1020 errorStream << "Could not match types for the "1021 << operandNames[iter.index()]1022 << " operands; expected one of ";1023 for (const auto &x : iter.value()) {1024 errorStream << x.size() << "x" << x[0] << " ";1025 }1026 errorStream << "but got ";1027 llvm::interleaveComma(operandTySeg, errorStream);1028 return emitOpError(errorMessage);1029 }1030 }1031 1032 // Check the result type1033 if (!llvm::any_of(expectedResult, [&](Type expectedResultType) {1034 return expectedResultType == getResult().getType();1035 })) {1036 errorStream1037 << "Could not match allowed types for the result; expected one of ";1038 llvm::interleaveComma(expectedResult, errorStream);1039 errorStream << " but got " << getResult().getType();1040 return emitOpError(errorMessage);1041 }1042 1043 // Ensure that binary MMA variants have a b1 MMA operation defined.1044 if (getMultiplicandAPtxType() == MMATypes::b1 && !getB1Op()) {1045 return emitOpError("op requires " + getB1OpAttrName().strref() +1046 " attribute");1047 }1048 1049 // Ensure int4/int8 MMA variants specify the accum overflow behavior1050 // attribute.1051 if (isInt4PtxType(*getMultiplicandAPtxType()) ||1052 isInt8PtxType(*getMultiplicandAPtxType())) {1053 if (!getIntOverflowBehavior())1054 return emitOpError("op requires " +1055 getIntOverflowBehaviorAttrName().strref() +1056 " attribute");1057 }1058 1059 // Validate layout combinations. According to the operation description, most1060 // MMA operations require layoutA=row and layoutB=col. Only m8n8k4 with f161061 // can use other layout combinations.1062 bool isM8N8K4_F16 =1063 (mmaShape[0] == 8 && mmaShape[1] == 8 && mmaShape[2] == 4 &&1064 getMultiplicandAPtxType() == MMATypes::f16);1065 1066 if (!isM8N8K4_F16) {1067 // For all other shapes/types, layoutA must be row and layoutB must be col1068 if (getLayoutA() != MMALayout::row || getLayoutB() != MMALayout::col) {1069 return emitOpError("requires layoutA = #nvvm.mma_layout<row> and "1070 "layoutB = #nvvm.mma_layout<col> for shape <")1071 << mmaShape[0] << ", " << mmaShape[1] << ", " << mmaShape[2]1072 << "> with element types "1073 << stringifyEnum(*getMultiplicandAPtxType()) << " and "1074 << stringifyEnum(*getMultiplicandBPtxType())1075 << ". Only m8n8k4 with f16 supports other layouts.";1076 }1077 }1078 1079 return success();1080}1081 1082MMATypes MmaSpOp::accumPtxType() {1083 std::optional<mlir::NVVM::MMATypes> val = MmaOp::inferOperandMMAType(1084 getODSOperands(2).getTypes().front(), /*isAccumulator=*/true);1085 assert(val.has_value() && "accumulator PTX type should always be inferrable");1086 return val.value();1087}1088 1089MMATypes MmaSpOp::resultPtxType() {1090 std::optional<mlir::NVVM::MMATypes> val =1091 MmaOp::inferOperandMMAType(getResult().getType(), /*isAccumulator=*/true);1092 assert(val.has_value() && "result PTX type should always be inferrable");1093 return val.value();1094}1095 1096mlir::NVVM::IDArgPair1097MmaSpOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,1098 llvm::IRBuilderBase &builder) {1099 auto thisOp = cast<NVVM::MmaSpOp>(op);1100 1101 // Get operands1102 llvm::SmallVector<llvm::Value *> args;1103 for (mlir::Value v : thisOp.getOperands())1104 args.push_back(mt.lookupValue(v));1105 1106 // Get intrinsic ID using the existing getIntrinsicID method1107 auto intId = MmaSpOp::getIntrinsicID(1108 thisOp.getShape().getM(), thisOp.getShape().getN(),1109 thisOp.getShape().getK(), thisOp.getIntOverflowBehavior(),1110 thisOp.getOrderedMetadata(), thisOp.getKind(),1111 *thisOp.getMultiplicandAPtxType(), *thisOp.getMultiplicandBPtxType(),1112 thisOp.accumPtxType(), thisOp.resultPtxType());1113 1114 return {intId, args};1115}1116 1117void MmaSpOp::print(OpAsmPrinter &p) {1118 SmallVector<Type, 4> regTypes;1119 struct OperandFragment {1120 StringRef operandName;1121 StringRef ptxTypeAttr;1122 SmallVector<Value, 4> regs;1123 explicit OperandFragment(StringRef name, StringRef ptxTypeName)1124 : operandName(name), ptxTypeAttr(ptxTypeName) {}1125 };1126 1127 std::array<OperandFragment, 5> frags{1128 OperandFragment("A", getMultiplicandAPtxTypeAttrName()),1129 OperandFragment("B", getMultiplicandBPtxTypeAttrName()),1130 OperandFragment("C", ""), OperandFragment("sparseMetadata", ""),1131 OperandFragment("selector", "")};1132 SmallVector<StringRef, 4> ignoreAttrNames{1133 mlir::NVVM::MmaSpOp::getOperandSegmentSizeAttr()};1134 1135 // Handle variadic operands A, B, C1136 for (unsigned fragIdx = 0; fragIdx < 3; fragIdx++) {1137 auto &frag = frags[fragIdx];1138 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);1139 for (auto operandIdx = varOperandSpec.first;1140 operandIdx < varOperandSpec.first + varOperandSpec.second;1141 operandIdx++) {1142 frag.regs.push_back(this->getOperand(operandIdx));1143 if (operandIdx == varOperandSpec.first) {1144 regTypes.push_back(this->getOperand(operandIdx).getType());1145 }1146 }1147 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(1148 regTypes.back(), /*isAccumulator=*/fragIdx >= 2);1149 if (inferredType)1150 ignoreAttrNames.push_back(frag.ptxTypeAttr);1151 }1152 1153 // Handle sparse metadata and selector (single operands)1154 frags[3].regs.push_back(getSparseMetadata());1155 frags[4].regs.push_back(getSparsitySelector());1156 1157 auto printMmaSpOperand = [&](const OperandFragment &frag) -> void {1158 p << " " << frag.operandName;1159 p << "[";1160 p.printOperands(frag.regs);1161 p << "]";1162 };1163 1164 for (const auto &frag : frags)1165 printMmaSpOperand(frag);1166 1167 p.printOptionalAttrDict((*this)->getAttrs(), ignoreAttrNames);1168 p << " : ";1169 p << "(";1170 for (int i = 0; i < 3; ++i) {1171 p << regTypes[i];1172 if (i < 2)1173 p << ", ";1174 }1175 p << ") -> " << getResult().getType();1176}1177 1178void MmaSpOp::build(1179 OpBuilder &builder, OperationState &result, Type resultType,1180 ValueRange operandA, ValueRange operandB, ValueRange operandC,1181 Value sparseMetadata, Value sparsitySelector, ArrayRef<int64_t> shape,1182 std::optional<MMAIntOverflow> intOverflow,1183 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {1184 1185 assert(shape.size() == 3 && "expected shape to have size 3 (m, n, k)");1186 MLIRContext *ctx = builder.getContext();1187 result.addAttribute(1188 "shape", builder.getAttr<MMAShapeAttr>(shape[0], shape[1], shape[2]));1189 1190 result.addOperands(operandA);1191 result.addOperands(operandB);1192 result.addOperands(operandC);1193 result.addOperands(sparseMetadata);1194 result.addOperands(sparsitySelector);1195 1196 if (multiplicandPtxTypes) {1197 result.addAttribute("multiplicandAPtxType",1198 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));1199 result.addAttribute("multiplicandBPtxType",1200 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));1201 } else {1202 if (auto res = MmaOp::inferOperandMMAType(operandA[0].getType(), false))1203 result.addAttribute("multiplicandAPtxType", MMATypesAttr::get(ctx, *res));1204 if (auto res = MmaOp::inferOperandMMAType(operandB[0].getType(), false))1205 result.addAttribute("multiplicandBPtxType", MMATypesAttr::get(ctx, *res));1206 }1207 1208 if (intOverflow.has_value())1209 result.addAttribute("intOverflowBehavior",1210 MMAIntOverflowAttr::get(ctx, *intOverflow));1211 1212 result.addTypes(resultType);1213 result.addAttribute(1214 MmaSpOp::getOperandSegmentSizeAttr(),1215 builder.getDenseI32ArrayAttr({static_cast<int32_t>(operandA.size()),1216 static_cast<int32_t>(operandB.size()),1217 static_cast<int32_t>(operandC.size()), 1,1218 1})); // sparseMetadata and sparsitySelector1219}1220 1221ParseResult MmaSpOp::parse(OpAsmParser &parser, OperationState &result) {1222 struct OperandFragment {1223 std::optional<MMATypes> elemtype;1224 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;1225 SmallVector<Type> regTypes;1226 };1227 1228 Builder &builder = parser.getBuilder();1229 std::array<OperandFragment, 6> frags; // A, B, C, sparseMetadata, selector1230 1231 NamedAttrList namedAttributes;1232 1233 // A helper to parse the operand segments.1234 auto parseMmaSpOperand = [&](StringRef operandName,1235 OperandFragment &frag) -> LogicalResult {1236 if (parser.parseKeyword(operandName).failed())1237 return failure();1238 if (parser1239 .parseOperandList(frag.regs, OpAsmParser::Delimiter::OptionalSquare)1240 .failed())1241 return failure();1242 return success();1243 };1244 1245 // Parse the operand segments.1246 if (parseMmaSpOperand("A", frags[0]).failed())1247 return failure();1248 if (parseMmaSpOperand("B", frags[1]).failed())1249 return failure();1250 if (parseMmaSpOperand("C", frags[2]).failed())1251 return failure();1252 if (parseMmaSpOperand("sparseMetadata", frags[3]).failed())1253 return failure();1254 if (parseMmaSpOperand("selector", frags[4]).failed())1255 return failure();1256 1257 if (parser.parseOptionalAttrDict(namedAttributes).failed())1258 return failure();1259 1260 // Parse the type specification and resolve operands.1261 SmallVector<Type, 3> operandTypes;1262 if (failed(parser.parseColon()))1263 return failure();1264 if (failed(parser.parseLParen()))1265 return failure();1266 if (failed(parser.parseTypeList(operandTypes)))1267 return failure();1268 if (failed(parser.parseRParen()))1269 return failure();1270 if (operandTypes.size() != 3)1271 return parser.emitError(1272 parser.getNameLoc(),1273 "expected one type for each operand segment but got " +1274 Twine(operandTypes.size()) + " types");1275 for (const auto &iter : llvm::enumerate(operandTypes)) {1276 auto &frag = frags[iter.index()];1277 frag.regTypes.resize(frag.regs.size(), iter.value());1278 if (failed(parser.resolveOperands(frag.regs, frag.regTypes,1279 parser.getNameLoc(), result.operands)))1280 return failure();1281 frag.elemtype =1282 MmaOp::inferOperandMMAType(frag.regTypes[0],1283 /*isAccumulator*/ iter.index() >= 2);1284 }1285 1286 Type resultType;1287 if (parser.parseArrow() || parser.parseType(resultType))1288 return failure();1289 frags[5].elemtype =1290 MmaOp::inferOperandMMAType(resultType, /*isAccumulator*/ true);1291 1292 // Resolve sparse metadata and selector (assume i32 type)1293 Type i32Type = builder.getIntegerType(32);1294 if (parser1295 .resolveOperands(frags[3].regs, i32Type, parser.getCurrentLocation(),1296 result.operands)1297 .failed())1298 return failure();1299 if (parser1300 .resolveOperands(frags[4].regs, i32Type, parser.getCurrentLocation(),1301 result.operands)1302 .failed())1303 return failure();1304 1305 std::array<StringRef, 2> names{"multiplicandAPtxType",1306 "multiplicandBPtxType"};1307 for (unsigned idx = 0; idx < names.size(); idx++) {1308 const auto &frag = frags[idx];1309 std::optional<NamedAttribute> attr = namedAttributes.getNamed(names[idx]);1310 if (!frag.elemtype.has_value() && !attr.has_value()) {1311 return parser.emitError(1312 parser.getNameLoc(),1313 "attribute " + names[idx] +1314 " is not provided explicitly and cannot be inferred");1315 }1316 if (!attr.has_value())1317 result.addAttribute(1318 names[idx], MMATypesAttr::get(parser.getContext(), *frag.elemtype));1319 }1320 1321 result.addTypes(resultType);1322 if (!namedAttributes.empty())1323 result.addAttributes(namedAttributes);1324 result.addAttribute(MmaSpOp::getOperandSegmentSizeAttr(),1325 builder.getDenseI32ArrayAttr({1326 static_cast<int32_t>(frags[0].regs.size()),1327 static_cast<int32_t>(frags[1].regs.size()),1328 static_cast<int32_t>(frags[2].regs.size()),1329 1, // sparseMetadata1330 1 // sparsitySelector1331 }));1332 return success();1333}1334 1335LogicalResult MmaSpOp::verify() {1336 MLIRContext *context = getContext();1337 auto f16Ty = Float16Type::get(context);1338 auto i32Ty = IntegerType::get(context, 32);1339 auto f16x2Ty = VectorType::get(2, f16Ty);1340 auto f32Ty = Float32Type::get(context);1341 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(1342 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});1343 1344 auto s32x4StructTy =1345 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});1346 auto f32x8StructTy =1347 LLVM::LLVMStructType::getLiteral(context, SmallVector<Type>(8, f32Ty));1348 auto f16x2x2StructTy =1349 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});1350 auto f32x4StructTy =1351 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});1352 auto s32x2StructTy =1353 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});1354 1355 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),1356 getShapeAttr().getK()};1357 1358 // These variables define the set of allowed data types for matrices A, B, C,1359 // and result.1360 using AllowedShapes = SmallVector<std::array<int64_t, 3>, 2>;1361 using AllowedTypes = SmallVector<SmallVector<Type, 4>, 2>;1362 AllowedShapes allowedShapes;1363 AllowedTypes expectedA;1364 AllowedTypes expectedB;1365 AllowedTypes expectedC;1366 SmallVector<Type> expectedResult;1367 1368 // When M = 16, we just need to calculate the number of 8xk tiles, where1369 // k is a factor that depends on the data type.1370 if (mmaShape[0] == 16) {1371 int64_t kFactor;1372 Type multiplicandFragType;1373 switch (*getMultiplicandAPtxType()) {1374 case MMATypes::tf32:1375 kFactor = 4;1376 multiplicandFragType = i32Ty;1377 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(1378 context, {f32Ty, f32Ty, f32Ty, f32Ty}));1379 // Sparse MMA supports m16n8k8 and m16n8k16 for tf321380 allowedShapes.push_back({16, 8, 8});1381 allowedShapes.push_back({16, 8, 16});1382 break;1383 case MMATypes::bf16:1384 kFactor = 8;1385 multiplicandFragType = i32Ty;1386 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(1387 context, {f32Ty, f32Ty, f32Ty, f32Ty}));1388 // Sparse MMA supports m16n8k16 and m16n8k32 for bf161389 allowedShapes.push_back({16, 8, 16});1390 allowedShapes.push_back({16, 8, 32});1391 break;1392 case MMATypes::f16:1393 kFactor = 8;1394 multiplicandFragType = f16x2Ty;1395 expectedResult.push_back(f16x2x2StructTy);1396 expectedResult.push_back(f32x4StructTy);1397 // Sparse MMA supports m16n8k16 and m16n8k32 for f161398 allowedShapes.push_back({16, 8, 16});1399 allowedShapes.push_back({16, 8, 32});1400 break;1401 case MMATypes::s4:1402 case MMATypes::u4:1403 kFactor = 32;1404 // Sparse MMA supports m16n8k64 and m16n8k128 for s4/u41405 allowedShapes.push_back({16, 8, 64});1406 allowedShapes.push_back({16, 8, 128});1407 break;1408 case MMATypes::s8:1409 case MMATypes::u8:1410 kFactor = 16;1411 // Sparse MMA supports m16n8k32 and m16n8k64 for s8/u81412 allowedShapes.push_back({16, 8, 32});1413 allowedShapes.push_back({16, 8, 64});1414 break;1415 case MMATypes::e4m3:1416 case MMATypes::e5m2:1417 case MMATypes::e3m2:1418 case MMATypes::e2m3:1419 case MMATypes::e2m1:1420 kFactor = 32;1421 multiplicandFragType = i32Ty;1422 expectedResult.push_back(f16x2x2StructTy);1423 expectedResult.push_back(f32x4StructTy);1424 // Sparse MMA supports m16n8k64 for FP8 types1425 allowedShapes.push_back({16, 8, 64});1426 break;1427 default:1428 return emitError("invalid shape or multiplicand type: " +1429 stringifyEnum(getMultiplicandAPtxType().value()));1430 }1431 1432 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {1433 expectedResult.push_back(s32x4StructTy);1434 expectedC.emplace_back(4, i32Ty);1435 multiplicandFragType = i32Ty;1436 } else if (*getMultiplicandAPtxType() >= MMATypes::e4m3 &&1437 *getMultiplicandAPtxType() <= MMATypes::e2m1) {1438 // FP8 types1439 expectedC.emplace_back(2, f16x2Ty);1440 expectedC.emplace_back(4, f32Ty);1441 } else {1442 expectedC.emplace_back(2, f16x2Ty);1443 expectedC.emplace_back(4, f32Ty);1444 }1445 1446 // For sparse MMA, A operand is compressed (2:4 sparsity means half the1447 // elements)1448 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor) / 2;1449 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);1450 expectedA.emplace_back(unitA, multiplicandFragType);1451 expectedB.emplace_back(unitB, multiplicandFragType);1452 1453 if (resultPtxType() != accumPtxType())1454 return emitOpError("ctype does not match dtype");1455 }1456 1457 // In the M=8 case, there is only 1 possible case per data type.1458 if (mmaShape[0] == 8) {1459 if (*getMultiplicandAPtxType() == MMATypes::f16) {1460 expectedA.emplace_back(2, f16x2Ty);1461 expectedB.emplace_back(2, f16x2Ty);1462 expectedResult.push_back(f16x2x4StructTy);1463 expectedResult.push_back(f32x8StructTy);1464 expectedC.emplace_back(4, f16x2Ty);1465 expectedC.emplace_back(8, f32Ty);1466 allowedShapes.push_back({8, 8, 4});1467 }1468 if (*getMultiplicandAPtxType() == MMATypes::f64) {1469 Type f64Ty = Float64Type::get(context);1470 expectedA.emplace_back(1, f64Ty);1471 expectedB.emplace_back(1, f64Ty);1472 expectedC.emplace_back(2, f64Ty);1473 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(1474 context, SmallVector<Type>(2, f64Ty)));1475 allowedShapes.push_back({8, 8, 4});1476 }1477 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {1478 expectedA.push_back({i32Ty});1479 expectedB.push_back({i32Ty});1480 expectedC.push_back({i32Ty, i32Ty});1481 expectedResult.push_back(s32x2StructTy);1482 if (isInt4PtxType(getMultiplicandAPtxType().value()))1483 allowedShapes.push_back({8, 8, 32});1484 if (isInt8PtxType(getMultiplicandAPtxType().value()))1485 allowedShapes.push_back({8, 8, 16});1486 }1487 }1488 1489 std::string errorMessage;1490 llvm::raw_string_ostream errorStream(errorMessage);1491 1492 // Check that we matched an existing shape/dtype combination.1493 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||1494 !llvm::is_contained(allowedShapes, mmaShape)) {1495 errorStream << "unimplemented variant for MMA shape <";1496 llvm::interleaveComma(mmaShape, errorStream);1497 errorStream << ">";1498 return emitOpError(errorMessage);1499 }1500 1501 // Verify the operand types for segments of A, B, and C operands.1502 std::array<StringRef, 3> operandNames{"A", "B", "C"};1503 for (const auto &iter : llvm::enumerate(1504 SmallVector<AllowedTypes, 3>{expectedA, expectedB, expectedC})) {1505 auto spec = this->getODSOperandIndexAndLength(iter.index());1506 SmallVector<Type, 4> operandTySeg(operand_type_begin() + spec.first,1507 operand_type_begin() + spec.first +1508 spec.second);1509 bool match = llvm::is_contained(iter.value(), operandTySeg);1510 1511 if (!match) {1512 errorStream << "Could not match types for the "1513 << operandNames[iter.index()]1514 << " operands; expected one of ";1515 for (const auto &x : iter.value()) {1516 errorStream << x.size() << "x" << x[0] << " ";1517 }1518 errorStream << "but got ";1519 llvm::interleaveComma(operandTySeg, errorStream);1520 return emitOpError(errorMessage);1521 }1522 }1523 1524 // Check the result type1525 if (!llvm::any_of(expectedResult, [&](Type expectedResultType) {1526 return expectedResultType == getResult().getType();1527 })) {1528 errorStream1529 << "Could not match allowed types for the result; expected one of ";1530 llvm::interleaveComma(expectedResult, errorStream);1531 errorStream << " but got " << getResult().getType();1532 return emitOpError(errorMessage);1533 }1534 1535 // Ensure int4/int8 MMA variants specify the accum overflow behavior1536 // attribute.1537 if (isInt4PtxType(*getMultiplicandAPtxType()) ||1538 isInt8PtxType(*getMultiplicandAPtxType())) {1539 if (!getIntOverflowBehavior())1540 return emitOpError("op requires " +1541 getIntOverflowBehaviorAttrName().strref() +1542 " attribute");1543 }1544 1545 // Validate sparse metadata type (should be i32)1546 if (!getSparseMetadata().getType().isInteger(32)) {1547 return emitOpError() << "sparse metadata must be i32 type";1548 }1549 1550 // Validate sparsity selector type (should be i32)1551 if (!getSparsitySelector().getType().isInteger(32)) {1552 return emitOpError() << "sparsity selector must be i32 type";1553 }1554 1555 return success();1556}1557 1558LogicalResult ShflOp::verify() {1559 auto returnStructType = llvm::dyn_cast<LLVM::LLVMStructType>(getType());1560 1561 auto verifyTypeError = [&](Twine desc, Type expectedType,1562 Type actualType) -> LogicalResult {1563 return emitOpError("expected " + desc + " to be of type ")1564 << expectedType << " but got " << actualType << " instead";1565 };1566 1567 if (returnStructType) {1568 if (!getReturnValueAndIsValid())1569 return emitOpError("\"return_value_and_is_valid\" attribute must be "1570 "specified when the return type is a struct type");1571 1572 if (returnStructType.getBody().size() != 2)1573 return emitOpError("expected return type to be a two-element struct");1574 1575 llvm::ArrayRef<Type> returnStruct = returnStructType.getBody();1576 auto resultType = returnStruct[0];1577 if (resultType != getVal().getType())1578 return verifyTypeError("first element in the returned struct",1579 getVal().getType(), resultType);1580 1581 auto predicateType = returnStruct[1];1582 if (!predicateType.isInteger(1))1583 return verifyTypeError("second element in the returned struct",1584 mlir::IntegerType::get(getContext(), 1),1585 predicateType);1586 } else {1587 if (getReturnValueAndIsValid())1588 return emitOpError("expected return type to be a two-element struct");1589 1590 if (getType() != getVal().getType())1591 return verifyTypeError("return type", getVal().getType(), getType());1592 }1593 return success();1594}1595 1596std::pair<mlir::Type, unsigned> NVVM::inferMMAType(NVVM::MMATypes type,1597 NVVM::MMAFrag frag, int nRow,1598 int nCol,1599 MLIRContext *context) {1600 unsigned numberElements = 0;1601 Type elementType;1602 OpBuilder builder(context);1603 Type f16x2 = VectorType::get(2, builder.getF16Type());1604 if (type == NVVM::MMATypes::f16) {1605 elementType = f16x2;1606 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)1607 numberElements = 8;1608 else1609 numberElements = 4;1610 } else if (type == NVVM::MMATypes::f32) {1611 elementType = builder.getF32Type();1612 numberElements = 8;1613 } else if (type == NVVM::MMATypes::f64) {1614 elementType = builder.getF64Type();1615 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)1616 numberElements = 1;1617 else1618 numberElements = 2;1619 } else if (type == NVVM::MMATypes::tf32) {1620 elementType = builder.getI32Type();1621 numberElements = 4;1622 } else if (type == NVVM::MMATypes::s8 || type == NVVM::MMATypes::u8) {1623 elementType = builder.getI32Type();1624 int parallelSize = 0;1625 if (frag == NVVM::MMAFrag::a)1626 parallelSize = nRow;1627 if (frag == NVVM::MMAFrag::b)1628 parallelSize = nCol;1629 1630 // m == 16 && n == 16 && k == 161631 if (parallelSize == 16)1632 numberElements = 2;1633 // m == 8 && n == 32 && k == 16 or m == 32 && n == 8 && k == 161634 else if (parallelSize == 8)1635 numberElements = 1;1636 else if (parallelSize == 32)1637 numberElements = 4;1638 } else if (type == NVVM::MMATypes::s32) {1639 elementType = builder.getI32Type();1640 numberElements = 8;1641 }1642 assert(numberElements != 0 && elementType != nullptr);1643 return std::make_pair(elementType, numberElements);1644}1645 1646static std::pair<mlir::Type, unsigned>1647inferMMATypeFromMNK(NVVM::MMATypes type, NVVM::MMAFrag frag, int m, int n,1648 int k, MLIRContext *context) {1649 int nRow, nCol;1650 if (frag == NVVM::MMAFrag::a) {1651 nRow = m;1652 nCol = k;1653 } else if (frag == NVVM::MMAFrag::b) {1654 nRow = k;1655 nCol = n;1656 } else {1657 nRow = m;1658 nCol = n;1659 }1660 assert(nRow && nCol);1661 return inferMMAType(type, frag, nRow, nCol, context);1662}1663 1664LogicalResult NVVM::WMMALoadOp::verify() {1665 unsigned addressSpace =1666 llvm::cast<LLVM::LLVMPointerType>(getPtr().getType()).getAddressSpace();1667 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&1668 addressSpace != NVVMMemorySpace::Shared)1669 return emitOpError("expected source pointer in memory "1670 "space 0, 1, 3");1671 1672 if (NVVM::WMMALoadOp::getIntrinsicID(getM(), getN(), getK(), getLayout(),1673 getEltype(), getFrag()) == 0)1674 return emitOpError() << "invalid attribute combination";1675 std::pair<Type, unsigned> typeInfo = inferMMATypeFromMNK(1676 getEltype(), getFrag(), getM(), getN(), getK(), getContext());1677 // Special case for f64 fragments1678 Type f64Ty = Float64Type::get(getContext());1679 if (typeInfo.first == f64Ty && typeInfo.second == 1) {1680 if (getType() != f64Ty)1681 return emitOpError("expected destination type to be f64");1682 return success();1683 }1684 // Everything else is a struct1685 Type dstType = LLVM::LLVMStructType::getLiteral(1686 getContext(), SmallVector<Type, 8>(typeInfo.second, typeInfo.first));1687 if (getType() != dstType)1688 return emitOpError("expected destination type is a structure of ")1689 << typeInfo.second << " elements of type " << typeInfo.first;1690 return success();1691}1692 1693LogicalResult NVVM::WMMAStoreOp::verify() {1694 unsigned addressSpace =1695 llvm::cast<LLVM::LLVMPointerType>(getPtr().getType()).getAddressSpace();1696 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&1697 addressSpace != NVVMMemorySpace::Shared)1698 return emitOpError("expected operands to be a source pointer in memory "1699 "space 0, 1, 3");1700 1701 if (NVVM::WMMAStoreOp::getIntrinsicID(getM(), getN(), getK(), getLayout(),1702 getEltype()) == 0)1703 return emitOpError() << "invalid attribute combination";1704 std::pair<Type, unsigned> typeInfo = inferMMATypeFromMNK(1705 getEltype(), NVVM::MMAFrag::c, getM(), getN(), getK(), getContext());1706 if (getArgs().size() != typeInfo.second)1707 return emitOpError() << "expected " << typeInfo.second << " data operands";1708 if (llvm::any_of(getArgs(), [&typeInfo](Value operands) {1709 return operands.getType() != typeInfo.first;1710 }))1711 return emitOpError() << "expected data operands of type " << typeInfo.first;1712 return success();1713}1714 1715LogicalResult NVVM::WMMAMmaOp::verify() {1716 if (NVVM::WMMAMmaOp::getIntrinsicID(getM(), getN(), getK(), getLayoutA(),1717 getLayoutB(), getEltypeA(),1718 getEltypeB()) == 0)1719 return emitOpError() << "invalid attribute combination";1720 std::pair<Type, unsigned> typeInfoA = inferMMATypeFromMNK(1721 getEltypeA(), NVVM::MMAFrag::a, getM(), getN(), getK(), getContext());1722 std::pair<Type, unsigned> typeInfoB = inferMMATypeFromMNK(1723 getEltypeA(), NVVM::MMAFrag::b, getM(), getN(), getK(), getContext());1724 std::pair<Type, unsigned> typeInfoC = inferMMATypeFromMNK(1725 getEltypeB(), NVVM::MMAFrag::c, getM(), getN(), getK(), getContext());1726 SmallVector<Type, 32> arguments;1727 arguments.append(typeInfoA.second, typeInfoA.first);1728 arguments.append(typeInfoB.second, typeInfoB.first);1729 arguments.append(typeInfoC.second, typeInfoC.first);1730 unsigned numArgs = arguments.size();1731 if (getArgs().size() != numArgs)1732 return emitOpError() << "expected " << numArgs << " arguments";1733 for (unsigned i = 0; i < numArgs; i++) {1734 if (getArgs()[i].getType() != arguments[i])1735 return emitOpError() << "expected argument " << i << " to be of type "1736 << arguments[i];1737 }1738 Type dstType = LLVM::LLVMStructType::getLiteral(1739 getContext(), SmallVector<Type, 8>(typeInfoC.second, typeInfoC.first));1740 if (getType() != dstType)1741 return emitOpError("expected destination type is a structure of ")1742 << typeInfoC.second << " elements of type " << typeInfoC.first;1743 return success();1744}1745 1746LogicalResult NVVM::LdMatrixOp::verify() {1747 uint32_t num = getNum(), m = getShape().getM(), n = getShape().getN();1748 if (m == 8 && n == 8) {1749 if (num != 1 && num != 2 && num != 4) {1750 return emitOpError("expected num attribute to be 1, 2 or 4 for 8x8 "1751 "matrix");1752 }1753 if (getEltType() != LdStMatrixEltType::B16) {1754 return emitOpError("expected element type to be b16 for 8x8 matrix");1755 }1756 } else if (m == 8 && n == 16) {1757 if (num != 1 && num != 2 && num != 4) {1758 return emitOpError("expected num attribute to be 1, 2 or 4 for 8x16 "1759 "matrix");1760 }1761 if (getLayout() != MMALayout::row) {1762 return emitOpError("expected layout to be row for 8x16 matrix");1763 }1764 if (getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&1765 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {1766 return emitOpError("expected element type to be b8x16.b4x16_p64 or "1767 "b8x16.b6x16_p32 for 8x16 matrix");1768 }1769 } else if (m == 16 && n == 16) {1770 if (num != 1 && num != 2) {1771 return emitOpError("expected num attribute to be 1 or 2 for 16x16 "1772 "matrix");1773 }1774 if (getLayout() != MMALayout::col) {1775 return emitOpError("expected layout to be col for 16x16 matrix");1776 }1777 if (getEltType() != LdStMatrixEltType::B8 &&1778 getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&1779 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {1780 return emitOpError("expected element type to be b8, b8x16.b4x16_p64 or "1781 "b8x16.b6x16_p32 for 16x16 matrix");1782 }1783 } else {1784 return emitOpError("expected shape to be 8x8, 8x16 or 16x16");1785 }1786 1787 Type i32 = IntegerType::get(getContext(), 32);1788 uint32_t numElements = (m == 16 && n == 16 ? num * 2 : num);1789 if (numElements == 1 && getType() != i32)1790 return emitOpError("expected destination type is i32");1791 if (numElements == 2 || numElements == 4) {1792 Type dstType = LLVM::LLVMStructType::getLiteral(1793 getContext(), SmallVector<Type>(numElements, i32));1794 if (getType() != dstType)1795 return emitOpError("expected destination type is a structure of ")1796 << numElements << " elements of type i32";1797 }1798 1799 return success();1800}1801 1802LogicalResult NVVM::StMatrixOp::verify() {1803 int numMatrix = getSources().size();1804 if (numMatrix != 1 && numMatrix != 2 && numMatrix != 4)1805 return emitOpError("expected num attribute to be 1, 2 or 4");1806 1807 int m = getShape().getM(), n = getShape().getN();1808 if (m == 8 && n == 8) {1809 if (getEltType() != NVVM::LdStMatrixEltType::B16) {1810 return emitOpError("expected element type to be B16 for 8x8 matrix");1811 }1812 } else if (m == 16 && n == 8) {1813 if (getEltType() != NVVM::LdStMatrixEltType::B8) {1814 return emitOpError("expected element type to be B8 for 16x8 matrix");1815 }1816 if (getLayout() != NVVM::MMALayout::col) {1817 return emitOpError("expected layout to be col for 16x8 matrix");1818 }1819 } else {1820 return emitOpError("expected shape to be 8x8 or 16x8");1821 }1822 1823 return success();1824}1825 1826static FailureOr<int> getAllowedSizeK(NVVM::WGMMATypes typeA) {1827 if (typeA == NVVM::WGMMATypes::tf32)1828 return 8;1829 if (typeA == NVVM::WGMMATypes::f16 || typeA == NVVM::WGMMATypes::bf16)1830 return 16;1831 if (typeA == NVVM::WGMMATypes::s8 || typeA == NVVM::WGMMATypes::u8)1832 return 32;1833 if (typeA == NVVM::WGMMATypes::e4m3 || typeA == NVVM::WGMMATypes::e5m2)1834 return 32;1835 if (typeA == NVVM::WGMMATypes::b1)1836 return 256;1837 return failure();1838}1839 1840static LogicalResult isAllowedWGMMADataType(NVVM::WGMMATypes typeD,1841 NVVM::WGMMATypes typeA,1842 NVVM::WGMMATypes typeB) {1843 switch (typeA) {1844 case NVVM::WGMMATypes::f16:1845 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&1846 typeB == NVVM::WGMMATypes::f16)1847 return success();1848 break;1849 case NVVM::WGMMATypes::tf32:1850 if (typeD == NVVM::WGMMATypes::f32 && typeB == NVVM::WGMMATypes::tf32)1851 return success();1852 break;1853 case NVVM::WGMMATypes::u8:1854 case NVVM::WGMMATypes::s8:1855 if (typeD == NVVM::WGMMATypes::s32 &&1856 (typeB == NVVM::WGMMATypes::u8 || typeB == NVVM::WGMMATypes::s8))1857 return success();1858 break;1859 case NVVM::WGMMATypes::b1:1860 if (typeD == NVVM::WGMMATypes::s32 && typeB == NVVM::WGMMATypes::b1)1861 return success();1862 break;1863 case NVVM::WGMMATypes::bf16:1864 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&1865 typeB == NVVM::WGMMATypes::bf16)1866 return success();1867 break;1868 case NVVM::WGMMATypes::e4m3:1869 case NVVM::WGMMATypes::e5m2:1870 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&1871 (typeB == NVVM::WGMMATypes::e5m2 || typeB == NVVM::WGMMATypes::e4m3))1872 return success();1873 break;1874 case WGMMATypes::f32:1875 case WGMMATypes::s32:1876 llvm_unreachable("unsupported input types");1877 break;1878 }1879 return failure();1880}1881 1882static LogicalResult isAllowedSizeN(int sizeN, NVVM::WGMMATypes typeA) {1883 SmallVector<int> allowedN = {8, 16, 24, 32, 40, 48, 56, 64,1884 72, 80, 88, 96, 104, 112, 120, 128,1885 136, 144, 152, 160, 168, 176, 184, 192,1886 200, 208, 216, 224, 232, 240, 248, 256};1887 SmallVector<int> allowedNshort = {8, 16, 24, 32, 48, 64,1888 80, 96, 112, 128, 144, 160,1889 176, 192, 208, 224, 240, 256};1890 switch (typeA) {1891 case WGMMATypes::f16:1892 case WGMMATypes::tf32:1893 case WGMMATypes::bf16:1894 case WGMMATypes::e4m3:1895 case WGMMATypes::e5m2:1896 if (llvm::is_contained(allowedN, sizeN))1897 return success();1898 break;1899 case WGMMATypes::u8:1900 case WGMMATypes::s8:1901 case WGMMATypes::b1:1902 if (llvm::is_contained(allowedNshort, sizeN))1903 return success();1904 break;1905 case WGMMATypes::f32:1906 case WGMMATypes::s32:1907 llvm_unreachable("unsupported input types");1908 break;1909 }1910 return failure();1911}1912 1913LogicalResult NVVM::WgmmaMmaAsyncOp::verify() {1914 Value outValue = getResults();1915 auto stype = dyn_cast<LLVM::LLVMStructType>(outValue.getType());1916 if (!stype)1917 return emitOpError() << "expected results to be struct";1918 int outputSize = stype.getBody().size();1919 WGMMATypes typeD = getTypeD();1920 WGMMATypes typeA = getTypeA();1921 WGMMATypes typeB = getTypeB();1922 1923 for (Type t : stype.getBody()) {1924 if (t != stype.getBody().front())1925 return emitOpError()1926 << "all elements in struct must be same type but there is " << t;1927 }1928 1929 if (typeD != WGMMATypes::f32 && typeD != WGMMATypes::f16 &&1930 typeD != WGMMATypes::s32) {1931 return emitOpError() << "does not support the given output type "1932 << NVVM::stringifyWGMMATypes(typeD);1933 }1934 if (typeD == WGMMATypes::s32 &&1935 (getScaleA() == WGMMAScaleIn::neg || getScaleB() == WGMMAScaleIn::neg)) {1936 return emitOpError() << "has s32 output, scaleA and scaleB cannot be neg";1937 }1938 1939 if (failed(isAllowedWGMMADataType(typeD, typeA, typeB))) {1940 return emitOpError() << NVVM::stringifyWGMMATypes(typeD)1941 << " += " << NVVM::stringifyWGMMATypes(typeA) << " * "1942 << NVVM::stringifyWGMMATypes(typeB)1943 << ", it is not supported.";1944 }1945 1946 // Check M1947 if (getShape().getM() != 64)1948 return emitOpError() << "shape 'm' must be 64";1949 1950 // Check K1951 FailureOr<int> allowedK = getAllowedSizeK(typeA);1952 if (failed(allowedK) || allowedK.value() != getShape().getK())1953 return emitOpError() << "shape 'k' must be " << allowedK.value()1954 << " for input type "1955 << NVVM::stringifyWGMMATypes(typeA);1956 1957 // Check N1958 if (failed(isAllowedSizeN(getShape().getN(), typeA))) {1959 return emitOpError() << "has input type "1960 << NVVM::stringifyWGMMATypes(typeA) << " n is set to "1961 << getShape().getN() << ", it is not supported.";1962 }1963 1964 // Check transpose (only available for f16/bf16)1965 // Matrices A should be stored in row-major and B in column-major.1966 // Only f16/bf16 matrices can be stored in either column-major or row-major1967 // by setting the transpose value(imm-trans-a,imm-trans-b) in PTX code.1968 if ((typeA != WGMMATypes::f16 && typeA != WGMMATypes::bf16) &&1969 (getLayoutA() == mlir::NVVM::MMALayout::col ||1970 getLayoutB() == mlir::NVVM::MMALayout::row)) {1971 return emitOpError()1972 << "given layouts layout_a = " << stringifyMMALayout(getLayoutA())1973 << " and layout_b = " << stringifyMMALayout(getLayoutB())1974 << " for input types " << stringifyWGMMATypes(typeA) << " and "1975 << stringifyWGMMATypes(typeB)1976 << " requires transpose. However, this is only supported for: "1977 << stringifyMMATypes(MMATypes::f16) << " and "1978 << stringifyMMATypes(MMATypes::bf16);1979 }1980 1981 // Check result registers1982 int expectedOutput = 0;1983 if (typeD == WGMMATypes::f32 || typeD == WGMMATypes::s32)1984 expectedOutput = getShape().getN() / 2;1985 if (typeD == WGMMATypes::f16)1986 expectedOutput = getShape().getN() / 4;1987 if (outputSize != expectedOutput) {1988 return emitOpError() << "results " << expectedOutput1989 << ", however output struct has " << outputSize1990 << " elements";1991 }1992 // Check satfinite (only available for s32 accumulator)1993 if (typeD != WGMMATypes::s32 &&1994 getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==1995 NVVM::MMAIntOverflow::satfinite) {1996 return emitOpError()1997 << " `satfinite` can be only used with s32 accumulator, however "1998 "the current accumulator is "1999 << NVVM::stringifyWGMMATypes(typeD);2000 }2001 2002 return success();2003}2004 2005std::string NVVM::WgmmaMmaAsyncOp::getPtx() {2006 2007 int m = getShape().getM(), n = getShape().getN(), k = getShape().getK();2008 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;2009 2010 StringRef outputTypeName = stringifyWGMMATypes(getTypeD());2011 2012 int expectedOutputRegisters = 0;2013 if (getTypeD() == WGMMATypes::f16)2014 expectedOutputRegisters = getShape().getN() / 4;2015 else2016 expectedOutputRegisters = getShape().getN() / 2;2017 2018 std::string ptx;2019 llvm::raw_string_ostream ss(ptx);2020 2021 ss << "{\n"2022 ".reg .pred p;\n"2023 "setp.ne.b32 p, $"2024 << ((expectedOutputRegisters * 2) + 2)2025 << ", 0;\n"2026 "wgmma.mma_async.sync.aligned.m"2027 << m << "n" << n << "k" << k << "." << outputTypeName << "."2028 << stringifyWGMMATypes(getTypeA()) << "."2029 << stringifyWGMMATypes(getTypeB());2030 if (getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==2031 NVVM::MMAIntOverflow::satfinite)2032 ss << ".satfinite";2033 ss << " {";2034 int regCnt = 0;2035 for (; regCnt < expectedOutputRegisters; ++regCnt) {2036 ss << "$" << regCnt;2037 if (regCnt != expectedOutputRegisters - 1)2038 ss << ", ";2039 }2040 2041 ss << "},";2042 // Need to map read/write registers correctly.2043 regCnt = (regCnt * 2);2044 ss << " $" << (regCnt) << ","2045 << " $" << (regCnt + 1) << ","2046 << " p";2047 if (getTypeD() != WGMMATypes::s32) {2048 ss << ", $" << (regCnt + 3) << ", $" << (regCnt + 4);2049 }2050 // Don't add transpose parameters unless needed.2051 if (isF16) {2052 ss << ", $" << (regCnt + 5) << ", $" << (regCnt + 6);2053 }2054 ss << ";\n"2055 << "}\n";2056 return ptx;2057}2058 2059bool NVVM::WgmmaMmaAsyncOp::getAsmValues(2060 RewriterBase &rewriter,2061 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>2062 &asmValues) {2063 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;2064 if (getResults())2065 asmValues.push_back({getResults(), mlir::NVVM::PTXRegisterMod::Write});2066 if (getInouts())2067 asmValues.push_back({getInouts(), mlir::NVVM::PTXRegisterMod::ReadWrite});2068 asmValues.push_back({getDescriptorA(), mlir::NVVM::PTXRegisterMod::Read});2069 asmValues.push_back({getDescriptorB(), mlir::NVVM::PTXRegisterMod::Read});2070 asmValues.push_back({makeConstantI32(rewriter, static_cast<int>(getScaleD())),2071 mlir::NVVM::PTXRegisterMod::Read});2072 if (getTypeD() != WGMMATypes::s32) {2073 asmValues.push_back(2074 {makeConstantI32(rewriter,2075 getScaleA() == NVVM::WGMMAScaleIn::neg ? -1 : 1),2076 mlir::NVVM::PTXRegisterMod::Read});2077 asmValues.push_back(2078 {makeConstantI32(rewriter,2079 getScaleB() == NVVM::WGMMAScaleIn::neg ? -1 : 1),2080 mlir::NVVM::PTXRegisterMod::Read});2081 }2082 if (isF16) {2083 asmValues.push_back(2084 {makeConstantI32(rewriter, static_cast<int>(getLayoutA())),2085 mlir::NVVM::PTXRegisterMod::Read});2086 asmValues.push_back(2087 {makeConstantI32(rewriter, 1 - static_cast<int>(getLayoutB())),2088 mlir::NVVM::PTXRegisterMod::Read});2089 }2090 return true; // Has manual mapping2091}2092 2093LogicalResult NVVM::FenceProxyOp::verify() {2094 if (getKind() == NVVM::ProxyKind::TENSORMAP)2095 return emitOpError() << "tensormap proxy is not a supported proxy kind";2096 if (getKind() == NVVM::ProxyKind::GENERIC)2097 return emitOpError() << "generic proxy not a supported proxy kind";2098 if (getKind() == NVVM::ProxyKind::async_shared && !getSpace().has_value()) {2099 return emitOpError() << "async_shared fence requires space attribute";2100 }2101 if (getKind() != NVVM::ProxyKind::async_shared && getSpace().has_value()) {2102 return emitOpError() << "only async_shared fence can have space attribute";2103 }2104 return success();2105}2106 2107LogicalResult NVVM::FenceProxyAcquireOp::verify() {2108 if (getFromProxy() != NVVM::ProxyKind::GENERIC)2109 return emitOpError("uni-directional proxies only support generic for "2110 "from_proxy attribute");2111 2112 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)2113 return emitOpError("uni-directional proxies only support tensormap "2114 "for to_proxy attribute");2115 2116 return success();2117}2118 2119LogicalResult NVVM::FenceProxyReleaseOp::verify() {2120 if (getFromProxy() != NVVM::ProxyKind::GENERIC)2121 return emitOpError("uni-directional proxies only support generic for "2122 "from_proxy attribute");2123 2124 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)2125 return emitOpError("uni-directional proxies only support tensormap "2126 "for to_proxy attribute");2127 2128 return success();2129}2130 2131LogicalResult NVVM::SetMaxRegisterOp::verify() {2132 if (getRegCount() % 8)2133 return emitOpError("new register size must be multiple of 8");2134 if (getRegCount() < 24 || getRegCount() > 256)2135 return emitOpError("new register size must be in between 24 to 256");2136 return success();2137}2138 2139LogicalResult NVVM::BarrierOp::verify() {2140 if (getNumberOfThreads() && !getBarrierId())2141 return emitOpError(2142 "barrier id is missing, it should be set between 0 to 15");2143 2144 if (getBarrierId() && (getReductionOp() || getReductionPredicate()))2145 return emitOpError("reduction are only available when id is 0");2146 2147 if ((getReductionOp() && !getReductionPredicate()) ||2148 (!getReductionOp() && getReductionPredicate()))2149 return emitOpError("reduction predicate and reduction operation must be "2150 "specified together");2151 2152 return success();2153}2154 2155LogicalResult NVVM::Tcgen05CpOp::verify() {2156 auto mc = getMulticast();2157 2158 using SH = Tcgen05CpShape;2159 using MC = Tcgen05CpMulticast;2160 switch (getShape()) {2161 case SH::SHAPE_128x256b:2162 case SH::SHAPE_128x128b:2163 case SH::SHAPE_4x256b:2164 if (mc != MC::NONE)2165 return emitError("Invalid multicast type for tcgen05.cp Op");2166 break;2167 case SH::SHAPE_64x128b:2168 if (mc != MC::WARPX2_01_23 && mc != MC::WARPX2_02_13)2169 return emitError("Shape 64x128b requires multicast warpx2_01_23 or "2170 "warpx2_02_13 for tcgen05.cp Op");2171 break;2172 case SH::SHAPE_32x128b:2173 if (mc != MC::WARPX4)2174 return emitError(2175 "Shape 32x128b requires multicast warpx4 for tcgen05.cp Op");2176 break;2177 }2178 return success();2179}2180 2181LogicalResult NVVM::MatchSyncOp::verify() {2182 if (getKind() == NVVM::MatchSyncKind::all) {2183 auto type = llvm::dyn_cast<LLVM::LLVMStructType>(getType());2184 if (!type || type.getBody().size() != 2 ||2185 !type.getBody()[0].isInteger(32) || !type.getBody()[1].isInteger(1)) {2186 return emitOpError("match.sync 'all' returns a two element struct with "2187 "first element as i32 and second element as i1");2188 }2189 } else {2190 if (!getType().isInteger(32)) {2191 return emitOpError("match.sync 'any' returns an i32");2192 }2193 }2194 return success();2195}2196 2197LogicalResult NVVM::VoteSyncOp::verify() {2198 if (getKind() == NVVM::VoteSyncKind::ballot) {2199 if (!getType().isInteger(32)) {2200 return emitOpError("vote.sync 'ballot' returns an i32");2201 }2202 } else {2203 if (!getType().isInteger(1)) {2204 return emitOpError("vote.sync 'any', 'all' and 'uni' returns an i1");2205 }2206 }2207 return success();2208}2209 2210LogicalResult NVVM::PrefetchOp::verify() {2211 using MemSpace = NVVM::NVVMMemorySpace;2212 using CacheLevel = NVVM::PrefetchCacheLevel;2213 2214 unsigned addressSpace =2215 llvm::cast<LLVM::LLVMPointerType>(getAddr().getType()).getAddressSpace();2216 std::optional<NVVM::CacheEvictionPriority> evictPriority = getEvictPriority();2217 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = getCacheLevel();2218 2219 if (getTensormap() && cacheLevel)2220 return emitOpError("cannot specify both tensormap and cache level");2221 2222 if (getTensormap()) {2223 if (addressSpace != MemSpace::Generic &&2224 addressSpace != MemSpace::Constant) {2225 return emitOpError(2226 "prefetch tensormap requires a generic or constant pointer");2227 }2228 2229 if (evictPriority) {2230 return emitOpError(2231 "prefetch tensormap does not support eviction priority");2232 }2233 2234 if (getInParamSpace() && addressSpace != MemSpace::Generic) {2235 return emitOpError(2236 "in_param_space can only be specified for a generic pointer");2237 }2238 2239 } else if (cacheLevel) {2240 if (addressSpace != MemSpace::Generic && addressSpace != MemSpace::Global &&2241 addressSpace != MemSpace::Local) {2242 return emitOpError("prefetch to cache level requires a generic, global, "2243 "or local pointer");2244 }2245 2246 if (getUniform()) {2247 if (*cacheLevel != CacheLevel::L1) {2248 return emitOpError(2249 "unsupported cache level, the only supported uniform "2250 "cache level is L1");2251 }2252 2253 if (addressSpace != MemSpace::Generic) {2254 return emitOpError(2255 "prefetch to uniform cache requires a generic pointer");2256 }2257 }2258 2259 if (evictPriority) {2260 if (*cacheLevel != CacheLevel::L2)2261 return emitOpError(2262 "cache eviction priority supported only for cache level L2");2263 2264 if (addressSpace != MemSpace::Global)2265 return emitOpError("cache eviction priority requires a global pointer");2266 2267 if (*evictPriority != NVVM::CacheEvictionPriority::EvictNormal &&2268 *evictPriority != NVVM::CacheEvictionPriority::EvictLast)2269 return emitOpError(2270 "unsupported cache eviction priority, only evict_last and "2271 "evict_normal are supported");2272 }2273 2274 if (getPredicate())2275 return emitOpError("predicate supported only on prefetch tensormap");2276 2277 } else {2278 return emitOpError(2279 "requires specification of either cache level or tensormap");2280 }2281 2282 return success();2283}2284 2285LogicalResult NVVM::ClusterLaunchControlQueryCancelOp::verify() {2286 switch (getQueryType()) {2287 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:2288 if (!getType().isInteger(1))2289 return emitOpError("is_canceled query type returns an i1");2290 break;2291 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:2292 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:2293 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:2294 if (!getType().isInteger(32)) {2295 return emitOpError("get_first_cta_id_x, get_first_cta_id_y, "2296 "get_first_cta_id_z query types return an i32");2297 }2298 break;2299 }2300 return success();2301}2302 2303LogicalResult NVVM::ReduxOp::verify() {2304 mlir::Type reduxType = getType();2305 2306 if (!reduxType.isF32()) {2307 if (getAbs())2308 return emitOpError("abs attribute is supported only for f32 type");2309 if (getNan())2310 return emitOpError("nan attribute is supported only for f32 type");2311 }2312 2313 NVVM::ReduxKind kind = getKind();2314 switch (kind) {2315 case NVVM::ReduxKind::ADD:2316 case NVVM::ReduxKind::AND:2317 case NVVM::ReduxKind::OR:2318 case NVVM::ReduxKind::XOR:2319 case NVVM::ReduxKind::MAX:2320 case NVVM::ReduxKind::MIN:2321 case NVVM::ReduxKind::UMAX:2322 case NVVM::ReduxKind::UMIN:2323 if (!reduxType.isInteger(32))2324 return emitOpError("'")2325 << stringifyEnum(kind) << "' redux kind unsupported with "2326 << reduxType << " type. Only supported type is 'i32'.";2327 break;2328 case NVVM::ReduxKind::FMIN:2329 case NVVM::ReduxKind::FMAX:2330 if (!reduxType.isF32())2331 return emitOpError("'")2332 << stringifyEnum(kind) << "' redux kind unsupported with "2333 << reduxType << " type. Only supported type is 'f32'.";2334 break;2335 }2336 2337 return success();2338}2339 2340/// Packs the given `field` into the `result`.2341/// The `result` is 64-bits and each `field` can be 32-bits or narrower.2342static llvm::Value *2343packValInto64Bits(llvm::IRBuilderBase &builder,2344 llvm::Value *result, // the `result` (unset bits are zero)2345 llvm::Value *field, // `field` to pack into `result`2346 unsigned sizeInBits, // Size of `field` in bits2347 unsigned start) { // Starting bit within `result`2348 field = builder.CreateZExtOrBitCast(field, builder.getInt32Ty());2349 2350 unsigned mask = (sizeInBits < 32 ? ((1u << sizeInBits) - 1) : 0xffffffffu);2351 if (mask != 0xffffffffu)2352 field = builder.CreateAnd(field, builder.getInt32(mask));2353 2354 field = builder.CreateZExtOrBitCast(field, builder.getInt64Ty());2355 field = builder.CreateShl(field, start);2356 2357 return builder.CreateOr(result, field);2358}2359 2360void Tcgen05MmaSmemDescOp::createSmemDescriptor(Operation &op,2361 LLVM::ModuleTranslation &mt,2362 llvm::IRBuilderBase &builder) {2363 auto thisOp = cast<NVVM::Tcgen05MmaSmemDescOp>(op);2364 llvm::Value *smemDesc = builder.getInt64(0);2365 2366 smemDesc = packValInto64Bits(builder, smemDesc,2367 mt.lookupValue(thisOp.getStartAddr()), 14, 0);2368 smemDesc = packValInto64Bits(2369 builder, smemDesc, mt.lookupValue(thisOp.getLeadingDimOffset()), 14, 16);2370 smemDesc = packValInto64Bits(2371 builder, smemDesc, mt.lookupValue(thisOp.getStrideDimOffset()), 14, 32);2372 2373 smemDesc = packValInto64Bits(builder, smemDesc, builder.getInt32(1), 3, 46);2374 smemDesc = packValInto64Bits(builder, smemDesc,2375 mt.lookupValue(thisOp.getBaseOffset()), 3, 49);2376 smemDesc = packValInto64Bits(2377 builder, smemDesc, mt.lookupValue(thisOp.getLeadingDimMode()), 1, 52);2378 smemDesc = packValInto64Bits(builder, smemDesc,2379 mt.lookupValue(thisOp.getSwizzleMode()), 3, 61);2380 2381 mt.mapValue(thisOp.getRes()) = smemDesc;2382}2383 2384//===----------------------------------------------------------------------===//2385// getPtx methods2386//===----------------------------------------------------------------------===//2387 2388std::string NVVM::MBarrierInitOp::getPtx() {2389 bool isShared = isPtrInSharedCTASpace(getAddr());2390 return isShared ? std::string("mbarrier.init.shared.b64 [%0], %1;")2391 : std::string("mbarrier.init.b64 [%0], %1;");2392}2393 2394std::string NVVM::MBarrierArriveExpectTxOp::getPtx() {2395 bool isShared = isPtrInSharedCTASpace(getAddr());2396 return isShared2397 ? std::string("mbarrier.arrive.expect_tx.shared.b64 _, [%0], %1;")2398 : std::string("mbarrier.arrive.expect_tx.b64 _, [%0], %1;");2399}2400 2401std::string NVVM::MBarrierTryWaitParityOp::getPtx() {2402 bool isShared = isPtrInSharedCTASpace(getAddr());2403 llvm::StringRef space = isShared ? ".shared" : "";2404 2405 return llvm::formatv("{\n\t"2406 ".reg .pred P1; \n\t"2407 "LAB_WAIT: \n\t"2408 "mbarrier.try_wait.parity{0}.b64 P1, [%0], %1, %2; \n\t"2409 "@P1 bra.uni DONE; \n\t"2410 "bra.uni LAB_WAIT; \n\t"2411 "DONE: \n\t"2412 "}",2413 space);2414}2415 2416//===----------------------------------------------------------------------===//2417// getIntrinsicID/getIntrinsicIDAndArgs methods2418//===----------------------------------------------------------------------===//2419 2420mlir::NVVM::IDArgPair NVVM::BarrierOp::getIntrinsicIDAndArgs(2421 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2422 auto thisOp = cast<NVVM::BarrierOp>(op);2423 llvm::Value *barrierId = thisOp.getBarrierId()2424 ? mt.lookupValue(thisOp.getBarrierId())2425 : builder.getInt32(0);2426 llvm::Intrinsic::ID id;2427 llvm::SmallVector<llvm::Value *> args;2428 if (thisOp.getNumberOfThreads()) {2429 id = llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_count;2430 args.push_back(barrierId);2431 args.push_back(mt.lookupValue(thisOp.getNumberOfThreads()));2432 } else if (thisOp.getReductionOp()) {2433 switch (*thisOp.getReductionOp()) {2434 case NVVM::BarrierReduction::AND:2435 id = llvm::Intrinsic::nvvm_barrier0_and;2436 break;2437 case NVVM::BarrierReduction::OR:2438 id = llvm::Intrinsic::nvvm_barrier0_or;2439 break;2440 case NVVM::BarrierReduction::POPC:2441 id = llvm::Intrinsic::nvvm_barrier0_popc;2442 break;2443 }2444 args.push_back(mt.lookupValue(thisOp.getReductionPredicate()));2445 } else {2446 id = llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_all;2447 args.push_back(barrierId);2448 }2449 2450 return {id, std::move(args)};2451}2452 2453mlir::NVVM::IDArgPair MBarrierInitOp::getIntrinsicIDAndArgs(2454 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2455 auto thisOp = cast<NVVM::MBarrierInitOp>(op);2456 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());2457 llvm::Intrinsic::ID id = isShared ? llvm::Intrinsic::nvvm_mbarrier_init_shared2458 : llvm::Intrinsic::nvvm_mbarrier_init;2459 2460 // Fill the Intrinsic Args2461 llvm::SmallVector<llvm::Value *> args;2462 args.push_back(mt.lookupValue(thisOp.getAddr()));2463 args.push_back(mt.lookupValue(thisOp.getCount()));2464 2465 return {id, std::move(args)};2466}2467 2468mlir::NVVM::IDArgPair MBarrierInvalOp::getIntrinsicIDAndArgs(2469 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2470 auto thisOp = cast<NVVM::MBarrierInvalOp>(op);2471 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());2472 llvm::Intrinsic::ID id = isShared2473 ? llvm::Intrinsic::nvvm_mbarrier_inval_shared2474 : llvm::Intrinsic::nvvm_mbarrier_inval;2475 2476 return {id, {mt.lookupValue(thisOp.getAddr())}};2477}2478 2479mlir::NVVM::IDArgPair MBarrierExpectTxOp::getIntrinsicIDAndArgs(2480 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2481 auto thisOp = cast<NVVM::MBarrierExpectTxOp>(op);2482 2483 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());2484 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;2485 // bit-0: Space2486 // bit-1: Scope2487 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);2488 2489 static constexpr llvm::Intrinsic::ID IDs[] = {2490 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cta,2491 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cluster,2492 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cta,2493 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cluster};2494 2495 // Fill the Intrinsic Args2496 llvm::SmallVector<llvm::Value *> args;2497 args.push_back(mt.lookupValue(thisOp.getAddr()));2498 args.push_back(mt.lookupValue(thisOp.getTxcount()));2499 2500 return {IDs[index], std::move(args)};2501}2502 2503mlir::NVVM::IDArgPair MBarrierCompleteTxOp::getIntrinsicIDAndArgs(2504 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2505 auto thisOp = cast<NVVM::MBarrierCompleteTxOp>(op);2506 2507 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());2508 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;2509 // bit-0: Space2510 // bit-1: Scope2511 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);2512 2513 static constexpr llvm::Intrinsic::ID IDs[] = {2514 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cta,2515 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cluster,2516 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cta,2517 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cluster};2518 2519 // Fill the Intrinsic Args2520 llvm::SmallVector<llvm::Value *> args;2521 args.push_back(mt.lookupValue(thisOp.getAddr()));2522 args.push_back(mt.lookupValue(thisOp.getTxcount()));2523 2524 return {IDs[index], std::move(args)};2525}2526 2527mlir::NVVM::IDArgPair MBarrierArriveOp::getIntrinsicIDAndArgs(2528 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2529 auto thisOp = cast<NVVM::MBarrierArriveOp>(op);2530 2531 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());2532 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;2533 // bit-0: Space2534 // bit-1: Scope2535 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);2536 2537 static constexpr llvm::Intrinsic::ID IDs[] = {2538 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta,2539 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cluster,2540 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cta,2541 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cluster};2542 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {2543 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cta,2544 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cluster,2545 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cta,2546 llvm::Intrinsic::2547 nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cluster};2548 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];2549 2550 // Tidy-up the Intrinsic Args2551 bool needCast = isPtrInGenericSpace(thisOp.getAddr());2552 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());2553 if (needCast)2554 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);2555 2556 // When count is not explicitly specified, the default is 1.2557 llvm::LLVMContext &ctx = mt.getLLVMContext();2558 bool hasCount = static_cast<bool>(thisOp.getCount());2559 llvm::Value *count =2560 hasCount ? mt.lookupValue(thisOp.getCount())2561 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);2562 2563 return {id, {mbar, count}};2564}2565 2566mlir::NVVM::IDArgPair MBarrierArriveDropOp::getIntrinsicIDAndArgs(2567 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2568 auto thisOp = cast<NVVM::MBarrierArriveDropOp>(op);2569 2570 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());2571 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;2572 // bit-0: Space2573 // bit-1: Scope2574 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);2575 2576 static constexpr llvm::Intrinsic::ID IDs[] = {2577 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cta,2578 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cluster,2579 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cta,2580 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cluster};2581 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {2582 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cta,2583 llvm::Intrinsic::2584 nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cluster,2585 llvm::Intrinsic::2586 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cta,2587 llvm::Intrinsic::2588 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cluster};2589 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];2590 2591 // Tidy-up the Intrinsic Args2592 bool needCast = isPtrInGenericSpace(thisOp.getAddr());2593 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());2594 if (needCast)2595 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);2596 2597 // When count is not explicitly specified, the default is 1.2598 llvm::LLVMContext &ctx = mt.getLLVMContext();2599 bool hasCount = static_cast<bool>(thisOp.getCount());2600 llvm::Value *count =2601 hasCount ? mt.lookupValue(thisOp.getCount())2602 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);2603 2604 return {id, {mbar, count}};2605}2606 2607bool MBarrierArriveExpectTxOp::getAsmValues(2608 RewriterBase &rewriter,2609 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>2610 &asmValues) {2611 // Add all the operands but not the attrs to the asmValues list.2612 // The attrs here are used to generate the right variants for2613 // intrinsics-lowering. So, we ignore them while generating inline-PTX.2614 for (auto val : getOperands())2615 asmValues.push_back({val, mlir::NVVM::PTXRegisterMod::Read});2616 2617 return false;2618}2619 2620mlir::NVVM::IDArgPair MBarrierArriveExpectTxOp::getIntrinsicIDAndArgs(2621 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2622 auto thisOp = cast<NVVM::MBarrierArriveExpectTxOp>(op);2623 2624 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());2625 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;2626 // bit-0: Space2627 // bit-1: Scope2628 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);2629 2630 // clang-format off2631 static constexpr llvm::Intrinsic::ID IDs[] = {2632 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cta,2633 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cluster,2634 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cta,2635 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cluster};2636 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {2637 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cta,2638 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cluster,2639 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cta,2640 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cluster};2641 // clang-format on2642 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];2643 2644 // Tidy-up the Intrinsic Args2645 llvm::Value *txcount = mt.lookupValue(thisOp.getTxcount());2646 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());2647 bool needCast = isPtrInGenericSpace(thisOp.getAddr());2648 if (needCast)2649 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);2650 2651 return {id, {mbar, txcount}};2652}2653 2654mlir::NVVM::IDArgPair MBarrierArriveDropExpectTxOp::getIntrinsicIDAndArgs(2655 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2656 auto thisOp = cast<NVVM::MBarrierArriveDropExpectTxOp>(op);2657 2658 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());2659 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;2660 // bit-0: Space2661 // bit-1: Scope2662 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);2663 2664 // clang-format off2665 static constexpr llvm::Intrinsic::ID IDs[] = {2666 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cta,2667 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cluster,2668 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cta,2669 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cluster};2670 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {2671 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cta,2672 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cluster,2673 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cta,2674 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cluster};2675 // clang-format on2676 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];2677 2678 // Tidy-up the Intrinsic Args2679 llvm::Value *txcount = mt.lookupValue(thisOp.getTxcount());2680 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());2681 bool needCast = isPtrInGenericSpace(thisOp.getAddr());2682 if (needCast)2683 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);2684 2685 return {id, {mbar, txcount}};2686}2687 2688mlir::NVVM::IDArgPair MBarrierArriveNocompleteOp::getIntrinsicIDAndArgs(2689 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2690 auto thisOp = cast<NVVM::MBarrierArriveNocompleteOp>(op);2691 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());2692 llvm::Intrinsic::ID id =2693 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete_shared2694 : llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete;2695 // Fill the Intrinsic Args2696 llvm::SmallVector<llvm::Value *> args;2697 args.push_back(mt.lookupValue(thisOp.getAddr()));2698 args.push_back(mt.lookupValue(thisOp.getCount()));2699 2700 return {id, std::move(args)};2701}2702 2703mlir::NVVM::IDArgPair MBarrierArriveDropNocompleteOp::getIntrinsicIDAndArgs(2704 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2705 auto thisOp = cast<NVVM::MBarrierArriveDropNocompleteOp>(op);2706 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());2707 llvm::Intrinsic::ID id =2708 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete_shared2709 : llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete;2710 // Fill the Intrinsic Args2711 llvm::SmallVector<llvm::Value *> args;2712 args.push_back(mt.lookupValue(thisOp.getAddr()));2713 args.push_back(mt.lookupValue(thisOp.getCount()));2714 2715 return {id, std::move(args)};2716}2717 2718mlir::NVVM::IDArgPair MBarrierTestWaitOp::getIntrinsicIDAndArgs(2719 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2720 auto thisOp = cast<NVVM::MBarrierTestWaitOp>(op);2721 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());2722 llvm::Intrinsic::ID id = isShared2723 ? llvm::Intrinsic::nvvm_mbarrier_test_wait_shared2724 : llvm::Intrinsic::nvvm_mbarrier_test_wait;2725 // Fill the Intrinsic Args2726 llvm::SmallVector<llvm::Value *> args;2727 args.push_back(mt.lookupValue(thisOp.getAddr()));2728 args.push_back(mt.lookupValue(thisOp.getState()));2729 2730 return {id, std::move(args)};2731}2732 2733mlir::NVVM::IDArgPair CpAsyncMBarrierArriveOp::getIntrinsicIDAndArgs(2734 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2735 auto thisOp = cast<NVVM::CpAsyncMBarrierArriveOp>(op);2736 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());2737 2738 llvm::Intrinsic::ID id;2739 if (thisOp.getNoinc()) {2740 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc_shared2741 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc;2742 } else {2743 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_shared2744 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive;2745 }2746 2747 return {id, {mt.lookupValue(thisOp.getAddr())}};2748}2749 2750#define CP_ASYNC_ID_IMPL(mod, size, suffix) \2751 llvm::Intrinsic::nvvm_cp_async_##mod##_shared_global_##size##suffix2752 2753#define GET_CP_ASYNC_ID(mod, size, has_cpsize) \2754 has_cpsize ? CP_ASYNC_ID_IMPL(mod, size, _s) : CP_ASYNC_ID_IMPL(mod, size, )2755 2756llvm::Intrinsic::ID2757CpAsyncOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,2758 llvm::SmallVector<llvm::Value *> &args) {2759 llvm::Intrinsic::ID id;2760 2761 auto cpAsyncOp = cast<NVVM::CpAsyncOp>(op);2762 bool hasCpSize = static_cast<bool>(cpAsyncOp.getCpSize());2763 switch (cpAsyncOp.getSize()) {2764 case 4:2765 id = GET_CP_ASYNC_ID(ca, 4, hasCpSize);2766 break;2767 case 8:2768 id = GET_CP_ASYNC_ID(ca, 8, hasCpSize);2769 break;2770 case 16:2771 id = (cpAsyncOp.getModifier() == NVVM::LoadCacheModifierKind::CG)2772 ? GET_CP_ASYNC_ID(cg, 16, hasCpSize)2773 : GET_CP_ASYNC_ID(ca, 16, hasCpSize);2774 break;2775 default:2776 llvm_unreachable("Invalid copy size in CpAsyncOp.");2777 }2778 2779 // Fill the Intrinsic Args2780 args.push_back(mt.lookupValue(cpAsyncOp.getDst()));2781 args.push_back(mt.lookupValue(cpAsyncOp.getSrc()));2782 if (hasCpSize)2783 args.push_back(mt.lookupValue(cpAsyncOp.getCpSize()));2784 2785 return id;2786}2787 2788mlir::NVVM::IDArgPair CpAsyncBulkPrefetchOp::getIntrinsicIDAndArgs(2789 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2790 auto thisOp = cast<NVVM::CpAsyncBulkPrefetchOp>(op);2791 llvm::SmallVector<llvm::Value *> args;2792 llvm::Intrinsic::ID id = llvm::Intrinsic::nvvm_cp_async_bulk_prefetch_L2;2793 2794 // Fill the Intrinsic Args2795 args.push_back(mt.lookupValue(thisOp.getSrcMem()));2796 args.push_back(mt.lookupValue(thisOp.getSize()));2797 2798 mlir::Value cacheHint = thisOp.getL2CacheHint();2799 const bool hasCacheHint = static_cast<bool>(cacheHint);2800 llvm::Value *i64Unused =2801 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);2802 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);2803 args.push_back(builder.getInt1(hasCacheHint));2804 2805 return {id, std::move(args)};2806}2807 2808mlir::NVVM::IDArgPair CpAsyncBulkGlobalToSharedClusterOp::getIntrinsicIDAndArgs(2809 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2810 auto thisOp = cast<NVVM::CpAsyncBulkGlobalToSharedClusterOp>(op);2811 llvm::SmallVector<llvm::Value *> args;2812 2813 // Fill the Intrinsic Args: dst, mbar, src, size.2814 args.push_back(mt.lookupValue(thisOp.getDstMem()));2815 args.push_back(mt.lookupValue(thisOp.getMbar()));2816 args.push_back(mt.lookupValue(thisOp.getSrcMem()));2817 args.push_back(mt.lookupValue(thisOp.getSize()));2818 2819 // Multicast mask for shared::cluster only, if available.2820 mlir::Value multicastMask = thisOp.getMulticastMask();2821 const bool hasMulticastMask = static_cast<bool>(multicastMask);2822 const bool isSharedCTA = isPtrInSharedCTASpace(thisOp.getDstMem());2823 if (!isSharedCTA) {2824 llvm::Value *i16Unused = llvm::ConstantInt::get(builder.getInt16Ty(), 0);2825 args.push_back(hasMulticastMask ? mt.lookupValue(multicastMask)2826 : i16Unused);2827 }2828 2829 // Cache hint, if available.2830 mlir::Value cacheHint = thisOp.getL2CacheHint();2831 const bool hasCacheHint = static_cast<bool>(cacheHint);2832 llvm::Value *i64Unused = llvm::ConstantInt::get(builder.getInt64Ty(), 0);2833 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);2834 2835 // Flag arguments for multicast and cachehint.2836 if (!isSharedCTA)2837 args.push_back(builder.getInt1(hasMulticastMask));2838 args.push_back(builder.getInt1(hasCacheHint));2839 2840 llvm::Intrinsic::ID id =2841 isSharedCTA2842 ? llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cta2843 : llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cluster;2844 2845 return {id, std::move(args)};2846}2847 2848mlir::NVVM::IDArgPair CpAsyncBulkSharedCTAToGlobalOp::getIntrinsicIDAndArgs(2849 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2850 auto thisOp = cast<NVVM::CpAsyncBulkSharedCTAToGlobalOp>(op);2851 llvm::SmallVector<llvm::Value *> args;2852 llvm::Intrinsic::ID id =2853 llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global;2854 2855 // Fill the Intrinsic Args2856 args.push_back(mt.lookupValue(thisOp.getDstMem()));2857 args.push_back(mt.lookupValue(thisOp.getSrcMem()));2858 args.push_back(mt.lookupValue(thisOp.getSize()));2859 2860 mlir::Value cacheHint = thisOp.getL2CacheHint();2861 const bool hasCacheHint = static_cast<bool>(cacheHint);2862 llvm::Value *i64Unused =2863 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);2864 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);2865 args.push_back(builder.getInt1(hasCacheHint));2866 2867 // Choose the bytemask variant2868 if (mlir::Value byteMask = thisOp.getByteMask()) {2869 args.push_back(mt.lookupValue(byteMask));2870 id = llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global_bytemask;2871 }2872 2873 return {id, std::move(args)};2874}2875 2876bool CpAsyncBulkTensorGlobalToSharedClusterOp::getAsmValues(2877 RewriterBase &rewriter,2878 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>2879 &asmValues) {2880 // Add all the operands but not the attrs to the asmValues list.2881 // The attrs here are used to generate the right variants for2882 // intrinsics-lowering. So, we ignore them while generating inline-PTX.2883 for (auto val : getOperands())2884 asmValues.push_back({val, mlir::NVVM::PTXRegisterMod::Read});2885 2886 return false;2887}2888 2889mlir::NVVM::IDArgPair2890CpAsyncBulkTensorGlobalToSharedClusterOp::getIntrinsicIDAndArgs(2891 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {2892 auto thisOp = cast<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(op);2893 const bool isCTAOnly = thisOp.getIsCTAOnly();2894 llvm::SmallVector<llvm::Value *> args;2895 2896 // Fill the Intrinsic Args2897 args.push_back(mt.lookupValue(thisOp.getDstMem()));2898 args.push_back(mt.lookupValue(thisOp.getMbar()));2899 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));2900 2901 // Coordinates and im2col-offsets2902 for (mlir::Value v : thisOp.getCoordinates())2903 args.push_back(mt.lookupValue(v));2904 for (mlir::Value v : thisOp.getIm2colOffsets())2905 args.push_back(mt.lookupValue(v));2906 2907 // MulticastMask, if available2908 mlir::Value mcMask = thisOp.getMulticastMask();2909 const bool hasMC = static_cast<bool>(mcMask);2910 llvm::Value *i16Zero =2911 llvm::ConstantInt::get(llvm::Type::getInt16Ty(mt.getLLVMContext()), 0);2912 2913 // CacheHint, if available2914 mlir::Value cacheHint = thisOp.getL2CacheHint();2915 const bool hasCacheHint = static_cast<bool>(cacheHint);2916 llvm::Value *i64Zero =2917 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);2918 2919 // Flag argument CTAGroup2920 // CTA_1/2 is mapped to values 1 and 2 for the intrinsics.2921 // Hence, the +1 to getGroup().2922 const int32_t val =2923 thisOp.getGroup() ? (static_cast<int32_t>(*thisOp.getGroup()) + 1) : 0;2924 llvm::Value *cg =2925 llvm::ConstantInt::get(llvm::Type::getInt32Ty(mt.getLLVMContext()), val);2926 2927 if (!isCTAOnly) {2928 // For shared::cluster, all the arguments that we build are applicable.2929 args.push_back(hasMC ? mt.lookupValue(mcMask) : i16Zero);2930 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Zero);2931 args.push_back(builder.getInt1(hasMC));2932 args.push_back(builder.getInt1(hasCacheHint));2933 args.push_back(cg);2934 } else {2935 // For shared::cta, only cache-hint is applicable.2936 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Zero);2937 args.push_back(builder.getInt1(hasCacheHint));2938 }2939 2940 constexpr size_t numDims = 5; // 1D to 5D2941 constexpr size_t numModes = 5; // Tile, Im2col, w, w_128, gather42942 using rowTy = std::array<llvm::Intrinsic::ID, numDims + 1>;2943 using TableTy = std::array<rowTy, numModes>;2944 static constexpr TableTy IDTable{2945 {{notIntrinsic, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_1d,2946 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_2d,2947 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_3d,2948 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_4d,2949 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_5d},2950 {notIntrinsic, notIntrinsic, notIntrinsic,2951 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_3d,2952 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_4d,2953 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_5d},2954 {notIntrinsic, notIntrinsic, notIntrinsic,2955 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_3d,2956 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_4d,2957 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_5d},2958 {notIntrinsic, notIntrinsic, notIntrinsic,2959 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_3d,2960 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_4d,2961 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_5d},2962 {notIntrinsic, notIntrinsic, notIntrinsic, notIntrinsic, notIntrinsic,2963 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_gather4_2d}}};2964 2965 static constexpr TableTy IDTableCTA{2966 {{notIntrinsic,2967 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_1d,2968 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_2d,2969 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_3d,2970 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_4d,2971 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_5d},2972 {notIntrinsic, notIntrinsic, notIntrinsic,2973 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_3d,2974 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_4d,2975 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_5d},2976 {notIntrinsic, notIntrinsic, notIntrinsic,2977 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_3d,2978 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_4d,2979 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_5d},2980 {notIntrinsic, notIntrinsic, notIntrinsic,2981 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_3d,2982 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_4d,2983 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_5d},2984 {notIntrinsic, notIntrinsic, notIntrinsic, notIntrinsic, notIntrinsic,2985 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_gather4_2d}}};2986 2987 static_assert(2988 (getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1) &&2989 (getMaxEnumValForTMALoadMode() == std::size(IDTableCTA) - 1),2990 "TMALoadModes must match number of rows in IDTable and IDTableCTA");2991 size_t mode = static_cast<size_t>(thisOp.getMode());2992 size_t dim = thisOp.getCoordinates().size();2993 auto id = isCTAOnly ? IDTableCTA[mode][dim] : IDTable[mode][dim];2994 assert(id != notIntrinsic &&2995 "Invalid intrinsic for CpAsyncBulkTensorGlobalToSharedClusterOp.");2996 2997 return {id, std::move(args)};2998}2999 3000mlir::NVVM::IDArgPair CpAsyncBulkTensorPrefetchOp::getIntrinsicIDAndArgs(3001 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3002 auto thisOp = cast<NVVM::CpAsyncBulkTensorPrefetchOp>(op);3003 llvm::SmallVector<llvm::Value *> args;3004 3005 // Fill the Intrinsic Args3006 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));3007 3008 for (auto v : thisOp.getCoordinates())3009 args.push_back(mt.lookupValue(v));3010 for (auto v : thisOp.getIm2colOffsets())3011 args.push_back(mt.lookupValue(v));3012 3013 mlir::Value cacheHint = thisOp.getL2CacheHint();3014 const bool hasCacheHint = static_cast<bool>(cacheHint);3015 llvm::Value *i64Unused =3016 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);3017 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);3018 args.push_back(builder.getInt1(hasCacheHint));3019 3020 const unsigned NI = llvm::Intrinsic::not_intrinsic;3021 static constexpr llvm::Intrinsic::ID IDTable[][6] = {3022 {NI, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_1d,3023 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_2d,3024 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_3d,3025 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_4d,3026 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_5d},3027 {NI, NI, NI,3028 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_3d,3029 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_4d,3030 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_5d},3031 {NI, NI, NI,3032 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_3d,3033 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_4d,3034 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_5d},3035 {NI, NI, NI,3036 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_3d,3037 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_4d,3038 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_5d},3039 {NI, NI, NI, NI, NI,3040 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_gather4_2d}};3041 3042 static_assert(getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1,3043 "TMALoadModes must match number of rows in IDTable");3044 size_t mode = static_cast<size_t>(thisOp.getMode());3045 size_t dim = thisOp.getCoordinates().size();3046 llvm::Intrinsic::ID id = IDTable[mode][dim];3047 if (id == llvm::Intrinsic::not_intrinsic)3048 llvm_unreachable("Invalid intrinsic for CpAsyncBulkTensorPrefetchOp.");3049 3050 return {id, std::move(args)};3051}3052 3053mlir::NVVM::IDArgPair3054CpAsyncBulkTensorSharedCTAToGlobalOp::getIntrinsicIDAndArgs(3055 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3056 auto thisOp = cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(op);3057 llvm::SmallVector<llvm::Value *> args;3058 3059 // Fill the Intrinsic Args3060 args.push_back(mt.lookupValue(thisOp.getSrcMem()));3061 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));3062 3063 for (auto v : thisOp.getCoordinates())3064 args.push_back(mt.lookupValue(v));3065 3066 mlir::Value cacheHint = thisOp.getL2CacheHint();3067 const bool hasCacheHint = static_cast<bool>(cacheHint);3068 llvm::Value *i64Unused =3069 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);3070 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);3071 args.push_back(builder.getInt1(hasCacheHint));3072 3073 const unsigned NI = llvm::Intrinsic::not_intrinsic;3074 static constexpr llvm::Intrinsic::ID IDTable[][6] = {3075 {NI, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_tile_1d,3076 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_tile_2d,3077 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_tile_3d,3078 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_tile_4d,3079 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_tile_5d},3080 {NI, NI, NI, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_im2col_3d,3081 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_im2col_4d,3082 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_im2col_5d},3083 {NI, NI, NI, NI, NI,3084 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_2d}};3085 3086 static_assert(getMaxEnumValForTMAStoreMode() == std::size(IDTable) - 1,3087 "TMAStoreModes must match number of rows in IDTable");3088 size_t mode = static_cast<size_t>(thisOp.getMode());3089 size_t dim = thisOp.getCoordinates().size();3090 llvm::Intrinsic::ID id = IDTable[mode][dim];3091 if (id == llvm::Intrinsic::not_intrinsic)3092 llvm_unreachable(3093 "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOp.");3094 3095 return {id, std::move(args)};3096}3097 3098NVVM::IDArgPair CpAsyncBulkTensorReduceOp::getIntrinsicIDAndArgs(3099 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3100 auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOp>(op);3101 llvm::LLVMContext &ctx = mt.getLLVMContext();3102 3103 llvm::SmallVector<llvm::Value *> args;3104 3105 // Arguments to the intrinsic:3106 // shared_mem_ptr, tmaDesc, tensorDims3107 // cache_hint(if applicable) and flag(boolean)3108 args.push_back(mt.lookupValue(thisOp.getSrcMem()));3109 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));3110 3111 for (Value v : thisOp.getCoordinates())3112 args.push_back(mt.lookupValue(v));3113 3114 mlir::Value cacheHint = thisOp.getL2CacheHint();3115 const bool hasCacheHint = static_cast<bool>(cacheHint);3116 llvm::Value *i64ZeroValue =3117 llvm::ConstantInt::get(llvm::Type::getInt64Ty(ctx), 0);3118 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64ZeroValue);3119 args.push_back(builder.getInt1(hasCacheHint));3120 3121 const llvm::Intrinsic::ID notIntrinsic = llvm::Intrinsic::not_intrinsic;3122 3123 constexpr unsigned numRedKinds = 8; // ADD, MIN, MAX, INC, DEC, AND, OR, XOR3124 constexpr unsigned numLayouts = 2; // TILE, IM2COL3125 constexpr unsigned maxDim = 5; // 1D to 5D3126 using row = std::array<llvm::Intrinsic::ID, maxDim + 1>;3127 using layoutTable = std::array<row, numLayouts>;3128 using fullTable = std::array<layoutTable, numRedKinds>;3129 static constexpr fullTable IDTable{3130 {// RedTy::ADD3131 {{{{notIntrinsic,3132 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_tile_1d,3133 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_tile_2d,3134 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_tile_3d,3135 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_tile_4d,3136 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_tile_5d}},3137 {{notIntrinsic, notIntrinsic, notIntrinsic,3138 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_im2col_3d,3139 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_im2col_4d,3140 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_add_im2col_5d}}}},3141 // RedTy::MIN3142 {{{{notIntrinsic,3143 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_tile_1d,3144 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_tile_2d,3145 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_tile_3d,3146 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_tile_4d,3147 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_tile_5d}},3148 {{notIntrinsic, notIntrinsic, notIntrinsic,3149 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_im2col_3d,3150 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_im2col_4d,3151 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_min_im2col_5d}}}},3152 // RedTy::MAX3153 {{{{notIntrinsic,3154 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_tile_1d,3155 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_tile_2d,3156 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_tile_3d,3157 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_tile_4d,3158 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_tile_5d}},3159 {{notIntrinsic, notIntrinsic, notIntrinsic,3160 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_im2col_3d,3161 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_im2col_4d,3162 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_max_im2col_5d}}}},3163 // RedTy::INC3164 {{{{notIntrinsic,3165 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_tile_1d,3166 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_tile_2d,3167 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_tile_3d,3168 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_tile_4d,3169 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_tile_5d}},3170 {{notIntrinsic, notIntrinsic, notIntrinsic,3171 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_im2col_3d,3172 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_im2col_4d,3173 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_inc_im2col_5d}}}},3174 // RedTy::DEC3175 {{{{notIntrinsic,3176 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_tile_1d,3177 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_tile_2d,3178 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_tile_3d,3179 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_tile_4d,3180 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_tile_5d}},3181 {{notIntrinsic, notIntrinsic, notIntrinsic,3182 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_im2col_3d,3183 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_im2col_4d,3184 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_dec_im2col_5d}}}},3185 // RedTy::AND3186 {{{{notIntrinsic,3187 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_tile_1d,3188 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_tile_2d,3189 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_tile_3d,3190 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_tile_4d,3191 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_tile_5d}},3192 {{notIntrinsic, notIntrinsic, notIntrinsic,3193 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_im2col_3d,3194 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_im2col_4d,3195 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_and_im2col_5d}}}},3196 // RedTy::OR3197 {{{{notIntrinsic,3198 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_tile_1d,3199 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_tile_2d,3200 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_tile_3d,3201 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_tile_4d,3202 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_tile_5d}},3203 {{notIntrinsic, notIntrinsic, notIntrinsic,3204 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_im2col_3d,3205 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_im2col_4d,3206 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_or_im2col_5d}}}},3207 // RedTy::XOR3208 {{{{notIntrinsic,3209 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_tile_1d,3210 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_tile_2d,3211 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_tile_3d,3212 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_tile_4d,3213 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_tile_5d}},3214 {{notIntrinsic, notIntrinsic, notIntrinsic,3215 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_im2col_3d,3216 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_reduce_xor_im2col_4d,3217 llvm::Intrinsic::3218 nvvm_cp_async_bulk_tensor_reduce_xor_im2col_5d}}}}}};3219 3220 static_assert(getMaxEnumValForTMAReduxKind() == std::size(IDTable) - 1,3221 "TMAReduxKinds must match number of rows in IDTable");3222 3223 size_t redKind = static_cast<size_t>(thisOp.getRedKind());3224 size_t mode = static_cast<size_t>(thisOp.getMode());3225 size_t dim = thisOp.getCoordinates().size();3226 3227 assert(redKind < IDTable.size() &&3228 "Invalid redKind for CpAsyncBulkTensorReduceOp");3229 assert(mode < IDTable[redKind].size() &&3230 "Invalid mode for CpAsyncBulkTensorReduceOp");3231 assert(dim < IDTable[redKind][mode].size() &&3232 "Invalid dim for CpAsyncBulkTensorReduceOp");3233 3234 llvm::Intrinsic::ID intrinsicID = IDTable[redKind][mode][dim];3235 3236 assert(intrinsicID != notIntrinsic &&3237 "Invalid intrinsic for CpAsyncBulkTensorReduceOp.");3238 3239 return {intrinsicID, std::move(args)};3240}3241 3242#define _none3243 3244#define CVT_F2TF32_ID_IMPL(rnd, relu, sf) \3245 hasRelu ? llvm::Intrinsic::nvvm_f2tf32_##rnd##relu##sf \3246 : llvm::Intrinsic::nvvm_f2tf32_##rnd##sf3247 3248#define GET_CVT_F2TF32_ID(rnd, relu, sf) \3249 hasSatFinite ? CVT_F2TF32_ID_IMPL(rnd, relu, sf) \3250 : CVT_F2TF32_ID_IMPL(rnd, relu, )3251 3252llvm::Intrinsic::ID3253ConvertFloatToTF32Op::getIntrinsicID(NVVM::FPRoundingMode rnd,3254 NVVM::SaturationMode sat, bool hasRelu) {3255 using RndMode = NVVM::FPRoundingMode;3256 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);3257 switch (rnd) {3258 case RndMode::RN:3259 return GET_CVT_F2TF32_ID(rn, _relu, _satfinite);3260 case RndMode::RZ:3261 return GET_CVT_F2TF32_ID(rz, _relu, _satfinite);3262 case RndMode::RNA:3263 return GET_CVT_F2TF32_ID(rna, _none, _satfinite);3264 default:3265 llvm_unreachable("Invalid RoundingMode for CvtFloatToTF32Op");3266 }3267}3268 3269NVVM::IDArgPair3270ConvertF32x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF4x2Op op,3271 LLVM::ModuleTranslation &mt,3272 llvm::IRBuilderBase &builder) {3273 llvm::SmallVector<llvm::Value *> args;3274 args.push_back(mt.lookupValue(op.getA()));3275 args.push_back(mt.lookupValue(op.getB()));3276 3277 bool hasRelu = op.getRelu();3278 3279 llvm::Intrinsic::ID intId =3280 hasRelu ? llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_relu_satfinite3281 : llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_satfinite;3282 3283 return {intId, std::move(args)};3284}3285 3286#define GET_F32x2_TO_F6x2_ID(type, has_relu) \3287 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu_satfinite \3288 : llvm::Intrinsic::nvvm_ff_to_##type##_rn_satfinite3289 3290llvm::Intrinsic::ID ConvertF32x2ToF6x2Op::getIntrinsicID(mlir::Type dstTy,3291 bool hasRelu) {3292 return llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(dstTy)3293 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {3294 return GET_F32x2_TO_F6x2_ID(e2m3x2, hasRelu);3295 })3296 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {3297 return GET_F32x2_TO_F6x2_ID(e3m2x2, hasRelu);3298 })3299 .Default([](mlir::Type) {3300 llvm_unreachable("Invalid conversion in ConvertF32x2ToF6x2Op");3301 return llvm::Intrinsic::not_intrinsic;3302 });3303}3304 3305#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf) \3306 has_satf ? llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd##_satfinite \3307 : llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd3308 3309#define GET_F32x2_TO_F8X2_S_ID(type, has_relu) \3310 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu \3311 : llvm::Intrinsic::nvvm_ff_to_##type##_rn3312 3313llvm::Intrinsic::ID3314ConvertF32x2ToF8x2Op::getIntrinsicID(mlir::Type dstTy, NVVM::FPRoundingMode rnd,3315 NVVM::SaturationMode sat, bool hasRelu) {3316 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);3317 bool hasRoundingModeRZ = (rnd == NVVM::FPRoundingMode::RZ);3318 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);3319 3320 return llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(dstTy)3321 .Case<mlir::Float8E4M3FNType>([&](mlir::Float8E4M3FNType) {3322 return GET_F32x2_TO_F8X2_S_ID(e4m3x2, hasRelu);3323 })3324 .Case<mlir::Float8E5M2Type>([&](mlir::Float8E5M2Type) {3325 return GET_F32x2_TO_F8X2_S_ID(e5m2x2, hasRelu);3326 })3327 .Case<mlir::Float8E8M0FNUType>([&](mlir::Float8E8M0FNUType) {3328 if (hasRoundingModeRZ)3329 return GET_F32x2_TO_F8X2_US_ID(rz, hasSatFinite);3330 else if (hasRoundingModeRP)3331 return GET_F32x2_TO_F8X2_US_ID(rp, hasSatFinite);3332 3333 llvm_unreachable("Invalid conversion in ConvertF32x2ToF8x2Op");3334 })3335 .Default([](mlir::Type) {3336 llvm_unreachable("Invalid conversion in ConvertF32x2ToF8x2Op");3337 return llvm::Intrinsic::not_intrinsic;3338 });3339}3340 3341#define GET_F16x2_TO_F8X2_ID(type, has_relu) \3342 has_relu ? llvm::Intrinsic::nvvm_f16x2_to_##type##_rn_relu \3343 : llvm::Intrinsic::nvvm_f16x2_to_##type##_rn3344 3345llvm::Intrinsic::ID ConvertF16x2ToF8x2Op::getIntrinsicID(mlir::Type dstTy,3346 bool hasRelu) {3347 return llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(dstTy)3348 .Case<mlir::Float8E4M3FNType>([&](mlir::Float8E4M3FNType) {3349 return GET_F16x2_TO_F8X2_ID(e4m3x2, hasRelu);3350 })3351 .Case<mlir::Float8E5M2Type>([&](mlir::Float8E5M2Type) {3352 return GET_F16x2_TO_F8X2_ID(e5m2x2, hasRelu);3353 })3354 .Default([](mlir::Type) {3355 llvm_unreachable("Invalid conversion in ConvertF16x2ToF8x2Op");3356 return llvm::Intrinsic::not_intrinsic;3357 });3358}3359 3360#define GET_BF16X2_TO_F8X2_ID(rnd, has_satf) \3361 has_satf ? llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_##rnd##_satfinite \3362 : llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_##rnd3363 3364llvm::Intrinsic::ID3365ConvertBF16x2ToF8x2Op::getIntrinsicID(NVVM::FPRoundingMode rnd,3366 NVVM::SaturationMode sat) {3367 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);3368 switch (rnd) {3369 case NVVM::FPRoundingMode::RZ:3370 return GET_BF16X2_TO_F8X2_ID(rz, hasSatFinite);3371 case NVVM::FPRoundingMode::RP:3372 return GET_BF16X2_TO_F8X2_ID(rp, hasSatFinite);3373 default:3374 llvm_unreachable("Invalid rounding mode for CvtBF16x2ToF8x2Op");3375 }3376}3377 3378NVVM::IDArgPair ConvertF8x2ToF16x2Op::getIntrinsicIDAndArgs(3379 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3380 auto curOp = cast<NVVM::ConvertF8x2ToF16x2Op>(op);3381 3382 bool hasRelu = curOp.getRelu();3383 3384 llvm::Intrinsic::ID intId =3385 llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(curOp.getSrcType())3386 .Case<Float8E4M3FNType>([&](Float8E4M3FNType type) {3387 return hasRelu ? llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn_relu3388 : llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn;3389 })3390 .Case<Float8E5M2Type>([&](Float8E5M2Type type) {3391 return hasRelu ? llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn_relu3392 : llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn;3393 })3394 .Default([](mlir::Type type) {3395 llvm_unreachable("Invalid type for ConvertF8x2ToF16x2Op");3396 return llvm::Intrinsic::not_intrinsic;3397 });3398 3399 llvm::Value *packedI16 =3400 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),3401 llvm::Type::getInt16Ty(builder.getContext()));3402 3403 return {intId, {packedI16}};3404}3405 3406NVVM::IDArgPair ConvertF8x2ToBF16x2Op::getIntrinsicIDAndArgs(3407 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3408 auto curOp = cast<NVVM::ConvertF8x2ToBF16x2Op>(op);3409 3410 llvm::Intrinsic::ID intId = llvm::Intrinsic::nvvm_ue8m0x2_to_bf16x2;3411 llvm::Value *packedI16 =3412 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),3413 llvm::Type::getInt16Ty(builder.getContext()));3414 3415 return {intId, {packedI16}};3416}3417 3418NVVM::IDArgPair ConvertF6x2ToF16x2Op::getIntrinsicIDAndArgs(3419 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3420 auto curOp = cast<NVVM::ConvertF6x2ToF16x2Op>(op);3421 3422 bool hasRelu = curOp.getRelu();3423 3424 llvm::Intrinsic::ID intId =3425 llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(curOp.getSrcType())3426 .Case<Float6E2M3FNType>([&](Float6E2M3FNType type) {3427 return hasRelu ? llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn_relu3428 : llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn;3429 })3430 .Case<Float6E3M2FNType>([&](Float6E3M2FNType type) {3431 return hasRelu ? llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn_relu3432 : llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn;3433 })3434 .Default([](mlir::Type type) {3435 llvm_unreachable("Invalid type for ConvertF6x2ToF16x2Op");3436 return llvm::Intrinsic::not_intrinsic;3437 });3438 3439 llvm::Value *packedI16 =3440 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),3441 llvm::Type::getInt16Ty(builder.getContext()));3442 3443 return {intId, {packedI16}};3444}3445 3446NVVM::IDArgPair ConvertF4x2ToF16x2Op::getIntrinsicIDAndArgs(3447 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3448 auto curOp = cast<NVVM::ConvertF4x2ToF16x2Op>(op);3449 3450 bool hasRelu = curOp.getRelu();3451 3452 llvm::Intrinsic::ID intId =3453 llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(curOp.getSrcType())3454 .Case<Float4E2M1FNType>([&](Float4E2M1FNType type) {3455 return hasRelu ? llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn_relu3456 : llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn;3457 })3458 .Default([](mlir::Type type) {3459 llvm_unreachable("Invalid type for ConvertF4x2ToF16x2Op");3460 return llvm::Intrinsic::not_intrinsic;3461 });3462 3463 llvm::Value *extendedI16 =3464 builder.CreateZExt(mt.lookupValue(curOp.getSrc()),3465 llvm::Type::getInt16Ty(builder.getContext()));3466 3467 return {intId, {extendedI16}};3468}3469 3470llvm::Intrinsic::ID3471Tcgen05AllocOp::getIntrinsicIDAndArgs(Operation &op,3472 LLVM::ModuleTranslation &mt,3473 llvm::SmallVector<llvm::Value *> &args) {3474 auto curOp = cast<NVVM::Tcgen05AllocOp>(op);3475 unsigned as = llvm::cast<LLVM::LLVMPointerType>(curOp.getAddr().getType())3476 .getAddressSpace();3477 bool isShared = as == NVVMMemorySpace::Shared;3478 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;3479 3480 llvm::Intrinsic::ID id;3481 if (isShared) {3482 id = is2CTAMode ? llvm::Intrinsic::nvvm_tcgen05_alloc_shared_cg23483 : llvm::Intrinsic::nvvm_tcgen05_alloc_shared_cg1;3484 } else {3485 id = is2CTAMode ? llvm::Intrinsic::nvvm_tcgen05_alloc_cg23486 : llvm::Intrinsic::nvvm_tcgen05_alloc_cg1;3487 }3488 3489 // Fill the Intrinsic Args3490 args.push_back(mt.lookupValue(curOp.getAddr()));3491 args.push_back(mt.lookupValue(curOp.getNCols()));3492 3493 return id;3494}3495 3496llvm::Intrinsic::ID Tcgen05DeallocOp::getIntrinsicIDAndArgs(3497 Operation &op, LLVM::ModuleTranslation &mt,3498 llvm::SmallVector<llvm::Value *> &args) {3499 auto curOp = cast<NVVM::Tcgen05DeallocOp>(op);3500 auto id = (curOp.getGroup() == CTAGroupKind::CTA_1)3501 ? llvm::Intrinsic::nvvm_tcgen05_dealloc_cg13502 : llvm::Intrinsic::nvvm_tcgen05_dealloc_cg2;3503 3504 // Fill the Intrinsic Args3505 args.push_back(mt.lookupValue(curOp.getTaddr()));3506 args.push_back(mt.lookupValue(curOp.getNCols()));3507 3508 return id;3509}3510 3511#define TCGEN05_COMMIT_IMPL(cg, is_shared, mc) \3512 is_shared ? llvm::Intrinsic::nvvm_tcgen05_commit##mc##_shared##_##cg \3513 : llvm::Intrinsic::nvvm_tcgen05_commit##mc##_##cg3514 3515#define GET_TCGEN05_COMMIT_ID(cta_group, is_shared, has_mc) \3516 has_mc ? TCGEN05_COMMIT_IMPL(cta_group, is_shared, _mc) \3517 : TCGEN05_COMMIT_IMPL(cta_group, is_shared, )3518 3519llvm::Intrinsic::ID3520Tcgen05CommitOp::getIntrinsicIDAndArgs(Operation &op,3521 LLVM::ModuleTranslation &mt,3522 llvm::SmallVector<llvm::Value *> &args) {3523 auto curOp = cast<NVVM::Tcgen05CommitOp>(op);3524 unsigned as = llvm::cast<LLVM::LLVMPointerType>(curOp.getAddr().getType())3525 .getAddressSpace();3526 bool isShared = as == NVVMMemorySpace::Shared;3527 bool hasMulticast = static_cast<bool>(curOp.getMulticastMask());3528 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;3529 3530 llvm::Intrinsic::ID id =3531 is2CTAMode ? GET_TCGEN05_COMMIT_ID(cg2, isShared, hasMulticast)3532 : GET_TCGEN05_COMMIT_ID(cg1, isShared, hasMulticast);3533 3534 // Fill the Intrinsic Args3535 args.push_back(mt.lookupValue(curOp.getAddr()));3536 if (hasMulticast)3537 args.push_back(mt.lookupValue(curOp.getMulticastMask()));3538 3539 return id;3540}3541 3542#define TCGEN05_CP_IMPL(shape_mc, src_fmt, cg) \3543 llvm::Intrinsic::nvvm_tcgen05_cp##shape_mc##src_fmt##cg3544 3545#define TCGEN05_CP_2CTA(shape_mc, src_fmt, is_2cta) \3546 is_2cta ? TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg2) \3547 : TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg1)3548 3549#define GET_TCGEN05_CP_ID(shape_mc, src_fmt, is_2cta) \3550 [&]() -> auto { \3551 if ((src_fmt) == Tcgen05CpSrcFormat::B6x16_P32) \3552 return TCGEN05_CP_2CTA(shape_mc, _b6x16_p32, is_2cta); \3553 if ((src_fmt) == Tcgen05CpSrcFormat::B4x16_P64) \3554 return TCGEN05_CP_2CTA(shape_mc, _b4x16_p64, is_2cta); \3555 return TCGEN05_CP_2CTA(shape_mc, , is_2cta); \3556 }()3557 3558NVVM::IDArgPair3559ConvertF32x2ToF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF16x2Op &op,3560 LLVM::ModuleTranslation &mt,3561 llvm::IRBuilderBase &builder) {3562 static constexpr llvm::Intrinsic::ID rndRNIds[] = {3563 llvm::Intrinsic::nvvm_ff2f16x2_rn,3564 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu,3565 llvm::Intrinsic::nvvm_ff2f16x2_rn_satfinite,3566 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu_satfinite,3567 };3568 static constexpr llvm::Intrinsic::ID rndRZIds[] = {3569 llvm::Intrinsic::nvvm_ff2f16x2_rz,3570 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu,3571 llvm::Intrinsic::nvvm_ff2f16x2_rz_satfinite,3572 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu_satfinite,3573 };3574 static constexpr llvm::Intrinsic::ID rndRSIds[] = {3575 llvm::Intrinsic::nvvm_ff2f16x2_rs,3576 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu,3577 llvm::Intrinsic::nvvm_ff2f16x2_rs_satfinite,3578 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu_satfinite,3579 };3580 3581 unsigned hasRelu = op.getRelu() ? 1 : 0;3582 unsigned hasSatFinite =3583 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;3584 // idx: bit-0 - relu3585 // bit-1 - satfinite3586 unsigned idx = (hasSatFinite << 1) | hasRelu;3587 3588 llvm::SmallVector<llvm::Value *> args;3589 args.push_back(mt.lookupValue(op.getSrcHi()));3590 args.push_back(mt.lookupValue(op.getSrcLo()));3591 if (op.getRandomBits())3592 args.push_back(mt.lookupValue(op.getRandomBits()));3593 3594 switch (op.getRnd()) {3595 case FPRoundingMode::RN:3596 return {rndRNIds[idx], std::move(args)};3597 case FPRoundingMode::RZ:3598 return {rndRZIds[idx], std::move(args)};3599 case FPRoundingMode::RS:3600 return {rndRSIds[idx], std::move(args)};3601 default:3602 llvm_unreachable("Invalid rounding mode for ConvertF32x2ToF16x2Op");3603 }3604}3605 3606NVVM::IDArgPair3607ConvertF32x2ToBF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToBF16x2Op &op,3608 LLVM::ModuleTranslation &mt,3609 llvm::IRBuilderBase &builder) {3610 static constexpr llvm::Intrinsic::ID rndRNIds[] = {3611 llvm::Intrinsic::nvvm_ff2bf16x2_rn,3612 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu,3613 llvm::Intrinsic::nvvm_ff2bf16x2_rn_satfinite,3614 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu_satfinite,3615 };3616 static constexpr llvm::Intrinsic::ID rndRZIds[] = {3617 llvm::Intrinsic::nvvm_ff2bf16x2_rz,3618 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu,3619 llvm::Intrinsic::nvvm_ff2bf16x2_rz_satfinite,3620 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu_satfinite,3621 };3622 static constexpr llvm::Intrinsic::ID rndRSIds[] = {3623 llvm::Intrinsic::nvvm_ff2bf16x2_rs,3624 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu,3625 llvm::Intrinsic::nvvm_ff2bf16x2_rs_satfinite,3626 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu_satfinite,3627 };3628 3629 unsigned hasRelu = op.getRelu() ? 1 : 0;3630 unsigned hasSatFinite =3631 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;3632 // idx: bit-0 - relu3633 // bit-1 - satfinite3634 unsigned idx = (hasSatFinite << 1) | hasRelu;3635 3636 llvm::SmallVector<llvm::Value *> args;3637 args.push_back(mt.lookupValue(op.getSrcHi()));3638 args.push_back(mt.lookupValue(op.getSrcLo()));3639 if (op.getRandomBits())3640 args.push_back(mt.lookupValue(op.getRandomBits()));3641 3642 switch (op.getRnd()) {3643 case FPRoundingMode::RN:3644 return {rndRNIds[idx], std::move(args)};3645 case FPRoundingMode::RZ:3646 return {rndRZIds[idx], std::move(args)};3647 case FPRoundingMode::RS:3648 return {rndRSIds[idx], std::move(args)};3649 default:3650 llvm_unreachable("Invalid rounding mode for ConvertF32x2ToBF16x2Op");3651 }3652}3653 3654llvm::Intrinsic::ID ConvertF32x4ToF8x4Op::getIntrinsicID() {3655 mlir::Type dstTy = getDstTy();3656 bool hasRelu = getRelu();3657 3658 return llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(dstTy)3659 .Case<mlir::Float8E4M3FNType>([&](mlir::Float8E4M3FNType) {3660 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_relu_satfinite3661 : llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_satfinite;3662 })3663 .Case<mlir::Float8E5M2Type>([&](mlir::Float8E5M2Type) {3664 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_relu_satfinite3665 : llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_satfinite;3666 })3667 .Default([](mlir::Type) {3668 llvm_unreachable("Invalid F8 type in ConvertF32x4ToF8x4Op");3669 return llvm::Intrinsic::not_intrinsic;3670 });3671}3672 3673llvm::Intrinsic::ID ConvertF32x4ToF6x4Op::getIntrinsicID() {3674 mlir::Type dstTy = getDstTy();3675 bool hasRelu = getRelu();3676 3677 return llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(dstTy)3678 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {3679 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_relu_satfinite3680 : llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_satfinite;3681 })3682 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {3683 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_relu_satfinite3684 : llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_satfinite;3685 })3686 .Default([](mlir::Type) {3687 llvm_unreachable("Invalid F6 type in ConvertF32x4ToF6x4Op");3688 return llvm::Intrinsic::not_intrinsic;3689 });3690}3691 3692llvm::Intrinsic::ID ConvertF32x4ToF4x4Op::getIntrinsicID() {3693 mlir::Type dstTy = getDstTy();3694 bool hasRelu = getRelu();3695 3696 return llvm::TypeSwitch<mlir::Type, llvm::Intrinsic::ID>(dstTy)3697 .Case<mlir::Float4E2M1FNType>([&](mlir::Float4E2M1FNType) {3698 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_relu_satfinite3699 : llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_satfinite;3700 })3701 .Default([](mlir::Type) {3702 llvm_unreachable("Invalid F4 type in ConvertF32x4ToF4x4Op");3703 return llvm::Intrinsic::not_intrinsic;3704 });3705}3706 3707llvm::Intrinsic::ID Tcgen05CpOp::getIntrinsicID(Operation &op) {3708 auto curOp = cast<NVVM::Tcgen05CpOp>(op);3709 bool is2CTA = curOp.getGroup() == CTAGroupKind::CTA_2;3710 auto srcFmt = curOp.getSrcFormat();3711 auto mc = curOp.getMulticast();3712 3713 switch (curOp.getShape()) {3714 case Tcgen05CpShape::SHAPE_128x256b:3715 return GET_TCGEN05_CP_ID(_128x256b, srcFmt, is2CTA);3716 case Tcgen05CpShape::SHAPE_128x128b:3717 return GET_TCGEN05_CP_ID(_128x128b, srcFmt, is2CTA);3718 case Tcgen05CpShape::SHAPE_4x256b:3719 return GET_TCGEN05_CP_ID(_4x256b, srcFmt, is2CTA);3720 case Tcgen05CpShape::SHAPE_32x128b:3721 return GET_TCGEN05_CP_ID(_32x128b_warpx4, srcFmt, is2CTA);3722 case Tcgen05CpShape::SHAPE_64x128b:3723 return (mc == Tcgen05CpMulticast::WARPX2_01_23)3724 ? GET_TCGEN05_CP_ID(_64x128b_warpx2_01_23, srcFmt, is2CTA)3725 : GET_TCGEN05_CP_ID(_64x128b_warpx2_02_13, srcFmt, is2CTA);3726 }3727 llvm_unreachable("Invalid shape in tcgen05 cp Op");3728}3729 3730// Returns the valid vector length for a given shape and vector length, the3731// function models the table mentioned in the tcgen05.{ld, st} Op description3732static unsigned isValidVectorLength(NVVM::Tcgen05LdStShape shape,3733 unsigned vecLen) {3734 if (shape == NVVM::Tcgen05LdStShape::SHAPE_16X128B)3735 return vecLen >= 2;3736 if (shape == NVVM::Tcgen05LdStShape::SHAPE_16X256B)3737 return vecLen >= 4;3738 return true;3739}3740 3741LogicalResult Tcgen05LdOp::verify() {3742 LogicalResult result = success();3743 if (getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())3744 result = emitError("shape 16x32bx2 requires offset argument");3745 3746 if (getShape() != NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && getOffset())3747 result = emitError("offset argument is only supported for shape 16x32bx2");3748 3749 auto resTy = getRes().getType();3750 unsigned resLen = isa<VectorType>(resTy)3751 ? llvm::cast<VectorType>(resTy).getNumElements()3752 : 1;3753 if (!isValidVectorLength(getShape(), resLen))3754 result = emitError(llvm::formatv("invalid result type length {0} for shape "3755 "{1} in tcgen05.ld Op",3756 resLen, stringifyEnum(getShape())));3757 3758 return result;3759}3760 3761LogicalResult Tcgen05StOp::verify() {3762 LogicalResult result = success();3763 if (getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())3764 result = emitError("shape 16x32bx2 requires offset argument");3765 3766 auto valTy = getVal().getType();3767 unsigned valLen = isa<VectorType>(valTy)3768 ? llvm::cast<VectorType>(valTy).getNumElements()3769 : 1;3770 if (!isValidVectorLength(getShape(), valLen))3771 result = emitError(llvm::formatv("invalid input length {0} for shape "3772 "{1} in tcgen05.st Op",3773 valLen, stringifyEnum(getShape())));3774 3775 return result;3776}3777 3778/// Infer the result ranges for the NVVM SpecialRangeableRegisterOp that might3779/// have ConstantRangeAttr.3780static void nvvmInferResultRanges(Operation *op, Value result,3781 ArrayRef<::mlir::ConstantIntRanges> argRanges,3782 SetIntRangeFn setResultRanges) {3783 if (auto rangeAttr = op->getAttrOfType<LLVM::ConstantRangeAttr>("range")) {3784 setResultRanges(result, {rangeAttr.getLower(), rangeAttr.getUpper(),3785 rangeAttr.getLower(), rangeAttr.getUpper()});3786 }3787}3788 3789/// Verify the range attribute satisfies LLVM ConstantRange constructor3790/// requirements for NVVM SpecialRangeableRegisterOp.3791static LogicalResult3792verifyConstantRangeAttr(Operation *op,3793 std::optional<LLVM::ConstantRangeAttr> rangeAttr) {3794 if (!rangeAttr)3795 return success();3796 3797 const llvm::APInt &lower = rangeAttr->getLower();3798 const llvm::APInt &upper = rangeAttr->getUpper();3799 3800 // Check LLVM ConstantRange constructor condition3801 if (lower == upper && !lower.isMaxValue() && !lower.isMinValue()) {3802 unsigned bitWidth = lower.getBitWidth();3803 llvm::APInt minVal = llvm::APInt::getMinValue(bitWidth);3804 llvm::APInt maxVal = llvm::APInt::getMaxValue(bitWidth);3805 return op->emitOpError(3806 "invalid range attribute: Lower == Upper, but they aren't min (")3807 << llvm::toString(minVal, 10, false) << ") or max ("3808 << llvm::toString(maxVal, 10, false)3809 << ") value! This is an invalid constant range.";3810 }3811 3812 return success();3813}3814 3815static llvm::Value *getAsPackedI32(llvm::Value *arg,3816 llvm::IRBuilderBase &builder) {3817 return builder.CreateBitCast(arg,3818 llvm::Type::getInt32Ty(builder.getContext()));3819}3820 3821NVVM::IDArgPair DotAccumulate4WayOp::getIntrinsicIDAndArgs(3822 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3823 auto curOp = cast<NVVM::DotAccumulate4WayOp>(op);3824 3825 llvm::SmallVector<llvm::Value *> args;3826 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getA()), builder));3827 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getB()), builder));3828 args.push_back(mt.lookupValue(curOp.getC()));3829 3830 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;3831 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;3832 unsigned type = (isASigned << 1) | isBSigned;3833 const llvm::Intrinsic::ID ids[] = {3834 llvm::Intrinsic::nvvm_idp4a_u_u,3835 llvm::Intrinsic::nvvm_idp4a_u_s,3836 llvm::Intrinsic::nvvm_idp4a_s_u,3837 llvm::Intrinsic::nvvm_idp4a_s_s,3838 };3839 return {ids[type], args};3840}3841 3842NVVM::IDArgPair DotAccumulate2WayOp::getIntrinsicIDAndArgs(3843 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3844 auto curOp = cast<NVVM::DotAccumulate2WayOp>(op);3845 3846 llvm::SmallVector<llvm::Value *> args;3847 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getA()), builder));3848 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getB()), builder));3849 args.push_back(builder.getInt1(curOp.getBHi()));3850 args.push_back(mt.lookupValue(curOp.getC()));3851 3852 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;3853 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;3854 unsigned type = (isASigned << 1) | isBSigned;3855 const llvm::Intrinsic::ID ids[] = {3856 llvm::Intrinsic::nvvm_idp2a_u_u,3857 llvm::Intrinsic::nvvm_idp2a_u_s,3858 llvm::Intrinsic::nvvm_idp2a_s_u,3859 llvm::Intrinsic::nvvm_idp2a_s_s,3860 };3861 return {ids[type], args};3862}3863 3864static llvm::Value *getParamCastedAddr(llvm::Value *addr,3865 llvm::IRBuilderBase &builder) {3866 return builder.CreateAddrSpaceCast(3867 addr,3868 llvm::PointerType::get(builder.getContext(),3869 llvm::NVPTXAS::AddressSpace::ADDRESS_SPACE_PARAM));3870}3871 3872NVVM::IDArgPair3873PrefetchOp::getIntrinsicIDAndArgs(NVVM::PrefetchOp &op,3874 LLVM::ModuleTranslation &mt,3875 llvm::IRBuilderBase &builder) {3876 using MemSpace = NVVM::NVVMMemorySpace;3877 using CacheLevel = NVVM::PrefetchCacheLevel;3878 3879 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = op.getCacheLevel();3880 std::optional<NVVM::CacheEvictionPriority> evictPriority =3881 op.getEvictPriority();3882 unsigned addressSpace =3883 llvm::cast<LLVM::LLVMPointerType>(op.getAddr().getType())3884 .getAddressSpace();3885 3886 llvm::SmallVector<llvm::Value *> args;3887 llvm::Value *addr = mt.lookupValue(op.getAddr());3888 args.push_back(op.getInParamSpace() ? getParamCastedAddr(addr, builder)3889 : addr);3890 3891 if (op.getTensormap())3892 return {llvm::Intrinsic::nvvm_prefetch_tensormap, args};3893 3894 assert(cacheLevel && "expected cache level for non-tensormap prefetch");3895 3896 if (op.getUniform() && *cacheLevel == CacheLevel::L1)3897 return {llvm::Intrinsic::nvvm_prefetchu_L1, args};3898 3899 if (evictPriority && *cacheLevel == CacheLevel::L2) {3900 switch (*evictPriority) {3901 case NVVM::CacheEvictionPriority::EvictLast:3902 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_last, args};3903 case NVVM::CacheEvictionPriority::EvictNormal:3904 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_normal, args};3905 default:3906 llvm_unreachable("Invalid cache eviction priority");3907 }3908 }3909 3910 switch (static_cast<MemSpace>(addressSpace)) {3911 case MemSpace::Generic:3912 return *cacheLevel == CacheLevel::L13913 ? NVVM::IDArgPair({llvm::Intrinsic::nvvm_prefetch_L1, args})3914 : NVVM::IDArgPair({llvm::Intrinsic::nvvm_prefetch_L2, args});3915 case MemSpace::Global:3916 return *cacheLevel == CacheLevel::L13917 ? NVVM::IDArgPair(3918 {llvm::Intrinsic::nvvm_prefetch_global_L1, args})3919 : NVVM::IDArgPair(3920 {llvm::Intrinsic::nvvm_prefetch_global_L2, args});3921 case MemSpace::Local:3922 return *cacheLevel == CacheLevel::L13923 ? NVVM::IDArgPair(3924 {llvm::Intrinsic::nvvm_prefetch_local_L1, args})3925 : NVVM::IDArgPair(3926 {llvm::Intrinsic::nvvm_prefetch_local_L2, args});3927 default:3928 llvm_unreachable("Invalid pointer address space");3929 }3930}3931 3932bool NVVM::InlinePtxOp::getAsmValues(3933 RewriterBase &rewriter,3934 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>3935 &asmValues) {3936 for (auto arg : getReadWriteArgs())3937 asmValues.push_back({arg, mlir::NVVM::PTXRegisterMod::ReadWrite});3938 for (auto arg : getResults())3939 asmValues.push_back({arg, mlir::NVVM::PTXRegisterMod::Write});3940 for (auto arg : getReadOnlyArgs())3941 asmValues.push_back({arg, mlir::NVVM::PTXRegisterMod::Read});3942 if (getPredicate())3943 asmValues.push_back({getPredicate(), mlir::NVVM::PTXRegisterMod::Read});3944 return false; // No manual mapping needed3945}3946 3947NVVM::IDArgPair ClusterLaunchControlTryCancelOp::getIntrinsicIDAndArgs(3948 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3949 auto curOp = cast<NVVM::ClusterLaunchControlTryCancelOp>(op);3950 llvm::SmallVector<llvm::Value *> args;3951 args.push_back(mt.lookupValue(curOp.getSmemAddress()));3952 args.push_back(mt.lookupValue(curOp.getMbarrier()));3953 3954 llvm::Intrinsic::ID intrinsicID =3955 curOp.getMulticast()3956 ? llvm::Intrinsic::3957 nvvm_clusterlaunchcontrol_try_cancel_async_multicast_shared3958 : llvm::Intrinsic::nvvm_clusterlaunchcontrol_try_cancel_async_shared;3959 3960 return {intrinsicID, args};3961}3962 3963NVVM::IDArgPair ClusterLaunchControlQueryCancelOp::getIntrinsicIDAndArgs(3964 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {3965 auto curOp = cast<NVVM::ClusterLaunchControlQueryCancelOp>(op);3966 llvm::SmallVector<llvm::Value *> args;3967 args.push_back(mt.lookupValue(curOp.getTryCancelResponse()));3968 3969 llvm::Intrinsic::ID intrinsicID;3970 3971 switch (curOp.getQueryType()) {3972 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:3973 intrinsicID =3974 llvm::Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_is_canceled;3975 break;3976 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:3977 intrinsicID = llvm::Intrinsic::3978 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_x;3979 break;3980 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:3981 intrinsicID = llvm::Intrinsic::3982 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_y;3983 break;3984 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:3985 intrinsicID = llvm::Intrinsic::3986 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_z;3987 break;3988 }3989 return {intrinsicID, args};3990}3991 3992mlir::NVVM::IDArgPair3993PermuteOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,3994 llvm::IRBuilderBase &builder) {3995 auto thisOp = cast<NVVM::PermuteOp>(op);3996 NVVM::PermuteMode mode = thisOp.getMode();3997 3998 static constexpr llvm::Intrinsic::ID IDs[] = {3999 llvm::Intrinsic::nvvm_prmt, llvm::Intrinsic::nvvm_prmt_f4e,4000 llvm::Intrinsic::nvvm_prmt_b4e, llvm::Intrinsic::nvvm_prmt_rc8,4001 llvm::Intrinsic::nvvm_prmt_ecl, llvm::Intrinsic::nvvm_prmt_ecr,4002 llvm::Intrinsic::nvvm_prmt_rc16};4003 4004 unsigned modeIndex = static_cast<unsigned>(mode);4005 llvm::SmallVector<llvm::Value *> args;4006 args.push_back(mt.lookupValue(thisOp.getLo()));4007 4008 // Only first 3 modes (Default, f4e, b4e) need the hi operand.4009 if (modeIndex < 3)4010 args.push_back(mt.lookupValue(thisOp.getHi()));4011 4012 args.push_back(mt.lookupValue(thisOp.getSelector()));4013 4014 return {IDs[modeIndex], args};4015}4016 4017//===----------------------------------------------------------------------===//4018// NVVM tcgen05.mma functions4019//===----------------------------------------------------------------------===//4020 4021mlir::NVVM::IDArgPair4022Tcgen05MMAOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,4023 llvm::IRBuilderBase &builder) {4024 4025 auto thisOp = cast<NVVM::Tcgen05MMAOp>(op);4026 llvm::SmallVector<llvm::Value *> args;4027 4028 args.push_back(mt.lookupValue(thisOp.getMatrixD()));4029 4030 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());4031 const bool isATensor = isa<llvm::PointerType>(A->getType());4032 args.push_back(A);4033 4034 args.push_back(mt.lookupValue(thisOp.getMatrixB()));4035 args.push_back(mt.lookupValue(thisOp.getIdesc()));4036 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));4037 4038 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;4039 using CtaGroupArray = std::array<EnableAShiftArray, 2>;4040 using IsATensorArray = std::array<CtaGroupArray, 2>;4041 using HasScaleInputDArray = std::array<IsATensorArray, 2>;4042 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;4043 4044 // [hasDisableOutputLane][hasScaleInputD][isATensor][CtaGroup][EnableAShift]4045 static constexpr HasDisableOutputLaneArray tcgen05MMAIDs = {4046 { // without diable output lane4047 {{// without scale input D4048 {{4049 // shared4050 {{// cg14051 {llvm::Intrinsic::nvvm_tcgen05_mma_shared, notIntrinsic},4052 // cg24053 {llvm::Intrinsic::nvvm_tcgen05_mma_shared, notIntrinsic}}},4054 {{// tensor4055 {4056 // cg14057 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,4058 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,4059 },4060 {4061 // cg24062 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,4063 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,4064 }}},4065 }},4066 // with scale input D4067 {{ // shared4068 {{// cg14069 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d, notIntrinsic},4070 // cg24071 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d, notIntrinsic}}},4072 {{// tensor4073 {4074 // cg14075 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,4076 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,4077 },4078 {4079 // cg24080 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,4081 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,4082 }}}}}}},4083 // with disable output lane4084 {{ // without scale input D4085 {{ // shared4086 {{// cg14087 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1,4088 notIntrinsic},4089 // cg24090 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2,4091 notIntrinsic}}},4092 {{// cg14093 {4094 llvm::Intrinsic::4095 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1,4096 llvm::Intrinsic::4097 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift,4098 },4099 // cg24100 {4101 llvm::Intrinsic::4102 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2,4103 llvm::Intrinsic::4104 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift,4105 }}}}},4106 // with scale input D4107 {{ // shared4108 {{// cg14109 {llvm::Intrinsic::4110 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1,4111 notIntrinsic},4112 // cg24113 {llvm::Intrinsic::4114 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2,4115 notIntrinsic}}},4116 // tensor4117 {{// cg14118 {llvm::Intrinsic::4119 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1,4120 llvm::Intrinsic::4121 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift},4122 // cg24123 {4124 llvm::Intrinsic::4125 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2,4126 llvm::Intrinsic::4127 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift,4128 }}}}}}}}};4129 4130 llvm::Value *ScaleInputD = mt.lookupValue(thisOp.getScaleInputD());4131 bool hasScaleInputD = ScaleInputD != nullptr;4132 4133 llvm::Value *DisableOutputLane =4134 mt.lookupValue(thisOp.getDisableOutputLane());4135 bool hasDisableOutputLane = DisableOutputLane != nullptr;4136 4137 const unsigned ctaGroup =4138 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()));4139 4140 llvm::Intrinsic::ID ID =4141 tcgen05MMAIDs[hasDisableOutputLane][hasScaleInputD][isATensor]4142 [ctaGroup - 1][thisOp.getAShift()];4143 4144 assert(ID != notIntrinsic && "Invalid intrinsic for Tcgen05MMAOp.");4145 4146 if (hasScaleInputD)4147 args.push_back(ScaleInputD);4148 4149 if (hasDisableOutputLane)4150 args.push_back(DisableOutputLane);4151 4152 args.push_back(builder.getInt32(static_cast<unsigned>(thisOp.getKind())));4153 4154 if (!hasDisableOutputLane)4155 args.push_back(builder.getInt32(ctaGroup));4156 4157 args.push_back(4158 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));4159 4160 return {ID, args};4161}4162 4163static LogicalResult4164verifyTcgen05MMAOp(bool isATensor, mlir::Value disableOutputLane,4165 NVVM::CTAGroupKind ctaGroup, bool hasAShift,4166 NVVM::Tcgen05MMACollectorOp collectorOp, Location loc) {4167 4168 if (disableOutputLane) {4169 mlir::VectorType disableOutputLaneType =4170 cast<mlir::VectorType>(disableOutputLane.getType());4171 if ((ctaGroup == NVVM::CTAGroupKind::CTA_1 &&4172 disableOutputLaneType.getNumElements() != 4) ||4173 (ctaGroup == NVVM::CTAGroupKind::CTA_2 &&4174 disableOutputLaneType.getNumElements() != 8))4175 return emitError(loc) << "Disable Output Lane of length "4176 << disableOutputLaneType.getNumElements()4177 << " is incompatible with CtaGroupAttr";4178 }4179 4180 if (hasAShift && !isATensor)4181 return emitError(4182 loc, "A-shift can be applied only when matrix A is in tensor memory");4183 4184 if (hasAShift == true && (collectorOp == Tcgen05MMACollectorOp::FILL ||4185 collectorOp == Tcgen05MMACollectorOp::USE))4186 return emitError(4187 loc, "Cannot use collector buffer operation fill or use with ashift");4188 4189 return success();4190}4191 4192LogicalResult Tcgen05MMAOp::verify() {4193 return verifyTcgen05MMAOp(isa<LLVM::LLVMPointerType>(getMatrixA().getType()),4194 getDisableOutputLane(), getCtaGroup(), getAShift(),4195 getCollectorOp(), getLoc());4196}4197 4198//===----------------------------------------------------------------------===//4199// NVVM tcgen05.mma.sp functions4200//===----------------------------------------------------------------------===//4201 4202mlir::NVVM::IDArgPair Tcgen05MMASparseOp::getIntrinsicIDAndArgs(4203 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {4204 4205 auto thisOp = cast<NVVM::Tcgen05MMASparseOp>(op);4206 llvm::SmallVector<llvm::Value *> args;4207 4208 args.push_back(mt.lookupValue(thisOp.getMatrixD()));4209 4210 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());4211 bool isATensor = isa<llvm::PointerType>(A->getType());4212 args.push_back(A);4213 4214 args.push_back(mt.lookupValue(thisOp.getMatrixB()));4215 args.push_back(mt.lookupValue(thisOp.getIdesc()));4216 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));4217 args.push_back(mt.lookupValue(thisOp.getSparseMetadata()));4218 4219 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;4220 using CtaGroupArray = std::array<EnableAShiftArray, 2>;4221 using IsATensorArray = std::array<CtaGroupArray, 2>;4222 using HasScaleInputDArray = std::array<IsATensorArray, 2>;4223 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;4224 4225 // [hasDisableOutputLane][hasScaleInputD][isATensor][CtaGroup][EnableAShift]4226 static constexpr HasDisableOutputLaneArray tcgen05MMASparseIDs = {4227 { // without diable output lane4228 {{// without scale input D4229 {{4230 // shared4231 {{// cg14232 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared, notIntrinsic},4233 // cg24234 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared, notIntrinsic}}},4235 {{// tensor4236 {4237 // cg14238 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,4239 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,4240 },4241 {4242 // cg24243 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,4244 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,4245 }}},4246 }},4247 // with scale input D4248 {{ // shared4249 {{// cg14250 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,4251 notIntrinsic},4252 // cg24253 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,4254 notIntrinsic}}},4255 {{// tensor4256 {4257 // cg14258 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,4259 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,4260 },4261 {4262 // cg24263 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,4264 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,4265 }}}}}}},4266 // with disable output lane4267 {{ // without scale input D4268 {{ // shared4269 {{// cg14270 {llvm::Intrinsic::4271 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1,4272 notIntrinsic},4273 // cg24274 {llvm::Intrinsic::4275 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2,4276 notIntrinsic}}},4277 {{// cg14278 {4279 llvm::Intrinsic::4280 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1,4281 llvm::Intrinsic::4282 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift,4283 },4284 // cg24285 {4286 llvm::Intrinsic::4287 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2,4288 llvm::Intrinsic::4289 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift,4290 }}}}},4291 // with scale input D4292 {{ // shared4293 {{// cg14294 {llvm::Intrinsic::4295 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1,4296 notIntrinsic},4297 // cg24298 {llvm::Intrinsic::4299 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2,4300 notIntrinsic}}},4301 // tensor4302 {{// cg14303 {llvm::Intrinsic::4304 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1,4305 llvm::Intrinsic::4306 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift},4307 // cg24308 {4309 llvm::Intrinsic::4310 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2,4311 llvm::Intrinsic::4312 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift,4313 }}}}}}}}};4314 4315 llvm::Value *ScaleInputD = mt.lookupValue(thisOp.getScaleInputD());4316 bool hasScaleInputD = ScaleInputD != nullptr;4317 4318 llvm::Value *DisableOutputLane =4319 mt.lookupValue(thisOp.getDisableOutputLane());4320 bool hasDisableOutputLane = DisableOutputLane != nullptr;4321 4322 unsigned ctaGroup =4323 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()));4324 4325 llvm::Intrinsic::ID ID =4326 tcgen05MMASparseIDs[hasDisableOutputLane][hasScaleInputD][isATensor]4327 [ctaGroup - 1][thisOp.getAShift()];4328 4329 assert(ID != notIntrinsic && "Invalid intrinsic for Tcgen05MMASparseOp.");4330 4331 if (hasScaleInputD)4332 args.push_back(ScaleInputD);4333 4334 if (hasDisableOutputLane)4335 args.push_back(DisableOutputLane);4336 4337 args.push_back(builder.getInt32(static_cast<unsigned>(thisOp.getKind())));4338 4339 if (!hasDisableOutputLane)4340 args.push_back(builder.getInt32(ctaGroup));4341 4342 args.push_back(4343 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));4344 4345 return {ID, args};4346}4347 4348LogicalResult Tcgen05MMASparseOp::verify() {4349 return verifyTcgen05MMAOp(isa<LLVM::LLVMPointerType>(getMatrixA().getType()),4350 getDisableOutputLane(), getCtaGroup(), getAShift(),4351 getCollectorOp(), getLoc());4352}4353 4354//===----------------------------------------------------------------------===//4355// NVVM tcgen05.mma.block_scale functions4356//===----------------------------------------------------------------------===//4357 4358mlir::NVVM::IDArgPair Tcgen05MMABlockScaleOp::getIntrinsicIDAndArgs(4359 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {4360 4361 auto thisOp = cast<NVVM::Tcgen05MMABlockScaleOp>(op);4362 llvm::SmallVector<llvm::Value *> args;4363 4364 args.push_back(mt.lookupValue(thisOp.getMatrixD()));4365 4366 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());4367 bool isATensor = isa<llvm::PointerType>(A->getType());4368 args.push_back(A);4369 4370 args.push_back(mt.lookupValue(thisOp.getMatrixB()));4371 args.push_back(mt.lookupValue(thisOp.getIdesc()));4372 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));4373 args.push_back(mt.lookupValue(thisOp.getScaleA()));4374 args.push_back(mt.lookupValue(thisOp.getScaleB()));4375 args.push_back(builder.getInt32(4376 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()))));4377 args.push_back(4378 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));4379 4380 auto kind = thisOp.getKind();4381 auto blockScale = thisOp.getBlockScale();4382 llvm::Intrinsic::ID ID = [&]() {4383 if (kind == NVVM::Tcgen05MMABlockScaleKind::MXF8F6F4) {4384 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {4385 return isATensor ? llvm::Intrinsic::4386 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale4387 : llvm::Intrinsic::4388 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale;4389 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {4390 return isATensor4391 ? llvm::Intrinsic::4392 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale_block324393 : llvm::Intrinsic::4394 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale_block32;4395 }4396 } else if (kind == NVVM::Tcgen05MMABlockScaleKind::MXF4) {4397 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {4398 return isATensor4399 ? llvm::Intrinsic::nvvm_tcgen05_mma_tensor_mxf4_block_scale4400 : llvm::Intrinsic::nvvm_tcgen05_mma_shared_mxf4_block_scale;4401 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {4402 return isATensor ? llvm::Intrinsic::4403 nvvm_tcgen05_mma_tensor_mxf4_block_scale_block324404 : llvm::Intrinsic::4405 nvvm_tcgen05_mma_shared_mxf4_block_scale_block32;4406 }4407 } else if (kind == NVVM::Tcgen05MMABlockScaleKind::MXF4NVF4) {4408 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {4409 return isATensor4410 ? llvm::Intrinsic::4411 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block324412 : llvm::Intrinsic::4413 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block32;4414 4415 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {4416 return isATensor4417 ? llvm::Intrinsic::4418 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block164419 : llvm::Intrinsic::4420 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block16;4421 }4422 }4423 llvm_unreachable("Invalid tcgen05.mma.block_scale attributes");4424 }();4425 4426 return {ID, args};4427}4428 4429static LogicalResult4430verifyTcgen05MMABlockScaleOp(NVVM::Tcgen05MMACollectorOp collectorOp,4431 NVVM::Tcgen05MMABlockScaleKind kind,4432 NVVM::Tcgen05MMABlockScale blockScale,4433 Location loc) {4434 4435 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT &&4436 kind == Tcgen05MMABlockScaleKind::MXF4NVF4)4437 return emitError(loc, "mxf4nvf4 requires block scale attribute");4438 4439 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16 &&4440 kind != Tcgen05MMABlockScaleKind::MXF4NVF4)4441 return emitError(loc,4442 llvm::formatv("{} kind does not support block16 attribute",4443 stringifyEnum(kind)));4444 4445 return success();4446}4447 4448LogicalResult Tcgen05MMABlockScaleOp::verify() {4449 return verifyTcgen05MMABlockScaleOp(getCollectorOp(), getKind(),4450 getBlockScale(), getLoc());4451}4452 4453//===----------------------------------------------------------------------===//4454// NVVM tcgen05.mma.sp.block_scale functions4455//===----------------------------------------------------------------------===//4456 4457mlir::NVVM::IDArgPair Tcgen05MMASparseBlockScaleOp::getIntrinsicIDAndArgs(4458 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {4459 4460 auto thisOp = cast<NVVM::Tcgen05MMASparseBlockScaleOp>(op);4461 llvm::SmallVector<llvm::Value *> args;4462 4463 args.push_back(mt.lookupValue(thisOp.getMatrixD()));4464 4465 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());4466 bool isATensor = isa<llvm::PointerType>(A->getType());4467 args.push_back(A);4468 4469 args.push_back(mt.lookupValue(thisOp.getMatrixB()));4470 args.push_back(mt.lookupValue(thisOp.getIdesc()));4471 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));4472 args.push_back(mt.lookupValue(thisOp.getSparseMetadata()));4473 args.push_back(mt.lookupValue(thisOp.getScaleA()));4474 args.push_back(mt.lookupValue(thisOp.getScaleB()));4475 args.push_back(builder.getInt32(4476 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()))));4477 args.push_back(4478 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));4479 4480 auto kind = thisOp.getKind();4481 auto blockScale = thisOp.getBlockScale();4482 llvm::Intrinsic::ID ID = [&]() {4483 if (kind == NVVM::Tcgen05MMABlockScaleKind::MXF8F6F4) {4484 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {4485 return isATensor ? llvm::Intrinsic::4486 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale4487 : llvm::Intrinsic::4488 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale;4489 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {4490 return isATensor4491 ? llvm::Intrinsic::4492 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale_block324493 : llvm::Intrinsic::4494 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale_block32;4495 }4496 } else if (kind == NVVM::Tcgen05MMABlockScaleKind::MXF4) {4497 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {4498 return isATensor ? llvm::Intrinsic::4499 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale4500 : llvm::Intrinsic::4501 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale;4502 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {4503 return isATensor4504 ? llvm::Intrinsic::4505 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale_block324506 : llvm::Intrinsic::4507 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale_block32;4508 }4509 } else if (kind == NVVM::Tcgen05MMABlockScaleKind::MXF4NVF4) {4510 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {4511 return isATensor4512 ? llvm::Intrinsic::4513 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block324514 : llvm::Intrinsic::4515 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block32;4516 4517 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {4518 return isATensor4519 ? llvm::Intrinsic::4520 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block164521 : llvm::Intrinsic::4522 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block16;4523 }4524 }4525 llvm_unreachable("Invalid tcgen05.mma.sp.block_scale attributes");4526 }();4527 4528 return {ID, args};4529}4530 4531LogicalResult Tcgen05MMASparseBlockScaleOp::verify() {4532 return verifyTcgen05MMABlockScaleOp(getCollectorOp(), getKind(),4533 getBlockScale(), getLoc());4534}4535 4536//===----------------------------------------------------------------------===//4537// NVVM tcgen05.mma.ws functions4538//===----------------------------------------------------------------------===//4539 4540mlir::NVVM::IDArgPair Tcgen05MMAWsOp::getIntrinsicIDAndArgs(4541 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {4542 4543 auto thisOp = cast<NVVM::Tcgen05MMAWsOp>(op);4544 llvm::SmallVector<llvm::Value *> args;4545 4546 args.push_back(mt.lookupValue(thisOp.getMatrixD()));4547 4548 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());4549 bool isATensor = isa<llvm::PointerType>(A->getType());4550 args.push_back(A);4551 4552 args.push_back(mt.lookupValue(thisOp.getMatrixB()));4553 args.push_back(mt.lookupValue(thisOp.getIdesc()));4554 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));4555 4556 mlir::Value ZeroColMask = thisOp.getZeroColMask();4557 llvm::Intrinsic::ID ID = notIntrinsic;4558 if (ZeroColMask) {4559 args.push_back(mt.lookupValue(ZeroColMask));4560 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor_zero_col_mask4561 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared_zero_col_mask;4562 } else4563 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor4564 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared;4565 4566 args.push_back(builder.getInt32(static_cast<unsigned>(thisOp.getKind())));4567 args.push_back(4568 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorBBuffer())));4569 args.push_back(4570 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));4571 4572 return {ID, args};4573}4574 4575//===----------------------------------------------------------------------===//4576// NVVM tcgen05.mma.ws.sp functions4577//===----------------------------------------------------------------------===//4578 4579mlir::NVVM::IDArgPair Tcgen05MMAWsSparseOp::getIntrinsicIDAndArgs(4580 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {4581 4582 auto thisOp = cast<NVVM::Tcgen05MMAWsSparseOp>(op);4583 llvm::SmallVector<llvm::Value *> args;4584 4585 args.push_back(mt.lookupValue(thisOp.getMatrixD()));4586 4587 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());4588 bool isATensor = isa<llvm::PointerType>(A->getType());4589 args.push_back(A);4590 4591 args.push_back(mt.lookupValue(thisOp.getMatrixB()));4592 args.push_back(mt.lookupValue(thisOp.getIdesc()));4593 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));4594 args.push_back(mt.lookupValue(thisOp.getSparseMetadata()));4595 4596 mlir::Value ZeroColMask = thisOp.getZeroColMask();4597 llvm::Intrinsic::ID ID = notIntrinsic;4598 if (ZeroColMask) {4599 args.push_back(mt.lookupValue(ZeroColMask));4600 ID = isATensor4601 ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor_zero_col_mask4602 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared_zero_col_mask;4603 } else4604 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor4605 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared;4606 4607 args.push_back(builder.getInt32(static_cast<unsigned>(thisOp.getKind())));4608 args.push_back(4609 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorBBuffer())));4610 args.push_back(4611 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));4612 4613 return {ID, args};4614}4615 4616//===----------------------------------------------------------------------===//4617// NVVMDialect initialization, type parsing, and registration.4618//===----------------------------------------------------------------------===//4619 4620// TODO: This should be the llvm.nvvm dialect once this is supported.4621void NVVMDialect::initialize() {4622 addOperations<4623#define GET_OP_LIST4624#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"4625 >();4626 addAttributes<4627#define GET_ATTRDEF_LIST4628#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"4629 >();4630 4631 // Support unknown operations because not all NVVM operations are4632 // registered.4633 allowUnknownOperations();4634 declarePromisedInterface<ConvertToLLVMPatternInterface, NVVMDialect>();4635 declarePromisedInterface<gpu::TargetAttrInterface, NVVMTargetAttr>();4636}4637 4638LogicalResult NVVMDialect::verifyOperationAttribute(Operation *op,4639 NamedAttribute attr) {4640 StringAttr attrName = attr.getName();4641 // Kernel function attribute should be attached to functions.4642 if (attrName == NVVMDialect::getKernelFuncAttrName()) {4643 if (!isa<LLVM::LLVMFuncOp>(op)) {4644 return op->emitError() << "'" << NVVMDialect::getKernelFuncAttrName()4645 << "' attribute attached to unexpected op";4646 }4647 }4648 // If maxntid / reqntid / cluster_dim exist, it must be an array with max 34649 // dim4650 if (attrName == NVVMDialect::getMaxntidAttrName() ||4651 attrName == NVVMDialect::getReqntidAttrName() ||4652 attrName == NVVMDialect::getClusterDimAttrName()) {4653 auto values = llvm::dyn_cast<DenseI32ArrayAttr>(attr.getValue());4654 if (!values || values.empty() || values.size() > 3) {4655 return op->emitError()4656 << "'" << attrName4657 << "' attribute must be integer array with maximum 3 index";4658 }4659 }4660 // If minctasm / maxnreg / cluster_max_blocks exist, it must be an integer4661 // attribute4662 if (attrName == NVVMDialect::getMinctasmAttrName() ||4663 attrName == NVVMDialect::getMaxnregAttrName() ||4664 attrName == NVVMDialect::getClusterMaxBlocksAttrName()) {4665 if (!llvm::dyn_cast<IntegerAttr>(attr.getValue())) {4666 return op->emitError()4667 << "'" << attrName << "' attribute must be integer constant";4668 }4669 }4670 // blocksareclusters must be used along with reqntid and cluster_dim4671 if (attrName == NVVMDialect::getBlocksAreClustersAttrName()) {4672 if (!op->hasAttr(NVVMDialect::getReqntidAttrName()) ||4673 !op->hasAttr(NVVMDialect::getClusterDimAttrName())) {4674 return op->emitError()4675 << "'" << attrName << "' attribute must be used along with "4676 << "'" << NVVMDialect::getReqntidAttrName() << "' and "4677 << "'" << NVVMDialect::getClusterDimAttrName() << "'";4678 }4679 }4680 4681 return success();4682}4683 4684LogicalResult NVVMDialect::verifyRegionArgAttribute(Operation *op,4685 unsigned regionIndex,4686 unsigned argIndex,4687 NamedAttribute argAttr) {4688 auto funcOp = dyn_cast<FunctionOpInterface>(op);4689 if (!funcOp)4690 return success();4691 4692 bool isKernel = op->hasAttr(NVVMDialect::getKernelFuncAttrName());4693 StringAttr attrName = argAttr.getName();4694 if (attrName == NVVM::NVVMDialect::getGridConstantAttrName()) {4695 if (!isKernel) {4696 return op->emitError()4697 << "'" << attrName4698 << "' attribute must be present only on kernel arguments";4699 }4700 if (!isa<UnitAttr>(argAttr.getValue()))4701 return op->emitError() << "'" << attrName << "' must be a unit attribute";4702 if (!funcOp.getArgAttr(argIndex, LLVM::LLVMDialect::getByValAttrName())) {4703 return op->emitError()4704 << "'" << attrName4705 << "' attribute requires the argument to also have attribute '"4706 << LLVM::LLVMDialect::getByValAttrName() << "'";4707 }4708 }4709 4710 return success();4711}4712 4713//===----------------------------------------------------------------------===//4714// NVVM Address Space Attr4715//===----------------------------------------------------------------------===//4716 4717unsigned NVVMMemorySpaceAttr::getAddressSpace() const {4718 return static_cast<unsigned>(getValue());4719}4720 4721bool NVVMMemorySpaceAttr::isValidLoad(4722 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,4723 const ::mlir::DataLayout *dataLayout,4724 function_ref<InFlightDiagnostic()> emitError) const {4725 return LLVM::detail::isValidLoadStoreImpl(type, ordering, alignment,4726 dataLayout, emitError);4727}4728 4729bool NVVMMemorySpaceAttr::isValidStore(4730 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,4731 const ::mlir::DataLayout *dataLayout,4732 function_ref<InFlightDiagnostic()> emitError) const {4733 return LLVM::detail::isValidLoadStoreImpl(type, ordering, alignment,4734 dataLayout, emitError);4735}4736 4737bool NVVMMemorySpaceAttr::isValidAtomicOp(4738 ptr::AtomicBinOp op, Type type, ptr::AtomicOrdering ordering,4739 std::optional<int64_t> alignment, const ::mlir::DataLayout *dataLayout,4740 function_ref<InFlightDiagnostic()> emitError) const {4741 // TODO: update this method once `ptr.atomic_rmw` is implemented.4742 assert(false && "unimplemented, see TODO in the source.");4743 return false;4744}4745 4746bool NVVMMemorySpaceAttr::isValidAtomicXchg(4747 Type type, ptr::AtomicOrdering successOrdering,4748 ptr::AtomicOrdering failureOrdering, std::optional<int64_t> alignment,4749 const ::mlir::DataLayout *dataLayout,4750 function_ref<InFlightDiagnostic()> emitError) const {4751 // TODO: update this method once `ptr.atomic_cmpxchg` is implemented.4752 assert(false && "unimplemented, see TODO in the source.");4753 return false;4754}4755 4756bool NVVMMemorySpaceAttr::isValidAddrSpaceCast(4757 Type tgt, Type src, function_ref<InFlightDiagnostic()> emitError) const {4758 // TODO: update this method once the `ptr.addrspace_cast` op is added to the4759 // dialect.4760 assert(false && "unimplemented, see TODO in the source.");4761 return false;4762}4763 4764bool NVVMMemorySpaceAttr::isValidPtrIntCast(4765 Type intLikeTy, Type ptrLikeTy,4766 function_ref<InFlightDiagnostic()> emitError) const {4767 // TODO: update this method once the int-cast ops are added to the `ptr`4768 // dialect.4769 assert(false && "unimplemented, see TODO in the source.");4770 return false;4771}4772 4773//===----------------------------------------------------------------------===//4774// NVVM target attribute.4775//===----------------------------------------------------------------------===//4776LogicalResult4777NVVMTargetAttr::verify(function_ref<InFlightDiagnostic()> emitError,4778 int optLevel, StringRef triple, StringRef chip,4779 StringRef features, DictionaryAttr flags,4780 ArrayAttr files, bool verifyTarget) {4781 if (optLevel < 0 || optLevel > 3) {4782 emitError() << "The optimization level must be a number between 0 and 3.";4783 return failure();4784 }4785 if (triple.empty()) {4786 emitError() << "The target triple cannot be empty.";4787 return failure();4788 }4789 if (chip.empty()) {4790 emitError() << "The target chip cannot be empty.";4791 return failure();4792 }4793 if (files && !llvm::all_of(files, [](::mlir::Attribute attr) {4794 return mlir::isa_and_nonnull<StringAttr>(attr);4795 })) {4796 emitError() << "All the elements in the `link` array must be strings.";4797 return failure();4798 }4799 return success();4800}4801 4802LogicalResult NVVMTargetAttr::verifyTarget(Operation *gpuModule) {4803 if (!getVerifyTarget())4804 return success();4805 4806 auto gpuModuleOp = llvm::dyn_cast<gpu::GPUModuleOp>(gpuModule);4807 if (!gpuModuleOp) {4808 return emitError(gpuModule->getLoc(),4809 "NVVM target attribute must be attached to a GPU module");4810 }4811 4812 const NVVMCheckSMVersion targetSMVersion =4813 NVVMCheckSMVersion::getTargetSMVersionFromStr(getChip());4814 if (!targetSMVersion.isMinimumSMVersion()) {4815 return emitError(gpuModule->getLoc(),4816 "Minimum NVVM target SM version is sm_20");4817 }4818 4819 if (gpuModuleOp4820 ->walk([&](Operation *op) {4821 if (auto reqOp = llvm::dyn_cast<NVVM::RequiresSMInterface>(op)) {4822 const NVVMCheckSMVersion requirement =4823 reqOp.getRequiredMinSMVersion();4824 if (!requirement.isCompatibleWith(targetSMVersion)) {4825 op->emitOpError() << "is not supported on " << getChip();4826 return WalkResult::interrupt();4827 }4828 }4829 return WalkResult::advance();4830 })4831 .wasInterrupted())4832 return failure();4833 4834 return success();4835}4836 4837#define GET_OP_CLASSES4838#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"4839 4840#define GET_ATTRDEF_CLASSES4841#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"4842