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.getTf32Enabled().value_or(
false);
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
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,
575struct NVGPUMmaSparseSyncLowering
577 using ConvertOpToLLVMPattern<nvgpu::MmaSparseSyncOp>::ConvertOpToLLVMPattern;
580 matchAndRewrite(nvgpu::MmaSparseSyncOp op, OpAdaptor adaptor,
581 ConversionPatternRewriter &rewriter)
const override {
582 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
585 VectorType aType = op.getMatrixA().getType();
586 VectorType bType = op.getMatrixB().getType();
587 VectorType cType = op.getMatrixC().getType();
589 FailureOr<NVVM::MMATypes> ptxTypeA = getNvvmMmaType(aType);
591 return op->emitOpError(
"failed to deduce operand PTX types");
592 FailureOr<NVVM::MMATypes> ptxTypeB = getNvvmMmaType(bType);
594 return op->emitOpError(
"failed to deduce operand PTX types");
595 std::optional<NVVM::MMATypes> ptxTypeC =
596 NVVM::MmaOp::inferOperandMMAType(cType.getElementType(),
599 return op->emitError(
600 "could not infer the PTX type for the accumulator/result");
603 bool tf32Enabled = op.getTf32Enabled().value_or(
false);
604 if (aType.getElementType().isF32() && !tf32Enabled)
608 std::optional<NVVM::MMAIntOverflow> overflow(std::nullopt);
609 if (isa<IntegerType>(aType.getElementType()))
610 overflow = NVVM::MMAIntOverflow::satfinite;
612 SmallVector<Value> matA =
614 SmallVector<Value> matB =
616 SmallVector<Value> matC =
619 Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
621 typeConverter->convertType(op->getResultTypes()[0]));
624 Value sparseMetadata = adaptor.getSparseMetadata();
625 if (sparseMetadata.
getType() != VectorType::get(2, rewriter.getI16Type()))
626 return op->emitOpError() <<
"Expected metadata type to be LLVM "
627 "VectorType of 2 i16 elements";
629 LLVM::BitcastOp::create(
b, rewriter.getI32Type(), sparseMetadata);
631 FailureOr<LLVM::InlineAsmOp> intrinsicResult = emitMmaSparseSyncOpAsm(
632 b, *ptxTypeA, *ptxTypeB, *ptxTypeC, *ptxTypeC, overflow, matA, matB,
633 matC, sparseMetadata, op.getSparsitySelector(), op.getMmaShapeAsArray(),
635 if (
failed(intrinsicResult))
638 assert((*intrinsicResult).getNumResults() == 1 &&
639 "expected inline asm op returns a single LLVM struct type");
642 (*intrinsicResult)->getResult(0), rewriter));
647struct NVGPUAsyncCopyLowering
649 using ConvertOpToLLVMPattern<
650 nvgpu::DeviceAsyncCopyOp>::ConvertOpToLLVMPattern;
653 matchAndRewrite(nvgpu::DeviceAsyncCopyOp op, OpAdaptor adaptor,
654 ConversionPatternRewriter &rewriter)
const override {
655 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
656 Location loc = op.getLoc();
657 auto dstMemrefType = cast<MemRefType>(op.getDst().getType());
660 adaptor.getDst(), adaptor.getDstIndices());
661 FailureOr<unsigned> dstAddressSpace =
662 getTypeConverter()->getMemRefAddressSpace(dstMemrefType);
663 if (
failed(dstAddressSpace))
664 return rewriter.notifyMatchFailure(
665 loc,
"destination memref address space not convertible to integer");
667 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
668 FailureOr<unsigned> srcAddressSpace =
669 getTypeConverter()->getMemRefAddressSpace(srcMemrefType);
670 if (
failed(srcAddressSpace))
671 return rewriter.notifyMatchFailure(
672 loc,
"source memref address space not convertible to integer");
676 adaptor.getSrcIndices());
678 auto srcPointerGlobalType = LLVM::LLVMPointerType::get(
679 op->getContext(),
static_cast<unsigned>(NVVM::NVVMMemorySpace::Global));
680 scrPtr = LLVM::AddrSpaceCastOp::create(
b, srcPointerGlobalType, scrPtr);
681 int64_t dstElements = adaptor.getDstElements().getZExtValue();
682 int64_t sizeInBytes =
683 (dstMemrefType.getElementTypeBitWidth() * dstElements) / 8;
688 Value srcBytes = adaptor.getSrcElements();
695 LLVM::ConstantOp::create(
b,
b.getI32Type(),
b.getI32IntegerAttr(3));
696 Value bitwidth = LLVM::ConstantOp::create(
698 b.getI32IntegerAttr(srcMemrefType.getElementTypeBitWidth()));
699 Value srcElementsI32 = LLVM::TruncOp::create(
b,
b.getI32Type(), srcBytes);
700 srcBytes = LLVM::LShrOp::create(
701 b, LLVM::MulOp::create(
b, bitwidth, srcElementsI32), c3I32);
705 NVVM::LoadCacheModifierKind cacheModifier =
706 (op.getBypassL1().value_or(
false) && sizeInBytes == 16)
707 ? NVVM::LoadCacheModifierKind::CG
708 : NVVM::LoadCacheModifierKind::CA;
710 NVVM::CpAsyncOp::create(
711 b, dstPtr, scrPtr, rewriter.getI32IntegerAttr(sizeInBytes),
712 NVVM::LoadCacheModifierKindAttr::get(op->getContext(), cacheModifier),
717 LLVM::ConstantOp::create(
b, IntegerType::get(op.getContext(), 32),
718 rewriter.getI32IntegerAttr(0));
719 rewriter.replaceOp(op, zero);
724struct NVGPUAsyncCreateGroupLowering
726 using ConvertOpToLLVMPattern<
727 nvgpu::DeviceAsyncCreateGroupOp>::ConvertOpToLLVMPattern;
730 matchAndRewrite(nvgpu::DeviceAsyncCreateGroupOp op, OpAdaptor adaptor,
731 ConversionPatternRewriter &rewriter)
const override {
732 NVVM::CpAsyncCommitGroupOp::create(rewriter, op.getLoc());
734 Value zero = LLVM::ConstantOp::create(rewriter, op->getLoc(),
735 IntegerType::get(op.getContext(), 32),
736 rewriter.getI32IntegerAttr(0));
737 rewriter.replaceOp(op, zero);
742struct NVGPUAsyncWaitLowering
744 using ConvertOpToLLVMPattern<
745 nvgpu::DeviceAsyncWaitOp>::ConvertOpToLLVMPattern;
748 matchAndRewrite(nvgpu::DeviceAsyncWaitOp op, OpAdaptor adaptor,
749 ConversionPatternRewriter &rewriter)
const override {
751 int32_t numGroups = adaptor.getNumGroups().value_or(0);
752 NVVM::CpAsyncWaitGroupOp::create(rewriter, op.getLoc(), numGroups);
753 rewriter.eraseOp(op);
759struct NVGPUMBarrierCreateLowering
761 using ConvertOpToLLVMPattern<nvgpu::MBarrierCreateOp>::ConvertOpToLLVMPattern;
763 template <
typename moduleT>
764 memref::GlobalOp generateGlobalBarrier(ConversionPatternRewriter &rewriter,
765 Operation *funcOp, moduleT moduleOp,
766 MemRefType barrierType)
const {
767 SymbolTable symbolTable(moduleOp);
768 OpBuilder::InsertionGuard guard(rewriter);
769 rewriter.setInsertionPoint(&moduleOp.front());
770 auto global = memref::GlobalOp::create(
771 rewriter, funcOp->
getLoc(),
"__mbarrier",
772 rewriter.getStringAttr(
"private"),
776 rewriter.getI64IntegerAttr(8));
777 symbolTable.insert(global);
782 matchAndRewrite(nvgpu::MBarrierCreateOp op, OpAdaptor adaptor,
783 ConversionPatternRewriter &rewriter)
const override {
786 rewriter.getContext(), op.getBarriers().getType());
788 memref::GlobalOp global;
790 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
792 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
794 rewriter.setInsertionPoint(op);
795 rewriter.replaceOpWithNewOp<memref::GetGlobalOp>(op, barrierType,
802template <
typename SourceOp>
805 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
807 Value getMbarrierPtr(ImplicitLocOpBuilder &
b,
808 nvgpu::MBarrierGroupType mbarType, Value memrefDesc,
810 ConversionPatternRewriter &rewriter)
const {
811 MemRefType mbarrierMemrefType =
814 rewriter,
b.getLoc(), mbarrierMemrefType, memrefDesc, {mbarId});
818struct NVGPUMBarrierGetLowering
819 :
public MBarrierBasePattern<nvgpu::MBarrierGetOp> {
820 using MBarrierBasePattern<nvgpu::MBarrierGetOp>::MBarrierBasePattern;
823 matchAndRewrite(nvgpu::MBarrierGetOp op, OpAdaptor adaptor,
824 ConversionPatternRewriter &rewriter)
const override {
825 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
826 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
827 rewriter.setInsertionPoint(op);
828 Value barrier = getMbarrierPtr(
b, mbarrierType, adaptor.getBarriers(),
829 adaptor.getMbarId(), rewriter);
830 Type resType = op.getMbarrierPointer().getType();
831 rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(op, resType, barrier);
837struct NVGPUMBarrierInitLowering
838 :
public MBarrierBasePattern<nvgpu::MBarrierInitOp> {
839 using MBarrierBasePattern<nvgpu::MBarrierInitOp>::MBarrierBasePattern;
842 matchAndRewrite(nvgpu::MBarrierInitOp op, OpAdaptor adaptor,
843 ConversionPatternRewriter &rewriter)
const override {
844 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
845 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
846 rewriter.setInsertionPoint(op);
847 Value barrier = getMbarrierPtr(
b, mbarrierType, adaptor.getBarriers(),
848 adaptor.getMbarId(), rewriter);
850 rewriter.replaceOpWithNewOp<NVVM::MBarrierInitOp>(op, barrier, count, 0,
851 adaptor.getPredicate());
857struct NVGPUMBarrierArriveLowering
858 :
public MBarrierBasePattern<nvgpu::MBarrierArriveOp> {
859 using MBarrierBasePattern<nvgpu::MBarrierArriveOp>::MBarrierBasePattern;
861 matchAndRewrite(nvgpu::MBarrierArriveOp op, OpAdaptor adaptor,
862 ConversionPatternRewriter &rewriter)
const override {
863 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
865 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
866 adaptor.getMbarId(), rewriter);
867 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveOp>(op, barrier, Value{});
874struct NVGPUMBarrierArriveNoCompleteLowering
875 :
public MBarrierBasePattern<nvgpu::MBarrierArriveNoCompleteOp> {
876 using MBarrierBasePattern<
877 nvgpu::MBarrierArriveNoCompleteOp>::MBarrierBasePattern;
879 matchAndRewrite(nvgpu::MBarrierArriveNoCompleteOp op, OpAdaptor adaptor,
880 ConversionPatternRewriter &rewriter)
const override {
881 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
883 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
884 adaptor.getMbarId(), rewriter);
885 Type tokenType = getTypeConverter()->convertType(
886 nvgpu::MBarrierTokenType::get(op->getContext()));
888 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveNocompleteOp>(
889 op, tokenType, barrier, count);
895struct NVGPUMBarrierTestWaitLowering
896 :
public MBarrierBasePattern<nvgpu::MBarrierTestWaitOp> {
897 using MBarrierBasePattern<nvgpu::MBarrierTestWaitOp>::MBarrierBasePattern;
899 matchAndRewrite(nvgpu::MBarrierTestWaitOp op, OpAdaptor adaptor,
900 ConversionPatternRewriter &rewriter)
const override {
901 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
903 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
904 adaptor.getMbarId(), rewriter);
905 Type retType = rewriter.getI1Type();
906 rewriter.replaceOpWithNewOp<NVVM::MBarrierTestWaitOp>(op, retType, barrier,
912struct NVGPUMBarrierArriveExpectTxLowering
913 :
public MBarrierBasePattern<nvgpu::MBarrierArriveExpectTxOp> {
914 using MBarrierBasePattern<
915 nvgpu::MBarrierArriveExpectTxOp>::MBarrierBasePattern;
917 matchAndRewrite(nvgpu::MBarrierArriveExpectTxOp op, OpAdaptor adaptor,
918 ConversionPatternRewriter &rewriter)
const override {
919 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
921 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
922 adaptor.getMbarId(), rewriter);
923 Value txcount =
truncToI32(
b, adaptor.getTxcount());
924 NVVM::MBarrierArriveExpectTxOp::create(
925 rewriter, op->getLoc(), barrier, txcount,
926 NVVM::MemScopeKind::CTA,
928 adaptor.getPredicate());
929 rewriter.eraseOp(op);
934struct NVGPUMBarrierTryWaitParityLowering
935 :
public MBarrierBasePattern<nvgpu::MBarrierTryWaitParityOp> {
936 using MBarrierBasePattern<
937 nvgpu::MBarrierTryWaitParityOp>::MBarrierBasePattern;
939 matchAndRewrite(nvgpu::MBarrierTryWaitParityOp op, OpAdaptor adaptor,
940 ConversionPatternRewriter &rewriter)
const override {
941 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
943 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
944 adaptor.getMbarId(), rewriter);
947 LLVM::ZExtOp::create(
b,
b.getI32Type(), adaptor.getPhaseParity());
948 rewriter.replaceOpWithNewOp<NVVM::MBarrierTryWaitParityOp>(op, barrier,
954struct NVGPUTmaAsyncLoadOpLowering
955 :
public MBarrierBasePattern<nvgpu::TmaAsyncLoadOp> {
956 using MBarrierBasePattern<nvgpu::TmaAsyncLoadOp>::MBarrierBasePattern;
958 matchAndRewrite(nvgpu::TmaAsyncLoadOp op, OpAdaptor adaptor,
959 ConversionPatternRewriter &rewriter)
const override {
960 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
961 auto srcMemrefType = cast<MemRefType>(op.getDst().getType());
963 adaptor.getDst(), {});
967 auto ptrSharedClusterType = LLVM::LLVMPointerType::get(
969 static_cast<unsigned>(NVVM::NVVMMemorySpace::SharedCluster));
970 dest = LLVM::AddrSpaceCastOp::create(
b, ptrSharedClusterType, dest);
973 getMbarrierPtr(
b, op.getBarriers().getType(), adaptor.getBarriers(),
974 adaptor.getMbarId(), rewriter);
976 SmallVector<Value> coords = adaptor.getCoordinates();
977 for (
auto [index, value] : llvm::enumerate(coords)) {
982 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(
983 op, dest, adaptor.getTensorMapDescriptor(), coords, barrier,
984 ValueRange{}, adaptor.getMulticastMask(), Value{},
985 NVVM::TMALoadMode::TILE,
988 adaptor.getPredicate());
993struct NVGPUTmaAsyncStoreOpLowering
994 :
public MBarrierBasePattern<nvgpu::TmaAsyncStoreOp> {
995 using MBarrierBasePattern<nvgpu::TmaAsyncStoreOp>::MBarrierBasePattern;
997 matchAndRewrite(nvgpu::TmaAsyncStoreOp op, OpAdaptor adaptor,
998 ConversionPatternRewriter &rewriter)
const override {
999 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1000 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
1002 adaptor.getSrc(), {});
1003 SmallVector<Value> coords = adaptor.getCoordinates();
1004 for (
auto [index, value] : llvm::enumerate(coords)) {
1009 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(
1010 op, adaptor.getTensorMapDescriptor(), dest, coords, Value{},
1011 NVVM::TMAStoreMode::TILE,
1012 adaptor.getPredicate());
1017struct NVGPUGenerateWarpgroupDescriptorLowering
1019 using ConvertOpToLLVMPattern<
1020 nvgpu::WarpgroupGenerateDescriptorOp>::ConvertOpToLLVMPattern;
1023 matchAndRewrite(nvgpu::WarpgroupGenerateDescriptorOp op, OpAdaptor adaptor,
1024 ConversionPatternRewriter &rewriter)
const override {
1026 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1028 nvgpu::TensorMapSwizzleKind swizzleKind =
1029 op.getTensorMap().getType().getSwizzle();
1032 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 128
1033 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 64
1034 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 32
1037 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 1
1038 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 2
1039 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 3
1042 auto ti64 =
b.getIntegerType(64);
1043 auto makeConst = [&](uint64_t index) -> Value {
1044 return LLVM::ConstantOp::create(
b, ti64,
b.getI64IntegerAttr(index));
1046 auto shiftLeft = [&](Value value,
unsigned shift) -> Value {
1047 return LLVM::ShlOp::create(
b, ti64, value, makeConst(shift));
1049 auto shiftRight = [&](Value value,
unsigned shift) -> Value {
1050 return LLVM::LShrOp::create(
b, ti64, value, makeConst(shift));
1052 auto insertBit = [&](Value desc, Value val,
int startBit) {
1053 return LLVM::OrOp::create(
b, ti64, desc, shiftLeft(val, startBit));
1056 int64_t sizeN = op.getTensorMap().
getType().getTensor().getDimSize(0);
1057 uint64_t strideDimVal = (layout << 3) >>
exclude4LSB;
1058 uint64_t leadDimVal = (sizeN * layout) >>
exclude4LSB;
1059 uint64_t offsetVal = 0;
1061 Value strideDim = makeConst(strideDimVal);
1062 Value leadDim = makeConst(leadDimVal);
1065 rewriter, op->getLoc(), cast<MemRefType>(op.getTensor().getType()),
1066 adaptor.getTensor(), {});
1067 Value basePtr = LLVM::PtrToIntOp::create(
b, ti64, baseAddr);
1069 Value basePtr14bit = shiftRight(shiftLeft(basePtr, 46), 50);
1071 int startSwizzleBit = 62, startOffsetBit = 49, startStrideBit = 32,
1072 startLeadBit = 16, startBaseAddrBit = 0;
1073 Value dsc = makeConst(0);
1075 dsc = insertBit(dsc, makeConst(swizzle), startSwizzleBit);
1077 dsc = insertBit(dsc, makeConst(offsetVal), startOffsetBit);
1079 dsc = insertBit(dsc, strideDim, startStrideBit);
1081 dsc = insertBit(dsc, leadDim, startLeadBit);
1083 dsc = insertBit(dsc, basePtr14bit, startBaseAddrBit);
1085 LDBG() <<
"Generating warpgroup.descriptor: " <<
"leading_off:"
1086 << leadDimVal <<
"\t" <<
"stride_off :" << strideDimVal <<
"\t"
1087 <<
"base_offset:" << offsetVal <<
"\t" <<
"layout_type:" << swizzle
1088 <<
" (" << nvgpu::stringifyTensorMapSwizzleKind(swizzleKind)
1089 <<
")\n start_addr : " << baseAddr;
1091 rewriter.replaceOp(op, dsc);
1097 return LLVM::ConstantOp::create(
b,
b.getIntegerType(64),
1098 b.getI64IntegerAttr(
index));
1105 enum CUtensorMapDataTypeEnum {
1106 CU_TENSOR_MAP_DATA_TYPE_UINT8 = 0,
1107 CU_TENSOR_MAP_DATA_TYPE_UINT16,
1108 CU_TENSOR_MAP_DATA_TYPE_UINT32,
1109 CU_TENSOR_MAP_DATA_TYPE_INT32,
1110 CU_TENSOR_MAP_DATA_TYPE_UINT64,
1111 CU_TENSOR_MAP_DATA_TYPE_INT64,
1112 CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
1113 CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
1114 CU_TENSOR_MAP_DATA_TYPE_FLOAT64,
1115 CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
1116 CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ,
1117 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32,
1118 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ
1122 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT8);
1124 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT16);
1126 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT32);
1128 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_UINT64);
1130 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_INT32);
1132 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_INT64);
1134 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_FLOAT16);
1136 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_FLOAT32);
1138 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_FLOAT64);
1140 return makeI64Const(
b, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16);
1142 llvm_unreachable(
"Not supported data type");
1145struct NVGPUTmaCreateDescriptorOpLowering
1147 using ConvertOpToLLVMPattern<
1148 nvgpu::TmaCreateDescriptorOp>::ConvertOpToLLVMPattern;
1150 matchAndRewrite(nvgpu::TmaCreateDescriptorOp op, OpAdaptor adaptor,
1151 ConversionPatternRewriter &rewriter)
const override {
1152 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1153 auto llvmPointerType = LLVM::LLVMPointerType::get(op->getContext());
1154 Type llvmInt64Type = IntegerType::get(op->getContext(), 64);
1156 Value tensorElementType =
1157 elementTypeAsLLVMConstant(
b, op.getTensor().getType().getElementType());
1158 auto promotedOperands = getTypeConverter()->promoteOperands(
1159 b.getLoc(), op->getOperands(), adaptor.getOperands(),
b);
1161 Value boxArrayPtr = LLVM::AllocaOp::create(
1162 b, llvmPointerType, llvmInt64Type, makeI64Const(
b, 5));
1163 for (
auto [index, value] : llvm::enumerate(adaptor.getBoxDimensions())) {
1164 Value gep = LLVM::GEPOp::create(
b, llvmPointerType, llvmPointerType,
1165 boxArrayPtr, makeI64Const(
b, index));
1166 LLVM::StoreOp::create(
b, value, gep);
1169 nvgpu::TensorMapDescriptorType desc = op.getTensorMap().
getType();
1171 SmallVector<Value> arguments;
1172 arguments.push_back(promotedOperands[0]);
1173 arguments.push_back(promotedOperands[1]);
1174 arguments.push_back(tensorElementType);
1175 arguments.push_back(
1176 makeI64Const(
b, (
int)desc.getInterleave()));
1177 arguments.push_back(makeI64Const(
b, (
int)desc.getSwizzle()));
1178 arguments.push_back(makeI64Const(
b, (
int)desc.getL2promo()));
1179 arguments.push_back(makeI64Const(
b, (
int)desc.getOob()));
1180 arguments.push_back(boxArrayPtr);
1183 SmallVector<Type> argTypes = {
1193 FunctionCallBuilder hostRegisterCallBuilder = {
1194 "mgpuTensorMapEncodeTiledMemref", llvmPointerType, argTypes};
1196 hostRegisterCallBuilder.
create(
b.getLoc(),
b, arguments).getResult();
1198 rewriter.replaceOp(op, tensorMap);
1203struct NVGPUWarpgroupMmaOpLowering
1205 using ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaOp>::ConvertOpToLLVMPattern;
1227 class WarpgroupGemm {
1228 nvgpu::WarpgroupMmaOp op;
1229 ImplicitLocOpBuilder b;
1233 int64_t totalM, totalN, totalK;
1236 int wgmmaM = 0, wgmmaN = 0, wgmmaK = 0;
1239 int iterationM = 0, iterationN = 0, iterationK = 0;
1244 void findWgmmaShape(int64_t sizeM, int64_t sizeN, Type inputElemType) {
1247 if (inputElemType.
isTF32()) {
1249 }
else if (inputElemType.
isF16() || inputElemType.
isBF16()) {
1251 }
else if (isa<Float8E4M3FNType, Float8E5M2Type>(inputElemType) ||
1254 }
else if (inputElemType.
isInteger(1)) {
1257 llvm_unreachable(
"msg: not supported K shape");
1259 LDBG() <<
"Generating WgmmaMmaAsyncOp shape[m = " << wgmmaM
1260 <<
", n = " << wgmmaN <<
", k = " << wgmmaK <<
"]";
1264 NVVM::WGMMATypesAttr generateWgmmaType(Type type,
1265 bool useF32 =
false)
const {
1266 auto getWgmmaType = [=](Type elemType) {
1268 return useF32 ? NVVM::WGMMATypes::f32 : NVVM::WGMMATypes::tf32;
1269 if (elemType.
isF16())
1270 return NVVM::WGMMATypes::f16;
1272 return NVVM::WGMMATypes::bf16;
1273 if (isa<Float8E4M3FNType>(elemType))
1274 return NVVM::WGMMATypes::e4m3;
1275 if (isa<Float8E5M2Type>(elemType))
1276 return NVVM::WGMMATypes::e5m2;
1278 return NVVM::WGMMATypes::b1;
1280 return NVVM::WGMMATypes::s8;
1282 return NVVM::WGMMATypes::u8;
1284 return NVVM::WGMMATypes::s32;
1285 llvm_unreachable(
"unsupported type");
1287 return NVVM::WGMMATypesAttr::get(op->getContext(), getWgmmaType(type));
1292 generateWgmmaLayout(std::optional<bool> transpose)
const {
1293 if (transpose.value_or(
false))
1294 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::col);
1295 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::row);
1299 NVVM::MMAShapeAttr generateWgmmaShape()
const {
1300 return NVVM::MMAShapeAttr::get(op->getContext(), wgmmaM, wgmmaN, wgmmaK);
1304 NVVM::WGMMAScaleOutAttr generateScaleOut()
const {
1305 return NVVM::WGMMAScaleOutAttr::get(op->getContext(),
1306 NVVM::WGMMAScaleOut::one);
1309 NVVM::WGMMAScaleInAttr generateScaleIn()
const {
1310 return NVVM::WGMMAScaleInAttr::get(op->getContext(),
1311 NVVM::WGMMAScaleIn::one);
1315 Value makeAdd(Value
lhs, Value
rhs) {
1316 return LLVM::AddOp::create(b,
lhs.getType(),
lhs,
rhs);
1337 Value iterateDescriptorA(Value desc,
int i,
int j,
int k) {
1338 MemRefType matrixTypeA = op.getDescriptorA().getType().getTensor();
1339 Type elemA = matrixTypeA.getElementType();
1341 int tileShapeA = matrixTypeA.getDimSize(1);
1342 int incrementVal = ((wgmmaK * k) + (totalK * tileShapeA * i)) *
byte;
1344 LDBG() <<
"\t\t[m: " << i <<
" n: " << j <<
" k: " << k
1345 <<
"] [wgmma descriptors] Descriptor A + " << incrementVal
1349 return makeAdd(desc, makeI64Const(b, incrementVal));
1363 Value iterateDescriptorB(Value desc,
int i,
int j,
int k) {
1364 MemRefType matrixTypeB = op.getDescriptorB().getType().getTensor();
1365 Type elemB = matrixTypeB.getElementType();
1367 int incrementVal = matrixTypeB.getDimSize(0) * wgmmaK * k * byte;
1369 LDBG() <<
"Descriptor B + " << incrementVal;
1372 return makeAdd(desc, makeI64Const(b, incrementVal));
1377 Value generateWgmma(
int i,
int j,
int k, Value matrixC) {
1378 LDBG() <<
"\t wgmma." <<
"m" << wgmmaM <<
"n" << wgmmaN <<
"k" << wgmmaK
1379 <<
"(A[" << (iterationM * wgmmaM) <<
":"
1380 << (iterationM * wgmmaM) + wgmmaM <<
"][" << (iterationK * wgmmaK)
1381 <<
":" << (iterationK * wgmmaK + wgmmaK) <<
"] * " <<
" B["
1382 << (iterationK * wgmmaK) <<
":" << (iterationK * wgmmaK + wgmmaK)
1383 <<
"][" << 0 <<
":" << wgmmaN <<
"])";
1385 Value descriptorA = iterateDescriptorA(adaptor.getDescriptorA(), i, j, k);
1386 Value descriptorB = iterateDescriptorB(adaptor.getDescriptorB(), i, j, k);
1388 Type elemA = op.getDescriptorA().getType().getTensor().getElementType();
1389 NVVM::WGMMATypesAttr itypeA = generateWgmmaType(elemA);
1391 Type elemB = op.getDescriptorB().getType().getTensor().getElementType();
1392 NVVM::WGMMATypesAttr itypeB = generateWgmmaType(elemB);
1394 Type elemD = op.getMatrixC().getType().getFragmented().getElementType();
1395 NVVM::WGMMATypesAttr itypeD = generateWgmmaType(elemD,
true);
1397 NVVM::MMAShapeAttr shape = generateWgmmaShape();
1398 NVVM::WGMMAScaleOutAttr scaleOut = generateScaleOut();
1399 NVVM::WGMMAScaleInAttr scaleIn = generateScaleIn();
1400 NVVM::MMALayoutAttr layoutA = generateWgmmaLayout(op.getTransposeA());
1401 NVVM::MMALayoutAttr layoutB = generateWgmmaLayout(!op.getTransposeB());
1403 auto overflow = NVVM::MMAIntOverflowAttr::get(
1404 op->getContext(), NVVM::MMAIntOverflow::wrapped);
1406 return NVVM::WgmmaMmaAsyncOp::create(
1407 b, matrixC.
getType(), matrixC, descriptorA, descriptorB, shape,
1408 itypeA, itypeB, itypeD, scaleOut, scaleIn, scaleIn, layoutA, layoutB,
1413 Value generateWgmmaGroup() {
1415 LLVM::PoisonOp::create(b, adaptor.getMatrixC().getType());
1418 SmallVector<Value> wgmmaResults;
1419 for (
int i = 0; i < iterationM; ++i) {
1421 LLVM::ExtractValueOp::create(b, adaptor.getMatrixC(), i);
1422 for (
int j = 0; j < iterationN; ++j)
1423 for (
int k = 0; k < iterationK; ++k)
1424 matrixC = generateWgmma(i, j, k, matrixC);
1425 wgmmaResults.push_back(matrixC);
1427 for (
auto [idx, matrix] : llvm::enumerate(wgmmaResults)) {
1428 wgmmaResult = LLVM::InsertValueOp::create(b, wgmmaResult.
getType(),
1429 wgmmaResult, matrix, idx);
1435 WarpgroupGemm(nvgpu::WarpgroupMmaOp op, ImplicitLocOpBuilder &b,
1437 : op(op), b(b), adaptor(adaptor) {
1439 totalM = op.getDescriptorA().
getType().getTensor().getDimSize(0);
1440 totalN = op.getDescriptorB().
getType().getTensor().getDimSize(1);
1441 totalK = op.getDescriptorA().
getType().getTensor().getDimSize(1);
1442 LDBG() <<
"===--- GEMM D[" << totalM <<
"][" << totalN <<
"] += A["
1443 << totalM <<
"][" << totalK <<
"] * B[" << totalK <<
"][" << totalN
1449 op.getDescriptorA().getType().getTensor().getElementType());
1452 iterationM = totalM / wgmmaM;
1453 iterationN = totalN / wgmmaN;
1454 iterationK = totalK / wgmmaK;
1462 Value generateWarpgroupMma() {
1463 NVVM::WgmmaFenceAlignedOp::create(b);
1464 Value wgmmaResult = generateWgmmaGroup();
1465 NVVM::WgmmaGroupSyncAlignedOp::create(b);
1466 NVVM::WgmmaWaitGroupSyncOp::create(b, op.getWaitGroup());
1471 matchAndRewrite(nvgpu::WarpgroupMmaOp op, OpAdaptor adaptor,
1472 ConversionPatternRewriter &rewriter)
const override {
1473 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1476 WarpgroupGemm warpgroupGemm(op,
b, adaptor);
1479 Value wgmmaResult = warpgroupGemm.generateWarpgroupMma();
1482 rewriter.replaceOp(op, wgmmaResult);
1487struct NVGPUWarpgroupMmaStoreOpLowering
1489 using ConvertOpToLLVMPattern<
1490 nvgpu::WarpgroupMmaStoreOp>::ConvertOpToLLVMPattern;
1528 void storeFragmentedMatrix(ImplicitLocOpBuilder &
b, Value matrixD,
1531 Type i32 =
b.getI32Type();
1533 auto makeConst = [&](int32_t index) -> Value {
1534 return LLVM::ConstantOp::create(
b, i32,
b.getI32IntegerAttr(index));
1536 Value c1 = makeConst(1);
1537 Value c2 = makeConst(2);
1538 Value c4 = makeConst(4);
1539 Value c8 = makeConst(8);
1540 Value c16 = makeConst(16);
1543 auto makeMul = [&](Value
lhs, Value
rhs) -> Value {
1544 return LLVM::MulOp::create(
b,
lhs.getType(),
lhs,
rhs);
1546 auto makeAdd = [&](Value
lhs, Value
rhs) -> Value {
1547 return LLVM::AddOp::create(
b,
lhs.getType(),
lhs,
rhs);
1550 auto makeExtractAndStore = [&](
int i, Value wgmmaResult, Value x, Value y,
1552 Type it =
b.getIndexType();
1553 Value idx = arith::IndexCastOp::create(
b, it, x);
1554 Value idy0 = arith::IndexCastOp::create(
b, it, y);
1555 Value idy1 = arith::IndexCastOp::create(
b, it, makeAdd(y, c1));
1556 Value d0 = LLVM::ExtractValueOp::create(
b, wgmmaResult, i);
1557 Value d1 = LLVM::ExtractValueOp::create(
b, wgmmaResult, i + 1);
1558 memref::StoreOp::create(
b, d0, memref,
ValueRange{idx, idy0});
1559 memref::StoreOp::create(
b, d1, memref,
ValueRange{idx, idy1});
1562 Value tidx = NVVM::ThreadIdXOp::create(
b, i32);
1563 Value laneId = LLVM::URemOp::create(
b, i32, tidx, warpSize);
1564 Value warpId = LLVM::UDivOp::create(
b, i32, tidx, warpSize);
1565 Value lane4Id = LLVM::UDivOp::create(
b, i32, laneId, c4);
1566 Value lane4modId = LLVM::URemOp::create(
b, i32, laneId, c4);
1568 Value tj = makeMul(lane4modId, c2);
1569 Value ti = makeAdd(lane4Id, makeMul(warpId, c16));
1571 ti = makeAdd(ti, makeConst(offset));
1573 auto structType = cast<LLVM::LLVMStructType>(matrixD.
getType());
1576 constexpr unsigned numAdjacentRegisters = 2;
1578 constexpr unsigned numStackedMatrices = 2;
1580 size_t storeCount = (structType.getBody().size() /
1581 (numStackedMatrices * numAdjacentRegisters));
1583 for (
size_t i = 0; i < numStackedMatrices; ++i) {
1584 Value idx = makeAdd(ti, makeMul(makeConst(i), c8));
1585 for (
size_t j = 0; j < storeCount; ++j) {
1586 Value idy = makeAdd(tj, makeMul(makeConst(j), c8));
1587 size_t structIndex = (i * numAdjacentRegisters) +
1588 (j * (numStackedMatrices * numAdjacentRegisters));
1589 makeExtractAndStore(structIndex, matrixD, idx, idy, dstMemref);
1595 matchAndRewrite(nvgpu::WarpgroupMmaStoreOp op, OpAdaptor adaptor,
1596 ConversionPatternRewriter &rewriter)
const override {
1598 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1599 Value matriDValue = adaptor.getMatrixD();
1600 auto stype = cast<LLVM::LLVMStructType>(matriDValue.
getType());
1601 for (
auto [idx, matrixD] : llvm::enumerate(stype.getBody())) {
1602 auto structType = cast<LLVM::LLVMStructType>(matrixD);
1603 Value innerStructValue =
1604 LLVM::ExtractValueOp::create(
b, matriDValue, idx);
1605 storeFragmentedMatrix(
b, innerStructValue, op.getDstMemref(), offset);
1606 offset += structType.getBody().size();
1608 rewriter.eraseOp(op);
1613struct NVGPUWarpgroupMmaInitAccumulatorOpLowering
1615 using ConvertOpToLLVMPattern<
1616 nvgpu::WarpgroupMmaInitAccumulatorOp>::ConvertOpToLLVMPattern;
1618 matchAndRewrite(nvgpu::WarpgroupMmaInitAccumulatorOp op, OpAdaptor adaptor,
1619 ConversionPatternRewriter &rewriter)
const override {
1620 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1621 LLVM::LLVMStructType packStructType = cast<LLVM::LLVMStructType>(
1622 getTypeConverter()->convertType(op.getMatrixC().getType()));
1623 Type elemType = cast<LLVM::LLVMStructType>(packStructType.getBody().front())
1626 Value zero = LLVM::ConstantOp::create(
b, elemType,
b.getZeroAttr(elemType));
1627 Value packStruct = LLVM::PoisonOp::create(
b, packStructType);
1628 SmallVector<Value> innerStructs;
1630 for (
auto [idx, s] : llvm::enumerate(packStructType.getBody())) {
1631 auto structType = cast<LLVM::LLVMStructType>(s);
1632 Value structValue = LLVM::ExtractValueOp::create(
b, packStruct, idx);
1633 for (
unsigned i = 0; i < structType.getBody().size(); ++i) {
1634 structValue = LLVM::InsertValueOp::create(
b, structType, structValue,
1635 zero, ArrayRef<int64_t>({i}));
1637 innerStructs.push_back(structValue);
1640 for (
auto [idx, matrix] : llvm::enumerate(innerStructs)) {
1641 packStruct = LLVM::InsertValueOp::create(
b, packStruct.
getType(),
1642 packStruct, matrix, idx);
1644 rewriter.replaceOp(op, packStruct);
1649struct NVGPUTmaFenceOpLowering
1651 using ConvertOpToLLVMPattern<nvgpu::TmaFenceOp>::ConvertOpToLLVMPattern;
1653 matchAndRewrite(nvgpu::TmaFenceOp op, OpAdaptor adaptor,
1654 ConversionPatternRewriter &rewriter)
const override {
1655 MLIRContext *ctx = op.getContext();
1656 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1657 auto i32Ty =
b.getI32Type();
1658 Value tensormapSize =
1659 LLVM::ConstantOp::create(
b, i32Ty, rewriter.getI32IntegerAttr(128));
1662 NVVM::MemScopeKindAttr::get(ctx, ::mlir::NVVM::MemScopeKind::SYS);
1664 rewriter.replaceOpWithNewOp<NVVM::FenceProxyAcquireOp>(
1665 op, memscope, adaptor.getTensorMapDescriptor(), tensormapSize);
1671struct NVGPUTmaPrefetchOpLowering
1673 using ConvertOpToLLVMPattern<nvgpu::TmaPrefetchOp>::ConvertOpToLLVMPattern;
1675 matchAndRewrite(nvgpu::TmaPrefetchOp op, OpAdaptor adaptor,
1676 ConversionPatternRewriter &rewriter)
const override {
1677 rewriter.replaceOpWithNewOp<NVVM::PrefetchOp>(
1678 op,
nullptr,
nullptr,
1679 adaptor.getTensorMapDescriptor(), adaptor.getPredicate(),
1680 mlir::UnitAttr::get(op.getContext()));
1686 using ConvertOpToLLVMPattern<nvgpu::RcpOp>::ConvertOpToLLVMPattern;
1688 matchAndRewrite(nvgpu::RcpOp op, OpAdaptor adaptor,
1689 ConversionPatternRewriter &rewriter)
const override {
1690 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
1691 auto i64Ty =
b.getI64Type();
1692 auto f32Ty =
b.getF32Type();
1693 VectorType inTy = op.getIn().getType();
1695 auto convert1DVec = [&](Type llvm1DVectorTy, Value inVec) {
1696 Value ret1DVec = LLVM::PoisonOp::create(
b, llvm1DVectorTy);
1697 int numElems = llvm::cast<VectorType>(llvm1DVectorTy).getNumElements();
1698 for (
int i = 0; i < numElems; i++) {
1699 Value idx = LLVM::ConstantOp::create(
b, i64Ty,
b.getI64IntegerAttr(i));
1700 Value elem = LLVM::ExtractElementOp::create(
b, inVec, idx);
1701 Value dst = NVVM::RcpApproxFtzF32Op::create(
b, f32Ty, elem);
1702 ret1DVec = LLVM::InsertElementOp::create(
b, ret1DVec, dst, idx);
1706 if (inTy.getRank() == 1) {
1707 rewriter.replaceOp(op, convert1DVec(inTy, adaptor.getIn()));
1711 op.getOperation(), adaptor.getOperands(), *(this->getTypeConverter()),
1712 [&](Type llvm1DVectorTy,
ValueRange operands) -> Value {
1713 OpAdaptor adaptor(operands);
1714 return convert1DVec(llvm1DVectorTy, adaptor.getIn());
1724enum class FPKind { F32, BF16, F16, F8, F6, F4 };
1728static int getEffectiveBitWidth(
int bitWidth) {
1729 return bitWidth == 6 ? 8 : bitWidth;
1732static std::optional<FPKind> classifyFPType(
Type t) {
1733 static constexpr auto isConvertibleF8Type = [](
Type t) {
1734 return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
1736 static constexpr auto isConvertibleF6Type = [](
Type t) {
1737 return isa<Float6E2M3FNType, Float6E3M2FNType>(t);
1739 static constexpr auto isConvertibleF4Type = [](
Type t) {
1740 return isa<Float4E2M1FNType>(t);
1746 return FPKind::BF16;
1749 if (isConvertibleF8Type(t))
1751 if (isConvertibleF6Type(t))
1753 if (isConvertibleF4Type(t))
1756 return std::nullopt;
1760enum class FPTruncConvOp {
1774struct FPTruncTableEntry {
1777 FPTruncConvOp convOp;
1780static constexpr FPTruncTableEntry kFPTruncTable[] = {
1782 {FPKind::F32, FPKind::F16, FPTruncConvOp::F32x2_TO_F16x2},
1783 {FPKind::F32, FPKind::BF16, FPTruncConvOp::F32x2_TO_BF16x2},
1784 {FPKind::F32, FPKind::F8, FPTruncConvOp::F32x2_TO_F8x2},
1785 {FPKind::F32, FPKind::F6, FPTruncConvOp::F32x2_TO_F6x2},
1786 {FPKind::F32, FPKind::F4, FPTruncConvOp::F32x2_TO_F4x2},
1788 {FPKind::F16, FPKind::F8, FPTruncConvOp::F16x2_TO_F8x2},
1789 {FPKind::F16, FPKind::F6, FPTruncConvOp::F16x2_TO_F6x2},
1790 {FPKind::F16, FPKind::F4, FPTruncConvOp::F16x2_TO_F4x2},
1792 {FPKind::BF16, FPKind::F8, FPTruncConvOp::BF16x2_TO_F8x2},
1793 {FPKind::BF16, FPKind::F6, FPTruncConvOp::BF16x2_TO_F6x2},
1794 {FPKind::BF16, FPKind::F4, FPTruncConvOp::BF16x2_TO_F4x2},
1799template <
typename TableEntry,
size_t N>
1800static std::optional<TableEntry>
1801lookupConvOp(
const TableEntry (&table)[N],
Type srcElemType,
Type dstElemType) {
1802 std::optional<FPKind> srcKind = classifyFPType(srcElemType);
1803 std::optional<FPKind> dstKind = classifyFPType(dstElemType);
1804 if (!srcKind || !dstKind)
1805 return std::nullopt;
1806 for (
const TableEntry &entry : table) {
1807 if (entry.src == *srcKind && entry.dst == *dstKind)
1810 return std::nullopt;
1816 idx < cast<VectorType>(srcVec.
getType()).getNumElements() &&
1817 "extractElement: index out of bounds");
1818 IntegerType i64Ty =
b.getI64Type();
1819 return b.create<LLVM::ExtractElementOp>(
1820 srcVec,
b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(idx)));
1825 Value srcI32Vec,
int baseIdx) {
1826 FloatType f32Ty =
b.getF32Type();
1827 Value elem0 = extractElement(
b, srcI32Vec, baseIdx);
1828 Value elem1 = extractElement(
b, srcI32Vec, baseIdx + 1);
1829 return {
b.create<LLVM::BitcastOp>(f32Ty, elem0),
1830 b.create<LLVM::BitcastOp>(f32Ty, elem1)};
1836 int idx, VectorType vecTy) {
1837 Value elem = extractElement(
b, srcI32Vec, idx);
1838 return b.create<LLVM::BitcastOp>(vecTy, elem);
1843template <
typename ConvertOp,
typename... Args>
1845 int srcBaseIdx,
Type resultTy, Args &&...args) {
1846 auto [lo, hi] = extractF32Pair(
b, srcI32Vec, srcBaseIdx);
1847 return b.create<ConvertOp>(resultTy, hi, lo, std::forward<Args>(args)...);
1852template <
typename ConvertOp,
typename... Args>
1854 int srcBaseIdx,
Type srcElemTy,
Type resultTy,
1856 Value src = extractAndBitcast(
b, srcI32Vec, srcBaseIdx,
1857 VectorType::get(2, srcElemTy));
1858 return b.create<ConvertOp>(resultTy, src, std::forward<Args>(args)...);
1862static Value createTruncConversion(
1864 Value srcI32Vec,
int srcBaseIdx, NVVM::FPRoundingModeAttr rndAttr,
1865 NVVM::SaturationModeAttr satAttr,
BoolAttr reluAttr,
Type dstElemType,
1867 IntegerType i8Ty =
b.getI8Type();
1868 IntegerType i16Ty =
b.getI16Type();
1869 IntegerType i32Ty =
b.getI32Type();
1870 TypeAttr dstTyAttr = TypeAttr::get(dstElemType);
1871 TypeAttr actualDstTyAttr = TypeAttr::get(actualDstFloatType);
1874 case FPTruncConvOp::F32x2_TO_F16x2: {
1875 auto [lo, hi] = extractF32Pair(
b, srcI32Vec, srcBaseIdx);
1876 Value r =
b.create<NVVM::ConvertF32x2ToF16x2Op>(
1877 VectorType::get(2,
b.getF16Type()), hi, lo, randomBits, rndAttr,
1879 return b.create<LLVM::BitcastOp>(i32Ty, r);
1881 case FPTruncConvOp::F32x2_TO_BF16x2: {
1882 auto [lo, hi] = extractF32Pair(
b, srcI32Vec, srcBaseIdx);
1883 Value r =
b.create<NVVM::ConvertF32x2ToBF16x2Op>(
1884 VectorType::get(2,
b.getBF16Type()), hi, lo, randomBits, rndAttr,
1886 return b.create<LLVM::BitcastOp>(i32Ty, r);
1888 case FPTruncConvOp::F32x2_TO_F8x2:
1889 return convertFromF32Pair<NVVM::ConvertF32x2ToF8x2Op>(
1890 b, srcI32Vec, srcBaseIdx, i16Ty, rndAttr, satAttr, reluAttr, dstTyAttr);
1891 case FPTruncConvOp::F32x2_TO_F6x2:
1892 return convertFromF32Pair<NVVM::ConvertF32x2ToF6x2Op>(
1893 b, srcI32Vec, srcBaseIdx, i16Ty, reluAttr, actualDstTyAttr);
1894 case FPTruncConvOp::F32x2_TO_F4x2:
1895 return convertFromF32Pair<NVVM::ConvertF32x2ToF4x2Op>(
1896 b, srcI32Vec, srcBaseIdx, i8Ty, reluAttr, dstTyAttr);
1897 case FPTruncConvOp::F16x2_TO_F8x2:
1898 return convertFromPacked<NVVM::ConvertF16x2ToF8x2Op>(
1899 b, srcI32Vec, srcBaseIdx,
b.getF16Type(), i16Ty, reluAttr, dstTyAttr);
1900 case FPTruncConvOp::F16x2_TO_F6x2:
1901 return convertFromPacked<NVVM::ConvertF16x2ToF6x2Op>(
1902 b, srcI32Vec, srcBaseIdx,
b.getF16Type(), i16Ty, reluAttr,
1904 case FPTruncConvOp::F16x2_TO_F4x2:
1905 return convertFromPacked<NVVM::ConvertF16x2ToF4x2Op>(
1906 b, srcI32Vec, srcBaseIdx,
b.getF16Type(), i8Ty, reluAttr,
1908 case FPTruncConvOp::BF16x2_TO_F8x2:
1909 return convertFromPacked<NVVM::ConvertBF16x2ToF8x2Op>(
1910 b, srcI32Vec, srcBaseIdx,
b.getBF16Type(), i16Ty, rndAttr, satAttr,
1911 reluAttr, dstTyAttr);
1912 case FPTruncConvOp::BF16x2_TO_F6x2:
1913 return convertFromPacked<NVVM::ConvertBF16x2ToF6x2Op>(
1914 b, srcI32Vec, srcBaseIdx,
b.getBF16Type(), i16Ty, reluAttr,
1916 case FPTruncConvOp::BF16x2_TO_F4x2:
1917 return convertFromPacked<NVVM::ConvertBF16x2ToF4x2Op>(
1918 b, srcI32Vec, srcBaseIdx,
b.getBF16Type(), i8Ty, reluAttr,
1921 llvm_unreachable(
"unhandled FPTruncConvOp");
1924static LogicalResult lowerTruncf(nvgpu::TruncfOp op,
1925 nvgpu::TruncfOp::Adaptor adaptor,
1926 ConversionPatternRewriter &rewriter,
1930 IntegerType i32Ty =
b.getI32Type();
1931 IntegerType i64Ty =
b.getI64Type();
1932 static constexpr int regBits = 32;
1934 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
1935 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
1936 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
1937 return rewriter.notifyMatchFailure(
1938 op,
"expected 1-D vector; canonicalize pattern handles other shapes");
1940 auto srcElemType = srcType.getElementType();
1941 auto dstElemType = dstType.getElementType();
1942 int srcBW = srcType.getElementTypeBitWidth();
1943 int dstBW = dstType.getElementTypeBitWidth();
1944 int numElems = srcType.getNumElements();
1946 NVVM::FPRoundingModeAttr rndModeAttr = op.getRndAttr();
1947 NVVM::SaturationModeAttr satModeAttr = op.getSatAttr();
1948 auto reluBoolAttr = op.getReluAttr();
1949 Value randomBits = adaptor.getRandomBits();
1950 Type actualDstFloatType = dstElemType;
1956 Value input = adaptor.getIn();
1959 Type convertedType = typeConverter->convertType(dstType);
1960 assert(convertedType &&
"failed to convert type");
1961 Value result =
b.create<LLVM::FPTruncOp>(convertedType, input);
1962 rewriter.replaceOp(op,
result);
1965 auto f32VecTy = VectorType::get(srcType.getShape(),
b.getF32Type());
1966 input =
b.create<LLVM::FPTruncOp>(f32VecTy, input);
1968 srcElemType =
b.getF32Type();
1973 int effectiveDstBW = getEffectiveBitWidth(dstBW);
1975 int srcI32Elems = numElems * srcBW / regBits;
1976 int dstI32Elems = numElems * effectiveDstBW / regBits;
1978 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), input);
1980 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
1983 auto convEntry = lookupConvOp(kFPTruncTable, srcElemType, dstElemType);
1985 return rewriter.notifyMatchFailure(
1986 op,
"unsupported type combination for truncation");
1987 FPTruncConvOp convOp = convEntry->convOp;
1990 auto getNumSrcI32PerConvert = [](FPKind src) {
1991 return src == FPKind::F32 ? 2 : 1;
1993 int numSrcI32PerConv = getNumSrcI32PerConvert(convEntry->src);
1996 const int srcStep = srcBW / effectiveDstBW;
1997 const int resultBW =
1999 const int numConvsPerI32 = regBits / resultBW;
2001 for (
int srcIdx = 0, dstIdx = 0; dstIdx < dstI32Elems;
2002 srcIdx += srcStep, dstIdx++) {
2004 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(dstIdx));
2007 if (numConvsPerI32 == 1) {
2009 dstValue = createTruncConversion(
2010 b, ctx, convOp, srcI32Vec, srcIdx, rndModeAttr, satModeAttr,
2011 reluBoolAttr, dstElemType, actualDstFloatType, randomBits);
2014 auto subResultType = IntegerType::get(ctx, resultBW);
2015 auto subVecTy = VectorType::get(numConvsPerI32, subResultType);
2016 Value subVec =
b.create<LLVM::UndefOp>(subVecTy);
2018 int insertIdx = numConvsPerI32 - 1;
2019 int curStep = srcStep;
2020 while (curStep > 0) {
2021 curStep -= numSrcI32PerConv;
2022 Value subResult = createTruncConversion(
2023 b, ctx, convOp, srcI32Vec, srcIdx + curStep, rndModeAttr,
2024 satModeAttr, reluBoolAttr, dstElemType, actualDstFloatType,
2026 subVec =
b.create<LLVM::InsertElementOp>(
2028 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(insertIdx)));
2032 dstValue =
b.create<LLVM::BitcastOp>(i32Ty, subVec);
2036 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2040 Type convertedType = typeConverter->convertType(dstType);
2041 assert(convertedType &&
"failed to convert type");
2042 if (convEntry->dst == FPKind::F6) {
2043 IntegerType i8Ty =
b.getI8Type();
2044 auto i8VecTy = VectorType::get(numElems, i8Ty);
2045 Value i8Vec =
b.create<LLVM::BitcastOp>(i8VecTy, dstI32Vec);
2046 Value truncVec =
b.create<LLVM::TruncOp>(convertedType, i8Vec);
2047 rewriter.replaceOp(op, truncVec);
2049 auto dstVec =
b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2050 rewriter.replaceOp(op, dstVec);
2056 using ConvertOpToLLVMPattern<nvgpu::TruncfOp>::ConvertOpToLLVMPattern;
2059 matchAndRewrite(nvgpu::TruncfOp op, OpAdaptor adaptor,
2060 ConversionPatternRewriter &rewriter)
const override {
2061 return lowerTruncf(op, adaptor, rewriter, getTypeConverter());
2070enum class FPExtConvOp {
2079struct FPExtTableEntry {
2085static constexpr FPExtTableEntry kFPExtTable[] = {
2086 {FPKind::F8, FPKind::F16, FPExtConvOp::F8x2_TO_F16x2},
2087 {FPKind::F8, FPKind::BF16, FPExtConvOp::F8x2_TO_BF16x2},
2088 {FPKind::F6, FPKind::F16, FPExtConvOp::F6x2_TO_F16x2},
2089 {FPKind::F6, FPKind::BF16, FPExtConvOp::F6x2_TO_BF16x2},
2090 {FPKind::F4, FPKind::F16, FPExtConvOp::F4x2_TO_F16x2},
2091 {FPKind::F4, FPKind::BF16, FPExtConvOp::F4x2_TO_BF16x2},
2098 FPExtConvOp convOp,
Value src,
2101 IntegerType i32Ty =
b.getI32Type();
2102 auto srcTyAttr = TypeAttr::get(actualSrcFloatType);
2105 case FPExtConvOp::F8x2_TO_F16x2: {
2106 Value r = NVVM::ConvertF8x2ToF16x2Op::create(
2107 b, VectorType::get(2,
b.getF16Type()), src, srcTyAttr, reluAttr);
2108 return b.create<LLVM::BitcastOp>(i32Ty, r);
2110 case FPExtConvOp::F8x2_TO_BF16x2: {
2111 Value r = NVVM::ConvertF8x2ToBF16x2Op::create(
2112 b, VectorType::get(2,
b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2113 return b.create<LLVM::BitcastOp>(i32Ty, r);
2115 case FPExtConvOp::F6x2_TO_F16x2: {
2116 Value r = NVVM::ConvertF6x2ToF16x2Op::create(
2117 b, VectorType::get(2,
b.getF16Type()), src, srcTyAttr, reluAttr);
2118 return b.create<LLVM::BitcastOp>(i32Ty, r);
2120 case FPExtConvOp::F6x2_TO_BF16x2: {
2121 Value r = NVVM::ConvertF6x2ToBF16x2Op::create(
2122 b, VectorType::get(2,
b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2123 return b.create<LLVM::BitcastOp>(i32Ty, r);
2125 case FPExtConvOp::F4x2_TO_F16x2: {
2126 Value r = NVVM::ConvertF4x2ToF16x2Op::create(
2127 b, VectorType::get(2,
b.getF16Type()), src, srcTyAttr, reluAttr);
2128 return b.create<LLVM::BitcastOp>(i32Ty, r);
2130 case FPExtConvOp::F4x2_TO_BF16x2: {
2131 Value r = NVVM::ConvertF4x2ToBF16x2Op::create(
2132 b, VectorType::get(2,
b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2133 return b.create<LLVM::BitcastOp>(i32Ty, r);
2136 llvm_unreachable(
"unhandled FPExtConvOp");
2139static LogicalResult lowerExtf(nvgpu::ExtfOp op, nvgpu::ExtfOp::Adaptor adaptor,
2140 ConversionPatternRewriter &rewriter,
2144 IntegerType i8Ty =
b.getI8Type();
2145 IntegerType i16Ty =
b.getI16Type();
2146 IntegerType i32Ty =
b.getI32Type();
2147 IntegerType i64Ty =
b.getI64Type();
2149 static constexpr int regBits = 32;
2150 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
2151 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
2152 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
2153 return rewriter.notifyMatchFailure(
2154 op,
"expected 1-D vector; canonicalize pattern handles other shapes");
2156 auto srcElemType = srcType.getElementType();
2157 auto dstElemType = dstType.getElementType();
2158 int srcBW = srcType.getElementTypeBitWidth();
2159 int dstBW = dstType.getElementTypeBitWidth();
2160 int numElems = srcType.getNumElements();
2162 auto reluBoolAttr = op.getReluAttr();
2163 Type actualSrcFloatType = srcElemType;
2165 assert(dstBW == 16 || dstBW == 32 || dstBW == 64);
2168 if (srcBW >= 16 && dstBW >= 32) {
2170 if (srcElemType != dstElemType) {
2171 Type convertedType = typeConverter->convertType(dstType);
2172 assert(convertedType &&
"failed to convert type");
2175 rewriter.replaceOp(op,
result);
2181 bool needsFinalFPExt = (dstBW >= 32);
2182 Type intermediateDstElem = dstElemType;
2183 if (needsFinalFPExt && llvm::isa<Float8E8M0FNUType>(srcElemType))
2184 intermediateDstElem =
b.getBF16Type();
2185 else if (needsFinalFPExt)
2186 intermediateDstElem =
b.getF16Type();
2187 int intermediateDstBW = needsFinalFPExt ? 16 : dstBW;
2190 int effectiveSrcBW = getEffectiveBitWidth(srcBW);
2194 Value inputVec = adaptor.getIn();
2196 auto i8VecTy = VectorType::get(numElems, i8Ty);
2197 inputVec =
b.create<LLVM::ZExtOp>(i8VecTy, inputVec);
2200 int srcI32Elems = numElems * effectiveSrcBW / regBits;
2201 int dstI32Elems = numElems * intermediateDstBW / regBits;
2203 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), inputVec);
2205 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
2208 auto convEntry = lookupConvOp(kFPExtTable, srcElemType, intermediateDstElem);
2210 return rewriter.notifyMatchFailure(
2211 op,
"unsupported type combination for extension");
2212 FPExtConvOp convOp = convEntry->convOp;
2213 Value extScaleFactor;
2216 for (
int srcIdx = 0, dstIdx = 0; srcIdx < srcI32Elems; srcIdx++) {
2217 Value srcI32 =
b.create<LLVM::ExtractElementOp>(
2219 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(srcIdx)));
2221 if (effectiveSrcBW == 8) {
2224 b.create<LLVM::BitcastOp>(VectorType::get(2, i16Ty), srcI32);
2225 for (
int half = 0; half < 2; half++) {
2226 Value halfI16 =
b.create<LLVM::ExtractElementOp>(
2228 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(half)));
2230 b.create<LLVM::BitcastOp>(VectorType::get(2, i8Ty), halfI16);
2232 createExtConversion(
b, ctx, convOp, src, reluBoolAttr,
2233 actualSrcFloatType, extScaleFactor);
2235 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(dstIdx));
2237 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2242 Value i8Vec =
b.create<LLVM::BitcastOp>(VectorType::get(4, i8Ty), srcI32);
2243 for (
int byteIdx = 0; byteIdx < 4; byteIdx++) {
2244 Value src =
b.create<LLVM::ExtractElementOp>(
2246 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(byteIdx)));
2248 createExtConversion(
b, ctx, convOp, src, reluBoolAttr,
2249 actualSrcFloatType, extScaleFactor);
2251 b.create<LLVM::ConstantOp>(i64Ty,
b.getI64IntegerAttr(dstIdx));
2253 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2260 Type convertedType = typeConverter->convertType(dstType);
2261 assert(convertedType &&
"failed to convert type");
2263 if (needsFinalFPExt) {
2264 auto intermediateVecTy = VectorType::get(numElems, intermediateDstElem);
2265 Value intermediateVec =
2266 b.create<LLVM::BitcastOp>(intermediateVecTy, dstI32Vec);
2267 result =
b.create<LLVM::FPExtOp>(convertedType, intermediateVec);
2269 result =
b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2271 rewriter.replaceOp(op,
result);
2276 using ConvertOpToLLVMPattern<nvgpu::ExtfOp>::ConvertOpToLLVMPattern;
2279 matchAndRewrite(nvgpu::ExtfOp op, OpAdaptor adaptor,
2280 ConversionPatternRewriter &rewriter)
const override {
2281 return lowerExtf(op, adaptor, rewriter, getTypeConverter());
2285static int64_t computePaddedElems(
int64_t numElems,
int srcBW,
int dstBW,
2287 static constexpr int regBits = 32;
2288 int effSrcBW = getEffectiveBitWidth(srcBW);
2289 int effDstBW = getEffectiveBitWidth(dstBW);
2290 auto ceilDiv = [](
int64_t x,
int64_t y) {
return (x + y - 1) / y; };
2292 std::max(ceilDiv(numElems * effSrcBW, regBits) * regBits / effSrcBW,
2293 ceilDiv(numElems * effDstBW, regBits) * regBits / effDstBW);
2294 return ceilDiv(padded, step) * step;
2300template <
typename CvtOp,
bool IsTrunc>
2302 using OpRewritePattern<CvtOp>::OpRewritePattern;
2304 LogicalResult matchAndRewrite(CvtOp op,
2305 PatternRewriter &rewriter)
const override {
2306 Type inType = op.getIn().getType();
2307 Type outType = op.getOut().getType();
2313 int effSrcBW = getEffectiveBitWidth(srcBW);
2314 int effDstBW = getEffectiveBitWidth(dstBW);
2316 bool isScalar = !isa<VectorType>(inType);
2317 auto srcVecTy = dyn_cast<VectorType>(inType);
2318 bool isMultiRank = srcVecTy && srcVecTy.getRank() > 1;
2319 int64_t numElems = isScalar ? 1 : srcVecTy.getNumElements();
2320 int step = IsTrunc ? effSrcBW / effDstBW : effDstBW / effSrcBW;
2321 int64_t paddedElems = computePaddedElems(numElems, srcBW, dstBW, step);
2322 bool needsPad = (paddedElems != numElems);
2324 if (!isScalar && !isMultiRank && !needsPad)
2327 ImplicitLocOpBuilder
b(op->getLoc(), rewriter);
2328 Value input = op.getIn();
2331 input = vector::BroadcastOp::create(
b, VectorType::get({1}, srcElemTy),
2334 input = vector::ShapeCastOp::create(
2335 b, VectorType::get({numElems}, srcElemTy), input);
2338 auto paddedTy = VectorType::get({paddedElems}, srcElemTy);
2339 Value zero = arith::ConstantOp::create(
2341 input = vector::InsertStridedSliceOp::create(
2342 b, input, zero, SmallVector<int64_t>{0}, SmallVector<int64_t>{1});
2346 VectorType::get({needsPad ? paddedElems : numElems}, dstElemTy);
2348 if constexpr (IsTrunc) {
2349 cvt = CvtOp::create(
b, cvtDstTy, input, op.getRndAttr(), op.getSatAttr(),
2350 op.getReluAttr(), op.getRandomBits());
2353 CvtOp::create(
b, cvtDstTy, input, op.getRndAttr(), op.getReluAttr());
2358 result = vector::ExtractStridedSliceOp::create(
2359 b,
result, SmallVector<int64_t>{0}, SmallVector<int64_t>{numElems},
2360 SmallVector<int64_t>{1});
2365 vector::ShapeCastOp::create(
b, cast<VectorType>(outType),
result);
2369 result = vector::ExtractOp::create(
b,
result, SmallVector<int64_t>{0});
2377using NVGPUTruncfCanonicalizePattern =
2378 NVGPUFPCanonicalizePattern<nvgpu::TruncfOp, true>;
2379using NVGPUExtfCanonicalizePattern =
2380 NVGPUFPCanonicalizePattern<nvgpu::ExtfOp, false>;
2390 typeConverter, [](gpu::AddressSpace space) ->
unsigned {
2392 case gpu::AddressSpace::Global:
2393 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Global);
2394 case gpu::AddressSpace::Workgroup:
2395 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Shared);
2396 case gpu::AddressSpace::Private:
2398 case gpu::AddressSpace::Constant:
2399 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Constant);
2401 llvm_unreachable(
"unknown address space enum value");
2408 NVGPUMBarrierCreateLowering,
2409 NVGPUMBarrierInitLowering,
2410 NVGPUMBarrierGetLowering,
2411 NVGPUMBarrierArriveLowering,
2412 NVGPUMBarrierArriveNoCompleteLowering,
2413 NVGPUMBarrierTestWaitLowering,
2414 NVGPUMBarrierTryWaitParityLowering,
2415 NVGPUTmaAsyncLoadOpLowering,
2416 NVGPUTmaAsyncStoreOpLowering,
2417 NVGPUTmaCreateDescriptorOpLowering,
2418 NVGPUTmaPrefetchOpLowering,
2419 NVGPUTmaFenceOpLowering,
2420 NVGPUMBarrierArriveExpectTxLowering,
2421 NVGPUGenerateWarpgroupDescriptorLowering,
2422 NVGPUWarpgroupMmaOpLowering,
2423 NVGPUWarpgroupMmaStoreOpLowering,
2424 NVGPUWarpgroupMmaInitAccumulatorOpLowering,
2425 NVGPUTruncfOpLowering,
2426 NVGPUExtfOpLowering,
2427 MmaSyncOptoNVVM, MmaLdMatrixOpToNVVM, NVGPUAsyncCopyLowering,
2428 NVGPUAsyncCreateGroupLowering, NVGPUAsyncWaitLowering,
2429 NVGPUMmaSparseSyncLowering, NVGPURcpOpLowering>(converter);
2431 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...