[Mlir-commits] [mlir] [mlir][SPIR-V] Preserve poison mask in scalarized vector.shuffle lowering (PR #214686)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Aug 7 03:42:43 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
---
Full diff: https://github.com/llvm/llvm-project/pull/214686.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/VectorToSPIRV/VectorToSPIRV.cpp (+5-1)
- (modified) mlir/test/Conversion/VectorToSPIRV/vector-to-spirv.mlir (+25)
``````````diff
diff --git a/mlir/lib/Conversion/VectorToSPIRV/VectorToSPIRV.cpp b/mlir/lib/Conversion/VectorToSPIRV/VectorToSPIRV.cpp
index 78693e924c4d9..62fe707b921bd 100644
--- a/mlir/lib/Conversion/VectorToSPIRV/VectorToSPIRV.cpp
+++ b/mlir/lib/Conversion/VectorToSPIRV/VectorToSPIRV.cpp
@@ -589,8 +589,12 @@ struct VectorShuffleOpConvert final
// When at least one of the operands or the result becomes a scalar after
// type conversion for SPIR-V, extract all the required elements and
// construct the result vector.
- auto getElementAtIdx = [&rewriter, loc = shuffleOp.getLoc()](
+ Type scalarType = oldV1Type.getElementType();
+ auto getElementAtIdx = [&rewriter, loc = shuffleOp.getLoc(), scalarType](
Value scalarOrVec, int32_t idx) -> Value {
+ if (idx == vector::ShuffleOp::kPoisonIndex)
+ return spirv::UndefOp::create(rewriter, loc, scalarType);
+
if (auto vecTy = dyn_cast<VectorType>(scalarOrVec.getType()))
return spirv::CompositeExtractOp::create(rewriter, loc, scalarOrVec,
idx);
diff --git a/mlir/test/Conversion/VectorToSPIRV/vector-to-spirv.mlir b/mlir/test/Conversion/VectorToSPIRV/vector-to-spirv.mlir
index f904dd9d35c37..b31e6842ef739 100644
--- a/mlir/test/Conversion/VectorToSPIRV/vector-to-spirv.mlir
+++ b/mlir/test/Conversion/VectorToSPIRV/vector-to-spirv.mlir
@@ -604,6 +604,31 @@ func.func @shuffle(%v0 : vector<4xi32>, %v1: vector<4xi32>) -> vector<4xi32> {
// -----
+// CHECK-LABEL: func @shuffle_poison_scalar_operand
+// CHECK-SAME: %[[ARG0:.+]]: vector<4xi32>, %[[ARG1:.+]]: vector<4xi32>
+// CHECK: %[[UNDEF:.+]] = spirv.Undef : i32
+// CHECK: %[[RES:.+]] = builtin.unrealized_conversion_cast %[[UNDEF]] : i32 to vector<1xi32>
+// CHECK: return %[[RES]] : vector<1xi32>
+func.func @shuffle_poison_scalar_operand(%v0 : vector<4xi32>, %v1: vector<4xi32>) -> vector<1xi32> {
+ %shuffle = vector.shuffle %v0, %v1 [-1] : vector<4xi32>, vector<4xi32>
+ return %shuffle : vector<1xi32>
+}
+
+// -----
+
+// CHECK-LABEL: func @shuffle_poison_scalar_result
+// CHECK-SAME: %[[ARG0:.+]]: vector<1xi32>, %[[ARG1:.+]]: vector<1xi32>
+// CHECK-DAG: %[[V1:.+]] = builtin.unrealized_conversion_cast %[[ARG1]] : vector<1xi32> to i32
+// CHECK-DAG: %[[UNDEF:.+]] = spirv.Undef : i32
+// CHECK: %[[RES:.+]] = spirv.CompositeConstruct %[[UNDEF]], %[[V1]] : (i32, i32) -> vector<2xi32>
+// CHECK: return %[[RES]]
+func.func @shuffle_poison_scalar_result(%v0 : vector<1xi32>, %v1: vector<1xi32>) -> vector<2xi32> {
+ %shuffle = vector.shuffle %v0, %v1 [-1, 1] : vector<1xi32>, vector<1xi32>
+ return %shuffle : vector<2xi32>
+}
+
+// -----
+
// CHECK-LABEL: func @interleave
// CHECK-SAME: (%[[ARG0:.+]]: vector<2xf32>, %[[ARG1:.+]]: vector<2xf32>)
// CHECK: %[[SHUFFLE:.*]] = spirv.VectorShuffle [0 : i32, 2 : i32, 1 : i32, 3 : i32] %[[ARG0]], %[[ARG1]] : vector<2xf32>, vector<2xf32> -> vector<4xf32>
``````````
</details>
https://github.com/llvm/llvm-project/pull/214686
More information about the Mlir-commits
mailing list