[Mlir-commits] [mlir] [NVGPU][NVVM] Add FP8 (e4m3/e5m2) support to dense mma.sync (PR #207307)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Thu Jul 2 19:32:23 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-mlir-llvm

Author: WMC (weimin023)

<details>
<summary>Changes</summary>

## Description:

`nvgpu.mma.sync` / `nvvm.mma.sync` (warp-level, dense) don't support FP8 (e4m3/e5m2) operands, even though the underlying PTX/NVPTX intrinsics have supported FP8 `mma.sync.m16n8k32` since sm_89. FP8 currently only exists on the sparse mma path and on warpgroup-level `wgmma`; the warp-level dense path was never updated.

This adds FP8 (e4m3/e5m2) support to `nvgpu.mma.sync` and its lowering to `nvvm.mma.sync`, covering `m16n8k16` and `m16n8k32`.

## Verification
`$ mlir-opt mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir -convert-nvgpu-to-nvvm -split-input-file`

```
%0 = nvvm.mma.sync A[...] B[...] C[...]
     {multiplicandAPtxType = #nvvm.mma_type<e4m3>,
      multiplicandBPtxType = #nvvm.mma_type<e4m3>,
      shape = #nvvm.shape<m = 16, n = 8, k = 32>}
     : (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>

```

Existing FP16/BF16 tests in `nvgpu-to-nvvm.mlir` still pass (no regression).



---
Full diff: https://github.com/llvm/llvm-project/pull/207307.diff


5 Files Affected:

- (modified) mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td (+5-1) 
- (modified) mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp (+4) 
- (modified) mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp (+9) 
- (modified) mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp (+4-3) 
- (added) mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm-fp8.mlir (+12) 


``````````diff
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/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/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/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp b/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
index 2b2686a9163bf..2a93a3ba3d235 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>
+}

``````````

</details>


https://github.com/llvm/llvm-project/pull/207307


More information about the Mlir-commits mailing list