[Mlir-commits] [mlir] [mlir][tosa][spirv] Add TOSA to SPIR-V TOSA pass plumbing (PR #196539)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Fri May 8 07:02:59 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-mlir

Author: Davide Grohmann (davidegrohmann)

<details>
<summary>Changes</summary>

Introduce the initial TosaToSPIRVTosa conversion pass and library wiring. This slice converts func.func regions to spirv.ARM.Graph inside spirv.module, rewrites graph input/result types to SPIR-V ARM tensor types, maps func.return to spirv.ARM.GraphOutputs, and adds focused tests for type conversion, descriptor bindings, and nested containers.


---

Patch is 24.25 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/196539.diff


10 Files Affected:

- (modified) mlir/include/mlir/Conversion/Passes.h (+1) 
- (modified) mlir/include/mlir/Conversion/Passes.td (+19) 
- (added) mlir/include/mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h (+35) 
- (modified) mlir/lib/Conversion/CMakeLists.txt (+1) 
- (added) mlir/lib/Conversion/TosaToSPIRVTosa/CMakeLists.txt (+21) 
- (added) mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp (+184) 
- (added) mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaPass.cpp (+160) 
- (added) mlir/test/Conversion/TosaToSPIRVTosa/descriptor-set-and-bindings.mlir (+18) 
- (added) mlir/test/Conversion/TosaToSPIRVTosa/op-nesting.mlir (+28) 
- (added) mlir/test/Conversion/TosaToSPIRVTosa/type-conversions.mlir (+44) 


``````````diff
diff --git a/mlir/include/mlir/Conversion/Passes.h b/mlir/include/mlir/Conversion/Passes.h
index a54b98004c3b6..82c7670296e52 100644
--- a/mlir/include/mlir/Conversion/Passes.h
+++ b/mlir/include/mlir/Conversion/Passes.h
@@ -76,6 +76,7 @@
 #include "mlir/Conversion/TosaToLinalg/TosaToLinalg.h"
 #include "mlir/Conversion/TosaToMLProgram/TosaToMLProgram.h"
 #include "mlir/Conversion/TosaToSCF/TosaToSCF.h"
+#include "mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h"
 #include "mlir/Conversion/TosaToTensor/TosaToTensor.h"
 #include "mlir/Conversion/UBToLLVM/UBToLLVM.h"
 #include "mlir/Conversion/UBToSPIRV/UBToSPIRV.h"
diff --git a/mlir/include/mlir/Conversion/Passes.td b/mlir/include/mlir/Conversion/Passes.td
index d401b56c7602d..dda756ddab152 100644
--- a/mlir/include/mlir/Conversion/Passes.td
+++ b/mlir/include/mlir/Conversion/Passes.td
@@ -1410,6 +1410,25 @@ def TosaToSCFPass : Pass<"tosa-to-scf"> {
   }];
 }
 
+//===----------------------------------------------------------------------===//
+// TOSA to SPIR-V Graph/TOSA
+//===----------------------------------------------------------------------===//
+
+def TosaToSPIRVTosa : Pass<"tosa-to-spirv-tosa"> {
+  let summary = "Lower TOSA IR to SPIR-V Graph/TOSA operations";
+  let dependentDialects = [
+    "spirv::SPIRVDialect",
+  ];
+  let description = [{
+    Converts TOSA programs to the SPIR-V Graph/TOSA representation by
+    wrapping converted functions in `spirv.module` and `spirv.ARM.Graph`,
+    and rewriting TOSA tensor and shape types to the corresponding SPIR-V ARM
+    tensor types.
+  }];
+
+  let constructor = "tosa::createTosaToSPIRVTosa()";
+}
+
 //===----------------------------------------------------------------------===//
 // TosaToTensor
 //===----------------------------------------------------------------------===//
