[Mlir-commits] [mlir] d4dd90a - [mlir][tosa] Add constant block scaled support (#205506)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Thu Jul 9 09:09:48 PDT 2026
Author: Luke Hutton
Date: 2026-07-09T17:09:44+01:00
New Revision: d4dd90ae4d7a6b559ebc2ef5b75d2eacafadb48f
URL: https://github.com/llvm/llvm-project/commit/d4dd90ae4d7a6b559ebc2ef5b75d2eacafadb48f
DIFF: https://github.com/llvm/llvm-project/commit/d4dd90ae4d7a6b559ebc2ef5b75d2eacafadb48f.diff
LOG: [mlir][tosa] Add constant block scaled support (#205506)
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.
Added:
Modified:
mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
mlir/test/Dialect/Tosa/availability.mlir
mlir/test/Dialect/Tosa/ops.mlir
mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
mlir/test/Dialect/Tosa/verifier.mlir
Removed:
################################################################################
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
index ef644845ed937..291097a7f7514 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaComplianceData.h.inc
@@ -1305,7 +1305,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.h b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
index dfecda940d8c7..b1404d9c700ae 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.h
@@ -92,8 +92,6 @@ class TosaShapeOperatorWithSameRanks
}
};
-LogicalResult verifyBlockScaledTensorType(mlir::Type type);
-
} // namespace tosa
} // namespace OpTrait
@@ -143,6 +141,10 @@ Type getStorageElementTypeOrSelf(Type type);
// Returns the storage element type for a given value
Type getStorageElementTypeOrSelf(Value value);
+// Verify a block scaled tensor type is valid
+LogicalResult verifyBlockScaledTensorType(mlir::Type type,
+ bool allowScaleValues);
+
} // namespace tosa
} // namespace mlir
diff --git a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
index 3518a62800963..5e5760e33b5f6 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaOps.td
@@ -2970,9 +2970,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 63be7b69b6398..1518024ecaa19 100644
--- a/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
+++ b/mlir/include/mlir/Dialect/Tosa/IR/TosaTypesBase.td
@@ -109,7 +109,12 @@ 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 = [{
@@ -121,17 +126,26 @@ def Tosa_BlockScaled : Tosa_Type<"BlockScaled", "block_scaled"> {
specifying the block size along the innermost dimension.
`scale_type` - The type of the scale value associated with each block.
`value_type` - The type of the data values in each block.
+ `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
EnumParameter<Tosa_BlockShape>:$block_shape,
Tosa_MXFPScale:$scale_type,
- Tosa_MXFPValue:$value_type
+ Tosa_MXFPValue:$value_type,
+ OptionalArrayRefParameter<"Attribute">:$scale_values
);
let assemblyFormat =
- "`<` $block_shape```:```$scale_type```:```$value_type `>`";
+ "`<` $block_shape```:```$scale_type```:```$value_type "
+ "(`,` `{` custom<ScaleValues>($scale_values, ref($scale_type))^ `}`)? "
+ "`>`";
+
+ let extraClassDeclaration = [{
+ bool hasScaleValues() const { return !getScaleValues().empty(); }
+ }];
}
//===----------------------------------------------------------------------===//
@@ -158,8 +172,7 @@ def AtLeastRankOne : And<[
CPred<"::llvm::cast<::mlir::RankedTensorType>($_self).getRank() >= 1">]>;
def IsValidBlockScaledTensorType
- : CPred<"::mlir::succeeded(::mlir::OpTrait::tosa::"
- "verifyBlockScaledTensorType($_self))">;
+ : CPred<"::mlir::succeeded(::mlir::tosa::verifyBlockScaledTensorType($_self, false))">;
class TosaTensorOf<list<Type> allowedTypes,
string summary = "tosa-conformant tensor">
diff --git a/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp b/mlir/lib/Dialect/Tosa/IR/TosaOps.cpp
index 829c0b41a68cf..bf9d98ffe35a0 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,31 +741,114 @@ LogicalResult mlir::tosa::mxint8Type::convertFromAttribute(
// TOSA block scaling utilities.
//===----------------------------------------------------------------------===//
-LogicalResult OpTrait::tosa::verifyBlockScaledTensorType(mlir::Type type) {
+LogicalResult mlir::tosa::verifyBlockScaledTensorType(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 failure();
+
if (!tensorType.hasRank())
return success();
if (tensorType.getRank() == 0)
return failure();
- 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 = tensorType.getNumElements() / blockSize;
+ if (elemType.getScaleValues().size() != numBlocks)
+ return failure();
+ }
+
+ const int64_t blockedDimension = tensorShape.back();
+ if (ShapedType::isDynamic(blockedDimension))
+ return success();
if (blockedDimension % blockSize != 0)
return failure();
return success();
}
+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);
+ });
+}
+
+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 {
+ // Block scaled values are stored as a single byte. This is because possible
+ // value data types are either 8-bit or sub-byte. Sub-byte types are aligned
+ // to 8-bits.
+ 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,
+ result);
+}
+
//===----------------------------------------------------------------------===//
// TOSA Operator Verifiers.
//===----------------------------------------------------------------------===//
@@ -851,7 +935,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());
@@ -860,17 +944,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 (!attrBlockScaledType.hasScaleValues())
+ return op.emitOpError(
+ "attribute block scaled type must have scale values");
+
+ if (failed(verifyBlockScaledTensorType(attrType, true)))
+ return op.emitOpError("block scaled attribute type is not valid, got ")
+ << attrType;
+
+ 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 c358460bbc662..bc014fb9b928e 100644
--- a/mlir/test/Dialect/Tosa/availability.mlir
+++ b/mlir/test/Dialect/Tosa/availability.mlir
@@ -660,7 +660,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 ec3d0c881d857..4e71d512def2d 100644
--- a/mlir/test/Dialect/Tosa/ops.mlir
+++ b/mlir/test/Dialect/Tosa/ops.mlir
@@ -1938,3 +1938,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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_splat
+func.func @test_block_scaled_const_splat() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN, {1.0, 1.0}>> : 0.0 : f4E2M1FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_splat_mxint8
+func.func @test_block_scaled_const_splat_mxint8() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.0, 1.0}>> : 0 : i8>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+}
+
+// -----
+
+// CHECK-LABEL: test_block_scaled_const_scale_values
+func.func @test_block_scaled_const_scale_values() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0, 2.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+// 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0 : f8E8M0FNU, 2.0 : f8E8M0FNU}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+// 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ %0 = "tosa.const"() <{values = dense<tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0, 2.0, 4.0, 8.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+// 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {2.0, 4.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ %1 = tosa.cast %0 : (tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>) -> 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 a5e4d4397f290..bea262a11a391 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p0-invalid.mlir
@@ -193,6 +193,24 @@ func.func @test_const_fp6e3m2(%arg0 : index) -> tensor<4xf6E3M2FN> {
// -----
+func.func @test_const_block_scaled_types() -> (tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>) {
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN, {1.0}>> : 0.0 : f4E2M1FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN, {1.0}>> : 0.0 : f6E2M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN>>
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN, {1.0}>> : 0.0 : f6E3M2FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN>>
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2, {1.0}>> : 0.0 : f8E5M2>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2>>
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.0}>> : 0 : i8>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+ return %0, %1, %2, %3, %4, %5 : tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+}
+
+// -----
+
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 e7a4fe6735457..d290ff71f052d 100644
--- a/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
+++ b/mlir/test/Dialect/Tosa/tosa-validation-version-1p1-valid.mlir
@@ -152,6 +152,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<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>) {
+ %0 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN, {1.0}>> : 0.0 : f4E2M1FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>
+ %1 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN, {1.0}>> : 0.0 : f6E2M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN>>
+ %2 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN, {1.0}>> : 0.0 : f6E3M2FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN>>
+ %3 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ %4 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2, {1.0}>> : 0.0 : f8E5M2>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2>>
+ %5 = "tosa.const"() <{values = dense<tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.0}>> : 0 : i8>}> : () -> tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+ return %0, %1, %2, %3, %4, %5 : tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f4E2M1FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E2M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f6E3M2FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E5M2>>, tensor<1x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+}
+
+// -----
+
// 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 478532e52c438..c31effeb6a1d2 100644
--- a/mlir/test/Dialect/Tosa/verifier.mlir
+++ b/mlir/test/Dialect/Tosa/verifier.mlir
@@ -2268,3 +2268,71 @@ 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ // expected-error at +1 {{incompatible attribute for element type}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> : 0 : ui8>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_splat_fp64() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ // expected-error at +1 {{incompatible attribute for element type}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> : 0.0>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_integer_scale_value() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1 : i32, 2 : i32}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_negative_scale_value() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ // expected-error at +1 {{scale value must be non-negative, got -1.000000e+00}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0 : f8E8M0FNU, -1.0 : f8E8M0FNU}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_scale_value_non_float_explicit_type() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>> {
+ // 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<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {1.0 : i32}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN>>
+}
+
+// -----
+
+!mxint8 = !tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>
+!mxint8_scale = !tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.0, 2.0, 4.0}>
+
+func.func @test_block_scaled_const_invalid_num_scales() -> tensor<2x32x!mxint8> {
+ // expected-error at +1 {{'tosa.const' op block scaled attribute type is not valid, got 'tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.000000e+00, 2.000000e+00, 4.000000e+00}>>'}}
+ %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<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>> {
+ // expected-error at +1 {{'tosa.const' op block scaled attribute type is not valid, got 'tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.000000e+00, 2.000000e+00}>>'}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8, {1.0, 2.0}>> : 0 : i8>}> : () -> tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+ return %0 : tensor<2x64x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:!tosa.mxint8>>
+}
+
+// -----
+
+func.func @test_block_scaled_const_cast_scale_values_propagate() -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {2.0, 4.0}>> {
+ // expected-error at +1 {{'tosa.const' op result #0 must be tosa-conformant tensor of number values, but got 'tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {2.000000e+00, 4.000000e+00}>>'}}
+ %0 = "tosa.const"() <{values = dense<tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {2.0, 4.0}>> : 0.0 : f8E4M3FN>}> : () -> tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {2.0, 4.0}>>
+ return %0 : tensor<2x32x!tosa.block_scaled<BLOCK_SHAPE_32:f8E8M0FNU:f8E4M3FN, {2.0, 4.0}>>
+}
More information about the Mlir-commits
mailing list