[llvm] [AArch64] Add predicate-as-counter loop rewrite pass (PR #220960)

Benjamin Maxwell via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 8 07:02:54 PDT 2026


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

>From d955bd9c5bf6f2bce839caa3d9cca7e1c7804a7d Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Thu, 3 Sep 2026 13:30:04 +0000
Subject: [PATCH 1/5] [AArch64] Add predicate-as-counter loop rewrite pass

This patch adds an AArch64 IR loop pass that rewrites wide loop-carried
`llvm.get.active.lane.mask` phis to predicate-as-counter `whilelo` phis
 (for SVE2.1 or streaming SME2 targets).

Users of the original mask are preserved by materializing vector
predicates with `aarch64.sve.pext`.

The element size and vector scale (VLx2 or VLx4) of the
predicate-as-counter is inferred from the mask load/store users within
the loop. These could be optimized to multi-vector loads/stores (though
that is not included in this patch).

For example, a loop like:

```
entry:
  %step = vscale x 64
  %mask.entry = @get.active.lane.mask(0, %n)

loop:
  %iv   = phi i64 [0, entry], [%iv.next, loop]
  %mask = phi <vscale x 64 x i1> [%mask.entry, entry],
                                 [%mask.next, loop]
  ...

  %iv.next   = %iv + %step
  %mask.next = @get.active.lane.mask(%iv.next, %n)
  br first.active(%mask.next), loop, exit
```

Could be rewritten to:

```
entry:
  %step = vscale x 64
  %mask.entry = @whilelo.c8(0, %n, VLx4)

loop:
  %iv   = phi i64 [0, entry], [%iv.next, loop]
  %mask = phi target("aarch64.svcount") [%mask.entry, entry],
                                        [%mask.next, loop]
  ...

  %iv.next   = %iv + %step
  %mask.next = @whilelo.c8(%iv.next, %n, VLx4)
  br first.active(@pext(%mask.next, 0)), loop, exit
```

Assisted-by: Codex
---
 llvm/lib/Target/AArch64/AArch64.h             |   2 +
 .../AArch64PredicateAsCounterLoopRewrites.cpp | 458 ++++++++++++
 .../Target/AArch64/AArch64TargetMachine.cpp   |  10 +
 llvm/lib/Target/AArch64/CMakeLists.txt        |   1 +
 .../predicate-as-counter-loop-rewrites.ll     | 673 ++++++++++++++++++
 5 files changed, 1144 insertions(+)
 create mode 100644 llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
 create mode 100644 llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll

diff --git a/llvm/lib/Target/AArch64/AArch64.h b/llvm/lib/Target/AArch64/AArch64.h
index bcf40a02e3cf5..f2e49d6578f9c 100644
--- a/llvm/lib/Target/AArch64/AArch64.h
+++ b/llvm/lib/Target/AArch64/AArch64.h
@@ -76,6 +76,7 @@ FunctionPass *createAArch64PTrueCoalescingLegacyPass();
 FunctionPass *createAArch64CleanupLocalDynamicTLSPass();
 
 FunctionPass *createAArch64CollectLOHPass();
+Pass *createAArch64PredicateAsCounterLoopRewritesPass();
 FunctionPass *createSMEPeepholeOptPass();
 FunctionPass *createMachineSMEABIPass(CodeGenOptLevel);
 FunctionPass *createAArch64SRLTDefineSuperRegsLegacyPass();
@@ -167,6 +168,7 @@ void initializeAArch64A57FPLoadBalancingLegacyPass(PassRegistry &);
 void initializeAArch64AdvSIMDScalarLegacyPass(PassRegistry &);
 void initializeAArch64AsmPrinterPass(PassRegistry &);
 void initializeAArch64PointerAuthLegacyPass(PassRegistry &);
+void initializeAArch64PredicateAsCounterLoopRewritesPass(PassRegistry &);
 void initializeAArch64BranchTargetsLegacyPass(PassRegistry &);
 void initializeAArch64CFIFixupPass(PassRegistry&);
 void initializeAArch64CollectLOHLegacyPass(PassRegistry &);
diff --git a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
new file mode 100644
index 0000000000000..b701e6c176146
--- /dev/null
+++ b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
@@ -0,0 +1,458 @@
+//===- AArch64PredicateAsCounterLoopRewrites.cpp --------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM
+// Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+// Rewrites IR for loop-carried wide masks that can be represented as
+// predicate-as-counter values. This applies when the mask is used by load/store
+// operations that can be mapped to multi-vector instructions (with +sve2p1).
+//
+// For example, a loop like:
+//
+//   entry:
+//     %step = vscale x 64
+//     %mask.entry = @get.active.lane.mask(0, %n)
+//
+//   loop:
+//     %iv   = phi i64 [0, entry], [%iv.next, loop]
+//     %mask = phi <vscale x 64 x i1> [%mask.entry, entry],
+//                                    [%mask.next, loop]
+//
+//     %load = load <vscale x 64 x i8> %src[%iv], %mask
+//     store <vscale x 64 x i8> %load, %dst[%iv], %mask
+//
+//     %iv.next   = %iv + %step
+//     %mask.next = @get.active.lane.mask(%iv.next, %n)
+//     br first.active(%mask.next), loop, exit
+//
+// Could be rewritten to:
+//
+//   entry:
+//     %step = vscale x 64
+//     %mask.entry = @whilelo.c8(0, %n, VLx4)
+//
+//   loop:
+//     %iv   = phi i64 [0, entry], [%iv.next, loop]
+//     %mask = phi target("aarch64.svcount") [%mask.entry, entry],
+//                                           [%mask.next, loop]
+//
+//     %load = @ld1.pn.x4 <4 x <vscale x 16 x i8>> %src[%iv], %mask
+//     @st1.pn.x4 <4 x <vscale x 16 x i8>> %load, %dest[%iv], %mask
+//
+//     %iv.next   = %iv + %step
+//     %mask.next = @whilelo.c8(%iv.next, %n, VLx4)
+//     br first.active(@pext(%mask.next, 0)), loop, exit
+//
+// This replaces the `get.active.lane.mask` intrinsics with AArch64
+// predicate-as-counter `whilelo` intrinsics and updates the mask phi to use the
+// `aarch64.svcount` target type. Within the loop, load/store users are mapped
+// to multi-vector load/store intrinsics where possible. Users that cannot be
+// mapped to multi-vector instructions materialize vector masks using the `pext`
+// intrinsic (which extracts vector predicates from a predicate-as-counter).
+//
+// This pass may be a temporary solution that is removed if we gain support
+// for target-specific VPlan transforms in the loop vectorizer.
+//
+//===----------------------------------------------------------------------===//
+
+#include "AArch64.h"
+#include "AArch64Subtarget.h"
+#include "AArch64TargetMachine.h"
+#include "llvm/ADT/DenseMap.h"
+#include "llvm/ADT/SmallVector.h"
+#include "llvm/ADT/Statistic.h"
+#include "llvm/Analysis/LoopInfo.h"
+#include "llvm/Analysis/LoopPass.h"
+#include "llvm/CodeGen/TargetPassConfig.h"
+#include "llvm/IR/Attributes.h"
+#include "llvm/IR/DataLayout.h"
+#include "llvm/IR/IRBuilder.h"
+#include "llvm/IR/IntrinsicInst.h"
+#include "llvm/IR/Intrinsics.h"
+#include "llvm/IR/IntrinsicsAArch64.h"
+#include "llvm/IR/Module.h"
+#include "llvm/InitializePasses.h"
+#include "llvm/Pass.h"
+#include "llvm/Support/Debug.h"
+#include "llvm/Transforms/Utils.h"
+#include "llvm/Transforms/Utils/Local.h"
+#include <optional>
+
+using namespace llvm;
+
+#define DEBUG_TYPE "aarch64-predicate-as-counter-loop-rewrites"
+namespace {
+
+STATISTIC(LoopsRewritten, "Number of loops rewritten");
+
+struct MaskRewriteCandidate {
+  /// The preheader block for the loop.
+  BasicBlock *Preheader = nullptr;
+  /// The latch block for the loop.
+  BasicBlock *Latch = nullptr;
+  /// The mask phi node (used by masked operations within the loop).
+  PHINode *MaskPhi = nullptr;
+  /// The initial value for the mask (incoming value from the preheader).
+  IntrinsicInst *StartMask = nullptr;
+  /// The updated value for the mask (incoming value from the loop latch).
+  IntrinsicInst *NextMask = nullptr;
+  /// The multi-vector scale for the predicate-as-counter (2 or 4).
+  unsigned VectorScale = 0;
+  /// The element size (in bits) for the predicate-as-counter.
+  unsigned ElementSizeInBits = 0;
+};
+
+static void logLoopBailout(const Loop &L, const Twine &Reason) {
+  LLVM_DEBUG({
+    dbgs() << "PAC loop rewrite: skipping loop with header ";
+    L.getHeader()->printAsOperand(dbgs(), /*PrintType=*/false);
+    dbgs() << ": " << Reason << "\n";
+  });
+}
+
+static void logMatchFailure(const PHINode &Phi, const Twine &Reason) {
+  LLVM_DEBUG({
+    dbgs() << "PAC loop rewrite: failed to match mask phi ";
+    Phi.printAsOperand(dbgs(), /*PrintType=*/false);
+    dbgs() << ": " << Reason << "\n";
+  });
+}
+
+/// Returns the scalar size in bits for a type according to the data layout.
+static unsigned getScalarSizeInBits(const DataLayout &DL, Type *Ty) {
+  return DL.getTypeSizeInBits(Ty->getScalarType()).getFixedValue();
+}
+
+/// Returns the SVE element count for \p ElementSizeInBits.
+static ElementCount getSVEElementCount(unsigned ElementSizeInBits) {
+  return ElementCount::getScalable(AArch64::SVEBitsPerBlock /
+                                   ElementSizeInBits);
+}
+
+/// Returns the most common scalar access size of masked loads/stores in loop
+/// \p L where \p MaskPhi is used as the mask. Ties are broken in favor of
+/// larger access sizes.
+static std::optional<unsigned>
+getMostCommonMaskedMemAccessSizeInBits(const Loop &L, PHINode &MaskPhi) {
+  const DataLayout &DL = MaskPhi.getModule()->getDataLayout();
+  DenseMap<unsigned, unsigned> AccessSizeCounts;
+  std::optional<unsigned> BestAccessSizeInBits;
+  unsigned BestAccessSizeCount = 0;
+
+  for (User *U : MaskPhi.users()) {
+    auto *II = dyn_cast<IntrinsicInst>(U);
+    if (!II || !L.contains(II))
+      continue;
+
+    Intrinsic::ID IID = II->getIntrinsicID();
+    if (IID != Intrinsic::masked_load && IID != Intrinsic::masked_store)
+      continue;
+
+    unsigned MaskOpIdx = IID == Intrinsic::masked_load ? 1 : 2;
+    if (II->getArgOperand(MaskOpIdx) != &MaskPhi)
+      continue;
+
+    unsigned AccessSizeInBits = getScalarSizeInBits(DL, II->getAccessType());
+    unsigned AccessSizeCount = ++AccessSizeCounts[AccessSizeInBits];
+
+    if (!BestAccessSizeInBits || AccessSizeCount > BestAccessSizeCount ||
+        (AccessSizeCount == BestAccessSizeCount &&
+         AccessSizeInBits > *BestAccessSizeInBits)) {
+      BestAccessSizeInBits = AccessSizeInBits;
+      BestAccessSizeCount = AccessSizeCount;
+    }
+  }
+
+  return BestAccessSizeInBits;
+}
+
+/// Returns the predicate-as-counter whilelo intrinsic ID for
+/// \p ElementSizeInBits.
+static Intrinsic::ID getWhileLOIntrinsic(unsigned ElementSizeInBits) {
+  switch (ElementSizeInBits) {
+  case 8:
+    return Intrinsic::aarch64_sve_whilelo_c8;
+  case 16:
+    return Intrinsic::aarch64_sve_whilelo_c16;
+  case 32:
+    return Intrinsic::aarch64_sve_whilelo_c32;
+  case 64:
+    return Intrinsic::aarch64_sve_whilelo_c64;
+  default:
+    llvm_unreachable("unsupported predicate-as-counter element size");
+  }
+}
+
+/// Expands the predicate-as-counter mask \p Count into a wide vector mask.
+/// Masks are extracted from the counter using the paired pext intrinsics,
+/// then concatenated to form the wide mask value.
+static Value *buildWideMask(IRBuilder<> &Builder, const MaskRewriteCandidate &C,
+                            Value *Count) {
+  ElementCount LegalEC = getSVEElementCount(C.ElementSizeInBits);
+  Module *M = Builder.GetInsertBlock()->getModule();
+  Type *LegalMaskTy = VectorType::get(Builder.getInt1Ty(), LegalEC);
+  FunctionCallee PExtX2 = Intrinsic::getOrInsertDeclaration(
+      M, Intrinsic::aarch64_sve_pext_x2, {LegalMaskTy});
+
+  Value *WideMask = PoisonValue::get(C.MaskPhi->getType());
+  for (unsigned PairOffset = 0; PairOffset != C.VectorScale / 2; ++PairOffset) {
+    auto *Pair = Builder.CreateCall(
+        PExtX2, {Count, Builder.getInt32(PairOffset)}, "pac.pext.pair");
+    for (unsigned SliceInPair = 0; SliceInPair != 2; ++SliceInPair) {
+      Value *Part = Builder.CreateExtractValue(Pair, SliceInPair, "pac.pext");
+      unsigned Slice = PairOffset * 2 + SliceInPair;
+      WideMask = Builder.CreateInsertVector(
+          C.MaskPhi->getType(), WideMask, Part,
+          Slice * LegalEC.getKnownMinValue(), "pac.mask");
+    }
+  }
+
+  return WideMask;
+}
+
+/// Creates a predicate-as-counter whilelo for \p ElementSizeInBits between
+/// \p Start and \p End with multi-vector \p VectorScale (2 or 4).
+static Value *createWhileLO(IRBuilder<> &Builder, unsigned ElementSizeInBits,
+                            Value *Start, Value *End, unsigned VectorScale) {
+  if (Start->getType()->getIntegerBitWidth() < 64) {
+    Start = Builder.CreateZExt(Start, Builder.getInt64Ty());
+    End = Builder.CreateZExt(End, Builder.getInt64Ty());
+  }
+  Module *M = Builder.GetInsertBlock()->getModule();
+  auto ID = getWhileLOIntrinsic(ElementSizeInBits);
+  FunctionCallee WhileLO = Intrinsic::getOrInsertDeclaration(M, ID);
+  return Builder.CreateCall(
+      WhileLO, {Start, End, Builder.getInt32(VectorScale)}, "pac.mask");
+}
+class AArch64PredicateAsCounterLoopRewrites : public LoopPass {
+public:
+  static char ID;
+
+  AArch64PredicateAsCounterLoopRewrites() : LoopPass(ID) {}
+
+  void getAnalysisUsage(AnalysisUsage &AU) const override {
+    AU.addRequired<TargetPassConfig>();
+    // Require loop simplify to ensure loops have a preheader.
+    AU.addRequiredID(LoopSimplifyID);
+    AU.addPreservedID(LoopSimplifyID);
+    AU.setPreservesCFG();
+  }
+
+  bool runOnLoop(Loop *L, LPPassManager &) override;
+
+private:
+  std::optional<MaskRewriteCandidate> matchMaskPhi(Loop &L, PHINode &Phi) const;
+  bool rewriteCandidate(const MaskRewriteCandidate &C, Loop &L) const;
+};
+
+} // end anonymous namespace
+
+char AArch64PredicateAsCounterLoopRewrites::ID = 0;
+
+INITIALIZE_PASS_BEGIN(AArch64PredicateAsCounterLoopRewrites, DEBUG_TYPE,
+                      "AArch64 Predicate As Counter Loop Rewrites", false,
+                      false)
+INITIALIZE_PASS_DEPENDENCY(TargetPassConfig)
+INITIALIZE_PASS_DEPENDENCY(LoopSimplify)
+INITIALIZE_PASS_END(AArch64PredicateAsCounterLoopRewrites, DEBUG_TYPE,
+                    "AArch64 Predicate As Counter Loop Rewrites", false, false)
+
+Pass *llvm::createAArch64PredicateAsCounterLoopRewritesPass() {
+  return new AArch64PredicateAsCounterLoopRewrites();
+}
+
+bool AArch64PredicateAsCounterLoopRewrites::runOnLoop(Loop *L,
+                                                      LPPassManager &) {
+  if (skipLoop(L)) {
+    logLoopBailout(*L, "skipLoop requested the loop to be skipped");
+    return false;
+  }
+
+  Function &F = *L->getHeader()->getParent();
+  auto &TPC = getAnalysis<TargetPassConfig>();
+  const AArch64Subtarget *ST =
+      TPC.getTM<AArch64TargetMachine>().getSubtargetImpl(F);
+  if (!ST->isSVEorStreamingSVEAvailable()) {
+    logLoopBailout(*L, "SVE or streaming SVE is unavailable");
+    return false;
+  }
+  if (!ST->hasSVE2p1() && !(ST->hasSME2() && ST->isStreaming())) {
+    logLoopBailout(*L, "neither SVE2.1 nor SME2 is available");
+    return false;
+  }
+
+  bool Changed = false;
+  BasicBlock *Header = L->getHeader();
+  for (PHINode &Phi : make_early_inc_range(Header->phis())) {
+    auto Candidate = matchMaskPhi(*L, Phi);
+    if (Candidate)
+      Changed |= rewriteCandidate(*Candidate, *L);
+  }
+
+  if (Changed)
+    ++LoopsRewritten;
+
+  return Changed;
+}
+
+static IntrinsicInst *getGetActiveLaneMask(Value *V) {
+  auto *II = dyn_cast<IntrinsicInst>(V);
+  return II && II->getIntrinsicID() == Intrinsic::get_active_lane_mask
+             ? II
+             : nullptr;
+}
+
+std::optional<MaskRewriteCandidate>
+AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
+                                                    PHINode &Phi) const {
+  BasicBlock *Preheader = L.getLoopPreheader();
+  BasicBlock *Latch = L.getLoopLatch();
+  if (!Preheader) {
+    logMatchFailure(Phi, "loop has no preheader");
+    return std::nullopt;
+  }
+  if (!Latch) {
+    logMatchFailure(Phi, "loop has no latch");
+    return std::nullopt;
+  }
+
+  auto *PhiTy = dyn_cast<ScalableVectorType>(Phi.getType());
+  if (!PhiTy || !PhiTy->getElementType()->isIntegerTy(1)) {
+    logMatchFailure(Phi, "phi type is not a scalable i1 vector mask");
+    return std::nullopt;
+  }
+  if (Phi.getNumIncomingValues() != 2) {
+    logMatchFailure(Phi, Twine("phi has ")
+                             .concat(Twine(Phi.getNumIncomingValues()))
+                             .concat(" incoming values; expected 2"));
+    return std::nullopt;
+  }
+
+  auto *StartMask =
+      getGetActiveLaneMask(Phi.getIncomingValueForBlock(Preheader));
+  auto *NextMask = getGetActiveLaneMask(Phi.getIncomingValueForBlock(Latch));
+  if (!StartMask) {
+    logMatchFailure(Phi,
+                    "preheader incoming value is not get_active_lane_mask");
+    return std::nullopt;
+  }
+  if (!NextMask) {
+    logMatchFailure(Phi, "latch incoming value is not get_active_lane_mask");
+    return std::nullopt;
+  }
+
+  unsigned WideMaskElements = PhiTy->getMinNumElements();
+  if (!isPowerOf2_32(WideMaskElements)) {
+    logMatchFailure(Phi, Twine("wide mask element count is not a power of 2: ")
+                             .concat(Twine(WideMaskElements)));
+    return std::nullopt;
+  }
+
+  if (StartMask->getArgOperand(0)->getType()->getIntegerBitWidth() > 64) {
+    logMatchFailure(Phi, "start mask induction operand is wider than i64");
+    return std::nullopt;
+  }
+  if (NextMask->getArgOperand(0)->getType()->getIntegerBitWidth() > 64) {
+    logMatchFailure(Phi, "next mask induction operand is wider than i64");
+    return std::nullopt;
+  }
+
+  std::optional<unsigned> PreferredMaskElementSizeInBits =
+      getMostCommonMaskedMemAccessSizeInBits(L, Phi);
+  if (!PreferredMaskElementSizeInBits) {
+    logMatchFailure(Phi, "mask phi has no masked load/store users in the loop");
+    return std::nullopt;
+  }
+
+  if (!is_contained({8u, 16u, 32u, 64u}, *PreferredMaskElementSizeInBits)) {
+    logMatchFailure(Phi, Twine("unsupported element size in bits: ")
+                             .concat(Twine(*PreferredMaskElementSizeInBits)));
+    return std::nullopt;
+  }
+
+  unsigned SVEMaskElements =
+      getSVEElementCount(*PreferredMaskElementSizeInBits).getKnownMinValue();
+  if (WideMaskElements <= SVEMaskElements) {
+    logMatchFailure(Phi, Twine("wide mask element count ")
+                             .concat(Twine(WideMaskElements))
+                             .concat(" is not wider than the legal mask width ")
+                             .concat(Twine(SVEMaskElements)));
+    return std::nullopt;
+  }
+
+  unsigned VectorScale = WideMaskElements / SVEMaskElements;
+  if (VectorScale != 2 && VectorScale != 4) {
+    logMatchFailure(Phi, Twine("unsupported predicate-as-counter scale: ")
+                             .concat(Twine(VectorScale)));
+    return std::nullopt;
+  }
+
+  return MaskRewriteCandidate{Preheader,
+                              Latch,
+                              &Phi,
+                              StartMask,
+                              NextMask,
+                              VectorScale,
+                              *PreferredMaskElementSizeInBits};
+}
+
+bool AArch64PredicateAsCounterLoopRewrites::rewriteCandidate(
+    const MaskRewriteCandidate &C, Loop &L) const {
+  IRBuilder<> Builder(C.StartMask);
+  Value *NewStart =
+      createWhileLO(Builder, C.ElementSizeInBits, C.StartMask->getArgOperand(0),
+                    C.StartMask->getArgOperand(1), C.VectorScale);
+  Builder.SetInsertPoint(C.NextMask);
+  Value *NewNext =
+      createWhileLO(Builder, C.ElementSizeInBits, C.NextMask->getArgOperand(0),
+                    C.NextMask->getArgOperand(1), C.VectorScale);
+
+  Builder.SetInsertPoint(C.MaskPhi);
+  auto *NewPhi =
+      Builder.CreatePHI(NewStart->getType(), 2, C.MaskPhi->getName() + ".pn");
+  NewPhi->addIncoming(NewStart, C.Preheader);
+  NewPhi->addIncoming(NewNext, C.Latch);
+
+  auto RewriteUses = [&](Instruction *OldMask, Value *Count,
+                         function_ref<bool(Use & U)> Predicate = nullptr) {
+    SmallVector<Use *, 8> UsesToRewrite;
+    for (Use &U : OldMask->uses()) {
+      if (!Predicate || Predicate(U))
+        UsesToRewrite.push_back(&U);
+    }
+
+    Value *WideMask = nullptr;
+    for (Use *U : UsesToRewrite) {
+      if (!WideMask) {
+        BasicBlock::iterator InsertPt = OldMask->getIterator();
+        if (isa<PHINode>(OldMask))
+          InsertPt = OldMask->getParent()->getFirstNonPHIIt();
+
+        Builder.SetInsertPoint(InsertPt);
+        Builder.SetCurrentDebugLocation(OldMask->getDebugLoc());
+        WideMask = buildWideMask(Builder, C, Count);
+      }
+
+      U->set(WideMask);
+    }
+  };
+
+  // For the start/next mask, we only replace non-phi in-loop users. Due to CSE,
+  // these masks could be used by other loops, and replacing them with
+  // predicate-as-counter masks could prevent matching those loops.
+  auto IsNonPhiUseInLoop = [&](Use &U) {
+    auto *UseInst = cast<Instruction>(U.getUser());
+    return !isa<PHINode>(UseInst) && L.contains(UseInst);
+  };
+
+  RewriteUses(C.MaskPhi, NewPhi);
+  RewriteUses(C.StartMask, NewStart, IsNonPhiUseInLoop);
+  RewriteUses(C.NextMask, NewNext, IsNonPhiUseInLoop);
+
+  RecursivelyDeleteTriviallyDeadInstructions(C.MaskPhi);
+  return true;
+}
diff --git a/llvm/lib/Target/AArch64/AArch64TargetMachine.cpp b/llvm/lib/Target/AArch64/AArch64TargetMachine.cpp
index 8195ad04c3556..72f9125e27c52 100644
--- a/llvm/lib/Target/AArch64/AArch64TargetMachine.cpp
+++ b/llvm/lib/Target/AArch64/AArch64TargetMachine.cpp
@@ -231,6 +231,12 @@ static cl::opt<bool> EnableSVEShuffleOpt(
              "instructions like tbl or the bottom/top variants"),
     cl::init(true), cl::Hidden);
 