diff --git a/mlir/include/mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h b/mlir/include/mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h
new file mode 100644
index 0000000000000..22ac6a677b08e
--- /dev/null
+++ b/mlir/include/mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h
@@ -0,0 +1,35 @@
+//===-- TosaToSPIRVTosa.h - TOSA to SPIR-V Graph/TOSA patterns --*- C++ -*-===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+// Provides pass and patterns to lower TOSA IR to SPIR-V Graph/TOSA
+// operations.
+//
+//===----------------------------------------------------------------------===//
+
+#ifndef MLIR_CONVERSION_TOSATOSPIRVTOSA_TOSATOSPIRVTOSA_H
+#define MLIR_CONVERSION_TOSATOSPIRVTOSA_TOSATOSPIRVTOSA_H
+
+#include "mlir/Dialect/SPIRV/Transforms/SPIRVConversion.h"
+#include "mlir/Pass/Pass.h"
+
+namespace mlir {
+
+#define GEN_PASS_DECL_TOSATOSPIRVTOSA
+#include "mlir/Conversion/Passes.h.inc"
+
+namespace tosa {
+
+std::unique_ptr<Pass> createTosaToSPIRVTosa();
+
+void populateTosaToSPIRVTosaConversionPatterns(
+    SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns);
+
+} // namespace tosa
+} // namespace mlir
+
+#endif // MLIR_CONVERSION_TOSATOSPIRVTOSA_TOSATOSPIRVTOSA_H
diff --git a/mlir/lib/Conversion/CMakeLists.txt b/mlir/lib/Conversion/CMakeLists.txt
index e17988b12cade..f5e0bcf613e59 100644
--- a/mlir/lib/Conversion/CMakeLists.txt
+++ b/mlir/lib/Conversion/CMakeLists.txt
@@ -69,6 +69,7 @@ add_subdirectory(TosaToArith)
 add_subdirectory(TosaToLinalg)
 add_subdirectory(TosaToMLProgram)
 add_subdirectory(TosaToSCF)
+add_subdirectory(TosaToSPIRVTosa)
 add_subdirectory(TosaToTensor)
 add_subdirectory(UBToLLVM)
 add_subdirectory(UBToSPIRV)
