[Mlir-commits] [mlir] [mlir][tosa] Add new block-scaled tensor type and support for MXFP CAST (PR #203583)
Luke Hutton
llvmlistbot at llvm.org
Thu Jun 25 09:43:14 PDT 2026
https://github.com/lhutton1 updated https://github.com/llvm/llvm-project/pull/203583
>From e9755287e97fb9bab0030b391da1a9f5846996f0 Mon Sep 17 00:00:00 2001
From: Luke Hutton <luke.hutton at arm.com>
Date: Tue, 19 May 2026 10:44:43 +0100
Subject: [PATCH 1/2] [mlir][tosa] Add new block-scaled tensor type and support
for MXFP CAST
This commit adds a new compound-block scaled tensor type and uses this
type to implement support for MXFP in the CAST operation, as per
the following specification changes:
https://github.com/arm/tosa-specification/pull/50,
https://github.com/arm/tosa-specification/pull/53.
The new block-scaled type is closely modeled after the `quant` dialect
type and supports the following parameters:
- value type - The type of the data values in each block.
- scale type - The type of the scale value associated with each block.
- block shape - The size and axis of each block.
Support for specifying constant scale values has not been added in this
commit, but will be added in a later one.
Example syntax for the new block-scaled type:
```
tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
```
As a pre-requisite for supporting MXFP CAST, the commit also adds new
extensions intended to split the exisiting EXT-MXFP extension into
separate extensions for each block-scaled type. This allows for more
fine-grained control over which block-scaled types are supported by a
given target. See specification change
https://github.com/arm/tosa-specification/pull/26 for details.
Finally, support for casting to/from the new block-scaled type has been
added, aligning with the behavior specified in the CAST operation
specification.
Change-Id: Idfab945402a95704ffeefbcee44172c3dde7552d
---
.../Dialect/Tosa/IR/TosaComplianceData.h.inc | 120 +++++++++++++++++-
.../mlir/Dialect/Tosa/IR/TosaOpBase.td | 61 +++++++--
mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h | 2 +
mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td | 69 +++-------
.../Dialect/Tosa/IR/TosaProfileCompliance.h | 30 ++++-
.../mlir/Dialect/Tosa/IR/TosaTypesBase.td | 117 +++++++++++------
mlir/lib/Dialect/Tosa/IR/TargetEnv.cpp | 14 ++
.../Dialect/Tosa/IR/TosaCanonicalizations.cpp | 44 ++++++-
mlir/lib/Dialect/Tosa/IR/TosaOps.cpp | 65 ++++++++++
.../Tosa/Transforms/TosaProfileCompliance.cpp | 92 ++++++++++----
.../Tosa/Transforms/TosaValidation.cpp | 2 +-
mlir/test/Dialect/Tosa/availability.mlir | 6 +-
mlir/test/Dialect/Tosa/canonicalize.mlir | 67 ++++++++++
mlir/test/Dialect/Tosa/invalid.mlir | 32 +++++
mlir/test/Dialect/Tosa/invalid_extension.mlir | 28 ++++
mlir/test/Dialect/Tosa/ops.mlir | 26 ++++
.../test/Dialect/Tosa/tosa-attach-target.mlir | 19 ++-
.../tosa-validation-version-1p0-invalid.mlir | 8 ++
.../tosa-validation-version-1p1-valid.mlir | 30 ++++-
mlir/test/Dialect/Tosa/verifier.mlir | 42 ++++++
20 files changed, 729 insertions(+), 145 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
index 50bb9f69c6242..26890abd187bd 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
@@ -1023,7 +1023,125 @@ extensionComplianceMap = {
{{{fp8e5m2T, fp16T}, SpecificationVersion::V_1_0},
{{fp8e5m2T, fp32T}, SpecificationVersion::V_1_0},
{{fp16T, fp8e5m2T}, SpecificationVersion::V_1_0},
- {{fp32T, fp8e5m2T}, SpecificationVersion::V_1_0}}}}},
+ {{fp32T, fp8e5m2T}, SpecificationVersion::V_1_0}}},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e4m3T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e4m3T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e4m3T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e4m3T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp8e4m3T, bs32_fp8ue8m0_fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e4m3T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e4m3, Extension::mx_common, Extension::mx_int8},
+ {{{fp8e4m3T, bs32_fp8ue8m0_mxint8T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp8e5m2T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp8e5m2T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp8e5m2T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp8e5m2T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp8e5m2T, bs32_fp8ue8m0_fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp8e5m2T},
+ SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::fp8e5m2, Extension::mx_common, Extension::mx_int8},
+ {{{fp8e5m2T, bs32_fp8ue8m0_mxint8T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{fp16T, bs32_fp8ue8m0_fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp16T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, fp32T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{fp16T, bs32_fp8ue8m0_fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp16T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, fp32T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{fp16T, bs32_fp8ue8m0_fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp16T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, fp32T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{fp16T, bs32_fp8ue8m0_fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp16T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, fp32T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{fp16T, bs32_fp8ue8m0_fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp16T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, fp32T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_int8},
+ {{{fp16T, bs32_fp8ue8m0_mxint8T}, SpecificationVersion::V_1_1_DRAFT},
+ {{fp32T, bs32_fp8ue8m0_mxint8T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp16T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, fp32T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{bf16T, bs32_fp8ue8m0_fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e4m3T, bf16T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{bf16T, bs32_fp8ue8m0_fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp8e5m2T, bf16T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{bf16T, bs32_fp8ue8m0_fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e3m2T, bf16T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{bf16T, bs32_fp8ue8m0_fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp6e2m3T, bf16T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{bf16T, bs32_fp8ue8m0_fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_fp4e2m1T, bf16T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::bf16, Extension::mx_common, Extension::mx_int8},
+ {{{bf16T, bs32_fp8ue8m0_mxint8T}, SpecificationVersion::V_1_1_DRAFT},
+ {{bs32_fp8ue8m0_mxint8T, bf16T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf}}},
{"tosa.cast_from_block_scaled",
{{{Extension::bf16, Extension::mxfp},
{{{fp4e2m1T, fp8ue8m0T, bf16T}, SpecificationVersion::V_1_1_DRAFT},
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaOpBase.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaOpBase.td
index 591073e9985ae..308e98ac42435 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOpBase.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOpBase.td
@@ -254,6 +254,13 @@ class Tosa_I32EnumAttr<string name, string description, string mnemonic,
// MXFP : Microscaling formats.
// MXFP_CONV : Microscaling format convolution.
// SHAPE : Shape calcuation operators.
+// MX_COMMON : Base for MXFP microscaling formats.
+// MX_FP4E2M1 : Microscaling format FP4E2M1.
+// MX_FP6E2M3 : Microscaling format FP6E2M3.
+// MX_FP6E3M2 : Microscaling format FP6E3M2.
+// MX_FP8E4M3 : Microscaling format FP8E4M3.
+// MX_FP8E5M2 : Microscaling format FP8E5M2.
+// MX_INT8 : Microscaling format INT8.
//===----------------------------------------------------------------------===//
def Tosa_NONE : I32EnumAttrCase<"none", 0>;
@@ -289,23 +296,35 @@ def Tosa_EXT_MXFP : I32EnumAttrCase<"mxfp", 12>;
def Tosa_EXT_INT64 : I32EnumAttrCase<"int64", 13>;
def Tosa_EXT_MXFP_CONV : I32EnumAttrCase<"mxfp_conv", 14>;
def Tosa_EXT_SHAPE : I32EnumAttrCase<"shape", 15>;
+def Tosa_EXT_MX_COMMON : I32EnumAttrCase<"mx_common", 16>;
+def Tosa_EXT_MX_FP4E2M1 : I32EnumAttrCase<"mx_fp4e2m1", 17>;
+def Tosa_EXT_MX_FP6E2M3 : I32EnumAttrCase<"mx_fp6e2m3", 18>;
+def Tosa_EXT_MX_FP6E3M2 : I32EnumAttrCase<"mx_fp6e3m2", 19>;
+def Tosa_EXT_MX_FP8E4M3 : I32EnumAttrCase<"mx_fp8e4m3", 20>;
+def Tosa_EXT_MX_FP8E5M2 : I32EnumAttrCase<"mx_fp8e5m2", 21>;
+def Tosa_EXT_MX_INT8 : I32EnumAttrCase<"mx_int8", 22>;
def Tosa_ExtensionAttr
- : Tosa_I32EnumAttr<"Extension", "supported TOSA extensions", "ext", [
- Tosa_EXT_NONE, Tosa_EXT_INT16, Tosa_EXT_INT4, Tosa_EXT_BF16,
- Tosa_EXT_FP8E4M3, Tosa_EXT_FP8E5M2, Tosa_EXT_FFT, Tosa_EXT_VARIABLE,
- Tosa_EXT_CONTROLFLOW, Tosa_EXT_DOUBLEROUND, Tosa_EXT_INEXACTROUND,
- Tosa_EXT_DYNAMIC, Tosa_EXT_MXFP, Tosa_EXT_INT64, Tosa_EXT_MXFP_CONV,
- Tosa_EXT_SHAPE,
- ]> {
+ : Tosa_I32EnumAttr<
+ "Extension", "supported TOSA extensions", "ext",
+ [Tosa_EXT_NONE, Tosa_EXT_INT16, Tosa_EXT_INT4, Tosa_EXT_BF16,
+ Tosa_EXT_FP8E4M3, Tosa_EXT_FP8E5M2, Tosa_EXT_FFT, Tosa_EXT_VARIABLE,
+ Tosa_EXT_CONTROLFLOW, Tosa_EXT_DOUBLEROUND, Tosa_EXT_INEXACTROUND,
+ Tosa_EXT_DYNAMIC, Tosa_EXT_MXFP, Tosa_EXT_INT64, Tosa_EXT_MXFP_CONV,
+ Tosa_EXT_SHAPE, 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 = [{
- static llvm::SmallVector<Extension, 14> getAllValues() {
+ static llvm::SmallVector<Extension, 22> getAllValues() {
return {
Extension::int16, Extension::int4, Extension::bf16,
Extension::fp8e4m3, Extension::fp8e5m2, Extension::fft,
Extension::variable, Extension::controlflow, Extension::doubleround,
Extension::inexactround, Extension::dynamic, Extension::mxfp,
- Extension::int64, Extension::mxfp_conv, Extension::shape
+ Extension::int64, Extension::mxfp_conv, Extension::shape,
+ Extension::mx_common, Extension::mx_fp4e2m1, Extension::mx_fp6e2m3,
+ Extension::mx_fp6e3m2, Extension::mx_fp8e4m3, Extension::mx_fp8e5m2,
+ Extension::mx_int8
};
}
}];
@@ -484,6 +503,7 @@ def Tosa_RoundingModeAttr
: Tosa_I32EnumAttr<"RoundingMode", "Supported rounding modes", "rounding_mode",
[Tosa_ROUNDING_SINGLE_ROUND, Tosa_ROUNDING_INEXACT_ROUND, Tosa_ROUNDING_DOUBLE_ROUND]>;
+// Block_size attr is deprecated and will be removed in the future
def Tosa_BLOCK_SIZE_1 : I32EnumAttrCase<"BLOCK_SIZE_1", 1>;
def Tosa_BLOCK_SIZE_32 : I32EnumAttrCase<"BLOCK_SIZE_32", 32>;
@@ -497,6 +517,29 @@ def Tosa_BlockSizeAttr
}];
}
+def Tosa_BLOCK_SHAPE_32 : I32EnumAttrCase<"BLOCK_SHAPE_32", 32>;
+
+def Tosa_BlockShape
+ : Tosa_I32Enum<
+ "BlockShape",
+ "Block shape for the block_scaled formats."
+ "The names follow the convention of BLOCK_SHAPE_M where M "
+ "is the block size of the innermost dimension. Similarly, "
+ "BLOCK_SHAPE_MxN indicates block size of N in the innermost "
+ "dimension and block size of M in the next innermost dimension. "
+ "As of now, only 1 dimension (innermost) is supported for "
+ "block-scaling.",
+ [Tosa_BLOCK_SHAPE_32]>;
+
+def Tosa_BlockShapeAttr
+ : EnumAttr<Tosa_Dialect, Tosa_BlockShape, "block_shape"> {
+ let extraClassDeclaration = [{
+ static uint32_t getBlockShapeValue(BlockShape blockShape) {
+ return static_cast<uint32_t>(blockShape);
+ }
+ }];
+}
+
//===----------------------------------------------------------------------===//
// TOSA Interfaces.
//===----------------------------------------------------------------------===//
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
index e0626368175ee..a5c6037692ff1 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
@@ -92,6 +92,8 @@ class TosaShapeOperatorWithSameRanks
}
};
+LogicalResult verifyBlockScaledTensorType(Operation &op, mlir::Type type);
+
} // namespace tosa
} // namespace OpTrait
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
index a99fb2fcae547..32f8bf08cbc84 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
@@ -2689,57 +2689,6 @@ def Tosa_CastOp: Tosa_Op<"cast", [Pure, SameOperandsAndResultShape,
let description = [{
Casts a tensor from one data type to another.
- * This table is showing the supported conversions from the TOSA Specification.
- * The MLIR dialect here can be used to represent other conversions.
-
- | Mode | Input | Output |
- |--------------------------|---------|---------|
- | fp16 to fp32 | float16 | float32 |
- | fp16 to int 16 | float16 | int16 |
- | fp16 to int 32 | float16 | int32 |
- | fp16 to int 8 | float16 | int8 |
- | fp32 to fp16 | float32 | float16 |
- | fp32 to int 16 | float32 | int16 |
- | fp32 to int 32 | float32 | int32 |
- | fp32 to int 8 | float32 | int8 |
- | int 16 to fp16 | int16 | float16 |
- | int 16 to fp32 | int16 | float32 |
- | int 32 to fp16 | int32 | float16 |
- | int 32 to fp32 | int32 | float32 |
- | int 8 to fp16 | int8 | float16 |
- | int 8 to fp32 | int8 | float32 |
- | bool to int 16 | Boolean | int16 |
- | bool to int 32 | Boolean | int32 |
- | bool to int 8 | Boolean | int8 |
- | int 16 to bool | int16 | Boolean |
- | int 16 to int 32 | int16 | int32 |
- | int 16 to int 8 | int16 | int8 |
- | int 32 to bool | int32 | Boolean |
- | int 32 to int 16 | int32 | int16 |
- | int 32 to int 8 | int32 | int8 |
- | int 8 to bool | int8 | Boolean |
- | int 8 to int 16 | int8 | int16 |
- | int 8 to int 32 | int8 | int32 |
- | bf16 to fp32 | bf16 | float32 |
- | bf16 to int 16 | bf16 | int16 |
- | bf16 to int 32 | bf16 | int32 |
- | bf16 to int 8 | bf16 | int8 |
- | fp32 to bf16 | float32 | bf16 |
- | int 16 to bf16 | int16 | bf16 |
- | int 32 to bf16 | int32 | bf16 |
- | int 8 to bf16 | int8 | bf16 |
- | bf16 to fp8e4m3 | bf16 | fp8e4m3 |
- | fp8e4m3 to bf16 | fp8e4m3 | bf16 |
- | bf16 to fp8e5m2 | bf16 | fp8e5m2 |
- | fp8e5m2 to bf16 | fp8e5m2 | bf16 |
- | fp16 to fp8e4m3 | float16 | fp8e4m3 |
- | fp32 to fp8e4m3 | float32 | fp8e4m3 |
- | fp8e4m3 to fp16 | fp8e4m3 | float16 |
- | fp8e4m3 to fp32 | fp8e4m3 | float32 |
- | fp16 to fp8e5m2 | float16 | fp8e5m2 |
- | fp32 to fp8e5m2 | float32 | fp8e5m2 |
- | fp8e5m2 to fp16 | fp8e5m2 | float16 |
- | fp8e5m2 to fp32 | fp8e5m2 | float32 |
}];
let arguments = (ins
@@ -2750,19 +2699,25 @@ def Tosa_CastOp: Tosa_Op<"cast", [Pure, SameOperandsAndResultShape,
Tosa_Tensor:$output
);
- list<Availability> availability = [
- Profile<[Tosa_PRO_INT, Tosa_PRO_FP]>,
- Extension<[Tosa_EXT_FP8E4M3, Tosa_EXT_FP8E5M2, Tosa_EXT_BF16, Tosa_EXT_INT64]>,
+ list<Availability> availability =
+ [Profile<[Tosa_PRO_INT, Tosa_PRO_FP]>,
+ Extension<[Tosa_EXT_FP8E4M3, Tosa_EXT_FP8E5M2, Tosa_EXT_BF16,
+ Tosa_EXT_INT64, 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 assemblyFormat = "operands attr-dict `:` functional-type(operands, results)";
let hasFolder = 1;
let hasCanonicalizer = 1;
+ let hasVerifier = 1;
}
//===----------------------------------------------------------------------===//
// Operator: cast_from_block_scaled
+//
+// Note: This operation is deprecated. It will be removed in the future.
//===----------------------------------------------------------------------===//
def Tosa_CastFromBlockScaledOp: Tosa_InferShapedTypeOp<"cast_from_block_scaled", [Pure]> {
let summary = "Apply scales from a scale tensor to the values in a value tensor";
@@ -2771,6 +2726,8 @@ def Tosa_CastFromBlockScaledOp: Tosa_InferShapedTypeOp<"cast_from_block_scaled",
Apply the scales from a scale tensor to the values in a value tensor, casting
the result to the output type. The block dimension must be the last dimension
of the tensor.
+
+ Note: This operation is deprecated. It will be removed in the future.
}];
let arguments = (ins
@@ -2794,6 +2751,8 @@ def Tosa_CastFromBlockScaledOp: Tosa_InferShapedTypeOp<"cast_from_block_scaled",
//===----------------------------------------------------------------------===//
// Operator: cast_to_block_scaled
+//
+// Note: This operation is deprecated. It will be removed in the future.
//===----------------------------------------------------------------------===//
def Tosa_CastToBlockScaledOp : Tosa_InferShapedTypeOp<"cast_to_block_scaled", [Pure]> {
let summary = "Calculate scale tensor values per block, output to separate scale and data tensors.";
@@ -2803,6 +2762,8 @@ def Tosa_CastToBlockScaledOp : Tosa_InferShapedTypeOp<"cast_to_block_scaled", [P
scaled data values from an input tensor. The output tensors are cast to the
specified scale and value types. The block dimension will be the last dimension
of the tensor.
+
+ Note: This operation is deprecated. It will be removed in the future.
}];
let arguments = (ins
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaProfileCompliance.h b/mlir/include/mlir/Dialect/Tosa/IR/TosaProfileCompliance.h
index 0135a651be481..4e14b27a421e4 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaProfileCompliance.h
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaProfileCompliance.h
@@ -23,10 +23,22 @@ using namespace mlir::tosa;
// Type Compilance Definition
//===----------------------------------------------------------------------===//
-typedef struct {
+struct TypeInfo {
+ TypeInfo(mlir::TypeID typeID, uint32_t bitWidth)
+ : typeID(typeID), bitWidth(bitWidth), valueTypeID(mlir::TypeID()),
+ scaleTypeID(mlir::TypeID()), blockSize(0) {}
+
+ TypeInfo(mlir::TypeID typeID, uint32_t bitWidth, mlir::TypeID valueTypeID,
+ mlir::TypeID scaleTypeID, uint32_t blockSize)
+ : typeID(typeID), bitWidth(bitWidth), valueTypeID(valueTypeID),
+ scaleTypeID(scaleTypeID), blockSize(blockSize) {}
+
mlir::TypeID typeID;
uint32_t bitWidth;
-} TypeInfo;
+ mlir::TypeID valueTypeID;
+ mlir::TypeID scaleTypeID;
+ uint32_t blockSize;
+};
enum CheckCondition {
invalid,
@@ -70,6 +82,14 @@ class ProfileInfoDepot {
private:
TypeInfo convertTypeToInfo(Type type) {
+ if (auto blockScaledTy = dyn_cast<tosa::BlockScaledType>(type)) {
+ Type valueTy = blockScaledTy.getValueType();
+ Type scaleTy = blockScaledTy.getScaleType();
+ return {
+ type.getTypeID(), tosa::getBitWidth(valueTy), valueTy.getTypeID(),
+ scaleTy.getTypeID(),
+ BlockShapeAttr::getBlockShapeValue(blockScaledTy.getBlockShape())};
+ }
return {type.getTypeID(), tosa::getBitWidth(type)};
}
@@ -128,7 +148,9 @@ class TosaProfileCompliance {
const SmallVector<ArrayRef<T>> &specDefinedProfileSet);
bool isSameTypeInfo(TypeInfo a, TypeInfo b) {
- return a.typeID == b.typeID && a.bitWidth == b.bitWidth;
+ return a.typeID == b.typeID && a.bitWidth == b.bitWidth &&
+ a.valueTypeID == b.valueTypeID && a.scaleTypeID == b.scaleTypeID &&
+ a.blockSize == b.blockSize;
}
// Find the required profiles or extensions from the compliance info according
@@ -145,7 +167,7 @@ class TosaProfileCompliance {
SmallVector<StringRef>
stringifyProfile(const SmallVector<ArrayRef<T>> &profileSet);
- static llvm::SmallString<7> stringifyTypeInfo(const TypeInfo &typeInfo);
+ static llvm::SmallString<32> stringifyTypeInfo(const TypeInfo &typeInfo);
private:
template <typename T>
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
index 10ddd3438aedd..0a43f14865f56 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
@@ -103,15 +103,39 @@ def Tosa_MXInt8
}];
}
+def Tosa_MXFPValue
+ : AnyTypeOf<[F8E4M3FN, F8E5M2, F4E2M1FN, F6E2M3FN, F6E3M2FN, Tosa_MXInt8],
+ "micro-scaling format number">;
+def Tosa_MXFPScale
+ : AnyTypeOf<[F8E8M0FNU], "micro-scaling format scale number">;
+
+def Tosa_BlockScaled : Tosa_Type<"BlockScaled", "block_scaled"> {
+ let summary = "Block scaled tensor element type.";
+
+ let description = [{
+ This does not specify an implementation type. A tensor of this type is a
+ tensor which has block-scaled quantization applied.
+
+ This compound type is made up of 3 components:
+ `value_type` - The type of the data values in each block.
+ `scale_type` - The type of the scale value associated with each block.
+ `block_shape` - The size and axis of each block. Currently only supports
+ specifying the block size along the innermost dimension.
+ }];
+
+ let parameters = (ins Tosa_MXFPValue:$value_type, Tosa_MXFPScale:$scale_type,
+ EnumParameter<Tosa_BlockShape>:$block_shape);
+
+ let assemblyFormat =
+ "`<` $value_type```:```$scale_type```:```$block_shape `>`";
+}
+
//===----------------------------------------------------------------------===//
// Multi-category types.
//===----------------------------------------------------------------------===//
-def Tosa_AnyNumber : AnyTypeOf<[Tosa_Int, Tosa_QuantizedInt, AnyFloat, Tosa_MXInt8],
- "number">;
-
-def Tosa_MXFPNumber : AnyTypeOf<[F8E4M3FN, F8E5M2, F4E2M1FN, F6E2M3FN, F6E3M2FN, Tosa_MXInt8],
- "micro-scaling format number">;
-def Tosa_MXFPScaleNumber : AnyTypeOf<[F8E8M0FNU], "micro-scaling format scale number">;
+def Tosa_AnyNumber : AnyTypeOf<[Tosa_Int, Tosa_QuantizedInt, AnyFloat,
+ Tosa_MXInt8, Tosa_BlockScaled],
+ "number">;
//===----------------------------------------------------------------------===//
// TOSA Tensor Conformance
@@ -129,16 +153,29 @@ def AtLeastRankOne : And<[
IsRankedTensorTypePred,
CPred<"::llvm::cast<::mlir::RankedTensorType>($_self).getRank() >= 1">]>;
-class TosaTensorOf<
- list<Type> allowedTypes, string summary = "tosa-conformant tensor">
- : TensorOf<allowedTypes, [Or<[HasNo0Dimensions, IsUnrankedTensorTypePred]>], summary>;
-
-class TosaRankedTensorOf<
- list<Type> allowedTypes, list<Pred> preds = [], string summary = "tosa-conformant ranked tensor">
- : RankedTensorOf<allowedTypes, !listconcat([HasNo0Dimensions], preds), summary>;
-
-class TosaUnrankedTensorOf<list<Type> allowedTypes, list<Pred> preds = [], string summary = "tosa-conformant unranked tensor">
- : UnrankedTensorOf<allowedTypes, preds, summary>;
+def IsValidBlockScaledTensorType
+ : CPred<"::mlir::succeeded(::mlir::OpTrait::tosa::"
+ "verifyBlockScaledTensorType($_op, $_self))">;
+
+class TosaTensorOf<list<Type> allowedTypes,
+ string summary = "tosa-conformant tensor">
+ : TensorOf<allowedTypes,
+ [Or<[HasNo0Dimensions, IsUnrankedTensorTypePred]>,
+ IsValidBlockScaledTensorType],
+ summary>;
+
+class TosaRankedTensorOf<list<Type> allowedTypes, list<Pred> preds = [],
+ string summary = "tosa-conformant ranked tensor">
+ : RankedTensorOf<
+ allowedTypes,
+ !listconcat([HasNo0Dimensions, IsValidBlockScaledTensorType], preds),
+ summary>;
+
+class TosaUnrankedTensorOf<list<Type> allowedTypes, list<Pred> preds = [],
+ string summary = "tosa-conformant unranked tensor">
+ : UnrankedTensorOf<allowedTypes,
+ !listconcat([IsValidBlockScaledTensorType], preds),
+ summary>;
class TosaTensorRankOf<list<Type> allowedTypes, list<int> ranks>
: TosaRankedTensorOf<allowedTypes,
@@ -217,32 +254,28 @@ def Tosa_IndexTensor2D : AnyTypeOf<[
def Tosa_TensorAtLeast1D : AnyTypeOf<[
Tosa_UnrankedTensor, TosaRankedTensorOf<[Tosa_AnyNumber], [AtLeastRankOne]>], "tosa-conformant tensor of at least rank 1", "::mlir::TensorType">;
-def Tosa_MXFPDataTensor3D : AnyTypeOf<[
- TosaUnrankedTensorOf<[Tosa_MXFPNumber]>,
- TosaTensorRankOf<[Tosa_MXFPNumber], [3]>
-]>;
-def Tosa_MXFPScaleTensor3D : AnyTypeOf<[
- TosaUnrankedTensorOf<[Tosa_MXFPScaleNumber]>,
- TosaTensorRankOf<[Tosa_MXFPScaleNumber], [3]>
-]>;
-def Tosa_MXFPDataTensor4D : AnyTypeOf<[
- TosaUnrankedTensorOf<[Tosa_MXFPNumber]>,
- TosaTensorRankOf<[Tosa_MXFPNumber], [4]>
-]>;
-def Tosa_MXFPScaleTensor4D : AnyTypeOf<[
- TosaUnrankedTensorOf<[Tosa_MXFPScaleNumber]>,
- TosaTensorRankOf<[Tosa_MXFPScaleNumber], [4]>
-]>;
-def Tosa_MXFPDataTensorAtLeast1D : AnyTypeOf<[
- TosaUnrankedTensorOf<[Tosa_MXFPNumber]>,
- TosaRankedTensorOf<[Tosa_MXFPNumber], [AtLeastRankOne]>],
- "tosa-conformant tensor of at least rank 1", "::mlir::TensorType"
->;
-def Tosa_MXFPScaleTensorAtLeast1D : AnyTypeOf<[
- TosaUnrankedTensorOf<[Tosa_MXFPScaleNumber]>,
- TosaRankedTensorOf<[Tosa_MXFPScaleNumber], [AtLeastRankOne]>],
- "tosa-conformant tensor of at least rank 1", "::mlir::TensorType"
->;
+def Tosa_MXFPDataTensor3D
+ : AnyTypeOf<[TosaUnrankedTensorOf<[Tosa_MXFPValue]>,
+ TosaTensorRankOf<[Tosa_MXFPValue], [3]>]>;
+def Tosa_MXFPScaleTensor3D
+ : AnyTypeOf<[TosaUnrankedTensorOf<[Tosa_MXFPScale]>,
+ TosaTensorRankOf<[Tosa_MXFPScale], [3]>]>;
+def Tosa_MXFPDataTensor4D
+ : AnyTypeOf<[TosaUnrankedTensorOf<[Tosa_MXFPValue]>,
+ TosaTensorRankOf<[Tosa_MXFPValue], [4]>]>;
+def Tosa_MXFPScaleTensor4D
+ : AnyTypeOf<[TosaUnrankedTensorOf<[Tosa_MXFPScale]>,
+ TosaTensorRankOf<[Tosa_MXFPScale], [4]>]>;
+def Tosa_MXFPDataTensorAtLeast1D
+ : AnyTypeOf<[TosaUnrankedTensorOf<[Tosa_MXFPValue]>,
+ TosaRankedTensorOf<[Tosa_MXFPValue], [AtLeastRankOne]>],
+ "tosa-conformant tensor of at least rank 1",
+ "::mlir::TensorType">;
+def Tosa_MXFPScaleTensorAtLeast1D
+ : AnyTypeOf<[TosaUnrankedTensorOf<[Tosa_MXFPScale]>,
+ TosaRankedTensorOf<[Tosa_MXFPScale], [AtLeastRankOne]>],
+ "tosa-conformant tensor of at least rank 1",
+ "::mlir::TensorType">;
//===----------------------------------------------------------------------===//
// Generic scalar, vector, or tensor of a particular type.
diff --git a/mlir/lib/Dialect/Tosa/IR/TargetEnv.cpp b/mlir/lib/Dialect/Tosa/IR/TargetEnv.cpp
index dc18fcaa04c8a..56e4901811dcb 100644
--- a/mlir/lib/Dialect/Tosa/IR/TargetEnv.cpp
+++ b/mlir/lib/Dialect/Tosa/IR/TargetEnv.cpp
@@ -56,6 +56,13 @@ TosaSpecificationVersion getMinVersion(const Extension &extension) {
case Extension::int64:
case Extension::mxfp_conv:
case Extension::shape:
+ case Extension::mx_common:
+ case Extension::mx_fp4e2m1:
+ case Extension::mx_fp6e2m3:
+ case Extension::mx_fp6e3m2:
+ case Extension::mx_fp8e4m3:
+ case Extension::mx_fp8e5m2:
+ case Extension::mx_int8:
return TosaSpecificationVersion(1, 1, true);
case Extension::none:
return TosaSpecificationVersion(0, 0);
@@ -76,6 +83,13 @@ SmallVector<Profile, 2> getCooperativeProfiles(Extension ext) {
case Extension::fft:
case Extension::mxfp:
case Extension::mxfp_conv:
+ case Extension::mx_common:
+ case Extension::mx_fp4e2m1:
+ case Extension::mx_fp6e2m3:
+ case Extension::mx_fp6e3m2:
+ case Extension::mx_fp8e4m3:
+ case Extension::mx_fp8e5m2:
+ case Extension::mx_int8:
return {Profile::pro_fp};
case Extension::variable:
case Extension::controlflow:
diff --git a/mlir/lib/Dialect/Tosa/IR/TosaCanonicalizations.cpp b/mlir/lib/Dialect/Tosa/IR/TosaCanonicalizations.cpp
index 63d40ed4a95a5..99315a787f12e 100644
--- a/mlir/lib/Dialect/Tosa/IR/TosaCanonicalizations.cpp
+++ b/mlir/lib/Dialect/Tosa/IR/TosaCanonicalizations.cpp
@@ -1184,12 +1184,50 @@ struct NonNarrowingCastsOptimization : public OpRewritePattern<tosa::CastOp> {
}
};
+struct CancellingBlockScaledCastsOptimization
+ : public OpRewritePattern<tosa::CastOp> {
+ using OpRewritePattern<tosa::CastOp>::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(tosa::CastOp castOp,
+ PatternRewriter &rewriter) const override {
+ const Value outerInput = castOp.getInput();
+ auto innerCastOp = outerInput.getDefiningOp<tosa::CastOp>();
+ if (!innerCastOp)
+ return rewriter.notifyMatchFailure(castOp,
+ "input must be a cast operation");
+
+ const Value innerInput = innerCastOp.getInput();
+ const auto innerInputTy = llvm::cast<ShapedType>(innerInput.getType());
+ const auto innerOutputTy = llvm::cast<ShapedType>(innerCastOp.getType());
+ const auto outerOutputTy = llvm::cast<ShapedType>(castOp.getType());
+
+ if (!llvm::isa<tosa::BlockScaledType>(innerInputTy.getElementType()))
+ return rewriter.notifyMatchFailure(
+ castOp, "inner cast input must have block scaled element type");
+
+ if (innerInputTy != outerOutputTy)
+ return rewriter.notifyMatchFailure(
+ castOp, "inner input type must match outer output type");
+
+ const Type innerOutputElemType = innerOutputTy.getElementType();
+ const bool isLosslessCast = isa<Float32Type>(innerOutputElemType);
+ if (!isLosslessCast)
+ return rewriter.notifyMatchFailure(
+ castOp, "avoid cancelling casts that could by lossy");
+
+ rewriter.replaceOp(castOp, innerInput);
+
+ return success();
+ }
+};
+
void CastOp::getCanonicalizationPatterns(RewritePatternSet &results,
MLIRContext *context) {
- results.add<NonNarrowingCastsOptimization>(context);
+ results.add<NonNarrowingCastsOptimization,
+ CancellingBlockScaledCastsOptimization>(context);
}
-struct CancellingBlockScaledCastsOptimization
+struct CancellingCastToFromBlockScaledOptimization
: public OpRewritePattern<tosa::CastToBlockScaledOp> {
using OpRewritePattern<tosa::CastToBlockScaledOp>::OpRewritePattern;
@@ -1235,7 +1273,7 @@ struct CancellingBlockScaledCastsOptimization
void CastToBlockScaledOp::getCanonicalizationPatterns(
RewritePatternSet &results, MLIRContext *context) {
- results.add<CancellingBlockScaledCastsOptimization>(context);
+ results.add<CancellingCastToFromBlockScaledOptimization>(context);
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
index f05399cf6b00b..486431a562714 100644
--- a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
+++ b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
@@ -628,6 +628,8 @@ Value mlir::tosa::createPadConstTensor(OpBuilder &builder, Location loc,
}
unsigned mlir::tosa::getBitWidth(Type type) {
+ if (auto blockScaledTy = dyn_cast<tosa::BlockScaledType>(type))
+ return getBitWidth(blockScaledTy.getValueType());
if (dyn_cast<tosa::mxint8Type>(type))
return 8;
return type.getIntOrFloatBitWidth();
@@ -734,6 +736,41 @@ LogicalResult mlir::tosa::mxint8Type::convertFromAttribute(
return cast<IntegerType>(attrType).convertFromAttribute(attr, result);
}
+//===----------------------------------------------------------------------===//
+// TOSA block scaling utilities.
+//===----------------------------------------------------------------------===//
+
+LogicalResult OpTrait::tosa::verifyBlockScaledTensorType(Operation &op,
+ mlir::Type type) {
+ const auto tensorType = llvm::cast<ShapedType>(type);
+ const BlockScaledType elemType =
+ llvm::dyn_cast<BlockScaledType>(tensorType.getElementType());
+ if (!elemType)
+ return success();
+
+ if (!tensorType.hasRank())
+ return success();
+
+ if (tensorType.getRank() == 0)
+ return op.emitError()
+ << "tensor type " << type
+ << " does not support block scaling on scalar tensors";
+
+ const int64_t blockedDimension = tensorType.getShape().back();
+ if (ShapedType::isDynamic(blockedDimension))
+ return success();
+
+ const uint32_t blockSize =
+ BlockShapeAttr::getBlockShapeValue(elemType.getBlockShape());
+ if (blockedDimension % blockSize != 0)
+ return op.emitError()
+ << "tensor type " << type
+ << " blocked dimension must be a multiple of block size, got "
+ << blockedDimension << " and block size " << blockSize;
+
+ return success();
+}
+
//===----------------------------------------------------------------------===//
// TOSA Operator Verifiers.
//===----------------------------------------------------------------------===//
@@ -5004,6 +5041,34 @@ LogicalResult RescaleOp::inferReturnTypeComponents(
return success();
}
+LogicalResult CastOp::verify() {
+ const ShapedType inputType = llvm::cast<ShapedType>(getInput().getType());
+ const ShapedType outputType = llvm::cast<ShapedType>(getType());
+ const Type inputElementType = inputType.getElementType();
+ const Type outputElementType = outputType.getElementType();
+
+ const bool inputIsBlockScaled = llvm::isa<BlockScaledType>(inputElementType);
+ const bool outputIsBlockScaled =
+ llvm::isa<BlockScaledType>(outputElementType);
+ if (!inputIsBlockScaled && !outputIsBlockScaled)
+ return success();
+
+ if (inputIsBlockScaled && outputIsBlockScaled)
+ return emitOpError()
+ << "requires exactly one of input or output to have block scaled "
+ "element type";
+
+ const Type scalarElementType =
+ inputIsBlockScaled ? outputElementType : inputElementType;
+ if (!llvm::isa<FloatType>(scalarElementType))
+ return emitOpError()
+ << "requires non-block-scaled element type to be floating-point "
+ "when casting to or from block scaled element type, got "
+ << scalarElementType;
+
+ return success();
+}
+
LogicalResult CastFromBlockScaledOp::inferReturnTypeComponents(
MLIRContext *context, ::std::optional<Location> location,
CastFromBlockScaledOp::Adaptor adaptor,
diff --git a/mlir/lib/Dialect/Tosa/Transforms/TosaProfileCompliance.cpp b/mlir/lib/Dialect/Tosa/Transforms/TosaProfileCompliance.cpp
index 9511d4da89dbd..9fbdfcc1d6690 100644
--- a/mlir/lib/Dialect/Tosa/Transforms/TosaProfileCompliance.cpp
+++ b/mlir/lib/Dialect/Tosa/Transforms/TosaProfileCompliance.cpp
@@ -8,6 +8,7 @@
#include "mlir/Dialect/Tosa/IR/TosaProfileCompliance.h"
#include "llvm/ADT/StringExtras.h"
+#include "llvm/Support/raw_ostream.h"
using namespace mlir;
using namespace mlir::tosa;
@@ -27,12 +28,37 @@ TosaProfileCompliance::TosaProfileCompliance() {
const TypeInfo fp8e5m2T = {mlir::Float8E5M2Type::getTypeID(), 8};
// micro-scaling formats
+ // Note: these types exist to suppport the deprecated block_scaled operations
+ // and can be removed once those operations are removed.
const TypeInfo fp6e2m3T = {mlir::Float6E2M3FNType::getTypeID(), 6};
const TypeInfo fp6e3m2T = {mlir::Float6E3M2FNType::getTypeID(), 6};
const TypeInfo fp4e2m1T = {mlir::Float4E2M1FNType::getTypeID(), 4};
const TypeInfo fp8ue8m0T = {mlir::Float8E8M0FNUType::getTypeID(), 8};
const TypeInfo mxint8T = {mlir::tosa::mxint8Type::getTypeID(), 8};
+ // Block scaled formats
+ const TypeID blockScaledID = mlir::tosa::BlockScaledType::getTypeID();
+ const TypeID fp4e2m1ID = mlir::Float4E2M1FNType::getTypeID();
+ const TypeID fp6e2m3ID = mlir::Float6E2M3FNType::getTypeID();
+ const TypeID fp6e3m2ID = mlir::Float6E3M2FNType::getTypeID();
+ const TypeID fp8e4m3ID = mlir::Float8E4M3FNType::getTypeID();
+ const TypeID fp8e5m2ID = mlir::Float8E5M2Type::getTypeID();
+ const TypeID fp8ue8m0ID = mlir::Float8E8M0FNUType::getTypeID();
+ const TypeID mxint8ID = mlir::tosa::mxint8Type::getTypeID();
+
+ const TypeInfo bs32_fp8ue8m0_fp4e2m1T = {blockScaledID, 4, fp4e2m1ID,
+ fp8ue8m0ID, 32};
+ const TypeInfo bs32_fp8ue8m0_fp6e2m3T = {blockScaledID, 6, fp6e2m3ID,
+ fp8ue8m0ID, 32};
+ const TypeInfo bs32_fp8ue8m0_fp6e3m2T = {blockScaledID, 6, fp6e3m2ID,
+ fp8ue8m0ID, 32};
+ const TypeInfo bs32_fp8ue8m0_fp8e4m3T = {blockScaledID, 8, fp8e4m3ID,
+ fp8ue8m0ID, 32};
+ const TypeInfo bs32_fp8ue8m0_fp8e5m2T = {blockScaledID, 8, fp8e5m2ID,
+ fp8ue8m0ID, 32};
+ const TypeInfo bs32_fp8ue8m0_mxint8T = {blockScaledID, 8, mxint8ID,
+ fp8ue8m0ID, 32};
+
// The profile-based compliance content below is auto-generated by a script
// in https://github.com/arm/tosa-specification
#include "mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc"
@@ -669,31 +695,47 @@ SmallVector<StringRef> TosaProfileCompliance::stringifyProfile(
return debugStrings;
}
-llvm::SmallString<7>
+llvm::SmallString<32>
TosaProfileCompliance::stringifyTypeInfo(const TypeInfo &typeInfo) {
- if (typeInfo.typeID == mlir::IntegerType::getTypeID()) {
- return {"i" + llvm::utostr(typeInfo.bitWidth)};
- }
- if (typeInfo.typeID == mlir::Float16Type::getTypeID()) {
- return {"f16"};
- } else if (typeInfo.typeID == mlir::Float32Type::getTypeID()) {
- return {"f32"};
- } else if (typeInfo.typeID == mlir::BFloat16Type::getTypeID()) {
- return {"bf16"};
- } else if (typeInfo.typeID == mlir::Float8E4M3FNType::getTypeID()) {
- return {"fp8e4m3"};
- } else if (typeInfo.typeID == mlir::Float8E5M2Type::getTypeID()) {
- return {"fp8e5m2"};
- } else if (typeInfo.typeID == mlir::Float6E2M3FNType::getTypeID()) {
- return {"fp6e2m3"};
- } else if (typeInfo.typeID == mlir::Float6E3M2FNType::getTypeID()) {
- return {"fp6e3m2"};
- } else if (typeInfo.typeID == mlir::Float4E2M1FNType::getTypeID()) {
- return {"fp4e2m1"};
- } else if (typeInfo.typeID == mlir::Float8E8M0FNUType::getTypeID()) {
- return {"fp8e8m0"};
- } else if (typeInfo.typeID == tosa::mxint8Type::getTypeID()) {
- return {"mxint8"};
+ const auto stringifyScalarTypeInfo =
+ [](const TypeInfo &typeInfo) -> llvm::SmallString<32> {
+ if (typeInfo.typeID == mlir::IntegerType::getTypeID()) {
+ return {"i" + llvm::utostr(typeInfo.bitWidth)};
+ }
+ if (typeInfo.typeID == mlir::Float16Type::getTypeID()) {
+ return {"f16"};
+ } else if (typeInfo.typeID == mlir::Float32Type::getTypeID()) {
+ return {"f32"};
+ } else if (typeInfo.typeID == mlir::BFloat16Type::getTypeID()) {
+ return {"bf16"};
+ } else if (typeInfo.typeID == mlir::Float8E4M3FNType::getTypeID()) {
+ return {"fp8e4m3"};
+ } else if (typeInfo.typeID == mlir::Float8E5M2Type::getTypeID()) {
+ return {"fp8e5m2"};
+ } else if (typeInfo.typeID == mlir::Float6E2M3FNType::getTypeID()) {
+ return {"fp6e2m3"};
+ } else if (typeInfo.typeID == mlir::Float6E3M2FNType::getTypeID()) {
+ return {"fp6e3m2"};
+ } else if (typeInfo.typeID == mlir::Float4E2M1FNType::getTypeID()) {
+ return {"fp4e2m1"};
+ } else if (typeInfo.typeID == mlir::Float8E8M0FNUType::getTypeID()) {
+ return {"fp8e8m0"};
+ } else if (typeInfo.typeID == tosa::mxint8Type::getTypeID()) {
+ return {"mxint8"};
+ }
+ llvm_unreachable("unknown type");
+ };
+
+ if (typeInfo.typeID == tosa::BlockScaledType::getTypeID()) {
+ TypeInfo valueInfo = {typeInfo.valueTypeID, typeInfo.bitWidth};
+ TypeInfo scaleInfo = {typeInfo.scaleTypeID, 8};
+ llvm::SmallString<32> result;
+ llvm::raw_svector_ostream os(result);
+ os << "bs" << typeInfo.blockSize << "_"
+ << stringifyScalarTypeInfo(scaleInfo) << "_"
+ << stringifyScalarTypeInfo(valueInfo);
+ return result;
}
- llvm_unreachable("unknown type");
+
+ return stringifyScalarTypeInfo(typeInfo);
}
diff --git a/mlir/lib/Dialect/Tosa/Transforms/TosaValidation.cpp b/mlir/lib/Dialect/Tosa/Transforms/TosaValidation.cpp
index 34ac84d237f45..8fa2ddf856c7d 100644
--- a/mlir/lib/Dialect/Tosa/Transforms/TosaValidation.cpp
+++ b/mlir/lib/Dialect/Tosa/Transforms/TosaValidation.cpp
@@ -1482,7 +1482,7 @@ bool TosaValidation::isValidElementType(Type type, const bool allowUnsigned) {
}
} else if (isa<tosa::shapeType>(type))
return true;
- else if (isa<tosa::mxint8Type>(type))
+ else if (isa<tosa::mxint8Type, tosa::BlockScaledType>(type))
return true;
return false;
}
diff --git a/mlir/test/Dialect/Tosa/availability.mlir b/mlir/test/Dialect/Tosa/availability.mlir
index 81294f2c0c308..450b4556c7b16 100644
--- a/mlir/test/Dialect/Tosa/availability.mlir
+++ b/mlir/test/Dialect/Tosa/availability.mlir
@@ -627,10 +627,10 @@ func.func @test_resize(%arg0: tensor<1x32x32x8xf32>) -> tensor<1x64x64x8xf32> {
}
// -----
-// CHECK-LABEL: cast
-func.func @test_cast1(%arg0: tensor<13x21x3xi32>) -> tensor<13x21x3xf32> {
+// CHECK-LABEL: test_cast
+func.func @test_cast(%arg0: tensor<13x21x3xi32>) -> tensor<13x21x3xf32> {
// CHECK: profiles: [ [pro_int, pro_fp] ]
- // CHECK: extensions: [ [fp8e4m3, fp8e5m2, bf16, int64] ]
+ // CHECK: extensions: [ [fp8e4m3, fp8e5m2, bf16, int64, mx_common, mx_fp4e2m1, mx_fp6e2m3, mx_fp6e3m2, mx_fp8e4m3, mx_fp8e5m2, mx_int8] ]
%0 = tosa.cast %arg0 : (tensor<13x21x3xi32>) -> tensor<13x21x3xf32>
return %0 : tensor<13x21x3xf32>
}
diff --git a/mlir/test/Dialect/Tosa/canonicalize.mlir b/mlir/test/Dialect/Tosa/canonicalize.mlir
index 2cd040f056db8..d454ca831cf0d 100644
--- a/mlir/test/Dialect/Tosa/canonicalize.mlir
+++ b/mlir/test/Dialect/Tosa/canonicalize.mlir
@@ -1651,6 +1651,73 @@ func.func @test_canonicalize_non_narrowing_cast_i8_to_f8E4M3FN_unsupported(%arg0
// -----
+// CHECK-LABEL: @test_canonicalize_cast_from_cast_to_block_scaled_type_f4E2M1
+// CHECK: return %arg0 : tensor<15x3x2x256x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+func.func @test_canonicalize_cast_from_cast_to_block_scaled_type_f4E2M1(%arg0: tensor<15x3x2x256x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<15x3x2x256x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = tosa.cast %arg0 : (tensor<15x3x2x256x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<15x3x2x256xf32>
+ %1 = tosa.cast %0 : (tensor<15x3x2x256xf32>) -> tensor<15x3x2x256x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %1 : tensor<15x3x2x256x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: @test_canonicalize_cast_from_cast_to_block_scaled_type_f8E5M2
+// CHECK: return %arg0 : tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+func.func @test_canonicalize_cast_from_cast_to_block_scaled_type_f8E5M2(%arg0: tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = tosa.cast %arg0 : (tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<160xf32>
+ %1 = tosa.cast %0 : (tensor<160xf32>) -> tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %1 : tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_different_types_f8E5M2_f6E2M3
+// CHECK: %[[values:.+]] = tosa.cast %arg0
+// CHECK: %[[block_scaled:.+]] = tosa.cast %[[values]]
+// CHECK: return %[[block_scaled]] : tensor<160x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+func.func @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_different_types_f8E5M2_f6E2M3(%arg0: tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<160x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = tosa.cast %arg0 : (tensor<160x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<160xf32>
+ %1 = tosa.cast %0 : (tensor<160xf32>) -> tensor<160x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %1 : tensor<160x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_different_types_f6E2M3_f6E3M2
+// CHECK: %[[values:.+]] = tosa.cast %arg0
+// CHECK: %[[block_scaled:.+]] = tosa.cast %[[values]]
+// CHECK: return %[[block_scaled]] : tensor<32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+func.func @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_different_types_f6E2M3_f6E3M2(%arg0: tensor<32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = tosa.cast %arg0 : (tensor<32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<32xf32>
+ %1 = tosa.cast %0 : (tensor<32xf32>) -> tensor<32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %1 : tensor<32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_unranked
+// CHECK: %[[values:.+]] = tosa.cast %arg0
+// CHECK: %[[block_scaled:.+]] = tosa.cast %[[values]]
+// CHECK: return %[[block_scaled]] : tensor<*x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+func.func @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_unranked(%arg0: tensor<3x64x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<*x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = tosa.cast %arg0 : (tensor<3x64x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<*xf32>
+ %1 = tosa.cast %0 : (tensor<*xf32>) -> tensor<*x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %1 : tensor<*x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_f8E5M2_f8E4M3
+// CHECK: tosa.cast
+// CHECK: tosa.cast
+func.func @test_do_not_canonicalize_cast_from_cast_to_block_scaled_type_f8E5M2_f8E4M3(%arg0: tensor<15x3x2x256x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<15x3x2x256x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = tosa.cast %arg0 : (tensor<15x3x2x256x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<15x3x2x256xf8E4M3FN>
+ %1 = tosa.cast %0 : (tensor<15x3x2x256xf8E4M3FN>) -> tensor<15x3x2x256x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %1 : tensor<15x3x2x256x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
// CHECK-LABEL: @test_canonicalize_cast_from_cast_to_block_scaled_f4E2M1
// CHECK: return %arg0, %arg1 : tensor<15x3x2x256xf4E2M1FN>, tensor<15x3x2x8xf8E8M0FNU>
func.func @test_canonicalize_cast_from_cast_to_block_scaled_f4E2M1(%arg0: tensor<15x3x2x256xf4E2M1FN>, %arg1: tensor<15x3x2x8xf8E8M0FNU>) -> (tensor<15x3x2x256xf4E2M1FN>, tensor<15x3x2x8xf8E8M0FNU>) {
diff --git a/mlir/test/Dialect/Tosa/invalid.mlir b/mlir/test/Dialect/Tosa/invalid.mlir
index 5e8111061cb3a..eca5c9ef1b14e 100644
--- a/mlir/test/Dialect/Tosa/invalid.mlir
+++ b/mlir/test/Dialect/Tosa/invalid.mlir
@@ -2266,3 +2266,35 @@ func.func @test_shape_func_output() -> !tosa.shape<4> {
%cst = tosa.const_shape {values = dense<[1, 2, 3, 4]> : tensor<4xindex>} : () -> !tosa.shape<4>
return %cst : !tosa.shape<4>
}
+
+// -----
+
+func.func @test_cast_f32_plain_fp4(%arg0: tensor<4x32xf32>) -> tensor<4x32xf4E2M1FN> {
+ // expected-error at +1 {{'tosa.cast' op illegal: operation operand/result data types did not align with any profile or extension, got (f32,fp4e2m1)}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xf32>) -> tensor<4x32xf4E2M1FN>
+ return %0 : tensor<4x32xf4E2M1FN>
+}
+
+// -----
+
+func.func @test_cast_fp4_block_scaled(%arg0: tensor<4x32xf4E2M1FN>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op illegal: operation operand/result data types did not align with any profile or extension, got (fp4e2m1,bs32_fp8e8m0_fp4e2m1), did you mean (fp8e4m3,bs32_fp8e8m0_fp4e2m1)? Otherwise, please refer to the 'supported data types' for 'tosa.cast' in the specification.}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xf4E2M1FN>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_cast_block_scaled_fp6e2m3(%arg0: tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf6E2M3FN> {
+ // expected-error at +1 {{'tosa.cast' op illegal: operation operand/result data types did not align with any profile or extension, got (bs32_fp8e8m0_fp4e2m1,fp6e2m3), did you mean (bs32_fp8e8m0_fp4e2m1,fp8e4m3)? Otherwise, please refer to the 'supported data types' for 'tosa.cast' in the specification.}}
+ %0 = tosa.cast %arg0 : (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf6E2M3FN>
+ return %0 : tensor<4x32xf6E2M3FN>
+}
+
+// -----
+
+func.func @test_cast_fp6e3m2_block_scaled(%arg0: tensor<4x32xf6E3M2FN>) -> tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op illegal: operation operand/result data types did not align with any profile or extension, got (fp6e3m2,bs32_fp8e8m0_mxint8), did you mean (fp8e4m3,bs32_fp8e8m0_mxint8)? Otherwise, please refer to the 'supported data types' for 'tosa.cast' in the specification.}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xf6E3M2FN>) -> tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
diff --git a/mlir/test/Dialect/Tosa/invalid_extension.mlir b/mlir/test/Dialect/Tosa/invalid_extension.mlir
index 295d7172bc2c4..9d4e59a9a9c6f 100644
--- a/mlir/test/Dialect/Tosa/invalid_extension.mlir
+++ b/mlir/test/Dialect/Tosa/invalid_extension.mlir
@@ -295,6 +295,34 @@ func.func @test_cast_f32_bf16(%arg0: tensor<13x21x3xf32>) -> tensor<13x21x3xbf16
return %0 : tensor<13x21x3xbf16>
}
+// -----
+func.func @test_cast_f32_block_scaled(%arg0: tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op illegal: requires all of [mx_common, mx_fp4e2m1] profiles/extensions to be specified in the target environment}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+func.func @test_cast_block_scaled_f32(%arg0: tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf32> {
+ // expected-error at +1 {{'tosa.cast' op illegal: requires all of [mx_common, mx_fp4e2m1] profiles/extensions to be specified in the target environment}}
+ %0 = tosa.cast %arg0 : (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf32>
+ return %0 : tensor<4x32xf32>
+}
+
+// -----
+func.func @test_cast_bf16_block_scaled(%arg0: tensor<4x32xbf16>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op illegal: requires all of [bf16, mx_common, mx_fp4e2m1] profiles/extensions to be specified in the target environment}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xbf16>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+func.func @test_cast_fp8_block_scaled(%arg0: tensor<4x32xf8E4M3FN>) -> tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op illegal: requires all of [fp8e4m3, mx_common, mx_int8] profiles/extensions to be specified in the target environment}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xf8E4M3FN>) -> tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
// -----
func.func @test_fft2d(%arg0: tensor<1x4x8xf32>, %arg1: tensor<1x4x8xf32>) -> (tensor<1x4x8xf32>, tensor<1x4x8xf32>) {
// expected-error at +1 {{'tosa.fft2d' op illegal: requires any of [fft] 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 c7f4c67072d06..86212177d34be 100644
--- a/mlir/test/Dialect/Tosa/ops.mlir
+++ b/mlir/test/Dialect/Tosa/ops.mlir
@@ -1059,6 +1059,32 @@ func.func @test_cast3(%arg0: tensor<13x21x3xi32>) -> tensor<13x21x3x!quant.unifo
return %0 : tensor<13x21x3x!quant.uniform<i16:f32, 0.078431375324726104:128>>
}
+// -----
+// CHECK-LABEL: test_cast_to_block_scaled
+func.func @test_cast_to_block_scaled(%arg0: tensor<4x32xf32>, %arg1: tensor<4x32xbf16>, %arg2: tensor<4x32xf8E4M3FN>) -> (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) {
+ %0 = tosa.cast %arg0 : (tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %1 = tosa.cast %arg1 : (tensor<4x32xbf16>) -> tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %2 = tosa.cast %arg2 : (tensor<4x32xf8E4M3FN>) -> tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0, %1, %2 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+// CHECK-LABEL: test_cast_from_block_scaled
+func.func @test_cast_from_block_scaled(%arg0: tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, %arg1: tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, %arg2: tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) -> (tensor<4x32xf32>, tensor<4x32xbf16>, tensor<4x32xf8E5M2>) {
+ %0 = tosa.cast %arg0 : (tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf32>
+ %1 = tosa.cast %arg1 : (tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xbf16>
+ %2 = tosa.cast %arg2 : (tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf8E5M2>
+ return %0, %1, %2 : tensor<4x32xf32>, tensor<4x32xbf16>, tensor<4x32xf8E5M2>
+}
+
+// -----
+// CHECK-LABEL: test_cast_block_scaled_dynamic
+func.func @test_cast_block_scaled_dynamic(%arg0: tensor<?x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, %arg1: tensor<4x?x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> (tensor<?x32xf32>, tensor<4x?xf32>) {
+ %0 = tosa.cast %arg0 : (tensor<?x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<?x32xf32>
+ %1 = tosa.cast %arg1 : (tensor<4x?x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x?xf32>
+ return %0, %1 : tensor<?x32xf32>, tensor<4x?xf32>
+}
+
// -----
// CHECK-LABEL: rescale
func.func @test_rescale(%arg0: tensor<13x21x3x!quant.uniform<u8:f32, 0.015655439347028732:127>>) -> tensor<13x21x3x!quant.uniform<i8:f32, 0.015655439347028732:-1>> {
diff --git a/mlir/test/Dialect/Tosa/tosa-attach-target.mlir b/mlir/test/Dialect/Tosa/tosa-attach-target.mlir
index a0c59c0c4bb3b..558a599fbbc1e 100644
--- a/mlir/test/Dialect/Tosa/tosa-attach-target.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-attach-target.mlir
@@ -1,11 +1,26 @@
-// RUN: mlir-opt %s -split-input-file -tosa-attach-target="profiles=pro_int,pro_fp extensions=int16,int4,bf16,fp8e4m3,fp8e5m2,fft,variable,controlflow,doubleround,inexactround,dynamic level=none" | FileCheck %s --check-prefix=CHECK-ALL
+// DEFINE: %{core_extensions} = int16,int4,bf16,fp8e4m3,fp8e5m2,fft,variable
+// DEFINE: %{rounding_extensions} = controlflow,doubleround,inexactround,dynamic
+// DEFINE: %{mx_float_extensions} = mx_common,mx_fp4e2m1,mx_fp6e2m3,mx_fp6e3m2
+// DEFINE: %{mx_int_extensions} = mx_fp8e4m3,mx_fp8e5m2,mx_int8
+// DEFINE: %{all_extensions} = %{core_extensions},%{rounding_extensions},%{mx_float_extensions},%{mx_int_extensions}
+// DEFINE: %{all_target} = specification_version=1.1.draft level=none \
+// DEFINE: profiles=pro_int,pro_fp extensions=%{all_extensions}
+
+// RUN: mlir-opt %s -split-input-file -tosa-attach-target="%{all_target}" | FileCheck %s --check-prefix=CHECK-ALL
// RUN: mlir-opt %s -split-input-file -tosa-attach-target="level=8k" | FileCheck %s --check-prefix=CHECK-LVL-8K
// RUN: mlir-opt %s -split-input-file -tosa-attach-target | FileCheck %s --check-prefix=CHECK-DEFAULT
// RUN: mlir-opt %s -split-input-file -tosa-attach-target="specification_version=1.1.draft" | FileCheck %s --check-prefix=CHECK-VERSION-1P1
// -----
-// CHECK-ALL: module attributes {tosa.target_env = #tosa.target_env<specification_version = "1.0", level = none, profiles = [pro_int, pro_fp], extensions = [int16, int4, bf16, fp8e4m3, fp8e5m2, fft, variable, controlflow, doubleround, inexactround, dynamic]>}
+// CHECK-ALL: module attributes {
+// CHECK-ALL-SAME: tosa.target_env = #tosa.target_env<specification_version = "1.1.draft",
+// CHECK-ALL-SAME: level = none,
+// CHECK-ALL-SAME: profiles = [pro_int, pro_fp],
+// CHECK-ALL-SAME: extensions = [int16, int4, bf16, fp8e4m3, fp8e5m2, fft,
+// CHECK-ALL-SAME: variable, controlflow, doubleround, inexactround, dynamic,
+// CHECK-ALL-SAME: mx_common, mx_fp4e2m1, mx_fp6e2m3, mx_fp6e3m2,
+// CHECK-ALL-SAME: mx_fp8e4m3, mx_fp8e5m2, mx_int8]>}
// CHECK-LVL-8K: module attributes {tosa.target_env = #tosa.target_env<specification_version = "1.0", level = "8k", profiles = [], extensions = []>}
// CHECK-DEFAULT: module attributes {tosa.target_env = #tosa.target_env<specification_version = "1.0", level = "8k", profiles = [], extensions = []>}
// CHECK-VERSION-1P1: module attributes {tosa.target_env = #tosa.target_env<specification_version = "1.1.draft", level = "8k", profiles = [], extensions = []>}
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 7ff883e8e5431..c7bd051399ae8 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
@@ -135,6 +135,14 @@ func.func @test_cast_i64_bool(%arg0: tensor<13x21x3xi64>) -> tensor<13x21x3xi1>
// -----
+func.func @test_cast_fp32_block_scaled(%arg0: tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
func.func @test_dyanmic_dims(%arg0: tensor<?x8x16xi8>) -> tensor<?x16xi32> {
// expected-error at +1 {{'tosa.argmax' op failed level check: operand shape dimension cannot be dynamic when targeting TOSA specification version 1.0 or below}}
%0 = tosa.argmax %arg0 { axis = 1 : i32 } : (tensor<?x8x16xi8>) -> tensor<?x16xi32>
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 1bac8bacbaf40..f45ba0a5b7ef5 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
@@ -1,4 +1,4 @@
-// RUN: mlir-opt %s -split-input-file -verify-diagnostics -tosa-attach-target="specification_version=1.1.draft profiles=pro_int,pro_fp extensions=int16,int4,bf16,fp8e4m3,fp8e5m2,fft,variable,controlflow,doubleround,inexactround,mxfp,int64,mxfp_conv,shape" -tosa-validate="strict-op-spec-alignment" | FileCheck %s
+// RUN: mlir-opt %s -split-input-file -verify-diagnostics -tosa-attach-target="specification_version=1.1.draft profiles=pro_int,pro_fp extensions=int16,int4,bf16,fp8e4m3,fp8e5m2,fft,variable,controlflow,doubleround,inexactround,mxfp,int64,mxfp_conv,shape,mx_common,mx_fp4e2m1,mx_fp6e2m3,mx_fp6e3m2,mx_fp8e4m3,mx_fp8e5m2,mx_int8" -tosa-validate="strict-op-spec-alignment" | FileCheck %s
// -----
@@ -352,6 +352,34 @@ func.func @test_cast_i64_bool(%arg0: tensor<13x21x3xi64>) -> tensor<13x21x3xi1>
// -----
+// CHECK-LABEL: test_cast_to_block_scaled_types
+func.func @test_cast_to_block_scaled_types(%fp16: tensor<4x32xf16>, %fp32: tensor<4x32xf32>, %bf16: tensor<4x32xbf16>, %fp8e4m3: tensor<4x32xf8E4M3FN>, %fp8e5m2: tensor<4x32xf8E5M2>) -> (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) {
+ %0 = tosa.cast %fp32 : (tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %1 = tosa.cast %fp32 : (tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %2 = tosa.cast %fp32 : (tensor<4x32xf32>) -> tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %3 = tosa.cast %fp16 : (tensor<4x32xf16>) -> tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %4 = tosa.cast %bf16 : (tensor<4x32xbf16>) -> tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %5 = tosa.cast %fp8e4m3 : (tensor<4x32xf8E4M3FN>) -> tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %6 = tosa.cast %fp8e5m2 : (tensor<4x32xf8E5M2>) -> tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0, %1, %2, %3, %4, %5 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_cast_from_block_scaled_types
+func.func @test_cast_from_block_scaled_types(%fp4: tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, %fp6e2m3: tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, %fp6e3m2: tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, %fp8e4m3: tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, %fp8e5m2: tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, %mxint8: tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) -> (tensor<4x32xf32>, tensor<4x32xf16>, tensor<4x32xbf16>, tensor<4x32xf8E4M3FN>, tensor<4x32xf8E5M2>) {
+ %0 = tosa.cast %fp4 : (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf32>
+ %1 = tosa.cast %fp6e2m3 : (tensor<4x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf32>
+ %2 = tosa.cast %fp6e3m2 : (tensor<4x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf32>
+ %3 = tosa.cast %fp8e4m3 : (tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf16>
+ %4 = tosa.cast %fp8e5m2 : (tensor<4x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xbf16>
+ %5 = tosa.cast %mxint8 : (tensor<4x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf8E4M3FN>
+ %6 = tosa.cast %fp4 : (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xf8E5M2>
+ return %0, %3, %4, %5, %6 : tensor<4x32xf32>, tensor<4x32xf16>, tensor<4x32xbf16>, tensor<4x32xf8E4M3FN>, tensor<4x32xf8E5M2>
+}
+
+// -----
+
// CHECK-LABEL: test_dynamic_dims
func.func @test_dynamic_dims(%arg0: tensor<?x8x16xi8>) -> tensor<?x16xi32> {
%0 = tosa.argmax %arg0 { axis = 1 : i32 } : (tensor<?x8x16xi8>) -> tensor<?x16xi32>
diff --git a/mlir/test/Dialect/Tosa/verifier.mlir b/mlir/test/Dialect/Tosa/verifier.mlir
index 9b5faf575971b..14a48d5e10ea0 100644
--- a/mlir/test/Dialect/Tosa/verifier.mlir
+++ b/mlir/test/Dialect/Tosa/verifier.mlir
@@ -1517,6 +1517,48 @@ func.func @test_cast_to_block_scaled_block_size_mismatch(%arg0: tensor<4x32xf32>
// -----
+func.func @test_cast_i8_block_scaled(%arg0: tensor<4x32xi8>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op requires non-block-scaled element type to be floating-point when casting to or from block scaled element type, got 'i8'}}
+ %0 = tosa.cast %arg0 : (tensor<4x32xi8>) -> tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_cast_block_scaled_i32(%arg0: tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xi32> {
+ // expected-error at +1 {{'tosa.cast' op requires non-block-scaled element type to be floating-point when casting to or from block scaled element type, got 'i32'}}
+ %0 = tosa.cast %arg0 : (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32xi32>
+ return %0 : tensor<4x32xi32>
+}
+
+// -----
+
+func.func @test_cast_between_block_scaled(%arg0: tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{'tosa.cast' op requires exactly one of input or output to have block scaled element type}}
+ %0 = tosa.cast %arg0 : (tensor<4x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<4x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_block_scaled_cast_invalid_block_shape(%arg0: tensor<1x16x31x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<1x16x31xf32> {
+ // expected-error at +2 {{tensor type 'tensor<1x16x31x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>' blocked dimension must be a multiple of block size, got 31 and block size 32}}
+ // expected-error at +1 {{'tosa.cast' op operand #0 must be tosa-conformant tensor of number values, but got 'tensor<1x16x31x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>'}}
+ %0 = tosa.cast %arg0 : (tensor<1x16x31x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<1x16x31xf32>
+ return %0 : tensor<1x16x31xf32>
+}
+
+// -----
+
+func.func @test_block_scaled_cast_scalar(%arg0: tensor<!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<f32> {
+ // expected-error at +2 {{tensor type 'tensor<!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>' does not support block scaling on scalar tensors}}
+ // expected-error at +1 {{'tosa.cast' op operand #0 must be tosa-conformant tensor of number values, but got 'tensor<!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>'}}
+ %0 = tosa.cast %arg0 : (tensor<!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<f32>
+ return %0 : tensor<f32>
+}
+
+// -----
+
func.func @test_clamp_quantized(%arg0:tensor<?x112x112x32x!quant.uniform<u8:f32, 0.023529412224888802:-128>>) -> (tensor<?x112x112x32x!quant.uniform<u8:f32, 0.023529412224888802:-128>>) {
// expected-error at +1 {{'tosa.clamp' op min/max attributes types are incompatible with input/output element types.}}
%0 = tosa.clamp %arg0 {max_val = 127 : i8, min_val = -128 : i8} : (tensor<?x112x112x32x!quant.uniform<u8:f32, 0.023529412224888802:-128>>) -> tensor<?x112x112x32x!quant.uniform<u8:f32, 0.023529412224888802:-128>>
>From 0e04404380f3d33127f411645b7240079d6ef584 Mon Sep 17 00:00:00 2001
From: Luke Hutton <luke.hutton at arm.com>
Date: Mon, 1 Jun 2026 20:38:58 +0100
Subject: [PATCH 2/2] [mlir][tosa] Add constant block scaled support
This commit adds support for block scaled tensors.
In particular, the block scaled type has been extended with the
`DenseElementTypeInterface` to allow block scaled data values to
be specified in dense element attributes.
The block scaled type has also been extended to allow optional scale
values to be specified by the type. For now, scale values are not
expected to be propagated beyond their use in the attribute input of
a constant operation. In the future, we may want to propagate these
values to allow certain optimizations.
The `tosa.const` operation has also been updated in the validation pass.
Change-Id: I879f9b8742b784fda976a879813e68b928f6b77e
---
.../Dialect/Tosa/IR/TosaComplianceData.h.inc | 20 ++-
mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td | 10 +-
.../mlir/Dialect/Tosa/IR/TosaTypesBase.td | 22 ++-
mlir/lib/Dialect/Tosa/IR/TosaOps.cpp | 149 ++++++++++++++++--
mlir/test/Dialect/Tosa/availability.mlir | 2 +-
mlir/test/Dialect/Tosa/ops.mlir | 72 +++++++++
.../tosa-validation-version-1p0-invalid.mlir | 18 +++
.../tosa-validation-version-1p1-valid.mlir | 13 ++
mlir/test/Dialect/Tosa/verifier.mlir | 69 ++++++++
9 files changed, 353 insertions(+), 22 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
index 26890abd187bd..25fec58e9d980 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
@@ -1191,7 +1191,25 @@ extensionComplianceMap = {
{{fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT},
{{fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT},
{{fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT},
- {{mxint8T}, SpecificationVersion::V_1_1_DRAFT}}}}},
+ {{mxint8T}, SpecificationVersion::V_1_1_DRAFT}}},
+ {{Extension::mx_common, Extension::mx_fp8e4m3},
+ {{{bs32_fp8ue8m0_fp8e4m3T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp8e5m2},
+ {{{bs32_fp8ue8m0_fp8e5m2T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e3m2},
+ {{{bs32_fp8ue8m0_fp6e3m2T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp6e2m3},
+ {{{bs32_fp8ue8m0_fp6e2m3T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_fp4e2m1},
+ {{{bs32_fp8ue8m0_fp4e2m1T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf},
+ {{Extension::mx_common, Extension::mx_int8},
+ {{{bs32_fp8ue8m0_mxint8T}, SpecificationVersion::V_1_1_DRAFT}},
+ allOf}}},
{"tosa.identity",
{{{Extension::int4}, {{{i4T, i4T}, SpecificationVersion::V_1_0}}},
{{Extension::int16}, {{{i48T, i48T}, SpecificationVersion::V_1_0}}},
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
index 32f8bf08cbc84..eda3279cd0534 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
@@ -2895,9 +2895,13 @@ def Tosa_ConstOp : Tosa_Op<"const", [ConstantLike, Pure,
Tosa_Tensor:$output
);
- 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, Tosa_EXT_MXFP, Tosa_EXT_INT64]>,
+ 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, Tosa_EXT_MXFP,
+ Tosa_EXT_INT64, 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 hasFolder = 1;
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
index 0a43f14865f56..79d716cb2229c 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
@@ -109,25 +109,39 @@ def Tosa_MXFPValue
def Tosa_MXFPScale
: AnyTypeOf<[F8E8M0FNU], "micro-scaling format scale number">;
-def Tosa_BlockScaled : Tosa_Type<"BlockScaled", "block_scaled"> {
+def Tosa_BlockScaled
+ : Tosa_Type<"BlockScaled", "block_scaled",
+ [DeclareTypeInterfaceMethods<
+ DenseElementTypeInterface, ["getDenseElementBitSize",
+ "convertToAttribute",
+ "convertFromAttribute"]>]> {
let summary = "Block scaled tensor element type.";
let description = [{
This does not specify an implementation type. A tensor of this type is a
tensor which has block-scaled quantization applied.
- This compound type is made up of 3 components:
+ This compound type is made up of 3 concrete components and 1 optional:
`value_type` - The type of the data values in each block.
`scale_type` - The type of the scale value associated with each block.
`block_shape` - The size and axis of each block. Currently only supports
specifying the block size along the innermost dimension.
+ `scale_values` - Optional array of per-block scale values. If not provided,
+ the type is assumed to be dynamically quantized at runtime.
}];
let parameters = (ins Tosa_MXFPValue:$value_type, Tosa_MXFPScale:$scale_type,
- EnumParameter<Tosa_BlockShape>:$block_shape);
+ EnumParameter<Tosa_BlockShape>:$block_shape,
+ OptionalArrayRefParameter<"Attribute">:$scale_values);
let assemblyFormat =
- "`<` $value_type```:```$scale_type```:```$block_shape `>`";
+ "`<` $value_type```:```$scale_type```:```$block_shape "
+ "(`,` `{` custom<ScaleValues>($scale_values, ref($scale_type))^ `}`)? "
+ "`>`";
+
+ let extraClassDeclaration = [{
+ bool hasScaleValues() const { return !getScaleValues().empty(); }
+ }];
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
index 486431a562714..aaa65601068f5 100644
--- a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
+++ b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
@@ -20,6 +20,7 @@
#include "mlir/Dialect/Tosa/Utils/ShapeUtils.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
#include "mlir/Dialect/Utils/VerificationUtils.h"
+#include "mlir/IR/BuiltinTypeInterfaces.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/DialectImplementation.h"
#include "mlir/IR/Matchers.h"
@@ -740,14 +741,55 @@ LogicalResult mlir::tosa::mxint8Type::convertFromAttribute(
// TOSA block scaling utilities.
//===----------------------------------------------------------------------===//
-LogicalResult OpTrait::tosa::verifyBlockScaledTensorType(Operation &op,
- mlir::Type type) {
+static ParseResult parseScaleValues(AsmParser &parser,
+ SmallVector<Attribute> &scaleValues,
+ Type scaleType) {
+ const auto parseScaleValue = [&]() -> ParseResult {
+ const SMLoc loc = parser.getCurrentLocation();
+
+ double floatValue;
+ if (parser.parseFloat(floatValue))
+ return failure();
+
+ if (floatValue < 0.0)
+ return parser.emitError(loc, "scale value must be non-negative, got ")
+ << floatValue;
+
+ Type attrType = scaleType;
+ if (succeeded(parser.parseOptionalColon()) && parser.parseType(attrType))
+ return failure();
+
+ if (attrType != scaleType)
+ return parser.emitError(loc, "parsed attribute type ")
+ << attrType << " does not match expected scale type " << scaleType;
+
+ scaleValues.push_back(FloatAttr::get(attrType, floatValue));
+ return success();
+ };
+
+ return parser.parseCommaSeparatedList(parseScaleValue);
+}
+
+static void printScaleValues(AsmPrinter &printer,
+ ArrayRef<Attribute> scaleValues, Type) {
+ llvm::interleaveComma(scaleValues, printer, [&](Attribute scaleValue) {
+ printer.printAttributeWithoutType(scaleValue);
+ });
+}
+
+static LogicalResult verifyBlockScaledTensorType(Operation &op, mlir::Type type,
+ bool allowScaleValues) {
const auto tensorType = llvm::cast<ShapedType>(type);
const BlockScaledType elemType =
llvm::dyn_cast<BlockScaledType>(tensorType.getElementType());
if (!elemType)
return success();
+ if (!allowScaleValues && elemType.hasScaleValues())
+ return op.emitError()
+ << "tensor type " << type
+ << " does not support scale values for this operation";
+
if (!tensorType.hasRank())
return success();
@@ -756,12 +798,24 @@ LogicalResult OpTrait::tosa::verifyBlockScaledTensorType(Operation &op,
<< "tensor type " << type
<< " does not support block scaling on scalar tensors";
- const int64_t blockedDimension = tensorType.getShape().back();
- if (ShapedType::isDynamic(blockedDimension))
- return success();
-
+ const ArrayRef<int64_t> tensorShape = tensorType.getShape();
const uint32_t blockSize =
BlockShapeAttr::getBlockShapeValue(elemType.getBlockShape());
+
+ if (allowScaleValues && elemType.hasScaleValues() &&
+ tensorType.hasStaticShape()) {
+ const size_t numBlocks =
+ llvm::accumulate(tensorShape, int64_t(1), std::multiplies<int64_t>());
+ if (elemType.getScaleValues().size() != numBlocks / blockSize)
+ return op.emitError()
+ << "tensor type " << type << " has " << numBlocks / blockSize
+ << " blocks but got " << elemType.getScaleValues().size()
+ << " scale values";
+ }
+
+ const int64_t blockedDimension = tensorShape.back();
+ if (ShapedType::isDynamic(blockedDimension))
+ return success();
if (blockedDimension % blockSize != 0)
return op.emitError()
<< "tensor type " << type
@@ -771,6 +825,43 @@ LogicalResult OpTrait::tosa::verifyBlockScaledTensorType(Operation &op,
return success();
}
+LogicalResult OpTrait::tosa::verifyBlockScaledTensorType(Operation &op,
+ mlir::Type type) {
+ return ::verifyBlockScaledTensorType(op, type, /*allowScaleValues=*/false);
+}
+
+size_t mlir::tosa::BlockScaledType::getDenseElementBitSize() const {
+ const Type valueType = getValueType();
+ if (isa<tosa::mxint8Type>(valueType))
+ return 8;
+ return valueType.getIntOrFloatBitWidth();
+}
+
+Attribute
+mlir::tosa::BlockScaledType::convertToAttribute(ArrayRef<char> rawData) const {
+ assert(rawData.size() == 1 && "expected 1 byte for block_scaled element");
+ const Type valueType = getValueType();
+ if (const auto mxint8Value = dyn_cast<tosa::mxint8Type>(valueType))
+ return mxint8Value.convertToAttribute(rawData);
+ if (!isa<FloatType>(valueType))
+ return {};
+ return mlir::detail::convertFloatTypeToAttribute(valueType, rawData);
+}
+
+LogicalResult mlir::tosa::BlockScaledType::convertFromAttribute(
+ Attribute attr, SmallVectorImpl<char> &result) const {
+ const Type valueType = getValueType();
+ if (const auto mxint8Value = dyn_cast<tosa::mxint8Type>(valueType))
+ return mxint8Value.convertFromAttribute(attr, result);
+
+ const auto floatAttr = dyn_cast<FloatAttr>(attr);
+ if (!floatAttr || floatAttr.getType() != valueType)
+ return failure();
+ const APFloat value = floatAttr.getValue();
+ return mlir::detail::convertFloatTypeFromAttribute(
+ valueType, FloatAttr::get(valueType, value), result);
+}
+
//===----------------------------------------------------------------------===//
// TOSA Operator Verifiers.
//===----------------------------------------------------------------------===//
@@ -857,7 +948,7 @@ static LogicalResult verifyConvOp(T op) {
}
LogicalResult tosa::ConstOp::verify() {
-
+ Operation &op = *getOperation();
auto attrType = llvm::dyn_cast<TensorType>(getValuesAttr().getType());
auto outputType = llvm::dyn_cast<TensorType>(getOutput().getType());
@@ -866,17 +957,49 @@ LogicalResult tosa::ConstOp::verify() {
return failure();
}
- if (auto result = llvm::dyn_cast<mlir::quant::QuantizedType>(
- outputType.getElementType())) {
- if (getStorageElementTypeFromQuantized(result) == attrType.getElementType())
+ const Type attrElemType = attrType.getElementType();
+ const Type resultElemType = outputType.getElementType();
+
+ if (auto result =
+ llvm::dyn_cast<mlir::quant::QuantizedType>(resultElemType)) {
+ if (getStorageElementTypeFromQuantized(result) == attrElemType)
return success();
}
- if (attrType.getElementType() != outputType.getElementType()) {
- emitOpError("expected same attr/result element types");
- return failure();
+ if (auto attrBlockScaledType =
+ llvm::dyn_cast<mlir::tosa::BlockScaledType>(attrElemType)) {
+ if (failed(verifyBlockScaledTensorType(op, attrType, true)) ||
+ failed(verifyBlockScaledTensorType(op, outputType, false)))
+ return failure();
+
+ if (!attrBlockScaledType.hasScaleValues())
+ return op.emitOpError(
+ "attribute block scaled type must have scale values");
+
+ const BlockScaledType resultBlockScaledType =
+ llvm::dyn_cast<mlir::tosa::BlockScaledType>(resultElemType);
+ if (!resultBlockScaledType)
+ return op.emitOpError(
+ "result type must be block scaled type if attribute is block "
+ "scaled type");
+
+ if (attrBlockScaledType.getValueType() !=
+ resultBlockScaledType.getValueType() ||
+ attrBlockScaledType.getScaleType() !=
+ resultBlockScaledType.getScaleType() ||
+ attrBlockScaledType.getBlockShape() !=
+ resultBlockScaledType.getBlockShape())
+ return op.emitOpError(
+ "expected block scaled element type to be compatible "
+ "between attr and result, got ")
+ << attrBlockScaledType << " vs. " << resultBlockScaledType;
+
+ return success();
}
+ if (attrElemType != resultElemType)
+ return emitOpError("expected same attr/result element types");
+
return success();
}
diff --git a/mlir/test/Dialect/Tosa/availability.mlir b/mlir/test/Dialect/Tosa/availability.mlir
index 450b4556c7b16..c2975b5ac40ac 100644
--- a/mlir/test/Dialect/Tosa/availability.mlir
+++ b/mlir/test/Dialect/Tosa/availability.mlir
@@ -650,7 +650,7 @@ func.func @test_rescale(%arg0: tensor<13x21x3x!quant.uniform<u8:f32, 0.015655439
// CHECK-LABEL: test_const
func.func @test_const(%arg0 : index) -> tensor<4xi32> {
// CHECK: profiles: [ [pro_int, pro_fp] ]
- // CHECK: extensions: [ [int4, int16, fp8e4m3, fp8e5m2, bf16, mxfp, int64] ]
+ // CHECK: extensions: [ [int4, int16, fp8e4m3, fp8e5m2, bf16, mxfp, int64, mx_common, mx_fp4e2m1, mx_fp6e2m3, mx_fp6e3m2, mx_fp8e4m3, mx_fp8e5m2, mx_int8] ]
%0 = "tosa.const"() {values = dense<[3, 0, 1, 2]> : tensor<4xi32>} : () -> tensor<4xi32>
return %0 : tensor<4xi32>
}
diff --git a/mlir/test/Dialect/Tosa/ops.mlir b/mlir/test/Dialect/Tosa/ops.mlir
index 86212177d34be..bdaefc3073d0a 100644
--- a/mlir/test/Dialect/Tosa/ops.mlir
+++ b/mlir/test/Dialect/Tosa/ops.mlir
@@ -1810,3 +1810,75 @@ func.func @test_assert_equal_shape() {
tosa.assert_equal_shape %0, %1 {allow_broadcast = true} : (!tosa.shape<2>, !tosa.shape<2>) -> ()
return
}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const
+func.func @test_block_scaled_const() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 1.0}>> : [[0.0 : f8E4M3FN, 1.0 : f8E4M3FN, 0.001953125 : f8E4M3FN, 0.0078125 : f8E4M3FN,
+ 2.0 : f8E4M3FN, 2.25 : f8E4M3FN, 2.5 : f8E4M3FN, 2.75 : f8E4M3FN,
+ 15.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN],
+ [0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN,
+ 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN, 0.0 : f8E4M3FN]]>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_splat
+func.func @test_block_scaled_const_splat() -> tensor<2x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 1.0}>> : 0.0 : f4E2M1FN>}> : () -> tensor<2x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_splat_mxint8
+func.func @test_block_scaled_const_splat_mxint8() -> tensor<2x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 1.0}>> : 0 : i8>}> : () -> tensor<2x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_scale_values
+func.func @test_block_scaled_const_scale_values() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 2.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_scale_values_explicit_type
+func.func @test_block_scaled_const_scale_values_explicit_type() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0 : f8E8M0FNU, 2.0 : f8E8M0FNU}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_scale_values_wide_inner_dim
+func.func @test_block_scaled_const_scale_values_wide_inner_dim() -> tensor<2x64x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x64x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 2.0, 4.0, 8.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x64x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x64x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_cast_scale_values_no_propagate
+func.func @test_block_scaled_const_cast_scale_values_no_propagate() -> tensor<2x32xf32> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.0, 4.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %1 = tosa.cast %0 : (tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>) -> tensor<2x32xf32>
+ return %1 : tensor<2x32xf32>
+}
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 c7bd051399ae8..873d22682219c 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
@@ -174,6 +174,24 @@ func.func @test_const_fp6e3m2(%arg0 : index) -> tensor<4xf6E3M2FN> {
// -----
+func.func @test_const_block_scaled_types() -> (tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) {
+ // expected-error at +1 {{'tosa.const' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %0 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f4E2M1FN>}> : () -> tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ // expected-error at +1 {{'tosa.const' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %1 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f6E2M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ // expected-error at +1 {{'tosa.const' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %2 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f6E3M2FN>}> : () -> tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ // expected-error at +1 {{'tosa.const' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %3 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ // expected-error at +1 {{'tosa.const' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %4 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f8E5M2>}> : () -> tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+ // expected-error at +1 {{'tosa.const' op illegal: requires specification version compatible with 1.1.draft (got 1.0)}}
+ %5 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0 : i8>}> : () -> tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0, %1, %2, %3, %4, %5 : tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
func.func @test_cast_from_block_scaled(%arg0: tensor<4x32xf8E5M2>, %arg1: tensor<4x1xf8E8M0FNU>) -> tensor<4x32xf32> {
// expected-error at +1 {{'tosa.cast_from_block_scaled' op illegal: requires specification version compatible with 1.1.draft (got 1.0) and requires any of [mxfp] profiles/extensions to be specified in the target environment}}
%0 = tosa.cast_from_block_scaled %arg0, %arg1 {block_size = #tosa.block_size<BLOCK_SIZE_32> : i32} : (tensor<4x32xf8E5M2>, tensor<4x1xf8E8M0FNU>) -> tensor<4x32xf32>
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 f45ba0a5b7ef5..6659ce8b9a3ba 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
@@ -142,6 +142,19 @@ func.func @test_const_mxint8() -> tensor<2x!tosa.mxint8> {
// -----
+// CHECK-LABEL: test_const_block_scaled_types
+func.func @test_const_block_scaled_types() -> (tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>) {
+ %0 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f4E2M1FN>}> : () -> tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %1 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f6E2M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %2 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f6E3M2FN>}> : () -> tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %3 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %4 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0.0 : f8E5M2>}> : () -> tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>
+ %5 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.0}>> : 0 : i8>}> : () -> tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0, %1, %2, %3, %4, %5 : tensor<1x32x!tosa.block_scaled<f4E2M1FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E2M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f6E3M2FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<f8E5M2:f8E8M0FNU:BLOCK_SHAPE_32>>, tensor<1x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
// CHECK-LABEL: test_matmul_t_block_scaled_mxint8
func.func @test_matmul_t_block_scaled_mxint8(%arg0: tensor<4x8x32x!tosa.mxint8>, %arg1: tensor<4x8x1xf8E8M0FNU>, %arg2: tensor<4x16x32x!tosa.mxint8>, %arg3: tensor<4x16x1xf8E8M0FNU>) -> tensor<4x8x16xf32> {
%0 = tosa.matmul_t_block_scaled %arg0, %arg1, %arg2, %arg3 {block_size = #tosa.block_size<BLOCK_SIZE_32>} : (tensor<4x8x32x!tosa.mxint8>, tensor<4x8x1xf8E8M0FNU>, tensor<4x16x32x!tosa.mxint8>, tensor<4x16x1xf8E8M0FNU>) -> tensor<4x8x16xf32>
diff --git a/mlir/test/Dialect/Tosa/verifier.mlir b/mlir/test/Dialect/Tosa/verifier.mlir
index 14a48d5e10ea0..86767ca231e0f 100644
--- a/mlir/test/Dialect/Tosa/verifier.mlir
+++ b/mlir/test/Dialect/Tosa/verifier.mlir
@@ -2115,3 +2115,72 @@ func.func @test_const_mxint8_int64(%arg0 : index) -> tensor<2x!tosa.mxint8> {
%0 = "tosa.const"() {values = dense<tensor<2x!tosa.mxint8> : [127, 245]>} : () -> tensor<2x!tosa.mxint8>
return %0 : tensor<2x!tosa.mxint8>
}
+
+// -----
+
+func.func @test_block_scaled_const_splat_ui8() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{incompatible attribute for element type}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> : 0 : ui8>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_splat_fp64() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{incompatible attribute for element type}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> : 0.0>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_integer_scale_value() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +2 {{unexpected decimal integer literal for a floating point value}}
+ // expected-note at +1 {{add a trailing dot to make the literal a float}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1 : i32, 2 : i32}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_negative_scale_value() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{scale value must be non-negative, got -1.000000e+00}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0 : f8E8M0FNU, -1.0 : f8E8M0FNU}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_scale_value_non_float_explicit_type() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{parsed attribute type 'i32' does not match expected scale type 'f8E8M0FNU'}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {1.0 : i32}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+!mxint8 = !tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>
+!mxint8_scale = !tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 2.0, 4.0}>
+
+func.func @test_block_scaled_const_invalid_num_scales() -> tensor<2x32x!mxint8> {
+ // expected-error at +1 {{tensor type 'tensor<2x32x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.000000e+00, 2.000000e+00, 4.000000e+00}>>' has 2 blocks but got 3 scale values}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!mxint8_scale> : 0 : i8>}> : () -> tensor<2x32x!mxint8>
+ return %0 : tensor<2x32x!mxint8>
+}
+
+// -----
+
+func.func @test_block_scaled_const_invalid_num_scales_wide_inner_dim() -> tensor<2x64x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>> {
+ // expected-error at +1 {{tensor type 'tensor<2x64x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.000000e+00, 2.000000e+00}>>' has 4 blocks but got 2 scale values}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x64x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32, {1.0, 2.0}>> : 0 : i8>}> : () -> tensor<2x64x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+ return %0 : tensor<2x64x!tosa.block_scaled<!tosa.mxint8:f8E8M0FNU:BLOCK_SHAPE_32>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_cast_scale_values_propagate() -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.0, 4.0}>> {
+ // expected-error at +2 {{tensor type 'tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.000000e+00, 4.000000e+00}>>' does not support scale values for this operation}}
+ // expected-error at +1 {{'tosa.const' op result #0 must be tosa-conformant tensor of number values, but got 'tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.000000e+00, 4.000000e+00}>>'}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.0, 4.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.0, 4.0}>>
+ return %0 : tensor<2x32x!tosa.block_scaled<f8E4M3FN:f8E8M0FNU:BLOCK_SHAPE_32, {2.0, 4.0}>>
+}
More information about the Mlir-commits
mailing list