[llvm] [SDAG][AArch64] Fold extract from pext to use status flags (PR #206443)

via llvm-commits llvm-commits at lists.llvm.org
Mon Jun 29 02:59:51 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-aarch64

Author: Benjamin Maxwell (MacDue)

<details>
<summary>Changes</summary>

This folds extracting the first bit from the first segment of a predicate-as-counter to use the "first active" status. E.g.:

```
%pn:aarch64svcount, %flags:FlagsVT = WHILELO_PRED_COUNTER(a, b, VLx4)
%first_pred:nxv4i1 = pext(%pn, 0)
%more:i1 = extractelement(%first_pred, 0)
```
->
```
%pn:aarch64svcount, %flags:FlagsVT = WHILELO_PRED_COUNTER(a, b, VLx4)
%more = CSET(%flags, FIRST_ACTIVE)
```

Assisted-by: Codex (adding test variations)

---
Full diff: https://github.com/llvm/llvm-project/pull/206443.diff


2 Files Affected:

- (modified) llvm/lib/Target/AArch64/AArch64ISelLowering.cpp (+61) 
- (added) llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll (+204) 


``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index e883c8bb5e96e..278dcf99d844e 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -21970,6 +21970,65 @@ performExtractLastActiveCombine(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
                      Vec);
 }
 
