[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