[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
Mon Jul 13 00:30:41 PDT 2026


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

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

---
 .../Target/AArch64/AArch64ISelLowering.cpp    | 53 ++++++++++++++++++
 .../CodeGen/AArch64/cmtst-select-pow2-mask.ll | 55 +++++++++++++++++++
 2 files changed, 108 insertions(+)
 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 f33d953eec747..15e963241a70a 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -16423,6 +16423,59 @@ static SDValue tryLowerToBSL(SDValue N, SelectionDAG &DAG) {
                            N0->getOperand(1 - i), N1->getOperand(1 - j));
     }
 
+  // Fold: or(and(xor(setcc(and(X,Mask), Mask, eq), -1), A), and(setcc(and(X,Mask), Mask, eq), B))
+  // --> BSP(setcc(and(X,Mask), 0, ne), A, B)
+  // (X & Mask) == Mask, for a power-of-2 Mask, is equivalent to (X & Mask) != 0.
+  // The latter lowers to CMTST (one instruction) instead of AND+CMEQ (two instructions).
+  for (int i = 1; i >= 0; --i)
+    for (int j = 1; j >= 0; --j) {
+      SDValue NotMask = N0->getOperand(i);
+      SDValue A = N0->getOperand(1 - i);
+      SDValue Mask = N1->getOperand(j);
+      SDValue B = N1->getOperand(1 - j);
+
+      if (NotMask.getOpcode() != ISD::XOR ||
+          !ISD::isBuildVectorAllOnes(NotMask.getOperand(1).getNode()))
+        continue;
+      if (Mask != NotMask.getOperand(0))
+        continue;
+      if (Mask.getOpcode() != ISD::SETCC)
+        continue;
+
+      ISD::CondCode CC = cast<CondCodeSDNode>(Mask.getOperand(2))->get();
+      if (CC != ISD::SETEQ && CC != ISD::SETNE)
+        continue;
+
+      SDValue InnerAND = Mask.getOperand(0);
+      SDValue CmpRHS = Mask.getOperand(1);
+
+      if (InnerAND.getOpcode() != ISD::AND)
+        continue;
+
+      APInt SplatVal;
+      bool Op0IsPow2 = ISD::isConstantSplatVector(
+                           InnerAND.getOperand(0).getNode(), SplatVal) &&
+                       SplatVal.isPowerOf2();
+      bool Op1IsPow2 = !Op0IsPow2 &&
+                       ISD::isConstantSplatVector(
+                           InnerAND.getOperand(1).getNode(), SplatVal) &&
+                       SplatVal.isPowerOf2();
+      if (!Op0IsPow2 && !Op1IsPow2)
+        continue;
+
+      bool RHSIsZero = ISD::isBuildVectorAllZeros(CmpRHS.getNode());
+      APInt RHSSplat;
+      bool RHSIsMask = !RHSIsZero &&
+                       ISD::isConstantSplatVector(CmpRHS.getNode(), RHSSplat) &&
+                       RHSSplat == SplatVal;
+      if (!RHSIsZero && !RHSIsMask)
+        continue;
+
+      SDValue Zero = DAG.getConstant(0, DL, CmpRHS.getValueType());
+      SDValue NewMask =
+          DAG.getSetCC(DL, Mask.getValueType(), InnerAND, Zero, ISD::SETNE);
+      return DAG.getNode(AArch64ISD::BSP, DL, VT, NewMask, A, B);
+    }
   return SDValue();
 }
 
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
+}



More information about the llvm-commits mailing list