1//===-------- LoopIdiomVectorize.cpp - Loop idiom vectorization -----------===//
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// This pass implements a pass that recognizes certain loop idioms and
10// transforms them into more optimized versions of the same loop. In cases
11// where this happens, it can be a significant performance win.
12//
13// We currently support two loops:
14//
15// 1. A loop that finds the first mismatched byte in an array and returns the
16// index, i.e. something like:
17//
18// while (++i != n) {
19// if (a[i] != b[i])
20// break;
21// }
22//
23// In this example we can actually vectorize the loop despite the early exit,
24// although the loop vectorizer does not support it. It requires some extra
25// checks to deal with the possibility of faulting loads when crossing page
26// boundaries. However, even with these checks it is still profitable to do the
27// transformation.
28//
29// TODO List:
30//
31// * Add support for the inverse case where we scan for a matching element.
32// * Permit 64-bit induction variable types.
33// * Recognize loops that increment the IV *after* comparing bytes.
34// * Allow 32-bit sign-extends of the IV used by the GEP.
35//
36// 2. A loop that finds the first matching character in an array among a set of
37// possible matches, e.g.:
38//
39// for (; first != last; ++first)
40// for (s_it = s_first; s_it != s_last; ++s_it)
41// if (*first == *s_it)
42// return first;
43// return last;
44//
45// This corresponds to std::find_first_of (for arrays of bytes) from the C++
46// standard library. This function can be implemented efficiently for targets
47// that support @llvm.experimental.vector.match. For example, on AArch64 targets
48// that implement SVE2, this lower to a MATCH instruction, which enables us to
49// perform up to 16x16=256 comparisons in one go. This can lead to very
50// significant speedups.
51//
52// TODO:
53//
54// * Add support for `find_first_not_of' loops (i.e. with not-equal comparison).
55// * Make VF a configurable parameter (right now we assume 128-bit vectors).
56// * Potentially adjust the cost model to let the transformation kick-in even if
57// @llvm.experimental.vector.match doesn't have direct support in hardware.
58//
59//===----------------------------------------------------------------------===//
60//
61// NOTE: This Pass matches really specific loop patterns because it's only
62// supposed to be a temporary solution until our LoopVectorizer is powerful
63// enough to vectorize them automatically.
64//
65//===----------------------------------------------------------------------===//
66
67#include "llvm/Transforms/Vectorize/LoopIdiomVectorize.h"
68#include "llvm/Analysis/DomTreeUpdater.h"
69#include "llvm/Analysis/LoopPass.h"
70#include "llvm/Analysis/OptimizationRemarkEmitter.h"
71#include "llvm/Analysis/TargetTransformInfo.h"
72#include "llvm/IR/Dominators.h"
73#include "llvm/IR/IRBuilder.h"
74#include "llvm/IR/Intrinsics.h"
75#include "llvm/IR/MDBuilder.h"
76#include "llvm/IR/PatternMatch.h"
77#include "llvm/Transforms/Utils/BasicBlockUtils.h"
78#include "llvm/Transforms/Vectorize/LoopVectorizationLegality.h"
79
80using namespace llvm;
81using namespace PatternMatch;
82
83#define DEBUG_TYPE "loop-idiom-vectorize"
84
85static cl::opt<bool> DisableAll("disable-loop-idiom-vectorize-all", cl::Hidden,
86 cl::init(Val: false),
87 cl::desc("Disable Loop Idiom Vectorize Pass."));
88
89static cl::opt<LoopIdiomVectorizeStyle>
90 LITVecStyle("loop-idiom-vectorize-style", cl::Hidden,
91 cl::desc("The vectorization style for loop idiom transform."),
92 cl::values(clEnumValN(LoopIdiomVectorizeStyle::Masked, "masked",
93 "Use masked vector intrinsics"),
94 clEnumValN(LoopIdiomVectorizeStyle::Predicated,
95 "predicated", "Use VP intrinsics")),
96 cl::init(Val: LoopIdiomVectorizeStyle::Masked));
97
98static cl::opt<bool>
99 DisableByteCmp("disable-loop-idiom-vectorize-bytecmp", cl::Hidden,
100 cl::init(Val: false),
101 cl::desc("Proceed with Loop Idiom Vectorize Pass, but do "
102 "not convert byte-compare loop(s)."));
103
104static cl::opt<unsigned>
105 ByteCmpVF("loop-idiom-vectorize-bytecmp-vf", cl::Hidden,
106 cl::desc("The vectorization factor for byte-compare patterns."),
107 cl::init(Val: 16));
108
109static cl::opt<bool>
110 DisableFindFirstByte("disable-loop-idiom-vectorize-find-first-byte",
111 cl::Hidden, cl::init(Val: false),
112 cl::desc("Do not convert find-first-byte loop(s)."));
113
114static cl::opt<bool>
115 VerifyLoops("loop-idiom-vectorize-verify", cl::Hidden, cl::init(Val: false),
116 cl::desc("Verify loops generated Loop Idiom Vectorize Pass."));
117
118namespace {
119class LoopIdiomVectorize {
120 LoopIdiomVectorizeStyle VectorizeStyle;
121 unsigned ByteCompareVF;
122 Loop *CurLoop = nullptr;
123 DominatorTree *DT;
124 LoopInfo *LI;
125 const TargetTransformInfo *TTI;
126 const DataLayout *DL;
127
128 /// Interface to emit optimization remarks.
129 OptimizationRemarkEmitter &ORE;
130
131 // Blocks that will be used for inserting vectorized code.
132 BasicBlock *EndBlock = nullptr;
133 BasicBlock *VectorLoopPreheaderBlock = nullptr;
134 BasicBlock *VectorLoopStartBlock = nullptr;
135 BasicBlock *VectorLoopMismatchBlock = nullptr;
136 BasicBlock *VectorLoopIncBlock = nullptr;
137
138public:
139 LoopIdiomVectorize(LoopIdiomVectorizeStyle S, unsigned VF, DominatorTree *DT,
140 LoopInfo *LI, const TargetTransformInfo *TTI,
141 const DataLayout *DL, OptimizationRemarkEmitter &ORE)
142 : VectorizeStyle(S), ByteCompareVF(VF), DT(DT), LI(LI), TTI(TTI), DL(DL),
143 ORE(ORE) {}
144
145 bool run(Loop *L);
146
147private:
148 /// \name Countable Loop Idiom Handling
149 /// @{
150
151 bool runOnLoopBlock(BasicBlock *BB, const SCEV *BECount,
152 SmallVectorImpl<BasicBlock *> &ExitBlocks);
153
154 bool recognizeByteCompare();
155
156 Value *expandFindMismatch(IRBuilder<> &Builder, DomTreeUpdater &DTU,
157 GetElementPtrInst *GEPA, GetElementPtrInst *GEPB,
158 Instruction *Index, Value *Start, Value *MaxLen);
159
160 Value *createMaskedFindMismatch(IRBuilder<> &Builder, DomTreeUpdater &DTU,
161 GetElementPtrInst *GEPA,
162 GetElementPtrInst *GEPB, Value *ExtStart,
163 Value *ExtEnd);
164 Value *createPredicatedFindMismatch(IRBuilder<> &Builder, DomTreeUpdater &DTU,
165 GetElementPtrInst *GEPA,
166 GetElementPtrInst *GEPB, Value *ExtStart,
167 Value *ExtEnd);
168
169 void transformByteCompare(GetElementPtrInst *GEPA, GetElementPtrInst *GEPB,
170 PHINode *IndPhi, Value *MaxLen, Instruction *Index,
171 Value *Start, bool IncIdx, BasicBlock *FoundBB,
172 BasicBlock *EndBB);
173
174 bool recognizeFindFirstByte();
175
176 Value *expandFindFirstByte(IRBuilder<> &Builder, DomTreeUpdater &DTU,
177 unsigned VF, Type *CharTy, Value *IndPhi,
178 BasicBlock *ExitSucc, BasicBlock *ExitFail,
179 Value *SearchStart, Value *SearchEnd,
180 Value *NeedleStart, Value *NeedleEnd);
181
182 void transformFindFirstByte(PHINode *IndPhi, unsigned VF, Type *CharTy,
183 BasicBlock *ExitSucc, BasicBlock *ExitFail,
184 Value *SearchStart, Value *SearchEnd,
185 Value *NeedleStart, Value *NeedleEnd);
186 /// @}
187};
188} // anonymous namespace
189
190PreservedAnalyses LoopIdiomVectorizePass::run(Loop &L, LoopAnalysisManager &AM,
191 LoopStandardAnalysisResults &AR,
192 LPMUpdater &) {
193 if (DisableAll)
194 return PreservedAnalyses::all();
195
196 const auto *DL = &L.getHeader()->getDataLayout();
197
198 LoopIdiomVectorizeStyle VecStyle = VectorizeStyle;
199 if (LITVecStyle.getNumOccurrences())
200 VecStyle = LITVecStyle;
201
202 unsigned BCVF = ByteCompareVF;
203 if (ByteCmpVF.getNumOccurrences())
204 BCVF = ByteCmpVF;
205
206 Function &F = *L.getHeader()->getParent();
207 auto &FAMP = AM.getResult<FunctionAnalysisManagerLoopProxy>(IR&: L, ExtraArgs&: AR);
208 auto *ORE = FAMP.getCachedResult<OptimizationRemarkEmitterAnalysis>(IR&: F);
209
210 std::optional<OptimizationRemarkEmitter> ORELocal;
211 if (!ORE) {
212 ORELocal.emplace(args: &F);
213 ORE = &*ORELocal;
214 }
215
216 LoopIdiomVectorize LIV(VecStyle, BCVF, &AR.DT, &AR.LI, &AR.TTI, DL, *ORE);
217 if (!LIV.run(L: &L))
218 return PreservedAnalyses::all();
219
220 return PreservedAnalyses::none();
221}
222
223//===----------------------------------------------------------------------===//
224//
225// Implementation of LoopIdiomVectorize
226//
227//===----------------------------------------------------------------------===//
228
229bool LoopIdiomVectorize::run(Loop *L) {
230 CurLoop = L;
231
232 Function &F = *L->getHeader()->getParent();
233 if (DisableAll || F.hasOptSize())
234 return false;
235
236 // Bail if vectorization is disabled on loop.
237 LoopVectorizeHints Hints(L, /*InterleaveOnlyWhenForced=*/true, ORE);
238 if (!Hints.allowVectorization(F: &F, L, /*VectorizeOnlyWhenForced=*/false)) {
239 LLVM_DEBUG(dbgs() << DEBUG_TYPE << " is disabled on " << L->getName()
240 << " due to vectorization hints\n");
241 return false;
242 }
243
244 if (F.hasFnAttribute(Kind: Attribute::NoImplicitFloat)) {
245 LLVM_DEBUG(dbgs() << DEBUG_TYPE << " is disabled on " << F.getName()
246 << " due to its NoImplicitFloat attribute");
247 return false;
248 }
249
250 // If the loop could not be converted to canonical form, it must have an
251 // indirectbr in it, just give up.
252 if (!L->getLoopPreheader())
253 return false;
254
255 LLVM_DEBUG(dbgs() << DEBUG_TYPE " Scanning: F[" << F.getName() << "] Loop %"
256 << CurLoop->getHeader()->getName() << "\n");
257
258 if (recognizeByteCompare())
259 return true;
260
261 if (recognizeFindFirstByte())
262 return true;
263
264 return false;
265}
266
267static void fixSuccessorPhis(Loop *L, Value *ScalarRes, Value *VectorRes,
268 BasicBlock *SuccBB, BasicBlock *IncBB) {
269 for (PHINode &PN : SuccBB->phis()) {
270 // Look through the incoming values to find ScalarRes, meaning this is a
271 // PHI collecting the results of the transformation.
272 bool ResPhi = false;
273 for (Value *Op : PN.incoming_values())
274 if (Op == ScalarRes) {
275 ResPhi = true;
276 break;
277 }
278
279 // Any PHI that depended upon the result of the transformation needs a new
280 // incoming value from IncBB.
281 if (ResPhi)
282 PN.addIncoming(V: VectorRes, BB: IncBB);
283 else {
284 // There should be no other outside uses of other values in the
285 // original loop. Any incoming values should either:
286 // 1. Be for blocks outside the loop, which aren't interesting. Or ..
287 // 2. These are from blocks in the loop with values defined outside
288 // the loop. We should a similar incoming value from CmpBB.
289 for (BasicBlock *BB : PN.blocks())
290 if (L->contains(BB)) {
291 PN.addIncoming(V: PN.getIncomingValueForBlock(BB), BB: IncBB);
292 break;
293 }
294 }
295 }
296}
297
298bool LoopIdiomVectorize::recognizeByteCompare() {
299 // Currently the transformation only works on scalable vector types, although
300 // there is no fundamental reason why it cannot be made to work for fixed
301 // width too.
302
303 // We also need to know the minimum page size for the target in order to
304 // generate runtime memory checks to ensure the vector version won't fault.
305 if (!TTI->supportsScalableVectors() || !TTI->getMinPageSize().has_value() ||
306 DisableByteCmp)
307 return false;
308
309 BasicBlock *Header = CurLoop->getHeader();
310
311 // In LoopIdiomVectorize::run we have already checked that the loop
312 // has a preheader so we can assume it's in a canonical form.
313 if (CurLoop->getNumBackEdges() != 1 || CurLoop->getNumBlocks() != 2)
314 return false;
315
316 PHINode *PN = dyn_cast<PHINode>(Val: &Header->front());
317 if (!PN || PN->getNumIncomingValues() != 2)
318 return false;
319
320 auto LoopBlocks = CurLoop->getBlocks();
321 // The first block in the loop should contain only 4 instructions, e.g.
322 //
323 // while.cond:
324 // %res.phi = phi i32 [ %start, %ph ], [ %inc, %while.body ]
325 // %inc = add i32 %res.phi, 1
326 // %cmp.not = icmp eq i32 %inc, %n
327 // br i1 %cmp.not, label %while.end, label %while.body
328 //
329 if (LoopBlocks[0]->size() > 4)
330 return false;
331
332 // The second block should contain 7 instructions, e.g.
333 //
334 // while.body:
335 // %idx = zext i32 %inc to i64
336 // %idx.a = getelementptr inbounds i8, ptr %a, i64 %idx
337 // %load.a = load i8, ptr %idx.a
338 // %idx.b = getelementptr inbounds i8, ptr %b, i64 %idx
339 // %load.b = load i8, ptr %idx.b
340 // %cmp.not.ld = icmp eq i8 %load.a, %load.b
341 // br i1 %cmp.not.ld, label %while.cond, label %while.end
342 //
343 if (LoopBlocks[1]->size() > 7)
344 return false;
345
346 // The incoming value to the PHI node from the loop should be an add of 1.
347 Value *StartIdx = nullptr;
348 Instruction *Index = nullptr;
349 if (!CurLoop->contains(BB: PN->getIncomingBlock(i: 0))) {
350 StartIdx = PN->getIncomingValue(i: 0);
351 Index = dyn_cast<Instruction>(Val: PN->getIncomingValue(i: 1));
352 } else {
353 StartIdx = PN->getIncomingValue(i: 1);
354 Index = dyn_cast<Instruction>(Val: PN->getIncomingValue(i: 0));
355 }
356
357 // Limit to 32-bit types for now
358 if (!Index || !Index->getType()->isIntegerTy(BitWidth: 32) ||
359 !match(V: Index, P: m_c_Add(L: m_Specific(V: PN), R: m_One())))
360 return false;
361
362 // If we match the pattern, PN and Index will be replaced with the result of
363 // the cttz.elts intrinsic. If any other instructions are used outside of
364 // the loop, we cannot replace it.
365 for (BasicBlock *BB : LoopBlocks)
366 for (Instruction &I : *BB)
367 if (&I != PN && &I != Index)
368 for (User *U : I.users())
369 if (!CurLoop->contains(Inst: cast<Instruction>(Val: U)))
370 return false;
371
372 // Match the branch instruction for the header
373 Value *MaxLen;
374 BasicBlock *EndBB, *WhileBB;
375 if (!match(V: Header->getTerminator(),
376 P: m_Br(C: m_SpecificICmp(MatchPred: ICmpInst::ICMP_EQ, L: m_Specific(V: Index),
377 R: m_Value(V&: MaxLen)),
378 T: m_BasicBlock(V&: EndBB), F: m_BasicBlock(V&: WhileBB))) ||
379 !CurLoop->contains(BB: WhileBB))
380 return false;
381
382 // WhileBB should contain the pattern of load & compare instructions. Match
383 // the pattern and find the GEP instructions used by the loads.
384 BasicBlock *FoundBB;
385 BasicBlock *TrueBB;
386 Value *LoadA, *LoadB;
387 if (!match(V: WhileBB->getTerminator(),
388 P: m_Br(C: m_SpecificICmp(MatchPred: ICmpInst::ICMP_EQ, L: m_Value(V&: LoadA),
389 R: m_Value(V&: LoadB)),
390 T: m_BasicBlock(V&: TrueBB), F: m_BasicBlock(V&: FoundBB))) ||
391 !CurLoop->contains(BB: TrueBB))
392 return false;
393
394 Value *A, *B;
395 if (!match(V: LoadA, P: m_Load(Op: m_Value(V&: A))) || !match(V: LoadB, P: m_Load(Op: m_Value(V&: B))))
396 return false;
397
398 LoadInst *LoadAI = cast<LoadInst>(Val: LoadA);
399 LoadInst *LoadBI = cast<LoadInst>(Val: LoadB);
400 if (!LoadAI->isSimple() || !LoadBI->isSimple())
401 return false;
402
403 GetElementPtrInst *GEPA = dyn_cast<GetElementPtrInst>(Val: A);
404 GetElementPtrInst *GEPB = dyn_cast<GetElementPtrInst>(Val: B);
405
406 if (!GEPA || !GEPB)
407 return false;
408
409 Value *PtrA = GEPA->getPointerOperand();
410 Value *PtrB = GEPB->getPointerOperand();
411
412 // Check we are loading i8 values from two loop invariant pointers
413 if (!CurLoop->isLoopInvariant(V: PtrA) || !CurLoop->isLoopInvariant(V: PtrB) ||
414 !GEPA->getResultElementType()->isIntegerTy(BitWidth: 8) ||
415 !GEPB->getResultElementType()->isIntegerTy(BitWidth: 8) ||
416 !LoadAI->getType()->isIntegerTy(BitWidth: 8) ||
417 !LoadBI->getType()->isIntegerTy(BitWidth: 8) || PtrA == PtrB)
418 return false;
419
420 // Check that the index to the GEPs is the index we found earlier
421 if (GEPA->getNumIndices() > 1 || GEPB->getNumIndices() > 1)
422 return false;
423
424 Value *IdxA = GEPA->getOperand(i_nocapture: GEPA->getNumIndices());
425 Value *IdxB = GEPB->getOperand(i_nocapture: GEPB->getNumIndices());
426 if (IdxA != IdxB || !match(V: IdxA, P: m_ZExt(Op: m_Specific(V: Index))))
427 return false;
428
429 // We only ever expect the pre-incremented index value to be used inside the
430 // loop.
431 if (!PN->hasOneUse())
432 return false;
433
434 // Ensure that when the Found and End blocks are identical the PHIs have the
435 // supported format. We don't currently allow cases like this:
436 // while.cond:
437 // ...
438 // br i1 %cmp.not, label %while.end, label %while.body
439 //
440 // while.body:
441 // ...
442 // br i1 %cmp.not2, label %while.cond, label %while.end
443 //
444 // while.end:
445 // %final_ptr = phi ptr [ %c, %while.body ], [ %d, %while.cond ]
446 //
447 // Where the incoming values for %final_ptr are unique and from each of the
448 // loop blocks, but not actually defined in the loop. This requires extra
449 // work setting up the byte.compare block, i.e. by introducing a select to
450 // choose the correct value.
451 // TODO: We could add support for this in future.
452 if (FoundBB == EndBB) {
453 for (PHINode &EndPN : EndBB->phis()) {
454 Value *WhileCondVal = EndPN.getIncomingValueForBlock(BB: Header);
455 Value *WhileBodyVal = EndPN.getIncomingValueForBlock(BB: WhileBB);
456
457 // The value of the index when leaving the while.cond block is always the
458 // same as the end value (MaxLen) so we permit either. The value when
459 // leaving the while.body block should only be the index. Otherwise for
460 // any other values we only allow ones that are same for both blocks.
461 if (WhileCondVal != WhileBodyVal &&
462 ((WhileCondVal != Index && WhileCondVal != MaxLen) ||
463 (WhileBodyVal != Index)))
464 return false;
465 }
466 }
467
468 LLVM_DEBUG(dbgs() << "FOUND IDIOM IN LOOP: \n"
469 << *(EndBB->getParent()) << "\n\n");
470
471 // The index is incremented before the GEP/Load pair so we need to
472 // add 1 to the start value.
473 transformByteCompare(GEPA, GEPB, IndPhi: PN, MaxLen, Index, Start: StartIdx, /*IncIdx=*/true,
474 FoundBB, EndBB);
475 return true;
476}
477
478Value *LoopIdiomVectorize::createMaskedFindMismatch(
479 IRBuilder<> &Builder, DomTreeUpdater &DTU, GetElementPtrInst *GEPA,
480 GetElementPtrInst *GEPB, Value *ExtStart, Value *ExtEnd) {
481 Type *I64Type = Builder.getInt64Ty();
482 Type *ResType = Builder.getInt32Ty();
483 Type *LoadType = Builder.getInt8Ty();
484 Value *PtrA = GEPA->getPointerOperand();
485 Value *PtrB = GEPB->getPointerOperand();
486
487 ScalableVectorType *PredVTy =
488 ScalableVectorType::get(ElementType: Builder.getInt1Ty(), MinNumElts: ByteCompareVF);
489
490 Value *InitialPred = Builder.CreateIntrinsic(
491 ID: Intrinsic::get_active_lane_mask, OverloadTypes: {PredVTy, I64Type}, Args: {ExtStart, ExtEnd});
492
493 Value *VecLen = Builder.CreateVScale(Ty: I64Type);
494 VecLen =
495 Builder.CreateMul(LHS: VecLen, RHS: ConstantInt::get(Ty: I64Type, V: ByteCompareVF), Name: "",
496 /*HasNUW=*/true, /*HasNSW=*/true);
497
498 Value *PFalse = Builder.CreateVectorSplat(EC: PredVTy->getElementCount(),
499 V: Builder.getInt1(V: false));
500
501 Builder.CreateBr(Dest: VectorLoopStartBlock);
502
503 DTU.applyUpdates(Updates: {{DominatorTree::Insert, VectorLoopPreheaderBlock,
504 VectorLoopStartBlock}});
505
506 // Set up the first vector loop block by creating the PHIs, doing the vector
507 // loads and comparing the vectors.
508 Builder.SetInsertPoint(VectorLoopStartBlock);
509 PHINode *LoopPred = Builder.CreatePHI(Ty: PredVTy, NumReservedValues: 2, Name: "mismatch_vec_loop_pred");
510 LoopPred->addIncoming(V: InitialPred, BB: VectorLoopPreheaderBlock);
511 PHINode *VectorIndexPhi = Builder.CreatePHI(Ty: I64Type, NumReservedValues: 2, Name: "mismatch_vec_index");
512 VectorIndexPhi->addIncoming(V: ExtStart, BB: VectorLoopPreheaderBlock);
513 Type *VectorLoadType =
514 ScalableVectorType::get(ElementType: Builder.getInt8Ty(), MinNumElts: ByteCompareVF);
515 Value *Passthru = ConstantInt::getNullValue(Ty: VectorLoadType);
516
517 Value *VectorLhsGep =
518 Builder.CreateGEP(Ty: LoadType, Ptr: PtrA, IdxList: VectorIndexPhi, Name: "", NW: GEPA->isInBounds());
519 Value *VectorLhsLoad = Builder.CreateMaskedLoad(Ty: VectorLoadType, Ptr: VectorLhsGep,
520 Alignment: Align(1), Mask: LoopPred, PassThru: Passthru);
521
522 Value *VectorRhsGep =
523 Builder.CreateGEP(Ty: LoadType, Ptr: PtrB, IdxList: VectorIndexPhi, Name: "", NW: GEPB->isInBounds());
524 Value *VectorRhsLoad = Builder.CreateMaskedLoad(Ty: VectorLoadType, Ptr: VectorRhsGep,
525 Alignment: Align(1), Mask: LoopPred, PassThru: Passthru);
526
527 Value *VectorMatchCmp = Builder.CreateICmpNE(LHS: VectorLhsLoad, RHS: VectorRhsLoad);
528 VectorMatchCmp = Builder.CreateSelect(C: LoopPred, True: VectorMatchCmp, False: PFalse);
529 Value *VectorMatchHasActiveLanes = Builder.CreateOrReduce(Src: VectorMatchCmp);
530 Builder.CreateCondBr(Cond: VectorMatchHasActiveLanes, True: VectorLoopMismatchBlock,
531 False: VectorLoopIncBlock);
532
533 DTU.applyUpdates(
534 Updates: {{DominatorTree::Insert, VectorLoopStartBlock, VectorLoopMismatchBlock},
535 {DominatorTree::Insert, VectorLoopStartBlock, VectorLoopIncBlock}});
536
537 // Increment the index counter and calculate the predicate for the next
538 // iteration of the loop. We branch back to the start of the loop if there
539 // is at least one active lane.
540 Builder.SetInsertPoint(VectorLoopIncBlock);
541 Value *NewVectorIndexPhi =
542 Builder.CreateAdd(LHS: VectorIndexPhi, RHS: VecLen, Name: "",
543 /*HasNUW=*/true, /*HasNSW=*/true);
544 VectorIndexPhi->addIncoming(V: NewVectorIndexPhi, BB: VectorLoopIncBlock);
545 Value *NewPred =
546 Builder.CreateIntrinsic(ID: Intrinsic::get_active_lane_mask,
547 OverloadTypes: {PredVTy, I64Type}, Args: {NewVectorIndexPhi, ExtEnd});
548 LoopPred->addIncoming(V: NewPred, BB: VectorLoopIncBlock);
549
550 Value *PredHasActiveLanes =
551 Builder.CreateExtractElement(Vec: NewPred, Idx: uint64_t(0));
552 Builder.CreateCondBr(Cond: PredHasActiveLanes, True: VectorLoopStartBlock, False: EndBlock);
553
554 DTU.applyUpdates(
555 Updates: {{DominatorTree::Insert, VectorLoopIncBlock, VectorLoopStartBlock},
556 {DominatorTree::Insert, VectorLoopIncBlock, EndBlock}});
557
558 // If we found a mismatch then we need to calculate which lane in the vector
559 // had a mismatch and add that on to the current loop index.
560 Builder.SetInsertPoint(VectorLoopMismatchBlock);
561 PHINode *FoundPred = Builder.CreatePHI(Ty: PredVTy, NumReservedValues: 1, Name: "mismatch_vec_found_pred");
562 FoundPred->addIncoming(V: VectorMatchCmp, BB: VectorLoopStartBlock);
563 PHINode *LastLoopPred =
564 Builder.CreatePHI(Ty: PredVTy, NumReservedValues: 1, Name: "mismatch_vec_last_loop_pred");
565 LastLoopPred->addIncoming(V: LoopPred, BB: VectorLoopStartBlock);
566 PHINode *VectorFoundIndex =
567 Builder.CreatePHI(Ty: I64Type, NumReservedValues: 1, Name: "mismatch_vec_found_index");
568 VectorFoundIndex->addIncoming(V: VectorIndexPhi, BB: VectorLoopStartBlock);
569
570 Value *PredMatchCmp = Builder.CreateAnd(LHS: LastLoopPred, RHS: FoundPred);
571 Value *Ctz = Builder.CreateCountTrailingZeroElems(ResTy: ResType, Mask: PredMatchCmp);
572 Ctz = Builder.CreateZExt(V: Ctz, DestTy: I64Type);
573 Value *VectorLoopRes64 = Builder.CreateAdd(LHS: VectorFoundIndex, RHS: Ctz, Name: "",
574 /*HasNUW=*/true, /*HasNSW=*/true);
575 return Builder.CreateTrunc(V: VectorLoopRes64, DestTy: ResType);
576}
577
578Value *LoopIdiomVectorize::createPredicatedFindMismatch(
579 IRBuilder<> &Builder, DomTreeUpdater &DTU, GetElementPtrInst *GEPA,
580 GetElementPtrInst *GEPB, Value *ExtStart, Value *ExtEnd) {
581 Type *I64Type = Builder.getInt64Ty();
582 Type *I32Type = Builder.getInt32Ty();
583 Type *ResType = I32Type;
584 Type *LoadType = Builder.getInt8Ty();
585 Value *PtrA = GEPA->getPointerOperand();
586 Value *PtrB = GEPB->getPointerOperand();
587
588 auto *JumpToVectorLoop = UncondBrInst::Create(Target: VectorLoopStartBlock);
589 Builder.Insert(I: JumpToVectorLoop);
590
591 DTU.applyUpdates(Updates: {{DominatorTree::Insert, VectorLoopPreheaderBlock,
592 VectorLoopStartBlock}});
593
594 // Set up the first Vector loop block by creating the PHIs, doing the vector
595 // loads and comparing the vectors.
596 Builder.SetInsertPoint(VectorLoopStartBlock);
597 auto *VectorIndexPhi = Builder.CreatePHI(Ty: I64Type, NumReservedValues: 2, Name: "mismatch_vector_index");
598 VectorIndexPhi->addIncoming(V: ExtStart, BB: VectorLoopPreheaderBlock);
599
600 // Calculate AVL by subtracting the vector loop index from the trip count
601 Value *AVL = Builder.CreateSub(LHS: ExtEnd, RHS: VectorIndexPhi, Name: "avl", /*HasNUW=*/true,
602 /*HasNSW=*/true);
603
604 auto *VectorLoadType = ScalableVectorType::get(ElementType: LoadType, MinNumElts: ByteCompareVF);
605 auto *VF = ConstantInt::get(Ty: I32Type, V: ByteCompareVF);
606
607 Value *VL = Builder.CreateIntrinsic(ID: Intrinsic::experimental_get_vector_length,
608 OverloadTypes: {I64Type}, Args: {AVL, VF, Builder.getTrue()});
609 Value *GepOffset = VectorIndexPhi;
610
611 Value *VectorLhsGep =
612 Builder.CreateGEP(Ty: LoadType, Ptr: PtrA, IdxList: GepOffset, Name: "", NW: GEPA->isInBounds());
613 VectorType *TrueMaskTy =
614 VectorType::get(ElementType: Builder.getInt1Ty(), EC: VectorLoadType->getElementCount());
615 Value *AllTrueMask = Constant::getAllOnesValue(Ty: TrueMaskTy);
616 Value *VectorLhsLoad = Builder.CreateIntrinsic(
617 ID: Intrinsic::vp_load, OverloadTypes: {VectorLoadType, VectorLhsGep->getType()},
618 Args: {VectorLhsGep, AllTrueMask, VL}, FMFSource: nullptr, Name: "lhs.load");
619
620 Value *VectorRhsGep =
621 Builder.CreateGEP(Ty: LoadType, Ptr: PtrB, IdxList: GepOffset, Name: "", NW: GEPB->isInBounds());
622 Value *VectorRhsLoad = Builder.CreateIntrinsic(
623 ID: Intrinsic::vp_load, OverloadTypes: {VectorLoadType, VectorLhsGep->getType()},
624 Args: {VectorRhsGep, AllTrueMask, VL}, FMFSource: nullptr, Name: "rhs.load");
625
626 Value *VectorMatchCmp =
627 Builder.CreateICmpNE(LHS: VectorLhsLoad, RHS: VectorRhsLoad, Name: "mismatch.cmp");
628 Value *CTZ = Builder.CreateIntrinsic(
629 ID: Intrinsic::vp_cttz_elts, OverloadTypes: {ResType, VectorMatchCmp->getType()},
630 Args: {VectorMatchCmp, /*ZeroIsPoison=*/Builder.getInt1(V: false), AllTrueMask,
631 VL});
632 Value *MismatchFound = Builder.CreateICmpNE(LHS: CTZ, RHS: VL);
633 auto *VectorEarlyExit = CondBrInst::Create(
634 Cond: MismatchFound, IfTrue: VectorLoopMismatchBlock, IfFalse: VectorLoopIncBlock);
635 Builder.Insert(I: VectorEarlyExit);
636
637 DTU.applyUpdates(
638 Updates: {{DominatorTree::Insert, VectorLoopStartBlock, VectorLoopMismatchBlock},
639 {DominatorTree::Insert, VectorLoopStartBlock, VectorLoopIncBlock}});
640
641 // Increment the index counter and calculate the predicate for the next
642 // iteration of the loop. We branch back to the start of the loop if there
643 // is at least one active lane.
644 Builder.SetInsertPoint(VectorLoopIncBlock);
645 Value *VL64 = Builder.CreateZExt(V: VL, DestTy: I64Type);
646 Value *NewVectorIndexPhi =
647 Builder.CreateAdd(LHS: VectorIndexPhi, RHS: VL64, Name: "",
648 /*HasNUW=*/true, /*HasNSW=*/true);
649 VectorIndexPhi->addIncoming(V: NewVectorIndexPhi, BB: VectorLoopIncBlock);
650 Value *ExitCond = Builder.CreateICmpNE(LHS: NewVectorIndexPhi, RHS: ExtEnd);
651 auto *VectorLoopBranchBack =
652 CondBrInst::Create(Cond: ExitCond, IfTrue: VectorLoopStartBlock, IfFalse: EndBlock);
653 Builder.Insert(I: VectorLoopBranchBack);
654
655 DTU.applyUpdates(
656 Updates: {{DominatorTree::Insert, VectorLoopIncBlock, VectorLoopStartBlock},
657 {DominatorTree::Insert, VectorLoopIncBlock, EndBlock}});
658
659 // If we found a mismatch then we need to calculate which lane in the vector
660 // had a mismatch and add that on to the current loop index.
661 Builder.SetInsertPoint(VectorLoopMismatchBlock);
662
663 // Add LCSSA phis for CTZ and VectorIndexPhi.
664 auto *CTZLCSSAPhi = Builder.CreatePHI(Ty: CTZ->getType(), NumReservedValues: 1, Name: "ctz");
665 CTZLCSSAPhi->addIncoming(V: CTZ, BB: VectorLoopStartBlock);
666 auto *VectorIndexLCSSAPhi =
667 Builder.CreatePHI(Ty: VectorIndexPhi->getType(), NumReservedValues: 1, Name: "mismatch_vector_index");
668 VectorIndexLCSSAPhi->addIncoming(V: VectorIndexPhi, BB: VectorLoopStartBlock);
669
670 Value *CTZI64 = Builder.CreateZExt(V: CTZLCSSAPhi, DestTy: I64Type);
671 Value *VectorLoopRes64 = Builder.CreateAdd(LHS: VectorIndexLCSSAPhi, RHS: CTZI64, Name: "",
672 /*HasNUW=*/true, /*HasNSW=*/true);
673 return Builder.CreateTrunc(V: VectorLoopRes64, DestTy: ResType);
674}
675
676Value *LoopIdiomVectorize::expandFindMismatch(
677 IRBuilder<> &Builder, DomTreeUpdater &DTU, GetElementPtrInst *GEPA,
678 GetElementPtrInst *GEPB, Instruction *Index, Value *Start, Value *MaxLen) {
679 Value *PtrA = GEPA->getPointerOperand();
680 Value *PtrB = GEPB->getPointerOperand();
681
682 // Get the arguments and types for the intrinsic.
683 BasicBlock *Preheader = CurLoop->getLoopPreheader();
684 Instruction *PHBranch = Preheader->getTerminator();
685 LLVMContext &Ctx = PHBranch->getContext();
686 Type *LoadType = Type::getInt8Ty(C&: Ctx);
687 Type *ResType = Builder.getInt32Ty();
688
689 // Split block in the original loop preheader.
690 EndBlock = SplitBlock(Old: Preheader, SplitPt: PHBranch, DT, LI, MSSAU: nullptr, BBName: "mismatch_end");
691
692 // Create the blocks that we're going to need:
693 // 1. A block for checking the zero-extended length exceeds 0
694 // 2. A block to check that the start and end addresses of a given array
695 // lie on the same page.
696 // 3. The vector loop preheader.
697 // 4. The first vector loop block.
698 // 5. The vector loop increment block.
699 // 6. A block we can jump to from the vector loop when a mismatch is found.
700 // 7. The first block of the scalar loop itself, containing PHIs , loads
701 // and cmp.
702 // 8. A scalar loop increment block to increment the PHIs and go back
703 // around the loop.
704
705 BasicBlock *MinItCheckBlock = BasicBlock::Create(
706 Context&: Ctx, Name: "mismatch_min_it_check", Parent: EndBlock->getParent(), InsertBefore: EndBlock);
707
708 // Update the terminator added by SplitBlock to branch to the first block
709 Preheader->getTerminator()->setSuccessor(Idx: 0, BB: MinItCheckBlock);
710
711 BasicBlock *MemCheckBlock = BasicBlock::Create(
712 Context&: Ctx, Name: "mismatch_mem_check", Parent: EndBlock->getParent(), InsertBefore: EndBlock);
713
714 VectorLoopPreheaderBlock = BasicBlock::Create(
715 Context&: Ctx, Name: "mismatch_vec_loop_preheader", Parent: EndBlock->getParent(), InsertBefore: EndBlock);
716
717 VectorLoopStartBlock = BasicBlock::Create(Context&: Ctx, Name: "mismatch_vec_loop",
718 Parent: EndBlock->getParent(), InsertBefore: EndBlock);
719
720 VectorLoopIncBlock = BasicBlock::Create(Context&: Ctx, Name: "mismatch_vec_loop_inc",
721 Parent: EndBlock->getParent(), InsertBefore: EndBlock);
722
723 VectorLoopMismatchBlock = BasicBlock::Create(Context&: Ctx, Name: "mismatch_vec_loop_found",
724 Parent: EndBlock->getParent(), InsertBefore: EndBlock);
725
726 BasicBlock *LoopPreHeaderBlock = BasicBlock::Create(
727 Context&: Ctx, Name: "mismatch_loop_pre", Parent: EndBlock->getParent(), InsertBefore: EndBlock);
728
729 BasicBlock *LoopStartBlock =
730 BasicBlock::Create(Context&: Ctx, Name: "mismatch_loop", Parent: EndBlock->getParent(), InsertBefore: EndBlock);
731
732 BasicBlock *LoopIncBlock = BasicBlock::Create(
733 Context&: Ctx, Name: "mismatch_loop_inc", Parent: EndBlock->getParent(), InsertBefore: EndBlock);
734
735 DTU.applyUpdates(Updates: {{DominatorTree::Insert, Preheader, MinItCheckBlock},
736 {DominatorTree::Delete, Preheader, EndBlock}});
737
738 // Update LoopInfo with the new vector & scalar loops.
739 auto VectorLoop = LI->AllocateLoop();
740 auto ScalarLoop = LI->AllocateLoop();
741
742 if (CurLoop->getParentLoop()) {
743 CurLoop->getParentLoop()->addBasicBlockToLoop(NewBB: MinItCheckBlock, LI&: *LI);
744 CurLoop->getParentLoop()->addBasicBlockToLoop(NewBB: MemCheckBlock, LI&: *LI);
745 CurLoop->getParentLoop()->addBasicBlockToLoop(NewBB: VectorLoopPreheaderBlock,
746 LI&: *LI);
747 CurLoop->getParentLoop()->addChildLoop(NewChild: VectorLoop);
748 CurLoop->getParentLoop()->addBasicBlockToLoop(NewBB: VectorLoopMismatchBlock, LI&: *LI);
749 CurLoop->getParentLoop()->addBasicBlockToLoop(NewBB: LoopPreHeaderBlock, LI&: *LI);
750 CurLoop->getParentLoop()->addChildLoop(NewChild: ScalarLoop);
751 } else {
752 LI->addTopLevelLoop(New: VectorLoop);
753 LI->addTopLevelLoop(New: ScalarLoop);
754 }
755
756 // Add the new basic blocks to their associated loops.
757 VectorLoop->addBasicBlockToLoop(NewBB: VectorLoopStartBlock, LI&: *LI);
758 VectorLoop->addBasicBlockToLoop(NewBB: VectorLoopIncBlock, LI&: *LI);
759
760 ScalarLoop->addBasicBlockToLoop(NewBB: LoopStartBlock, LI&: *LI);
761 ScalarLoop->addBasicBlockToLoop(NewBB: LoopIncBlock, LI&: *LI);
762
763 // Set up some types and constants that we intend to reuse.
764 Type *I64Type = Builder.getInt64Ty();
765
766 // Check the zero-extended iteration count > 0
767 Builder.SetInsertPoint(MinItCheckBlock);
768 Value *ExtStart = Builder.CreateZExt(V: Start, DestTy: I64Type);
769 Value *ExtEnd = Builder.CreateZExt(V: MaxLen, DestTy: I64Type);
770 // This check doesn't really cost us very much.
771
772 Value *LimitCheck = Builder.CreateICmpULE(LHS: Start, RHS: MaxLen);
773 CondBrInst *MinItCheckBr =
774 CondBrInst::Create(Cond: LimitCheck, IfTrue: MemCheckBlock, IfFalse: LoopPreHeaderBlock);
775 MinItCheckBr->setMetadata(
776 KindID: LLVMContext::MD_prof,
777 Node: MDBuilder(MinItCheckBr->getContext()).createBranchWeights(TrueWeight: 99, FalseWeight: 1));
778 Builder.Insert(I: MinItCheckBr);
779
780 DTU.applyUpdates(
781 Updates: {{DominatorTree::Insert, MinItCheckBlock, MemCheckBlock},
782 {DominatorTree::Insert, MinItCheckBlock, LoopPreHeaderBlock}});
783
784 // For each of the arrays, check the start/end addresses are on the same
785 // page.
786 Builder.SetInsertPoint(MemCheckBlock);
787
788 // The early exit in the original loop means that when performing vector
789 // loads we are potentially reading ahead of the early exit. So we could
790 // fault if crossing a page boundary. Therefore, we create runtime memory
791 // checks based on the minimum page size as follows:
792 // 1. Calculate the addresses of the first memory accesses in the loop,
793 // i.e. LhsStart and RhsStart.
794 // 2. Get the last accessed addresses in the loop, i.e. LhsEnd and RhsEnd.
795 // 3. Determine which pages correspond to all the memory accesses, i.e
796 // LhsStartPage, LhsEndPage, RhsStartPage, RhsEndPage.
797 // 4. If LhsStartPage == LhsEndPage and RhsStartPage == RhsEndPage, then
798 // we know we won't cross any page boundaries in the loop so we can
799 // enter the vector loop! Otherwise we fall back on the scalar loop.
800 Value *LhsStartGEP = Builder.CreateGEP(Ty: LoadType, Ptr: PtrA, IdxList: ExtStart);
801 Value *RhsStartGEP = Builder.CreateGEP(Ty: LoadType, Ptr: PtrB, IdxList: ExtStart);
802 Value *RhsStart = Builder.CreatePtrToInt(V: RhsStartGEP, DestTy: I64Type);
803 Value *LhsStart = Builder.CreatePtrToInt(V: LhsStartGEP, DestTy: I64Type);
804 Value *LhsEndGEP = Builder.CreateGEP(Ty: LoadType, Ptr: PtrA, IdxList: ExtEnd);
805 Value *RhsEndGEP = Builder.CreateGEP(Ty: LoadType, Ptr: PtrB, IdxList: ExtEnd);
806 Value *LhsEnd = Builder.CreatePtrToInt(V: LhsEndGEP, DestTy: I64Type);
807 Value *RhsEnd = Builder.CreatePtrToInt(V: RhsEndGEP, DestTy: I64Type);
808
809 const uint64_t MinPageSize = TTI->getMinPageSize().value();
810 const uint64_t AddrShiftAmt = llvm::Log2_64(Value: MinPageSize);
811 Value *LhsStartPage = Builder.CreateLShr(LHS: LhsStart, RHS: AddrShiftAmt);
812 Value *LhsEndPage = Builder.CreateLShr(LHS: LhsEnd, RHS: AddrShiftAmt);
813 Value *RhsStartPage = Builder.CreateLShr(LHS: RhsStart, RHS: AddrShiftAmt);
814 Value *RhsEndPage = Builder.CreateLShr(LHS: RhsEnd, RHS: AddrShiftAmt);
815 Value *LhsPageCmp = Builder.CreateICmpNE(LHS: LhsStartPage, RHS: LhsEndPage);
816 Value *RhsPageCmp = Builder.CreateICmpNE(LHS: RhsStartPage, RHS: RhsEndPage);
817
818 Value *CombinedPageCmp = Builder.CreateOr(LHS: LhsPageCmp, RHS: RhsPageCmp);
819 CondBrInst *CombinedPageCmpCmpBr = CondBrInst::Create(
820 Cond: CombinedPageCmp, IfTrue: LoopPreHeaderBlock, IfFalse: VectorLoopPreheaderBlock);
821 CombinedPageCmpCmpBr->setMetadata(
822 KindID: LLVMContext::MD_prof, Node: MDBuilder(CombinedPageCmpCmpBr->getContext())
823 .createBranchWeights(TrueWeight: 10, FalseWeight: 90));
824 Builder.Insert(I: CombinedPageCmpCmpBr);
825
826 DTU.applyUpdates(
827 Updates: {{DominatorTree::Insert, MemCheckBlock, LoopPreHeaderBlock},
828 {DominatorTree::Insert, MemCheckBlock, VectorLoopPreheaderBlock}});
829
830 // Set up the vector loop preheader, i.e. calculate initial loop predicate,
831 // zero-extend MaxLen to 64-bits, determine the number of vector elements
832 // processed in each iteration, etc.
833 Builder.SetInsertPoint(VectorLoopPreheaderBlock);
834
835 // At this point we know two things must be true:
836 // 1. Start <= End
837 // 2. ExtMaxLen <= MinPageSize due to the page checks.
838 // Therefore, we know that we can use a 64-bit induction variable that
839 // starts from 0 -> ExtMaxLen and it will not overflow.
840 Value *VectorLoopRes = nullptr;
841 switch (VectorizeStyle) {
842 case LoopIdiomVectorizeStyle::Masked:
843 VectorLoopRes =
844 createMaskedFindMismatch(Builder, DTU, GEPA, GEPB, ExtStart, ExtEnd);
845 break;
846 case LoopIdiomVectorizeStyle::Predicated:
847 VectorLoopRes = createPredicatedFindMismatch(Builder, DTU, GEPA, GEPB,
848 ExtStart, ExtEnd);
849 break;
850 }
851
852 Builder.CreateBr(Dest: EndBlock);
853
854 DTU.applyUpdates(
855 Updates: {{DominatorTree::Insert, VectorLoopMismatchBlock, EndBlock}});
856
857 // Generate code for scalar loop.
858 Builder.SetInsertPoint(LoopPreHeaderBlock);
859 Builder.CreateBr(Dest: LoopStartBlock);
860
861 DTU.applyUpdates(
862 Updates: {{DominatorTree::Insert, LoopPreHeaderBlock, LoopStartBlock}});
863
864 Builder.SetInsertPoint(LoopStartBlock);
865 PHINode *IndexPhi = Builder.CreatePHI(Ty: ResType, NumReservedValues: 2, Name: "mismatch_index");
866 IndexPhi->addIncoming(V: Start, BB: LoopPreHeaderBlock);
867
868 // Otherwise compare the values
869 // Load bytes from each array and compare them.
870 Value *GepOffset = Builder.CreateZExt(V: IndexPhi, DestTy: I64Type);
871
872 Value *LhsGep =
873 Builder.CreateGEP(Ty: LoadType, Ptr: PtrA, IdxList: GepOffset, Name: "", NW: GEPA->isInBounds());
874 Value *LhsLoad = Builder.CreateLoad(Ty: LoadType, Ptr: LhsGep);
875
876 Value *RhsGep =
877 Builder.CreateGEP(Ty: LoadType, Ptr: PtrB, IdxList: GepOffset, Name: "", NW: GEPB->isInBounds());
878 Value *RhsLoad = Builder.CreateLoad(Ty: LoadType, Ptr: RhsGep);
879
880 Value *MatchCmp = Builder.CreateICmpEQ(LHS: LhsLoad, RHS: RhsLoad);
881 // If we have a mismatch then exit the loop ...
882 Builder.CreateCondBr(Cond: MatchCmp, True: LoopIncBlock, False: EndBlock);
883
884 DTU.applyUpdates(Updates: {{DominatorTree::Insert, LoopStartBlock, LoopIncBlock},
885 {DominatorTree::Insert, LoopStartBlock, EndBlock}});
886
887 // Have we reached the maximum permitted length for the loop?
888 Builder.SetInsertPoint(LoopIncBlock);
889 Value *PhiInc = Builder.CreateAdd(LHS: IndexPhi, RHS: ConstantInt::get(Ty: ResType, V: 1), Name: "",
890 /*HasNUW=*/Index->hasNoUnsignedWrap(),
891 /*HasNSW=*/Index->hasNoSignedWrap());
892 IndexPhi->addIncoming(V: PhiInc, BB: LoopIncBlock);
893 Value *IVCmp = Builder.CreateICmpEQ(LHS: PhiInc, RHS: MaxLen);
894 Builder.CreateCondBr(Cond: IVCmp, True: EndBlock, False: LoopStartBlock);
895
896 DTU.applyUpdates(Updates: {{DominatorTree::Insert, LoopIncBlock, EndBlock},
897 {DominatorTree::Insert, LoopIncBlock, LoopStartBlock}});
898
899 // In the end block we need to insert a PHI node to deal with three cases:
900 // 1. We didn't find a mismatch in the scalar loop, so we return MaxLen.
901 // 2. We exitted the scalar loop early due to a mismatch and need to return
902 // the index that we found.
903 // 3. We didn't find a mismatch in the vector loop, so we return MaxLen.
904 // 4. We exitted the vector loop early due to a mismatch and need to return
905 // the index that we found.
906 Builder.SetInsertPoint(EndBlock->getFirstInsertionPt());
907 PHINode *ResPhi = Builder.CreatePHI(Ty: ResType, NumReservedValues: 4, Name: "mismatch_result");
908 ResPhi->addIncoming(V: MaxLen, BB: LoopIncBlock);
909 ResPhi->addIncoming(V: IndexPhi, BB: LoopStartBlock);
910 ResPhi->addIncoming(V: MaxLen, BB: VectorLoopIncBlock);
911 ResPhi->addIncoming(V: VectorLoopRes, BB: VectorLoopMismatchBlock);
912
913 Value *FinalRes = Builder.CreateTrunc(V: ResPhi, DestTy: ResType);
914
915 if (VerifyLoops) {
916 ScalarLoop->verifyLoop();
917 VectorLoop->verifyLoop();
918 if (!VectorLoop->isRecursivelyLCSSAForm(DT: *DT, LI: *LI))
919 report_fatal_error(reason: "Loops must remain in LCSSA form!");
920 if (!ScalarLoop->isRecursivelyLCSSAForm(DT: *DT, LI: *LI))
921 report_fatal_error(reason: "Loops must remain in LCSSA form!");
922 }
923
924 return FinalRes;
925}
926
927void LoopIdiomVectorize::transformByteCompare(GetElementPtrInst *GEPA,
928 GetElementPtrInst *GEPB,
929 PHINode *IndPhi, Value *MaxLen,
930 Instruction *Index, Value *Start,
931 bool IncIdx, BasicBlock *FoundBB,
932 BasicBlock *EndBB) {
933
934 // Insert the byte compare code at the end of the preheader block
935 BasicBlock *Preheader = CurLoop->getLoopPreheader();
936 BasicBlock *Header = CurLoop->getHeader();
937 UncondBrInst *PHBranch = cast<UncondBrInst>(Val: Preheader->getTerminator());
938 IRBuilder<> Builder(PHBranch);
939 DomTreeUpdater DTU(DT, DomTreeUpdater::UpdateStrategy::Lazy);
940 Builder.SetCurrentDebugLocation(PHBranch->getDebugLoc());
941
942 // Increment the pointer if this was done before the loads in the loop.
943 if (IncIdx)
944 Start = Builder.CreateAdd(LHS: Start, RHS: ConstantInt::get(Ty: Start->getType(), V: 1));
945
946 Value *ByteCmpRes =
947 expandFindMismatch(Builder, DTU, GEPA, GEPB, Index, Start, MaxLen);
948
949 // Replaces uses of index & induction Phi with intrinsic (we already
950 // checked that the the first instruction of Header is the Phi above).
951 assert(IndPhi->hasOneUse() && "Index phi node has more than one use!");
952 Index->replaceAllUsesWith(V: ByteCmpRes);
953
954 // If no mismatch was found, we can jump to the end block. Create a
955 // new basic block for the compare instruction.
956 auto *CmpBB = BasicBlock::Create(Context&: Preheader->getContext(), Name: "byte.compare",
957 Parent: Preheader->getParent());
958 CmpBB->moveBefore(MovePos: EndBB);
959
960 // Replace the branch in the preheader with an always-true conditional branch.
961 // This ensures there is still a reference to the original loop.
962 Builder.CreateCondBr(Cond: Builder.getTrue(), True: CmpBB, False: Header);
963 PHBranch->eraseFromParent();
964
965 BasicBlock *MismatchEnd = cast<Instruction>(Val: ByteCmpRes)->getParent();
966 DTU.applyUpdates(Updates: {{DominatorTree::Insert, MismatchEnd, CmpBB}});
967
968 // Create the branch to either the end or found block depending on the value
969 // returned by the intrinsic.
970 Builder.SetInsertPoint(CmpBB);
971 if (FoundBB != EndBB) {
972 Value *FoundCmp = Builder.CreateICmpEQ(LHS: ByteCmpRes, RHS: MaxLen);
973 Builder.CreateCondBr(Cond: FoundCmp, True: EndBB, False: FoundBB);
974 DTU.applyUpdates(Updates: {{DominatorTree::Insert, CmpBB, FoundBB},
975 {DominatorTree::Insert, CmpBB, EndBB}});
976
977 } else {
978 Builder.CreateBr(Dest: FoundBB);
979 DTU.applyUpdates(Updates: {{DominatorTree::Insert, CmpBB, FoundBB}});
980 }
981
982 // Ensure all Phis in the successors of CmpBB have an incoming value from it.
983 fixSuccessorPhis(L: CurLoop, ScalarRes: ByteCmpRes, VectorRes: ByteCmpRes, SuccBB: EndBB, IncBB: CmpBB);
984 if (EndBB != FoundBB)
985 fixSuccessorPhis(L: CurLoop, ScalarRes: ByteCmpRes, VectorRes: ByteCmpRes, SuccBB: FoundBB, IncBB: CmpBB);
986
987 // The new CmpBB block isn't part of the loop, but will need to be added to
988 // the outer loop if there is one.
989 if (!CurLoop->isOutermost())
990 CurLoop->getParentLoop()->addBasicBlockToLoop(NewBB: CmpBB, LI&: *LI);
991
992 if (VerifyLoops && CurLoop->getParentLoop()) {
993 CurLoop->getParentLoop()->verifyLoop();
994 if (!CurLoop->getParentLoop()->isRecursivelyLCSSAForm(DT: *DT, LI: *LI))
995 report_fatal_error(reason: "Loops must remain in LCSSA form!");
996 }
997}
998
999bool LoopIdiomVectorize::recognizeFindFirstByte() {
1000 // Currently the transformation only works on scalable vector types, although
1001 // there is no fundamental reason why it cannot be made to work for fixed
1002 // vectors. We also need to know the target's minimum page size in order to
1003 // generate runtime memory checks to ensure the vector version won't fault.
1004 if (!TTI->supportsScalableVectors() || !TTI->getMinPageSize().has_value() ||
1005 DisableFindFirstByte)
1006 return false;
1007
1008 // We exclude loops with trip counts > minimum page size via runtime checks,
1009 // so make sure that the minimum page size is something sensible such that
1010 // induction variables cannot overflow.
1011 if (uint64_t(*TTI->getMinPageSize()) >
1012 (std::numeric_limits<uint64_t>::max() / 2))
1013 return false;
1014
1015 // Define some constants we need throughout.
1016 BasicBlock *Header = CurLoop->getHeader();
1017 LLVMContext &Ctx = Header->getContext();
1018
1019 // We are expecting the four blocks defined below: Header, MatchBB, InnerBB,
1020 // and OuterBB. For now, we will bail our for almost anything else. The Four
1021 // blocks contain one nested loop.
1022 if (CurLoop->getNumBackEdges() != 1 || CurLoop->getNumBlocks() != 4 ||
1023 CurLoop->getSubLoops().size() != 1)
1024 return false;
1025
1026 auto *InnerLoop = CurLoop->getSubLoops().front();
1027 Function &F = *InnerLoop->getHeader()->getParent();
1028
1029 // Bail if vectorization is disabled on inner loop.
1030 LoopVectorizeHints Hints(InnerLoop, /*InterleaveOnlyWhenForced=*/true, ORE);
1031 if (!Hints.allowVectorization(F: &F, L: InnerLoop,
1032 /*VectorizeOnlyWhenForced=*/false)) {
1033 LLVM_DEBUG(dbgs() << DEBUG_TYPE << " is disabled on inner loop "
1034 << InnerLoop->getName()
1035 << " due to vectorization hints\n");
1036 return false;
1037 }
1038
1039 PHINode *IndPhi = dyn_cast<PHINode>(Val: &Header->front());
1040 if (!IndPhi || IndPhi->getNumIncomingValues() != 2)
1041 return false;
1042
1043 // Check instruction counts.
1044 auto LoopBlocks = CurLoop->getBlocks();
1045 if (LoopBlocks[0]->size() > 3 || LoopBlocks[1]->size() > 4 ||
1046 LoopBlocks[2]->size() > 3 || LoopBlocks[3]->size() > 3)
1047 return false;
1048
1049 // Check that no instruction other than IndPhi has outside uses.
1050 for (BasicBlock *BB : LoopBlocks)
1051 for (Instruction &I : *BB)
1052 if (&I != IndPhi)
1053 for (User *U : I.users())
1054 if (!CurLoop->contains(Inst: cast<Instruction>(Val: U)))
1055 return false;
1056
1057 // Match the branch instruction in the header. We are expecting an
1058 // unconditional branch to the inner loop.
1059 //
1060 // Header:
1061 // %14 = phi ptr [ %24, %OuterBB ], [ %3, %Header.preheader ]
1062 // %15 = load i8, ptr %14, align 1
1063 // br label %MatchBB
1064 BasicBlock *MatchBB;
1065 if (!match(V: Header->getTerminator(), P: m_UnconditionalBr(Succ&: MatchBB)) ||
1066 !InnerLoop->contains(BB: MatchBB))
1067 return false;
1068
1069 // MatchBB should be the entrypoint into the inner loop containing the
1070 // comparison between a search element and a needle.
1071 //
1072 // MatchBB:
1073 // %20 = phi ptr [ %7, %Header ], [ %17, %InnerBB ]
1074 // %21 = load i8, ptr %20, align 1
1075 // %22 = icmp eq i8 %15, %21
1076 // br i1 %22, label %ExitSucc, label %InnerBB
1077 BasicBlock *ExitSucc, *InnerBB;
1078 Value *LoadSearch, *LoadNeedle;
1079 CmpPredicate MatchPred;
1080 if (!match(V: MatchBB->getTerminator(),
1081 P: m_Br(C: m_ICmp(Pred&: MatchPred, L: m_Value(V&: LoadSearch), R: m_Value(V&: LoadNeedle)),
1082 T: m_BasicBlock(V&: ExitSucc), F: m_BasicBlock(V&: InnerBB))) ||
1083 MatchPred != ICmpInst::ICMP_EQ || !InnerLoop->contains(BB: InnerBB))
1084 return false;
1085
1086 // We expect outside uses of `IndPhi' in ExitSucc (and only there).
1087 for (User *U : IndPhi->users())
1088 if (!CurLoop->contains(Inst: cast<Instruction>(Val: U))) {
1089 auto *PN = dyn_cast<PHINode>(Val: U);
1090 if (!PN || PN->getParent() != ExitSucc)
1091 return false;
1092 }
1093
1094 // Match the loads and check they are simple. The loads come from two PHIs,
1095 // each with two incoming values.
1096 PHINode *PSearch, *PNeedle;
1097 if (!match(V: LoadSearch, P: m_Load(Op: m_Phi(PN&: PSearch))) ||
1098 !match(V: LoadNeedle, P: m_Load(Op: m_Phi(PN&: PNeedle))) ||
1099 !cast<LoadInst>(Val: LoadSearch)->isSimple() ||
1100 !cast<LoadInst>(Val: LoadNeedle)->isSimple())
1101 return false;
1102
1103 // Check we are loading valid characters.
1104 Type *CharTy = LoadSearch->getType();
1105 if (!CharTy->isIntegerTy() || LoadNeedle->getType() != CharTy)
1106 return false;
1107
1108 // Pick the vectorisation factor based on CharTy, work out the cost of the
1109 // match intrinsic and decide if we should use it.
1110 // Note: For the time being we assume 128-bit vectors.
1111 unsigned VF = 128 / CharTy->getIntegerBitWidth();
1112 SmallVector<Type *> Args = {
1113 ScalableVectorType::get(ElementType: CharTy, MinNumElts: VF), FixedVectorType::get(ElementType: CharTy, NumElts: VF),
1114 ScalableVectorType::get(ElementType: Type::getInt1Ty(C&: Ctx), MinNumElts: VF)};
1115 IntrinsicCostAttributes Attrs(Intrinsic::experimental_vector_match, Args[2],
1116 Args);
1117 if (TTI->getIntrinsicInstrCost(ICA: Attrs, CostKind: TTI::TCK_SizeAndLatency) > 4)
1118 return false;
1119
1120 if (PSearch->getNumIncomingValues() != 2 ||
1121 PNeedle->getNumIncomingValues() != 2)
1122 return false;
1123
1124 // One PHI comes from the outer loop (PSearch), the other one from the inner
1125 // loop (PNeedle). PSearch effectively corresponds to IndPhi.
1126 if (InnerLoop->contains(Inst: PSearch))
1127 std::swap(a&: PSearch, b&: PNeedle);
1128 if (PSearch != &Header->front() || PNeedle != &MatchBB->front())
1129 return false;
1130
1131 // The incoming values of both PHI nodes should be a gep of 1.
1132 Value *SearchStart = PSearch->getIncomingValue(i: 0);
1133 Value *SearchIndex = PSearch->getIncomingValue(i: 1);
1134 if (CurLoop->contains(BB: PSearch->getIncomingBlock(i: 0)))
1135 std::swap(a&: SearchStart, b&: SearchIndex);
1136
1137 Value *NeedleStart = PNeedle->getIncomingValue(i: 0);
1138 Value *NeedleIndex = PNeedle->getIncomingValue(i: 1);
1139 if (InnerLoop->contains(BB: PNeedle->getIncomingBlock(i: 0)))
1140 std::swap(a&: NeedleStart, b&: NeedleIndex);
1141
1142 // Match the GEPs.
1143 if (!match(V: SearchIndex, P: m_GEP(Ops: m_Specific(V: PSearch), Ops: m_One())) ||
1144 !match(V: NeedleIndex, P: m_GEP(Ops: m_Specific(V: PNeedle), Ops: m_One())))
1145 return false;
1146
1147 // Check the GEPs result type matches `CharTy'.
1148 GetElementPtrInst *GEPSearch = cast<GetElementPtrInst>(Val: SearchIndex);
1149 GetElementPtrInst *GEPNeedle = cast<GetElementPtrInst>(Val: NeedleIndex);
1150 if (GEPSearch->getResultElementType() != CharTy ||
1151 GEPNeedle->getResultElementType() != CharTy)
1152 return false;
1153
1154 // InnerBB should increment the address of the needle pointer.
1155 //
1156 // InnerBB:
1157 // %17 = getelementptr inbounds i8, ptr %20, i64 1
1158 // %18 = icmp eq ptr %17, %10
1159 // br i1 %18, label %OuterBB, label %MatchBB
1160 BasicBlock *OuterBB;
1161 Value *NeedleEnd;
1162 if (!match(V: InnerBB->getTerminator(),
1163 P: m_Br(C: m_SpecificICmp(MatchPred: ICmpInst::ICMP_EQ, L: m_Specific(V: GEPNeedle),
1164 R: m_Value(V&: NeedleEnd)),
1165 T: m_BasicBlock(V&: OuterBB), F: m_Specific(V: MatchBB))) ||
1166 !CurLoop->contains(BB: OuterBB))
1167 return false;
1168
1169 // OuterBB should increment the address of the search element pointer.
1170 //
1171 // OuterBB:
1172 // %24 = getelementptr inbounds i8, ptr %14, i64 1
1173 // %25 = icmp eq ptr %24, %6
1174 // br i1 %25, label %ExitFail, label %Header
1175 BasicBlock *ExitFail;
1176 Value *SearchEnd;
1177 if (!match(V: OuterBB->getTerminator(),
1178 P: m_Br(C: m_SpecificICmp(MatchPred: ICmpInst::ICMP_EQ, L: m_Specific(V: GEPSearch),
1179 R: m_Value(V&: SearchEnd)),
1180 T: m_BasicBlock(V&: ExitFail), F: m_Specific(V: Header))))
1181 return false;
1182
1183 if (!CurLoop->isLoopInvariant(V: SearchStart) ||
1184 !CurLoop->isLoopInvariant(V: SearchEnd) ||
1185 !CurLoop->isLoopInvariant(V: NeedleStart) ||
1186 !CurLoop->isLoopInvariant(V: NeedleEnd))
1187 return false;
1188
1189 LLVM_DEBUG(dbgs() << "Found idiom in loop: \n" << *CurLoop << "\n\n");
1190
1191 transformFindFirstByte(IndPhi, VF, CharTy, ExitSucc, ExitFail, SearchStart,
1192 SearchEnd, NeedleStart, NeedleEnd);
1193 return true;
1194}
1195
1196Value *LoopIdiomVectorize::expandFindFirstByte(
1197 IRBuilder<> &Builder, DomTreeUpdater &DTU, unsigned VF, Type *CharTy,
1198 Value *IndPhi, BasicBlock *ExitSucc, BasicBlock *ExitFail,
1199 Value *SearchStart, Value *SearchEnd, Value *NeedleStart,
1200 Value *NeedleEnd) {
1201 // Set up some types and constants that we intend to reuse.
1202 auto *I64Ty = Builder.getInt64Ty();
1203 auto *PredVTy = ScalableVectorType::get(ElementType: Builder.getInt1Ty(), MinNumElts: VF);
1204 auto *CharVTy = ScalableVectorType::get(ElementType: CharTy, MinNumElts: VF);
1205 auto *ConstVF = ConstantInt::get(Ty: I64Ty, V: VF);
1206
1207 // Other common arguments.
1208 BasicBlock *Preheader = CurLoop->getLoopPreheader();
1209 LLVMContext &Ctx = Preheader->getContext();
1210 Value *Passthru = ConstantInt::getNullValue(Ty: CharVTy);
1211
1212 // Split block in the original loop preheader.
1213 // SPH is the new preheader to the old scalar loop.
1214 BasicBlock *SPH = SplitBlock(Old: Preheader, SplitPt: Preheader->getTerminator(), DT, LI,
1215 MSSAU: nullptr, BBName: "scalar_preheader");
1216
1217 // Create the blocks that we're going to use.
1218 //
1219 // We will have the following loops:
1220 // (O) Outer loop where we iterate over the elements of the search array.
1221 // (I) Inner loop where we iterate over the elements of the needle array.
1222 //
1223 // Overall, the blocks do the following:
1224 // (0) Check if the arrays can't cross page boundaries. If so go to (1),
1225 // otherwise fall back to the original scalar loop.
1226 // (1) Load the search array. Go to (2).
1227 // (2) (a) Load the needle array.
1228 // (b) Splat the first element to the inactive lanes.
1229 // (c) Accumulate any matches found. If we haven't reached the end of the
1230 // needle array loop back to (2), otherwise go to (3).
1231 // (3) Test if we found any match. If so go to (4), otherwise go to (5).
1232 // (4) Compute the index of the first match and exit.
1233 // (5) Check if we've reached the end of the search array. If not loop back to
1234 // (1), otherwise exit.
1235 // Blocks (0,4) are not part of any loop. Blocks (1,3,5) and (2) belong to the
1236 // outer and inner loops, respectively.
1237 BasicBlock *BB0 = BasicBlock::Create(Context&: Ctx, Name: "mem_check", Parent: SPH->getParent(), InsertBefore: SPH);
1238 BasicBlock *BB1 =
1239 BasicBlock::Create(Context&: Ctx, Name: "find_first_vec_header", Parent: SPH->getParent(), InsertBefore: SPH);
1240 BasicBlock *BB2 =
1241 BasicBlock::Create(Context&: Ctx, Name: "needle_check_vec", Parent: SPH->getParent(), InsertBefore: SPH);
1242 BasicBlock *BB3 =
1243 BasicBlock::Create(Context&: Ctx, Name: "match_check_vec", Parent: SPH->getParent(), InsertBefore: SPH);
1244 BasicBlock *BB4 =
1245 BasicBlock::Create(Context&: Ctx, Name: "calculate_match", Parent: SPH->getParent(), InsertBefore: SPH);
1246 BasicBlock *BB5 =
1247 BasicBlock::Create(Context&: Ctx, Name: "search_check_vec", Parent: SPH->getParent(), InsertBefore: SPH);
1248
1249 // Update LoopInfo with the new loops.
1250 auto OuterLoop = LI->AllocateLoop();
1251 auto InnerLoop = LI->AllocateLoop();
1252
1253 if (auto ParentLoop = CurLoop->getParentLoop()) {
1254 ParentLoop->addBasicBlockToLoop(NewBB: BB0, LI&: *LI);
1255 ParentLoop->addChildLoop(NewChild: OuterLoop);
1256 } else {
1257 LI->addTopLevelLoop(New: OuterLoop);
1258 }
1259
1260 // BB4 branches only to ExitSucc, so it belongs to the innermost enclosing
1261 // loop that contains ExitSucc. That is not necessarily CurLoop's parent: a
1262 // match can exit several levels out, or out of every loop.
1263 Loop *ExitLoop = CurLoop->getParentLoop();
1264 while (ExitLoop && !ExitLoop->contains(BB: ExitSucc))
1265 ExitLoop = ExitLoop->getParentLoop();
1266 if (ExitLoop)
1267 ExitLoop->addBasicBlockToLoop(NewBB: BB4, LI&: *LI);
1268
1269 // Add the inner loop to the outer.
1270 OuterLoop->addChildLoop(NewChild: InnerLoop);
1271
1272 // Add the new basic blocks to the corresponding loops.
1273 OuterLoop->addBasicBlockToLoop(NewBB: BB1, LI&: *LI);
1274 OuterLoop->addBasicBlockToLoop(NewBB: BB3, LI&: *LI);
1275 OuterLoop->addBasicBlockToLoop(NewBB: BB5, LI&: *LI);
1276 InnerLoop->addBasicBlockToLoop(NewBB: BB2, LI&: *LI);
1277
1278 // Update the terminator added by SplitBlock to branch to the first block.
1279 Preheader->getTerminator()->setSuccessor(Idx: 0, BB: BB0);
1280 DTU.applyUpdates(Updates: {{DominatorTree::Delete, Preheader, SPH},
1281 {DominatorTree::Insert, Preheader, BB0}});
1282
1283 // (0) Check if we could be crossing a page boundary; if so, fallback to the
1284 // old scalar loops. Also create a predicate of VF elements to be used in the
1285 // vector loops.
1286 Builder.SetInsertPoint(BB0);
1287 Value *ISearchStart =
1288 Builder.CreatePtrToInt(V: SearchStart, DestTy: I64Ty, Name: "search_start_int");
1289 Value *ISearchEnd =
1290 Builder.CreatePtrToInt(V: SearchEnd, DestTy: I64Ty, Name: "search_end_int");
1291 Value *SearchIdxInit = Constant::getNullValue(Ty: I64Ty);
1292 Value *SearchTripCount =
1293 Builder.CreateZExt(V: Builder.CreatePtrDiff(ElemTy: CharTy, LHS: SearchEnd, RHS: SearchStart,
1294 Name: "search_trip_count"),
1295 DestTy: I64Ty);
1296 Value *INeedleStart =
1297 Builder.CreatePtrToInt(V: NeedleStart, DestTy: I64Ty, Name: "needle_start_int");
1298 Value *INeedleEnd =
1299 Builder.CreatePtrToInt(V: NeedleEnd, DestTy: I64Ty, Name: "needle_end_int");
1300 Value *NeedleIdxInit = Constant::getNullValue(Ty: I64Ty);
1301 Value *NeedleTripCount =
1302 Builder.CreateZExt(V: Builder.CreatePtrDiff(ElemTy: CharTy, LHS: NeedleEnd, RHS: NeedleStart,
1303 Name: "needle_trip_count"),
1304 DestTy: I64Ty);
1305 Value *PredVF =
1306 Builder.CreateIntrinsic(ID: Intrinsic::get_active_lane_mask, OverloadTypes: {PredVTy, I64Ty},
1307 Args: {ConstantInt::get(Ty: I64Ty, V: 0), ConstVF});
1308
1309 const uint64_t MinPageSize = TTI->getMinPageSize().value();
1310 const uint64_t AddrShiftAmt = llvm::Log2_64(Value: MinPageSize);
1311 Value *SearchStartPage =
1312 Builder.CreateLShr(LHS: ISearchStart, RHS: AddrShiftAmt, Name: "search_start_page");
1313 Value *SearchEndPage =
1314 Builder.CreateLShr(LHS: ISearchEnd, RHS: AddrShiftAmt, Name: "search_end_page");
1315 Value *NeedleStartPage =
1316 Builder.CreateLShr(LHS: INeedleStart, RHS: AddrShiftAmt, Name: "needle_start_page");
1317 Value *NeedleEndPage =
1318 Builder.CreateLShr(LHS: INeedleEnd, RHS: AddrShiftAmt, Name: "needle_end_page");
1319 Value *SearchPageCmp =
1320 Builder.CreateICmpNE(LHS: SearchStartPage, RHS: SearchEndPage, Name: "search_page_cmp");
1321 Value *NeedlePageCmp =
1322 Builder.CreateICmpNE(LHS: NeedleStartPage, RHS: NeedleEndPage, Name: "needle_page_cmp");
1323
1324 Value *CombinedPageCmp =
1325 Builder.CreateOr(LHS: SearchPageCmp, RHS: NeedlePageCmp, Name: "combined_page_cmp");
1326 CondBrInst *CombinedPageBr = Builder.CreateCondBr(Cond: CombinedPageCmp, True: SPH, False: BB1);
1327 CombinedPageBr->setMetadata(KindID: LLVMContext::MD_prof,
1328 Node: MDBuilder(Ctx).createBranchWeights(TrueWeight: 10, FalseWeight: 90));
1329 DTU.applyUpdates(
1330 Updates: {{DominatorTree::Insert, BB0, SPH}, {DominatorTree::Insert, BB0, BB1}});
1331
1332 // (1) Load the search array and branch to the inner loop.
1333 Builder.SetInsertPoint(BB1);
1334 PHINode *SearchIdx = Builder.CreatePHI(Ty: I64Ty, NumReservedValues: 2, Name: "search_idx");
1335 Value *PredSearch = Builder.CreateIntrinsic(
1336 ID: Intrinsic::get_active_lane_mask, OverloadTypes: {PredVTy, I64Ty},
1337 Args: {SearchIdx, SearchTripCount}, FMFSource: nullptr, Name: "search_pred");
1338 PredSearch = Builder.CreateAnd(LHS: PredVF, RHS: PredSearch, Name: "search_masked");
1339 Value *Search = Builder.CreateGEP(Ty: CharTy, Ptr: SearchStart, IdxList: SearchIdx, Name: "psearch");
1340 Value *LoadSearch = Builder.CreateMaskedLoad(
1341 Ty: CharVTy, Ptr: Search, Alignment: Align(1), Mask: PredSearch, PassThru: Passthru, Name: "search_load_vec");
1342 Value *MatchInit = Constant::getNullValue(Ty: PredVTy);
1343 Builder.CreateBr(Dest: BB2);
1344 DTU.applyUpdates(Updates: {{DominatorTree::Insert, BB1, BB2}});
1345
1346 // (2) Inner loop.
1347 Builder.SetInsertPoint(BB2);
1348 PHINode *NeedleIdx = Builder.CreatePHI(Ty: I64Ty, NumReservedValues: 2, Name: "needle_idx");
1349 PHINode *Match = Builder.CreatePHI(Ty: PredVTy, NumReservedValues: 2, Name: "pmatch");
1350
1351 // (2.a) Load the needle array.
1352 Value *PredNeedle = Builder.CreateIntrinsic(
1353 ID: Intrinsic::get_active_lane_mask, OverloadTypes: {PredVTy, I64Ty},
1354 Args: {NeedleIdx, NeedleTripCount}, FMFSource: nullptr, Name: "needle_pred");
1355 PredNeedle = Builder.CreateAnd(LHS: PredVF, RHS: PredNeedle, Name: "needle_masked");
1356 Value *Needle = Builder.CreateGEP(Ty: CharTy, Ptr: NeedleStart, IdxList: NeedleIdx, Name: "pneedle");
1357 Value *LoadNeedle = Builder.CreateMaskedLoad(
1358 Ty: CharVTy, Ptr: Needle, Alignment: Align(1), Mask: PredNeedle, PassThru: Passthru, Name: "needle_load_vec");
1359
1360 // (2.b) Splat the first element to the inactive lanes.
1361 Value *Needle0 =
1362 Builder.CreateExtractElement(Vec: LoadNeedle, Idx: uint64_t(0), Name: "needle0");
1363 Value *Needle0Splat = Builder.CreateVectorSplat(EC: ElementCount::getScalable(MinVal: VF),
1364 V: Needle0, Name: "needle0");
1365 LoadNeedle = Builder.CreateSelect(C: PredNeedle, True: LoadNeedle, False: Needle0Splat,
1366 Name: "needle_splat");
1367 LoadNeedle = Builder.CreateExtractVector(
1368 DstType: FixedVectorType::get(ElementType: CharTy, NumElts: VF), SrcVec: LoadNeedle, Idx: uint64_t(0), Name: "needle_vec");
1369
1370 // (2.c) Accumulate matches.
1371 Value *MatchSeg = Builder.CreateIntrinsic(
1372 ID: Intrinsic::experimental_vector_match, OverloadTypes: {CharVTy, LoadNeedle->getType()},
1373 Args: {LoadSearch, LoadNeedle, PredSearch}, FMFSource: nullptr, Name: "match_segment");
1374 Value *MatchAcc = Builder.CreateOr(LHS: Match, RHS: MatchSeg, Name: "match_accumulator");
1375 Value *NextNeedleIdx =
1376 Builder.CreateAdd(LHS: NeedleIdx, RHS: ConstVF, Name: "needle_idx_next");
1377 Builder.CreateCondBr(Cond: Builder.CreateICmpULT(LHS: NextNeedleIdx, RHS: NeedleTripCount),
1378 True: BB2, False: BB3);
1379 DTU.applyUpdates(
1380 Updates: {{DominatorTree::Insert, BB2, BB2}, {DominatorTree::Insert, BB2, BB3}});
1381
1382 // (3) Check if we found a match.
1383 Builder.SetInsertPoint(BB3);
1384 PHINode *MatchPredAccLCSSA = Builder.CreatePHI(Ty: PredVTy, NumReservedValues: 1, Name: "match_pred");
1385 Value *IfAnyMatch = Builder.CreateOrReduce(Src: MatchPredAccLCSSA);
1386 Builder.CreateCondBr(Cond: IfAnyMatch, True: BB4, False: BB5);
1387 DTU.applyUpdates(
1388 Updates: {{DominatorTree::Insert, BB3, BB4}, {DominatorTree::Insert, BB3, BB5}});
1389
1390 // (4) We found a match. Compute the index of its location and exit.
1391 Builder.SetInsertPoint(BB4);
1392 PHINode *MatchLCSSA =
1393 Builder.CreatePHI(Ty: SearchStart->getType(), NumReservedValues: 1, Name: "match_start");
1394 PHINode *MatchPredLCSSA = Builder.CreatePHI(Ty: PredVTy, NumReservedValues: 1, Name: "match_vec");
1395 Value *MatchCnt = Builder.CreateIntrinsic(
1396 ID: Intrinsic::experimental_cttz_elts, OverloadTypes: {I64Ty, PredVTy},
1397 Args: {MatchPredLCSSA, /*ZeroIsPoison=*/Builder.getInt1(V: true)}, FMFSource: nullptr,
1398 Name: "match_idx");
1399 Value *MatchVal =
1400 Builder.CreateGEP(Ty: CharTy, Ptr: MatchLCSSA, IdxList: MatchCnt, Name: "match_res");
1401 Builder.CreateBr(Dest: ExitSucc);
1402 DTU.applyUpdates(Updates: {{DominatorTree::Insert, BB4, ExitSucc}});
1403
1404 // (5) Check if we've reached the end of the search array.
1405 Builder.SetInsertPoint(BB5);
1406 Value *NextSearchIdx =
1407 Builder.CreateAdd(LHS: SearchIdx, RHS: ConstVF, Name: "search_idx_next");
1408 Builder.CreateCondBr(Cond: Builder.CreateICmpULT(LHS: NextSearchIdx, RHS: SearchTripCount),
1409 True: BB1, False: ExitFail);
1410 DTU.applyUpdates(Updates: {{DominatorTree::Insert, BB5, BB1},
1411 {DominatorTree::Insert, BB5, ExitFail}});
1412
1413 // Set up the PHI nodes.
1414 SearchIdx->addIncoming(V: SearchIdxInit, BB: BB0);
1415 SearchIdx->addIncoming(V: NextSearchIdx, BB: BB5);
1416 NeedleIdx->addIncoming(V: NeedleIdxInit, BB: BB1);
1417 NeedleIdx->addIncoming(V: NextNeedleIdx, BB: BB2);
1418 Match->addIncoming(V: MatchInit, BB: BB1);
1419 Match->addIncoming(V: MatchAcc, BB: BB2);
1420 // These are needed to retain LCSSA form.
1421 MatchPredAccLCSSA->addIncoming(V: MatchAcc, BB: BB2);
1422 MatchLCSSA->addIncoming(V: Search, BB: BB3);
1423 MatchPredLCSSA->addIncoming(V: MatchPredAccLCSSA, BB: BB3);
1424
1425 // Ensure all Phis in the successors of BB4/BB5 have an incoming value from
1426 // them.
1427 fixSuccessorPhis(L: CurLoop, ScalarRes: IndPhi, VectorRes: MatchVal, SuccBB: ExitSucc, IncBB: BB4);
1428 if (ExitSucc != ExitFail)
1429 fixSuccessorPhis(L: CurLoop, ScalarRes: IndPhi, VectorRes: MatchVal, SuccBB: ExitFail, IncBB: BB5);
1430
1431 if (VerifyLoops) {
1432 OuterLoop->verifyLoop();
1433 InnerLoop->verifyLoop();
1434 if (!OuterLoop->isRecursivelyLCSSAForm(DT: *DT, LI: *LI))
1435 report_fatal_error(reason: "Loops must remain in LCSSA form!");
1436 }
1437
1438 return MatchVal;
1439}
1440
1441void LoopIdiomVectorize::transformFindFirstByte(
1442 PHINode *IndPhi, unsigned VF, Type *CharTy, BasicBlock *ExitSucc,
1443 BasicBlock *ExitFail, Value *SearchStart, Value *SearchEnd,
1444 Value *NeedleStart, Value *NeedleEnd) {
1445 // Insert the find first byte code at the end of the preheader block.
1446 BasicBlock *Preheader = CurLoop->getLoopPreheader();
1447 UncondBrInst *PHBranch = cast<UncondBrInst>(Val: Preheader->getTerminator());
1448 IRBuilder<> Builder(PHBranch);
1449 DomTreeUpdater DTU(DT, DomTreeUpdater::UpdateStrategy::Lazy);
1450 Builder.SetCurrentDebugLocation(PHBranch->getDebugLoc());
1451
1452 expandFindFirstByte(Builder, DTU, VF, CharTy, IndPhi, ExitSucc, ExitFail,
1453 SearchStart, SearchEnd, NeedleStart, NeedleEnd);
1454
1455 if (VerifyLoops && CurLoop->getParentLoop()) {
1456 CurLoop->getParentLoop()->verifyLoop();
1457 if (!CurLoop->getParentLoop()->isRecursivelyLCSSAForm(DT: *DT, LI: *LI))
1458 report_fatal_error(reason: "Loops must remain in LCSSA form!");
1459 }
1460}
1461