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,
672 }
else if (chipset.majorVersion < 12) {
673 ROCDL::SBarrierOp::create(rewriter, loc);
675 ROCDL::BarrierSignalOp::create(rewriter, loc, -1);
676 ROCDL::BarrierWaitOp::create(rewriter, loc, -1);
679 auto acqFence = LLVM::FenceOp::create(rewriter, loc,
680 LLVM::AtomicOrdering::acquire, scope);
681 acqFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);
682 rewriter.replaceOp(op, acqFence);
688 SchedBarrierOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
689 : ConvertOpToLLVMPattern<SchedBarrierOp>(converter), chipset(chipset) {}
694 matchAndRewrite(SchedBarrierOp op, SchedBarrierOp::Adaptor adaptor,
695 ConversionPatternRewriter &rewriter)
const override {
696 rewriter.replaceOpWithNewOp<ROCDL::SchedBarrier>(op, op.getOptsAttr());
720 bool allowBf16 =
true) {
722 if (
auto vectorType = dyn_cast<VectorType>(inputType)) {
723 if (vectorType.getElementType().isBF16() && !allowBf16)
724 return LLVM::BitcastOp::create(
725 rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);
726 if (vectorType.getElementType().isInteger(8) &&
727 vectorType.getNumElements() <= 8)
728 return LLVM::BitcastOp::create(
730 rewriter.getIntegerType(vectorType.getNumElements() * 8), input);
731 if (isa<IntegerType>(vectorType.getElementType()) &&
732 vectorType.getElementTypeBitWidth() <= 8) {
733 int64_t numWords = llvm::divideCeil(
734 vectorType.getNumElements() * vectorType.getElementTypeBitWidth(),
736 return LLVM::BitcastOp::create(
737 rewriter, loc, VectorType::get(numWords, rewriter.getI32Type()),
747 bool allowBf16 =
true) {
749 auto vectorType = cast<VectorType>(inputType);
751 if (vectorType.getElementType().isBF16() && !allowBf16)
752 return LLVM::BitcastOp::create(
753 rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);
755 if (isa<IntegerType>(vectorType.getElementType()) &&
756 vectorType.getElementTypeBitWidth() <= 8) {
757 int64_t numWords = llvm::divideCeil(
758 vectorType.getNumElements() * vectorType.getElementTypeBitWidth(), 32);
759 Type castType = (numWords > 1)
760 ?
Type{VectorType::get(numWords, rewriter.getI32Type())}
761 : rewriter.getI32Type();
762 return LLVM::BitcastOp::create(rewriter, loc, castType, input);
780 .Case([&](IntegerType) {
782 return LLVM::ZExtOp::create(rewriter, loc, rewriter.getI32Type(),
785 .Case([&](VectorType vectorType) {
787 int64_t numElements = vectorType.getNumElements();
788 assert((numElements == 4 || numElements == 8) &&
789 "scale operand must be a vector of length 4 or 8");
790 IntegerType outputType =
791 (numElements == 4) ? rewriter.getI32Type() : rewriter.getI64Type();
792 return LLVM::BitcastOp::create(rewriter, loc, outputType, input);
794 .DefaultUnreachable(
"unexpected input type for scale operand");
798static std::optional<ROCDL::WMMAMatrixScaleFormat>
801 .Case([](Float8E8M0FNUType) {
return ROCDL::WMMAMatrixScaleFormat::e8; })
802 .Case([](Float8E4M3FNType) {
return ROCDL::WMMAMatrixScaleFormat::e4m3; })
803 .Default(std::nullopt);
808static std::optional<StringRef>
810 if (m == 16 && n == 16 && k == 128)
812 ? ROCDL::wmma_scale16_f32_16x16x128_f8f6f4::getOperationName()
813 : ROCDL::wmma_scale_f32_16x16x128_f8f6f4::getOperationName();
815 if (m == 32 && n == 16 && k == 128)
816 return isScale16 ? ROCDL::wmma_scale16_f32_32x16x128_f4::getOperationName()
817 : ROCDL::wmma_scale_f32_32x16x128_f4::getOperationName();
831 ConversionPatternRewriter &rewriter,
Location loc,
836 auto vectorType = dyn_cast<VectorType>(inputType);
838 operands.push_back(llvmInput);
841 Type elemType = vectorType.getElementType();
843 operands.push_back(llvmInput);
850 auto mlirInputType = cast<VectorType>(mlirInput.
getType());
851 bool isInputInteger = mlirInputType.getElementType().isInteger();
852 if (isInputInteger) {
854 bool localIsUnsigned = isUnsigned;
856 localIsUnsigned =
true;
858 localIsUnsigned =
false;
861 NamedAttribute(attrName, rewriter.getBoolAttr(!localIsUnsigned)));
866 Type i32 = rewriter.getI32Type();
867 Type intrinsicInType = numBits <= 32
868 ? (
Type)rewriter.getIntegerType(numBits)
869 : (
Type)VectorType::get(numBits / 32, i32);
870 auto llvmIntrinsicInType = typeConverter->convertType(intrinsicInType);
871 Value castInput = rewriter.createOrFold<LLVM::BitcastOp>(
872 loc, llvmIntrinsicInType, llvmInput);
877 castInput = LLVM::ZExtOp::create(rewriter, loc, i32, castInput);
878 operands.push_back(castInput);
891 Value output, int32_t subwordOffset,
895 auto vectorType = dyn_cast<VectorType>(inputType);
896 Type elemType = vectorType.getElementType();
897 operands.push_back(output);
909 return (chipset ==
kGfx942 && isa<Float8E5M2FNUZType>(type)) ||
910 (
hasOcpFp8(chipset) && isa<Float8E5M2Type>(type));
916 return (chipset ==
kGfx942 && isa<Float8E4M3FNUZType>(type)) ||
917 (
hasOcpFp8(chipset) && isa<Float8E4M3FNType>(type));
925 uint32_t m = mfma.getM(), n = mfma.getN(), k = mfma.getK(),
926 b = mfma.getBlocks();
931 if (mfma.getReducePrecision() && chipset >=
kGfx942) {
932 if (m == 32 && n == 32 && k == 4 &&
b == 1)
933 return ROCDL::mfma_f32_32x32x4_xf32::getOperationName();
934 if (m == 16 && n == 16 && k == 8 &&
b == 1)
935 return ROCDL::mfma_f32_16x16x8_xf32::getOperationName();
937 if (m == 32 && n == 32 && k == 1 &&
b == 2)
938 return ROCDL::mfma_f32_32x32x1f32::getOperationName();
939 if (m == 16 && n == 16 && k == 1 &&
b == 4)
940 return ROCDL::mfma_f32_16x16x1f32::getOperationName();
941 if (m == 4 && n == 4 && k == 1 &&
b == 16)
942 return ROCDL::mfma_f32_4x4x1f32::getOperationName();
943 if (m == 32 && n == 32 && k == 2 &&
b == 1)
944 return ROCDL::mfma_f32_32x32x2f32::getOperationName();
945 if (m == 16 && n == 16 && k == 4 &&
b == 1)
946 return ROCDL::mfma_f32_16x16x4f32::getOperationName();
951 if (m == 32 && n == 32 && k == 16 &&
b == 1)
952 return ROCDL::mfma_f32_32x32x16_f16::getOperationName();
953 if (m == 16 && n == 16 && k == 32 &&
b == 1)
954 return ROCDL::mfma_f32_16x16x32_f16::getOperationName();
956 if (m == 32 && n == 32 && k == 4 &&
b == 2)
957 return ROCDL::mfma_f32_32x32x4f16::getOperationName();
958 if (m == 16 && n == 16 && k == 4 &&
b == 4)
959 return ROCDL::mfma_f32_16x16x4f16::getOperationName();
960 if (m == 4 && n == 4 && k == 4 &&
b == 16)
961 return ROCDL::mfma_f32_4x4x4f16::getOperationName();
962 if (m == 32 && n == 32 && k == 8 &&
b == 1)
963 return ROCDL::mfma_f32_32x32x8f16::getOperationName();
964 if (m == 16 && n == 16 && k == 16 &&
b == 1)
965 return ROCDL::mfma_f32_16x16x16f16::getOperationName();
970 if (m == 32 && n == 32 && k == 16 &&
b == 1)
971 return ROCDL::mfma_f32_32x32x16_bf16::getOperationName();
972 if (m == 16 && n == 16 && k == 32 &&
b == 1)
973 return ROCDL::mfma_f32_16x16x32_bf16::getOperationName();
976 if (m == 32 && n == 32 && k == 4 &&
b == 2)
977 return ROCDL::mfma_f32_32x32x4bf16_1k::getOperationName();
978 if (m == 16 && n == 16 && k == 4 &&
b == 4)
979 return ROCDL::mfma_f32_16x16x4bf16_1k::getOperationName();
980 if (m == 4 && n == 4 && k == 4 &&
b == 16)
981 return ROCDL::mfma_f32_4x4x4bf16_1k::getOperationName();
982 if (m == 32 && n == 32 && k == 8 &&
b == 1)
983 return ROCDL::mfma_f32_32x32x8bf16_1k::getOperationName();
984 if (m == 16 && n == 16 && k == 16 &&
b == 1)
985 return ROCDL::mfma_f32_16x16x16bf16_1k::getOperationName();
987 if (m == 32 && n == 32 && k == 2 &&
b == 2)
988 return ROCDL::mfma_f32_32x32x2bf16::getOperationName();
989 if (m == 16 && n == 16 && k == 2 &&
b == 4)
990 return ROCDL::mfma_f32_16x16x2bf16::getOperationName();
991 if (m == 4 && n == 4 && k == 2 &&
b == 16)
992 return ROCDL::mfma_f32_4x4x2bf16::getOperationName();
993 if (m == 32 && n == 32 && k == 4 &&
b == 1)
994 return ROCDL::mfma_f32_32x32x4bf16::getOperationName();
995 if (m == 16 && n == 16 && k == 8 &&
b == 1)
996 return ROCDL::mfma_f32_16x16x8bf16::getOperationName();
1001 if (m == 32 && n == 32 && k == 32 &&
b == 1)
1002 return ROCDL::mfma_i32_32x32x32_i8::getOperationName();
1003 if (m == 16 && n == 16 && k == 64 &&
b == 1)
1004 return ROCDL::mfma_i32_16x16x64_i8::getOperationName();
1006 if (m == 32 && n == 32 && k == 4 &&
b == 2)
1007 return ROCDL::mfma_i32_32x32x4i8::getOperationName();
1008 if (m == 16 && n == 16 && k == 4 &&
b == 4)
1009 return ROCDL::mfma_i32_16x16x4i8::getOperationName();
1010 if (m == 4 && n == 4 && k == 4 &&
b == 16)
1011 return ROCDL::mfma_i32_4x4x4i8::getOperationName();
1012 if (m == 32 && n == 32 && k == 8 &&
b == 1)
1013 return ROCDL::mfma_i32_32x32x8i8::getOperationName();
1014 if (m == 16 && n == 16 && k == 16 &&
b == 1)
1015 return ROCDL::mfma_i32_16x16x16i8::getOperationName();
1016 if (m == 32 && n == 32 && k == 16 &&
b == 1 && chipset >=
kGfx942)
1017 return ROCDL::mfma_i32_32x32x16_i8::getOperationName();
1018 if (m == 16 && n == 16 && k == 32 &&
b == 1 && chipset >=
kGfx942)
1019 return ROCDL::mfma_i32_16x16x32_i8::getOperationName();
1023 if (m == 16 && n == 16 && k == 4 &&
b == 1)
1024 return ROCDL::mfma_f64_16x16x4f64::getOperationName();
1025 if (m == 4 && n == 4 && k == 4 &&
b == 4)
1026 return ROCDL::mfma_f64_4x4x4f64::getOperationName();
1033 cast<VectorType>(mfma.getSourceB().getType()).getElementType();
1034 if (m == 16 && n == 16 && k == 32 &&
b == 1) {
1036 return ROCDL::mfma_f32_16x16x32_bf8_bf8::getOperationName();
1038 return ROCDL::mfma_f32_16x16x32_bf8_fp8::getOperationName();
1040 if (m == 32 && n == 32 && k == 16 &&
b == 1) {
1042 return ROCDL::mfma_f32_32x32x16_bf8_bf8::getOperationName();
1044 return ROCDL::mfma_f32_32x32x16_bf8_fp8::getOperationName();
1050 cast<VectorType>(mfma.getSourceB().getType()).getElementType();
1051 if (m == 16 && n == 16 && k == 32 &&
b == 1) {
1053 return ROCDL::mfma_f32_16x16x32_fp8_bf8::getOperationName();
1055 return ROCDL::mfma_f32_16x16x32_fp8_fp8::getOperationName();
1057 if (m == 32 && n == 32 && k == 16 &&
b == 1) {
1059 return ROCDL::mfma_f32_32x32x16_fp8_bf8::getOperationName();
1061 return ROCDL::mfma_f32_32x32x16_fp8_fp8::getOperationName();
1065 return std::nullopt;
1068static std::optional<ROCDL::MatrixFormat>
1072 .Case([](Float8E4M3FNType) {
return ROCDL::MatrixFormat::fp8_e4m3; })
1073 .Case([](Float8E5M2Type) {
return ROCDL::MatrixFormat::fp8_e5m2; })
1074 .Case([](Float6E2M3FNType) {
return ROCDL::MatrixFormat::fp6_e2m3; })
1075 .Case([](Float6E3M2FNType) {
return ROCDL::MatrixFormat::fp6_e3m2; })
1076 .Case([](Float4E2M1FNType) {
return ROCDL::MatrixFormat::fp4_e2m1; })
1077 .Default(std::nullopt);
1088 std::tuple<StringRef, ROCDL::MatrixFormat, ROCDL::MatrixFormat>;
1090static std::optional<ScaledMFMAIntrinsic>
1092 uint32_t n, uint32_t k, uint32_t
b,
Chipset chipset) {
1098 return std::nullopt;
1099 if (!isa<Float32Type>(destType))
1100 return std::nullopt;
1102 std::optional<ROCDL::MatrixFormat> aTypeCode =
1104 std::optional<ROCDL::MatrixFormat> bTypeCode =
1106 if (!aTypeCode || !bTypeCode)
1107 return std::nullopt;
1109 if (m == 32 && n == 32 && k == 64 &&
b == 1)
1110 return std::tuple{ROCDL::mfma_scale_f32_32x32x64_f8f6f4::getOperationName(),
1111 *aTypeCode, *bTypeCode};
1112 if (m == 16 && n == 16 && k == 128 &&
b == 1)
1114 ROCDL::mfma_scale_f32_16x16x128_f8f6f4::getOperationName(), *aTypeCode,
1117 return std::nullopt;
1120static std::optional<ScaledMFMAIntrinsic>
1123 mfma.getSourceA().getType(), mfma.getSourceB().getType(),
1124 mfma.getDestC().getType(), mfma.getM(), mfma.getN(), mfma.getK(),
1125 mfma.getBlocks(), chipset);
1128static std::optional<ScaledMFMAIntrinsic>
1131 smfma.getSourceB().getType(),
1132 smfma.getDestC().getType(), smfma.getM(),
1133 smfma.getN(), smfma.getK(), 1u, chipset);
1138static std::optional<StringRef>
1140 Type elemDestType, uint32_t k,
bool isRDNA3) {
1141 using fp8 = Float8E4M3FNType;
1142 using bf8 = Float8E5M2Type;
1147 if (elemSourceType.
isF16() && elemDestType.
isF32())
1148 return ROCDL::wmma_f32_16x16x16_f16::getOperationName();
1149 if (elemSourceType.
isBF16() && elemDestType.
isF32())
1150 return ROCDL::wmma_f32_16x16x16_bf16::getOperationName();
1151 if (elemSourceType.
isF16() && elemDestType.
isF16())
1152 return ROCDL::wmma_f16_16x16x16_f16::getOperationName();
1154 return ROCDL::wmma_bf16_16x16x16_bf16::getOperationName();
1156 return ROCDL::wmma_i32_16x16x16_iu8::getOperationName();
1161 return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();
1162 return std::nullopt;
1166 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType) &&
1167 elemDestType.
isF32())
1168 return ROCDL::wmma_f32_16x16x16_fp8_fp8::getOperationName();
1169 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType) &&
1170 elemDestType.
isF32())
1171 return ROCDL::wmma_f32_16x16x16_fp8_bf8::getOperationName();
1172 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType) &&
1173 elemDestType.
isF32())
1174 return ROCDL::wmma_f32_16x16x16_bf8_bf8::getOperationName();
1175 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType) &&
1176 elemDestType.
isF32())
1177 return ROCDL::wmma_f32_16x16x16_bf8_fp8::getOperationName();
1179 return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();
1181 return std::nullopt;
1185 if (k == 32 && !isRDNA3) {
1187 return ROCDL::wmma_i32_16x16x32_iu4::getOperationName();
1190 return std::nullopt;
1196 Type elemBSourceType,
1199 using fp8 = Float8E4M3FNType;
1200 using bf8 = Float8E5M2Type;
1203 if (elemSourceType.
isF32() && elemDestType.
isF32())
1204 return ROCDL::wmma_f32_16x16x4_f32::getOperationName();
1206 return std::nullopt;
1210 if (elemSourceType.
isF16() && elemDestType.
isF32())
1211 return ROCDL::wmma_f32_16x16x32_f16::getOperationName();
1212 if (elemSourceType.
isBF16() && elemDestType.
isF32())
1213 return ROCDL::wmma_f32_16x16x32_bf16::getOperationName();
1214 if (elemSourceType.
isF16() && elemDestType.
isF16())
1215 return ROCDL::wmma_f16_16x16x32_f16::getOperationName();
1217 return ROCDL::wmma_bf16_16x16x32_bf16::getOperationName();
1219 return std::nullopt;
1223 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1224 if (elemDestType.
isF32())
1225 return ROCDL::wmma_f32_16x16x64_fp8_fp8::getOperationName();
1226 if (elemDestType.
isF16())
1227 return ROCDL::wmma_f16_16x16x64_fp8_fp8::getOperationName();
1229 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1230 if (elemDestType.
isF32())
1231 return ROCDL::wmma_f32_16x16x64_fp8_bf8::getOperationName();
1232 if (elemDestType.
isF16())
1233 return ROCDL::wmma_f16_16x16x64_fp8_bf8::getOperationName();
1235 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1236 if (elemDestType.
isF32())
1237 return ROCDL::wmma_f32_16x16x64_bf8_bf8::getOperationName();
1238 if (elemDestType.
isF16())
1239 return ROCDL::wmma_f16_16x16x64_bf8_bf8::getOperationName();
1241 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1242 if (elemDestType.
isF32())
1243 return ROCDL::wmma_f32_16x16x64_bf8_fp8::getOperationName();
1244 if (elemDestType.
isF16())
1245 return ROCDL::wmma_f16_16x16x64_bf8_fp8::getOperationName();
1248 return ROCDL::wmma_i32_16x16x64_iu8::getOperationName();
1250 return std::nullopt;
1254 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1255 if (elemDestType.
isF32())
1256 return ROCDL::wmma_f32_16x16x128_fp8_fp8::getOperationName();
1257 if (elemDestType.
isF16())
1258 return ROCDL::wmma_f16_16x16x128_fp8_fp8::getOperationName();
1260 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1261 if (elemDestType.
isF32())
1262 return ROCDL::wmma_f32_16x16x128_fp8_bf8::getOperationName();
1263 if (elemDestType.
isF16())
1264 return ROCDL::wmma_f16_16x16x128_fp8_bf8::getOperationName();
1266 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1267 if (elemDestType.
isF32())
1268 return ROCDL::wmma_f32_16x16x128_bf8_bf8::getOperationName();
1269 if (elemDestType.
isF16())
1270 return ROCDL::wmma_f16_16x16x128_bf8_bf8::getOperationName();
1272 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1273 if (elemDestType.
isF32())
1274 return ROCDL::wmma_f32_16x16x128_bf8_fp8::getOperationName();
1275 if (elemDestType.
isF16())
1276 return ROCDL::wmma_f16_16x16x128_bf8_fp8::getOperationName();
1279 return std::nullopt;
1282 return std::nullopt;
1290 bool isGfx950 = chipset >=
kGfx950;
1294 uint32_t m = op.getM(), n = op.getN(), k = op.getK();
1299 if (m == 16 && n == 16 && k == 32) {
1301 return ROCDL::smfmac_f32_16x16x32_f16::getOperationName();
1303 return ROCDL::smfmac_f32_16x16x32_bf16::getOperationName();
1306 if (m == 16 && n == 16 && k == 64) {
1309 return ROCDL::smfmac_f32_16x16x64_f16::getOperationName();
1311 return ROCDL::smfmac_f32_16x16x64_bf16::getOperationName();
1315 return ROCDL::smfmac_i32_16x16x64_i8::getOperationName();
1316 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1317 return ROCDL::smfmac_f32_16x16x64_fp8_fp8::getOperationName();
1318 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1319 return ROCDL::smfmac_f32_16x16x64_fp8_bf8::getOperationName();
1320 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1321 return ROCDL::smfmac_f32_16x16x64_bf8_fp8::getOperationName();
1322 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1323 return ROCDL::smfmac_f32_16x16x64_bf8_bf8::getOperationName();
1326 if (m == 16 && n == 16 && k == 128 && isGfx950) {
1329 return ROCDL::smfmac_i32_16x16x128_i8::getOperationName();
1330 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1331 return ROCDL::smfmac_f32_16x16x128_fp8_fp8::getOperationName();
1332 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1333 return ROCDL::smfmac_f32_16x16x128_fp8_bf8::getOperationName();
1334 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1335 return ROCDL::smfmac_f32_16x16x128_bf8_fp8::getOperationName();
1336 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1337 return ROCDL::smfmac_f32_16x16x128_bf8_bf8::getOperationName();
1340 if (m == 32 && n == 32 && k == 16) {
1342 return ROCDL::smfmac_f32_32x32x16_f16::getOperationName();
1344 return ROCDL::smfmac_f32_32x32x16_bf16::getOperationName();
1347 if (m == 32 && n == 32 && k == 32) {
1350 return ROCDL::smfmac_f32_32x32x32_f16::getOperationName();
1352 return ROCDL::smfmac_f32_32x32x32_bf16::getOperationName();
1356 return ROCDL::smfmac_i32_32x32x32_i8::getOperationName();
1357 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1358 return ROCDL::smfmac_f32_32x32x32_fp8_fp8::getOperationName();
1359 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1360 return ROCDL::smfmac_f32_32x32x32_fp8_bf8::getOperationName();
1361 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1362 return ROCDL::smfmac_f32_32x32x32_bf8_fp8::getOperationName();
1363 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1364 return ROCDL::smfmac_f32_32x32x32_bf8_bf8::getOperationName();
1367 if (m == 32 && n == 32 && k == 64 && isGfx950) {
1370 return ROCDL::smfmac_i32_32x32x64_i8::getOperationName();
1371 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1372 return ROCDL::smfmac_f32_32x32x64_fp8_fp8::getOperationName();
1373 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1374 return ROCDL::smfmac_f32_32x32x64_fp8_bf8::getOperationName();
1375 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.
isF32())
1376 return ROCDL::smfmac_f32_32x32x64_bf8_fp8::getOperationName();
1377 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.
isF32())
1378 return ROCDL::smfmac_f32_32x32x64_bf8_bf8::getOperationName();
1381 return std::nullopt;
1389 auto sourceVectorType = cast<VectorType>(wmma.getSourceA().getType());
1390 auto sourceBVectorType = cast<VectorType>(wmma.getSourceB().getType());
1391 auto destVectorType = cast<VectorType>(wmma.getDestC().getType());
1392 Type elemSourceType = sourceVectorType.getElementType();
1393 Type elemBSourceType = sourceBVectorType.getElementType();
1394 Type elemDestType = destVectorType.getElementType();
1396 const uint32_t k = wmma.getK();
1401 if (isRDNA3 || isRDNA4)
1410 return std::nullopt;
1423static std::optional<SparseWMMAOpInfo>
1429 uint32_t m = swmmac.getM(), n = swmmac.getN(), k = swmmac.getK();
1431 if ((m != 16) || (n != 16))
1432 return std::nullopt;
1439 ROCDL::swmmac_f32_16x16x32_f16::getOperationName(),
false,
false,
1443 ROCDL::swmmac_f32_16x16x32_bf16::getOperationName(),
false,
false,
1447 ROCDL::swmmac_f16_16x16x32_f16::getOperationName(),
false,
false,
1451 ROCDL::swmmac_bf16_16x16x32_bf16::getOperationName(),
false,
false,
1456 ROCDL::swmmac_i32_16x16x32_iu8::getOperationName(),
true,
false,
1461 ROCDL::swmmac_i32_16x16x32_iu4::getOperationName(),
true,
false,
1466 ROCDL::swmmac_f32_16x16x32_fp8_fp8::getOperationName(),
false,
1471 ROCDL::swmmac_f32_16x16x32_fp8_bf8::getOperationName(),
false,
1476 ROCDL::swmmac_f32_16x16x32_bf8_fp8::getOperationName(),
false,
1480 ROCDL::swmmac_f32_16x16x32_bf8_bf8::getOperationName(),
false,
1487 ROCDL::swmmac_i32_16x16x64_iu4::getOperationName(),
true,
false,
1492 const bool isGFX1250 = chipset ==
kGfx1250;
1493 const bool isWavesize64 = swmmac.getWave64();
1494 if (isGFX1250 && !isWavesize64) {
1498 ROCDL::swmmac_f32_16x16x64_f16::getOperationName(),
true,
true,
1502 ROCDL::swmmac_f32_16x16x64_bf16::getOperationName(),
true,
true,
1506 ROCDL::swmmac_f16_16x16x64_f16::getOperationName(),
true,
true,
1510 ROCDL::swmmac_bf16_16x16x64_bf16::getOperationName(),
true,
true,
1517 ROCDL::swmmac_f32_16x16x128_fp8_fp8::getOperationName(),
false,
1522 ROCDL::swmmac_f32_16x16x128_fp8_bf8::getOperationName(),
false,
1527 ROCDL::swmmac_f32_16x16x128_bf8_fp8::getOperationName(),
false,
1531 ROCDL::swmmac_f32_16x16x128_bf8_bf8::getOperationName(),
false,
1536 ROCDL::swmmac_f16_16x16x128_fp8_fp8::getOperationName(),
false,
1541 ROCDL::swmmac_f16_16x16x128_fp8_bf8::getOperationName(),
false,
1546 ROCDL::swmmac_f16_16x16x128_bf8_fp8::getOperationName(),
false,
1550 ROCDL::swmmac_f16_16x16x128_bf8_bf8::getOperationName(),
false,
1555 ROCDL::swmmac_f16_16x16x128_bf8_bf8::getOperationName(),
false,
1560 ROCDL::swmmac_i32_16x16x128_iu8::getOperationName(),
true,
true,
1565 return std::nullopt;
1570 MFMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1571 : ConvertOpToLLVMPattern<MFMAOp>(converter), chipset(chipset) {}
1576 matchAndRewrite(MFMAOp op, MFMAOpAdaptor adaptor,
1577 ConversionPatternRewriter &rewriter)
const override {
1578 Location loc = op.getLoc();
1580 Type outType = typeConverter->convertType(op.getDestD().getType());
1581 Type intrinsicOutType = outType;
1582 if (
auto outVecType = dyn_cast<VectorType>(outType))
1583 if (outVecType.getElementType().isBF16())
1584 intrinsicOutType = outVecType.clone(rewriter.getI16Type());
1586 if (chipset.majorVersion != 9 || chipset <
kGfx908)
1587 return op->emitOpError(
"MFMA only supported on gfx908+");
1588 uint32_t getBlgpField =
static_cast<uint32_t
>(op.getBlgp());
1589 if (op.getNegateA() || op.getNegateB() || op.getNegateC()) {
1591 return op.emitOpError(
"negation unsupported on older than gfx942");
1593 op.getNegateA() | (op.getNegateB() << 1) | (op.getNegateC() << 2);
1596 std::optional<ScaledMFMAIntrinsic> maybeScaledIntrinsic =
1598 if (!maybeIntrinsic.has_value() && !maybeScaledIntrinsic.has_value())
1599 return op.emitOpError(
"no intrinsic matching MFMA size on given chipset");
1602 !maybeIntrinsic.has_value() && maybeScaledIntrinsic.has_value();
1604 (adaptor.getAbid() > 0 || getBlgpField > 0 || op.getCbsz() > 0)) {
1605 return op.emitOpError(
1606 "non-default abid, blgp, and cbsz aren't supported on MFMAs that can "
1607 "be scaled as those fields are used for type information");
1610 StringRef intrinsicName =
1611 isScaled ? std::get<0>(*maybeScaledIntrinsic) : *maybeIntrinsic;
1614 bool allowBf16 = [&]() {
1619 return intrinsicName.contains(
"16x16x32.bf16") ||
1620 intrinsicName.contains(
"32x32x16.bf16");
1622 OperationState loweredOp(loc, intrinsicName);
1623 loweredOp.addTypes(intrinsicOutType);
1625 rewriter, loc, adaptor.getSourceA(), allowBf16),
1627 rewriter, loc, adaptor.getSourceB(), allowBf16),
1628 adaptor.getDestC()});
1631 auto [_scaledName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;
1632 loweredOp.addOperands({zero, zero});
1633 loweredOp.addAttributes(
1635 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), aTypeCode)},
1637 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), bTypeCode)},
1638 {
"opselA", rewriter.getI32IntegerAttr(0)},
1639 {
"opselB", rewriter.getI32IntegerAttr(0)}});
1641 Attribute blgpAttr =
1643 ? Attribute(ROCDL::MFMANegModifierAttr::get(
1644 rewriter.getContext(),
1645 static_cast<ROCDL::MFMANegModifier
>(getBlgpField)))
1646 : Attribute(ROCDL::MFMAPermBAttr::
get(
1648 static_cast<ROCDL::MFMAPermB>(getBlgpField)));
1649 loweredOp.addAttributes(
1650 {{
"cbsz", rewriter.getI32IntegerAttr(op.getCbsz())},
1651 {
"abid", rewriter.getI32IntegerAttr(op.getAbid())},
1652 {
"blgp", blgpAttr}});
1654 Value lowered = rewriter.create(loweredOp)->getResult(0);
1655 if (outType != intrinsicOutType)
1656 lowered = LLVM::BitcastOp::create(rewriter, loc, outType, lowered);
1657 rewriter.replaceOp(op, lowered);
1663 ScaledMFMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1664 : ConvertOpToLLVMPattern(converter), chipset(chipset) {}
1669 matchAndRewrite(ScaledMFMAOp op, ScaledMFMAOpAdaptor adaptor,
1670 ConversionPatternRewriter &rewriter)
const override {
1671 Location loc = op.getLoc();
1672 Type intrinsicOutType = typeConverter->convertType(op.getDestD().getType());
1674 if (chipset.majorVersion != 9 || chipset <
kGfx950)
1675 return op->emitOpError(
"scaled MFMA only supported on gfx908+");
1676 std::optional<ScaledMFMAIntrinsic> maybeScaledIntrinsic =
1678 if (!maybeScaledIntrinsic.has_value())
1679 return op.emitOpError(
1680 "no intrinsic matching scaled MFMA size on given chipset");
1682 auto [intrinsicName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;
1683 OperationState loweredOp(loc, intrinsicName);
1684 loweredOp.addTypes(intrinsicOutType);
1685 loweredOp.addOperands(
1688 adaptor.getDestC()});
1689 loweredOp.addOperands(
1694 loweredOp.addAttributes(
1696 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), aTypeCode)},
1698 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), bTypeCode)},
1699 {
"opselA", rewriter.getI32IntegerAttr(adaptor.getScalesIdxA())},
1700 {
"opselB", rewriter.getI32IntegerAttr(adaptor.getScalesIdxB())}});
1702 Value lowered = rewriter.create(loweredOp)->getResult(0);
1703 rewriter.replaceOp(op, lowered);
1709 SparseMFMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1710 : ConvertOpToLLVMPattern<SparseMFMAOp>(converter), chipset(chipset) {}
1715 matchAndRewrite(SparseMFMAOp op, SparseMFMAOpAdaptor adaptor,
1716 ConversionPatternRewriter &rewriter)
const override {
1717 Location loc = op.getLoc();
1719 typeConverter->convertType<VectorType>(op.getDestC().
getType());
1721 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1724 if (chipset.majorVersion != 9 || chipset <
kGfx942)
1725 return op->emitOpError(
"sparse MFMA (smfmac) only supported on gfx942+");
1728 if (!maybeIntrinsic.has_value())
1729 return op.emitOpError(
1730 "no intrinsic matching sparse MFMA on the given chipset");
1733 ROCDL::smfmac_f32_16x16x32_bf16::getOperationName() ||
1735 ROCDL::smfmac_f32_32x32x16_bf16::getOperationName());
1736 bool isGfx950 = (chipset >=
kGfx950) && !isGfx942BF16;
1742 Value c = adaptor.getDestC();
1746 Value sparseIdx = adaptor.getSparseIdx();
1747 Type i32Type = rewriter.getI32Type();
1748 if (sparseIdx.
getType() != i32Type)
1749 sparseIdx = LLVM::BitcastOp::create(rewriter, loc, i32Type, sparseIdx);
1751 OperationState loweredOp(loc, maybeIntrinsic.value());
1752 loweredOp.addTypes(outType);
1753 loweredOp.addOperands({a,
b, c, sparseIdx});
1754 loweredOp.addAttributes(
1755 {{
"cbsz", rewriter.getI32IntegerAttr(op.getCbsz())},
1756 {
"abid", rewriter.getI32IntegerAttr(op.getAbid())}});
1757 Value lowered = rewriter.create(loweredOp)->getResult(0);
1758 rewriter.replaceOp(op, lowered);
1764 WMMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1765 : ConvertOpToLLVMPattern<WMMAOp>(converter), chipset(chipset) {}
1770 matchAndRewrite(WMMAOp op, WMMAOpAdaptor adaptor,
1771 ConversionPatternRewriter &rewriter)
const override {
1772 Location loc = op.getLoc();
1774 typeConverter->convertType<VectorType>(op.getDestD().
getType());
1776 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1778 if (chipset.majorVersion != 11 && chipset.majorVersion != 12)
1779 return op->emitOpError(
"WMMA only supported on gfx11 and gfx12");
1781 bool isGFX1250 = chipset >=
kGfx1250;
1786 auto aType = cast<VectorType>(adaptor.getSourceA().getType());
1787 auto bType = cast<VectorType>(adaptor.getSourceB().getType());
1788 auto destCType = cast<VectorType>(adaptor.getDestC().getType());
1789 bool castAToI16 = aType.getElementType().isBF16() && !isGFX1250;
1790 bool castBToI16 = bType.getElementType().isBF16() && !isGFX1250;
1791 bool castDestCToI16 = destCType.getElementType().isBF16() && !isGFX1250;
1792 bool castOutToI16 = outType.getElementType().
isBF16() && !isGFX1250;
1793 VectorType rawOutType = outType;
1795 rawOutType = outType.clone(rewriter.getI16Type());
1796 Value a = adaptor.getSourceA();
1798 a = LLVM::BitcastOp::create(rewriter, loc,
1799 aType.clone(rewriter.getI16Type()), a);
1800 Value
b = adaptor.getSourceB();
1802 b = LLVM::BitcastOp::create(rewriter, loc,
1803 bType.clone(rewriter.getI16Type()),
b);
1804 Value destC = adaptor.getDestC();
1806 destC = LLVM::BitcastOp::create(
1807 rewriter, loc, destCType.clone(rewriter.getI16Type()), destC);
1811 if (!maybeIntrinsic.has_value())
1812 return op.emitOpError(
"no intrinsic matching WMMA on the given chipset");
1814 if (chipset.majorVersion >= 12 && op.getSubwordOffset() != 0)
1815 return op.emitOpError(
"subwordOffset not supported on gfx12+");
1817 SmallVector<Value, 4> operands;
1818 SmallVector<NamedAttribute, 4> attrs;
1820 op.getSourceA(), operands, attrs,
"signA");
1822 op.getSourceB(), operands, attrs,
"signB");
1824 op.getSubwordOffset(), op.getClamp(), operands,
1827 OperationState loweredOp(loc, *maybeIntrinsic);
1828 loweredOp.addTypes(rawOutType);
1829 loweredOp.addOperands(operands);
1830 loweredOp.addAttributes(attrs);
1831 Operation *lowered = rewriter.create(loweredOp);
1833 Operation *maybeCastBack = lowered;
1834 if (rawOutType != outType)
1835 maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,
1837 rewriter.replaceOp(op, maybeCastBack->
getResults());
1843enum class DotFamily {
1852static std::optional<std::pair<StringRef, DotFamily>>
1853dotOpToIntrinsic(DotOp op,
Chipset chipset) {
1854 Type aElem = cast<VectorType>(op.getSourceA().getType()).getElementType();
1855 Type bElem = cast<VectorType>(op.getSourceB().getType()).getElementType();
1856 Type dest = op.getDestC().getType();
1857 bool uA = op.getUnsignedA();
1858 bool uB = op.getUnsignedB();
1863 return {{ROCDL::fdot2::getOperationName(), DotFamily::Clamp}};
1865 return {{ROCDL::fdot2_f16_f16::getOperationName(), DotFamily::NoClamp}};
1866 return std::nullopt;
1872 return {{ROCDL::fdot2_f32_bf16::getOperationName(), DotFamily::Clamp}};
1874 return {{ROCDL::fdot2_bf16_bf16::getOperationName(), DotFamily::NoClamp}};
1875 return std::nullopt;
1879 if (isa<IntegerType>(aElem) && isa<IntegerType>(bElem) &&
1881 bool mixedSign = (uA != uB);
1886 return std::nullopt;
1888 switch (elemWidth) {
1890 name = ROCDL::sudot4::getOperationName();
1893 name = ROCDL::sudot8::getOperationName();
1896 return std::nullopt;
1898 return {{name, DotFamily::Sudot}};
1902 bool supported =
false;
1903 switch (elemWidth) {
1906 name = uA ? ROCDL::udot2::getOperationName()
1907 :
ROCDL::sdot2::getOperationName();
1912 name = uA ? ROCDL::udot4::getOperationName()
1913 :
ROCDL::sdot4::getOperationName();
1918 name = uA ? ROCDL::udot8::getOperationName()
1919 :
ROCDL::sdot8::getOperationName();
1922 return std::nullopt;
1925 return std::nullopt;
1926 return {{name, DotFamily::Clamp}};
1930 bool aIsFp8 = isa<Float8E4M3FNType>(aElem);
1931 bool aIsBf8 = isa<Float8E5M2Type>(aElem);
1932 bool bIsFp8 = isa<Float8E4M3FNType>(bElem);
1933 bool bIsBf8 = isa<Float8E5M2Type>(bElem);
1934 if ((aIsFp8 || aIsBf8) && (bIsFp8 || bIsBf8) && dest.
isF32()) {
1936 return std::nullopt;
1938 if (aIsFp8 && bIsFp8)
1939 name = ROCDL::dot4_f32_fp8_fp8::getOperationName();
1940 else if (aIsFp8 && bIsBf8)
1941 name = ROCDL::dot4_f32_fp8_bf8::getOperationName();
1942 else if (aIsBf8 && bIsFp8)
1943 name = ROCDL::dot4_f32_bf8_fp8::getOperationName();
1945 name = ROCDL::dot4_f32_bf8_bf8::getOperationName();
1946 return {{name, DotFamily::NoClamp}};
1949 return std::nullopt;
1953 DotOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
1954 : ConvertOpToLLVMPattern<DotOp>(converter), chipset(chipset) {}
1959 matchAndRewrite(DotOp op, DotOpAdaptor adaptor,
1960 ConversionPatternRewriter &rewriter)
const override {
1961 Location loc = op.getLoc();
1963 std::optional<std::pair<StringRef, DotFamily>> maybeIntrinsic =
1964 dotOpToIntrinsic(op, chipset);
1965 if (!maybeIntrinsic)
1966 return op.emitOpError(
"no intrinsic matching dot on the given chipset: ")
1967 << op.getSourceA().getType() <<
" * " << op.getSourceB().getType()
1968 <<
" + " << op.getDestC().getType();
1970 auto [intrinsicName, family] = maybeIntrinsic.value();
1974 Value c = adaptor.getDestC();
1976 SmallVector<NamedAttribute, 3> attrs;
1977 if (family == DotFamily::Sudot) {
1978 attrs.push_back(rewriter.getNamedAttr(
1979 "signA", rewriter.getBoolAttr(!op.getUnsignedA())));
1980 attrs.push_back(rewriter.getNamedAttr(
1981 "signB", rewriter.getBoolAttr(!op.getUnsignedB())));
1984 if (family != DotFamily::NoClamp && op.getClamp())
1986 rewriter.getNamedAttr(
"clamp", rewriter.getBoolAttr(
true)));
1988 Type resultType = typeConverter->convertType(op.getDestD().getType());
1990 OperationState loweredOp(loc, intrinsicName);
1991 loweredOp.addTypes(resultType);
1992 loweredOp.addOperands({a,
b, c});
1993 loweredOp.addAttributes(attrs);
1994 Operation *lowered = rewriter.create(loweredOp);
1995 rewriter.replaceOp(op, lowered->
getResults());
2001 SparseWMMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2002 : ConvertOpToLLVMPattern<SparseWMMAOp>(converter), chipset(chipset) {}
2007 matchAndRewrite(SparseWMMAOp op, SparseWMMAOpAdaptor adaptor,
2008 ConversionPatternRewriter &rewriter)
const override {
2009 Location loc = op.getLoc();
2011 typeConverter->convertType<VectorType>(op.getDestD().
getType());
2013 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2015 std::optional<SparseWMMAOpInfo> maybeIntrinsic =
2018 if (!maybeIntrinsic.has_value())
2019 return op.emitOpError(
2020 "no intrinsic matching Sparse WMMA on the given chipset");
2021 SparseWMMAOpInfo intrinsic = maybeIntrinsic.value();
2023 SmallVector<NamedAttribute> attrs;
2025 if ((op.getUnsignedA() || op.getUnsignedB()) && !intrinsic.
useSign)
2026 return op->emitOpError(
"intrinsic doesn't support unsign");
2028 if (
auto attr = op.getUnsignedAAttr())
2029 attrs.push_back({
"signA", attr});
2030 if (
auto attr = op.getUnsignedBAttr())
2031 attrs.push_back({
"signB", attr});
2034 if ((op.getReuseA() || op.getReuseB()) && !intrinsic.
useReuse)
2035 return op->emitOpError(
"intrinsic doesn't support reuse");
2037 if (
auto attr = op.getReuseAAttr())
2038 attrs.push_back({
"reuseA", attr});
2039 if (
auto attr = op.getReuseBAttr())
2040 attrs.push_back({
"reuseB", attr});
2043 if (op.getClamp() && !intrinsic.
useClamp)
2044 return op->emitOpError(
"intrinsic doesn't support clamp");
2045 if (intrinsic.
useClamp && op.getClampAttr())
2046 attrs.push_back({
"clamp", op.getClampAttr()});
2048 const bool isGFX1250orHigher =
2049 chipset.majorVersion == 12 && chipset.minorVersion >= 5;
2054 Value c = adaptor.getDestC();
2055 VectorType rawOutType = outType;
2056 if (!isGFX1250orHigher) {
2058 rawOutType = cast<VectorType>(c.
getType());
2062 Value sparseIdx = LLVM::BitcastOp::create(
2063 rewriter, loc, rewriter.getI32Type(), adaptor.getSparseIdx());
2065 OperationState loweredOp(loc, intrinsic.
name);
2066 loweredOp.addTypes(rawOutType);
2067 loweredOp.addOperands({a,
b, c, sparseIdx});
2068 loweredOp.addAttributes(attrs);
2069 Operation *lowered = rewriter.create(loweredOp);
2071 Operation *maybeCastBack = lowered;
2072 if (rawOutType != outType)
2073 maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,
2075 rewriter.replaceOp(op, maybeCastBack->
getResults());
2082 ScaledWMMAOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2083 : ConvertOpToLLVMPattern<ScaledWMMAOp>(converter), chipset(chipset) {}
2088 matchAndRewrite(ScaledWMMAOp op, ScaledWMMAOpAdaptor adaptor,
2089 ConversionPatternRewriter &rewriter)
const override {
2090 Location loc = op.getLoc();
2092 typeConverter->convertType<VectorType>(op.getDestD().
getType());
2094 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2097 return op->emitOpError(
"WMMA scale only supported on gfx1250+");
2099 int64_t m = op.getM();
2100 int64_t n = op.getN();
2101 int64_t k = op.getK();
2106 std::optional<ROCDL::MatrixFormat> aFmtCode =
2108 std::optional<ROCDL::MatrixFormat> bFmtCode =
2111 if (!aFmtCode || !bFmtCode)
2112 return op.emitOpError(
"unsupported element types for scaled_wmma");
2115 auto scaleAVecType = cast<VectorType>(op.getScaleA().getType());
2116 auto scaleBVecType = cast<VectorType>(op.getScaleB().getType());
2118 if (scaleAVecType.getNumElements() != scaleBVecType.getNumElements())
2119 return op.emitOpError(
"scaleA and scaleB must have equal vector length");
2122 Type scaleAElemType = scaleAVecType.getElementType();
2123 Type scaleBElemType = scaleBVecType.getElementType();
2125 std::optional<ROCDL::WMMAMatrixScaleFormat> scaleAFmt =
2127 std::optional<ROCDL::WMMAMatrixScaleFormat> scaleBFmt =
2130 if (!scaleAFmt || !scaleBFmt)
2131 return op.emitOpError(
"unsupported scale element types");
2134 bool isScale16 = (scaleAVecType.getNumElements() == 8);
2135 std::optional<StringRef> intrinsicName =
2138 return op.emitOpError(
"unsupported scaled_wmma dimensions: ")
2139 << m <<
"x" << n <<
"x" << k;
2141 SmallVector<NamedAttribute, 8> attrs;
2144 bool is32x16 = (m == 32 && n == 16 && k == 128);
2146 attrs.emplace_back(
"fmtA", ROCDL::MatrixFormatAttr::get(
2147 rewriter.getContext(), *aFmtCode));
2148 attrs.emplace_back(
"fmtB", ROCDL::MatrixFormatAttr::get(
2149 rewriter.getContext(), *bFmtCode));
2154 "modC", ROCDL::WMMACModifierAttr::get(rewriter.getContext(),
2155 ROCDL::WMMACModifier::none));
2159 attrs.emplace_back(
"scaleAType", ROCDL::WMMAMatrixScaleAttr::get(
2160 rewriter.getContext(),
2161 static_cast<ROCDL::WMMAMatrixScale
>(
2162 op.getAFirstScaleLane() / 16)));
2163 attrs.emplace_back(
"fmtScaleA", ROCDL::WMMAMatrixScaleFormatAttr::get(
2164 rewriter.getContext(), *scaleAFmt));
2165 attrs.emplace_back(
"scaleBType", ROCDL::WMMAMatrixScaleAttr::get(
2166 rewriter.getContext(),
2167 static_cast<ROCDL::WMMAMatrixScale
>(
2168 op.getBFirstScaleLane() / 16)));
2169 attrs.emplace_back(
"fmtScaleB", ROCDL::WMMAMatrixScaleFormatAttr::get(
2170 rewriter.getContext(), *scaleBFmt));
2173 attrs.emplace_back(
"reuseA", rewriter.getBoolAttr(
false));
2174 attrs.emplace_back(
"reuseB", rewriter.getBoolAttr(
false));
2187 OperationState loweredOp(loc, *intrinsicName);
2188 loweredOp.addTypes(outType);
2189 loweredOp.addOperands(
2190 {sourceA, sourceB, adaptor.getDestC(), packedScaleA, packedScaleB});
2191 loweredOp.addAttributes(attrs);
2193 Operation *lowered = rewriter.create(loweredOp);
2194 rewriter.replaceOp(op, lowered->
getResults());
2200struct TransposeLoadOpLowering
2202 TransposeLoadOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2203 : ConvertOpToLLVMPattern<TransposeLoadOp>(converter), chipset(chipset) {}
2208 matchAndRewrite(TransposeLoadOp op, TransposeLoadOpAdaptor adaptor,
2209 ConversionPatternRewriter &rewriter)
const override {
2211 return op.emitOpError(
2212 "transpose_load is only supported on gfx950 and gfx1250+");
2214 Location loc = op.getLoc();
2215 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2219 size_t srcElementSize =
2220 srcMemRefType.getElementType().getIntOrFloatBitWidth();
2221 if (srcElementSize < 8)
2222 return op.emitOpError(
"Expect source memref to have at least 8 bits "
2223 "element size, got ")
2226 auto resultType = cast<VectorType>(op.getResult().getType());
2229 (adaptor.getSrcIndices()));
2231 size_t numElements = resultType.getNumElements();
2232 size_t elementTypeSize =
2235 Type llvmResultType = typeConverter->convertType(resultType);
2238 Type rocdlResultType =
2239 elementTypeSize < 16
2240 ? VectorType::get((numElements * elementTypeSize) / 32,
2241 rewriter.getIntegerType(32))
2244 auto emitNumElementsError = [&](
size_t expected, StringRef chipsetName) {
2245 return op.emitOpError()
2246 << elementTypeSize <<
"-bit transpose_load requires " << expected
2247 <<
" elements on " << chipsetName;
2252 switch (elementTypeSize) {
2254 if (numElements != 16)
2255 return emitNumElementsError(16,
"gfx1250+");
2257 ROCDL::DsLoadTr4_B64::create(rewriter, loc, rocdlResultType, srcPtr,
2264 if (numElements != 16)
2265 return emitNumElementsError(16,
"gfx1250+");
2267 ROCDL::DsLoadTr6_B96::create(rewriter, loc, rocdlResultType, srcPtr,
2274 if (numElements != 8)
2275 return emitNumElementsError(8,
"gfx1250+");
2277 ROCDL::DsLoadTr8_B64::create(rewriter, loc, rocdlResultType, srcPtr,
2284 if (numElements != 8)
2285 return emitNumElementsError(8,
"gfx1250+");
2286 intrinsic = ROCDL::DsLoadTr16_B128::create(
2287 rewriter, loc, rocdlResultType, srcPtr,
2293 return op.emitOpError(
"Unsupported element size for transpose load");
2296 switch (elementTypeSize) {
2298 if (numElements != 16)
2299 return emitNumElementsError(16,
"gfx950");
2300 intrinsic = ROCDL::ds_read_tr4_b64::create(
2301 rewriter, loc, rocdlResultType, srcPtr,
2307 if (numElements != 16)
2308 return emitNumElementsError(16,
"gfx950");
2309 intrinsic = ROCDL::ds_read_tr6_b96::create(
2310 rewriter, loc, rocdlResultType, srcPtr,
2316 if (numElements != 8)
2317 return emitNumElementsError(8,
"gfx950");
2318 intrinsic = ROCDL::ds_read_tr8_b64::create(
2319 rewriter, loc, rocdlResultType, srcPtr,
2325 if (numElements != 4)
2326 return emitNumElementsError(4,
"gfx950");
2327 intrinsic = ROCDL::ds_read_tr16_b64::create(
2328 rewriter, loc, rocdlResultType, srcPtr,
2334 return op.emitOpError(
"Unsupported element size for transpose load");
2338 assert(intrinsic &&
"expected ROCDL transpose load intrinsic");
2339 if (intrinsic.
getType() == llvmResultType) {
2340 rewriter.replaceOp(op, intrinsic);
2343 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, intrinsic);
2348struct GlobalTransposeLoadOpLowering
2350 GlobalTransposeLoadOpLowering(
const LLVMTypeConverter &converter,
2352 : ConvertOpToLLVMPattern<GlobalTransposeLoadOp>(converter),
2358 matchAndRewrite(GlobalTransposeLoadOp op,
2359 GlobalTransposeLoadOpAdaptor adaptor,
2360 ConversionPatternRewriter &rewriter)
const override {
2362 return op.emitOpError(
2363 "global_transpose_load is only supported on gfx1200+");
2365 Location loc = op.getLoc();
2366 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2367 auto resultType = cast<VectorType>(op.getResult().getType());
2370 rewriter, loc, srcMemRefType, adaptor.getSrc(), adaptor.getSrcIndices(),
2371 LLVM::GEPNoWrapFlags::inbounds | LLVM::GEPNoWrapFlags::nuw);
2373 size_t numElements = resultType.getNumElements();
2374 size_t elementTypeSize =
2379 Type rocdlResultType =
2380 elementTypeSize < 16
2381 ? VectorType::get((numElements * elementTypeSize) / 32,
2382 rewriter.getIntegerType(32))
2383 : typeConverter->convertType(resultType);
2384 Type llvmResultType = typeConverter->convertType(resultType);
2386 switch (elementTypeSize) {
2388 assert(numElements == 16);
2390 return op.emitOpError(
"4-bit global_transpose_load requires gfx1250+");
2391 auto rocdlOp = ROCDL::GlobalLoadTr4_B64::create(
2394 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2398 assert(numElements == 16);
2400 return op.emitOpError(
"6-bit global_transpose_load requires gfx1250+");
2401 auto rocdlOp = ROCDL::GlobalLoadTr6_B96::create(
2404 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2408 assert(numElements == 8);
2409 auto rocdlOp = ROCDL::GlobalLoadTr8_B64::create(
2412 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2416 assert(numElements == 8);
2417 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadTr8_B128>(
2422 return op.emitOpError(
2423 "unsupported element size for global transpose load");
2430 GatherToLDSOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2431 : ConvertOpToLLVMPattern<GatherToLDSOp>(converter), chipset(chipset) {}
2436 matchAndRewrite(GatherToLDSOp op, GatherToLDSOpAdaptor adaptor,
2437 ConversionPatternRewriter &rewriter)
const override {
2438 if (chipset.majorVersion < 9 || chipset.majorVersion > 10)
2439 return op.emitOpError(
"pre-gfx9 and post-gfx10 not supported");
2441 Location loc = op.getLoc();
2443 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2444 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2449 Type transferType = op.getTransferType();
2450 int loadWidth = [&]() ->
int {
2451 if (
auto transferVectorType = dyn_cast<VectorType>(transferType)) {
2452 return (transferVectorType.getNumElements() *
2453 transferVectorType.getElementTypeBitWidth()) /
2460 if (!llvm::is_contained({1, 2, 4, 12, 16}, loadWidth))
2461 return op.emitOpError(
"chipset unsupported element size");
2463 if (chipset !=
kGfx950 && llvm::is_contained({12, 16}, loadWidth))
2464 return op.emitOpError(
"Gather to LDS instructions with 12-byte and "
2465 "16-byte load widths are only supported on gfx950");
2469 (adaptor.getSrcIndices()));
2472 (adaptor.getDstIndices()));
2474 if (op.getAsync()) {
2475 rewriter.replaceOpWithNewOp<ROCDL::LoadAsyncToLDSOp>(
2476 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2477 rewriter.getI32IntegerAttr(0),
2481 rewriter.replaceOpWithNewOp<ROCDL::LoadToLDSOp>(
2482 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2483 rewriter.getI32IntegerAttr(0),
2492struct GlobalLoadAsyncToLDSOpLowering
2494 GlobalLoadAsyncToLDSOpLowering(
const LLVMTypeConverter &converter,
2496 : ConvertOpToLLVMPattern<GlobalLoadAsyncToLDSOp>(converter),
2502 matchAndRewrite(GlobalLoadAsyncToLDSOp op,
2503 GlobalLoadAsyncToLDSOpAdaptor adaptor,
2504 ConversionPatternRewriter &rewriter)
const override {
2506 return op.emitOpError(
2507 "global_load_async_to_lds is only supported on gfx1250+");
2509 Location loc = op.getLoc();
2510 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2511 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2513 Type transferType = op.getTransferType();
2515 isa<VectorType>(transferType)
2516 ? cast<VectorType>(transferType).getNumElements() *
2517 cast<VectorType>(transferType).getElementTypeBitWidth()
2522 adaptor.getSrcIndices());
2525 adaptor.getDstIndices());
2528 Value mask = adaptor.getMask();
2529 int64_t nullptrVal =
2530 llvm::AMDGPU::getNullPointerValue(llvm::AMDGPUAS::LOCAL_ADDRESS);
2534 LLVM::IntToPtrOp::create(rewriter, loc, dstPtr.
getType(), nullInt);
2535 dstPtr = LLVM::SelectOp::create(rewriter, loc, mask, dstPtr, nullPtr);
2538 auto offset = rewriter.getI32IntegerAttr(0);
2539 Attribute aux = rewriter.getI32IntegerAttr(0);
2541 switch (transferBits) {
2543 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB8Op>(
2548 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB32Op>(
2553 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB64Op>(
2558 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB128Op>(
2563 return op.emitOpError(
"unsupported transfer width");
2570struct ExtPackedFp8OpLowering final
2572 ExtPackedFp8OpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2573 : ConvertOpToLLVMPattern<amdgpu::ExtPackedFp8Op>(converter),
2578 matchAndRewrite(ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2579 ConversionPatternRewriter &rewriter)
const override;
2582struct ScaledExtPackedMatrixOpLowering final
2584 ScaledExtPackedMatrixOpLowering(
const LLVMTypeConverter &converter,
2586 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedMatrixOp>(converter),
2591 matchAndRewrite(ScaledExtPackedMatrixOp op,
2592 ScaledExtPackedMatrixOpAdaptor adaptor,
2593 ConversionPatternRewriter &rewriter)
const override;
2596struct PackedTrunc2xFp8OpLowering final
2598 PackedTrunc2xFp8OpLowering(
const LLVMTypeConverter &converter,
2600 : ConvertOpToLLVMPattern<amdgpu::PackedTrunc2xFp8Op>(converter),
2605 matchAndRewrite(PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
2606 ConversionPatternRewriter &rewriter)
const override;
2609struct PackedStochRoundFp8OpLowering final
2611 PackedStochRoundFp8OpLowering(
const LLVMTypeConverter &converter,
2613 : ConvertOpToLLVMPattern<amdgpu::PackedStochRoundFp8Op>(converter),
2618 matchAndRewrite(PackedStochRoundFp8Op op,
2619 PackedStochRoundFp8OpAdaptor adaptor,
2620 ConversionPatternRewriter &rewriter)
const override;
2623struct ScaledExtPackedOpLowering final
2625 ScaledExtPackedOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2626 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedOp>(converter),
2631 matchAndRewrite(ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2632 ConversionPatternRewriter &rewriter)
const override;
2635struct PackedScaledTruncOpLowering final
2637 PackedScaledTruncOpLowering(
const LLVMTypeConverter &converter,
2639 : ConvertOpToLLVMPattern<amdgpu::PackedScaledTruncOp>(converter),
2644 matchAndRewrite(PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2645 ConversionPatternRewriter &rewriter)
const override;
2650LogicalResult ExtPackedFp8OpLowering::matchAndRewrite(
2651 ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2652 ConversionPatternRewriter &rewriter)
const {
2653 Location loc = op.getLoc();
2655 return rewriter.notifyMatchFailure(
2656 loc,
"Fp8 conversion instructions are not available on target "
2657 "architecture and their emulation is not implemented");
2659 getTypeConverter()->convertType(VectorType::get(4, rewriter.getI8Type()));
2660 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2661 Type f32 = getTypeConverter()->convertType(op.getResult().getType());
2663 Value source = adaptor.getSource();
2664 auto sourceVecType = dyn_cast<VectorType>(op.getSource().getType());
2665 auto resultVecType = dyn_cast<VectorType>(op.getResult().getType());
2668 if (!sourceVecType || sourceVecType.getNumElements() < 4) {
2669 Value longVec = LLVM::UndefOp::create(rewriter, loc, v4i8);
2670 if (!sourceVecType) {
2671 longVec = LLVM::InsertElementOp::create(
2674 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2676 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2678 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2683 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2684 if (resultVecType) {
2686 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Bf8Op>(op, f32, i32Source,
2689 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Fp8Op>(op, f32, i32Source,
2694 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Bf8Op>(op, f32, i32Source,
2697 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Fp8Op>(op, f32, i32Source,
2704int32_t getScaleSel(int32_t blockSize,
unsigned bitWidth, int32_t scaleWaveHalf,
2705 int32_t firstScaleByte) {
2711 assert(llvm::is_contained({16, 32}, blockSize));
2712 assert(llvm::is_contained({4u, 6u, 8u}, bitWidth));
2714 const bool isFp8 = bitWidth == 8;
2715 const bool isBlock16 = blockSize == 16;
2718 int32_t bit0 = isBlock16;
2719 assert(llvm::is_contained({0, 1, 2}, firstScaleByte));
2720 int32_t bit1 = (firstScaleByte == 2) << 1;
2721 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2722 int32_t bit2 = scaleWaveHalf << 2;
2723 return bit2 | bit1 | bit0;
2726 int32_t bit0 = isBlock16;
2728 assert(llvm::is_contained({0, 1, 2, 3}, firstScaleByte));
2729 int32_t bits2and1 = firstScaleByte << 1;
2730 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2731 int32_t bit3 = scaleWaveHalf << 3;
2732 int32_t bits = bit3 | bits2and1 | bit0;
2734 assert(!llvm::is_contained(
2735 {0b0011, 0b0101, 0b0111, 0b1000, 0b1001, 0b1011, 0b1111}, bits));
2739static std::optional<StringRef>
2740scaledExtPacked816ToIntrinsic(Type srcElemType, Type destElemType) {
2741 using fp4 = Float4E2M1FNType;
2742 using fp8 = Float8E4M3FNType;
2743 using bf8 = Float8E5M2Type;
2744 using fp6 = Float6E2M3FNType;
2745 using bf6 = Float6E3M2FNType;
2746 if (isa<fp4>(srcElemType)) {
2747 if (destElemType.
isF16())
2748 return ROCDL::CvtPkScalePk8F16Fp4Op::getOperationName();
2749 if (destElemType.
isBF16())
2750 return ROCDL::CvtPkScalePk8Bf16Fp4Op::getOperationName();
2751 if (destElemType.
isF32())
2752 return ROCDL::CvtPkScalePk8F32Fp4Op::getOperationName();
2753 return std::nullopt;
2755 if (isa<fp8>(srcElemType)) {
2756 if (destElemType.
isF16())
2757 return ROCDL::CvtPkScalePk8F16Fp8Op::getOperationName();
2758 if (destElemType.
isBF16())
2759 return ROCDL::CvtPkScalePk8Bf16Fp8Op::getOperationName();
2760 if (destElemType.
isF32())
2761 return ROCDL::CvtPkScalePk8F32Fp8Op::getOperationName();
2762 return std::nullopt;
2764 if (isa<bf8>(srcElemType)) {
2765 if (destElemType.
isF16())
2766 return ROCDL::CvtPkScalePk8F16Bf8Op::getOperationName();
2767 if (destElemType.
isBF16())
2768 return ROCDL::CvtPkScalePk8Bf16Bf8Op::getOperationName();
2769 if (destElemType.
isF32())
2770 return ROCDL::CvtPkScalePk8F32Bf8Op::getOperationName();
2771 return std::nullopt;
2773 if (isa<fp6>(srcElemType)) {
2774 if (destElemType.
isF16())
2775 return ROCDL::CvtPkScalePk16F16Fp6Op::getOperationName();
2776 if (destElemType.
isBF16())
2777 return ROCDL::CvtPkScalePk16Bf16Fp6Op::getOperationName();
2778 if (destElemType.
isF32())
2779 return ROCDL::CvtPkScalePk16F32Fp6Op::getOperationName();
2780 return std::nullopt;
2782 if (isa<bf6>(srcElemType)) {
2783 if (destElemType.
isF16())
2784 return ROCDL::CvtPkScalePk16F16Bf6Op::getOperationName();
2785 if (destElemType.
isBF16())
2786 return ROCDL::CvtPkScalePk16Bf16Bf6Op::getOperationName();
2787 if (destElemType.
isF32())
2788 return ROCDL::CvtPkScalePk16F32Bf6Op::getOperationName();
2789 return std::nullopt;
2791 llvm_unreachable(
"invalid combination of element types for packed conversion "
2795LogicalResult ScaledExtPackedMatrixOpLowering::matchAndRewrite(
2796 ScaledExtPackedMatrixOp op, ScaledExtPackedMatrixOpAdaptor adaptor,
2797 ConversionPatternRewriter &rewriter)
const {
2798 using fp4 = Float4E2M1FNType;
2799 using fp8 = Float8E4M3FNType;
2800 using bf8 = Float8E5M2Type;
2801 using fp6 = Float6E2M3FNType;
2802 using bf6 = Float6E3M2FNType;
2803 Location loc = op.getLoc();
2805 return rewriter.notifyMatchFailure(
2807 "Scaled fp packed conversion instructions are not available on target "
2808 "architecture and their emulation is not implemented");
2812 int32_t scaleWaveHalf = op.getFirstScaleLane() / 16;
2813 int32_t firstScaleByte = op.getFirstScaleByte();
2814 int32_t blockSize = op.getBlockSize();
2815 auto sourceType = cast<VectorType>(op.getSource().getType());
2816 auto srcElemType = cast<FloatType>(sourceType.getElementType());
2817 unsigned bitWidth = srcElemType.getWidth();
2819 auto targetType = cast<VectorType>(op.getResult().getType());
2820 auto destElemType = cast<FloatType>(targetType.getElementType());
2822 IntegerType i32 = rewriter.getI32Type();
2823 Value source = adaptor.getSource();
2824 Type llvmResultType = typeConverter->convertType(op.getResult().getType());
2825 Type packedType =
nullptr;
2826 if (isa<fp4>(srcElemType)) {
2828 packedType = getTypeConverter()->convertType(packedType);
2829 }
else if (isa<fp8, bf8>(srcElemType)) {
2830 packedType = VectorType::get(2, i32);
2831 packedType = getTypeConverter()->convertType(packedType);
2832 }
else if (isa<fp6, bf6>(srcElemType)) {
2833 packedType = VectorType::get(3, i32);
2834 packedType = getTypeConverter()->convertType(packedType);
2836 llvm_unreachable(
"invalid element type for packed scaled ext");
2839 if (!packedType || !llvmResultType) {
2840 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2843 std::optional<StringRef> maybeIntrinsic =
2844 scaledExtPacked816ToIntrinsic(srcElemType, destElemType);
2845 if (!maybeIntrinsic.has_value())
2846 return op.emitOpError(
2847 "no intrinsic matching packed scaled conversion on the given chipset");
2850 getScaleSel(blockSize, bitWidth, scaleWaveHalf, firstScaleByte);
2852 LLVM::BitcastOp::create(rewriter, loc, i32, adaptor.getScale());
2853 Value castedSource =
2854 LLVM::BitcastOp::create(rewriter, loc, packedType, source);
2856 OperationState loweredOp(loc, *maybeIntrinsic);
2857 loweredOp.addTypes({llvmResultType});
2858 loweredOp.addOperands({castedSource, castedScale});
2860 SmallVector<NamedAttribute, 1> attrs;
2862 NamedAttribute(
"scaleSel", rewriter.getI32IntegerAttr(scaleSel)));
2864 loweredOp.addAttributes(attrs);
2865 Operation *lowered = rewriter.create(loweredOp);
2866 rewriter.replaceOp(op, lowered);
2871LogicalResult ScaledExtPackedOpLowering::matchAndRewrite(
2872 ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2873 ConversionPatternRewriter &rewriter)
const {
2874 Location loc = op.getLoc();
2876 return rewriter.notifyMatchFailure(
2877 loc,
"Scaled fp conversion instructions are not available on target "
2878 "architecture and their emulation is not implemented");
2879 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2881 Value source = adaptor.getSource();
2882 Value scale = adaptor.getScale();
2884 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2885 Type sourceElemType = sourceVecType.getElementType();
2886 VectorType destVecType = cast<VectorType>(op.getResult().getType());
2887 Type destElemType = destVecType.getElementType();
2889 VectorType packedVecType;
2890 if (isa<Float8E5M2Type, Float8E4M3FNType>(sourceElemType)) {
2891 VectorType v4i8 = VectorType::get(4, rewriter.getI8Type());
2892 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v4i8));
2893 }
else if (isa<Float4E2M1FNType>(sourceElemType)) {
2894 VectorType v8i4 = VectorType::get(8, rewriter.getI4Type());
2895 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v8i4));
2897 llvm_unreachable(
"invalid element type for scaled ext");
2901 if (sourceVecType.getNumElements() < packedVecType.getNumElements()) {
2902 Value longVec = LLVM::ZeroOp::create(rewriter, loc, packedVecType);
2903 if (!sourceVecType) {
2904 longVec = LLVM::InsertElementOp::create(
2907 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2909 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2911 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2916 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2918 if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isF32())
2919 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Bf8Op>(
2920 op, destVecType, i32Source, scale, op.getIndex());
2921 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isF16())
2922 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Bf8Op>(
2923 op, destVecType, i32Source, scale, op.getIndex());
2924 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isBF16())
2925 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Bf8Op>(
2926 op, destVecType, i32Source, scale, op.getIndex());
2927 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isF32())
2928 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp8Op>(
2929 op, destVecType, i32Source, scale, op.getIndex());
2930 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isF16())
2931 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp8Op>(
2932 op, destVecType, i32Source, scale, op.getIndex());
2933 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isBF16())
2934 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp8Op>(
2935 op, destVecType, i32Source, scale, op.getIndex());
2936 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isF32())
2937 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp4Op>(
2938 op, destVecType, i32Source, scale, op.getIndex());
2939 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isF16())
2940 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp4Op>(
2941 op, destVecType, i32Source, scale, op.getIndex());
2942 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isBF16())
2943 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp4Op>(
2944 op, destVecType, i32Source, scale, op.getIndex());
2951LogicalResult PackedScaledTruncOpLowering::matchAndRewrite(
2952 PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2953 ConversionPatternRewriter &rewriter)
const {
2954 Location loc = op.getLoc();
2956 return rewriter.notifyMatchFailure(
2957 loc,
"Scaled fp conversion instructions are not available on target "
2958 "architecture and their emulation is not implemented");
2959 Type v2i16 = getTypeConverter()->convertType(
2960 VectorType::get(2, rewriter.getI16Type()));
2961 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2963 Type resultType = op.getResult().getType();
2965 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2966 Type sourceElemType = sourceVecType.getElementType();
2968 Type intResultType = isa<Float4E2M1FNType>(resultElemType) ? i32 : v2i16;
2970 Value source = adaptor.getSource();
2971 Value scale = adaptor.getScale();
2972 Value existing = adaptor.getExisting();
2974 existing = LLVM::BitcastOp::create(rewriter, loc, intResultType, existing);
2976 existing = LLVM::ZeroOp::create(rewriter, loc, intResultType);
2978 if (sourceVecType.getNumElements() < 2) {
2980 Value elem0 = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2981 VectorType v2 = VectorType::get(2, sourceElemType);
2982 source = LLVM::ZeroOp::create(rewriter, loc, v2);
2983 source = LLVM::InsertElementOp::create(rewriter, loc, source, elem0, c0);
2986 Value sourceA, sourceB;
2987 if (sourceElemType.
isF32()) {
2990 sourceA = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2991 sourceB = LLVM::ExtractElementOp::create(rewriter, loc, source, c1);
2995 if (sourceElemType.
isF32() && isa<Float8E5M2Type>(resultElemType))
2996 result = ROCDL::CvtScaleF32PkBf8F32Op::create(rewriter, loc, intResultType,
2997 existing, sourceA, sourceB,
2998 scale, op.getIndex());
2999 else if (sourceElemType.
isF16() && isa<Float8E5M2Type>(resultElemType))
3000 result = ROCDL::CvtScaleF32PkBf8F16Op::create(
3001 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3002 else if (sourceElemType.
isBF16() && isa<Float8E5M2Type>(resultElemType))
3003 result = ROCDL::CvtScaleF32PkBf8Bf16Op::create(
3004 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3005 else if (sourceElemType.
isF32() && isa<Float8E4M3FNType>(resultElemType))
3006 result = ROCDL::CvtScaleF32PkFp8F32Op::create(rewriter, loc, intResultType,
3007 existing, sourceA, sourceB,
3008 scale, op.getIndex());
3009 else if (sourceElemType.
isF16() && isa<Float8E4M3FNType>(resultElemType))
3010 result = ROCDL::CvtScaleF32PkFp8F16Op::create(
3011 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3012 else if (sourceElemType.
isBF16() && isa<Float8E4M3FNType>(resultElemType))
3013 result = ROCDL::CvtScaleF32PkFp8Bf16Op::create(
3014 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3015 else if (sourceElemType.
isF32() && isa<Float4E2M1FNType>(resultElemType))
3016 result = ROCDL::CvtScaleF32PkFp4F32Op::create(rewriter, loc, intResultType,
3017 existing, sourceA, sourceB,
3018 scale, op.getIndex());
3019 else if (sourceElemType.
isF16() && isa<Float4E2M1FNType>(resultElemType))
3020 result = ROCDL::CvtScaleF32PkFp4F16Op::create(
3021 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3022 else if (sourceElemType.
isBF16() && isa<Float4E2M1FNType>(resultElemType))
3023 result = ROCDL::CvtScaleF32PkFp4Bf16Op::create(
3024 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3028 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3029 op, getTypeConverter()->convertType(resultType),
result);
3033LogicalResult PackedTrunc2xFp8OpLowering::matchAndRewrite(
3034 PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
3035 ConversionPatternRewriter &rewriter)
const {
3036 Location loc = op.getLoc();
3038 return rewriter.notifyMatchFailure(
3039 loc,
"Fp8 conversion instructions are not available on target "
3040 "architecture and their emulation is not implemented");
3041 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3043 Type resultType = op.getResult().getType();
3046 Value sourceA = adaptor.getSourceA();
3047 Value sourceB = adaptor.getSourceB();
3049 sourceB = LLVM::UndefOp::create(rewriter, loc, sourceA.
getType());
3050 Value existing = adaptor.getExisting();
3052 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3054 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3058 result = ROCDL::CvtPkBf8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3059 existing, op.getWordIndex());
3061 result = ROCDL::CvtPkFp8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3062 existing, op.getWordIndex());
3064 return op.emitOpError(
3065 "no truncation to result type available on given chipset");
3067 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3068 op, getTypeConverter()->convertType(resultType),
result);
3072LogicalResult PackedStochRoundFp8OpLowering::matchAndRewrite(
3073 PackedStochRoundFp8Op op, PackedStochRoundFp8OpAdaptor adaptor,
3074 ConversionPatternRewriter &rewriter)
const {
3075 Location loc = op.getLoc();
3077 return rewriter.notifyMatchFailure(
3078 loc,
"Fp8 conversion instructions are not available on target "
3079 "architecture and their emulation is not implemented");
3080 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3082 Type resultType = op.getResult().getType();
3085 Value source = adaptor.getSource();
3086 Value stoch = adaptor.getStochiasticParam();
3087 Value existing = adaptor.getExisting();
3089 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3091 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3095 result = ROCDL::CvtSrBf8F32Op::create(rewriter, loc, i32, source, stoch,
3096 existing, op.getStoreIndex());
3098 result = ROCDL::CvtSrFp8F32Op::create(rewriter, loc, i32, source, stoch,
3099 existing, op.getStoreIndex());
3101 return op.emitOpError(
3102 "no stochastic rounding to result type available on given chipset");
3104 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3105 op, getTypeConverter()->convertType(resultType),
result);
3111struct AMDGPUDPPLowering :
public ConvertOpToLLVMPattern<DPPOp> {
3112 AMDGPUDPPLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3113 : ConvertOpToLLVMPattern<DPPOp>(converter), chipset(chipset) {}
3117 matchAndRewrite(DPPOp DppOp, DPPOp::Adaptor adaptor,
3118 ConversionPatternRewriter &rewriter)
const override {
3121 Location loc = DppOp.getLoc();
3122 Value src = adaptor.getSrc();
3123 Value old = adaptor.getOld();
3126 Type llvmType =
nullptr;
3128 llvmType = rewriter.getI32Type();
3129 }
else if (isa<FloatType>(srcType)) {
3131 ? rewriter.getF32Type()
3132 : rewriter.getF64Type();
3133 }
else if (isa<IntegerType>(srcType)) {
3135 ? rewriter.getI32Type()
3136 : rewriter.getI64Type();
3138 auto llvmSrcIntType = typeConverter->convertType(
3142 auto convertOperand = [&](Value operand, Type operandType) {
3143 if (operandType.getIntOrFloatBitWidth() <= 16) {
3144 if (llvm::isa<FloatType>(operandType)) {
3146 LLVM::BitcastOp::create(rewriter, loc, llvmSrcIntType, operand);
3148 auto llvmVecType = typeConverter->convertType(mlir::VectorType::get(
3149 32 / operandType.getIntOrFloatBitWidth(), llvmSrcIntType));
3150 Value undefVec = LLVM::UndefOp::create(rewriter, loc, llvmVecType);
3152 LLVM::InsertElementOp::create(rewriter, loc, undefVec, operand,
3154 operand = LLVM::BitcastOp::create(rewriter, loc, llvmType, operand);
3159 src = convertOperand(src, srcType);
3160 old = convertOperand(old, oldType);
3163 enum DppCtrl :
unsigned {
3172 ROW_HALF_MIRROR = 0x141,
3177 auto kind = DppOp.getKind();
3178 auto permArgument = DppOp.getPermArgument();
3179 uint32_t DppCtrl = 0;
3183 case DPPPerm::quad_perm: {
3184 auto quadPermAttr = cast<ArrayAttr>(*permArgument);
3186 for (
auto elem : quadPermAttr.getAsRange<IntegerAttr>()) {
3187 uint32_t num = elem.getInt();
3188 DppCtrl |= num << (i * 2);
3193 case DPPPerm::row_shl: {
3194 auto intAttr = cast<IntegerAttr>(*permArgument);
3195 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHL0;
3198 case DPPPerm::row_shr: {
3199 auto intAttr = cast<IntegerAttr>(*permArgument);
3200 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHR0;
3203 case DPPPerm::row_ror: {
3204 auto intAttr = cast<IntegerAttr>(*permArgument);
3205 DppCtrl = intAttr.getInt() + DppCtrl::ROW_ROR0;
3208 case DPPPerm::wave_shl:
3209 DppCtrl = DppCtrl::WAVE_SHL1;
3211 case DPPPerm::wave_shr:
3212 DppCtrl = DppCtrl::WAVE_SHR1;
3214 case DPPPerm::wave_rol:
3215 DppCtrl = DppCtrl::WAVE_ROL1;
3217 case DPPPerm::wave_ror:
3218 DppCtrl = DppCtrl::WAVE_ROR1;
3220 case DPPPerm::row_mirror:
3221 DppCtrl = DppCtrl::ROW_MIRROR;
3223 case DPPPerm::row_half_mirror:
3224 DppCtrl = DppCtrl::ROW_HALF_MIRROR;
3226 case DPPPerm::row_bcast_15:
3227 DppCtrl = DppCtrl::BCAST15;
3229 case DPPPerm::row_bcast_31:
3230 DppCtrl = DppCtrl::BCAST31;
3236 auto rowMask = DppOp.getRowMask();
3237 auto bankMask = DppOp.getBankMask();
3238 bool boundCtrl = DppOp.getBoundCtrl();
3242 ROCDL::DPPUpdateOp::create(rewriter, loc, llvmType, old, src, DppCtrl,
3243 rowMask, bankMask, boundCtrl);
3245 Value
result = dppMovOp.getRes();
3247 result = LLVM::TruncOp::create(rewriter, loc, llvmSrcIntType,
result);
3248 if (!llvm::isa<IntegerType>(srcType)) {
3249 result = LLVM::BitcastOp::create(rewriter, loc, srcType,
result);
3260struct AMDGPUSwizzleBitModeLowering
3261 :
public ConvertOpToLLVMPattern<SwizzleBitModeOp> {
3265 matchAndRewrite(SwizzleBitModeOp op, OpAdaptor adaptor,
3266 ConversionPatternRewriter &rewriter)
const override {
3267 Location loc = op.getLoc();
3268 Type i32 = rewriter.getI32Type();
3269 Value src = adaptor.getSrc();
3270 SmallVector<Value> decomposed;
3272 return rewriter.notifyMatchFailure(op,
3273 "failed to decompose value to i32");
3274 unsigned andMask = op.getAndMask();
3275 unsigned orMask = op.getOrMask();
3276 unsigned xorMask = op.getXorMask();
3280 unsigned mask = andMask | (orMask << 5) | (xorMask << 10);
3282 SmallVector<Value> swizzled;
3283 for (Value v : decomposed) {
3285 ROCDL::DsSwizzleOp::create(rewriter, loc, v.getType(), v, maskValue);
3286 swizzled.emplace_back(res);
3290 rewriter.replaceOp(op,
result);
3295struct AMDGPUPermlaneLowering :
public ConvertOpToLLVMPattern<PermlaneSwapOp> {
3298 AMDGPUPermlaneLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3299 : ConvertOpToLLVMPattern<PermlaneSwapOp>(converter), chipset(chipset) {}
3303 matchAndRewrite(PermlaneSwapOp op, OpAdaptor adaptor,
3304 ConversionPatternRewriter &rewriter)
const override {
3306 return op->emitOpError(
"permlane_swap is only supported on gfx950+");
3308 Location loc = op.getLoc();
3309 Type i32 = rewriter.getI32Type();
3310 Value src = adaptor.getSrc();
3311 unsigned rowLength = op.getRowLength();
3312 bool fi = op.getFetchInactive();
3313 bool boundctrl = op.getBoundCtrl();
3315 SmallVector<Value> decomposed;
3317 return rewriter.notifyMatchFailure(op,
3318 "failed to decompose value to i32");
3320 SmallVector<Value> permuted;
3321 for (Value v : decomposed) {
3323 Type i32pair = LLVM::LLVMStructType::getLiteral(
3324 rewriter.getContext(), {v.getType(), v.getType()});
3326 if (rowLength == 16)
3327 res = ROCDL::Permlane16SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3329 else if (rowLength == 32)
3330 res = ROCDL::Permlane32SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3333 llvm_unreachable(
"unsupported row length");
3335 Value vdst0 = LLVM::ExtractValueOp::create(rewriter, loc, res, {0});
3336 Value vdst1 = LLVM::ExtractValueOp::create(rewriter, loc, res, {1});
3338 Value isEqual = LLVM::ICmpOp::create(rewriter, loc,
3339 LLVM::ICmpPredicate::eq, vdst0, v);
3344 LLVM::SelectOp::create(rewriter, loc, isEqual, vdst1, vdst0);
3345 permuted.emplace_back(vdstNew);
3349 rewriter.replaceOp(op,
result);
3354struct AMDGPUPermlaneVarLowering
3355 :
public ConvertOpToLLVMPattern<PermlaneVarOp> {
3358 AMDGPUPermlaneVarLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3359 : ConvertOpToLLVMPattern<PermlaneVarOp>(converter), chipset(chipset) {}
3363 matchAndRewrite(PermlaneVarOp op, OpAdaptor adaptor,
3364 ConversionPatternRewriter &rewriter)
const override {
3366 return op->emitOpError(
"permlane_var is only supported on GFX12+");
3368 Location loc = op.getLoc();
3369 Type i32 = rewriter.getI32Type();
3370 Value src = adaptor.getSrc();
3371 Value selector = adaptor.getSelector();
3372 bool cross = op.getCross();
3373 bool fi = op.getFetchInactive();
3374 bool boundCtrl = op.getBoundCtrl();
3376 SmallVector<Value> decomposed;
3378 return rewriter.notifyMatchFailure(op,
3379 "failed to decompose value to i32");
3381 SmallVector<Value> permuted;
3382 for (Value v : decomposed) {
3385 res = ROCDL::PermlaneX16VarOp::create(rewriter, loc, i32, v, v,
3386 selector, fi, boundCtrl);
3388 res = ROCDL::Permlane16VarOp::create(rewriter, loc, i32, v, v, selector,
3390 permuted.emplace_back(res);
3394 rewriter.replaceOp(op,
result);
3407constexpr int32_t kDsBarrierPendingCountBitWidth = 29;
3408constexpr int32_t kDsBarrierPhasePos = kDsBarrierPendingCountBitWidth;
3409constexpr int32_t kDsBarrierInitCountPos = 32;
3410constexpr int32_t kDsBarrierPendingCountMask =
3411 (1 << kDsBarrierPendingCountBitWidth) - 1;
3413struct DsBarrierInitOpLowering
3414 :
public ConvertOpToLLVMPattern<DsBarrierInitOp> {
3417 DsBarrierInitOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3418 : ConvertOpToLLVMPattern<DsBarrierInitOp>(converter), chipset(chipset) {}
3421 matchAndRewrite(DsBarrierInitOp op, OpAdaptor adaptor,
3422 ConversionPatternRewriter &rewriter)
const override {
3424 return op->emitOpError(
"only supported on gfx1250+");
3426 Location loc = op.getLoc();
3427 Type i64 = rewriter.getI64Type();
3429 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3431 adaptor.getBase(), adaptor.getIndices());
3438 LLVM::SubOp::create(rewriter, loc, adaptor.getParticipants(),
3445 Value maskedCount32 =
3446 LLVM::AndOp::create(rewriter, loc, initCount, countMask);
3447 Value maskedCount = LLVM::ZExtOp::create(rewriter, loc, i64, maskedCount32);
3449 Value initCountShifted = LLVM::ShlOp::create(
3450 rewriter, loc, maskedCount,
3452 Value barrierState =
3453 LLVM::OrOp::create(rewriter, loc, initCountShifted, maskedCount);
3455 LLVM::StoreOp::create(
3456 rewriter, loc, barrierState, ptr, 8,
false,
3458 false, LLVM::AtomicOrdering::release,
3461 rewriter.eraseOp(op);
3466struct DsBarrierPollStateOpLowering
3467 :
public ConvertOpToLLVMPattern<DsBarrierPollStateOp> {
3470 DsBarrierPollStateOpLowering(
const LLVMTypeConverter &converter,
3472 : ConvertOpToLLVMPattern<DsBarrierPollStateOp>(converter),
3476 matchAndRewrite(DsBarrierPollStateOp op, OpAdaptor adaptor,
3477 ConversionPatternRewriter &rewriter)
const override {
3479 return op->emitOpError(
"only supported on gfx1250+");
3481 Location loc = op.getLoc();
3482 Type i64 = rewriter.getI64Type();
3484 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3486 adaptor.getBase(), adaptor.getIndices());
3490 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
3491 op, i64, ptr, 8,
false,
3493 false, LLVM::AtomicOrdering::acquire,
3499struct DsAsyncBarrierArriveOpLowering
3500 :
public ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp> {
3503 DsAsyncBarrierArriveOpLowering(
const LLVMTypeConverter &converter,
3505 : ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp>(converter),
3509 matchAndRewrite(DsAsyncBarrierArriveOp op, OpAdaptor adaptor,
3510 ConversionPatternRewriter &rewriter)
const override {
3512 return op->emitOpError(
"only supported on gfx1250+");
3514 Location loc = op.getLoc();
3516 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3518 adaptor.getBase(), adaptor.getIndices());
3520 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicAsyncBarrierArriveOp>(
3521 op, ptr,
nullptr,
nullptr,
3527struct DsBarrierArriveOpLowering
3528 :
public ConvertOpToLLVMPattern<DsBarrierArriveOp> {
3531 DsBarrierArriveOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3532 : ConvertOpToLLVMPattern<DsBarrierArriveOp>(converter), chipset(chipset) {
3536 matchAndRewrite(DsBarrierArriveOp op, OpAdaptor adaptor,
3537 ConversionPatternRewriter &rewriter)
const override {
3539 return op->emitOpError(
"only supported on gfx1250+");
3541 Location loc = op.getLoc();
3542 Type i64 = rewriter.getI64Type();
3544 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3546 adaptor.getBase(), adaptor.getIndices());
3548 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicBarrierArriveRtnOp>(
3549 op, i64, ptr, adaptor.getCount(),
nullptr,
3555struct DsBarrierStatePhaseOpLowering
3556 :
public ConvertOpToLLVMPattern<DsBarrierStatePhaseOp> {
3560 matchAndRewrite(DsBarrierStatePhaseOp op, OpAdaptor adaptor,
3561 ConversionPatternRewriter &rewriter)
const override {
3562 Location loc = op.getLoc();
3563 Type i32 = rewriter.getI32Type();
3565 Value state = adaptor.getState();
3567 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3568 Value phase = LLVM::LShrOp::create(
3569 rewriter, loc, noInitCount,
3572 rewriter.replaceOp(op, phase);
3577struct DsBarrierStatePendingCountOpLowering
3578 :
public ConvertOpToLLVMPattern<DsBarrierStatePendingCountOp> {
3582 matchAndRewrite(DsBarrierStatePendingCountOp op, OpAdaptor adaptor,
3583 ConversionPatternRewriter &rewriter)
const override {
3584 Location loc = op.getLoc();
3585 Type i32 = rewriter.getI32Type();
3587 Value state = adaptor.getState();
3589 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3590 Value pendingCount = LLVM::AndOp::create(
3591 rewriter, loc, noInitCount,
3593 static_cast<uint32_t
>(kDsBarrierPendingCountMask)));
3595 rewriter.replaceOp(op, pendingCount);
3600struct DsBarrierStateInitCountOpLowering
3601 :
public ConvertOpToLLVMPattern<DsBarrierStateInitCountOp> {
3605 matchAndRewrite(DsBarrierStateInitCountOp op, OpAdaptor adaptor,
3606 ConversionPatternRewriter &rewriter)
const override {
3607 Location loc = op.getLoc();
3608 Type i32 = rewriter.getI32Type();
3610 Value state = adaptor.getState();
3612 Value initCountI64 = LLVM::LShrOp::create(
3613 rewriter, loc, state,
3615 Value initCount = LLVM::TruncOp::create(rewriter, loc, i32, initCountI64);
3617 rewriter.replaceOp(op, initCount);
3622struct DsBarrierStatePhaseParityLowering
3623 :
public ConvertOpToLLVMPattern<DsBarrierStatePhaseParity> {
3627 matchAndRewrite(DsBarrierStatePhaseParity op, OpAdaptor adaptor,
3628 ConversionPatternRewriter &rewriter)
const override {
3629 Location loc = op.getLoc();
3630 Type i1 = rewriter.getI1Type();
3632 Value state = adaptor.getState();
3635 LLVM::TruncOp::create(rewriter, loc, rewriter.getI32Type(), state);
3636 Value phase = LLVM::LShrOp::create(
3637 rewriter, loc, noInitCount,
3639 Value parity = LLVM::TruncOp::create(rewriter, loc, i1, phase);
3641 rewriter.replaceOp(op, parity);
3650static Value setValueAtOffset(ConversionPatternRewriter &rewriter, Location loc,
3651 Value accumulator, Value value, int64_t shift) {
3656 value = LLVM::ShlOp::create(rewriter, loc, value, shiftAmount);
3662 constexpr bool isDisjoint =
true;
3663 return LLVM::OrOp::create(rewriter, loc, accumulator, value, isDisjoint);
3666template <
typename BaseOp>
3667struct AMDGPUMakeDmaBaseLowering :
public ConvertOpToLLVMPattern<BaseOp> {
3668 using ConvertOpToLLVMPattern<BaseOp>::ConvertOpToLLVMPattern;
3671 AMDGPUMakeDmaBaseLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3672 : ConvertOpToLLVMPattern<BaseOp>(converter), chipset(chipset) {}
3676 matchAndRewrite(BaseOp op, Adaptor adaptor,
3677 ConversionPatternRewriter &rewriter)
const override {
3679 return op->emitOpError(
"make_dma_base is only supported on gfx1250");
3681 Location loc = op.getLoc();
3683 constexpr int32_t constlen = 4;
3684 Value consts[constlen];
3685 for (int64_t i = 0; i < constlen; ++i)
3688 constexpr int32_t sgprslen = constlen;
3689 Value sgprs[sgprslen];
3690 for (int64_t i = 0; i < sgprslen; ++i) {
3691 sgprs[i] = consts[0];
3694 sgprs[0] = consts[1];
3696 if constexpr (BaseOp::isGather()) {
3697 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 30);
3699 auto type = cast<TDMGatherBaseType>(op.getResult().getType());
3700 Type indexType = type.getIndexType();
3702 assert(llvm::is_contained({16u, 32u}, indexSize) &&
3703 "expected index_size to be 16 or 32");
3704 unsigned idx = (indexSize / 16) - 1;
3707 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 31);
3710 ValueRange ldsIndices = adaptor.getLdsIndices();
3711 Value lds = adaptor.getLds();
3712 auto ldsMemRefType = cast<MemRefType>(op.getLds().getType());
3715 rewriter, loc, ldsMemRefType, lds, ldsIndices);
3717 ValueRange globalIndices = adaptor.getGlobalIndices();
3718 Value global = adaptor.getGlobal();
3719 auto globalMemRefType = cast<MemRefType>(op.getGlobal().getType());
3722 rewriter, loc, globalMemRefType, global, globalIndices);
3724 Type i32 = rewriter.getI32Type();
3725 Type i64 = rewriter.getI64Type();
3727 sgprs[1] = LLVM::PtrToIntOp::create(rewriter, loc, i32, ldsPtr);
3728 Value castForGlobalAddr =
3729 LLVM::PtrToIntOp::create(rewriter, loc, i64, globalPtr);
3731 sgprs[2] = LLVM::TruncOp::create(rewriter, loc, i32, castForGlobalAddr);
3733 Value shift = LLVM::LShrOp::create(rewriter, loc, castForGlobalAddr,
3736 Value highHalf = LLVM::TruncOp::create(rewriter, loc, i32, shift);
3739 highHalf = LLVM::AndOp::create(rewriter, loc, highHalf, mask);
3741 sgprs[3] = setValueAtOffset(rewriter, loc, highHalf, consts[2], 30);
3743 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
3744 assert(v4i32 &&
"expected type conversion to succeed");
3745 Value
result = LLVM::PoisonOp::create(rewriter, loc, v4i32);
3747 for (
auto [sgpr, constant] : llvm::zip_equal(sgprs, consts))
3749 LLVM::InsertElementOp::create(rewriter, loc,
result, sgpr, constant);
3751 rewriter.replaceOp(op,
result);
3756template <
typename DescriptorOp>
3757struct AMDGPULowerDescriptor :
public ConvertOpToLLVMPattern<DescriptorOp> {
3758 using ConvertOpToLLVMPattern<DescriptorOp>::ConvertOpToLLVMPattern;
3761 AMDGPULowerDescriptor(
const LLVMTypeConverter &converter, Chipset chipset)
3762 : ConvertOpToLLVMPattern<DescriptorOp>(converter), chipset(chipset) {}
3765 Value getDGroup0(OpAdaptor &adaptor)
const {
return adaptor.getBase(); }
3767 Value setWorkgroupMask(DescriptorOp op, OpAdaptor &adaptor,
3768 ConversionPatternRewriter &rewriter, Location loc,
3769 Value sgpr0)
const {
3770 Value mask = op.getWorkgroupMask();
3774 Type i16 = rewriter.getI16Type();
3775 mask = LLVM::BitcastOp::create(rewriter, loc, i16, mask);
3776 Type i32 = rewriter.getI32Type();
3777 Value extendedMask = LLVM::ZExtOp::create(rewriter, loc, i32, mask);
3778 return setValueAtOffset(rewriter, loc, sgpr0, extendedMask, 0);
3781 Value setDataSize(DescriptorOp op, OpAdaptor &adaptor,
3782 ConversionPatternRewriter &rewriter, Location loc,
3783 Value sgpr0, ArrayRef<Value> consts)
const {
3784 unsigned elementTypeWidthInBits = op.getElementTypeWidth();
3785 assert(llvm::is_contained({8u, 16u, 32u, 64u}, elementTypeWidthInBits) &&
3786 "expected type width to be 8, 16, 32, or 64.");
3787 int64_t idx = llvm::Log2_32(elementTypeWidthInBits / 8);
3788 Value size = consts[idx];
3789 return setValueAtOffset(rewriter, loc, sgpr0, size, 16);
3792 Value setAtomicBarrier(DescriptorOp op, OpAdaptor &adaptor,
3793 ConversionPatternRewriter &rewriter, Location loc,
3794 Value sgpr0, ArrayRef<Value> consts)
const {
3795 if (!adaptor.getAtomicBarrierAddress())
3798 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 18);
3801 Value setIterateEnable(DescriptorOp op, OpAdaptor &adaptor,
3802 ConversionPatternRewriter &rewriter, Location loc,
3803 Value sgpr0, ArrayRef<Value> consts)
const {
3804 if (!adaptor.getGlobalIncrement())
3809 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 19);
3812 Value setPadEnable(DescriptorOp op, OpAdaptor &adaptor,
3813 ConversionPatternRewriter &rewriter, Location loc,
3814 Value sgpr0, ArrayRef<Value> consts)
const {
3815 if (!op.getPadAmount())
3818 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 20);
3821 Value setEarlyTimeout(DescriptorOp op, OpAdaptor &adaptor,
3822 ConversionPatternRewriter &rewriter, Location loc,
3823 Value sgpr0, ArrayRef<Value> consts)
const {
3824 if (!op.getWorkgroupMask())
3827 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 21);
3830 Value setPadInterval(DescriptorOp op, OpAdaptor &adaptor,
3831 ConversionPatternRewriter &rewriter, Location loc,
3832 Value sgpr0, ArrayRef<Value> consts)
const {
3833 if (!op.getPadAmount())
3842 IntegerType i32 = rewriter.getI32Type();
3843 Value padInterval = adaptor.getPadInterval();
3844 padInterval = LLVM::CountTrailingZerosOp::create(rewriter, loc, i32,
3845 padInterval,
false);
3846 padInterval = LLVM::SubOp::create(rewriter, loc, padInterval, consts[1]);
3848 return setValueAtOffset(rewriter, loc, sgpr0, padInterval, 22);
3851 Value setPadAmount(DescriptorOp op, OpAdaptor &adaptor,
3852 ConversionPatternRewriter &rewriter, Location loc,
3853 Value sgpr0, ArrayRef<Value> consts)
const {
3854 if (!op.getPadAmount())
3863 Value padAmount = adaptor.getPadAmount();
3864 padAmount = LLVM::SubOp::create(rewriter, loc, padAmount, consts[1]);
3866 return setValueAtOffset(rewriter, loc, sgpr0, padAmount, 25);
3869 Value setAtomicBarrierAddress(DescriptorOp op, OpAdaptor &adaptor,
3870 ConversionPatternRewriter &rewriter,
3871 Location loc, Value sgpr1,
3872 ArrayRef<Value> consts)
const {
3873 if (!adaptor.getAtomicBarrierAddress())
3876 Value atomicBarrierAddress = adaptor.getAtomicBarrierAddress();
3877 auto barrierAddressTy =
3878 cast<MemRefType>(op.getAtomicBarrierAddress().getType());
3879 ValueRange atomicBarrierIndices = adaptor.getAtomicBarrierIndices();
3881 rewriter, loc, barrierAddressTy, atomicBarrierAddress,
3882 atomicBarrierIndices);
3883 IntegerType i32 = rewriter.getI32Type();
3889 atomicBarrierAddress =
3890 LLVM::PtrToIntOp::create(rewriter, loc, i32, atomicBarrierAddress);
3891 atomicBarrierAddress =
3892 LLVM::LShrOp::create(rewriter, loc, atomicBarrierAddress, consts[3]);
3894 atomicBarrierAddress =
3895 LLVM::AndOp::create(rewriter, loc, atomicBarrierAddress, mask);
3896 return setValueAtOffset(rewriter, loc, sgpr1, atomicBarrierAddress, 32);
3899 std::pair<Value, Value> setTensorDimX(DescriptorOp op, OpAdaptor &adaptor,
3900 ConversionPatternRewriter &rewriter,
3901 Location loc, Value sgpr1, Value sgpr2,
3902 ArrayRef<Value> consts, uint64_t dimX,
3903 uint32_t offset)
const {
3904 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
3905 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
3906 SmallVector<OpFoldResult> mixedGlobalSizes =
3908 if (mixedGlobalSizes.size() <= dimX)
3909 return {sgpr1, sgpr2};
3911 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
3918 if (
auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
3922 IntegerType i32 = rewriter.getI32Type();
3923 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
3924 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
3927 sgpr1 = setValueAtOffset(rewriter, loc, sgpr1, tensorDimX, offset);
3930 Value tensorDimXHigh = LLVM::LShrOp::create(rewriter, loc, tensorDimX, c16);
3931 sgpr2 = setValueAtOffset(rewriter, loc, sgpr2, tensorDimXHigh, offset + 16);
3932 return {sgpr1, sgpr2};
3935 std::pair<Value, Value> setTensorDim0(DescriptorOp op, OpAdaptor &adaptor,
3936 ConversionPatternRewriter &rewriter,
3937 Location loc, Value sgpr1, Value sgpr2,
3938 ArrayRef<Value> consts)
const {
3939 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, 0,
3943 std::pair<Value, Value> setTensorDim1(DescriptorOp op, OpAdaptor &adaptor,
3944 ConversionPatternRewriter &rewriter,
3945 Location loc, Value sgpr2, Value sgpr3,
3946 ArrayRef<Value> consts)
const {
3947 return setTensorDimX(op, adaptor, rewriter, loc, sgpr2, sgpr3, consts, 1,
3951 Value setTileDimX(DescriptorOp op, OpAdaptor &adaptor,
3952 ConversionPatternRewriter &rewriter, Location loc,
3953 Value sgpr, ArrayRef<Value> consts,
size_t dimX,
3954 int64_t offset)
const {
3955 ArrayRef<int64_t> sharedStaticSizes = adaptor.getSharedStaticSizes();
3956 ValueRange sharedDynamicSizes = adaptor.getSharedDynamicSizes();
3957 SmallVector<OpFoldResult> mixedSharedSizes =
3959 if (mixedSharedSizes.size() <= dimX)
3962 OpFoldResult tileDimXOpFoldResult = *(mixedSharedSizes.rbegin() + dimX);
3971 if (
auto attr = dyn_cast<Attribute>(tileDimXOpFoldResult)) {
3975 IntegerType i32 = rewriter.getI32Type();
3976 tileDimX = cast<Value>(tileDimXOpFoldResult);
3977 tileDimX = LLVM::TruncOp::create(rewriter, loc, i32, tileDimX);
3980 return setValueAtOffset(rewriter, loc, sgpr, tileDimX, offset);
3983 Value setTileDim0(DescriptorOp op, OpAdaptor &adaptor,
3984 ConversionPatternRewriter &rewriter, Location loc,
3985 Value sgpr3, ArrayRef<Value> consts)
const {
3986 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, 0, 112);
3989 Value setTileDim1(DescriptorOp op, OpAdaptor &adaptor,
3990 ConversionPatternRewriter &rewriter, Location loc,
3991 Value sgpr4, ArrayRef<Value> consts)
const {
3992 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 1, 128);
3995 Value setValidIndices(DescriptorOp op, OpAdaptor &adaptor,
3996 ConversionPatternRewriter &rewriter, Location loc,
3997 Value sgpr4, ArrayRef<Value> consts)
const {
3998 auto type = cast<VectorType>(op.getIndices().getType());
3999 ArrayRef<int64_t> shape = type.getShape();
4000 assert(shape.size() == 1 &&
"expected shape to be of rank 1.");
4001 unsigned length = shape.back();
4002 assert(0 < length && length <= 16 &&
"expected length to be at most 16.");
4004 return setValueAtOffset(rewriter, loc, sgpr4, value, 128);
4007 Value setTileDim1OrValidIndices(DescriptorOp op, OpAdaptor &adaptor,
4008 ConversionPatternRewriter &rewriter,
4009 Location loc, Value sgpr4,
4010 ArrayRef<Value> consts)
const {
4011 if constexpr (DescriptorOp::isGather())
4012 return setValidIndices(op, adaptor, rewriter, loc, sgpr4, consts);
4013 return setTileDim1(op, adaptor, rewriter, loc, sgpr4, consts);
4016 Value setTileDim2(DescriptorOp op, OpAdaptor &adaptor,
4017 ConversionPatternRewriter &rewriter, Location loc,
4018 Value sgpr4, ArrayRef<Value> consts)
const {
4020 if constexpr (DescriptorOp::isGather())
4022 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 2, 144);
4025 std::pair<Value, Value>
4026 setTensorDimXStride(DescriptorOp op, OpAdaptor &adaptor,
4027 ConversionPatternRewriter &rewriter, Location loc,
4028 Value sgprY, Value sgprZ, ArrayRef<Value> consts,
4029 size_t dimX, int64_t offset)
const {
4030 ArrayRef<int64_t> globalStaticStrides = adaptor.getGlobalStaticStrides();
4031 ValueRange globalDynamicStrides = adaptor.getGlobalDynamicStrides();
4032 SmallVector<OpFoldResult> mixedGlobalStrides =
4033 getMixedValues(globalStaticStrides, globalDynamicStrides, rewriter);
4035 if (mixedGlobalStrides.size() <= (dimX + 1))
4036 return {sgprY, sgprZ};
4038 OpFoldResult tensorDimXStrideOpFoldResult =
4039 *(mixedGlobalStrides.rbegin() + dimX + 1);
4044 Value tensorDimXStride;
4045 if (
auto attr = dyn_cast<Attribute>(tensorDimXStrideOpFoldResult))
4049 tensorDimXStride = cast<Value>(tensorDimXStrideOpFoldResult);
4051 constexpr int64_t first48bits = (1ll << 48) - 1;
4054 LLVM::AndOp::create(rewriter, loc, mask, tensorDimXStride);
4055 IntegerType i32 = rewriter.getI32Type();
4056 Value tensorDimXStrideLow =
4057 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStride);
4058 sgprY = setValueAtOffset(rewriter, loc, sgprY, tensorDimXStrideLow, offset);
4060 int64_t shift = (offset % 32) == 0 ? 32 : offset % 32;
4062 Value tensorDimXStrideHigh =
4063 LLVM::LShrOp::create(rewriter, loc, tensorDimXStride, shiftVal);
4064 tensorDimXStrideHigh =
4065 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStrideHigh);
4066 sgprZ = setValueAtOffset(rewriter, loc, sgprZ, tensorDimXStrideHigh,
4068 return {sgprY, sgprZ};
4071 std::pair<Value, Value>
4072 setTensorDim0Stride(DescriptorOp op, OpAdaptor &adaptor,
4073 ConversionPatternRewriter &rewriter, Location loc,
4074 Value sgpr5, Value sgpr6, ArrayRef<Value> consts)
const {
4075 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4079 std::pair<Value, Value>
4080 setTensorDim1Stride(DescriptorOp op, OpAdaptor &adaptor,
4081 ConversionPatternRewriter &rewriter, Location loc,
4082 Value sgpr5, Value sgpr6, ArrayRef<Value> consts)
const {
4084 if constexpr (DescriptorOp::isGather())
4085 return {sgpr5, sgpr6};
4086 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4090 Value getDGroup1(DescriptorOp op, OpAdaptor &adaptor,
4091 ConversionPatternRewriter &rewriter, Location loc,
4092 ArrayRef<Value> consts)
const {
4094 for (int64_t i = 0; i < 8; ++i) {
4095 sgprs[i] = consts[0];
4098 sgprs[0] = setWorkgroupMask(op, adaptor, rewriter, loc, sgprs[0]);
4099 sgprs[0] = setDataSize(op, adaptor, rewriter, loc, sgprs[0], consts);
4100 sgprs[0] = setAtomicBarrier(op, adaptor, rewriter, loc, sgprs[0], consts);
4101 sgprs[0] = setIterateEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4102 sgprs[0] = setPadEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4103 sgprs[0] = setEarlyTimeout(op, adaptor, rewriter, loc, sgprs[0], consts);
4104 sgprs[0] = setPadInterval(op, adaptor, rewriter, loc, sgprs[0], consts);
4105 sgprs[0] = setPadAmount(op, adaptor, rewriter, loc, sgprs[0], consts);
4108 setAtomicBarrierAddress(op, adaptor, rewriter, loc, sgprs[1], consts);
4109 std::tie(sgprs[1], sgprs[2]) =
4110 setTensorDim0(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4111 std::tie(sgprs[2], sgprs[3]) =
4112 setTensorDim1(op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4114 sgprs[3] = setTileDim0(op, adaptor, rewriter, loc, sgprs[3], consts);
4116 setTileDim1OrValidIndices(op, adaptor, rewriter, loc, sgprs[4], consts);
4117 sgprs[4] = setTileDim2(op, adaptor, rewriter, loc, sgprs[4], consts);
4118 std::tie(sgprs[5], sgprs[6]) = setTensorDim0Stride(
4119 op, adaptor, rewriter, loc, sgprs[5], sgprs[6], consts);
4120 std::tie(sgprs[6], sgprs[7]) = setTensorDim1Stride(
4121 op, adaptor, rewriter, loc, sgprs[6], sgprs[7], consts);
4123 IntegerType i32 = rewriter.getI32Type();
4124 Type v8i32 = this->typeConverter->convertType(VectorType::get(8, i32));
4125 assert(v8i32 &&
"expected type conversion to succeed");
4126 Value dgroup1 = LLVM::PoisonOp::create(rewriter, loc, v8i32);
4128 for (
auto [sgpr, constant] : llvm::zip_equal(sgprs, consts)) {
4130 LLVM::InsertElementOp::create(rewriter, loc, dgroup1, sgpr, constant);
4136 Value setTensorDimX(DescriptorOp op, OpAdaptor &adaptor,
4137 ConversionPatternRewriter &rewriter, Location loc,
4138 Value sgpr0, ArrayRef<Value> consts, int64_t dimX,
4139 int64_t offset)
const {
4140 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
4141 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
4142 SmallVector<OpFoldResult> mixedGlobalSizes =
4144 if (mixedGlobalSizes.size() <=
static_cast<unsigned long>(dimX))
4147 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
4149 if (
auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
4153 IntegerType i32 = rewriter.getI32Type();
4154 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
4155 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
4158 return setValueAtOffset(rewriter, loc, sgpr0, tensorDimX, offset);
4161 Value setTensorDim2(DescriptorOp op, OpAdaptor &adaptor,
4162 ConversionPatternRewriter &rewriter, Location loc,
4163 Value sgpr0, ArrayRef<Value> consts)
const {
4164 return setTensorDimX(op, adaptor, rewriter, loc, sgpr0, consts, 2, 0);
4167 Value truncateAndSetValueAtOffset(ConversionPatternRewriter &rewriter,
4168 Location loc, Value accumulator,
4169 Value value, int64_t shift)
const {
4171 IntegerType i32 = rewriter.getI32Type();
4172 value = LLVM::TruncOp::create(rewriter, loc, i32, value);
4173 return setValueAtOffset(rewriter, loc, accumulator, value, shift);
4176 Value setLDSAddrIncrement(DescriptorOp op, OpAdaptor &adaptor,
4177 ConversionPatternRewriter &rewriter, Location loc,
4178 Value sgpr1, ArrayRef<Value> consts,
4179 int64_t offset)
const {
4180 Value ldsAddrIncrement = adaptor.getLdsIncrement();
4181 return setValueAtOffset(rewriter, loc, sgpr1, ldsAddrIncrement, offset);
4184 std::pair<Value, Value>
4185 setGlobalAddrIncrement(DescriptorOp op, OpAdaptor &adaptor,
4186 ConversionPatternRewriter &rewriter, Location loc,
4187 Value sgpr2, Value sgpr3, ArrayRef<Value> consts,
4188 int64_t offset)
const {
4189 Value globalAddrIncrement = adaptor.getGlobalIncrement();
4190 sgpr2 = truncateAndSetValueAtOffset(rewriter, loc, sgpr2,
4191 globalAddrIncrement, offset);
4193 globalAddrIncrement =
4194 LLVM::LShrOp::create(rewriter, loc, globalAddrIncrement, shift);
4195 constexpr int64_t first16BitsHigh = (1ll << 16) - 1;
4196 sgpr3 = truncateAndSetValueAtOffset(rewriter, loc, sgpr3,
4197 globalAddrIncrement, offset + 32);
4199 sgpr3 = LLVM::AndOp::create(rewriter, loc, sgpr3, mask);
4200 return {sgpr2, sgpr3};
4203 Value setTensorDim3OrLDSAddrIncrement(DescriptorOp op, OpAdaptor &adaptor,
4204 ConversionPatternRewriter &rewriter,
4205 Location loc, Value sgpr1,
4206 ArrayRef<Value> consts)
const {
4207 Value ldsIncrement = op.getLdsIncrement();
4208 constexpr int64_t dim = 3;
4209 constexpr int64_t offset = 32;
4211 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, consts, dim,
4213 return setLDSAddrIncrement(op, adaptor, rewriter, loc, sgpr1, consts,
4217 std::pair<Value, Value> setTensorDim2StrideOrGlobalAddrIncrement(
4218 DescriptorOp op, OpAdaptor &adaptor, ConversionPatternRewriter &rewriter,
4219 Location loc, Value sgpr2, Value sgpr3, ArrayRef<Value> consts)
const {
4220 Value globalIncrement = op.getGlobalIncrement();
4221 constexpr int32_t dim = 2;
4222 constexpr int32_t offset = 64;
4223 if (!globalIncrement)
4224 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4225 consts, dim, offset);
4226 return setGlobalAddrIncrement(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4230 Value setIterateCount(DescriptorOp op, OpAdaptor &adaptor,
4231 ConversionPatternRewriter &rewriter, Location loc,
4232 Value sgpr3, ArrayRef<Value> consts,
4233 int32_t offset)
const {
4234 Value iterationCount = adaptor.getIterationCount();
4235 IntegerType i32 = rewriter.getI32Type();
4242 iterationCount = LLVM::TruncOp::create(rewriter, loc, i32, iterationCount);
4244 LLVM::SubOp::create(rewriter, loc, iterationCount, consts[1]);
4245 return setValueAtOffset(rewriter, loc, sgpr3, iterationCount, offset);
4248 Value setTileDim3OrIterateCount(DescriptorOp op, OpAdaptor &adaptor,
4249 ConversionPatternRewriter &rewriter,
4250 Location loc, Value sgpr3,
4251 ArrayRef<Value> consts)
const {
4252 Value iterateCount = op.getIterationCount();
4253 constexpr int32_t dim = 2;
4254 constexpr int32_t offset = 112;
4256 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, dim,
4259 return setIterateCount(op, adaptor, rewriter, loc, sgpr3, consts, offset);
4262 Value getDGroup2(DescriptorOp op, OpAdaptor &adaptor,
4263 ConversionPatternRewriter &rewriter, Location loc,
4264 ArrayRef<Value> consts)
const {
4265 if constexpr (DescriptorOp::isGather())
4266 return getDGroup2Gather(op, adaptor, rewriter, loc, consts);
4267 return getDGroup2NonGather(op, adaptor, rewriter, loc, consts);
4270 Value getDGroup2NonGather(DescriptorOp op, OpAdaptor &adaptor,
4271 ConversionPatternRewriter &rewriter, Location loc,
4272 ArrayRef<Value> consts)
const {
4273 IntegerType i32 = rewriter.getI32Type();
4274 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4275 assert(v4i32 &&
"expected type conversion to succeed.");
4277 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4278 if (onlyNeedsTwoDescriptors)
4279 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4281 constexpr int64_t sgprlen = 4;
4282 Value sgprs[sgprlen];
4283 for (
int i = 0; i < sgprlen; ++i)
4284 sgprs[i] = consts[0];
4286 sgprs[0] = setTensorDim2(op, adaptor, rewriter, loc, sgprs[0], consts);
4287 sgprs[1] = setTensorDim3OrLDSAddrIncrement(op, adaptor, rewriter, loc,
4289 std::tie(sgprs[2], sgprs[3]) = setTensorDim2StrideOrGlobalAddrIncrement(
4290 op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4292 setTileDim3OrIterateCount(op, adaptor, rewriter, loc, sgprs[3], consts);
4294 Value dgroup2 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4295 for (
auto [sgpr, constant] : llvm::zip(sgprs, consts))
4297 LLVM::InsertElementOp::create(rewriter, loc, dgroup2, sgpr, constant);
4302 Value getGatherIndices(DescriptorOp op, OpAdaptor &adaptor,
4303 ConversionPatternRewriter &rewriter, Location loc,
4304 ArrayRef<Value> consts,
bool firstHalf)
const {
4305 IntegerType i32 = rewriter.getI32Type();
4306 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4307 assert(v4i32 &&
"expected type conversion to succeed.");
4309 Value
indices = adaptor.getIndices();
4310 auto vectorType = cast<VectorType>(
indices.getType());
4311 unsigned length = vectorType.getShape().back();
4312 Type elementType = vectorType.getElementType();
4313 unsigned maxLength = elementType == i32 ? 4 : 8;
4314 int32_t offset = firstHalf ? 0 : maxLength;
4315 unsigned discountedLength =
4316 std::max(
static_cast<int32_t
>(length - offset), 0);
4318 unsigned targetSize = std::min(maxLength, discountedLength);
4320 SmallVector<Value> indicesVector;
4321 for (
unsigned i = offset; i < targetSize + offset; ++i) {
4323 if (i < consts.size())
4327 Value elem = LLVM::ExtractElementOp::create(rewriter, loc,
indices, idx);
4328 indicesVector.push_back(elem);
4331 SmallVector<Value> indicesI32Vector;
4332 if (elementType == i32) {
4333 indicesI32Vector = std::move(indicesVector);
4335 for (
unsigned i = 0; i < targetSize; ++i) {
4336 Value index = indicesVector[i];
4337 indicesI32Vector.push_back(
4338 LLVM::ZExtOp::create(rewriter, loc, i32, index));
4340 if ((targetSize % 2) != 0)
4342 indicesI32Vector.push_back(consts[0]);
4345 SmallVector<Value> indicesToInsert;
4346 if (elementType == i32) {
4347 indicesToInsert = std::move(indicesI32Vector);
4349 unsigned size = indicesI32Vector.size() / 2;
4350 for (
unsigned i = 0; i < size; ++i) {
4351 Value first = indicesI32Vector[2 * i];
4352 Value second = indicesI32Vector[2 * i + 1];
4353 Value joined = setValueAtOffset(rewriter, loc, first, second, 16);
4354 indicesToInsert.push_back(joined);
4358 Value dgroup = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4359 for (
auto [sgpr, constant] : llvm::zip_first(indicesToInsert, consts))
4361 LLVM::InsertElementOp::create(rewriter, loc, dgroup, sgpr, constant);
4366 Value getDGroup2Gather(DescriptorOp op, OpAdaptor &adaptor,
4367 ConversionPatternRewriter &rewriter, Location loc,
4368 ArrayRef<Value> consts)
const {
4369 return getGatherIndices(op, adaptor, rewriter, loc, consts,
true);
4372 std::pair<Value, Value>
4373 setTensorDim3Stride(DescriptorOp op, OpAdaptor &adaptor,
4374 ConversionPatternRewriter &rewriter, Location loc,
4375 Value sgpr0, Value sgpr1, ArrayRef<Value> consts)
const {
4376 constexpr int32_t dim = 3;
4377 constexpr int32_t offset = 0;
4378 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr0, sgpr1, consts,
4382 std::pair<Value, Value> setTensorDim4(DescriptorOp op, OpAdaptor &adaptor,
4383 ConversionPatternRewriter &rewriter,
4384 Location loc, Value sgpr1, Value sgpr2,
4385 ArrayRef<Value> consts)
const {
4386 constexpr int32_t dim = 4;
4387 constexpr int32_t offset = 48;
4388 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, dim,
4392 Value setTileDim4(DescriptorOp op, OpAdaptor &adaptor,
4393 ConversionPatternRewriter &rewriter, Location loc,
4394 Value sgpr2, ArrayRef<Value> consts)
const {
4395 constexpr int32_t dim = 4;
4396 constexpr int32_t offset = 80;
4397 return setTileDimX(op, adaptor, rewriter, loc, sgpr2, consts, dim, offset);
4400 Value getDGroup3(DescriptorOp op, OpAdaptor &adaptor,
4401 ConversionPatternRewriter &rewriter, Location loc,
4402 ArrayRef<Value> consts)
const {
4403 if constexpr (DescriptorOp::isGather())
4404 return getDGroup3Gather(op, adaptor, rewriter, loc, consts);
4405 return getDGroup3NonGather(op, adaptor, rewriter, loc, consts);
4408 Value getDGroup3NonGather(DescriptorOp op, OpAdaptor &adaptor,
4409 ConversionPatternRewriter &rewriter, Location loc,
4410 ArrayRef<Value> consts)
const {
4411 IntegerType i32 = rewriter.getI32Type();
4412 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4413 assert(v4i32 &&
"expected type conversion to succeed.");
4414 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4415 if (onlyNeedsTwoDescriptors)
4416 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4418 constexpr int32_t sgprlen = 4;
4419 Value sgprs[sgprlen];
4420 for (
int i = 0; i < sgprlen; ++i)
4421 sgprs[i] = consts[0];
4423 std::tie(sgprs[0], sgprs[1]) = setTensorDim3Stride(
4424 op, adaptor, rewriter, loc, sgprs[0], sgprs[1], consts);
4425 std::tie(sgprs[1], sgprs[2]) =
4426 setTensorDim4(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4427 sgprs[2] = setTileDim4(op, adaptor, rewriter, loc, sgprs[2], consts);
4429 Value dgroup3 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4430 for (
auto [sgpr, constant] : llvm::zip(sgprs, consts))
4432 LLVM::InsertElementOp::create(rewriter, loc, dgroup3, sgpr, constant);
4437 Value getDGroup3Gather(DescriptorOp op, OpAdaptor &adaptor,
4438 ConversionPatternRewriter &rewriter, Location loc,
4439 ArrayRef<Value> consts)
const {
4440 return getGatherIndices(op, adaptor, rewriter, loc, consts,
false);
4444 matchAndRewrite(DescriptorOp op, OpAdaptor adaptor,
4445 ConversionPatternRewriter &rewriter)
const override {
4447 return op->emitOpError(
4448 "make_dma_descriptor is only supported on gfx1250");
4450 Location loc = op.getLoc();
4452 SmallVector<Value> consts;
4453 for (int64_t i = 0; i < 8; ++i)
4456 Value dgroup0 = this->getDGroup0(adaptor);
4457 Value dgroup1 = this->getDGroup1(op, adaptor, rewriter, loc, consts);
4458 Value dgroup2 = this->getDGroup2(op, adaptor, rewriter, loc, consts);
4459 Value dgroup3 = this->getDGroup3(op, adaptor, rewriter, loc, consts);
4460 SmallVector<Value> results = {dgroup0, dgroup1, dgroup2, dgroup3};
4461 rewriter.replaceOpWithMultiple(op, {results});
4466template <
typename SourceOp,
typename TargetOp>
4467struct AMDGPUTensorLoadStoreOpLowering
4468 :
public ConvertOpToLLVMPattern<SourceOp> {
4469 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
4471 AMDGPUTensorLoadStoreOpLowering(
const LLVMTypeConverter &converter,
4473 : ConvertOpToLLVMPattern<SourceOp>(converter), chipset(chipset) {}
4477 matchAndRewrite(SourceOp op, Adaptor adaptor,
4478 ConversionPatternRewriter &rewriter)
const override {
4480 return op->emitOpError(
"is only supported on gfx1250");
4485 auto v8i32 = VectorType::get(8, rewriter.getI32Type());
4486 Value dgroup4 = LLVM::ZeroOp::create(rewriter, op.getLoc(), v8i32);
4487 Attribute cachePolicy = rewriter.getI32IntegerAttr(0);
4488 rewriter.replaceOpWithNewOp<TargetOp>(op, desc[0], desc[1], desc[2],
4489 desc[3], dgroup4, cachePolicy,
4497struct GlobalPrefetchOpLowering
4498 :
public ConvertOpToLLVMPattern<GlobalPrefetchOp> {
4499 GlobalPrefetchOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
4500 : ConvertOpToLLVMPattern<GlobalPrefetchOp>(converter), chipset(chipset) {}
4503 matchAndRewrite(GlobalPrefetchOp op, GlobalPrefetchOpAdaptor adaptor,
4504 ConversionPatternRewriter &rewriter)
const override {
4506 return op->emitOpError(
"is only supported on gfx1250+");
4508 const bool isSpeculative = op.getSpeculative();
4510 op.getTemporalHint(), op.getCacheScope(), isSpeculative);
4513 Attribute cachePolicy = ROCDL::Gfx12CachePolicyAttr::get(
4514 rewriter.getContext(),
4515 static_cast<ROCDL::Gfx12CachePolicy
>(immArgValue));
4518 Value memRef = adaptor.getSrc();
4519 MemRefDescriptor descriptor(memRef);
4520 MemRefType memRefType = op.getSrc().getType();
4521 Location loc = op->getLoc();
4522 auto inboundsFlags = isSpeculative ? LLVM::GEPNoWrapFlags::none
4523 : LLVM::GEPNoWrapFlags::inbounds |
4524 LLVM::GEPNoWrapFlags::nuw;
4526 rewriter, loc, memRefType, descriptor,
indices, inboundsFlags);
4528 rewriter.replaceOpWithNewOp<ROCDL::GlobalPrefetchOp>(
4529 op, prefetchPtr, cachePolicy, mlir::ArrayAttr{}, mlir::ArrayAttr{},
4538struct ConvertAMDGPUToROCDLPass
4539 :
public impl::ConvertAMDGPUToROCDLPassBase<ConvertAMDGPUToROCDLPass> {
4542 void runOnOperation()
override {
4545 if (
failed(maybeChipset)) {
4546 emitError(UnknownLoc::get(ctx),
"Invalid chipset name: " + chipset);
4547 return signalPassFailure();
4550 RewritePatternSet patterns(ctx);
4551 LLVMTypeConverter converter(ctx);
4556 target.addIllegalDialect<::mlir::amdgpu::AMDGPUDialect>();
4557 target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
4558 target.addLegalDialect<::mlir::ROCDL::ROCDLDialect>();
4559 if (
failed(applyPartialConversion(getOperation(),
target,
4560 std::move(patterns))))
4561 signalPassFailure();
4569 typeConverter, [](gpu::AddressSpace space) {
4571 case gpu::AddressSpace::Global:
4572 return ROCDL::ROCDLDialect::kGlobalMemoryAddressSpace;
4573 case gpu::AddressSpace::Workgroup:
4574 return ROCDL::ROCDLDialect::kSharedMemoryAddressSpace;
4575 case gpu::AddressSpace::Private:
4576 return ROCDL::ROCDLDialect::kPrivateMemoryAddressSpace;
4577 case gpu::AddressSpace::Constant:
4578 return ROCDL::ROCDLDialect::kConstantMemoryAddressSpace;
4580 llvm_unreachable(
"unknown address space enum value");
4583 return LLVM::LLVMPointerType::get(
4584 type.getContext(), ROCDL::ROCDLDialect::kBarrierAddressSpace);
4590 typeConverter.addTypeAttributeConversion(
4592 -> TypeConverter::AttributeConversionResult {
4594 Type i64 = IntegerType::get(ctx, 64);
4595 switch (as.getValue()) {
4596 case amdgpu::AddressSpace::FatRawBuffer:
4597 return IntegerAttr::get(i64, 7);
4598 case amdgpu::AddressSpace::BufferRsrc:
4599 return IntegerAttr::get(i64, 8);
4600 case amdgpu::AddressSpace::FatStructuredBuffer:
4601 return IntegerAttr::get(i64, 9);
4603 return TypeConverter::AttributeConversionResult::abort();
4605 typeConverter.addConversion([&](DsBarrierStateType type) ->
Type {
4606 return IntegerType::get(type.
getContext(), 64);
4608 typeConverter.addConversion([&](TDMBaseType type) ->
Type {
4610 return typeConverter.convertType(VectorType::get(4, i32));
4612 typeConverter.addConversion([&](TDMGatherBaseType type) ->
Type {
4614 return typeConverter.convertType(VectorType::get(4, i32));
4616 typeConverter.addConversion(
4617 [&](TDMDescriptorType type,
4620 Type v4i32 = typeConverter.convertType(VectorType::get(4, i32));
4621 Type v8i32 = typeConverter.convertType(VectorType::get(8, i32));
4622 llvm::append_values(
result, v4i32, v8i32, v4i32, v4i32);
4632 if (inputs.size() != 1)
4635 if (!isa<TDMDescriptorType>(inputs[0].
getType()))
4638 auto cast = UnrealizedConversionCastOp::create(builder, loc, types, inputs);
4639 return cast.getResults();
4642 typeConverter.addTargetMaterialization(addUnrealizedCast);
4650 .
add<FatRawBufferCastLowering,
4651 RawBufferOpLowering<RawBufferLoadOp, ROCDL::RawPtrBufferLoadOp>,
4652 RawBufferOpLowering<RawBufferStoreOp, ROCDL::RawPtrBufferStoreOp>,
4653 RawBufferOpLowering<RawBufferAtomicFaddOp,
4654 ROCDL::RawPtrBufferAtomicFaddOp>,
4655 RawBufferOpLowering<RawBufferAtomicFmaxOp,
4656 ROCDL::RawPtrBufferAtomicFmaxOp>,
4657 RawBufferOpLowering<RawBufferAtomicSmaxOp,
4658 ROCDL::RawPtrBufferAtomicSmaxOp>,
4659 RawBufferOpLowering<RawBufferAtomicUminOp,
4660 ROCDL::RawPtrBufferAtomicUminOp>,
4661 RawBufferOpLowering<RawBufferAtomicCmpswapOp,
4662 ROCDL::RawPtrBufferAtomicCmpSwap>,
4663 AMDGPUDPPLowering, MemoryCounterWaitOpLowering, LDSBarrierOpLowering,
4664 SchedBarrierOpLowering, MFMAOpLowering, ScaledMFMAOpLowering,
4665 SparseMFMAOpLowering, WMMAOpLowering, ScaledWMMAOpLowering,
4666 SparseWMMAOpLowering, DotOpLowering, ExtPackedFp8OpLowering,
4667 ScaledExtPackedMatrixOpLowering, ScaledExtPackedOpLowering,
4668 PackedScaledTruncOpLowering, PackedTrunc2xFp8OpLowering,
4669 PackedStochRoundFp8OpLowering, GatherToLDSOpLowering,
4670 GlobalLoadAsyncToLDSOpLowering, TransposeLoadOpLowering,
4671 GlobalTransposeLoadOpLowering, AMDGPUPermlaneLowering,
4672 AMDGPUPermlaneVarLowering, AMDGPUMakeDmaBaseLowering<MakeDmaBaseOp>,
4673 AMDGPUMakeDmaBaseLowering<MakeGatherDmaBaseOp>,
4674 AMDGPULowerDescriptor<MakeDmaDescriptorOp>,
4675 AMDGPULowerDescriptor<MakeGatherDmaDescriptorOp>,
4676 AMDGPUTensorLoadStoreOpLowering<TensorLoadToLDSOp,
4677 ROCDL::TensorLoadToLDSOp>,
4678 AMDGPUTensorLoadStoreOpLowering<TensorStoreFromLDSOp,
4679 ROCDL::TensorStoreFromLDSOp>,
4680 DsBarrierInitOpLowering, DsBarrierPollStateOpLowering,
4681 DsAsyncBarrierArriveOpLowering, DsBarrierArriveOpLowering,
4682 GlobalPrefetchOpLowering>(converter, chipset);
4683 patterns.
add<AMDGPUSwizzleBitModeLowering, DsBarrierStatePhaseOpLowering,
4684 DsBarrierStatePendingCountOpLowering,
4685 DsBarrierStateInitCountOpLowering,
4686 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.