[llvm] [AArch64] Fold vector select with power-of-2 bit-test to CMTST+BSP (PR #209100)

Mugundan S via llvm-commits llvm-commits at lists.llvm.org
Fri Aug 14 23:46:21 PDT 2026


https://github.com/MGN-GIT updated https://github.com/llvm/llvm-project/pull/209100

>From 036b89edcbb4a07027a42cae0eeb94a978e8034e Mon Sep 17 00:00:00 2001
From: Greenie0701 <smugundan12a at gmail.com>
Date: Mon, 13 Jul 2026 12:52:50 +0530
Subject: [PATCH 1/2] [AArch64] Fold vector select with power-of-2 bit-test to
 CMTST+BSP

---
 .../Target/AArch64/AArch64ISelLowering.cpp    | 59 +++++++++++++++++++
 llvm/lib/Target/AArch64/AArch64InstrInfo.td   | 10 +++-
 .../CodeGen/AArch64/cmtst-select-pow2-mask.ll | 55 +++++++++++++++++
 3 files changed, 121 insertions(+), 3 deletions(-)
 create mode 100644 llvm/test/CodeGen/AArch64/cmtst-select-pow2-mask.ll

diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 46ffc9287dc62..8e6e2a0746254 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -16757,6 +16757,28 @@ static SDValue tryLowerToBSL(SDValue N, SelectionDAG &DAG) {
                            N0->getOperand(1 - i), N1->getOperand(1 - j));
     }
 
+  // Fold: or(and(xor(AArch64ISD::CMTST(X,M), allones), A),
+  // and(AArch64ISD::CMTST(X,M), B)) --> BSP(CMTST(X,M), A, B)
+  // This absorbs the NOT produced by performSETCCCombine when it folds
+  // setcc(and(X,Mask), 0, seteq) --> NOT(CMTST(X,Mask)).
+  for (int i = 1; i >= 0; --i)
+    for (int j = 1; j >= 0; --j) {
+      SDValue NotCMTST = N0->getOperand(i);
+      SDValue A = N0->getOperand(1 - i);
+      SDValue CMTST = N1->getOperand(j);
+      SDValue B = N1->getOperand(1 - j);
+
+      if (NotCMTST.getOpcode() != ISD::XOR ||
+          !ISD::isBuildVectorAllOnes(NotCMTST.getOperand(1).getNode()))
+        continue;
+      if (NotCMTST.getOperand(0) != CMTST)
+        continue;
+      if (CMTST.getOpcode() != AArch64ISD::CMTST)
+        continue;
+
+      return DAG.getNode(AArch64ISD::BSP, DL, VT, CMTST, A, B);
+    }
+
   return SDValue();
 }
 
@@ -29261,6 +29283,43 @@ static SDValue performSETCCCombine(SDNode *N,
       SplatLHSVal.isOne())
     return DAG.getSetCC(DL, VT, DAG.getConstant(0, DL, CmpVT), RHS, ISD::SETGE);
 
+  // Fold setcc(and(X, Mask), Mask/0, eq/ne) --> [not] AArch64ISD::CMTST(X,
+  // Mask) for a power of 2 splat Mask, replacing AND+CMEQ with a single CMTST.
+  // Any NOT folds away when the result feeds a BSL/BIF/BIT select.
+  if (!DCI.isBeforeLegalize() && CmpVT.isFixedLengthVector() &&
+      (Cond == ISD::SETEQ || Cond == ISD::SETNE) &&
+      LHS.getOpcode() == ISD::AND) {
+    APInt SplatVal;
+    SDValue X, MaskOp;
+    if (ISD::isConstantSplatVector(LHS.getOperand(1).getNode(), SplatVal) &&
+        SplatVal.isPowerOf2()) {
+      X = LHS.getOperand(0);
+      MaskOp = LHS.getOperand(1);
+    } else if (ISD::isConstantSplatVector(LHS.getOperand(0).getNode(),
+                                          SplatVal) &&
+               SplatVal.isPowerOf2()) {
+      X = LHS.getOperand(1);
+      MaskOp = LHS.getOperand(0);
+    }
+    if (X.getNode()) {
+      bool RHSIsZero = ISD::isBuildVectorAllZeros(RHS.getNode());
+      APInt RHSSplat;
+      bool RHSIsMask = !RHSIsZero &&
+                       ISD::isConstantSplatVector(RHS.getNode(), RHSSplat) &&
+                       RHSSplat == SplatVal;
+      if (RHSIsZero || RHSIsMask) {
+        SDValue CMTSTNode =
+            DAG.getNode(AArch64ISD::CMTST, DL, CmpVT, X, MaskOp);
+        // CMTST gives all-ones where (X & Mask) != 0, i.e. SETNE(AND, 0).
+        // Invert when the original condition is the opposite sense.
+        bool Invert = (Cond == ISD::SETEQ) ? RHSIsZero : RHSIsMask;
+        if (Invert)
+          return DAG.getNOT(DL, CMTSTNode, CmpVT);
+        return CMTSTNode;
+      }
+    }
+  }
+
   return SDValue();
 }
 
