[llvm] [AArch64] Fold CSET + BR_CC into a conditional branch (PR #207398)

Benjamin Maxwell via llvm-commits llvm-commits at lists.llvm.org
Mon Jul 13 05:24:05 PDT 2026


https://github.com/MacDue updated https://github.com/llvm/llvm-project/pull/207398

>From 5f5ecef3f5a4ddc482da3854670ce105e8971a5e Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Fri, 12 Jun 2026 13:44:21 +0000
Subject: [PATCH 1/2] [AArch64] Fold CSET + BR_CC into a conditional branch

Fold CSETs into conditional branches (rather than lowering to a CSET
+ TB[N]Z). This is most useful for lowering loop conditions based on
predicate-as-counter whiles to:

```
whilelo pn8.s, x0, x1, vlx4
b.mi .Lloop
```
---
 .../lib/Target/AArch64/AArch64ISelLowering.cpp | 18 ++++++++++++++++++
 .../CodeGen/AArch64/branch-cond-split-fcmp.ll  |  6 ++----
 .../CodeGen/AArch64/sve2p1-while-pn-folds.ll   |  6 ++----
 3 files changed, 22 insertions(+), 8 deletions(-)

diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index d22d3ab16914b..9e49ef6516097 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -11725,6 +11725,24 @@ SDValue AArch64TargetLowering::LowerBR_CC(SDValue Op, SelectionDAG &DAG) const {
                        Overflow);
   }
 
