[Mlir-commits] [mlir] bbf8a33 - [mlir][spirv][tosa] Extend TOSA to SPIR-V TOSA op conversion (#200009)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Thu May 28 23:48:57 PDT 2026
Author: Davide Grohmann
Date: 2026-05-29T08:48:52+02:00
New Revision: bbf8a33ad5f51b5641e7665e5f0d2f0079b454f6
URL: https://github.com/llvm/llvm-project/commit/bbf8a33ad5f51b5641e7665e5f0d2f0079b454f6
DIFF: https://github.com/llvm/llvm-project/commit/bbf8a33ad5f51b5641e7665e5f0d2f0079b454f6.diff
LOG: [mlir][spirv][tosa] Extend TOSA to SPIR-V TOSA op conversion (#200009)
Add conversion patterns for additional TOSA 1.0 operations targeting
the SPIR-V TOSA extended instruction set.
Introduce a common TosaOpConvert pattern with small replacer classes
to share result type conversion while keeping op-specific replacement
logic explicit.
The newly covered operations include:
* elementwise ops such as clamp, arithmetic_right_shift, mul, table,
negate, and select
* reductions and argmax
* gather, scatter, resize
* reshape, reverse, slice, tile, transpose
* cast and const_shape
Also add conversion tests.
---------
Signed-off-by: Davide Grohmann <davide.grohmann at arm.com>
Added:
Modified:
mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaOps.cpp
mlir/test/Conversion/TosaToSPIRVTosa/tosa-to-spirv.mlir
Removed:
################################################################################
diff --git a/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaOps.cpp b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaOps.cpp
index bb4f4b76c9b29..948a0b277cd86 100644
--- a/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaOps.cpp
+++ b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaOps.cpp
@@ -21,18 +21,26 @@ namespace mlir::tosa {
namespace {
template <typename OpAdaptor>
-Value getInput1(OpAdaptor adaptor) {
- return adaptor.getInput1();
+spirv::TosaExtNaNPropagationModeType getNanMode(OpAdaptor adaptor) {
+ return static_cast<spirv::TosaExtNaNPropagationModeType>(
+ adaptor.getNanMode());
}
-Value getInput1(tosa::ErfOpAdaptor adaptor) { return adaptor.getInput(); }
-
-Value getInput1(tosa::SigmoidOpAdaptor adaptor) { return adaptor.getInput(); }
+template <typename OpAdaptor>
+spirv::TosaExtResizeModeType getResizeMode(OpAdaptor adaptor) {
+ return static_cast<spirv::TosaExtResizeModeType>(adaptor.getMode());
+}
-Value getInput1(tosa::TanhOpAdaptor adaptor) { return adaptor.getInput(); }
+DenseIntElementsAttr getI32TensorArmAttr(ArrayRef<int32_t> values,
+ ConversionPatternRewriter &rewriter) {
+ return DenseIntElementsAttr::get(
+ spirv::TensorArmType::get(static_cast<int64_t>(values.size()),
+ IntegerType::get(rewriter.getContext(), 32)),
+ values);
+}
-template <typename SourceOp, typename TargetOp>
-struct UnaryElementwiseOpConvert final : public OpConversionPattern<SourceOp> {
+template <typename SourceOp, auto Replace>
+struct TosaOpConvert final : public OpConversionPattern<SourceOp> {
using OpConversionPattern<SourceOp>::OpConversionPattern;
LogicalResult
@@ -41,86 +49,302 @@ struct UnaryElementwiseOpConvert final : public OpConversionPattern<SourceOp> {
Type type = this->getTypeConverter()->convertType(op.getType());
if (!type)
return rewriter.notifyMatchFailure(op, "type conversion failed");
- rewriter.replaceOpWithNewOp<TargetOp>(op, type, getInput1(adaptor));
- return success();
+ return Replace(op, adaptor, type, rewriter);
}
};
template <typename SourceOp, typename TargetOp>
-struct BinaryElementwiseOpConvert final : public OpConversionPattern<SourceOp> {
- using OpConversionPattern<SourceOp>::OpConversionPattern;
+LogicalResult replaceUnaryInput1(SourceOp op,
+ typename SourceOp::Adaptor adaptor, Type type,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<TargetOp>(op, type, adaptor.getInput1());
+ return success();
+}
- LogicalResult
- matchAndRewrite(SourceOp op, typename SourceOp::Adaptor adaptor,
- ConversionPatternRewriter &rewriter) const override {
- Type type = this->getTypeConverter()->convertType(op.getType());
- if (!type)
- return rewriter.notifyMatchFailure(op, "type conversion failed");
- rewriter.replaceOpWithNewOp<TargetOp>(op, type, adaptor.getInput1(),
- adaptor.getInput2());
- return success();
- }
-};
+template <typename SourceOp, typename TargetOp>
+LogicalResult replaceUnaryInput(SourceOp op, typename SourceOp::Adaptor adaptor,
+ Type type,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<TargetOp>(op, type, adaptor.getInput());
+ return success();
+}
template <typename SourceOp, typename TargetOp>
-struct BinaryNanModeElementwiseOpConvert final
- : public OpConversionPattern<SourceOp> {
- using OpConversionPattern<SourceOp>::OpConversionPattern;
+LogicalResult
+replaceBinaryElementwise(SourceOp op, typename SourceOp::Adaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<TargetOp>(op, type, adaptor.getInput1(),
+ adaptor.getInput2());
+ return success();
+}
- LogicalResult
- matchAndRewrite(SourceOp op, typename SourceOp::Adaptor adaptor,
- ConversionPatternRewriter &rewriter) const override {
- auto nanMode =
- static_cast<spirv::TosaExtNaNPropagationModeType>(adaptor.getNanMode());
- Type type = this->getTypeConverter()->convertType(op.getType());
- if (!type)
- return rewriter.notifyMatchFailure(op, "type conversion failed");
- rewriter.replaceOpWithNewOp<TargetOp>(
- op, type, nanMode, adaptor.getInput1(), adaptor.getInput2());
- return success();
- }
-};
+template <typename SourceOp, typename TargetOp>
+LogicalResult
+replaceBinaryNanModeElementwise(SourceOp op, typename SourceOp::Adaptor adaptor,
+ Type type,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<TargetOp>(
+ op, type, getNanMode(adaptor), adaptor.getInput1(), adaptor.getInput2());
+ return success();
+}
+
+template <typename SourceOp, typename TargetOp>
+LogicalResult replaceReduction(SourceOp op, typename SourceOp::Adaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<TargetOp>(op, type, adaptor.getAxis(),
+ adaptor.getInput());
+ return success();
+}
+
+template <typename SourceOp, typename TargetOp>
+LogicalResult
+replaceNanModeReduction(SourceOp op, typename SourceOp::Adaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<TargetOp>(
+ op, type, adaptor.getAxis(), getNanMode(adaptor), adaptor.getInput());
+ return success();
+}
+
+LogicalResult replaceClamp(tosa::ClampOp op, tosa::ClampOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaClampOp>(
+ op, type, adaptor.getMinVal(), adaptor.getMaxVal(), getNanMode(adaptor),
+ adaptor.getInput());
+ return success();
+}
+
+LogicalResult
+replaceArithmeticRightShift(tosa::ArithmeticRightShiftOp op,
+ tosa::ArithmeticRightShiftOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaArithmeticRightShiftOp>(
+ op, type, adaptor.getRound(), adaptor.getInput1(), adaptor.getInput2());
+ return success();
+}
+
+LogicalResult replaceMul(tosa::MulOp op, tosa::MulOpAdaptor adaptor, Type type,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaMulOp>(
+ op, type, adaptor.getInput1(), adaptor.getInput2(), adaptor.getShift());
+ return success();
+}
+
+LogicalResult replaceTable(tosa::TableOp op, tosa::TableOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaTableOp>(op, type, adaptor.getInput1(),
+ adaptor.getTable());
+ return success();
+}
+
+LogicalResult replaceNegate(tosa::NegateOp op, tosa::NegateOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaNegateOp>(
+ op, type, adaptor.getInput1(), adaptor.getInput1Zp(),
+ adaptor.getOutputZp());
+ return success();
+}
+
+LogicalResult replaceSelect(tosa::SelectOp op, tosa::SelectOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaSelectOp>(
+ op, type, adaptor.getInput1(), adaptor.getInput2(), adaptor.getInput3());
+ return success();
+}
+
+LogicalResult replaceReshape(tosa::ReshapeOp op, tosa::ReshapeOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaReshapeOp>(
+ op, type, adaptor.getInput1(), adaptor.getShape());
+ return success();
+}
+
+LogicalResult replaceReverse(tosa::ReverseOp op, tosa::ReverseOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaReverseOp>(op, type, adaptor.getAxis(),
+ adaptor.getInput1());
+ return success();
+}
+
+LogicalResult replaceSlice(tosa::SliceOp op, tosa::SliceOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaSliceOp>(
+ op, type, adaptor.getInput1(), adaptor.getStart(), adaptor.getSize());
+ return success();
+}
+
+LogicalResult replaceTile(tosa::TileOp op, tosa::TileOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaTileOp>(op, type, adaptor.getInput1(),
+ adaptor.getMultiples());
+ return success();
+}
+
+LogicalResult replaceTranspose(tosa::TransposeOp op,
+ tosa::TransposeOpAdaptor adaptor, Type type,
+ ConversionPatternRewriter &rewriter) {
+ DenseIntElementsAttr perms =
+ getI32TensorArmAttr(adaptor.getPerms(), rewriter);
+ rewriter.replaceOpWithNewOp<spirv::TosaTransposeOp>(op, type, perms,
+ adaptor.getInput1());
+ return success();
+}
+
+LogicalResult replaceGather(tosa::GatherOp op, tosa::GatherOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaGatherOp>(
+ op, type, adaptor.getValues(), adaptor.getIndices());
+ return success();
+}
+
+LogicalResult replaceScatter(tosa::ScatterOp op, tosa::ScatterOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaScatterOp>(
+ op, type, adaptor.getValuesIn(), adaptor.getIndices(),
+ adaptor.getInput());
+ return success();
+}
+
+LogicalResult replaceResize(tosa::ResizeOp op, tosa::ResizeOpAdaptor adaptor,
+ Type type, ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<spirv::TosaResizeOp>(
+ op, type, getResizeMode(adaptor), adaptor.getInput(), adaptor.getScale(),
+ adaptor.getOffset(), adaptor.getBorder());
+ return success();
+}
+
+LogicalResult replaceConstShape(tosa::ConstShapeOp op,
+ tosa::ConstShapeOpAdaptor adaptor, Type type,
+ ConversionPatternRewriter &rewriter) {
+ SmallVector<int32_t> values;
+ for (const APInt &value : adaptor.getValues().getValues<APInt>())
+ values.push_back(value.getSExtValue());
+
+ rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
+ op, type, getI32TensorArmAttr(values, rewriter));
+ return success();
+}
} // namespace
void populateTosaToSPIRVTosaOpsConversionPatterns(
SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns) {
patterns.add<
- UnaryElementwiseOpConvert<tosa::ErfOp, spirv::TosaErfOp>,
- UnaryElementwiseOpConvert<tosa::SigmoidOp, spirv::TosaSigmoidOp>,
- UnaryElementwiseOpConvert<tosa::TanhOp, spirv::TosaTanhOp>,
- BinaryElementwiseOpConvert<tosa::AddOp, spirv::TosaAddOp>,
- BinaryElementwiseOpConvert<tosa::BitwiseAndOp, spirv::TosaBitwiseAndOp>,
- BinaryElementwiseOpConvert<tosa::BitwiseOrOp, spirv::TosaBitwiseOrOp>,
- BinaryElementwiseOpConvert<tosa::BitwiseXorOp, spirv::TosaBitwiseXorOp>,
- BinaryElementwiseOpConvert<tosa::IntDivOp, spirv::TosaIntDivOp>,
- BinaryElementwiseOpConvert<tosa::LogicalAndOp, spirv::TosaLogicalAndOp>,
- BinaryElementwiseOpConvert<tosa::LogicalLeftShiftOp,
- spirv::TosaLogicalLeftShiftOp>,
- BinaryElementwiseOpConvert<tosa::LogicalRightShiftOp,
- spirv::TosaLogicalRightShiftOp>,
- BinaryElementwiseOpConvert<tosa::LogicalOrOp, spirv::TosaLogicalOrOp>,
- BinaryElementwiseOpConvert<tosa::LogicalXorOp, spirv::TosaLogicalXorOp>,
- BinaryNanModeElementwiseOpConvert<tosa::MaximumOp, spirv::TosaMaximumOp>,
- BinaryNanModeElementwiseOpConvert<tosa::MinimumOp, spirv::TosaMinimumOp>,
- BinaryElementwiseOpConvert<tosa::PowOp, spirv::TosaPowOp>,
- BinaryElementwiseOpConvert<tosa::SubOp, spirv::TosaSubOp>,
- UnaryElementwiseOpConvert<tosa::AbsOp, spirv::TosaAbsOp>,
- UnaryElementwiseOpConvert<tosa::BitwiseNotOp, spirv::TosaBitwiseNotOp>,
- UnaryElementwiseOpConvert<tosa::CeilOp, spirv::TosaCeilOp>,
- UnaryElementwiseOpConvert<tosa::ClzOp, spirv::TosaClzOp>,
- UnaryElementwiseOpConvert<tosa::CosOp, spirv::TosaCosOp>,
- UnaryElementwiseOpConvert<tosa::ExpOp, spirv::TosaExpOp>,
- UnaryElementwiseOpConvert<tosa::FloorOp, spirv::TosaFloorOp>,
- UnaryElementwiseOpConvert<tosa::LogOp, spirv::TosaLogOp>,
- UnaryElementwiseOpConvert<tosa::LogicalNotOp, spirv::TosaLogicalNotOp>,
- UnaryElementwiseOpConvert<tosa::ReciprocalOp, spirv::TosaReciprocalOp>,
- UnaryElementwiseOpConvert<tosa::RsqrtOp, spirv::TosaRsqrtOp>,
- UnaryElementwiseOpConvert<tosa::SinOp, spirv::TosaSinOp>,
- BinaryElementwiseOpConvert<tosa::EqualOp, spirv::TosaEqualOp>,
- BinaryElementwiseOpConvert<tosa::GreaterOp, spirv::TosaGreaterOp>,
- BinaryElementwiseOpConvert<tosa::GreaterEqualOp,
- spirv::TosaGreaterEqualOp>>(
+ TosaOpConvert<tosa::ArgMaxOp, replaceNanModeReduction<
+ tosa::ArgMaxOp, spirv::TosaArgMaxOp>>,
+ TosaOpConvert<tosa::ClampOp, replaceClamp>,
+ TosaOpConvert<tosa::ErfOp,
+ replaceUnaryInput<tosa::ErfOp, spirv::TosaErfOp>>,
+ TosaOpConvert<tosa::SigmoidOp,
+ replaceUnaryInput<tosa::SigmoidOp, spirv::TosaSigmoidOp>>,
+ TosaOpConvert<tosa::TanhOp,
+ replaceUnaryInput<tosa::TanhOp, spirv::TosaTanhOp>>,
+ TosaOpConvert<tosa::AddOp,
+ replaceBinaryElementwise<tosa::AddOp, spirv::TosaAddOp>>,
+ TosaOpConvert<tosa::ArithmeticRightShiftOp, replaceArithmeticRightShift>,
+ TosaOpConvert<tosa::BitwiseAndOp,
+ replaceBinaryElementwise<tosa::BitwiseAndOp,
+ spirv::TosaBitwiseAndOp>>,
+ TosaOpConvert<
+ tosa::BitwiseOrOp,
+ replaceBinaryElementwise<tosa::BitwiseOrOp, spirv::TosaBitwiseOrOp>>,
+ TosaOpConvert<tosa::BitwiseXorOp,
+ replaceBinaryElementwise<tosa::BitwiseXorOp,
+ spirv::TosaBitwiseXorOp>>,
+ TosaOpConvert<tosa::IntDivOp, replaceBinaryElementwise<
+ tosa::IntDivOp, spirv::TosaIntDivOp>>,
+ TosaOpConvert<tosa::LogicalAndOp,
+ replaceBinaryElementwise<tosa::LogicalAndOp,
+ spirv::TosaLogicalAndOp>>,
+ TosaOpConvert<tosa::LogicalLeftShiftOp,
+ replaceBinaryElementwise<tosa::LogicalLeftShiftOp,
+ spirv::TosaLogicalLeftShiftOp>>,
+ TosaOpConvert<tosa::LogicalRightShiftOp,
+ replaceBinaryElementwise<tosa::LogicalRightShiftOp,
+ spirv::TosaLogicalRightShiftOp>>,
+ TosaOpConvert<
+ tosa::LogicalOrOp,
+ replaceBinaryElementwise<tosa::LogicalOrOp, spirv::TosaLogicalOrOp>>,
+ TosaOpConvert<tosa::LogicalXorOp,
+ replaceBinaryElementwise<tosa::LogicalXorOp,
+ spirv::TosaLogicalXorOp>>,
+ TosaOpConvert<tosa::MaximumOp,
+ replaceBinaryNanModeElementwise<tosa::MaximumOp,
+ spirv::TosaMaximumOp>>,
+ TosaOpConvert<tosa::MinimumOp,
+ replaceBinaryNanModeElementwise<tosa::MinimumOp,
+ spirv::TosaMinimumOp>>,
+ TosaOpConvert<tosa::MulOp, replaceMul>,
+ TosaOpConvert<tosa::PowOp,
+ replaceBinaryElementwise<tosa::PowOp, spirv::TosaPowOp>>,
+ TosaOpConvert<tosa::SubOp,
+ replaceBinaryElementwise<tosa::SubOp, spirv::TosaSubOp>>,
+ TosaOpConvert<tosa::TableOp, replaceTable>,
+ TosaOpConvert<tosa::AbsOp,
+ replaceUnaryInput1<tosa::AbsOp, spirv::TosaAbsOp>>,
+ TosaOpConvert<
+ tosa::BitwiseNotOp,
+ replaceUnaryInput1<tosa::BitwiseNotOp, spirv::TosaBitwiseNotOp>>,
+ TosaOpConvert<tosa::CeilOp,
+ replaceUnaryInput1<tosa::CeilOp, spirv::TosaCeilOp>>,
+ TosaOpConvert<tosa::ClzOp,
+ replaceUnaryInput1<tosa::ClzOp, spirv::TosaClzOp>>,
+ TosaOpConvert<tosa::CosOp,
+ replaceUnaryInput1<tosa::CosOp, spirv::TosaCosOp>>,
+ TosaOpConvert<tosa::ExpOp,
+ replaceUnaryInput1<tosa::ExpOp, spirv::TosaExpOp>>,
+ TosaOpConvert<tosa::FloorOp,
+ replaceUnaryInput1<tosa::FloorOp, spirv::TosaFloorOp>>,
+ TosaOpConvert<tosa::LogOp,
+ replaceUnaryInput1<tosa::LogOp, spirv::TosaLogOp>>,
+ TosaOpConvert<
+ tosa::LogicalNotOp,
+ replaceUnaryInput1<tosa::LogicalNotOp, spirv::TosaLogicalNotOp>>,
+ TosaOpConvert<tosa::NegateOp, replaceNegate>,
+ TosaOpConvert<
+ tosa::ReciprocalOp,
+ replaceUnaryInput1<tosa::ReciprocalOp, spirv::TosaReciprocalOp>>,
+ TosaOpConvert<tosa::RsqrtOp,
+ replaceUnaryInput1<tosa::RsqrtOp, spirv::TosaRsqrtOp>>,
+ TosaOpConvert<tosa::SinOp,
+ replaceUnaryInput1<tosa::SinOp, spirv::TosaSinOp>>,
+ TosaOpConvert<tosa::SelectOp, replaceSelect>,
+ TosaOpConvert<tosa::EqualOp, replaceBinaryElementwise<
+ tosa::EqualOp, spirv::TosaEqualOp>>,
+ TosaOpConvert<
+ tosa::GreaterOp,
+ replaceBinaryElementwise<tosa::GreaterOp, spirv::TosaGreaterOp>>,
+ TosaOpConvert<tosa::GreaterEqualOp,
+ replaceBinaryElementwise<tosa::GreaterEqualOp,
+ spirv::TosaGreaterEqualOp>>,
+ TosaOpConvert<
+ tosa::ReduceAllOp,
+ replaceReduction<tosa::ReduceAllOp, spirv::TosaReduceAllOp>>,
+ TosaOpConvert<
+ tosa::ReduceAnyOp,
+ replaceReduction<tosa::ReduceAnyOp, spirv::TosaReduceAnyOp>>,
+ TosaOpConvert<
+ tosa::ReduceMaxOp,
+ replaceNanModeReduction<tosa::ReduceMaxOp, spirv::TosaReduceMaxOp>>,
+ TosaOpConvert<
+ tosa::ReduceMinOp,
+ replaceNanModeReduction<tosa::ReduceMinOp, spirv::TosaReduceMinOp>>,
+ TosaOpConvert<
+ tosa::ReduceProductOp,
+ replaceReduction<tosa::ReduceProductOp, spirv::TosaReduceProductOp>>,
+ TosaOpConvert<
+ tosa::ReduceSumOp,
+ replaceReduction<tosa::ReduceSumOp, spirv::TosaReduceSumOp>>,
+ TosaOpConvert<tosa::ReshapeOp, replaceReshape>,
+ TosaOpConvert<tosa::ReverseOp, replaceReverse>,
+ TosaOpConvert<tosa::SliceOp, replaceSlice>,
+ TosaOpConvert<tosa::TileOp, replaceTile>,
+ TosaOpConvert<tosa::TransposeOp, replaceTranspose>,
+ TosaOpConvert<tosa::GatherOp, replaceGather>,
+ TosaOpConvert<tosa::ScatterOp, replaceScatter>,
+ TosaOpConvert<tosa::ResizeOp, replaceResize>,
+ TosaOpConvert<tosa::CastOp,
+ replaceUnaryInput<tosa::CastOp, spirv::TosaCastOp>>,
+ TosaOpConvert<tosa::ConstShapeOp, replaceConstShape>>(
typeConverter, patterns.getContext());
}
diff --git a/mlir/test/Conversion/TosaToSPIRVTosa/tosa-to-spirv.mlir b/mlir/test/Conversion/TosaToSPIRVTosa/tosa-to-spirv.mlir
index 4e54e5cdd3634..a175baf62eda1 100644
--- a/mlir/test/Conversion/TosaToSPIRVTosa/tosa-to-spirv.mlir
+++ b/mlir/test/Conversion/TosaToSPIRVTosa/tosa-to-spirv.mlir
@@ -1,5 +1,31 @@
// RUN: mlir-opt --split-input-file --tosa-to-spirv-tosa %s | FileCheck %s
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ArgMax
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @argmax_int
+func.func @argmax_int(%arg0: tensor<2x3x4xi8>) -> tensor<2x4xi32> {
+ // CHECK: %[[ARGMAX:.*]] = spirv.Tosa.ArgMax axis = 1, nan_mode = <Propagate>, %arg0 : !spirv.arm.tensor<2x3x4xi8> -> !spirv.arm.tensor<2x4xi32>
+ %res = tosa.argmax %arg0 {axis = 1 : i32, nan_mode = PROPAGATE} : (tensor<2x3x4xi8>) -> tensor<2x4xi32>
+ return %res : tensor<2x4xi32>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Clamp
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @clamp_int
+func.func @clamp_int(%arg0: tensor<4x8xi8>) -> tensor<4x8xi8> {
+ // CHECK: %[[CLAMP:.*]] = spirv.Tosa.Clamp min_val = -2 : i8, max_val = 3 : i8, nan_mode = <Propagate>, %arg0 : !spirv.arm.tensor<4x8xi8> -> !spirv.arm.tensor<4x8xi8>
+ %res = tosa.clamp %arg0 {min_val = -2 : i8, max_val = 3 : i8, nan_mode = PROPAGATE} : (tensor<4x8xi8>) -> tensor<4x8xi8>
+ return %res : tensor<4x8xi8>
+}
+
+// -----
+
//===----------------------------------------------------------------------===//
// spirv.TOSA.Erf
//===----------------------------------------------------------------------===//
@@ -52,6 +78,19 @@ func.func @add_int(%arg0: tensor<4x7x3x10xi32>, %arg1: tensor<4x7x3x1xi32>) -> t
// -----
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ArithmeticRightShift
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @arithmetic_right_shift_int
+func.func @arithmetic_right_shift_int(%arg0: tensor<1x4xi16>, %arg1: tensor<3x4xi16>) -> tensor<3x4xi16> {
+ // CHECK: %[[SHIFT:.*]] = spirv.Tosa.ArithmeticRightShift round = true, %arg0, %arg1 : !spirv.arm.tensor<1x4xi16>, !spirv.arm.tensor<3x4xi16> -> !spirv.arm.tensor<3x4xi16>
+ %res = tosa.arithmetic_right_shift %arg0, %arg1 {round = true} : (tensor<1x4xi16>, tensor<3x4xi16>) -> tensor<3x4xi16>
+ return %res : tensor<3x4xi16>
+}
+
+// -----
+
//===----------------------------------------------------------------------===//
// spirv.TOSA.BitwiseAnd
//===----------------------------------------------------------------------===//
@@ -195,6 +234,19 @@ func.func @minimum_int(%arg0: tensor<15x2x10x11xi32>, %arg1: tensor<15x1x10x11xi
// -----
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Mul
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @mul_int
+func.func @mul_int(%arg0: tensor<2x4xi32>, %arg1: tensor<2x1xi32>, %arg2: tensor<1xi8>) -> tensor<2x4xi32> {
+ // CHECK: %[[MUL:.*]] = spirv.Tosa.Mul %arg0, %arg1, %arg2 : !spirv.arm.tensor<2x4xi32>, !spirv.arm.tensor<2x1xi32>, !spirv.arm.tensor<1xi8> -> !spirv.arm.tensor<2x4xi32>
+ %res = tosa.mul %arg0, %arg1, %arg2 : (tensor<2x4xi32>, tensor<2x1xi32>, tensor<1xi8>) -> tensor<2x4xi32>
+ return %res : tensor<2x4xi32>
+}
+
+// -----
+
//===----------------------------------------------------------------------===//
// spirv.TOSA.Pow
//===----------------------------------------------------------------------===//
@@ -221,6 +273,19 @@ func.func @sub_int(%arg0: tensor<6x10x6x6xi32>, %arg1: tensor<1x10x6x6xi32>) ->
// -----
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Table
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @table_int
+func.func @table_int(%arg0: tensor<3x2xi8>, %arg1: tensor<256xi8>) -> tensor<3x2xi8> {
+ // CHECK: %[[TABLE:.*]] = spirv.Tosa.Table %arg0, %arg1 : !spirv.arm.tensor<3x2xi8>, !spirv.arm.tensor<256xi8> -> !spirv.arm.tensor<3x2xi8>
+ %res = tosa.table %arg0, %arg1 : (tensor<3x2xi8>, tensor<256xi8>) -> tensor<3x2xi8>
+ return %res : tensor<3x2xi8>
+}
+
+// -----
+
//===----------------------------------------------------------------------===//
// spirv.TOSA.Abs
//===----------------------------------------------------------------------===//
@@ -338,6 +403,19 @@ func.func @logicalnot_any(%arg0: tensor<54x26x10xi1>) -> tensor<54x26x10xi1> {
// -----
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Negate
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @negate_int
+func.func @negate_int(%arg0: tensor<2x3xi8>, %arg1: tensor<1xi8>, %arg2: tensor<1xi8>) -> tensor<2x3xi8> {
+ // CHECK: %[[NEGATE:.*]] = spirv.Tosa.Negate %arg0, %arg1, %arg2 : !spirv.arm.tensor<2x3xi8>, !spirv.arm.tensor<1xi8>, !spirv.arm.tensor<1xi8> -> !spirv.arm.tensor<2x3xi8>
+ %res = tosa.negate %arg0, %arg1, %arg2 : (tensor<2x3xi8>, tensor<1xi8>, tensor<1xi8>) -> tensor<2x3xi8>
+ return %res : tensor<2x3xi8>
+}
+
+// -----
+
//===----------------------------------------------------------------------===//
// spirv.TOSA.Reciprocal
//===----------------------------------------------------------------------===//
@@ -377,6 +455,19 @@ func.func @sin_fp(%arg0: tensor<49x38x58xf16>) -> tensor<49x38x58xf16> {
// -----
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Select
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @select_int
+func.func @select_int(%arg0: tensor<2x1xi1>, %arg1: tensor<2x4xi8>, %arg2: tensor<2x4xi8>) -> tensor<2x4xi8> {
+ // CHECK: %[[SELECT:.*]] = spirv.Tosa.Select %arg0, %arg1, %arg2 : !spirv.arm.tensor<2x1xi1>, !spirv.arm.tensor<2x4xi8>, !spirv.arm.tensor<2x4xi8> -> !spirv.arm.tensor<2x4xi8>
+ %res = tosa.select %arg0, %arg1, %arg2 : (tensor<2x1xi1>, tensor<2x4xi8>, tensor<2x4xi8>) -> tensor<2x4xi8>
+ return %res : tensor<2x4xi8>
+}
+
+// -----
+
//===----------------------------------------------------------------------===//
// spirv.TOSA.Equal
//===----------------------------------------------------------------------===//
@@ -414,3 +505,211 @@ func.func @greaterequal_int(%arg0: tensor<10x17x7x1xi32>, %arg1: tensor<10x17x7x
return %res : tensor<10x17x7x16xi1>
}
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ReduceAll
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reduce_all
+func.func @reduce_all(%arg0: tensor<2x3x4xi1>) -> tensor<2x1x4xi1> {
+ // CHECK: %[[REDUCE:.*]] = spirv.Tosa.ReduceAll axis = 1, %arg0 : !spirv.arm.tensor<2x3x4xi1> -> !spirv.arm.tensor<2x1x4xi1>
+ %res = tosa.reduce_all %arg0 {axis = 1 : i32} : (tensor<2x3x4xi1>) -> tensor<2x1x4xi1>
+ return %res : tensor<2x1x4xi1>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ReduceAny
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reduce_any
+func.func @reduce_any(%arg0: tensor<2x3x4xi1>) -> tensor<2x1x4xi1> {
+ // CHECK: %[[REDUCE:.*]] = spirv.Tosa.ReduceAny axis = 1, %arg0 : !spirv.arm.tensor<2x3x4xi1> -> !spirv.arm.tensor<2x1x4xi1>
+ %res = tosa.reduce_any %arg0 {axis = 1 : i32} : (tensor<2x3x4xi1>) -> tensor<2x1x4xi1>
+ return %res : tensor<2x1x4xi1>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ReduceMax
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reduce_max_int
+func.func @reduce_max_int(%arg0: tensor<2x3x4xi8>) -> tensor<2x1x4xi8> {
+ // CHECK: %[[REDUCE:.*]] = spirv.Tosa.ReduceMax axis = 1, nan_mode = <Propagate>, %arg0 : !spirv.arm.tensor<2x3x4xi8> -> !spirv.arm.tensor<2x1x4xi8>
+ %res = tosa.reduce_max %arg0 {axis = 1 : i32, nan_mode = PROPAGATE} : (tensor<2x3x4xi8>) -> tensor<2x1x4xi8>
+ return %res : tensor<2x1x4xi8>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ReduceMin
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reduce_min_int
+func.func @reduce_min_int(%arg0: tensor<2x3x4xi8>) -> tensor<2x1x4xi8> {
+ // CHECK: %[[REDUCE:.*]] = spirv.Tosa.ReduceMin axis = 1, nan_mode = <Propagate>, %arg0 : !spirv.arm.tensor<2x3x4xi8> -> !spirv.arm.tensor<2x1x4xi8>
+ %res = tosa.reduce_min %arg0 {axis = 1 : i32, nan_mode = PROPAGATE} : (tensor<2x3x4xi8>) -> tensor<2x1x4xi8>
+ return %res : tensor<2x1x4xi8>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ReduceProduct
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reduce_product_fp
+func.func @reduce_product_fp(%arg0: tensor<2x3x4xf32>) -> tensor<2x1x4xf32> {
+ // CHECK: %[[REDUCE:.*]] = spirv.Tosa.ReduceProduct axis = 1, %arg0 : !spirv.arm.tensor<2x3x4xf32> -> !spirv.arm.tensor<2x1x4xf32>
+ %res = tosa.reduce_product %arg0 {axis = 1 : i32} : (tensor<2x3x4xf32>) -> tensor<2x1x4xf32>
+ return %res : tensor<2x1x4xf32>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.ReduceSum
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reduce_sum_int
+func.func @reduce_sum_int(%arg0: tensor<2x3x4xi32>) -> tensor<2x1x4xi32> {
+ // CHECK: %[[REDUCE:.*]] = spirv.Tosa.ReduceSum axis = 1, %arg0 : !spirv.arm.tensor<2x3x4xi32> -> !spirv.arm.tensor<2x1x4xi32>
+ %res = tosa.reduce_sum %arg0 {axis = 1 : i32} : (tensor<2x3x4xi32>) -> tensor<2x1x4xi32>
+ return %res : tensor<2x1x4xi32>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Reshape
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reshape_int
+func.func @reshape_int(%arg0: tensor<25x6x29x35xi16>) -> tensor<125x6x7x29xi16> {
+ %shape = "tosa.const_shape"() <{values = dense<[125, 6, 7, 29]> : tensor<4xindex>}> : () -> !tosa.shape<4>
+ // CHECK: %[[SHAPE:.*]] = spirv.Constant dense<[125, 6, 7, 29]> : !spirv.arm.tensor<4xi32>
+ // CHECK: %[[RESHAPE:.*]] = spirv.Tosa.Reshape %arg0, %[[SHAPE]] : !spirv.arm.tensor<25x6x29x35xi16>, !spirv.arm.tensor<4xi32> -> !spirv.arm.tensor<125x6x7x29xi16>
+ %res = tosa.reshape %arg0, %shape : (tensor<25x6x29x35xi16>, !tosa.shape<4>) -> tensor<125x6x7x29xi16>
+ return %res : tensor<125x6x7x29xi16>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Reverse
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @reverse_int
+func.func @reverse_int(%arg0: tensor<20x5x28x31xi32>) -> tensor<20x5x28x31xi32> {
+ // CHECK: %[[REVERSE:.*]] = spirv.Tosa.Reverse axis = 2, %arg0 : !spirv.arm.tensor<20x5x28x31xi32> -> !spirv.arm.tensor<20x5x28x31xi32>
+ %res = tosa.reverse %arg0 {axis = 2 : i32} : (tensor<20x5x28x31xi32>) -> tensor<20x5x28x31xi32>
+ return %res : tensor<20x5x28x31xi32>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Slice
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @slice_int
+func.func @slice_int(%arg0: tensor<32x19x41xi8>) -> tensor<21x5x2xi8> {
+ %start = "tosa.const_shape"() <{values = dense<[8, 11, 39]> : tensor<3xindex>}> : () -> !tosa.shape<3>
+ %size = "tosa.const_shape"() <{values = dense<[21, 5, 2]> : tensor<3xindex>}> : () -> !tosa.shape<3>
+ // CHECK: %[[START:.*]] = spirv.Constant dense<[8, 11, 39]> : !spirv.arm.tensor<3xi32>
+ // CHECK: %[[SIZE:.*]] = spirv.Constant dense<[21, 5, 2]> : !spirv.arm.tensor<3xi32>
+ // CHECK: %[[SLICE:.*]] = spirv.Tosa.Slice %arg0, %[[START]], %[[SIZE]] : !spirv.arm.tensor<32x19x41xi8>, !spirv.arm.tensor<3xi32>, !spirv.arm.tensor<3xi32> -> !spirv.arm.tensor<21x5x2xi8>
+ %res = tosa.slice %arg0, %start, %size : (tensor<32x19x41xi8>, !tosa.shape<3>, !tosa.shape<3>) -> tensor<21x5x2xi8>
+ return %res : tensor<21x5x2xi8>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Tile
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @tile_int
+func.func @tile_int(%arg0: tensor<10x28x21xi16>) -> tensor<10x28x63xi16> {
+ %multiples = "tosa.const_shape"() <{values = dense<[1, 1, 3]> : tensor<3xindex>}> : () -> !tosa.shape<3>
+ // CHECK: %[[MULTIPLES:.*]] = spirv.Constant dense<[1, 1, 3]> : !spirv.arm.tensor<3xi32>
+ // CHECK: %[[TILE:.*]] = spirv.Tosa.Tile %arg0, %[[MULTIPLES]] : !spirv.arm.tensor<10x28x21xi16>, !spirv.arm.tensor<3xi32> -> !spirv.arm.tensor<10x28x63xi16>
+ %res = tosa.tile %arg0, %multiples : (tensor<10x28x21xi16>, !tosa.shape<3>) -> tensor<10x28x63xi16>
+ return %res : tensor<10x28x63xi16>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Transpose
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @transpose_int
+func.func @transpose_int(%arg0: tensor<14x28x1x61xi16>) -> tensor<1x14x28x61xi16> {
+ // CHECK: %[[TRANSPOSE:.*]] = spirv.Tosa.Transpose perms = [2, 0, 1, 3], %arg0 : !spirv.arm.tensor<14x28x1x61xi16> -> !spirv.arm.tensor<1x14x28x61xi16>
+ %res = tosa.transpose %arg0 {perms = array<i32: 2, 0, 1, 3>} : (tensor<14x28x1x61xi16>) -> tensor<1x14x28x61xi16>
+ return %res : tensor<1x14x28x61xi16>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Gather
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @gather_int
+func.func @gather_int(%arg0: tensor<31x11x45xi32>, %arg1: tensor<31x15xi32>) -> tensor<31x15x45xi32> {
+ // CHECK: %[[GATHER:.*]] = spirv.Tosa.Gather %arg0, %arg1 : !spirv.arm.tensor<31x11x45xi32>, !spirv.arm.tensor<31x15xi32> -> !spirv.arm.tensor<31x15x45xi32>
+ %res = tosa.gather %arg0, %arg1 : (tensor<31x11x45xi32>, tensor<31x15xi32>) -> tensor<31x15x45xi32>
+ return %res : tensor<31x15x45xi32>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Scatter
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @scatter_int
+func.func @scatter_int(%arg0: tensor<34x28x54xi32>, %arg1: tensor<34x18xi32>, %arg2: tensor<34x18x54xi32>) -> tensor<34x28x54xi32> {
+ // CHECK: %[[SCATTER:.*]] = spirv.Tosa.Scatter %arg0, %arg1, %arg2 : !spirv.arm.tensor<34x28x54xi32>, !spirv.arm.tensor<34x18xi32>, !spirv.arm.tensor<34x18x54xi32> -> !spirv.arm.tensor<34x28x54xi32>
+ %res = tosa.scatter %arg0, %arg1, %arg2 : (tensor<34x28x54xi32>, tensor<34x18xi32>, tensor<34x18x54xi32>) -> tensor<34x28x54xi32>
+ return %res : tensor<34x28x54xi32>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Resize
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @resize_int
+func.func @resize_int(%arg0: tensor<1x1x31x55xi8>) -> tensor<1x1x278x55xi8> {
+ %scale = "tosa.const_shape"() <{values = dense<[16, 1, 9, 1]> : tensor<4xindex>}> : () -> !tosa.shape<4>
+ %offset = "tosa.const_shape"() <{values = dense<0> : tensor<2xindex>}> : () -> !tosa.shape<2>
+ %border = "tosa.const_shape"() <{values = dense<[0, 7]> : tensor<2xindex>}> : () -> !tosa.shape<2>
+ // CHECK: %[[SCALE:.*]] = spirv.Constant dense<[16, 1, 9, 1]> : !spirv.arm.tensor<4xi32>
+ // CHECK: %[[OFFSET:.*]] = spirv.Constant dense<0> : !spirv.arm.tensor<2xi32>
+ // CHECK: %[[BORDER:.*]] = spirv.Constant dense<[0, 7]> : !spirv.arm.tensor<2xi32>
+ // CHECK: %[[RESIZE:.*]] = spirv.Tosa.Resize mode = <NearestNeighbor>, %arg0, %[[SCALE]], %[[OFFSET]], %[[BORDER]] : !spirv.arm.tensor<1x1x31x55xi8>, !spirv.arm.tensor<4xi32>, !spirv.arm.tensor<2xi32>, !spirv.arm.tensor<2xi32> -> !spirv.arm.tensor<1x1x278x55xi8>
+ %res = tosa.resize %arg0, %scale, %offset, %border {mode = NEAREST_NEIGHBOR} : (tensor<1x1x31x55xi8>, !tosa.shape<4>, !tosa.shape<2>, !tosa.shape<2>) -> tensor<1x1x278x55xi8>
+ return %res : tensor<1x1x278x55xi8>
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.TOSA.Cast
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: spirv.ARM.Graph @cast_int
+func.func @cast_int(%arg0: tensor<2x3xi8>) -> tensor<2x3xi32> {
+ // CHECK: %[[CAST:.*]] = spirv.Tosa.Cast %arg0 : !spirv.arm.tensor<2x3xi8> -> !spirv.arm.tensor<2x3xi32>
+ %res = tosa.cast %arg0 : (tensor<2x3xi8>) -> tensor<2x3xi32>
+ return %res : tensor<2x3xi32>
+}
More information about the Mlir-commits
mailing list