[llvm] [AArch64] Add predicate-as-counter loop rewrite pass (PR #220960)
Paul Walker via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 11 05:38:03 PDT 2026
================
@@ -0,0 +1,440 @@
+//===- 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-pn-loop-rewrites"
+namespace {
+
+STATISTIC(LoopsRewritten, "Number of loops rewritten");
+
+struct MaskRewriteCandidate {
+ /// 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() << "PN 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() << "PN 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 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>
+getLargestMaskedMemAccessSizeInBits(const Loop &L, PHINode &MaskPhi) {
+ const DataLayout &DL = MaskPhi.getModule()->getDataLayout();
+ std::optional<unsigned> LargestAccessSizeInBits;
+
+ 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());
+ if (!LargestAccessSizeInBits || AccessSizeInBits > LargestAccessSizeInBits)
+ LargestAccessSizeInBits = AccessSizeInBits;
+ }
+
+ return LargestAccessSizeInBits;
+}
+
+/// 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)}, "pn.pext.pair");
+ for (unsigned SliceInPair = 0; SliceInPair != 2; ++SliceInPair) {
+ Value *Part = Builder.CreateExtractValue(Pair, SliceInPair, "pn.pext");
+ unsigned Slice = PairOffset * 2 + SliceInPair;
+ WideMask = Builder.CreateInsertVector(
+ C.MaskPhi->getType(), WideMask, Part,
+ Slice * LegalEC.getKnownMinValue(), "pn.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)}, "pn.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;
+ }
+
+ 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())) {
+ if (std::optional<MaskRewriteCandidate> Candidate = matchMaskPhi(*L, Phi))
+ 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 {
+ auto *PhiTy = dyn_cast<ScalableVectorType>(Phi.getType());
+ if (!PhiTy || !PhiTy->getElementType()->isIntegerTy(1)) {
+ logMatchFailure(Phi, "phi type is not a scalable i1 vector mask");
----------------
paulwalker-arm wrote:
This has the potential to generate a lot of noise? As the lowest bar to entry for the transformation I think logging the positive should be sufficient?
https://github.com/llvm/llvm-project/pull/220960
More information about the llvm-commits
mailing list