1//===- SLPMemoryUtils.cpp - SLP pointer/stride helpers --------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#include "SLPMemoryUtils.h"
10#include "SLPCompatibilityAnalysis.h"
11#include "SLPCostAnalysis.h"
12#include "SLPTypeUtils.h"
13#include "SLPUtils.h"
14
15#include "llvm/ADT/APInt.h"
16#include "llvm/ADT/MapVector.h"
17#include "llvm/ADT/STLExtras.h"
18#include "llvm/ADT/Sequence.h"
19#include "llvm/ADT/SmallPtrSet.h"
20#include "llvm/Analysis/Loads.h"
21#include "llvm/Analysis/LoopAccessAnalysis.h"
22#include "llvm/Analysis/ScalarEvolution.h"
23#include "llvm/Analysis/ScalarEvolutionExpressions.h"
24#include "llvm/Analysis/TargetTransformInfo.h"
25#include "llvm/Analysis/ValueTracking.h"
26#include "llvm/IR/Constants.h"
27#include "llvm/IR/DataLayout.h"
28#include "llvm/IR/DerivedTypes.h"
29#include "llvm/IR/IRBuilder.h"
30#include "llvm/IR/Instructions.h"
31#include "llvm/IR/Intrinsics.h"
32#include "llvm/Support/InstructionCost.h"
33
34#include <algorithm>
35#include <limits>
36#include <optional>
37#include <set>
38#include <tuple>
39#include <utility>
40
41using namespace llvm;
42
43namespace llvm::slpvectorizer {
44
45Value *createWidenedStridedCast(IRBuilderBase &Builder, Value *V, Type *DstTy,
46 const DataLayout &DL) {
47 bool ToPtr = cast<VectorType>(Val: DstTy)->getElementType()->isPointerTy();
48 if (ToPtr == cast<VectorType>(Val: V->getType())->getElementType()->isPointerTy())
49 return Builder.CreateBitOrPointerCast(V, DestTy: DstTy);
50 if (ToPtr)
51 return Builder.CreateIntToPtr(
52 V: Builder.CreateBitCast(V, DestTy: DL.getIntPtrType(DstTy)), DestTy: DstTy);
53 return Builder.CreateBitCast(
54 V: Builder.CreatePtrToInt(V, DestTy: DL.getIntPtrType(V->getType())), DestTy: DstTy);
55}
56
57ConstantInt *getStrideBytesIfConstant(Value *Stride, Type *ScalarTy,
58 const DataLayout &DL, bool IsReverse) {
59 auto *CI = dyn_cast_or_null<ConstantInt>(Val: Stride);
60 if (!CI)
61 return nullptr;
62
63 uint64_t ElementSize = DL.getTypeAllocSize(Ty: ScalarTy).getFixedValue();
64 APInt Bytes = CI->getValue() * ElementSize;
65 return ConstantInt::get(Context&: CI->getContext(), V: IsReverse ? -Bytes : Bytes);
66}
67
68bool arePointersCompatible(Value *Ptr1, Value *Ptr2,
69 const TargetLibraryInfo &TLI, unsigned MaxDepth,
70 bool CompareOpcodes) {
71 if (getUnderlyingObject(V: Ptr1, MaxLookup: MaxDepth) !=
72 getUnderlyingObject(V: Ptr2, MaxLookup: MaxDepth))
73 return false;
74 auto *GEP1 = dyn_cast<GetElementPtrInst>(Val: Ptr1);
75 auto *GEP2 = dyn_cast<GetElementPtrInst>(Val: Ptr2);
76 return (!GEP1 || GEP1->getNumOperands() == 2) &&
77 (!GEP2 || GEP2->getNumOperands() == 2) &&
78 (((!GEP1 || isConstant(V: GEP1->getOperand(i_nocapture: 1))) &&
79 (!GEP2 || isConstant(V: GEP2->getOperand(i_nocapture: 1)))) ||
80 !CompareOpcodes ||
81 (GEP1 && GEP2 &&
82 getSameOpcode(VL: {GEP1->getOperand(i_nocapture: 1), GEP2->getOperand(i_nocapture: 1)}, TLI)));
83}
84
85/// Calculates minimal alignment as a common alignment.
86template <typename T> Align computeCommonAlignment(ArrayRef<Value *> VL) {
87 Align CommonAlignment = cast<T>(VL.consume_front())->getAlign();
88 for (Value *V : VL)
89 CommonAlignment = std::min(CommonAlignment, cast<T>(V)->getAlign());
90 return CommonAlignment;
91}
92
93template Align computeCommonAlignment<LoadInst>(ArrayRef<Value *>);
94template Align computeCommonAlignment<StoreInst>(ArrayRef<Value *>);
95
96const SCEV *calculateRtStride(ArrayRef<Value *> PointerOps, Type *ElemTy,
97 const DataLayout &DL, ScalarEvolution &SE,
98 SmallVectorImpl<unsigned> &SortedIndices) {
99 SmallVector<const SCEV *> SCEVs;
100 const SCEV *PtrSCEVLowest = nullptr;
101 const SCEV *PtrSCEVHighest = nullptr;
102 // Find lower/upper pointers from the PointerOps (i.e. with lowest and highest
103 // addresses).
104 for (Value *Ptr : PointerOps) {
105 const SCEV *PtrSCEV = SE.getSCEV(V: Ptr);
106 if (!PtrSCEV)
107 return nullptr;
108 SCEVs.push_back(Elt: PtrSCEV);
109 if (!PtrSCEVLowest && !PtrSCEVHighest) {
110 PtrSCEVLowest = PtrSCEVHighest = PtrSCEV;
111 continue;
112 }
113 const SCEV *Diff = SE.getMinusSCEV(LHS: PtrSCEV, RHS: PtrSCEVLowest);
114 if (isa<SCEVCouldNotCompute>(Val: Diff))
115 return nullptr;
116 if (Diff->isNonConstantNegative()) {
117 PtrSCEVLowest = PtrSCEV;
118 continue;
119 }
120 const SCEV *Diff1 = SE.getMinusSCEV(LHS: PtrSCEVHighest, RHS: PtrSCEV);
121 if (isa<SCEVCouldNotCompute>(Val: Diff1))
122 return nullptr;
123 if (Diff1->isNonConstantNegative()) {
124 PtrSCEVHighest = PtrSCEV;
125 continue;
126 }
127 }
128 // Dist = PtrSCEVHighest - PtrSCEVLowest;
129 const SCEV *Dist = SE.getMinusSCEV(LHS: PtrSCEVHighest, RHS: PtrSCEVLowest);
130 if (isa<SCEVCouldNotCompute>(Val: Dist))
131 return nullptr;
132 int Size = DL.getTypeStoreSize(Ty: ElemTy);
133 auto TryGetStride = [&](const SCEV *Dist,
134 const SCEV *Multiplier) -> const SCEV * {
135 if (const auto *M = dyn_cast<SCEVMulExpr>(Val: Dist)) {
136 if (M->getOperand(i: 0) == Multiplier)
137 return M->getOperand(i: 1);
138 if (M->getOperand(i: 1) == Multiplier)
139 return M->getOperand(i: 0);
140 return nullptr;
141 }
142 if (Multiplier == Dist)
143 return SE.getConstant(Ty: Dist->getType(), V: 1);
144 return SE.getUDivExactExpr(LHS: Dist, RHS: Multiplier);
145 };
146 // Stride_in_elements = Dist / element_size * (num_elems - 1).
147 const SCEV *Stride = nullptr;
148 if (Size != 1 || SCEVs.size() > 1) {
149 const SCEV *Sz = SE.getConstant(Ty: Dist->getType(), V: Size * (SCEVs.size() - 1));
150 Stride = TryGetStride(Dist, Sz);
151 if (!Stride)
152 return nullptr;
153 }
154 if (!Stride || isa<SCEVConstant>(Val: Stride))
155 return nullptr;
156 // Iterate through all pointers and check if all distances are
157 // unique multiple of Stride.
158 using DistOrdPair = std::pair<int64_t, int>;
159 auto Compare = llvm::less_first();
160 std::set<DistOrdPair, decltype(Compare)> Offsets(Compare);
161 bool IsConsecutive = true;
162 for (const auto [Idx, PtrSCEV] : enumerate(First&: SCEVs)) {
163 unsigned Dist = 0;
164 if (PtrSCEV != PtrSCEVLowest) {
165 const SCEV *Diff = SE.getMinusSCEV(LHS: PtrSCEV, RHS: PtrSCEVLowest);
166 const SCEV *Coeff = TryGetStride(Diff, Stride);
167 if (!Coeff)
168 return nullptr;
169 const auto *SC = dyn_cast<SCEVConstant>(Val: Coeff);
170 if (!SC || isa<SCEVCouldNotCompute>(Val: SC))
171 return nullptr;
172 if (!SE.getMinusSCEV(LHS: PtrSCEV, RHS: SE.getAddExpr(LHS: PtrSCEVLowest,
173 RHS: SE.getMulExpr(LHS: Stride, RHS: SC)))
174 ->isZero())
175 return nullptr;
176 Dist = SC->getAPInt().getZExtValue();
177 }
178 // If the strides are not the same or repeated, we can't vectorize.
179 if ((Dist / Size) * Size != Dist || (Dist / Size) >= SCEVs.size())
180 return nullptr;
181 auto Res = Offsets.emplace(args&: Dist, args&: Idx);
182 if (!Res.second)
183 return nullptr;
184 // Consecutive order if the inserted element is the last one.
185 IsConsecutive = IsConsecutive && std::next(x: Res.first) == Offsets.end();
186 }
187 SortedIndices.clear();
188 if (!IsConsecutive) {
189 // Fill SortedIndices array only if it is non-consecutive.
190 SortedIndices.resize(N: PointerOps.size());
191 for (const auto [Idx, Pair] : enumerate(First&: Offsets))
192 SortedIndices[Idx] = Pair.second;
193 }
194 return Stride;
195}
196
197/// Builds compress-like mask for shuffles for the given \p PointerOps, ordered
198/// with \p Order.
199/// \return true if the mask represents strided access, false - otherwise.
200static bool buildCompressMask(ArrayRef<Value *> PointerOps,
201 ArrayRef<unsigned> Order, Type *ScalarTy,
202 const DataLayout &DL, ScalarEvolution &SE,
203 SmallVectorImpl<int> &CompressMask) {
204 const unsigned Sz = PointerOps.size();
205 CompressMask.assign(NumElts: Sz, Elt: PoisonMaskElem);
206 // The first element always set.
207 CompressMask[0] = 0;
208 // Check if the mask represents strided access.
209 std::optional<unsigned> Stride = 0;
210 Value *Ptr0 = Order.empty() ? PointerOps.front() : PointerOps[Order.front()];
211 for (unsigned I : seq<unsigned>(Begin: 1, End: Sz)) {
212 Value *Ptr = Order.empty() ? PointerOps[I] : PointerOps[Order[I]];
213 std::optional<int64_t> OptPos =
214 getPointersDiff(ElemTyA: ScalarTy, PtrA: Ptr0, ElemTyB: ScalarTy, PtrB: Ptr, DL, SE);
215 if (!OptPos || OptPos > std::numeric_limits<unsigned>::max())
216 return false;
217 unsigned Pos = static_cast<unsigned>(*OptPos);
218 CompressMask[I] = Pos;
219 if (!Stride)
220 continue;
221 if (*Stride == 0) {
222 *Stride = Pos;
223 continue;
224 }
225 if (Pos != *Stride * I)
226 Stride.reset();
227 }
228 return Stride.has_value();
229}
230
231/// Checks if the \p VL can be transformed to a (masked)load + compress or
232/// (masked) interleaved load.
233bool isMaskedLoadCompress(
234 ArrayRef<Value *> VL, ArrayRef<Value *> PointerOps,
235 ArrayRef<unsigned> Order, const TargetTransformInfo &TTI,
236 const DataLayout &DL, ScalarEvolution &SE, AssumptionCache &AC,
237 const DominatorTree &DT, const TargetLibraryInfo &TLI,
238 const TargetTransformInfo::TargetCostKind CostKind,
239 const function_ref<bool(Value *)> AreAllUsersVectorized, bool ReVec,
240 bool &IsMasked, unsigned &InterleaveFactor,
241 SmallVectorImpl<int> &CompressMask, VectorType *&LoadVecTy) {
242 InterleaveFactor = 0;
243 Type *ScalarTy = VL.front()->getType();
244 const size_t Sz = VL.size();
245 auto *VecTy = cast<VectorType>(Val: getWidenedType(ScalarTy, VF: Sz));
246 SmallVector<int> Mask;
247 if (!Order.empty())
248 inversePermutation(Indices: Order, Mask);
249 // Check external uses.
250 for (const auto [I, V] : enumerate(First&: VL)) {
251 if (AreAllUsersVectorized(V))
252 continue;
253 InstructionCost ExtractCost =
254 TTI.getVectorInstrCost(Opcode: Instruction::ExtractElement, Val: VecTy, CostKind,
255 Index: Mask.empty() ? I : Mask[I]);
256 InstructionCost ScalarCost =
257 TTI.getInstructionCost(U: cast<Instruction>(Val: V), CostKind);
258 if (ExtractCost <= ScalarCost)
259 return false;
260 }
261 Value *Ptr0;
262 Value *PtrN;
263 if (Order.empty()) {
264 Ptr0 = PointerOps.front();
265 PtrN = PointerOps.back();
266 } else {
267 Ptr0 = PointerOps[Order.front()];
268 PtrN = PointerOps[Order.back()];
269 }
270 std::optional<int64_t> Diff =
271 getPointersDiff(ElemTyA: ScalarTy, PtrA: Ptr0, ElemTyB: ScalarTy, PtrB: PtrN, DL, SE);
272 if (!Diff)
273 return false;
274 const size_t MaxRegSize =
275 TTI.getRegisterBitWidth(K: TargetTransformInfo::RGK_FixedWidthVector)
276 .getFixedValue();
277 // Check for very large distances between elements.
278 if (*Diff / Sz >= MaxRegSize / 8)
279 return false;
280 LoadVecTy = cast<FixedVectorType>(Val: getWidenedType(ScalarTy, VF: *Diff + 1));
281 auto *LI = cast<LoadInst>(Val: Order.empty() ? VL.front() : VL[Order.front()]);
282 Align CommonAlignment = LI->getAlign();
283 SimplifyQuery SQ(
284 DL, &TLI, &DT, &AC,
285 cast<LoadInst>(Val: Order.empty() ? VL.back() : VL[Order.back()]));
286 IsMasked = !isSafeToLoadUnconditionally(V: Ptr0, Ty: LoadVecTy, Alignment: CommonAlignment, SQ);
287 if (IsMasked && !TTI.isLegalMaskedLoad(DataType: LoadVecTy, Alignment: CommonAlignment,
288 AddressSpace: LI->getPointerAddressSpace()))
289 return false;
290 // TODO: perform the analysis of each scalar load for better
291 // safe-load-unconditionally analysis.
292 bool IsStrided =
293 buildCompressMask(PointerOps, Order, ScalarTy, DL, SE, CompressMask);
294 assert(CompressMask.size() >= 2 && "At least two elements are required");
295 SmallVector<Value *> OrderedPointerOps(PointerOps);
296 if (!Order.empty())
297 reorderScalars(Scalars&: OrderedPointerOps, Mask);
298 auto [ScalarGEPCost, VectorGEPCost] =
299 getGEPCosts(TTI, Ptrs: OrderedPointerOps, BasePtr: OrderedPointerOps.front(),
300 Opcode: Instruction::Load, CostKind, ScalarTy, VecTy: LoadVecTy);
301 // The cost of scalar loads.
302 InstructionCost ScalarLoadsCost =
303 accumulate(Range&: VL, Init: InstructionCost(),
304 Op: [&](InstructionCost C, Value *V) {
305 return C + TTI.getInstructionCost(U: cast<Instruction>(Val: V),
306 CostKind);
307 }) +
308 ScalarGEPCost;
309 APInt DemandedElts = APInt::getAllOnes(numBits: Sz);
310 InstructionCost GatherCost =
311 getScalarizationOverhead(TTI, ReVec, ScalarTy, Ty: VecTy, DemandedElts,
312 /*Insert=*/true,
313 /*Extract=*/false, CostKind) +
314 ScalarLoadsCost;
315 InstructionCost LoadCost = 0;
316 if (IsMasked) {
317 LoadCost = TTI.getMemIntrinsicInstrCost(
318 MICA: MemIntrinsicCostAttributes(Intrinsic::masked_load, LoadVecTy,
319 CommonAlignment,
320 LI->getPointerAddressSpace()),
321 CostKind);
322 } else {
323 LoadCost =
324 TTI.getMemoryOpCost(Opcode: Instruction::Load, Src: LoadVecTy, Alignment: CommonAlignment,
325 AddressSpace: LI->getPointerAddressSpace(), CostKind,
326 OpdInfo: TTI::getOperandInfo(V: LI->getPointerOperand()));
327 }
328 if (IsStrided && !IsMasked && Order.empty()) {
329 // Check for potential segmented(interleaved) loads.
330 VectorType *AlignedLoadVecTy = cast<VectorType>(Val: getWidenedType(
331 ScalarTy,
332 VF: getFullVectorNumberOfElements(TTI, Ty: ScalarTy, Sz: *Diff + 1, ReVec)));
333 SimplifyQuery SQ(DL, &TLI, &DT, &AC, cast<LoadInst>(Val: VL.back()));
334 if (!isSafeToLoadUnconditionally(V: Ptr0, Ty: AlignedLoadVecTy, Alignment: CommonAlignment,
335 SQ))
336 AlignedLoadVecTy = LoadVecTy;
337 if (TTI.isLegalInterleavedAccessType(VTy: AlignedLoadVecTy, Factor: CompressMask[1],
338 Alignment: CommonAlignment,
339 AddrSpace: LI->getPointerAddressSpace())) {
340 InstructionCost InterleavedCost =
341 VectorGEPCost + TTI.getInterleavedMemoryOpCost(
342 Opcode: Instruction::Load, VecTy: AlignedLoadVecTy,
343 Factor: CompressMask[1], Indices: {}, Alignment: CommonAlignment,
344 AddressSpace: LI->getPointerAddressSpace(), CostKind, UseMaskForCond: IsMasked);
345 if (InterleavedCost < GatherCost) {
346 InterleaveFactor = CompressMask[1];
347 LoadVecTy = AlignedLoadVecTy;
348 return true;
349 }
350 }
351 }
352 // Estimating the compression shuffle cost below can be extremely expensive
353 // for a very wide LoadVecTy, which is split into a large number of vector
354 // registers (see processShuffleMasks). The shuffle cost is always
355 // non-negative, so if the load cost alone already reaches the gather cost the
356 // masked-load-compress cannot be profitable. Bail out before the costly
357 // shuffle cost estimation in that case.
358 if (VectorGEPCost + LoadCost >= GatherCost)
359 return false;
360 InstructionCost CompressCost = getShuffleCost(
361 TTI, Kind: TTI::SK_PermuteSingleSrc, Tp: LoadVecTy, CostKind, Mask: CompressMask);
362 if (!Order.empty()) {
363 SmallVector<int> NewMask(Sz, PoisonMaskElem);
364 for (unsigned I : seq<unsigned>(Size: Sz)) {
365 NewMask[I] = CompressMask[Mask[I]];
366 }
367 CompressMask.swap(RHS&: NewMask);
368 }
369 InstructionCost TotalVecCost = VectorGEPCost + LoadCost + CompressCost;
370 return TotalVecCost < GatherCost;
371}
372
373/// Checks if the \p VL can be transformed to a (masked)load + compress or
374/// (masked) interleaved load.
375bool isMaskedLoadCompress(
376 ArrayRef<Value *> VL, ArrayRef<Value *> PointerOps,
377 ArrayRef<unsigned> Order, const TargetTransformInfo &TTI,
378 const DataLayout &DL, ScalarEvolution &SE, AssumptionCache &AC,
379 const DominatorTree &DT, const TargetLibraryInfo &TLI,
380 const TargetTransformInfo::TargetCostKind CostKind,
381 const function_ref<bool(Value *)> AreAllUsersVectorized, bool ReVec) {
382 bool IsMasked;
383 unsigned InterleaveFactor;
384 SmallVector<int> CompressMask;
385 VectorType *LoadVecTy;
386 return isMaskedLoadCompress(VL, PointerOps, Order, TTI, DL, SE, AC, DT, TLI,
387 CostKind, AreAllUsersVectorized, ReVec, IsMasked,
388 InterleaveFactor, CompressMask, LoadVecTy);
389}
390
391/// Checks if the stores \p VL with pointers \p PointerOps can be lowered as a
392/// single masked store. On success \p StoreVecTy is the widened store type and
393/// \p ReuseShuffleIndices is the expand mask that places each stored value at
394/// its element offset from the base (poison in the gaps).
395bool isMaskedStoreCompress(ArrayRef<Value *> VL, ArrayRef<Value *> PointerOps,
396 ArrayRef<unsigned> Order,
397 const TargetTransformInfo &TTI, const DataLayout &DL,
398 ScalarEvolution &SE, Align CommonAlignment,
399 SmallVectorImpl<int> &ReuseShuffleIndices,
400 FixedVectorType *&StoreVecTy) {
401 Type *ScalarTy = cast<StoreInst>(Val: VL.front())->getValueOperand()->getType();
402 const size_t Sz = VL.size();
403 // Only simple scalar element types are supported.
404 if (Sz < 2 || (!ScalarTy->isIntOrPtrTy() && !ScalarTy->isFloatingPointTy()))
405 return false;
406 Value *Ptr0 = Order.empty() ? PointerOps.front() : PointerOps[Order.front()];
407 Value *PtrN = Order.empty() ? PointerOps.back() : PointerOps[Order.back()];
408 std::optional<int64_t> Diff =
409 getPointersDiff(ElemTyA: ScalarTy, PtrA: Ptr0, ElemTyB: ScalarTy, PtrB: PtrN, DL, SE);
410 if (!Diff || *Diff <= 0)
411 return false;
412 // Avoid widened vectors with very large gaps between the stored elements.
413 const unsigned MaxRegSize =
414 TTI.getRegisterBitWidth(K: TargetTransformInfo::RGK_FixedWidthVector)
415 .getFixedValue();
416 const unsigned ScalarBits = DL.getTypeSizeInBits(Ty: ScalarTy).getFixedValue();
417 if (ScalarBits == 0 ||
418 static_cast<uint64_t>(*Diff) / Sz >= MaxRegSize / ScalarBits)
419 return false;
420 StoreVecTy = cast<FixedVectorType>(Val: getWidenedType(ScalarTy, VF: *Diff + 1));
421 unsigned AS = cast<StoreInst>(Val: VL.front())->getPointerAddressSpace();
422 if (!TTI.isLegalMaskedStore(DataType: StoreVecTy, Alignment: CommonAlignment, AddressSpace: AS,
423 MaskKind: TTI::ConstantMask))
424 return false;
425 // Build the expand mask: store I (in address-sorted order) is placed at its
426 // element offset from the base, other widened lanes are poison.
427 ReuseShuffleIndices.assign(NumElts: *Diff + 1, Elt: PoisonMaskElem);
428 int64_t Prev = -1;
429 for (unsigned I : seq<unsigned>(Size: Sz)) {
430 Value *Ptr = Order.empty() ? PointerOps[I] : PointerOps[Order[I]];
431 std::optional<int64_t> Off =
432 getPointersDiff(ElemTyA: ScalarTy, PtrA: Ptr0, ElemTyB: ScalarTy, PtrB: Ptr, DL, SE);
433 if (!Off || *Off <= Prev || *Off > *Diff)
434 return false;
435 ReuseShuffleIndices[*Off] = static_cast<int>(I);
436 Prev = *Off;
437 }
438 return true;
439}
440
441bool clusterSortPtrAccesses(ArrayRef<Value *> VL, ArrayRef<BasicBlock *> BBs,
442 Type *ElemTy, const DataLayout &DL,
443 ScalarEvolution &SE, unsigned MaxDepth,
444 SmallVectorImpl<unsigned> &SortedIndices) {
445 assert(
446 all_of(VL, [](const Value *V) { return V->getType()->isPointerTy(); }) &&
447 "Expected list of pointer operands.");
448 // Map from bases to a vector of (Ptr, Offset, OrigIdx), which we insert each
449 // Ptr into, sort and return the sorted indices with values next to one
450 // another.
451 SmallMapVector<
452 std::pair<BasicBlock *, Value *>,
453 SmallVector<SmallVector<std::tuple<Value *, int64_t, unsigned>>>, 8>
454 Bases;
455 Bases
456 .try_emplace(Key: std::make_pair(x: BBs.front(),
457 y: getUnderlyingObject(V: VL.front(), MaxLookup: MaxDepth)))
458 .first->second.emplace_back()
459 .emplace_back(Args: VL.front(), Args: 0U, Args: 0U);
460
461 SortedIndices.clear();
462 for (auto [Cnt, Ptr] : enumerate(First: VL.drop_front())) {
463 auto Key = std::make_pair(x: BBs[Cnt + 1], y: getUnderlyingObject(V: Ptr, MaxLookup: MaxDepth));
464 bool Found = any_of(Range&: Bases.try_emplace(Key).first->second,
465 P: [&, &Cnt = Cnt, &Ptr = Ptr](auto &Base) {
466 std::optional<int64_t> Diff =
467 getPointersDiff(ElemTy, std::get<0>(Base.front()),
468 ElemTy, Ptr, DL, SE,
469 /*StrictCheck=*/true);
470 if (!Diff)
471 return false;
472
473 Base.emplace_back(Ptr, *Diff, Cnt + 1);
474 return true;
475 });
476
477 if (!Found) {
478 // If we haven't found enough to usefully cluster, return early.
479 if (Bases.size() > VL.size() / 2 - 1)
480 return false;
481
482 // Not found already - add a new Base
483 Bases.find(Key)->second.emplace_back().emplace_back(Args: Ptr, Args: 0, Args: Cnt + 1);
484 }
485 }
486
487 if (Bases.size() == VL.size())
488 return false;
489
490 if (Bases.size() == 1 && (Bases.front().second.size() == 1 ||
491 Bases.front().second.size() == VL.size()))
492 return false;
493
494 // For each of the bases sort the pointers by Offset and check if any of the
495 // base become consecutively allocated.
496 auto ComparePointers = [MaxDepth](Value *Ptr1, Value *Ptr2) {
497 SmallPtrSet<Value *, 13> FirstPointers;
498 SmallPtrSet<Value *, 13> SecondPointers;
499 Value *P1 = Ptr1;
500 Value *P2 = Ptr2;
501 unsigned Depth = 0;
502 while (!FirstPointers.contains(Ptr: P2) && !SecondPointers.contains(Ptr: P1)) {
503 if (P1 == P2 || Depth > MaxDepth)
504 return false;
505 FirstPointers.insert(Ptr: P1);
506 SecondPointers.insert(Ptr: P2);
507 P1 = getUnderlyingObject(V: P1, /*MaxLookup=*/1);
508 P2 = getUnderlyingObject(V: P2, /*MaxLookup=*/1);
509 ++Depth;
510 }
511 assert((FirstPointers.contains(P2) || SecondPointers.contains(P1)) &&
512 "Unable to find matching root.");
513 return FirstPointers.contains(Ptr: P2) && !SecondPointers.contains(Ptr: P1);
514 };
515 for (auto &Base : Bases) {
516 for (auto &Vec : Base.second) {
517 if (Vec.size() > 1) {
518 stable_sort(Range&: Vec, C: llvm::less_second());
519 int64_t InitialOffset = std::get<1>(t&: Vec[0]);
520 bool AnyConsecutive =
521 all_of(Range: enumerate(First&: Vec), P: [InitialOffset](const auto &P) {
522 return std::get<1>(P.value()) ==
523 int64_t(P.index()) + InitialOffset;
524 });
525 // Fill SortedIndices array only if it looks worth-while to sort the
526 // ptrs.
527 if (!AnyConsecutive)
528 return false;
529 }
530 }
531 stable_sort(Range&: Base.second, C: [&](const auto &V1, const auto &V2) {
532 return ComparePointers(std::get<0>(V1.front()), std::get<0>(V2.front()));
533 });
534 }
535
536 for (auto &T : Bases)
537 for (const auto &Vec : T.second)
538 for (const auto &P : Vec)
539 SortedIndices.push_back(Elt: std::get<2>(t: P));
540
541 assert(SortedIndices.size() == VL.size() &&
542 "Expected SortedIndices to be the size of VL");
543 return true;
544}
545
546} // namespace llvm::slpvectorizer
547