[Mlir-commits] [mlir] [mlir][SPIR-V] Convert scf.index_switch to spirv.mlir.selection with spirv.Switch (PR #200573)
Arseniy Obolenskiy
llvmlistbot at llvm.org
Sat May 30 05:59:53 PDT 2026
https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/200573
None
>From 01320c0cd45a009e25140907e79baa300356c105 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Sat, 30 May 2026 14:58:59 +0200
Subject: [PATCH] [mlir][SPIR-V] Convert scf.index_switch to
spirv.mlir.selection with spirv.Switch
---
mlir/lib/Conversion/SCFToSPIRV/SCFToSPIRV.cpp | 92 +++++++++++++++++-
.../Conversion/SCFToSPIRV/index-switch.mlir | 96 +++++++++++++++++++
2 files changed, 184 insertions(+), 4 deletions(-)
create mode 100644 mlir/test/Conversion/SCFToSPIRV/index-switch.mlir
diff --git a/mlir/lib/Conversion/SCFToSPIRV/SCFToSPIRV.cpp b/mlir/lib/Conversion/SCFToSPIRV/SCFToSPIRV.cpp
index d5140f3faa6ff..fb5148b230f47 100644
--- a/mlir/lib/Conversion/SCFToSPIRV/SCFToSPIRV.cpp
+++ b/mlir/lib/Conversion/SCFToSPIRV/SCFToSPIRV.cpp
@@ -293,6 +293,90 @@ struct IfOpConversion : SCFToSPIRVPattern<scf::IfOp> {
}
};
+//===----------------------------------------------------------------------===//
+// scf::IndexSwitchOp
+//===----------------------------------------------------------------------===//
+
+/// Pattern to convert a scf::IndexSwitchOp within kernel functions into
+/// spirv::SelectionOp with a spirv::SwitchOp header.
+struct IndexSwitchOpConversion final : SCFToSPIRVPattern<scf::IndexSwitchOp> {
+ using SCFToSPIRVPattern::SCFToSPIRVPattern;
+
+ LogicalResult
+ matchAndRewrite(scf::IndexSwitchOp switchOp, OpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ Location loc = switchOp.getLoc();
+
+ // Compute return types.
+ SmallVector<Type, 8> returnTypes;
+ for (auto result : switchOp.getResults()) {
+ auto convertedType = typeConverter.convertType(result.getType());
+ if (!convertedType)
+ return rewriter.notifyMatchFailure(
+ loc,
+ llvm::formatv("failed to convert type '{0}'", result.getType()));
+ returnTypes.push_back(convertedType);
+ }
+
+ // The selector must be a SPIR-V integer; spirv.Switch literals are
+ // interpreted with the selector's bit width.
+ Value selector = adaptor.getArg();
+ auto selectorType = dyn_cast<IntegerType>(selector.getType());
+ if (!selectorType)
+ return rewriter.notifyMatchFailure(loc,
+ "selector type is not an integer");
+ unsigned selectorWidth = selectorType.getWidth();
+
+ // Create the `spirv.mlir.selection` op, its header block, and merge block.
+ auto selectionControl = spirv::SelectionControl::None;
+ if (auto attr = switchOp->getAttrOfType<spirv::SelectionControlAttr>(
+ spirv::getSelectionControlAttrName()))
+ selectionControl = attr.getValue();
+ auto selectionOp =
+ spirv::SelectionOp::create(rewriter, loc, selectionControl);
+ auto *mergeBlock = rewriter.createBlock(&selectionOp.getBody(),
+ selectionOp.getBody().end());
+ spirv::MergeOp::create(rewriter, loc);
+
+ OpBuilder::InsertionGuard guard(rewriter);
+ auto *headerBlock = rewriter.createBlock(&selectionOp.getBody().front());
+
+ // Inline each case region before the merge block and branch to it.
+ SmallVector<APInt> caseLiterals;
+ SmallVector<Block *> caseBlocks;
+ ArrayRef<int64_t> cases = switchOp.getCases();
+ for (auto [caseValue, caseRegion] :
+ llvm::zip_equal(cases, switchOp.getCaseRegions())) {
+ Block *caseBlock = &caseRegion.front();
+ rewriter.setInsertionPointToEnd(&caseRegion.back());
+ spirv::BranchOp::create(rewriter, loc, mergeBlock);
+ rewriter.inlineRegionBefore(caseRegion, mergeBlock);
+ caseLiterals.push_back(
+ APInt(selectorWidth, caseValue, /*isSigned=*/true));
+ caseBlocks.push_back(caseBlock);
+ }
+
+ // Inline the default region before the merge block and branch to it.
+ Region &defaultRegion = switchOp.getDefaultRegion();
+ Block *defaultBlock = &defaultRegion.front();
+ rewriter.setInsertionPointToEnd(&defaultRegion.back());
+ spirv::BranchOp::create(rewriter, loc, mergeBlock);
+ rewriter.inlineRegionBefore(defaultRegion, mergeBlock);
+
+ // Create the `spirv.Switch` terminator for the header block. The case
+ // regions carry their results through variables, so the branches take no
+ // operands.
+ SmallVector<ValueRange> caseOperands(caseBlocks.size(), ValueRange());
+ rewriter.setInsertionPointToEnd(headerBlock);
+ spirv::SwitchOp::create(rewriter, loc, selector, defaultBlock, ValueRange(),
+ caseLiterals, caseBlocks, caseOperands);
+
+ replaceSCFOutputValue(switchOp, selectionOp, rewriter, scfToSPIRVContext,
+ returnTypes);
+ return success();
+ }
+};
+
//===----------------------------------------------------------------------===//
// scf::YieldOp
//===----------------------------------------------------------------------===//
@@ -311,7 +395,7 @@ struct TerminatorOpConversion final : SCFToSPIRVPattern<scf::YieldOp> {
// TODO: Implement conversion for the remaining `scf` ops.
if (parent->getDialect()->getNamespace() ==
scf::SCFDialect::getDialectNamespace() &&
- !isa<scf::IfOp, scf::ForOp, scf::WhileOp>(parent))
+ !isa<scf::IfOp, scf::ForOp, scf::WhileOp, scf::IndexSwitchOp>(parent))
return rewriter.notifyMatchFailure(
terminatorOp,
llvm::formatv("conversion not supported for parent op: '{0}'",
@@ -459,7 +543,7 @@ struct WhileOpConversion final : SCFToSPIRVPattern<scf::WhileOp> {
void mlir::populateSCFToSPIRVPatterns(const SPIRVTypeConverter &typeConverter,
ScfToSPIRVContext &scfToSPIRVContext,
RewritePatternSet &patterns) {
- patterns.add<ForOpConversion, IfOpConversion, TerminatorOpConversion,
- WhileOpConversion>(patterns.getContext(), typeConverter,
- scfToSPIRVContext.getImpl());
+ patterns.add<ForOpConversion, IfOpConversion, IndexSwitchOpConversion,
+ TerminatorOpConversion, WhileOpConversion>(
+ patterns.getContext(), typeConverter, scfToSPIRVContext.getImpl());
}
diff --git a/mlir/test/Conversion/SCFToSPIRV/index-switch.mlir b/mlir/test/Conversion/SCFToSPIRV/index-switch.mlir
new file mode 100644
index 0000000000000..5e017838e51be
--- /dev/null
+++ b/mlir/test/Conversion/SCFToSPIRV/index-switch.mlir
@@ -0,0 +1,96 @@
+// RUN: mlir-opt -convert-scf-to-spirv %s -o - | FileCheck %s
+
+module attributes {
+ spirv.target_env = #spirv.target_env<
+ #spirv.vce<v1.0, [Shader], [SPV_KHR_storage_buffer_storage_class]>, #spirv.resource_limits<>>
+} {
+
+// CHECK-LABEL: @switch_no_result
+func.func @switch_no_result(%arg0 : memref<10xf32, #spirv.storage_class<StorageBuffer>>, %arg1 : index) {
+ %value = arith.constant 0.0 : f32
+ %i = arith.constant 0 : index
+
+ // CHECK: spirv.mlir.selection {
+ // CHECK-NEXT: spirv.Switch {{%.*}} : i32, [
+ // CHECK-NEXT: default: [[DEFAULT:\^.*]],
+ // CHECK-NEXT: 2: [[CASE:\^.*]]
+ // CHECK-NEXT: ]
+ // CHECK-NEXT: [[CASE]]:
+ // CHECK: spirv.Branch [[MERGE:\^.*]]
+ // CHECK-NEXT: [[DEFAULT]]:
+ // CHECK-NEXT: spirv.Branch [[MERGE]]
+ // CHECK-NEXT: [[MERGE]]:
+ // CHECK-NEXT: spirv.mlir.merge
+ // CHECK-NEXT: }
+ // CHECK-NEXT: spirv.Return
+
+ scf.index_switch %arg1
+ case 2 {
+ memref.store %value, %arg0[%i] : memref<10xf32, #spirv.storage_class<StorageBuffer>>
+ scf.yield
+ }
+ default {
+ scf.yield
+ }
+ return
+}
+
+// CHECK-LABEL: @switch_yield
+func.func @switch_yield(%arg1 : index) -> i32 {
+ // CHECK: %[[VAR:.*]] = spirv.Variable : !spirv.ptr<i32, Function>
+ // CHECK: spirv.mlir.selection {
+ // CHECK-NEXT: spirv.Switch {{%.*}} : i32, [
+ // CHECK-NEXT: default: [[DEFAULT:\^.*]],
+ // CHECK-NEXT: 2: [[CASE2:\^.*]],
+ // CHECK-NEXT: 5: [[CASE5:\^.*]]
+ // CHECK-NEXT: ]
+ // CHECK-NEXT: [[CASE2]]:
+ // CHECK: %[[C10:.*]] = spirv.Constant 10 : i32
+ // CHECK: spirv.Store "Function" %[[VAR]], %[[C10]] : i32
+ // CHECK: spirv.Branch [[MERGE:\^.*]]
+ // CHECK-NEXT: [[CASE5]]:
+ // CHECK: %[[C20:.*]] = spirv.Constant 20 : i32
+ // CHECK: spirv.Store "Function" %[[VAR]], %[[C20]] : i32
+ // CHECK: spirv.Branch [[MERGE]]
+ // CHECK-NEXT: [[DEFAULT]]:
+ // CHECK: %[[C30:.*]] = spirv.Constant 30 : i32
+ // CHECK: spirv.Store "Function" %[[VAR]], %[[C30]] : i32
+ // CHECK: spirv.Branch [[MERGE]]
+ // CHECK-NEXT: [[MERGE]]:
+ // CHECK-NEXT: spirv.mlir.merge
+ // CHECK-NEXT: }
+ // CHECK: %[[OUT:.*]] = spirv.Load "Function" %[[VAR]] : i32
+ // CHECK: spirv.ReturnValue %[[OUT]] : i32
+ %0 = scf.index_switch %arg1 -> i32
+ case 2 {
+ %c10 = arith.constant 10 : i32
+ scf.yield %c10 : i32
+ }
+ case 5 {
+ %c20 = arith.constant 20 : i32
+ scf.yield %c20 : i32
+ }
+ default {
+ %c30 = arith.constant 30 : i32
+ scf.yield %c30 : i32
+ }
+ return %0 : i32
+}
+
+// CHECK-LABEL: @switch_selection_control
+func.func @switch_selection_control(%arg0 : memref<10xf32, #spirv.storage_class<StorageBuffer>>, %arg1 : index) {
+ %value = arith.constant 0.0 : f32
+ %i = arith.constant 0 : index
+ // CHECK: spirv.mlir.selection control(Flatten) {
+ scf.index_switch %arg1 {spirv.selection_control = #spirv.selection_control<Flatten>}
+ case 2 {
+ memref.store %value, %arg0[%i] : memref<10xf32, #spirv.storage_class<StorageBuffer>>
+ scf.yield
+ }
+ default {
+ scf.yield
+ }
+ return
+}
+
+} // end module
More information about the Mlir-commits
mailing list