76#include "llvm/IR/IntrinsicsAArch64.h"
87#define DEBUG_TYPE "aarch64-pn-loop-rewrites"
90STATISTIC(LoopsRewritten,
"Number of loops rewritten");
92struct MaskRewriteCandidate {
100 unsigned VectorScale = 0;
102 unsigned ElementSizeInBits = 0;
105static void logLoopBailout(
const Loop &L,
const Twine &Reason) {
107 dbgs() <<
"PN loop rewrite: skipping loop with header ";
108 L.getHeader()->printAsOperand(
dbgs(),
false);
109 dbgs() <<
": " << Reason <<
"\n";
113static void logMatchFailure(
const PHINode &Phi,
const Twine &Reason) {
115 dbgs() <<
"PN loop rewrite: failed to match mask phi ";
116 Phi.printAsOperand(
dbgs(),
false);
117 dbgs() <<
": " << Reason <<
"\n";
123 return DL.getTypeSizeInBits(Ty->getScalarType()).getFixedValue();
127static ElementCount getSVEElementCount(
unsigned ElementSizeInBits) {
135static unsigned getLargestMaskedMemAccessSizeInBits(
const Loop &L,
138 unsigned LargestAccessSizeInBits = 0;
142 if (!
II || !L.contains(
II))
146 if (IID != Intrinsic::masked_load && IID != Intrinsic::masked_store)
149 unsigned MaskOpIdx = IID == Intrinsic::masked_load ? 1 : 2;
150 if (
II->getArgOperand(MaskOpIdx) != &MaskPhi)
154 if (AccessSizeInBits > LargestAccessSizeInBits)
155 LargestAccessSizeInBits = AccessSizeInBits;
158 return LargestAccessSizeInBits;
163static Intrinsic::ID getWhileLOIntrinsic(
unsigned ElementSizeInBits) {
164 switch (ElementSizeInBits) {
166 return Intrinsic::aarch64_sve_whilelo_c8;
168 return Intrinsic::aarch64_sve_whilelo_c16;
170 return Intrinsic::aarch64_sve_whilelo_c32;
172 return Intrinsic::aarch64_sve_whilelo_c64;
183 ElementCount LegalEC = getSVEElementCount(
C.ElementSizeInBits);
187 for (
unsigned PairOffset = 0; PairOffset !=
C.VectorScale / 2; ++PairOffset) {
189 Builder.CreateIntrinsic(Intrinsic::aarch64_sve_pext_x2, {LegalMaskTy},
190 {
Count, Builder.getInt32(PairOffset)},
192 for (
unsigned SliceInPair = 0; SliceInPair != 2; ++SliceInPair) {
193 Value *Part = Builder.CreateExtractValue(Pair, SliceInPair,
"pn.pext");
194 unsigned Slice = PairOffset * 2 + SliceInPair;
195 WideMask = Builder.CreateInsertVector(
196 C.MaskPhi->getType(), WideMask, Part,
206static Value *createWhileLO(
IRBuilder<> &Builder,
unsigned ElementSizeInBits,
207 Value *Start,
Value *End,
unsigned VectorScale) {
208 if (Start->getType()->getIntegerBitWidth() < 64) {
209 Start = Builder.CreateZExt(Start, Builder.getInt64Ty());
210 End = Builder.CreateZExt(End, Builder.getInt64Ty());
213 Intrinsic::ID WhileLO = getWhileLOIntrinsic(ElementSizeInBits);
214 return Builder.CreateIntrinsic(WhileLO,
215 {Start, End, Builder.getInt32(VectorScale)},
224static bool tryRewriteExtractElement(
Instruction &UserI,
225 const MaskRewriteCandidate &
C,
231 ElementCount LegalEC = getSVEElementCount(
C.ElementSizeInBits);
237 Builder.SetCurrentDebugLocation(EEI->getDebugLoc());
239 auto *ExtractMask = Builder.CreateIntrinsic(
240 Intrinsic::aarch64_sve_pext,
242 {
Count, Builder.getInt32(0)}, {},
"pn.pext");
244 Value *Extracted = Builder.CreateExtractElement(
245 ExtractMask, EEI->getIndexOperand(), EEI->getName() +
".pn");
248 EEI->eraseFromParent();
252class AArch64PredicateAsCounterLoopRewrites :
public LoopPass {
256 AArch64PredicateAsCounterLoopRewrites() :
LoopPass(ID) {}
269 std::optional<MaskRewriteCandidate> matchMaskPhi(
Loop &L,
PHINode &Phi)
const;
270 bool rewriteCandidate(
const MaskRewriteCandidate &
C,
Loop &L)
const;
275char AArch64PredicateAsCounterLoopRewrites::ID = 0;
278 "AArch64 Predicate As Counter Loop Rewrites",
false,
286 return new AArch64PredicateAsCounterLoopRewrites();
289bool AArch64PredicateAsCounterLoopRewrites::runOnLoop(
Loop *L,
292 logLoopBailout(*L,
"skipLoop requested the loop to be skipped");
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,
"neither SVE2.1 nor SME2 is available");
305 if (!
L->getLoopPreheader()) {
306 logLoopBailout(*L,
"loop has no preheader");
309 if (!
L->getLoopLatch()) {
310 logLoopBailout(*L,
"loop has no latch");
317 if (std::optional<MaskRewriteCandidate> Candidate = matchMaskPhi(*L, Phi))
318 Changed |= rewriteCandidate(*Candidate, *L);
329 return II &&
II->getIntrinsicID() == Intrinsic::get_active_lane_mask
334std::optional<MaskRewriteCandidate>
335AArch64PredicateAsCounterLoopRewrites::matchMaskPhi(
Loop &L,
336 PHINode &Phi)
const {
338 if (!PhiTy || !PhiTy->getElementType()->isIntegerTy(1))
341 if (
Phi.getNumIncomingValues() != 2) {
342 logMatchFailure(Phi, Twine(
"phi has ") + Twine(
Phi.getNumIncomingValues()) +
343 " incoming values; expected 2");
347 Value *StartValue =
Phi.getIncomingValueForBlock(
L.getLoopPreheader());
348 Value *NextValue =
Phi.getIncomingValueForBlock(
L.getLoopLatch());
353 "preheader incoming value is not get_active_lane_mask");
357 logMatchFailure(Phi,
"latch incoming value is not get_active_lane_mask");
361 unsigned WideMaskElements = PhiTy->getMinNumElements();
364 Twine(
"wide mask element count is not a power of 2: ") +
365 Twine(WideMaskElements));
370 logMatchFailure(Phi,
"start mask induction operand is wider than i64");
374 logMatchFailure(Phi,
"next mask induction operand is wider than i64");
378 unsigned PreferredMaskElementSizeInBits =
379 getLargestMaskedMemAccessSizeInBits(L, Phi);
381 if (!
is_contained({8u, 16u, 32u, 64u}, PreferredMaskElementSizeInBits)) {
382 logMatchFailure(Phi, Twine(
"unsupported element size in bits: ") +
383 Twine(PreferredMaskElementSizeInBits));
387 unsigned SVEMaskElements =
389 if (WideMaskElements <= SVEMaskElements) {
390 logMatchFailure(Phi, Twine(
"wide mask element count ") +
391 Twine(WideMaskElements) +
392 " is not wider than the legal mask width " +
393 Twine(SVEMaskElements));
397 unsigned VectorScale = WideMaskElements / SVEMaskElements;
398 if (VectorScale != 2 && VectorScale != 4) {
399 logMatchFailure(Phi, Twine(
"unsupported predicate-as-counter scale: ") +
404 return MaskRewriteCandidate{&
Phi, StartMask, NextMask, VectorScale,
405 PreferredMaskElementSizeInBits};
408bool AArch64PredicateAsCounterLoopRewrites::rewriteCandidate(
409 const MaskRewriteCandidate &
C,
Loop &L)
const {
412 createWhileLO(Builder,
C.ElementSizeInBits,
C.StartMask->getArgOperand(0),
413 C.StartMask->getArgOperand(1),
C.VectorScale);
414 Builder.SetInsertPoint(
C.NextMask);
416 createWhileLO(Builder,
C.ElementSizeInBits,
C.NextMask->getArgOperand(0),
417 C.NextMask->getArgOperand(1),
C.VectorScale);
419 Builder.SetInsertPoint(
C.MaskPhi);
421 Builder.CreatePHI(NewStart->
getType(), 2,
C.MaskPhi->getName() +
".pn");
422 NewPhi->addIncoming(NewStart,
L.getLoopPreheader());
423 NewPhi->addIncoming(NewNext,
L.getLoopLatch());
426 function_ref<bool(Use &U)>
Predicate =
nullptr) {
428 for (Use &U : OldMask->
uses()) {
433 Value *WideMask =
nullptr;
434 for (Use *U : UsesToRewrite) {
436 if (tryRewriteExtractElement(*UserI,
C,
Count))
442 InsertPt = OldMask->
getParent()->getFirstNonPHIIt();
444 Builder.SetInsertPoint(InsertPt);
445 Builder.SetCurrentDebugLocation(OldMask->
getDebugLoc());
446 WideMask = buildWideMask(Builder,
C,
Count);
456 auto IsNonPhiUseInLoop = [&](
Use &
U) {
461 RewriteUses(
C.MaskPhi, NewPhi);
462 RewriteUses(
C.StartMask, NewStart, IsNonPhiUseInLoop);
463 RewriteUses(
C.NextMask, NewNext, IsNonPhiUseInLoop);
static IntrinsicInst * getGetActiveLaneMask(Value *V)
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
This file contains the simple types necessary to represent the attributes associated with functions a...
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
This file defines the DenseMap class.
Module.h This file contains the declarations for the Module class.
uint64_t IntrinsicInst * II
#define INITIALIZE_PASS_DEPENDENCY(depName)
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
This file defines the SmallVector class.
This file defines the 'Statistic' class, which is designed to be an easy way to expose various metric...
#define STATISTIC(VARNAME, DESC)
Target-Independent Code Generator Pass Configuration Options pass.
Represent the analysis usage information of a pass.
LLVM_ABI AnalysisUsage & addRequiredID(const void *ID)
AnalysisUsage & addPreservedID(const void *ID)
AnalysisUsage & addRequired()
LLVM_ABI void setPreservesCFG()
This function should be called by the pass, iff they do not:
InstListType::iterator iterator
Instruction iterators...
Value * getArgOperand(unsigned i) const
A parsed version of the target data layout string in and methods for querying it.
static constexpr ElementCount getScalable(ScalarTy MinVal)
This provides a uniform API for creating instructions and inserting them into a basic block: either a...
const DebugLoc & getDebugLoc() const
Return the debug location for this node as a DebugLoc.
LLVM_ABI const Module * getModule() const
Return the module owning the function this instruction belongs to or nullptr it the function does not...
iterator_range< user_iterator > users()
A wrapper class for inspecting calls to intrinsic functions.
Represents a single loop in the control flow graph.
const DataLayout & getDataLayout() const
Get the data layout for the module's target platform.
Pass interface - Implemented by all 'passes'.
static LLVM_ABI PoisonValue * get(Type *T)
Static factory methods - Return an 'poison' object of the specified type.
void push_back(const T &Elt)
Target-Independent Code Generator Pass Configuration Options.
Twine - A lightweight data structure for efficiently representing the concatenation of temporary valu...
The instances of the Type class are immutable: once they are created, they are never changed.
LLVM_ABI unsigned getIntegerBitWidth() const
LLVM Value Representation.
Type * getType() const
All values are typed, get the type of this value.
LLVM_ABI void replaceAllUsesWith(Value *V)
Change all uses of this to point to a new Value.
iterator_range< use_iterator > uses()
static LLVM_ABI VectorType * get(Type *ElementType, ElementCount EC)
This static method is the primary way to construct an VectorType.
constexpr ScalarTy getKnownMinValue() const
Returns the minimum value this quantity can represent.
const ParentTy * getParent() const
self_iterator getIterator()
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
static constexpr unsigned SVEBitsPerBlock
@ BasicBlock
Various leaf nodes.
Predicate
Predicate - These are "(BI << 5) | BO" for various predicates.
NodeAddr< PhiNode * > Phi
NodeAddr< UseNode * > Use
friend class Instruction
Iterator for Instructions in a `BasicBlock.
This is an optimization pass for GlobalISel generic memory operations.
LLVM_ABI bool RecursivelyDeleteTriviallyDeadInstructions(Value *V, const TargetLibraryInfo *TLI=nullptr, MemorySSAUpdater *MSSAU=nullptr, std::function< void(Value *)> AboutToDeleteCallback=std::function< void(Value *)>())
If the specified value is a trivially dead instruction, delete it.
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
iterator_range< early_inc_iterator_impl< detail::IterOfRange< RangeT > > > make_early_inc_range(RangeT &&Range)
Make a range that does early increment to allow mutation of the underlying range without disrupting i...
Pass * createAArch64PredicateAsCounterLoopRewritesPass()
LLVM_ABI char & LoopSimplifyID
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Value
constexpr bool isPowerOf2_32(uint32_t Value)
Return true if the argument is a power of two > 0.
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
IRBuilder(LLVMContext &, FolderTy, InserterTy) -> IRBuilder< FolderTy, InserterTy >
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Count
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
bool is_contained(R &&Range, const E &Element)
Returns true if Element is found in Range.