30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/TypeSwitch.h"
32#include "llvm/Support/AMDGPUAddrSpace.h"
33#include "llvm/Support/Casting.h"
34#include "llvm/Support/ErrorHandling.h"
39#define GEN_PASS_DEF_CONVERTAMDGPUTOROCDLPASS
40#include "mlir/Conversion/Passes.h.inc"
58 return chipset >=
Chipset(9, 0, 6);
100 if (chipset ==
Chipset(9, 5, 0))
115 IntegerType destTy = rewriter.getIntegerType(width);
117 auto valTy = cast<IntegerType>(val.
getType());
120 return valTy.getWidth() > width
121 ?
Value(LLVM::TruncOp::create(rewriter, loc, destTy, val))
122 :
Value(LLVM::ZExtOp::create(rewriter, loc, destTy, val));
133 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), value);
144 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), value);
151 IntegerType i32 = rewriter.getI32Type();
153 for (
auto [i, increment, stride] : llvm::enumerate(
indices, strides)) {
156 ShapedType::isDynamic(stride)
158 memRefDescriptor.
stride(rewriter, loc, i))
159 : LLVM::ConstantOp::create(rewriter, loc, i32, stride);
160 increment = LLVM::MulOp::create(rewriter, loc, increment, strideValue);
172 MemRefType memrefType,
177 constexpr int64_t first45bits = (1ll << 45) - 1;
180 if (memrefType.hasStaticShape() &&
181 !llvm::any_of(strides, ShapedType::isDynamic)) {
182 int64_t size = memrefType.getRank() == 0 ? 1 : 0;
184 for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i)
185 size = std::max(
shape[i] * strides[i], size);
186 size = size * elementByteWidth;
190 for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i) {
191 Value size = memrefDescriptor.
size(rewriter, loc, i);
192 Value stride = memrefDescriptor.
stride(rewriter, loc, i);
193 Value maxThisDim = LLVM::MulOp::create(rewriter, loc, size, stride);
195 ? LLVM::UMaxOp::create(rewriter, loc, maxIndex, maxThisDim)
200 return LLVM::MulOp::create(rewriter, loc, maxIndexI64, byteWidthConst);
206 Value cacheSwizzleStride =
nullptr,
207 unsigned addressSpace = 8) {
211 Type i16 = rewriter.getI16Type();
214 Value cacheStrideZext =
215 LLVM::ZExtOp::create(rewriter, loc, i16, cacheSwizzleStride);
216 Value swizzleBit = LLVM::ConstantOp::create(
217 rewriter, loc, i16, rewriter.getI16IntegerAttr(1 << 14));
218 stride = LLVM::OrOp::create(rewriter, loc, cacheStrideZext, swizzleBit,
221 stride = LLVM::ConstantOp::create(rewriter, loc, i16,
222 rewriter.getI16IntegerAttr(0));
251 flags |= (7 << 12) | (4 << 15);
254 uint32_t oob = boundsCheck ? 3 : 2;
255 flags |= (oob << 28);
263 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
264 Value resource = rewriter.createOrFold<ROCDL::MakeBufferRsrcOp>(
265 loc, rsrcType, basePointer, stride, numRecords, flagsConst);
270struct FatRawBufferCastLowering
272 FatRawBufferCastLowering(
const LLVMTypeConverter &converter, Chipset chipset)
273 : ConvertOpToLLVMPattern<FatRawBufferCastOp>(converter),
279 matchAndRewrite(FatRawBufferCastOp op, FatRawBufferCastOpAdaptor adaptor,
280 ConversionPatternRewriter &rewriter)
const override {
281 Location loc = op.getLoc();
282 Value memRef = adaptor.getSource();
283 Value unconvertedMemref = op.getSource();
284 MemRefType memrefType = cast<MemRefType>(unconvertedMemref.
getType());
285 MemRefDescriptor descriptor(memRef);
287 DataLayout dataLayout = DataLayout::closest(op);
288 int64_t elementByteWidth =
291 int64_t unusedOffset = 0;
292 SmallVector<int64_t, 5> strideVals;
293 if (
failed(memrefType.getStridesAndOffset(strideVals, unusedOffset)))
294 return op.emitOpError(
"Can't lower non-stride-offset memrefs");
296 Value numRecords = adaptor.getValidBytes();
299 getNumRecords(rewriter, loc, memrefType, descriptor, strideVals,
300 elementByteWidth, chipset, adaptor.getBoundsCheck());
303 adaptor.getResetOffset()
304 ? descriptor.bufferPtr(rewriter, loc, *getTypeConverter(),
306 : descriptor.alignedPtr(rewriter, loc);
309 adaptor.getResetOffset()
311 : descriptor.offset(rewriter, loc);
313 bool hasSizes = memrefType.getRank() > 0;
316 Value sizes = hasSizes
317 ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,
321 hasSizes ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,
326 rewriter, loc, basePointer, numRecords, adaptor.getBoundsCheck(),
327 chipset, adaptor.getCacheSwizzleStride(), 7);
329 Value
result = MemRefDescriptor::poison(
331 getTypeConverter()->convertType(op.getResult().getType()));
333 result = LLVM::InsertValueOp::create(rewriter, loc,
result, fatPtr, pos);
334 result = LLVM::InsertValueOp::create(rewriter, loc,
result, fatPtr,
336 result = LLVM::InsertValueOp::create(rewriter, loc,
result, offset,
339 result = LLVM::InsertValueOp::create(rewriter, loc,
result, sizes,
341 result = LLVM::InsertValueOp::create(rewriter, loc,
result, strides,
344 rewriter.replaceOp(op,
result);
350template <
typename GpuOp,
typename Intrinsic>
352 RawBufferOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
353 : ConvertOpToLLVMPattern<GpuOp>(converter), chipset(chipset) {}
356 static constexpr uint32_t maxVectorOpWidth = 128;
359 matchAndRewrite(GpuOp gpuOp,
typename GpuOp::Adaptor adaptor,
360 ConversionPatternRewriter &rewriter)
const override {
361 Location loc = gpuOp.getLoc();
362 Value memref = adaptor.getMemref();
363 Value unconvertedMemref = gpuOp.getMemref();
364 MemRefType memrefType = cast<MemRefType>(unconvertedMemref.
getType());
366 if (chipset.majorVersion < 9)
367 return gpuOp.emitOpError(
"raw buffer ops require GCN or higher");
369 Value storeData = adaptor.getODSOperands(0)[0];
370 if (storeData == memref)
374 wantedDataType = storeData.
getType();
376 wantedDataType = gpuOp.getODSResults(0)[0].getType();
378 Value atomicCmpData = Value();
381 Value maybeCmpData = adaptor.getODSOperands(1)[0];
382 if (maybeCmpData != memref)
383 atomicCmpData = maybeCmpData;
386 Type llvmWantedDataType = this->typeConverter->convertType(wantedDataType);
388 Type i32 = rewriter.getI32Type();
391 DataLayout dataLayout = DataLayout::closest(gpuOp);
392 int64_t elementByteWidth =
401 Type llvmBufferValType = llvmWantedDataType;
403 if (
auto floatType = dyn_cast<FloatType>(wantedDataType))
404 llvmBufferValType = this->getTypeConverter()->convertType(
405 rewriter.getIntegerType(floatType.getWidth()));
407 if (
auto dataVector = dyn_cast<VectorType>(wantedDataType)) {
408 uint32_t vecLen = dataVector.getNumElements();
411 uint32_t totalBits = elemBits * vecLen;
413 isa_and_present<RawBufferAtomicFaddOp>(*gpuOp) && vecLen == 2;
414 if (totalBits > maxVectorOpWidth)
415 return gpuOp.emitOpError(
416 "Total width of loads or stores must be no more than " +
417 Twine(maxVectorOpWidth) +
" bits, but we call for " +
419 " bits. This should've been caught in validation");
420 if (!usePackedFp16 && elemBits < 32) {
421 if (totalBits > 32) {
422 if (totalBits % 32 != 0)
423 return gpuOp.emitOpError(
"Load or store of more than 32-bits that "
424 "doesn't fit into words. Can't happen\n");
425 llvmBufferValType = this->typeConverter->convertType(
426 VectorType::get(totalBits / 32, i32));
428 llvmBufferValType = this->typeConverter->convertType(
429 rewriter.getIntegerType(totalBits));
433 if (
auto vecType = dyn_cast<VectorType>(llvmBufferValType)) {
436 if (vecType.getNumElements() == 1)
437 llvmBufferValType = vecType.getElementType();
440 SmallVector<Value, 6> args;
442 if (llvmBufferValType != llvmWantedDataType) {
443 Value castForStore = LLVM::BitcastOp::create(
444 rewriter, loc, llvmBufferValType, storeData);
445 args.push_back(castForStore);
447 args.push_back(storeData);
452 if (llvmBufferValType != llvmWantedDataType) {
453 Value castForCmp = LLVM::BitcastOp::create(
454 rewriter, loc, llvmBufferValType, atomicCmpData);
455 args.push_back(castForCmp);
457 args.push_back(atomicCmpData);
463 SmallVector<int64_t, 5> strides;
464 if (
failed(memrefType.getStridesAndOffset(strides, offset)))
465 return gpuOp.emitOpError(
"Can't lower non-stride-offset memrefs");
467 MemRefDescriptor memrefDescriptor(memref);
469 Value ptr = memrefDescriptor.bufferPtr(
470 rewriter, loc, *this->getTypeConverter(), memrefType);
472 getNumRecords(rewriter, loc, memrefType, memrefDescriptor, strides,
473 elementByteWidth, chipset, adaptor.getBoundsCheck());
475 adaptor.getBoundsCheck(), chipset);
476 args.push_back(resource);
480 adaptor.getIndices(), strides);
481 if (std::optional<int32_t> indexOffset = adaptor.getIndexOffset();
482 indexOffset && *indexOffset > 0) {
484 voffset = voffset ? LLVM::AddOp::create(rewriter, loc, voffset,
488 voffset = LLVM::MulOp::create(rewriter, loc, voffset, byteWidthConst);
489 args.push_back(voffset);
492 Value sgprOffset = adaptor.getSgprOffset();
495 sgprOffset = LLVM::MulOp::create(rewriter, loc, sgprOffset, byteWidthConst);
496 args.push_back(sgprOffset);
498 llvm::SmallVector<Type, 1> resultTypes(gpuOp->getNumResults(),
500 typename Intrinsic::Properties properties;
501 properties.aux = rewriter.getI32IntegerAttr(0);
503 Intrinsic::create(rewriter, loc, resultTypes, args, properties);
506 if (llvmBufferValType != llvmWantedDataType) {
507 replacement = LLVM::BitcastOp::create(rewriter, loc, llvmWantedDataType,
512 rewriter.eraseOp(gpuOp);
529static FailureOr<unsigned> encodeWaitcnt(
Chipset chipset,
unsigned vmcnt,
530 unsigned expcnt,
unsigned lgkmcnt) {
532 vmcnt = std::min(15u, vmcnt);
533 expcnt = std::min(7u, expcnt);
534 lgkmcnt = std::min(15u, lgkmcnt);
535 return vmcnt | (expcnt << 4) | (lgkmcnt << 8);
538 vmcnt = std::min(63u, vmcnt);
539 expcnt = std::min(7u, expcnt);
540 lgkmcnt = std::min(15u, lgkmcnt);
541 unsigned lowBits = vmcnt & 0xF;
542 unsigned highBits = (vmcnt >> 4) << 14;
543 unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);
544 return lowBits | highBits | otherCnts;
547 vmcnt = std::min(63u, vmcnt);
548 expcnt = std::min(7u, expcnt);
549 lgkmcnt = std::min(63u, lgkmcnt);
550 unsigned lowBits = vmcnt & 0xF;
551 unsigned highBits = (vmcnt >> 4) << 14;
552 unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);
553 return lowBits | highBits | otherCnts;
556 vmcnt = std::min(63u, vmcnt);
557 expcnt = std::min(7u, expcnt);
558 lgkmcnt = std::min(63u, lgkmcnt);
559 return (vmcnt << 10) | expcnt | (lgkmcnt << 4);
564struct MemoryCounterWaitOpLowering
574 matchAndRewrite(MemoryCounterWaitOp op, OpAdaptor adaptor,
575 ConversionPatternRewriter &rewriter)
const override {
576 if (
chipset.majorVersion >= 12) {
578 if (std::optional<int> ds = adaptor.getDs())
579 ROCDL::WaitDscntOp::create(rewriter, loc, *ds);
581 if (std::optional<int>
load = adaptor.getLoad())
582 ROCDL::WaitLoadcntOp::create(rewriter, loc, *
load);
584 if (std::optional<int> store = adaptor.getStore())
585 ROCDL::WaitStorecntOp::create(rewriter, loc, *store);
587 if (std::optional<int> exp = adaptor.getExp())
588 ROCDL::WaitExpcntOp::create(rewriter, loc, *exp);
590 if (std::optional<int>
tensor = adaptor.getTensor())
591 ROCDL::WaitTensorcntOp::create(rewriter, loc, *
tensor);
593 rewriter.eraseOp(op);
597 if (adaptor.getTensor())
598 return op.emitOpError(
"unsupported chipset");
600 auto getVal = [](
Attribute attr) ->
unsigned {
602 return cast<IntegerAttr>(attr).getInt();
607 unsigned ds = getVal(adaptor.getDsAttr());
608 unsigned exp = getVal(adaptor.getExpAttr());
610 unsigned vmcnt = 1024;
612 Attribute store = adaptor.getStoreAttr();
614 vmcnt = getVal(
load) + getVal(store);
616 vmcnt = getVal(
load);
618 vmcnt = getVal(store);
621 FailureOr<unsigned> waitcnt = encodeWaitcnt(chipset, vmcnt, exp, ds);
623 return op.emitOpError(
"unsupported chipset");
625 rewriter.replaceOpWithNewOp<ROCDL::SWaitcntOp>(op, *waitcnt);
631 LDSBarrierOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
632 : ConvertOpToLLVMPattern<LDSBarrierOp>(converter), chipset(chipset) {}
637 matchAndRewrite(LDSBarrierOp op, LDSBarrierOp::Adaptor adaptor,
638 ConversionPatternRewriter &rewriter)
const override {
639 Location loc = op.getLoc();
642 bool requiresInlineAsm = chipset <
kGfx90a;
645 rewriter.getAttr<LLVM::MMRATagAttr>(
"amdgpu-synchronize-as",
"local");
654 StringRef scope =
"workgroup";
656 auto relFence = LLVM::FenceOp::create(rewriter, loc,
657 LLVM::AtomicOrdering::release, scope);
658 relFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);
659 if (requiresInlineAsm) {
660 auto asmDialectAttr = LLVM::AsmDialectAttr::get(rewriter.getContext(),
661 LLVM::AsmDialect::AD_ATT);
662 const char *asmStr =
";;;WARNING: BREAKS DEBUG WATCHES\ns_barrier";
663 const char *constraints =
"";
664 LLVM::InlineAsmOp::create(
667 asmStr, constraints,
true,
668 false, LLVM::TailCallKind::None,
671 }
else if (chipset.majorVersion < 12) {
672 ROCDL::SBarrierOp::create(rewriter, loc);
674 ROCDL::BarrierSignalOp::create(rewriter, loc, -1);
675 ROCDL::BarrierWaitOp::create(rewriter, loc, -1);
678 auto acqFence = LLVM::FenceOp::create(rewriter, loc,
679 LLVM::AtomicOrdering::acquire, scope);
680 acqFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);
681 rewriter.replaceOp(op, acqFence);
687 SchedBarrierOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
688 : ConvertOpToLLVMPattern<SchedBarrierOp>(converter), chipset(chipset) {}
693 matchAndRewrite(SchedBarrierOp op, SchedBarrierOp::Adaptor adaptor,
694 ConversionPatternRewriter &rewriter)
const override {
695 rewriter.replaceOpWithNewOp<ROCDL::SchedBarrier>(op, op.getOptsAttr());
719 bool allowBf16 =
true) {
721 if (
auto vectorType = dyn_cast<VectorType>(inputType)) {
722 if (vectorType.getElementType().isBF16() && !allowBf16)
723 return LLVM::BitcastOp::create(
724 rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);
725 if (vectorType.getElementType().isInteger(8) &&
726 vectorType.getNumElements() <= 8)
727 return LLVM::BitcastOp::create(
729 rewriter.getIntegerType(vectorType.getNumElements() * 8), input);
730 if (isa<IntegerType>(vectorType.getElementType()) &&
731 vectorType.getElementTypeBitWidth() <= 8) {
732 int64_t numWords = llvm::divideCeil(
733 vectorType.getNumElements() * vectorType.getElementTypeBitWidth(),
735 return LLVM::BitcastOp::create(
736 rewriter, loc, VectorType::get(numWords, rewriter.getI32Type()),
746 bool allowBf16 =
true) {
748 auto vectorType = cast<VectorType>(inputType);
750 if (vectorType.getElementType().isBF16() && !allowBf16)
751 return LLVM::BitcastOp::create(
752 rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);
754 if (isa<IntegerType>(vectorType.getElementType()) &&
755 vectorType.getElementTypeBitWidth() <= 8) {
756 int64_t numWords = llvm::divideCeil(
757 vectorType.getNumElements() * vectorType.getElementTypeBitWidth(), 32);
758 Type castType = (numWords > 1)
759 ?
Type{VectorType::get(numWords, rewriter.getI32Type())}
760 : rewriter.getI32Type();
761 return LLVM::BitcastOp::create(rewriter, loc, castType, input);
779 .Case([&](IntegerType) {
781 return LLVM::ZExtOp::create(rewriter, loc, rewriter.getI32Type(),
784 .Case([&](VectorType vectorType) {
786 int64_t numElements = vectorType.getNumElements();
787 assert((numElements == 4 || numElements == 8) &&
788 "scale operand must be a vector of length 4 or 8");
789 IntegerType outputType =
790 (numElements == 4) ? rewriter.getI32Type() : rewriter.getI64Type();
791 return LLVM::BitcastOp::create(rewriter, loc, outputType, input);
793 .DefaultUnreachable(
"unexpected input type for scale operand");
797static std::optional<ROCDL::WMMAMatrixScaleFormat>
800 .Case([](Float8E8M0FNUType) {
return ROCDL::WMMAMatrixScaleFormat::e8; })
801 .Case([](Float8E4M3FNType) {
return ROCDL::WMMAMatrixScaleFormat::e4m3; })
802 .Default(std::nullopt);
807static std::optional<StringRef>
809 if (m == 16 && n == 16 && k == 128)
811 ? ROCDL::wmma_scale16_f32_16x16x128_f8f6f4::getOperationName()
812 : ROCDL::wmma_scale_f32_16x16x128_f8f6f4::getOperationName();
814 if (m == 32 && n == 16 && k == 128)
815 return isScale16 ? ROCDL::wmma_scale16_f32_32x16x128_f4::getOperationName()
816 : ROCDL::wmma_scale_f32_32x16x128_f4::getOperationName();
830 ConversionPatternRewriter &rewriter,
Location loc,
835 auto vectorType = dyn_cast<VectorType>(inputType);
837 operands.push_back(llvmInput);
840 Type elemType = vectorType.getElementType();
842 operands.push_back(llvmInput);
849 auto mlirInputType = cast<VectorType>(mlirInput.
getType());
850 bool isInputInteger = mlirInputType.getElementType().isInteger();
851 if (isInputInteger) {
853 bool localIsUnsigned = isUnsigned;
855 localIsUnsigned =
true;
857 localIsUnsigned =
false;
860 NamedAttribute(attrName, rewriter.getBoolAttr(!localIsUnsigned)));
865 Type i32 = rewriter.getI32Type();
866 Type intrinsicInType = numBits <= 32
867 ? (
Type)rewriter.getIntegerType(numBits)
868 : (
Type)VectorType::get(numBits / 32, i32);
869 auto llvmIntrinsicInType = typeConverter->convertType(intrinsicInType);
870 Value castInput = rewriter.createOrFold<LLVM::BitcastOp>(
871 loc, llvmIntrinsicInType, llvmInput);
876 castInput = LLVM::ZExtOp::create(rewriter, loc, i32, castInput);
877 operands.push_back(castInput);
890 Value output, int32_t subwordOffset,
894 auto vectorType = dyn_cast<VectorType>(inputType);
895 Type elemType = vectorType.getElementType();
896 operands.push_back(output);
908 return (chipset ==
kGfx942 && isa<Float8E5M2FNUZType>(type)) ||
909 (
hasOcpFp8(chipset) && isa<Float8E5M2Type>(type));
915 return (chipset ==
kGfx942 && isa<Float8E4M3FNUZType>(type)) ||
916 (
hasOcpFp8(chipset) && isa<Float8E4M3FNType>(type));
924 uint32_t m = mfma.getM(), n = mfma.getN(), k = mfma.getK(),
925 b = mfma.getBlocks();
930 if (mfma.getReducePrecision() && chipset >=
kGfx942) {
931 if (m == 32 && n == 32 && k == 4 &&
b == 1)
932 return ROCDL::mfma_f32_32x32x4_xf32::getOperationName();
933 if (m == 16 && n == 16 && k == 8 &&
b == 1)
934 return ROCDL::mfma_f32_16x16x8_xf32::getOperationName();
936 if (m == 32 && n == 32 && k == 1 &&
b == 2)
937 return ROCDL::mfma_f32_32x32x1f32::getOperationName();
938 if (m == 16 && n == 16 && k == 1 &&
b == 4)
939 return ROCDL::mfma_f32_16x16x1f32::getOperationName();
940 if (m == 4 && n == 4 && k == 1 &&
b == 16)
941 return ROCDL::mfma_f32_4x4x1f32::getOperationName();
942 if (m == 32 && n == 32 && k == 2 &&
b == 1)
943 return ROCDL::mfma_f32_32x32x2f32::getOperationName();
944 if (m == 16 && n == 16 && k == 4 &&
b == 1)
945 return ROCDL::mfma_f32_16x16x4f32::getOperationName();
950 if (m == 32 && n == 32 && k == 16 &&
b == 1)
951 return ROCDL::mfma_f32_32x32x16_f16::getOperationName();
952 if (m == 16 && n == 16 && k == 32 &&
b == 1)
953 return ROCDL::mfma_f32_16x16x32_f16::getOperationName();
955 if (m == 32 && n == 32 && k == 4 &&
b == 2)
956 return ROCDL::mfma_f32_32x32x4f16::getOperationName();
957 if (m == 16 && n == 16 && k == 4 &&
b == 4)
958 return ROCDL::mfma_f32_16x16x4f16::getOperationName();
959 if (m == 4 && n == 4 && k == 4 &&
b == 16)
960 return ROCDL::mfma_f32_4x4x4f16::getOperationName();
961 if (m == 32 && n == 32 && k == 8 &&
b == 1)
962 return ROCDL::mfma_f32_32x32x8f16::getOperationName();
963 if (m == 16 && n == 16 && k == 16 &&
b == 1)
964 return ROCDL::mfma_f32_16x16x16f16::getOperationName();
969 if (m == 32 && n == 32 && k == 16 &&
b == 1)
970 return ROCDL::mfma_f32_32x32x16_bf16::getOperationName();
971 if (m == 16 && n == 16 && k == 32 &&
b == 1)
972 return ROCDL::mfma_f32_16x16x32_bf16::getOperationName();
975 if (m == 32 && n == 32 && k == 4 &&
b == 2)
976 return ROCDL::mfma_f32_32x32x4bf16_1k::getOperationName();
977 if (m == 16 && n == 16 && k == 4 &&
b == 4)
978 return ROCDL::mfma_f32_16x16x4bf16_1k::getOperationName();
979 if (m == 4 && n == 4 && k == 4 &&
b == 16)
980 return ROCDL::mfma_f32_4x4x4bf16_1k::getOperationName();
981 if (m == 32 && n == 32 && k == 8 &&
b == 1)
982 return ROCDL::mfma_f32_32x32x8bf16_1k::getOperationName();
983 if (m == 16 && n == 16 && k == 16 &&
b == 1)
984 return ROCDL::mfma_f32_16x16x16bf16_1k::getOperationName();
986 if (m == 32 && n == 32 && k == 2 &&
b == 2)
987 return ROCDL::mfma_f32_32x32x2bf16::getOperationName();
988 if (m == 16 && n == 16 && k == 2 &&
b == 4)
989 return ROCDL::mfma_f32_16x16x2bf16::getOperationName();
990 if (m == 4 && n == 4 && k == 2 &&
b == 16)
991 return ROCDL::mfma_f32_4x4x2bf16::getOperationName();
992 if (m == 32 && n == 32 && k == 4 &&
b == 1)
993 return ROCDL::mfma_f32_32x32x4bf16::getOperationName();
994 if (m == 16 && n == 16 && k == 8 &&
b == 1)
995 return ROCDL::mfma_f32_16x16x8bf16::getOperationName();
1000 if (m == 32 && n == 32 && k == 32 &&
b == 1)
1001 return ROCDL::mfma_i32_32x32x32_i8::getOperationName();
1002 if (m == 16 && n == 16 && k == 64 &&
b == 1)
1003 return ROCDL::mfma_i32_16x16x64_i8::getOperationName();
1005 if (m == 32 && n == 32 && k == 4 &&
b == 2)
1006 return ROCDL::mfma_i32_32x32x4i8::getOperationName();
1007 if (m == 16 && n == 16 && k == 4 &&
b == 4)
1008 return ROCDL::mfma_i32_16x16x4i8::getOperationName();
1009 if (m == 4 && n == 4 && k == 4 &&
b == 16)
1010 return ROCDL::mfma_i32_4x4x4i8::getOperationName();
1011 if (m == 32 && n == 32 && k == 8 &&
b == 1)
1012 return ROCDL::mfma_i32_32x32x8i8::getOperationName();
1013 if (m == 16 && n == 16 && k == 16 &&
b == 1)
1014 return ROCDL::mfma_i32_16x16x16i8::getOperationName();
1015 if (m == 32 && n == 32 && k == 16 &&
b == 1 && chipset >=
kGfx942)
1016 return ROCDL::mfma_i32_32x32x16_i8::getOperationName();
1017 if (m == 16 && n == 16 && k == 32 &&
b == 1 && chipset >=
kGfx942)
1018 return ROCDL::mfma_i32_16x16x32_i8::getOperationName();
1022 if (m == 16 && n == 16 && k == 4 &&
b == 1)
1023 return ROCDL::mfma_f64_16x16x4f64::getOperationName();
1024 if (m == 4 && n == 4 && k == 4 &&
b == 4)
1025 return ROCDL::mfma_f64_4x4x4f64::getOperationName();
1032 cast<VectorType>(mfma.getSourceB().getType()).getElementType();
1033 if (m == 16 && n == 16 && k == 32 &&
b == 1) {
1035 return ROCDL::mfma_f32_16x16x32_bf8_bf8::getOperationName();
1037 return ROCDL::mfma_f32_16x16x32_bf8_fp8::getOperationName();
1039 if (m == 32 && n == 32 && k == 16 &&
b == 1) {
1041 return ROCDL::mfma_f32_32x32x16_bf8_bf8::getOperationName();
1043 return ROCDL::mfma_f32_32x32x16_bf8_fp8::getOperationName();
1049 cast<VectorType>(mfma.getSourceB().getType()).getElementType();
1050 if (m == 16 && n == 16 && k == 32 &&
b == 1) {
1052 return ROCDL::mfma_f32_16x16x32_fp8_bf8::getOperationName();
1054 return ROCDL::mfma_f32_16x16x32_fp8_fp8::getOperationName();
1056 if (m == 32 && n == 32 && k == 16 &&
b == 1) {
1058 return ROCDL::mfma_f32_32x32x16_fp8_bf8::getOperationName();
1060 return ROCDL::mfma_f32_32x32x16_fp8_fp8::getOperationName();
1064 return std::nullopt;
1067static std::optional<ROCDL::MatrixFormat>
1071 .Case([](Float8E4M3FNType) {
return ROCDL::MatrixFormat::fp8_e4m3; })
1072 .Case([](Float8E5M2Type) {
return ROCDL::MatrixFormat::fp8_e5m2; })
1073 .Case([](Float6E2M3FNType) {
return ROCDL::MatrixFormat::fp6_e2m3; })
1074 .Case([](Float6E3M2FNType) {
return ROCDL::MatrixFormat::fp6_e3m2; })
1075 .Case([](Float4E2M1FNType) {
return ROCDL::MatrixFormat::fp4_e2m1; })
1076 .Default(std::nullopt);
1087 std::tuple<StringRef, ROCDL::MatrixFormat, ROCDL::MatrixFormat>;
1089static std::optional<ScaledMFMAIntrinsic>
1091 uint32_t n, uint32_t k, uint32_t
b,
Chipset chipset) {
1097 return std::nullopt;
1098 if (!isa<Float32Type>(destType))
1099 return std::nullopt;
1101 std::optional<ROCDL::MatrixFormat> aTypeCode =
1103 std::optional<ROCDL::MatrixFormat> bTypeCode =
1105 if (!aTypeCode || !bTypeCode)
1106 return std::nullopt;
1108 if (m == 32 && n == 32 && k == 64 &&
b == 1)
1109 return std::tuple{ROCDL::mfma_scale_f32_32x32x64_f8f6f4::getOperationName(),
1110 *aTypeCode, *bTypeCode};
1111 if (m == 16 && n == 16 && k == 128 &&
b == 1)
1113 ROCDL::mfma_scale_f32_16x16x128_f8f6f4::getOperationName(), *aTypeCode,
1116 return std::nullopt;
1119static std::optional<ScaledMFMAIntrinsic>
1122 mfma.getSourceA().getType(), mfma.getSourceB().getType(),
1123 mfma.getDestC().getType(), mfma.getM(), mfma.getN(), mfma.getK(),
1124 mfma.getBlocks(), chipset);
1127static std::optional<ScaledMFMAIntrinsic>
1130 smfma.getSourceB().getType(),
1131 smfma.getDestC().getType(), smfma.getM(),
1132 smfma.getN(), smfma.getK(), 1u, chipset);
1137static std::optional<StringRef>
1139 Type elemDestType, uint32_t k,
bool isRDNA3) {
1140 using fp8 = Float8E4M3FNType;
1141 using bf8 = Float8E5M2Type;
1146 if (elemSourceType.
isF16() && elemDestType.
isF32())
1147 return ROCDL::wmma_f32_16x16x16_f16::getOperationName();
1148 if (elemSourceType.
isBF16() && elemDestType.
isF32())
1149 return ROCDL::wmma_f32_16x16x16_bf16::getOperationName();
1150 if (elemSourceType.
isF16() && elemDestType.
isF16())
1151 return ROCDL::wmma_f16_16x16x16_f16::getOperationName();
1153 return ROCDL::wmma_bf16_16x16x16_bf16::getOperationName();
1155 return ROCDL::wmma_i32_16x16x16_iu8::getOperationName();
1160 return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();
1161 return std::nullopt;
1165 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType) &&
1166 elemDestType.
isF32())
1167 return ROCDL::wmma_f32_16x16x16_fp8_fp8::getOperationName();
1168 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType) &&
1169 elemDestType.
isF32())
1170 return ROCDL::wmma_f32_16x16x16_fp8_bf8::getOperationName();
1171 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType) &&
1172 elemDestType.
isF32())
1173 return ROCDL::wmma_f32_16x16x16_bf8_bf8::getOperationName();
1174 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType) &&
1175 elemDestType.
isF32())
1176 return ROCDL::wmma_f32_16x16x16_bf8_fp8::getOperationName();
1178 return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();
1180 return std::nullopt;
1184 if (k == 32 && !isRDNA3) {
1186 return ROCDL::wmma_i32_16x16x32_iu4::getOperationName();
1189 return std::nullopt;
1195 Type elemBSourceType,
1198 using fp8 = Float8E4M3FNType;
1199 using bf8 = Float8E5M2Type;
1202 if (elemSourceType.
isF32() && elemDestType.
isF32())
1203 return ROCDL::wmma_f32_16x16x4_f32::getOperationName();
1205 return std::nullopt;
1209 if (elemSourceType.
isF16() && elemDestType.
isF32())
1210 return ROCDL::wmma_f32_16x16x32_f16::getOperationName();
1211 if (elemSourceType.
isBF16() && elemDestType.
isF32())
1212 return ROCDL::wmma_f32_16x16x32_bf16::getOperationName();
1213 if (elemSourceType.
isF16() && elemDestType.
isF16())
1214 return ROCDL::wmma_f16_16x16x32_f16::getOperationName();
1216 return ROCDL::wmma_bf16_16x16x32_bf16::getOperationName();
1218 return std::nullopt;
1222 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1223 if (elemDestType.
isF32())
1224 return ROCDL::wmma_f32_16x16x64_fp8_fp8::getOperationName();
1225 if (elemDestType.
isF16())
1226 return ROCDL::wmma_f16_16x16x64_fp8_fp8::getOperationName();
1228 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1229 if (elemDestType.
isF32())
1230 return ROCDL::wmma_f32_16x16x64_fp8_bf8::getOperationName();
1231 if (elemDestType.
isF16())
1232 return ROCDL::wmma_f16_16x16x64_fp8_bf8::getOperationName();
1234 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1235 if (elemDestType.
isF32())
1236 return ROCDL::wmma_f32_16x16x64_bf8_bf8::getOperationName();
1237 if (elemDestType.
isF16())
1238 return ROCDL::wmma_f16_16x16x64_bf8_bf8::getOperationName();
1240 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1241 if (elemDestType.
isF32())
1242 return ROCDL::wmma_f32_16x16x64_bf8_fp8::getOperationName();
1243 if (elemDestType.
isF16())
1244 return ROCDL::wmma_f16_16x16x64_bf8_fp8::getOperationName();
1247 return ROCDL::wmma_i32_16x16x64_iu8::getOperationName();
1249 return std::nullopt;
1253 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1254 if (elemDestType.
isF32())
1255 return ROCDL::wmma_f32_16x16x128_fp8_fp8::getOperationName();
1256 if (elemDestType.
isF16())
1257 return ROCDL::wmma_f16_16x16x128_fp8_fp8::getOperationName();
1259 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1260 if (elemDestType.
isF32())
1261 return ROCDL::wmma_f32_16x16x128_fp8_bf8::getOperationName();
1262 if (elemDestType.
isF16())
1263 return ROCDL::wmma_f16_16x16x128_fp8_bf8::getOperationName();
1265 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1266 if (elemDestType.
isF32())
1267 return ROCDL::wmma_f32_16x16x128_bf8_bf8::getOperationName();
1268 if (elemDestType.
isF16())
1269 return ROCDL::wmma_f16_16x16x128_bf8_bf8::getOperationName();
1271 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1272 if (elemDestType.
isF32())
1273 return ROCDL::wmma_f32_16x16x128_bf8_fp8::getOperationName();
1274 if (elemDestType.
isF16())
1275 return ROCDL::wmma_f16_16x16x128_bf8_fp8::getOperationName();
1278 return std::nullopt;
1281 return std::nullopt;
1289 bool isGfx950 = chipset >=
kGfx950;
1293 uint32_t m = op.getM(), n = op.getN(), k = op.getK();
1298 if (m == 16 && n == 16 && k == 32) {
1300 return ROCDL::smfmac_f32_16x16x32_f16::getOperationName();
1302 return ROCDL::smfmac_f32_16x16x32_bf16::getOperationName();
1305 if (m == 16 && n == 16 && k == 64) {
1308 return ROCDL::smfmac_f32_16x16x64_f16::getOperationName();
1310 return ROCDL::smfmac_f32_16x16x64_bf16::getOperationName();
1314 return ROCDL::smfmac_i32_16x16x64_i8::getOperationName();
1315 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1316 return ROCDL::smfmac_f32_16x16x64_fp8_fp8::getOperationName();
1317 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1318 return ROCDL::smfmac_f32_16x16x64_fp8_bf8::getOperationName();
1319 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1320 return ROCDL::smfmac_f32_16x16x64_bf8_fp8::getOperationName();
1321 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1322 return ROCDL::smfmac_f32_16x16x64_bf8_bf8::getOperationName();
1325 if (m == 16 && n == 16 && k == 128 && isGfx950) {
1328 return ROCDL::smfmac_i32_16x16x128_i8::getOperationName();
1329 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1330 return ROCDL::smfmac_f32_16x16x128_fp8_fp8::getOperationName();
1331 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1332 return ROCDL::smfmac_f32_16x16x128_fp8_bf8::getOperationName();
1333 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1334 return ROCDL::smfmac_f32_16x16x128_bf8_fp8::getOperationName();
1335 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1336 return ROCDL::smfmac_f32_16x16x128_bf8_bf8::getOperationName();
1339 if (m == 32 && n == 32 && k == 16) {
1341 return ROCDL::smfmac_f32_32x32x16_f16::getOperationName();
1343 return ROCDL::smfmac_f32_32x32x16_bf16::getOperationName();
1346 if (m == 32 && n == 32 && k == 32) {
1349 return ROCDL::smfmac_f32_32x32x32_f16::getOperationName();
1351 return ROCDL::smfmac_f32_32x32x32_bf16::getOperationName();
1355 return ROCDL::smfmac_i32_32x32x32_i8::getOperationName();
1356 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1357 return ROCDL::smfmac_f32_32x32x32_fp8_fp8::getOperationName();
1358 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1359 return ROCDL::smfmac_f32_32x32x32_fp8_bf8::getOperationName();
1360 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1361 return ROCDL::smfmac_f32_32x32x32_bf8_fp8::getOperationName();
1362 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1363 return ROCDL::smfmac_f32_32x32x32_bf8_bf8::getOperationName();
1366 if (m == 32 && n == 32 && k == 64 && isGfx950) {
1369 return ROCDL::smfmac_i32_32x32x64_i8::getOperationName();
1370 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1371 return ROCDL::smfmac_f32_32x32x64_fp8_fp8::getOperationName();
1372 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1373 return ROCDL::smfmac_f32_32x32x64_fp8_bf8::getOperationName();
1374 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1375 return ROCDL::smfmac_f32_32x32x64_bf8_fp8::getOperationName();
1376 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1377 return ROCDL::smfmac_f32_32x32x64_bf8_bf8::getOperationName();
1380 return std::nullopt;
1388 auto sourceVectorType = cast<VectorType>(wmma.getSourceA().getType());
1389 auto sourceBVectorType = cast<VectorType>(wmma.getSourceB().getType());
1390 auto destVectorType = cast<VectorType>(wmma.getDestC().getType());
1391 Type elemSourceType = sourceVectorType.getElementType();
1392 Type elemBSourceType = sourceBVectorType.getElementType();
1393 Type elemDestType = destVectorType.getElementType();
1395 const uint32_t k = wmma.getK();
1400 if (isRDNA3 || isRDNA4)
1409 return std::nullopt;
1422static std::optional<SparseWMMAOpInfo>
1428 uint32_t m = swmmac.getM(), n = swmmac.getN(), k = swmmac.getK();
1430 if ((m != 16) || (n != 16))
1431 return std::nullopt;
1438 ROCDL::swmmac_f32_16x16x32_f16::getOperationName(),
false,
false,
1442 ROCDL::swmmac_f32_16x16x32_bf16::getOperationName(),
false,
false,
1446 ROCDL::swmmac_f16_16x16x32_f16::getOperationName(),
false,
false,
1450 ROCDL::swmmac_bf16_16x16x32_bf16::getOperationName(),
false,
false,
1455 ROCDL::swmmac_i32_16x16x32_iu8::getOperationName(),
true,
false,
1460 ROCDL::swmmac_i32_16x16x32_iu4::getOperationName(),
true,
false,
1465 ROCDL::swmmac_f32_16x16x32_fp8_fp8::getOperationName(),
false,
1470 ROCDL::swmmac_f32_16x16x32_fp8_bf8::getOperationName(),
false,
1475 ROCDL::swmmac_f32_16x16x32_bf8_fp8::getOperationName(),
false,
1479 ROCDL::swmmac_f32_16x16x32_bf8_bf8::getOperationName(),
false,
1486 ROCDL::swmmac_i32_16x16x64_iu4::getOperationName(),
true,
false,
1491 const bool isGFX1250 = chipset ==
kGfx1250;
1492 const bool isWavesize64 = swmmac.getWave64();
1493 if (isGFX1250 && !isWavesize64) {
1497 ROCDL::swmmac_f32_16x16x64_f16::getOperationName(),
true,
true,
1501 ROCDL::swmmac_f32_16x16x64_bf16::getOperationName(),
true,
true,
1505 ROCDL::swmmac_f16_16x16x64_f16::getOperationName(),
true,
true,
1509 ROCDL::swmmac_bf16_16x16x64_bf16::getOperationName(),
true,
true,
1516 ROCDL::swmmac_f32_16x16x128_fp8_fp8::getOperationName(),
false,
1521 ROCDL::swmmac_f32_16x16x128_fp8_bf8::getOperationName(),
false,
1526 ROCDL::swmmac_f32_16x16x128_bf8_fp8::getOperationName(),
false,
1530 ROCDL::swmmac_f32_16x16x128_bf8_bf8::getOperationName(),
false,
1535 ROCDL::swmmac_f16_16x16x128_fp8_fp8::getOperationName(),
false,
1540 ROCDL::swmmac_f16_16x16x128_fp8_bf8::getOperationName(),
false,
1545 ROCDL::swmmac_f16_16x16x128_bf8_fp8::getOperationName(),
false,
1549 ROCDL::swmmac_f16_16x16x128_bf8_bf8::getOperationName(),
false,
1554 ROCDL::swmmac_f16_16x16x128_bf8_bf8::getOperationName(),
false,
1559 ROCDL::swmmac_i32_16x16x128_iu8::getOperationName(),
true,
true,
1564 return std::nullopt;
1569 MFMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1570 : ConvertOpToLLVMPattern<MFMAOp>(converter), chipset(chipset) {}
1575 matchAndRewrite(MFMAOp op, MFMAOpAdaptor adaptor,
1576 ConversionPatternRewriter &rewriter)
const override {
1577 Location loc = op.getLoc();
1579 Type outType = typeConverter->convertType(op.getDestD().getType());
1580 Type intrinsicOutType = outType;
1581 if (
auto outVecType = dyn_cast<VectorType>(outType))
1582 if (outVecType.getElementType().isBF16())
1583 intrinsicOutType = outVecType.clone(rewriter.getI16Type());
1585 if (chipset.majorVersion != 9 || chipset <
kGfx908)
1586 return op->emitOpError(
"MFMA only supported on gfx908+");
1587 uint32_t getBlgpField =
static_cast<uint32_t
>(op.getBlgp());
1588 if (op.getNegateA() || op.getNegateB() || op.getNegateC()) {
1590 return op.emitOpError(
"negation unsupported on older than gfx942");
1592 op.getNegateA() | (op.getNegateB() << 1) | (op.getNegateC() << 2);
1595 std::optional<ScaledMFMAIntrinsic> maybeScaledIntrinsic =
1597 if (!maybeIntrinsic.has_value() && !maybeScaledIntrinsic.has_value())
1598 return op.emitOpError(
"no intrinsic matching MFMA size on given chipset");
1601 !maybeIntrinsic.has_value() && maybeScaledIntrinsic.has_value();
1603 (adaptor.getAbid() > 0 || getBlgpField > 0 || op.getCbsz() > 0)) {
1604 return op.emitOpError(
1605 "non-default abid, blgp, and cbsz aren't supported on MFMAs that can "
1606 "be scaled as those fields are used for type information");
1609 StringRef intrinsicName =
1610 isScaled ? std::get<0>(*maybeScaledIntrinsic) : *maybeIntrinsic;
1613 bool allowBf16 = [&]() {
1618 return intrinsicName.contains(
"16x16x32.bf16") ||
1619 intrinsicName.contains(
"32x32x16.bf16");
1621 OperationState loweredOp(loc, intrinsicName);
1622 loweredOp.addTypes(intrinsicOutType);
1624 rewriter, loc, adaptor.getSourceA(), allowBf16),
1626 rewriter, loc, adaptor.getSourceB(), allowBf16),
1627 adaptor.getDestC()});
1630 auto [_scaledName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;
1631 loweredOp.addOperands({zero, zero});
1632 loweredOp.addAttributes(
1634 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), aTypeCode)},
1636 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), bTypeCode)},
1637 {
"opselA", rewriter.getI32IntegerAttr(0)},
1638 {
"opselB", rewriter.getI32IntegerAttr(0)}});
1640 Attribute blgpAttr =
1642 ? Attribute(ROCDL::MFMANegModifierAttr::get(
1643 rewriter.getContext(),
1644 static_cast<ROCDL::MFMANegModifier
>(getBlgpField)))
1645 : Attribute(ROCDL::MFMAPermBAttr::
get(
1647 static_cast<ROCDL::MFMAPermB>(getBlgpField)));
1648 loweredOp.addAttributes(
1649 {{
"cbsz", rewriter.getI32IntegerAttr(op.getCbsz())},
1650 {
"abid", rewriter.getI32IntegerAttr(op.getAbid())},
1651 {
"blgp", blgpAttr}});
1653 Value lowered = rewriter.create(loweredOp)->getResult(0);
1654 if (outType != intrinsicOutType)
1655 lowered = LLVM::BitcastOp::create(rewriter, loc, outType, lowered);
1656 rewriter.replaceOp(op, lowered);
1662 ScaledMFMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1663 : ConvertOpToLLVMPattern(converter), chipset(chipset) {}
1668 matchAndRewrite(ScaledMFMAOp op, ScaledMFMAOpAdaptor adaptor,
1669 ConversionPatternRewriter &rewriter)
const override {
1670 Location loc = op.getLoc();
1671 Type intrinsicOutType = typeConverter->convertType(op.getDestD().getType());
1673 if (chipset.majorVersion != 9 || chipset <
kGfx950)
1674 return op->emitOpError(
"scaled MFMA only supported on gfx908+");
1675 std::optional<ScaledMFMAIntrinsic> maybeScaledIntrinsic =
1677 if (!maybeScaledIntrinsic.has_value())
1678 return op.emitOpError(
1679 "no intrinsic matching scaled MFMA size on given chipset");
1681 auto [intrinsicName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;
1682 OperationState loweredOp(loc, intrinsicName);
1683 loweredOp.addTypes(intrinsicOutType);
1684 loweredOp.addOperands(
1687 adaptor.getDestC()});
1688 loweredOp.addOperands(
1693 loweredOp.addAttributes(
1695 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), aTypeCode)},
1697 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), bTypeCode)},
1698 {
"opselA", rewriter.getI32IntegerAttr(adaptor.getScalesIdxA())},
1699 {
"opselB", rewriter.getI32IntegerAttr(adaptor.getScalesIdxB())}});
1701 Value lowered = rewriter.create(loweredOp)->getResult(0);
1702 rewriter.replaceOp(op, lowered);
1708 SparseMFMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1709 : ConvertOpToLLVMPattern<SparseMFMAOp>(converter), chipset(chipset) {}
1714 matchAndRewrite(SparseMFMAOp op, SparseMFMAOpAdaptor adaptor,
1715 ConversionPatternRewriter &rewriter)
const override {
1716 Location loc = op.getLoc();
1718 typeConverter->convertType<VectorType>(op.getDestC().
getType());
1720 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1723 if (chipset.majorVersion != 9 || chipset <
kGfx942)
1724 return op->emitOpError(
"sparse MFMA (smfmac) only supported on gfx942+");
1727 if (!maybeIntrinsic.has_value())
1728 return op.emitOpError(
1729 "no intrinsic matching sparse MFMA on the given chipset");
1732 ROCDL::smfmac_f32_16x16x32_bf16::getOperationName() ||
1734 ROCDL::smfmac_f32_32x32x16_bf16::getOperationName());
1735 bool isGfx950 = (chipset >=
kGfx950) && !isGfx942BF16;
1741 Value c = adaptor.getDestC();
1745 Value sparseIdx = adaptor.getSparseIdx();
1746 Type i32Type = rewriter.getI32Type();
1747 if (sparseIdx.
getType() != i32Type)
1748 sparseIdx = LLVM::BitcastOp::create(rewriter, loc, i32Type, sparseIdx);
1750 OperationState loweredOp(loc, maybeIntrinsic.value());
1751 loweredOp.addTypes(outType);
1752 loweredOp.addOperands({a,
b, c, sparseIdx});
1753 loweredOp.addAttributes(
1754 {{
"cbsz", rewriter.getI32IntegerAttr(op.getCbsz())},
1755 {
"abid", rewriter.getI32IntegerAttr(op.getAbid())}});
1756 Value lowered = rewriter.create(loweredOp)->getResult(0);
1757 rewriter.replaceOp(op, lowered);
1763 WMMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1764 : ConvertOpToLLVMPattern<WMMAOp>(converter), chipset(chipset) {}
1769 matchAndRewrite(WMMAOp op, WMMAOpAdaptor adaptor,
1770 ConversionPatternRewriter &rewriter)
const override {
1771 Location loc = op.getLoc();
1773 typeConverter->convertType<VectorType>(op.getDestD().
getType());
1775 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1777 if (chipset.majorVersion != 11 && chipset.majorVersion != 12)
1778 return op->emitOpError(
"WMMA only supported on gfx11 and gfx12");
1780 bool isGFX1250 = chipset >=
kGfx1250;
1785 auto aType = cast<VectorType>(adaptor.getSourceA().getType());
1786 auto bType = cast<VectorType>(adaptor.getSourceB().getType());
1787 auto destCType = cast<VectorType>(adaptor.getDestC().getType());
1788 bool castAToI16 = aType.getElementType().isBF16() && !isGFX1250;
1789 bool castBToI16 = bType.getElementType().isBF16() && !isGFX1250;
1790 bool castDestCToI16 = destCType.getElementType().isBF16() && !isGFX1250;
1791 bool castOutToI16 = outType.getElementType().
isBF16() && !isGFX1250;
1792 VectorType rawOutType = outType;
1794 rawOutType = outType.clone(rewriter.getI16Type());
1795 Value a = adaptor.getSourceA();
1797 a = LLVM::BitcastOp::create(rewriter, loc,
1798 aType.clone(rewriter.getI16Type()), a);
1799 Value
b = adaptor.getSourceB();
1801 b = LLVM::BitcastOp::create(rewriter, loc,
1802 bType.clone(rewriter.getI16Type()),
b);
1803 Value destC = adaptor.getDestC();
1805 destC = LLVM::BitcastOp::create(
1806 rewriter, loc, destCType.clone(rewriter.getI16Type()), destC);
1810 if (!maybeIntrinsic.has_value())
1811 return op.emitOpError(
"no intrinsic matching WMMA on the given chipset");
1813 if (chipset.majorVersion >= 12 && op.getSubwordOffset() != 0)
1814 return op.emitOpError(
"subwordOffset not supported on gfx12+");
1816 SmallVector<Value, 4> operands;
1817 SmallVector<NamedAttribute, 4> attrs;
1819 op.getSourceA(), operands, attrs,
"signA");
1821 op.getSourceB(), operands, attrs,
"signB");
1823 op.getSubwordOffset(), op.getClamp(), operands,
1826 OperationState loweredOp(loc, *maybeIntrinsic);
1827 loweredOp.addTypes(rawOutType);
1828 loweredOp.addOperands(operands);
1829 loweredOp.addAttributes(attrs);
1830 Operation *lowered = rewriter.create(loweredOp);
1832 Operation *maybeCastBack = lowered;
1833 if (rawOutType != outType)
1834 maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,
1836 rewriter.replaceOp(op, maybeCastBack->
getResults());
1842enum class DotFamily {
1851static std::optional<std::pair<StringRef, DotFamily>>
1852dotOpToIntrinsic(DotOp op,
Chipset chipset) {
1853 Type aElem = cast<VectorType>(op.getSourceA().getType()).getElementType();
1854 Type bElem = cast<VectorType>(op.getSourceB().getType()).getElementType();
1855 Type dest = op.getDestC().getType();
1856 bool uA = op.getUnsignedA();
1857 bool uB = op.getUnsignedB();
1862 return {{ROCDL::fdot2::getOperationName(), DotFamily::Clamp}};
1864 return {{ROCDL::fdot2_f16_f16::getOperationName(), DotFamily::NoClamp}};
1865 return std::nullopt;
1871 return {{ROCDL::fdot2_f32_bf16::getOperationName(), DotFamily::Clamp}};
1873 return {{ROCDL::fdot2_bf16_bf16::getOperationName(), DotFamily::NoClamp}};
1874 return std::nullopt;
1878 if (isa<IntegerType>(aElem) && isa<IntegerType>(bElem) &&
1880 bool mixedSign = (uA != uB);
1885 return std::nullopt;
1887 switch (elemWidth) {
1889 name = ROCDL::sudot4::getOperationName();
1892 name = ROCDL::sudot8::getOperationName();
1895 return std::nullopt;
1897 return {{name, DotFamily::Sudot}};
1901 bool supported =
false;
1902 switch (elemWidth) {
1905 name = uA ? ROCDL::udot2::getOperationName()
1906 :
ROCDL::sdot2::getOperationName();
1911 name = uA ? ROCDL::udot4::getOperationName()
1912 :
ROCDL::sdot4::getOperationName();
1917 name = uA ? ROCDL::udot8::getOperationName()
1918 :
ROCDL::sdot8::getOperationName();
1921 return std::nullopt;
1924 return std::nullopt;
1925 return {{name, DotFamily::Clamp}};
1929 bool aIsFp8 = isa<Float8E4M3FNType>(aElem);
1930 bool aIsBf8 = isa<Float8E5M2Type>(aElem);
1931 bool bIsFp8 = isa<Float8E4M3FNType>(bElem);
1932 bool bIsBf8 = isa<Float8E5M2Type>(bElem);
1933 if ((aIsFp8 || aIsBf8) && (bIsFp8 || bIsBf8) && dest.
isF32()) {
1935 return std::nullopt;
1937 if (aIsFp8 && bIsFp8)
1938 name = ROCDL::dot4_f32_fp8_fp8::getOperationName();
1939 else if (aIsFp8 && bIsBf8)
1940 name = ROCDL::dot4_f32_fp8_bf8::getOperationName();
1941 else if (aIsBf8 && bIsFp8)
1942 name = ROCDL::dot4_f32_bf8_fp8::getOperationName();
1944 name = ROCDL::dot4_f32_bf8_bf8::getOperationName();
1945 return {{name, DotFamily::NoClamp}};
1948 return std::nullopt;
1952 DotOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1953 : ConvertOpToLLVMPattern<DotOp>(converter), chipset(chipset) {}
1958 matchAndRewrite(DotOp op, DotOpAdaptor adaptor,
1959 ConversionPatternRewriter &rewriter)
const override {
1960 Location loc = op.getLoc();
1962 std::optional<std::pair<StringRef, DotFamily>> maybeIntrinsic =
1963 dotOpToIntrinsic(op, chipset);
1964 if (!maybeIntrinsic)
1965 return op.emitOpError(
"no intrinsic matching dot on the given chipset: ")
1966 << op.getSourceA().getType() <<
" * " << op.getSourceB().getType()
1967 <<
" + " << op.getDestC().getType();
1969 auto [intrinsicName, family] = maybeIntrinsic.value();
1973 Value c = adaptor.getDestC();
1975 SmallVector<NamedAttribute, 3> attrs;
1976 if (family == DotFamily::Sudot) {
1977 attrs.push_back(rewriter.getNamedAttr(
1978 "signA", rewriter.getBoolAttr(!op.getUnsignedA())));
1979 attrs.push_back(rewriter.getNamedAttr(
1980 "signB", rewriter.getBoolAttr(!op.getUnsignedB())));
1983 if (family != DotFamily::NoClamp && op.getClamp())
1985 rewriter.getNamedAttr(
"clamp", rewriter.getBoolAttr(
true)));
1987 Type resultType = typeConverter->convertType(op.getDestD().getType());
1989 OperationState loweredOp(loc, intrinsicName);
1990 loweredOp.addTypes(resultType);
1991 loweredOp.addOperands({a,
b, c});
1992 loweredOp.addAttributes(attrs);
1993 Operation *lowered = rewriter.create(loweredOp);
1994 rewriter.replaceOp(op, lowered->
getResults());
2000 SparseWMMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2001 : ConvertOpToLLVMPattern<SparseWMMAOp>(converter), chipset(chipset) {}
2006 matchAndRewrite(SparseWMMAOp op, SparseWMMAOpAdaptor adaptor,
2007 ConversionPatternRewriter &rewriter)
const override {
2008 Location loc = op.getLoc();
2010 typeConverter->convertType<VectorType>(op.getDestD().
getType());
2012 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2014 std::optional<SparseWMMAOpInfo> maybeIntrinsic =
2017 if (!maybeIntrinsic.has_value())
2018 return op.emitOpError(
2019 "no intrinsic matching Sparse WMMA on the given chipset");
2020 SparseWMMAOpInfo intrinsic = maybeIntrinsic.value();
2022 SmallVector<NamedAttribute> attrs;
2024 if ((op.getUnsignedA() || op.getUnsignedB()) && !intrinsic.
useSign)
2025 return op->emitOpError(
"intrinsic doesn't support unsign");
2027 if (
auto attr = op.getUnsignedAAttr())
2028 attrs.push_back({
"signA", attr});
2029 if (
auto attr = op.getUnsignedBAttr())
2030 attrs.push_back({
"signB", attr});
2033 if ((op.getReuseA() || op.getReuseB()) && !intrinsic.
useReuse)
2034 return op->emitOpError(
"intrinsic doesn't support reuse");
2036 if (
auto attr = op.getReuseAAttr())
2037 attrs.push_back({
"reuseA", attr});
2038 if (
auto attr = op.getReuseBAttr())
2039 attrs.push_back({
"reuseB", attr});
2042 if (op.getClamp() && !intrinsic.
useClamp)
2043 return op->emitOpError(
"intrinsic doesn't support clamp");
2044 if (intrinsic.
useClamp && op.getClampAttr())
2045 attrs.push_back({
"clamp", op.getClampAttr()});
2047 const bool isGFX1250orHigher =
2048 chipset.majorVersion == 12 && chipset.minorVersion >= 5;
2053 Value c = adaptor.getDestC();
2054 VectorType rawOutType = outType;
2055 if (!isGFX1250orHigher) {
2057 rawOutType = cast<VectorType>(c.
getType());
2061 Value sparseIdx = LLVM::BitcastOp::create(
2062 rewriter, loc, rewriter.getI32Type(), adaptor.getSparseIdx());
2064 OperationState loweredOp(loc, intrinsic.
name);
2065 loweredOp.addTypes(rawOutType);
2066 loweredOp.addOperands({a,
b, c, sparseIdx});
2067 loweredOp.addAttributes(attrs);
2068 Operation *lowered = rewriter.create(loweredOp);
2070 Operation *maybeCastBack = lowered;
2071 if (rawOutType != outType)
2072 maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,
2074 rewriter.replaceOp(op, maybeCastBack->
getResults());
2081 ScaledWMMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2082 : ConvertOpToLLVMPattern<ScaledWMMAOp>(converter), chipset(chipset) {}
2087 matchAndRewrite(ScaledWMMAOp op, ScaledWMMAOpAdaptor adaptor,
2088 ConversionPatternRewriter &rewriter)
const override {
2089 Location loc = op.getLoc();
2091 typeConverter->convertType<VectorType>(op.getDestD().
getType());
2093 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2096 return op->emitOpError(
"WMMA scale only supported on gfx1250+");
2098 int64_t m = op.getM();
2099 int64_t n = op.getN();
2100 int64_t k = op.getK();
2105 std::optional<ROCDL::MatrixFormat> aFmtCode =
2107 std::optional<ROCDL::MatrixFormat> bFmtCode =
2110 if (!aFmtCode || !bFmtCode)
2111 return op.emitOpError(
"unsupported element types for scaled_wmma");
2114 auto scaleAVecType = cast<VectorType>(op.getScaleA().getType());
2115 auto scaleBVecType = cast<VectorType>(op.getScaleB().getType());
2117 if (scaleAVecType.getNumElements() != scaleBVecType.getNumElements())
2118 return op.emitOpError(
"scaleA and scaleB must have equal vector length");
2121 Type scaleAElemType = scaleAVecType.getElementType();
2122 Type scaleBElemType = scaleBVecType.getElementType();
2124 std::optional<ROCDL::WMMAMatrixScaleFormat> scaleAFmt =
2126 std::optional<ROCDL::WMMAMatrixScaleFormat> scaleBFmt =
2129 if (!scaleAFmt || !scaleBFmt)
2130 return op.emitOpError(
"unsupported scale element types");
2133 bool isScale16 = (scaleAVecType.getNumElements() == 8);
2134 std::optional<StringRef> intrinsicName =
2137 return op.emitOpError(
"unsupported scaled_wmma dimensions: ")
2138 << m <<
"x" << n <<
"x" << k;
2140 SmallVector<NamedAttribute, 8> attrs;
2143 bool is32x16 = (m == 32 && n == 16 && k == 128);
2145 attrs.emplace_back(
"fmtA", ROCDL::MatrixFormatAttr::get(
2146 rewriter.getContext(), *aFmtCode));
2147 attrs.emplace_back(
"fmtB", ROCDL::MatrixFormatAttr::get(
2148 rewriter.getContext(), *bFmtCode));
2153 "modC", ROCDL::WMMACModifierAttr::get(rewriter.getContext(),
2154 ROCDL::WMMACModifier::none));
2158 attrs.emplace_back(
"scaleAType", ROCDL::WMMAMatrixScaleAttr::get(
2159 rewriter.getContext(),
2160 static_cast<ROCDL::WMMAMatrixScale
>(
2161 op.getAFirstScaleLane() / 16)));
2162 attrs.emplace_back(
"fmtScaleA", ROCDL::WMMAMatrixScaleFormatAttr::get(
2163 rewriter.getContext(), *scaleAFmt));
2164 attrs.emplace_back(
"scaleBType", ROCDL::WMMAMatrixScaleAttr::get(
2165 rewriter.getContext(),
2166 static_cast<ROCDL::WMMAMatrixScale
>(
2167 op.getBFirstScaleLane() / 16)));
2168 attrs.emplace_back(
"fmtScaleB", ROCDL::WMMAMatrixScaleFormatAttr::get(
2169 rewriter.getContext(), *scaleBFmt));
2172 attrs.emplace_back(
"reuseA", rewriter.getBoolAttr(
false));
2173 attrs.emplace_back(
"reuseB", rewriter.getBoolAttr(
false));
2186 OperationState loweredOp(loc, *intrinsicName);
2187 loweredOp.addTypes(outType);
2188 loweredOp.addOperands(
2189 {sourceA, sourceB, adaptor.getDestC(), packedScaleA, packedScaleB});
2190 loweredOp.addAttributes(attrs);
2192 Operation *lowered = rewriter.create(loweredOp);
2193 rewriter.replaceOp(op, lowered->
getResults());
2199struct TransposeLoadOpLowering
2201 TransposeLoadOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2202 : ConvertOpToLLVMPattern<TransposeLoadOp>(converter), chipset(chipset) {}
2207 matchAndRewrite(TransposeLoadOp op, TransposeLoadOpAdaptor adaptor,
2208 ConversionPatternRewriter &rewriter)
const override {
2210 return op.emitOpError(
2211 "transpose_load is only supported on gfx950 and gfx1250+");
2213 Location loc = op.getLoc();
2214 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2218 size_t srcElementSize =
2219 srcMemRefType.getElementType().getIntOrFloatBitWidth();
2220 if (srcElementSize < 8)
2221 return op.emitOpError(
"Expect source memref to have at least 8 bits "
2222 "element size, got ")
2225 auto resultType = cast<VectorType>(op.getResult().getType());
2228 (adaptor.getSrcIndices()));
2230 size_t numElements = resultType.getNumElements();
2231 size_t elementTypeSize =
2234 Type llvmResultType = typeConverter->convertType(resultType);
2237 Type rocdlResultType =
2238 elementTypeSize < 16
2239 ? VectorType::get((numElements * elementTypeSize) / 32,
2240 rewriter.getIntegerType(32))
2243 auto emitNumElementsError = [&](
size_t expected, StringRef chipsetName) {
2244 return op.emitOpError()
2245 << elementTypeSize <<
"-bit transpose_load requires " << expected
2246 <<
" elements on " << chipsetName;
2251 switch (elementTypeSize) {
2253 if (numElements != 16)
2254 return emitNumElementsError(16,
"gfx1250+");
2256 ROCDL::DsLoadTr4_B64::create(rewriter, loc, rocdlResultType, srcPtr)
2261 if (numElements != 16)
2262 return emitNumElementsError(16,
"gfx1250+");
2264 ROCDL::DsLoadTr6_B96::create(rewriter, loc, rocdlResultType, srcPtr)
2269 if (numElements != 8)
2270 return emitNumElementsError(8,
"gfx1250+");
2272 ROCDL::DsLoadTr8_B64::create(rewriter, loc, rocdlResultType, srcPtr)
2277 if (numElements != 8)
2278 return emitNumElementsError(8,
"gfx1250+");
2279 intrinsic = ROCDL::DsLoadTr16_B128::create(rewriter, loc,
2280 rocdlResultType, srcPtr)
2285 return op.emitOpError(
"Unsupported element size for transpose load");
2288 switch (elementTypeSize) {
2290 if (numElements != 16)
2291 return emitNumElementsError(16,
"gfx950");
2292 intrinsic = ROCDL::ds_read_tr4_b64::create(rewriter, loc,
2293 rocdlResultType, srcPtr)
2298 if (numElements != 16)
2299 return emitNumElementsError(16,
"gfx950");
2300 intrinsic = ROCDL::ds_read_tr6_b96::create(rewriter, loc,
2301 rocdlResultType, srcPtr)
2306 if (numElements != 8)
2307 return emitNumElementsError(8,
"gfx950");
2308 intrinsic = ROCDL::ds_read_tr8_b64::create(rewriter, loc,
2309 rocdlResultType, srcPtr)
2314 if (numElements != 4)
2315 return emitNumElementsError(4,
"gfx950");
2316 intrinsic = ROCDL::ds_read_tr16_b64::create(rewriter, loc,
2317 rocdlResultType, srcPtr)
2322 return op.emitOpError(
"Unsupported element size for transpose load");
2326 assert(intrinsic &&
"expected ROCDL transpose load intrinsic");
2327 if (intrinsic.
getType() == llvmResultType) {
2328 rewriter.replaceOp(op, intrinsic);
2331 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, intrinsic);
2336struct GlobalTransposeLoadOpLowering
2338 GlobalTransposeLoadOpLowering(
const LLVMTypeConverter &converter,
2340 : ConvertOpToLLVMPattern<GlobalTransposeLoadOp>(converter),
2346 matchAndRewrite(GlobalTransposeLoadOp op,
2347 GlobalTransposeLoadOpAdaptor adaptor,
2348 ConversionPatternRewriter &rewriter)
const override {
2350 return op.emitOpError(
2351 "global_transpose_load is only supported on gfx1200+");
2353 Location loc = op.getLoc();
2354 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2355 auto resultType = cast<VectorType>(op.getResult().getType());
2358 rewriter, loc, srcMemRefType, adaptor.getSrc(), adaptor.getSrcIndices(),
2359 LLVM::GEPNoWrapFlags::inbounds | LLVM::GEPNoWrapFlags::nuw);
2361 size_t numElements = resultType.getNumElements();
2362 size_t elementTypeSize =
2367 Type rocdlResultType =
2368 elementTypeSize < 16
2369 ? VectorType::get((numElements * elementTypeSize) / 32,
2370 rewriter.getIntegerType(32))
2371 : typeConverter->convertType(resultType);
2372 Type llvmResultType = typeConverter->convertType(resultType);
2374 switch (elementTypeSize) {
2376 assert(numElements == 16);
2378 return op.emitOpError(
"4-bit global_transpose_load requires gfx1250+");
2379 auto rocdlOp = ROCDL::GlobalLoadTr4_B64::create(rewriter, loc,
2380 rocdlResultType, srcPtr);
2381 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2385 assert(numElements == 16);
2387 return op.emitOpError(
"6-bit global_transpose_load requires gfx1250+");
2388 auto rocdlOp = ROCDL::GlobalLoadTr6_B96::create(rewriter, loc,
2389 rocdlResultType, srcPtr);
2390 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2394 assert(numElements == 8);
2395 auto rocdlOp = ROCDL::GlobalLoadTr8_B64::create(rewriter, loc,
2396 rocdlResultType, srcPtr);
2397 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2401 assert(numElements == 8);
2402 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadTr8_B128>(op, llvmResultType,
2407 return op.emitOpError(
2408 "unsupported element size for global transpose load");
2415 GatherToLDSOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2416 : ConvertOpToLLVMPattern<GatherToLDSOp>(converter), chipset(chipset) {}
2421 matchAndRewrite(GatherToLDSOp op, GatherToLDSOpAdaptor adaptor,
2422 ConversionPatternRewriter &rewriter)
const override {
2423 if (chipset.majorVersion < 9 || chipset.majorVersion > 10)
2424 return op.emitOpError(
"pre-gfx9 and post-gfx10 not supported");
2426 Location loc = op.getLoc();
2428 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2429 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2434 Type transferType = op.getTransferType();
2435 int loadWidth = [&]() ->
int {
2436 if (
auto transferVectorType = dyn_cast<VectorType>(transferType)) {
2437 return (transferVectorType.getNumElements() *
2438 transferVectorType.getElementTypeBitWidth()) /
2445 if (!llvm::is_contained({1, 2, 4, 12, 16}, loadWidth))
2446 return op.emitOpError(
"chipset unsupported element size");
2448 if (chipset !=
kGfx950 && llvm::is_contained({12, 16}, loadWidth))
2449 return op.emitOpError(
"Gather to LDS instructions with 12-byte and "
2450 "16-byte load widths are only supported on gfx950");
2454 (adaptor.getSrcIndices()));
2457 (adaptor.getDstIndices()));
2459 if (op.getAsync()) {
2460 rewriter.replaceOpWithNewOp<ROCDL::LoadAsyncToLDSOp>(
2461 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2462 rewriter.getI32IntegerAttr(0),
2466 rewriter.replaceOpWithNewOp<ROCDL::LoadToLDSOp>(
2467 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2468 rewriter.getI32IntegerAttr(0),
2477struct GlobalLoadAsyncToLDSOpLowering
2479 GlobalLoadAsyncToLDSOpLowering(
const LLVMTypeConverter &converter,
2481 : ConvertOpToLLVMPattern<GlobalLoadAsyncToLDSOp>(converter),
2487 matchAndRewrite(GlobalLoadAsyncToLDSOp op,
2488 GlobalLoadAsyncToLDSOpAdaptor adaptor,
2489 ConversionPatternRewriter &rewriter)
const override {
2491 return op.emitOpError(
2492 "global_load_async_to_lds is only supported on gfx1250+");
2494 Location loc = op.getLoc();
2495 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2496 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2498 Type transferType = op.getTransferType();
2500 isa<VectorType>(transferType)
2501 ? cast<VectorType>(transferType).getNumElements() *
2502 cast<VectorType>(transferType).getElementTypeBitWidth()
2507 adaptor.getSrcIndices());
2510 adaptor.getDstIndices());
2513 Value mask = adaptor.getMask();
2514 int64_t nullptrVal =
2515 llvm::AMDGPU::getNullPointerValue(llvm::AMDGPUAS::LOCAL_ADDRESS);
2519 LLVM::IntToPtrOp::create(rewriter, loc, dstPtr.
getType(), nullInt);
2520 dstPtr = LLVM::SelectOp::create(rewriter, loc, mask, dstPtr, nullPtr);
2523 auto offset = rewriter.getI32IntegerAttr(0);
2524 Attribute aux = rewriter.getI32IntegerAttr(0);
2526 switch (transferBits) {
2528 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB8Op>(
2533 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB32Op>(
2538 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB64Op>(
2543 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB128Op>(
2548 return op.emitOpError(
"unsupported transfer width");
2555struct ExtPackedFp8OpLowering final
2557 ExtPackedFp8OpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2558 : ConvertOpToLLVMPattern<amdgpu::ExtPackedFp8Op>(converter),
2563 matchAndRewrite(ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2564 ConversionPatternRewriter &rewriter)
const override;
2567struct ScaledExtPackedMatrixOpLowering final
2569 ScaledExtPackedMatrixOpLowering(
const LLVMTypeConverter &converter,
2571 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedMatrixOp>(converter),
2576 matchAndRewrite(ScaledExtPackedMatrixOp op,
2577 ScaledExtPackedMatrixOpAdaptor adaptor,
2578 ConversionPatternRewriter &rewriter)
const override;
2581struct PackedTrunc2xFp8OpLowering final
2583 PackedTrunc2xFp8OpLowering(
const LLVMTypeConverter &converter,
2585 : ConvertOpToLLVMPattern<amdgpu::PackedTrunc2xFp8Op>(converter),
2590 matchAndRewrite(PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
2591 ConversionPatternRewriter &rewriter)
const override;
2594struct PackedStochRoundFp8OpLowering final
2596 PackedStochRoundFp8OpLowering(
const LLVMTypeConverter &converter,
2598 : ConvertOpToLLVMPattern<amdgpu::PackedStochRoundFp8Op>(converter),
2603 matchAndRewrite(PackedStochRoundFp8Op op,
2604 PackedStochRoundFp8OpAdaptor adaptor,
2605 ConversionPatternRewriter &rewriter)
const override;
2608struct ScaledExtPackedOpLowering final
2610 ScaledExtPackedOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2611 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedOp>(converter),
2616 matchAndRewrite(ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2617 ConversionPatternRewriter &rewriter)
const override;
2620struct PackedScaledTruncOpLowering final
2622 PackedScaledTruncOpLowering(
const LLVMTypeConverter &converter,
2624 : ConvertOpToLLVMPattern<amdgpu::PackedScaledTruncOp>(converter),
2629 matchAndRewrite(PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2630 ConversionPatternRewriter &rewriter)
const override;
2635LogicalResult ExtPackedFp8OpLowering::matchAndRewrite(
2636 ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2637 ConversionPatternRewriter &rewriter)
const {
2638 Location loc = op.getLoc();
2640 return rewriter.notifyMatchFailure(
2641 loc,
"Fp8 conversion instructions are not available on target "
2642 "architecture and their emulation is not implemented");
2644 getTypeConverter()->convertType(VectorType::get(4, rewriter.getI8Type()));
2645 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2646 Type f32 = getTypeConverter()->convertType(op.getResult().getType());
2648 Value source = adaptor.getSource();
2649 auto sourceVecType = dyn_cast<VectorType>(op.getSource().getType());
2650 auto resultVecType = dyn_cast<VectorType>(op.getResult().getType());
2653 if (!sourceVecType || sourceVecType.getNumElements() < 4) {
2654 Value longVec = LLVM::UndefOp::create(rewriter, loc, v4i8);
2655 if (!sourceVecType) {
2656 longVec = LLVM::InsertElementOp::create(
2659 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2661 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2663 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2668 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2669 if (resultVecType) {
2671 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Bf8Op>(op, f32, i32Source,
2674 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Fp8Op>(op, f32, i32Source,
2679 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Bf8Op>(op, f32, i32Source,
2682 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Fp8Op>(op, f32, i32Source,
2689int32_t getScaleSel(int32_t blockSize,
unsigned bitWidth, int32_t scaleWaveHalf,
2690 int32_t firstScaleByte) {
2696 assert(llvm::is_contained({16, 32}, blockSize));
2697 assert(llvm::is_contained({4u, 6u, 8u}, bitWidth));
2699 const bool isFp8 = bitWidth == 8;
2700 const bool isBlock16 = blockSize == 16;
2703 int32_t bit0 = isBlock16;
2704 assert(llvm::is_contained({0, 1, 2}, firstScaleByte));
2705 int32_t bit1 = (firstScaleByte == 2) << 1;
2706 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2707 int32_t bit2 = scaleWaveHalf << 2;
2708 return bit2 | bit1 | bit0;
2711 int32_t bit0 = isBlock16;
2713 assert(llvm::is_contained({0, 1, 2, 3}, firstScaleByte));
2714 int32_t bits2and1 = firstScaleByte << 1;
2715 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2716 int32_t bit3 = scaleWaveHalf << 3;
2717 int32_t bits = bit3 | bits2and1 | bit0;
2719 assert(!llvm::is_contained(
2720 {0b0011, 0b0101, 0b0111, 0b1000, 0b1001, 0b1011, 0b1111}, bits));
2724static std::optional<StringRef>
2725scaledExtPacked816ToIntrinsic(Type srcElemType, Type destElemType) {
2726 using fp4 = Float4E2M1FNType;
2727 using fp8 = Float8E4M3FNType;
2728 using bf8 = Float8E5M2Type;
2729 using fp6 = Float6E2M3FNType;
2730 using bf6 = Float6E3M2FNType;
2731 if (isa<fp4>(srcElemType)) {
2732 if (destElemType.
isF16())
2733 return ROCDL::CvtPkScalePk8F16Fp4Op::getOperationName();
2734 if (destElemType.
isBF16())
2735 return ROCDL::CvtPkScalePk8Bf16Fp4Op::getOperationName();
2736 if (destElemType.
isF32())
2737 return ROCDL::CvtPkScalePk8F32Fp4Op::getOperationName();
2738 return std::nullopt;
2740 if (isa<fp8>(srcElemType)) {
2741 if (destElemType.
isF16())
2742 return ROCDL::CvtPkScalePk8F16Fp8Op::getOperationName();
2743 if (destElemType.
isBF16())
2744 return ROCDL::CvtPkScalePk8Bf16Fp8Op::getOperationName();
2745 if (destElemType.
isF32())
2746 return ROCDL::CvtPkScalePk8F32Fp8Op::getOperationName();
2747 return std::nullopt;
2749 if (isa<bf8>(srcElemType)) {
2750 if (destElemType.
isF16())
2751 return ROCDL::CvtPkScalePk8F16Bf8Op::getOperationName();
2752 if (destElemType.
isBF16())
2753 return ROCDL::CvtPkScalePk8Bf16Bf8Op::getOperationName();
2754 if (destElemType.
isF32())
2755 return ROCDL::CvtPkScalePk8F32Bf8Op::getOperationName();
2756 return std::nullopt;
2758 if (isa<fp6>(srcElemType)) {
2759 if (destElemType.
isF16())
2760 return ROCDL::CvtPkScalePk16F16Fp6Op::getOperationName();
2761 if (destElemType.
isBF16())
2762 return ROCDL::CvtPkScalePk16Bf16Fp6Op::getOperationName();
2763 if (destElemType.
isF32())
2764 return ROCDL::CvtPkScalePk16F32Fp6Op::getOperationName();
2765 return std::nullopt;
2767 if (isa<bf6>(srcElemType)) {
2768 if (destElemType.
isF16())
2769 return ROCDL::CvtPkScalePk16F16Bf6Op::getOperationName();
2770 if (destElemType.
isBF16())
2771 return ROCDL::CvtPkScalePk16Bf16Bf6Op::getOperationName();
2772 if (destElemType.
isF32())
2773 return ROCDL::CvtPkScalePk16F32Bf6Op::getOperationName();
2774 return std::nullopt;
2776 llvm_unreachable(
"invalid combination of element types for packed conversion "
2780LogicalResult ScaledExtPackedMatrixOpLowering::matchAndRewrite(
2781 ScaledExtPackedMatrixOp op, ScaledExtPackedMatrixOpAdaptor adaptor,
2782 ConversionPatternRewriter &rewriter)
const {
2783 using fp4 = Float4E2M1FNType;
2784 using fp8 = Float8E4M3FNType;
2785 using bf8 = Float8E5M2Type;
2786 using fp6 = Float6E2M3FNType;
2787 using bf6 = Float6E3M2FNType;
2788 Location loc = op.getLoc();
2790 return rewriter.notifyMatchFailure(
2792 "Scaled fp packed conversion instructions are not available on target "
2793 "architecture and their emulation is not implemented");
2797 int32_t scaleWaveHalf = op.getFirstScaleLane() / 16;
2798 int32_t firstScaleByte = op.getFirstScaleByte();
2799 int32_t blockSize = op.getBlockSize();
2800 auto sourceType = cast<VectorType>(op.getSource().getType());
2801 auto srcElemType = cast<FloatType>(sourceType.getElementType());
2802 unsigned bitWidth = srcElemType.getWidth();
2804 auto targetType = cast<VectorType>(op.getResult().getType());
2805 auto destElemType = cast<FloatType>(targetType.getElementType());
2807 IntegerType i32 = rewriter.getI32Type();
2808 Value source = adaptor.getSource();
2809 Type llvmResultType = typeConverter->convertType(op.getResult().getType());
2810 Type packedType =
nullptr;
2811 if (isa<fp4>(srcElemType)) {
2813 packedType = getTypeConverter()->convertType(packedType);
2814 }
else if (isa<fp8, bf8>(srcElemType)) {
2815 packedType = VectorType::get(2, i32);
2816 packedType = getTypeConverter()->convertType(packedType);
2817 }
else if (isa<fp6, bf6>(srcElemType)) {
2818 packedType = VectorType::get(3, i32);
2819 packedType = getTypeConverter()->convertType(packedType);
2821 llvm_unreachable(
"invalid element type for packed scaled ext");
2824 if (!packedType || !llvmResultType) {
2825 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2828 std::optional<StringRef> maybeIntrinsic =
2829 scaledExtPacked816ToIntrinsic(srcElemType, destElemType);
2830 if (!maybeIntrinsic.has_value())
2831 return op.emitOpError(
2832 "no intrinsic matching packed scaled conversion on the given chipset");
2835 getScaleSel(blockSize, bitWidth, scaleWaveHalf, firstScaleByte);
2837 LLVM::BitcastOp::create(rewriter, loc, i32, adaptor.getScale());
2838 Value castedSource =
2839 LLVM::BitcastOp::create(rewriter, loc, packedType, source);
2841 OperationState loweredOp(loc, *maybeIntrinsic);
2842 loweredOp.addTypes({llvmResultType});
2843 loweredOp.addOperands({castedSource, castedScale});
2845 SmallVector<NamedAttribute, 1> attrs;
2847 NamedAttribute(
"scaleSel", rewriter.getI32IntegerAttr(scaleSel)));
2849 loweredOp.addAttributes(attrs);
2850 Operation *lowered = rewriter.create(loweredOp);
2851 rewriter.replaceOp(op, lowered);
2856LogicalResult ScaledExtPackedOpLowering::matchAndRewrite(
2857 ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2858 ConversionPatternRewriter &rewriter)
const {
2859 Location loc = op.getLoc();
2861 return rewriter.notifyMatchFailure(
2862 loc,
"Scaled fp conversion instructions are not available on target "
2863 "architecture and their emulation is not implemented");
2864 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2866 Value source = adaptor.getSource();
2867 Value scale = adaptor.getScale();
2869 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2870 Type sourceElemType = sourceVecType.getElementType();
2871 VectorType destVecType = cast<VectorType>(op.getResult().getType());
2872 Type destElemType = destVecType.getElementType();
2874 VectorType packedVecType;
2875 if (isa<Float8E5M2Type, Float8E4M3FNType>(sourceElemType)) {
2876 VectorType v4i8 = VectorType::get(4, rewriter.getI8Type());
2877 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v4i8));
2878 }
else if (isa<Float4E2M1FNType>(sourceElemType)) {
2879 VectorType v8i4 = VectorType::get(8, rewriter.getI4Type());
2880 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v8i4));
2882 llvm_unreachable(
"invalid element type for scaled ext");
2886 if (sourceVecType.getNumElements() < packedVecType.getNumElements()) {
2887 Value longVec = LLVM::ZeroOp::create(rewriter, loc, packedVecType);
2888 if (!sourceVecType) {
2889 longVec = LLVM::InsertElementOp::create(
2892 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2894 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2896 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2901 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2903 if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isF32())
2904 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Bf8Op>(
2905 op, destVecType, i32Source, scale, op.getIndex());
2906 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isF16())
2907 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Bf8Op>(
2908 op, destVecType, i32Source, scale, op.getIndex());
2909 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isBF16())
2910 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Bf8Op>(
2911 op, destVecType, i32Source, scale, op.getIndex());
2912 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isF32())
2913 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp8Op>(
2914 op, destVecType, i32Source, scale, op.getIndex());
2915 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isF16())
2916 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp8Op>(
2917 op, destVecType, i32Source, scale, op.getIndex());
2918 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isBF16())
2919 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp8Op>(
2920 op, destVecType, i32Source, scale, op.getIndex());
2921 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isF32())
2922 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp4Op>(
2923 op, destVecType, i32Source, scale, op.getIndex());
2924 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isF16())
2925 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp4Op>(
2926 op, destVecType, i32Source, scale, op.getIndex());
2927 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isBF16())
2928 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp4Op>(
2929 op, destVecType, i32Source, scale, op.getIndex());
2936LogicalResult PackedScaledTruncOpLowering::matchAndRewrite(
2937 PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2938 ConversionPatternRewriter &rewriter)
const {
2939 Location loc = op.getLoc();
2941 return rewriter.notifyMatchFailure(
2942 loc,
"Scaled fp conversion instructions are not available on target "
2943 "architecture and their emulation is not implemented");
2944 Type v2i16 = getTypeConverter()->convertType(
2945 VectorType::get(2, rewriter.getI16Type()));
2946 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2948 Type resultType = op.getResult().getType();
2950 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2951 Type sourceElemType = sourceVecType.getElementType();
2953 Type intResultType = isa<Float4E2M1FNType>(resultElemType) ? i32 : v2i16;
2955 Value source = adaptor.getSource();
2956 Value scale = adaptor.getScale();
2957 Value existing = adaptor.getExisting();
2959 existing = LLVM::BitcastOp::create(rewriter, loc, intResultType, existing);
2961 existing = LLVM::ZeroOp::create(rewriter, loc, intResultType);
2963 if (sourceVecType.getNumElements() < 2) {
2965 Value elem0 = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2966 VectorType v2 = VectorType::get(2, sourceElemType);
2967 source = LLVM::ZeroOp::create(rewriter, loc, v2);
2968 source = LLVM::InsertElementOp::create(rewriter, loc, source, elem0, c0);
2971 Value sourceA, sourceB;
2972 if (sourceElemType.
isF32()) {
2975 sourceA = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2976 sourceB = LLVM::ExtractElementOp::create(rewriter, loc, source, c1);
2980 if (sourceElemType.
isF32() && isa<Float8E5M2Type>(resultElemType))
2981 result = ROCDL::CvtScaleF32PkBf8F32Op::create(rewriter, loc, intResultType,
2982 existing, sourceA, sourceB,
2983 scale, op.getIndex());
2984 else if (sourceElemType.
isF16() && isa<Float8E5M2Type>(resultElemType))
2985 result = ROCDL::CvtScaleF32PkBf8F16Op::create(
2986 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2987 else if (sourceElemType.
isBF16() && isa<Float8E5M2Type>(resultElemType))
2988 result = ROCDL::CvtScaleF32PkBf8Bf16Op::create(
2989 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2990 else if (sourceElemType.
isF32() && isa<Float8E4M3FNType>(resultElemType))
2991 result = ROCDL::CvtScaleF32PkFp8F32Op::create(rewriter, loc, intResultType,
2992 existing, sourceA, sourceB,
2993 scale, op.getIndex());
2994 else if (sourceElemType.
isF16() && isa<Float8E4M3FNType>(resultElemType))
2995 result = ROCDL::CvtScaleF32PkFp8F16Op::create(
2996 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2997 else if (sourceElemType.
isBF16() && isa<Float8E4M3FNType>(resultElemType))
2998 result = ROCDL::CvtScaleF32PkFp8Bf16Op::create(
2999 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3000 else if (sourceElemType.
isF32() && isa<Float4E2M1FNType>(resultElemType))
3001 result = ROCDL::CvtScaleF32PkFp4F32Op::create(rewriter, loc, intResultType,
3002 existing, sourceA, sourceB,
3003 scale, op.getIndex());
3004 else if (sourceElemType.
isF16() && isa<Float4E2M1FNType>(resultElemType))
3005 result = ROCDL::CvtScaleF32PkFp4F16Op::create(
3006 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3007 else if (sourceElemType.
isBF16() && isa<Float4E2M1FNType>(resultElemType))
3008 result = ROCDL::CvtScaleF32PkFp4Bf16Op::create(
3009 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3013 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3014 op, getTypeConverter()->convertType(resultType),
result);
3018LogicalResult PackedTrunc2xFp8OpLowering::matchAndRewrite(
3019 PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
3020 ConversionPatternRewriter &rewriter)
const {
3021 Location loc = op.getLoc();
3023 return rewriter.notifyMatchFailure(
3024 loc,
"Fp8 conversion instructions are not available on target "
3025 "architecture and their emulation is not implemented");
3026 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3028 Type resultType = op.getResult().getType();
3031 Value sourceA = adaptor.getSourceA();
3032 Value sourceB = adaptor.getSourceB();
3034 sourceB = LLVM::UndefOp::create(rewriter, loc, sourceA.
getType());
3035 Value existing = adaptor.getExisting();
3037 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3039 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3043 result = ROCDL::CvtPkBf8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3044 existing, op.getWordIndex());
3046 result = ROCDL::CvtPkFp8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3047 existing, op.getWordIndex());
3049 return op.emitOpError(
3050 "no truncation to result type available on given chipset");
3052 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3053 op, getTypeConverter()->convertType(resultType),
result);
3057LogicalResult PackedStochRoundFp8OpLowering::matchAndRewrite(
3058 PackedStochRoundFp8Op op, PackedStochRoundFp8OpAdaptor adaptor,
3059 ConversionPatternRewriter &rewriter)
const {
3060 Location loc = op.getLoc();
3062 return rewriter.notifyMatchFailure(
3063 loc,
"Fp8 conversion instructions are not available on target "
3064 "architecture and their emulation is not implemented");
3065 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3067 Type resultType = op.getResult().getType();
3070 Value source = adaptor.getSource();
3071 Value stoch = adaptor.getStochiasticParam();
3072 Value existing = adaptor.getExisting();
3074 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3076 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3080 result = ROCDL::CvtSrBf8F32Op::create(rewriter, loc, i32, source, stoch,
3081 existing, op.getStoreIndex());
3083 result = ROCDL::CvtSrFp8F32Op::create(rewriter, loc, i32, source, stoch,
3084 existing, op.getStoreIndex());
3086 return op.emitOpError(
3087 "no stochastic rounding to result type available on given chipset");
3089 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3090 op, getTypeConverter()->convertType(resultType),
result);
3096struct AMDGPUDPPLowering :
public ConvertOpToLLVMPattern<DPPOp> {
3097 AMDGPUDPPLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3098 : ConvertOpToLLVMPattern<DPPOp>(converter), chipset(chipset) {}
3102 matchAndRewrite(DPPOp DppOp, DPPOp::Adaptor adaptor,
3103 ConversionPatternRewriter &rewriter)
const override {
3106 Location loc = DppOp.getLoc();
3107 Value src = adaptor.getSrc();
3108 Value old = adaptor.getOld();
3111 Type llvmType =
nullptr;
3113 llvmType = rewriter.getI32Type();
3114 }
else if (isa<FloatType>(srcType)) {
3116 ? rewriter.getF32Type()
3117 : rewriter.getF64Type();
3118 }
else if (isa<IntegerType>(srcType)) {
3120 ? rewriter.getI32Type()
3121 : rewriter.getI64Type();
3123 auto llvmSrcIntType = typeConverter->convertType(
3127 auto convertOperand = [&](Value operand, Type operandType) {
3128 if (operandType.getIntOrFloatBitWidth() <= 16) {
3129 if (llvm::isa<FloatType>(operandType)) {
3131 LLVM::BitcastOp::create(rewriter, loc, llvmSrcIntType, operand);
3133 auto llvmVecType = typeConverter->convertType(mlir::VectorType::get(
3134 32 / operandType.getIntOrFloatBitWidth(), llvmSrcIntType));
3135 Value undefVec = LLVM::UndefOp::create(rewriter, loc, llvmVecType);
3137 LLVM::InsertElementOp::create(rewriter, loc, undefVec, operand,
3139 operand = LLVM::BitcastOp::create(rewriter, loc, llvmType, operand);
3144 src = convertOperand(src, srcType);
3145 old = convertOperand(old, oldType);
3148 enum DppCtrl :
unsigned {
3157 ROW_HALF_MIRROR = 0x141,
3162 auto kind = DppOp.getKind();
3163 auto permArgument = DppOp.getPermArgument();
3164 uint32_t DppCtrl = 0;
3168 case DPPPerm::quad_perm: {
3169 auto quadPermAttr = cast<ArrayAttr>(*permArgument);
3171 for (
auto elem : quadPermAttr.getAsRange<IntegerAttr>()) {
3172 uint32_t num = elem.getInt();
3173 DppCtrl |= num << (i * 2);
3178 case DPPPerm::row_shl: {
3179 auto intAttr = cast<IntegerAttr>(*permArgument);
3180 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHL0;
3183 case DPPPerm::row_shr: {
3184 auto intAttr = cast<IntegerAttr>(*permArgument);
3185 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHR0;
3188 case DPPPerm::row_ror: {
3189 auto intAttr = cast<IntegerAttr>(*permArgument);
3190 DppCtrl = intAttr.getInt() + DppCtrl::ROW_ROR0;
3193 case DPPPerm::wave_shl:
3194 DppCtrl = DppCtrl::WAVE_SHL1;
3196 case DPPPerm::wave_shr:
3197 DppCtrl = DppCtrl::WAVE_SHR1;
3199 case DPPPerm::wave_rol:
3200 DppCtrl = DppCtrl::WAVE_ROL1;
3202 case DPPPerm::wave_ror:
3203 DppCtrl = DppCtrl::WAVE_ROR1;
3205 case DPPPerm::row_mirror:
3206 DppCtrl = DppCtrl::ROW_MIRROR;
3208 case DPPPerm::row_half_mirror:
3209 DppCtrl = DppCtrl::ROW_HALF_MIRROR;
3211 case DPPPerm::row_bcast_15:
3212 DppCtrl = DppCtrl::BCAST15;
3214 case DPPPerm::row_bcast_31:
3215 DppCtrl = DppCtrl::BCAST31;
3221 auto rowMask = DppOp->getAttrOfType<IntegerAttr>(
"row_mask").getInt();
3222 auto bankMask = DppOp->getAttrOfType<IntegerAttr>(
"bank_mask").getInt();
3223 bool boundCtrl = DppOp->getAttrOfType<BoolAttr>(
"bound_ctrl").getValue();
3227 ROCDL::DPPUpdateOp::create(rewriter, loc, llvmType, old, src, DppCtrl,
3228 rowMask, bankMask, boundCtrl);
3230 Value
result = dppMovOp.getRes();
3232 result = LLVM::TruncOp::create(rewriter, loc, llvmSrcIntType,
result);
3233 if (!llvm::isa<IntegerType>(srcType)) {
3234 result = LLVM::BitcastOp::create(rewriter, loc, srcType,
result);
3245struct AMDGPUSwizzleBitModeLowering
3246 :
public ConvertOpToLLVMPattern<SwizzleBitModeOp> {
3250 matchAndRewrite(SwizzleBitModeOp op, OpAdaptor adaptor,
3251 ConversionPatternRewriter &rewriter)
const override {
3252 Location loc = op.getLoc();
3253 Type i32 = rewriter.getI32Type();
3254 Value src = adaptor.getSrc();
3255 SmallVector<Value> decomposed;
3257 return rewriter.notifyMatchFailure(op,
3258 "failed to decompose value to i32");
3259 unsigned andMask = op.getAndMask();
3260 unsigned orMask = op.getOrMask();
3261 unsigned xorMask = op.getXorMask();
3265 unsigned mask = andMask | (orMask << 5) | (xorMask << 10);
3267 SmallVector<Value> swizzled;
3268 for (Value v : decomposed) {
3270 ROCDL::DsSwizzleOp::create(rewriter, loc, v.getType(), v, maskValue);
3271 swizzled.emplace_back(res);
3275 rewriter.replaceOp(op,
result);
3280struct AMDGPUPermlaneLowering :
public ConvertOpToLLVMPattern<PermlaneSwapOp> {
3283 AMDGPUPermlaneLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3284 : ConvertOpToLLVMPattern<PermlaneSwapOp>(converter), chipset(chipset) {}
3288 matchAndRewrite(PermlaneSwapOp op, OpAdaptor adaptor,
3289 ConversionPatternRewriter &rewriter)
const override {
3291 return op->emitOpError(
"permlane_swap is only supported on gfx950+");
3293 Location loc = op.getLoc();
3294 Type i32 = rewriter.getI32Type();
3295 Value src = adaptor.getSrc();
3296 unsigned rowLength = op.getRowLength();
3297 bool fi = op.getFetchInactive();
3298 bool boundctrl = op.getBoundCtrl();
3300 SmallVector<Value> decomposed;
3302 return rewriter.notifyMatchFailure(op,
3303 "failed to decompose value to i32");
3305 SmallVector<Value> permuted;
3306 for (Value v : decomposed) {
3308 Type i32pair = LLVM::LLVMStructType::getLiteral(
3309 rewriter.getContext(), {v.getType(), v.getType()});
3311 if (rowLength == 16)
3312 res = ROCDL::Permlane16SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3314 else if (rowLength == 32)
3315 res = ROCDL::Permlane32SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3318 llvm_unreachable(
"unsupported row length");
3320 Value vdst0 = LLVM::ExtractValueOp::create(rewriter, loc, res, {0});
3321 Value vdst1 = LLVM::ExtractValueOp::create(rewriter, loc, res, {1});
3323 Value isEqual = LLVM::ICmpOp::create(rewriter, loc,
3324 LLVM::ICmpPredicate::eq, vdst0, v);
3329 LLVM::SelectOp::create(rewriter, loc, isEqual, vdst1, vdst0);
3330 permuted.emplace_back(vdstNew);
3334 rewriter.replaceOp(op,
result);
3339struct AMDGPUPermlaneVarLowering
3340 :
public ConvertOpToLLVMPattern<PermlaneVarOp> {
3343 AMDGPUPermlaneVarLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3344 : ConvertOpToLLVMPattern<PermlaneVarOp>(converter), chipset(chipset) {}
3348 matchAndRewrite(PermlaneVarOp op, OpAdaptor adaptor,
3349 ConversionPatternRewriter &rewriter)
const override {
3351 return op->emitOpError(
"permlane_var is only supported on GFX12+");
3353 Location loc = op.getLoc();
3354 Type i32 = rewriter.getI32Type();
3355 Value src = adaptor.getSrc();
3356 Value selector = adaptor.getSelector();
3357 bool cross = op.getCross();
3358 bool fi = op.getFetchInactive();
3359 bool boundCtrl = op.getBoundCtrl();
3361 SmallVector<Value> decomposed;
3363 return rewriter.notifyMatchFailure(op,
3364 "failed to decompose value to i32");
3366 SmallVector<Value> permuted;
3367 for (Value v : decomposed) {
3370 res = ROCDL::PermlaneX16VarOp::create(rewriter, loc, i32, v, v,
3371 selector, fi, boundCtrl);
3373 res = ROCDL::Permlane16VarOp::create(rewriter, loc, i32, v, v, selector,
3375 permuted.emplace_back(res);
3379 rewriter.replaceOp(op,
result);
3392constexpr int32_t kDsBarrierPendingCountBitWidth = 29;
3393constexpr int32_t kDsBarrierPhasePos = kDsBarrierPendingCountBitWidth;
3394constexpr int32_t kDsBarrierInitCountPos = 32;
3395constexpr int32_t kDsBarrierPendingCountMask =
3396 (1 << kDsBarrierPendingCountBitWidth) - 1;
3398struct DsBarrierInitOpLowering
3399 :
public ConvertOpToLLVMPattern<DsBarrierInitOp> {
3402 DsBarrierInitOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3403 : ConvertOpToLLVMPattern<DsBarrierInitOp>(converter), chipset(chipset) {}
3406 matchAndRewrite(DsBarrierInitOp op, OpAdaptor adaptor,
3407 ConversionPatternRewriter &rewriter)
const override {
3409 return op->emitOpError(
"only supported on gfx1250+");
3411 Location loc = op.getLoc();
3412 Type i64 = rewriter.getI64Type();
3414 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3416 adaptor.getBase(), adaptor.getIndices());
3423 LLVM::SubOp::create(rewriter, loc, adaptor.getParticipants(),
3430 Value maskedCount32 =
3431 LLVM::AndOp::create(rewriter, loc, initCount, countMask);
3432 Value maskedCount = LLVM::ZExtOp::create(rewriter, loc, i64, maskedCount32);
3434 Value initCountShifted = LLVM::ShlOp::create(
3435 rewriter, loc, maskedCount,
3437 Value barrierState =
3438 LLVM::OrOp::create(rewriter, loc, initCountShifted, maskedCount);
3440 LLVM::StoreOp::create(
3441 rewriter, loc, barrierState, ptr, 8,
false,
3443 false, LLVM::AtomicOrdering::release,
3446 rewriter.eraseOp(op);
3451struct DsBarrierPollStateOpLowering
3452 :
public ConvertOpToLLVMPattern<DsBarrierPollStateOp> {
3455 DsBarrierPollStateOpLowering(
const LLVMTypeConverter &converter,
3457 : ConvertOpToLLVMPattern<DsBarrierPollStateOp>(converter),
3461 matchAndRewrite(DsBarrierPollStateOp op, OpAdaptor adaptor,
3462 ConversionPatternRewriter &rewriter)
const override {
3464 return op->emitOpError(
"only supported on gfx1250+");
3466 Location loc = op.getLoc();
3467 Type i64 = rewriter.getI64Type();
3469 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3471 adaptor.getBase(), adaptor.getIndices());
3475 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
3476 op, i64, ptr, 8,
false,
3478 false, LLVM::AtomicOrdering::acquire,
3484struct DsAsyncBarrierArriveOpLowering
3485 :
public ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp> {
3488 DsAsyncBarrierArriveOpLowering(
const LLVMTypeConverter &converter,
3490 : ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp>(converter),
3494 matchAndRewrite(DsAsyncBarrierArriveOp op, OpAdaptor adaptor,
3495 ConversionPatternRewriter &rewriter)
const override {
3497 return op->emitOpError(
"only supported on gfx1250+");
3499 Location loc = op.getLoc();
3501 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3503 adaptor.getBase(), adaptor.getIndices());
3505 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicAsyncBarrierArriveOp>(
3506 op, ptr,
nullptr,
nullptr,
3512struct DsBarrierArriveOpLowering
3513 :
public ConvertOpToLLVMPattern<DsBarrierArriveOp> {
3516 DsBarrierArriveOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3517 : ConvertOpToLLVMPattern<DsBarrierArriveOp>(converter), chipset(chipset) {
3521 matchAndRewrite(DsBarrierArriveOp op, OpAdaptor adaptor,
3522 ConversionPatternRewriter &rewriter)
const override {
3524 return op->emitOpError(
"only supported on gfx1250+");
3526 Location loc = op.getLoc();
3527 Type i64 = rewriter.getI64Type();
3529 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3531 adaptor.getBase(), adaptor.getIndices());
3533 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicBarrierArriveRtnOp>(
3534 op, i64, ptr, adaptor.getCount(),
nullptr,
3540struct DsBarrierStatePhaseOpLowering
3541 :
public ConvertOpToLLVMPattern<DsBarrierStatePhaseOp> {
3545 matchAndRewrite(DsBarrierStatePhaseOp op, OpAdaptor adaptor,
3546 ConversionPatternRewriter &rewriter)
const override {
3547 Location loc = op.getLoc();
3548 Type i32 = rewriter.getI32Type();
3550 Value state = adaptor.getState();
3552 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3553 Value phase = LLVM::LShrOp::create(
3554 rewriter, loc, noInitCount,
3557 rewriter.replaceOp(op, phase);
3562struct DsBarrierStatePendingCountOpLowering
3563 :
public ConvertOpToLLVMPattern<DsBarrierStatePendingCountOp> {
3567 matchAndRewrite(DsBarrierStatePendingCountOp op, OpAdaptor adaptor,
3568 ConversionPatternRewriter &rewriter)
const override {
3569 Location loc = op.getLoc();
3570 Type i32 = rewriter.getI32Type();
3572 Value state = adaptor.getState();
3574 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3575 Value pendingCount = LLVM::AndOp::create(
3576 rewriter, loc, noInitCount,
3578 static_cast<uint32_t
>(kDsBarrierPendingCountMask)));
3580 rewriter.replaceOp(op, pendingCount);
3585struct DsBarrierStateInitCountOpLowering
3586 :
public ConvertOpToLLVMPattern<DsBarrierStateInitCountOp> {
3590 matchAndRewrite(DsBarrierStateInitCountOp op, OpAdaptor adaptor,
3591 ConversionPatternRewriter &rewriter)
const override {
3592 Location loc = op.getLoc();
3593 Type i32 = rewriter.getI32Type();
3595 Value state = adaptor.getState();
3597 Value initCountI64 = LLVM::LShrOp::create(
3598 rewriter, loc, state,
3600 Value initCount = LLVM::TruncOp::create(rewriter, loc, i32, initCountI64);
3602 rewriter.replaceOp(op, initCount);
3607struct DsBarrierStatePhaseParityLowering
3608 :
public ConvertOpToLLVMPattern<DsBarrierStatePhaseParity> {
3612 matchAndRewrite(DsBarrierStatePhaseParity op, OpAdaptor adaptor,
3613 ConversionPatternRewriter &rewriter)
const override {
3614 Location loc = op.getLoc();
3615 Type i1 = rewriter.getI1Type();
3617 Value state = adaptor.getState();
3620 LLVM::TruncOp::create(rewriter, loc, rewriter.getI32Type(), state);
3621 Value phase = LLVM::LShrOp::create(
3622 rewriter, loc, noInitCount,
3624 Value parity = LLVM::TruncOp::create(rewriter, loc, i1, phase);
3626 rewriter.replaceOp(op, parity);
3635static Value setValueAtOffset(ConversionPatternRewriter &rewriter, Location loc,
3636 Value accumulator, Value value, int64_t shift) {
3641 value = LLVM::ShlOp::create(rewriter, loc, value, shiftAmount);
3647 constexpr bool isDisjoint =
true;
3648 return LLVM::OrOp::create(rewriter, loc, accumulator, value, isDisjoint);
3651template <
typename BaseOp>
3652struct AMDGPUMakeDmaBaseLowering :
public ConvertOpToLLVMPattern<BaseOp> {
3653 using ConvertOpToLLVMPattern<BaseOp>::ConvertOpToLLVMPattern;
3656 AMDGPUMakeDmaBaseLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3657 : ConvertOpToLLVMPattern<BaseOp>(converter), chipset(chipset) {}
3661 matchAndRewrite(BaseOp op, Adaptor adaptor,
3662 ConversionPatternRewriter &rewriter)
const override {
3664 return op->emitOpError(
"make_dma_base is only supported on gfx1250");
3666 Location loc = op.getLoc();
3668 constexpr int32_t constlen = 4;
3669 Value consts[constlen];
3670 for (int64_t i = 0; i < constlen; ++i)
3673 constexpr int32_t sgprslen = constlen;
3674 Value sgprs[sgprslen];
3675 for (int64_t i = 0; i < sgprslen; ++i) {
3676 sgprs[i] = consts[0];
3679 sgprs[0] = consts[1];
3681 if constexpr (BaseOp::isGather()) {
3682 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 30);
3684 auto type = cast<TDMGatherBaseType>(op.getResult().getType());
3685 Type indexType = type.getIndexType();
3687 assert(llvm::is_contained({16u, 32u}, indexSize) &&
3688 "expected index_size to be 16 or 32");
3689 unsigned idx = (indexSize / 16) - 1;
3692 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 31);
3695 ValueRange ldsIndices = adaptor.getLdsIndices();
3696 Value lds = adaptor.getLds();
3697 auto ldsMemRefType = cast<MemRefType>(op.getLds().getType());
3700 rewriter, loc, ldsMemRefType, lds, ldsIndices);
3702 ValueRange globalIndices = adaptor.getGlobalIndices();
3703 Value global = adaptor.getGlobal();
3704 auto globalMemRefType = cast<MemRefType>(op.getGlobal().getType());
3707 rewriter, loc, globalMemRefType, global, globalIndices);
3709 Type i32 = rewriter.getI32Type();
3710 Type i64 = rewriter.getI64Type();
3712 sgprs[1] = LLVM::PtrToIntOp::create(rewriter, loc, i32, ldsPtr);
3713 Value castForGlobalAddr =
3714 LLVM::PtrToIntOp::create(rewriter, loc, i64, globalPtr);
3716 sgprs[2] = LLVM::TruncOp::create(rewriter, loc, i32, castForGlobalAddr);
3718 Value shift = LLVM::LShrOp::create(rewriter, loc, castForGlobalAddr,
3721 Value highHalf = LLVM::TruncOp::create(rewriter, loc, i32, shift);
3724 highHalf = LLVM::AndOp::create(rewriter, loc, highHalf, mask);
3726 sgprs[3] = setValueAtOffset(rewriter, loc, highHalf, consts[2], 30);
3728 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
3729 assert(v4i32 &&
"expected type conversion to succeed");
3730 Value
result = LLVM::PoisonOp::create(rewriter, loc, v4i32);
3732 for (
auto [sgpr, constant] : llvm::zip_equal(sgprs, consts))
3734 LLVM::InsertElementOp::create(rewriter, loc,
result, sgpr, constant);
3736 rewriter.replaceOp(op,
result);
3741template <
typename DescriptorOp>
3742struct AMDGPULowerDescriptor :
public ConvertOpToLLVMPattern<DescriptorOp> {
3743 using ConvertOpToLLVMPattern<DescriptorOp>::ConvertOpToLLVMPattern;
3746 AMDGPULowerDescriptor(
const LLVMTypeConverter &converter, Chipset chipset)
3747 : ConvertOpToLLVMPattern<DescriptorOp>(converter), chipset(chipset) {}
3750 Value getDGroup0(OpAdaptor adaptor)
const {
return adaptor.getBase(); }
3752 Value setWorkgroupMask(DescriptorOp op, OpAdaptor adaptor,
3753 ConversionPatternRewriter &rewriter, Location loc,
3754 Value sgpr0)
const {
3755 Value mask = op.getWorkgroupMask();
3759 Type i16 = rewriter.getI16Type();
3760 mask = LLVM::BitcastOp::create(rewriter, loc, i16, mask);
3761 Type i32 = rewriter.getI32Type();
3762 Value extendedMask = LLVM::ZExtOp::create(rewriter, loc, i32, mask);
3763 return setValueAtOffset(rewriter, loc, sgpr0, extendedMask, 0);
3766 Value setDataSize(DescriptorOp op, OpAdaptor adaptor,
3767 ConversionPatternRewriter &rewriter, Location loc,
3768 Value sgpr0, ArrayRef<Value> consts)
const {
3769 unsigned elementTypeWidthInBits = op.getElementTypeWidth();
3770 assert(llvm::is_contained({8u, 16u, 32u, 64u}, elementTypeWidthInBits) &&
3771 "expected type width to be 8, 16, 32, or 64.");
3772 int64_t idx = llvm::Log2_32(elementTypeWidthInBits / 8);
3773 Value size = consts[idx];
3774 return setValueAtOffset(rewriter, loc, sgpr0, size, 16);
3777 Value setAtomicBarrier(DescriptorOp op, OpAdaptor adaptor,
3778 ConversionPatternRewriter &rewriter, Location loc,
3779 Value sgpr0, ArrayRef<Value> consts)
const {
3780 if (!adaptor.getAtomicBarrierAddress())
3783 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 18);
3786 Value setIterateEnable(DescriptorOp op, OpAdaptor adaptor,
3787 ConversionPatternRewriter &rewriter, Location loc,
3788 Value sgpr0, ArrayRef<Value> consts)
const {
3789 if (!adaptor.getGlobalIncrement())
3794 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 19);
3797 Value setPadEnable(DescriptorOp op, OpAdaptor adaptor,
3798 ConversionPatternRewriter &rewriter, Location loc,
3799 Value sgpr0, ArrayRef<Value> consts)
const {
3800 if (!op.getPadAmount())
3803 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 20);
3806 Value setEarlyTimeout(DescriptorOp op, OpAdaptor adaptor,
3807 ConversionPatternRewriter &rewriter, Location loc,
3808 Value sgpr0, ArrayRef<Value> consts)
const {
3809 if (!op.getWorkgroupMask())
3812 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 21);
3815 Value setPadInterval(DescriptorOp op, OpAdaptor adaptor,
3816 ConversionPatternRewriter &rewriter, Location loc,
3817 Value sgpr0, ArrayRef<Value> consts)
const {
3818 if (!op.getPadAmount())
3827 IntegerType i32 = rewriter.getI32Type();
3828 Value padInterval = adaptor.getPadInterval();
3829 padInterval = LLVM::CountTrailingZerosOp::create(rewriter, loc, i32,
3830 padInterval,
false);
3831 padInterval = LLVM::SubOp::create(rewriter, loc, padInterval, consts[1]);
3833 return setValueAtOffset(rewriter, loc, sgpr0, padInterval, 22);
3836 Value setPadAmount(DescriptorOp op, OpAdaptor adaptor,
3837 ConversionPatternRewriter &rewriter, Location loc,
3838 Value sgpr0, ArrayRef<Value> consts)
const {
3839 if (!op.getPadAmount())
3848 Value padAmount = adaptor.getPadAmount();
3849 padAmount = LLVM::SubOp::create(rewriter, loc, padAmount, consts[1]);
3851 return setValueAtOffset(rewriter, loc, sgpr0, padAmount, 25);
3854 Value setAtomicBarrierAddress(DescriptorOp op, OpAdaptor adaptor,
3855 ConversionPatternRewriter &rewriter,
3856 Location loc, Value sgpr1,
3857 ArrayRef<Value> consts)
const {
3858 if (!adaptor.getAtomicBarrierAddress())
3861 Value atomicBarrierAddress = adaptor.getAtomicBarrierAddress();
3862 auto barrierAddressTy =
3863 cast<MemRefType>(op.getAtomicBarrierAddress().getType());
3864 ValueRange atomicBarrierIndices = adaptor.getAtomicBarrierIndices();
3866 rewriter, loc, barrierAddressTy, atomicBarrierAddress,
3867 atomicBarrierIndices);
3868 IntegerType i32 = rewriter.getI32Type();
3874 atomicBarrierAddress =
3875 LLVM::PtrToIntOp::create(rewriter, loc, i32, atomicBarrierAddress);
3876 atomicBarrierAddress =
3877 LLVM::LShrOp::create(rewriter, loc, atomicBarrierAddress, consts[3]);
3879 atomicBarrierAddress =
3880 LLVM::AndOp::create(rewriter, loc, atomicBarrierAddress, mask);
3881 return setValueAtOffset(rewriter, loc, sgpr1, atomicBarrierAddress, 32);
3884 std::pair<Value, Value> setTensorDimX(DescriptorOp op, OpAdaptor adaptor,
3885 ConversionPatternRewriter &rewriter,
3886 Location loc, Value sgpr1, Value sgpr2,
3887 ArrayRef<Value> consts, uint64_t dimX,
3888 uint32_t offset)
const {
3889 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
3890 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
3891 SmallVector<OpFoldResult> mixedGlobalSizes =
3893 if (mixedGlobalSizes.size() <= dimX)
3894 return {sgpr1, sgpr2};
3896 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
3903 if (
auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
3907 IntegerType i32 = rewriter.getI32Type();
3908 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
3909 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
3912 sgpr1 = setValueAtOffset(rewriter, loc, sgpr1, tensorDimX, offset);
3915 Value tensorDimXHigh = LLVM::LShrOp::create(rewriter, loc, tensorDimX, c16);
3916 sgpr2 = setValueAtOffset(rewriter, loc, sgpr2, tensorDimXHigh, offset + 16);
3917 return {sgpr1, sgpr2};
3920 std::pair<Value, Value> setTensorDim0(DescriptorOp op, OpAdaptor adaptor,
3921 ConversionPatternRewriter &rewriter,
3922 Location loc, Value sgpr1, Value sgpr2,
3923 ArrayRef<Value> consts)
const {
3924 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, 0,
3928 std::pair<Value, Value> setTensorDim1(DescriptorOp op, OpAdaptor adaptor,
3929 ConversionPatternRewriter &rewriter,
3930 Location loc, Value sgpr2, Value sgpr3,
3931 ArrayRef<Value> consts)
const {
3932 return setTensorDimX(op, adaptor, rewriter, loc, sgpr2, sgpr3, consts, 1,
3936 Value setTileDimX(DescriptorOp op, OpAdaptor adaptor,
3937 ConversionPatternRewriter &rewriter, Location loc,
3938 Value sgpr, ArrayRef<Value> consts,
size_t dimX,
3939 int64_t offset)
const {
3940 ArrayRef<int64_t> sharedStaticSizes = adaptor.getSharedStaticSizes();
3941 ValueRange sharedDynamicSizes = adaptor.getSharedDynamicSizes();
3942 SmallVector<OpFoldResult> mixedSharedSizes =
3944 if (mixedSharedSizes.size() <= dimX)
3947 OpFoldResult tileDimXOpFoldResult = *(mixedSharedSizes.rbegin() + dimX);
3956 if (
auto attr = dyn_cast<Attribute>(tileDimXOpFoldResult)) {
3960 IntegerType i32 = rewriter.getI32Type();
3961 tileDimX = cast<Value>(tileDimXOpFoldResult);
3962 tileDimX = LLVM::TruncOp::create(rewriter, loc, i32, tileDimX);
3965 return setValueAtOffset(rewriter, loc, sgpr, tileDimX, offset);
3968 Value setTileDim0(DescriptorOp op, OpAdaptor adaptor,
3969 ConversionPatternRewriter &rewriter, Location loc,
3970 Value sgpr3, ArrayRef<Value> consts)
const {
3971 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, 0, 112);
3974 Value setTileDim1(DescriptorOp op, OpAdaptor adaptor,
3975 ConversionPatternRewriter &rewriter, Location loc,
3976 Value sgpr4, ArrayRef<Value> consts)
const {
3977 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 1, 128);
3980 Value setValidIndices(DescriptorOp op, OpAdaptor adaptor,
3981 ConversionPatternRewriter &rewriter, Location loc,
3982 Value sgpr4, ArrayRef<Value> consts)
const {
3983 auto type = cast<VectorType>(op.getIndices().getType());
3984 ArrayRef<int64_t> shape = type.getShape();
3985 assert(shape.size() == 1 &&
"expected shape to be of rank 1.");
3986 unsigned length = shape.back();
3987 assert(0 < length && length <= 16 &&
"expected length to be at most 16.");
3989 return setValueAtOffset(rewriter, loc, sgpr4, value, 128);
3992 Value setTileDim1OrValidIndices(DescriptorOp op, OpAdaptor adaptor,
3993 ConversionPatternRewriter &rewriter,
3994 Location loc, Value sgpr4,
3995 ArrayRef<Value> consts)
const {
3996 if constexpr (DescriptorOp::isGather())
3997 return setValidIndices(op, adaptor, rewriter, loc, sgpr4, consts);
3998 return setTileDim1(op, adaptor, rewriter, loc, sgpr4, consts);
4001 Value setTileDim2(DescriptorOp op, OpAdaptor adaptor,
4002 ConversionPatternRewriter &rewriter, Location loc,
4003 Value sgpr4, ArrayRef<Value> consts)
const {
4005 if constexpr (DescriptorOp::isGather())
4007 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 2, 144);
4010 std::pair<Value, Value>
4011 setTensorDimXStride(DescriptorOp op, OpAdaptor adaptor,
4012 ConversionPatternRewriter &rewriter, Location loc,
4013 Value sgprY, Value sgprZ, ArrayRef<Value> consts,
4014 size_t dimX, int64_t offset)
const {
4015 ArrayRef<int64_t> globalStaticStrides = adaptor.getGlobalStaticStrides();
4016 ValueRange globalDynamicStrides = adaptor.getGlobalDynamicStrides();
4017 SmallVector<OpFoldResult> mixedGlobalStrides =
4018 getMixedValues(globalStaticStrides, globalDynamicStrides, rewriter);
4020 if (mixedGlobalStrides.size() <= (dimX + 1))
4021 return {sgprY, sgprZ};
4023 OpFoldResult tensorDimXStrideOpFoldResult =
4024 *(mixedGlobalStrides.rbegin() + dimX + 1);
4029 Value tensorDimXStride;
4030 if (
auto attr = dyn_cast<Attribute>(tensorDimXStrideOpFoldResult))
4034 tensorDimXStride = cast<Value>(tensorDimXStrideOpFoldResult);
4036 constexpr int64_t first48bits = (1ll << 48) - 1;
4039 LLVM::AndOp::create(rewriter, loc, mask, tensorDimXStride);
4040 IntegerType i32 = rewriter.getI32Type();
4041 Value tensorDimXStrideLow =
4042 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStride);
4043 sgprY = setValueAtOffset(rewriter, loc, sgprY, tensorDimXStrideLow, offset);
4045 int64_t shift = (offset % 32) == 0 ? 32 : offset % 32;
4047 Value tensorDimXStrideHigh =
4048 LLVM::LShrOp::create(rewriter, loc, tensorDimXStride, shiftVal);
4049 tensorDimXStrideHigh =
4050 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStrideHigh);
4051 sgprZ = setValueAtOffset(rewriter, loc, sgprZ, tensorDimXStrideHigh,
4053 return {sgprY, sgprZ};
4056 std::pair<Value, Value>
4057 setTensorDim0Stride(DescriptorOp op, OpAdaptor adaptor,
4058 ConversionPatternRewriter &rewriter, Location loc,
4059 Value sgpr5, Value sgpr6, ArrayRef<Value> consts)
const {
4060 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4064 std::pair<Value, Value>
4065 setTensorDim1Stride(DescriptorOp op, OpAdaptor adaptor,
4066 ConversionPatternRewriter &rewriter, Location loc,
4067 Value sgpr5, Value sgpr6, ArrayRef<Value> consts)
const {
4069 if constexpr (DescriptorOp::isGather())
4070 return {sgpr5, sgpr6};
4071 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4075 Value getDGroup1(DescriptorOp op, OpAdaptor adaptor,
4076 ConversionPatternRewriter &rewriter, Location loc,
4077 ArrayRef<Value> consts)
const {
4079 for (int64_t i = 0; i < 8; ++i) {
4080 sgprs[i] = consts[0];
4083 sgprs[0] = setWorkgroupMask(op, adaptor, rewriter, loc, sgprs[0]);
4084 sgprs[0] = setDataSize(op, adaptor, rewriter, loc, sgprs[0], consts);
4085 sgprs[0] = setAtomicBarrier(op, adaptor, rewriter, loc, sgprs[0], consts);
4086 sgprs[0] = setIterateEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4087 sgprs[0] = setPadEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4088 sgprs[0] = setEarlyTimeout(op, adaptor, rewriter, loc, sgprs[0], consts);
4089 sgprs[0] = setPadInterval(op, adaptor, rewriter, loc, sgprs[0], consts);
4090 sgprs[0] = setPadAmount(op, adaptor, rewriter, loc, sgprs[0], consts);
4093 setAtomicBarrierAddress(op, adaptor, rewriter, loc, sgprs[1], consts);
4094 std::tie(sgprs[1], sgprs[2]) =
4095 setTensorDim0(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4096 std::tie(sgprs[2], sgprs[3]) =
4097 setTensorDim1(op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4099 sgprs[3] = setTileDim0(op, adaptor, rewriter, loc, sgprs[3], consts);
4101 setTileDim1OrValidIndices(op, adaptor, rewriter, loc, sgprs[4], consts);
4102 sgprs[4] = setTileDim2(op, adaptor, rewriter, loc, sgprs[4], consts);
4103 std::tie(sgprs[5], sgprs[6]) = setTensorDim0Stride(
4104 op, adaptor, rewriter, loc, sgprs[5], sgprs[6], consts);
4105 std::tie(sgprs[6], sgprs[7]) = setTensorDim1Stride(
4106 op, adaptor, rewriter, loc, sgprs[6], sgprs[7], consts);
4108 IntegerType i32 = rewriter.getI32Type();
4109 Type v8i32 = this->typeConverter->convertType(VectorType::get(8, i32));
4110 assert(v8i32 &&
"expected type conversion to succeed");
4111 Value dgroup1 = LLVM::PoisonOp::create(rewriter, loc, v8i32);
4113 for (
auto [sgpr, constant] : llvm::zip_equal(sgprs, consts)) {
4115 LLVM::InsertElementOp::create(rewriter, loc, dgroup1, sgpr, constant);
4121 Value setTensorDimX(DescriptorOp op, OpAdaptor adaptor,
4122 ConversionPatternRewriter &rewriter, Location loc,
4123 Value sgpr0, ArrayRef<Value> consts, int64_t dimX,
4124 int64_t offset)
const {
4125 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
4126 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
4127 SmallVector<OpFoldResult> mixedGlobalSizes =
4129 if (mixedGlobalSizes.size() <=
static_cast<unsigned long>(dimX))
4132 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
4134 if (
auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
4138 IntegerType i32 = rewriter.getI32Type();
4139 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
4140 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
4143 return setValueAtOffset(rewriter, loc, sgpr0, tensorDimX, offset);
4146 Value setTensorDim2(DescriptorOp op, OpAdaptor adaptor,
4147 ConversionPatternRewriter &rewriter, Location loc,
4148 Value sgpr0, ArrayRef<Value> consts)
const {
4149 return setTensorDimX(op, adaptor, rewriter, loc, sgpr0, consts, 2, 0);
4152 Value truncateAndSetValueAtOffset(ConversionPatternRewriter &rewriter,
4153 Location loc, Value accumulator,
4154 Value value, int64_t shift)
const {
4156 IntegerType i32 = rewriter.getI32Type();
4157 value = LLVM::TruncOp::create(rewriter, loc, i32, value);
4158 return setValueAtOffset(rewriter, loc, accumulator, value, shift);
4161 Value setLDSAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4162 ConversionPatternRewriter &rewriter, Location loc,
4163 Value sgpr1, ArrayRef<Value> consts,
4164 int64_t offset)
const {
4165 Value ldsAddrIncrement = adaptor.getLdsIncrement();
4166 return setValueAtOffset(rewriter, loc, sgpr1, ldsAddrIncrement, offset);
4169 std::pair<Value, Value>
4170 setGlobalAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4171 ConversionPatternRewriter &rewriter, Location loc,
4172 Value sgpr2, Value sgpr3, ArrayRef<Value> consts,
4173 int64_t offset)
const {
4174 Value globalAddrIncrement = adaptor.getGlobalIncrement();
4175 sgpr2 = truncateAndSetValueAtOffset(rewriter, loc, sgpr2,
4176 globalAddrIncrement, offset);
4178 globalAddrIncrement =
4179 LLVM::LShrOp::create(rewriter, loc, globalAddrIncrement, shift);
4180 constexpr int64_t first16BitsHigh = (1ll << 16) - 1;
4181 sgpr3 = truncateAndSetValueAtOffset(rewriter, loc, sgpr3,
4182 globalAddrIncrement, offset + 32);
4184 sgpr3 = LLVM::AndOp::create(rewriter, loc, sgpr3, mask);
4185 return {sgpr2, sgpr3};
4188 Value setTensorDim3OrLDSAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4189 ConversionPatternRewriter &rewriter,
4190 Location loc, Value sgpr1,
4191 ArrayRef<Value> consts)
const {
4192 Value ldsIncrement = op.getLdsIncrement();
4193 constexpr int64_t dim = 3;
4194 constexpr int64_t offset = 32;
4196 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, consts, dim,
4198 return setLDSAddrIncrement(op, adaptor, rewriter, loc, sgpr1, consts,
4202 std::pair<Value, Value> setTensorDim2StrideOrGlobalAddrIncrement(
4203 DescriptorOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter,
4204 Location loc, Value sgpr2, Value sgpr3, ArrayRef<Value> consts)
const {
4205 Value globalIncrement = op.getGlobalIncrement();
4206 constexpr int32_t dim = 2;
4207 constexpr int32_t offset = 64;
4208 if (!globalIncrement)
4209 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4210 consts, dim, offset);
4211 return setGlobalAddrIncrement(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4215 Value setIterateCount(DescriptorOp op, OpAdaptor adaptor,
4216 ConversionPatternRewriter &rewriter, Location loc,
4217 Value sgpr3, ArrayRef<Value> consts,
4218 int32_t offset)
const {
4219 Value iterationCount = adaptor.getIterationCount();
4220 IntegerType i32 = rewriter.getI32Type();
4227 iterationCount = LLVM::TruncOp::create(rewriter, loc, i32, iterationCount);
4229 LLVM::SubOp::create(rewriter, loc, iterationCount, consts[1]);
4230 return setValueAtOffset(rewriter, loc, sgpr3, iterationCount, offset);
4233 Value setTileDim3OrIterateCount(DescriptorOp op, OpAdaptor adaptor,
4234 ConversionPatternRewriter &rewriter,
4235 Location loc, Value sgpr3,
4236 ArrayRef<Value> consts)
const {
4237 Value iterateCount = op.getIterationCount();
4238 constexpr int32_t dim = 2;
4239 constexpr int32_t offset = 112;
4241 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, dim,
4244 return setIterateCount(op, adaptor, rewriter, loc, sgpr3, consts, offset);
4247 Value getDGroup2(DescriptorOp op, OpAdaptor adaptor,
4248 ConversionPatternRewriter &rewriter, Location loc,
4249 ArrayRef<Value> consts)
const {
4250 if constexpr (DescriptorOp::isGather())
4251 return getDGroup2Gather(op, adaptor, rewriter, loc, consts);
4252 return getDGroup2NonGather(op, adaptor, rewriter, loc, consts);
4255 Value getDGroup2NonGather(DescriptorOp op, OpAdaptor adaptor,
4256 ConversionPatternRewriter &rewriter, Location loc,
4257 ArrayRef<Value> consts)
const {
4258 IntegerType i32 = rewriter.getI32Type();
4259 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4260 assert(v4i32 &&
"expected type conversion to succeed.");
4262 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4263 if (onlyNeedsTwoDescriptors)
4264 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4266 constexpr int64_t sgprlen = 4;
4267 Value sgprs[sgprlen];
4268 for (
int i = 0; i < sgprlen; ++i)
4269 sgprs[i] = consts[0];
4271 sgprs[0] = setTensorDim2(op, adaptor, rewriter, loc, sgprs[0], consts);
4272 sgprs[1] = setTensorDim3OrLDSAddrIncrement(op, adaptor, rewriter, loc,
4274 std::tie(sgprs[2], sgprs[3]) = setTensorDim2StrideOrGlobalAddrIncrement(
4275 op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4277 setTileDim3OrIterateCount(op, adaptor, rewriter, loc, sgprs[3], consts);
4279 Value dgroup2 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4280 for (
auto [sgpr, constant] : llvm::zip(sgprs, consts))
4282 LLVM::InsertElementOp::create(rewriter, loc, dgroup2, sgpr, constant);
4287 Value getGatherIndices(DescriptorOp op, OpAdaptor adaptor,
4288 ConversionPatternRewriter &rewriter, Location loc,
4289 ArrayRef<Value> consts,
bool firstHalf)
const {
4290 IntegerType i32 = rewriter.getI32Type();
4291 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4292 assert(v4i32 &&
"expected type conversion to succeed.");
4294 Value
indices = adaptor.getIndices();
4295 auto vectorType = cast<VectorType>(
indices.getType());
4296 unsigned length = vectorType.getShape().back();
4297 Type elementType = vectorType.getElementType();
4298 unsigned maxLength = elementType == i32 ? 4 : 8;
4299 int32_t offset = firstHalf ? 0 : maxLength;
4300 unsigned discountedLength =
4301 std::max(
static_cast<int32_t
>(length - offset), 0);
4303 unsigned targetSize = std::min(maxLength, discountedLength);
4305 SmallVector<Value> indicesVector;
4306 for (
unsigned i = offset; i < targetSize + offset; ++i) {
4308 if (i < consts.size())
4312 Value elem = LLVM::ExtractElementOp::create(rewriter, loc,
indices, idx);
4313 indicesVector.push_back(elem);
4316 SmallVector<Value> indicesI32Vector;
4317 if (elementType == i32) {
4318 indicesI32Vector = indicesVector;
4320 for (
unsigned i = 0; i < targetSize; ++i) {
4321 Value index = indicesVector[i];
4322 indicesI32Vector.push_back(
4323 LLVM::ZExtOp::create(rewriter, loc, i32, index));
4325 if ((targetSize % 2) != 0)
4327 indicesI32Vector.push_back(consts[0]);
4330 SmallVector<Value> indicesToInsert;
4331 if (elementType == i32) {
4332 indicesToInsert = indicesI32Vector;
4334 unsigned size = indicesI32Vector.size() / 2;
4335 for (
unsigned i = 0; i < size; ++i) {
4336 Value first = indicesI32Vector[2 * i];
4337 Value second = indicesI32Vector[2 * i + 1];
4338 Value joined = setValueAtOffset(rewriter, loc, first, second, 16);
4339 indicesToInsert.push_back(joined);
4343 Value dgroup = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4344 for (
auto [sgpr, constant] : llvm::zip_first(indicesToInsert, consts))
4346 LLVM::InsertElementOp::create(rewriter, loc, dgroup, sgpr, constant);
4351 Value getDGroup2Gather(DescriptorOp op, OpAdaptor adaptor,
4352 ConversionPatternRewriter &rewriter, Location loc,
4353 ArrayRef<Value> consts)
const {
4354 return getGatherIndices(op, adaptor, rewriter, loc, consts,
true);
4357 std::pair<Value, Value>
4358 setTensorDim3Stride(DescriptorOp op, OpAdaptor adaptor,
4359 ConversionPatternRewriter &rewriter, Location loc,
4360 Value sgpr0, Value sgpr1, ArrayRef<Value> consts)
const {
4361 constexpr int32_t dim = 3;
4362 constexpr int32_t offset = 0;
4363 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr0, sgpr1, consts,
4367 std::pair<Value, Value> setTensorDim4(DescriptorOp op, OpAdaptor adaptor,
4368 ConversionPatternRewriter &rewriter,
4369 Location loc, Value sgpr1, Value sgpr2,
4370 ArrayRef<Value> consts)
const {
4371 constexpr int32_t dim = 4;
4372 constexpr int32_t offset = 48;
4373 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, dim,
4377 Value setTileDim4(DescriptorOp op, OpAdaptor adaptor,
4378 ConversionPatternRewriter &rewriter, Location loc,
4379 Value sgpr2, ArrayRef<Value> consts)
const {
4380 constexpr int32_t dim = 4;
4381 constexpr int32_t offset = 80;
4382 return setTileDimX(op, adaptor, rewriter, loc, sgpr2, consts, dim, offset);
4385 Value getDGroup3(DescriptorOp op, OpAdaptor adaptor,
4386 ConversionPatternRewriter &rewriter, Location loc,
4387 ArrayRef<Value> consts)
const {
4388 if constexpr (DescriptorOp::isGather())
4389 return getDGroup3Gather(op, adaptor, rewriter, loc, consts);
4390 return getDGroup3NonGather(op, adaptor, rewriter, loc, consts);
4393 Value getDGroup3NonGather(DescriptorOp op, OpAdaptor adaptor,
4394 ConversionPatternRewriter &rewriter, Location loc,
4395 ArrayRef<Value> consts)
const {
4396 IntegerType i32 = rewriter.getI32Type();
4397 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4398 assert(v4i32 &&
"expected type conversion to succeed.");
4399 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4400 if (onlyNeedsTwoDescriptors)
4401 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4403 constexpr int32_t sgprlen = 4;
4404 Value sgprs[sgprlen];
4405 for (
int i = 0; i < sgprlen; ++i)
4406 sgprs[i] = consts[0];
4408 std::tie(sgprs[0], sgprs[1]) = setTensorDim3Stride(
4409 op, adaptor, rewriter, loc, sgprs[0], sgprs[1], consts);
4410 std::tie(sgprs[1], sgprs[2]) =
4411 setTensorDim4(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4412 sgprs[2] = setTileDim4(op, adaptor, rewriter, loc, sgprs[2], consts);
4414 Value dgroup3 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4415 for (
auto [sgpr, constant] : llvm::zip(sgprs, consts))
4417 LLVM::InsertElementOp::create(rewriter, loc, dgroup3, sgpr, constant);
4422 Value getDGroup3Gather(DescriptorOp op, OpAdaptor adaptor,
4423 ConversionPatternRewriter &rewriter, Location loc,
4424 ArrayRef<Value> consts)
const {
4425 return getGatherIndices(op, adaptor, rewriter, loc, consts,
false);
4429 matchAndRewrite(DescriptorOp op, OpAdaptor adaptor,
4430 ConversionPatternRewriter &rewriter)
const override {
4432 return op->emitOpError(
4433 "make_dma_descriptor is only supported on gfx1250");
4435 Location loc = op.getLoc();
4437 SmallVector<Value> consts;
4438 for (int64_t i = 0; i < 8; ++i)
4441 Value dgroup0 = this->getDGroup0(adaptor);
4442 Value dgroup1 = this->getDGroup1(op, adaptor, rewriter, loc, consts);
4443 Value dgroup2 = this->getDGroup2(op, adaptor, rewriter, loc, consts);
4444 Value dgroup3 = this->getDGroup3(op, adaptor, rewriter, loc, consts);
4445 SmallVector<Value> results = {dgroup0, dgroup1, dgroup2, dgroup3};
4446 rewriter.replaceOpWithMultiple(op, {results});
4451template <
typename SourceOp,
typename TargetOp>
4452struct AMDGPUTensorLoadStoreOpLowering
4453 :
public ConvertOpToLLVMPattern<SourceOp> {
4454 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
4456 AMDGPUTensorLoadStoreOpLowering(
const LLVMTypeConverter &converter,
4458 : ConvertOpToLLVMPattern<SourceOp>(converter), chipset(chipset) {}
4462 matchAndRewrite(SourceOp op, Adaptor adaptor,
4463 ConversionPatternRewriter &rewriter)
const override {
4465 return op->emitOpError(
"is only supported on gfx1250");
4470 auto v8i32 = VectorType::get(8, rewriter.getI32Type());
4471 Value dgroup4 = LLVM::ZeroOp::create(rewriter, op.getLoc(), v8i32);
4472 Attribute cachePolicy = rewriter.getI32IntegerAttr(0);
4473 rewriter.replaceOpWithNewOp<TargetOp>(op, desc[0], desc[1], desc[2],
4474 desc[3], dgroup4, cachePolicy,
4482struct GlobalPrefetchOpLowering
4483 :
public ConvertOpToLLVMPattern<GlobalPrefetchOp> {
4484 GlobalPrefetchOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
4485 : ConvertOpToLLVMPattern<GlobalPrefetchOp>(converter), chipset(chipset) {}
4488 matchAndRewrite(GlobalPrefetchOp op, GlobalPrefetchOpAdaptor adaptor,
4489 ConversionPatternRewriter &rewriter)
const override {
4491 return op->emitOpError(
"is only supported on gfx1250+");
4493 const bool isSpeculative = op.getSpeculative();
4495 op.getTemporalHint(), op.getCacheScope(), isSpeculative);
4498 Attribute cachePolicy = ROCDL::Gfx12CachePolicyAttr::get(
4499 rewriter.getContext(),
4500 static_cast<ROCDL::Gfx12CachePolicy
>(immArgValue));
4503 Value memRef = adaptor.getSrc();
4504 MemRefDescriptor descriptor(memRef);
4505 MemRefType memRefType = op.getSrc().getType();
4506 Location loc = op->getLoc();
4507 auto inboundsFlags = isSpeculative ? LLVM::GEPNoWrapFlags::none
4508 : LLVM::GEPNoWrapFlags::inbounds |
4509 LLVM::GEPNoWrapFlags::nuw;
4511 rewriter, loc, memRefType, descriptor,
indices, inboundsFlags);
4513 rewriter.replaceOpWithNewOp<ROCDL::GlobalPrefetchOp>(
4514 op, prefetchPtr, cachePolicy, mlir::ArrayAttr{}, mlir::ArrayAttr{},
4523struct ConvertAMDGPUToROCDLPass
4524 :
public impl::ConvertAMDGPUToROCDLPassBase<ConvertAMDGPUToROCDLPass> {
4527 void runOnOperation()
override {
4530 if (
failed(maybeChipset)) {
4531 emitError(UnknownLoc::get(ctx),
"Invalid chipset name: " + chipset);
4532 return signalPassFailure();
4535 RewritePatternSet patterns(ctx);
4536 LLVMTypeConverter converter(ctx);
4541 target.addIllegalDialect<::mlir::amdgpu::AMDGPUDialect>();
4542 target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
4543 target.addLegalDialect<::mlir::ROCDL::ROCDLDialect>();
4544 if (
failed(applyPartialConversion(getOperation(),
target,
4545 std::move(patterns))))
4546 signalPassFailure();
4554 typeConverter, [](gpu::AddressSpace space) {
4556 case gpu::AddressSpace::Global:
4557 return ROCDL::ROCDLDialect::kGlobalMemoryAddressSpace;
4558 case gpu::AddressSpace::Workgroup:
4559 return ROCDL::ROCDLDialect::kSharedMemoryAddressSpace;
4560 case gpu::AddressSpace::Private:
4561 return ROCDL::ROCDLDialect::kPrivateMemoryAddressSpace;
4562 case gpu::AddressSpace::Constant:
4563 return ROCDL::ROCDLDialect::kConstantMemoryAddressSpace;
4565 llvm_unreachable(
"unknown address space enum value");
4568 return LLVM::LLVMPointerType::get(
4569 type.getContext(), ROCDL::ROCDLDialect::kSharedMemoryAddressSpace);
4575 typeConverter.addTypeAttributeConversion(
4577 -> TypeConverter::AttributeConversionResult {
4579 Type i64 = IntegerType::get(ctx, 64);
4580 switch (as.getValue()) {
4581 case amdgpu::AddressSpace::FatRawBuffer:
4582 return IntegerAttr::get(i64, 7);
4583 case amdgpu::AddressSpace::BufferRsrc:
4584 return IntegerAttr::get(i64, 8);
4585 case amdgpu::AddressSpace::FatStructuredBuffer:
4586 return IntegerAttr::get(i64, 9);
4588 return TypeConverter::AttributeConversionResult::abort();
4590 typeConverter.addConversion([&](DsBarrierStateType type) ->
Type {
4591 return IntegerType::get(type.
getContext(), 64);
4593 typeConverter.addConversion([&](TDMBaseType type) ->
Type {
4595 return typeConverter.convertType(VectorType::get(4, i32));
4597 typeConverter.addConversion([&](TDMGatherBaseType type) ->
Type {
4599 return typeConverter.convertType(VectorType::get(4, i32));
4601 typeConverter.addConversion(
4602 [&](TDMDescriptorType type,
4605 Type v4i32 = typeConverter.convertType(VectorType::get(4, i32));
4606 Type v8i32 = typeConverter.convertType(VectorType::get(8, i32));
4607 llvm::append_values(
result, v4i32, v8i32, v4i32, v4i32);
4617 if (inputs.size() != 1)
4620 if (!isa<TDMDescriptorType>(inputs[0].
getType()))
4623 auto cast = UnrealizedConversionCastOp::create(builder, loc, types, inputs);
4624 return cast.getResults();
4627 typeConverter.addTargetMaterialization(addUnrealizedCast);
4635 .
add<FatRawBufferCastLowering,
4636 RawBufferOpLowering<RawBufferLoadOp, ROCDL::RawPtrBufferLoadOp>,
4637 RawBufferOpLowering<RawBufferStoreOp, ROCDL::RawPtrBufferStoreOp>,
4638 RawBufferOpLowering<RawBufferAtomicFaddOp,
4639 ROCDL::RawPtrBufferAtomicFaddOp>,
4640 RawBufferOpLowering<RawBufferAtomicFmaxOp,
4641 ROCDL::RawPtrBufferAtomicFmaxOp>,
4642 RawBufferOpLowering<RawBufferAtomicSmaxOp,
4643 ROCDL::RawPtrBufferAtomicSmaxOp>,
4644 RawBufferOpLowering<RawBufferAtomicUminOp,
4645 ROCDL::RawPtrBufferAtomicUminOp>,
4646 RawBufferOpLowering<RawBufferAtomicCmpswapOp,
4647 ROCDL::RawPtrBufferAtomicCmpSwap>,
4648 AMDGPUDPPLowering, MemoryCounterWaitOpLowering, LDSBarrierOpLowering,
4649 SchedBarrierOpLowering, MFMAOpLowering, ScaledMFMAOpLowering,
4650 SparseMFMAOpLowering, WMMAOpLowering, ScaledWMMAOpLowering,
4651 SparseWMMAOpLowering, DotOpLowering, ExtPackedFp8OpLowering,
4652 ScaledExtPackedMatrixOpLowering, ScaledExtPackedOpLowering,
4653 PackedScaledTruncOpLowering, PackedTrunc2xFp8OpLowering,
4654 PackedStochRoundFp8OpLowering, GatherToLDSOpLowering,
4655 GlobalLoadAsyncToLDSOpLowering, TransposeLoadOpLowering,
4656 GlobalTransposeLoadOpLowering, AMDGPUPermlaneLowering,
4657 AMDGPUPermlaneVarLowering, AMDGPUMakeDmaBaseLowering<MakeDmaBaseOp>,
4658 AMDGPUMakeDmaBaseLowering<MakeGatherDmaBaseOp>,
4659 AMDGPULowerDescriptor<MakeDmaDescriptorOp>,
4660 AMDGPULowerDescriptor<MakeGatherDmaDescriptorOp>,
4661 AMDGPUTensorLoadStoreOpLowering<TensorLoadToLDSOp,
4662 ROCDL::TensorLoadToLDSOp>,
4663 AMDGPUTensorLoadStoreOpLowering<TensorStoreFromLDSOp,
4664 ROCDL::TensorStoreFromLDSOp>,
4665 DsBarrierInitOpLowering, DsBarrierPollStateOpLowering,
4666 DsAsyncBarrierArriveOpLowering, DsBarrierArriveOpLowering,
4667 GlobalPrefetchOpLowering>(converter, chipset);
4668 patterns.
add<AMDGPUSwizzleBitModeLowering, DsBarrierStatePhaseOpLowering,
4669 DsBarrierStatePendingCountOpLowering,
4670 DsBarrierStateInitCountOpLowering,
4671 DsBarrierStatePhaseParityLowering>(converter);
static bool typeIsExpectedFp8ForChipset(Chipset chipset, Type type)
Return true if type is the E4M3FN variant of an 8-bit float that is supported by the _fp8 instruction...
constexpr Chipset kGfx942
static std::optional< StringRef > wmmaOpToIntrinsicRDNA(Type elemSourceType, Type elemBSourceType, Type elemDestType, uint32_t k, bool isRDNA3)
Returns the rocdl intrinsic corresponding to a WMMA operation wmma for RDNA3/4 architectures.
static bool hasDot10Insts(const Chipset &chipset)
static bool hasDot7Insts(const Chipset &chipset)
static std::optional< SparseWMMAOpInfo > sparseWMMAOpToIntrinsic(SparseWMMAOp swmmac, Chipset chipset)
static std::optional< StringRef > mfmaOpToIntrinsic(MFMAOp mfma, Chipset chipset)
Return the rocdl intrinsic corresponding to a MFMA operation mfma if one exists.
static Value convertUnsignedToInt(ConversionPatternRewriter &rewriter, Location loc, Value val, unsigned width)
Zero-extend or truncate the unsigned number val to width bits.
constexpr Chipset kGfx908
static void wmmaPushInputOperand(ConversionPatternRewriter &rewriter, Location loc, const TypeConverter *typeConverter, bool isUnsigned, Value llvmInput, Value mlirInput, SmallVectorImpl< Value > &operands, SmallVectorImpl< NamedAttribute > &attrs, StringRef attrName)
Push an input operand.
static std::optional< ScaledMFMAIntrinsic > mfmaOpToScaledIntrinsic(Type aType, Type bType, Type destType, uint32_t m, uint32_t n, uint32_t k, uint32_t b, Chipset chipset)
constexpr Chipset kGfx1250
static Value castScaleOperand(ConversionPatternRewriter &rewriter, Location loc, Value input)
Converts the scaled MFMA/WMMA operands, scalesA and scalesB, from MLIR AMDGPU dialect convention to R...
constexpr Chipset kGfx90a
static std::optional< StringRef > getScaledWmmaIntrinsicName(int64_t m, int64_t n, int64_t k, bool isScale16)
Determines the ROCDL intrinsic name for scaled WMMA based on dimensions and scale block size (16 or 3...
static void wmmaPushOutputOperand(ConversionPatternRewriter &rewriter, Location loc, const TypeConverter *typeConverter, Value output, int32_t subwordOffset, bool clamp, SmallVectorImpl< Value > &operands, SmallVectorImpl< NamedAttribute > &attrs)
Push the output operand.
static bool typeIsExpectedBf8ForChipset(Chipset chipset, Type type)
Return true if type is the E5M2 variant of an 8-bit float that is supported by the _bf8 instructions ...
static std::optional< StringRef > wmmaOpToIntrinsic(WMMAOp wmma, Chipset chipset)
Returns the rocdl intrinsic corresponding to a WMMA operation wmma if one exists.
static bool hasDot11Insts(const Chipset &chipset)
static std::optional< StringRef > smfmacOpToIntrinsic(SparseMFMAOp op, Chipset chipset)
Returns the rocdl intrinsic corresponding to a SparseMFMA (smfmac) operation if one exists.
static Value makeBufferRsrc(ConversionPatternRewriter &rewriter, Location loc, Value basePointer, Value numRecords, bool boundsCheck, amdgpu::Chipset chipset, Value cacheSwizzleStride=nullptr, unsigned addressSpace=8)
static Value createI64Constant(ConversionPatternRewriter &rewriter, Location loc, int64_t value)
static bool hasDot9Insts(const Chipset &chipset)
static std::optional< StringRef > wmmaOpToIntrinsicGfx1250(Type elemSourceType, Type elemBSourceType, Type elemDestType, uint32_t k)
Return the rocdl intrinsic corresponding to a WMMA operation wmma for the gfx1250 architecture.
constexpr Chipset kGfx1200
static Value getNumRecords(ConversionPatternRewriter &rewriter, Location loc, MemRefType memrefType, MemRefDescriptor &memrefDescriptor, ArrayRef< int64_t > strides, int64_t elementByteWidth, amdgpu::Chipset chipset, bool boundsCheck)
Compute the contents of the num_records field for a given memref descriptor - that is,...
static Value packSmallFloatVectorOperand(ConversionPatternRewriter &rewriter, Location loc, Value input, bool allowBf16=true)
Pack small float vector operands (fp4/fp6/fp8/bf16) into the format expected by scaled matrix multipl...
static bool has45BitNumRecordsBufferResource(const Chipset &chipset)
static std::optional< ROCDL::WMMAMatrixScaleFormat > getWmmaScaleFormat(Type elemType)
Maps f8 scale element types to WMMA scale format codes.
static Value convertPackedVectorOperand(ConversionPatternRewriter &rewriter, Location loc, Value input, bool allowBf16=true)
Converts packed vector operands to the expected ROCDL types.
static Value getLinearIndexI32(ConversionPatternRewriter &rewriter, Location loc, MemRefDescriptor &memRefDescriptor, ValueRange indices, ArrayRef< int64_t > strides)
Returns the linear index used to access an element in the memref.
static Value convertUnsignedToI32(ConversionPatternRewriter &rewriter, Location loc, Value val)
Convert an unsigned number val to i32.
static bool hasDot8Insts(const Chipset &chipset)
static bool hasDot2Insts(const Chipset &chipset)
static Value createI32Constant(ConversionPatternRewriter &rewriter, Location loc, int32_t value)
static std::optional< ROCDL::MatrixFormat > smallFloatTypeToMatrixFormat(Type mlirElemType)
std::tuple< StringRef, ROCDL::MatrixFormat, ROCDL::MatrixFormat > ScaledMFMAIntrinsic
If there is a scaled MFMA instruction for the input element types aType and bType,...
static bool hasDot12Insts(const Chipset &chipset)
static Value convertUnsignedToI64(ConversionPatternRewriter &rewriter, Location loc, Value val)
Convert an unsigned number val to i64.
constexpr Chipset kGfx950
static bool hasDot1Insts(const Chipset &chipset)
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static constexpr unsigned kSizePosInMemRefDescriptor
static constexpr unsigned kStridePosInMemRefDescriptor
static constexpr unsigned kOffsetPosInMemRefDescriptor
static constexpr unsigned kAllocatedPtrPosInMemRefDescriptor
static constexpr unsigned kAlignedPtrPosInMemRefDescriptor
static Value clamp(ImplicitLocOpBuilder &builder, Value value, Value lowerBound, Value upperBound)
Attributes are known-constant values of operations.
This class provides a shared interface for ranked and unranked memref types.
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
typename SourceOp::template GenericAdaptor< ArrayRef< ValueRange > > OneToNOpAdaptor
typename SourceOp::Adaptor OpAdaptor
Value getStridedElementPtr(ConversionPatternRewriter &rewriter, Location loc, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none) const
Convenience wrapper for the corresponding helper utility.
llvm::TypeSize getTypeSizeInBits(Type t) const
Returns the size in bits of the given type in the current scope.
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
Helper class to produce LLVM dialect operations extracting or inserting elements of a MemRef descript...
Value stride(OpBuilder &builder, Location loc, unsigned pos)
Builds IR extracting the pos-th size from the descriptor.
Value size(OpBuilder &builder, Location loc, unsigned pos)
Builds IR extracting the pos-th size from the descriptor.
NamedAttribute represents a combination of a name and an Attribute value.
This class helps build Operations.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
result_range getResults()
unsigned getNumResults()
Return the number of results held by this operation.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isSignedInteger() const
Return true if this is a signed integer type (with the specified width).
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
::mlir::Pass::Option< std::string > chipset
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
LogicalResult decomposeValue(OpBuilder &builder, Location loc, Value src, Type dstType, SmallVectorImpl< Value > &result, bool permitVariablySizedScalars=false)
Decomposes a src value into a set of values of type dstType through series of bitcasts and vector ops...
Value createIndexAttrConstant(OpBuilder &builder, Location loc, Type resultType, int64_t value)
Creates an llvm.mlir.constant producing value as resultType, which is expected to be the converted in...
Value composeValue(OpBuilder &builder, Location loc, ValueRange src, Type dstType)
Composes a set of src values into a single value of type dstType through series of bitcasts and vecto...
int32_t getGlobalPrefetchLLVMEncoding(amdgpu::LoadTemporalHint hint, amdgpu::Scope scope, bool isSpeculative)
bool hasOcpFp8(const Chipset &chipset)
void populateCommonGPUTypeAndAttributeConversions(TypeConverter &typeConverter)
Remap common GPU memory spaces (Workgroup, Private, etc) to LLVM address spaces.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void populateGpuMemorySpaceAttributeConversions(TypeConverter &typeConverter, const MemorySpaceMapping &mapping)
Populates memory space attribute conversion rules for lowering gpu.address_space to integer values.
void populateAMDGPUToROCDLConversionPatterns(LLVMTypeConverter &converter, RewritePatternSet &patterns, amdgpu::Chipset chipset)
Note: This function will also add conversions for the AMDGPU-specific address spaces and types,...
llvm::TypeSwitch< T, ResultT > TypeSwitch
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
void populateAMDGPUTypeAndAttributeConversions(TypeConverter &typeConverter)
Remap AMDGPU memory spaces to LLVM address spaces by mapping amdgpu::AddressSpace::fat_raw_buffer to ...
Returns the rocdl intrinsic corresponding to a SparseWMMA operation swmmac if one exists.
Represents the amdgpu gfx chipset version, e.g., gfx90a, gfx942, gfx1103.
static FailureOr< Chipset > parse(StringRef name)
Parses the chipset version string and returns the chipset on success, and failure otherwise.