+static cl::opt<bool> EnablePredicateAsCounterLoopRewrites(
+    "aarch64-enable-predicate-as-counter-loop-rewrites",
+    cl::desc("Enable rewriting loops with wide loop-carried masks to use "
+             "predicate-as-counter"),
+    cl::init(false), cl::Hidden);
+
 extern "C" LLVM_ABI LLVM_EXTERNAL_VISIBILITY void
 LLVMInitializeAArch64Target() {
   // Register the target.
@@ -258,6 +264,7 @@ LLVMInitializeAArch64Target() {
   initializeAArch64PTrueCoalescingLegacyPass(PR);
   initializeAArch64SIMDInstrOptLegacyPass(PR);
   initializeAArch64O0PreLegalizerCombinerLegacyPass(PR);
+  initializeAArch64PredicateAsCounterLoopRewritesPass(PR);
   initializeAArch64PreLegalizerCombinerLegacyPass(PR);
   initializeAArch64PointerAuthLegacyPass(PR);
   initializeAArch64PostCoalescerLegacyPass(PR);
@@ -650,6 +657,9 @@ void AArch64PassConfig::addIRPasses() {
   // ourselves.
   addPass(createAtomicExpandLegacyPass());
 
+  if (EnablePredicateAsCounterLoopRewrites)
+    addPass(createAArch64PredicateAsCounterLoopRewritesPass());
+
   // Cmpxchg instructions are often used with a subsequent comparison to
   // determine whether it succeeded. We can exploit existing control-flow in
   // ldrex/strex loops to simplify this, but it needs tidying up.
diff --git a/llvm/lib/Target/AArch64/CMakeLists.txt b/llvm/lib/Target/AArch64/CMakeLists.txt
index 12a2214f8e58e..5c4740baf79de 100644
--- a/llvm/lib/Target/AArch64/CMakeLists.txt
+++ b/llvm/lib/Target/AArch64/CMakeLists.txt
@@ -79,6 +79,7 @@ add_llvm_target(AArch64CodeGen
   AArch64PromoteConstant.cpp
   AArch64PTrueCoalescing.cpp
   AArch64PBQPRegAlloc.cpp
+  AArch64PredicateAsCounterLoopRewrites.cpp
   AArch64RegisterInfo.cpp
   AArch64SMEAttributes.cpp
   AArch64SLSHardening.cpp
diff --git a/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll b/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll
new file mode 100644
index 0000000000000..b064c81107795
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll
@@ -0,0 +1,673 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=aarch64-linux-gnu -aarch64-enable-predicate-as-counter-loop-rewrites -mattr=+sve2p1 -aarch64-enable-subreg-liveness-tracking < %s | FileCheck %s
+
+target triple = "aarch64-unknown-linux-gnu"
+
+; Tests the baseline i64 masked load/store rewrite to predicate-as-counter with vlx4.
+define void @rewrite_masked_load_store_i64_vlx4(ptr %x, i64 %n) #0 {
+; CHECK-LABEL: rewrite_masked_load_store_i64_vlx4:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    whilelo pn8.d, xzr, x1, vlx4
+; CHECK-NEXT:  .LBB0_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p2.d, p3.d }, pn8[1]
+; CHECK-NEXT:    pext { p0.d, p1.d }, pn8[0]
+; CHECK-NEXT:    whilelo pn8.d, x8, x1, vlx4
+; CHECK-NEXT:    ld1d { z0.d }, p3/z, [x0, #3, mul vl]
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn8[0]
+; CHECK-NEXT:    ld1d { z1.d }, p2/z, [x0, #2, mul vl]
+; CHECK-NEXT:    uzp1 p4.s, p4.s, p5.s
+; CHECK-NEXT:    pext { p5.d, p6.d }, pn8[1]
+; CHECK-NEXT:    ld1d { z2.d }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    uzp1 p5.s, p5.s, p6.s
+; CHECK-NEXT:    ld1d { z3.d }, p0/z, [x0]
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    add z0.d, z0.d, #1 // =0x1
+; CHECK-NEXT:    add z1.d, z1.d, #1 // =0x1
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    add z2.d, z2.d, #1 // =0x1
+; CHECK-NEXT:    add z3.d, z3.d, #1 // =0x1
+; CHECK-NEXT:    st1d { z0.d }, p3, [x0, #3, mul vl]
+; CHECK-NEXT:    mov z0.h, p4/z, #1 // =0x1
+; CHECK-NEXT:    st1d { z1.d }, p2, [x0, #2, mul vl]
+; CHECK-NEXT:    st1d { z2.d }, p1, [x0, #1, mul vl]
+; CHECK-NEXT:    st1d { z3.d }, p0, [x0]
+; CHECK-NEXT:    incb x0, all, mul #4
+; CHECK-NEXT:    fmov w9, s0
+; CHECK-NEXT:    tbnz w9, #0, .LBB0_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr = getelementptr i64, ptr %x, i64 %base
+  %v = call <vscale x 8 x i64> @llvm.masked.load.nxv8i64(ptr %ptr, i32 8, <vscale x 8 x i1> %update.mask, <vscale x 8 x i64> poison)
+  %add = add <vscale x 8 x i64> %v, splat (i64 1)
+  call void @llvm.masked.store.nxv8i64(<vscale x 8 x i64> %add, ptr %ptr, i32 8, <vscale x 8 x i1> %update.mask)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests rewriting a loop whose header has multiple outside predecessors before LoopSimplify creates a preheader.
+define void @rewrite_masked_load_store_i64_multi_pred_header(ptr %x, i64 %n, i1 %c) #0 {
+; CHECK-LABEL: rewrite_masked_load_store_i64_multi_pred_header:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    whilelo pn8.d, xzr, x1, vlx4
+; CHECK-NEXT:  .LBB1_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p2.d, p3.d }, pn8[1]
+; CHECK-NEXT:    pext { p0.d, p1.d }, pn8[0]
+; CHECK-NEXT:    whilelo pn8.d, x8, x1, vlx4
+; CHECK-NEXT:    ld1d { z0.d }, p3/z, [x0, #3, mul vl]
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn8[0]
+; CHECK-NEXT:    ld1d { z1.d }, p2/z, [x0, #2, mul vl]
+; CHECK-NEXT:    uzp1 p4.s, p4.s, p5.s
+; CHECK-NEXT:    pext { p5.d, p6.d }, pn8[1]
+; CHECK-NEXT:    ld1d { z2.d }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    uzp1 p5.s, p5.s, p6.s
+; CHECK-NEXT:    ld1d { z3.d }, p0/z, [x0]
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    add z0.d, z0.d, #1 // =0x1
+; CHECK-NEXT:    add z1.d, z1.d, #1 // =0x1
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    add z2.d, z2.d, #1 // =0x1
+; CHECK-NEXT:    add z3.d, z3.d, #1 // =0x1
+; CHECK-NEXT:    st1d { z0.d }, p3, [x0, #3, mul vl]
+; CHECK-NEXT:    mov z0.h, p4/z, #1 // =0x1
+; CHECK-NEXT:    st1d { z1.d }, p2, [x0, #2, mul vl]
+; CHECK-NEXT:    st1d { z2.d }, p1, [x0, #1, mul vl]
+; CHECK-NEXT:    st1d { z3.d }, p0, [x0]
+; CHECK-NEXT:    incb x0, all, mul #4
+; CHECK-NEXT:    fmov w9, s0
+; CHECK-NEXT:    tbnz w9, #0, .LBB1_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br i1 %c, label %left, label %right
+
+left:
+  br label %loop
+
+right:
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %left ], [ 0, %right ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %left ], [ %start.mask, %right ], [ %new.mask, %loop ]
+  %ptr = getelementptr i64, ptr %x, i64 %base
+  %v = call <vscale x 8 x i64> @llvm.masked.load.nxv8i64(ptr %ptr, i32 8, <vscale x 8 x i1> %update.mask, <vscale x 8 x i64> poison)
+  %add = add <vscale x 8 x i64> %v, splat (i64 1)
+  call void @llvm.masked.store.nxv8i64(<vscale x 8 x i64> %add, ptr %ptr, i32 8, <vscale x 8 x i1> %update.mask)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests that a shared start mask remains available so a later loop can also be rewritten.
+define void @shared_start_mask_between_loops(ptr %x, ptr %y, i64 %n) #0 {
+; CHECK-LABEL: shared_start_mask_between_loops:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    whilelo pn8.d, xzr, x2, vlx4
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    mov p9.b, p8.b
+; CHECK-NEXT:    mov x9, x8
+; CHECK-NEXT:  .LBB2_1: // %loop1
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p2.d, p3.d }, pn9[1]
+; CHECK-NEXT:    pext { p0.d, p1.d }, pn9[0]
+; CHECK-NEXT:    whilelo pn9.d, x9, x2, vlx4
+; CHECK-NEXT:    ld1d { z0.d }, p3/z, [x0, #3, mul vl]
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn9[0]
+; CHECK-NEXT:    ld1d { z1.d }, p2/z, [x0, #2, mul vl]
+; CHECK-NEXT:    uzp1 p4.s, p4.s, p5.s
+; CHECK-NEXT:    pext { p5.d, p6.d }, pn9[1]
+; CHECK-NEXT:    ld1d { z2.d }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    uzp1 p5.s, p5.s, p6.s
+; CHECK-NEXT:    ld1d { z3.d }, p0/z, [x0]
+; CHECK-NEXT:    inch x9
+; CHECK-NEXT:    add z0.d, z0.d, #1 // =0x1
+; CHECK-NEXT:    add z1.d, z1.d, #1 // =0x1
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    add z2.d, z2.d, #1 // =0x1
+; CHECK-NEXT:    add z3.d, z3.d, #1 // =0x1
+; CHECK-NEXT:    st1d { z0.d }, p3, [x0, #3, mul vl]
+; CHECK-NEXT:    mov z0.h, p4/z, #1 // =0x1
+; CHECK-NEXT:    st1d { z1.d }, p2, [x0, #2, mul vl]
+; CHECK-NEXT:    st1d { z2.d }, p1, [x0, #1, mul vl]
+; CHECK-NEXT:    st1d { z3.d }, p0, [x0]
+; CHECK-NEXT:    incb x0, all, mul #4
+; CHECK-NEXT:    fmov w10, s0
+; CHECK-NEXT:    tbnz w10, #0, .LBB2_1
+; CHECK-NEXT:  .LBB2_2: // %loop2
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.d, p1.d }, pn8[0]
+; CHECK-NEXT:    pext { p2.d, p3.d }, pn8[1]
+; CHECK-NEXT:    whilelo pn8.d, x8, x2, vlx4
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn8[0]
+; CHECK-NEXT:    ld1d { z0.d }, p3/z, [x1, #3, mul vl]
+; CHECK-NEXT:    ld1d { z1.d }, p2/z, [x1, #2, mul vl]
+; CHECK-NEXT:    uzp1 p4.s, p4.s, p5.s
+; CHECK-NEXT:    pext { p5.d, p6.d }, pn8[1]
+; CHECK-NEXT:    ld1d { z2.d }, p1/z, [x1, #1, mul vl]
+; CHECK-NEXT:    uzp1 p5.s, p5.s, p6.s
+; CHECK-NEXT:    ld1d { z3.d }, p0/z, [x1]
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    st1d { z0.d }, p3, [x1, #3, mul vl]
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    st1d { z1.d }, p2, [x1, #2, mul vl]
+; CHECK-NEXT:    st1d { z2.d }, p1, [x1, #1, mul vl]
+; CHECK-NEXT:    mov z0.h, p4/z, #1 // =0x1
+; CHECK-NEXT:    st1d { z3.d }, p0, [x1]
+; CHECK-NEXT:    incb x1, all, mul #4
+; CHECK-NEXT:    fmov w9, s0
+; CHECK-NEXT:    tbnz w9, #0, .LBB2_2
+; CHECK-NEXT:  // %bb.3: // %exit
+; CHECK-NEXT:    ret
+entry:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop1
+
+loop1:
+  %base1 = phi i64 [ 0, %entry ], [ %next.base1, %loop1 ]
+  %mask1 = phi <vscale x 8 x i1> [ %start.mask, %entry ], [ %new.mask1, %loop1 ]
+  %ptr1 = getelementptr i64, ptr %x, i64 %base1
+  %v1 = call <vscale x 8 x i64> @llvm.masked.load.nxv8i64(ptr %ptr1, i32 8, <vscale x 8 x i1> %mask1, <vscale x 8 x i64> poison)
+  %add1 = add <vscale x 8 x i64> %v1, splat (i64 1)
+  call void @llvm.masked.store.nxv8i64(<vscale x 8 x i64> %add1, ptr %ptr1, i32 8, <vscale x 8 x i1> %mask1)
+  %next.base1 = add i64 %base1, %step
+  %new.mask1 = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base1, i64 %n)
+  %more1 = extractelement <vscale x 8 x i1> %new.mask1, i64 0
+  br i1 %more1, label %loop1, label %between
+
+between:
+  br label %loop2
+
+loop2:
+  %base2 = phi i64 [ 0, %between ], [ %next.base2, %loop2 ]
+  %mask2 = phi <vscale x 8 x i1> [ %start.mask, %between ], [ %new.mask2, %loop2 ]
+  %ptr2 = getelementptr ptr, ptr %y, i64 %base2
+  %v2 = call <vscale x 8 x ptr> @llvm.masked.load.nxv8p0.p0(ptr %ptr2, i32 8, <vscale x 8 x i1> %mask2, <vscale x 8 x ptr> poison)
+  call void @llvm.masked.store.nxv8p0.p0(<vscale x 8 x ptr> %v2, ptr %ptr2, i32 8, <vscale x 8 x i1> %mask2)
+  %next.base2 = add i64 %base2, %step
+  %new.mask2 = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base2, i64 %n)
+  %more2 = extractelement <vscale x 8 x i1> %new.mask2, i64 0
+  br i1 %more2, label %loop2, label %exit
+
+exit:
+  ret void
+}
+
+; Tests the i8 masked load/store rewrite when the induction operands are i32 and must be extended.
+define void @rewrite_masked_load_store_i8_i32_induction(ptr %x, i32 %n) #0 {
+; CHECK-LABEL: rewrite_masked_load_store_i8_i32_induction:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    mov w8, w1
+; CHECK-NEXT:    rdvl x9, #2
+; CHECK-NEXT:    whilelo pn8.b, xzr, x8, vlx2
+; CHECK-NEXT:  .LBB3_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    mov x10, x9
+; CHECK-NEXT:    pext { p0.b, p1.b }, pn8[0]
+; CHECK-NEXT:    mov w12, w9
+; CHECK-NEXT:    decb x10, all, mul #2
+; CHECK-NEXT:    whilelo pn8.b, x12, x8, vlx2
+; CHECK-NEXT:    incb x9, all, mul #2
+; CHECK-NEXT:    pext { p2.b, p3.b }, pn8[0]
+; CHECK-NEXT:    mov z2.b, p2/z, #1 // =0x1
+; CHECK-NEXT:    sxtw x10, w10
+; CHECK-NEXT:    ld1b { z0.b }, p0/z, [x0, x10]
+; CHECK-NEXT:    add x11, x0, x10
+; CHECK-NEXT:    ld1b { z1.b }, p1/z, [x11, #1, mul vl]
+; CHECK-NEXT:    add z0.b, z0.b, #1 // =0x1
+; CHECK-NEXT:    add z1.b, z1.b, #1 // =0x1
+; CHECK-NEXT:    st1b { z0.b }, p0, [x0, x10]
+; CHECK-NEXT:    fmov w10, s2
+; CHECK-NEXT:    st1b { z1.b }, p1, [x11, #1, mul vl]
+; CHECK-NEXT:    tbnz w10, #0, .LBB3_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 32 x i1> @llvm.get.active.lane.mask.nxv32i1.i32(i32 0, i32 %n)
+  %vscale = call i32 @llvm.vscale.i32()
+  %step = shl i32 %vscale, 5
+  br label %loop
+
+loop:
+  %base = phi i32 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 32 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr = getelementptr i8, ptr %x, i32 %base
+  %v = call <vscale x 32 x i8> @llvm.masked.load.nxv32i8(ptr %ptr, i32 1, <vscale x 32 x i1> %update.mask, <vscale x 32 x i8> poison)
+  %add = add <vscale x 32 x i8> %v, splat (i8 1)
+  call void @llvm.masked.store.nxv32i8(<vscale x 32 x i8> %add, ptr %ptr, i32 1, <vscale x 32 x i1> %update.mask)
+  %next.base = add i32 %base, %step
+  %new.mask = call <vscale x 32 x i1> @llvm.get.active.lane.mask.nxv32i1.i32(i32 %next.base, i32 %n)
+  %more = extractelement <vscale x 32 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests the i32 masked load/store rewrite to predicate-as-counter with vlx2.
+define void @rewrite_masked_load_store_i32_vlx2(ptr %x, i64 %n) #0 {
+; CHECK-LABEL: rewrite_masked_load_store_i32_vlx2:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    whilelo pn8.s, xzr, x1, vlx2
+; CHECK-NEXT:  .LBB4_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.s, p1.s }, pn8[0]
+; CHECK-NEXT:    whilelo pn8.s, x8, x1, vlx2
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    ld1w { z0.s }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    ld1w { z1.s }, p0/z, [x0]
+; CHECK-NEXT:    pext { p2.s, p3.s }, pn8[0]
+; CHECK-NEXT:    uzp1 p2.h, p2.h, p3.h
+; CHECK-NEXT:    add z0.s, z0.s, #1 // =0x1
+; CHECK-NEXT:    add z1.s, z1.s, #1 // =0x1
+; CHECK-NEXT:    mov z2.h, p2/z, #1 // =0x1
+; CHECK-NEXT:    st1w { z0.s }, p1, [x0, #1, mul vl]
+; CHECK-NEXT:    fmov w9, s2
+; CHECK-NEXT:    st1w { z1.s }, p0, [x0]
+; CHECK-NEXT:    incb x0, all, mul #2
+; CHECK-NEXT:    tbnz w9, #0, .LBB4_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr = getelementptr i32, ptr %x, i64 %base
+  %v = call <vscale x 8 x i32> @llvm.masked.load.nxv8i32(ptr %ptr, i32 4, <vscale x 8 x i1> %update.mask, <vscale x 8 x i32> poison)
+  %add = add <vscale x 8 x i32> %v, splat (i32 1)
+  call void @llvm.masked.store.nxv8i32(<vscale x 8 x i32> %add, ptr %ptr, i32 4, <vscale x 8 x i1> %update.mask)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests choosing an i64 predicate-as-counter rewrite while materializing a wide mask for mismatched i32 users.
+define void @mixed_data_vector_types(ptr %x, ptr %y, i64 %n) #0 {
+; CHECK-LABEL: mixed_data_vector_types:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    mov x8, xzr
+; CHECK-NEXT:    whilelo pn8.d, xzr, x2, vlx4
+; CHECK-NEXT:  .LBB5_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.d, p1.d }, pn8[0]
+; CHECK-NEXT:    add x9, x0, x8, lsl #3
+; CHECK-NEXT:    add x10, x1, x8, lsl #2
+; CHECK-NEXT:    uzp1 p2.s, p0.s, p1.s
+; CHECK-NEXT:    ld1d { z0.d }, p0/z, [x0, x8, lsl #3]
+; CHECK-NEXT:    ld1d { z5.d }, p1/z, [x9, #1, mul vl]
+; CHECK-NEXT:    ld1w { z1.s }, p2/z, [x1, x8, lsl #2]
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    pext { p2.d, p3.d }, pn8[1]
+; CHECK-NEXT:    ld1d { z3.d }, p3/z, [x9, #3, mul vl]
+; CHECK-NEXT:    ld1d { z4.d }, p2/z, [x9, #2, mul vl]
+; CHECK-NEXT:    whilelo pn8.d, x8, x2, vlx4
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn8[0]
+; CHECK-NEXT:    uzp1 p0.s, p4.s, p5.s
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn8[1]
+; CHECK-NEXT:    uzp1 p4.s, p4.s, p5.s
+; CHECK-NEXT:    uzp1 p0.h, p0.h, p4.h
+; CHECK-NEXT:    uzp1 p4.s, p2.s, p3.s
+; CHECK-NEXT:    mov z2.h, p0/z, #1 // =0x1
+; CHECK-NEXT:    ld1w { z6.s }, p4/z, [x10, #1, mul vl]
+; CHECK-NEXT:    fmov w9, s2
+; CHECK-NEXT:    // fake_use: $z0
+; CHECK-NEXT:    // fake_use: $z5
+; CHECK-NEXT:    // fake_use: $z4
+; CHECK-NEXT:    // fake_use: $z3
+; CHECK-NEXT:    // fake_use: $z1
+; CHECK-NEXT:    // fake_use: $z6
+; CHECK-NEXT:    tbnz w9, #0, .LBB5_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr64 = getelementptr i64, ptr %x, i64 %base
+  %ptr32 = getelementptr i32, ptr %y, i64 %base
+  %v64 = call <vscale x 8 x i64> @llvm.masked.load.nxv8i64(ptr %ptr64, i32 8, <vscale x 8 x i1> %update.mask, <vscale x 8 x i64> poison)
+  %v32 = call <vscale x 8 x i32> @llvm.masked.load.nxv8i32(ptr %ptr32, i32 4, <vscale x 8 x i1> %update.mask, <vscale x 8 x i32> poison)
+  call void (...) @llvm.fake.use(<vscale x 8 x i64> %v64)
+  call void (...) @llvm.fake.use(<vscale x 8 x i32> %v32)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests preferring the most common masked memory access size when selecting the predicate-as-counter element size.
+define void @prefer_more_common_masked_access_size(ptr %x16, ptr %y32, i64 %n) #0 {
+; CHECK-LABEL: prefer_more_common_masked_access_size:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    movi v0.2d, #0000000000000000
+; CHECK-NEXT:    mov x8, xzr
+; CHECK-NEXT:    whilelo pn8.h, xzr, x2, vlx2
+; CHECK-NEXT:  .LBB6_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.h, p1.h }, pn8[0]
+; CHECK-NEXT:    add x9, x0, x8, lsl #1
+; CHECK-NEXT:    punpklo p2.h, p0.b
+; CHECK-NEXT:    st1h { z0.h }, p0, [x0, x8, lsl #1]
+; CHECK-NEXT:    st1h { z0.h }, p1, [x9, #1, mul vl]
+; CHECK-NEXT:    add x9, x1, x8, lsl #2
+; CHECK-NEXT:    punpkhi p0.h, p0.b
+; CHECK-NEXT:    ld1w { z1.s }, p2/z, [x1, x8, lsl #2]
+; CHECK-NEXT:    incb x8
+; CHECK-NEXT:    punpkhi p2.h, p1.b
+; CHECK-NEXT:    punpklo p1.h, p1.b
+; CHECK-NEXT:    ld1w { z5.s }, p0/z, [x9, #1, mul vl]
+; CHECK-NEXT:    ld1w { z3.s }, p2/z, [x9, #3, mul vl]
+; CHECK-NEXT:    whilelo pn8.h, x8, x2, vlx2
+; CHECK-NEXT:    ld1w { z4.s }, p1/z, [x9, #2, mul vl]
+; CHECK-NEXT:    pext { p3.h, p4.h }, pn8[0]
+; CHECK-NEXT:    uzp1 p3.b, p3.b, p4.b
+; CHECK-NEXT:    mov z2.b, p3/z, #1 // =0x1
+; CHECK-NEXT:    fmov w9, s2
+; CHECK-NEXT:    // fake_use: $z1
+; CHECK-NEXT:    // fake_use: $z5
+; CHECK-NEXT:    // fake_use: $z4
+; CHECK-NEXT:    // fake_use: $z3
+; CHECK-NEXT:    tbnz w9, #0, .LBB6_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 4
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 16 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr16 = getelementptr i16, ptr %x16, i64 %base
+  %ptr32 = getelementptr i32, ptr %y32, i64 %base
+  call void @llvm.masked.store.nxv16i16(<vscale x 16 x i16> zeroinitializer, ptr %ptr16, i32 2, <vscale x 16 x i1> %update.mask)
+  call void @llvm.masked.store.nxv16i16(<vscale x 16 x i16> zeroinitializer, ptr %ptr16, i32 2, <vscale x 16 x i1> %update.mask)
+  %v32 = call <vscale x 16 x i32> @llvm.masked.load.nxv16i32(ptr %ptr32, i32 4, <vscale x 16 x i1> %update.mask, <vscale x 16 x i32> poison)
+  call void (...) @llvm.fake.use(<vscale x 16 x i32> %v32)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 16 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests breaking equal access-size counts in favor of the larger access size.
+define void @prefer_larger_access_size_on_tie(ptr %x16, ptr %y32, i64 %n) #0 {
+; CHECK-LABEL: prefer_larger_access_size_on_tie:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    movi v0.2d, #0000000000000000
+; CHECK-NEXT:    mov x8, xzr
+; CHECK-NEXT:    whilelo pn8.s, xzr, x2, vlx4
+; CHECK-NEXT:  .LBB7_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.s, p1.s }, pn8[0]
+; CHECK-NEXT:    pext { p2.s, p3.s }, pn8[1]
+; CHECK-NEXT:    add x9, x0, x8, lsl #1
+; CHECK-NEXT:    uzp1 p4.h, p0.h, p1.h
+; CHECK-NEXT:    uzp1 p5.h, p2.h, p3.h
+; CHECK-NEXT:    st1h { z0.h }, p4, [x0, x8, lsl #1]
+; CHECK-NEXT:    st1h { z0.h }, p5, [x9, #1, mul vl]
+; CHECK-NEXT:    add x9, x1, x8, lsl #2
+; CHECK-NEXT:    ld1w { z1.s }, p0/z, [x1, x8, lsl #2]
+; CHECK-NEXT:    incb x8
+; CHECK-NEXT:    ld1w { z3.s }, p3/z, [x9, #3, mul vl]
+; CHECK-NEXT:    ld1w { z4.s }, p2/z, [x9, #2, mul vl]
+; CHECK-NEXT:    ld1w { z5.s }, p1/z, [x9, #1, mul vl]
+; CHECK-NEXT:    whilelo pn8.s, x8, x2, vlx4
+; CHECK-NEXT:    pext { p4.s, p5.s }, pn8[0]
+; CHECK-NEXT:    uzp1 p0.h, p4.h, p5.h
+; CHECK-NEXT:    pext { p4.s, p5.s }, pn8[1]
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    uzp1 p0.b, p0.b, p4.b
+; CHECK-NEXT:    mov z2.b, p0/z, #1 // =0x1
+; CHECK-NEXT:    fmov w9, s2
+; CHECK-NEXT:    // fake_use: $z1
+; CHECK-NEXT:    // fake_use: $z5
+; CHECK-NEXT:    // fake_use: $z4
+; CHECK-NEXT:    // fake_use: $z3
+; CHECK-NEXT:    tbnz w9, #0, .LBB7_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 4
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 16 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr16 = getelementptr i16, ptr %x16, i64 %base
+  %ptr32 = getelementptr i32, ptr %y32, i64 %base
+  call void @llvm.masked.store.nxv16i16(<vscale x 16 x i16> zeroinitializer, ptr %ptr16, i32 2, <vscale x 16 x i1> %update.mask)
+  %v32 = call <vscale x 16 x i32> @llvm.masked.load.nxv16i32(ptr %ptr32, i32 4, <vscale x 16 x i1> %update.mask, <vscale x 16 x i32> poison)
+  call void (...) @llvm.fake.use(<vscale x 16 x i32> %v32)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 16 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Tests rewriting extractelement when the extracted lane is within the first legal predicate section.
+define void @rewrite_extractelement_within_first_section_i32(ptr %x, i64 %n) #0 {
+; CHECK-LABEL: rewrite_extractelement_within_first_section_i32:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    whilelo pn8.s, xzr, x1, vlx2
+; CHECK-NEXT:  .LBB8_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.s, p1.s }, pn8[0]
+; CHECK-NEXT:    whilelo pn8.s, x8, x1, vlx2
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    pext { p2.s, p3.s }, pn8[0]
+; CHECK-NEXT:    ld1w { z1.s }, p0/z, [x0]
+; CHECK-NEXT:    uzp1 p2.h, p2.h, p3.h
+; CHECK-NEXT:    mov z0.h, p2/z, #1 // =0x1
+; CHECK-NEXT:    umov w9, v0.h[3]
+; CHECK-NEXT:    ld1w { z0.s }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    incb x0, all, mul #2
+; CHECK-NEXT:    // fake_use: $z1
+; CHECK-NEXT:    // fake_use: $z0
+; CHECK-NEXT:    tbnz w9, #0, .LBB8_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr = getelementptr i32, ptr %x, i64 %base
+  %v = call <vscale x 8 x i32> @llvm.masked.load.nxv8i32(ptr %ptr, i32 4, <vscale x 8 x i1> %update.mask, <vscale x 8 x i32> poison)
+  call void (...) @llvm.fake.use(<vscale x 8 x i32> %v)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 3
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Negative test: masked loads with a non-poison passthru value are not rewritten to pn loads.
+define void @masked_load_passthru_non_poison(ptr %x, i64 %n) #0 {
+; CHECK-LABEL: masked_load_passthru_non_poison:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    whilelo pn8.d, xzr, x1, vlx4
+; CHECK-NEXT:  .LBB9_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p2.d, p3.d }, pn8[1]
+; CHECK-NEXT:    pext { p0.d, p1.d }, pn8[0]
+; CHECK-NEXT:    whilelo pn8.d, x8, x1, vlx4
+; CHECK-NEXT:    ld1d { z0.d }, p3/z, [x0, #3, mul vl]
+; CHECK-NEXT:    pext { p4.d, p5.d }, pn8[0]
+; CHECK-NEXT:    ld1d { z1.d }, p2/z, [x0, #2, mul vl]
+; CHECK-NEXT:    uzp1 p4.s, p4.s, p5.s
+; CHECK-NEXT:    pext { p5.d, p6.d }, pn8[1]
+; CHECK-NEXT:    ld1d { z2.d }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    uzp1 p5.s, p5.s, p6.s
+; CHECK-NEXT:    ld1d { z3.d }, p0/z, [x0]
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    add z0.d, z0.d, #1 // =0x1
+; CHECK-NEXT:    add z1.d, z1.d, #1 // =0x1
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    add z2.d, z2.d, #1 // =0x1
+; CHECK-NEXT:    add z3.d, z3.d, #1 // =0x1
+; CHECK-NEXT:    st1d { z0.d }, p3, [x0, #3, mul vl]
+; CHECK-NEXT:    mov z0.h, p4/z, #1 // =0x1
+; CHECK-NEXT:    st1d { z1.d }, p2, [x0, #2, mul vl]
+; CHECK-NEXT:    st1d { z2.d }, p1, [x0, #1, mul vl]
+; CHECK-NEXT:    st1d { z3.d }, p0, [x0]
+; CHECK-NEXT:    incb x0, all, mul #4
+; CHECK-NEXT:    fmov w9, s0
+; CHECK-NEXT:    tbnz w9, #0, .LBB9_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr = getelementptr i64, ptr %x, i64 %base
+  %v = call <vscale x 8 x i64> @llvm.masked.load.nxv8i64(ptr %ptr, i32 8, <vscale x 8 x i1> %update.mask, <vscale x 8 x i64> zeroinitializer)
+  %add = add <vscale x 8 x i64> %v, splat (i64 1)
+  call void @llvm.masked.store.nxv8i64(<vscale x 8 x i64> %add, ptr %ptr, i32 8, <vscale x 8 x i1> %update.mask)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 0
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+; Negative test: extractelement is not rewritten when the extracted lane is outside the first predicate section.
+define void @extractelement_outside_first_section_i32(ptr %x, i64 %n) #0 {
+; CHECK-LABEL: extractelement_outside_first_section_i32:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    cnth x8
+; CHECK-NEXT:    whilelo pn8.s, xzr, x1, vlx2
+; CHECK-NEXT:  .LBB10_1: // %loop
+; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
+; CHECK-NEXT:    pext { p0.s, p1.s }, pn8[0]
+; CHECK-NEXT:    whilelo pn8.s, x8, x1, vlx2
+; CHECK-NEXT:    inch x8
+; CHECK-NEXT:    pext { p2.s, p3.s }, pn8[0]
+; CHECK-NEXT:    ld1w { z1.s }, p0/z, [x0]
+; CHECK-NEXT:    uzp1 p2.h, p2.h, p3.h
+; CHECK-NEXT:    mov z0.h, p2/z, #1 // =0x1
+; CHECK-NEXT:    umov w9, v0.h[6]
+; CHECK-NEXT:    ld1w { z0.s }, p1/z, [x0, #1, mul vl]
+; CHECK-NEXT:    incb x0, all, mul #2
+; CHECK-NEXT:    // fake_use: $z1
+; CHECK-NEXT:    // fake_use: $z0
+; CHECK-NEXT:    tbnz w9, #0, .LBB10_1
+; CHECK-NEXT:  // %bb.2: // %exit
+; CHECK-NEXT:    ret
+entry:
+  br label %preheader
+
+preheader:
+  %start.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 0, i64 %n)
+  %vscale = call i64 @llvm.vscale.i64()
+  %step = shl i64 %vscale, 3
+  br label %loop
+
+loop:
+  %base = phi i64 [ 0, %preheader ], [ %next.base, %loop ]
+  %update.mask = phi <vscale x 8 x i1> [ %start.mask, %preheader ], [ %new.mask, %loop ]
+  %ptr = getelementptr i32, ptr %x, i64 %base
+  %v = call <vscale x 8 x i32> @llvm.masked.load.nxv8i32(ptr %ptr, i32 4, <vscale x 8 x i1> %update.mask, <vscale x 8 x i32> poison)
+  call void (...) @llvm.fake.use(<vscale x 8 x i32> %v)
+  %next.base = add i64 %base, %step
+  %new.mask = call <vscale x 8 x i1> @llvm.get.active.lane.mask.nxv8i1.i64(i64 %next.base, i64 %n)
+  %more = extractelement <vscale x 8 x i1> %new.mask, i64 6
+  br i1 %more, label %loop, label %exit
+
+exit:
+  ret void
+}
+
+attributes #0 = { "target-features"="+sve2p1" }

>From 589769cb7d6c51c4de520f769922b4598c5f5977 Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Thu, 3 Sep 2026 15:01:32 +0000
Subject: [PATCH 2/5] Fix whitespace

---
 .../lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp | 1 +
 1 file changed, 1 insertion(+)

diff --git a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
index b701e6c176146..c46df4bd3f84a 100644
--- a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
+++ b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
@@ -228,6 +228,7 @@ static Value *createWhileLO(IRBuilder<> &Builder, unsigned ElementSizeInBits,
   return Builder.CreateCall(
       WhileLO, {Start, End, Builder.getInt32(VectorScale)}, "pac.mask");
 }
+
 class AArch64PredicateAsCounterLoopRewrites : public LoopPass {
 public:
   static char ID;

>From 1cbe7d42b4aeae4cf06d6a583c58ebea2460e67a Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Thu, 3 Sep 2026 15:13:38 +0000
Subject: [PATCH 3/5] Fix format

---
 .../Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp    | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
index c46df4bd3f84a..94c648913fb5b 100644
--- a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
+++ b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
@@ -419,7 +419,7 @@ bool AArch64PredicateAsCounterLoopRewrites::rewriteCandidate(
   NewPhi->addIncoming(NewNext, C.Latch);
 
   auto RewriteUses = [&](Instruction *OldMask, Value *Count,
-                         function_ref<bool(Use & U)> Predicate = nullptr) {
+                         function_ref<bool(Use &U)> Predicate = nullptr) {
     SmallVector<Use *, 8> UsesToRewrite;
     for (Use &U : OldMask->uses()) {
       if (!Predicate || Predicate(U))

>From cd81f2de50641e6747a559f081b0f504bf2cb04d Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Tue, 8 Sep 2026 13:28:01 +0000
Subject: [PATCH 4/5] Fixups

---
 .../AArch64PredicateAsCounterLoopRewrites.cpp | 89 ++++++++-----------
 .../predicate-as-counter-loop-rewrites.ll     | 34 +++----
 2 files changed, 53 insertions(+), 70 deletions(-)

diff --git a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
index 94c648913fb5b..d8f6597529ff0 100644
--- a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
+++ b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
@@ -84,16 +84,12 @@
 
 using namespace llvm;
 
-#define DEBUG_TYPE "aarch64-predicate-as-counter-loop-rewrites"
+#define DEBUG_TYPE "aarch64-pn-loop-rewrites"
 namespace {
 
 STATISTIC(LoopsRewritten, "Number of loops rewritten");
 
 struct MaskRewriteCandidate {
-  /// The preheader block for the loop.
-  BasicBlock *Preheader = nullptr;
-  /// The latch block for the loop.
-  BasicBlock *Latch = nullptr;
   /// The mask phi node (used by masked operations within the loop).
   PHINode *MaskPhi = nullptr;
   /// The initial value for the mask (incoming value from the preheader).
@@ -108,7 +104,7 @@ struct MaskRewriteCandidate {
 
 static void logLoopBailout(const Loop &L, const Twine &Reason) {
   LLVM_DEBUG({
-    dbgs() << "PAC loop rewrite: skipping loop with header ";
+    dbgs() << "PN loop rewrite: skipping loop with header ";
     L.getHeader()->printAsOperand(dbgs(), /*PrintType=*/false);
     dbgs() << ": " << Reason << "\n";
   });
@@ -116,7 +112,7 @@ static void logLoopBailout(const Loop &L, const Twine &Reason) {
 
 static void logMatchFailure(const PHINode &Phi, const Twine &Reason) {
   LLVM_DEBUG({
-    dbgs() << "PAC loop rewrite: failed to match mask phi ";
+    dbgs() << "PN loop rewrite: failed to match mask phi ";
     Phi.printAsOperand(dbgs(), /*PrintType=*/false);
     dbgs() << ": " << Reason << "\n";
   });
@@ -133,15 +129,13 @@ static ElementCount getSVEElementCount(unsigned ElementSizeInBits) {
                                    ElementSizeInBits);
 }
 
-/// Returns the most common scalar access size of masked loads/stores in loop
-/// \p L where \p MaskPhi is used as the mask. Ties are broken in favor of
-/// larger access sizes.
+/// Returns the largest scalar access size of masked loads/stores in loop \p L
+/// where \p MaskPhi is used as the mask. TODO: This heuristic may need
+/// refinement based on the frequency of different access sizes.
 static std::optional<unsigned>
-getMostCommonMaskedMemAccessSizeInBits(const Loop &L, PHINode &MaskPhi) {
+getLargestMaskedMemAccessSizeInBits(const Loop &L, PHINode &MaskPhi) {
   const DataLayout &DL = MaskPhi.getModule()->getDataLayout();
-  DenseMap<unsigned, unsigned> AccessSizeCounts;
-  std::optional<unsigned> BestAccessSizeInBits;
-  unsigned BestAccessSizeCount = 0;
+  std::optional<unsigned> LargestAccessSizeInBits;
 
   for (User *U : MaskPhi.users()) {
     auto *II = dyn_cast<IntrinsicInst>(U);
@@ -157,17 +151,11 @@ getMostCommonMaskedMemAccessSizeInBits(const Loop &L, PHINode &MaskPhi) {
       continue;
 
     unsigned AccessSizeInBits = getScalarSizeInBits(DL, II->getAccessType());
-    unsigned AccessSizeCount = ++AccessSizeCounts[AccessSizeInBits];
-
-    if (!BestAccessSizeInBits || AccessSizeCount > BestAccessSizeCount ||
-        (AccessSizeCount == BestAccessSizeCount &&
-         AccessSizeInBits > *BestAccessSizeInBits)) {
-      BestAccessSizeInBits = AccessSizeInBits;
-      BestAccessSizeCount = AccessSizeCount;
-    }
+    if (!LargestAccessSizeInBits || AccessSizeInBits > LargestAccessSizeInBits)
+      LargestAccessSizeInBits = AccessSizeInBits;
   }
 
-  return BestAccessSizeInBits;
+  return LargestAccessSizeInBits;
 }
 
 /// Returns the predicate-as-counter whilelo intrinsic ID for
@@ -201,13 +189,13 @@ static Value *buildWideMask(IRBuilder<> &Builder, const MaskRewriteCandidate &C,
   Value *WideMask = PoisonValue::get(C.MaskPhi->getType());
   for (unsigned PairOffset = 0; PairOffset != C.VectorScale / 2; ++PairOffset) {
     auto *Pair = Builder.CreateCall(
-        PExtX2, {Count, Builder.getInt32(PairOffset)}, "pac.pext.pair");
+        PExtX2, {Count, Builder.getInt32(PairOffset)}, "pn.pext.pair");
     for (unsigned SliceInPair = 0; SliceInPair != 2; ++SliceInPair) {
-      Value *Part = Builder.CreateExtractValue(Pair, SliceInPair, "pac.pext");
+      Value *Part = Builder.CreateExtractValue(Pair, SliceInPair, "pn.pext");
       unsigned Slice = PairOffset * 2 + SliceInPair;
       WideMask = Builder.CreateInsertVector(
           C.MaskPhi->getType(), WideMask, Part,
-          Slice * LegalEC.getKnownMinValue(), "pac.mask");
+          Slice * LegalEC.getKnownMinValue(), "pn.mask");
     }
   }
 
@@ -226,7 +214,7 @@ static Value *createWhileLO(IRBuilder<> &Builder, unsigned ElementSizeInBits,
   auto ID = getWhileLOIntrinsic(ElementSizeInBits);
   FunctionCallee WhileLO = Intrinsic::getOrInsertDeclaration(M, ID);
   return Builder.CreateCall(
-      WhileLO, {Start, End, Builder.getInt32(VectorScale)}, "pac.mask");
+      WhileLO, {Start, End, Builder.getInt32(VectorScale)}, "pn.mask");
 }
 
 class AArch64PredicateAsCounterLoopRewrites : public LoopPass {
@@ -286,11 +274,19 @@ bool AArch64PredicateAsCounterLoopRewrites::runOnLoop(Loop *L,
     return false;
   }
 
+  if (!L->getLoopPreheader()) {
+    logLoopBailout(*L, "loop has no preheader");
+    return false;
+  }
+  if (!L->getLoopLatch()) {
+    logLoopBailout(*L, "loop has no latch");
+    return false;
+  }
+
   bool Changed = false;
   BasicBlock *Header = L->getHeader();
   for (PHINode &Phi : make_early_inc_range(Header->phis())) {
-    auto Candidate = matchMaskPhi(*L, Phi);
-    if (Candidate)
+    if (std::optional<MaskRewriteCandidate> Candidate = matchMaskPhi(*L, Phi))
       Changed |= rewriteCandidate(*Candidate, *L);
   }
 
@@ -310,17 +306,6 @@ static IntrinsicInst *getGetActiveLaneMask(Value *V) {
 std::optional<MaskRewriteCandidate>
 AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
                                                     PHINode &Phi) const {
-  BasicBlock *Preheader = L.getLoopPreheader();
-  BasicBlock *Latch = L.getLoopLatch();
-  if (!Preheader) {
-    logMatchFailure(Phi, "loop has no preheader");
-    return std::nullopt;
-  }
-  if (!Latch) {
-    logMatchFailure(Phi, "loop has no latch");
-    return std::nullopt;
-  }
-
   auto *PhiTy = dyn_cast<ScalableVectorType>(Phi.getType());
   if (!PhiTy || !PhiTy->getElementType()->isIntegerTy(1)) {
     logMatchFailure(Phi, "phi type is not a scalable i1 vector mask");
@@ -333,9 +318,10 @@ AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
     return std::nullopt;
   }
 
-  auto *StartMask =
-      getGetActiveLaneMask(Phi.getIncomingValueForBlock(Preheader));
-  auto *NextMask = getGetActiveLaneMask(Phi.getIncomingValueForBlock(Latch));
+  Value *StartValue = Phi.getIncomingValueForBlock(L.getLoopPreheader());
+  Value *NextValue = Phi.getIncomingValueForBlock(L.getLoopLatch());
+  IntrinsicInst *StartMask = getGetActiveLaneMask(StartValue);
+  IntrinsicInst *NextMask = getGetActiveLaneMask(NextValue);
   if (!StartMask) {
     logMatchFailure(Phi,
                     "preheader incoming value is not get_active_lane_mask");
@@ -363,7 +349,7 @@ AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
   }
 
   std::optional<unsigned> PreferredMaskElementSizeInBits =
-      getMostCommonMaskedMemAccessSizeInBits(L, Phi);
+      getLargestMaskedMemAccessSizeInBits(L, Phi);
   if (!PreferredMaskElementSizeInBits) {
     logMatchFailure(Phi, "mask phi has no masked load/store users in the loop");
     return std::nullopt;
@@ -376,7 +362,7 @@ AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
   }
 
   unsigned SVEMaskElements =
-      getSVEElementCount(*PreferredMaskElementSizeInBits).getKnownMinValue();
+      AArch64::SVEBitsPerBlock / *PreferredMaskElementSizeInBits;
   if (WideMaskElements <= SVEMaskElements) {
     logMatchFailure(Phi, Twine("wide mask element count ")
                              .concat(Twine(WideMaskElements))
@@ -392,12 +378,7 @@ AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
     return std::nullopt;
   }
 
-  return MaskRewriteCandidate{Preheader,
-                              Latch,
-                              &Phi,
-                              StartMask,
-                              NextMask,
-                              VectorScale,
+  return MaskRewriteCandidate{&Phi, StartMask, NextMask, VectorScale,
                               *PreferredMaskElementSizeInBits};
 }
 
@@ -415,11 +396,11 @@ bool AArch64PredicateAsCounterLoopRewrites::rewriteCandidate(
   Builder.SetInsertPoint(C.MaskPhi);
   auto *NewPhi =
       Builder.CreatePHI(NewStart->getType(), 2, C.MaskPhi->getName() + ".pn");
-  NewPhi->addIncoming(NewStart, C.Preheader);
-  NewPhi->addIncoming(NewNext, C.Latch);
+  NewPhi->addIncoming(NewStart, L.getLoopPreheader());
+  NewPhi->addIncoming(NewNext, L.getLoopLatch());
 
   auto RewriteUses = [&](Instruction *OldMask, Value *Count,
-                         function_ref<bool(Use &U)> Predicate = nullptr) {
+                         function_ref<bool(Use & U)> Predicate = nullptr) {
     SmallVector<Use *, 8> UsesToRewrite;
     for (Use &U : OldMask->uses()) {
       if (!Predicate || Predicate(U))
diff --git a/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll b/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll
index b064c81107795..4e66407d971c8 100644
--- a/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll
+++ b/llvm/test/CodeGen/AArch64/predicate-as-counter-loop-rewrites.ll
@@ -397,27 +397,29 @@ define void @prefer_more_common_masked_access_size(ptr %x16, ptr %y32, i64 %n) #
 ; CHECK:       // %bb.0: // %entry
 ; CHECK-NEXT:    movi v0.2d, #0000000000000000
 ; CHECK-NEXT:    mov x8, xzr
-; CHECK-NEXT:    whilelo pn8.h, xzr, x2, vlx2
+; CHECK-NEXT:    whilelo pn8.s, xzr, x2, vlx4
 ; CHECK-NEXT:  .LBB6_1: // %loop
 ; CHECK-NEXT:    // =>This Inner Loop Header: Depth=1
-; CHECK-NEXT:    pext { p0.h, p1.h }, pn8[0]
+; CHECK-NEXT:    pext { p0.s, p1.s }, pn8[0]
+; CHECK-NEXT:    pext { p2.s, p3.s }, pn8[1]
 ; CHECK-NEXT:    add x9, x0, x8, lsl #1
-; CHECK-NEXT:    punpklo p2.h, p0.b
-; CHECK-NEXT:    st1h { z0.h }, p0, [x0, x8, lsl #1]
-; CHECK-NEXT:    st1h { z0.h }, p1, [x9, #1, mul vl]
+; CHECK-NEXT:    uzp1 p4.h, p0.h, p1.h
+; CHECK-NEXT:    uzp1 p5.h, p2.h, p3.h
+; CHECK-NEXT:    st1h { z0.h }, p4, [x0, x8, lsl #1]
+; CHECK-NEXT:    st1h { z0.h }, p5, [x9, #1, mul vl]
 ; CHECK-NEXT:    add x9, x1, x8, lsl #2
-; CHECK-NEXT:    punpkhi p0.h, p0.b
-; CHECK-NEXT:    ld1w { z1.s }, p2/z, [x1, x8, lsl #2]
+; CHECK-NEXT:    ld1w { z1.s }, p0/z, [x1, x8, lsl #2]
 ; CHECK-NEXT:    incb x8
-; CHECK-NEXT:    punpkhi p2.h, p1.b
-; CHECK-NEXT:    punpklo p1.h, p1.b
-; CHECK-NEXT:    ld1w { z5.s }, p0/z, [x9, #1, mul vl]
-; CHECK-NEXT:    ld1w { z3.s }, p2/z, [x9, #3, mul vl]
-; CHECK-NEXT:    whilelo pn8.h, x8, x2, vlx2
-; CHECK-NEXT:    ld1w { z4.s }, p1/z, [x9, #2, mul vl]
-; CHECK-NEXT:    pext { p3.h, p4.h }, pn8[0]
-; CHECK-NEXT:    uzp1 p3.b, p3.b, p4.b
-; CHECK-NEXT:    mov z2.b, p3/z, #1 // =0x1
+; CHECK-NEXT:    ld1w { z3.s }, p3/z, [x9, #3, mul vl]
+; CHECK-NEXT:    ld1w { z4.s }, p2/z, [x9, #2, mul vl]
+; CHECK-NEXT:    ld1w { z5.s }, p1/z, [x9, #1, mul vl]
+; CHECK-NEXT:    whilelo pn8.s, x8, x2, vlx4
+; CHECK-NEXT:    pext { p4.s, p5.s }, pn8[0]
+; CHECK-NEXT:    uzp1 p0.h, p4.h, p5.h
+; CHECK-NEXT:    pext { p4.s, p5.s }, pn8[1]
+; CHECK-NEXT:    uzp1 p4.h, p4.h, p5.h
+; CHECK-NEXT:    uzp1 p0.b, p0.b, p4.b
+; CHECK-NEXT:    mov z2.b, p0/z, #1 // =0x1
 ; CHECK-NEXT:    fmov w9, s2
 ; CHECK-NEXT:    // fake_use: $z1
 ; CHECK-NEXT:    // fake_use: $z5

>From 842f7afbb68049b6360a0f0525d693d8fa5411a0 Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Tue, 8 Sep 2026 14:02:29 +0000
Subject: [PATCH 5/5] Format

---
 .../Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp    | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
index d8f6597529ff0..e8a3f2f6ba0a6 100644
--- a/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
+++ b/llvm/lib/Target/AArch64/AArch64PredicateAsCounterLoopRewrites.cpp
@@ -400,7 +400,7 @@ bool AArch64PredicateAsCounterLoopRewrites::rewriteCandidate(
   NewPhi->addIncoming(NewNext, L.getLoopLatch());
 
   auto RewriteUses = [&](Instruction *OldMask, Value *Count,
-                         function_ref<bool(Use & U)> Predicate = nullptr) {
+                         function_ref<bool(Use &U)> Predicate = nullptr) {
     SmallVector<Use *, 8> UsesToRewrite;
     for (Use &U : OldMask->uses()) {
       if (!Predicate || Predicate(U))



More information about the llvm-commits mailing list