[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