[llvm] [SLP] Don't truncate compare constants that change value at the demoted width (PR #209030)

via llvm-commits llvm-commits at lists.llvm.org
Sun Jul 12 10:34:03 PDT 2026


https://github.com/aokblast updated https://github.com/llvm/llvm-project/pull/209030

>From 03af1bf70067be1323deeed00d448d24780da57a Mon Sep 17 00:00:00 2001
From: ShengYi Hung <aokblast at FreeBSD.org>
Date: Mon, 13 Jul 2026 00:26:12 +0800
Subject: [PATCH 1/2] [SLP] Don't truncate compare constants that change value
 at the demoted width

SLP preserves the original bit width when narrowing compare operands,
but doesn't always account for the required size convertion in LLVM IR.
This can produce incorrect compare constants after truncation.

Fix this by using getSignificantBits() for sign-extended operands,
forbidding truncation for signed predicates on zero-extended operands,
and checking the correct operand in the LBW > RBW case.
---
 .../Transforms/Vectorize/SLPVectorizer.cpp    | 61 ++++++++++++----
 .../minbitwidth-icmp-signed-const-trunc.ll    | 72 +++++++++++++++++++
 2 files changed, 120 insertions(+), 13 deletions(-)
 create mode 100644 llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll

diff --git a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
index 48fb4beba6935..2eff5519fac1a 100644
--- a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
+++ b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
@@ -24060,19 +24060,54 @@ Value *BoUpSLP::vectorizeTree(TreeEntry *E) {
         const unsigned RBW = cast<VectorType>(R->getType())
                                  ->getElementType()
                                  ->getIntegerBitWidth();
-        if ((LBW < RBW && (!allConstant(E->getOperand(1)) ||
-                           any_of(
-                               E->getOperand(1),
-                               [&](Value *V) {
-                                 auto *CI = dyn_cast<ConstantInt>(V);
-                                 return !CI ||
-                                        CI->getValue().getActiveBits() > LBW;
-                               }))) ||
-            (LBW > RBW && allConstant(E->getOperand(0)) &&
-             all_of(E->getOperand(1), [&](Value *V) {
-               auto *CI = dyn_cast<ConstantInt>(V);
-               return CI && CI->getValue().getActiveBits() <= RBW;
-             }))) {
+        // Preserve the original bits whenever possible. However, for
+        // icmp signed_lhs, rhs, we must extend lhs to zext(signed_lhs)
+        // instead of truncate rhs to preserve the correct comparison
+        // result. Track whether the operation is signed and force the
+        // extension when it is.
+        auto NarrowOperandNonNeg = [&](unsigned OpIdx, unsigned BW) {
+          return all_of(E->getOperand(OpIdx), [&](Value *V) {
+            if (isa<PoisonValue>(V))
+              return true;
+            unsigned OrigBW = V->getType()->getScalarSizeInBits();
+            return MaskedValueIsZero(V, APInt::getOneBitSet(OrigBW, BW - 1),
+                                     SimplifyQuery(*DL));
+          });
+        };
+        auto ConstFitsNarrow = [&](Value *V, unsigned BW, bool OpSigned,
+                                   bool NarrowNonNeg) {
+          auto *CI = dyn_cast<ConstantInt>(V);
+          if (!CI)
+            return false;
+          if (OpSigned)
+            return CI->getValue().getSignificantBits() <= BW;
+          if (!cast<CmpInst>(VL0)->isSigned())
+            return CI->getValue().getActiveBits() <= BW;
+          return NarrowNonNeg && CI->getValue().getSignificantBits() <= BW;
+        };
+        bool ExpandNarrowOperand;
+        if (LBW < RBW) {
+          bool OpSigned = GetOperandSignedness(0);
+          bool NarrowNonNeg = !OpSigned && cast<CmpInst>(VL0)->isSigned() &&
+                              NarrowOperandNonNeg(0, LBW);
+          ExpandNarrowOperand =
+              !allConstant(E->getOperand(1)) ||
+              any_of(E->getOperand(1), [&](Value *V) {
+                return !ConstFitsNarrow(V, LBW, OpSigned, NarrowNonNeg);
+              });
+        } else if (LBW > RBW) {
+          bool OpSigned = GetOperandSignedness(1);
+          bool NarrowNonNeg = !OpSigned && cast<CmpInst>(VL0)->isSigned() &&
+                              NarrowOperandNonNeg(1, RBW);
+          ExpandNarrowOperand =
+              allConstant(E->getOperand(0)) &&
+              all_of(E->getOperand(0), [&](Value *V) {
+                return ConstFitsNarrow(V, RBW, OpSigned, NarrowNonNeg);
+              });
+        } else {
+          ExpandNarrowOperand = false;
+        }
+        if (ExpandNarrowOperand) {
           Type *CastTy = R->getType();
           L = Builder.CreateIntCast(L, CastTy, GetOperandSignedness(0));
         } else {
diff --git a/llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll b/llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll
new file mode 100644
index 0000000000000..3acec2fae1711
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll
@@ -0,0 +1,72 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 5
+; RUN: opt -passes=slp-vectorizer -S -mtriple=x86_64-unknown-linux-gnu < %s | FileCheck %s
+
+; The xor values are sign-extends from i16, so the xor tree is demoted to
+; i16. The compare constant 58593 fits in 16 unsigned bits but not in 16
+; signed bits, so it must not be truncated for the signed predicate:
+; truncating it to i16 gives -6943 and inverts the compare result
+; (sext(i16) slt 58593 is always true). The compare must stay in i32.
+
+define i1 @test(i16 %a, i16 %c) {
+; CHECK-LABEL: define i1 @test(
+; CHECK-SAME: i16 [[A:%.*]], i16 [[C:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[CONV:%.*]] = sext i16 [[A]] to i32
+; CHECK-NEXT:    [[OR:%.*]] = or i32 [[CONV]], 5
+; CHECK-NEXT:    [[TMP0:%.*]] = insertelement <8 x i16> poison, i16 [[C]], i32 0
+; CHECK-NEXT:    [[TMP1:%.*]] = shufflevector <8 x i16> [[TMP0]], <8 x i16> poison, <8 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP2:%.*]] = add <8 x i16> [[TMP1]], <i16 1, i16 2, i16 3, i16 4, i16 5, i16 6, i16 7, i16 8>
+; CHECK-NEXT:    [[TMP3:%.*]] = trunc i32 [[OR]] to i16
+; CHECK-NEXT:    [[TMP4:%.*]] = insertelement <8 x i16> poison, i16 [[TMP3]], i32 0
+; CHECK-NEXT:    [[TMP5:%.*]] = shufflevector <8 x i16> [[TMP4]], <8 x i16> poison, <8 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP6:%.*]] = xor <8 x i16> [[TMP5]], [[TMP2]]
+; CHECK-NEXT:    [[TMP7:%.*]] = sext <8 x i16> [[TMP6]] to <8 x i32>
+; CHECK-NEXT:    [[TMP8:%.*]] = icmp slt <8 x i32> [[TMP7]], splat (i32 58593)
+; CHECK-NEXT:    [[TMP9:%.*]] = freeze <8 x i1> [[TMP8]]
+; CHECK-NEXT:    [[TMP10:%.*]] = call i1 @llvm.vector.reduce.and.v8i1(<8 x i1> [[TMP9]])
+; CHECK-NEXT:    ret i1 [[TMP10]]
+;
+entry:
+  %conv = sext i16 %a to i32
+  %or = or i32 %conv, 5
+  %inc1 = add i16 %c, 1
+  %sext1 = sext i16 %inc1 to i32
+  %xor1 = xor i32 %or, %sext1
+  %cmp1 = icmp slt i32 %xor1, 58593
+  %inc2 = add i16 %c, 2
+  %sext2 = sext i16 %inc2 to i32
+  %xor2 = xor i32 %or, %sext2
+  %cmp2 = icmp slt i32 %xor2, 58593
+  %inc3 = add i16 %c, 3
+  %sext3 = sext i16 %inc3 to i32
+  %xor3 = xor i32 %or, %sext3
+  %cmp3 = icmp slt i32 %xor3, 58593
+  %inc4 = add i16 %c, 4
+  %sext4 = sext i16 %inc4 to i32
+  %xor4 = xor i32 %or, %sext4
+  %cmp4 = icmp slt i32 %xor4, 58593
+  %inc5 = add i16 %c, 5
+  %sext5 = sext i16 %inc5 to i32
+  %xor5 = xor i32 %or, %sext5
+  %cmp5 = icmp slt i32 %xor5, 58593
+  %inc6 = add i16 %c, 6
+  %sext6 = sext i16 %inc6 to i32
+  %xor6 = xor i32 %or, %sext6
+  %cmp6 = icmp slt i32 %xor6, 58593
+  %inc7 = add i16 %c, 7
+  %sext7 = sext i16 %inc7 to i32
+  %xor7 = xor i32 %or, %sext7
+  %cmp7 = icmp slt i32 %xor7, 58593
+  %inc8 = add i16 %c, 8
+  %sext8 = sext i16 %inc8 to i32
+  %xor8 = xor i32 %or, %sext8
+  %cmp8 = icmp slt i32 %xor8, 58593
+  %and1 = select i1 %cmp1, i1 %cmp2, i1 false
+  %and2 = select i1 %and1, i1 %cmp3, i1 false
+  %and3 = select i1 %and2, i1 %cmp4, i1 false
+  %and4 = select i1 %and3, i1 %cmp5, i1 false
+  %and5 = select i1 %and4, i1 %cmp6, i1 false
+  %and6 = select i1 %and5, i1 %cmp7, i1 false
+  %and7 = select i1 %and6, i1 %cmp8, i1 false
+  ret i1 %and7
+}

>From 780ac308849504d1527ca7cf44e1115696c31bab Mon Sep 17 00:00:00 2001
From: ShengYi Hung <aokblast at FreeBSD.org>
Date: Mon, 13 Jul 2026 00:26:12 +0800
Subject: [PATCH 2/2] [SLP] Don't truncate compare constants that change value
 at the demoted width

SLP preserves the original bit width when narrowing compare operands,
but doesn't always account for the required size convertion in LLVM IR.
This can produce incorrect compare constants after truncation.

Fix this by using getSignificantBits() for sign-extended operands,
forbidding truncation for signed predicates on zero-extended operands,
and checking the correct operand in the LBW > RBW case.
---
 .../Transforms/Vectorize/SLPVectorizer.cpp    | 61 ++++++++++++----
 .../minbitwidth-icmp-signed-const-trunc.ll    | 72 +++++++++++++++++++
 2 files changed, 120 insertions(+), 13 deletions(-)
 create mode 100644 llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll

diff --git a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
index 48fb4beba6935..2eff5519fac1a 100644
--- a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
+++ b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
@@ -24060,19 +24060,54 @@ Value *BoUpSLP::vectorizeTree(TreeEntry *E) {
         const unsigned RBW = cast<VectorType>(R->getType())
                                  ->getElementType()
                                  ->getIntegerBitWidth();
-        if ((LBW < RBW && (!allConstant(E->getOperand(1)) ||
-                           any_of(
-                               E->getOperand(1),
-                               [&](Value *V) {
-                                 auto *CI = dyn_cast<ConstantInt>(V);
-                                 return !CI ||
-                                        CI->getValue().getActiveBits() > LBW;
-                               }))) ||
-            (LBW > RBW && allConstant(E->getOperand(0)) &&
-             all_of(E->getOperand(1), [&](Value *V) {
-               auto *CI = dyn_cast<ConstantInt>(V);
-               return CI && CI->getValue().getActiveBits() <= RBW;
-             }))) {
+        // Preserve the original bits whenever possible. However, for
+        // icmp signed_lhs, rhs, we must extend lhs to zext(signed_lhs)
+        // instead of truncate rhs to preserve the correct comparison
+        // result. Track whether the operation is signed and force the
+        // extension when it is.
+        auto NarrowOperandNonNeg = [&](unsigned OpIdx, unsigned BW) {
+          return all_of(E->getOperand(OpIdx), [&](Value *V) {
+            if (isa<PoisonValue>(V))
+              return true;
+            unsigned OrigBW = V->getType()->getScalarSizeInBits();
+            return MaskedValueIsZero(V, APInt::getOneBitSet(OrigBW, BW - 1),
+                                     SimplifyQuery(*DL));
+          });
+        };
+        auto ConstFitsNarrow = [&](Value *V, unsigned BW, bool OpSigned,
+                                   bool NarrowNonNeg) {
+          auto *CI = dyn_cast<ConstantInt>(V);
+          if (!CI)
+            return false;
+          if (OpSigned)
+            return CI->getValue().getSignificantBits() <= BW;
+          if (!cast<CmpInst>(VL0)->isSigned())
+            return CI->getValue().getActiveBits() <= BW;
+          return NarrowNonNeg && CI->getValue().getSignificantBits() <= BW;
+        };
+        bool ExpandNarrowOperand;
+        if (LBW < RBW) {
+          bool OpSigned = GetOperandSignedness(0);
+          bool NarrowNonNeg = !OpSigned && cast<CmpInst>(VL0)->isSigned() &&
+                              NarrowOperandNonNeg(0, LBW);
+          ExpandNarrowOperand =
+              !allConstant(E->getOperand(1)) ||
+              any_of(E->getOperand(1), [&](Value *V) {
+                return !ConstFitsNarrow(V, LBW, OpSigned, NarrowNonNeg);
+              });
+        } else if (LBW > RBW) {
+          bool OpSigned = GetOperandSignedness(1);
+          bool NarrowNonNeg = !OpSigned && cast<CmpInst>(VL0)->isSigned() &&
+                              NarrowOperandNonNeg(1, RBW);
+          ExpandNarrowOperand =
+              allConstant(E->getOperand(0)) &&
+              all_of(E->getOperand(0), [&](Value *V) {
+                return ConstFitsNarrow(V, RBW, OpSigned, NarrowNonNeg);
+              });
+        } else {
+          ExpandNarrowOperand = false;
+        }
+        if (ExpandNarrowOperand) {
           Type *CastTy = R->getType();
           L = Builder.CreateIntCast(L, CastTy, GetOperandSignedness(0));
         } else {
diff --git a/llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll b/llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll
new file mode 100644
index 0000000000000..c8972802bbf51
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/X86/minbitwidth-icmp-signed-const-trunc.ll
@@ -0,0 +1,72 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 5
+; RUN: opt -passes=slp-vectorizer -S -mtriple=x86_64-unknown-linux-gnu < %s | FileCheck %s
+
+; The xor values are sign-extends from i16, so the xor tree is demoted to
+; i16. The compare constant 58593 fits in 16 unsigned bits but not in 16
+; signed bits, so it must not be truncated for the signed predicate:
+; truncating it to i16 gives -6943 and inverts the compare result
+; (sext(i16) slt 58593 is always true). The compare must stay in i32.
+
+define i1 @test(i16 %a, i16 %c) {
+; CHECK-LABEL: define i1 @test(
+; CHECK-SAME: i16 [[A:%.*]], i16 [[C:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[CONV:%.*]] = sext i16 [[A]] to i32
+; CHECK-NEXT:    [[OR:%.*]] = or i32 [[CONV]], 5
+; CHECK-NEXT:    [[TMP0:%.*]] = insertelement <8 x i16> poison, i16 [[C]], i64 0
+; CHECK-NEXT:    [[TMP1:%.*]] = shufflevector <8 x i16> [[TMP0]], <8 x i16> poison, <8 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP2:%.*]] = add <8 x i16> [[TMP1]], <i16 1, i16 2, i16 3, i16 4, i16 5, i16 6, i16 7, i16 8>
+; CHECK-NEXT:    [[TMP3:%.*]] = trunc i32 [[OR]] to i16
+; CHECK-NEXT:    [[TMP4:%.*]] = insertelement <8 x i16> poison, i16 [[TMP3]], i64 0
+; CHECK-NEXT:    [[TMP5:%.*]] = shufflevector <8 x i16> [[TMP4]], <8 x i16> poison, <8 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP6:%.*]] = xor <8 x i16> [[TMP5]], [[TMP2]]
+; CHECK-NEXT:    [[TMP7:%.*]] = sext <8 x i16> [[TMP6]] to <8 x i32>
+; CHECK-NEXT:    [[TMP8:%.*]] = icmp slt <8 x i32> [[TMP7]], splat (i32 58593)
+; CHECK-NEXT:    [[TMP9:%.*]] = freeze <8 x i1> [[TMP8]]
+; CHECK-NEXT:    [[TMP10:%.*]] = call i1 @llvm.vector.reduce.and.v8i1(<8 x i1> [[TMP9]])
+; CHECK-NEXT:    ret i1 [[TMP10]]
+;
+entry:
+  %conv = sext i16 %a to i32
+  %or = or i32 %conv, 5
+  %inc1 = add i16 %c, 1
+  %sext1 = sext i16 %inc1 to i32
+  %xor1 = xor i32 %or, %sext1
+  %cmp1 = icmp slt i32 %xor1, 58593
+  %inc2 = add i16 %c, 2
+  %sext2 = sext i16 %inc2 to i32
+  %xor2 = xor i32 %or, %sext2
+  %cmp2 = icmp slt i32 %xor2, 58593
+  %inc3 = add i16 %c, 3
+  %sext3 = sext i16 %inc3 to i32
+  %xor3 = xor i32 %or, %sext3
+  %cmp3 = icmp slt i32 %xor3, 58593
+  %inc4 = add i16 %c, 4
+  %sext4 = sext i16 %inc4 to i32
+  %xor4 = xor i32 %or, %sext4
+  %cmp4 = icmp slt i32 %xor4, 58593
+  %inc5 = add i16 %c, 5
+  %sext5 = sext i16 %inc5 to i32
+  %xor5 = xor i32 %or, %sext5
+  %cmp5 = icmp slt i32 %xor5, 58593
+  %inc6 = add i16 %c, 6
+  %sext6 = sext i16 %inc6 to i32
+  %xor6 = xor i32 %or, %sext6
+  %cmp6 = icmp slt i32 %xor6, 58593
+  %inc7 = add i16 %c, 7
+  %sext7 = sext i16 %inc7 to i32
+  %xor7 = xor i32 %or, %sext7
+  %cmp7 = icmp slt i32 %xor7, 58593
+  %inc8 = add i16 %c, 8
+  %sext8 = sext i16 %inc8 to i32
+  %xor8 = xor i32 %or, %sext8
+  %cmp8 = icmp slt i32 %xor8, 58593
+  %and1 = select i1 %cmp1, i1 %cmp2, i1 false
+  %and2 = select i1 %and1, i1 %cmp3, i1 false
+  %and3 = select i1 %and2, i1 %cmp4, i1 false
+  %and4 = select i1 %and3, i1 %cmp5, i1 false
+  %and5 = select i1 %and4, i1 %cmp6, i1 false
+  %and6 = select i1 %and5, i1 %cmp7, i1 false
+  %and7 = select i1 %and6, i1 %cmp8, i1 false
+  ret i1 %and7
+}



More information about the llvm-commits mailing list