+static bool hasSVE2p1OrStreamingSME2(const AArch64Subtarget *Subtarget) {
+  return (Subtarget->isSVEorStreamingSVEAvailable() &&
+          Subtarget->hasSVE2p1()) ||
+         (Subtarget->isStreaming() && Subtarget->hasSME2());
+}
+
+static auto m_PredicateAsCounterWhile() {
+  using namespace llvm::SDPatternMatch;
+  return m_AnyOf(m_SpecificOpc(AArch64ISD::WHILEGE_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILEGT_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILELT_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILELE_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILEHS_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILEHI_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILELO_PRED_COUNTER),
+                 m_SpecificOpc(AArch64ISD::WHILELS_PRED_COUNTER));
+}
+
+/// Folds extracting the first lane from the first segment of a
+/// predicate-as-counter while to a conditional set (CSET) based on the
+/// "FIRST_ACTIVE" status flag from the while.
+///
+///   %while = WHILE_*_PRED_COUNTER .. ; predicate-as-counter while
+///   %first.segment = pext(%while, 0) ; predicate extract of segment 0
+///   %first.active = extract_elt(%first_segment, 0) ; extract first lane
+///
+///   ->
+///
+///   %while = WHILE_*_PRED_COUNTER .. ; predicate-as-counter while
+///   %first.active = cset(%while, FIRST_ACTIVE)
+static SDValue
+perfomPextFirstTrueVectorCombine(SDNode *N,
+                                 TargetLowering::DAGCombinerInfo &DCI,
+                                 const AArch64Subtarget *Subtarget) {
+  using namespace llvm::SDPatternMatch;
+  assert(N->getOpcode() == ISD::EXTRACT_VECTOR_ELT);
+  if (DCI.isBeforeLegalize() || !hasSVE2p1OrStreamingSME2(Subtarget))
+    return SDValue();
+
+  SDValue N0 = N->getOperand(0);
+  EVT VT = N0.getValueType();
+
+  if (!VT.isScalableVectorOf(MVT::i1) || !isNullConstant(N->getOperand(1)))
+    return SDValue();
+
+  SelectionDAG &DAG = DCI.DAG;
+
+  // Fold: extract_elt(pext(WHILE_*_PRED_COUNTER, 0), 0)
+  // to cset(WHILE_*_PRED_COUNTER, FIRST_ACTIVE).
+  if (sd_match(N0, m_IntrinsicWOChain<Intrinsic::aarch64_sve_pext>(
+                       m_PredicateAsCounterWhile(), m_Zero()))) {
+    SDValue WhilePredCounter = N0->getOperand(1);
+    SDValue Flags = SDValue(WhilePredCounter.getNode(), 1);
+    return getSETCC(AArch64CC::CondCode::FIRST_ACTIVE, Flags, SDLoc(N), DAG);
+  }
+
+  return SDValue();
+}
+
 static SDValue
 performExtractVectorEltCombine(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
                                const AArch64Subtarget *Subtarget) {
@@ -21980,6 +22039,8 @@ performExtractVectorEltCombine(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
     return Res;
   if (SDValue Res = performExtractLastActiveCombine(N, DCI, Subtarget))
     return Res;
+  if (SDValue Res = perfomPextFirstTrueVectorCombine(N, DCI, Subtarget))
+    return Res;
 
   SelectionDAG &DAG = DCI.DAG;
   SDValue N0 = N->getOperand(0), N1 = N->getOperand(1);
diff --git a/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll b/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll
new file mode 100644
index 0000000000000..e4c0307fd0dcb
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/sve2p1-while-pn-folds.ll
@@ -0,0 +1,204 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=aarch64-linux-unknown -mattr=+sve2p1 -o - < %s | FileCheck %s -check-prefixes=CHECK,CHECK-SVE
+; RUN: llc -mtriple=aarch64-linux-unknown -mattr=+sme2 -force-streaming -o - < %s | FileCheck %s -check-prefixes=CHECK,CHECK-SME
+
+; Tests extracting the first lane from the first segment of a
+; predicate-as-counter while folds to a conditional set (CSET) based on the
+; "FIRST_ACTIVE" status flag from the while.
+;
+; %while = WHILE_*_PRED_COUNTER .. ; predicate-as-counter while
+; %first.segment = pext(%while, 0) ; predicate extract of segment 0
+; %first.active = extract_elt(%first_segment, 0) ; extract first lane
+;
+; ->
+;
+; %while = WHILE_*_PRED_COUNTER .. ; predicate-as-counter while
+; %first.active = cset(%while, FIRST_ACTIVE)
+
+define i1 @whilege_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilege_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilege pn8.b, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilege.c8(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 16 x i1> @llvm.aarch64.sve.pext.nxv16i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 16 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilegt_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilegt_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilegt pn8.h, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilegt.c16(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 8 x i1> @llvm.aarch64.sve.pext.nxv8i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 8 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilelt_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilelt_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilelt pn8.s, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilelt.c32(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 4 x i1> @llvm.aarch64.sve.pext.nxv4i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 4 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilele_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilele_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilele pn8.d, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilele.c64(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 2 x i1> @llvm.aarch64.sve.pext.nxv2i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 2 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilehs_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilehs_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilehs pn8.b, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilehs.c8(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 16 x i1> @llvm.aarch64.sve.pext.nxv16i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 16 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilehi_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilehi_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilehi pn8.h, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilehi.c16(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 8 x i1> @llvm.aarch64.sve.pext.nxv8i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 8 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilelo_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilelo_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilelo pn8.s, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilelo.c32(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 4 x i1> @llvm.aarch64.sve.pext.nxv4i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 4 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+define i1 @whilels_first_active(i64 %a, i64 %b) {
+; CHECK-LABEL: whilels_first_active:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilels pn8.d, x0, x1, vlx4
+; CHECK-NEXT:    cset w0, mi
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilels.c64(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 2 x i1> @llvm.aarch64.sve.pext.nxv2i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 2 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+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:  // %bb.1: // %then
+; CHECK-NEXT:    //APP
+; CHECK-NEXT:    //NO_APP
+; CHECK-NEXT:  .LBB8_2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilege.c8(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 16 x i1> @llvm.aarch64.sve.pext.nxv16i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 16 x i1> %pext, i64 0
+  br i1 %bit, label %then, label %exit
+
+then:
+  tail call void asm sideeffect "", ""()
+  br label %exit
+
+exit:
+  ret void
+}
+
+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:  // %bb.1: // %then
+; CHECK-NEXT:    //APP
+; CHECK-NEXT:    //NO_APP
+; CHECK-NEXT:  .LBB9_2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilelo.c32(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 4 x i1> @llvm.aarch64.sve.pext.nxv4i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 4 x i1> %pext, i64 0
+  br i1 %bit, label %exit, label %then
+
+then:
+  tail call void asm sideeffect "", ""()
+  br label %exit
+
+exit:
+  ret void
+}
+
+; Negative test: a nonzero pext offset does not match the first-active fold.
+define i1 @whilelo_pext_nonzero_offset(i64 %a, i64 %b) {
+; CHECK-LABEL: whilelo_pext_nonzero_offset:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    whilelo pn8.s, x0, x1, vlx4
+; CHECK-NEXT:    pext p0.s, pn8[1]
+; CHECK-NEXT:    mov z0.s, p0/z, #1 // =0x1
+; CHECK-NEXT:    fmov w8, s0
+; CHECK-NEXT:    and w0, w8, #0x1
+; CHECK-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilelo.c32(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 4 x i1> @llvm.aarch64.sve.pext.nxv4i1(target("aarch64.svcount") %while, i32 1)
+  %bit = extractelement <vscale x 4 x i1> %pext, i64 0
+  ret i1 %bit
+}
+
+; Negative test: extracting any element other than lane 0 does not match the first-active fold.
+define i1 @whilelt_extractelement_nonzero_index(i64 %a, i64 %b) {
+; CHECK-SVE-LABEL: whilelt_extractelement_nonzero_index:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelt pn8.h, x0, x1, vlx4
+; CHECK-SVE-NEXT:    pext p0.h, pn8[0]
+; CHECK-SVE-NEXT:    mov z0.h, p0/z, #1 // =0x1
+; CHECK-SVE-NEXT:    umov w8, v0.h[1]
+; CHECK-SVE-NEXT:    and w0, w8, #0x1
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SME-LABEL: whilelt_extractelement_nonzero_index:
+; CHECK-SME:       // %bb.0:
+; CHECK-SME-NEXT:    whilelt pn8.h, x0, x1, vlx4
+; CHECK-SME-NEXT:    pext p0.h, pn8[0]
+; CHECK-SME-NEXT:    mov z0.h, p0/z, #1 // =0x1
+; CHECK-SME-NEXT:    mov z0.h, z0.h[1]
+; CHECK-SME-NEXT:    fmov w8, s0
+; CHECK-SME-NEXT:    and w0, w8, #0x1
+; CHECK-SME-NEXT:    ret
+  %while = call target("aarch64.svcount") @llvm.aarch64.sve.whilelt.c16(i64 %a, i64 %b, i32 4)
+  %pext = call <vscale x 8 x i1> @llvm.aarch64.sve.pext.nxv8i1(target("aarch64.svcount") %while, i32 0)
+  %bit = extractelement <vscale x 8 x i1> %pext, i64 1
+  ret i1 %bit
+}

``````````

</details>


https://github.com/llvm/llvm-project/pull/206443


More information about the llvm-commits mailing list