diff --git a/mlir/lib/Conversion/TosaToSPIRVTosa/CMakeLists.txt b/mlir/lib/Conversion/TosaToSPIRVTosa/CMakeLists.txt
new file mode 100644
index 0000000000000..630278447fa42
--- /dev/null
+++ b/mlir/lib/Conversion/TosaToSPIRVTosa/CMakeLists.txt
@@ -0,0 +1,21 @@
+add_mlir_conversion_library(MLIRTosaToSPIRVTosa
+  TosaToSPIRVTosa.cpp
+  TosaToSPIRVTosaPass.cpp
+
+  ADDITIONAL_HEADER_DIRS
+  ${MLIR_MAIN_INCLUDE_DIR}/mlir/Dialect/Tosa
+  ${MLIR_MAIN_INCLUDE_DIR}/mlir/IR
+
+  DEPENDS
+  MLIRConversionPassIncGen
+
+  LINK_LIBS PUBLIC
+  MLIRFuncDialect
+  MLIRIR
+  MLIRPass
+  MLIRSPIRVDialect
+  MLIRSPIRVConversion
+  MLIRSupport
+  MLIRTransformUtils
+  MLIRTosaDialect
+)
diff --git a/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp
new file mode 100644
index 0000000000000..20de63fa4848f
--- /dev/null
+++ b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp
@@ -0,0 +1,184 @@
+//===- TosaToSPIRVTosa.cpp - TOSA to SPIR-V Graph/TOSA patterns -----------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+// This file implements patterns to convert TOSA IR to SPIR-V Graph/TOSA.
+//
+//===----------------------------------------------------------------------===//
+
+#include "mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h"
+#include "mlir/Dialect/Func/IR/FuncOps.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVOps.h"
+#include "mlir/Transforms/DialectConversion.h"
+#include "llvm/ADT/STLExtras.h"
+
+#define DEBUG_TYPE "tosa-to-spirv-tosa-pattern"
+
+namespace mlir {
+namespace tosa {
+namespace {
+
+constexpr StringLiteral graphARMInterfaceVarABIAttrName =
+    "spv.grapharm.interface_var_abi";
+
+void copyFuncAttrsToGraph(func::FuncOp funcOp, func::FuncOpAdaptor adaptor,
+                          spirv::GraphARMOp graphOp) {
+  for (NamedAttribute attr : adaptor.getAttributes()) {
+    StringRef attrName = attr.getName().getValue();
+    if (attrName == SymbolTable::getSymbolAttrName() ||
+        attrName == funcOp.getFunctionTypeAttrName() ||
+        attrName == funcOp.getArgAttrsAttrName() ||
+        attrName == funcOp.getResAttrsAttrName() ||
+        attrName == graphOp.getEntryPointAttrName())
+      continue;
+
+    graphOp->setAttr(attr.getName(), attr.getValue());
+  }
+}
+
+struct FuncGraphConvert final : public OpConversionPattern<func::FuncOp> {
+  using OpConversionPattern<func::FuncOp>::OpConversionPattern;
+
+private:
+  void normalizeInterfaceVarABIAttr(spirv::GraphARMOp graphOp,
+                                    MLIRContext *context, unsigned index,
+                                    bool isResult,
+                                    uint32_t defaultDescriptorSet,
+                                    uint32_t defaultBinding) const {
+    auto abiInfo =
+        isResult ? graphOp.getResultAttrOfType<spirv::InterfaceVarABIAttr>(
+                       index, graphARMInterfaceVarABIAttrName)
+                 : graphOp.getArgAttrOfType<spirv::InterfaceVarABIAttr>(
+                       index, graphARMInterfaceVarABIAttrName);
+
+    if (!abiInfo) {
+      abiInfo = isResult
+                    ? graphOp.getResultAttrOfType<spirv::InterfaceVarABIAttr>(
+                          index, spirv::getInterfaceVarABIAttrName())
+                    : graphOp.getArgAttrOfType<spirv::InterfaceVarABIAttr>(
+                          index, spirv::getInterfaceVarABIAttrName());
+    }
+
+    if (!abiInfo) {
+      abiInfo = spirv::InterfaceVarABIAttr::get(
+          defaultDescriptorSet, defaultBinding, std::nullopt, context);
+    }
+
+    if (isResult) {
+      graphOp.setResultAttr(index, spirv::getInterfaceVarABIAttrName(),
+                            abiInfo);
+      graphOp.removeResultAttr(index, graphARMInterfaceVarABIAttrName);
+    } else {
+      graphOp.setArgAttr(index, spirv::getInterfaceVarABIAttrName(), abiInfo);
+      graphOp.removeArgAttr(index, graphARMInterfaceVarABIAttrName);
+    }
+  }
+
+  void normalizeInterfaceVarABIAttrs(spirv::GraphARMOp graphOp,
+                                     MLIRContext *context, unsigned inputs,
+                                     unsigned outputs,
+                                     uint32_t descriptorSet) const {
+    for (auto argIndex : llvm::seq<unsigned>(0, inputs)) {
+      normalizeInterfaceVarABIAttr(graphOp, context, argIndex, false,
+                                   descriptorSet, argIndex);
+    }
+    for (auto resIndex : llvm::seq<unsigned>(0, outputs)) {
+      normalizeInterfaceVarABIAttr(graphOp, context, resIndex, true,
+                                   descriptorSet, resIndex + inputs);
+    }
+  }
+
+public:
+  LogicalResult
+  matchAndRewrite(func::FuncOp funcOp, func::FuncOpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    MLIRContext *context = rewriter.getContext();
+
+    StringRef name = adaptor.getSymName();
+
+    bool entryPoint = !isa<func::FuncOp>(funcOp->getParentOp());
+    if (entryPoint) {
+      auto spvModule = spirv::ModuleOp::create(
+          rewriter, funcOp.getLoc(), spirv::AddressingModel::Logical,
+          spirv::MemoryModel::Vulkan, std::nullopt,
+          ("_spirv_tosa_" + name).str());
+
+      rewriter.setInsertionPoint(spvModule.getBody(), spvModule.begin());
+    }
+
+    FunctionType ftype = adaptor.getFunctionType();
+    ArrayAttr argAttrs = adaptor.getArgAttrsAttr();
+    ArrayAttr resAttrs = adaptor.getResAttrsAttr();
+
+    TypeConverter::SignatureConversion signatureConverter(ftype.getNumInputs());
+    if (failed(typeConverter->convertSignatureArgs(ftype.getInputs(),
+                                                   signatureConverter))) {
+      return funcOp.emitError("failed to convert function argument types");
+    }
+
+    // Update the signature of the function.
+    SmallVector<Type, 2> newResultTypes;
+    if (failed(getTypeConverter()->convertTypes(ftype.getResults(),
+                                                newResultTypes))) {
+      return funcOp.emitError("failed to convert function result types");
+    }
+
+    auto graphTy = GraphType::get(
+        context, signatureConverter.getConvertedTypes(), newResultTypes);
+    auto entryPointAttr = BoolAttr::get(context, entryPoint);
+    auto graphOp =
+        spirv::GraphARMOp::create(rewriter, funcOp.getLoc(), graphTy, argAttrs,
+                                  resAttrs, entryPointAttr, name);
+    copyFuncAttrsToGraph(funcOp, adaptor, graphOp);
+    rewriter.inlineRegionBefore(funcOp.getBody(), graphOp.getBody(),
+                                graphOp.end());
+    if (failed(rewriter.convertRegionTypes(
+            &graphOp.getBody(), *getTypeConverter(), &signatureConverter))) {
+      return funcOp.emitError("failed to convert function regions");
+    }
+
+    if (entryPoint) {
+      uint32_t descriptorSet = 0;
+      if (auto descriptorSetAttr =
+              funcOp->getAttrOfType<IntegerAttr>("descriptor_set")) {
+        descriptorSet = static_cast<uint32_t>(descriptorSetAttr.getUInt());
+      }
+
+      normalizeInterfaceVarABIAttrs(graphOp, context, ftype.getNumInputs(),
+                                    ftype.getNumResults(), descriptorSet);
+    }
+
+    rewriter.eraseOp(funcOp);
+    return success();
+  }
+};
+
+/// Converts func.return to spirv.ARM.GraphOutputs.
+class ReturnGraphOutputConvert final
+    : public OpConversionPattern<func::ReturnOp> {
+  using OpConversionPattern<func::ReturnOp>::OpConversionPattern;
+
+  LogicalResult
+  matchAndRewrite(func::ReturnOp returnOp, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    rewriter.replaceOpWithNewOp<spirv::GraphOutputsARMOp>(
+        returnOp, adaptor.getOperands());
+    return success();
+  }
+};
+
+} // namespace
+
+void populateTosaToSPIRVTosaConversionPatterns(
+    SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns) {
+  patterns.add<FuncGraphConvert, ReturnGraphOutputConvert>(
+      typeConverter, patterns.getContext());
+}
+
+} // namespace tosa
+} // namespace mlir
diff --git a/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaPass.cpp b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaPass.cpp
new file mode 100644
index 0000000000000..817709cac9a41
--- /dev/null
+++ b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosaPass.cpp
@@ -0,0 +1,160 @@
+//===- TosaToSPIRVTosaPass.cpp - Lower TOSA to SPIR-V Graph/TOSA ----------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+// This pass lowers TOSA IR to the SPIR-V Graph/TOSA representation.
+//
+//===----------------------------------------------------------------------===//
+
+#include "mlir/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.h"
+
+#include "mlir/Dialect/Func/IR/FuncOps.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h"
+#include "mlir/Dialect/SPIRV/Transforms/SPIRVConversion.h"
+#include "mlir/Dialect/Tosa/IR/TosaOps.h"
+#include "mlir/IR/TypeUtilities.h"
+#include "mlir/Transforms/DialectConversion.h"
+
+#include <algorithm>
+
+namespace mlir {
+#define GEN_PASS_DEF_TOSATOSPIRVTOSA
+#include "mlir/Conversion/Passes.h.inc"
+} // namespace mlir
+
+namespace mlir {
+namespace tosa {
+namespace {
+
+spirv::VerCapExtAttr getVerCapExtAttr(MLIRContext *context) {
+  return spirv::VerCapExtAttr::get(
+      spirv::Version::V_1_5,
+      {
+          spirv::Capability::VulkanMemoryModel,
+          spirv::Capability::Shader,
+          spirv::Capability::Int8,
+          spirv::Capability::Int16,
+          spirv::Capability::Int64,
+          spirv::Capability::Float16,
+          spirv::Capability::BFloat16TypeKHR,
+          spirv::Capability::Float8EXT,
+          spirv::Capability::TensorsARM,
+          spirv::Capability::GraphARM,
+          spirv::Capability::ReplicatedCompositesEXT,
+      },
+      {
+          spirv::Extension::SPV_ARM_tensors,
+          spirv::Extension::SPV_ARM_graph,
+          spirv::Extension::SPV_KHR_vulkan_memory_model,
+          spirv::Extension::SPV_EXT_replicated_composites,
+          spirv::Extension::SPV_KHR_bfloat16,
+          spirv::Extension::SPV_EXT_float8,
+      },
+      context);
+}
+
+struct TosaToSPIRVTosa : public impl::TosaToSPIRVTosaBase<TosaToSPIRVTosa> {
+  void runOnOperation() override {
+    MLIRContext *context = &getContext();
+    RewritePatternSet patterns(context);
+    Operation *op = getOperation();
+
+    auto targetAttr = spirv::lookupTargetEnv(op);
+    if (!targetAttr) {
+      targetAttr = spirv::TargetEnvAttr::get(
+          getVerCapExtAttr(context), spirv::getDefaultResourceLimits(context),
+          spirv::ClientAPI::Unknown, spirv::Vendor::Unknown,
+          spirv::DeviceType::Unknown, spirv::TargetEnvAttr::kUnknownDeviceID);
+    }
+
+    std::unique_ptr<ConversionTarget> target =
+        SPIRVConversionTarget::get(targetAttr);
+
+    target->addLegalDialect<spirv::SPIRVDialect>();
+    target->addIllegalDialect<tosa::TosaDialect>();
+
+    SPIRVTypeConverter typeConverter(targetAttr);
+    typeConverter.addConversion([this](IntegerType integerType) {
+      return this->convertIntegerType(integerType);
+    });
+    typeConverter.addConversion([this](TensorType tensorType) {
+      return this->convertTensorType(tensorType);
+    });
+    typeConverter.addConversion([this](tosa::shapeType shapeType) {
+      return this->convertShapeType(shapeType);
+    });
+
+    populateTosaToSPIRVTosaConversionPatterns(typeConverter, patterns);
+
+    auto targetEnvAttrName = spirv::getTargetEnvAttrName();
+    auto symbolTableOp = SymbolTable::getNearestSymbolTable(op);
+    if (symbolTableOp && !symbolTableOp->hasAttrOfType<spirv::TargetEnvAttr>(
+                             targetEnvAttrName)) {
+      symbolTableOp->setAttr(targetEnvAttrName, targetAttr);
+    }
+
+    FrozenRewritePatternSet frozenPatterns(std::move(patterns));
+
+    if (failed(applyPartialConversion(op, *target, frozenPatterns))) {
+      signalPassFailure();
+    }
+  }
+
+private:
+  IntegerType convertIntegerType(IntegerType integerType) {
+    if (integerType.getWidth() == 48) {
+      return IntegerType::get(&getContext(), 64, integerType.getSignedness());
+    }
+
+    if (integerType.getWidth() == 4) {
+      return IntegerType::get(&getContext(), 8, integerType.getSignedness());
+    }
+
+    return integerType;
+  }
+
+  SmallVector<int64_t> convertShape(ArrayRef<int64_t> shape) {
+    bool requiresRankConversion =
+        llvm::all_of(shape, [](int64_t dim) { return dim == 0; });
+    if (requiresRankConversion)
+      return SmallVector<int64_t>({1});
+    bool isPartiallyDynamic =
+        llvm::any_of(shape, [](int64_t dim) { return dim < 0; }) &&
+        llvm::any_of(shape, [](int64_t dim) { return dim > 0; });
+    if (isPartiallyDynamic)
+      return SmallVector<int64_t>(shape.size(), ShapedType::kDynamic);
+    return SmallVector<int64_t>(shape);
+  }
+
+  spirv::TensorArmType convertTensorType(TensorType tensorType) {
+    Type elementType = getElementTypeOrSelf(tensorType);
+    if (elementType.isIndex())
+      elementType = IntegerType::get(&getContext(), 32);
+    if (auto integerType = dyn_cast<IntegerType>(elementType))
+      elementType = convertIntegerType(integerType);
+
+    SmallVector<int64_t> shape = tensorType.hasRank()
+                                     ? convertShape(tensorType.getShape())
+                                     : SmallVector<int64_t>();
+
+    return spirv::TensorArmType::get(shape, elementType);
+  }
+
+  spirv::TensorArmType convertShapeType(tosa::shapeType shapeType) {
+    const auto rank = std::max(shapeType.getRank(), 1);
+    return spirv::TensorArmType::get({rank},
+                                     IntegerType::get(&getContext(), 32));
+  }
+};
+} // namespace
+
+std::unique_ptr<Pass> createTosaToSPIRVTosa() {
+  return std::make_unique<TosaToSPIRVTosa>();
+}
+
+} // namespace tosa
+} // namespace mlir
diff --git a/mlir/test/Conversion/TosaToSPIRVTosa/descriptor-set-and-bindings.mlir b/mlir/test/Conversion/TosaToSPIRVTosa/descriptor-set-and-bindings.mlir
new file mode 100644
index 0000000000000..7a716069022ec
--- /dev/null
+++ b/mlir/test/Conversion/TosaToSPIRVTosa/descriptor-set-and-bindings.mlir
@@ -0,0 +1,18 @@
+// RUN: mlir-opt --split-input-file --tosa-to-spirv-tosa -verify-diagnostics %s | FileCheck %s
+// CHECK-NOT: spv.grapharm.interface_var_abi
+
+// CHECK: spirv.module @_spirv_tosa_descriptor_set Logical Vulkan {
+// CHECK: spirv.ARM.Graph @descriptor_set(%[[ARG0:.*]]: !spirv.arm.tensor<1xi8> {spirv.interface_var_abi = #spirv.interface_var_abi<(42, 0)>}) -> (!spirv.arm.tensor<1xi8> {spirv.interface_var_abi = #spirv.interface_var_abi<(42, 1)>}) attributes {descriptor_set = 42 : ui32, entry_point = true} {
+func.func @descriptor_set(%arg0: tensor<1xi8>) -> tensor<1xi8> attributes {descriptor_set = 42 : ui32} {
+  // CHECK: spirv.ARM.GraphOutputs %[[ARG0]] : !spirv.arm.tensor<1xi8>
+  return %arg0 : tensor<1xi8>
+}
+
+// -----
+
+// CHECK: spirv.module @_spirv_tosa_custom_grapharm_abi Logical Vulkan {
+// CHECK: spirv.ARM.Graph @custom_grapharm_abi(%[[ARG0:.*]]: !spirv.arm.tensor<1xi8> {spirv.interface_var_abi = #spirv.interface_var_abi<(3, 9)>}) -> (!spirv.arm.tensor<1xi8> {spirv.interface_var_abi = #spirv.interface_var_abi<(7, 11)>}) attributes {descriptor_set = 42 : ui32, entry_point = true} {
+func.func @custom_grapharm_abi(%arg0: tensor<1xi8> {spv.grapharm.interface_var_abi = #spirv.interface_var_abi<(3, 9)>}) -> (tensor<1xi8> {spv.grapharm.interface_var_abi = #spirv.interface_var_abi<(7, 11)>}) attributes {descriptor_set = 42 : ui32} {
+  // CHECK: spirv.ARM.GraphOutputs %[[ARG0]] : !spirv.arm.tensor<1xi8>
+  return %arg0 : tensor<1xi8>
+}
diff --git a/mlir/test/Conversion/TosaToSPIRVTosa/op-nesting.mlir b/mlir/test/Conversion/TosaToSPIRVTosa/op-nesting.mlir
new file mode 100644
index 0000000000000..9d2a5650b184c
--- /dev/null
+++ b/mlir/test/Conversion/TosaToSPIRVTosa/op-nesting.mlir
@@ -0,0 +1,28 @@
+// RUN: mlir-opt --split-input-file --tosa-to-spirv-tosa -verify-diagnostics %s | FileCheck %s
+
+// CHECK: gpu.module @random_container
+gpu.module @random_container {
+  // CHECK: spirv.module @_spirv_tosa_nested Logical...
[truncated]

``````````

</details>


https://github.com/llvm/llvm-project/pull/196539


More information about the Mlir-commits mailing list