[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