1//===- AArch64PredicateAsCounterLoopRewrites.cpp --------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM
4// Exceptions.
5// See https://llvm.org/LICENSE.txt for license information.
6// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
7//
8//===----------------------------------------------------------------------===//
9//
10// Rewrites IR for loop-carried wide masks that can be represented as
11// predicate-as-counter values. This applies when the mask is used by load/store
12// operations that can be mapped to multi-vector instructions (with +sve2p1).
13//
14// For example, a loop like:
15//
16// entry:
17// %step = vscale x 64
18// %mask.entry = @get.active.lane.mask(0, %n)
19//
20// loop:
21// %iv = phi i64 [0, entry], [%iv.next, loop]
22// %mask = phi <vscale x 64 x i1> [%mask.entry, entry],
23// [%mask.next, loop]
24//
25// %load = load <vscale x 64 x i8> %src[%iv], %mask
26// store <vscale x 64 x i8> %load, %dst[%iv], %mask
27//
28// %iv.next = %iv + %step
29// %mask.next = @get.active.lane.mask(%iv.next, %n)
30// br first.active(%mask.next), loop, exit
31//
32// Could be rewritten to:
33//
34// entry:
35// %step = vscale x 64
36// %mask.entry = @whilelo.c8(0, %n, VLx4)
37//
38// loop:
39// %iv = phi i64 [0, entry], [%iv.next, loop]
40// %mask = phi target("aarch64.svcount") [%mask.entry, entry],
41// [%mask.next, loop]
42//
43// %load = @ld1.pn.x4 <4 x <vscale x 16 x i8>> %src[%iv], %mask
44// @st1.pn.x4 <4 x <vscale x 16 x i8>> %load, %dest[%iv], %mask
45//
46// %iv.next = %iv + %step
47// %mask.next = @whilelo.c8(%iv.next, %n, VLx4)
48// br first.active(@pext(%mask.next, 0)), loop, exit
49//
50// This replaces the `get.active.lane.mask` intrinsics with AArch64
51// predicate-as-counter `whilelo` intrinsics and updates the mask phi to use the
52// `aarch64.svcount` target type. Within the loop, load/store users are mapped
53// to multi-vector load/store intrinsics where possible. Users that cannot be
54// mapped to multi-vector instructions materialize vector masks using the `pext`
55// intrinsic (which extracts vector predicates from a predicate-as-counter).
56//
57// This pass may be a temporary solution that is removed if we gain support
58// for target-specific VPlan transforms in the loop vectorizer.
59//
60//===----------------------------------------------------------------------===//
61
62#include "AArch64.h"
63#include "AArch64Subtarget.h"
64#include "AArch64TargetMachine.h"
65#include "llvm/ADT/DenseMap.h"
66#include "llvm/ADT/SmallVector.h"
67#include "llvm/ADT/Statistic.h"
68#include "llvm/Analysis/LoopInfo.h"
69#include "llvm/Analysis/LoopPass.h"
70#include "llvm/CodeGen/TargetPassConfig.h"
71#include "llvm/IR/Attributes.h"
72#include "llvm/IR/DataLayout.h"
73#include "llvm/IR/IRBuilder.h"
74#include "llvm/IR/IntrinsicInst.h"
75#include "llvm/IR/Intrinsics.h"
76#include "llvm/IR/IntrinsicsAArch64.h"
77#include "llvm/IR/Module.h"
78#include "llvm/InitializePasses.h"
79#include "llvm/Pass.h"
80#include "llvm/Support/Debug.h"
81#include "llvm/Transforms/Utils.h"
82#include "llvm/Transforms/Utils/Local.h"
83#include <optional>
84
85using namespace llvm;
86
87#define DEBUG_TYPE "aarch64-pn-loop-rewrites"
88namespace {
89
90STATISTIC(LoopsRewritten, "Number of loops rewritten");
91
92struct MaskRewriteCandidate {
93 /// The mask phi node (used by masked operations within the loop).
94 PHINode *MaskPhi = nullptr;
95 /// The initial value for the mask (incoming value from the preheader).
96 IntrinsicInst *StartMask = nullptr;
97 /// The updated value for the mask (incoming value from the loop latch).
98 IntrinsicInst *NextMask = nullptr;
99 /// The multi-vector scale for the predicate-as-counter (2 or 4).
100 unsigned VectorScale = 0;
101 /// The element size (in bits) for the predicate-as-counter.
102 unsigned ElementSizeInBits = 0;
103};
104
105static void logLoopBailout(const Loop &L, const Twine &Reason) {
106 LLVM_DEBUG({
107 dbgs() << "PN loop rewrite: skipping loop with header ";
108 L.getHeader()->printAsOperand(dbgs(), /*PrintType=*/false);
109 dbgs() << ": " << Reason << "\n";
110 });
111}
112
113static void logMatchFailure(const PHINode &Phi, const Twine &Reason) {
114 LLVM_DEBUG({
115 dbgs() << "PN loop rewrite: failed to match mask phi ";
116 Phi.printAsOperand(dbgs(), /*PrintType=*/false);
117 dbgs() << ": " << Reason << "\n";
118 });
119}
120
121/// Returns the scalar size in bits for a type according to the data layout.
122static unsigned getScalarSizeInBits(const DataLayout &DL, Type *Ty) {
123 return DL.getTypeSizeInBits(Ty: Ty->getScalarType()).getFixedValue();
124}
125
126/// Returns the SVE element count for \p ElementSizeInBits.
127static ElementCount getSVEElementCount(unsigned ElementSizeInBits) {
128 return ElementCount::getScalable(MinVal: AArch64::SVEBitsPerBlock /
129 ElementSizeInBits);
130}
131
132/// Returns the largest scalar access size of masked loads/stores in loop \p L
133/// where \p MaskPhi is used as the mask. TODO: This heuristic may need
134/// refinement based on the frequency of different access sizes.
135static unsigned getLargestMaskedMemAccessSizeInBits(const Loop &L,
136 PHINode &MaskPhi) {
137 const DataLayout &DL = MaskPhi.getModule()->getDataLayout();
138 unsigned LargestAccessSizeInBits = 0;
139
140 for (User *U : MaskPhi.users()) {
141 auto *II = dyn_cast<IntrinsicInst>(Val: U);
142 if (!II || !L.contains(Inst: II))
143 continue;
144
145 Intrinsic::ID IID = II->getIntrinsicID();
146 if (IID != Intrinsic::masked_load && IID != Intrinsic::masked_store)
147 continue;
148
149 unsigned MaskOpIdx = IID == Intrinsic::masked_load ? 1 : 2;
150 if (II->getArgOperand(i: MaskOpIdx) != &MaskPhi)
151 continue;
152
153 unsigned AccessSizeInBits = getScalarSizeInBits(DL, Ty: II->getAccessType());
154 if (AccessSizeInBits > LargestAccessSizeInBits)
155 LargestAccessSizeInBits = AccessSizeInBits;
156 }
157
158 return LargestAccessSizeInBits;
159}
160
161/// Returns the predicate-as-counter whilelo intrinsic ID for
162/// \p ElementSizeInBits.
163static Intrinsic::ID getWhileLOIntrinsic(unsigned ElementSizeInBits) {
164 switch (ElementSizeInBits) {
165 case 8:
166 return Intrinsic::aarch64_sve_whilelo_c8;
167 case 16:
168 return Intrinsic::aarch64_sve_whilelo_c16;
169 case 32:
170 return Intrinsic::aarch64_sve_whilelo_c32;
171 case 64:
172 return Intrinsic::aarch64_sve_whilelo_c64;
173 default:
174 llvm_unreachable("unsupported predicate-as-counter element size");
175 }
176}
177
178/// Expands the predicate-as-counter mask \p Count into a wide vector mask.
179/// Masks are extracted from the counter using the paired pext intrinsics,
180/// then concatenated to form the wide mask value.
181static Value *buildWideMask(IRBuilder<> &Builder, const MaskRewriteCandidate &C,
182 Value *Count) {
183 ElementCount LegalEC = getSVEElementCount(ElementSizeInBits: C.ElementSizeInBits);
184 Type *LegalMaskTy = VectorType::get(ElementType: Builder.getInt1Ty(), EC: LegalEC);
185
186 Value *WideMask = PoisonValue::get(T: C.MaskPhi->getType());
187 for (unsigned PairOffset = 0; PairOffset != C.VectorScale / 2; ++PairOffset) {
188 auto *Pair =
189 Builder.CreateIntrinsic(ID: Intrinsic::aarch64_sve_pext_x2, OverloadTypes: {LegalMaskTy},
190 Args: {Count, Builder.getInt32(C: PairOffset)},
191 /*FMFSource=*/{}, Name: "pn.pext.pair");
192 for (unsigned SliceInPair = 0; SliceInPair != 2; ++SliceInPair) {
193 Value *Part = Builder.CreateExtractValue(Agg: Pair, Idxs: SliceInPair, Name: "pn.pext");
194 unsigned Slice = PairOffset * 2 + SliceInPair;
195 WideMask = Builder.CreateInsertVector(
196 DstType: C.MaskPhi->getType(), SrcVec: WideMask, SubVec: Part,
197 Idx: Slice * LegalEC.getKnownMinValue(), Name: "pn.mask");
198 }
199 }
200
201 return WideMask;
202}
203
204/// Creates a predicate-as-counter whilelo for \p ElementSizeInBits between
205/// \p Start and \p End with multi-vector \p VectorScale (2 or 4).
206static Value *createWhileLO(IRBuilder<> &Builder, unsigned ElementSizeInBits,
207 Value *Start, Value *End, unsigned VectorScale) {
208 if (Start->getType()->getIntegerBitWidth() < 64) {
209 Start = Builder.CreateZExt(V: Start, DestTy: Builder.getInt64Ty());
210 End = Builder.CreateZExt(V: End, DestTy: Builder.getInt64Ty());
211 }
212
213 Intrinsic::ID WhileLO = getWhileLOIntrinsic(ElementSizeInBits);
214 return Builder.CreateIntrinsic(ID: WhileLO,
215 Args: {Start, End, Builder.getInt32(C: VectorScale)},
216 /*FMFSource=*/{}, Name: "pn.mask");
217}
218
219/// If \p UserI is an extractelement use of the original loop mask, attempt to
220/// rewrite it to `extractelement(pext(count, 0), idx)` if the extract index is
221/// known to be within the first mask section (which means pext index = 0). If
222/// the user of the extract is a branch and the source of the mask is a whilelo,
223/// this form can be optimized into checking the status flags.
224static bool tryRewriteExtractElement(Instruction &UserI,
225 const MaskRewriteCandidate &C,
226 Value *Count) {
227 auto *EEI = dyn_cast<ExtractElementInst>(Val: &UserI);
228 if (!EEI)
229 return false;
230
231 ElementCount LegalEC = getSVEElementCount(ElementSizeInBits: C.ElementSizeInBits);
232 auto *Idx = dyn_cast<ConstantInt>(Val: EEI->getIndexOperand());
233 if (!Idx || Idx->getValue().uge(RHS: LegalEC.getKnownMinValue()))
234 return false;
235
236 IRBuilder<> Builder(EEI);
237 Builder.SetCurrentDebugLocation(EEI->getDebugLoc());
238
239 auto *ExtractMask = Builder.CreateIntrinsic(
240 ID: Intrinsic::aarch64_sve_pext,
241 OverloadTypes: {VectorType::get(ElementType: Builder.getInt1Ty(), EC: LegalEC)},
242 Args: {Count, Builder.getInt32(C: 0)}, /*FMFSource=*/{}, Name: "pn.pext");
243
244 Value *Extracted = Builder.CreateExtractElement(
245 Vec: ExtractMask, Idx: EEI->getIndexOperand(), Name: EEI->getName() + ".pn");
246
247 EEI->replaceAllUsesWith(V: Extracted);
248 EEI->eraseFromParent();
249 return true;
250}
251
252class AArch64PredicateAsCounterLoopRewrites : public LoopPass {
253public:
254 static char ID;
255
256 AArch64PredicateAsCounterLoopRewrites() : LoopPass(ID) {}
257
258 void getAnalysisUsage(AnalysisUsage &AU) const override {
259 AU.addRequired<TargetPassConfig>();
260 // Require loop simplify to ensure loops have a preheader.
261 AU.addRequiredID(ID&: LoopSimplifyID);
262 AU.addPreservedID(ID&: LoopSimplifyID);
263 AU.setPreservesCFG();
264 }
265
266 bool runOnLoop(Loop *L, LPPassManager &) override;
267
268private:
269 std::optional<MaskRewriteCandidate> matchMaskPhi(Loop &L, PHINode &Phi) const;
270 bool rewriteCandidate(const MaskRewriteCandidate &C, Loop &L) const;
271};
272
273} // end anonymous namespace
274
275char AArch64PredicateAsCounterLoopRewrites::ID = 0;
276
277INITIALIZE_PASS_BEGIN(AArch64PredicateAsCounterLoopRewrites, DEBUG_TYPE,
278 "AArch64 Predicate As Counter Loop Rewrites", false,
279 false)
280INITIALIZE_PASS_DEPENDENCY(TargetPassConfig)
281INITIALIZE_PASS_DEPENDENCY(LoopSimplify)
282INITIALIZE_PASS_END(AArch64PredicateAsCounterLoopRewrites, DEBUG_TYPE,
283 "AArch64 Predicate As Counter Loop Rewrites", false, false)
284
285Pass *llvm::createAArch64PredicateAsCounterLoopRewritesPass() {
286 return new AArch64PredicateAsCounterLoopRewrites();
287}
288
289bool AArch64PredicateAsCounterLoopRewrites::runOnLoop(Loop *L,
290 LPPassManager &) {
291 if (skipLoop(L)) {
292 logLoopBailout(L: *L, Reason: "skipLoop requested the loop to be skipped");
293 return false;
294 }
295
296 Function &F = *L->getHeader()->getParent();
297 auto &TPC = getAnalysis<TargetPassConfig>();
298 const AArch64Subtarget *ST =
299 TPC.getTM<AArch64TargetMachine>().getSubtargetImpl(F);
300 if (!ST->hasSVE2p1() && !(ST->hasSME2() && ST->isStreaming())) {
301 logLoopBailout(L: *L, Reason: "neither SVE2.1 nor SME2 is available");
302 return false;
303 }
304
305 if (!L->getLoopPreheader()) {
306 logLoopBailout(L: *L, Reason: "loop has no preheader");
307 return false;
308 }
309 if (!L->getLoopLatch()) {
310 logLoopBailout(L: *L, Reason: "loop has no latch");
311 return false;
312 }
313
314 bool Changed = false;
315 BasicBlock *Header = L->getHeader();
316 for (PHINode &Phi : make_early_inc_range(Range: Header->phis())) {
317 if (std::optional<MaskRewriteCandidate> Candidate = matchMaskPhi(L&: *L, Phi))
318 Changed |= rewriteCandidate(C: *Candidate, L&: *L);
319 }
320
321 if (Changed)
322 ++LoopsRewritten;
323
324 return Changed;
325}
326
327static IntrinsicInst *getGetActiveLaneMask(Value *V) {
328 auto *II = dyn_cast<IntrinsicInst>(Val: V);
329 return II && II->getIntrinsicID() == Intrinsic::get_active_lane_mask
330 ? II
331 : nullptr;
332}
333
334std::optional<MaskRewriteCandidate>
335AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(Loop &L,
336 PHINode &Phi) const {
337 auto *PhiTy = dyn_cast<ScalableVectorType>(Val: Phi.getType());
338 if (!PhiTy || !PhiTy->getElementType()->isIntegerTy(BitWidth: 1))
339 return std::nullopt;
340
341 if (Phi.getNumIncomingValues() != 2) {
342 logMatchFailure(Phi, Reason: Twine("phi has ") + Twine(Phi.getNumIncomingValues()) +
343 " incoming values; expected 2");
344 return std::nullopt;
345 }
346
347 Value *StartValue = Phi.getIncomingValueForBlock(BB: L.getLoopPreheader());
348 Value *NextValue = Phi.getIncomingValueForBlock(BB: L.getLoopLatch());
349 IntrinsicInst *StartMask = getGetActiveLaneMask(V: StartValue);
350 IntrinsicInst *NextMask = getGetActiveLaneMask(V: NextValue);
351 if (!StartMask) {
352 logMatchFailure(Phi,
353 Reason: "preheader incoming value is not get_active_lane_mask");
354 return std::nullopt;
355 }
356 if (!NextMask) {
357 logMatchFailure(Phi, Reason: "latch incoming value is not get_active_lane_mask");
358 return std::nullopt;
359 }
360
361 unsigned WideMaskElements = PhiTy->getMinNumElements();
362 if (!isPowerOf2_32(Value: WideMaskElements)) {
363 logMatchFailure(Phi,
364 Reason: Twine("wide mask element count is not a power of 2: ") +
365 Twine(WideMaskElements));
366 return std::nullopt;
367 }
368
369 if (StartMask->getArgOperand(i: 0)->getType()->getIntegerBitWidth() > 64) {
370 logMatchFailure(Phi, Reason: "start mask induction operand is wider than i64");
371 return std::nullopt;
372 }
373 if (NextMask->getArgOperand(i: 0)->getType()->getIntegerBitWidth() > 64) {
374 logMatchFailure(Phi, Reason: "next mask induction operand is wider than i64");
375 return std::nullopt;
376 }
377
378 unsigned PreferredMaskElementSizeInBits =
379 getLargestMaskedMemAccessSizeInBits(L, MaskPhi&: Phi);
380
381 if (!is_contained(Set: {8u, 16u, 32u, 64u}, Element: PreferredMaskElementSizeInBits)) {
382 logMatchFailure(Phi, Reason: Twine("unsupported element size in bits: ") +
383 Twine(PreferredMaskElementSizeInBits));
384 return std::nullopt;
385 }
386
387 unsigned SVEMaskElements =
388 AArch64::SVEBitsPerBlock / PreferredMaskElementSizeInBits;
389 if (WideMaskElements <= SVEMaskElements) {
390 logMatchFailure(Phi, Reason: Twine("wide mask element count ") +
391 Twine(WideMaskElements) +
392 " is not wider than the legal mask width " +
393 Twine(SVEMaskElements));
394 return std::nullopt;
395 }
396
397 unsigned VectorScale = WideMaskElements / SVEMaskElements;
398 if (VectorScale != 2 && VectorScale != 4) {
399 logMatchFailure(Phi, Reason: Twine("unsupported predicate-as-counter scale: ") +
400 Twine(VectorScale));
401 return std::nullopt;
402 }
403
404 return MaskRewriteCandidate{.MaskPhi: &Phi, .StartMask: StartMask, .NextMask: NextMask, .VectorScale: VectorScale,
405 .ElementSizeInBits: PreferredMaskElementSizeInBits};
406}
407
408bool AArch64PredicateAsCounterLoopRewrites::rewriteCandidate(
409 const MaskRewriteCandidate &C, Loop &L) const {
410 IRBuilder<> Builder(C.StartMask);
411 Value *NewStart =
412 createWhileLO(Builder, ElementSizeInBits: C.ElementSizeInBits, Start: C.StartMask->getArgOperand(i: 0),
413 End: C.StartMask->getArgOperand(i: 1), VectorScale: C.VectorScale);
414 Builder.SetInsertPoint(C.NextMask);
415 Value *NewNext =
416 createWhileLO(Builder, ElementSizeInBits: C.ElementSizeInBits, Start: C.NextMask->getArgOperand(i: 0),
417 End: C.NextMask->getArgOperand(i: 1), VectorScale: C.VectorScale);
418
419 Builder.SetInsertPoint(C.MaskPhi);
420 auto *NewPhi =
421 Builder.CreatePHI(Ty: NewStart->getType(), NumReservedValues: 2, Name: C.MaskPhi->getName() + ".pn");
422 NewPhi->addIncoming(V: NewStart, BB: L.getLoopPreheader());
423 NewPhi->addIncoming(V: NewNext, BB: L.getLoopLatch());
424
425 auto RewriteUses = [&](Instruction *OldMask, Value *Count,
426 function_ref<bool(Use &U)> Predicate = nullptr) {
427 SmallVector<Use *, 8> UsesToRewrite;
428 for (Use &U : OldMask->uses()) {
429 if (!Predicate || Predicate(U))
430 UsesToRewrite.push_back(Elt: &U);
431 }
432
433 Value *WideMask = nullptr;
434 for (Use *U : UsesToRewrite) {
435 auto *UserI = cast<Instruction>(Val: U->getUser());
436 if (tryRewriteExtractElement(UserI&: *UserI, C, Count))
437 continue;
438
439 if (!WideMask) {
440 BasicBlock::iterator InsertPt = OldMask->getIterator();
441 if (isa<PHINode>(Val: OldMask))
442 InsertPt = OldMask->getParent()->getFirstNonPHIIt();
443
444 Builder.SetInsertPoint(InsertPt);
445 Builder.SetCurrentDebugLocation(OldMask->getDebugLoc());
446 WideMask = buildWideMask(Builder, C, Count);
447 }
448
449 U->set(WideMask);
450 }
451 };
452
453 // For the start/next mask, we only replace non-phi in-loop users. Due to CSE,
454 // these masks could be used by other loops, and replacing them with
455 // predicate-as-counter masks could prevent matching those loops.
456 auto IsNonPhiUseInLoop = [&](Use &U) {
457 auto *UseInst = cast<Instruction>(Val: U.getUser());
458 return !isa<PHINode>(Val: UseInst) && L.contains(Inst: UseInst);
459 };
460
461 RewriteUses(C.MaskPhi, NewPhi);
462 RewriteUses(C.StartMask, NewStart, IsNonPhiUseInLoop);
463 RewriteUses(C.NextMask, NewNext, IsNonPhiUseInLoop);
464
465 RecursivelyDeleteTriviallyDeadInstructions(V: C.MaskPhi);
466 return true;
467}
468