[Mlir-commits] [mlir] [NVGPU] Add FP8 (e4m3/e5m2) support to nvgpu.mma.sync (PR #207342)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jul 3 08:37:27 PDT 2026
https://github.com/weimin023 updated https://github.com/llvm/llvm-project/pull/207342
>From 6695b3e0d8bddeab555e1ca36802df2decb85e82 Mon Sep 17 00:00:00 2001
From: weimin023 <tnwilly at gmail.com>
Date: Fri, 3 Jul 2026 07:16:56 +0000
Subject: [PATCH 1/2] [NVVM] Add FP8 (e4m3/e5m2) support to dense mma.sync
Signed-off-by: weimin023 <tnwilly at gmail.com>
---
mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td | 6 +++++-
mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp | 9 ++++++++
mlir/test/Dialect/LLVMIR/nvvm.mlir | 24 +++++++++++++++++++++
mlir/test/Target/LLVMIR/nvvmir.mlir | 24 +++++++++++++++++++++
4 files changed, 62 insertions(+), 1 deletion(-)
diff --git a/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td b/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
index 40f7f15b694cb..67e0ef41d927f 100644
--- a/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
+++ b/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
@@ -2705,9 +2705,13 @@ class NVVM_MMA_OPS {
list<list<WMMA_REGS>> bit_mma_ops = MMA_OPS<
[GEOM<8,8,128>, GEOM<16,8,128>, GEOM<16,8,256>],
["b1"], [], ["s32"], []>.ret;
+ list<list<WMMA_REGS>> fp8_mma_ops = MMA_OPS<
+ [GEOM<16,8,16>, GEOM<16,8,32>],
+ ["e4m3", "e5m2"], ["e4m3", "e5m2"], ["f16", "f32"], []>.ret;
list<list<WMMA_REGS>> all_mma_sync_ops = !listconcat(
tf32_mma_ops, bf16_mma_ops, f64_mma_ops,
- fp_mma_ops, int_mma_ops, subint_mma_ops, bit_mma_ops);
+ fp_mma_ops, int_mma_ops, subint_mma_ops, bit_mma_ops,
+ fp8_mma_ops);
list<list<WMMA_REGS>> bf16_mma_sp_ops = MMA_OPS<
[GEOM<16,8,16>, GEOM<16,8,32>],
diff --git a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
index 14406ab55f138..b29cba96a410f 100644
--- a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
+++ b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
@@ -1007,6 +1007,15 @@ LogicalResult MmaOp::verify() {
expectedResult.push_back(f16x2x2StructTy);
expectedResult.push_back(f32x4StructTy);
break;
+ case MMATypes::e4m3:
+ case MMATypes::e5m2:
+ // FP8 (m16n8k16 / m16n8k32) packs 4 values per 32-bit register, same
+ // as s8/u8, but the accumulator is f16 or f32 (not integer).
+ kFactor = 16;
+ multiplicandFragType = i32Ty;
+ expectedResult.push_back(f16x2x2StructTy);
+ expectedResult.push_back(f32x4StructTy);
+ break;
case MMATypes::s4:
case MMATypes::u4:
kFactor = 32;
diff --git a/mlir/test/Dialect/LLVMIR/nvvm.mlir b/mlir/test/Dialect/LLVMIR/nvvm.mlir
index 72935ffd7b7ce..9445770ebb026 100644
--- a/mlir/test/Dialect/LLVMIR/nvvm.mlir
+++ b/mlir/test/Dialect/LLVMIR/nvvm.mlir
@@ -210,6 +210,30 @@ func.func @nvvm_mma_m16n8k16_bf16_bf16(%a0 : i32, %a1 : i32, %a2 : i32, %a3 : i3
llvm.return %0 : !llvm.struct<(f32, f32, f32, f32)>
}
+// CHECK-LABEL: @nvvm_mma_m16n8k16_e4m3_e4m3
+func.func @nvvm_mma_m16n8k16_e4m3_e4m3(%a0 : i32, %a1 : i32,
+ %b0 : i32,
+ %c0 : f32, %c1 : f32, %c2 : f32, %c3 : f32) -> !llvm.struct<(f32, f32, f32, f32)> {
+ // CHECK: nvvm.mma.sync A[{{.*}}] B[{{.*}}] C[{{.*}}] {layoutA = #nvvm.mma_layout<row>, layoutB = #nvvm.mma_layout<col>, multiplicandAPtxType = #nvvm.mma_type<e4m3>, multiplicandBPtxType = #nvvm.mma_type<e4m3>, shape = #nvvm.shape<m = 16, n = 8, k = 16>} : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ %0 = nvvm.mma.sync A[%a0, %a1] B[%b0] C[%c0, %c1, %c2, %c3]
+ {layoutA = #nvvm.mma_layout<row>, layoutB = #nvvm.mma_layout<col>,
+ multiplicandAPtxType = #nvvm.mma_type<e4m3>, multiplicandBPtxType = #nvvm.mma_type<e4m3>,
+ shape = #nvvm.shape<m = 16, n = 8, k = 16>} : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ llvm.return %0 : !llvm.struct<(f32, f32, f32, f32)>
+}
+
+// CHECK-LABEL: @nvvm_mma_m16n8k32_e4m3_e5m2
+func.func @nvvm_mma_m16n8k32_e4m3_e5m2(%a0 : i32, %a1 : i32, %a2 : i32, %a3 : i32,
+ %b0 : i32, %b1 : i32,
+ %c0 : f32, %c1 : f32, %c2 : f32, %c3 : f32) -> !llvm.struct<(f32, f32, f32, f32)> {
+ // CHECK: nvvm.mma.sync A[{{.*}}] B[{{.*}}] C[{{.*}}] {layoutA = #nvvm.mma_layout<row>, layoutB = #nvvm.mma_layout<col>, multiplicandAPtxType = #nvvm.mma_type<e4m3>, multiplicandBPtxType = #nvvm.mma_type<e5m2>, shape = #nvvm.shape<m = 16, n = 8, k = 32>} : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ %0 = nvvm.mma.sync A[%a0, %a1, %a2, %a3] B[%b0, %b1] C[%c0, %c1, %c2, %c3]
+ {layoutA = #nvvm.mma_layout<row>, layoutB = #nvvm.mma_layout<col>,
+ multiplicandAPtxType = #nvvm.mma_type<e4m3>, multiplicandBPtxType = #nvvm.mma_type<e5m2>,
+ shape = #nvvm.shape<m = 16, n = 8, k = 32>} : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ llvm.return %0 : !llvm.struct<(f32, f32, f32, f32)>
+}
+
// CHECK-LABEL: @nvvm_mma_m8n8k16_s8_s8
func.func @nvvm_mma_m8n8k16_s8_s8(%a0 : i32, %b0 : i32,
%c0 : i32, %c1 : i32) -> !llvm.struct<(i32, i32)> {
diff --git a/mlir/test/Target/LLVMIR/nvvmir.mlir b/mlir/test/Target/LLVMIR/nvvmir.mlir
index f2888025d8a08..f726efe3d39b7 100644
--- a/mlir/test/Target/LLVMIR/nvvmir.mlir
+++ b/mlir/test/Target/LLVMIR/nvvmir.mlir
@@ -293,6 +293,30 @@ llvm.func @nvvm_mma_m16n8k16_bf16_bf16(%a0 : i32, %a1 : i32, %a2 : i32, %a3 : i3
llvm.return %0 : !llvm.struct<(f32, f32, f32, f32)>
}
+// CHECK-LABEL: @nvvm_mma_m16n8k16_e4m3_e4m3
+llvm.func @nvvm_mma_m16n8k16_e4m3_e4m3(%a0 : i32, %a1 : i32,
+ %b0 : i32,
+ %c0 : f32, %c1 : f32, %c2 : f32, %c3 : f32) -> !llvm.struct<(f32, f32, f32, f32)> {
+ // CHECK: call { float, float, float, float } @llvm.nvvm.mma.m16n8k16.row.col.f32.e4m3.e4m3.f32
+ %0 = nvvm.mma.sync A[%a0, %a1] B[%b0] C[%c0, %c1, %c2, %c3]
+ {layoutA = #nvvm.mma_layout<row>, layoutB = #nvvm.mma_layout<col>,
+ multiplicandAPtxType = #nvvm.mma_type<e4m3>, multiplicandBPtxType = #nvvm.mma_type<e4m3>,
+ shape = #nvvm.shape<m = 16, n = 8, k = 16>} : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ llvm.return %0 : !llvm.struct<(f32, f32, f32, f32)>
+}
+
+// CHECK-LABEL: @nvvm_mma_m16n8k32_e4m3_e5m2
+llvm.func @nvvm_mma_m16n8k32_e4m3_e5m2(%a0 : i32, %a1 : i32, %a2 : i32, %a3 : i32,
+ %b0 : i32, %b1 : i32,
+ %c0 : f32, %c1 : f32, %c2 : f32, %c3 : f32) -> !llvm.struct<(f32, f32, f32, f32)> {
+ // CHECK: call { float, float, float, float } @llvm.nvvm.mma.m16n8k32.row.col.f32.e4m3.e5m2.f32
+ %0 = nvvm.mma.sync A[%a0, %a1, %a2, %a3] B[%b0, %b1] C[%c0, %c1, %c2, %c3]
+ {layoutA = #nvvm.mma_layout<row>, layoutB = #nvvm.mma_layout<col>,
+ multiplicandAPtxType = #nvvm.mma_type<e4m3>, multiplicandBPtxType = #nvvm.mma_type<e5m2>,
+ shape = #nvvm.shape<m = 16, n = 8, k = 32>} : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ llvm.return %0 : !llvm.struct<(f32, f32, f32, f32)>
+}
+
// f32 return type, f32 accumulate type
// CHECK-LABEL: @nvvm_mma_m16n8k16_f32_f32
llvm.func @nvvm_mma_m16n8k16_f32_f32(%a0 : vector<2xf16>, %a1 : vector<2xf16>,
>From eacc2df90b292589350c3fcdfec6327a54b273ff Mon Sep 17 00:00:00 2001
From: weimin023 <tnwilly at gmail.com>
Date: Fri, 3 Jul 2026 07:43:32 +0000
Subject: [PATCH 2/2] [NVGPU] Add FP8 (e4m3/e5m2) support to nvgpu.mma.sync
Signed-off-by: weimin023 <tnwilly at gmail.com>
---
mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp | 4 ++++
mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp | 7 ++++---
.../Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir | 12 ++++++++++++
3 files changed, 20 insertions(+), 3 deletions(-)
create mode 100644 mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir
diff --git a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
index e566449ffadff..c647fdc1e7a75 100644
--- a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
+++ b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
@@ -329,6 +329,10 @@ static FailureOr<NVVM::MMATypes> getNvvmMmaType(Type t) {
return NVVM::MMATypes::f64;
if (elType.isF32())
return NVVM::MMATypes::tf32;
+ if (elType.isF8E4M3FN())
+ return NVVM::MMATypes::e4m3;
+ if (elType.isF8E5M2())
+ return NVVM::MMATypes::e5m2;
return failure();
}
diff --git a/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp b/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
index 2b2686a9163bf..198fd958693fb 100644
--- a/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
+++ b/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
@@ -184,7 +184,8 @@ static LogicalResult verifyMmaSyncOp(Operation *op,
numElementA = 1;
numElementB = 1;
} else if (aType.isF32() || aType.isBF16() || aType.isF16() ||
- aType.isInteger(8) || aType.isInteger(4)) {
+ aType.isInteger(8) || aType.isInteger(4) ||
+ aType.isF8E4M3FN() || aType.isF8E5M2()) {
// 8-by-8-128b fundamental tensor core tile size
int operandBitwidth = aType.getIntOrFloatBitWidth();
shapeK = 128 / operandBitwidth; // 128b wide shapeK
@@ -193,8 +194,8 @@ static LogicalResult verifyMmaSyncOp(Operation *op,
numElementB = 32 / operandBitwidth; // 32b wide operand B
} else {
return op->emitError()
- << "expected input data type (i4,i8,f16,bf16,tf32,f64) "
- "supported by "
+ << "expected input data type (i4,i8,f16,bf16,tf32,f64,"
+ "f8E4M3FN,f8E5M2) supported by "
<< op->getName();
}
diff --git a/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir
new file mode 100644
index 0000000000000..90c81af872f84
--- /dev/null
+++ b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir
@@ -0,0 +1,12 @@
+// RUN: mlir-opt %s -convert-nvgpu-to-nvvm -split-input-file | FileCheck %s
+
+// Test that FP8 (e4m3) nvgpu.mma.sync lowers to nvvm.mma.sync with the
+// correct multiplicand PTX type.
+func.func @fp8_mma_sync(%arg0: vector<4x4xf8E4M3FN>, %arg1: vector<2x4xf8E4M3FN>, %arg2: vector<2x2xf32>) -> vector<2x2xf32> {
+ // CHECK: nvvm.mma.sync
+ // CHECK-SAME: multiplicandAPtxType = #nvvm.mma_type<e4m3>
+ // CHECK-SAME: multiplicandBPtxType = #nvvm.mma_type<e4m3>
+ // CHECK-SAME: shape = #nvvm.shape<m = 16, n = 8, k = 32>
+ %0 = nvgpu.mma.sync(%arg0, %arg1, %arg2) {mmaShape = [16, 8, 32]} : (vector<4x4xf8E4M3FN>, vector<2x4xf8E4M3FN>, vector<2x2xf32>) -> vector<2x2xf32>
+ return %0 : vector<2x2xf32>
+}
More information about the Mlir-commits
mailing list