+  // Fold CSET + BR_CC to a conditional branch (rather than keeping the CSET
+  // and emitting a TB[N]Z below).
+  {
+    using namespace llvm::SDPatternMatch;
+    SDValue Flags;
+    uint64_t InvCC;
+    auto m_CSET = m_Node(AArch64ISD::CSINC, m_Zero(), m_Zero(),
+                         m_ConstInt(InvCC), m_Value(Flags));
+    if ((CC == ISD::SETEQ || CC == ISD::SETNE) && isNullConstant(RHS) &&
+        sd_match(LHS, m_AnyOf(m_CSET, m_And(m_CSET, m_One())))) {
+      AArch64CC::CondCode BranchCC = AArch64CC::CondCode(InvCC);
+      if (CC == ISD::SETNE)
+        BranchCC = AArch64CC::getInvertedCondCode(BranchCC);
+      return DAG.getNode(AArch64ISD::BRCOND, DL, MVT::Other, Chain, Dest,
+                         getCondCode(DAG, BranchCC), Flags);
+    }
+  }
+
   if (LHS.getValueType().isInteger()) {
     assert((LHS.getValueType() == RHS.getValueType()) &&
            (LHS.getValueType() == MVT::i32 || LHS.getValueType() == MVT::i64));
diff --git a/llvm/test/CodeGen/AArch64/branch-cond-split-fcmp.ll b/llvm/test/CodeGen/AArch64/branch-cond-split-fcmp.ll
index 10c958abc71d6..26eda6c60a1b1 100644
--- a/llvm/test/CodeGen/AArch64/branch-cond-split-fcmp.ll
+++ b/llvm/test/CodeGen/AArch64/branch-cond-split-fcmp.ll
@@ -219,8 +219,7 @@ define i64 @test_or_fast(float %a, float %b) {
 ; CHECK-SD-NEXT:    movi d2, #0000000000000000
 ; CHECK-SD-NEXT:    fcmp s1, #0.0
 ; CHECK-SD-NEXT:    fccmp s0, s2, #0, eq
-; CHECK-SD-NEXT:    cset w8, eq
-; CHECK-SD-NEXT:    tbnz w8, #0, .LBB4_2
+; CHECK-SD-NEXT:    b.eq .LBB4_2
 ; CHECK-SD-NEXT:  // %bb.1:
 ; CHECK-SD-NEXT:    mov x0, xzr
 ; CHECK-SD-NEXT:    ret
@@ -320,8 +319,7 @@ define i64 @test_and_fast(float %a, float %b) {
 ; CHECK-SD-NEXT:    movi d2, #0000000000000000
 ; CHECK-SD-NEXT:    fcmp s1, #0.0
 ; CHECK-SD-NEXT:    fccmp s0, s2, #0, eq
-; CHECK-SD-NEXT:    cset w8, eq
-; CHECK-SD-NEXT:    tbz w8, #0, .LBB6_2
+; CHECK-SD-NEXT:    b.ne .LBB6_2
 ; CHECK-SD-NEXT:  // %bb.1: // %bb4
 ; CHECK-SD-NEXT:    str x30, [sp, #-16]! // 8-byte Folded Spill
 ; CHECK-SD-NEXT:    .cfi_def_cfa_offset 16
diff --git a/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll b/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll
index d58097f77f6d3..111f0925f1c75 100644
--- a/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll
+++ b/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll
@@ -118,8 +118,7 @@ define void @whilege_first_active_branch(i64 %a, i64 %b) {
 ; CHECK-LABEL: whilege_first_active_branch:
 ; CHECK:       // %bb.0: // %entry
 ; CHECK-NEXT:    whilege pn8.b, x0, x1, vlx4
-; CHECK-NEXT:    cset w8, mi
-; CHECK-NEXT:    cbz w8, .LBB8_2
+; CHECK-NEXT:    b.pl .LBB8_2
 ; CHECK-NEXT:  // %bb.1: // %then
 ; CHECK-NEXT:    //APP
 ; CHECK-NEXT:    //NO_APP
@@ -143,8 +142,7 @@ define void @whilelo_first_active_branch(i64 %a, i64 %b) {
 ; CHECK-LABEL: whilelo_first_active_branch:
 ; CHECK:       // %bb.0: // %entry
 ; CHECK-NEXT:    whilelo pn8.s, x0, x1, vlx4
-; CHECK-NEXT:    cset w8, mi
-; CHECK-NEXT:    tbnz w8, #0, .LBB9_2
+; CHECK-NEXT:    b.mi .LBB9_2
 ; CHECK-NEXT:  // %bb.1: // %then
 ; CHECK-NEXT:    //APP
 ; CHECK-NEXT:    //NO_APP

>From 7409ae864ee1277da14df954fe6b35fee332376a Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Mon, 13 Jul 2026 12:23:52 +0000
Subject: [PATCH 2/2] Comments

---
 llvm/lib/Target/AArch64/AArch64ISelLowering.cpp | 8 +++++---
 1 file changed, 5 insertions(+), 3 deletions(-)

diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 9e49ef6516097..c2f686966d81a 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -11730,12 +11730,14 @@ SDValue AArch64TargetLowering::LowerBR_CC(SDValue Op, SelectionDAG &DAG) const {
   {
     using namespace llvm::SDPatternMatch;
     SDValue Flags;
-    uint64_t InvCC;
+    uint64_t InverseCC;
+    // `CSET <Wd>, <cond>` is an alias of `CSINC <Wd>, WZR, WZR, invert(<cond>)`
     auto m_CSET = m_Node(AArch64ISD::CSINC, m_Zero(), m_Zero(),
-                         m_ConstInt(InvCC), m_Value(Flags));
+                         m_ConstInt(InverseCC), m_Value(Flags));
+    // Note: We look through `& 1` as the result of CSET is known to be 0 or 1.
     if ((CC == ISD::SETEQ || CC == ISD::SETNE) && isNullConstant(RHS) &&
         sd_match(LHS, m_AnyOf(m_CSET, m_And(m_CSET, m_One())))) {
-      AArch64CC::CondCode BranchCC = AArch64CC::CondCode(InvCC);
+      AArch64CC::CondCode BranchCC = AArch64CC::CondCode(InverseCC);
       if (CC == ISD::SETNE)
         BranchCC = AArch64CC::getInvertedCondCode(BranchCC);
       return DAG.getNode(AArch64ISD::BRCOND, DL, MVT::Other, Chain, Dest,



More information about the llvm-commits mailing list