brintos

brintos / llvm-project-archived public Read only

0
0
Text · 188.2 KiB · 4131252 Raw
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