[Mlir-commits] [mlir] 1f2772f - [mlir][tosa] Add support for MXFP conv2d (#210054)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jul 17 07:19:27 PDT 2026
Author: Iliyan Georgiev
Date: 2026-07-17T15:19:23+01:00
New Revision: 1f2772f8e26beb909c2a559f1fb08697e3e06909
URL: https://github.com/llvm/llvm-project/commit/1f2772f8e26beb909c2a559f1fb08697e3e06909
DIFF: https://github.com/llvm/llvm-project/commit/1f2772f8e26beb909c2a559f1fb08697e3e06909.diff
LOG: [mlir][tosa] Add support for MXFP conv2d (#210054)
- Adds profile compliance support for MXFP conv2d
- Relax conv2d accumulator constraits. They are checked as part of
validation pass
and this would allow experimentation with different types.
Signed-off-by: Iliyan Georgiev <iliyan.georgiev at arm.com>
Added:
Modified:
mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
mlir/test/Dialect/Tosa/availability.mlir
mlir/test/Dialect/Tosa/invalid.mlir
mlir/test/Dialect/Tosa/invalid_extension.mlir
mlir/test/Dialect/Tosa/ops.mlir
mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
mlir/test/Dialect/Tosa/verifier.mlir
Removed:
################################################################################
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
index f191da6583584..54cf204624937 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
@@ -615,16 +615,991 @@ extensionComplianceMap = {
{{Extension::fp8e4m3},
{{{fp8e4m3T, fp8e4m3T, fp16T, fp8e4m3T, fp8e4m3T, fp16T, fp16T},
SpecificationVersion::V_1_0},
+ {{fp8e4m3T, fp8e4m3T, fp16T, fp8e4m3T, fp8e4m3T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, fp16T, fp16T, fp8e4m3T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, fp8e4m3T, fp16T, fp16T, fp8e4m3T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
{{fp8e4m3T, fp8e4m3T, fp32T, fp8e4m3T, fp8e4m3T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, fp8e4m3T, fp32T, fp16T, fp8e4m3T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T, fp32T},
SpecificationVersion::V_1_1_DRAFT}}},
{{Extension::fp8e5m2},
{{{fp8e5m2T, fp8e5m2T, fp16T, fp8e5m2T, fp8e5m2T, fp16T, fp16T},
SpecificationVersion::V_1_0},
+ {{fp8e5m2T, fp8e5m2T, fp16T, fp8e5m2T, fp8e5m2T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp16T, fp16T, fp8e5m2T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, fp8e5m2T, fp16T, fp16T, fp8e5m2T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
{{fp8e5m2T, fp8e5m2T, fp32T, fp8e5m2T, fp8e5m2T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, fp8e5m2T, fp32T, fp16T, fp8e5m2T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T, fp32T},
SpecificationVersion::V_1_1_DRAFT}}},
{{Extension::bf16},
{{{bf16T, bf16T, bf16T, bf16T, bf16T, fp32T, bf16T},
- SpecificationVersion::V_1_0}}}}},
+ SpecificationVersion::V_1_0},
+ {{fp16T, fp16T, fp16T, fp16T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bf16T, fp16T, fp16T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bf16T, fp16T, fp16T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp16T, fp16T, bf16T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp16T, fp16T, bf16T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bf16T, fp16T, bf16T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bf16T, fp16T, bf16T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bf16T, fp32T, fp16T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp16T, fp32T, bf16T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bf16T, fp32T, bf16T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp32T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}}},
+ {{Extension::bf16, Extension::fp8e4m3},
+ {{{fp8e4m3T, fp8e4m3T, fp16T, fp8e4m3T, fp8e4m3T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, fp16T, fp16T, fp8e4m3T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bf16T, fp16T, fp8e4m3T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bf16T, fp16T, fp8e4m3T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, fp8e4m3T, fp16T, fp16T, fp8e4m3T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp8e4m3T, fp16T, bf16T, fp8e4m3T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp8e4m3T, fp16T, bf16T, fp8e4m3T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bf16T, fp32T, fp8e4m3T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp8e4m3T, fp32T, bf16T, fp8e4m3T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::fp8e5m2},
+ {{{fp8e4m3T, fp8e5m2T, fp16T, fp8e4m3T, fp8e5m2T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp8e4m3T, fp16T, fp8e5m2T, fp8e4m3T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::fp8e5m2},
+ {{{fp8e4m3T, fp8e5m2T, fp16T, fp8e4m3T, fp8e5m2T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp8e4m3T, fp16T, fp8e5m2T, fp8e4m3T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, fp8e5m2T, fp32T, fp8e4m3T, fp8e5m2T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp8e4m3T, fp32T, fp8e5m2T, fp8e4m3T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::mx_common,
+ Extension::mx_fp8e4m3},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp8e4m3T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp8e4m3T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp8e4m3T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::mx_common,
+ Extension::mx_fp8e5m2},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp8e4m3T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp8e4m3T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp8e4m3T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::mx_common,
+ Extension::mx_fp6e3m2},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp8e4m3T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp8e4m3T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp8e4m3T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::mx_common,
+ Extension::mx_fp6e2m3},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp8e4m3T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp8e4m3T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp8e4m3T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::mx_common,
+ Extension::mx_fp4e2m1},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp8e4m3T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp8e4m3T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp8e4m3T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e4m3, Extension::mx_common,
+ Extension::mx_int8},
+ {{{fp8e4m3T, bs32_fp8ue8m0_mxint8T, fp16T, fp8e4m3T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_int8},
+ {{{fp8e4m3T, bs32_fp8ue8m0_mxint8T, fp16T, fp8e4m3T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e4m3T, fp16T, fp32T, fp8e4m3T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e4m3T, bs32_fp8ue8m0_mxint8T, fp32T, fp8e4m3T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e4m3T, fp32T, fp32T, fp8e4m3T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2},
+ {{{fp8e5m2T, fp8e5m2T, fp16T, fp8e5m2T, fp8e5m2T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, fp16T, fp16T, fp8e5m2T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bf16T, fp16T, fp8e5m2T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bf16T, fp16T, fp8e5m2T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, fp8e5m2T, fp16T, fp16T, fp8e5m2T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp8e5m2T, fp16T, bf16T, fp8e5m2T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp8e5m2T, fp16T, bf16T, fp8e5m2T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bf16T, fp32T, fp8e5m2T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, fp8e5m2T, fp32T, bf16T, fp8e5m2T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2, Extension::mx_common,
+ Extension::mx_fp8e4m3},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp8e5m2T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp8e5m2T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp8e5m2T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2, Extension::mx_common,
+ Extension::mx_fp8e5m2},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp8e5m2T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp8e5m2T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp8e5m2T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2, Extension::mx_common,
+ Extension::mx_fp6e3m2},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp8e5m2T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp8e5m2T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp8e5m2T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2, Extension::mx_common,
+ Extension::mx_fp6e2m3},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp8e5m2T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp8e5m2T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp8e5m2T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2, Extension::mx_common,
+ Extension::mx_fp4e2m1},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp8e5m2T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp8e5m2T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp8e5m2T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::fp8e5m2, Extension::mx_common,
+ Extension::mx_int8},
+ {{{fp8e5m2T, bs32_fp8ue8m0_mxint8T, fp16T, fp8e5m2T, fp32T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, bf16T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_int8},
+ {{{fp8e5m2T, bs32_fp8ue8m0_mxint8T, fp16T, fp8e5m2T, fp32T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e5m2T, fp16T, fp32T, fp8e5m2T, fp32T,
+ fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp8e5m2T, bs32_fp8ue8m0_mxint8T, fp32T, fp8e5m2T, fp32T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e5m2T, fp32T, fp32T, fp8e5m2T, fp32T,
+ fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp16T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp8e4m3T, fp16T, bf16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp8e4m3T, fp16T, bf16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp16T, fp16T, fp32T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bf16T, fp16T, fp32T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bf16T, fp16T, fp32T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp8e4m3T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp16T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp16T, fp16T, fp32T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp16T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp8e5m2T, fp16T, bf16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp8e5m2T, fp16T, bf16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp16T, fp16T, fp32T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bf16T, fp16T, fp32T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bf16T, fp16T, fp32T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp8e5m2T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp16T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp16T, fp16T, fp32T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp16T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp6e3m2T, fp16T, bf16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp6e3m2T, fp16T, bf16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp16T, fp16T, fp32T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bf16T, fp16T, fp32T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bf16T, fp16T, fp32T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp6e3m2T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp16T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp16T, fp16T, fp32T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp16T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp6e2m3T, fp16T, bf16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp6e2m3T, fp16T, bf16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp16T, fp16T, fp32T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bf16T, fp16T, fp32T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bf16T, fp16T, fp32T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp6e2m3T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp16T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp16T, fp16T, fp32T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp16T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp4e2m1T, fp16T, bf16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp4e2m1T, fp16T, bf16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp16T, fp16T, fp32T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bf16T, fp16T, fp32T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bf16T, fp16T, fp32T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_fp4e2m1T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp16T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp16T, fp16T, fp32T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_int8},
+ {{{fp16T, bs32_fp8ue8m0_mxint8T, fp16T, fp16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_mxint8T, fp16T, bf16T, fp32T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_mxint8T, fp16T, bf16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp16T, fp16T, fp32T, fp16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bf16T, fp16T, fp32T, bf16T, bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bf16T, fp16T, fp32T, bf16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bf16T, bs32_fp8ue8m0_mxint8T, fp32T, bf16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bf16T, fp32T, fp32T, bf16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_int8},
+ {{{fp16T, bs32_fp8ue8m0_mxint8T, fp16T, fp16T, fp32T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp16T, fp16T, fp32T, fp16T, fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp16T, bs32_fp8ue8m0_mxint8T, fp32T, fp16T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T, fp16T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T, fp32T, fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e4m3,
+ Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e4m3, Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e3m2,
+ Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e3m2, Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e2m3,
+ Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3, Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1,
+ Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1, Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e4m3,
+ Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e4m3, Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp8e4m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp8e4m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e3m2,
+ Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e3m2, Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e2m3,
+ Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3, Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1,
+ Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1, Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e5m2,
+ Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e5m2, Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp8e5m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp8e5m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e2m3,
+ Extension::mx_fp6e3m2},
+ {{{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3, Extension::mx_fp6e3m2},
+ {{{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1,
+ Extension::mx_fp6e3m2},
+ {{{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1, Extension::mx_fp6e3m2},
+ {{{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e3m2,
+ Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e3m2, Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp6e3m2T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp6e3m2T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1,
+ Extension::mx_fp6e2m3},
+ {{{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1, Extension::mx_fp6e2m3},
+ {{{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e2m3,
+ Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3, Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp6e2m3T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp6e2m3T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1,
+ Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ bf16T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1, Extension::mx_int8},
+ {{{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_mxint8T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp4e2m1T, fp16T, fp32T, fp32T,
+ fp32T, fp16T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bs32_fp8ue8m0_mxint8T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bs32_fp8ue8m0_fp4e2m1T, fp32T, fp32T, fp32T,
+ fp32T, fp32T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf}}},
{"tosa.conv2d_block_scaled",
{{{Extension::mxfp_conv},
{{{fp4e2m1T, fp8ue8m0T, fp4e2m1T, fp8ue8m0T, fp32T, fp32T},
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
index 66521cbf73db8..7136d629a751c 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
@@ -68,7 +68,7 @@ def Tosa_ArgMaxOp : Tosa_InferShapedTypeOp<"argmax", [NoMemoryEffect]> {
// Accumulator types.
//===----------------------------------------------------------------------===//
-def Tosa_AccType : AnyTypeOf<[I<32>, I<48>, F16, F32]>;
+def Tosa_AccType : AnyTypeOf<[I<32>, I<48>, F16, F32, BF16]>;
//===----------------------------------------------------------------------===//
// Operator: avg_pool2d
@@ -203,7 +203,11 @@ def Tosa_Conv2DOp : Tosa_ConvOp<"conv2d"> {
list<Availability> availability = [
Profile<[Tosa_PRO_INT, Tosa_PRO_FP]>,
- Extension<[Tosa_EXT_INT4, Tosa_EXT_INT16, Tosa_EXT_FP8E4M3, Tosa_EXT_FP8E5M2, Tosa_EXT_BF16]>,
+ Extension<[Tosa_EXT_INT4, Tosa_EXT_INT16, Tosa_EXT_FP8E4M3,
+ Tosa_EXT_FP8E5M2, Tosa_EXT_BF16, Tosa_EXT_MX_COMMON,
+ Tosa_EXT_MX_FP4E2M1, Tosa_EXT_MX_FP6E2M3,
+ Tosa_EXT_MX_FP6E3M2, Tosa_EXT_MX_FP8E4M3,
+ Tosa_EXT_MX_FP8E5M2, Tosa_EXT_MX_INT8]>,
];
let extraClassDeclaration = [{
@@ -221,6 +225,8 @@ def Tosa_Conv2DOp : Tosa_ConvOp<"conv2d"> {
//===----------------------------------------------------------------------===//
// Operator: conv2d_block_scaled
+//
+// Note: This operation is deprecated. It will be removed in the future.
//===----------------------------------------------------------------------===//
def Tosa_Conv2DBlockScaledOp : Tosa_InferShapedTypeOp<"conv2d_block_scaled", [NoMemoryEffect]> {
let summary = "Performs two dimensional convolution using block scaled tensors.";
@@ -232,6 +238,8 @@ def Tosa_Conv2DBlockScaledOp : Tosa_InferShapedTypeOp<"conv2d_block_scaled", [No
This operation is not pure. Undefined behaviour may occur if the accumulated
result overflows.
+
+ // Note: This operation is deprecated. It will be removed in the future.
}];
let arguments = (ins
@@ -463,7 +471,7 @@ def Tosa_MatMulTOp : Tosa_InferShapedTypeOp<"matmul_t", [NoMemoryEffect]> {
let description = [{
Performs two dimensional matrix multiplications. `A` matrix is of shape
`N x H x C`. `B` matrix is of shape `D x W x C`. This is effectively a
- matrix multiply of `A` by the transposed `B` matrix. If the batched
+ matrix multiply of `A` by the transposed `B` matrix. If the batched
dimension of input `B` is of size 1, the `B` matrix is broadcast.
}];
@@ -2597,7 +2605,7 @@ def Tosa_RowGatherOp : Tosa_InferShapedTypeOp<"row_gather", [NoMemoryEffect]> {
let description = [{
Generate a tensor based on the input indices and row_count. The number of
- consecutive rows gathered for each index is specified in row_count. N is
+ consecutive rows gathered for each index is specified in row_count. N is
the number of batches, W is the number of indices in each batch, K is
the range of each index, and C is the number of data channels for each
index. The values tensor has shape [N, K, C] and the output tensor has shape
diff --git a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
index 0b6abc2df3a70..0dfcd23504019 100644
--- a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
+++ b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
@@ -888,21 +888,16 @@ static LogicalResult verifyConvOp(T op) {
return failure();
}
- if (isa<Float8E5M2Type>(inputEType) || isa<Float8E4M3FNType>(inputEType) ||
- isa<Float8E5M2Type>(weightEType) || isa<Float8E4M3FNType>(weightEType)) {
- if (inputEType != weightEType) {
- op.emitOpError(
- "expect both input and weight to have same element type, got ")
- << inputEType << " and " << weightEType;
- return failure();
- }
- }
+ const bool isInputBlockScaled = llvm::isa<BlockScaledType>(inputEType);
+ const bool isWeightBlockScaled = llvm::isa<BlockScaledType>(weightEType);
+ const bool isInputFloat = llvm::isa<FloatType>(inputEType);
+ const bool isWeightFloat = llvm::isa<FloatType>(weightEType);
- bool inputIsFloat = llvm::isa<FloatType>(inputEType);
- bool weightIsFloat = llvm::isa<FloatType>(weightEType);
+ const bool isInputBSorFloat = isInputBlockScaled || isInputFloat;
+ const bool isWeightBSorFloat = isWeightBlockScaled || isWeightFloat;
// Either both must be float or both non-float.
- if (inputIsFloat != weightIsFloat) {
+ if (isInputBSorFloat != isWeightBSorFloat) {
op.emitOpError(
"expect both input and weight to be float or not together, got ")
<< inputEType << " and " << weightEType;
@@ -910,18 +905,28 @@ static LogicalResult verifyConvOp(T op) {
}
auto inputZpEType = getStorageElementTypeOrSelf(op.getInputZp().getType());
- if (inputEType != inputZpEType) {
+ if (!isInputBlockScaled && inputEType != inputZpEType) {
return op.emitOpError("expect both input and its zero point are the same "
"element type, got ")
<< inputEType << " and " << inputZpEType;
}
+ if (isInputBlockScaled && !llvm::isa<Float32Type>(inputZpEType)) {
+ return op.emitOpError(
+ "expect block scaled input to have fp32 zero point, got ")
+ << inputEType << " and " << inputZpEType;
+ }
auto weightZpEType = getStorageElementTypeOrSelf(op.getWeightZp().getType());
- if (weightEType != weightZpEType) {
+ if (!isWeightBlockScaled && weightEType != weightZpEType) {
return op.emitOpError("expect both weight and its zero point are the same "
"element type, got ")
<< weightEType << " and " << weightZpEType;
}
+ if (isWeightBlockScaled && !llvm::isa<Float32Type>(weightZpEType)) {
+ return op.emitOpError(
+ "expect block scaled weight to have fp32 zero point, got ")
+ << weightEType << " and " << weightZpEType;
+ }
FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
if (succeeded(maybeIZp) && op.verifyInputZeroPoint(*maybeIZp).failed())
@@ -998,33 +1003,6 @@ static LogicalResult verifyConvOpModes(T op) {
if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(inputEType))
inputEType = getStorageElementTypeFromQuantized(quantType);
- auto accType = op.getAccType();
- if (inputEType.isInteger(8) && !accType.isInteger(32))
- return op.emitOpError("accumulator type for i8 tensor is not i32, got ")
- << accType;
-
- if (inputEType.isInteger(16) && !accType.isInteger(48))
- return op.emitOpError("accumulator type for i16 tensor is not i48, got ")
- << accType;
-
- if (isa<Float8E5M2Type, Float8E4M3Type>(inputEType) &&
- !(accType.isF16() || accType.isF32()))
- return op.emitOpError("accumulator type for f8 tensor is not f16/f32, got ")
- << accType;
-
- if (inputEType.isF16() && !(accType.isF16() || accType.isF32()))
- return op.emitOpError(
- "accumulator type for f16 tensor is not f16/f32, got ")
- << accType;
-
- if (inputEType.isBF16() && !accType.isF32())
- return op.emitOpError("accumulator type for bf16 tensor is not f32, got ")
- << accType;
-
- if (inputEType.isF32() && !accType.isF32())
- return op.emitOpError("accumulator type for f32 tensor is not f32, got ")
- << accType;
-
auto resultEType =
llvm::cast<ShapedType>(op.getResult().getType()).getElementType();
diff --git a/mlir/test/Dialect/Tosa/availability.mlir b/mlir/test/Dialect/Tosa/availability.mlir
index 04277c4ea9bfb..eb571f9a163d7 100644
--- a/mlir/test/Dialect/Tosa/availability.mlir
+++ b/mlir/test/Dialect/Tosa/availability.mlir
@@ -43,7 +43,7 @@ func.func @test_avg_pool2d_adaptive(%arg0: tensor<1x7x7x9xf32>) -> tensor<1x7x7x
// CHECK-LABEL: conv2d
func.func @test_conv2d(%arg0: tensor<1x4x4x4xf32>, %arg1: tensor<8x1x1x4xf32>, %arg2: tensor<8xf32>) -> tensor<1x4x4x8xf32> {
// CHECK: profiles: [ [pro_int, pro_fp] ]
- // CHECK: extensions: [ [int4, int16, fp8e4m3, fp8e5m2, bf16] ]
+ // CHECK: extensions: [ [int4, int16, fp8e4m3, fp8e5m2, bf16, mx_common, mx_fp4e2m1, mx_fp6e2m3, mx_fp6e3m2, mx_fp8e4m3, mx_fp8e5m2, mx_int8] ]
%input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
%weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
%0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = f32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x4xf32>, tensor<8x1x1x4xf32>, tensor<8xf32>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x4x4x8xf32>
diff --git a/mlir/test/Dialect/Tosa/invalid.mlir b/mlir/test/Dialect/Tosa/invalid.mlir
index d0336da15cee8..e559ea5f3c0ab 100644
--- a/mlir/test/Dialect/Tosa/invalid.mlir
+++ b/mlir/test/Dialect/Tosa/invalid.mlir
@@ -74,105 +74,6 @@ func.func @test_conv2d_weight_zp(%arg0: tensor<1x29x29x4xf16>, %arg1: tensor<16x
// -----
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xi8>, %arg1: tensor<16x3x3x4xi8>, %arg2: tensor<16xi8>) -> tensor<1x27x27x16xi8> {
- %zp = "tosa.const"() {values = dense<0> : tensor<1xi8>} : () -> tensor<1xi8>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for i8 tensor is not i32, got 'f16'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = f16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xi8>, tensor<16x3x3x4xi8>, tensor<16xi8>, tensor<1xi8>, tensor<1xi8>) -> tensor<1x27x27x16xi8>
- return %0 : tensor<1x27x27x16xi8>
-}
-
-// -----
-
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xi16>, %arg1: tensor<16x3x3x4xi8>, %arg2: tensor<16xi16>) -> tensor<1x27x27x16xi16> {
- %input_zp = "tosa.const"() {values = dense<0> : tensor<1xi16>} : () -> tensor<1xi16>
- %weight_zp = "tosa.const"() {values = dense<0> : tensor<1xi8>} : () -> tensor<1xi8>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for i16 tensor is not i48, got 'f16'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = f16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xi16>, tensor<16x3x3x4xi8>, tensor<16xi16>, tensor<1xi16>, tensor<1xi8>) -> tensor<1x27x27x16xi16>
- return %0 : tensor<1x27x27x16xi16>
-}
-
-// -----
-
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xf8E5M2>, %arg1: tensor<16x3x3x4xf8E5M2>, %arg2: tensor<16xf16>) -> tensor<1x27x27x16xf16> {
- %zp = "tosa.const"() {values = dense<0.0> : tensor<1xf8E5M2>} : () -> tensor<1xf8E5M2>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for f8 tensor is not f16/f32, got 'i32'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = i32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xf8E5M2>, tensor<16x3x3x4xf8E5M2>, tensor<16xf16>, tensor<1xf8E5M2>, tensor<1xf8E5M2>) -> tensor<1x27x27x16xf16>
- return %0 : tensor<1x27x27x16xf16>
-}
-
-// -----
-
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xf8E4M3>, %arg1: tensor<16x3x3x4xf8E4M3>, %arg2: tensor<16xf16>) -> tensor<1x27x27x16xf16> {
- %zp = "tosa.const"() {values = dense<0.0> : tensor<1xf8E4M3>} : () -> tensor<1xf8E4M3>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for f8 tensor is not f16/f32, got 'i32'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = i32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xf8E4M3>, tensor<16x3x3x4xf8E4M3>, tensor<16xf16>, tensor<1xf8E4M3>, tensor<1xf8E4M3>) -> tensor<1x27x27x16xf16>
- return %0 : tensor<1x27x27x16xf16>
-}
-
-// -----
-
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xf16>, %arg1: tensor<16x3x3x4xf16>, %arg2: tensor<16xf16>) -> tensor<1x27x27x16xf16> {
- %zp = "tosa.const"() {values = dense<0.0> : tensor<1xf16>} : () -> tensor<1xf16>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for f16 tensor is not f16/f32, got 'i32'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = i32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xf16>, tensor<16x3x3x4xf16>, tensor<16xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<1x27x27x16xf16>
- return %0 : tensor<1x27x27x16xf16>
-}
-
-// -----
-
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xbf16>, %arg1: tensor<16x3x3x4xbf16>, %arg2: tensor<16xbf16>) -> tensor<1x27x27x16xbf16> {
- %zp = "tosa.const"() {values = dense<0.0> : tensor<1xbf16>} : () -> tensor<1xbf16>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for bf16 tensor is not f32, got 'i32'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = i32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xbf16>, tensor<16x3x3x4xbf16>, tensor<16xbf16>, tensor<1xbf16>, tensor<1xbf16>) -> tensor<1x27x27x16xbf16>
- return %0 : tensor<1x27x27x16xbf16>
-}
-
-// -----
-
-func.func @test_conv2d_acc_type(%arg0: tensor<1x29x29x4xf32>, %arg1: tensor<16x3x3x4xf32>, %arg2: tensor<16xf32>) -> tensor<1x27x27x16xf32> {
- %zp = "tosa.const"() {values = dense<0.0> : tensor<1xf32>} : () -> tensor<1xf32>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for f32 tensor is not f32, got 'i32'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = i32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>}
- : (tensor<1x29x29x4xf32>, tensor<16x3x3x4xf32>, tensor<16xf32>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x27x27x16xf32>
- return %0 : tensor<1x27x27x16xf32>
-}
-
-// -----
-
-func.func @test_conv3d_acc_type(%arg0: tensor<1x4x8x21x17xi8>, %arg1: tensor<34x1x1x1x17xi8>, %arg2: tensor<34xi8>) -> tensor<1x4x8x21x34xi8> {
- %zp = "tosa.const"() {values = dense<0> : tensor<1xi8>} : () -> tensor<1xi8>
- // expected-error at +1 {{'tosa.conv3d' op accumulator type for i8 tensor is not i32, got 'f16'}}
- %0 = tosa.conv3d %arg0, %arg1, %arg2, %zp, %zp {acc_type = f16, dilation = array<i64: 1, 1, 1>, pad = array<i64: 0, 0, 0, 0, 0, 0>, stride = array<i64: 1, 1, 1>}
- : (tensor<1x4x8x21x17xi8>, tensor<34x1x1x1x17xi8>, tensor<34xi8>, tensor<1xi8>, tensor<1xi8>) -> tensor<1x4x8x21x34xi8>
- return %0 : tensor<1x4x8x21x34xi8>
-}
-
-// -----
-
-func.func @test_depthwise_conv2d_acc_type(%arg0: tensor<1x4x4x4xi8>, %arg1: tensor<1x1x4x2xi8>, %arg2: tensor<8xi8>) -> tensor<1x4x4x8xi8> {
- %zp = "tosa.const"() {values = dense<0> : tensor<1xi8>} : () -> tensor<1xi8>
- // expected-error at +1 {{'tosa.depthwise_conv2d' op accumulator type for i8 tensor is not i32, got 'f16'}}
- %0 = tosa.depthwise_conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = f16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>} : (tensor<1x4x4x4xi8>, tensor<1x1x4x2xi8>, tensor<8xi8>, tensor<1xi8>, tensor<1xi8>) -> tensor<1x4x4x8xi8>
- return %0 : tensor<1x4x4x8xi8>
-}
-
-// -----
-
-func.func @test_transpose_conv2d(%arg0: tensor<1x32x32x8xi8>, %arg1: tensor<16x1x1x8xi8>, %arg2: tensor<16xi8>) -> tensor<1x32x32x16xi8> {
- %zp = "tosa.const"() {values = dense<0> : tensor<1xi8>} : () -> tensor<1xi8>
- // expected-error at +1 {{'tosa.transpose_conv2d' op accumulator type for i8 tensor is not i32, got 'f16'}}
- %0 = tosa.transpose_conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = f16, out_pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>} : (tensor<1x32x32x8xi8>, tensor<16x1x1x8xi8>, tensor<16xi8>, tensor<1xi8>, tensor<1xi8>) -> tensor<1x32x32x16xi8>
- return %0 : tensor<1x32x32x16xi8>
-}
-
-// -----
-
func.func @test_transpose_conv2d_invalid_padding_top(%arg0: tensor<1x32x32x8xf32>, %arg1: tensor<16x1x1x8xf32>, %arg2: tensor<16xf32>, %arg3: tensor<1xf32>, %arg4: tensor<1xf32>) -> tensor<1x32x32x16xf32> {
// expected-error at +1 {{'tosa.transpose_conv2d' op expected out_pad_top > -KH, but got: out_pad_top=-3 and KH=1}}
%0 = tosa.transpose_conv2d %arg0, %arg1, %arg2, %arg3, %arg4 {acc_type = f32, out_pad = array<i64: -3, 0, 0, 0>, out_shape = array<i64: 1, 32, 32, 16>, stride = array<i64: 1, 1>} : (tensor<1x32x32x8xf32>, tensor<16x1x1x8xf32>, tensor<16xf32>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x32x32x16xf32>
@@ -243,15 +144,6 @@ func.func @test_transpose_conv2d_invalid_bias(%arg0: tensor<1x32x32x8xf32>, %arg
return %0 : tensor<1x32x32x16xf32>
}
-// -----
-// CHECK-LABEL: conv2d_quant_any_acc
-func.func @test_conv2d_quant_any_acc(%arg0: tensor<1x4x4x4x!quant.any<i8<-8:7>>>, %arg1: tensor<8x1x1x4x!quant.any<i8<-8:7>>>, %arg2: tensor<8x!quant.any<i8<-8:7>>>) -> tensor<1x4x4x8x!quant.any<i8<-8:7>>> {
- %zp = "tosa.const" () { values = dense<0> : tensor<1xi8> } : () -> tensor<1xi8>
- // expected-error at +1 {{'tosa.conv2d' op accumulator type for i8 tensor is not i32, got 'f32'}}
- %0 = tosa.conv2d %arg0, %arg1, %arg2, %zp, %zp {acc_type = f32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x4x!quant.any<i8<-8:7>>>, tensor<8x1x1x4x!quant.any<i8<-8:7>>>, tensor<8x!quant.any<i8<-8:7>>>, tensor<1xi8>, tensor<1xi8>) -> tensor<1x4x4x8x!quant.any<i8<-8:7>>>
- return %0 : tensor<1x4x4x8x!quant.any<i8<-8:7>>>
-}
-
// -----
// CHECK-LABEL: conv2d_quant_any
func.func @test_conv2d_quant_any(%arg0: tensor<1x4x4x4x!quant.any<i8<-8:7>>>, %arg1: tensor<8x1x1x4x!quant.any<i8<-8:7>>>, %arg2: tensor<8x!quant.any<i32<-8:7>>>) -> tensor<1x4x4x8x!quant.any<i32<-8:7>>> {
diff --git a/mlir/test/Dialect/Tosa/invalid_extension.mlir b/mlir/test/Dialect/Tosa/invalid_extension.mlir
index 94574dea21d01..fe17b67465fb3 100644
--- a/mlir/test/Dialect/Tosa/invalid_extension.mlir
+++ b/mlir/test/Dialect/Tosa/invalid_extension.mlir
@@ -25,6 +25,15 @@ func.func @test_conv2d(%arg0: tensor<1x4x4x4xi8>, %arg1: tensor<8x1x1x4xi4>, %ar
return %0 : tensor<1x4x4x8xi32>
}
+// -----
+func.func @test_conv2d_mxfp(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ // expected-error at +1 {{'tosa.conv2d' op illegal: requires all of [bf16, mx_common, mx_fp4e2m1] profiles/extensions to be specified in the target environment}}
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = bf16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
// -----
func.func @test_conv3d(%arg0: tensor<1x4x8x21x17xi16>, %arg1: tensor<34x1x1x1x17xi8>, %arg2: tensor<34xi48>, %arg3: tensor<1xi16>, %arg4: tensor<1xi8>) -> tensor<1x4x8x21x34xi48> {
// expected-error at +1 {{'tosa.conv3d' op illegal: requires any of [int16] profiles/extensions to be specified in the target environment}}
diff --git a/mlir/test/Dialect/Tosa/ops.mlir b/mlir/test/Dialect/Tosa/ops.mlir
index b2c9bcfdb2be1..998826da567f3 100644
--- a/mlir/test/Dialect/Tosa/ops.mlir
+++ b/mlir/test/Dialect/Tosa/ops.mlir
@@ -142,6 +142,15 @@ func.func @test_conv2d(%arg0: tensor<1x4x4x4xf32>, %arg1: tensor<8x1x1x4xf32>, %
return %0 : tensor<1x4x4x8xf32>
}
+// -----
+// CHECK-LABEL: conv2d_mxfp
+func.func @test_conv2d_mxfp(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = bf16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
// -----
// CHECK-LABEL: conv2d_unranked_input
func.func @test_conv2d_unranked_input(%arg0: tensor<*xf32>, %arg1: tensor<8x1x1x4xf32>, %arg2: tensor<8xf32>, %arg3: tensor<1xf32>, %arg4: tensor<1xf32>) -> tensor<1x4x4x8xf32> {
diff --git a/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir b/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
index b80705348b46b..4623ed167a1b8 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
@@ -42,6 +42,16 @@ func.func @test_conv2d_fp8_acc32(%arg0: tensor<1x4x4x4xf8E5M2>, %arg1: tensor<8x
// -----
+func.func @test_conv2d_mxfp(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ // expected-error at +1 {{'tosa.conv2d' op illegal: requires specification version compatible with 1.1.draft (got 1.0) and requires all of [bf16, mx_common, mx_fp4e2m1] profiles/extensions to be specified in the target environment}}
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = bf16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
+// -----
+
func.func @test_conv3d_fp8_acc32(%arg0: tensor<1x4x8x21x17xf8E5M2>, %arg1: tensor<34x1x1x1x17xf8E5M2>, %arg2: tensor<34xf32>) -> tensor<1x4x8x21x34xf32> {
%input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf8E5M2>}> : () -> tensor<1xf8E5M2>
%weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf8E5M2>}> : () -> tensor<1xf8E5M2>
diff --git a/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir b/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
index e3eb88b10bded..4c48d3f12d645 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
@@ -72,6 +72,24 @@ func.func @test_conv2d_fp8_acc32(%arg0: tensor<1x4x4x4xf8E5M2>, %arg1: tensor<8x
// -----
+// CHECK-LABEL: test_conv2d_mxfp
+func.func @test_conv2d_mxfp(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = bf16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
+// CHECK-LABEL: test_conv2d_mxfp_acc32
+func.func @test_conv2d_mxfp_acc32(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = f32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf32>, tensor<1xf32>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
+// -----
+
// CHECK-LABEL: test_conv3d_fp8_acc32
func.func @test_conv3d_fp8_acc32(%arg0: tensor<1x4x8x21x17xf8E5M2>, %arg1: tensor<34x1x1x1x17xf8E5M2>, %arg2: tensor<34xf32>) -> tensor<1x4x8x21x34xf32> {
%input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf8E5M2>}> : () -> tensor<1xf8E5M2>
@@ -575,7 +593,7 @@ func.func @test_assert_equal_shape() {
}
// -----
-func.func @test_maxpool2d_adaptive(%arg0: tensor<1x32x32x8xf32>) -> tensor<1x32x32x8xf32> {
+func.func @test_maxpool2d_adaptive(%arg0: tensor<1x32x32x8xf32>) -> tensor<1x32x32x8xf32> {
%kernel = tosa.const_shape {values = dense<[1, 1]> : tensor<2xindex>} : () -> !tosa.shape<2>
%stride = tosa.const_shape {values = dense<[1, 1]> : tensor<2xindex>} : () -> !tosa.shape<2>
%pad = tosa.const_shape {values = dense<[0, 0, 0, 0]> : tensor<4xindex>} : () -> !tosa.shape<4>
diff --git a/mlir/test/Dialect/Tosa/verifier.mlir b/mlir/test/Dialect/Tosa/verifier.mlir
index 6fc88e4bd7694..96ce19f34aee5 100644
--- a/mlir/test/Dialect/Tosa/verifier.mlir
+++ b/mlir/test/Dialect/Tosa/verifier.mlir
@@ -244,6 +244,26 @@ func.func @test_slice_output_shape_mismatch_dynamic(%arg0: tensor<?x5x6xf32>) {
// -----
+func.func @test_conv2d_mxfp_invalid_weight_zp(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16>
+ // expected-error at +1 {{'tosa.conv2d' op expect block scaled weight to have fp32 zero point}}
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = bf16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf32>, tensor<1xf16>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
+// -----
+
+func.func @test_conv2d_mxfp_invalid_input_zp(%arg0: tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg1: tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, %arg2: tensor<8xf16>) -> tensor<1x4x4x8xf16> {
+ %input_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16>
+ %weight_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf32>}> : () -> tensor<1xf32>
+ // expected-error at +1 {{'tosa.conv2d' op expect block scaled input to have fp32 zero point}}
+ %0 = tosa.conv2d %arg0, %arg1, %arg2, %input_zp, %weight_zp {acc_type = bf16, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, 0, 0>, stride = array<i64: 1, 1>, local_bound = true} : (tensor<1x4x4x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8x1x1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<8xf16>, tensor<1xf16>, tensor<1xf32>) -> tensor<1x4x4x8xf16>
+ return %0 : tensor<1x4x4x8xf16>
+}
+
+// -----
+
func.func @test_depthwise_conv2d_invalid_padding(%arg0: tensor<1x4x4x4xf32>, %arg1: tensor<1x1x8x4xf32>, %arg2: tensor<8xf32>, %arg3: tensor<1xf32>, %arg4: tensor<1xf32>) -> tensor<1x4x4x8xf32> {
// expected-error at +1 {{'tosa.depthwise_conv2d' op expect all padding values to be >= 0, got 0, 0, -1, 0}}
%0 = tosa.depthwise_conv2d %arg0, %arg1, %arg2, %arg3, %arg4 {acc_type = f32, dilation = array<i64: 1, 1>, pad = array<i64: 0, 0, -1, 0>, stride = array<i64: 1, 1>, local_bound = true}
More information about the Mlir-commits
mailing list