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 i32 = rewriter.getI32Type();
117 auto valTy = cast<IntegerType>(val.
getType());
120 return valTy.getWidth() > 32
121 ?
Value(LLVM::TruncOp::create(rewriter, loc, i32, val))
122 :
Value(LLVM::ZExtOp::create(rewriter, loc, i32, val));
127 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), value);
133 IntegerType i64 = rewriter.getI64Type();
135 auto valTy = cast<IntegerType>(val.
getType());
138 return valTy.getWidth() > 64
139 ?
Value(LLVM::TruncOp::create(rewriter, loc, i64, val))
140 :
Value(LLVM::ZExtOp::create(rewriter, loc, i64, val));
145 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), value);
152 IntegerType i32 = rewriter.getI32Type();
154 for (
auto [i, increment, stride] : llvm::enumerate(
indices, strides)) {
157 ShapedType::isDynamic(stride)
159 memRefDescriptor.
stride(rewriter, loc, i))
160 : LLVM::ConstantOp::create(rewriter, loc, i32, stride);
161 increment = LLVM::MulOp::create(rewriter, loc, increment, strideValue);
173 MemRefType memrefType,
178 constexpr int64_t first45bits = (1ll << 45) - 1;
181 if (memrefType.hasStaticShape() &&
182 !llvm::any_of(strides, ShapedType::isDynamic)) {
183 int64_t size = memrefType.getRank() == 0 ? 1 : 0;
185 for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i)
186 size = std::max(
shape[i] * strides[i], size);
187 size = size * elementByteWidth;
191 for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i) {
192 Value size = memrefDescriptor.
size(rewriter, loc, i);
193 Value stride = memrefDescriptor.
stride(rewriter, loc, i);
194 Value maxThisDim = LLVM::MulOp::create(rewriter, loc, size, stride);
196 ? LLVM::UMaxOp::create(rewriter, loc, maxIndex, maxThisDim)
201 return LLVM::MulOp::create(rewriter, loc, maxIndexI64, byteWidthConst);
207 Value cacheSwizzleStride =
nullptr,
208 unsigned addressSpace = 8) {
212 Type i16 = rewriter.getI16Type();
215 Value cacheStrideZext =
216 LLVM::ZExtOp::create(rewriter, loc, i16, cacheSwizzleStride);
217 Value swizzleBit = LLVM::ConstantOp::create(
218 rewriter, loc, i16, rewriter.getI16IntegerAttr(1 << 14));
219 stride = LLVM::OrOp::create(rewriter, loc, cacheStrideZext, swizzleBit,
222 stride = LLVM::ConstantOp::create(rewriter, loc, i16,
223 rewriter.getI16IntegerAttr(0));
252 flags |= (7 << 12) | (4 << 15);
255 uint32_t oob = boundsCheck ? 3 : 2;
256 flags |= (oob << 28);
264 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
265 Value resource = rewriter.createOrFold<ROCDL::MakeBufferRsrcOp>(
266 loc, rsrcType, basePointer, stride, numRecords, flagsConst);
271struct FatRawBufferCastLowering
273 FatRawBufferCastLowering(
const LLVMTypeConverter &converter, Chipset chipset)
274 : ConvertOpToLLVMPattern<FatRawBufferCastOp>(converter),
280 matchAndRewrite(FatRawBufferCastOp op, FatRawBufferCastOpAdaptor adaptor,
281 ConversionPatternRewriter &rewriter)
const override {
282 Location loc = op.getLoc();
283 Value memRef = adaptor.getSource();
284 Value unconvertedMemref = op.getSource();
285 MemRefType memrefType = cast<MemRefType>(unconvertedMemref.
getType());
286 MemRefDescriptor descriptor(memRef);
288 DataLayout dataLayout = DataLayout::closest(op);
289 int64_t elementByteWidth =
292 int64_t unusedOffset = 0;
293 SmallVector<int64_t, 5> strideVals;
294 if (
failed(memrefType.getStridesAndOffset(strideVals, unusedOffset)))
295 return op.emitOpError(
"Can't lower non-stride-offset memrefs");
297 Value numRecords = adaptor.getValidBytes();
300 getNumRecords(rewriter, loc, memrefType, descriptor, strideVals,
301 elementByteWidth, chipset, adaptor.getBoundsCheck());
304 adaptor.getResetOffset()
305 ? descriptor.bufferPtr(rewriter, loc, *getTypeConverter(),
307 : descriptor.alignedPtr(rewriter, loc);
309 Value offset = adaptor.getResetOffset()
310 ? LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
311 rewriter.getIndexAttr(0))
312 : descriptor.offset(rewriter, loc);
314 bool hasSizes = memrefType.getRank() > 0;
317 Value sizes = hasSizes
318 ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,
322 hasSizes ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,
327 rewriter, loc, basePointer, numRecords, adaptor.getBoundsCheck(),
328 chipset, adaptor.getCacheSwizzleStride(), 7);
330 Value
result = MemRefDescriptor::poison(
332 getTypeConverter()->convertType(op.getResult().getType()));
334 result = LLVM::InsertValueOp::create(rewriter, loc,
result, fatPtr, pos);
335 result = LLVM::InsertValueOp::create(rewriter, loc,
result, fatPtr,
337 result = LLVM::InsertValueOp::create(rewriter, loc,
result, offset,
340 result = LLVM::InsertValueOp::create(rewriter, loc,
result, sizes,
342 result = LLVM::InsertValueOp::create(rewriter, loc,
result, strides,
345 rewriter.replaceOp(op,
result);
351template <
typename GpuOp,
typename Intrinsic>
353 RawBufferOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
354 : ConvertOpToLLVMPattern<GpuOp>(converter), chipset(chipset) {}
357 static constexpr uint32_t maxVectorOpWidth = 128;
360 matchAndRewrite(GpuOp gpuOp,
typename GpuOp::Adaptor adaptor,
361 ConversionPatternRewriter &rewriter)
const override {
362 Location loc = gpuOp.getLoc();
363 Value memref = adaptor.getMemref();
364 Value unconvertedMemref = gpuOp.getMemref();
365 MemRefType memrefType = cast<MemRefType>(unconvertedMemref.
getType());
367 if (chipset.majorVersion < 9)
368 return gpuOp.emitOpError(
"raw buffer ops require GCN or higher");
370 Value storeData = adaptor.getODSOperands(0)[0];
371 if (storeData == memref)
375 wantedDataType = storeData.
getType();
377 wantedDataType = gpuOp.getODSResults(0)[0].getType();
379 Value atomicCmpData = Value();
382 Value maybeCmpData = adaptor.getODSOperands(1)[0];
383 if (maybeCmpData != memref)
384 atomicCmpData = maybeCmpData;
387 Type llvmWantedDataType = this->typeConverter->convertType(wantedDataType);
389 Type i32 = rewriter.getI32Type();
392 DataLayout dataLayout = DataLayout::closest(gpuOp);
393 int64_t elementByteWidth =
402 Type llvmBufferValType = llvmWantedDataType;
404 if (
auto floatType = dyn_cast<FloatType>(wantedDataType))
405 llvmBufferValType = this->getTypeConverter()->convertType(
406 rewriter.getIntegerType(floatType.getWidth()));
408 if (
auto dataVector = dyn_cast<VectorType>(wantedDataType)) {
409 uint32_t vecLen = dataVector.getNumElements();
412 uint32_t totalBits = elemBits * vecLen;
414 isa_and_present<RawBufferAtomicFaddOp>(*gpuOp) && vecLen == 2;
415 if (totalBits > maxVectorOpWidth)
416 return gpuOp.emitOpError(
417 "Total width of loads or stores must be no more than " +
418 Twine(maxVectorOpWidth) +
" bits, but we call for " +
420 " bits. This should've been caught in validation");
421 if (!usePackedFp16 && elemBits < 32) {
422 if (totalBits > 32) {
423 if (totalBits % 32 != 0)
424 return gpuOp.emitOpError(
"Load or store of more than 32-bits that "
425 "doesn't fit into words. Can't happen\n");
426 llvmBufferValType = this->typeConverter->convertType(
427 VectorType::get(totalBits / 32, i32));
429 llvmBufferValType = this->typeConverter->convertType(
430 rewriter.getIntegerType(totalBits));
434 if (
auto vecType = dyn_cast<VectorType>(llvmBufferValType)) {
437 if (vecType.getNumElements() == 1)
438 llvmBufferValType = vecType.getElementType();
441 SmallVector<Value, 6> args;
443 if (llvmBufferValType != llvmWantedDataType) {
444 Value castForStore = LLVM::BitcastOp::create(
445 rewriter, loc, llvmBufferValType, storeData);
446 args.push_back(castForStore);
448 args.push_back(storeData);
453 if (llvmBufferValType != llvmWantedDataType) {
454 Value castForCmp = LLVM::BitcastOp::create(
455 rewriter, loc, llvmBufferValType, atomicCmpData);
456 args.push_back(castForCmp);
458 args.push_back(atomicCmpData);
464 SmallVector<int64_t, 5> strides;
465 if (
failed(memrefType.getStridesAndOffset(strides, offset)))
466 return gpuOp.emitOpError(
"Can't lower non-stride-offset memrefs");
468 MemRefDescriptor memrefDescriptor(memref);
470 Value ptr = memrefDescriptor.bufferPtr(
471 rewriter, loc, *this->getTypeConverter(), memrefType);
473 getNumRecords(rewriter, loc, memrefType, memrefDescriptor, strides,
474 elementByteWidth, chipset, adaptor.getBoundsCheck());
476 adaptor.getBoundsCheck(), chipset);
477 args.push_back(resource);
481 adaptor.getIndices(), strides);
482 if (std::optional<int32_t> indexOffset = adaptor.getIndexOffset();
483 indexOffset && *indexOffset > 0) {
485 voffset = voffset ? LLVM::AddOp::create(rewriter, loc, voffset,
489 voffset = LLVM::MulOp::create(rewriter, loc, voffset, byteWidthConst);
490 args.push_back(voffset);
493 Value sgprOffset = adaptor.getSgprOffset();
496 sgprOffset = LLVM::MulOp::create(rewriter, loc, sgprOffset, byteWidthConst);
497 args.push_back(sgprOffset);
499 llvm::SmallVector<Type, 1> resultTypes(gpuOp->getNumResults(),
501 typename Intrinsic::Properties properties;
502 properties.aux = rewriter.getI32IntegerAttr(0);
504 Intrinsic::create(rewriter, loc, resultTypes, args, properties);
507 if (llvmBufferValType != llvmWantedDataType) {
508 replacement = LLVM::BitcastOp::create(rewriter, loc, llvmWantedDataType,
513 rewriter.eraseOp(gpuOp);
530static FailureOr<unsigned> encodeWaitcnt(
Chipset chipset,
unsigned vmcnt,
531 unsigned expcnt,
unsigned lgkmcnt) {
533 vmcnt = std::min(15u, vmcnt);
534 expcnt = std::min(7u, expcnt);
535 lgkmcnt = std::min(15u, lgkmcnt);
536 return vmcnt | (expcnt << 4) | (lgkmcnt << 8);
539 vmcnt = std::min(63u, vmcnt);
540 expcnt = std::min(7u, expcnt);
541 lgkmcnt = std::min(15u, lgkmcnt);
542 unsigned lowBits = vmcnt & 0xF;
543 unsigned highBits = (vmcnt >> 4) << 14;
544 unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);
545 return lowBits | highBits | otherCnts;
548 vmcnt = std::min(63u, vmcnt);
549 expcnt = std::min(7u, expcnt);
550 lgkmcnt = std::min(63u, lgkmcnt);
551 unsigned lowBits = vmcnt & 0xF;
552 unsigned highBits = (vmcnt >> 4) << 14;
553 unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);
554 return lowBits | highBits | otherCnts;
557 vmcnt = std::min(63u, vmcnt);
558 expcnt = std::min(7u, expcnt);
559 lgkmcnt = std::min(63u, lgkmcnt);
560 return (vmcnt << 10) | expcnt | (lgkmcnt << 4);
565struct MemoryCounterWaitOpLowering
575 matchAndRewrite(MemoryCounterWaitOp op, OpAdaptor adaptor,
576 ConversionPatternRewriter &rewriter)
const override {
579 if (std::optional<int> ds = adaptor.getDs())
580 ROCDL::WaitDscntOp::create(rewriter, loc, *ds);
582 if (std::optional<int>
load = adaptor.getLoad())
583 ROCDL::WaitLoadcntOp::create(rewriter, loc, *
load);
585 if (std::optional<int> store = adaptor.getStore())
586 ROCDL::WaitStorecntOp::create(rewriter, loc, *store);
588 if (std::optional<int> exp = adaptor.getExp())
589 ROCDL::WaitExpcntOp::create(rewriter, loc, *exp);
591 if (std::optional<int>
tensor = adaptor.getTensor())
592 ROCDL::WaitTensorcntOp::create(rewriter, loc, *
tensor);
594 rewriter.eraseOp(op);
598 if (adaptor.getTensor())
599 return op.emitOpError(
"unsupported chipset");
601 auto getVal = [](
Attribute attr) ->
unsigned {
603 return cast<IntegerAttr>(attr).getInt();
608 unsigned ds = getVal(adaptor.getDsAttr());
609 unsigned exp = getVal(adaptor.getExpAttr());
611 unsigned vmcnt = 1024;
615 vmcnt = getVal(
load) + getVal(store);
617 vmcnt = getVal(
load);
619 vmcnt = getVal(store);
622 FailureOr<unsigned> waitcnt = encodeWaitcnt(chipset, vmcnt, exp, ds);
624 return op.emitOpError(
"unsupported chipset");
626 rewriter.replaceOpWithNewOp<ROCDL::SWaitcntOp>(op, *waitcnt);
632 LDSBarrierOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
633 : ConvertOpToLLVMPattern<LDSBarrierOp>(converter), chipset(chipset) {}
638 matchAndRewrite(LDSBarrierOp op, LDSBarrierOp::Adaptor adaptor,
639 ConversionPatternRewriter &rewriter)
const override {
640 Location loc = op.getLoc();
643 bool requiresInlineAsm = chipset <
kGfx90a;
646 rewriter.getAttr<LLVM::MMRATagAttr>(
"amdgpu-synchronize-as",
"local");
655 StringRef scope =
"workgroup";
657 auto relFence = LLVM::FenceOp::create(rewriter, loc,
658 LLVM::AtomicOrdering::release, scope);
659 relFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);
660 if (requiresInlineAsm) {
661 auto asmDialectAttr = LLVM::AsmDialectAttr::get(rewriter.getContext(),
662 LLVM::AsmDialect::AD_ATT);
663 const char *asmStr =
";;;WARNING: BREAKS DEBUG WATCHES\ns_barrier";
664 const char *constraints =
"";
665 LLVM::InlineAsmOp::create(
668 asmStr, constraints,
true,
669 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)
2262 if (numElements != 16)
2263 return emitNumElementsError(16,
"gfx1250+");
2265 ROCDL::DsLoadTr6_B96::create(rewriter, loc, rocdlResultType, srcPtr)
2270 if (numElements != 8)
2271 return emitNumElementsError(8,
"gfx1250+");
2273 ROCDL::DsLoadTr8_B64::create(rewriter, loc, rocdlResultType, srcPtr)
2278 if (numElements != 8)
2279 return emitNumElementsError(8,
"gfx1250+");
2280 intrinsic = ROCDL::DsLoadTr16_B128::create(rewriter, loc,
2281 rocdlResultType, srcPtr)
2286 return op.emitOpError(
"Unsupported element size for transpose load");
2289 switch (elementTypeSize) {
2291 if (numElements != 16)
2292 return emitNumElementsError(16,
"gfx950");
2293 intrinsic = ROCDL::ds_read_tr4_b64::create(rewriter, loc,
2294 rocdlResultType, srcPtr)
2299 if (numElements != 16)
2300 return emitNumElementsError(16,
"gfx950");
2301 intrinsic = ROCDL::ds_read_tr6_b96::create(rewriter, loc,
2302 rocdlResultType, srcPtr)
2307 if (numElements != 8)
2308 return emitNumElementsError(8,
"gfx950");
2309 intrinsic = ROCDL::ds_read_tr8_b64::create(rewriter, loc,
2310 rocdlResultType, srcPtr)
2315 if (numElements != 4)
2316 return emitNumElementsError(4,
"gfx950");
2317 intrinsic = ROCDL::ds_read_tr16_b64::create(rewriter, loc,
2318 rocdlResultType, srcPtr)
2323 return op.emitOpError(
"Unsupported element size for transpose load");
2327 assert(intrinsic &&
"expected ROCDL transpose load intrinsic");
2328 if (intrinsic.
getType() == llvmResultType) {
2329 rewriter.replaceOp(op, intrinsic);
2332 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, intrinsic);
2337struct GlobalTransposeLoadOpLowering
2339 GlobalTransposeLoadOpLowering(
const LLVMTypeConverter &converter,
2341 : ConvertOpToLLVMPattern<GlobalTransposeLoadOp>(converter),
2347 matchAndRewrite(GlobalTransposeLoadOp op,
2348 GlobalTransposeLoadOpAdaptor adaptor,
2349 ConversionPatternRewriter &rewriter)
const override {
2351 return op.emitOpError(
2352 "global_transpose_load is only supported on gfx1200+");
2354 Location loc = op.getLoc();
2355 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2356 auto resultType = cast<VectorType>(op.getResult().getType());
2359 rewriter, loc, srcMemRefType, adaptor.getSrc(), adaptor.getSrcIndices(),
2360 LLVM::GEPNoWrapFlags::inbounds | LLVM::GEPNoWrapFlags::nuw);
2362 size_t numElements = resultType.getNumElements();
2363 size_t elementTypeSize =
2368 Type rocdlResultType =
2369 elementTypeSize < 16
2370 ? VectorType::get((numElements * elementTypeSize) / 32,
2371 rewriter.getIntegerType(32))
2372 : typeConverter->convertType(resultType);
2373 Type llvmResultType = typeConverter->convertType(resultType);
2375 switch (elementTypeSize) {
2377 assert(numElements == 16);
2379 return op.emitOpError(
"4-bit global_transpose_load requires gfx1250+");
2380 auto rocdlOp = ROCDL::GlobalLoadTr4_B64::create(rewriter, loc,
2381 rocdlResultType, srcPtr);
2382 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2386 assert(numElements == 16);
2388 return op.emitOpError(
"6-bit global_transpose_load requires gfx1250+");
2389 auto rocdlOp = ROCDL::GlobalLoadTr6_B96::create(rewriter, loc,
2390 rocdlResultType, srcPtr);
2391 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2395 assert(numElements == 8);
2396 auto rocdlOp = ROCDL::GlobalLoadTr8_B64::create(rewriter, loc,
2397 rocdlResultType, srcPtr);
2398 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2402 assert(numElements == 8);
2403 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadTr8_B128>(op, llvmResultType,
2408 return op.emitOpError(
2409 "unsupported element size for global transpose load");
2416 GatherToLDSOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2417 : ConvertOpToLLVMPattern<GatherToLDSOp>(converter), chipset(chipset) {}
2422 matchAndRewrite(GatherToLDSOp op, GatherToLDSOpAdaptor adaptor,
2423 ConversionPatternRewriter &rewriter)
const override {
2424 if (chipset.majorVersion < 9 || chipset.majorVersion > 10)
2425 return op.emitOpError(
"pre-gfx9 and post-gfx10 not supported");
2427 Location loc = op.getLoc();
2429 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2430 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2435 Type transferType = op.getTransferType();
2436 int loadWidth = [&]() ->
int {
2437 if (
auto transferVectorType = dyn_cast<VectorType>(transferType)) {
2438 return (transferVectorType.getNumElements() *
2439 transferVectorType.getElementTypeBitWidth()) /
2446 if (!llvm::is_contained({1, 2, 4, 12, 16}, loadWidth))
2447 return op.emitOpError(
"chipset unsupported element size");
2449 if (chipset !=
kGfx950 && llvm::is_contained({12, 16}, loadWidth))
2450 return op.emitOpError(
"Gather to LDS instructions with 12-byte and "
2451 "16-byte load widths are only supported on gfx950");
2455 (adaptor.getSrcIndices()));
2458 (adaptor.getDstIndices()));
2460 if (op.getAsync()) {
2461 rewriter.replaceOpWithNewOp<ROCDL::LoadAsyncToLDSOp>(
2462 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2463 rewriter.getI32IntegerAttr(0),
2467 rewriter.replaceOpWithNewOp<ROCDL::LoadToLDSOp>(
2468 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2469 rewriter.getI32IntegerAttr(0),
2478struct GlobalLoadAsyncToLDSOpLowering
2480 GlobalLoadAsyncToLDSOpLowering(
const LLVMTypeConverter &converter,
2482 : ConvertOpToLLVMPattern<GlobalLoadAsyncToLDSOp>(converter),
2488 matchAndRewrite(GlobalLoadAsyncToLDSOp op,
2489 GlobalLoadAsyncToLDSOpAdaptor adaptor,
2490 ConversionPatternRewriter &rewriter)
const override {
2492 return op.emitOpError(
2493 "global_load_async_to_lds is only supported on gfx1250+");
2495 Location loc = op.getLoc();
2496 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2497 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2499 Type transferType = op.getTransferType();
2501 isa<VectorType>(transferType)
2502 ? cast<VectorType>(transferType).getNumElements() *
2503 cast<VectorType>(transferType).getElementTypeBitWidth()
2508 adaptor.getSrcIndices());
2511 adaptor.getDstIndices());
2514 Value mask = adaptor.getMask();
2515 int64_t nullptrVal =
2516 llvm::AMDGPU::getNullPointerValue(llvm::AMDGPUAS::LOCAL_ADDRESS);
2520 LLVM::IntToPtrOp::create(rewriter, loc, dstPtr.
getType(), nullInt);
2521 dstPtr = LLVM::SelectOp::create(rewriter, loc, mask, dstPtr, nullPtr);
2524 auto offset = rewriter.getI32IntegerAttr(0);
2525 Attribute aux = rewriter.getI32IntegerAttr(0);
2527 switch (transferBits) {
2529 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB8Op>(
2534 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB32Op>(
2539 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB64Op>(
2544 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB128Op>(
2549 return op.emitOpError(
"unsupported transfer width");
2556struct ExtPackedFp8OpLowering final
2558 ExtPackedFp8OpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2559 : ConvertOpToLLVMPattern<amdgpu::ExtPackedFp8Op>(converter),
2564 matchAndRewrite(ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2565 ConversionPatternRewriter &rewriter)
const override;
2568struct ScaledExtPackedMatrixOpLowering final
2570 ScaledExtPackedMatrixOpLowering(
const LLVMTypeConverter &converter,
2572 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedMatrixOp>(converter),
2577 matchAndRewrite(ScaledExtPackedMatrixOp op,
2578 ScaledExtPackedMatrixOpAdaptor adaptor,
2579 ConversionPatternRewriter &rewriter)
const override;
2582struct PackedTrunc2xFp8OpLowering final
2584 PackedTrunc2xFp8OpLowering(
const LLVMTypeConverter &converter,
2586 : ConvertOpToLLVMPattern<amdgpu::PackedTrunc2xFp8Op>(converter),
2591 matchAndRewrite(PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
2592 ConversionPatternRewriter &rewriter)
const override;
2595struct PackedStochRoundFp8OpLowering final
2597 PackedStochRoundFp8OpLowering(
const LLVMTypeConverter &converter,
2599 : ConvertOpToLLVMPattern<amdgpu::PackedStochRoundFp8Op>(converter),
2604 matchAndRewrite(PackedStochRoundFp8Op op,
2605 PackedStochRoundFp8OpAdaptor adaptor,
2606 ConversionPatternRewriter &rewriter)
const override;
2609struct ScaledExtPackedOpLowering final
2611 ScaledExtPackedOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
2612 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedOp>(converter),
2617 matchAndRewrite(ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2618 ConversionPatternRewriter &rewriter)
const override;
2621struct PackedScaledTruncOpLowering final
2623 PackedScaledTruncOpLowering(
const LLVMTypeConverter &converter,
2625 : ConvertOpToLLVMPattern<amdgpu::PackedScaledTruncOp>(converter),
2630 matchAndRewrite(PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2631 ConversionPatternRewriter &rewriter)
const override;
2636LogicalResult ExtPackedFp8OpLowering::matchAndRewrite(
2637 ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2638 ConversionPatternRewriter &rewriter)
const {
2639 Location loc = op.getLoc();
2641 return rewriter.notifyMatchFailure(
2642 loc,
"Fp8 conversion instructions are not available on target "
2643 "architecture and their emulation is not implemented");
2645 getTypeConverter()->convertType(VectorType::get(4, rewriter.getI8Type()));
2646 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2647 Type f32 = getTypeConverter()->convertType(op.getResult().getType());
2649 Value source = adaptor.getSource();
2650 auto sourceVecType = dyn_cast<VectorType>(op.getSource().getType());
2651 auto resultVecType = dyn_cast<VectorType>(op.getResult().getType());
2654 if (!sourceVecType || sourceVecType.getNumElements() < 4) {
2655 Value longVec = LLVM::UndefOp::create(rewriter, loc, v4i8);
2656 if (!sourceVecType) {
2657 longVec = LLVM::InsertElementOp::create(
2660 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2662 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2664 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2669 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2670 if (resultVecType) {
2672 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Bf8Op>(op, f32, i32Source,
2675 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Fp8Op>(op, f32, i32Source,
2680 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Bf8Op>(op, f32, i32Source,
2683 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Fp8Op>(op, f32, i32Source,
2690int32_t getScaleSel(int32_t blockSize,
unsigned bitWidth, int32_t scaleWaveHalf,
2691 int32_t firstScaleByte) {
2697 assert(llvm::is_contained({16, 32}, blockSize));
2698 assert(llvm::is_contained({4u, 6u, 8u}, bitWidth));
2700 const bool isFp8 = bitWidth == 8;
2701 const bool isBlock16 = blockSize == 16;
2704 int32_t bit0 = isBlock16;
2705 assert(llvm::is_contained({0, 1, 2}, firstScaleByte));
2706 int32_t bit1 = (firstScaleByte == 2) << 1;
2707 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2708 int32_t bit2 = scaleWaveHalf << 2;
2709 return bit2 | bit1 | bit0;
2712 int32_t bit0 = isBlock16;
2714 assert(llvm::is_contained({0, 1, 2, 3}, firstScaleByte));
2715 int32_t bits2and1 = firstScaleByte << 1;
2716 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2717 int32_t bit3 = scaleWaveHalf << 3;
2718 int32_t bits = bit3 | bits2and1 | bit0;
2720 assert(!llvm::is_contained(
2721 {0b0011, 0b0101, 0b0111, 0b1000, 0b1001, 0b1011, 0b1111}, bits));
2725static std::optional<StringRef>
2726scaledExtPacked816ToIntrinsic(Type srcElemType, Type destElemType) {
2727 using fp4 = Float4E2M1FNType;
2728 using fp8 = Float8E4M3FNType;
2729 using bf8 = Float8E5M2Type;
2730 using fp6 = Float6E2M3FNType;
2731 using bf6 = Float6E3M2FNType;
2732 if (isa<fp4>(srcElemType)) {
2733 if (destElemType.
isF16())
2734 return ROCDL::CvtPkScalePk8F16Fp4Op::getOperationName();
2735 if (destElemType.
isBF16())
2736 return ROCDL::CvtPkScalePk8Bf16Fp4Op::getOperationName();
2737 if (destElemType.
isF32())
2738 return ROCDL::CvtPkScalePk8F32Fp4Op::getOperationName();
2739 return std::nullopt;
2741 if (isa<fp8>(srcElemType)) {
2742 if (destElemType.
isF16())
2743 return ROCDL::CvtPkScalePk8F16Fp8Op::getOperationName();
2744 if (destElemType.
isBF16())
2745 return ROCDL::CvtPkScalePk8Bf16Fp8Op::getOperationName();
2746 if (destElemType.
isF32())
2747 return ROCDL::CvtPkScalePk8F32Fp8Op::getOperationName();
2748 return std::nullopt;
2750 if (isa<bf8>(srcElemType)) {
2751 if (destElemType.
isF16())
2752 return ROCDL::CvtPkScalePk8F16Bf8Op::getOperationName();
2753 if (destElemType.
isBF16())
2754 return ROCDL::CvtPkScalePk8Bf16Bf8Op::getOperationName();
2755 if (destElemType.
isF32())
2756 return ROCDL::CvtPkScalePk8F32Bf8Op::getOperationName();
2757 return std::nullopt;
2759 if (isa<fp6>(srcElemType)) {
2760 if (destElemType.
isF16())
2761 return ROCDL::CvtPkScalePk16F16Fp6Op::getOperationName();
2762 if (destElemType.
isBF16())
2763 return ROCDL::CvtPkScalePk16Bf16Fp6Op::getOperationName();
2764 if (destElemType.
isF32())
2765 return ROCDL::CvtPkScalePk16F32Fp6Op::getOperationName();
2766 return std::nullopt;
2768 if (isa<bf6>(srcElemType)) {
2769 if (destElemType.
isF16())
2770 return ROCDL::CvtPkScalePk16F16Bf6Op::getOperationName();
2771 if (destElemType.
isBF16())
2772 return ROCDL::CvtPkScalePk16Bf16Bf6Op::getOperationName();
2773 if (destElemType.
isF32())
2774 return ROCDL::CvtPkScalePk16F32Bf6Op::getOperationName();
2775 return std::nullopt;
2777 llvm_unreachable(
"invalid combination of element types for packed conversion "
2781LogicalResult ScaledExtPackedMatrixOpLowering::matchAndRewrite(
2782 ScaledExtPackedMatrixOp op, ScaledExtPackedMatrixOpAdaptor adaptor,
2783 ConversionPatternRewriter &rewriter)
const {
2784 using fp4 = Float4E2M1FNType;
2785 using fp8 = Float8E4M3FNType;
2786 using bf8 = Float8E5M2Type;
2787 using fp6 = Float6E2M3FNType;
2788 using bf6 = Float6E3M2FNType;
2789 Location loc = op.getLoc();
2791 return rewriter.notifyMatchFailure(
2793 "Scaled fp packed conversion instructions are not available on target "
2794 "architecture and their emulation is not implemented");
2798 int32_t scaleWaveHalf = op.getFirstScaleLane() / 16;
2799 int32_t firstScaleByte = op.getFirstScaleByte();
2800 int32_t blockSize = op.getBlockSize();
2801 auto sourceType = cast<VectorType>(op.getSource().getType());
2802 auto srcElemType = cast<FloatType>(sourceType.getElementType());
2803 unsigned bitWidth = srcElemType.getWidth();
2805 auto targetType = cast<VectorType>(op.getResult().getType());
2806 auto destElemType = cast<FloatType>(targetType.getElementType());
2808 IntegerType i32 = rewriter.getI32Type();
2809 Value source = adaptor.getSource();
2810 Type llvmResultType = typeConverter->convertType(op.getResult().getType());
2811 Type packedType =
nullptr;
2812 if (isa<fp4>(srcElemType)) {
2814 packedType = getTypeConverter()->convertType(packedType);
2815 }
else if (isa<fp8, bf8>(srcElemType)) {
2816 packedType = VectorType::get(2, i32);
2817 packedType = getTypeConverter()->convertType(packedType);
2818 }
else if (isa<fp6, bf6>(srcElemType)) {
2819 packedType = VectorType::get(3, i32);
2820 packedType = getTypeConverter()->convertType(packedType);
2822 llvm_unreachable(
"invalid element type for packed scaled ext");
2825 if (!packedType || !llvmResultType) {
2826 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2829 std::optional<StringRef> maybeIntrinsic =
2830 scaledExtPacked816ToIntrinsic(srcElemType, destElemType);
2831 if (!maybeIntrinsic.has_value())
2832 return op.emitOpError(
2833 "no intrinsic matching packed scaled conversion on the given chipset");
2836 getScaleSel(blockSize, bitWidth, scaleWaveHalf, firstScaleByte);
2838 LLVM::BitcastOp::create(rewriter, loc, i32, adaptor.getScale());
2839 Value castedSource =
2840 LLVM::BitcastOp::create(rewriter, loc, packedType, source);
2842 OperationState loweredOp(loc, *maybeIntrinsic);
2843 loweredOp.addTypes({llvmResultType});
2844 loweredOp.addOperands({castedSource, castedScale});
2846 SmallVector<NamedAttribute, 1> attrs;
2848 NamedAttribute(
"scaleSel", rewriter.getI32IntegerAttr(scaleSel)));
2850 loweredOp.addAttributes(attrs);
2851 Operation *lowered = rewriter.create(loweredOp);
2852 rewriter.replaceOp(op, lowered);
2857LogicalResult ScaledExtPackedOpLowering::matchAndRewrite(
2858 ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2859 ConversionPatternRewriter &rewriter)
const {
2860 Location loc = op.getLoc();
2862 return rewriter.notifyMatchFailure(
2863 loc,
"Scaled fp conversion instructions are not available on target "
2864 "architecture and their emulation is not implemented");
2865 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2867 Value source = adaptor.getSource();
2868 Value scale = adaptor.getScale();
2870 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2871 Type sourceElemType = sourceVecType.getElementType();
2872 VectorType destVecType = cast<VectorType>(op.getResult().getType());
2873 Type destElemType = destVecType.getElementType();
2875 VectorType packedVecType;
2876 if (isa<Float8E5M2Type, Float8E4M3FNType>(sourceElemType)) {
2877 VectorType v4i8 = VectorType::get(4, rewriter.getI8Type());
2878 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v4i8));
2879 }
else if (isa<Float4E2M1FNType>(sourceElemType)) {
2880 VectorType v8i4 = VectorType::get(8, rewriter.getI4Type());
2881 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v8i4));
2883 llvm_unreachable(
"invalid element type for scaled ext");
2887 if (sourceVecType.getNumElements() < packedVecType.getNumElements()) {
2888 Value longVec = LLVM::ZeroOp::create(rewriter, loc, packedVecType);
2889 if (!sourceVecType) {
2890 longVec = LLVM::InsertElementOp::create(
2893 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2895 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2897 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2902 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2904 if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isF32())
2905 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Bf8Op>(
2906 op, destVecType, i32Source, scale, op.getIndex());
2907 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isF16())
2908 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Bf8Op>(
2909 op, destVecType, i32Source, scale, op.getIndex());
2910 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.
isBF16())
2911 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Bf8Op>(
2912 op, destVecType, i32Source, scale, op.getIndex());
2913 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isF32())
2914 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp8Op>(
2915 op, destVecType, i32Source, scale, op.getIndex());
2916 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isF16())
2917 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp8Op>(
2918 op, destVecType, i32Source, scale, op.getIndex());
2919 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.
isBF16())
2920 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp8Op>(
2921 op, destVecType, i32Source, scale, op.getIndex());
2922 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isF32())
2923 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp4Op>(
2924 op, destVecType, i32Source, scale, op.getIndex());
2925 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isF16())
2926 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp4Op>(
2927 op, destVecType, i32Source, scale, op.getIndex());
2928 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.
isBF16())
2929 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp4Op>(
2930 op, destVecType, i32Source, scale, op.getIndex());
2937LogicalResult PackedScaledTruncOpLowering::matchAndRewrite(
2938 PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2939 ConversionPatternRewriter &rewriter)
const {
2940 Location loc = op.getLoc();
2942 return rewriter.notifyMatchFailure(
2943 loc,
"Scaled fp conversion instructions are not available on target "
2944 "architecture and their emulation is not implemented");
2945 Type v2i16 = getTypeConverter()->convertType(
2946 VectorType::get(2, rewriter.getI16Type()));
2947 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2949 Type resultType = op.getResult().getType();
2951 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2952 Type sourceElemType = sourceVecType.getElementType();
2954 Type intResultType = isa<Float4E2M1FNType>(resultElemType) ? i32 : v2i16;
2956 Value source = adaptor.getSource();
2957 Value scale = adaptor.getScale();
2958 Value existing = adaptor.getExisting();
2960 existing = LLVM::BitcastOp::create(rewriter, loc, intResultType, existing);
2962 existing = LLVM::ZeroOp::create(rewriter, loc, intResultType);
2964 if (sourceVecType.getNumElements() < 2) {
2966 Value elem0 = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2967 VectorType v2 = VectorType::get(2, sourceElemType);
2968 source = LLVM::ZeroOp::create(rewriter, loc, v2);
2969 source = LLVM::InsertElementOp::create(rewriter, loc, source, elem0, c0);
2972 Value sourceA, sourceB;
2973 if (sourceElemType.
isF32()) {
2976 sourceA = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2977 sourceB = LLVM::ExtractElementOp::create(rewriter, loc, source, c1);
2981 if (sourceElemType.
isF32() && isa<Float8E5M2Type>(resultElemType))
2982 result = ROCDL::CvtScaleF32PkBf8F32Op::create(rewriter, loc, intResultType,
2983 existing, sourceA, sourceB,
2984 scale, op.getIndex());
2985 else if (sourceElemType.
isF16() && isa<Float8E5M2Type>(resultElemType))
2986 result = ROCDL::CvtScaleF32PkBf8F16Op::create(
2987 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2988 else if (sourceElemType.
isBF16() && isa<Float8E5M2Type>(resultElemType))
2989 result = ROCDL::CvtScaleF32PkBf8Bf16Op::create(
2990 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2991 else if (sourceElemType.
isF32() && isa<Float8E4M3FNType>(resultElemType))
2992 result = ROCDL::CvtScaleF32PkFp8F32Op::create(rewriter, loc, intResultType,
2993 existing, sourceA, sourceB,
2994 scale, op.getIndex());
2995 else if (sourceElemType.
isF16() && isa<Float8E4M3FNType>(resultElemType))
2996 result = ROCDL::CvtScaleF32PkFp8F16Op::create(
2997 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2998 else if (sourceElemType.
isBF16() && isa<Float8E4M3FNType>(resultElemType))
2999 result = ROCDL::CvtScaleF32PkFp8Bf16Op::create(
3000 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3001 else if (sourceElemType.
isF32() && isa<Float4E2M1FNType>(resultElemType))
3002 result = ROCDL::CvtScaleF32PkFp4F32Op::create(rewriter, loc, intResultType,
3003 existing, sourceA, sourceB,
3004 scale, op.getIndex());
3005 else if (sourceElemType.
isF16() && isa<Float4E2M1FNType>(resultElemType))
3006 result = ROCDL::CvtScaleF32PkFp4F16Op::create(
3007 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3008 else if (sourceElemType.
isBF16() && isa<Float4E2M1FNType>(resultElemType))
3009 result = ROCDL::CvtScaleF32PkFp4Bf16Op::create(
3010 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3014 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3015 op, getTypeConverter()->convertType(resultType),
result);
3019LogicalResult PackedTrunc2xFp8OpLowering::matchAndRewrite(
3020 PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
3021 ConversionPatternRewriter &rewriter)
const {
3022 Location loc = op.getLoc();
3024 return rewriter.notifyMatchFailure(
3025 loc,
"Fp8 conversion instructions are not available on target "
3026 "architecture and their emulation is not implemented");
3027 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3029 Type resultType = op.getResult().getType();
3032 Value sourceA = adaptor.getSourceA();
3033 Value sourceB = adaptor.getSourceB();
3035 sourceB = LLVM::UndefOp::create(rewriter, loc, sourceA.
getType());
3036 Value existing = adaptor.getExisting();
3038 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3040 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3044 result = ROCDL::CvtPkBf8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3045 existing, op.getWordIndex());
3047 result = ROCDL::CvtPkFp8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3048 existing, op.getWordIndex());
3050 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3051 op, getTypeConverter()->convertType(resultType),
result);
3055LogicalResult PackedStochRoundFp8OpLowering::matchAndRewrite(
3056 PackedStochRoundFp8Op op, PackedStochRoundFp8OpAdaptor adaptor,
3057 ConversionPatternRewriter &rewriter)
const {
3058 Location loc = op.getLoc();
3060 return rewriter.notifyMatchFailure(
3061 loc,
"Fp8 conversion instructions are not available on target "
3062 "architecture and their emulation is not implemented");
3063 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3065 Type resultType = op.getResult().getType();
3068 Value source = adaptor.getSource();
3069 Value stoch = adaptor.getStochiasticParam();
3070 Value existing = adaptor.getExisting();
3072 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3074 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3078 result = ROCDL::CvtSrBf8F32Op::create(rewriter, loc, i32, source, stoch,
3079 existing, op.getStoreIndex());
3081 result = ROCDL::CvtSrFp8F32Op::create(rewriter, loc, i32, source, stoch,
3082 existing, op.getStoreIndex());
3084 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3085 op, getTypeConverter()->convertType(resultType),
result);
3091struct AMDGPUDPPLowering :
public ConvertOpToLLVMPattern<DPPOp> {
3092 AMDGPUDPPLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3093 : ConvertOpToLLVMPattern<DPPOp>(converter), chipset(chipset) {}
3097 matchAndRewrite(DPPOp DppOp, DPPOp::Adaptor adaptor,
3098 ConversionPatternRewriter &rewriter)
const override {
3101 Location loc = DppOp.getLoc();
3102 Value src = adaptor.getSrc();
3103 Value old = adaptor.getOld();
3106 Type llvmType =
nullptr;
3108 llvmType = rewriter.getI32Type();
3109 }
else if (isa<FloatType>(srcType)) {
3111 ? rewriter.getF32Type()
3112 : rewriter.getF64Type();
3113 }
else if (isa<IntegerType>(srcType)) {
3115 ? rewriter.getI32Type()
3116 : rewriter.getI64Type();
3118 auto llvmSrcIntType = typeConverter->convertType(
3122 auto convertOperand = [&](Value operand, Type operandType) {
3123 if (operandType.getIntOrFloatBitWidth() <= 16) {
3124 if (llvm::isa<FloatType>(operandType)) {
3126 LLVM::BitcastOp::create(rewriter, loc, llvmSrcIntType, operand);
3128 auto llvmVecType = typeConverter->convertType(mlir::VectorType::get(
3129 32 / operandType.getIntOrFloatBitWidth(), llvmSrcIntType));
3130 Value undefVec = LLVM::UndefOp::create(rewriter, loc, llvmVecType);
3132 LLVM::InsertElementOp::create(rewriter, loc, undefVec, operand,
3134 operand = LLVM::BitcastOp::create(rewriter, loc, llvmType, operand);
3139 src = convertOperand(src, srcType);
3140 old = convertOperand(old, oldType);
3143 enum DppCtrl :
unsigned {
3152 ROW_HALF_MIRROR = 0x141,
3157 auto kind = DppOp.getKind();
3158 auto permArgument = DppOp.getPermArgument();
3159 uint32_t DppCtrl = 0;
3163 case DPPPerm::quad_perm: {
3164 auto quadPermAttr = cast<ArrayAttr>(*permArgument);
3166 for (
auto elem : quadPermAttr.getAsRange<IntegerAttr>()) {
3167 uint32_t num = elem.getInt();
3168 DppCtrl |= num << (i * 2);
3173 case DPPPerm::row_shl: {
3174 auto intAttr = cast<IntegerAttr>(*permArgument);
3175 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHL0;
3178 case DPPPerm::row_shr: {
3179 auto intAttr = cast<IntegerAttr>(*permArgument);
3180 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHR0;
3183 case DPPPerm::row_ror: {
3184 auto intAttr = cast<IntegerAttr>(*permArgument);
3185 DppCtrl = intAttr.getInt() + DppCtrl::ROW_ROR0;
3188 case DPPPerm::wave_shl:
3189 DppCtrl = DppCtrl::WAVE_SHL1;
3191 case DPPPerm::wave_shr:
3192 DppCtrl = DppCtrl::WAVE_SHR1;
3194 case DPPPerm::wave_rol:
3195 DppCtrl = DppCtrl::WAVE_ROL1;
3197 case DPPPerm::wave_ror:
3198 DppCtrl = DppCtrl::WAVE_ROR1;
3200 case DPPPerm::row_mirror:
3201 DppCtrl = DppCtrl::ROW_MIRROR;
3203 case DPPPerm::row_half_mirror:
3204 DppCtrl = DppCtrl::ROW_HALF_MIRROR;
3206 case DPPPerm::row_bcast_15:
3207 DppCtrl = DppCtrl::BCAST15;
3209 case DPPPerm::row_bcast_31:
3210 DppCtrl = DppCtrl::BCAST31;
3216 auto rowMask = DppOp->getAttrOfType<IntegerAttr>(
"row_mask").getInt();
3217 auto bankMask = DppOp->getAttrOfType<IntegerAttr>(
"bank_mask").getInt();
3218 bool boundCtrl = DppOp->getAttrOfType<BoolAttr>(
"bound_ctrl").getValue();
3222 ROCDL::DPPUpdateOp::create(rewriter, loc, llvmType, old, src, DppCtrl,
3223 rowMask, bankMask, boundCtrl);
3225 Value
result = dppMovOp.getRes();
3227 result = LLVM::TruncOp::create(rewriter, loc, llvmSrcIntType,
result);
3228 if (!llvm::isa<IntegerType>(srcType)) {
3229 result = LLVM::BitcastOp::create(rewriter, loc, srcType,
result);
3240struct AMDGPUSwizzleBitModeLowering
3241 :
public ConvertOpToLLVMPattern<SwizzleBitModeOp> {
3245 matchAndRewrite(SwizzleBitModeOp op, OpAdaptor adaptor,
3246 ConversionPatternRewriter &rewriter)
const override {
3247 Location loc = op.getLoc();
3248 Type i32 = rewriter.getI32Type();
3249 Value src = adaptor.getSrc();
3250 SmallVector<Value> decomposed;
3252 return rewriter.notifyMatchFailure(op,
3253 "failed to decompose value to i32");
3254 unsigned andMask = op.getAndMask();
3255 unsigned orMask = op.getOrMask();
3256 unsigned xorMask = op.getXorMask();
3260 unsigned mask = andMask | (orMask << 5) | (xorMask << 10);
3262 SmallVector<Value> swizzled;
3263 for (Value v : decomposed) {
3265 ROCDL::DsSwizzleOp::create(rewriter, loc, v.getType(), v, maskValue);
3266 swizzled.emplace_back(res);
3270 rewriter.replaceOp(op,
result);
3275struct AMDGPUPermlaneLowering :
public ConvertOpToLLVMPattern<PermlaneSwapOp> {
3278 AMDGPUPermlaneLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3279 : ConvertOpToLLVMPattern<PermlaneSwapOp>(converter), chipset(chipset) {}
3283 matchAndRewrite(PermlaneSwapOp op, OpAdaptor adaptor,
3284 ConversionPatternRewriter &rewriter)
const override {
3286 return op->emitOpError(
"permlane_swap is only supported on gfx950+");
3288 Location loc = op.getLoc();
3289 Type i32 = rewriter.getI32Type();
3290 Value src = adaptor.getSrc();
3291 unsigned rowLength = op.getRowLength();
3292 bool fi = op.getFetchInactive();
3293 bool boundctrl = op.getBoundCtrl();
3295 SmallVector<Value> decomposed;
3297 return rewriter.notifyMatchFailure(op,
3298 "failed to decompose value to i32");
3300 SmallVector<Value> permuted;
3301 for (Value v : decomposed) {
3303 Type i32pair = LLVM::LLVMStructType::getLiteral(
3304 rewriter.getContext(), {v.getType(), v.getType()});
3306 if (rowLength == 16)
3307 res = ROCDL::Permlane16SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3309 else if (rowLength == 32)
3310 res = ROCDL::Permlane32SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3313 llvm_unreachable(
"unsupported row length");
3315 Value vdst0 = LLVM::ExtractValueOp::create(rewriter, loc, res, {0});
3316 Value vdst1 = LLVM::ExtractValueOp::create(rewriter, loc, res, {1});
3318 Value isEqual = LLVM::ICmpOp::create(rewriter, loc,
3319 LLVM::ICmpPredicate::eq, vdst0, v);
3324 LLVM::SelectOp::create(rewriter, loc, isEqual, vdst1, vdst0);
3325 permuted.emplace_back(vdstNew);
3329 rewriter.replaceOp(op,
result);
3334struct AMDGPUPermlaneVarLowering
3335 :
public ConvertOpToLLVMPattern<PermlaneVarOp> {
3338 AMDGPUPermlaneVarLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3339 : ConvertOpToLLVMPattern<PermlaneVarOp>(converter), chipset(chipset) {}
3343 matchAndRewrite(PermlaneVarOp op, OpAdaptor adaptor,
3344 ConversionPatternRewriter &rewriter)
const override {
3346 return op->emitOpError(
"permlane_var is only supported on GFX12+");
3348 Location loc = op.getLoc();
3349 Type i32 = rewriter.getI32Type();
3350 Value src = adaptor.getSrc();
3351 Value selector = adaptor.getSelector();
3352 bool cross = op.getCross();
3353 bool fi = op.getFetchInactive();
3354 bool boundCtrl = op.getBoundCtrl();
3356 SmallVector<Value> decomposed;
3358 return rewriter.notifyMatchFailure(op,
3359 "failed to decompose value to i32");
3361 SmallVector<Value> permuted;
3362 for (Value v : decomposed) {
3365 res = ROCDL::PermlaneX16VarOp::create(rewriter, loc, i32, v, v,
3366 selector, fi, boundCtrl);
3368 res = ROCDL::Permlane16VarOp::create(rewriter, loc, i32, v, v, selector,
3370 permuted.emplace_back(res);
3374 rewriter.replaceOp(op,
result);
3387constexpr int32_t kDsBarrierPendingCountBitWidth = 29;
3388constexpr int32_t kDsBarrierPhasePos = kDsBarrierPendingCountBitWidth;
3389constexpr int32_t kDsBarrierInitCountPos = 32;
3390constexpr int32_t kDsBarrierPendingCountMask =
3391 (1 << kDsBarrierPendingCountBitWidth) - 1;
3393struct DsBarrierInitOpLowering
3394 :
public ConvertOpToLLVMPattern<DsBarrierInitOp> {
3397 DsBarrierInitOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3398 : ConvertOpToLLVMPattern<DsBarrierInitOp>(converter), chipset(chipset) {}
3401 matchAndRewrite(DsBarrierInitOp op, OpAdaptor adaptor,
3402 ConversionPatternRewriter &rewriter)
const override {
3404 return op->emitOpError(
"only supported on gfx1250+");
3406 Location loc = op.getLoc();
3407 Type i64 = rewriter.getI64Type();
3409 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3411 adaptor.getBase(), adaptor.getIndices());
3418 LLVM::SubOp::create(rewriter, loc, adaptor.getParticipants(),
3425 Value maskedCount32 =
3426 LLVM::AndOp::create(rewriter, loc, initCount, countMask);
3427 Value maskedCount = LLVM::ZExtOp::create(rewriter, loc, i64, maskedCount32);
3429 Value initCountShifted = LLVM::ShlOp::create(
3430 rewriter, loc, maskedCount,
3432 Value barrierState =
3433 LLVM::OrOp::create(rewriter, loc, initCountShifted, maskedCount);
3435 LLVM::StoreOp::create(
3436 rewriter, loc, barrierState, ptr, 8,
false,
3438 false, LLVM::AtomicOrdering::release,
3441 rewriter.eraseOp(op);
3446struct DsBarrierPollStateOpLowering
3447 :
public ConvertOpToLLVMPattern<DsBarrierPollStateOp> {
3450 DsBarrierPollStateOpLowering(
const LLVMTypeConverter &converter,
3452 : ConvertOpToLLVMPattern<DsBarrierPollStateOp>(converter),
3456 matchAndRewrite(DsBarrierPollStateOp op, OpAdaptor adaptor,
3457 ConversionPatternRewriter &rewriter)
const override {
3459 return op->emitOpError(
"only supported on gfx1250+");
3461 Location loc = op.getLoc();
3462 Type i64 = rewriter.getI64Type();
3464 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3466 adaptor.getBase(), adaptor.getIndices());
3470 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
3471 op, i64, ptr, 8,
false,
3473 false, LLVM::AtomicOrdering::acquire,
3479struct DsAsyncBarrierArriveOpLowering
3480 :
public ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp> {
3483 DsAsyncBarrierArriveOpLowering(
const LLVMTypeConverter &converter,
3485 : ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp>(converter),
3489 matchAndRewrite(DsAsyncBarrierArriveOp op, OpAdaptor adaptor,
3490 ConversionPatternRewriter &rewriter)
const override {
3492 return op->emitOpError(
"only supported on gfx1250+");
3494 Location loc = op.getLoc();
3496 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3498 adaptor.getBase(), adaptor.getIndices());
3500 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicAsyncBarrierArriveOp>(
3501 op, ptr,
nullptr,
nullptr,
3507struct DsBarrierArriveOpLowering
3508 :
public ConvertOpToLLVMPattern<DsBarrierArriveOp> {
3511 DsBarrierArriveOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3512 : ConvertOpToLLVMPattern<DsBarrierArriveOp>(converter), chipset(chipset) {
3516 matchAndRewrite(DsBarrierArriveOp op, OpAdaptor adaptor,
3517 ConversionPatternRewriter &rewriter)
const override {
3519 return op->emitOpError(
"only supported on gfx1250+");
3521 Location loc = op.getLoc();
3522 Type i64 = rewriter.getI64Type();
3524 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3526 adaptor.getBase(), adaptor.getIndices());
3528 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicBarrierArriveRtnOp>(
3529 op, i64, ptr, adaptor.getCount(),
nullptr,
3535struct DsBarrierStatePhaseOpLowering
3536 :
public ConvertOpToLLVMPattern<DsBarrierStatePhaseOp> {
3540 matchAndRewrite(DsBarrierStatePhaseOp op, OpAdaptor adaptor,
3541 ConversionPatternRewriter &rewriter)
const override {
3542 Location loc = op.getLoc();
3543 Type i32 = rewriter.getI32Type();
3545 Value state = adaptor.getState();
3547 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3548 Value phase = LLVM::LShrOp::create(
3549 rewriter, loc, noInitCount,
3552 rewriter.replaceOp(op, phase);
3557struct DsBarrierStatePendingCountOpLowering
3558 :
public ConvertOpToLLVMPattern<DsBarrierStatePendingCountOp> {
3562 matchAndRewrite(DsBarrierStatePendingCountOp op, OpAdaptor adaptor,
3563 ConversionPatternRewriter &rewriter)
const override {
3564 Location loc = op.getLoc();
3565 Type i32 = rewriter.getI32Type();
3567 Value state = adaptor.getState();
3569 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3570 Value pendingCount = LLVM::AndOp::create(
3571 rewriter, loc, noInitCount,
3573 static_cast<uint32_t
>(kDsBarrierPendingCountMask)));
3575 rewriter.replaceOp(op, pendingCount);
3580struct DsBarrierStateInitCountOpLowering
3581 :
public ConvertOpToLLVMPattern<DsBarrierStateInitCountOp> {
3585 matchAndRewrite(DsBarrierStateInitCountOp op, OpAdaptor adaptor,
3586 ConversionPatternRewriter &rewriter)
const override {
3587 Location loc = op.getLoc();
3588 Type i32 = rewriter.getI32Type();
3590 Value state = adaptor.getState();
3592 Value initCountI64 = LLVM::LShrOp::create(
3593 rewriter, loc, state,
3595 Value initCount = LLVM::TruncOp::create(rewriter, loc, i32, initCountI64);
3597 rewriter.replaceOp(op, initCount);
3602struct DsBarrierStatePhaseParityLowering
3603 :
public ConvertOpToLLVMPattern<DsBarrierStatePhaseParity> {
3607 matchAndRewrite(DsBarrierStatePhaseParity op, OpAdaptor adaptor,
3608 ConversionPatternRewriter &rewriter)
const override {
3609 Location loc = op.getLoc();
3610 Type i1 = rewriter.getI1Type();
3612 Value state = adaptor.getState();
3615 LLVM::TruncOp::create(rewriter, loc, rewriter.getI32Type(), state);
3616 Value phase = LLVM::LShrOp::create(
3617 rewriter, loc, noInitCount,
3619 Value parity = LLVM::TruncOp::create(rewriter, loc, i1, phase);
3621 rewriter.replaceOp(op, parity);
3630static Value setValueAtOffset(ConversionPatternRewriter &rewriter, Location loc,
3631 Value accumulator, Value value, int64_t shift) {
3636 value = LLVM::ShlOp::create(rewriter, loc, value, shiftAmount);
3642 constexpr bool isDisjoint =
true;
3643 return LLVM::OrOp::create(rewriter, loc, accumulator, value, isDisjoint);
3646template <
typename BaseOp>
3647struct AMDGPUMakeDmaBaseLowering :
public ConvertOpToLLVMPattern<BaseOp> {
3648 using ConvertOpToLLVMPattern<BaseOp>::ConvertOpToLLVMPattern;
3651 AMDGPUMakeDmaBaseLowering(
const LLVMTypeConverter &converter, Chipset chipset)
3652 : ConvertOpToLLVMPattern<BaseOp>(converter), chipset(chipset) {}
3656 matchAndRewrite(BaseOp op, Adaptor adaptor,
3657 ConversionPatternRewriter &rewriter)
const override {
3659 return op->emitOpError(
"make_dma_base is only supported on gfx1250");
3661 Location loc = op.getLoc();
3663 constexpr int32_t constlen = 4;
3664 Value consts[constlen];
3665 for (int64_t i = 0; i < constlen; ++i)
3668 constexpr int32_t sgprslen = constlen;
3669 Value sgprs[sgprslen];
3670 for (int64_t i = 0; i < sgprslen; ++i) {
3671 sgprs[i] = consts[0];
3674 sgprs[0] = consts[1];
3676 if constexpr (BaseOp::isGather()) {
3677 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 30);
3679 auto type = cast<TDMGatherBaseType>(op.getResult().getType());
3680 Type indexType = type.getIndexType();
3682 assert(llvm::is_contained({16u, 32u}, indexSize) &&
3683 "expected index_size to be 16 or 32");
3684 unsigned idx = (indexSize / 16) - 1;
3687 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 31);
3690 ValueRange ldsIndices = adaptor.getLdsIndices();
3691 Value lds = adaptor.getLds();
3692 auto ldsMemRefType = cast<MemRefType>(op.getLds().getType());
3695 rewriter, loc, ldsMemRefType, lds, ldsIndices);
3697 ValueRange globalIndices = adaptor.getGlobalIndices();
3698 Value global = adaptor.getGlobal();
3699 auto globalMemRefType = cast<MemRefType>(op.getGlobal().getType());
3702 rewriter, loc, globalMemRefType, global, globalIndices);
3704 Type i32 = rewriter.getI32Type();
3705 Type i64 = rewriter.getI64Type();
3707 sgprs[1] = LLVM::PtrToIntOp::create(rewriter, loc, i32, ldsPtr);
3708 Value castForGlobalAddr =
3709 LLVM::PtrToIntOp::create(rewriter, loc, i64, globalPtr);
3711 sgprs[2] = LLVM::TruncOp::create(rewriter, loc, i32, castForGlobalAddr);
3713 Value shift = LLVM::LShrOp::create(rewriter, loc, castForGlobalAddr,
3716 Value highHalf = LLVM::TruncOp::create(rewriter, loc, i32, shift);
3719 highHalf = LLVM::AndOp::create(rewriter, loc, highHalf, mask);
3721 sgprs[3] = setValueAtOffset(rewriter, loc, highHalf, consts[2], 30);
3723 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
3724 assert(v4i32 &&
"expected type conversion to succeed");
3725 Value
result = LLVM::PoisonOp::create(rewriter, loc, v4i32);
3727 for (
auto [sgpr, constant] : llvm::zip_equal(sgprs, consts))
3729 LLVM::InsertElementOp::create(rewriter, loc,
result, sgpr, constant);
3731 rewriter.replaceOp(op,
result);
3736template <
typename DescriptorOp>
3737struct AMDGPULowerDescriptor :
public ConvertOpToLLVMPattern<DescriptorOp> {
3738 using ConvertOpToLLVMPattern<DescriptorOp>::ConvertOpToLLVMPattern;
3741 AMDGPULowerDescriptor(
const LLVMTypeConverter &converter, Chipset chipset)
3742 : ConvertOpToLLVMPattern<DescriptorOp>(converter), chipset(chipset) {}
3745 Value getDGroup0(OpAdaptor adaptor)
const {
return adaptor.getBase(); }
3747 Value setWorkgroupMask(DescriptorOp op, OpAdaptor adaptor,
3748 ConversionPatternRewriter &rewriter, Location loc,
3749 Value sgpr0)
const {
3750 Value mask = op.getWorkgroupMask();
3754 Type i16 = rewriter.getI16Type();
3755 mask = LLVM::BitcastOp::create(rewriter, loc, i16, mask);
3756 Type i32 = rewriter.getI32Type();
3757 Value extendedMask = LLVM::ZExtOp::create(rewriter, loc, i32, mask);
3758 return setValueAtOffset(rewriter, loc, sgpr0, extendedMask, 0);
3761 Value setDataSize(DescriptorOp op, OpAdaptor adaptor,
3762 ConversionPatternRewriter &rewriter, Location loc,
3763 Value sgpr0, ArrayRef<Value> consts)
const {
3764 unsigned elementTypeWidthInBits = op.getElementTypeWidth();
3765 assert(llvm::is_contained({8u, 16u, 32u, 64u}, elementTypeWidthInBits) &&
3766 "expected type width to be 8, 16, 32, or 64.");
3767 int64_t idx = llvm::Log2_32(elementTypeWidthInBits / 8);
3768 Value size = consts[idx];
3769 return setValueAtOffset(rewriter, loc, sgpr0, size, 16);
3772 Value setAtomicBarrier(DescriptorOp op, OpAdaptor adaptor,
3773 ConversionPatternRewriter &rewriter, Location loc,
3774 Value sgpr0, ArrayRef<Value> consts)
const {
3775 if (!adaptor.getAtomicBarrierAddress())
3778 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 18);
3781 Value setIterateEnable(DescriptorOp op, OpAdaptor adaptor,
3782 ConversionPatternRewriter &rewriter, Location loc,
3783 Value sgpr0, ArrayRef<Value> consts)
const {
3784 if (!adaptor.getGlobalIncrement())
3789 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 19);
3792 Value setPadEnable(DescriptorOp op, OpAdaptor adaptor,
3793 ConversionPatternRewriter &rewriter, Location loc,
3794 Value sgpr0, ArrayRef<Value> consts)
const {
3795 if (!op.getPadAmount())
3798 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 20);
3801 Value setEarlyTimeout(DescriptorOp op, OpAdaptor adaptor,
3802 ConversionPatternRewriter &rewriter, Location loc,
3803 Value sgpr0, ArrayRef<Value> consts)
const {
3804 if (!op.getWorkgroupMask())
3807 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 21);
3810 Value setPadInterval(DescriptorOp op, OpAdaptor adaptor,
3811 ConversionPatternRewriter &rewriter, Location loc,
3812 Value sgpr0, ArrayRef<Value> consts)
const {
3813 if (!op.getPadAmount())
3822 IntegerType i32 = rewriter.getI32Type();
3823 Value padInterval = adaptor.getPadInterval();
3824 padInterval = LLVM::CountTrailingZerosOp::create(rewriter, loc, i32,
3825 padInterval,
false);
3826 padInterval = LLVM::SubOp::create(rewriter, loc, padInterval, consts[1]);
3828 return setValueAtOffset(rewriter, loc, sgpr0, padInterval, 22);
3831 Value setPadAmount(DescriptorOp op, OpAdaptor adaptor,
3832 ConversionPatternRewriter &rewriter, Location loc,
3833 Value sgpr0, ArrayRef<Value> consts)
const {
3834 if (!op.getPadAmount())
3843 Value padAmount = adaptor.getPadAmount();
3844 padAmount = LLVM::SubOp::create(rewriter, loc, padAmount, consts[1]);
3846 return setValueAtOffset(rewriter, loc, sgpr0, padAmount, 25);
3849 Value setAtomicBarrierAddress(DescriptorOp op, OpAdaptor adaptor,
3850 ConversionPatternRewriter &rewriter,
3851 Location loc, Value sgpr1,
3852 ArrayRef<Value> consts)
const {
3853 if (!adaptor.getAtomicBarrierAddress())
3856 Value atomicBarrierAddress = adaptor.getAtomicBarrierAddress();
3857 auto barrierAddressTy =
3858 cast<MemRefType>(op.getAtomicBarrierAddress().getType());
3859 ValueRange atomicBarrierIndices = adaptor.getAtomicBarrierIndices();
3861 rewriter, loc, barrierAddressTy, atomicBarrierAddress,
3862 atomicBarrierIndices);
3863 IntegerType i32 = rewriter.getI32Type();
3869 atomicBarrierAddress =
3870 LLVM::PtrToIntOp::create(rewriter, loc, i32, atomicBarrierAddress);
3871 atomicBarrierAddress =
3872 LLVM::LShrOp::create(rewriter, loc, atomicBarrierAddress, consts[3]);
3874 atomicBarrierAddress =
3875 LLVM::AndOp::create(rewriter, loc, atomicBarrierAddress, mask);
3876 return setValueAtOffset(rewriter, loc, sgpr1, atomicBarrierAddress, 32);
3879 std::pair<Value, Value> setTensorDimX(DescriptorOp op, OpAdaptor adaptor,
3880 ConversionPatternRewriter &rewriter,
3881 Location loc, Value sgpr1, Value sgpr2,
3882 ArrayRef<Value> consts, uint64_t dimX,
3883 uint32_t offset)
const {
3884 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
3885 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
3886 SmallVector<OpFoldResult> mixedGlobalSizes =
3888 if (mixedGlobalSizes.size() <= dimX)
3889 return {sgpr1, sgpr2};
3891 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
3898 if (
auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
3902 IntegerType i32 = rewriter.getI32Type();
3903 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
3904 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
3907 sgpr1 = setValueAtOffset(rewriter, loc, sgpr1, tensorDimX, offset);
3910 Value tensorDimXHigh = LLVM::LShrOp::create(rewriter, loc, tensorDimX, c16);
3911 sgpr2 = setValueAtOffset(rewriter, loc, sgpr2, tensorDimXHigh, offset + 16);
3912 return {sgpr1, sgpr2};
3915 std::pair<Value, Value> setTensorDim0(DescriptorOp op, OpAdaptor adaptor,
3916 ConversionPatternRewriter &rewriter,
3917 Location loc, Value sgpr1, Value sgpr2,
3918 ArrayRef<Value> consts)
const {
3919 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, 0,
3923 std::pair<Value, Value> setTensorDim1(DescriptorOp op, OpAdaptor adaptor,
3924 ConversionPatternRewriter &rewriter,
3925 Location loc, Value sgpr2, Value sgpr3,
3926 ArrayRef<Value> consts)
const {
3927 return setTensorDimX(op, adaptor, rewriter, loc, sgpr2, sgpr3, consts, 1,
3931 Value setTileDimX(DescriptorOp op, OpAdaptor adaptor,
3932 ConversionPatternRewriter &rewriter, Location loc,
3933 Value sgpr, ArrayRef<Value> consts,
size_t dimX,
3934 int64_t offset)
const {
3935 ArrayRef<int64_t> sharedStaticSizes = adaptor.getSharedStaticSizes();
3936 ValueRange sharedDynamicSizes = adaptor.getSharedDynamicSizes();
3937 SmallVector<OpFoldResult> mixedSharedSizes =
3939 if (mixedSharedSizes.size() <= dimX)
3942 OpFoldResult tileDimXOpFoldResult = *(mixedSharedSizes.rbegin() + dimX);
3951 if (
auto attr = dyn_cast<Attribute>(tileDimXOpFoldResult)) {
3955 IntegerType i32 = rewriter.getI32Type();
3956 tileDimX = cast<Value>(tileDimXOpFoldResult);
3957 tileDimX = LLVM::TruncOp::create(rewriter, loc, i32, tileDimX);
3960 return setValueAtOffset(rewriter, loc, sgpr, tileDimX, offset);
3963 Value setTileDim0(DescriptorOp op, OpAdaptor adaptor,
3964 ConversionPatternRewriter &rewriter, Location loc,
3965 Value sgpr3, ArrayRef<Value> consts)
const {
3966 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, 0, 112);
3969 Value setTileDim1(DescriptorOp op, OpAdaptor adaptor,
3970 ConversionPatternRewriter &rewriter, Location loc,
3971 Value sgpr4, ArrayRef<Value> consts)
const {
3972 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 1, 128);
3975 Value setValidIndices(DescriptorOp op, OpAdaptor adaptor,
3976 ConversionPatternRewriter &rewriter, Location loc,
3977 Value sgpr4, ArrayRef<Value> consts)
const {
3978 auto type = cast<VectorType>(op.getIndices().getType());
3979 ArrayRef<int64_t> shape = type.getShape();
3980 assert(shape.size() == 1 &&
"expected shape to be of rank 1.");
3981 unsigned length = shape.back();
3982 assert(0 < length && length <= 16 &&
"expected length to be at most 16.");
3984 return setValueAtOffset(rewriter, loc, sgpr4, value, 128);
3987 Value setTileDim1OrValidIndices(DescriptorOp op, OpAdaptor adaptor,
3988 ConversionPatternRewriter &rewriter,
3989 Location loc, Value sgpr4,
3990 ArrayRef<Value> consts)
const {
3991 if constexpr (DescriptorOp::isGather())
3992 return setValidIndices(op, adaptor, rewriter, loc, sgpr4, consts);
3993 return setTileDim1(op, adaptor, rewriter, loc, sgpr4, consts);
3996 Value setTileDim2(DescriptorOp op, OpAdaptor adaptor,
3997 ConversionPatternRewriter &rewriter, Location loc,
3998 Value sgpr4, ArrayRef<Value> consts)
const {
4000 if constexpr (DescriptorOp::isGather())
4002 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 2, 144);
4005 std::pair<Value, Value>
4006 setTensorDimXStride(DescriptorOp op, OpAdaptor adaptor,
4007 ConversionPatternRewriter &rewriter, Location loc,
4008 Value sgprY, Value sgprZ, ArrayRef<Value> consts,
4009 size_t dimX, int64_t offset)
const {
4010 ArrayRef<int64_t> globalStaticStrides = adaptor.getGlobalStaticStrides();
4011 ValueRange globalDynamicStrides = adaptor.getGlobalDynamicStrides();
4012 SmallVector<OpFoldResult> mixedGlobalStrides =
4013 getMixedValues(globalStaticStrides, globalDynamicStrides, rewriter);
4015 if (mixedGlobalStrides.size() <= (dimX + 1))
4016 return {sgprY, sgprZ};
4018 OpFoldResult tensorDimXStrideOpFoldResult =
4019 *(mixedGlobalStrides.rbegin() + dimX + 1);
4024 Value tensorDimXStride;
4025 if (
auto attr = dyn_cast<Attribute>(tensorDimXStrideOpFoldResult))
4029 tensorDimXStride = cast<Value>(tensorDimXStrideOpFoldResult);
4031 constexpr int64_t first48bits = (1ll << 48) - 1;
4034 LLVM::AndOp::create(rewriter, loc, mask, tensorDimXStride);
4035 IntegerType i32 = rewriter.getI32Type();
4036 Value tensorDimXStrideLow =
4037 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStride);
4038 sgprY = setValueAtOffset(rewriter, loc, sgprY, tensorDimXStrideLow, offset);
4040 int64_t shift = (offset % 32) == 0 ? 32 : offset % 32;
4042 Value tensorDimXStrideHigh =
4043 LLVM::LShrOp::create(rewriter, loc, tensorDimXStride, shiftVal);
4044 tensorDimXStrideHigh =
4045 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStrideHigh);
4046 sgprZ = setValueAtOffset(rewriter, loc, sgprZ, tensorDimXStrideHigh,
4048 return {sgprY, sgprZ};
4051 std::pair<Value, Value>
4052 setTensorDim0Stride(DescriptorOp op, OpAdaptor adaptor,
4053 ConversionPatternRewriter &rewriter, Location loc,
4054 Value sgpr5, Value sgpr6, ArrayRef<Value> consts)
const {
4055 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4059 std::pair<Value, Value>
4060 setTensorDim1Stride(DescriptorOp op, OpAdaptor adaptor,
4061 ConversionPatternRewriter &rewriter, Location loc,
4062 Value sgpr5, Value sgpr6, ArrayRef<Value> consts)
const {
4064 if constexpr (DescriptorOp::isGather())
4065 return {sgpr5, sgpr6};
4066 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4070 Value getDGroup1(DescriptorOp op, OpAdaptor adaptor,
4071 ConversionPatternRewriter &rewriter, Location loc,
4072 ArrayRef<Value> consts)
const {
4074 for (int64_t i = 0; i < 8; ++i) {
4075 sgprs[i] = consts[0];
4078 sgprs[0] = setWorkgroupMask(op, adaptor, rewriter, loc, sgprs[0]);
4079 sgprs[0] = setDataSize(op, adaptor, rewriter, loc, sgprs[0], consts);
4080 sgprs[0] = setAtomicBarrier(op, adaptor, rewriter, loc, sgprs[0], consts);
4081 sgprs[0] = setIterateEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4082 sgprs[0] = setPadEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4083 sgprs[0] = setEarlyTimeout(op, adaptor, rewriter, loc, sgprs[0], consts);
4084 sgprs[0] = setPadInterval(op, adaptor, rewriter, loc, sgprs[0], consts);
4085 sgprs[0] = setPadAmount(op, adaptor, rewriter, loc, sgprs[0], consts);
4088 setAtomicBarrierAddress(op, adaptor, rewriter, loc, sgprs[1], consts);
4089 std::tie(sgprs[1], sgprs[2]) =
4090 setTensorDim0(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4091 std::tie(sgprs[2], sgprs[3]) =
4092 setTensorDim1(op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4094 sgprs[3] = setTileDim0(op, adaptor, rewriter, loc, sgprs[3], consts);
4096 setTileDim1OrValidIndices(op, adaptor, rewriter, loc, sgprs[4], consts);
4097 sgprs[4] = setTileDim2(op, adaptor, rewriter, loc, sgprs[4], consts);
4098 std::tie(sgprs[5], sgprs[6]) = setTensorDim0Stride(
4099 op, adaptor, rewriter, loc, sgprs[5], sgprs[6], consts);
4100 std::tie(sgprs[6], sgprs[7]) = setTensorDim1Stride(
4101 op, adaptor, rewriter, loc, sgprs[6], sgprs[7], consts);
4103 IntegerType i32 = rewriter.getI32Type();
4104 Type v8i32 = this->typeConverter->convertType(VectorType::get(8, i32));
4105 assert(v8i32 &&
"expected type conversion to succeed");
4106 Value dgroup1 = LLVM::PoisonOp::create(rewriter, loc, v8i32);
4108 for (
auto [sgpr, constant] : llvm::zip_equal(sgprs, consts)) {
4110 LLVM::InsertElementOp::create(rewriter, loc, dgroup1, sgpr, constant);
4116 Value setTensorDimX(DescriptorOp op, OpAdaptor adaptor,
4117 ConversionPatternRewriter &rewriter, Location loc,
4118 Value sgpr0, ArrayRef<Value> consts, int64_t dimX,
4119 int64_t offset)
const {
4120 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
4121 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
4122 SmallVector<OpFoldResult> mixedGlobalSizes =
4124 if (mixedGlobalSizes.size() <=
static_cast<unsigned long>(dimX))
4127 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
4129 if (
auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
4133 IntegerType i32 = rewriter.getI32Type();
4134 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
4135 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
4138 return setValueAtOffset(rewriter, loc, sgpr0, tensorDimX, offset);
4141 Value setTensorDim2(DescriptorOp op, OpAdaptor adaptor,
4142 ConversionPatternRewriter &rewriter, Location loc,
4143 Value sgpr0, ArrayRef<Value> consts)
const {
4144 return setTensorDimX(op, adaptor, rewriter, loc, sgpr0, consts, 2, 0);
4147 Value truncateAndSetValueAtOffset(ConversionPatternRewriter &rewriter,
4148 Location loc, Value accumulator,
4149 Value value, int64_t shift)
const {
4151 IntegerType i32 = rewriter.getI32Type();
4152 value = LLVM::TruncOp::create(rewriter, loc, i32, value);
4153 return setValueAtOffset(rewriter, loc, accumulator, value, shift);
4156 Value setLDSAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4157 ConversionPatternRewriter &rewriter, Location loc,
4158 Value sgpr1, ArrayRef<Value> consts,
4159 int64_t offset)
const {
4160 Value ldsAddrIncrement = adaptor.getLdsIncrement();
4161 return setValueAtOffset(rewriter, loc, sgpr1, ldsAddrIncrement, offset);
4164 std::pair<Value, Value>
4165 setGlobalAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4166 ConversionPatternRewriter &rewriter, Location loc,
4167 Value sgpr2, Value sgpr3, ArrayRef<Value> consts,
4168 int64_t offset)
const {
4169 Value globalAddrIncrement = adaptor.getGlobalIncrement();
4170 sgpr2 = truncateAndSetValueAtOffset(rewriter, loc, sgpr2,
4171 globalAddrIncrement, offset);
4173 globalAddrIncrement =
4174 LLVM::LShrOp::create(rewriter, loc, globalAddrIncrement, shift);
4175 constexpr int64_t first16BitsHigh = (1ll << 16) - 1;
4176 sgpr3 = truncateAndSetValueAtOffset(rewriter, loc, sgpr3,
4177 globalAddrIncrement, offset + 32);
4179 sgpr3 = LLVM::AndOp::create(rewriter, loc, sgpr3, mask);
4180 return {sgpr2, sgpr3};
4183 Value setTensorDim3OrLDSAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4184 ConversionPatternRewriter &rewriter,
4185 Location loc, Value sgpr1,
4186 ArrayRef<Value> consts)
const {
4187 Value ldsIncrement = op.getLdsIncrement();
4188 constexpr int64_t dim = 3;
4189 constexpr int64_t offset = 32;
4191 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, consts, dim,
4193 return setLDSAddrIncrement(op, adaptor, rewriter, loc, sgpr1, consts,
4197 std::pair<Value, Value> setTensorDim2StrideOrGlobalAddrIncrement(
4198 DescriptorOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter,
4199 Location loc, Value sgpr2, Value sgpr3, ArrayRef<Value> consts)
const {
4200 Value globalIncrement = op.getGlobalIncrement();
4201 constexpr int32_t dim = 2;
4202 constexpr int32_t offset = 64;
4203 if (!globalIncrement)
4204 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4205 consts, dim, offset);
4206 return setGlobalAddrIncrement(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4210 Value setIterateCount(DescriptorOp op, OpAdaptor adaptor,
4211 ConversionPatternRewriter &rewriter, Location loc,
4212 Value sgpr3, ArrayRef<Value> consts,
4213 int32_t offset)
const {
4214 Value iterationCount = adaptor.getIterationCount();
4215 IntegerType i32 = rewriter.getI32Type();
4222 iterationCount = LLVM::TruncOp::create(rewriter, loc, i32, iterationCount);
4224 LLVM::SubOp::create(rewriter, loc, iterationCount, consts[1]);
4225 return setValueAtOffset(rewriter, loc, sgpr3, iterationCount, offset);
4228 Value setTileDim3OrIterateCount(DescriptorOp op, OpAdaptor adaptor,
4229 ConversionPatternRewriter &rewriter,
4230 Location loc, Value sgpr3,
4231 ArrayRef<Value> consts)
const {
4232 Value iterateCount = op.getIterationCount();
4233 constexpr int32_t dim = 2;
4234 constexpr int32_t offset = 112;
4236 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, dim,
4239 return setIterateCount(op, adaptor, rewriter, loc, sgpr3, consts, offset);
4242 Value getDGroup2(DescriptorOp op, OpAdaptor adaptor,
4243 ConversionPatternRewriter &rewriter, Location loc,
4244 ArrayRef<Value> consts)
const {
4245 if constexpr (DescriptorOp::isGather())
4246 return getDGroup2Gather(op, adaptor, rewriter, loc, consts);
4247 return getDGroup2NonGather(op, adaptor, rewriter, loc, consts);
4250 Value getDGroup2NonGather(DescriptorOp op, OpAdaptor adaptor,
4251 ConversionPatternRewriter &rewriter, Location loc,
4252 ArrayRef<Value> consts)
const {
4253 IntegerType i32 = rewriter.getI32Type();
4254 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4255 assert(v4i32 &&
"expected type conversion to succeed.");
4257 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4258 if (onlyNeedsTwoDescriptors)
4259 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4261 constexpr int64_t sgprlen = 4;
4262 Value sgprs[sgprlen];
4263 for (
int i = 0; i < sgprlen; ++i)
4264 sgprs[i] = consts[0];
4266 sgprs[0] = setTensorDim2(op, adaptor, rewriter, loc, sgprs[0], consts);
4267 sgprs[1] = setTensorDim3OrLDSAddrIncrement(op, adaptor, rewriter, loc,
4269 std::tie(sgprs[2], sgprs[3]) = setTensorDim2StrideOrGlobalAddrIncrement(
4270 op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4272 setTileDim3OrIterateCount(op, adaptor, rewriter, loc, sgprs[3], consts);
4274 Value dgroup2 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4275 for (
auto [sgpr, constant] : llvm::zip(sgprs, consts))
4277 LLVM::InsertElementOp::create(rewriter, loc, dgroup2, sgpr, constant);
4282 Value getGatherIndices(DescriptorOp op, OpAdaptor adaptor,
4283 ConversionPatternRewriter &rewriter, Location loc,
4284 ArrayRef<Value> consts,
bool firstHalf)
const {
4285 IntegerType i32 = rewriter.getI32Type();
4286 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4287 assert(v4i32 &&
"expected type conversion to succeed.");
4289 Value
indices = adaptor.getIndices();
4290 auto vectorType = cast<VectorType>(
indices.getType());
4291 unsigned length = vectorType.getShape().back();
4292 Type elementType = vectorType.getElementType();
4293 unsigned maxLength = elementType == i32 ? 4 : 8;
4294 int32_t offset = firstHalf ? 0 : maxLength;
4295 unsigned discountedLength =
4296 std::max(
static_cast<int32_t
>(length - offset), 0);
4298 unsigned targetSize = std::min(maxLength, discountedLength);
4300 SmallVector<Value> indicesVector;
4301 for (
unsigned i = offset; i < targetSize + offset; ++i) {
4303 if (i < consts.size())
4307 Value elem = LLVM::ExtractElementOp::create(rewriter, loc,
indices, idx);
4308 indicesVector.push_back(elem);
4311 SmallVector<Value> indicesI32Vector;
4312 if (elementType == i32) {
4313 indicesI32Vector = indicesVector;
4315 for (
unsigned i = 0; i < targetSize; ++i) {
4316 Value index = indicesVector[i];
4317 indicesI32Vector.push_back(
4318 LLVM::ZExtOp::create(rewriter, loc, i32, index));
4320 if ((targetSize % 2) != 0)
4322 indicesI32Vector.push_back(consts[0]);
4325 SmallVector<Value> indicesToInsert;
4326 if (elementType == i32) {
4327 indicesToInsert = indicesI32Vector;
4329 unsigned size = indicesI32Vector.size() / 2;
4330 for (
unsigned i = 0; i < size; ++i) {
4331 Value first = indicesI32Vector[2 * i];
4332 Value second = indicesI32Vector[2 * i + 1];
4333 Value joined = setValueAtOffset(rewriter, loc, first, second, 16);
4334 indicesToInsert.push_back(joined);
4338 Value dgroup = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4339 for (
auto [sgpr, constant] : llvm::zip_first(indicesToInsert, consts))
4341 LLVM::InsertElementOp::create(rewriter, loc, dgroup, sgpr, constant);
4346 Value getDGroup2Gather(DescriptorOp op, OpAdaptor adaptor,
4347 ConversionPatternRewriter &rewriter, Location loc,
4348 ArrayRef<Value> consts)
const {
4349 return getGatherIndices(op, adaptor, rewriter, loc, consts,
true);
4352 std::pair<Value, Value>
4353 setTensorDim3Stride(DescriptorOp op, OpAdaptor adaptor,
4354 ConversionPatternRewriter &rewriter, Location loc,
4355 Value sgpr0, Value sgpr1, ArrayRef<Value> consts)
const {
4356 constexpr int32_t dim = 3;
4357 constexpr int32_t offset = 0;
4358 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr0, sgpr1, consts,
4362 std::pair<Value, Value> setTensorDim4(DescriptorOp op, OpAdaptor adaptor,
4363 ConversionPatternRewriter &rewriter,
4364 Location loc, Value sgpr1, Value sgpr2,
4365 ArrayRef<Value> consts)
const {
4366 constexpr int32_t dim = 4;
4367 constexpr int32_t offset = 48;
4368 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, dim,
4372 Value setTileDim4(DescriptorOp op, OpAdaptor adaptor,
4373 ConversionPatternRewriter &rewriter, Location loc,
4374 Value sgpr2, ArrayRef<Value> consts)
const {
4375 constexpr int32_t dim = 4;
4376 constexpr int32_t offset = 80;
4377 return setTileDimX(op, adaptor, rewriter, loc, sgpr2, consts, dim, offset);
4380 Value getDGroup3(DescriptorOp op, OpAdaptor adaptor,
4381 ConversionPatternRewriter &rewriter, Location loc,
4382 ArrayRef<Value> consts)
const {
4383 if constexpr (DescriptorOp::isGather())
4384 return getDGroup3Gather(op, adaptor, rewriter, loc, consts);
4385 return getDGroup3NonGather(op, adaptor, rewriter, loc, consts);
4388 Value getDGroup3NonGather(DescriptorOp op, OpAdaptor adaptor,
4389 ConversionPatternRewriter &rewriter, Location loc,
4390 ArrayRef<Value> consts)
const {
4391 IntegerType i32 = rewriter.getI32Type();
4392 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4393 assert(v4i32 &&
"expected type conversion to succeed.");
4394 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4395 if (onlyNeedsTwoDescriptors)
4396 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4398 constexpr int32_t sgprlen = 4;
4399 Value sgprs[sgprlen];
4400 for (
int i = 0; i < sgprlen; ++i)
4401 sgprs[i] = consts[0];
4403 std::tie(sgprs[0], sgprs[1]) = setTensorDim3Stride(
4404 op, adaptor, rewriter, loc, sgprs[0], sgprs[1], consts);
4405 std::tie(sgprs[1], sgprs[2]) =
4406 setTensorDim4(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4407 sgprs[2] = setTileDim4(op, adaptor, rewriter, loc, sgprs[2], consts);
4409 Value dgroup3 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4410 for (
auto [sgpr, constant] : llvm::zip(sgprs, consts))
4412 LLVM::InsertElementOp::create(rewriter, loc, dgroup3, sgpr, constant);
4417 Value getDGroup3Gather(DescriptorOp op, OpAdaptor adaptor,
4418 ConversionPatternRewriter &rewriter, Location loc,
4419 ArrayRef<Value> consts)
const {
4420 return getGatherIndices(op, adaptor, rewriter, loc, consts,
false);
4424 matchAndRewrite(DescriptorOp op, OpAdaptor adaptor,
4425 ConversionPatternRewriter &rewriter)
const override {
4427 return op->emitOpError(
4428 "make_dma_descriptor is only supported on gfx1250");
4430 Location loc = op.getLoc();
4432 SmallVector<Value> consts;
4433 for (int64_t i = 0; i < 8; ++i)
4436 Value dgroup0 = this->getDGroup0(adaptor);
4437 Value dgroup1 = this->getDGroup1(op, adaptor, rewriter, loc, consts);
4438 Value dgroup2 = this->getDGroup2(op, adaptor, rewriter, loc, consts);
4439 Value dgroup3 = this->getDGroup3(op, adaptor, rewriter, loc, consts);
4440 SmallVector<Value> results = {dgroup0, dgroup1, dgroup2, dgroup3};
4441 rewriter.replaceOpWithMultiple(op, {results});
4446template <
typename SourceOp,
typename TargetOp>
4447struct AMDGPUTensorLoadStoreOpLowering
4448 :
public ConvertOpToLLVMPattern<SourceOp> {
4449 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
4451 AMDGPUTensorLoadStoreOpLowering(
const LLVMTypeConverter &converter,
4453 : ConvertOpToLLVMPattern<SourceOp>(converter), chipset(chipset) {}
4457 matchAndRewrite(SourceOp op, Adaptor adaptor,
4458 ConversionPatternRewriter &rewriter)
const override {
4460 return op->emitOpError(
"is only supported on gfx1250");
4465 auto v8i32 = VectorType::get(8, rewriter.getI32Type());
4466 Value dgroup4 = LLVM::ZeroOp::create(rewriter, op.getLoc(), v8i32);
4467 Attribute cachePolicy = rewriter.getI32IntegerAttr(0);
4468 rewriter.replaceOpWithNewOp<TargetOp>(op, desc[0], desc[1], desc[2],
4469 desc[3], dgroup4, cachePolicy,
4477struct GlobalPrefetchOpLowering
4478 :
public ConvertOpToLLVMPattern<GlobalPrefetchOp> {
4479 GlobalPrefetchOpLowering(
const LLVMTypeConverter &converter, Chipset chipset)
4480 : ConvertOpToLLVMPattern<GlobalPrefetchOp>(converter), chipset(chipset) {}
4483 matchAndRewrite(GlobalPrefetchOp op, GlobalPrefetchOpAdaptor adaptor,
4484 ConversionPatternRewriter &rewriter)
const override {
4486 return op->emitOpError(
"is only supported on gfx1250+");
4488 const bool isSpeculative = op.getSpeculative();
4490 op.getTemporalHint(), op.getCacheScope(), isSpeculative);
4493 Attribute cachePolicy = ROCDL::Gfx12CachePolicyAttr::get(
4494 rewriter.getContext(),
4495 static_cast<ROCDL::Gfx12CachePolicy
>(immArgValue));
4498 Value memRef = adaptor.getSrc();
4499 MemRefDescriptor descriptor(memRef);
4500 MemRefType memRefType = op.getSrc().getType();
4501 Location loc = op->getLoc();
4502 auto inboundsFlags = isSpeculative ? LLVM::GEPNoWrapFlags::none
4503 : LLVM::GEPNoWrapFlags::inbounds |
4504 LLVM::GEPNoWrapFlags::nuw;
4506 rewriter, loc, memRefType, descriptor,
indices, inboundsFlags);
4508 rewriter.replaceOpWithNewOp<ROCDL::GlobalPrefetchOp>(
4509 op, prefetchPtr, cachePolicy, mlir::ArrayAttr{}, mlir::ArrayAttr{},
4518struct ConvertAMDGPUToROCDLPass
4519 :
public impl::ConvertAMDGPUToROCDLPassBase<ConvertAMDGPUToROCDLPass> {
4522 void runOnOperation()
override {
4525 if (
failed(maybeChipset)) {
4526 emitError(UnknownLoc::get(ctx),
"Invalid chipset name: " + chipset);
4527 return signalPassFailure();
4530 RewritePatternSet patterns(ctx);
4531 LLVMTypeConverter converter(ctx);
4536 target.addIllegalDialect<::mlir::amdgpu::AMDGPUDialect>();
4537 target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
4538 target.addLegalDialect<::mlir::ROCDL::ROCDLDialect>();
4539 if (
failed(applyPartialConversion(getOperation(),
target,
4540 std::move(patterns))))
4541 signalPassFailure();
4549 typeConverter, [](gpu::AddressSpace space) {
4551 case gpu::AddressSpace::Global:
4552 return ROCDL::ROCDLDialect::kGlobalMemoryAddressSpace;
4553 case gpu::AddressSpace::Workgroup:
4554 return ROCDL::ROCDLDialect::kSharedMemoryAddressSpace;
4555 case gpu::AddressSpace::Private:
4556 return ROCDL::ROCDLDialect::kPrivateMemoryAddressSpace;
4557 case gpu::AddressSpace::Constant:
4558 return ROCDL::ROCDLDialect::kConstantMemoryAddressSpace;
4560 llvm_unreachable(
"unknown address space enum value");
4563 return LLVM::LLVMPointerType::get(
4564 type.getContext(), ROCDL::ROCDLDialect::kSharedMemoryAddressSpace);
4570 typeConverter.addTypeAttributeConversion(
4572 -> TypeConverter::AttributeConversionResult {
4574 Type i64 = IntegerType::get(ctx, 64);
4575 switch (as.getValue()) {
4576 case amdgpu::AddressSpace::FatRawBuffer:
4577 return IntegerAttr::get(i64, 7);
4578 case amdgpu::AddressSpace::BufferRsrc:
4579 return IntegerAttr::get(i64, 8);
4580 case amdgpu::AddressSpace::FatStructuredBuffer:
4581 return IntegerAttr::get(i64, 9);
4583 return TypeConverter::AttributeConversionResult::abort();
4585 typeConverter.addConversion([&](DsBarrierStateType type) ->
Type {
4586 return IntegerType::get(type.
getContext(), 64);
4588 typeConverter.addConversion([&](TDMBaseType type) ->
Type {
4590 return typeConverter.convertType(VectorType::get(4, i32));
4592 typeConverter.addConversion([&](TDMGatherBaseType type) ->
Type {
4594 return typeConverter.convertType(VectorType::get(4, i32));
4596 typeConverter.addConversion(
4597 [&](TDMDescriptorType type,
4600 Type v4i32 = typeConverter.convertType(VectorType::get(4, i32));
4601 Type v8i32 = typeConverter.convertType(VectorType::get(8, i32));
4602 llvm::append_values(
result, v4i32, v8i32, v4i32, v4i32);
4612 if (inputs.size() != 1)
4615 if (!isa<TDMDescriptorType>(inputs[0].
getType()))
4618 auto cast = UnrealizedConversionCastOp::create(builder, loc, types, inputs);
4619 return cast.getResults();
4622 typeConverter.addTargetMaterialization(addUnrealizedCast);
4630 .
add<FatRawBufferCastLowering,
4631 RawBufferOpLowering<RawBufferLoadOp, ROCDL::RawPtrBufferLoadOp>,
4632 RawBufferOpLowering<RawBufferStoreOp, ROCDL::RawPtrBufferStoreOp>,
4633 RawBufferOpLowering<RawBufferAtomicFaddOp,
4634 ROCDL::RawPtrBufferAtomicFaddOp>,
4635 RawBufferOpLowering<RawBufferAtomicFmaxOp,
4636 ROCDL::RawPtrBufferAtomicFmaxOp>,
4637 RawBufferOpLowering<RawBufferAtomicSmaxOp,
4638 ROCDL::RawPtrBufferAtomicSmaxOp>,
4639 RawBufferOpLowering<RawBufferAtomicUminOp,
4640 ROCDL::RawPtrBufferAtomicUminOp>,
4641 RawBufferOpLowering<RawBufferAtomicCmpswapOp,
4642 ROCDL::RawPtrBufferAtomicCmpSwap>,
4643 AMDGPUDPPLowering, MemoryCounterWaitOpLowering, LDSBarrierOpLowering,
4644 SchedBarrierOpLowering, MFMAOpLowering, ScaledMFMAOpLowering,
4645 SparseMFMAOpLowering, WMMAOpLowering, ScaledWMMAOpLowering,
4646 SparseWMMAOpLowering, DotOpLowering, ExtPackedFp8OpLowering,
4647 ScaledExtPackedMatrixOpLowering, ScaledExtPackedOpLowering,
4648 PackedScaledTruncOpLowering, PackedTrunc2xFp8OpLowering,
4649 PackedStochRoundFp8OpLowering, GatherToLDSOpLowering,
4650 GlobalLoadAsyncToLDSOpLowering, TransposeLoadOpLowering,
4651 GlobalTransposeLoadOpLowering, AMDGPUPermlaneLowering,
4652 AMDGPUPermlaneVarLowering, AMDGPUMakeDmaBaseLowering<MakeDmaBaseOp>,
4653 AMDGPUMakeDmaBaseLowering<MakeGatherDmaBaseOp>,
4654 AMDGPULowerDescriptor<MakeDmaDescriptorOp>,
4655 AMDGPULowerDescriptor<MakeGatherDmaDescriptorOp>,
4656 AMDGPUTensorLoadStoreOpLowering<TensorLoadToLDSOp,
4657 ROCDL::TensorLoadToLDSOp>,
4658 AMDGPUTensorLoadStoreOpLowering<TensorStoreFromLDSOp,
4659 ROCDL::TensorStoreFromLDSOp>,
4660 DsBarrierInitOpLowering, DsBarrierPollStateOpLowering,
4661 DsAsyncBarrierArriveOpLowering, DsBarrierArriveOpLowering,
4662 GlobalPrefetchOpLowering>(converter, chipset);
4663 patterns.
add<AMDGPUSwizzleBitModeLowering, DsBarrierStatePhaseOpLowering,
4664 DsBarrierStatePendingCountOpLowering,
4665 DsBarrierStateInitCountOpLowering,
4666 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.
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 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.