Guray Ozen 5c3150e584 [MLIR][NVVM] Introduce WGMMA Types
This work introduces `WGMMATypes` attributes for the `WgmmaMmaSyncOp`. This op, having been recently added to MLIR, previously used `MMATypes`. However, there arises a disparity in supported types between `MmaOp` and `WgmmaMmaSyncOp`. To address this discrepancy more effectively, a new set of attributes is introduced.

Furthermore, this patch refines and optimizing the verification mechanisms of `WgmmaMmaSyncOp` Op.

It also adds support for f8 types, including `e4m3` and `e5m2`, within the `WgmmaMmaSyncOp`.

Reviewed By: nicolasvasilache

Differential Revision: https://reviews.llvm.org/D157695
2023-08-12 12:47:45 +02:00

135 lines
5.3 KiB
MLIR

// RUN: mlir-opt --convert-nvvm-to-llvm --split-input-file -verify-diagnostics %s
!mat64f32 = !llvm.struct<(f32, f32, f32, f32, f32, f32, f32)>
func.func @wgmma_f32_f16_f16(%descA : i64, %descB : i64) -> !mat64f32{
%result = llvm.mlir.undef : !mat64f32
// expected-error @+1 {{'nvvm.wgmma.mma_async' op results 64, however output struct has 7 elements}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 128, k = 16>,
D [%result, <zero>],
A [<f16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !mat64f32 -> !mat64f32
return %res : !mat64f32
}
// -----
func.func @wgmma_f32_satfinite(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
// expected-error @+1 {{`satfinite` can be only used with s32 accumulator, however the current accumulator is 'f32'}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 16, k = 16>,
D [%result, <zero>, <satfinite>],
A [<f16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
-> !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
return
}
// -----
func.func @wgmma_f32_m32(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
// expected-error @+1 {{shape 'm' must be 64}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 32, n = 16, k = 16>,
D [%result, <zero>],
A [<f16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
-> !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
return
}
// -----
func.func @wgmma_f32_m32(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f32, f32, i32, f32, f32, f32, f32, f32)>
// expected-error @+1 {{op all elements in struct must be same type but there is 'i32'}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 16, k = 16>,
D [%result, <zero>],
A [<f16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(f32, f32, i32, f32, f32, f32, f32, f32)>
-> !llvm.struct<(f32, f32, i32, f32, f32, f32, f32, f32)>
return
}
// -----
func.func @wgmma_f32_m32(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
// expected-error @+1 {{op shape 'k' must be 16 for input type f16}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 16, k = 3>,
D [%result, <zero>],
A [<f16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
-> !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
return
}
// -----
func.func @wgmma_transpose(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
// expected-error @+1 {{op given layouts layout_a = col and layout_b = col for input types tf32 and tf32 requires transpose. However, this is only supported for: f16 and bf16}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 16, k = 8>,
D [%result, <zero>],
A [<tf32>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<tf32>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
-> !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
return
}
// -----
func.func @wgmma_transpose(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f16, f16, f16, f16)>
// expected-error @+1 {{'nvvm.wgmma.mma_async' op 'f16' += tf32 * tf32, it is not supported.}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 16, k = 8>,
D [%result, <zero>],
A [<tf32>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<tf32>, #nvvm.wgmma_scale_in<neg>, <col>]
:!llvm.struct<(f16, f16, f16, f16)>
-> !llvm.struct<(f16, f16, f16, f16)>
return
}
// -----
func.func @wgmma_f32_m32(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(i32, i32, i32, i32)>
// expected-error @+1 {{input struct and result struct must be the same type}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 8, k = 16>,
D [%result, <zero>],
A [<f16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(i32, i32, i32, i32)>
-> !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
return
}
// -----
func.func @wgmma_f32_m32(%descA : i64, %descB : i64) {
%result = llvm.mlir.undef : !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
// expected-error @+1 {{op 'f32' += bf16 * f16, it is not supported}}
%res = nvvm.wgmma.mma_async %descA, %descB,
#nvvm.shape<m = 64, n = 8, k = 16>,
D [%result, <zero>],
A [<bf16>, #nvvm.wgmma_scale_in<neg>, <col>],
B [<f16>, #nvvm.wgmma_scale_in<neg>, <col>]
: !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
-> !llvm.struct<(f32, f32, f32, f32, f32, f32, f32, f32)>
return
}