brintos

brintos / llvm-project-archived public Read only

0
0
Text · 97.0 KiB · b9a5e7d Raw
2339 lines · cpp
1//===- AMDGPUToROCDL.cpp - AMDGPU to ROCDL dialect conversion -------===//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#include "mlir/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.h"10 11#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"12#include "mlir/Conversion/LLVMCommon/Pattern.h"13#include "mlir/Conversion/LLVMCommon/TypeConverter.h"14#include "mlir/Dialect/AMDGPU/IR/AMDGPUDialect.h"15#include "mlir/Dialect/AMDGPU/Utils/Chipset.h"16#include "mlir/Dialect/LLVMIR/LLVMDialect.h"17#include "mlir/Dialect/LLVMIR/LLVMTypes.h"18#include "mlir/Dialect/LLVMIR/ROCDLDialect.h"19#include "mlir/IR/Attributes.h"20#include "mlir/IR/BuiltinAttributes.h"21#include "mlir/IR/BuiltinTypes.h"22#include "mlir/IR/TypeUtilities.h"23#include "mlir/Pass/Pass.h"24 25#include "../LLVMCommon/MemRefDescriptor.h"26 27#include "llvm/ADT/STLExtras.h"28#include "llvm/ADT/TypeSwitch.h"29#include "llvm/Support/Casting.h"30#include "llvm/Support/ErrorHandling.h"31#include <optional>32 33namespace mlir {34#define GEN_PASS_DEF_CONVERTAMDGPUTOROCDLPASS35#include "mlir/Conversion/Passes.h.inc"36} // namespace mlir37 38using namespace mlir;39using namespace mlir::amdgpu;40 41// Define commonly used chipsets versions for convenience.42constexpr Chipset kGfx908 = Chipset(9, 0, 8);43constexpr Chipset kGfx90a = Chipset(9, 0, 0xa);44constexpr Chipset kGfx942 = Chipset(9, 4, 2);45constexpr Chipset kGfx950 = Chipset(9, 5, 0);46constexpr Chipset kGfx1250 = Chipset(12, 5, 0);47 48/// Convert an unsigned number `val` to i32.49static Value convertUnsignedToI32(ConversionPatternRewriter &rewriter,50                                  Location loc, Value val) {51  IntegerType i32 = rewriter.getI32Type();52  // Force check that `val` is of int type.53  auto valTy = cast<IntegerType>(val.getType());54  if (i32 == valTy)55    return val;56  return valTy.getWidth() > 3257             ? Value(LLVM::TruncOp::create(rewriter, loc, i32, val))58             : Value(LLVM::ZExtOp::create(rewriter, loc, i32, val));59}60 61static Value createI32Constant(ConversionPatternRewriter &rewriter,62                               Location loc, int32_t value) {63  return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), value);64}65 66/// Convert an unsigned number `val` to i64.67static Value convertUnsignedToI64(ConversionPatternRewriter &rewriter,68                                  Location loc, Value val) {69  IntegerType i64 = rewriter.getI64Type();70  // Force check that `val` is of int type.71  auto valTy = cast<IntegerType>(val.getType());72  if (i64 == valTy)73    return val;74  return valTy.getWidth() > 6475             ? Value(LLVM::TruncOp::create(rewriter, loc, i64, val))76             : Value(LLVM::ZExtOp::create(rewriter, loc, i64, val));77}78 79static Value createI64Constant(ConversionPatternRewriter &rewriter,80                               Location loc, int64_t value) {81  return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), value);82}83 84/// Returns the linear index used to access an element in the memref.85static Value getLinearIndexI32(ConversionPatternRewriter &rewriter,86                               Location loc, MemRefDescriptor &memRefDescriptor,87                               ValueRange indices, ArrayRef<int64_t> strides) {88  IntegerType i32 = rewriter.getI32Type();89  Value index;90  for (auto [i, increment, stride] : llvm::enumerate(indices, strides)) {91    if (stride != 1) { // Skip if stride is 1.92      Value strideValue =93          ShapedType::isDynamic(stride)94              ? convertUnsignedToI32(rewriter, loc,95                                     memRefDescriptor.stride(rewriter, loc, i))96              : LLVM::ConstantOp::create(rewriter, loc, i32, stride);97      increment = LLVM::MulOp::create(rewriter, loc, increment, strideValue);98    }99    index = index ? LLVM::AddOp::create(rewriter, loc, index, increment)100                  : increment;101  }102  return index ? index : createI32Constant(rewriter, loc, 0);103}104 105/// Compute the contents of the `num_records` field for a given memref106/// descriptor - that is, the number of bytes that's one element past the107/// greatest possible valid index into the memref.108static Value getNumRecords(ConversionPatternRewriter &rewriter, Location loc,109                           MemRefType memrefType,110                           MemRefDescriptor &memrefDescriptor,111                           ArrayRef<int64_t> strides,112                           int64_t elementByteWidth) {113  if (memrefType.hasStaticShape() &&114      !llvm::any_of(strides, ShapedType::isDynamic)) {115    int64_t size = memrefType.getRank() == 0 ? 1 : 0;116    ArrayRef<int64_t> shape = memrefType.getShape();117    for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i)118      size = std::max(shape[i] * strides[i], size);119    size = size * elementByteWidth;120    return createI64Constant(rewriter, loc, size);121  }122  Value maxIndex;123  for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i) {124    Value size = memrefDescriptor.size(rewriter, loc, i);125    Value stride = memrefDescriptor.stride(rewriter, loc, i);126    Value maxThisDim = LLVM::MulOp::create(rewriter, loc, size, stride);127    maxIndex = maxIndex128                   ? LLVM::UMaxOp::create(rewriter, loc, maxIndex, maxThisDim)129                   : maxThisDim;130  }131  Value maxIndexI64 = convertUnsignedToI64(rewriter, loc, maxIndex);132  Value byteWidthConst = createI64Constant(rewriter, loc, elementByteWidth);133  return LLVM::MulOp::create(rewriter, loc, maxIndexI64, byteWidthConst);134}135 136static Value makeBufferRsrc(ConversionPatternRewriter &rewriter, Location loc,137                            Value basePointer, Value numRecords,138                            bool boundsCheck, amdgpu::Chipset chipset,139                            Value cacheSwizzleStride = nullptr,140                            unsigned addressSpace = 8) {141  // The stride value is generally 0. However, on MI-300 and onward, you can142  // enable a cache swizzling mode by setting bit 14 of the stride field143  // and setting that stride to a cache stride.144  Type i16 = rewriter.getI16Type();145  Value stride;146  if (chipset.majorVersion == 9 && chipset >= kGfx942 && cacheSwizzleStride) {147    Value cacheStrideZext =148        LLVM::ZExtOp::create(rewriter, loc, i16, cacheSwizzleStride);149    Value swizzleBit = LLVM::ConstantOp::create(150        rewriter, loc, i16, rewriter.getI16IntegerAttr(1 << 14));151    stride = LLVM::OrOp::create(rewriter, loc, cacheStrideZext, swizzleBit,152                                /*isDisjoint=*/true);153  } else {154    stride = LLVM::ConstantOp::create(rewriter, loc, i16,155                                      rewriter.getI16IntegerAttr(0));156  }157  // Get the number of elements.158  // Flag word:159  // bits 0-11: dst sel, ignored by these intrinsics160  // bits 12-14: data format (ignored, must be nonzero, 7=float)161  // bits 15-18: data format (ignored, must be nonzero, 4=32bit)162  // bit 19: In nested heap (0 here)163  // bit 20: Behavior on unmap (0 means  "return 0 / ignore")164  // bits 21-22: Index stride for swizzles (N/A)165  // bit 23: Add thread ID (0)166  // bit 24: Reserved to 1 (RDNA) or 0 (CDNA)167  // bits 25-26: Reserved (0)168  // bit 27: Buffer is non-volatile (CDNA only)169  // bits 28-29: Out of bounds select (0 = structured, 1 = check index, 2 =170  //  none, 3 = either swizzles or testing against offset field) RDNA only171  // bits 30-31: Type (must be 0)172  uint32_t flags = (7 << 12) | (4 << 15);173  if (chipset.majorVersion >= 10) {174    flags |= (1 << 24);175    uint32_t oob = boundsCheck ? 3 : 2;176    flags |= (oob << 28);177  }178  Value flagsConst = createI32Constant(rewriter, loc, flags);179  Type rsrcType =180      LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);181  Value resource = rewriter.createOrFold<ROCDL::MakeBufferRsrcOp>(182      loc, rsrcType, basePointer, stride, numRecords, flagsConst);183  return resource;184}185 186namespace {187struct FatRawBufferCastLowering188    : public ConvertOpToLLVMPattern<FatRawBufferCastOp> {189  FatRawBufferCastLowering(const LLVMTypeConverter &converter, Chipset chipset)190      : ConvertOpToLLVMPattern<FatRawBufferCastOp>(converter),191        chipset(chipset) {}192 193  Chipset chipset;194 195  LogicalResult196  matchAndRewrite(FatRawBufferCastOp op, FatRawBufferCastOpAdaptor adaptor,197                  ConversionPatternRewriter &rewriter) const override {198    Location loc = op.getLoc();199    Value memRef = adaptor.getSource();200    Value unconvertedMemref = op.getSource();201    MemRefType memrefType = cast<MemRefType>(unconvertedMemref.getType());202    MemRefDescriptor descriptor(memRef);203 204    DataLayout dataLayout = DataLayout::closest(op);205    int64_t elementByteWidth =206        dataLayout.getTypeSizeInBits(memrefType.getElementType()) / 8;207 208    int64_t unusedOffset = 0;209    SmallVector<int64_t, 5> strideVals;210    if (failed(memrefType.getStridesAndOffset(strideVals, unusedOffset)))211      return op.emitOpError("Can't lower non-stride-offset memrefs");212 213    Value numRecords = adaptor.getValidBytes();214    if (!numRecords)215      numRecords = getNumRecords(rewriter, loc, memrefType, descriptor,216                                 strideVals, elementByteWidth);217 218    Value basePointer =219        adaptor.getResetOffset()220            ? descriptor.bufferPtr(rewriter, loc, *getTypeConverter(),221                                   memrefType)222            : descriptor.alignedPtr(rewriter, loc);223 224    Value offset = adaptor.getResetOffset()225                       ? LLVM::ConstantOp::create(rewriter, loc, getIndexType(),226                                                  rewriter.getIndexAttr(0))227                       : descriptor.offset(rewriter, loc);228 229    bool hasSizes = memrefType.getRank() > 0;230    // No need to unpack() and pack() all the individual sizes and strides,231    // so we'll just extract the arrays.232    Value sizes = hasSizes233                      ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,234                                                     kSizePosInMemRefDescriptor)235                      : Value{};236    Value strides =237        hasSizes ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,238                                                kStridePosInMemRefDescriptor)239                 : Value{};240 241    Value fatPtr = makeBufferRsrc(242        rewriter, loc, basePointer, numRecords, adaptor.getBoundsCheck(),243        chipset, adaptor.getCacheSwizzleStride(), /*addressSpace=*/7);244 245    Value result = MemRefDescriptor::poison(246        rewriter, loc,247        getTypeConverter()->convertType(op.getResult().getType()));248    SmallVector<int64_t> pos{kAllocatedPtrPosInMemRefDescriptor};249    result = LLVM::InsertValueOp::create(rewriter, loc, result, fatPtr, pos);250    result = LLVM::InsertValueOp::create(rewriter, loc, result, fatPtr,251                                         kAlignedPtrPosInMemRefDescriptor);252    result = LLVM::InsertValueOp::create(rewriter, loc, result, offset,253                                         kOffsetPosInMemRefDescriptor);254    if (hasSizes) {255      result = LLVM::InsertValueOp::create(rewriter, loc, result, sizes,256                                           kSizePosInMemRefDescriptor);257      result = LLVM::InsertValueOp::create(rewriter, loc, result, strides,258                                           kStridePosInMemRefDescriptor);259    }260    rewriter.replaceOp(op, result);261    return success();262  }263};264 265/// Define lowering patterns for raw buffer ops266template <typename GpuOp, typename Intrinsic>267struct RawBufferOpLowering : public ConvertOpToLLVMPattern<GpuOp> {268  RawBufferOpLowering(const LLVMTypeConverter &converter, Chipset chipset)269      : ConvertOpToLLVMPattern<GpuOp>(converter), chipset(chipset) {}270 271  Chipset chipset;272  static constexpr uint32_t maxVectorOpWidth = 128;273 274  LogicalResult275  matchAndRewrite(GpuOp gpuOp, typename GpuOp::Adaptor adaptor,276                  ConversionPatternRewriter &rewriter) const override {277    Location loc = gpuOp.getLoc();278    Value memref = adaptor.getMemref();279    Value unconvertedMemref = gpuOp.getMemref();280    MemRefType memrefType = cast<MemRefType>(unconvertedMemref.getType());281 282    if (chipset.majorVersion < 9)283      return gpuOp.emitOpError("raw buffer ops require GCN or higher");284 285    Value storeData = adaptor.getODSOperands(0)[0];286    if (storeData == memref) // no write component to this op287      storeData = Value();288    Type wantedDataType;289    if (storeData)290      wantedDataType = storeData.getType();291    else292      wantedDataType = gpuOp.getODSResults(0)[0].getType();293 294    Value atomicCmpData = Value();295    // Operand index 1 of a load is the indices, trying to read them can crash.296    if (storeData) {297      Value maybeCmpData = adaptor.getODSOperands(1)[0];298      if (maybeCmpData != memref)299        atomicCmpData = maybeCmpData;300    }301 302    Type llvmWantedDataType = this->typeConverter->convertType(wantedDataType);303 304    Type i32 = rewriter.getI32Type();305 306    // Get the type size in bytes.307    DataLayout dataLayout = DataLayout::closest(gpuOp);308    int64_t elementByteWidth =309        dataLayout.getTypeSizeInBits(memrefType.getElementType()) / 8;310    Value byteWidthConst = createI32Constant(rewriter, loc, elementByteWidth);311 312    // If we want to load a vector<NxT> with total size <= 32313    // bits, use a scalar load and bitcast it. Similarly, if bitsize(T) < 32314    // and the total load size is >= 32, use a vector load of N / (bitsize(T) /315    // 32) x i32 and bitcast. Also, the CAS intrinsic requires integer operands,316    // so bitcast any floats to integers.317    Type llvmBufferValType = llvmWantedDataType;318    if (atomicCmpData) {319      if (auto floatType = dyn_cast<FloatType>(wantedDataType))320        llvmBufferValType = this->getTypeConverter()->convertType(321            rewriter.getIntegerType(floatType.getWidth()));322    }323    if (auto dataVector = dyn_cast<VectorType>(wantedDataType)) {324      uint32_t vecLen = dataVector.getNumElements();325      uint32_t elemBits =326          dataLayout.getTypeSizeInBits(dataVector.getElementType());327      uint32_t totalBits = elemBits * vecLen;328      bool usePackedFp16 =329          isa_and_present<RawBufferAtomicFaddOp>(*gpuOp) && vecLen == 2;330      if (totalBits > maxVectorOpWidth)331        return gpuOp.emitOpError(332            "Total width of loads or stores must be no more than " +333            Twine(maxVectorOpWidth) + " bits, but we call for " +334            Twine(totalBits) +335            " bits. This should've been caught in validation");336      if (!usePackedFp16 && elemBits < 32) {337        if (totalBits > 32) {338          if (totalBits % 32 != 0)339            return gpuOp.emitOpError("Load or store of more than 32-bits that "340                                     "doesn't fit into words. Can't happen\n");341          llvmBufferValType = this->typeConverter->convertType(342              VectorType::get(totalBits / 32, i32));343        } else {344          llvmBufferValType = this->typeConverter->convertType(345              rewriter.getIntegerType(totalBits));346        }347      }348    }349    if (auto vecType = dyn_cast<VectorType>(llvmBufferValType)) {350      // Buffer intrinsics doesn't support 1-element vectors, cast them to351      // scalars.352      if (vecType.getNumElements() == 1)353        llvmBufferValType = vecType.getElementType();354    }355 356    SmallVector<Value, 6> args;357    if (storeData) {358      if (llvmBufferValType != llvmWantedDataType) {359        Value castForStore = LLVM::BitcastOp::create(360            rewriter, loc, llvmBufferValType, storeData);361        args.push_back(castForStore);362      } else {363        args.push_back(storeData);364      }365    }366 367    if (atomicCmpData) {368      if (llvmBufferValType != llvmWantedDataType) {369        Value castForCmp = LLVM::BitcastOp::create(370            rewriter, loc, llvmBufferValType, atomicCmpData);371        args.push_back(castForCmp);372      } else {373        args.push_back(atomicCmpData);374      }375    }376 377    // Construct buffer descriptor from memref, attributes378    int64_t offset = 0;379    SmallVector<int64_t, 5> strides;380    if (failed(memrefType.getStridesAndOffset(strides, offset)))381      return gpuOp.emitOpError("Can't lower non-stride-offset memrefs");382 383    MemRefDescriptor memrefDescriptor(memref);384 385    Value ptr = memrefDescriptor.bufferPtr(386        rewriter, loc, *this->getTypeConverter(), memrefType);387    Value numRecords = getNumRecords(388        rewriter, loc, memrefType, memrefDescriptor, strides, elementByteWidth);389    Value resource = makeBufferRsrc(rewriter, loc, ptr, numRecords,390                                    adaptor.getBoundsCheck(), chipset);391    args.push_back(resource);392 393    // Indexing (voffset)394    Value voffset = getLinearIndexI32(rewriter, loc, memrefDescriptor,395                                      adaptor.getIndices(), strides);396    if (std::optional<int32_t> indexOffset = adaptor.getIndexOffset();397        indexOffset && *indexOffset > 0) {398      Value extraOffsetConst = createI32Constant(rewriter, loc, *indexOffset);399      voffset = voffset ? LLVM::AddOp::create(rewriter, loc, voffset,400                                              extraOffsetConst)401                        : extraOffsetConst;402    }403    voffset = LLVM::MulOp::create(rewriter, loc, voffset, byteWidthConst);404    args.push_back(voffset);405 406    // SGPR offset.407    Value sgprOffset = adaptor.getSgprOffset();408    if (!sgprOffset)409      sgprOffset = createI32Constant(rewriter, loc, 0);410    sgprOffset = LLVM::MulOp::create(rewriter, loc, sgprOffset, byteWidthConst);411    args.push_back(sgprOffset);412 413    // bit 0: GLC = 0 (atomics drop value, less coherency)414    // bits 1-2: SLC, DLC = 0 (similarly)415    // bit 3: swizzled (0 for raw)416    args.push_back(createI32Constant(rewriter, loc, 0));417 418    llvm::SmallVector<Type, 1> resultTypes(gpuOp->getNumResults(),419                                           llvmBufferValType);420    Operation *lowered = Intrinsic::create(rewriter, loc, resultTypes, args,421                                           ArrayRef<NamedAttribute>());422    if (lowered->getNumResults() == 1) {423      Value replacement = lowered->getResult(0);424      if (llvmBufferValType != llvmWantedDataType) {425        replacement = LLVM::BitcastOp::create(rewriter, loc, llvmWantedDataType,426                                              replacement);427      }428      rewriter.replaceOp(gpuOp, replacement);429    } else {430      rewriter.eraseOp(gpuOp);431    }432    return success();433  }434};435 436// TODO: AMDGPU backend already have all this bitpacking logic, we should move437// it to some common place.438///  Vmcnt, Expcnt and Lgkmcnt are decoded as follows:439///     Vmcnt = Waitcnt[3:0]        (pre-gfx9)440///     Vmcnt = Waitcnt[15:14,3:0]  (gfx9,10)441///     Vmcnt = Waitcnt[15:10]      (gfx11)442///     Expcnt = Waitcnt[6:4]       (pre-gfx11)443///     Expcnt = Waitcnt[2:0]       (gfx11)444///     Lgkmcnt = Waitcnt[11:8]     (pre-gfx10)445///     Lgkmcnt = Waitcnt[13:8]     (gfx10)446///     Lgkmcnt = Waitcnt[9:4]      (gfx11)447static FailureOr<unsigned> encodeWaitcnt(Chipset chipset, unsigned vmcnt,448                                         unsigned expcnt, unsigned lgkmcnt) {449  if (chipset.majorVersion < 9) {450    vmcnt = std::min(15u, vmcnt);451    expcnt = std::min(7u, expcnt);452    lgkmcnt = std::min(15u, lgkmcnt);453    return vmcnt | (expcnt << 4) | (lgkmcnt << 8);454  }455  if (chipset.majorVersion == 9) {456    vmcnt = std::min(63u, vmcnt);457    expcnt = std::min(7u, expcnt);458    lgkmcnt = std::min(15u, lgkmcnt);459    unsigned lowBits = vmcnt & 0xF;460    unsigned highBits = (vmcnt >> 4) << 14;461    unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);462    return lowBits | highBits | otherCnts;463  }464  if (chipset.majorVersion == 10) {465    vmcnt = std::min(63u, vmcnt);466    expcnt = std::min(7u, expcnt);467    lgkmcnt = std::min(63u, lgkmcnt);468    unsigned lowBits = vmcnt & 0xF;469    unsigned highBits = (vmcnt >> 4) << 14;470    unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);471    return lowBits | highBits | otherCnts;472  }473  if (chipset.majorVersion == 11) {474    vmcnt = std::min(63u, vmcnt);475    expcnt = std::min(7u, expcnt);476    lgkmcnt = std::min(63u, lgkmcnt);477    return (vmcnt << 10) | expcnt | (lgkmcnt << 4);478  }479  return failure();480}481 482struct MemoryCounterWaitOpLowering483    : public ConvertOpToLLVMPattern<MemoryCounterWaitOp> {484  MemoryCounterWaitOpLowering(const LLVMTypeConverter &converter,485                              Chipset chipset)486      : ConvertOpToLLVMPattern<MemoryCounterWaitOp>(converter),487        chipset(chipset) {}488 489  Chipset chipset;490 491  LogicalResult492  matchAndRewrite(MemoryCounterWaitOp op, OpAdaptor adaptor,493                  ConversionPatternRewriter &rewriter) const override {494    if (chipset.majorVersion >= 12) {495      Location loc = op.getLoc();496      if (std::optional<int> ds = adaptor.getDs())497        ROCDL::WaitDscntOp::create(rewriter, loc, *ds);498 499      if (std::optional<int> load = adaptor.getLoad())500        ROCDL::WaitLoadcntOp::create(rewriter, loc, *load);501 502      if (std::optional<int> store = adaptor.getStore())503        ROCDL::WaitStorecntOp::create(rewriter, loc, *store);504 505      if (std::optional<int> exp = adaptor.getExp())506        ROCDL::WaitExpcntOp::create(rewriter, loc, *exp);507 508      rewriter.eraseOp(op);509      return success();510    }511 512    auto getVal = [](Attribute attr) -> unsigned {513      if (attr)514        return cast<IntegerAttr>(attr).getInt();515 516      // This value will be clamped to the maximum value for the chipset.517      return 1024;518    };519    unsigned ds = getVal(adaptor.getDsAttr());520    unsigned exp = getVal(adaptor.getExpAttr());521 522    unsigned vmcnt = 1024;523    Attribute load = adaptor.getLoadAttr();524    Attribute store = adaptor.getStoreAttr();525    if (load && store) {526      vmcnt = getVal(load) + getVal(store);527    } else if (load) {528      vmcnt = getVal(load);529    } else if (store) {530      vmcnt = getVal(store);531    }532 533    FailureOr<unsigned> waitcnt = encodeWaitcnt(chipset, vmcnt, exp, ds);534    if (failed(waitcnt))535      return op.emitOpError("unsupported chipset");536 537    rewriter.replaceOpWithNewOp<ROCDL::SWaitcntOp>(op, *waitcnt);538    return success();539  }540};541 542struct LDSBarrierOpLowering : public ConvertOpToLLVMPattern<LDSBarrierOp> {543  LDSBarrierOpLowering(const LLVMTypeConverter &converter, Chipset chipset)544      : ConvertOpToLLVMPattern<LDSBarrierOp>(converter), chipset(chipset) {}545 546  Chipset chipset;547 548  LogicalResult549  matchAndRewrite(LDSBarrierOp op, LDSBarrierOp::Adaptor adaptor,550                  ConversionPatternRewriter &rewriter) const override {551    Location loc = op.getLoc();552    // This ensures that waits on global memory aren't introduced on553    // chips that don't have the BackOffBarrier feature enabled in LLVM.554    bool requiresInlineAsm = chipset < kGfx90a;555 556    Attribute mmra =557        rewriter.getAttr<LLVM::MMRATagAttr>("amdgpu-synchronize-as", "local");558    // Note: while there *is* a workgroup-one-as scope, this, when combined with559    // the MMRA, will lead to the fence having no effect. This is because the560    // codepaths for an atomic load or store will observe that a561    // one-address-space atomic to LDS requires no synchronization because562    // operations on LDS are totally ordered with respect to each other, and so563    // will not emit the correct waitcnt operations that these fences are564    // intended to produce. Therefore, we use a broader type of fence and rely565    // on the MMRA to relax it to the semantics we want.566    StringRef scope = "workgroup";567 568    auto relFence = LLVM::FenceOp::create(rewriter, loc,569                                          LLVM::AtomicOrdering::release, scope);570    relFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);571    if (requiresInlineAsm) {572      auto asmDialectAttr = LLVM::AsmDialectAttr::get(rewriter.getContext(),573                                                      LLVM::AsmDialect::AD_ATT);574      const char *asmStr = ";;;WARNING: BREAKS DEBUG WATCHES\ns_barrier";575      const char *constraints = "";576      LLVM::InlineAsmOp::create(577          rewriter, loc,578          /*resultTypes=*/TypeRange(), /*operands=*/ValueRange(),579          /*asm_string=*/asmStr, constraints, /*has_side_effects=*/true,580          /*is_align_stack=*/false, LLVM::TailCallKind::None,581          /*asm_dialect=*/asmDialectAttr,582          /*operand_attrs=*/ArrayAttr());583    } else if (chipset.majorVersion < 12) {584      ROCDL::SBarrierOp::create(rewriter, loc);585    } else {586      ROCDL::BarrierSignalOp::create(rewriter, loc, -1);587      ROCDL::BarrierWaitOp::create(rewriter, loc, -1);588    }589 590    auto acqFence = LLVM::FenceOp::create(rewriter, loc,591                                          LLVM::AtomicOrdering::acquire, scope);592    acqFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);593    rewriter.replaceOp(op, acqFence);594    return success();595  }596};597 598struct SchedBarrierOpLowering : public ConvertOpToLLVMPattern<SchedBarrierOp> {599  SchedBarrierOpLowering(const LLVMTypeConverter &converter, Chipset chipset)600      : ConvertOpToLLVMPattern<SchedBarrierOp>(converter), chipset(chipset) {}601 602  Chipset chipset;603 604  LogicalResult605  matchAndRewrite(SchedBarrierOp op, SchedBarrierOp::Adaptor adaptor,606                  ConversionPatternRewriter &rewriter) const override {607    rewriter.replaceOpWithNewOp<ROCDL::SchedBarrier>(op,608                                                     (uint32_t)op.getOpts());609    return success();610  }611};612 613} // namespace614 615/// Converts a MFMA vector operand from MLIR AMDGPU dialect convention to ROCDL616/// and LLVM AMDGPU intrinsics convention.617///618/// Specifically:619/// 1. If the element type is bfloat16, bitcast it to i16 unless rocdl intrinsic620/// allows bf16. Newer MFMAs support bf16 types on operand, check621/// IntrinsicsAMDGPU.td file for reference.622/// 2. If instead we have a more than 64-bit quantity, use a <N / 4 x i32>623/// instead, which is what the f8f6f4 intrinsics use.624/// 3. If `input` is a vector of N <= 8 bytes, bitcast it to a (N * 8)-bit625/// integer.626///627/// Note that the type of `input` has already been LLVM type converted:628/// therefore 8-bit and smaller floats are represented as their corresponding629/// `iN` integers.630static Value convertMFMAVectorOperand(ConversionPatternRewriter &rewriter,631                                      Location loc, Value input,632                                      bool allowBf16 = true) {633  Type inputType = input.getType();634  if (auto vectorType = dyn_cast<VectorType>(inputType)) {635    if (vectorType.getElementType().isBF16() && !allowBf16)636      return LLVM::BitcastOp::create(637          rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);638    if (vectorType.getElementType().isInteger(8) &&639        vectorType.getNumElements() <= 8)640      return LLVM::BitcastOp::create(641          rewriter, loc,642          rewriter.getIntegerType(vectorType.getNumElements() * 8), input);643    if (isa<IntegerType>(vectorType.getElementType()) &&644        vectorType.getElementTypeBitWidth() <= 8) {645      int64_t numWords = llvm::divideCeil(646          vectorType.getNumElements() * vectorType.getElementTypeBitWidth(),647          32);648      return LLVM::BitcastOp::create(649          rewriter, loc, VectorType::get(numWords, rewriter.getI32Type()),650          input);651    }652  }653  return input;654}655 656/// Converts the scaled MFMA operands, `scalesA` and `scalesB`, from MLIR AMDGPU657/// dialect convention to ROCDL and LLVM AMDGPU intrinsics convention.658///659/// Specifically:660/// 1. If `input` is a i8 value, zero extend it to i32661/// 2. If `input` is a vector of length 4 and type i8, cast it to i32662///663/// Note that the type of `input` has already been LLVM type converted:664/// therefore 8-bit and smaller floats are represented as their corresponding665/// `iN` integers.666static Value castMFMAScaleOperand(ConversionPatternRewriter &rewriter,667                                  Location loc, Value input) {668  Type inputType = input.getType();669  Type outputType = rewriter.getI32Type();670  if (auto intType = dyn_cast<IntegerType>(inputType))671    return LLVM::ZExtOp::create(rewriter, loc, outputType, input);672  return LLVM::BitcastOp::create(rewriter, loc, outputType, input);673}674 675/// Push an input operand. If it is a float type, nothing to do. If it is676/// an integer type, then we need to also push its signdness (1 for signed, 0677/// for unsigned) and we need to pack the input 16xi8 vector into a 4xi32678/// vector (or the 8xi8 vector into a 2xi32 one for gfx12+).679/// We also need to convert bfloat inputs to i16 to account for the bfloat680/// intrinsics having been defined before the AMD backend supported bfloat. We681/// similarly need to pack 8-bit float types into integers as if they were i8682/// (which they are for the backend's purposes).683static void wmmaPushInputOperand(684    ConversionPatternRewriter &rewriter, Location loc,685    const TypeConverter *typeConverter, bool isUnsigned, Value llvmInput,686    Value mlirInput, SmallVectorImpl<Value> &operands,687    SmallVectorImpl<NamedAttribute> &attrs, StringRef attrName) {688  Type inputType = llvmInput.getType();689  auto vectorType = dyn_cast<VectorType>(inputType);690  if (!vectorType) {691    operands.push_back(llvmInput);692    return;693  }694  Type elemType = vectorType.getElementType();695  if (elemType.getIntOrFloatBitWidth() > 8) {696    operands.push_back(llvmInput);697    return;698  }699 700  // We need to check the type of the input before conversion to properly test701  // for int8. This is because, in LLVM, fp8 type is converted to int8, so the702  // fp8/int8 information is lost during the conversion process.703  auto mlirInputType = cast<VectorType>(mlirInput.getType());704  bool isInputInteger = mlirInputType.getElementType().isInteger();705  if (isInputInteger) {706    // if element type is 8-bit signed or unsigned, ignore the isUnsigned flag707    bool localIsUnsigned = isUnsigned;708    if (elemType.isUnsignedInteger()) {709      localIsUnsigned = true;710    } else if (elemType.isSignedInteger()) {711      localIsUnsigned = false;712    }713    attrs.push_back(714        NamedAttribute(attrName, rewriter.getBoolAttr(!localIsUnsigned)));715  }716 717  int64_t numBits =718      vectorType.getNumElements() * elemType.getIntOrFloatBitWidth();719  Type i32 = rewriter.getI32Type();720  Type intrinsicInType = numBits <= 32721                             ? (Type)rewriter.getIntegerType(numBits)722                             : (Type)VectorType::get(numBits / 32, i32);723  auto llvmIntrinsicInType = typeConverter->convertType(intrinsicInType);724  Value castInput = rewriter.createOrFold<LLVM::BitcastOp>(725      loc, llvmIntrinsicInType, llvmInput);726  // The wave64-mode 16x16x16 intrinsics that take 4-bit integers only need727  // (256 / 64) * 4 = 16 bits of input (on gfx12+) but take i32 arguments.728  // Add in the zeros here.729  if (numBits < 32)730    castInput = LLVM::ZExtOp::create(rewriter, loc, i32, castInput);731  operands.push_back(castInput);732}733 734/// Push the output operand. For many cases this is only pushing the output in735/// the operand list. But when we have f16 -> f16 or bf16 -> bf16 intrinsics,736/// since the same numbers of VGPRs is used, we need to decide if to store the737/// result in the upper 16 bits of the VGPRs or in the lower part. To store the738/// result in the lower 16 bits, set subwordOffset to 1, otherwise result will739/// be stored it in the upper part. The subwordOffset must not be set for gfx12,740/// as the instructions have been changed to return fewer registers instead.741static void wmmaPushOutputOperand(ConversionPatternRewriter &rewriter,742                                  Location loc,743                                  const TypeConverter *typeConverter,744                                  Value output, int32_t subwordOffset,745                                  bool clamp, SmallVectorImpl<Value> &operands,746                                  SmallVectorImpl<NamedAttribute> &attrs) {747  Type inputType = output.getType();748  auto vectorType = dyn_cast<VectorType>(inputType);749  Type elemType = vectorType.getElementType();750  operands.push_back(output);751  if (elemType.isF16() || elemType.isBF16() || elemType.isInteger(16)) {752    attrs.push_back(753        NamedAttribute("opsel", rewriter.getBoolAttr(subwordOffset)));754  } else if (elemType.isInteger(32)) {755    attrs.push_back(NamedAttribute("clamp", rewriter.getBoolAttr(clamp)));756  }757}758 759/// Return true if `type` is the E5M2 variant of an 8-bit float that is760/// supported by the `_bf8` instructions on the given `chipset`.761static bool typeIsExpectedBf8ForChipset(Chipset chipset, Type type) {762  return (chipset == kGfx942 && isa<Float8E5M2FNUZType>(type)) ||763         (hasOcpFp8(chipset) && isa<Float8E5M2Type>(type));764}765 766/// Return true if `type` is the E4M3FN variant of an 8-bit float that is767/// supported by the `_fp8` instructions on the given `chipset`.768static bool typeIsExpectedFp8ForChipset(Chipset chipset, Type type) {769  return (chipset == kGfx942 && isa<Float8E4M3FNUZType>(type)) ||770         (hasOcpFp8(chipset) && isa<Float8E4M3FNType>(type));771}772 773/// Return the `rocdl` intrinsic corresponding to a MFMA operation `mfma`774/// if one exists. This includes checking to ensure the intrinsic is supported775/// on the architecture you are compiling for.776static std::optional<StringRef> mfmaOpToIntrinsic(MFMAOp mfma,777                                                  Chipset chipset) {778  uint32_t m = mfma.getM(), n = mfma.getN(), k = mfma.getK(),779           b = mfma.getBlocks();780  Type sourceElem = getElementTypeOrSelf(mfma.getSourceA().getType());781  Type destElem = getElementTypeOrSelf(mfma.getDestC().getType());782 783  if (sourceElem.isF32() && destElem.isF32()) {784    if (mfma.getReducePrecision() && chipset >= kGfx942) {785      if (m == 32 && n == 32 && k == 4 && b == 1)786        return ROCDL::mfma_f32_32x32x4_xf32::getOperationName();787      if (m == 16 && n == 16 && k == 8 && b == 1)788        return ROCDL::mfma_f32_16x16x8_xf32::getOperationName();789    }790    if (m == 32 && n == 32 && k == 1 && b == 2)791      return ROCDL::mfma_f32_32x32x1f32::getOperationName();792    if (m == 16 && n == 16 && k == 1 && b == 4)793      return ROCDL::mfma_f32_16x16x1f32::getOperationName();794    if (m == 4 && n == 4 && k == 1 && b == 16)795      return ROCDL::mfma_f32_4x4x1f32::getOperationName();796    if (m == 32 && n == 32 && k == 2 && b == 1)797      return ROCDL::mfma_f32_32x32x2f32::getOperationName();798    if (m == 16 && n == 16 && k == 4 && b == 1)799      return ROCDL::mfma_f32_16x16x4f32::getOperationName();800  }801 802  if (sourceElem.isF16() && destElem.isF32()) {803    if (chipset >= kGfx950) {804      if (m == 32 && n == 32 && k == 16 && b == 1)805        return ROCDL::mfma_f32_32x32x16_f16::getOperationName();806      if (m == 16 && n == 16 && k == 32 && b == 1)807        return ROCDL::mfma_f32_16x16x32_f16::getOperationName();808    }809    if (m == 32 && n == 32 && k == 4 && b == 2)810      return ROCDL::mfma_f32_32x32x4f16::getOperationName();811    if (m == 16 && n == 16 && k == 4 && b == 4)812      return ROCDL::mfma_f32_16x16x4f16::getOperationName();813    if (m == 4 && n == 4 && k == 4 && b == 16)814      return ROCDL::mfma_f32_4x4x4f16::getOperationName();815    if (m == 32 && n == 32 && k == 8 && b == 1)816      return ROCDL::mfma_f32_32x32x8f16::getOperationName();817    if (m == 16 && n == 16 && k == 16 && b == 1)818      return ROCDL::mfma_f32_16x16x16f16::getOperationName();819  }820 821  if (sourceElem.isBF16() && destElem.isF32()) {822    if (chipset >= kGfx950) {823      if (m == 32 && n == 32 && k == 16 && b == 1)824        return ROCDL::mfma_f32_32x32x16_bf16::getOperationName();825      if (m == 16 && n == 16 && k == 32 && b == 1)826        return ROCDL::mfma_f32_16x16x32_bf16::getOperationName();827    }828    if (chipset >= kGfx90a) {829      if (m == 32 && n == 32 && k == 4 && b == 2)830        return ROCDL::mfma_f32_32x32x4bf16_1k::getOperationName();831      if (m == 16 && n == 16 && k == 4 && b == 4)832        return ROCDL::mfma_f32_16x16x4bf16_1k::getOperationName();833      if (m == 4 && n == 4 && k == 4 && b == 16)834        return ROCDL::mfma_f32_4x4x4bf16_1k::getOperationName();835      if (m == 32 && n == 32 && k == 8 && b == 1)836        return ROCDL::mfma_f32_32x32x8bf16_1k::getOperationName();837      if (m == 16 && n == 16 && k == 16 && b == 1)838        return ROCDL::mfma_f32_16x16x16bf16_1k::getOperationName();839    }840    if (m == 32 && n == 32 && k == 2 && b == 2)841      return ROCDL::mfma_f32_32x32x2bf16::getOperationName();842    if (m == 16 && n == 16 && k == 2 && b == 4)843      return ROCDL::mfma_f32_16x16x2bf16::getOperationName();844    if (m == 4 && n == 4 && k == 2 && b == 16)845      return ROCDL::mfma_f32_4x4x2bf16::getOperationName();846    if (m == 32 && n == 32 && k == 4 && b == 1)847      return ROCDL::mfma_f32_32x32x4bf16::getOperationName();848    if (m == 16 && n == 16 && k == 8 && b == 1)849      return ROCDL::mfma_f32_16x16x8bf16::getOperationName();850  }851 852  if (sourceElem.isInteger(8) && destElem.isInteger(32)) {853    if (chipset >= kGfx950) {854      if (m == 32 && n == 32 && k == 32 && b == 1)855        return ROCDL::mfma_i32_32x32x32_i8::getOperationName();856      if (m == 16 && n == 16 && k == 64 && b == 1)857        return ROCDL::mfma_i32_16x16x64_i8::getOperationName();858    }859    if (m == 32 && n == 32 && k == 4 && b == 2)860      return ROCDL::mfma_i32_32x32x4i8::getOperationName();861    if (m == 16 && n == 16 && k == 4 && b == 4)862      return ROCDL::mfma_i32_16x16x4i8::getOperationName();863    if (m == 4 && n == 4 && k == 4 && b == 16)864      return ROCDL::mfma_i32_4x4x4i8::getOperationName();865    if (m == 32 && n == 32 && k == 8 && b == 1)866      return ROCDL::mfma_i32_32x32x8i8::getOperationName();867    if (m == 16 && n == 16 && k == 16 && b == 1)868      return ROCDL::mfma_i32_16x16x16i8::getOperationName();869    if (m == 32 && n == 32 && k == 16 && b == 1 && chipset >= kGfx942)870      return ROCDL::mfma_i32_32x32x16_i8::getOperationName();871    if (m == 16 && n == 16 && k == 32 && b == 1 && chipset >= kGfx942)872      return ROCDL::mfma_i32_16x16x32_i8::getOperationName();873  }874 875  if (sourceElem.isF64() && destElem.isF64() && chipset >= kGfx90a) {876    if (m == 16 && n == 16 && k == 4 && b == 1)877      return ROCDL::mfma_f64_16x16x4f64::getOperationName();878    if (m == 4 && n == 4 && k == 4 && b == 4)879      return ROCDL::mfma_f64_4x4x4f64::getOperationName();880  }881 882  if (destElem.isF32() && typeIsExpectedBf8ForChipset(chipset, sourceElem)) {883    // Known to be correct because there are no scalar f8 instructions and884    // because a length mismatch will have been caught by the verifier.885    Type sourceBElem =886        cast<VectorType>(mfma.getSourceB().getType()).getElementType();887    if (m == 16 && n == 16 && k == 32 && b == 1) {888      if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))889        return ROCDL::mfma_f32_16x16x32_bf8_bf8::getOperationName();890      if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))891        return ROCDL::mfma_f32_16x16x32_bf8_fp8::getOperationName();892    }893    if (m == 32 && n == 32 && k == 16 && b == 1) {894      if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))895        return ROCDL::mfma_f32_32x32x16_bf8_bf8::getOperationName();896      if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))897        return ROCDL::mfma_f32_32x32x16_bf8_fp8::getOperationName();898    }899  }900 901  if (destElem.isF32() && typeIsExpectedFp8ForChipset(chipset, sourceElem)) {902    Type sourceBElem =903        cast<VectorType>(mfma.getSourceB().getType()).getElementType();904    if (m == 16 && n == 16 && k == 32 && b == 1) {905      if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))906        return ROCDL::mfma_f32_16x16x32_fp8_bf8::getOperationName();907      if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))908        return ROCDL::mfma_f32_16x16x32_fp8_fp8::getOperationName();909    }910    if (m == 32 && n == 32 && k == 16 && b == 1) {911      if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))912        return ROCDL::mfma_f32_32x32x16_fp8_bf8::getOperationName();913      if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))914        return ROCDL::mfma_f32_32x32x16_fp8_fp8::getOperationName();915    }916  }917 918  return std::nullopt;919}920 921static std::optional<uint32_t> mfmaTypeSelectCode(Type mlirElemType) {922  return llvm::TypeSwitch<Type, std::optional<uint32_t>>(mlirElemType)923      .Case([](Float8E4M3FNType) { return 0u; })924      .Case([](Float8E5M2Type) { return 1u; })925      .Case([](Float6E2M3FNType) { return 2u; })926      .Case([](Float6E3M2FNType) { return 3u; })927      .Case([](Float4E2M1FNType) { return 4u; })928      .Default(std::nullopt);929}930 931/// If there is a scaled MFMA instruction for the input element types `aType`932/// and `bType`, output type `destType`, problem size M, N, K, and B (number of933/// blocks) on the given `chipset`, return a tuple consisting of the934/// OperationName of the intrinsic and the type codes that need to be passed to935/// that intrinsic. Note that this is also used to implement some un-scaled936/// MFMAs, since the compiler represents the ordinary instruction as a "scaled"937/// MFMA with a scale of 0.938static std::optional<std::tuple<StringRef, uint32_t, uint32_t>>939mfmaOpToScaledIntrinsic(Type aType, Type bType, Type destType, uint32_t m,940                        uint32_t n, uint32_t k, uint32_t b, Chipset chipset) {941  aType = getElementTypeOrSelf(aType);942  bType = getElementTypeOrSelf(bType);943  destType = getElementTypeOrSelf(destType);944 945  if (chipset < kGfx950)946    return std::nullopt;947  if (!isa<Float32Type>(destType))948    return std::nullopt;949 950  std::optional<uint32_t> aTypeCode = mfmaTypeSelectCode(aType);951  std::optional<uint32_t> bTypeCode = mfmaTypeSelectCode(bType);952  if (!aTypeCode || !bTypeCode)953    return std::nullopt;954 955  if (m == 32 && n == 32 && k == 64 && b == 1)956    return std::tuple{ROCDL::mfma_scale_f32_32x32x64_f8f6f4::getOperationName(),957                      *aTypeCode, *bTypeCode};958  if (m == 16 && n == 16 && k == 128 && b == 1)959    return std::tuple{960        ROCDL::mfma_scale_f32_16x16x128_f8f6f4::getOperationName(), *aTypeCode,961        *bTypeCode};962 963  return std::nullopt;964}965 966static std::optional<std::tuple<StringRef, uint32_t, uint32_t>>967mfmaOpToScaledIntrinsic(MFMAOp mfma, Chipset chipset) {968  return mfmaOpToScaledIntrinsic(969      mfma.getSourceA().getType(), mfma.getSourceB().getType(),970      mfma.getDestC().getType(), mfma.getM(), mfma.getN(), mfma.getK(),971      mfma.getBlocks(), chipset);972}973 974static std::optional<std::tuple<StringRef, uint32_t, uint32_t>>975mfmaOpToScaledIntrinsic(ScaledMFMAOp smfma, Chipset chipset) {976  return mfmaOpToScaledIntrinsic(smfma.getSourceA().getType(),977                                 smfma.getSourceB().getType(),978                                 smfma.getDestC().getType(), smfma.getM(),979                                 smfma.getN(), smfma.getK(), 1u, chipset);980}981 982/// Returns the `rocdl` intrinsic corresponding to a WMMA operation `wmma`983/// for RDNA3/4 architectures.984static std::optional<StringRef>985wmmaOpToIntrinsicRDNA(Type elemSourceType, Type elemBSourceType,986                      Type elemDestType, uint32_t k, bool isRDNA3) {987  using fp8 = Float8E4M3FNType;988  using bf8 = Float8E5M2Type;989 990  // Handle k == 16 for RDNA3/4.991  if (k == 16) {992    // Common patterns for RDNA3 and RDNA4.993    if (elemSourceType.isF16() && elemDestType.isF32())994      return ROCDL::wmma_f32_16x16x16_f16::getOperationName();995    if (elemSourceType.isBF16() && elemDestType.isF32())996      return ROCDL::wmma_f32_16x16x16_bf16::getOperationName();997    if (elemSourceType.isF16() && elemDestType.isF16())998      return ROCDL::wmma_f16_16x16x16_f16::getOperationName();999    if (elemSourceType.isBF16() && elemDestType.isBF16())1000      return ROCDL::wmma_bf16_16x16x16_bf16::getOperationName();1001    if (elemSourceType.isInteger(8) && elemDestType.isInteger(32))1002      return ROCDL::wmma_i32_16x16x16_iu8::getOperationName();1003 1004    // RDNA3 specific patterns.1005    if (isRDNA3) {1006      if (elemSourceType.isInteger(4) && elemDestType.isInteger(32))1007        return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();1008      return std::nullopt;1009    }1010 1011    // RDNA4 specific patterns (fp8/bf8).1012    if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType) &&1013        elemDestType.isF32())1014      return ROCDL::wmma_f32_16x16x16_fp8_fp8::getOperationName();1015    if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType) &&1016        elemDestType.isF32())1017      return ROCDL::wmma_f32_16x16x16_fp8_bf8::getOperationName();1018    if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType) &&1019        elemDestType.isF32())1020      return ROCDL::wmma_f32_16x16x16_bf8_bf8::getOperationName();1021    if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType) &&1022        elemDestType.isF32())1023      return ROCDL::wmma_f32_16x16x16_bf8_fp8::getOperationName();1024    if (elemSourceType.isInteger(4) && elemDestType.isInteger(32))1025      return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();1026 1027    return std::nullopt;1028  }1029 1030  // Handle k == 32 for RDNA4.1031  if (k == 32 && !isRDNA3) {1032    if (elemSourceType.isInteger(4) && elemDestType.isInteger(32))1033      return ROCDL::wmma_i32_16x16x32_iu4::getOperationName();1034  }1035 1036  return std::nullopt;1037}1038 1039/// Return the `rocdl` intrinsic corresponding to a WMMA operation `wmma`1040/// for the gfx1250 architecture.1041static std::optional<StringRef> wmmaOpToIntrinsicGfx1250(Type elemSourceType,1042                                                         Type elemBSourceType,1043                                                         Type elemDestType,1044                                                         uint32_t k) {1045  using fp8 = Float8E4M3FNType;1046  using bf8 = Float8E5M2Type;1047 1048  if (k == 4) {1049    if (elemSourceType.isF32() && elemDestType.isF32())1050      return ROCDL::wmma_f32_16x16x4_f32::getOperationName();1051 1052    return std::nullopt;1053  }1054 1055  if (k == 32) {1056    if (elemSourceType.isF16() && elemDestType.isF32())1057      return ROCDL::wmma_f32_16x16x32_f16::getOperationName();1058    if (elemSourceType.isBF16() && elemDestType.isF32())1059      return ROCDL::wmma_f32_16x16x32_bf16::getOperationName();1060    if (elemSourceType.isF16() && elemDestType.isF16())1061      return ROCDL::wmma_f16_16x16x32_f16::getOperationName();1062    if (elemSourceType.isBF16() && elemDestType.isBF16())1063      return ROCDL::wmma_bf16_16x16x32_bf16::getOperationName();1064 1065    return std::nullopt;1066  }1067 1068  if (k == 64) {1069    if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {1070      if (elemDestType.isF32())1071        return ROCDL::wmma_f32_16x16x64_fp8_fp8::getOperationName();1072      if (elemDestType.isF16())1073        return ROCDL::wmma_f16_16x16x64_fp8_fp8::getOperationName();1074    }1075    if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {1076      if (elemDestType.isF32())1077        return ROCDL::wmma_f32_16x16x64_fp8_bf8::getOperationName();1078      if (elemDestType.isF16())1079        return ROCDL::wmma_f16_16x16x64_fp8_bf8::getOperationName();1080    }1081    if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {1082      if (elemDestType.isF32())1083        return ROCDL::wmma_f32_16x16x64_bf8_bf8::getOperationName();1084      if (elemDestType.isF16())1085        return ROCDL::wmma_f16_16x16x64_bf8_bf8::getOperationName();1086    }1087    if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {1088      if (elemDestType.isF32())1089        return ROCDL::wmma_f32_16x16x64_bf8_fp8::getOperationName();1090      if (elemDestType.isF16())1091        return ROCDL::wmma_f16_16x16x64_bf8_fp8::getOperationName();1092    }1093    if (elemSourceType.isInteger(8) && elemDestType.isInteger(32))1094      return ROCDL::wmma_i32_16x16x64_iu8::getOperationName();1095 1096    return std::nullopt;1097  }1098 1099  if (k == 128) {1100    if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {1101      if (elemDestType.isF32())1102        return ROCDL::wmma_f32_16x16x128_fp8_fp8::getOperationName();1103      if (elemDestType.isF16())1104        return ROCDL::wmma_f16_16x16x128_fp8_fp8::getOperationName();1105    }1106    if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {1107      if (elemDestType.isF32())1108        return ROCDL::wmma_f32_16x16x128_fp8_bf8::getOperationName();1109      if (elemDestType.isF16())1110        return ROCDL::wmma_f16_16x16x128_fp8_bf8::getOperationName();1111    }1112    if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {1113      if (elemDestType.isF32())1114        return ROCDL::wmma_f32_16x16x128_bf8_bf8::getOperationName();1115      if (elemDestType.isF16())1116        return ROCDL::wmma_f16_16x16x128_bf8_bf8::getOperationName();1117    }1118    if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {1119      if (elemDestType.isF32())1120        return ROCDL::wmma_f32_16x16x128_bf8_fp8::getOperationName();1121      if (elemDestType.isF16())1122        return ROCDL::wmma_f16_16x16x128_bf8_fp8::getOperationName();1123    }1124 1125    return std::nullopt;1126  }1127 1128  return std::nullopt;1129}1130 1131/// Returns the `rocdl` intrinsic corresponding to a WMMA operation `wmma`1132/// if one exists. This includes checking to ensure the intrinsic is supported1133/// on the architecture you are compiling for.1134static std::optional<StringRef> wmmaOpToIntrinsic(WMMAOp wmma,1135                                                  Chipset chipset) {1136  auto sourceVectorType = cast<VectorType>(wmma.getSourceA().getType());1137  auto sourceBVectorType = cast<VectorType>(wmma.getSourceB().getType());1138  auto destVectorType = cast<VectorType>(wmma.getDestC().getType());1139  Type elemSourceType = sourceVectorType.getElementType();1140  Type elemBSourceType = sourceBVectorType.getElementType();1141  Type elemDestType = destVectorType.getElementType();1142 1143  const uint32_t k = wmma.getK();1144  const bool isRDNA3 = chipset.majorVersion == 11;1145  const bool isRDNA4 = chipset.majorVersion == 12 && chipset.minorVersion == 0;1146 1147  // Handle RDNA3 and RDNA4.1148  if (isRDNA3 || isRDNA4)1149    return wmmaOpToIntrinsicRDNA(elemSourceType, elemBSourceType, elemDestType,1150                                 k, isRDNA3);1151 1152  // Handle gfx1250.1153  if (chipset == kGfx1250)1154    return wmmaOpToIntrinsicGfx1250(elemSourceType, elemBSourceType,1155                                    elemDestType, k);1156 1157  return std::nullopt;1158}1159 1160namespace {1161struct MFMAOpLowering : public ConvertOpToLLVMPattern<MFMAOp> {1162  MFMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)1163      : ConvertOpToLLVMPattern<MFMAOp>(converter), chipset(chipset) {}1164 1165  Chipset chipset;1166 1167  LogicalResult1168  matchAndRewrite(MFMAOp op, MFMAOpAdaptor adaptor,1169                  ConversionPatternRewriter &rewriter) const override {1170    Location loc = op.getLoc();1171    Type outType = typeConverter->convertType(op.getDestD().getType());1172    Type intrinsicOutType = outType;1173    if (auto outVecType = dyn_cast<VectorType>(outType))1174      if (outVecType.getElementType().isBF16())1175        intrinsicOutType = outVecType.clone(rewriter.getI16Type());1176 1177    if (chipset.majorVersion != 9 || chipset < kGfx908)1178      return op->emitOpError("MFMA only supported on gfx908+");1179    uint32_t getBlgpField = static_cast<uint32_t>(op.getBlgp());1180    if (op.getNegateA() || op.getNegateB() || op.getNegateC()) {1181      if (chipset < kGfx942)1182        return op.emitOpError("negation unsupported on older than gfx942");1183      getBlgpField |=1184          op.getNegateA() | (op.getNegateB() << 1) | (op.getNegateC() << 2);1185    }1186    std::optional<StringRef> maybeIntrinsic = mfmaOpToIntrinsic(op, chipset);1187    std::optional<std::tuple<StringRef, uint32_t, uint32_t>>1188        maybeScaledIntrinsic = mfmaOpToScaledIntrinsic(op, chipset);1189    if (!maybeIntrinsic.has_value() && !maybeScaledIntrinsic.has_value())1190      return op.emitOpError("no intrinsic matching MFMA size on given chipset");1191 1192    bool isScaled =1193        !maybeIntrinsic.has_value() && maybeScaledIntrinsic.has_value();1194    if (isScaled &&1195        (adaptor.getAbid() > 0 || getBlgpField > 0 || op.getCbsz() > 0)) {1196      return op.emitOpError(1197          "non-default abid, blgp, and cbsz aren't supported on MFMAs that can "1198          "be scaled as those fields are used for type information");1199    }1200 1201    StringRef intrinsicName =1202        isScaled ? std::get<0>(*maybeScaledIntrinsic) : *maybeIntrinsic;1203    // Determine if we can use bf16 in the intrinsic. Newer MFMAs in gfx950+1204    // allows bf16 as the input. For reference check IntrinsicsAMDGPU.td file.1205    bool allowBf16 = [&]() {1206      if (chipset < kGfx950)1207        return false;1208      if (isScaled)1209        return true;1210      return intrinsicName.contains("16x16x32.bf16") ||1211             intrinsicName.contains("32x32x16.bf16");1212    }();1213    OperationState loweredOp(loc, intrinsicName);1214    loweredOp.addTypes(intrinsicOutType);1215    loweredOp.addOperands({convertMFMAVectorOperand(1216                               rewriter, loc, adaptor.getSourceA(), allowBf16),1217                           convertMFMAVectorOperand(1218                               rewriter, loc, adaptor.getSourceB(), allowBf16),1219                           adaptor.getDestC()});1220    if (isScaled) {1221      Value zero = createI32Constant(rewriter, loc, 0);1222      auto [_scaledName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;1223      loweredOp.addOperands({createI32Constant(rewriter, loc, aTypeCode),1224                             createI32Constant(rewriter, loc, bTypeCode),1225                             /*scale A byte=*/zero, /*scale A=*/zero,1226                             /*scale B byte=*/zero, /*scale B=*/zero});1227    } else {1228      loweredOp.addOperands({createI32Constant(rewriter, loc, op.getCbsz()),1229                             createI32Constant(rewriter, loc, op.getAbid()),1230                             createI32Constant(rewriter, loc, getBlgpField)});1231    };1232    Value lowered = rewriter.create(loweredOp)->getResult(0);1233    if (outType != intrinsicOutType)1234      lowered = LLVM::BitcastOp::create(rewriter, loc, outType, lowered);1235    rewriter.replaceOp(op, lowered);1236    return success();1237  }1238};1239 1240struct ScaledMFMAOpLowering : public ConvertOpToLLVMPattern<ScaledMFMAOp> {1241  ScaledMFMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)1242      : ConvertOpToLLVMPattern(converter), chipset(chipset) {}1243 1244  Chipset chipset;1245 1246  LogicalResult1247  matchAndRewrite(ScaledMFMAOp op, ScaledMFMAOpAdaptor adaptor,1248                  ConversionPatternRewriter &rewriter) const override {1249    Location loc = op.getLoc();1250    Type intrinsicOutType = typeConverter->convertType(op.getDestD().getType());1251 1252    if (chipset.majorVersion != 9 || chipset < kGfx950)1253      return op->emitOpError("scaled MFMA only supported on gfx908+");1254    std::optional<std::tuple<StringRef, uint32_t, uint32_t>>1255        maybeScaledIntrinsic = mfmaOpToScaledIntrinsic(op, chipset);1256    if (!maybeScaledIntrinsic.has_value())1257      return op.emitOpError(1258          "no intrinsic matching scaled MFMA size on given chipset");1259 1260    auto [intrinsicName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;1261    OperationState loweredOp(loc, intrinsicName);1262    loweredOp.addTypes(intrinsicOutType);1263    loweredOp.addOperands(1264        {convertMFMAVectorOperand(rewriter, loc, adaptor.getSourceA()),1265         convertMFMAVectorOperand(rewriter, loc, adaptor.getSourceB()),1266         adaptor.getDestC()});1267    Value scalesIdxA =1268        createI32Constant(rewriter, loc, adaptor.getScalesIdxA());1269    Value scalesIdxB =1270        createI32Constant(rewriter, loc, adaptor.getScalesIdxB());1271    loweredOp.addOperands(1272        {createI32Constant(rewriter, loc, aTypeCode),1273         createI32Constant(rewriter, loc, bTypeCode),1274         /*scales idx A=*/scalesIdxA,1275         /*scales A*/1276         castMFMAScaleOperand(rewriter, loc, adaptor.getScalesA()),1277         /*scales idx B=*/scalesIdxB,1278         /*scales B*/1279         castMFMAScaleOperand(rewriter, loc, adaptor.getScalesB())});1280    Value lowered = rewriter.create(loweredOp)->getResult(0);1281    rewriter.replaceOp(op, lowered);1282    return success();1283  }1284};1285 1286struct WMMAOpLowering : public ConvertOpToLLVMPattern<WMMAOp> {1287  WMMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)1288      : ConvertOpToLLVMPattern<WMMAOp>(converter), chipset(chipset) {}1289 1290  Chipset chipset;1291 1292  LogicalResult1293  matchAndRewrite(WMMAOp op, WMMAOpAdaptor adaptor,1294                  ConversionPatternRewriter &rewriter) const override {1295    Location loc = op.getLoc();1296    auto outType =1297        typeConverter->convertType<VectorType>(op.getDestD().getType());1298    if (!outType)1299      return rewriter.notifyMatchFailure(op, "type conversion failed");1300 1301    if (chipset.majorVersion != 11 && chipset.majorVersion != 12)1302      return op->emitOpError("WMMA only supported on gfx11 and gfx12");1303 1304    bool isGFX1250 = chipset >= kGfx1250;1305 1306    // The WMMA operations represent vectors of bf16s as vectors of i16s1307    // (except on gfx1250), so we need to bitcast bfloats to i16 and then1308    // bitcast them back.1309    auto aType = cast<VectorType>(adaptor.getSourceA().getType());1310    auto bType = cast<VectorType>(adaptor.getSourceB().getType());1311    auto destCType = cast<VectorType>(adaptor.getDestC().getType());1312    bool castAToI16 = aType.getElementType().isBF16() && !isGFX1250;1313    bool castBToI16 = bType.getElementType().isBF16() && !isGFX1250;1314    bool castDestCToI16 = destCType.getElementType().isBF16() && !isGFX1250;1315    bool castOutToI16 = outType.getElementType().isBF16() && !isGFX1250;1316    VectorType rawOutType = outType;1317    if (castOutToI16)1318      rawOutType = outType.clone(rewriter.getI16Type());1319    Value a = adaptor.getSourceA();1320    if (castAToI16)1321      a = LLVM::BitcastOp::create(rewriter, loc,1322                                  aType.clone(rewriter.getI16Type()), a);1323    Value b = adaptor.getSourceB();1324    if (castBToI16)1325      b = LLVM::BitcastOp::create(rewriter, loc,1326                                  bType.clone(rewriter.getI16Type()), b);1327    Value destC = adaptor.getDestC();1328    if (castDestCToI16)1329      destC = LLVM::BitcastOp::create(1330          rewriter, loc, destCType.clone(rewriter.getI16Type()), destC);1331 1332    std::optional<StringRef> maybeIntrinsic = wmmaOpToIntrinsic(op, chipset);1333 1334    if (!maybeIntrinsic.has_value())1335      return op.emitOpError("no intrinsic matching WMMA on the given chipset");1336 1337    if (chipset.majorVersion >= 12 && op.getSubwordOffset() != 0)1338      return op.emitOpError("subwordOffset not supported on gfx12+");1339 1340    SmallVector<Value, 4> operands;1341    SmallVector<NamedAttribute, 4> attrs;1342    wmmaPushInputOperand(rewriter, loc, typeConverter, op.getUnsignedA(), a,1343                         op.getSourceA(), operands, attrs, "signA");1344    wmmaPushInputOperand(rewriter, loc, typeConverter, op.getUnsignedB(), b,1345                         op.getSourceB(), operands, attrs, "signB");1346    wmmaPushOutputOperand(rewriter, loc, typeConverter, destC,1347                          op.getSubwordOffset(), op.getClamp(), operands,1348                          attrs);1349 1350    OperationState loweredOp(loc, *maybeIntrinsic);1351    loweredOp.addTypes(rawOutType);1352    loweredOp.addOperands(operands);1353    loweredOp.addAttributes(attrs);1354    Operation *lowered = rewriter.create(loweredOp);1355 1356    Operation *maybeCastBack = lowered;1357    if (rawOutType != outType)1358      maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,1359                                              lowered->getResult(0));1360    rewriter.replaceOp(op, maybeCastBack->getResults());1361 1362    return success();1363  }1364};1365 1366struct TransposeLoadOpLowering1367    : public ConvertOpToLLVMPattern<TransposeLoadOp> {1368  TransposeLoadOpLowering(const LLVMTypeConverter &converter, Chipset chipset)1369      : ConvertOpToLLVMPattern<TransposeLoadOp>(converter), chipset(chipset) {}1370 1371  Chipset chipset;1372 1373  LogicalResult1374  matchAndRewrite(TransposeLoadOp op, TransposeLoadOpAdaptor adaptor,1375                  ConversionPatternRewriter &rewriter) const override {1376    if (chipset != kGfx950)1377      return op.emitOpError("Non-gfx950 chipset not supported");1378 1379    Location loc = op.getLoc();1380    auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());1381 1382    // Elements in subbyte memrefs are stored non-contiguously,1383    // reject if source is sub-byte memref. Use emulated memrefs instead.1384    size_t srcElementSize =1385        srcMemRefType.getElementType().getIntOrFloatBitWidth();1386    if (srcElementSize < 8)1387      return op.emitOpError("Expect source memref to have at least 8 bits "1388                            "element size, got ")1389             << srcElementSize;1390 1391    auto resultType = cast<VectorType>(op.getResult().getType());1392    Value srcPtr =1393        getStridedElementPtr(rewriter, loc, srcMemRefType, adaptor.getSrc(),1394                             (adaptor.getSrcIndices()));1395 1396    size_t numElements = resultType.getNumElements();1397    size_t elementTypeSize =1398        resultType.getElementType().getIntOrFloatBitWidth();1399 1400    // ROCDL transpose load intrinsics return vectors of 32-bit integers, if1401    // the element size is smaller than 16 bits.1402    Type rocdlResultType = VectorType::get((numElements * elementTypeSize) / 32,1403                                           rewriter.getIntegerType(32));1404    Type llvmResultType = typeConverter->convertType(resultType);1405 1406    switch (elementTypeSize) {1407    case 4: {1408      assert(numElements == 16);1409      auto rocdlOp = ROCDL::ds_read_tr4_b64::create(rewriter, loc,1410                                                    rocdlResultType, srcPtr);1411      rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);1412      break;1413    }1414    case 6: {1415      assert(numElements == 16);1416      auto rocdlOp = ROCDL::ds_read_tr6_b96::create(rewriter, loc,1417                                                    rocdlResultType, srcPtr);1418      rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);1419      break;1420    }1421    case 8: {1422      assert(numElements == 8);1423      auto rocdlOp = ROCDL::ds_read_tr8_b64::create(rewriter, loc,1424                                                    rocdlResultType, srcPtr);1425      rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);1426      break;1427    }1428    case 16: {1429      assert(numElements == 4);1430      rewriter.replaceOpWithNewOp<ROCDL::ds_read_tr16_b64>(op, llvmResultType,1431                                                           srcPtr);1432      break;1433    }1434    default:1435      return op.emitOpError("Unsupported element size for transpose load");1436    }1437    return success();1438  }1439};1440 1441struct GatherToLDSOpLowering : public ConvertOpToLLVMPattern<GatherToLDSOp> {1442  GatherToLDSOpLowering(const LLVMTypeConverter &converter, Chipset chipset)1443      : ConvertOpToLLVMPattern<GatherToLDSOp>(converter), chipset(chipset) {}1444 1445  Chipset chipset;1446 1447  LogicalResult1448  matchAndRewrite(GatherToLDSOp op, GatherToLDSOpAdaptor adaptor,1449                  ConversionPatternRewriter &rewriter) const override {1450    if (chipset.majorVersion < 9 || chipset.majorVersion > 10)1451      return op.emitOpError("pre-gfx9 and post-gfx10 not supported");1452 1453    Location loc = op.getLoc();1454 1455    auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());1456    auto dstMemRefType = cast<MemRefType>(op.getDst().getType());1457 1458    // TODO: instead of only transfering one element per thread, we could1459    // augment it to transfer multiple elements per thread by issuing multiple1460    // `global_load_lds` instructions.1461    Type transferType = op.getTransferType();1462    int loadWidth = [&]() -> int {1463      if (auto transferVectorType = dyn_cast<VectorType>(transferType)) {1464        return (transferVectorType.getNumElements() *1465                transferVectorType.getElementTypeBitWidth()) /1466               8;1467      }1468      return transferType.getIntOrFloatBitWidth() / 8;1469    }();1470 1471    // Currently only 1, 2, 4, 12 and 16 byte loads are supported.1472    if (!llvm::is_contained({1, 2, 4, 12, 16}, loadWidth))1473      return op.emitOpError("chipset unsupported element size");1474 1475    if (chipset != kGfx950 && llvm::is_contained({12, 16}, loadWidth))1476      return op.emitOpError("Gather to LDS instructions with 12-byte and "1477                            "16-byte load widths are only supported on gfx950");1478 1479    Value srcPtr =1480        getStridedElementPtr(rewriter, loc, srcMemRefType, adaptor.getSrc(),1481                             (adaptor.getSrcIndices()));1482    Value dstPtr =1483        getStridedElementPtr(rewriter, loc, dstMemRefType, adaptor.getDst(),1484                             (adaptor.getDstIndices()));1485 1486    rewriter.replaceOpWithNewOp<ROCDL::LoadToLDSOp>(1487        op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),1488        /*offset=*/rewriter.getI32IntegerAttr(0),1489        /*aux=*/rewriter.getI32IntegerAttr(0), ArrayAttr{}, ArrayAttr{},1490        ArrayAttr{});1491 1492    return success();1493  }1494};1495 1496namespace {1497struct ExtPackedFp8OpLowering final1498    : public ConvertOpToLLVMPattern<ExtPackedFp8Op> {1499  ExtPackedFp8OpLowering(const LLVMTypeConverter &converter, Chipset chipset)1500      : ConvertOpToLLVMPattern<amdgpu::ExtPackedFp8Op>(converter),1501        chipset(chipset) {}1502  Chipset chipset;1503 1504  LogicalResult1505  matchAndRewrite(ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,1506                  ConversionPatternRewriter &rewriter) const override;1507};1508 1509struct ScaledExtPacked816OpLowering final1510    : public ConvertOpToLLVMPattern<ScaledExtPacked816Op> {1511  ScaledExtPacked816OpLowering(const LLVMTypeConverter &converter,1512                               Chipset chipset)1513      : ConvertOpToLLVMPattern<amdgpu::ScaledExtPacked816Op>(converter),1514        chipset(chipset) {}1515  Chipset chipset;1516 1517  LogicalResult1518  matchAndRewrite(ScaledExtPacked816Op op, ScaledExtPacked816OpAdaptor adaptor,1519                  ConversionPatternRewriter &rewriter) const override;1520};1521 1522struct PackedTrunc2xFp8OpLowering final1523    : public ConvertOpToLLVMPattern<PackedTrunc2xFp8Op> {1524  PackedTrunc2xFp8OpLowering(const LLVMTypeConverter &converter,1525                             Chipset chipset)1526      : ConvertOpToLLVMPattern<amdgpu::PackedTrunc2xFp8Op>(converter),1527        chipset(chipset) {}1528  Chipset chipset;1529 1530  LogicalResult1531  matchAndRewrite(PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,1532                  ConversionPatternRewriter &rewriter) const override;1533};1534 1535struct PackedStochRoundFp8OpLowering final1536    : public ConvertOpToLLVMPattern<PackedStochRoundFp8Op> {1537  PackedStochRoundFp8OpLowering(const LLVMTypeConverter &converter,1538                                Chipset chipset)1539      : ConvertOpToLLVMPattern<amdgpu::PackedStochRoundFp8Op>(converter),1540        chipset(chipset) {}1541  Chipset chipset;1542 1543  LogicalResult1544  matchAndRewrite(PackedStochRoundFp8Op op,1545                  PackedStochRoundFp8OpAdaptor adaptor,1546                  ConversionPatternRewriter &rewriter) const override;1547};1548 1549struct ScaledExtPackedOpLowering final1550    : public ConvertOpToLLVMPattern<ScaledExtPackedOp> {1551  ScaledExtPackedOpLowering(const LLVMTypeConverter &converter, Chipset chipset)1552      : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedOp>(converter),1553        chipset(chipset) {}1554  Chipset chipset;1555 1556  LogicalResult1557  matchAndRewrite(ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,1558                  ConversionPatternRewriter &rewriter) const override;1559};1560 1561struct PackedScaledTruncOpLowering final1562    : public ConvertOpToLLVMPattern<PackedScaledTruncOp> {1563  PackedScaledTruncOpLowering(const LLVMTypeConverter &converter,1564                              Chipset chipset)1565      : ConvertOpToLLVMPattern<amdgpu::PackedScaledTruncOp>(converter),1566        chipset(chipset) {}1567  Chipset chipset;1568 1569  LogicalResult1570  matchAndRewrite(PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,1571                  ConversionPatternRewriter &rewriter) const override;1572};1573 1574} // end namespace1575 1576LogicalResult ExtPackedFp8OpLowering::matchAndRewrite(1577    ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,1578    ConversionPatternRewriter &rewriter) const {1579  Location loc = op.getLoc();1580  if (!(chipset == kGfx942 || hasOcpFp8(chipset)))1581    return rewriter.notifyMatchFailure(1582        loc, "Fp8 conversion instructions are not available on target "1583             "architecture and their emulation is not implemented");1584  Type v4i8 =1585      getTypeConverter()->convertType(VectorType::get(4, rewriter.getI8Type()));1586  Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());1587  Type f32 = getTypeConverter()->convertType(op.getResult().getType());1588 1589  Value source = adaptor.getSource();1590  auto sourceVecType = dyn_cast<VectorType>(op.getSource().getType());1591  auto resultVecType = dyn_cast<VectorType>(op.getResult().getType());1592  Type sourceElemType = getElementTypeOrSelf(op.getSource());1593  // Extend to a v4i81594  if (!sourceVecType || sourceVecType.getNumElements() < 4) {1595    Value longVec = LLVM::UndefOp::create(rewriter, loc, v4i8);1596    if (!sourceVecType) {1597      longVec = LLVM::InsertElementOp::create(1598          rewriter, loc, longVec, source, createI32Constant(rewriter, loc, 0));1599    } else {1600      for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {1601        Value idx = createI32Constant(rewriter, loc, i);1602        Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);1603        longVec =1604            LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);1605      }1606    }1607    source = longVec;1608  }1609  Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);1610  if (resultVecType) {1611    if (typeIsExpectedBf8ForChipset(chipset, sourceElemType)) {1612      rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Bf8Op>(op, f32, i32Source,1613                                                        op.getIndex());1614    } else if (typeIsExpectedFp8ForChipset(chipset, sourceElemType)) {1615      rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Fp8Op>(op, f32, i32Source,1616                                                        op.getIndex());1617    }1618  } else {1619    if (typeIsExpectedBf8ForChipset(chipset, sourceElemType)) {1620      rewriter.replaceOpWithNewOp<ROCDL::CvtF32Bf8Op>(op, f32, i32Source,1621                                                      op.getIndex());1622    } else if (typeIsExpectedFp8ForChipset(chipset, sourceElemType)) {1623      rewriter.replaceOpWithNewOp<ROCDL::CvtF32Fp8Op>(op, f32, i32Source,1624                                                      op.getIndex());1625    }1626  }1627  return success();1628}1629 1630int32_t getScaleSel(int32_t blockSize, unsigned bitWidth,1631                    int32_t firstScaleLane, int32_t firstScaleByte) {1632  // When lowering amdgpu.scaled_ext_packed816 to rocdl.cvt.scale.pk*.f*.f*1633  // operations, the attributes blockSize, sourceType, firstScaleLane and1634  // firstScaleByte are merged into a single attribute scaleSel. This is how1635  // those values are merged together.1636  assert(llvm::is_contained({16, 32}, blockSize));1637  assert(llvm::is_contained(llvm::ArrayRef<unsigned>{4, 6, 8}, bitWidth));1638 1639  const bool is_fp8 = bitWidth == 8;1640  const bool is_block_16 = blockSize == 16;1641 1642  if (!is_fp8) {1643    int bit_0 = is_block_16;1644    assert(llvm::is_contained({0, 1, 2}, firstScaleByte));1645    int bit_1 = (firstScaleByte == 2) << 1;1646    assert(llvm::is_contained({0, 1}, firstScaleLane));1647    int bit_2 = firstScaleLane << 2;1648    return bit_2 | bit_1 | bit_0;1649  }1650 1651  int bit_0 = is_block_16;1652  // firstScaleByte is guaranteed to be defined by two bits.1653  assert(llvm::is_contained({0, 1, 2, 3}, firstScaleByte));1654  int bit_2_and_1 = firstScaleByte << 1;1655  assert(llvm::is_contained({0, 1}, firstScaleLane));1656  int bit_3 = firstScaleLane << 3;1657  int bits = bit_3 | bit_2_and_1 | bit_0;1658  // These are invalid cases.1659  assert(!llvm::is_contained(1660      {0b0011, 0b0101, 0b0111, 0b1000, 0b1001, 0b1011, 0b1111}, bits));1661  return bits;1662}1663 1664static std::optional<StringRef>1665scaledExtPacked816ToIntrinsic(Type srcElemType, Type destElemType) {1666  using fp4 = Float4E2M1FNType;1667  using fp8 = Float8E4M3FNType;1668  using bf8 = Float8E5M2Type;1669  using fp6 = Float6E2M3FNType;1670  using bf6 = Float6E3M2FNType;1671  if (isa<fp4>(srcElemType)) {1672    if (destElemType.isF16())1673      return ROCDL::CvtPkScalePk8F16Fp4Op::getOperationName();1674    if (destElemType.isBF16())1675      return ROCDL::CvtPkScalePk8Bf16Fp4Op::getOperationName();1676    if (destElemType.isF32())1677      return ROCDL::CvtPkScalePk8F32Fp4Op::getOperationName();1678    return std::nullopt;1679  }1680  if (isa<fp8>(srcElemType)) {1681    if (destElemType.isF16())1682      return ROCDL::CvtPkScalePk8F16Fp8Op::getOperationName();1683    if (destElemType.isBF16())1684      return ROCDL::CvtPkScalePk8Bf16Fp8Op::getOperationName();1685    if (destElemType.isF32())1686      return ROCDL::CvtPkScalePk8F32Fp8Op::getOperationName();1687    return std::nullopt;1688  }1689  if (isa<bf8>(srcElemType)) {1690    if (destElemType.isF16())1691      return ROCDL::CvtPkScalePk8F16Bf8Op::getOperationName();1692    if (destElemType.isBF16())1693      return ROCDL::CvtPkScalePk8Bf16Bf8Op::getOperationName();1694    if (destElemType.isF32())1695      return ROCDL::CvtPkScalePk8F32Bf8Op::getOperationName();1696    return std::nullopt;1697  }1698  if (isa<fp6>(srcElemType)) {1699    if (destElemType.isF16())1700      return ROCDL::CvtPkScalePk16F16Fp6Op::getOperationName();1701    if (destElemType.isBF16())1702      return ROCDL::CvtPkScalePk16Bf16Fp6Op::getOperationName();1703    if (destElemType.isF32())1704      return ROCDL::CvtPkScalePk16F32Fp6Op::getOperationName();1705    return std::nullopt;1706  }1707  if (isa<bf6>(srcElemType)) {1708    if (destElemType.isF16())1709      return ROCDL::CvtPkScalePk16F16Bf6Op::getOperationName();1710    if (destElemType.isBF16())1711      return ROCDL::CvtPkScalePk16Bf16Bf6Op::getOperationName();1712    if (destElemType.isF32())1713      return ROCDL::CvtPkScalePk16F32Bf6Op::getOperationName();1714    return std::nullopt;1715  }1716  llvm_unreachable("invalid combination of element types for packed conversion "1717                   "instructions");1718}1719 1720LogicalResult ScaledExtPacked816OpLowering::matchAndRewrite(1721    ScaledExtPacked816Op op, ScaledExtPacked816OpAdaptor adaptor,1722    ConversionPatternRewriter &rewriter) const {1723  using fp4 = Float4E2M1FNType;1724  using fp8 = Float8E4M3FNType;1725  using bf8 = Float8E5M2Type;1726  using fp6 = Float6E2M3FNType;1727  using bf6 = Float6E3M2FNType;1728  Location loc = op.getLoc();1729  if (chipset != kGfx1250) {1730    return rewriter.notifyMatchFailure(1731        loc,1732        "Scaled fp packed conversion instructions are not available on target "1733        "architecture and their emulation is not implemented");1734  }1735  int32_t firstScaleLane = op.getFirstScaleLane();1736  int32_t firstScaleByte = op.getFirstScaleByte();1737  int32_t blockSize = op.getBlockSize();1738  auto sourceType = cast<VectorType>(op.getSource().getType());1739  auto srcElemType = cast<FloatType>(sourceType.getElementType());1740  unsigned bitWidth = srcElemType.getWidth();1741 1742  auto targetType = cast<VectorType>(op.getResult().getType());1743  auto destElemType = cast<FloatType>(targetType.getElementType());1744 1745  IntegerType i32 = rewriter.getI32Type();1746  Value source = adaptor.getSource();1747  Type llvmResultType = typeConverter->convertType(op.getResult().getType());1748  Type packedType = nullptr;1749  if (isa<fp4>(srcElemType)) {1750    packedType = i32;1751    packedType = getTypeConverter()->convertType(packedType);1752  } else if (isa<fp8, bf8>(srcElemType)) {1753    packedType = VectorType::get(2, i32);1754    packedType = getTypeConverter()->convertType(packedType);1755  } else if (isa<fp6, bf6>(srcElemType)) {1756    packedType = VectorType::get(3, i32);1757    packedType = getTypeConverter()->convertType(packedType);1758  } else {1759    llvm_unreachable("invalid element type for packed scaled ext");1760  }1761 1762  if (!packedType || !llvmResultType) {1763    return rewriter.notifyMatchFailure(op, "type conversion failed");1764  }1765 1766  std::optional<StringRef> maybeIntrinsic =1767      scaledExtPacked816ToIntrinsic(srcElemType, destElemType);1768  if (!maybeIntrinsic.has_value())1769    return op.emitOpError(1770        "no intrinsic matching packed scaled conversion on the given chipset");1771 1772  int32_t scaleSel =1773      getScaleSel(blockSize, bitWidth, firstScaleLane, firstScaleByte);1774  Value castedScale =1775      LLVM::BitcastOp::create(rewriter, loc, i32, adaptor.getScale());1776  Value castedSource =1777      LLVM::BitcastOp::create(rewriter, loc, packedType, source);1778 1779  OperationState loweredOp(loc, *maybeIntrinsic);1780  loweredOp.addTypes({llvmResultType});1781  loweredOp.addOperands({castedSource, castedScale});1782 1783  SmallVector<NamedAttribute, 1> attrs;1784  attrs.push_back(1785      NamedAttribute("scaleSel", rewriter.getI32IntegerAttr(scaleSel)));1786 1787  loweredOp.addAttributes(attrs);1788  Operation *lowered = rewriter.create(loweredOp);1789  rewriter.replaceOp(op, lowered);1790 1791  return success();1792}1793 1794LogicalResult ScaledExtPackedOpLowering::matchAndRewrite(1795    ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,1796    ConversionPatternRewriter &rewriter) const {1797  Location loc = op.getLoc();1798  if (chipset != kGfx950)1799    return rewriter.notifyMatchFailure(1800        loc, "Scaled fp conversion instructions are not available on target "1801             "architecture and their emulation is not implemented");1802  Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());1803 1804  Value source = adaptor.getSource();1805  Value scale = adaptor.getScale();1806 1807  VectorType sourceVecType = cast<VectorType>(op.getSource().getType());1808  Type sourceElemType = sourceVecType.getElementType();1809  VectorType destVecType = cast<VectorType>(op.getResult().getType());1810  Type destElemType = destVecType.getElementType();1811 1812  VectorType packedVecType;1813  if (isa<Float8E5M2Type, Float8E4M3FNType>(sourceElemType)) {1814    VectorType v4i8 = VectorType::get(4, rewriter.getI8Type());1815    packedVecType = cast<VectorType>(getTypeConverter()->convertType(v4i8));1816  } else if (isa<Float4E2M1FNType>(sourceElemType)) {1817    VectorType v8i4 = VectorType::get(8, rewriter.getI4Type());1818    packedVecType = cast<VectorType>(getTypeConverter()->convertType(v8i4));1819  } else {1820    llvm_unreachable("invalid element type for scaled ext");1821  }1822 1823  // Extend to a packedVectorType1824  if (sourceVecType.getNumElements() < packedVecType.getNumElements()) {1825    Value longVec = LLVM::ZeroOp::create(rewriter, loc, packedVecType);1826    if (!sourceVecType) {1827      longVec = LLVM::InsertElementOp::create(1828          rewriter, loc, longVec, source, createI32Constant(rewriter, loc, 0));1829    } else {1830      for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {1831        Value idx = createI32Constant(rewriter, loc, i);1832        Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);1833        longVec =1834            LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);1835      }1836    }1837    source = longVec;1838  }1839  Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);1840 1841  if (isa<Float8E5M2Type>(sourceElemType) && destElemType.isF32())1842    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Bf8Op>(1843        op, destVecType, i32Source, scale, op.getIndex());1844  else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.isF16())1845    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Bf8Op>(1846        op, destVecType, i32Source, scale, op.getIndex());1847  else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.isBF16())1848    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Bf8Op>(1849        op, destVecType, i32Source, scale, op.getIndex());1850  else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.isF32())1851    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp8Op>(1852        op, destVecType, i32Source, scale, op.getIndex());1853  else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.isF16())1854    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp8Op>(1855        op, destVecType, i32Source, scale, op.getIndex());1856  else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.isBF16())1857    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp8Op>(1858        op, destVecType, i32Source, scale, op.getIndex());1859  else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.isF32())1860    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp4Op>(1861        op, destVecType, i32Source, scale, op.getIndex());1862  else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.isF16())1863    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp4Op>(1864        op, destVecType, i32Source, scale, op.getIndex());1865  else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.isBF16())1866    rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp4Op>(1867        op, destVecType, i32Source, scale, op.getIndex());1868  else1869    return failure();1870 1871  return success();1872}1873 1874LogicalResult PackedScaledTruncOpLowering::matchAndRewrite(1875    PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,1876    ConversionPatternRewriter &rewriter) const {1877  Location loc = op.getLoc();1878  if (chipset != kGfx950)1879    return rewriter.notifyMatchFailure(1880        loc, "Scaled fp conversion instructions are not available on target "1881             "architecture and their emulation is not implemented");1882  Type v2i16 = getTypeConverter()->convertType(1883      VectorType::get(2, rewriter.getI16Type()));1884  Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());1885 1886  Type resultType = op.getResult().getType();1887  Type resultElemType = getElementTypeOrSelf(resultType);1888  VectorType sourceVecType = cast<VectorType>(op.getSource().getType());1889  Type sourceElemType = sourceVecType.getElementType();1890 1891  Type intResultType = isa<Float4E2M1FNType>(resultElemType) ? i32 : v2i16;1892 1893  Value source = adaptor.getSource();1894  Value scale = adaptor.getScale();1895  Value existing = adaptor.getExisting();1896  if (existing)1897    existing = LLVM::BitcastOp::create(rewriter, loc, intResultType, existing);1898  else1899    existing = LLVM::ZeroOp::create(rewriter, loc, intResultType);1900 1901  if (sourceVecType.getNumElements() < 2) {1902    Value c0 = createI32Constant(rewriter, loc, 0);1903    Value elem0 = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);1904    VectorType v2 = VectorType::get(2, sourceElemType);1905    source = LLVM::ZeroOp::create(rewriter, loc, v2);1906    source = LLVM::InsertElementOp::create(rewriter, loc, source, elem0, c0);1907  }1908 1909  Value sourceA, sourceB;1910  if (sourceElemType.isF32()) {1911    Value c0 = createI32Constant(rewriter, loc, 0);1912    Value c1 = createI32Constant(rewriter, loc, 1);1913    sourceA = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);1914    sourceB = LLVM::ExtractElementOp::create(rewriter, loc, source, c1);1915  }1916 1917  Value result;1918  if (sourceElemType.isF32() && isa<Float8E5M2Type>(resultElemType))1919    result = ROCDL::CvtScaleF32PkBf8F32Op::create(rewriter, loc, intResultType,1920                                                  existing, sourceA, sourceB,1921                                                  scale, op.getIndex());1922  else if (sourceElemType.isF16() && isa<Float8E5M2Type>(resultElemType))1923    result = ROCDL::CvtScaleF32PkBf8F16Op::create(1924        rewriter, loc, intResultType, existing, source, scale, op.getIndex());1925  else if (sourceElemType.isBF16() && isa<Float8E5M2Type>(resultElemType))1926    result = ROCDL::CvtScaleF32PkBf8Bf16Op::create(1927        rewriter, loc, intResultType, existing, source, scale, op.getIndex());1928  else if (sourceElemType.isF32() && isa<Float8E4M3FNType>(resultElemType))1929    result = ROCDL::CvtScaleF32PkFp8F32Op::create(rewriter, loc, intResultType,1930                                                  existing, sourceA, sourceB,1931                                                  scale, op.getIndex());1932  else if (sourceElemType.isF16() && isa<Float8E4M3FNType>(resultElemType))1933    result = ROCDL::CvtScaleF32PkFp8F16Op::create(1934        rewriter, loc, intResultType, existing, source, scale, op.getIndex());1935  else if (sourceElemType.isBF16() && isa<Float8E4M3FNType>(resultElemType))1936    result = ROCDL::CvtScaleF32PkFp8Bf16Op::create(1937        rewriter, loc, intResultType, existing, source, scale, op.getIndex());1938  else if (sourceElemType.isF32() && isa<Float4E2M1FNType>(resultElemType))1939    result = ROCDL::CvtScaleF32PkFp4F32Op::create(rewriter, loc, intResultType,1940                                                  existing, sourceA, sourceB,1941                                                  scale, op.getIndex());1942  else if (sourceElemType.isF16() && isa<Float4E2M1FNType>(resultElemType))1943    result = ROCDL::CvtScaleF32PkFp4F16Op::create(1944        rewriter, loc, intResultType, existing, source, scale, op.getIndex());1945  else if (sourceElemType.isBF16() && isa<Float4E2M1FNType>(resultElemType))1946    result = ROCDL::CvtScaleF32PkFp4Bf16Op::create(1947        rewriter, loc, intResultType, existing, source, scale, op.getIndex());1948  else1949    return failure();1950 1951  result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(1952      op, getTypeConverter()->convertType(resultType), result);1953  return success();1954}1955 1956LogicalResult PackedTrunc2xFp8OpLowering::matchAndRewrite(1957    PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,1958    ConversionPatternRewriter &rewriter) const {1959  Location loc = op.getLoc();1960  if (!(chipset == kGfx942 || hasOcpFp8(chipset)))1961    return rewriter.notifyMatchFailure(1962        loc, "Fp8 conversion instructions are not available on target "1963             "architecture and their emulation is not implemented");1964  Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());1965 1966  Type resultType = op.getResult().getType();1967  Type resultElemType = getElementTypeOrSelf(resultType);1968 1969  Value sourceA = adaptor.getSourceA();1970  Value sourceB = adaptor.getSourceB();1971  if (!sourceB)1972    sourceB = LLVM::UndefOp::create(rewriter, loc, sourceA.getType());1973  Value existing = adaptor.getExisting();1974  if (existing)1975    existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);1976  else1977    existing = LLVM::UndefOp::create(rewriter, loc, i32);1978 1979  Value result;1980  if (typeIsExpectedBf8ForChipset(chipset, resultElemType))1981    result = ROCDL::CvtPkBf8F32Op::create(rewriter, loc, i32, sourceA, sourceB,1982                                          existing, op.getWordIndex());1983  else if (typeIsExpectedFp8ForChipset(chipset, resultElemType))1984    result = ROCDL::CvtPkFp8F32Op::create(rewriter, loc, i32, sourceA, sourceB,1985                                          existing, op.getWordIndex());1986 1987  result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(1988      op, getTypeConverter()->convertType(resultType), result);1989  return success();1990}1991 1992LogicalResult PackedStochRoundFp8OpLowering::matchAndRewrite(1993    PackedStochRoundFp8Op op, PackedStochRoundFp8OpAdaptor adaptor,1994    ConversionPatternRewriter &rewriter) const {1995  Location loc = op.getLoc();1996  if (!(chipset == kGfx942 || hasOcpFp8(chipset)))1997    return rewriter.notifyMatchFailure(1998        loc, "Fp8 conversion instructions are not available on target "1999             "architecture and their emulation is not implemented");2000  Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());2001 2002  Type resultType = op.getResult().getType();2003  Type resultElemType = getElementTypeOrSelf(resultType);2004 2005  Value source = adaptor.getSource();2006  Value stoch = adaptor.getStochiasticParam();2007  Value existing = adaptor.getExisting();2008  if (existing)2009    existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);2010  else2011    existing = LLVM::UndefOp::create(rewriter, loc, i32);2012 2013  Value result;2014  if (typeIsExpectedBf8ForChipset(chipset, resultElemType))2015    result = ROCDL::CvtSrBf8F32Op::create(rewriter, loc, i32, source, stoch,2016                                          existing, op.getStoreIndex());2017  else if (typeIsExpectedFp8ForChipset(chipset, resultElemType))2018    result = ROCDL::CvtSrFp8F32Op::create(rewriter, loc, i32, source, stoch,2019                                          existing, op.getStoreIndex());2020 2021  result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(2022      op, getTypeConverter()->convertType(resultType), result);2023  return success();2024}2025 2026// Implement the AMDGPU_DPPLowering class that will convert the amdgpu.dpp2027// operation into the corresponding ROCDL instructions.2028struct AMDGPUDPPLowering : public ConvertOpToLLVMPattern<DPPOp> {2029  AMDGPUDPPLowering(const LLVMTypeConverter &converter, Chipset chipset)2030      : ConvertOpToLLVMPattern<DPPOp>(converter), chipset(chipset) {}2031  Chipset chipset;2032 2033  LogicalResult2034  matchAndRewrite(DPPOp DppOp, DPPOp::Adaptor adaptor,2035                  ConversionPatternRewriter &rewriter) const override {2036 2037    // Convert the source operand to the corresponding LLVM type2038    Location loc = DppOp.getLoc();2039    Value src = adaptor.getSrc();2040    Value old = adaptor.getOld();2041    Type srcType = src.getType();2042    Type oldType = old.getType();2043    Type llvmType = nullptr;2044    if (srcType.getIntOrFloatBitWidth() < 32) {2045      llvmType = rewriter.getI32Type();2046    } else if (isa<FloatType>(srcType)) {2047      llvmType = (srcType.getIntOrFloatBitWidth() == 32)2048                     ? rewriter.getF32Type()2049                     : rewriter.getF64Type();2050    } else if (isa<IntegerType>(srcType)) {2051      llvmType = (srcType.getIntOrFloatBitWidth() == 32)2052                     ? rewriter.getI32Type()2053                     : rewriter.getI64Type();2054    }2055    auto llvmSrcIntType = typeConverter->convertType(2056        rewriter.getIntegerType(srcType.getIntOrFloatBitWidth()));2057 2058    // If the source type is less of 32, use bitcast to convert it to i32.2059    auto convertOperand = [&](Value operand, Type operandType) {2060      if (operandType.getIntOrFloatBitWidth() <= 16) {2061        if (llvm::isa<FloatType>(operandType)) {2062          operand =2063              LLVM::BitcastOp::create(rewriter, loc, llvmSrcIntType, operand);2064        }2065        auto llvmVecType = typeConverter->convertType(mlir::VectorType::get(2066            32 / operandType.getIntOrFloatBitWidth(), llvmSrcIntType));2067        Value undefVec = LLVM::UndefOp::create(rewriter, loc, llvmVecType);2068        operand =2069            LLVM::InsertElementOp::create(rewriter, loc, undefVec, operand,2070                                          createI32Constant(rewriter, loc, 0));2071        operand = LLVM::BitcastOp::create(rewriter, loc, llvmType, operand);2072      }2073      return operand;2074    };2075 2076    src = convertOperand(src, srcType);2077    old = convertOperand(old, oldType);2078 2079    // This is taken from the following file llvm/lib/Target/AMDGPU/SIDefines.h2080    enum DppCtrl : unsigned {2081      ROW_SHL0 = 0x100,2082      ROW_SHR0 = 0x110,2083      ROW_ROR0 = 0x120,2084      WAVE_SHL1 = 0x130,2085      WAVE_ROL1 = 0x134,2086      WAVE_SHR1 = 0x138,2087      WAVE_ROR1 = 0x13C,2088      ROW_MIRROR = 0x140,2089      ROW_HALF_MIRROR = 0x141,2090      BCAST15 = 0x142,2091      BCAST31 = 0x143,2092    };2093 2094    auto kind = DppOp.getKind();2095    auto permArgument = DppOp.getPermArgument();2096    uint32_t DppCtrl = 0;2097 2098    switch (kind) {2099 2100    case DPPPerm::quad_perm:2101      if (auto quadPermAttr = cast<ArrayAttr>(*permArgument)) {2102        int32_t i = 0;2103        for (auto elem : quadPermAttr.getAsRange<IntegerAttr>()) {2104          uint32_t num = elem.getInt();2105          DppCtrl |= num << (i * 2);2106          i++;2107        }2108      }2109      break;2110    case DPPPerm::row_shl:2111      if (auto intAttr = cast<IntegerAttr>(*permArgument)) {2112        DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHL0;2113      }2114      break;2115    case DPPPerm::row_shr:2116      if (auto intAttr = cast<IntegerAttr>(*permArgument)) {2117        DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHR0;2118      }2119      break;2120    case DPPPerm::row_ror:2121      if (auto intAttr = cast<IntegerAttr>(*permArgument)) {2122        DppCtrl = intAttr.getInt() + DppCtrl::ROW_ROR0;2123      }2124      break;2125    case DPPPerm::wave_shl:2126      DppCtrl = DppCtrl::WAVE_SHL1;2127      break;2128    case DPPPerm::wave_shr:2129      DppCtrl = DppCtrl::WAVE_SHR1;2130      break;2131    case DPPPerm::wave_rol:2132      DppCtrl = DppCtrl::WAVE_ROL1;2133      break;2134    case DPPPerm::wave_ror:2135      DppCtrl = DppCtrl::WAVE_ROR1;2136      break;2137    case DPPPerm::row_mirror:2138      DppCtrl = DppCtrl::ROW_MIRROR;2139      break;2140    case DPPPerm::row_half_mirror:2141      DppCtrl = DppCtrl::ROW_HALF_MIRROR;2142      break;2143    case DPPPerm::row_bcast_15:2144      DppCtrl = DppCtrl::BCAST15;2145      break;2146    case DPPPerm::row_bcast_31:2147      DppCtrl = DppCtrl::BCAST31;2148      break;2149    }2150 2151    // Check for row_mask, bank_mask, bound_ctrl if they exist and create2152    // constants2153    auto rowMask = DppOp->getAttrOfType<IntegerAttr>("row_mask").getInt();2154    auto bankMask = DppOp->getAttrOfType<IntegerAttr>("bank_mask").getInt();2155    bool boundCtrl = DppOp->getAttrOfType<BoolAttr>("bound_ctrl").getValue();2156 2157    // create a ROCDL_DPPMovOp instruction with the appropriate attributes2158    auto dppMovOp =2159        ROCDL::DPPUpdateOp::create(rewriter, loc, llvmType, old, src, DppCtrl,2160                                   rowMask, bankMask, boundCtrl);2161 2162    Value result = dppMovOp.getRes();2163    if (srcType.getIntOrFloatBitWidth() < 32) {2164      result = LLVM::TruncOp::create(rewriter, loc, llvmSrcIntType, result);2165      if (!llvm::isa<IntegerType>(srcType)) {2166        result = LLVM::BitcastOp::create(rewriter, loc, srcType, result);2167      }2168    }2169 2170    // We are replacing the AMDGPU_DPPOp instruction with the new2171    // ROCDL_DPPMovOp instruction2172    rewriter.replaceOp(DppOp, ValueRange(result));2173    return success();2174  }2175};2176 2177struct AMDGPUSwizzleBitModeLowering2178    : public ConvertOpToLLVMPattern<SwizzleBitModeOp> {2179  using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;2180 2181  LogicalResult2182  matchAndRewrite(SwizzleBitModeOp op, OpAdaptor adaptor,2183                  ConversionPatternRewriter &rewriter) const override {2184    Location loc = op.getLoc();2185    Type i32 = rewriter.getI32Type();2186    Value src = adaptor.getSrc();2187    SmallVector<Value> decomposed =2188        LLVM::decomposeValue(rewriter, loc, src, i32);2189    unsigned andMask = op.getAndMask();2190    unsigned orMask = op.getOrMask();2191    unsigned xorMask = op.getXorMask();2192 2193    // bit 15 is 0 for the BitMode swizzle.2194    // https://gpuopen.com/learn/amd-gcn-assembly-cross-lane-operations/2195    unsigned mask = andMask | (orMask << 5) | (xorMask << 10);2196    Value maskValue = createI32Constant(rewriter, loc, mask);2197    SmallVector<Value> swizzled;2198    for (Value v : decomposed) {2199      Value res =2200          ROCDL::DsSwizzleOp::create(rewriter, loc, v.getType(), v, maskValue);2201      swizzled.emplace_back(res);2202    }2203 2204    Value result = LLVM::composeValue(rewriter, loc, swizzled, src.getType());2205    rewriter.replaceOp(op, result);2206    return success();2207  }2208};2209 2210struct AMDGPUPermlaneLowering : public ConvertOpToLLVMPattern<PermlaneSwapOp> {2211  using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;2212 2213  AMDGPUPermlaneLowering(const LLVMTypeConverter &converter, Chipset chipset)2214      : ConvertOpToLLVMPattern<PermlaneSwapOp>(converter), chipset(chipset) {}2215  Chipset chipset;2216 2217  LogicalResult2218  matchAndRewrite(PermlaneSwapOp op, OpAdaptor adaptor,2219                  ConversionPatternRewriter &rewriter) const override {2220    if (chipset < kGfx950)2221      return op->emitOpError("permlane_swap is only supported on gfx950+");2222 2223    Location loc = op.getLoc();2224    Type i32 = rewriter.getI32Type();2225    Value src = adaptor.getSrc();2226    unsigned rowLength = op.getRowLength();2227    bool fi = op.getFetchInactive();2228    bool boundctrl = op.getBoundCtrl();2229 2230    SmallVector<Value> decomposed =2231        LLVM::decomposeValue(rewriter, loc, src, i32);2232 2233    SmallVector<Value> permuted;2234    for (Value v : decomposed) {2235      Value res;2236      Type i32pair = LLVM::LLVMStructType::getLiteral(2237          rewriter.getContext(), {v.getType(), v.getType()});2238 2239      if (rowLength == 16)2240        res = ROCDL::Permlane16SwapOp::create(rewriter, loc, i32pair, v, v, fi,2241                                              boundctrl);2242      else if (rowLength == 32)2243        res = ROCDL::Permlane32SwapOp::create(rewriter, loc, i32pair, v, v, fi,2244                                              boundctrl);2245      else2246        llvm_unreachable("unsupported row length");2247 2248      Value vdst0 = LLVM::ExtractValueOp::create(rewriter, loc, res, {0});2249      Value vdst1 = LLVM::ExtractValueOp::create(rewriter, loc, res, {1});2250 2251      Value isEqual = LLVM::ICmpOp::create(rewriter, loc,2252                                           LLVM::ICmpPredicate::eq, vdst0, v);2253 2254      // Per `permlane(16|32)` semantics: if the first extracted element equals2255      // 'v', the result is the second element; otherwise it is the first.2256      Value vdstNew =2257          LLVM::SelectOp::create(rewriter, loc, isEqual, vdst1, vdst0);2258      permuted.emplace_back(vdstNew);2259    }2260 2261    Value result = LLVM::composeValue(rewriter, loc, permuted, src.getType());2262    rewriter.replaceOp(op, result);2263    return success();2264  }2265};2266 2267struct ConvertAMDGPUToROCDLPass2268    : public impl::ConvertAMDGPUToROCDLPassBase<ConvertAMDGPUToROCDLPass> {2269  using Base::Base;2270 2271  void runOnOperation() override {2272    MLIRContext *ctx = &getContext();2273    FailureOr<Chipset> maybeChipset = Chipset::parse(chipset);2274    if (failed(maybeChipset)) {2275      emitError(UnknownLoc::get(ctx), "Invalid chipset name: " + chipset);2276      return signalPassFailure();2277    }2278 2279    RewritePatternSet patterns(ctx);2280    LLVMTypeConverter converter(ctx);2281    populateAMDGPUToROCDLConversionPatterns(converter, patterns, *maybeChipset);2282    LLVMConversionTarget target(getContext());2283    target.addIllegalDialect<::mlir::amdgpu::AMDGPUDialect>();2284    target.addLegalDialect<::mlir::LLVM::LLVMDialect>();2285    target.addLegalDialect<::mlir::ROCDL::ROCDLDialect>();2286    if (failed(applyPartialConversion(getOperation(), target,2287                                      std::move(patterns))))2288      signalPassFailure();2289  }2290};2291} // namespace2292 2293void mlir::populateAMDGPUMemorySpaceAttributeConversions(2294    TypeConverter &typeConverter) {2295  typeConverter.addTypeAttributeConversion(2296      [](BaseMemRefType type, amdgpu::AddressSpaceAttr as)2297          -> TypeConverter::AttributeConversionResult {2298        MLIRContext *ctx = as.getContext();2299        Type i64 = IntegerType::get(ctx, 64);2300        switch (as.getValue()) {2301        case amdgpu::AddressSpace::FatRawBuffer:2302          return IntegerAttr::get(i64, 7);2303        case amdgpu::AddressSpace::BufferRsrc:2304          return IntegerAttr::get(i64, 8);2305        case amdgpu::AddressSpace::FatStructuredBuffer:2306          return IntegerAttr::get(i64, 9);2307        }2308        return TypeConverter::AttributeConversionResult::abort();2309      });2310}2311 2312void mlir::populateAMDGPUToROCDLConversionPatterns(LLVMTypeConverter &converter,2313                                                   RewritePatternSet &patterns,2314                                                   Chipset chipset) {2315  populateAMDGPUMemorySpaceAttributeConversions(converter);2316  patterns2317      .add<FatRawBufferCastLowering,2318           RawBufferOpLowering<RawBufferLoadOp, ROCDL::RawPtrBufferLoadOp>,2319           RawBufferOpLowering<RawBufferStoreOp, ROCDL::RawPtrBufferStoreOp>,2320           RawBufferOpLowering<RawBufferAtomicFaddOp,2321                               ROCDL::RawPtrBufferAtomicFaddOp>,2322           RawBufferOpLowering<RawBufferAtomicFmaxOp,2323                               ROCDL::RawPtrBufferAtomicFmaxOp>,2324           RawBufferOpLowering<RawBufferAtomicSmaxOp,2325                               ROCDL::RawPtrBufferAtomicSmaxOp>,2326           RawBufferOpLowering<RawBufferAtomicUminOp,2327                               ROCDL::RawPtrBufferAtomicUminOp>,2328           RawBufferOpLowering<RawBufferAtomicCmpswapOp,2329                               ROCDL::RawPtrBufferAtomicCmpSwap>,2330           AMDGPUDPPLowering, MemoryCounterWaitOpLowering, LDSBarrierOpLowering,2331           SchedBarrierOpLowering, MFMAOpLowering, ScaledMFMAOpLowering,2332           WMMAOpLowering, ExtPackedFp8OpLowering, ScaledExtPacked816OpLowering,2333           ScaledExtPackedOpLowering, PackedScaledTruncOpLowering,2334           PackedTrunc2xFp8OpLowering, PackedStochRoundFp8OpLowering,2335           GatherToLDSOpLowering, TransposeLoadOpLowering,2336           AMDGPUPermlaneLowering>(converter, chipset);2337  patterns.add<AMDGPUSwizzleBitModeLowering>(converter);2338}2339