29#include "llvm/Support/Debug.h"
30#include "llvm/Support/DebugLog.h"
31#include "llvm/Support/ErrorHandling.h"
32#include "llvm/Support/raw_ostream.h"
35#define DEBUG_TYPE "nvgpu-to-nvvm"
38#define GEN_PASS_DEF_CONVERTNVGPUTONVVMPASS
39#include "mlir/Conversion/Passes.h.inc"
52 assert(llvm::isa<IntegerType>(type) &&
"expected an integer Value");
55 return LLVM::TruncOp::create(
b,
b.getI32Type(), value);
62 auto a = cast<LLVM::LLVMArrayType>(vectorResultType);
63 auto f16x2Ty = VectorType::get(2, Float16Type::get(ctx));
64 auto i32Ty = IntegerType::get(ctx, 32);
65 auto i32x2Ty = VectorType::get(2, i32Ty);
66 Type f64Ty = Float64Type::get(ctx);
67 Type f64x2Ty = VectorType::get(2, f64Ty);
68 Type f32Ty = Float32Type::get(ctx);
69 Type f32x2Ty = VectorType::get(2, f32Ty);
70 if (a.getElementType() == f16x2Ty) {
71 return LLVM::LLVMStructType::getLiteral(
74 if (a.getElementType() == i32x2Ty) {
75 return LLVM::LLVMStructType::getLiteral(
79 if (a.getElementType() == f64x2Ty) {
80 return LLVM::LLVMStructType::getLiteral(ctx, {f64Ty, f64Ty});
82 if (a.getElementType() == f32x2Ty) {
83 return LLVM::LLVMStructType::getLiteral(
87 if (a.getElementType() == VectorType::get(1, f32Ty)) {
88 return LLVM::LLVMStructType::getLiteral(
91 return vectorResultType;
103 auto structType = dyn_cast<LLVM::LLVMStructType>(intrinsicResultType);
104 auto arrayType = dyn_cast<LLVM::LLVMArrayType>(resultType);
109 Type i32x2Ty = VectorType::get(2, i32Ty);
110 Type f64x2Ty = VectorType::get(2, f64Ty);
111 Type f32x2Ty = VectorType::get(2, f32Ty);
112 Type f32x1Ty = VectorType::get(1, f32Ty);
114 auto makeConst = [&](int32_t
index) ->
Value {
115 return LLVM::ConstantOp::create(rewriter, loc, IntegerType::get(ctx, 32),
124 if (arrayType.getElementType() == f16x2Ty ||
125 arrayType.getElementType() == f32x1Ty) {
126 for (
unsigned i = 0; i < structType.getBody().size(); i++) {
128 LLVM::ExtractValueOp::create(rewriter, loc, intrinsicResult, i);
130 loc, arrayType.getElementType(), el);
131 elements.push_back(el);
139 if (arrayType.getElementType() == i32x2Ty ||
140 arrayType.getElementType() == f64x2Ty ||
141 arrayType.getElementType() == f32x2Ty) {
143 for (
unsigned i = 0, e = structType.getBody().size() / 2; i < e; i++) {
145 LLVM::PoisonOp::create(rewriter, loc, arrayType.getElementType());
147 LLVM::ExtractValueOp::create(rewriter, loc, intrinsicResult, i * 2);
148 Value x2 = LLVM::ExtractValueOp::create(rewriter, loc, intrinsicResult,
150 vec = LLVM::InsertElementOp::create(rewriter, loc, vec.
getType(), vec,
152 vec = LLVM::InsertElementOp::create(rewriter, loc, vec.
getType(), vec,
154 elements.push_back(vec);
159 Value result = LLVM::PoisonOp::create(rewriter, loc, arrayType);
160 for (
const auto &el : llvm::enumerate(elements)) {
161 result = LLVM::InsertValueOp::create(rewriter, loc,
result, el.value(),
167 return intrinsicResult;
177 NVVM::MMATypes operandPtxType) {
179 Type i32Ty =
b.getI32Type();
180 Type f64Ty =
b.getF64Type();
181 Type f32Ty =
b.getF32Type();
182 Type i64Ty =
b.getI64Type();
183 Type bf16x2Ty = VectorType::get(2,
b.getBF16Type());
184 Type i8x4Ty = VectorType::get(4,
b.getI8Type());
185 Type i4x8Ty = VectorType::get(8,
b.getIntegerType(4));
186 Type f32x1Ty = VectorType::get(1, f32Ty);
187 auto arrayTy = cast<LLVM::LLVMArrayType>(operand.
getType());
189 for (
unsigned i = 0, e = arrayTy.getNumElements(); i < e; ++i) {
190 Value toUse = LLVM::ExtractValueOp::create(
b, operand, i);
194 if (arrayTy.getElementType() == i8x4Ty ||
195 arrayTy.getElementType() == i4x8Ty ||
196 (arrayTy.getElementType() == bf16x2Ty &&
197 operandPtxType == NVVM::MMATypes::bf16) ||
198 (arrayTy.getElementType() == f32x1Ty &&
199 operandPtxType == NVVM::MMATypes::tf32)) {
200 result.push_back(LLVM::BitcastOp::create(
b, i32Ty, toUse));
207 VectorType innerArrayTy = dyn_cast<VectorType>(arrayTy.getElementType());
208 if (innerArrayTy && (innerArrayTy.getElementType() == i32Ty ||
209 innerArrayTy.getElementType() == f64Ty ||
210 innerArrayTy.getElementType() == f32Ty)) {
211 for (
unsigned idx = 0, innerSize = innerArrayTy.getNumElements();
212 idx < innerSize; idx++) {
213 result.push_back(LLVM::ExtractElementOp::create(
215 LLVM::ConstantOp::create(
b, i64Ty,
b.getI64IntegerAttr(idx))));
226 return (mlir::nvgpu::NVGPUDialect::isSharedMemoryAddressSpace(
227 barrierType.getMemorySpace()));
232 nvgpu::MBarrierGroupType barrierType) {
236 IntegerAttr::get(IntegerType::get(context, 64),
237 nvgpu::NVGPUDialect::kSharedMemoryAddressSpace);
245 nvgpu::MBarrierGroupType barrierType) {
247 MemRefLayoutAttrInterface layout;
248 return MemRefType::get({barrierType.getNumBarriers()},
249 IntegerType::get(context, 64), layout, memorySpace);
255 using ConvertOpToLLVMPattern<nvgpu::LdMatrixOp>::ConvertOpToLLVMPattern;
258 matchAndRewrite(nvgpu::LdMatrixOp op, OpAdaptor adaptor,
259 ConversionPatternRewriter &rewriter)
const override {
261 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
269 auto vectorResultType = dyn_cast<VectorType>(op->getResultTypes()[0]);
270 if (!vectorResultType) {
273 Type innerVectorType = VectorType::get(vectorResultType.getDimSize(1),
274 vectorResultType.getElementType());
276 int64_t num32BitRegs = vectorResultType.getDimSize(0);
278 Type ldMatrixResultType;
279 if (num32BitRegs > 1) {
280 ldMatrixResultType = LLVM::LLVMStructType::getLiteral(
281 ctx, SmallVector<Type>(num32BitRegs, rewriter.getI32Type()));
283 ldMatrixResultType = rewriter.getI32Type();
286 auto srcMemrefType = cast<MemRefType>(op.getSrcMemref().getType());
289 adaptor.getSrcMemref(), adaptor.getIndices());
290 auto shape = NVVM::LdStMatrixShapeAttr::get(rewriter.getContext(), 8, 8);
291 Value ldMatrixResult = NVVM::LdMatrixOp::create(
292 b, ldMatrixResultType, srcPtr,
294 op.getTranspose() ? NVVM::MMALayout::col
295 : NVVM::MMALayout::row,
296 shape, NVVM::LdStMatrixEltType::B16);
302 Type finalResultType = typeConverter->convertType(vectorResultType);
303 Value
result = LLVM::PoisonOp::create(
b, finalResultType);
304 for (int64_t i = 0, e = vectorResultType.getDimSize(0); i < e; i++) {
306 num32BitRegs > 1 ? LLVM::ExtractValueOp::create(
b, ldMatrixResult, i)
308 Value casted = LLVM::BitcastOp::create(
b, innerVectorType, i32Register);
312 rewriter.replaceOp(op,
result);
319static FailureOr<NVVM::MMATypes> getNvvmMmaType(
Type t) {
322 return NVVM::MMATypes::s8;
324 return NVVM::MMATypes::s4;
326 return NVVM::MMATypes::f16;
328 return NVVM::MMATypes::bf16;
330 return NVVM::MMATypes::f64;
332 return NVVM::MMATypes::tf32;
334 return NVVM::MMATypes::e4m3;
336 return NVVM::MMATypes::e5m2;
341 using ConvertOpToLLVMPattern<nvgpu::MmaSyncOp>::ConvertOpToLLVMPattern;
344 matchAndRewrite(nvgpu::MmaSyncOp op, OpAdaptor adaptor,
345 ConversionPatternRewriter &rewriter)
const override {
346 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
349 VectorType aType = op.getMatrixA().getType();
350 VectorType bType = op.getMatrixA().getType();
351 VectorType cType = op.getMatrixC().getType();
353 std::array<int64_t, 3> gemmShape = op.getMmaShapeAsArray();
356 bool tf32Enabled = op->hasAttr(op.getTf32EnabledAttrName());
357 if (aType.getElementType().isF32() && !tf32Enabled)
360 FailureOr<NVVM::MMATypes> ptxTypeA = getNvvmMmaType(aType);
362 return op->emitOpError(
"failed to deduce operand PTX types");
363 FailureOr<NVVM::MMATypes> ptxTypeB = getNvvmMmaType(bType);
365 return op->emitOpError(
"failed to deduce operand PTX types");
366 std::optional<NVVM::MMATypes> ptxTypeC =
367 NVVM::MmaOp::inferOperandMMAType(cType.getElementType(),
370 return op->emitError(
371 "could not infer the PTX type for the accumulator/result");
374 std::optional<NVVM::MMAIntOverflow> overflow(std::nullopt);
375 if (isa<IntegerType>(aType.getElementType()))
376 overflow = NVVM::MMAIntOverflow::satfinite;
378 SmallVector<Value> matA =
380 SmallVector<Value> matB =
382 SmallVector<Value> matC =
385 Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
387 typeConverter->convertType(op->getResultTypes()[0]));
388 Value intrinsicResult =
389 NVVM::MmaOp::create(
b, intrinsicResTy, matA, matB, matC,
394 std::array<NVVM::MMATypes, 2>{*ptxTypeA, *ptxTypeB},
396 std::array<NVVM::MMALayout, 2>{
397 NVVM::MMALayout::row, NVVM::MMALayout::col});
399 desiredRetTy, intrinsicResult,
405struct ConvertNVGPUToNVVMPass
406 :
public impl::ConvertNVGPUToNVVMPassBase<ConvertNVGPUToNVVMPass> {
409 void runOnOperation()
override {
419 converter.addConversion([&](nvgpu::DeviceAsyncTokenType type) -> Type {
420 return converter.convertType(IntegerType::get(type.getContext(), 32));
422 converter.addConversion([&](nvgpu::WarpgroupAccumulatorType type) -> Type {
423 Type elemType = type.getFragmented().getElementType();
424 int64_t sizeM = type.getFragmented().getDimSize(0);
425 int64_t sizeN = type.getFragmented().getDimSize(1);
429 numMembers = sizeN / 2;
430 else if (elemType.
isF16())
431 numMembers = sizeN / 4;
433 llvm_unreachable(
"unsupported type for warpgroup accumulator");
435 SmallVector<Type> innerStructBody;
436 for (
unsigned i = 0; i < numMembers; i++)
437 innerStructBody.push_back(elemType);
438 auto innerStructType =
439 LLVM::LLVMStructType::getLiteral(type.getContext(), innerStructBody);
441 SmallVector<Type> structBody;
443 structBody.push_back(innerStructType);
446 LLVM::LLVMStructType::getLiteral(type.getContext(), structBody);
447 return converter.convertType(convertedType);
449 converter.addConversion([&](nvgpu::MBarrierTokenType type) -> Type {
450 return converter.convertType(IntegerType::get(type.getContext(), 64));
452 converter.addConversion(
453 [&](nvgpu::WarpgroupMatrixDescriptorType type) -> Type {
454 return converter.convertType(IntegerType::get(type.getContext(), 64));
456 converter.addConversion([&](nvgpu::MBarrierGroupType type) -> Type {
457 return converter.convertType(
460 converter.addConversion([&](nvgpu::TensorMapDescriptorType type) -> Type {
461 return LLVM::LLVMPointerType::get(type.getContext());
465 target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
466 target.addLegalDialect<::mlir::arith::ArithDialect>();
467 target.addLegalDialect<::mlir::memref::MemRefDialect>();
468 target.addLegalDialect<::mlir::NVVM::NVVMDialect>();
469 target.addLegalDialect<::mlir::vector::VectorDialect>();
471 converter, patterns,
target);
472 if (
failed(applyPartialConversion(getOperation(),
target,
473 std::move(patterns))))
479static std::string buildMmaSparseAsmConstraintString(
unsigned matASize,
483 llvm::raw_string_ostream ss(str);
484 for (
unsigned i = 0; i < matCSize; i++)
486 for (
unsigned i = 0; i < matASize + matBSize + matCSize; i++)
498static std::string buildMmaSparseAsmString(
499 const std::array<int64_t, 3> &
shape,
unsigned matASize,
unsigned matBSize,
500 unsigned matCSize, NVVM::MMATypes ptxTypeA, NVVM::MMATypes ptxTypeB,
501 NVVM::MMATypes ptxTypeC, NVVM::MMATypes ptxTypeD,
502 std::optional<NVVM::MMAIntOverflow> overflow,
unsigned metaDataSelector) {
503 auto ptxTypeStr = [](NVVM::MMATypes ptxType) {
504 return NVVM::stringifyMMATypes(ptxType);
508 llvm::raw_string_ostream ss(asmStr);
509 ss <<
"mma.sp.sync.aligned.m" <<
shape[0] <<
"n" <<
shape[1] <<
"k"
510 <<
shape[2] <<
".row.col.";
513 ss << NVVM::stringifyMMAIntOverflow(*overflow) <<
".";
515 ss << ptxTypeStr(ptxTypeD) <<
"." << ptxTypeStr(ptxTypeA) <<
"."
516 << ptxTypeStr(ptxTypeB) <<
"." << ptxTypeStr(ptxTypeC) <<
" ";
517 unsigned asmArgIdx = 0;
521 for (
const auto arrSize : {matCSize, matASize, matBSize, matCSize}) {
523 for (
unsigned i = 0; i < arrSize; i++)
524 ss <<
"$" << asmArgIdx++ << (i < arrSize - 1 ?
"," :
"");
527 ss <<
"$" << asmArgIdx++ <<
",";
528 assert(metaDataSelector <= 1);
529 ss <<
"0x" << metaDataSelector <<
";";
535static FailureOr<LLVM::InlineAsmOp> emitMmaSparseSyncOpAsm(
537 NVVM::MMATypes ptxTypeC, NVVM::MMATypes ptxTypeD,
538 std::optional<NVVM::MMAIntOverflow> overflow,
ArrayRef<Value> unpackedAData,
540 int64_t metadataSelector,
const std::array<int64_t, 3> &
shape,
541 Type intrinsicResultType) {
542 auto asmDialectAttr =
543 LLVM::AsmDialectAttr::get(
b.getContext(), LLVM::AsmDialect::AD_ATT);
545 const unsigned matASize = unpackedAData.size();
546 const unsigned matBSize = unpackedB.size();
547 const unsigned matCSize = unpackedC.size();
549 std::string asmStr = buildMmaSparseAsmString(
550 shape, matASize, matBSize, matCSize, ptxTypeA, ptxTypeB, ptxTypeC,
551 ptxTypeD, overflow, metadataSelector);
552 std::string constraintStr =
553 buildMmaSparseAsmConstraintString(matASize, matBSize, matCSize);
556 asmVals.reserve(matASize + matBSize + matCSize + 1);
558 llvm::append_range(asmVals, args);
559 asmVals.push_back(indexData);
561 return LLVM::InlineAsmOp::create(
b,
568 LLVM::TailCallKind::None,
574struct NVGPUMmaSparseSyncLowering
576 using ConvertOpToLLVMPattern<nvgpu::MmaSparseSyncOp>::ConvertOpToLLVMPattern;
579 matchAndRewrite(nvgpu::MmaSparseSyncOp op, OpAdaptor adaptor,
580 ConversionPatternRewriter &rewriter)
const override {
581 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
584 VectorType aType = op.getMatrixA().getType();
585 VectorType bType = op.getMatrixB().getType();
586 VectorType cType = op.getMatrixC().getType();
588 FailureOr<NVVM::MMATypes> ptxTypeA = getNvvmMmaType(aType);
590 return op->emitOpError(
"failed to deduce operand PTX types");
591 FailureOr<NVVM::MMATypes> ptxTypeB = getNvvmMmaType(bType);
593 return op->emitOpError(
"failed to deduce operand PTX types");
594 std::optional<NVVM::MMATypes> ptxTypeC =
595 NVVM::MmaOp::inferOperandMMAType(cType.getElementType(),
598 return op->emitError(
599 "could not infer the PTX type for the accumulator/result");
602 bool tf32Enabled = op->hasAttr(op.getTf32EnabledAttrName());
603 if (aType.getElementType().isF32() && !tf32Enabled)
607 std::optional<NVVM::MMAIntOverflow> overflow(std::nullopt);
608 if (isa<IntegerType>(aType.getElementType()))
609 overflow = NVVM::MMAIntOverflow::satfinite;
611 SmallVector<Value> matA =
613 SmallVector<Value> matB =
615 SmallVector<Value> matC =
618 Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
620 typeConverter->convertType(op->getResultTypes()[0]));
623 Value sparseMetadata = adaptor.getSparseMetadata();
624 if (sparseMetadata.
getType() != VectorType::get(2, rewriter.getI16Type()))
625 return op->emitOpError() <<
"Expected metadata type to be LLVM "
626 "VectorType of 2 i16 elements";
628 LLVM::BitcastOp::create(
b, rewriter.getI32Type(), sparseMetadata);
630 FailureOr<LLVM::InlineAsmOp> intrinsicResult = emitMmaSparseSyncOpAsm(
631 b, *ptxTypeA, *ptxTypeB, *ptxTypeC, *ptxTypeC, overflow, matA, matB,
632 matC, sparseMetadata, op.getSparsitySelector(), op.getMmaShapeAsArray(),
634 if (
failed(intrinsicResult))
637 assert((*intrinsicResult).getNumResults() == 1 &&
638 "expected inline asm op returns a single LLVM struct type");
641 (*intrinsicResult)->getResult(0), rewriter));
646struct NVGPUAsyncCopyLowering
648 using ConvertOpToLLVMPattern<
649 nvgpu::DeviceAsyncCopyOp>::ConvertOpToLLVMPattern;
652 matchAndRewrite(nvgpu::DeviceAsyncCopyOp op, OpAdaptor adaptor,
653 ConversionPatternRewriter &rewriter)
const override {
654 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
655 Location loc = op.getLoc();
656 auto dstMemrefType = cast<MemRefType>(op.getDst().getType());
659 adaptor.getDst(), adaptor.getDstIndices());
660 FailureOr<unsigned> dstAddressSpace =
661 getTypeConverter()->getMemRefAddressSpace(dstMemrefType);
662 if (
failed(dstAddressSpace))
663 return rewriter.notifyMatchFailure(
664 loc,
"destination memref address space not convertible to integer");
666 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
667 FailureOr<unsigned> srcAddressSpace =
668 getTypeConverter()->getMemRefAddressSpace(srcMemrefType);
669 if (
failed(srcAddressSpace))
670 return rewriter.notifyMatchFailure(
671 loc,
"source memref address space not convertible to integer");
675 adaptor.getSrcIndices());
677 auto srcPointerGlobalType = LLVM::LLVMPointerType::get(
678 op->getContext(),
static_cast<unsigned>(NVVM::NVVMMemorySpace::Global));
679 scrPtr = LLVM::AddrSpaceCastOp::create(
b, srcPointerGlobalType, scrPtr);
680 int64_t dstElements = adaptor.getDstElements().getZExtValue();
681 int64_t sizeInBytes =
682 (dstMemrefType.getElementTypeBitWidth() * dstElements) / 8;
687 Value srcBytes = adaptor.getSrcElements();
694 LLVM::ConstantOp::create(
b,
b.getI32Type(),
b.getI32IntegerAttr(3));
695 Value bitwidth = LLVM::ConstantOp::create(
697 b.getI32IntegerAttr(srcMemrefType.getElementTypeBitWidth()));
698 Value srcElementsI32 = LLVM::TruncOp::create(
b,
b.getI32Type(), srcBytes);
699 srcBytes = LLVM::LShrOp::create(
700 b, LLVM::MulOp::create(
b, bitwidth, srcElementsI32), c3I32);
704 NVVM::LoadCacheModifierKind cacheModifier =
705 (op.getBypassL1().value_or(
false) && sizeInBytes == 16)
706 ? NVVM::LoadCacheModifierKind::CG
707 : NVVM::LoadCacheModifierKind::CA;
709 NVVM::CpAsyncOp::create(
710 b, dstPtr, scrPtr, rewriter.getI32IntegerAttr(sizeInBytes),
711 NVVM::LoadCacheModifierKindAttr::get(op->getContext(), cacheModifier),
716 LLVM::ConstantOp::create(
b, IntegerType::get(op.getContext(), 32),
717 rewriter.getI32IntegerAttr(0));
718 rewriter.replaceOp(op, zero);
723struct NVGPUAsyncCreateGroupLowering
725 using ConvertOpToLLVMPattern<
726 nvgpu::DeviceAsyncCreateGroupOp>::ConvertOpToLLVMPattern;
729 matchAndRewrite(nvgpu::DeviceAsyncCreateGroupOp op, OpAdaptor adaptor,
730 ConversionPatternRewriter &rewriter)
const override {
731 NVVM::CpAsyncCommitGroupOp::create(rewriter, op.getLoc());
733 Value zero = LLVM::ConstantOp::create(rewriter, op->getLoc(),
734 IntegerType::get(op.getContext(), 32),
735 rewriter.getI32IntegerAttr(0));
736 rewriter.replaceOp(op, zero);
741struct NVGPUAsyncWaitLowering
743 using ConvertOpToLLVMPattern<
744 nvgpu::DeviceAsyncWaitOp>::ConvertOpToLLVMPattern;
747 matchAndRewrite(nvgpu::DeviceAsyncWaitOp op, OpAdaptor adaptor,
748 ConversionPatternRewriter &rewriter)
const override {
750 int32_t numGroups = adaptor.getNumGroups().value_or(0);
751 NVVM::CpAsyncWaitGroupOp::create(rewriter, op.getLoc(), numGroups);
752 rewriter.eraseOp(op);
758struct NVGPUMBarrierCreateLowering
760 using ConvertOpToLLVMPattern<nvgpu::MBarrierCreateOp>::ConvertOpToLLVMPattern;
762 template <
typename moduleT>
763 memref::GlobalOp generateGlobalBarrier(ConversionPatternRewriter &rewriter,
764 Operation *funcOp, moduleT moduleOp,
765 MemRefType barrierType)
const {
766 SymbolTable symbolTable(moduleOp);
767 OpBuilder::InsertionGuard guard(rewriter);
768 rewriter.setInsertionPoint(&moduleOp.front());
769 auto global = memref::GlobalOp::create(
770 rewriter, funcOp->
getLoc(),
"__mbarrier",
771 rewriter.getStringAttr(
"private"),
775 rewriter.getI64IntegerAttr(8));
776 symbolTable.insert(global);
781 matchAndRewrite(nvgpu::MBarrierCreateOp op, OpAdaptor adaptor,
782 ConversionPatternRewriter &rewriter)
const override {
785 rewriter.getContext(), op.getBarriers().getType());
787 memref::GlobalOp global;
789 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
791 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
793 rewriter.setInsertionPoint(op);
794 rewriter.replaceOpWithNewOp<memref::GetGlobalOp>(op, barrierType,
801template <
typename SourceOp>
804 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
806 Value getMbarrierPtr(ImplicitLocOpBuilder &
b,
807 nvgpu::MBarrierGroupType mbarType, Value memrefDesc,
809 ConversionPatternRewriter &rewriter)
const {
810 MemRefType mbarrierMemrefType =
813 rewriter,
b.getLoc(), mbarrierMemrefType, memrefDesc, {mbarId});
817struct NVGPUMBarrierGetLowering
818 :
public MBarrierBasePattern<nvgpu::MBarrierGetOp> {
819 using MBarrierBasePattern<nvgpu::MBarrierGetOp>::MBarrierBasePattern;
822 matchAndRewrite(nvgpu::MBarrierGetOp op, OpAdaptor adaptor,
823 ConversionPatternRewriter &rewriter)
const override {
824 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
825 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
826 rewriter.setInsertionPoint(op);
827 Value barrier = getMbarrierPtr(
b, mbarrierType, adaptor.getBarriers(),
828 adaptor.getMbarId(), rewriter);
829 Type resType = op.getMbarrierPointer().getType();
830 rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(op, resType, barrier);
836struct NVGPUMBarrierInitLowering
837 :
public MBarrierBasePattern<nvgpu::MBarrierInitOp> {
838 using MBarrierBasePattern<nvgpu::MBarrierInitOp>::MBarrierBasePattern;
841 matchAndRewrite(nvgpu::MBarrierInitOp op, OpAdaptor adaptor,
842 ConversionPatternRewriter &rewriter)
const override {
843 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
844 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
845 rewriter.setInsertionPoint(op);
846 Value barrier = getMbarrierPtr(
b, mbarrierType, adaptor.getBarriers(),
847 adaptor.getMbarId(), rewriter);
849 rewriter.replaceOpWithNewOp<NVVM::MBarrierInitOp>(op, barrier, count,
850 adaptor.getPredicate());
856struct NVGPUMBarrierArriveLowering
857 :
public MBarrierBasePattern<nvgpu::MBarrierArriveOp> {
858 using MBarrierBasePattern<nvgpu::MBarrierArriveOp>::MBarrierBasePattern;
860 matchAndRewrite(nvgpu::MBarrierArriveOp op, OpAdaptor adaptor,
861 ConversionPatternRewriter &rewriter)
const override {
862 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
864 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
865 adaptor.getMbarId(), rewriter);
866 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveOp>(op, barrier);
873struct NVGPUMBarrierArriveNoCompleteLowering
874 :
public MBarrierBasePattern<nvgpu::MBarrierArriveNoCompleteOp> {
875 using MBarrierBasePattern<
876 nvgpu::MBarrierArriveNoCompleteOp>::MBarrierBasePattern;
878 matchAndRewrite(nvgpu::MBarrierArriveNoCompleteOp op, OpAdaptor adaptor,
879 ConversionPatternRewriter &rewriter)
const override {
880 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
882 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
883 adaptor.getMbarId(), rewriter);
884 Type tokenType = getTypeConverter()->convertType(
885 nvgpu::MBarrierTokenType::get(op->getContext()));
887 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveNocompleteOp>(
888 op, tokenType, barrier, count);
894struct NVGPUMBarrierTestWaitLowering
895 :
public MBarrierBasePattern<nvgpu::MBarrierTestWaitOp> {
896 using MBarrierBasePattern<nvgpu::MBarrierTestWaitOp>::MBarrierBasePattern;
898 matchAndRewrite(nvgpu::MBarrierTestWaitOp op, OpAdaptor adaptor,
899 ConversionPatternRewriter &rewriter)
const override {
900 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
902 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
903 adaptor.getMbarId(), rewriter);
904 Type retType = rewriter.getI1Type();
905 rewriter.replaceOpWithNewOp<NVVM::MBarrierTestWaitOp>(op, retType, barrier,
911struct NVGPUMBarrierArriveExpectTxLowering
912 :
public MBarrierBasePattern<nvgpu::MBarrierArriveExpectTxOp> {
913 using MBarrierBasePattern<
914 nvgpu::MBarrierArriveExpectTxOp>::MBarrierBasePattern;
916 matchAndRewrite(nvgpu::MBarrierArriveExpectTxOp op, OpAdaptor adaptor,
917 ConversionPatternRewriter &rewriter)
const override {
918 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
920 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
921 adaptor.getMbarId(), rewriter);
922 Value txcount =
truncToI32(
b, adaptor.getTxcount());
923 NVVM::MBarrierArriveExpectTxOp::create(
924 rewriter, op->getLoc(), barrier, txcount,
925 NVVM::MemScopeKind::CTA,
927 adaptor.getPredicate());
928 rewriter.eraseOp(op);
933struct NVGPUMBarrierTryWaitParityLowering
934 :
public MBarrierBasePattern<nvgpu::MBarrierTryWaitParityOp> {
935 using MBarrierBasePattern<
936 nvgpu::MBarrierTryWaitParityOp>::MBarrierBasePattern;
938 matchAndRewrite(nvgpu::MBarrierTryWaitParityOp op, OpAdaptor adaptor,
939 ConversionPatternRewriter &rewriter)
const override {
940 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
942 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
943 adaptor.getMbarId(), rewriter);
946 LLVM::ZExtOp::create(
b,
b.getI32Type(), adaptor.getPhaseParity());
947 rewriter.replaceOpWithNewOp<NVVM::MBarrierTryWaitParityOp>(op, barrier,
953struct NVGPUTmaAsyncLoadOpLowering
954 :
public MBarrierBasePattern<nvgpu::TmaAsyncLoadOp> {
955 using MBarrierBasePattern<nvgpu::TmaAsyncLoadOp>::MBarrierBasePattern;
957 matchAndRewrite(nvgpu::TmaAsyncLoadOp op, OpAdaptor adaptor,
958 ConversionPatternRewriter &rewriter)
const override {
959 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
960 auto srcMemrefType = cast<MemRefType>(op.getDst().getType());
962 adaptor.getDst(), {});
966 auto ptrSharedClusterType = LLVM::LLVMPointerType::get(
968 static_cast<unsigned>(NVVM::NVVMMemorySpace::SharedCluster));
969 dest = LLVM::AddrSpaceCastOp::create(
b, ptrSharedClusterType, dest);
972 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
973 adaptor.getMbarId(), rewriter);
975 SmallVector<Value> coords = adaptor.getCoordinates();
976 for (
auto [index, value] : llvm::enumerate(coords)) {
981 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(
982 op, dest, adaptor.getTensorMapDescriptor(), coords, barrier,
983 ValueRange{}, adaptor.getMulticastMask(), Value{},
984 NVVM::TMALoadMode::TILE,
987 adaptor.getPredicate());
992struct NVGPUTmaAsyncStoreOpLowering
993 :
public MBarrierBasePattern<nvgpu::TmaAsyncStoreOp> {
994 using MBarrierBasePattern<nvgpu::TmaAsyncStoreOp>::MBarrierBasePattern;
996 matchAndRewrite(nvgpu::TmaAsyncStoreOp op, OpAdaptor adaptor,
997 ConversionPatternRewriter &rewriter)
const override {
998 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
999 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
1001 adaptor.getSrc(), {});
1002 SmallVector<Value> coords = adaptor.getCoordinates();
1003 for (
auto [index, value] : llvm::enumerate(coords)) {
1008 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(
1009 op, adaptor.getTensorMapDescriptor(), dest, coords, Value{},
1010 NVVM::TMAStoreMode::TILE,
1011 adaptor.getPredicate());
1016struct NVGPUGenerateWarpgroupDescriptorLowering
1018 using ConvertOpToLLVMPattern<
1019 nvgpu::WarpgroupGenerateDescriptorOp>::ConvertOpToLLVMPattern;
1022 matchAndRewrite(nvgpu::WarpgroupGenerateDescriptorOp op, OpAdaptor adaptor,
1023 ConversionPatternRewriter &rewriter)
const override {
1025 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1027 nvgpu::TensorMapSwizzleKind swizzleKind =
1028 op.getTensorMap().getType().getSwizzle();
1031 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 128
1032 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 64
1033 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 32
1036 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 1
1037 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 2
1038 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 3
1041 auto ti64 =
b.getIntegerType(64);
1042 auto makeConst = [&](uint64_t index) -> Value {
1043 return LLVM::ConstantOp::create(
b, ti64,
b.getI64IntegerAttr(index));
1045 auto shiftLeft = [&](Value value,
unsigned shift) -> Value {
1046 return LLVM::ShlOp::create(
b, ti64, value, makeConst(shift));
1048 auto shiftRight = [&](Value value,
unsigned shift) -> Value {
1049 return LLVM::LShrOp::create(
b, ti64, value, makeConst(shift));
1051 auto insertBit = [&](Value desc, Value val,
int startBit) {
1052 return LLVM::OrOp::create(
b, ti64, desc, shiftLeft(val, startBit));
1055 int64_t sizeN = op.getTensorMap().
getType().getTensor().getDimSize(0);
1056 uint64_t strideDimVal = (layout << 3) >>
exclude4LSB;
1057 uint64_t leadDimVal = (sizeN * layout) >>
exclude4LSB;
1058 uint64_t offsetVal = 0;
1060 Value strideDim = makeConst(strideDimVal);
1061 Value leadDim = makeConst(leadDimVal);
1064 rewriter, op->getLoc(), cast<MemRefType>(op.getTensor().getType()),
1065 adaptor.getTensor(), {});
1066 Value basePtr = LLVM::PtrToIntOp::create(
b, ti64, baseAddr);
1068 Value basePtr14bit = shiftRight(shiftLeft(basePtr, 46), 50);
1070 int startSwizzleBit = 62, startOffsetBit = 49, startStrideBit = 32,
1071 startLeadBit = 16, startBaseAddrBit = 0;
1072 Value dsc = makeConst(0);
1074 dsc = insertBit(dsc, makeConst(swizzle), startSwizzleBit);
1076 dsc = insertBit(dsc, makeConst(offsetVal), startOffsetBit);
1078 dsc = insertBit(dsc, strideDim, startStrideBit);
1080 dsc = insertBit(dsc, leadDim, startLeadBit);
1082 dsc = insertBit(dsc, basePtr14bit, startBaseAddrBit);
1084 LDBG() <<
"Generating warpgroup.descriptor: " <<
"leading_off:"
1085 << leadDimVal <<
"\t" <<
"stride_off :" << strideDimVal <<
"\t"
1086 <<
"base_offset:" << offsetVal <<
"\t" <<
"layout_type:" << swizzle
1087 <<
" (" << nvgpu::stringifyTensorMapSwizzleKind(swizzleKind)
1088 <<
")\n start_addr : " << baseAddr;
1090 rewriter.replaceOp(op, dsc);
1096 return LLVM::ConstantOp::create(
b,
b.getIntegerType(64),
1097 b.getI32IntegerAttr(
index));
1104 enum CUtensorMapDataTypeEnum {
1105 CU_TENSOR_MAP_DATA_TYPE_UINT8 = 0,
1106 CU_TENSOR_MAP_DATA_TYPE_UINT16,
1107 CU_TENSOR_MAP_DATA_TYPE_UINT32,
1108 CU_TENSOR_MAP_DATA_TYPE_INT32,
1109 CU_TENSOR_MAP_DATA_TYPE_UINT64,
1110 CU_TENSOR_MAP_DATA_TYPE_INT64,
1111 CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
1112 CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
1113 CU_TENSOR_MAP_DATA_TYPE_FLOAT64,
1114 CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
1115 CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ,
1116 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32,
1117 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ
1121 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT8);
1123 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT16);
1125 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT32);
1127 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT64);
1129 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_INT32);
1131 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_INT64);
1133 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_FLOAT16);
1135 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_FLOAT32);
1137 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_FLOAT64);
1139 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16);
1141 llvm_unreachable(
"Not supported data type");
1144struct NVGPUTmaCreateDescriptorOpLowering
1146 using ConvertOpToLLVMPattern<
1147 nvgpu::TmaCreateDescriptorOp>::ConvertOpToLLVMPattern;
1149 matchAndRewrite(nvgpu::TmaCreateDescriptorOp op, OpAdaptor adaptor,
1150 ConversionPatternRewriter &rewriter)
const override {
1151 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1152 auto llvmPointerType = LLVM::LLVMPointerType::get(op->getContext());
1153 Type llvmInt64Type = IntegerType::get(op->getContext(), 64);
1155 Value tensorElementType =
1156 elementTypeAsLLVMConstant(
b, op.getTensor().getType().getElementType());
1157 auto promotedOperands = getTypeConverter()->promoteOperands(
1158 b.getLoc(), op->getOperands(), adaptor.getOperands(),
b);
1160 Value boxArrayPtr = LLVM::AllocaOp::create(
1161 b, llvmPointerType, llvmInt64Type, makeI64Const(
b, 5));
1162 for (
auto [index, value] : llvm::enumerate(adaptor.getBoxDimensions())) {
1163 Value gep = LLVM::GEPOp::create(
b, llvmPointerType, llvmPointerType,
1164 boxArrayPtr, makeI64Const(
b, index));
1165 LLVM::StoreOp::create(
b, value, gep);
1168 nvgpu::TensorMapDescriptorType desc = op.getTensorMap().
getType();
1170 SmallVector<Value> arguments;
1171 arguments.push_back(promotedOperands[0]);
1172 arguments.push_back(promotedOperands[1]);
1173 arguments.push_back(tensorElementType);
1174 arguments.push_back(
1175 makeI64Const(
b, (
int)desc.getInterleave()));
1176 arguments.push_back(makeI64Const(
b, (
int)desc.getSwizzle()));
1177 arguments.push_back(makeI64Const(
b, (
int)desc.getL2promo()));
1178 arguments.push_back(makeI64Const(
b, (
int)desc.getOob()));
1179 arguments.push_back(boxArrayPtr);
1182 SmallVector<Type> argTypes = {
1192 FunctionCallBuilder hostRegisterCallBuilder = {
1193 "mgpuTensorMapEncodeTiledMemref", llvmPointerType, argTypes};
1195 hostRegisterCallBuilder.
create(
b.getLoc(),
b, arguments).getResult();
1197 rewriter.replaceOp(op, tensorMap);
1202struct NVGPUWarpgroupMmaOpLowering
1204 using ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaOp>::ConvertOpToLLVMPattern;
1226 class WarpgroupGemm {
1227 nvgpu::WarpgroupMmaOp op;
1228 ImplicitLocOpBuilder b;
1232 int64_t totalM, totalN, totalK;
1235 int wgmmaM = 0, wgmmaN = 0, wgmmaK = 0;
1238 int iterationM = 0, iterationN = 0, iterationK = 0;
1243 void findWgmmaShape(int64_t sizeM, int64_t sizeN, Type inputElemType) {
1246 if (inputElemType.
isTF32()) {
1248 }
else if (inputElemType.
isF16() || inputElemType.
isBF16()) {
1250 }
else if (isa<Float8E4M3FNType, Float8E5M2Type>(inputElemType) ||
1253 }
else if (inputElemType.
isInteger(1)) {
1256 llvm_unreachable(
"msg: not supported K shape");
1258 LDBG() <<
"Generating WgmmaMmaAsyncOp shape[m = " << wgmmaM
1259 <<
", n = " << wgmmaN <<
", k = " << wgmmaK <<
"]";
1263 NVVM::WGMMATypesAttr generateWgmmaType(Type type,
1264 bool useF32 =
false)
const {
1265 auto getWgmmaType = [=](Type elemType) {
1267 return useF32 ? NVVM::WGMMATypes::f32 : NVVM::WGMMATypes::tf32;
1268 if (elemType.
isF16())
1269 return NVVM::WGMMATypes::f16;
1271 return NVVM::WGMMATypes::bf16;
1272 if (isa<Float8E4M3FNType>(elemType))
1273 return NVVM::WGMMATypes::e4m3;
1274 if (isa<Float8E5M2Type>(elemType))
1275 return NVVM::WGMMATypes::e5m2;
1277 return NVVM::WGMMATypes::b1;
1279 return NVVM::WGMMATypes::s8;
1281 return NVVM::WGMMATypes::u8;
1283 return NVVM::WGMMATypes::s32;
1284 llvm_unreachable(
"unsupported type");
1286 return NVVM::WGMMATypesAttr::get(op->getContext(), getWgmmaType(type));
1291 generateWgmmaLayout(std::optional<bool> transpose)
const {
1292 if (transpose.value_or(
false))
1293 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::col);
1294 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::row);
1298 NVVM::MMAShapeAttr generateWgmmaShape()
const {
1299 return NVVM::MMAShapeAttr::get(op->getContext(), wgmmaM, wgmmaN, wgmmaK);
1303 NVVM::WGMMAScaleOutAttr generateScaleOut()
const {
1304 return NVVM::WGMMAScaleOutAttr::get(op->getContext(),
1305 NVVM::WGMMAScaleOut::one);
1308 NVVM::WGMMAScaleInAttr generateScaleIn()
const {
1309 return NVVM::WGMMAScaleInAttr::get(op->getContext(),
1310 NVVM::WGMMAScaleIn::one);
1314 Value makeAdd(Value
lhs, Value
rhs) {
1315 return LLVM::AddOp::create(b,
lhs.getType(),
lhs,
rhs);
1336 Value iterateDescriptorA(Value desc,
int i,
int j,
int k) {
1337 MemRefType matrixTypeA = op.getDescriptorA().getType().getTensor();
1338 Type elemA = matrixTypeA.getElementType();
1340 int tileShapeA = matrixTypeA.getDimSize(1);
1341 int incrementVal = ((wgmmaK * k) + (totalK * tileShapeA * i)) *
byte;
1343 LDBG() <<
"\t\t[m: " << i <<
" n: " << j <<
" k: " << k
1344 <<
"] [wgmma descriptors] Descriptor A + " << incrementVal
1348 return makeAdd(desc, makeI64Const(b, incrementVal));
1362 Value iterateDescriptorB(Value desc,
int i,
int j,
int k) {
1363 MemRefType matrixTypeB = op.getDescriptorB().getType().getTensor();
1364 Type elemB = matrixTypeB.getElementType();
1366 int incrementVal = matrixTypeB.getDimSize(0) * wgmmaK * k * byte;
1368 LDBG() <<
"Descriptor B + " << incrementVal;
1371 return makeAdd(desc, makeI64Const(b, incrementVal));
1376 Value generateWgmma(
int i,
int j,
int k, Value matrixC) {
1377 LDBG() <<
"\t wgmma." <<
"m" << wgmmaM <<
"n" << wgmmaN <<
"k" << wgmmaK
1378 <<
"(A[" << (iterationM * wgmmaM) <<
":"
1379 << (iterationM * wgmmaM) + wgmmaM <<
"][" << (iterationK * wgmmaK)
1380 <<
":" << (iterationK * wgmmaK + wgmmaK) <<
"] * " <<
" B["
1381 << (iterationK * wgmmaK) <<
":" << (iterationK * wgmmaK + wgmmaK)
1382 <<
"][" << 0 <<
":" << wgmmaN <<
"])";
1384 Value descriptorA = iterateDescriptorA(adaptor.getDescriptorA(), i, j, k);
1385 Value descriptorB = iterateDescriptorB(adaptor.getDescriptorB(), i, j, k);
1387 Type elemA = op.getDescriptorA().getType().getTensor().getElementType();
1388 NVVM::WGMMATypesAttr itypeA = generateWgmmaType(elemA);
1390 Type elemB = op.getDescriptorB().getType().getTensor().getElementType();
1391 NVVM::WGMMATypesAttr itypeB = generateWgmmaType(elemB);
1393 Type elemD = op.getMatrixC().getType().getFragmented().getElementType();
1394 NVVM::WGMMATypesAttr itypeD = generateWgmmaType(elemD,
true);
1396 NVVM::MMAShapeAttr shape = generateWgmmaShape();
1397 NVVM::WGMMAScaleOutAttr scaleOut = generateScaleOut();
1398 NVVM::WGMMAScaleInAttr scaleIn = generateScaleIn();
1399 NVVM::MMALayoutAttr layoutA = generateWgmmaLayout(op.getTransposeA());
1400 NVVM::MMALayoutAttr layoutB = generateWgmmaLayout(!op.getTransposeB());
1402 auto overflow = NVVM::MMAIntOverflowAttr::get(
1403 op->getContext(), NVVM::MMAIntOverflow::wrapped);
1405 return NVVM::WgmmaMmaAsyncOp::create(
1406 b, matrixC.
getType(), matrixC, descriptorA, descriptorB, shape,
1407 itypeA, itypeB, itypeD, scaleOut, scaleIn, scaleIn, layoutA, layoutB,
1412 Value generateWgmmaGroup() {
1414 LLVM::PoisonOp::create(b, adaptor.getMatrixC().getType());
1417 SmallVector<Value> wgmmaResults;
1418 for (
int i = 0; i < iterationM; ++i) {
1420 LLVM::ExtractValueOp::create(b, adaptor.getMatrixC(), i);
1421 for (
int j = 0; j < iterationN; ++j)
1422 for (
int k = 0; k < iterationK; ++k)
1423 matrixC = generateWgmma(i, j, k, matrixC);
1424 wgmmaResults.push_back(matrixC);
1426 for (
auto [idx, matrix] : llvm::enumerate(wgmmaResults)) {
1427 wgmmaResult = LLVM::InsertValueOp::create(b, wgmmaResult.
getType(),
1428 wgmmaResult, matrix, idx);
1434 WarpgroupGemm(nvgpu::WarpgroupMmaOp op, ImplicitLocOpBuilder &b,
1436 : op(op), b(b), adaptor(adaptor) {
1438 totalM = op.getDescriptorA().
getType().getTensor().getDimSize(0);
1439 totalN = op.getDescriptorB().
getType().getTensor().getDimSize(1);
1440 totalK = op.getDescriptorA().
getType().getTensor().getDimSize(1);
1441 LDBG() <<
"===--- GEMM D[" << totalM <<
"][" << totalN <<
"] += A["
1442 << totalM <<
"][" << totalK <<
"] * B[" << totalK <<
"][" << totalN
1448 op.getDescriptorA().getType().getTensor().getElementType());
1451 iterationM = totalM / wgmmaM;
1452 iterationN = totalN / wgmmaN;
1453 iterationK = totalK / wgmmaK;
1461 Value generateWarpgroupMma() {
1462 NVVM::WgmmaFenceAlignedOp::create(b);
1463 Value wgmmaResult = generateWgmmaGroup();
1464 NVVM::WgmmaGroupSyncAlignedOp::create(b);
1465 NVVM::WgmmaWaitGroupSyncOp::create(b, op.getWaitGroup());
1470 matchAndRewrite(nvgpu::WarpgroupMmaOp op, OpAdaptor adaptor,
1471 ConversionPatternRewriter &rewriter)
const override {
1472 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1475 WarpgroupGemm warpgroupGemm(op,
b, adaptor);
1478 Value wgmmaResult = warpgroupGemm.generateWarpgroupMma();
1481 rewriter.replaceOp(op, wgmmaResult);
1486struct NVGPUWarpgroupMmaStoreOpLowering
1488 using ConvertOpToLLVMPattern<
1489 nvgpu::WarpgroupMmaStoreOp>::ConvertOpToLLVMPattern;
1527 void storeFragmentedMatrix(ImplicitLocOpBuilder &
b, Value matrixD,
1530 Type i32 =
b.getI32Type();
1532 auto makeConst = [&](int32_t index) -> Value {
1533 return LLVM::ConstantOp::create(
b, i32,
b.getI32IntegerAttr(index));
1535 Value c1 = makeConst(1);
1536 Value c2 = makeConst(2);
1537 Value c4 = makeConst(4);
1538 Value c8 = makeConst(8);
1539 Value c16 = makeConst(16);
1542 auto makeMul = [&](Value
lhs, Value
rhs) -> Value {
1543 return LLVM::MulOp::create(
b,
lhs.getType(),
lhs,
rhs);
1545 auto makeAdd = [&](Value
lhs, Value
rhs) -> Value {
1546 return LLVM::AddOp::create(
b,
lhs.getType(),
lhs,
rhs);
1549 auto makeExtractAndStore = [&](
int i, Value wgmmaResult, Value x, Value y,
1551 Type it =
b.getIndexType();
1552 Value idx = arith::IndexCastOp::create(
b, it, x);
1553 Value idy0 = arith::IndexCastOp::create(
b, it, y);
1554 Value idy1 = arith::IndexCastOp::create(
b, it, makeAdd(y, c1));
1555 Value d0 = LLVM::ExtractValueOp::create(
b, wgmmaResult, i);
1556 Value d1 = LLVM::ExtractValueOp::create(
b, wgmmaResult, i + 1);
1557 memref::StoreOp::create(
b, d0, memref,
ValueRange{idx, idy0});
1558 memref::StoreOp::create(
b, d1, memref,
ValueRange{idx, idy1});
1561 Value tidx = NVVM::ThreadIdXOp::create(
b, i32);
1562 Value laneId = LLVM::URemOp::create(
b, i32, tidx, warpSize);
1563 Value warpId = LLVM::UDivOp::create(
b, i32, tidx, warpSize);
1564 Value lane4Id = LLVM::UDivOp::create(
b, i32, laneId, c4);
1565 Value lane4modId = LLVM::URemOp::create(
b, i32, laneId, c4);
1567 Value tj = makeMul(lane4modId, c2);
1568 Value ti = makeAdd(lane4Id, makeMul(warpId, c16));
1570 ti = makeAdd(ti, makeConst(offset));
1572 auto structType = cast<LLVM::LLVMStructType>(matrixD.
getType());
1575 constexpr unsigned numAdjacentRegisters = 2;
1577 constexpr unsigned numStackedMatrices = 2;
1579 size_t storeCount = (structType.getBody().size() /
1580 (numStackedMatrices * numAdjacentRegisters));
1582 for (
size_t i = 0; i < numStackedMatrices; ++i) {
1583 Value idx = makeAdd(ti, makeMul(makeConst(i), c8));
1584 for (
size_t j = 0; j < storeCount; ++j) {
1585 Value idy = makeAdd(tj, makeMul(makeConst(j), c8));
1586 size_t structIndex = (i * numAdjacentRegisters) +
1587 (j * (numStackedMatrices * numAdjacentRegisters));
1588 makeExtractAndStore(structIndex, matrixD, idx, idy, dstMemref);
1594 matchAndRewrite(nvgpu::WarpgroupMmaStoreOp op, OpAdaptor adaptor,
1595 ConversionPatternRewriter &rewriter)
const override {
1597 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1598 Value matriDValue = adaptor.getMatrixD();
1599 auto stype = cast<LLVM::LLVMStructType>(matriDValue.
getType());
1600 for (
auto [idx, matrixD] : llvm::enumerate(stype.getBody())) {
1601 auto structType = cast<LLVM::LLVMStructType>(matrixD);
1602 Value innerStructValue =
1603 LLVM::ExtractValueOp::create(
b, matriDValue, idx);
1604 storeFragmentedMatrix(
b, innerStructValue, op.getDstMemref(), offset);
1605 offset += structType.getBody().size();
1607 rewriter.eraseOp(op);
1612struct NVGPUWarpgroupMmaInitAccumulatorOpLowering
1614 using ConvertOpToLLVMPattern<
1615 nvgpu::WarpgroupMmaInitAccumulatorOp>::ConvertOpToLLVMPattern;
1617 matchAndRewrite(nvgpu::WarpgroupMmaInitAccumulatorOp op, OpAdaptor adaptor,
1618 ConversionPatternRewriter &rewriter)
const override {
1619 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1620 LLVM::LLVMStructType packStructType = cast<LLVM::LLVMStructType>(
1621 getTypeConverter()->convertType(op.getMatrixC().getType()));
1622 Type elemType = cast<LLVM::LLVMStructType>(packStructType.getBody().front())
1625 Value zero = LLVM::ConstantOp::create(
b, elemType,
b.getZeroAttr(elemType));
1626 Value packStruct = LLVM::PoisonOp::create(
b, packStructType);
1627 SmallVector<Value> innerStructs;
1629 for (
auto [idx, s] : llvm::enumerate(packStructType.getBody())) {
1630 auto structType = cast<LLVM::LLVMStructType>(s);
1631 Value structValue = LLVM::ExtractValueOp::create(
b, packStruct, idx);
1632 for (
unsigned i = 0; i < structType.getBody().size(); ++i) {
1633 structValue = LLVM::InsertValueOp::create(
b, structType, structValue,
1634 zero, ArrayRef<int64_t>({i}));
1636 innerStructs.push_back(structValue);
1639 for (
auto [idx, matrix] : llvm::enumerate(innerStructs)) {
1640 packStruct = LLVM::InsertValueOp::create(
b, packStruct.
getType(),
1641 packStruct, matrix, idx);
1643 rewriter.replaceOp(op, packStruct);
1648struct NVGPUTmaFenceOpLowering
1650 using ConvertOpToLLVMPattern<nvgpu::TmaFenceOp>::ConvertOpToLLVMPattern;
1652 matchAndRewrite(nvgpu::TmaFenceOp op, OpAdaptor adaptor,
1653 ConversionPatternRewriter &rewriter)
const override {
1654 MLIRContext *ctx = op.getContext();
1655 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1656 auto i32Ty =
b.getI32Type();
1657 Value tensormapSize =
1658 LLVM::ConstantOp::create(
b, i32Ty, rewriter.getI32IntegerAttr(128));
1661 NVVM::MemScopeKindAttr::get(ctx, ::mlir::NVVM::MemScopeKind::SYS);
1663 rewriter.replaceOpWithNewOp<NVVM::FenceProxyAcquireOp>(
1664 op, memscope, adaptor.getTensorMapDescriptor(), tensormapSize);
1670struct NVGPUTmaPrefetchOpLowering
1672 using ConvertOpToLLVMPattern<nvgpu::TmaPrefetchOp>::ConvertOpToLLVMPattern;
1674 matchAndRewrite(nvgpu::TmaPrefetchOp op, OpAdaptor adaptor,
1675 ConversionPatternRewriter &rewriter)
const override {
1676 rewriter.replaceOpWithNewOp<NVVM::PrefetchOp>(
1677 op,
nullptr,
nullptr,
1678 adaptor.getTensorMapDescriptor(), adaptor.getPredicate(),
1679 mlir::UnitAttr::get(op.getContext()));
1685 using ConvertOpToLLVMPattern<nvgpu::RcpOp>::ConvertOpToLLVMPattern;
1687 matchAndRewrite(nvgpu::RcpOp op, OpAdaptor adaptor,
1688 ConversionPatternRewriter &rewriter)
const override {
1689 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1690 auto i64Ty =
b.getI64Type();
1691 auto f32Ty =
b.getF32Type();
1692 VectorType inTy = op.getIn().getType();
1694 auto convert1DVec = [&](Type llvm1DVectorTy, Value inVec) {
1695 Value ret1DVec = LLVM::PoisonOp::create(
b, llvm1DVectorTy);
1696 int numElems = llvm::cast<VectorType>(llvm1DVectorTy).getNumElements();
1697 for (
int i = 0; i < numElems; i++) {
1698 Value idx = LLVM::ConstantOp::create(
b, i64Ty,
b.getI64IntegerAttr(i));
1699 Value elem = LLVM::ExtractElementOp::create(
b, inVec, idx);
1700 Value dst = NVVM::RcpApproxFtzF32Op::create(
b, f32Ty, elem);
1701 ret1DVec = LLVM::InsertElementOp::create(
b, ret1DVec, dst, idx);
1705 if (inTy.getRank() == 1) {
1706 rewriter.replaceOp(op, convert1DVec(inTy, adaptor.getIn()));
1710 op.getOperation(), adaptor.getOperands(), *(this->getTypeConverter()),
1711 [&](Type llvm1DVectorTy,
ValueRange operands) -> Value {
1712 OpAdaptor adaptor(operands);
1713 return convert1DVec(llvm1DVectorTy, adaptor.getIn());
1723enum class FPKind { F32, BF16, F16, F8, F6, F4 };
1727static int getEffectiveBitWidth(
int bitWidth) {
1728 return bitWidth == 6 ? 8 : bitWidth;
1731static std::optional<FPKind> classifyFPType(
Type t) {
1732 static constexpr auto isConvertibleF8Type = [](
Type t) {
1733 return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
1735 static constexpr auto isConvertibleF6Type = [](
Type t) {
1736 return isa<Float6E2M3FNType, Float6E3M2FNType>(t);
1738 static constexpr auto isConvertibleF4Type = [](
Type t) {
1739 return isa<Float4E2M1FNType>(t);
1745 return FPKind::BF16;
1748 if (isConvertibleF8Type(t))
1750 if (isConvertibleF6Type(t))
1752 if (isConvertibleF4Type(t))
1755 return std::nullopt;
1759enum class FPTruncConvOp {
1773struct FPTruncTableEntry {
1776 FPTruncConvOp convOp;
1779static constexpr FPTruncTableEntry kFPTruncTable[] = {
1781 {FPKind::F32, FPKind::F16, FPTruncConvOp::F32x2_TO_F16x2},
1782 {FPKind::F32, FPKind::BF16, FPTruncConvOp::F32x2_TO_BF16x2},
1783 {FPKind::F32, FPKind::F8, FPTruncConvOp::F32x2_TO_F8x2},
1784 {FPKind::F32, FPKind::F6, FPTruncConvOp::F32x2_TO_F6x2},
1785 {FPKind::F32, FPKind::F4, FPTruncConvOp::F32x2_TO_F4x2},
1787 {FPKind::F16, FPKind::F8, FPTruncConvOp::F16x2_TO_F8x2},
1788 {FPKind::F16, FPKind::F6, FPTruncConvOp::F16x2_TO_F6x2},
1789 {FPKind::F16, FPKind::F4, FPTruncConvOp::F16x2_TO_F4x2},
1791 {FPKind::BF16, FPKind::F8, FPTruncConvOp::BF16x2_TO_F8x2},
1792 {FPKind::BF16, FPKind::F6, FPTruncConvOp::BF16x2_TO_F6x2},
1793 {FPKind::BF16, FPKind::F4, FPTruncConvOp::BF16x2_TO_F4x2},
1798template <
typename TableEntry,
size_t N>
1799static std::optional<TableEntry>
1800lookupConvOp(
const TableEntry (&table)[N],
Type srcElemType,
Type dstElemType) {
1801 std::optional<FPKind> srcKind = classifyFPType(srcElemType);
1802 std::optional<FPKind> dstKind = classifyFPType(dstElemType);
1803 if (!srcKind || !dstKind)
1804 return std::nullopt;
1805 for (
const TableEntry &entry : table) {
1806 if (entry.src == *srcKind && entry.dst == *dstKind)
1809 return std::nullopt;
1815 idx < cast<VectorType>(srcVec.
getType()).getNumElements() &&
1816 "extractElement: index out of bounds");
1817 IntegerType i64Ty =
b.getI64Type();
1818 return b.create<LLVM::ExtractElementOp>(
1819 srcVec,
b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(idx)));
1824 Value srcI32Vec,
int baseIdx) {
1825 FloatType f32Ty =
b.getF32Type();
1826 Value elem0 = extractElement(
b, srcI32Vec, baseIdx);
1827 Value elem1 = extractElement(
b, srcI32Vec, baseIdx + 1);
1828 return {
b.create<LLVM::BitcastOp>(f32Ty, elem0),
1829 b.create<LLVM::BitcastOp>(f32Ty, elem1)};
1835 int idx, VectorType vecTy) {
1836 Value elem = extractElement(
b, srcI32Vec, idx);
1837 return b.create<LLVM::BitcastOp>(vecTy, elem);
1842template <
typename ConvertOp,
typename... Args>
1844 int srcBaseIdx,
Type resultTy, Args &&...args) {
1845 auto [lo, hi] = extractF32Pair(
b, srcI32Vec, srcBaseIdx);
1846 return b.create<ConvertOp>(resultTy, hi, lo, std::forward<Args>(args)...);
1851template <
typename ConvertOp,
typename... Args>
1853 int srcBaseIdx,
Type srcElemTy,
Type resultTy,
1855 Value src = extractAndBitcast(
b, srcI32Vec, srcBaseIdx,
1856 VectorType::get(2, srcElemTy));
1857 return b.create<ConvertOp>(resultTy, src, std::forward<Args>(args)...);
1861static Value createTruncConversion(
1863 Value srcI32Vec,
int srcBaseIdx, NVVM::FPRoundingModeAttr rndAttr,
1864 NVVM::SaturationModeAttr satAttr,
BoolAttr reluAttr,
Type dstElemType,
1866 IntegerType i8Ty =
b.getI8Type();
1867 IntegerType i16Ty =
b.getI16Type();
1868 IntegerType i32Ty =
b.getI32Type();
1869 TypeAttr dstTyAttr = TypeAttr::get(dstElemType);
1870 TypeAttr actualDstTyAttr = TypeAttr::get(actualDstFloatType);
1873 case FPTruncConvOp::F32x2_TO_F16x2: {
1874 auto [lo, hi] = extractF32Pair(
b, srcI32Vec, srcBaseIdx);
1875 Value r =
b.create<NVVM::ConvertF32x2ToF16x2Op>(
1876 VectorType::get(2,
b.getF16Type()), hi, lo, randomBits, rndAttr,
1878 return b.create<LLVM::BitcastOp>(i32Ty, r);
1880 case FPTruncConvOp::F32x2_TO_BF16x2: {
1881 auto [lo, hi] = extractF32Pair(
b, srcI32Vec, srcBaseIdx);
1882 Value r =
b.create<NVVM::ConvertF32x2ToBF16x2Op>(
1883 VectorType::get(2,
b.getBF16Type()), hi, lo, randomBits, rndAttr,
1885 return b.create<LLVM::BitcastOp>(i32Ty, r);
1887 case FPTruncConvOp::F32x2_TO_F8x2:
1888 return convertFromF32Pair<NVVM::ConvertF32x2ToF8x2Op>(
1889 b, srcI32Vec, srcBaseIdx, i16Ty, rndAttr, satAttr, reluAttr, dstTyAttr);
1890 case FPTruncConvOp::F32x2_TO_F6x2:
1891 return convertFromF32Pair<NVVM::ConvertF32x2ToF6x2Op>(
1892 b, srcI32Vec, srcBaseIdx, i16Ty, reluAttr, actualDstTyAttr);
1893 case FPTruncConvOp::F32x2_TO_F4x2:
1894 return convertFromF32Pair<NVVM::ConvertF32x2ToF4x2Op>(
1895 b, srcI32Vec, srcBaseIdx, i8Ty, reluAttr, dstTyAttr);
1896 case FPTruncConvOp::F16x2_TO_F8x2:
1897 return convertFromPacked<NVVM::ConvertF16x2ToF8x2Op>(
1898 b, srcI32Vec, srcBaseIdx,
b.getF16Type(), i16Ty, reluAttr, dstTyAttr);
1899 case FPTruncConvOp::F16x2_TO_F6x2:
1900 return convertFromPacked<NVVM::ConvertF16x2ToF6x2Op>(
1901 b, srcI32Vec, srcBaseIdx,
b.getF16Type(), i16Ty, reluAttr,
1903 case FPTruncConvOp::F16x2_TO_F4x2:
1904 return convertFromPacked<NVVM::ConvertF16x2ToF4x2Op>(
1905 b, srcI32Vec, srcBaseIdx,
b.getF16Type(), i8Ty, reluAttr,
1907 case FPTruncConvOp::BF16x2_TO_F8x2:
1908 return convertFromPacked<NVVM::ConvertBF16x2ToF8x2Op>(
1909 b, srcI32Vec, srcBaseIdx,
b.getBF16Type(), i16Ty, rndAttr, satAttr,
1910 reluAttr, dstTyAttr);
1911 case FPTruncConvOp::BF16x2_TO_F6x2:
1912 return convertFromPacked<NVVM::ConvertBF16x2ToF6x2Op>(
1913 b, srcI32Vec, srcBaseIdx,
b.getBF16Type(), i16Ty, reluAttr,
1915 case FPTruncConvOp::BF16x2_TO_F4x2:
1916 return convertFromPacked<NVVM::ConvertBF16x2ToF4x2Op>(
1917 b, srcI32Vec, srcBaseIdx,
b.getBF16Type(), i8Ty, reluAttr,
1920 llvm_unreachable(
"unhandled FPTruncConvOp");
1923static LogicalResult lowerTruncf(nvgpu::TruncfOp op,
1924 nvgpu::TruncfOp::Adaptor adaptor,
1925 ConversionPatternRewriter &rewriter,
1929 IntegerType i32Ty =
b.getI32Type();
1930 IntegerType i64Ty =
b.getI64Type();
1931 static constexpr int regBits = 32;
1933 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
1934 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
1935 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
1936 return rewriter.notifyMatchFailure(
1937 op,
"expected 1-D vector; canonicalize pattern handles other shapes");
1939 auto srcElemType = srcType.getElementType();
1940 auto dstElemType = dstType.getElementType();
1941 int srcBW = srcType.getElementTypeBitWidth();
1942 int dstBW = dstType.getElementTypeBitWidth();
1943 int numElems = srcType.getNumElements();
1945 NVVM::FPRoundingModeAttr rndModeAttr = op.getRndAttr();
1946 NVVM::SaturationModeAttr satModeAttr = op.getSatAttr();
1947 auto reluBoolAttr = op.getReluAttr();
1948 Value randomBits = adaptor.getRandomBits();
1949 Type actualDstFloatType = dstElemType;
1955 Value input = adaptor.getIn();
1958 Type convertedType = typeConverter->convertType(dstType);
1959 assert(convertedType &&
"failed to convert type");
1960 Value result =
b.create<LLVM::FPTruncOp>(convertedType, input);
1961 rewriter.replaceOp(op,
result);
1964 auto f32VecTy = VectorType::get(srcType.getShape(),
b.getF32Type());
1965 input =
b.create<LLVM::FPTruncOp>(f32VecTy, input);
1967 srcElemType =
b.getF32Type();
1972 int effectiveDstBW = getEffectiveBitWidth(dstBW);
1974 int srcI32Elems = numElems * srcBW / regBits;
1975 int dstI32Elems = numElems * effectiveDstBW / regBits;
1977 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), input);
1979 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
1982 auto convEntry = lookupConvOp(kFPTruncTable, srcElemType, dstElemType);
1984 return rewriter.notifyMatchFailure(
1985 op,
"unsupported type combination for truncation");
1986 FPTruncConvOp convOp = convEntry->convOp;
1989 auto getNumSrcI32PerConvert = [](FPKind src) {
1990 return src == FPKind::F32 ? 2 : 1;
1992 int numSrcI32PerConv = getNumSrcI32PerConvert(convEntry->src);
1995 const int srcStep = srcBW / effectiveDstBW;
1996 const int resultBW =
1998 const int numConvsPerI32 = regBits / resultBW;
2000 for (
int srcIdx = 0, dstIdx = 0; dstIdx < dstI32Elems;
2001 srcIdx += srcStep, dstIdx++) {
2003 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(dstIdx));
2006 if (numConvsPerI32 == 1) {
2008 dstValue = createTruncConversion(
2009 b, ctx, convOp, srcI32Vec, srcIdx, rndModeAttr, satModeAttr,
2010 reluBoolAttr, dstElemType, actualDstFloatType, randomBits);
2013 auto subResultType = IntegerType::get(ctx, resultBW);
2014 auto subVecTy = VectorType::get(numConvsPerI32, subResultType);
2015 Value subVec =
b.create<LLVM::UndefOp>(subVecTy);
2017 int insertIdx = numConvsPerI32 - 1;
2018 int curStep = srcStep;
2019 while (curStep > 0) {
2020 curStep -= numSrcI32PerConv;
2021 Value subResult = createTruncConversion(
2022 b, ctx, convOp, srcI32Vec, srcIdx + curStep, rndModeAttr,
2023 satModeAttr, reluBoolAttr, dstElemType, actualDstFloatType,
2025 subVec =
b.create<LLVM::InsertElementOp>(
2027 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(insertIdx)));
2031 dstValue =
b.create<LLVM::BitcastOp>(i32Ty, subVec);
2035 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2039 Type convertedType = typeConverter->convertType(dstType);
2040 assert(convertedType &&
"failed to convert type");
2041 if (convEntry->dst == FPKind::F6) {
2042 IntegerType i8Ty =
b.getI8Type();
2043 auto i8VecTy = VectorType::get(numElems, i8Ty);
2044 Value i8Vec =
b.create<LLVM::BitcastOp>(i8VecTy, dstI32Vec);
2045 Value truncVec =
b.create<LLVM::TruncOp>(convertedType, i8Vec);
2046 rewriter.replaceOp(op, truncVec);
2048 auto dstVec =
b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2049 rewriter.replaceOp(op, dstVec);
2055 using ConvertOpToLLVMPattern<nvgpu::TruncfOp>::ConvertOpToLLVMPattern;
2058 matchAndRewrite(nvgpu::TruncfOp op, OpAdaptor adaptor,
2059 ConversionPatternRewriter &rewriter)
const override {
2060 return lowerTruncf(op, adaptor, rewriter, getTypeConverter());
2069enum class FPExtConvOp {
2078struct FPExtTableEntry {
2084static constexpr FPExtTableEntry kFPExtTable[] = {
2085 {FPKind::F8, FPKind::F16, FPExtConvOp::F8x2_TO_F16x2},
2086 {FPKind::F8, FPKind::BF16, FPExtConvOp::F8x2_TO_BF16x2},
2087 {FPKind::F6, FPKind::F16, FPExtConvOp::F6x2_TO_F16x2},
2088 {FPKind::F6, FPKind::BF16, FPExtConvOp::F6x2_TO_BF16x2},
2089 {FPKind::F4, FPKind::F16, FPExtConvOp::F4x2_TO_F16x2},
2090 {FPKind::F4, FPKind::BF16, FPExtConvOp::F4x2_TO_BF16x2},
2097 FPExtConvOp convOp,
Value src,
2100 IntegerType i32Ty =
b.getI32Type();
2101 auto srcTyAttr = TypeAttr::get(actualSrcFloatType);
2104 case FPExtConvOp::F8x2_TO_F16x2: {
2105 Value r = NVVM::ConvertF8x2ToF16x2Op::create(
2106 b, VectorType::get(2,
b.getF16Type()), src, srcTyAttr, reluAttr);
2107 return b.create<LLVM::BitcastOp>(i32Ty, r);
2109 case FPExtConvOp::F8x2_TO_BF16x2: {
2110 Value r = NVVM::ConvertF8x2ToBF16x2Op::create(
2111 b, VectorType::get(2,
b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2112 return b.create<LLVM::BitcastOp>(i32Ty, r);
2114 case FPExtConvOp::F6x2_TO_F16x2: {
2115 Value r = NVVM::ConvertF6x2ToF16x2Op::create(
2116 b, VectorType::get(2,
b.getF16Type()), src, srcTyAttr, reluAttr);
2117 return b.create<LLVM::BitcastOp>(i32Ty, r);
2119 case FPExtConvOp::F6x2_TO_BF16x2: {
2120 Value r = NVVM::ConvertF6x2ToBF16x2Op::create(
2121 b, VectorType::get(2,
b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2122 return b.create<LLVM::BitcastOp>(i32Ty, r);
2124 case FPExtConvOp::F4x2_TO_F16x2: {
2125 Value r = NVVM::ConvertF4x2ToF16x2Op::create(
2126 b, VectorType::get(2,
b.getF16Type()), src, srcTyAttr, reluAttr);
2127 return b.create<LLVM::BitcastOp>(i32Ty, r);
2129 case FPExtConvOp::F4x2_TO_BF16x2: {
2130 Value r = NVVM::ConvertF4x2ToBF16x2Op::create(
2131 b, VectorType::get(2,
b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2132 return b.create<LLVM::BitcastOp>(i32Ty, r);
2135 llvm_unreachable(
"unhandled FPExtConvOp");
2138static LogicalResult lowerExtf(nvgpu::ExtfOp op, nvgpu::ExtfOp::Adaptor adaptor,
2139 ConversionPatternRewriter &rewriter,
2143 IntegerType i8Ty =
b.getI8Type();
2144 IntegerType i16Ty =
b.getI16Type();
2145 IntegerType i32Ty =
b.getI32Type();
2146 IntegerType i64Ty =
b.getI64Type();
2148 static constexpr int regBits = 32;
2149 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
2150 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
2151 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
2152 return rewriter.notifyMatchFailure(
2153 op,
"expected 1-D vector; canonicalize pattern handles other shapes");
2155 auto srcElemType = srcType.getElementType();
2156 auto dstElemType = dstType.getElementType();
2157 int srcBW = srcType.getElementTypeBitWidth();
2158 int dstBW = dstType.getElementTypeBitWidth();
2159 int numElems = srcType.getNumElements();
2161 auto reluBoolAttr = op.getReluAttr();
2162 Type actualSrcFloatType = srcElemType;
2164 assert(dstBW == 16 || dstBW == 32 || dstBW == 64);
2167 if (srcBW >= 16 && dstBW >= 32) {
2169 if (srcElemType != dstElemType) {
2170 Type convertedType = typeConverter->convertType(dstType);
2171 assert(convertedType &&
"failed to convert type");
2174 rewriter.replaceOp(op,
result);
2180 bool needsFinalFPExt = (dstBW >= 32);
2181 Type intermediateDstElem = dstElemType;
2182 if (needsFinalFPExt && llvm::isa<Float8E8M0FNUType>(srcElemType))
2183 intermediateDstElem =
b.getBF16Type();
2184 else if (needsFinalFPExt)
2185 intermediateDstElem =
b.getF16Type();
2186 int intermediateDstBW = needsFinalFPExt ? 16 : dstBW;
2189 int effectiveSrcBW = getEffectiveBitWidth(srcBW);
2193 Value inputVec = adaptor.getIn();
2195 auto i8VecTy = VectorType::get(numElems, i8Ty);
2196 inputVec =
b.create<LLVM::ZExtOp>(i8VecTy, inputVec);
2199 int srcI32Elems = numElems * effectiveSrcBW / regBits;
2200 int dstI32Elems = numElems * intermediateDstBW / regBits;
2202 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), inputVec);
2204 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
2207 auto convEntry = lookupConvOp(kFPExtTable, srcElemType, intermediateDstElem);
2209 return rewriter.notifyMatchFailure(
2210 op,
"unsupported type combination for extension");
2211 FPExtConvOp convOp = convEntry->convOp;
2212 Value extScaleFactor;
2215 for (
int srcIdx = 0, dstIdx = 0; srcIdx < srcI32Elems; srcIdx++) {
2216 Value srcI32 =
b.create<LLVM::ExtractElementOp>(
2218 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(srcIdx)));
2220 if (effectiveSrcBW == 8) {
2223 b.create<LLVM::BitcastOp>(VectorType::get(2, i16Ty), srcI32);
2224 for (
int half = 0; half < 2; half++) {
2225 Value halfI16 =
b.create<LLVM::ExtractElementOp>(
2227 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(half)));
2229 b.create<LLVM::BitcastOp>(VectorType::get(2, i8Ty), halfI16);
2231 createExtConversion(
b, ctx, convOp, src, reluBoolAttr,
2232 actualSrcFloatType, extScaleFactor);
2234 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(dstIdx));
2236 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2241 Value i8Vec =
b.create<LLVM::BitcastOp>(VectorType::get(4, i8Ty), srcI32);
2242 for (
int byteIdx = 0; byteIdx < 4; byteIdx++) {
2243 Value src =
b.create<LLVM::ExtractElementOp>(
2245 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(byteIdx)));
2247 createExtConversion(
b, ctx, convOp, src, reluBoolAttr,
2248 actualSrcFloatType, extScaleFactor);
2250 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(dstIdx));
2252 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2259 Type convertedType = typeConverter->convertType(dstType);
2260 assert(convertedType &&
"failed to convert type");
2262 if (needsFinalFPExt) {
2263 auto intermediateVecTy = VectorType::get(numElems, intermediateDstElem);
2264 Value intermediateVec =
2265 b.create<LLVM::BitcastOp>(intermediateVecTy, dstI32Vec);
2266 result =
b.create<LLVM::FPExtOp>(convertedType, intermediateVec);
2268 result =
b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2270 rewriter.replaceOp(op,
result);
2275 using ConvertOpToLLVMPattern<nvgpu::ExtfOp>::ConvertOpToLLVMPattern;
2278 matchAndRewrite(nvgpu::ExtfOp op, OpAdaptor adaptor,
2279 ConversionPatternRewriter &rewriter)
const override {
2280 return lowerExtf(op, adaptor, rewriter, getTypeConverter());
2284static int64_t computePaddedElems(
int64_t numElems,
int srcBW,
int dstBW,
2286 static constexpr int regBits = 32;
2287 int effSrcBW = getEffectiveBitWidth(srcBW);
2288 int effDstBW = getEffectiveBitWidth(dstBW);
2289 auto ceilDiv = [](
int64_t x,
int64_t y) {
return (x + y - 1) / y; };
2291 std::max(ceilDiv(numElems * effSrcBW, regBits) * regBits / effSrcBW,
2292 ceilDiv(numElems * effDstBW, regBits) * regBits / effDstBW);
2293 return ceilDiv(padded, step) * step;
2299template <
typename CvtOp,
bool IsTrunc>
2301 using OpRewritePattern<CvtOp>::OpRewritePattern;
2303 LogicalResult matchAndRewrite(CvtOp op,
2304 PatternRewriter &rewriter)
const override {
2305 Type inType = op.getIn().getType();
2306 Type outType = op.getOut().getType();
2312 int effSrcBW = getEffectiveBitWidth(srcBW);
2313 int effDstBW = getEffectiveBitWidth(dstBW);
2315 bool isScalar = !isa<VectorType>(inType);
2316 auto srcVecTy = dyn_cast<VectorType>(inType);
2317 bool isMultiRank = srcVecTy && srcVecTy.getRank() > 1;
2318 int64_t numElems = isScalar ? 1 : srcVecTy.getNumElements();
2319 int step = IsTrunc ? effSrcBW / effDstBW : effDstBW / effSrcBW;
2320 int64_t paddedElems = computePaddedElems(numElems, srcBW, dstBW, step);
2321 bool needsPad = (paddedElems != numElems);
2323 if (!isScalar && !isMultiRank && !needsPad)
2326 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
2327 Value input = op.getIn();
2330 input = vector::BroadcastOp::create(
b, VectorType::get({1}, srcElemTy),
2333 input = vector::ShapeCastOp::create(
2334 b, VectorType::get({numElems}, srcElemTy), input);
2337 auto paddedTy = VectorType::get({paddedElems}, srcElemTy);
2338 Value zero = arith::ConstantOp::create(
2340 input = vector::InsertStridedSliceOp::create(
2341 b, input, zero, SmallVector<int64_t>{0}, SmallVector<int64_t>{1});
2345 VectorType::get({needsPad ? paddedElems : numElems}, dstElemTy);
2347 if constexpr (IsTrunc) {
2348 cvt = CvtOp::create(
b, cvtDstTy, input, op.getRndAttr(), op.getSatAttr(),
2349 op.getReluAttr(), op.getRandomBits());
2352 CvtOp::create(
b, cvtDstTy, input, op.getRndAttr(), op.getReluAttr());
2357 result = vector::ExtractStridedSliceOp::create(
2358 b,
result, SmallVector<int64_t>{0}, SmallVector<int64_t>{numElems},
2359 SmallVector<int64_t>{1});
2364 vector::ShapeCastOp::create(
b, cast<VectorType>(outType),
result);
2368 result = vector::ExtractOp::create(
b,
result, SmallVector<int64_t>{0});
2376using NVGPUTruncfCanonicalizePattern =
2377 NVGPUFPCanonicalizePattern<nvgpu::TruncfOp, true>;
2378using NVGPUExtfCanonicalizePattern =
2379 NVGPUFPCanonicalizePattern<nvgpu::ExtfOp, false>;
2389 typeConverter, [](gpu::AddressSpace space) ->
unsigned {
2391 case gpu::AddressSpace::Global:
2392 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Global);
2393 case gpu::AddressSpace::Workgroup:
2394 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Shared);
2395 case gpu::AddressSpace::Private:
2397 case gpu::AddressSpace::Constant:
2398 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Constant);
2400 llvm_unreachable(
"unknown address space enum value");
2407 NVGPUMBarrierCreateLowering,
2408 NVGPUMBarrierInitLowering,
2409 NVGPUMBarrierGetLowering,
2410 NVGPUMBarrierArriveLowering,
2411 NVGPUMBarrierArriveNoCompleteLowering,
2412 NVGPUMBarrierTestWaitLowering,
2413 NVGPUMBarrierTryWaitParityLowering,
2414 NVGPUTmaAsyncLoadOpLowering,
2415 NVGPUTmaAsyncStoreOpLowering,
2416 NVGPUTmaCreateDescriptorOpLowering,
2417 NVGPUTmaPrefetchOpLowering,
2418 NVGPUTmaFenceOpLowering,
2419 NVGPUMBarrierArriveExpectTxLowering,
2420 NVGPUGenerateWarpgroupDescriptorLowering,
2421 NVGPUWarpgroupMmaOpLowering,
2422 NVGPUWarpgroupMmaStoreOpLowering,
2423 NVGPUWarpgroupMmaInitAccumulatorOpLowering,
2424 NVGPUTruncfOpLowering,
2425 NVGPUExtfOpLowering,
2426 MmaSyncOptoNVVM, MmaLdMatrixOpToNVVM, NVGPUAsyncCopyLowering,
2427 NVGPUAsyncCreateGroupLowering, NVGPUAsyncWaitLowering,
2428 NVGPUMmaSparseSyncLowering, NVGPURcpOpLowering>(converter);
2430 patterns.
add<NVGPUTruncfCanonicalizePattern, NVGPUExtfCanonicalizePattern>(
constexpr int kWgmmaSizeM
M size of wgmma.mma_async instruction.
static Value truncToI32(ImplicitLocOpBuilder &b, Value value)
GPU has 32 bit registers, this function truncates values when larger width is not needed.
static SmallVector< Value > unpackOperandVector(ImplicitLocOpBuilder &b, Value operand, NVVM::MMATypes operandPtxType)
The gpu.mma.sync converter below expects matrix fragment operands to be given as 2D vectors where the...
static Type inferIntrinsicResultType(Type vectorResultType)
Returns the type for the intrinsic given the vectorResultType of the gpu.mma.sync operation.
constexpr int exclude4LSB
Number of bits that needs to be excluded when building matrix descriptor for wgmma operations.
static bool isMbarrierShared(nvgpu::MBarrierGroupType barrierType)
Returns whether mbarrier object has shared memory address space.
static Value convertIntrinsicResult(Location loc, Type intrinsicResultType, Type resultType, Value intrinsicResult, RewriterBase &rewriter)
Convert the SSA result of the NVVM intrinsic nvvm.mma.sync (which is always an LLVM struct) into a fr...
static llvm::ManagedStatic< PassManagerOptions > options
Attributes are known-constant values of operations.
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
IntegerAttr getI32IntegerAttr(int32_t value)
MLIRContext * getContext() const
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
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.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
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.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Location getLoc()
The source location the operation was defined or derived from.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
MLIRContext * getContext() const
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 coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
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 isSignlessInteger() const
Return true if this is a signless 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 represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
LogicalResult handleMultidimensionalVectors(Operation *op, ValueRange operands, const LLVMTypeConverter &typeConverter, std::function< Value(Type, ValueRange)> createOperand, ConversionPatternRewriter &rewriter)
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...
void populateCommonGPUTypeAndAttributeConversions(TypeConverter &typeConverter)
Remap common GPU memory spaces (Workgroup, Private, etc) to LLVM address spaces.
MemRefType getMBarrierMemrefType(MLIRContext *context, MBarrierGroupType barrierType)
Return the memref type that can be used to represent an mbarrier object.
Attribute getMbarrierMemorySpace(MLIRContext *context, MBarrierGroupType barrierType)
Returns the memory space attribute of the mbarrier object.
void populateSCFStructuralTypeConversionsAndLegality(const TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, PatternBenefit benefit=1)
Populates patterns for SCF structural type conversions and sets up the provided ConversionTarget with...
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
void populateNVGPUToNVVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns)
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
void populateGpuMemorySpaceAttributeConversions(TypeConverter &typeConverter, const MemorySpaceMapping &mapping)
Populates memory space attribute conversion rules for lowering gpu.address_space to integer values.
LLVM::CallOp create(Location loc, OpBuilder &builder, ArrayRef< Value > arguments) const
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...