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