[Mlir-commits] [mlir] dc19e4b - [mlir][NVGPUToNVVM] Support BF16 mma.sync lowering (#194203)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Mon Apr 27 01:43:36 PDT 2026
Author: Hao Ren
Date: 2026-04-27T10:43:31+02:00
New Revision: dc19e4b0b6c9a10633958046d2e27b597ba7e37e
URL: https://github.com/llvm/llvm-project/commit/dc19e4b0b6c9a10633958046d2e27b597ba7e37e
DIFF: https://github.com/llvm/llvm-project/commit/dc19e4b0b6c9a10633958046d2e27b597ba7e37e.diff
LOG: [mlir][NVGPUToNVVM] Support BF16 mma.sync lowering (#194203)
Let NVGPUToNVVM to recognize BF16 MMA operand element types
Pack `vector<2xbf16>` fragments to `i32` before emitting
`nvvm.mma.sync`.
This matches the PTX operand encoding for `m16n8k16` BF16 MMA
instructions.
Add a conversion test for `nvgpu.mma.sync` `bf16xbf16` to `f32`
lowering.
Co-authored-by: Hao Ren <rhao8608 at gmail.com>
Added:
Modified:
mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm.mlir
Removed:
################################################################################
diff --git a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
index c4175016ab30c..e566449ffadff 100644
--- a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
+++ b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
@@ -179,6 +179,7 @@ static SmallVector<Value> unpackOperandVector(ImplicitLocOpBuilder &b,
Type f64Ty = b.getF64Type();
Type f32Ty = b.getF32Type();
Type i64Ty = b.getI64Type();
+ Type bf16x2Ty = VectorType::get(2, b.getBF16Type());
Type i8x4Ty = VectorType::get(4, b.getI8Type());
Type i4x8Ty = VectorType::get(8, b.getIntegerType(4));
Type f32x1Ty = VectorType::get(1, f32Ty);
@@ -191,6 +192,8 @@ static SmallVector<Value> unpackOperandVector(ImplicitLocOpBuilder &b,
// scalar types.
if (arrayTy.getElementType() == i8x4Ty ||
arrayTy.getElementType() == i4x8Ty ||
+ (arrayTy.getElementType() == bf16x2Ty &&
+ operandPtxType == NVVM::MMATypes::bf16) ||
(arrayTy.getElementType() == f32x1Ty &&
operandPtxType == NVVM::MMATypes::tf32)) {
result.push_back(LLVM::BitcastOp::create(b, i32Ty, toUse));
@@ -320,6 +323,8 @@ static FailureOr<NVVM::MMATypes> getNvvmMmaType(Type t) {
return NVVM::MMATypes::s4;
if (elType.isF16())
return NVVM::MMATypes::f16;
+ if (elType.isBF16())
+ return NVVM::MMATypes::bf16;
if (elType.isF64())
return NVVM::MMATypes::f64;
if (elType.isF32())
diff --git a/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm.mlir b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm.mlir
index 50bea5a85022e..6b7fd1578c35a 100644
--- a/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm.mlir
+++ b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-to-nvvm.mlir
@@ -49,6 +49,29 @@ func.func @m16n8k16_fp16_fp32(%arg0: vector<4x2xf16>, %arg1: vector<2x2xf16>, %a
return %d : vector<2x2xf32>
}
+// CHECK-LABEL: @m16n8k16_bf16_fp32
+func.func @m16n8k16_bf16_fp32(%arg0: vector<4x2xbf16>, %arg1: vector<2x2xbf16>, %arg2: vector<2x2xf32>) -> vector<2x2xf32> {
+ // CHECK: llvm.extractvalue %{{.*}}[0] : !llvm.array<4 x vector<2xbf16>>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: llvm.extractvalue %{{.*}}[1] : !llvm.array<4 x vector<2xbf16>>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: llvm.extractvalue %{{.*}}[2] : !llvm.array<4 x vector<2xbf16>>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: llvm.extractvalue %{{.*}}[3] : !llvm.array<4 x vector<2xbf16>>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: llvm.extractvalue %{{.*}}[0] : !llvm.array<2 x vector<2xbf16>>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: llvm.extractvalue %{{.*}}[1] : !llvm.array<2 x vector<2xbf16>>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: [[d:%.+]] = nvvm.mma.sync A[{{%.+}}, {{%.+}}, {{%.+}}, {{%.+}}] B[{{%.+}}, {{%.+}}] C[{{%.+}}, {{%.+}}, {{%.+}}, {{%.+}}]
+ // CHECK-SAME: multiplicandAPtxType = #nvvm.mma_type<bf16>
+ // CHECK-SAME: multiplicandBPtxType = #nvvm.mma_type<bf16>
+ // CHECK-SAME: shape = #nvvm.shape<m = 16, n = 8, k = 16>
+ // CHECK-SAME: (i32, i32, f32) -> !llvm.struct<(f32, f32, f32, f32)>
+ %d = nvgpu.mma.sync (%arg0, %arg1, %arg2) {mmaShape = [16, 8, 16]} : (vector<4x2xbf16>, vector<2x2xbf16>, vector<2x2xf32>) -> vector<2x2xf32>
+ return %d : vector<2x2xf32>
+}
+
// CHECK-LABEL: @m16n8k8_fp16
func.func @m16n8k8_fp16(%arg0: vector<2x2xf16>, %arg1: vector<1x2xf16>, %arg2: vector<2x2xf16>) -> vector<2x2xf16> {
// CHECK: llvm.extractvalue %{{.*}}[0] : !llvm.array<2 x vector<2xf16>>
More information about the Mlir-commits
mailing list