[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