[llvm] [SPIRV][Matrix] Legalize ICmp and Fcmp (PR #218467)
Farzon Lotfi via llvm-commits
llvm-commits at lists.llvm.org
Mon Aug 24 21:03:51 PDT 2026
https://github.com/farzonl updated https://github.com/llvm/llvm-project/pull/218467
>From d29741df0d9df4ea913b1e5c0d9dcdb9db96c393 Mon Sep 17 00:00:00 2001
From: Farzon Lotfi <farzon at farzon.org>
Date: Mon, 24 Aug 2026 13:16:51 -0400
Subject: [PATCH 1/2] [SPIRV][Matrix] Legalize ICmp and Fcmp
fixes https://github.com/llvm/llvm-project/issues/218444
Simple fix we just need to apply the same fewerElementsIf change we have
done for other Global opcodes. This is going to feel like wack-a-mole
for a bit but the fix will be generally the same everytime.
---
llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp | 12 +++++--
.../legalization/matrix-wide-vector-shader.ll | 35 +++++++++++++++++++
2 files changed, 44 insertions(+), 3 deletions(-)
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index e45321640458b..e32d720aed8f7 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -445,6 +445,9 @@ SPIRVLegalizerInfo::SPIRVLegalizerInfo(const SPIRVSubtarget &ST) {
.unsupportedIf(LegalityPredicates::any(
all(typeIs(0, p9), typeInSet(1, allPtrs), typeIsNot(1, p9)),
all(typeInSet(0, allPtrs), typeIsNot(0, p9), typeIs(1, p9))))
+ .fewerElementsIf(vectorElementCountIsGreaterThan(1, MaxVectorSize),
+ LegalizeMutations::changeElementCountTo(
+ 1, ElementCount::getFixed(MaxVectorSize)))
.legalIf([IsExtendedInts](const LegalityQuery &Query) {
const LLT Ty = Query.Types[1];
return IsExtendedInts && Ty.isValid() && !Ty.isPointerOrPointerVector();
@@ -452,9 +455,12 @@ SPIRVLegalizerInfo::SPIRVLegalizerInfo(const SPIRVSubtarget &ST) {
.customIf(all(typeInSet(0, allBoolScalarsAndVectors),
typeInSet(1, allPtrsScalarsAndVectors)));
- getActionDefinitionsBuilder(G_FCMP).legalIf(
- all(typeInSet(0, allBoolScalarsAndVectors),
- typeInSet(1, allFloatScalarsAndVectors)));
+ getActionDefinitionsBuilder(G_FCMP)
+ .fewerElementsIf(vectorElementCountIsGreaterThan(1, MaxVectorSize),
+ LegalizeMutations::changeElementCountTo(
+ 1, ElementCount::getFixed(MaxVectorSize)))
+ .legalIf(all(typeInSet(0, allBoolScalarsAndVectors),
+ typeInSet(1, allFloatScalarsAndVectors)));
getActionDefinitionsBuilder({G_ATOMICRMW_OR, G_ATOMICRMW_ADD, G_ATOMICRMW_AND,
G_ATOMICRMW_MAX, G_ATOMICRMW_MIN,
diff --git a/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll b/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll
index 7f9591fe27755..73a7259f9a91f 100644
--- a/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll
+++ b/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll
@@ -4,12 +4,19 @@
@Ints9 = internal addrspace(10) global [9 x i32] poison
@Bools9 = internal addrspace(10) global [9 x i32] poison
@Ints12 = internal addrspace(10) global [12 x i32] poison
+ at Ints12B = internal addrspace(10) global [12 x i32] poison
@Bools12 = internal addrspace(10) global [12 x i32] poison
+ at Floats12 = internal addrspace(10) global [12 x float] poison
+ at Floats12B = internal addrspace(10) global [12 x float] poison
@Ints16 = internal addrspace(10) global [16 x i32] poison
@Bools16 = internal addrspace(10) global [16 x i32] poison
+; CHECK-DAG: %[[Bool:[0-9]+]] = OpTypeBool
; CHECK-DAG: %[[Int32:[0-9]+]] = OpTypeInt 32 0
; CHECK-DAG: %[[Vec4Int32:[0-9]+]] = OpTypeVector %[[Int32]] 4
+; CHECK-DAG: %[[Vec4Bool:[0-9]+]] = OpTypeVector %[[Bool]] 4
+; CHECK-DAG: %[[Float32:[0-9]+]] = OpTypeFloat 32
+; CHECK-DAG: %[[Vec4Float32:[0-9]+]] = OpTypeVector %[[Float32]] 4
; No vector wider than 4 lanes is ever materialized for shader targets.
; CHECK-NOT: OpTypeVector %[[Int32]] 8
@@ -204,6 +211,32 @@ define internal void @narrow_16elem_sext() {
ret void
}
+;--- G_ICMP/G_FCMP: split the flattened matrix into 4-lane comparisons ---
+
+; CHECK-LABEL: ; -- Begin function icmp_12elem
+; CHECK-COUNT-3: OpIEqual %[[Vec4Bool]] %{{[0-9]+}} %{{[0-9]+}}
+; CHECK-NOT: OpIEqual
+define internal void @icmp_12elem() {
+ %a = load <12 x i32>, ptr addrspace(10) @Ints12
+ %b = load <12 x i32>, ptr addrspace(10) @Ints12B
+ %cmp = icmp eq <12 x i32> %a, %b
+ %ext = zext <12 x i1> %cmp to <12 x i32>
+ store <12 x i32> %ext, ptr addrspace(10) @Bools12
+ ret void
+}
+
+; CHECK-LABEL: ; -- Begin function fcmp_12elem
+; CHECK-COUNT-3: OpFOrdEqual %[[Vec4Bool]] %{{[0-9]+}} %{{[0-9]+}}
+; CHECK-NOT: OpFOrdEqual
+define internal void @fcmp_12elem() {
+ %a = load <12 x float>, ptr addrspace(10) @Floats12
+ %b = load <12 x float>, ptr addrspace(10) @Floats12B
+ %cmp = fcmp oeq <12 x float> %a, %b
+ %ext = zext <12 x i1> %cmp to <12 x i32>
+ store <12 x i32> %ext, ptr addrspace(10) @Bools12
+ ret void
+}
+
define void @main() #0 {
call void @copy_bool3x3()
call void @copy_bool_12elem()
@@ -216,6 +249,8 @@ define void @main() #0 {
call void @bool4x4_sext()
call void @narrow_12elem_zext()
call void @narrow_16elem_sext()
+ call void @icmp_12elem()
+ call void @fcmp_12elem()
ret void
}
>From b7605792e6fad52a65ab142b8cf69358ca3bce47 Mon Sep 17 00:00:00 2001
From: Farzon Lotfi <farzon at farzon.org>
Date: Tue, 25 Aug 2026 00:03:34 -0400
Subject: [PATCH 2/2] add 3x3 matrix cases. Make the legalization clearer with
more check lines
---
.../legalization/matrix-wide-vector-shader.ll | 69 ++++++++++++++++++-
1 file changed, 67 insertions(+), 2 deletions(-)
diff --git a/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll b/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll
index 73a7259f9a91f..0d054b6ef3338 100644
--- a/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll
+++ b/llvm/test/CodeGen/SPIRV/legalization/matrix-wide-vector-shader.ll
@@ -2,7 +2,10 @@
; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3 %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
@Ints9 = internal addrspace(10) global [9 x i32] poison
+ at Ints9B = internal addrspace(10) global [9 x i32] poison
@Bools9 = internal addrspace(10) global [9 x i32] poison
+ at Floats9 = internal addrspace(10) global [9 x float] poison
+ at Floats9B = internal addrspace(10) global [9 x float] poison
@Ints12 = internal addrspace(10) global [12 x i32] poison
@Ints12B = internal addrspace(10) global [12 x i32] poison
@Bools12 = internal addrspace(10) global [12 x i32] poison
@@ -213,8 +216,38 @@ define internal void @narrow_16elem_sext() {
;--- G_ICMP/G_FCMP: split the flattened matrix into 4-lane comparisons ---
+; CHECK-LABEL: ; -- Begin function icmp_3x3
+; CHECK-COUNT-8: OpLoad %[[Int32]]
+; CHECK: %[[ICMP_AS:[0-9]+]] = OpLoad %[[Int32]]
+; CHECK-COUNT-8: OpLoad %[[Int32]]
+; CHECK: %[[ICMP_BS:[0-9]+]] = OpLoad %[[Int32]]
+; CHECK: %[[ICMP_A0:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_A1:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_B0:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_B1:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: OpIEqual %[[Vec4Bool]] %[[ICMP_A0]] %[[ICMP_B0]]
+; CHECK: OpIEqual %[[Vec4Bool]] %[[ICMP_A1]] %[[ICMP_B1]]
+; CHECK: OpIEqual %[[Bool]] %[[ICMP_AS]] %[[ICMP_BS]]
+; CHECK-NOT: OpIEqual
+define internal void @icmp_3x3() {
+ %a = load <9 x i32>, ptr addrspace(10) @Ints9
+ %b = load <9 x i32>, ptr addrspace(10) @Ints9B
+ %cmp = icmp eq <9 x i32> %a, %b
+ %ext = zext <9 x i1> %cmp to <9 x i32>
+ store <9 x i32> %ext, ptr addrspace(10) @Bools9
+ ret void
+}
+
; CHECK-LABEL: ; -- Begin function icmp_12elem
-; CHECK-COUNT-3: OpIEqual %[[Vec4Bool]] %{{[0-9]+}} %{{[0-9]+}}
+; CHECK: %[[ICMP_A0:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_A1:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_A2:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_B0:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_B1:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: %[[ICMP_B2:[0-9]+]] = OpCompositeConstruct %[[Vec4Int32]]
+; CHECK: OpIEqual %[[Vec4Bool]] %[[ICMP_A0]] %[[ICMP_B0]]
+; CHECK: OpIEqual %[[Vec4Bool]] %[[ICMP_A1]] %[[ICMP_B1]]
+; CHECK: OpIEqual %[[Vec4Bool]] %[[ICMP_A2]] %[[ICMP_B2]]
; CHECK-NOT: OpIEqual
define internal void @icmp_12elem() {
%a = load <12 x i32>, ptr addrspace(10) @Ints12
@@ -225,8 +258,38 @@ define internal void @icmp_12elem() {
ret void
}
+; CHECK-LABEL: ; -- Begin function fcmp_3x3
+; CHECK-COUNT-8: OpLoad %[[Float32]]
+; CHECK: %[[FCMP_AS:[0-9]+]] = OpLoad %[[Float32]]
+; CHECK-COUNT-8: OpLoad %[[Float32]]
+; CHECK: %[[FCMP_BS:[0-9]+]] = OpLoad %[[Float32]]
+; CHECK: %[[FCMP_A0:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_A1:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_B0:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_B1:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: OpFOrdEqual %[[Vec4Bool]] %[[FCMP_A0]] %[[FCMP_B0]]
+; CHECK: OpFOrdEqual %[[Vec4Bool]] %[[FCMP_A1]] %[[FCMP_B1]]
+; CHECK: OpFOrdEqual %[[Bool]] %[[FCMP_AS]] %[[FCMP_BS]]
+; CHECK-NOT: OpFOrdEqual
+define internal void @fcmp_3x3() {
+ %a = load <9 x float>, ptr addrspace(10) @Floats9
+ %b = load <9 x float>, ptr addrspace(10) @Floats9B
+ %cmp = fcmp oeq <9 x float> %a, %b
+ %ext = zext <9 x i1> %cmp to <9 x i32>
+ store <9 x i32> %ext, ptr addrspace(10) @Bools9
+ ret void
+}
+
; CHECK-LABEL: ; -- Begin function fcmp_12elem
-; CHECK-COUNT-3: OpFOrdEqual %[[Vec4Bool]] %{{[0-9]+}} %{{[0-9]+}}
+; CHECK: %[[FCMP_A0:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_A1:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_A2:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_B0:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_B1:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: %[[FCMP_B2:[0-9]+]] = OpCompositeConstruct %[[Vec4Float32]]
+; CHECK: OpFOrdEqual %[[Vec4Bool]] %[[FCMP_A0]] %[[FCMP_B0]]
+; CHECK: OpFOrdEqual %[[Vec4Bool]] %[[FCMP_A1]] %[[FCMP_B1]]
+; CHECK: OpFOrdEqual %[[Vec4Bool]] %[[FCMP_A2]] %[[FCMP_B2]]
; CHECK-NOT: OpFOrdEqual
define internal void @fcmp_12elem() {
%a = load <12 x float>, ptr addrspace(10) @Floats12
@@ -249,7 +312,9 @@ define void @main() #0 {
call void @bool4x4_sext()
call void @narrow_12elem_zext()
call void @narrow_16elem_sext()
+ call void @icmp_3x3()
call void @icmp_12elem()
+ call void @fcmp_3x3()
call void @fcmp_12elem()
ret void
}
More information about the llvm-commits
mailing list