diff --git a/llvm/lib/Target/AArch64/AArch64InstrInfo.td b/llvm/lib/Target/AArch64/AArch64InstrInfo.td
index e59428f0ea33c..789a8aecbd886 100644
--- a/llvm/lib/Target/AArch64/AArch64InstrInfo.td
+++ b/llvm/lib/Target/AArch64/AArch64InstrInfo.td
@@ -953,6 +953,9 @@ def AArch64vsri : SDNode<"AArch64ISD::VSRI", SDT_AArch64vshiftinsert>;
 // element must be identical.
 def AArch64bsp: SDNode<"AArch64ISD::BSP", SDT_AArch64trivec>;
 
+// AArch64ISD::CMTST node: result is all-ones per lane where (X & Y) != 0.
+def AArch64cmtst: SDNode<"AArch64ISD::CMTST", SDT_AArch64Zip>;
+
 def AArch64cmeq : PatFrag<(ops node:$lhs, node:$rhs),
                           (setcc node:$lhs, node:$rhs, SETEQ)>;
 def AArch64cmge : PatFrag<(ops node:$lhs, node:$rhs),
@@ -981,7 +984,7 @@ def AArch64cmlez : PatFrag<(ops node:$lhs),
 def AArch64cmltz : PatFrag<(ops node:$lhs),
                            (setcc immAllZerosV, node:$lhs, SETGT)>;
 
-def AArch64cmtst : PatFrag<(ops node:$LHS, node:$RHS),
+def AArch64cmtst_frag : PatFrag<(ops node:$LHS, node:$RHS),
                            (vnot (AArch64cmeqz (and node:$LHS, node:$RHS)))>;
 
 def AArch64fcmeqz : PatFrag<(ops node:$lhs),
@@ -6161,9 +6164,10 @@ defm CMGE    : SIMDThreeSameVector<0, 0b00111, "cmge", AArch64cmge>;
 defm CMGT    : SIMDThreeSameVector<0, 0b00110, "cmgt", AArch64cmgt>;
 defm CMHI    : SIMDThreeSameVector<1, 0b00110, "cmhi", AArch64cmhi>;
 defm CMHS    : SIMDThreeSameVector<1, 0b00111, "cmhs", AArch64cmhs>;
-defm CMTST   : SIMDThreeSameVector<0, 0b10001, "cmtst", AArch64cmtst>;
+defm CMTST   : SIMDThreeSameVector<0, 0b10001, "cmtst", AArch64cmtst_frag>;
 foreach VT = [ v8i8, v16i8, v4i16, v8i16, v2i32, v4i32, v2i64 ] in {
 def : Pat<(VT (vnot (AArch64cmeqz VT:$Rn))), (!cast<Instruction>("CMTST"#VT) VT:$Rn, VT:$Rn)>;
+def : Pat<(VT (AArch64cmtst VT:$Rn, VT:$Rm)), (!cast<Instruction>("CMTST"#VT) VT:$Rn, VT:$Rm)>;
 }
 defm FABD    : SIMDThreeSameVectorFP<1,1,0b010,"fabd", int_aarch64_neon_fabd>;
 let Predicates = [HasNEON] in {
@@ -6605,7 +6609,7 @@ defm CMGE     : SIMDThreeScalarD<0, 0b00111, "cmge", AArch64cmge>;
 defm CMGT     : SIMDThreeScalarD<0, 0b00110, "cmgt", AArch64cmgt>;
 defm CMHI     : SIMDThreeScalarD<1, 0b00110, "cmhi", AArch64cmhi>;
 defm CMHS     : SIMDThreeScalarD<1, 0b00111, "cmhs", AArch64cmhs>;
-defm CMTST    : SIMDThreeScalarD<0, 0b10001, "cmtst", AArch64cmtst>;
+defm CMTST    : SIMDThreeScalarD<0, 0b10001, "cmtst", AArch64cmtst_frag>;
 defm FABD     : SIMDFPThreeScalar<1, 1, 0b010, "fabd", int_aarch64_sisd_fabd>;
 def : Pat<(v1f64 (int_aarch64_neon_fabd (v1f64 FPR64:$Rn), (v1f64 FPR64:$Rm))),
           (FABD64 FPR64:$Rn, FPR64:$Rm)>;
diff --git a/llvm/test/CodeGen/AArch64/cmtst-select-pow2-mask.ll b/llvm/test/CodeGen/AArch64/cmtst-select-pow2-mask.ll
new file mode 100644
index 0000000000000..60c837be638be
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/cmtst-select-pow2-mask.ll
@@ -0,0 +1,55 @@
+; Test: (X & Mask) == Mask, for a power-of-2 Mask, folds to CMTST
+; instead of AND+CMEQ when used as a select condition inside a BSL/BSP fold.
+; RUN: llc -mtriple=aarch64-none-linux-gnu -mattr=+neon < %s | FileCheck %s
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 5
+
+define <16 x i8> @cmtst_select_v16i8_pow2(<16 x i8> %x, <16 x i8> %y) {
+; CHECK-LABEL: cmtst_select_v16i8_pow2:
+; CHECK:         movi v2.16b, #2
+; CHECK-NEXT:    cmtst v2.16b, v0.16b, v2.16b
+; CHECK-NEXT:    bif v0.16b, v1.16b, v2.16b
+; CHECK-NEXT:    ret
+  %mask = and <16 x i8> %x, splat(i8 2)
+  %cmp  = icmp eq <16 x i8> %mask, splat(i8 2)
+  %sel  = select <16 x i1> %cmp, <16 x i8> %x, <16 x i8> %y
+  ret <16 x i8> %sel
+}
+
+define <8 x i16> @cmtst_select_v8i16_pow2(<8 x i16> %x, <8 x i16> %y) {
+; CHECK-LABEL: cmtst_select_v8i16_pow2:
+; CHECK:         movi v2.8h, #4
+; CHECK-NEXT:    cmtst v2.8h, v0.8h, v2.8h
+; CHECK-NEXT:    bif v0.16b, v1.16b, v2.16b
+; CHECK-NEXT:    ret
+  %mask = and <8 x i16> %x, splat(i16 4)
+  %cmp  = icmp eq <8 x i16> %mask, splat(i16 4)
+  %sel  = select <8 x i1> %cmp, <8 x i16> %x, <8 x i16> %y
+  ret <8 x i16> %sel
+}
+
+define <4 x i32> @cmtst_select_v4i32_pow2(<4 x i32> %x, <4 x i32> %y) {
+; CHECK-LABEL: cmtst_select_v4i32_pow2:
+; CHECK:         movi v2.4s, #8
+; CHECK-NEXT:    cmtst v2.4s, v0.4s, v2.4s
+; CHECK-NEXT:    bif v0.16b, v1.16b, v2.16b
+; CHECK-NEXT:    ret
+  %mask = and <4 x i32> %x, splat(i32 8)
+  %cmp  = icmp eq <4 x i32> %mask, splat(i32 8)
+  %sel  = select <4 x i1> %cmp, <4 x i32> %x, <4 x i32> %y
+  ret <4 x i32> %sel
+}
+
+; Negative test - non-power-of-2 mask must NOT use CMTST; must fall back to
+; AND+CMEQ.
+define <16 x i8> @no_cmtst_non_pow2(<16 x i8> %x, <16 x i8> %y) {
+; CHECK-LABEL: no_cmtst_non_pow2:
+; CHECK:         movi v2.16b, #3
+; CHECK-NEXT:    and v3.16b, v0.16b, v2.16b
+; CHECK-NEXT:    cmeq v2.16b, v3.16b, v2.16b
+; CHECK-NEXT:    bif v0.16b, v1.16b, v2.16b
+; CHECK-NEXT:    ret
+  %mask = and <16 x i8> %x, splat(i8 3)
+  %cmp  = icmp eq <16 x i8> %mask, splat(i8 3)
+  %sel  = select <16 x i1> %cmp, <16 x i8> %x, <16 x i8> %y
+  ret <16 x i8> %sel
+}

>From df03c29678e4ffae0d2114bdf2c33c4f765bb80d Mon Sep 17 00:00:00 2001
From: Mugundan S <137760120+MGN-GIT at users.noreply.github.com>
Date: Mon, 3 Aug 2026 09:19:10 +0530
Subject: [PATCH 2/2] [AArch64] Fold vector select with power-of-2 bit-test to
 CMTST+BSP

---
 .../Target/AArch64/AArch64ISelLowering.cpp    | 44 ++++---------------
 1 file changed, 9 insertions(+), 35 deletions(-)

diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 8e6e2a0746254..b873b33ab3328 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -29283,41 +29283,15 @@ static SDValue performSETCCCombine(SDNode *N,
       SplatLHSVal.isOne())
     return DAG.getSetCC(DL, VT, DAG.getConstant(0, DL, CmpVT), RHS, ISD::SETGE);
 
-  // Fold setcc(and(X, Mask), Mask/0, eq/ne) --> [not] AArch64ISD::CMTST(X,
-  // Mask) for a power of 2 splat Mask, replacing AND+CMEQ with a single CMTST.
-  // Any NOT folds away when the result feeds a BSL/BIF/BIT select.
-  if (!DCI.isBeforeLegalize() && CmpVT.isFixedLengthVector() &&
-      (Cond == ISD::SETEQ || Cond == ISD::SETNE) &&
-      LHS.getOpcode() == ISD::AND) {
-    APInt SplatVal;
-    SDValue X, MaskOp;
-    if (ISD::isConstantSplatVector(LHS.getOperand(1).getNode(), SplatVal) &&
-        SplatVal.isPowerOf2()) {
-      X = LHS.getOperand(0);
-      MaskOp = LHS.getOperand(1);
-    } else if (ISD::isConstantSplatVector(LHS.getOperand(0).getNode(),
-                                          SplatVal) &&
-               SplatVal.isPowerOf2()) {
-      X = LHS.getOperand(1);
-      MaskOp = LHS.getOperand(0);
-    }
-    if (X.getNode()) {
-      bool RHSIsZero = ISD::isBuildVectorAllZeros(RHS.getNode());
-      APInt RHSSplat;
-      bool RHSIsMask = !RHSIsZero &&
-                       ISD::isConstantSplatVector(RHS.getNode(), RHSSplat) &&
-                       RHSSplat == SplatVal;
-      if (RHSIsZero || RHSIsMask) {
-        SDValue CMTSTNode =
-            DAG.getNode(AArch64ISD::CMTST, DL, CmpVT, X, MaskOp);
-        // CMTST gives all-ones where (X & Mask) != 0, i.e. SETNE(AND, 0).
-        // Invert when the original condition is the opposite sense.
-        bool Invert = (Cond == ISD::SETEQ) ? RHSIsZero : RHSIsMask;
-        if (Invert)
-          return DAG.getNOT(DL, CMTSTNode, CmpVT);
-        return CMTSTNode;
-      }
-    }
+  // Fold setcc(and(X, Y), 0, seteq) --> NOT(AArch64ISD::CMTST(X, Y))
+  // after DAG legalization. SETNE will have been legalized to SETEQ by now.
+  // The NOT folds away when the result feeds a BSP.
+  if (DCI.isAfterLegalizeDAG() && CmpVT.isFixedLengthVector() &&
+      Cond == ISD::SETEQ && LHS.getOpcode() == ISD::AND &&
+      ISD::isConstantSplatVectorAllZeros(RHS.getNode())) {
+    SDValue CMTSTNode = DAG.getNode(AArch64ISD::CMTST, DL, CmpVT,
+                                    LHS.getOperand(0), LHS.getOperand(1));
+    return DAG.getNOT(DL, CMTSTNode, CmpVT);
   }
 
   return SDValue();



More information about the llvm-commits mailing list