LLVM 24.0.0git
VectorCombine.cpp
Go to the documentation of this file.
1//===------- VectorCombine.cpp - Optimize partial vector operations -------===//
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 optimizes scalar/vector interactions using target cost models. The
10// transforms implemented here may not fit in traditional loop-based or SLP
11// vectorization passes.
12//
13//===----------------------------------------------------------------------===//
14
16#include "llvm/ADT/DenseMap.h"
17#include "llvm/ADT/STLExtras.h"
18#include "llvm/ADT/ScopeExit.h"
21#include "llvm/ADT/Statistic.h"
26#include "llvm/Analysis/Loads.h"
31#include "llvm/IR/Dominators.h"
32#include "llvm/IR/Function.h"
33#include "llvm/IR/IRBuilder.h"
42#include <numeric>
43#include <optional>
44#include <queue>
45#include <set>
46
47#define DEBUG_TYPE "vector-combine"
49
50using namespace llvm;
51using namespace llvm::PatternMatch;
52
53STATISTIC(NumVecLoad, "Number of vector loads formed");
54STATISTIC(NumVecCmp, "Number of vector compares formed");
55STATISTIC(NumVecBO, "Number of vector binops formed");
56STATISTIC(NumVecCmpBO, "Number of vector compare + binop formed");
57STATISTIC(NumShufOfBitcast, "Number of shuffles moved after bitcast");
58STATISTIC(NumScalarOps, "Number of scalar unary + binary ops formed");
59STATISTIC(NumScalarCmp, "Number of scalar compares formed");
60STATISTIC(NumScalarIntrinsic, "Number of scalar intrinsic calls formed");
61
63 "disable-vector-combine", cl::init(false), cl::Hidden,
64 cl::desc("Disable all vector combine transforms"));
65
67 "disable-binop-extract-shuffle", cl::init(false), cl::Hidden,
68 cl::desc("Disable binop extract to shuffle transforms"));
69
71 "vector-combine-max-scan-instrs", cl::init(30), cl::Hidden,
72 cl::desc("Max number of instructions to scan for vector combining."));
73
74static const unsigned InvalidIndex = std::numeric_limits<unsigned>::max();
75
76namespace {
77class VectorCombine {
78public:
79 VectorCombine(Function &F, const TargetTransformInfo &TTI,
82 bool TryEarlyFoldsOnly)
83 : F(F), Builder(F.getContext(), InstSimplifyFolder(*DL)), TTI(TTI),
84 DT(DT), AA(AA), DL(DL), CostKind(CostKind),
85 SQ(*DL, /*TLI=*/nullptr, &DT, &AC),
86 TryEarlyFoldsOnly(TryEarlyFoldsOnly) {}
87
88 bool run();
89
90private:
91 Function &F;
93 const TargetTransformInfo &TTI;
94 const DominatorTree &DT;
95 AAResults &AA;
96 const DataLayout *DL;
97 TTI::TargetCostKind CostKind;
98 const SimplifyQuery SQ;
99
100 /// If true, only perform beneficial early IR transforms. Do not introduce new
101 /// vector operations.
102 bool TryEarlyFoldsOnly;
103
104 InstructionWorklist Worklist;
105
106 /// Next instruction to iterate. It will be updated when it is erased by
107 /// RecursivelyDeleteTriviallyDeadInstructions.
108 Instruction *NextInst;
109
110 // TODO: Direct calls from the top-level "run" loop use a plain "Instruction"
111 // parameter. That should be updated to specific sub-classes because the
112 // run loop was changed to dispatch on opcode.
113 bool vectorizeLoadInsert(Instruction &I);
114 bool widenSubvectorLoad(Instruction &I);
115 ExtractElementInst *getShuffleExtract(ExtractElementInst *Ext0,
116 ExtractElementInst *Ext1,
117 unsigned PreferredExtractIndex) const;
118 bool isExtractExtractCheap(ExtractElementInst *Ext0, ExtractElementInst *Ext1,
119 const Instruction &I,
120 ExtractElementInst *&ConvertToShuffle,
121 unsigned PreferredExtractIndex);
122 Value *foldExtExtCmp(Value *V0, Value *V1, Value *ExtIndex, Instruction &I);
123 Value *foldExtExtBinop(Value *V0, Value *V1, Value *ExtIndex, Instruction &I);
124 bool foldExtractExtract(Instruction &I);
125 bool foldInsExtFNeg(Instruction &I);
126 bool foldInsExtBinop(Instruction &I);
127 bool foldInsExtVectorToShuffle(Instruction &I);
128 bool foldBitOpOfCastops(Instruction &I);
129 bool foldBitOpOfCastConstant(Instruction &I);
130 bool foldBitcastShuffle(Instruction &I);
131 bool scalarizeOpOrCmp(Instruction &I);
132 bool foldExtractedCmps(Instruction &I);
133 bool foldSelectsFromBitcast(Instruction &I);
134 bool foldBinopOfReductions(Instruction &I);
135 bool foldInsertElementsToStores(Instruction &I);
136 bool scalarizeLoad(Instruction &I);
137 bool scalarizeLoadExtract(LoadInst *LI, VectorType *VecTy, Value *Ptr);
138 bool scalarizeLoadBitcast(LoadInst *LI, VectorType *VecTy, Value *Ptr);
139 bool scalarizeExtExtract(Instruction &I);
140 bool foldConcatOfBoolMasks(Instruction &I);
141 bool foldPermuteOfBinops(Instruction &I);
142 bool foldShuffleOfBinops(Instruction &I);
143 bool foldShuffleOfSelects(Instruction &I);
144 bool foldShuffleOfCastops(Instruction &I);
145 bool foldShuffleOfShuffles(Instruction &I);
146 bool foldPermuteOfIntrinsic(Instruction &I);
147 bool foldShufflesOfLengthChangingShuffles(Instruction &I);
148 bool foldShuffleOfIntrinsics(Instruction &I);
149 bool foldShuffleToIdentity(Instruction &I);
150 bool foldShuffleFromReductions(Instruction &I);
151 bool foldShuffleChainsToReduce(Instruction &I);
152 bool foldCastFromReductions(Instruction &I);
153 bool foldSignBitReductionCmp(Instruction &I);
154 bool foldReductionZeroTest(Instruction &I);
155 bool foldICmpEqZeroVectorReduce(Instruction &I);
156 bool foldEquivalentReductionCmp(Instruction &I);
157 bool foldReduceAddCmpZero(Instruction &I);
158 bool foldSelectShuffle(Instruction &I, bool FromReduction = false);
159 bool foldInterleaveIntrinsics(Instruction &I);
160 bool foldDeinterleaveIntrinsics(Instruction &I);
161 bool foldBitcastOfVPLoad(Instruction &I);
162 bool foldBitOrderReverseAndSwap(Instruction &I);
163 bool shrinkType(Instruction &I);
164 bool shrinkLoadForShuffles(Instruction &I);
165 bool shrinkPhiOfShuffles(Instruction &I);
166 bool foldDeinterleaveInterleavePair(Instruction &I);
167
168 void replaceValue(Instruction &Old, Value &New, bool Erase = true) {
169 LLVM_DEBUG(dbgs() << "VC: Replacing: " << Old << '\n');
170 LLVM_DEBUG(dbgs() << " With: " << New << '\n');
171 Old.replaceAllUsesWith(&New);
172 if (auto *NewI = dyn_cast<Instruction>(&New)) {
173 New.takeName(&Old);
174 Worklist.pushUsersToWorkList(*NewI);
175 Worklist.pushValue(NewI);
176 }
177 if (Erase && isInstructionTriviallyDead(&Old)) {
178 eraseInstruction(Old);
179 } else {
180 Worklist.push(&Old);
181 }
182 }
183
184 void eraseInstruction(Instruction &I) {
185 LLVM_DEBUG(dbgs() << "VC: Erasing: " << I << '\n');
186 SmallVector<Value *> Ops(I.operands());
187 Worklist.remove(&I);
188 I.eraseFromParent();
189
190 // Push remaining users of the operands and then the operand itself - allows
191 // further folds that were hindered by OneUse limits.
192 SmallPtrSet<Value *, 4> Visited;
193 for (Value *Op : Ops) {
194 if (!Visited.contains(Op)) {
195 if (auto *OpI = dyn_cast<Instruction>(Op)) {
197 OpI, nullptr, nullptr, [&](Value *V) {
198 if (auto *I = dyn_cast<Instruction>(V)) {
199 LLVM_DEBUG(dbgs() << "VC: Erased: " << *I << '\n');
200 Worklist.remove(I);
201 if (I == NextInst)
202 NextInst = NextInst->getNextNode();
203 Visited.insert(I);
204 }
205 }))
206 continue;
207 Worklist.pushUsersToWorkList(*OpI);
208 Worklist.pushValue(OpI);
209 }
210 }
211 }
212 }
213};
214} // namespace
215
216/// Return the source operand of a potentially bitcasted value. If there is no
217/// bitcast, return the input value itself.
219 while (auto *BitCast = dyn_cast<BitCastInst>(V))
220 V = BitCast->getOperand(0);
221 return V;
222}
223
224/// Helper to peek through bitcasts to the same value.
225static bool isEquivBitcast(Value *X, Value *Y) {
226 return X->getType() == Y->getType() &&
228}
229
231 // Do not widen load if atomic/volatile or under asan/hwasan/memtag/tsan.
232 // The widened load may load data from dirty regions or create data races
233 // non-existent in the source.
234 if (!Load || !Load->isSimple() || !Load->hasOneUse() ||
235 Load->getFunction()->hasFnAttribute(Attribute::SanitizeMemTag) ||
237 return false;
238
239 // We are potentially transforming byte-sized (8-bit) memory accesses, so make
240 // sure we have all of our type-based constraints in place for this target.
241 Type *ScalarTy = Load->getType()->getScalarType();
242 uint64_t ScalarSize = ScalarTy->getPrimitiveSizeInBits();
243 unsigned MinVectorSize = TTI.getMinVectorRegisterBitWidth();
244 if (!ScalarSize || !MinVectorSize || MinVectorSize % ScalarSize != 0 ||
245 ScalarSize % 8 != 0)
246 return false;
247
248 return true;
249}
250
251bool VectorCombine::vectorizeLoadInsert(Instruction &I) {
252 // Match insert into fixed vector of scalar value.
253 // TODO: Handle non-zero insert index.
254 Value *Scalar;
255 if (!match(&I,
257 return false;
258
259 // Optionally match an extract from another vector.
260 Value *X;
261 bool HasExtract = match(Scalar, m_ExtractElt(m_Value(X), m_ZeroInt()));
262 if (!HasExtract)
263 X = Scalar;
264
265 auto *Load = dyn_cast<LoadInst>(X);
266 if (!canWidenLoad(Load, TTI))
267 return false;
268
269 Type *ScalarTy = Scalar->getType();
270 uint64_t ScalarSize = ScalarTy->getPrimitiveSizeInBits();
271 unsigned MinVectorSize = TTI.getMinVectorRegisterBitWidth();
272
273 // Check safety of replacing the scalar load with a larger vector load.
274 // We use minimal alignment (maximum flexibility) because we only care about
275 // the dereferenceable region. When calculating cost and creating a new op,
276 // we may use a larger value based on alignment attributes.
277 Value *SrcPtr = Load->getPointerOperand()->stripPointerCasts();
278 assert(isa<PointerType>(SrcPtr->getType()) && "Expected a pointer type");
279
280 unsigned MinVecNumElts = MinVectorSize / ScalarSize;
281 auto *MinVecTy = VectorType::get(ScalarTy, MinVecNumElts, false);
282 unsigned OffsetEltIndex = 0;
283 Align Alignment = Load->getAlign();
284 if (!isSafeToLoadUnconditionally(SrcPtr, MinVecTy, Align(1), *DL, Load, SQ.AC,
285 SQ.DT)) {
286 // It is not safe to load directly from the pointer, but we can still peek
287 // through gep offsets and check if it safe to load from a base address with
288 // updated alignment. If it is, we can shuffle the element(s) into place
289 // after loading.
290 unsigned OffsetBitWidth = DL->getIndexTypeSizeInBits(SrcPtr->getType());
291 APInt Offset(OffsetBitWidth, 0);
293
294 // We want to shuffle the result down from a high element of a vector, so
295 // the offset must be positive.
296 if (Offset.isNegative())
297 return false;
298
299 // The offset must be a multiple of the scalar element to shuffle cleanly
300 // in the element's size.
301 uint64_t ScalarSizeInBytes = ScalarSize / 8;
302 if (Offset.urem(ScalarSizeInBytes) != 0)
303 return false;
304
305 // If we load MinVecNumElts, will our target element still be loaded?
306 APInt OffsetEltIndexAP = Offset.udiv(ScalarSizeInBytes);
307 if (OffsetEltIndexAP.uge(MinVecNumElts))
308 return false;
309 OffsetEltIndex = OffsetEltIndexAP.getZExtValue();
310
311 if (!isSafeToLoadUnconditionally(SrcPtr, MinVecTy, Align(1), *DL, Load,
312 SQ.AC, SQ.DT))
313 return false;
314
315 // Update alignment with offset value. Note that the offset could be negated
316 // to more accurately represent "(new) SrcPtr - Offset = (old) SrcPtr", but
317 // negation does not change the result of the alignment calculation.
318 Alignment = commonAlignment(Alignment, Offset.getZExtValue());
319 }
320
321 // Original pattern: insertelt undef, load [free casts of] PtrOp, 0
322 // Use the greater of the alignment on the load or its source pointer.
323 Alignment = std::max(SrcPtr->getPointerAlignment(*DL), Alignment);
324 Type *LoadTy = Load->getType();
325 unsigned AS = Load->getPointerAddressSpace();
326 InstructionCost OldCost =
327 TTI.getMemoryOpCost(Instruction::Load, LoadTy, Alignment, AS, CostKind);
328 APInt DemandedElts = APInt::getOneBitSet(MinVecNumElts, 0);
329 OldCost +=
330 TTI.getScalarizationOverhead(MinVecTy, DemandedElts,
331 /* Insert */ true, HasExtract, CostKind);
332
333 // New pattern: load VecPtr
334 InstructionCost NewCost =
335 TTI.getMemoryOpCost(Instruction::Load, MinVecTy, Alignment, AS, CostKind);
336 // Optionally, we are shuffling the loaded vector element(s) into place.
337 // For the mask set everything but element 0 to undef to prevent poison from
338 // propagating from the extra loaded memory. This will also optionally
339 // shrink/grow the vector from the loaded size to the output size.
340 // We assume this operation has no cost in codegen if there was no offset.
341 // Note that we could use freeze to avoid poison problems, but then we might
342 // still need a shuffle to change the vector size.
343 auto *Ty = cast<FixedVectorType>(I.getType());
344 unsigned OutputNumElts = Ty->getNumElements();
345 SmallVector<int, 16> Mask(OutputNumElts, PoisonMaskElem);
346 assert(OffsetEltIndex < MinVecNumElts && "Address offset too big");
347 Mask[0] = OffsetEltIndex;
348 if (OffsetEltIndex)
349 NewCost += TTI.getShuffleCost(TTI::SK_PermuteSingleSrc, Ty, MinVecTy, Mask,
350 CostKind);
351
352 // We can aggressively convert to the vector form because the backend can
353 // invert this transform if it does not result in a performance win.
354 if (OldCost < NewCost || !NewCost.isValid())
355 return false;
356
357 // It is safe and potentially profitable to load a vector directly:
358 // inselt undef, load Scalar, 0 --> load VecPtr
359 IRBuilder<> Builder(Load);
360 Value *CastedPtr =
361 Builder.CreatePointerBitCastOrAddrSpaceCast(SrcPtr, Builder.getPtrTy(AS));
362 Value *VecLd = Builder.CreateAlignedLoad(MinVecTy, CastedPtr, Alignment);
363 VecLd = Builder.CreateShuffleVector(VecLd, Mask);
364
365 replaceValue(I, *VecLd);
366 ++NumVecLoad;
367 return true;
368}
369
370/// If we are loading a vector and then inserting it into a larger vector with
371/// undefined elements, try to load the larger vector and eliminate the insert.
372/// This removes a shuffle in IR and may allow combining of other loaded values.
373bool VectorCombine::widenSubvectorLoad(Instruction &I) {
374 // Match subvector insert of fixed vector.
375 auto *Shuf = cast<ShuffleVectorInst>(&I);
376 if (!Shuf->isIdentityWithPadding())
377 return false;
378
379 // Allow a non-canonical shuffle mask that is choosing elements from op1.
380 unsigned NumOpElts =
381 cast<FixedVectorType>(Shuf->getOperand(0)->getType())->getNumElements();
382 unsigned OpIndex = any_of(Shuf->getShuffleMask(), [&NumOpElts](int M) {
383 return M >= (int)(NumOpElts);
384 });
385
386 auto *Load = dyn_cast<LoadInst>(Shuf->getOperand(OpIndex));
387 if (!canWidenLoad(Load, TTI))
388 return false;
389
390 // We use minimal alignment (maximum flexibility) because we only care about
391 // the dereferenceable region. When calculating cost and creating a new op,
392 // we may use a larger value based on alignment attributes.
393 auto *Ty = cast<FixedVectorType>(I.getType());
394 Value *SrcPtr = Load->getPointerOperand()->stripPointerCasts();
395 assert(isa<PointerType>(SrcPtr->getType()) && "Expected a pointer type");
396 Align Alignment = Load->getAlign();
397 if (!isSafeToLoadUnconditionally(SrcPtr, Ty, Align(1), *DL, Load, SQ.AC,
398 SQ.DT))
399 return false;
400
401 Alignment = std::max(SrcPtr->getPointerAlignment(*DL), Alignment);
402 Type *LoadTy = Load->getType();
403 unsigned AS = Load->getPointerAddressSpace();
404
405 // Original pattern: insert_subvector (load PtrOp)
406 // This conservatively assumes that the cost of a subvector insert into an
407 // undef value is 0. We could add that cost if the cost model accurately
408 // reflects the real cost of that operation.
409 InstructionCost OldCost =
410 TTI.getMemoryOpCost(Instruction::Load, LoadTy, Alignment, AS, CostKind);
411
412 // New pattern: load PtrOp
413 InstructionCost NewCost =
414 TTI.getMemoryOpCost(Instruction::Load, Ty, Alignment, AS, CostKind);
415
416 // We can aggressively convert to the vector form because the backend can
417 // invert this transform if it does not result in a performance win.
418 if (OldCost < NewCost || !NewCost.isValid())
419 return false;
420
421 IRBuilder<> Builder(Load);
422 Value *CastedPtr =
423 Builder.CreatePointerBitCastOrAddrSpaceCast(SrcPtr, Builder.getPtrTy(AS));
424 Value *VecLd = Builder.CreateAlignedLoad(Ty, CastedPtr, Alignment);
425 replaceValue(I, *VecLd);
426 ++NumVecLoad;
427 return true;
428}
429
430/// Determine which, if any, of the inputs should be replaced by a shuffle
431/// followed by extract from a different index.
432ExtractElementInst *VectorCombine::getShuffleExtract(
433 ExtractElementInst *Ext0, ExtractElementInst *Ext1,
434 unsigned PreferredExtractIndex = InvalidIndex) const {
435 auto *Index0C = dyn_cast<ConstantInt>(Ext0->getIndexOperand());
436 auto *Index1C = dyn_cast<ConstantInt>(Ext1->getIndexOperand());
437 assert(Index0C && Index1C && "Expected constant extract indexes");
438
439 unsigned Index0 = Index0C->getZExtValue();
440 unsigned Index1 = Index1C->getZExtValue();
441
442 // If the extract indexes are identical, no shuffle is needed.
443 if (Index0 == Index1)
444 return nullptr;
445
446 Type *VecTy = Ext0->getVectorOperand()->getType();
447 assert(VecTy == Ext1->getVectorOperand()->getType() && "Need matching types");
448 InstructionCost Cost0 =
449 TTI.getVectorInstrCost(*Ext0, VecTy, CostKind, Index0);
450 InstructionCost Cost1 =
451 TTI.getVectorInstrCost(*Ext1, VecTy, CostKind, Index1);
452
453 // If both costs are invalid no shuffle is needed
454 if (!Cost0.isValid() && !Cost1.isValid())
455 return nullptr;
456
457 // We are extracting from 2 different indexes, so one operand must be shuffled
458 // before performing a vector operation and/or extract. The more expensive
459 // extract will be replaced by a shuffle.
460 if (Cost0 > Cost1)
461 return Ext0;
462 if (Cost1 > Cost0)
463 return Ext1;
464
465 // If the costs are equal and there is a preferred extract index, shuffle the
466 // opposite operand.
467 if (PreferredExtractIndex == Index0)
468 return Ext1;
469 if (PreferredExtractIndex == Index1)
470 return Ext0;
471
472 // Otherwise, replace the extract with the higher index.
473 return Index0 > Index1 ? Ext0 : Ext1;
474}
475
476/// Compare the relative costs of 2 extracts followed by scalar operation vs.
477/// vector operation(s) followed by extract. Return true if the existing
478/// instructions are cheaper than a vector alternative. Otherwise, return false
479/// and if one of the extracts should be transformed to a shufflevector, set
480/// \p ConvertToShuffle to that extract instruction.
481bool VectorCombine::isExtractExtractCheap(ExtractElementInst *Ext0,
482 ExtractElementInst *Ext1,
483 const Instruction &I,
484 ExtractElementInst *&ConvertToShuffle,
485 unsigned PreferredExtractIndex) {
486 auto *Ext0IndexC = dyn_cast<ConstantInt>(Ext0->getIndexOperand());
487 auto *Ext1IndexC = dyn_cast<ConstantInt>(Ext1->getIndexOperand());
488 assert(Ext0IndexC && Ext1IndexC && "Expected constant extract indexes");
489
490 unsigned Opcode = I.getOpcode();
491 Value *Ext0Src = Ext0->getVectorOperand();
492 Value *Ext1Src = Ext1->getVectorOperand();
493 Type *ScalarTy = Ext0->getType();
494 auto *VecTy = cast<VectorType>(Ext0Src->getType());
495 InstructionCost ScalarOpCost, VectorOpCost;
496
497 // Get cost estimates for scalar and vector versions of the operation.
498 bool IsBinOp = Instruction::isBinaryOp(Opcode);
499 if (IsBinOp) {
500 ScalarOpCost = TTI.getArithmeticInstrCost(Opcode, ScalarTy, CostKind);
501 VectorOpCost = TTI.getArithmeticInstrCost(Opcode, VecTy, CostKind);
502 } else {
503 assert((Opcode == Instruction::ICmp || Opcode == Instruction::FCmp) &&
504 "Expected a compare");
505 CmpInst::Predicate Pred = cast<CmpInst>(I).getPredicate();
506 ScalarOpCost = TTI.getCmpSelInstrCost(
507 Opcode, ScalarTy, CmpInst::makeCmpResultType(ScalarTy), Pred, CostKind);
508 VectorOpCost = TTI.getCmpSelInstrCost(
509 Opcode, VecTy, CmpInst::makeCmpResultType(VecTy), Pred, CostKind);
510 }
511
512 // Get cost estimates for the extract elements. These costs will factor into
513 // both sequences.
514 unsigned Ext0Index = Ext0IndexC->getZExtValue();
515 unsigned Ext1Index = Ext1IndexC->getZExtValue();
516
517 InstructionCost Extract0Cost =
518 TTI.getVectorInstrCost(*Ext0, VecTy, CostKind, Ext0Index);
519 InstructionCost Extract1Cost =
520 TTI.getVectorInstrCost(*Ext1, VecTy, CostKind, Ext1Index);
521
522 // A more expensive extract will always be replaced by a splat shuffle.
523 // For example, if Ext0 is more expensive:
524 // opcode (extelt V0, Ext0), (ext V1, Ext1) -->
525 // extelt (opcode (splat V0, Ext0), V1), Ext1
526 // TODO: Evaluate whether that always results in lowest cost. Alternatively,
527 // check the cost of creating a broadcast shuffle and shuffling both
528 // operands to element 0.
529 unsigned BestExtIndex = Extract0Cost > Extract1Cost ? Ext0Index : Ext1Index;
530 unsigned BestInsIndex = Extract0Cost > Extract1Cost ? Ext1Index : Ext0Index;
531 InstructionCost CheapExtractCost = std::min(Extract0Cost, Extract1Cost);
532
533 // Extra uses of the extracts mean that we include those costs in the
534 // vector total because those instructions will not be eliminated.
535 InstructionCost OldCost, NewCost;
536 if (Ext0Src == Ext1Src && Ext0Index == Ext1Index) {
537 // Handle a special case. If the 2 extracts are identical, adjust the
538 // formulas to account for that. The extra use charge allows for either the
539 // CSE'd pattern or an unoptimized form with identical values:
540 // opcode (extelt V, C), (extelt V, C) --> extelt (opcode V, V), C
541 bool HasUseTax = Ext0 == Ext1 ? !Ext0->hasNUses(2)
542 : !Ext0->hasOneUse() || !Ext1->hasOneUse();
543 OldCost = CheapExtractCost + ScalarOpCost;
544 NewCost = VectorOpCost + CheapExtractCost + HasUseTax * CheapExtractCost;
545 } else {
546 // Handle the general case. Each extract is actually a different value:
547 // opcode (extelt V0, C0), (extelt V1, C1) --> extelt (opcode V0, V1), C
548 OldCost = Extract0Cost + Extract1Cost + ScalarOpCost;
549 NewCost = VectorOpCost + CheapExtractCost +
550 !Ext0->hasOneUse() * Extract0Cost +
551 !Ext1->hasOneUse() * Extract1Cost;
552 }
553
554 ConvertToShuffle = getShuffleExtract(Ext0, Ext1, PreferredExtractIndex);
555 if (ConvertToShuffle) {
556 if (IsBinOp && DisableBinopExtractShuffle)
557 return true;
558
559 // If we are extracting from 2 different indexes, then one operand must be
560 // shuffled before performing the vector operation. The shuffle mask is
561 // poison except for 1 lane that is being translated to the remaining
562 // extraction lane. Therefore, it is a splat shuffle. Ex:
563 // ShufMask = { poison, poison, 0, poison }
564 // TODO: The cost model has an option for a "broadcast" shuffle
565 // (splat-from-element-0), but no option for a more general splat.
566 if (auto *FixedVecTy = dyn_cast<FixedVectorType>(VecTy)) {
567 SmallVector<int> ShuffleMask(FixedVecTy->getNumElements(),
569 ShuffleMask[BestInsIndex] = BestExtIndex;
571 VecTy, VecTy, ShuffleMask, CostKind, 0,
572 nullptr, {ConvertToShuffle});
573 } else {
575 VecTy, VecTy, {}, CostKind, 0, nullptr,
576 {ConvertToShuffle});
577 }
578 }
579
580 LLVM_DEBUG(dbgs() << "Found a binop of extractions: " << I << "\n OldCost: "
581 << OldCost << " vs NewCost: " << NewCost << "\n");
582
583 // Aggressively form a vector op if the cost is equal because the transform
584 // may enable further optimization.
585 // Codegen can reverse this transform (scalarize) if it was not profitable.
586 return OldCost < NewCost;
587}
588
589/// Create a shuffle that translates (shifts) 1 element from the input vector
590/// to a new element location.
591static Value *createShiftShuffle(Value *Vec, unsigned OldIndex,
592 unsigned NewIndex, IRBuilderBase &Builder) {
593 // The shuffle mask is poison except for 1 lane that is being translated
594 // to the new element index. Example for OldIndex == 2 and NewIndex == 0:
595 // ShufMask = { 2, poison, poison, poison }
596 auto *VecTy = cast<FixedVectorType>(Vec->getType());
597 SmallVector<int, 32> ShufMask(VecTy->getNumElements(), PoisonMaskElem);
598 ShufMask[NewIndex] = OldIndex;
599 return Builder.CreateShuffleVector(Vec, ShufMask, "shift");
600}
601
602/// Given an extract element instruction with constant index operand, shuffle
603/// the source vector (shift the scalar element) to a NewIndex for extraction.
604/// Return null if the input can be constant folded, so that we are not creating
605/// unnecessary instructions.
606static Value *translateExtract(ExtractElementInst *ExtElt, unsigned NewIndex,
607 IRBuilderBase &Builder) {
608 // Shufflevectors can only be created for fixed-width vectors.
609 Value *X = ExtElt->getVectorOperand();
610 if (!isa<FixedVectorType>(X->getType()))
611 return nullptr;
612
613 // If the extract can be constant-folded, this code is unsimplified. Defer
614 // to other passes to handle that.
615 Value *C = ExtElt->getIndexOperand();
616 assert(isa<ConstantInt>(C) && "Expected a constant index operand");
617 if (isa<Constant>(X))
618 return nullptr;
619
620 Value *Shuf = createShiftShuffle(X, cast<ConstantInt>(C)->getZExtValue(),
621 NewIndex, Builder);
622 return Shuf;
623}
624
625/// Try to reduce extract element costs by converting scalar compares to vector
626/// compares followed by extract.
627/// cmp (ext0 V0, ExtIndex), (ext1 V1, ExtIndex)
628Value *VectorCombine::foldExtExtCmp(Value *V0, Value *V1, Value *ExtIndex,
629 Instruction &I) {
630 assert(isa<CmpInst>(&I) && "Expected a compare");
631
632 // cmp Pred (extelt V0, ExtIndex), (extelt V1, ExtIndex)
633 // --> extelt (cmp Pred V0, V1), ExtIndex
634 ++NumVecCmp;
635 CmpInst::Predicate Pred = cast<CmpInst>(&I)->getPredicate();
636 Value *VecCmp = Builder.CreateCmp(Pred, V0, V1);
637 return Builder.CreateExtractElement(VecCmp, ExtIndex, "foldExtExtCmp");
638}
639
640/// Try to reduce extract element costs by converting scalar binops to vector
641/// binops followed by extract.
642/// bo (ext0 V0, ExtIndex), (ext1 V1, ExtIndex)
643Value *VectorCombine::foldExtExtBinop(Value *V0, Value *V1, Value *ExtIndex,
644 Instruction &I) {
645 assert(isa<BinaryOperator>(&I) && "Expected a binary operator");
646
647 // bo (extelt V0, ExtIndex), (extelt V1, ExtIndex)
648 // --> extelt (bo V0, V1), ExtIndex
649 ++NumVecBO;
650 Value *VecBO = Builder.CreateBinOp(cast<BinaryOperator>(&I)->getOpcode(), V0,
651 V1, "foldExtExtBinop");
652
653 // All IR flags are safe to back-propagate because any potential poison
654 // created in unused vector elements is discarded by the extract.
655 if (auto *VecBOInst = dyn_cast<Instruction>(VecBO))
656 VecBOInst->copyIRFlags(&I);
657
658 return Builder.CreateExtractElement(VecBO, ExtIndex, "foldExtExtBinop");
659}
660
661/// Match an instruction with extracted vector operands.
662bool VectorCombine::foldExtractExtract(Instruction &I) {
663 // It is not safe to transform things like div, urem, etc. because we may
664 // create undefined behavior when executing those on unknown vector elements.
666 return false;
667
668 Instruction *I0, *I1;
669 CmpPredicate Pred = CmpInst::BAD_ICMP_PREDICATE;
670 if (!match(&I, m_Cmp(Pred, m_Instruction(I0), m_Instruction(I1))) &&
672 return false;
673
674 Value *V0, *V1;
675 uint64_t C0, C1;
676 if (!match(I0, m_ExtractElt(m_Value(V0), m_ConstantInt(C0))) ||
678 V0->getType() != V1->getType())
679 return false;
680
681 // For fixed-width vectors, reject out-of-bounds extract indexes
682 if (auto *FixedVecTy = dyn_cast<FixedVectorType>(V0->getType())) {
683 unsigned NumElts = FixedVecTy->getNumElements();
684 if (C0 >= NumElts || C1 >= NumElts)
685 return false;
686 }
687
688 // If the scalar value 'I' is going to be re-inserted into a vector, then try
689 // to create an extract to that same element. The extract/insert can be
690 // reduced to a "select shuffle".
691 // TODO: If we add a larger pattern match that starts from an insert, this
692 // probably becomes unnecessary.
693 auto *Ext0 = cast<ExtractElementInst>(I0);
694 auto *Ext1 = cast<ExtractElementInst>(I1);
695 uint64_t InsertIndex = InvalidIndex;
696 if (I.hasOneUse())
697 match(I.user_back(),
698 m_InsertElt(m_Value(), m_Value(), m_ConstantInt(InsertIndex)));
699
700 ExtractElementInst *ExtractToChange;
701 if (isExtractExtractCheap(Ext0, Ext1, I, ExtractToChange, InsertIndex))
702 return false;
703
704 Value *ExtOp0 = Ext0->getVectorOperand();
705 Value *ExtOp1 = Ext1->getVectorOperand();
706
707 if (ExtractToChange) {
708 unsigned CheapExtractIdx = ExtractToChange == Ext0 ? C1 : C0;
709 Value *NewExtOp =
710 translateExtract(ExtractToChange, CheapExtractIdx, Builder);
711 if (!NewExtOp)
712 return false;
713 if (ExtractToChange == Ext0)
714 ExtOp0 = NewExtOp;
715 else
716 ExtOp1 = NewExtOp;
717 }
718
719 Value *ExtIndex = ExtractToChange == Ext0 ? Ext1->getIndexOperand()
720 : Ext0->getIndexOperand();
721 Value *NewExt = Pred != CmpInst::BAD_ICMP_PREDICATE
722 ? foldExtExtCmp(ExtOp0, ExtOp1, ExtIndex, I)
723 : foldExtExtBinop(ExtOp0, ExtOp1, ExtIndex, I);
724 Worklist.push(Ext0);
725 Worklist.push(Ext1);
726 replaceValue(I, *NewExt);
727 return true;
728}
729
730/// Try to replace an extract + scalar fneg + insert with a vector fneg +
731/// shuffle.
732bool VectorCombine::foldInsExtFNeg(Instruction &I) {
733 // Match an insert (op (extract)) pattern.
734 Value *DstVec;
735 uint64_t ExtIdx, InsIdx;
736 Instruction *FNeg;
737 if (!match(&I, m_InsertElt(m_Value(DstVec), m_OneUse(m_Instruction(FNeg)),
738 m_ConstantInt(InsIdx))))
739 return false;
740
741 // Note: This handles the canonical fneg instruction and "fsub -0.0, X".
742 Value *SrcVec;
743 Instruction *Extract;
744 if (!match(FNeg, m_FNeg(m_CombineAnd(
745 m_Instruction(Extract),
746 m_ExtractElt(m_Value(SrcVec), m_ConstantInt(ExtIdx))))))
747 return false;
748
749 auto *DstVecTy = cast<FixedVectorType>(DstVec->getType());
750 auto *DstVecScalarTy = DstVecTy->getScalarType();
751 auto *SrcVecTy = dyn_cast<FixedVectorType>(SrcVec->getType());
752 if (!SrcVecTy || DstVecScalarTy != SrcVecTy->getScalarType())
753 return false;
754
755 // Ignore if insert/extract index is out of bounds or destination vector has
756 // one element
757 unsigned NumDstElts = DstVecTy->getNumElements();
758 unsigned NumSrcElts = SrcVecTy->getNumElements();
759 if (ExtIdx > NumSrcElts || InsIdx >= NumDstElts || NumDstElts == 1)
760 return false;
761
762 // We are inserting the negated element into the same lane that we extracted
763 // from. This is equivalent to a select-shuffle that chooses all but the
764 // negated element from the destination vector.
765 SmallVector<int> Mask(NumDstElts);
766 std::iota(Mask.begin(), Mask.end(), 0);
767 Mask[InsIdx] = (ExtIdx % NumDstElts) + NumDstElts;
768 InstructionCost OldCost =
769 TTI.getArithmeticInstrCost(Instruction::FNeg, DstVecScalarTy, CostKind) +
770 TTI.getVectorInstrCost(I, DstVecTy, CostKind, InsIdx);
771
772 // If the extract has one use, it will be eliminated, so count it in the
773 // original cost. If it has more than one use, ignore the cost because it will
774 // be the same before/after.
775 if (Extract->hasOneUse())
776 OldCost += TTI.getVectorInstrCost(*Extract, SrcVecTy, CostKind, ExtIdx);
777
778 InstructionCost NewCost =
779 TTI.getArithmeticInstrCost(Instruction::FNeg, SrcVecTy, CostKind) +
781 DstVecTy, Mask, CostKind);
782
783 bool NeedLenChg = SrcVecTy->getNumElements() != NumDstElts;
784 // If the lengths of the two vectors are not equal,
785 // we need to add a length-change vector. Add this cost.
786 SmallVector<int> SrcMask;
787 if (NeedLenChg) {
788 SrcMask.assign(NumDstElts, PoisonMaskElem);
789 SrcMask[ExtIdx % NumDstElts] = ExtIdx;
791 DstVecTy, SrcVecTy, SrcMask, CostKind);
792 }
793
794 LLVM_DEBUG(dbgs() << "Found an insertion of (extract)fneg : " << I
795 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
796 << "\n");
797 if (NewCost > OldCost)
798 return false;
799
800 Value *NewShuf, *LenChgShuf = nullptr;
801 // insertelt DstVec, (fneg (extractelt SrcVec, Index)), Index
802 Value *VecFNeg = Builder.CreateFNegFMF(SrcVec, FNeg);
803 if (NeedLenChg) {
804 // shuffle DstVec, (shuffle (fneg SrcVec), poison, SrcMask), Mask
805 LenChgShuf = Builder.CreateShuffleVector(VecFNeg, SrcMask);
806 NewShuf = Builder.CreateShuffleVector(DstVec, LenChgShuf, Mask);
807 Worklist.pushValue(LenChgShuf);
808 } else {
809 // shuffle DstVec, (fneg SrcVec), Mask
810 NewShuf = Builder.CreateShuffleVector(DstVec, VecFNeg, Mask);
811 }
812
813 Worklist.pushValue(VecFNeg);
814 replaceValue(I, *NewShuf);
815 return true;
816}
817
818/// Try to fold insert(binop(x,y),binop(a,b),idx)
819/// --> binop(insert(x,a,idx),insert(y,b,idx))
820bool VectorCombine::foldInsExtBinop(Instruction &I) {
821 BinaryOperator *VecBinOp, *SclBinOp;
823 if (!match(&I,
824 m_InsertElt(m_OneUse(m_BinOp(VecBinOp)),
825 m_OneUse(m_BinOp(SclBinOp)), m_ConstantInt(Index))))
826 return false;
827
828 // TODO: Add support for addlike etc.
829 Instruction::BinaryOps BinOpcode = VecBinOp->getOpcode();
830 if (BinOpcode != SclBinOp->getOpcode())
831 return false;
832
833 auto *ResultTy = dyn_cast<FixedVectorType>(I.getType());
834 if (!ResultTy)
835 return false;
836
837 // TODO: Attempt to detect m_ExtractElt for scalar operands and convert to
838 // shuffle?
839
841 TTI.getInstructionCost(VecBinOp, CostKind) +
843 InstructionCost NewCost =
844 TTI.getArithmeticInstrCost(BinOpcode, ResultTy, CostKind) +
845 TTI.getVectorInstrCost(Instruction::InsertElement, ResultTy, CostKind,
846 Index, VecBinOp->getOperand(0),
847 SclBinOp->getOperand(0)) +
848 TTI.getVectorInstrCost(Instruction::InsertElement, ResultTy, CostKind,
849 Index, VecBinOp->getOperand(1),
850 SclBinOp->getOperand(1));
851
852 LLVM_DEBUG(dbgs() << "Found an insertion of two binops: " << I
853 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
854 << "\n");
855 if (NewCost > OldCost)
856 return false;
857
858 Value *NewIns0 = Builder.CreateInsertElement(VecBinOp->getOperand(0),
859 SclBinOp->getOperand(0), Index);
860 Value *NewIns1 = Builder.CreateInsertElement(VecBinOp->getOperand(1),
861 SclBinOp->getOperand(1), Index);
862 Value *NewBO = Builder.CreateBinOp(BinOpcode, NewIns0, NewIns1);
863
864 // Intersect flags from the old binops.
865 if (auto *NewInst = dyn_cast<Instruction>(NewBO)) {
866 NewInst->copyIRFlags(VecBinOp);
867 NewInst->andIRFlags(SclBinOp);
868 }
869
870 Worklist.pushValue(NewIns0);
871 Worklist.pushValue(NewIns1);
872 replaceValue(I, *NewBO);
873 return true;
874}
875
876/// Match: bitop(castop(x), castop(y)) -> castop(bitop(x, y))
877/// Supports: bitcast, trunc, sext, zext
878bool VectorCombine::foldBitOpOfCastops(Instruction &I) {
879 // Check if this is a bitwise logic operation
880 auto *BinOp = dyn_cast<BinaryOperator>(&I);
881 if (!BinOp || !BinOp->isBitwiseLogicOp())
882 return false;
883
884 // Get the cast instructions
885 auto *LHSCast = dyn_cast<CastInst>(BinOp->getOperand(0));
886 auto *RHSCast = dyn_cast<CastInst>(BinOp->getOperand(1));
887 if (!LHSCast || !RHSCast) {
888 LLVM_DEBUG(dbgs() << " One or both operands are not cast instructions\n");
889 return false;
890 }
891
892 // Both casts must be the same type
893 Instruction::CastOps CastOpcode = LHSCast->getOpcode();
894 if (CastOpcode != RHSCast->getOpcode())
895 return false;
896
897 // Only handle supported cast operations
898 switch (CastOpcode) {
899 case Instruction::BitCast:
900 case Instruction::Trunc:
901 case Instruction::SExt:
902 case Instruction::ZExt:
903 break;
904 default:
905 return false;
906 }
907
908 Value *LHSSrc = LHSCast->getOperand(0);
909 Value *RHSSrc = RHSCast->getOperand(0);
910
911 // Source types must match
912 if (LHSSrc->getType() != RHSSrc->getType())
913 return false;
914
915 auto *SrcTy = LHSSrc->getType();
916 auto *DstTy = I.getType();
917 // Bitcasts can handle scalar/vector mixes, such as i16 -> <16 x i1>.
918 // Other casts only handle vector types with integer elements.
919 if (CastOpcode != Instruction::BitCast &&
920 (!isa<FixedVectorType>(SrcTy) || !isa<FixedVectorType>(DstTy)))
921 return false;
922
923 // Only integer scalar/vector values are legal for bitwise logic operations.
924 if (!SrcTy->getScalarType()->isIntegerTy() ||
925 !DstTy->getScalarType()->isIntegerTy())
926 return false;
927
928 // Cost Check :
929 // OldCost = bitlogic + 2*casts
930 // NewCost = bitlogic + cast
931
932 // Calculate specific costs for each cast with instruction context
934 CastOpcode, DstTy, SrcTy, TTI::CastContextHint::None, CostKind, LHSCast);
936 CastOpcode, DstTy, SrcTy, TTI::CastContextHint::None, CostKind, RHSCast);
937
938 InstructionCost OldCost =
939 TTI.getArithmeticInstrCost(BinOp->getOpcode(), DstTy, CostKind) +
940 LHSCastCost + RHSCastCost;
941
942 // For new cost, we can't provide an instruction (it doesn't exist yet)
943 InstructionCost GenericCastCost = TTI.getCastInstrCost(
944 CastOpcode, DstTy, SrcTy, TTI::CastContextHint::None, CostKind);
945
946 InstructionCost NewCost =
947 TTI.getArithmeticInstrCost(BinOp->getOpcode(), SrcTy, CostKind) +
948 GenericCastCost;
949
950 // Account for multi-use casts using specific costs
951 if (!LHSCast->hasOneUse())
952 NewCost += LHSCastCost;
953 if (!RHSCast->hasOneUse())
954 NewCost += RHSCastCost;
955
956 LLVM_DEBUG(dbgs() << "foldBitOpOfCastops: OldCost=" << OldCost
957 << " NewCost=" << NewCost << "\n");
958
959 if (NewCost > OldCost)
960 return false;
961
962 // Create the operation on the source type
963 Value *NewOp = Builder.CreateBinOp(BinOp->getOpcode(), LHSSrc, RHSSrc,
964 BinOp->getName() + ".inner");
965 if (auto *NewBinOp = dyn_cast<BinaryOperator>(NewOp))
966 NewBinOp->copyIRFlags(BinOp);
967
968 Worklist.pushValue(NewOp);
969
970 // Create the cast operation directly to ensure we get a new instruction
971 Instruction *NewCast = CastInst::Create(CastOpcode, NewOp, I.getType());
972
973 // Preserve cast instruction flags
974 NewCast->copyIRFlags(LHSCast);
975 NewCast->andIRFlags(RHSCast);
976
977 // Insert the new instruction
978 Value *Result = Builder.Insert(NewCast);
979
980 replaceValue(I, *Result);
981 return true;
982}
983
984/// Match:
985// bitop(castop(x), C) ->
986// bitop(castop(x), castop(InvC)) ->
987// castop(bitop(x, InvC))
988// Supports: bitcast
989bool VectorCombine::foldBitOpOfCastConstant(Instruction &I) {
991 Constant *C;
992
993 // Check if this is a bitwise logic operation
995 return false;
996
997 // Get the cast instructions
998 auto *LHSCast = dyn_cast<CastInst>(LHS);
999 if (!LHSCast)
1000 return false;
1001
1002 Instruction::CastOps CastOpcode = LHSCast->getOpcode();
1003
1004 // Only handle supported cast operations
1005 switch (CastOpcode) {
1006 case Instruction::BitCast:
1007 case Instruction::ZExt:
1008 case Instruction::SExt:
1009 case Instruction::Trunc:
1010 break;
1011 default:
1012 return false;
1013 }
1014
1015 Value *LHSSrc = LHSCast->getOperand(0);
1016
1017 auto *SrcTy = LHSSrc->getType();
1018 auto *DstTy = I.getType();
1019 // Bitcasts can handle scalar/vector mixes, such as i16 -> <16 x i1>.
1020 // Other casts only handle vector types with integer elements.
1021 if (CastOpcode != Instruction::BitCast &&
1022 (!isa<FixedVectorType>(SrcTy) || !isa<FixedVectorType>(DstTy)))
1023 return false;
1024
1025 // Only integer scalar/vector values are legal for bitwise logic operations.
1026 if (!SrcTy->getScalarType()->isIntegerTy() ||
1027 !DstTy->getScalarType()->isIntegerTy())
1028 return false;
1029
1030 // Find the constant InvC, such that castop(InvC) equals to C.
1031 PreservedCastFlags RHSFlags;
1032 Constant *InvC = getLosslessInvCast(C, SrcTy, CastOpcode, *DL, &RHSFlags);
1033 if (!InvC)
1034 return false;
1035
1036 // Cost Check :
1037 // OldCost = bitlogic + cast
1038 // NewCost = bitlogic + cast
1039
1040 // Calculate specific costs for each cast with instruction context
1041 InstructionCost LHSCastCost = TTI.getCastInstrCost(
1042 CastOpcode, DstTy, SrcTy, TTI::CastContextHint::None, CostKind, LHSCast);
1043
1044 InstructionCost OldCost =
1045 TTI.getArithmeticInstrCost(I.getOpcode(), DstTy, CostKind) + LHSCastCost;
1046
1047 // For new cost, we can't provide an instruction (it doesn't exist yet)
1048 InstructionCost GenericCastCost = TTI.getCastInstrCost(
1049 CastOpcode, DstTy, SrcTy, TTI::CastContextHint::None, CostKind);
1050
1051 InstructionCost NewCost =
1052 TTI.getArithmeticInstrCost(I.getOpcode(), SrcTy, CostKind) +
1053 GenericCastCost;
1054
1055 // Account for multi-use casts using specific costs
1056 if (!LHSCast->hasOneUse())
1057 NewCost += LHSCastCost;
1058
1059 LLVM_DEBUG(dbgs() << "foldBitOpOfCastConstant: OldCost=" << OldCost
1060 << " NewCost=" << NewCost << "\n");
1061
1062 if (NewCost > OldCost)
1063 return false;
1064
1065 // Create the operation on the source type
1066 Value *NewOp = Builder.CreateBinOp((Instruction::BinaryOps)I.getOpcode(),
1067 LHSSrc, InvC, I.getName() + ".inner");
1068 if (auto *NewBinOp = dyn_cast<BinaryOperator>(NewOp))
1069 NewBinOp->copyIRFlags(&I);
1070
1071 Worklist.pushValue(NewOp);
1072
1073 // Create the cast operation directly to ensure we get a new instruction
1074 Instruction *NewCast = CastInst::Create(CastOpcode, NewOp, I.getType());
1075
1076 // Preserve cast instruction flags
1077 if (RHSFlags.NNeg)
1078 NewCast->setNonNeg();
1079 if (RHSFlags.NUW)
1080 NewCast->setHasNoUnsignedWrap();
1081 if (RHSFlags.NSW)
1082 NewCast->setHasNoSignedWrap();
1083
1084 NewCast->andIRFlags(LHSCast);
1085
1086 // Insert the new instruction
1087 Value *Result = Builder.Insert(NewCast);
1088
1089 replaceValue(I, *Result);
1090 return true;
1091}
1092
1093/// If this is a bitcast of a shuffle, try to bitcast the source vector to the
1094/// destination type followed by shuffle. This can enable further transforms by
1095/// moving bitcasts or shuffles together.
1096bool VectorCombine::foldBitcastShuffle(Instruction &I) {
1097 Value *V0, *V1;
1098 ArrayRef<int> Mask;
1099 if (!match(&I, m_BitCast(m_OneUse(
1100 m_Shuffle(m_Value(V0), m_Value(V1), m_Mask(Mask))))))
1101 return false;
1102
1103 // 1) Do not fold bitcast shuffle for scalable type. First, shuffle cost for
1104 // scalable type is unknown; Second, we cannot reason if the narrowed shuffle
1105 // mask for scalable type is a splat or not.
1106 // 2) Disallow non-vector casts.
1107 // TODO: We could allow any shuffle.
1108 auto *DestTy = dyn_cast<FixedVectorType>(I.getType());
1109 auto *SrcTy = dyn_cast<FixedVectorType>(V0->getType());
1110 if (!DestTy || !SrcTy)
1111 return false;
1112
1113 unsigned DestEltSize = DestTy->getScalarSizeInBits();
1114 unsigned SrcEltSize = SrcTy->getScalarSizeInBits();
1115 if (SrcTy->getPrimitiveSizeInBits() % DestEltSize != 0)
1116 return false;
1117
1118 bool IsUnary = isa<UndefValue>(V1);
1119
1120 // For binary shuffles, only fold bitcast(shuffle(X,Y))
1121 // if it won't increase the number of bitcasts.
1122 if (!IsUnary) {
1125 if (!(BCTy0 && BCTy0->getElementType() == DestTy->getElementType()) &&
1126 !(BCTy1 && BCTy1->getElementType() == DestTy->getElementType()))
1127 return false;
1128 }
1129
1130 SmallVector<int, 16> NewMask;
1131 if (DestEltSize <= SrcEltSize) {
1132 // The bitcast is from wide to narrow/equal elements. The shuffle mask can
1133 // always be expanded to the equivalent form choosing narrower elements.
1134 if (SrcEltSize % DestEltSize != 0)
1135 return false;
1136 unsigned ScaleFactor = SrcEltSize / DestEltSize;
1137 narrowShuffleMaskElts(ScaleFactor, Mask, NewMask);
1138 } else {
1139 // The bitcast is from narrow elements to wide elements. The shuffle mask
1140 // must choose consecutive elements to allow casting first.
1141 if (DestEltSize % SrcEltSize != 0)
1142 return false;
1143 unsigned ScaleFactor = DestEltSize / SrcEltSize;
1144 if (!widenShuffleMaskElts(ScaleFactor, Mask, NewMask))
1145 return false;
1146 }
1147
1148 // Bitcast the shuffle src - keep its original width but using the destination
1149 // scalar type.
1150 unsigned NumSrcElts = SrcTy->getPrimitiveSizeInBits() / DestEltSize;
1151 auto *NewShuffleTy =
1152 FixedVectorType::get(DestTy->getScalarType(), NumSrcElts);
1153 auto *OldShuffleTy =
1154 FixedVectorType::get(SrcTy->getScalarType(), Mask.size());
1155 unsigned NumOps = IsUnary ? 1 : 2;
1156
1157 // The new shuffle must not cost more than the old shuffle.
1161
1162 InstructionCost NewCost =
1163 TTI.getShuffleCost(SK, DestTy, NewShuffleTy, NewMask, CostKind) +
1164 (NumOps * TTI.getCastInstrCost(Instruction::BitCast, NewShuffleTy, SrcTy,
1165 TargetTransformInfo::CastContextHint::None,
1166 CostKind));
1167 InstructionCost OldCost =
1168 TTI.getShuffleCost(SK, OldShuffleTy, SrcTy, Mask, CostKind) +
1169 TTI.getCastInstrCost(Instruction::BitCast, DestTy, OldShuffleTy,
1170 TargetTransformInfo::CastContextHint::None,
1171 CostKind);
1172
1173 LLVM_DEBUG(dbgs() << "Found a bitcasted shuffle: " << I << "\n OldCost: "
1174 << OldCost << " vs NewCost: " << NewCost << "\n");
1175
1176 if (NewCost > OldCost || !NewCost.isValid())
1177 return false;
1178
1179 // bitcast (shuf V0, V1, MaskC) --> shuf (bitcast V0), (bitcast V1), MaskC'
1180 ++NumShufOfBitcast;
1181 Value *CastV0 = Builder.CreateBitCast(peekThroughBitcasts(V0), NewShuffleTy);
1182 Value *CastV1 = Builder.CreateBitCast(peekThroughBitcasts(V1), NewShuffleTy);
1183 Value *Shuf = Builder.CreateShuffleVector(CastV0, CastV1, NewMask);
1184 replaceValue(I, *Shuf);
1185 return true;
1186}
1187
1188/// Match a vector op/compare/intrinsic with at least one
1189/// inserted scalar operand and convert to scalar op/cmp/intrinsic followed
1190/// by insertelement.
1191bool VectorCombine::scalarizeOpOrCmp(Instruction &I) {
1192 auto *UO = dyn_cast<UnaryOperator>(&I);
1193 auto *BO = dyn_cast<BinaryOperator>(&I);
1194 auto *CI = dyn_cast<CmpInst>(&I);
1195 auto *II = dyn_cast<IntrinsicInst>(&I);
1196 if (!UO && !BO && !CI && !II)
1197 return false;
1198
1199 // TODO: Allow intrinsics with different argument types
1200 if (II) {
1201 if (!isTriviallyVectorizable(II->getIntrinsicID()))
1202 return false;
1203 for (auto [Idx, Arg] : enumerate(II->args()))
1204 if (Arg->getType() != II->getType() &&
1205 !isVectorIntrinsicWithScalarOpAtArg(II->getIntrinsicID(), Idx, &TTI))
1206 return false;
1207 }
1208
1209 // Do not convert the vector condition of a vector select into a scalar
1210 // condition. That may cause problems for codegen because of differences in
1211 // boolean formats and register-file transfers.
1212 // TODO: Can we account for that in the cost model?
1213 if (CI)
1214 for (User *U : I.users())
1215 if (match(U, m_Select(m_Specific(&I), m_Value(), m_Value())))
1216 return false;
1217
1218 // Match constant vectors or scalars being inserted into constant vectors:
1219 // vec_op [VecC0 | (inselt VecC0, V0, Index)], ...
1220 SmallVector<Value *> VecCs, ScalarOps;
1221 std::optional<uint64_t> Index;
1222
1223 auto Ops = II ? II->args() : I.operands();
1224 for (auto [OpNum, Op] : enumerate(Ops)) {
1225 Constant *VecC;
1226 Value *V;
1227 uint64_t InsIdx = 0;
1228 if (match(Op.get(), m_InsertElt(m_Constant(VecC), m_Value(V),
1229 m_ConstantInt(InsIdx)))) {
1230 // Bail if any inserts are out of bounds.
1231 VectorType *OpTy = cast<VectorType>(Op->getType());
1232 if (OpTy->getElementCount().getKnownMinValue() <= InsIdx)
1233 return false;
1234 // All inserts must have the same index.
1235 // TODO: Deal with mismatched index constants and variable indexes?
1236 if (!Index)
1237 Index = InsIdx;
1238 else if (InsIdx != *Index)
1239 return false;
1240 VecCs.push_back(VecC);
1241 ScalarOps.push_back(V);
1242 } else if (II && isVectorIntrinsicWithScalarOpAtArg(II->getIntrinsicID(),
1243 OpNum, &TTI)) {
1244 VecCs.push_back(Op.get());
1245 ScalarOps.push_back(Op.get());
1246 } else if (match(Op.get(), m_Constant(VecC))) {
1247 VecCs.push_back(VecC);
1248 ScalarOps.push_back(nullptr);
1249 } else {
1250 return false;
1251 }
1252 }
1253
1254 // Bail if all operands are constant.
1255 if (!Index.has_value())
1256 return false;
1257
1258 VectorType *VecTy = cast<VectorType>(I.getType());
1259 Type *ScalarTy = VecTy->getScalarType();
1260 assert(VecTy->isVectorTy() &&
1261 (ScalarTy->isIntegerTy() || ScalarTy->isFloatingPointTy() ||
1262 ScalarTy->isPointerTy()) &&
1263 "Unexpected types for insert element into binop or cmp");
1264
1265 unsigned Opcode = I.getOpcode();
1266 InstructionCost ScalarOpCost, VectorOpCost;
1267 if (CI) {
1268 CmpInst::Predicate Pred = CI->getPredicate();
1269 ScalarOpCost = TTI.getCmpSelInstrCost(
1270 Opcode, ScalarTy, CmpInst::makeCmpResultType(ScalarTy), Pred, CostKind);
1271 VectorOpCost = TTI.getCmpSelInstrCost(
1272 Opcode, VecTy, CmpInst::makeCmpResultType(VecTy), Pred, CostKind);
1273 } else if (UO || BO) {
1274 ScalarOpCost = TTI.getArithmeticInstrCost(Opcode, ScalarTy, CostKind);
1275 VectorOpCost = TTI.getArithmeticInstrCost(Opcode, VecTy, CostKind);
1276 } else {
1277 IntrinsicCostAttributes ScalarICA(
1278 II->getIntrinsicID(), ScalarTy,
1279 SmallVector<Type *>(II->arg_size(), ScalarTy));
1280 ScalarOpCost = TTI.getIntrinsicInstrCost(ScalarICA, CostKind);
1281 IntrinsicCostAttributes VectorICA(
1282 II->getIntrinsicID(), VecTy,
1283 SmallVector<Type *>(II->arg_size(), VecTy));
1284 VectorOpCost = TTI.getIntrinsicInstrCost(VectorICA, CostKind);
1285 }
1286
1287 // Fold the vector constants in the original vectors into a new base vector to
1288 // get more accurate cost modelling.
1289 Value *NewVecC = nullptr;
1290 if (CI)
1291 NewVecC = simplifyCmpInst(CI->getPredicate(), VecCs[0], VecCs[1], SQ);
1292 else if (UO)
1293 NewVecC =
1294 simplifyUnOp(UO->getOpcode(), VecCs[0], UO->getFastMathFlags(), SQ);
1295 else if (BO)
1296 NewVecC = simplifyBinOp(BO->getOpcode(), VecCs[0], VecCs[1], SQ);
1297 else if (II)
1298 NewVecC = simplifyCall(II, II->getCalledOperand(), VecCs, SQ);
1299
1300 if (!NewVecC)
1301 return false;
1302
1303 // Get cost estimate for the insert element. This cost will factor into
1304 // both sequences.
1305 InstructionCost OldCost = VectorOpCost;
1306 InstructionCost NewCost =
1307 ScalarOpCost + TTI.getVectorInstrCost(Instruction::InsertElement, VecTy,
1308 CostKind, *Index, NewVecC);
1309
1310 for (auto [Idx, Op, VecC, Scalar] : enumerate(Ops, VecCs, ScalarOps)) {
1311 if (!Scalar || (II && isVectorIntrinsicWithScalarOpAtArg(
1312 II->getIntrinsicID(), Idx, &TTI)))
1313 continue;
1315 Instruction::InsertElement, VecTy, CostKind, *Index, VecC, Scalar);
1316 OldCost += InsertCost;
1317 NewCost += !Op->hasOneUse() * InsertCost;
1318 }
1319
1320 // We want to scalarize unless the vector variant actually has lower cost.
1321 if (OldCost < NewCost || !NewCost.isValid())
1322 return false;
1323
1324 // vec_op (inselt VecC0, V0, Index), (inselt VecC1, V1, Index) -->
1325 // inselt NewVecC, (scalar_op V0, V1), Index
1326 if (CI)
1327 ++NumScalarCmp;
1328 else if (UO || BO)
1329 ++NumScalarOps;
1330 else
1331 ++NumScalarIntrinsic;
1332
1333 // For constant cases, extract the scalar element, this should constant fold.
1334 for (auto [OpIdx, Scalar, VecC] : enumerate(ScalarOps, VecCs))
1335 if (!Scalar)
1336 ScalarOps[OpIdx] = ConstantExpr::getExtractElement(
1337 cast<Constant>(VecC), Builder.getInt64(*Index));
1338
1339 Value *Scalar;
1340 if (CI)
1341 Scalar = Builder.CreateCmp(CI->getPredicate(), ScalarOps[0], ScalarOps[1]);
1342 else if (UO || BO)
1343 Scalar = Builder.CreateNAryOp(Opcode, ScalarOps);
1344 else
1345 Scalar = Builder.CreateIntrinsic(ScalarTy, II->getIntrinsicID(), ScalarOps);
1346
1347 Scalar->setName(I.getName() + ".scalar");
1348
1349 // All IR flags are safe to back-propagate. There is no potential for extra
1350 // poison to be created by the scalar instruction.
1351 if (auto *ScalarInst = dyn_cast<Instruction>(Scalar))
1352 ScalarInst->copyIRFlags(&I);
1353
1354 Value *Insert = Builder.CreateInsertElement(NewVecC, Scalar, *Index);
1355 replaceValue(I, *Insert);
1356 return true;
1357}
1358
1359/// Try to combine a scalar binop + 2 scalar compares of extracted elements of
1360/// a vector into vector operations followed by extract. Note: The SLP pass
1361/// may miss this pattern because of implementation problems.
1362bool VectorCombine::foldExtractedCmps(Instruction &I) {
1363 auto *BI = dyn_cast<BinaryOperator>(&I);
1364
1365 // We are looking for a scalar binop of booleans.
1366 // binop i1 (cmp Pred I0, C0), (cmp Pred I1, C1)
1367 if (!BI || !I.getType()->isIntegerTy(1))
1368 return false;
1369
1370 // The compare predicates should match, and each compare should have a
1371 // constant operand.
1372 Value *B0 = I.getOperand(0), *B1 = I.getOperand(1);
1373 Instruction *I0, *I1;
1374 Constant *C0, *C1;
1375 CmpPredicate P0, P1;
1376 if (!match(B0, m_Cmp(P0, m_Instruction(I0), m_Constant(C0))) ||
1377 !match(B1, m_Cmp(P1, m_Instruction(I1), m_Constant(C1))))
1378 return false;
1379
1380 auto MatchingPred = CmpPredicate::getMatching(P0, P1);
1381 if (!MatchingPred)
1382 return false;
1383
1384 // The compare operands must be extracts of the same vector with constant
1385 // extract indexes.
1386 Value *X;
1387 uint64_t Index0, Index1;
1388 if (!match(I0, m_ExtractElt(m_Value(X), m_ConstantInt(Index0))) ||
1389 !match(I1, m_ExtractElt(m_Specific(X), m_ConstantInt(Index1))))
1390 return false;
1391
1392 auto *Ext0 = cast<ExtractElementInst>(I0);
1393 auto *Ext1 = cast<ExtractElementInst>(I1);
1394 ExtractElementInst *ConvertToShuf = getShuffleExtract(Ext0, Ext1, CostKind);
1395 if (!ConvertToShuf)
1396 return false;
1397 assert((ConvertToShuf == Ext0 || ConvertToShuf == Ext1) &&
1398 "Unknown ExtractElementInst");
1399
1400 // The original scalar pattern is:
1401 // binop i1 (cmp Pred (ext X, Index0), C0), (cmp Pred (ext X, Index1), C1)
1402 CmpInst::Predicate Pred = *MatchingPred;
1403 unsigned CmpOpcode =
1404 CmpInst::isFPPredicate(Pred) ? Instruction::FCmp : Instruction::ICmp;
1405 auto *VecTy = dyn_cast<FixedVectorType>(X->getType());
1406 if (!VecTy)
1407 return false;
1408
1409 if (Index0 >= VecTy->getNumElements() || Index1 >= VecTy->getNumElements())
1410 return false;
1411
1412 InstructionCost Ext0Cost =
1413 TTI.getVectorInstrCost(*Ext0, VecTy, CostKind, Index0);
1414 InstructionCost Ext1Cost =
1415 TTI.getVectorInstrCost(*Ext1, VecTy, CostKind, Index1);
1417 CmpOpcode, I0->getType(), CmpInst::makeCmpResultType(I0->getType()), Pred,
1418 CostKind);
1419
1420 InstructionCost OldCost =
1421 Ext0Cost + Ext1Cost + CmpCost * 2 +
1422 TTI.getArithmeticInstrCost(I.getOpcode(), I.getType(), CostKind);
1423
1424 // The proposed vector pattern is:
1425 // vcmp = cmp Pred X, VecC
1426 // ext (binop vNi1 vcmp, (shuffle vcmp, Index1)), Index0
1427 int CheapIndex = ConvertToShuf == Ext0 ? Index1 : Index0;
1428 int ExpensiveIndex = ConvertToShuf == Ext0 ? Index0 : Index1;
1431 CmpOpcode, VecTy, CmpInst::makeCmpResultType(VecTy), Pred, CostKind);
1432 SmallVector<int, 32> ShufMask(VecTy->getNumElements(), PoisonMaskElem);
1433 ShufMask[CheapIndex] = ExpensiveIndex;
1435 CmpTy, ShufMask, CostKind);
1436 NewCost += TTI.getArithmeticInstrCost(I.getOpcode(), CmpTy, CostKind);
1437 NewCost += TTI.getVectorInstrCost(*Ext0, CmpTy, CostKind, CheapIndex);
1438 NewCost += Ext0->hasOneUse() ? 0 : Ext0Cost;
1439 NewCost += Ext1->hasOneUse() ? 0 : Ext1Cost;
1440
1441 // Aggressively form vector ops if the cost is equal because the transform
1442 // may enable further optimization.
1443 // Codegen can reverse this transform (scalarize) if it was not profitable.
1444 if (OldCost < NewCost || !NewCost.isValid())
1445 return false;
1446
1447 // Create a vector constant from the 2 scalar constants.
1448 SmallVector<Constant *, 32> CmpC(VecTy->getNumElements(),
1449 PoisonValue::get(VecTy->getElementType()));
1450 CmpC[Index0] = C0;
1451 CmpC[Index1] = C1;
1452 Value *VCmp = Builder.CreateCmp(Pred, X, ConstantVector::get(CmpC));
1453 Value *Shuf = createShiftShuffle(VCmp, ExpensiveIndex, CheapIndex, Builder);
1454 Value *LHS = ConvertToShuf == Ext0 ? Shuf : VCmp;
1455 Value *RHS = ConvertToShuf == Ext0 ? VCmp : Shuf;
1456 Value *VecLogic = Builder.CreateBinOp(BI->getOpcode(), LHS, RHS);
1457 Value *NewExt = Builder.CreateExtractElement(VecLogic, CheapIndex);
1458 replaceValue(I, *NewExt);
1459 ++NumVecCmpBO;
1460 return true;
1461}
1462
1463/// Try to fold scalar selects that select between extracted elements and zero
1464/// into extracting from a vector select. This is rooted at the bitcast.
1465///
1466/// This pattern arises when a vector is bitcast to a smaller element type,
1467/// elements are extracted, and then conditionally selected with zero:
1468///
1469/// %bc = bitcast <4 x i32> %src to <16 x i8>
1470/// %e0 = extractelement <16 x i8> %bc, i32 0
1471/// %s0 = select i1 %cond, i8 %e0, i8 0
1472/// %e1 = extractelement <16 x i8> %bc, i32 1
1473/// %s1 = select i1 %cond, i8 %e1, i8 0
1474/// ...
1475///
1476/// Transforms to:
1477/// %sel = select i1 %cond, <4 x i32> %src, <4 x i32> zeroinitializer
1478/// %bc = bitcast <4 x i32> %sel to <16 x i8>
1479/// %e0 = extractelement <16 x i8> %bc, i32 0
1480/// %e1 = extractelement <16 x i8> %bc, i32 1
1481/// ...
1482///
1483/// This is profitable because vector select on wider types produces fewer
1484/// select/cndmask instructions than scalar selects on each element.
1485bool VectorCombine::foldSelectsFromBitcast(Instruction &I) {
1486 auto *BC = dyn_cast<BitCastInst>(&I);
1487 if (!BC)
1488 return false;
1489
1490 FixedVectorType *SrcVecTy = dyn_cast<FixedVectorType>(BC->getSrcTy());
1491 FixedVectorType *DstVecTy = dyn_cast<FixedVectorType>(BC->getDestTy());
1492 if (!SrcVecTy || !DstVecTy)
1493 return false;
1494
1495 // Source must be 32-bit or 64-bit elements, destination must be smaller
1496 // integer elements. Zero in all these types is all-bits-zero.
1497 Type *SrcEltTy = SrcVecTy->getElementType();
1498 Type *DstEltTy = DstVecTy->getElementType();
1499 unsigned SrcEltBits = SrcEltTy->getPrimitiveSizeInBits();
1500 unsigned DstEltBits = DstEltTy->getPrimitiveSizeInBits();
1501
1502 if (SrcEltBits != 32 && SrcEltBits != 64)
1503 return false;
1504
1505 if (!DstEltTy->isIntegerTy() || DstEltBits >= SrcEltBits)
1506 return false;
1507
1508 // Check profitability using TTI before collecting users.
1509 Type *CondTy = CmpInst::makeCmpResultType(DstEltTy);
1510 Type *VecCondTy = CmpInst::makeCmpResultType(SrcVecTy);
1511
1512 InstructionCost ScalarSelCost =
1513 TTI.getCmpSelInstrCost(Instruction::Select, DstEltTy, CondTy,
1515 InstructionCost VecSelCost =
1516 TTI.getCmpSelInstrCost(Instruction::Select, SrcVecTy, VecCondTy,
1518
1519 // We need at least this many selects for vectorization to be profitable.
1520 // VecSelCost < ScalarSelCost * NumSelects => NumSelects > VecSelCost /
1521 // ScalarSelCost
1522 if (!ScalarSelCost.isValid() || ScalarSelCost == 0)
1523 return false;
1524
1525 unsigned MinSelects = (VecSelCost.getValue() / ScalarSelCost.getValue()) + 1;
1526
1527 // Quick check: if bitcast doesn't have enough users, bail early.
1528 if (!BC->hasNUsesOrMore(MinSelects))
1529 return false;
1530
1531 // Collect all select users that match the pattern, grouped by condition.
1532 // Pattern: select i1 %cond, (extractelement %bc, idx), 0
1533 DenseMap<Value *, SmallVector<SelectInst *, 8>> CondToSelects;
1534
1535 for (User *U : BC->users()) {
1536 auto *Ext = dyn_cast<ExtractElementInst>(U);
1537 if (!Ext)
1538 continue;
1539
1540 for (User *ExtUser : Ext->users()) {
1541 Value *Cond;
1542 // Match: select i1 %cond, %ext, 0
1543 if (match(ExtUser, m_Select(m_Value(Cond), m_Specific(Ext), m_Zero())) &&
1544 Cond->getType()->isIntegerTy(1))
1545 CondToSelects[Cond].push_back(cast<SelectInst>(ExtUser));
1546 }
1547 }
1548
1549 if (CondToSelects.empty())
1550 return false;
1551
1552 bool MadeChange = false;
1553 Value *SrcVec = BC->getOperand(0);
1554
1555 // Process each group of selects with the same condition.
1556 for (auto [Cond, Selects] : CondToSelects) {
1557 // Only profitable if vector select cost < total scalar select cost.
1558 if (Selects.size() < MinSelects) {
1559 LLVM_DEBUG(dbgs() << "VectorCombine: foldSelectsFromBitcast not "
1560 << "profitable (VecCost=" << VecSelCost
1561 << ", ScalarCost=" << ScalarSelCost
1562 << ", NumSelects=" << Selects.size() << ")\n");
1563 continue;
1564 }
1565
1566 // Create the vector select and bitcast once for this condition.
1567 auto InsertPt = std::next(BC->getIterator());
1568
1569 if (auto *CondInst = dyn_cast<Instruction>(Cond))
1570 if (DT.dominates(BC, CondInst))
1571 InsertPt = std::next(CondInst->getIterator());
1572
1573 Builder.SetInsertPoint(InsertPt);
1574 Value *VecSel =
1575 Builder.CreateSelect(Cond, SrcVec, Constant::getNullValue(SrcVecTy));
1576 Value *NewBC = Builder.CreateBitCast(VecSel, DstVecTy);
1577
1578 // Replace each scalar select with an extract from the new bitcast.
1579 for (SelectInst *Sel : Selects) {
1580 auto *Ext = cast<ExtractElementInst>(Sel->getTrueValue());
1581 Value *Idx = Ext->getIndexOperand();
1582
1583 Builder.SetInsertPoint(Sel);
1584 Value *NewExt = Builder.CreateExtractElement(NewBC, Idx);
1585 replaceValue(*Sel, *NewExt);
1586 MadeChange = true;
1587 }
1588
1589 LLVM_DEBUG(dbgs() << "VectorCombine: folded " << Selects.size()
1590 << " selects into vector select\n");
1591 }
1592
1593 return MadeChange;
1594}
1595
1598 const TargetTransformInfo &TTI,
1599 InstructionCost &CostBeforeReduction,
1600 InstructionCost &CostAfterReduction) {
1601 Instruction *Op0, *Op1;
1602 auto *RedOp = dyn_cast<Instruction>(II.getOperand(0));
1603 auto *VecRedTy = cast<VectorType>(II.getOperand(0)->getType());
1604 unsigned ReductionOpc =
1605 getArithmeticReductionInstruction(II.getIntrinsicID());
1606 if (RedOp && match(RedOp, m_ZExtOrSExt(m_Value()))) {
1607 bool IsUnsigned = isa<ZExtInst>(RedOp);
1608 auto *ExtType = cast<VectorType>(RedOp->getOperand(0)->getType());
1609
1610 CostBeforeReduction =
1611 TTI.getCastInstrCost(RedOp->getOpcode(), VecRedTy, ExtType,
1613 CostAfterReduction =
1614 TTI.getExtendedReductionCost(ReductionOpc, IsUnsigned, II.getType(),
1615 ExtType, FastMathFlags(), CostKind);
1616 return;
1617 }
1618 if (RedOp && II.getIntrinsicID() == Intrinsic::vector_reduce_add &&
1619 match(RedOp,
1621 match(Op0, m_ZExtOrSExt(m_Value())) &&
1622 Op0->getOpcode() == Op1->getOpcode() &&
1623 Op0->getOperand(0)->getType() == Op1->getOperand(0)->getType() &&
1624 (Op0->getOpcode() == RedOp->getOpcode() || Op0 == Op1)) {
1625 // Matched reduce.add(ext(mul(ext(A), ext(B)))
1626 bool IsUnsigned = isa<ZExtInst>(Op0);
1627 auto *ExtType = cast<VectorType>(Op0->getOperand(0)->getType());
1628 VectorType *MulType = VectorType::get(Op0->getType(), VecRedTy);
1629
1630 InstructionCost ExtCost =
1631 TTI.getCastInstrCost(Op0->getOpcode(), MulType, ExtType,
1633 InstructionCost MulCost =
1634 TTI.getArithmeticInstrCost(Instruction::Mul, MulType, CostKind);
1635 InstructionCost Ext2Cost =
1636 TTI.getCastInstrCost(RedOp->getOpcode(), VecRedTy, MulType,
1638
1639 CostBeforeReduction = ExtCost * 2 + MulCost + Ext2Cost;
1640 CostAfterReduction = TTI.getMulAccReductionCost(
1641 IsUnsigned, ReductionOpc, II.getType(), ExtType, CostKind);
1642 return;
1643 }
1644 CostAfterReduction = TTI.getArithmeticReductionCost(ReductionOpc, VecRedTy,
1645 std::nullopt, CostKind);
1646}
1647
1648bool VectorCombine::foldBinopOfReductions(Instruction &I) {
1649 Instruction::BinaryOps BinOpOpc = cast<BinaryOperator>(&I)->getOpcode();
1650 Intrinsic::ID ReductionIID = getReductionForBinop(BinOpOpc);
1651 if (BinOpOpc == Instruction::Sub)
1652 ReductionIID = Intrinsic::vector_reduce_add;
1653 if (ReductionIID == Intrinsic::not_intrinsic)
1654 return false;
1655 // FP reductions have a start-value operand that this fold doesn't handle.
1656 if (ReductionIID == Intrinsic::vector_reduce_fadd ||
1657 ReductionIID == Intrinsic::vector_reduce_fmul)
1658 return false;
1659
1660 auto checkIntrinsicAndGetItsArgument = [](Value *V,
1661 Intrinsic::ID IID) -> Value * {
1662 auto *II = dyn_cast<IntrinsicInst>(V);
1663 if (!II)
1664 return nullptr;
1665 if (II->getIntrinsicID() == IID && II->hasOneUse())
1666 return II->getArgOperand(0);
1667 return nullptr;
1668 };
1669
1670 Value *V0 = checkIntrinsicAndGetItsArgument(I.getOperand(0), ReductionIID);
1671 if (!V0)
1672 return false;
1673 Value *V1 = checkIntrinsicAndGetItsArgument(I.getOperand(1), ReductionIID);
1674 if (!V1)
1675 return false;
1676
1677 auto *VTy = cast<VectorType>(V0->getType());
1678 if (V1->getType() != VTy)
1679 return false;
1680 const auto &II0 = *cast<IntrinsicInst>(I.getOperand(0));
1681 const auto &II1 = *cast<IntrinsicInst>(I.getOperand(1));
1682 unsigned ReductionOpc =
1683 getArithmeticReductionInstruction(II0.getIntrinsicID());
1684
1685 InstructionCost OldCost = 0;
1686 InstructionCost NewCost = 0;
1687 InstructionCost CostOfRedOperand0 = 0;
1688 InstructionCost CostOfRed0 = 0;
1689 InstructionCost CostOfRedOperand1 = 0;
1690 InstructionCost CostOfRed1 = 0;
1691 analyzeCostOfVecReduction(II0, CostKind, TTI, CostOfRedOperand0, CostOfRed0);
1692 analyzeCostOfVecReduction(II1, CostKind, TTI, CostOfRedOperand1, CostOfRed1);
1693 OldCost = CostOfRed0 + CostOfRed1 + TTI.getInstructionCost(&I, CostKind);
1694 NewCost =
1695 CostOfRedOperand0 + CostOfRedOperand1 +
1696 TTI.getArithmeticInstrCost(BinOpOpc, VTy, CostKind) +
1697 TTI.getArithmeticReductionCost(ReductionOpc, VTy, std::nullopt, CostKind);
1698 if (NewCost >= OldCost || !NewCost.isValid())
1699 return false;
1700
1701 LLVM_DEBUG(dbgs() << "Found two mergeable reductions: " << I
1702 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
1703 << "\n");
1704 Value *VectorBO;
1705 if (BinOpOpc == Instruction::Or)
1706 VectorBO = Builder.CreateOr(V0, V1, "",
1707 cast<PossiblyDisjointInst>(I).isDisjoint());
1708 else
1709 VectorBO = Builder.CreateBinOp(BinOpOpc, V0, V1);
1710
1711 Value *Rdx = Builder.CreateIntrinsic(ReductionIID, {VTy}, {VectorBO});
1712 replaceValue(I, *Rdx);
1713 return true;
1714}
1715
1716// Check if memory is modified, freed, or synchronized between two instrs in
1717// the same BB.
1720 const MemoryLocation &Loc, AAResults &AA) {
1721 unsigned NumScanned = 0;
1722 if (std::any_of(Begin, End, [&](const Instruction &Instr) {
1723 return isModSet(AA.getModRefInfo(&Instr, Loc)) ||
1724 ++NumScanned > MaxInstrsToScan;
1725 }))
1726 return true;
1727
1728 // willNotFreeBetween expects instructions rather than iterators. An empty
1729 // range cannot free or synchronize, so avoid dereferencing its end.
1730 return Begin != End && !willNotFreeBetween(&*Begin, &*End);
1731}
1732
1733namespace {
1734/// Helper class to indicate whether a vector index can be safely scalarized and
1735/// if a freeze needs to be inserted.
1736class ScalarizationResult {
1737 enum class StatusTy { Unsafe, Safe, SafeWithFreeze };
1738
1739 StatusTy Status;
1740 Value *ToFreeze;
1741
1742 ScalarizationResult(StatusTy Status, Value *ToFreeze = nullptr)
1743 : Status(Status), ToFreeze(ToFreeze) {}
1744
1745public:
1746 ScalarizationResult(const ScalarizationResult &Other) = default;
1747 ~ScalarizationResult() {
1748 assert(!ToFreeze && "freeze() not called with ToFreeze being set");
1749 }
1750
1751 static ScalarizationResult unsafe() { return {StatusTy::Unsafe}; }
1752 static ScalarizationResult safe() { return {StatusTy::Safe}; }
1753 static ScalarizationResult safeWithFreeze(Value *ToFreeze) {
1754 return {StatusTy::SafeWithFreeze, ToFreeze};
1755 }
1756
1757 /// Returns true if the index can be scalarize without requiring a freeze.
1758 bool isSafe() const { return Status == StatusTy::Safe; }
1759 /// Returns true if the index cannot be scalarized.
1760 bool isUnsafe() const { return Status == StatusTy::Unsafe; }
1761 /// Returns true if the index can be scalarize, but requires inserting a
1762 /// freeze.
1763 bool isSafeWithFreeze() const { return Status == StatusTy::SafeWithFreeze; }
1764
1765 /// Reset the state of Unsafe and clear ToFreze if set.
1766 void discard() {
1767 ToFreeze = nullptr;
1768 Status = StatusTy::Unsafe;
1769 }
1770
1771 /// Freeze the ToFreeze and update the use in \p User to use it.
1772 void freeze(IRBuilderBase &Builder, Instruction &UserI) {
1773 assert(isSafeWithFreeze() &&
1774 "should only be used when freezing is required");
1775 assert(is_contained(ToFreeze->users(), &UserI) &&
1776 "UserI must be a user of ToFreeze");
1777 IRBuilder<>::InsertPointGuard Guard(Builder);
1778 Builder.SetInsertPoint(cast<Instruction>(&UserI));
1779 Value *Frozen =
1780 Builder.CreateFreeze(ToFreeze, ToFreeze->getName() + ".frozen");
1781 for (Use &U : make_early_inc_range((UserI.operands())))
1782 if (U.get() == ToFreeze)
1783 U.set(Frozen);
1784
1785 ToFreeze = nullptr;
1786 }
1787};
1788} // namespace
1789
1790/// Check if it is legal to scalarize a memory access to \p VecTy at index \p
1791/// Idx. \p Idx must access a valid vector element.
1792static ScalarizationResult canScalarizeAccess(VectorType *VecTy, Value *Idx,
1793 const SimplifyQuery &SQ) {
1794 // We do checks for both fixed vector types and scalable vector types.
1795 // This is the number of elements of fixed vector types,
1796 // or the minimum number of elements of scalable vector types.
1797 uint64_t NumElements = VecTy->getElementCount().getKnownMinValue();
1798 unsigned IntWidth = Idx->getType()->getScalarSizeInBits();
1799
1800 if (auto *C = dyn_cast<ConstantInt>(Idx)) {
1801 if (C->getValue().ult(NumElements))
1802 return ScalarizationResult::safe();
1803 return ScalarizationResult::unsafe();
1804 }
1805
1806 // Always unsafe if the index type can't handle all inbound values.
1807 if (!llvm::isUIntN(IntWidth, NumElements))
1808 return ScalarizationResult::unsafe();
1809
1810 APInt Zero(IntWidth, 0);
1811 APInt MaxElts(IntWidth, NumElements);
1812 ConstantRange ValidIndices(Zero, MaxElts);
1813 ConstantRange IdxRange(IntWidth, true);
1814
1815 if (isGuaranteedNotToBePoison(Idx, SQ.AC, SQ.CxtI, SQ.DT)) {
1816 if (ValidIndices.contains(
1817 computeConstantRange(Idx, /*ForSigned=*/false, SQ)))
1818 return ScalarizationResult::safe();
1819 return ScalarizationResult::unsafe();
1820 }
1821
1822 // If the index may be poison, check if we can insert a freeze before the
1823 // range of the index is restricted.
1824 Value *IdxBase;
1825 ConstantInt *CI;
1826 if (match(Idx, m_And(m_Value(IdxBase), m_ConstantInt(CI)))) {
1827 IdxRange = IdxRange.binaryAnd(CI->getValue());
1828 } else if (match(Idx, m_URem(m_Value(IdxBase), m_ConstantInt(CI)))) {
1829 IdxRange = IdxRange.urem(CI->getValue());
1830 }
1831
1832 if (ValidIndices.contains(IdxRange))
1833 return ScalarizationResult::safeWithFreeze(IdxBase);
1834 return ScalarizationResult::unsafe();
1835}
1836
1837/// Return the GEP index type if the unsigned vector index \p Idx can be
1838/// represented by an inbounds GEP. A null result means that the maximum byte
1839/// offset cannot be represented by the pointer's signed GEP index type.
1840///
1841/// unsigned lane range
1842/// |
1843/// v
1844/// MaxByteOffset = MaxLane * element store size
1845/// |
1846/// +-- unavailable or outside signed GEP range --> reject
1847/// |
1848/// v
1849/// valid range --> use the pointer's GEP index type
1851 Type *PtrTy,
1852 const DataLayout &DL) {
1853 auto *GEPIndexTy = cast<IntegerType>(DL.getIndexType(PtrTy));
1854 unsigned GEPBits = GEPIndexTy->getBitWidth();
1855 uint64_t NumElements = VecTy->getElementCount().getKnownMinValue();
1856
1857 uint64_t MaxLane = NumElements - 1;
1858 if (auto *C = dyn_cast<ConstantInt>(Idx)) {
1859 if (C->getValue().uge(NumElements))
1860 return nullptr;
1861 MaxLane = C->getZExtValue();
1862 }
1863
1864 Type *ElemTy = VecTy->getElementType();
1865 if (!DL.typeSizeEqualsStoreSize(ElemTy))
1866 return nullptr;
1867
1868 TypeSize ElemStride = DL.getTypeStoreSize(ElemTy);
1869 if (ElemStride.isScalable())
1870 return nullptr;
1871
1872 // Compare both values in a common width:
1873 //
1874 // MaxLane (uint64_t) * ElemStride (uint64_t) signed_max(GEPBits)
1875 // | |
1876 // v v
1877 // ByteOffset (up to 128 bits) sext to WideBits
1878 // \ /
1879 // +------------ ugt ------------+
1880 // |
1881 // greater -> reject
1882 //
1883 // WideBits = max(GEPBits, 128) prevents the multiplication from wrapping
1884 // and preserves the GEP limit during the comparison.
1885 unsigned WideBits = std::max(GEPBits, 128u);
1886 APInt MaxLaneValue(WideBits, MaxLane);
1887 APInt ByteOffset = MaxLaneValue;
1888 ByteOffset *= APInt(WideBits, ElemStride.getFixedValue());
1889 APInt MaxGEPOffset = APInt::getSignedMaxValue(GEPBits).sext(WideBits);
1890 // Reject offsets outside the GEP's positive signed range. Compare as
1891 // unsigned because the full 128-bit product may set its sign bit.
1892 if (ByteOffset.ugt(MaxGEPOffset))
1893 return nullptr;
1894
1895 return GEPIndexTy;
1896}
1897
1898/// Materialize an index for a scalarized GEP after profitability is known.
1899/// Vector element indices are unsigned, but GEP sign-extends narrow integer
1900/// indices. Widen a narrow index explicitly so its unsigned value is retained.
1902 IRBuilderBase &Builder) {
1903 unsigned SrcBits = Idx->getType()->getIntegerBitWidth();
1904 unsigned DstBits = GEPIndexTy->getBitWidth();
1905 if (SrcBits >= DstBits)
1906 return Idx;
1907
1908 return Builder.CreateZExt(Idx, GEPIndexTy, Idx->getName() + ".gepidx");
1909}
1910
1911/// The memory operation on a vector of \p ScalarType had alignment of
1912/// \p VectorAlignment. Compute the maximal, but conservatively correct,
1913/// alignment that will be valid for the memory operation on a single scalar
1914/// element of the same type with index \p Idx.
1916 Type *ScalarType, Value *Idx,
1917 const DataLayout &DL) {
1918 if (auto *C = dyn_cast<ConstantInt>(Idx))
1919 return commonAlignment(VectorAlignment,
1920 C->getZExtValue() * DL.getTypeStoreSize(ScalarType));
1921 return commonAlignment(VectorAlignment, DL.getTypeStoreSize(ScalarType));
1922}
1923
1924/// Fold a vector store fed by a single-use insertelement chain into scalar
1925/// stores.
1926///
1927/// Before:
1928///
1929/// %p --> vector load --> insert %x, lane 1 --> insert %y, lane 3
1930/// |
1931/// v
1932/// vector store to %p
1933///
1934/// Vector lanes: [ 0 ] [ 1 ] [ 2 ] [ 3 ]
1935/// Stored value: [ old | x | old | y ] (one vector store)
1936///
1937/// After:
1938///
1939/// +--> GEP(%p, lane 1) --> store %x
1940/// %p -------------+
1941/// +--> GEP(%p, lane 3) --> store %y
1942///
1943/// Vector lanes: [ 0 ] [ 1 ] [ 2 ] [ 3 ]
1944/// Scalar stores: x y
1945/// store@1 store@3
1946///
1947/// Step 1. Gate:
1948/// target supports vector-element GEP addressing
1949///
1950/// Step 2. Trace:
1951/// vector store <-- insertelement <-- ... <-- insertelement <-- load
1952///
1953/// Steps 3-5. Validate:
1954/// reject unprofitable full overwrites; require simple accesses, a
1955/// common address/block, no memory write in between, and scalarizable
1956/// indices.
1957bool VectorCombine::foldInsertElementsToStores(Instruction &I) {
1958 // Step 1: The target must support addressing a vector element with a GEP.
1960 return false;
1961
1962 auto *SI = cast<StoreInst>(&I);
1963 if (!SI->isSimple() || !isa<VectorType>(SI->getValueOperand()->getType()))
1964 return false;
1965
1966 // Step 2: Collect a single-use insertelement chain, starting at the vector
1967 // store and walking back to the candidate load.
1968 Value *Source = SI->getValueOperand();
1969 SmallVector<std::pair<Value *, Value *>, 4> InsertElements;
1970 Value *Base = Source;
1971 while (auto *Insert = dyn_cast<InsertElementInst>(Base)) {
1972 if (!Insert->hasOneUse())
1973 break;
1974 Value *InsertVal = Insert->getOperand(1);
1975 Value *Idx = Insert->getOperand(2);
1976 InsertElements.push_back({InsertVal, Idx});
1977 Base = Insert->getOperand(0);
1978 }
1979
1980 if (InsertElements.empty())
1981 return false;
1982
1983 // The backwards walk collected the inserts in reverse program order. Restore
1984 // it now so later scalar stores preserve writes to duplicate/equal indices.
1985 std::reverse(InsertElements.begin(), InsertElements.end());
1986 auto *Load = dyn_cast<LoadInst>(Base);
1987 if (!Load)
1988 return false;
1989 auto *VecTy = cast<VectorType>(SI->getValueOperand()->getType());
1990
1991 // Step 3: Avoid replacing a complete overwrite with scalar stores when every
1992 // lane receives the same value; keeping the vector operation is preferable.
1993 if (auto *FVT = dyn_cast<FixedVectorType>(VecTy)) {
1994 if (InsertElements.size() == FVT->getNumElements()) {
1995 Value *FirstVal = InsertElements.front().first;
1996 if (all_of(InsertElements,
1997 [FirstVal](const auto &Elt) { return Elt.first == FirstVal; }))
1998 return false;
1999 }
2000 }
2001 Value *SrcAddr = Load->getPointerOperand()->stripPointerCasts();
2002 // Step 4: Establish the load/store update is legal: both accesses are simple,
2003 // have the same base address and block, have scalar elements whose type size
2004 // equals their store size, and no intervening operation modifies the updated
2005 // memory.
2006 if (!Load->isSimple() || Load->getParent() != SI->getParent() ||
2007 !DL->typeSizeEqualsStoreSize(Load->getType()->getScalarType()) ||
2008 SrcAddr != SI->getPointerOperand()->stripPointerCasts())
2009 return false;
2010
2011 if (isMemModifiedBetween(Load->getIterator(), SI->getIterator(),
2012 MemoryLocation::get(SI), AA))
2013 return false;
2014
2015 // Step 5: Validate every index before changing IR. A safe-with-freeze result
2016 // is recorded by ScalarizationResult, so discard it until profitability is
2017 // known; otherwise a rejected candidate could leave a freeze behind.
2018 for (auto [InsertVal, Idx] : InsertElements) {
2019 auto ScalarizableIdx =
2020 canScalarizeAccess(VecTy, Idx, SQ.getWithInstruction(&I));
2021 if (ScalarizableIdx.isUnsafe())
2022 return false;
2023
2024 auto GEPIndex =
2025 getScalarizedGEPIndexInfo(VecTy, Idx, SI->getPointerOperandType(), *DL);
2026 if (!GEPIndex) {
2027 ScalarizableIdx.discard();
2028 return false;
2029 }
2030
2031 // We are only checking legality here. Do not mutate IR before the
2032 // profitability check, but also do not leave a pending ToFreeze behind.
2033 ScalarizableIdx.discard();
2034 }
2035
2037 Instruction::Store, SI->getValueOperand()->getType(), SI->getAlign(),
2038 SI->getPointerAddressSpace(), CostKind);
2039
2040 if (Load->hasOneUse())
2041 OldCost += TTI.getMemoryOpCost(Instruction::Load, Load->getType(),
2042 Load->getAlign(),
2043 Load->getPointerAddressSpace(), CostKind);
2044
2045 for (auto [InsertVal, Idx] : InsertElements) {
2046 int Index = -1;
2047 if (auto *CIdx = dyn_cast<ConstantInt>(Idx))
2048 Index = CIdx->getZExtValue();
2049
2050 OldCost += TTI.getVectorInstrCost(Instruction::InsertElement, VecTy,
2051 CostKind, Index);
2052 }
2053
2054 InstructionCost NewCost = 0;
2055 // This transform replaces insertelement operations on a single vector with
2056 // GEPs and scalar stores, so assume constant-index GEP offsets stay within
2057 // addressing-mode ranges that getGEPCost considers TCC_Free. Cost only GEPs
2058 // with dynamic indices.
2059 for (auto [InsertVal, Idx] : InsertElements) {
2060 if (isa<ConstantInt>(Idx))
2061 continue;
2062 const Value *GEPIndices[] = {ConstantInt::get(Idx->getType(), 0), Idx};
2063 NewCost += TTI.getGEPCost(VecTy, SI->getPointerOperand(), GEPIndices,
2064 InsertVal->getType(), CostKind);
2065 }
2066
2067 for (auto [InsertVal, Idx] : InsertElements) {
2068 Align ScalarOpAlignment = computeAlignmentAfterScalarization(
2069 std::max(SI->getAlign(), Load->getAlign()), InsertVal->getType(), Idx,
2070 *DL);
2071
2072 NewCost += TTI.getMemoryOpCost(Instruction::Store, InsertVal->getType(),
2073 ScalarOpAlignment,
2074 SI->getPointerAddressSpace(), CostKind);
2075 }
2076
2077 LLVM_DEBUG(dbgs() << "Found an insert-elements vector store scalarization "
2078 "candidate: "
2079 << I << "\n"
2080 << " NumInserts: " << InsertElements.size() << "\n"
2081 << " OldCost: " << OldCost << " vs NewCost: " << NewCost
2082 << "\n");
2083
2084 if (OldCost <= NewCost)
2085 return false;
2086
2087 for (auto [InsertVal, Idx] : InsertElements) {
2088 auto ScalarizableIdx =
2089 canScalarizeAccess(VecTy, Idx, SQ.getWithInstruction(&I));
2090 assert(!ScalarizableIdx.isUnsafe() && "already checked above");
2091
2092 if (ScalarizableIdx.isSafeWithFreeze())
2093 ScalarizableIdx.freeze(Builder, *cast<Instruction>(Idx));
2094 }
2095
2096 Worklist.push(Load);
2097 StoreInst *LastStore = nullptr;
2098 for (auto [InsertVal, Idx] : InsertElements) {
2099 auto ScalarizableIdx =
2100 canScalarizeAccess(VecTy, Idx, SQ.getWithInstruction(&I));
2101 if (ScalarizableIdx.isUnsafe())
2102 return false;
2103
2104 IntegerType *GEPIndexTy =
2105 getScalarizedGEPIndexInfo(VecTy, Idx, SI->getPointerOperandType(), *DL);
2106
2107 Value *GEPIdx = materializeScalarizedGEPIndex(Idx, GEPIndexTy, Builder);
2108 Value *GEP = Builder.CreateInBoundsGEP(
2109 SI->getValueOperand()->getType(), SI->getPointerOperand(),
2110 {ConstantInt::get(GEPIdx->getType(), 0), GEPIdx});
2111
2112 LastStore = Builder.CreateStore(InsertVal, GEP);
2113 LastStore->copyMetadata(*SI);
2114
2115 // The new GEP may change the pointer operand, so !invariant.group cannot
2116 // be transferred to the scalar store.
2117 LastStore->setMetadata(LLVMContext::MD_invariant_group, nullptr);
2118 Align ScalarOpAlignment = computeAlignmentAfterScalarization(
2119 std::max(SI->getAlign(), Load->getAlign()), InsertVal->getType(), Idx,
2120 *DL);
2121 LastStore->setAlignment(ScalarOpAlignment);
2122 }
2123
2124 replaceValue(I, *LastStore);
2126 return true;
2127}
2128
2129/// Try to scalarize vector loads feeding extractelement or bitcast
2130/// instructions.
2131bool VectorCombine::scalarizeLoad(Instruction &I) {
2132 Value *Ptr;
2133 if (!match(&I, m_Load(m_Value(Ptr))))
2134 return false;
2135
2136 auto *LI = cast<LoadInst>(&I);
2137 auto *VecTy = cast<VectorType>(LI->getType());
2138
2139 // The isSimple() check could be isUnordered(), but for now we cowardly
2140 // refuse to handle even unordered atomics.
2141 if (!LI->isSimple() || !DL->typeSizeEqualsStoreSize(VecTy->getScalarType()))
2142 return false;
2143
2144 bool AllExtracts = true;
2145 bool AllBitcasts = true;
2146 Instruction *LastCheckedInst = LI;
2147 unsigned NumInstChecked = 0;
2148
2149 // Check what type of users we have (must either all be extracts or
2150 // bitcasts) and ensure no memory modifications between the load and
2151 // its users.
2152 for (User *U : LI->users()) {
2153 auto *UI = dyn_cast<Instruction>(U);
2154 if (!UI || UI->getParent() != LI->getParent())
2155 return false;
2156
2157 // If any user is waiting to be erased, then bail out as this will
2158 // distort the cost calculation and possibly lead to infinite loops.
2159 if (UI->use_empty())
2160 return false;
2161
2162 if (!isa<ExtractElementInst>(UI))
2163 AllExtracts = false;
2164 if (!isa<BitCastInst>(UI))
2165 AllBitcasts = false;
2166
2167 // Check if any instruction between the load and the user may modify memory.
2168 if (LastCheckedInst->comesBefore(UI)) {
2169 for (Instruction &I :
2170 make_range(std::next(LI->getIterator()), UI->getIterator())) {
2171 // Bail out if we reached the check limit or the instruction may write
2172 // to memory.
2173 if (NumInstChecked == MaxInstrsToScan || I.mayWriteToMemory())
2174 return false;
2175 NumInstChecked++;
2176 }
2177 LastCheckedInst = UI;
2178 }
2179 }
2180
2181 if (AllExtracts)
2182 return scalarizeLoadExtract(LI, VecTy, Ptr);
2183 if (AllBitcasts)
2184 return scalarizeLoadBitcast(LI, VecTy, Ptr);
2185 return false;
2186}
2187
2188/// Try to scalarize vector loads feeding extractelement instructions.
2189bool VectorCombine::scalarizeLoadExtract(LoadInst *LI, VectorType *VecTy,
2190 Value *Ptr) {
2192 return false;
2193
2194 DenseMap<ExtractElementInst *, ScalarizationResult> NeedFreeze;
2195 DenseMap<ExtractElementInst *, IntegerType *> GEPIndexInfos;
2196 llvm::scope_exit FailureGuard([&]() {
2197 // If the transform is aborted, discard the ScalarizationResults.
2198 for (auto &Pair : NeedFreeze)
2199 Pair.second.discard();
2200 });
2201
2202 InstructionCost OriginalCost =
2203 TTI.getMemoryOpCost(Instruction::Load, VecTy, LI->getAlign(),
2205 InstructionCost ScalarizedCost = 0;
2206
2207 for (User *U : LI->users()) {
2208 auto *UI = cast<ExtractElementInst>(U);
2209
2210 auto ScalarIdx = canScalarizeAccess(VecTy, UI->getIndexOperand(),
2211 SQ.getWithInstruction(LI));
2212 if (ScalarIdx.isUnsafe())
2213 return false;
2214
2215 IntegerType *GEPIndex = getScalarizedGEPIndexInfo(
2216 VecTy, UI->getIndexOperand(), LI->getPointerOperandType(), *DL);
2217 if (!GEPIndex) {
2218 ScalarIdx.discard();
2219 return false;
2220 }
2221
2222 GEPIndexInfos.try_emplace(UI, GEPIndex);
2223
2224 if (ScalarIdx.isSafeWithFreeze()) {
2225 NeedFreeze.try_emplace(UI, ScalarIdx);
2226 ScalarIdx.discard();
2227 }
2228
2229 auto *Index = dyn_cast<ConstantInt>(UI->getIndexOperand());
2230 OriginalCost +=
2231 TTI.getVectorInstrCost(Instruction::ExtractElement, VecTy, CostKind,
2232 Index ? Index->getZExtValue() : -1);
2233 ScalarizedCost +=
2234 TTI.getMemoryOpCost(Instruction::Load, VecTy->getElementType(),
2236 ScalarizedCost += TTI.getAddressComputationCost(LI->getPointerOperandType(),
2237 nullptr, nullptr, CostKind);
2238 if (!Index && UI->getIndexOperand()->getType()->getIntegerBitWidth() <
2239 GEPIndex->getBitWidth())
2240 ScalarizedCost += TTI.getCastInstrCost(
2241 Instruction::ZExt, GEPIndex, UI->getIndexOperand()->getType(),
2243 }
2244
2245 LLVM_DEBUG(dbgs() << "Found all extractions of a vector load: " << *LI
2246 << "\n LoadExtractCost: " << OriginalCost
2247 << " vs ScalarizedCost: " << ScalarizedCost << "\n");
2248
2249 if (ScalarizedCost >= OriginalCost)
2250 return false;
2251
2252 // Ensure we add the load back to the worklist BEFORE its users so they can
2253 // erased in the correct order.
2254 Worklist.push(LI);
2255
2256 Type *ElemType = VecTy->getElementType();
2257
2258 // Replace extracts with narrow scalar loads.
2259 for (User *U : LI->users()) {
2260 auto *EI = cast<ExtractElementInst>(U);
2261 Value *Idx = EI->getIndexOperand();
2262
2263 // Insert 'freeze' for poison indexes.
2264 if (auto It = NeedFreeze.find(EI); It != NeedFreeze.end())
2265 It->second.freeze(Builder, *cast<Instruction>(Idx));
2266
2267 Builder.SetInsertPoint(EI);
2268 auto It = GEPIndexInfos.find(EI);
2269 assert(It != GEPIndexInfos.end() &&
2270 "Missing scalarized GEP index information");
2271 Value *GEPIdx = materializeScalarizedGEPIndex(Idx, It->second, Builder);
2272 Value *GEP = Builder.CreateInBoundsGEP(
2273 VecTy, Ptr, {ConstantInt::get(GEPIdx->getType(), 0), GEPIdx});
2274 auto *NewLoad = cast<LoadInst>(
2275 Builder.CreateLoad(ElemType, GEP, EI->getName() + ".scalar"));
2276
2277 Align ScalarOpAlignment =
2278 computeAlignmentAfterScalarization(LI->getAlign(), ElemType, Idx, *DL);
2279 NewLoad->setAlignment(ScalarOpAlignment);
2280
2281 if (auto *ConstIdx = dyn_cast<ConstantInt>(Idx)) {
2282 size_t Offset = ConstIdx->getZExtValue() * DL->getTypeStoreSize(ElemType);
2283 AAMDNodes OldAAMD = LI->getAAMetadata();
2284 NewLoad->setAAMetadata(OldAAMD.adjustForAccess(Offset, ElemType, *DL));
2285 }
2286
2287 replaceValue(*EI, *NewLoad, false);
2288 }
2289
2290 FailureGuard.release();
2291 return true;
2292}
2293
2294/// Try to scalarize vector loads feeding bitcast instructions.
2295bool VectorCombine::scalarizeLoadBitcast(LoadInst *LI, VectorType *VecTy,
2296 Value *Ptr) {
2297 InstructionCost OriginalCost =
2298 TTI.getMemoryOpCost(Instruction::Load, VecTy, LI->getAlign(),
2300
2301 Type *TargetScalarType = nullptr;
2302 unsigned VecBitWidth = DL->getTypeSizeInBits(VecTy);
2303
2304 for (User *U : LI->users()) {
2305 auto *BC = cast<BitCastInst>(U);
2306
2307 Type *DestTy = BC->getDestTy();
2308 if (!DestTy->isIntegerTy() && !DestTy->isFloatingPointTy())
2309 return false;
2310
2311 unsigned DestBitWidth = DL->getTypeSizeInBits(DestTy);
2312 if (DestBitWidth != VecBitWidth)
2313 return false;
2314
2315 // All bitcasts must target the same scalar type.
2316 if (!TargetScalarType)
2317 TargetScalarType = DestTy;
2318 else if (TargetScalarType != DestTy)
2319 return false;
2320
2321 OriginalCost +=
2322 TTI.getCastInstrCost(Instruction::BitCast, TargetScalarType, VecTy,
2324 }
2325
2326 if (!TargetScalarType)
2327 return false;
2328
2329 assert(!LI->user_empty() && "Unexpected load without bitcast users");
2330 InstructionCost ScalarizedCost =
2331 TTI.getMemoryOpCost(Instruction::Load, TargetScalarType, LI->getAlign(),
2333
2334 LLVM_DEBUG(dbgs() << "Found vector load feeding only bitcasts: " << *LI
2335 << "\n OriginalCost: " << OriginalCost
2336 << " vs ScalarizedCost: " << ScalarizedCost << "\n");
2337
2338 if (ScalarizedCost >= OriginalCost)
2339 return false;
2340
2341 // Ensure we add the load back to the worklist BEFORE its users so they can
2342 // erased in the correct order.
2343 Worklist.push(LI);
2344
2345 Builder.SetInsertPoint(LI);
2346 auto *ScalarLoad =
2347 Builder.CreateLoad(TargetScalarType, Ptr, LI->getName() + ".scalar");
2348 ScalarLoad->setAlignment(LI->getAlign());
2349 ScalarLoad->copyMetadata(*LI);
2350
2351 // Replace all bitcast users with the scalar load.
2352 for (User *U : LI->users()) {
2353 auto *BC = cast<BitCastInst>(U);
2354 replaceValue(*BC, *ScalarLoad, false);
2355 }
2356
2357 return true;
2358}
2359
2360bool VectorCombine::scalarizeExtExtract(Instruction &I) {
2362 return false;
2363 auto *Ext = dyn_cast<ZExtInst>(&I);
2364 if (!Ext)
2365 return false;
2366
2367 // Try to convert a vector zext feeding only extracts to a set of scalar
2368 // (Src << ExtIdx *Size) & (Size -1)
2369 // if profitable .
2370 auto *SrcTy = dyn_cast<FixedVectorType>(Ext->getOperand(0)->getType());
2371 if (!SrcTy)
2372 return false;
2373 auto *DstTy = cast<FixedVectorType>(Ext->getType());
2374
2375 Type *ScalarDstTy = DstTy->getElementType();
2376 if (DL->getTypeSizeInBits(SrcTy) != DL->getTypeSizeInBits(ScalarDstTy))
2377 return false;
2378
2379 InstructionCost VectorCost =
2380 TTI.getCastInstrCost(Instruction::ZExt, DstTy, SrcTy,
2382 unsigned ExtCnt = 0;
2383 bool ExtLane0 = false;
2384 for (User *U : Ext->users()) {
2385 uint64_t Idx;
2386 if (!match(U, m_ExtractElt(m_Value(), m_ConstantInt(Idx))))
2387 return false;
2388 if (cast<Instruction>(U)->use_empty())
2389 continue;
2390 ExtCnt += 1;
2391 ExtLane0 |= !Idx;
2392 VectorCost += TTI.getVectorInstrCost(Instruction::ExtractElement, DstTy,
2393 CostKind, Idx, U);
2394 }
2395
2396 InstructionCost ScalarCost =
2397 ExtCnt * TTI.getArithmeticInstrCost(
2398 Instruction::And, ScalarDstTy, CostKind,
2401 (ExtCnt - ExtLane0) *
2403 Instruction::LShr, ScalarDstTy, CostKind,
2406 if (ScalarCost > VectorCost)
2407 return false;
2408
2409 Value *ScalarV = Ext->getOperand(0);
2410 if (!isGuaranteedNotToBePoison(ScalarV, SQ.AC, dyn_cast<Instruction>(ScalarV),
2411 SQ.DT)) {
2412 // Check wether all lanes are extracted, all extracts trigger UB
2413 // on poison, and the last extract (and hence all previous ones)
2414 // are guaranteed to execute if Ext executes. If so, we do not
2415 // need to insert a freeze.
2416 SmallDenseSet<ConstantInt *, 8> ExtractedLanes;
2417 bool AllExtractsTriggerUB = true;
2418 ExtractElementInst *LastExtract = nullptr;
2419 BasicBlock *ExtBB = Ext->getParent();
2420 for (User *U : Ext->users()) {
2421 auto *Extract = cast<ExtractElementInst>(U);
2422 if (Extract->getParent() != ExtBB || !programUndefinedIfPoison(Extract)) {
2423 AllExtractsTriggerUB = false;
2424 break;
2425 }
2426 ExtractedLanes.insert(cast<ConstantInt>(Extract->getIndexOperand()));
2427 if (!LastExtract || LastExtract->comesBefore(Extract))
2428 LastExtract = Extract;
2429 }
2430 if (ExtractedLanes.size() != DstTy->getNumElements() ||
2431 !AllExtractsTriggerUB ||
2433 LastExtract->getIterator()))
2434 ScalarV = Builder.CreateFreeze(ScalarV);
2435 }
2436 ScalarV = Builder.CreateBitCast(
2437 ScalarV,
2438 IntegerType::get(SrcTy->getContext(), DL->getTypeSizeInBits(SrcTy)));
2439 uint64_t SrcEltSizeInBits = DL->getTypeSizeInBits(SrcTy->getElementType());
2440 uint64_t TotalBits = DL->getTypeSizeInBits(SrcTy);
2441 APInt EltBitMask = APInt::getLowBitsSet(TotalBits, SrcEltSizeInBits);
2442 Type *PackedTy = IntegerType::get(SrcTy->getContext(), TotalBits);
2443 Value *Mask = ConstantInt::get(PackedTy, EltBitMask);
2444 for (User *U : Ext->users()) {
2445 auto *Extract = cast<ExtractElementInst>(U);
2446 uint64_t Idx =
2447 cast<ConstantInt>(Extract->getIndexOperand())->getZExtValue();
2448 uint64_t ShiftAmt =
2449 DL->isBigEndian()
2450 ? (TotalBits - SrcEltSizeInBits - Idx * SrcEltSizeInBits)
2451 : (Idx * SrcEltSizeInBits);
2452 Value *LShr = Builder.CreateLShr(ScalarV, ShiftAmt);
2453 Value *And = Builder.CreateAnd(LShr, Mask);
2454 U->replaceAllUsesWith(And);
2455 }
2456 return true;
2457}
2458
2459/// Try to fold "(or (zext (bitcast X)), (shl (zext (bitcast Y)), C))"
2460/// to "(bitcast (concat X, Y))"
2461/// where X/Y are bitcasted from i1 mask vectors.
2462bool VectorCombine::foldConcatOfBoolMasks(Instruction &I) {
2463 Type *Ty = I.getType();
2464 if (!Ty->isIntegerTy())
2465 return false;
2466
2467 // TODO: Add big endian test coverage
2468 if (DL->isBigEndian())
2469 return false;
2470
2471 // Restrict to disjoint cases so the mask vectors aren't overlapping.
2472 Instruction *X, *Y;
2474 return false;
2475
2476 // Allow both sources to contain shl, to handle more generic pattern:
2477 // "(or (shl (zext (bitcast X)), C1), (shl (zext (bitcast Y)), C2))"
2478 Value *SrcX;
2479 uint64_t ShAmtX = 0;
2480 if (!match(X, m_OneUse(m_ZExt(m_OneUse(m_BitCast(m_Value(SrcX)))))) &&
2481 !match(X, m_OneUse(
2483 m_ConstantInt(ShAmtX)))))
2484 return false;
2485
2486 Value *SrcY;
2487 uint64_t ShAmtY = 0;
2488 if (!match(Y, m_OneUse(m_ZExt(m_OneUse(m_BitCast(m_Value(SrcY)))))) &&
2489 !match(Y, m_OneUse(
2491 m_ConstantInt(ShAmtY)))))
2492 return false;
2493
2494 // Canonicalize larger shift to the RHS.
2495 if (ShAmtX > ShAmtY) {
2496 std::swap(X, Y);
2497 std::swap(SrcX, SrcY);
2498 std::swap(ShAmtX, ShAmtY);
2499 }
2500
2501 // Ensure both sources are matching vXi1 bool mask types, and that the shift
2502 // difference is the mask width so they can be easily concatenated together.
2503 uint64_t ShAmtDiff = ShAmtY - ShAmtX;
2504 unsigned NumSHL = (ShAmtX > 0) + (ShAmtY > 0);
2505 unsigned BitWidth = Ty->getPrimitiveSizeInBits();
2506 auto *MaskTy = dyn_cast<FixedVectorType>(SrcX->getType());
2507 if (!MaskTy || SrcX->getType() != SrcY->getType() ||
2508 !MaskTy->getElementType()->isIntegerTy(1) ||
2509 MaskTy->getNumElements() != ShAmtDiff ||
2510 MaskTy->getNumElements() > (BitWidth / 2))
2511 return false;
2512
2513 auto *ConcatTy = FixedVectorType::getDoubleElementsVectorType(MaskTy);
2514 auto *ConcatIntTy =
2515 Type::getIntNTy(Ty->getContext(), ConcatTy->getNumElements());
2516 auto *MaskIntTy = Type::getIntNTy(Ty->getContext(), ShAmtDiff);
2517
2518 SmallVector<int, 32> ConcatMask(ConcatTy->getNumElements());
2519 std::iota(ConcatMask.begin(), ConcatMask.end(), 0);
2520
2521 // TODO: Is it worth supporting multi use cases?
2522 InstructionCost OldCost = 0;
2523 OldCost += TTI.getArithmeticInstrCost(Instruction::Or, Ty, CostKind);
2524 OldCost +=
2525 NumSHL * TTI.getArithmeticInstrCost(Instruction::Shl, Ty, CostKind);
2526 OldCost += 2 * TTI.getCastInstrCost(Instruction::ZExt, Ty, MaskIntTy,
2528 OldCost += 2 * TTI.getCastInstrCost(Instruction::BitCast, MaskIntTy, MaskTy,
2530
2531 InstructionCost NewCost = 0;
2533 MaskTy, ConcatMask, CostKind);
2534 NewCost += TTI.getCastInstrCost(Instruction::BitCast, ConcatIntTy, ConcatTy,
2536 if (Ty != ConcatIntTy)
2537 NewCost += TTI.getCastInstrCost(Instruction::ZExt, Ty, ConcatIntTy,
2539 if (ShAmtX > 0)
2540 NewCost += TTI.getArithmeticInstrCost(Instruction::Shl, Ty, CostKind);
2541
2542 LLVM_DEBUG(dbgs() << "Found a concatenation of bitcasted bool masks: " << I
2543 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
2544 << "\n");
2545
2546 if (NewCost > OldCost)
2547 return false;
2548
2549 // Build bool mask concatenation, bitcast back to scalar integer, and perform
2550 // any residual zero-extension or shifting.
2551 Value *Concat = Builder.CreateShuffleVector(SrcX, SrcY, ConcatMask);
2552 Worklist.pushValue(Concat);
2553
2554 Value *Result = Builder.CreateBitCast(Concat, ConcatIntTy);
2555
2556 if (Ty != ConcatIntTy) {
2557 Worklist.pushValue(Result);
2558 Result = Builder.CreateZExt(Result, Ty);
2559 }
2560
2561 if (ShAmtX > 0) {
2562 Worklist.pushValue(Result);
2563 Result = Builder.CreateShl(Result, ShAmtX);
2564 }
2565
2566 replaceValue(I, *Result);
2567 return true;
2568}
2569
2570/// Try to convert "shuffle (binop (shuffle, shuffle)), undef"
2571/// --> "binop (shuffle), (shuffle)".
2572bool VectorCombine::foldPermuteOfBinops(Instruction &I) {
2573 BinaryOperator *BinOp;
2574 ArrayRef<int> OuterMask;
2575 if (!match(&I, m_Shuffle(m_BinOp(BinOp), m_Undef(), m_Mask(OuterMask))))
2576 return false;
2577
2578 // Don't introduce poison into div/rem.
2579 if (BinOp->isIntDivRem() && llvm::is_contained(OuterMask, PoisonMaskElem))
2580 return false;
2581
2582 Value *Op00, *Op01, *Op10, *Op11;
2583 ArrayRef<int> Mask0, Mask1;
2584 bool Match0 = match(BinOp->getOperand(0),
2585 m_Shuffle(m_Value(Op00), m_Value(Op01), m_Mask(Mask0)));
2586 bool Match1 = match(BinOp->getOperand(1),
2587 m_Shuffle(m_Value(Op10), m_Value(Op11), m_Mask(Mask1)));
2588 if (!Match0 && !Match1)
2589 return false;
2590
2591 Op00 = Match0 ? Op00 : BinOp->getOperand(0);
2592 Op01 = Match0 ? Op01 : BinOp->getOperand(0);
2593 Op10 = Match1 ? Op10 : BinOp->getOperand(1);
2594 Op11 = Match1 ? Op11 : BinOp->getOperand(1);
2595
2596 Instruction::BinaryOps Opcode = BinOp->getOpcode();
2597 auto *ShuffleDstTy = dyn_cast<FixedVectorType>(I.getType());
2598 auto *BinOpTy = dyn_cast<FixedVectorType>(BinOp->getType());
2599 auto *Op0Ty = dyn_cast<FixedVectorType>(Op00->getType());
2600 auto *Op1Ty = dyn_cast<FixedVectorType>(Op10->getType());
2601 if (!ShuffleDstTy || !BinOpTy || !Op0Ty || !Op1Ty)
2602 return false;
2603
2604 unsigned NumSrcElts = BinOpTy->getNumElements();
2605
2606 // Don't accept shuffles that reference the second operand in
2607 // div/rem or if its an undef arg.
2608 if ((BinOp->isIntDivRem() || !isa<PoisonValue>(I.getOperand(1))) &&
2609 any_of(OuterMask, [NumSrcElts](int M) { return M >= (int)NumSrcElts; }))
2610 return false;
2611
2612 // Merge outer / inner (or identity if no match) shuffles.
2613 SmallVector<int> NewMask0, NewMask1;
2614 for (int M : OuterMask) {
2615 if (M < 0 || M >= (int)NumSrcElts) {
2616 NewMask0.push_back(PoisonMaskElem);
2617 NewMask1.push_back(PoisonMaskElem);
2618 } else {
2619 NewMask0.push_back(Match0 ? Mask0[M] : M);
2620 NewMask1.push_back(Match1 ? Mask1[M] : M);
2621 }
2622 }
2623
2624 unsigned NumOpElts = Op0Ty->getNumElements();
2625 bool IsIdentity0 = ShuffleDstTy == Op0Ty &&
2626 all_of(NewMask0, [NumOpElts](int M) { return M < (int)NumOpElts; }) &&
2627 ShuffleVectorInst::isIdentityMask(NewMask0, NumOpElts);
2628 bool IsIdentity1 = ShuffleDstTy == Op1Ty &&
2629 all_of(NewMask1, [NumOpElts](int M) { return M < (int)NumOpElts; }) &&
2630 ShuffleVectorInst::isIdentityMask(NewMask1, NumOpElts);
2631
2632 InstructionCost NewCost = 0;
2633 // Try to merge shuffles across the binop if the new shuffles are not costly.
2634 InstructionCost BinOpCost =
2635 TTI.getArithmeticInstrCost(Opcode, BinOpTy, CostKind);
2636 InstructionCost OldCost =
2638 ShuffleDstTy, BinOpTy, OuterMask, CostKind,
2639 0, nullptr, {BinOp}, &I);
2640 if (!BinOp->hasOneUse())
2641 NewCost += BinOpCost;
2642
2643 if (Match0) {
2645 TargetTransformInfo::SK_PermuteTwoSrc, BinOpTy, Op0Ty, Mask0, CostKind,
2646 0, nullptr, {Op00, Op01}, cast<Instruction>(BinOp->getOperand(0)));
2647 OldCost += Shuf0Cost;
2648 if (!BinOp->hasOneUse() || !BinOp->getOperand(0)->hasOneUse())
2649 NewCost += Shuf0Cost;
2650 }
2651 if (Match1) {
2653 TargetTransformInfo::SK_PermuteTwoSrc, BinOpTy, Op1Ty, Mask1, CostKind,
2654 0, nullptr, {Op10, Op11}, cast<Instruction>(BinOp->getOperand(1)));
2655 OldCost += Shuf1Cost;
2656 if (!BinOp->hasOneUse() || !BinOp->getOperand(1)->hasOneUse())
2657 NewCost += Shuf1Cost;
2658 }
2659
2660 NewCost += TTI.getArithmeticInstrCost(Opcode, ShuffleDstTy, CostKind);
2661
2662 if (!IsIdentity0)
2663 NewCost +=
2665 Op0Ty, NewMask0, CostKind, 0, nullptr, {Op00, Op01});
2666 if (!IsIdentity1)
2667 NewCost +=
2669 Op1Ty, NewMask1, CostKind, 0, nullptr, {Op10, Op11});
2670
2671 LLVM_DEBUG(dbgs() << "Found a shuffle feeding a shuffled binop: " << I
2672 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
2673 << "\n");
2674
2675 // If costs are equal, still fold as we reduce instruction count.
2676 if (NewCost > OldCost)
2677 return false;
2678
2679 Value *LHS =
2680 IsIdentity0 ? Op00 : Builder.CreateShuffleVector(Op00, Op01, NewMask0);
2681 Value *RHS =
2682 IsIdentity1 ? Op10 : Builder.CreateShuffleVector(Op10, Op11, NewMask1);
2683 Value *NewBO = Builder.CreateBinOp(Opcode, LHS, RHS);
2684
2685 // Intersect flags from the old binops.
2686 if (auto *NewInst = dyn_cast<Instruction>(NewBO))
2687 NewInst->copyIRFlags(BinOp);
2688
2689 Worklist.pushValue(LHS);
2690 Worklist.pushValue(RHS);
2691 replaceValue(I, *NewBO);
2692 return true;
2693}
2694
2695/// Try to convert "shuffle (binop), (binop)" into "binop (shuffle), (shuffle)".
2696/// Try to convert "shuffle (cmpop), (cmpop)" into "cmpop (shuffle), (shuffle)".
2697bool VectorCombine::foldShuffleOfBinops(Instruction &I) {
2698 ArrayRef<int> OldMask;
2699 Instruction *LHS, *RHS;
2701 m_Mask(OldMask))))
2702 return false;
2703
2704 // TODO: Add support for addlike etc.
2705 if (LHS->getOpcode() != RHS->getOpcode())
2706 return false;
2707
2708 Value *X, *Y, *Z, *W;
2709 bool IsCommutative = false;
2710 CmpPredicate PredLHS = CmpInst::BAD_ICMP_PREDICATE;
2711 CmpPredicate PredRHS = CmpInst::BAD_ICMP_PREDICATE;
2712 if (match(LHS, m_BinOp(m_Value(X), m_Value(Y))) &&
2713 match(RHS, m_BinOp(m_Value(Z), m_Value(W)))) {
2714 auto *BO = cast<BinaryOperator>(LHS);
2715 // Don't introduce poison into div/rem.
2716 if (llvm::is_contained(OldMask, PoisonMaskElem) && BO->isIntDivRem())
2717 return false;
2718 IsCommutative = BinaryOperator::isCommutative(BO->getOpcode());
2719 } else if (match(LHS, m_Cmp(PredLHS, m_Value(X), m_Value(Y))) &&
2720 match(RHS, m_Cmp(PredRHS, m_Value(Z), m_Value(W))) &&
2721 (CmpInst::Predicate)PredLHS == (CmpInst::Predicate)PredRHS) {
2722 IsCommutative = cast<CmpInst>(LHS)->isCommutative();
2723 } else
2724 return false;
2725
2726 auto *ShuffleDstTy = dyn_cast<FixedVectorType>(I.getType());
2727 auto *BinResTy = dyn_cast<FixedVectorType>(LHS->getType());
2728 auto *BinOpTy = dyn_cast<FixedVectorType>(X->getType());
2729 if (!ShuffleDstTy || !BinResTy || !BinOpTy || X->getType() != Z->getType())
2730 return false;
2731
2732 bool SameBinOp = LHS == RHS;
2733 unsigned NumSrcElts = BinOpTy->getNumElements();
2734
2735 // If we have something like "add X, Y" and "add Z, X", swap ops to match.
2736 if (IsCommutative && X != Z && Y != W && (X == W || Y == Z))
2737 std::swap(X, Y);
2738
2739 auto ConvertToUnary = [NumSrcElts](int &M) {
2740 if (M >= (int)NumSrcElts)
2741 M -= NumSrcElts;
2742 };
2743
2744 SmallVector<int> NewMask0(OldMask);
2746 TTI::OperandValueInfo Op0Info = TTI.commonOperandInfo(X, Z);
2747 if (X == Z) {
2748 llvm::for_each(NewMask0, ConvertToUnary);
2750 Z = PoisonValue::get(BinOpTy);
2751 }
2752
2753 SmallVector<int> NewMask1(OldMask);
2755 TTI::OperandValueInfo Op1Info = TTI.commonOperandInfo(Y, W);
2756 if (Y == W) {
2757 llvm::for_each(NewMask1, ConvertToUnary);
2759 W = PoisonValue::get(BinOpTy);
2760 }
2761
2762 // Try to replace a binop with a shuffle if the shuffle is not costly.
2763 // When SameBinOp, only count the binop cost once.
2766
2767 InstructionCost OldCost = LHSCost;
2768 if (!SameBinOp) {
2769 OldCost += RHSCost;
2770 }
2772 ShuffleDstTy, BinResTy, OldMask, CostKind, 0,
2773 nullptr, {LHS, RHS}, &I);
2774
2775 // Handle shuffle(binop(shuffle(x),y),binop(z,shuffle(w))) style patterns
2776 // where one use shuffles have gotten split across the binop/cmp. These
2777 // often allow a major reduction in total cost that wouldn't happen as
2778 // individual folds.
2779 auto MergeInner = [&](Value *&Op, int Offset, MutableArrayRef<int> Mask,
2780 TTI::TargetCostKind CostKind) -> bool {
2781 Value *InnerOp;
2782 ArrayRef<int> InnerMask;
2783 if (match(Op, m_OneUse(m_Shuffle(m_Value(InnerOp), m_Undef(),
2784 m_Mask(InnerMask)))) &&
2785 InnerOp->getType() == Op->getType() &&
2786 all_of(InnerMask,
2787 [NumSrcElts](int M) { return M < (int)NumSrcElts; })) {
2788 for (int &M : Mask)
2789 if (Offset <= M && M < (int)(Offset + NumSrcElts)) {
2790 M = InnerMask[M - Offset];
2791 M = 0 <= M ? M + Offset : M;
2792 }
2794 Op = InnerOp;
2795 return true;
2796 }
2797 return false;
2798 };
2799 bool ReducedInstCount = false;
2800 ReducedInstCount |= MergeInner(X, 0, NewMask0, CostKind);
2801 ReducedInstCount |= MergeInner(Y, 0, NewMask1, CostKind);
2802 ReducedInstCount |= MergeInner(Z, NumSrcElts, NewMask0, CostKind);
2803 ReducedInstCount |= MergeInner(W, NumSrcElts, NewMask1, CostKind);
2804 bool SingleSrcBinOp = (X == Y) && (Z == W) && (NewMask0 == NewMask1);
2805 // SingleSrcBinOp only reduces instruction count if we also eliminate the
2806 // original binop(s). If binops have multiple uses, they won't be eliminated.
2807 ReducedInstCount |= SingleSrcBinOp && LHS->hasOneUser() && RHS->hasOneUser();
2808
2809 // For concat shuffles of i1 vectors where both binops are one-use, the
2810 // transform keeps the same instruction count but canonicalises to a single
2811 // wider binop, enabling downstream folds (e.g. NOT(XOR(concat(a,b),
2812 // concat(c,d))) -> XNOR(concat(a,b),concat(c,d)) on AVX-512 mask regs).
2813 // Restrict to BinaryOperator (not CmpInst) since narrow comparisons may
2814 // be cheaper than wide ones on some targets (e.g. AVX-512 vpcmpeq).
2815 ReducedInstCount |= cast<ShuffleVectorInst>(&I)->isConcat() &&
2816 I.getType()->getScalarType()->isIntegerTy(1) &&
2818 RHS->hasOneUser();
2819
2820 auto *ShuffleCmpTy =
2821 FixedVectorType::get(BinOpTy->getElementType(), ShuffleDstTy);
2823 SK0, ShuffleCmpTy, BinOpTy, NewMask0, CostKind, 0, nullptr, {X, Z});
2824 if (!SingleSrcBinOp)
2825 NewCost += TTI.getShuffleCost(SK1, ShuffleCmpTy, BinOpTy, NewMask1,
2826 CostKind, 0, nullptr, {Y, W});
2827
2828 if (PredLHS == CmpInst::BAD_ICMP_PREDICATE) {
2829 NewCost += TTI.getArithmeticInstrCost(LHS->getOpcode(), ShuffleDstTy,
2830 CostKind, Op0Info, Op1Info);
2831 } else {
2832 NewCost +=
2833 TTI.getCmpSelInstrCost(LHS->getOpcode(), ShuffleCmpTy, ShuffleDstTy,
2834 PredLHS, CostKind, Op0Info, Op1Info);
2835 }
2836 // If LHS/RHS have other uses, we need to account for the cost of keeping
2837 // the original instructions. When SameBinOp, only add the cost once.
2838 if (!LHS->hasOneUser())
2839 NewCost += LHSCost;
2840 if (!SameBinOp && !RHS->hasOneUser())
2841 NewCost += RHSCost;
2842
2843 LLVM_DEBUG(dbgs() << "Found a shuffle feeding two binops: " << I
2844 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
2845 << "\n");
2846
2847 // If either shuffle will constant fold away, then fold for the same cost as
2848 // we will reduce the instruction count.
2849 ReducedInstCount |= (isa<Constant>(X) && isa<Constant>(Z)) ||
2850 (isa<Constant>(Y) && isa<Constant>(W));
2851 if (ReducedInstCount ? (NewCost > OldCost) : (NewCost >= OldCost))
2852 return false;
2853
2854 Value *Shuf0 = Builder.CreateShuffleVector(X, Z, NewMask0);
2855 Value *Shuf1 =
2856 SingleSrcBinOp ? Shuf0 : Builder.CreateShuffleVector(Y, W, NewMask1);
2857 Value *NewBO = PredLHS == CmpInst::BAD_ICMP_PREDICATE
2858 ? Builder.CreateBinOp(
2859 cast<BinaryOperator>(LHS)->getOpcode(), Shuf0, Shuf1)
2860 : Builder.CreateCmp(PredLHS, Shuf0, Shuf1);
2861
2862 // Intersect flags from the old binops.
2863 if (auto *NewInst = dyn_cast<Instruction>(NewBO)) {
2864 NewInst->copyIRFlags(LHS);
2865 NewInst->andIRFlags(RHS);
2866 }
2867
2868 Worklist.pushValue(Shuf0);
2869 Worklist.pushValue(Shuf1);
2870 replaceValue(I, *NewBO);
2871 return true;
2872}
2873
2874/// Try to convert,
2875/// (shuffle(select(c1,t1,f1)), (select(c2,t2,f2)), m) into
2876/// (select (shuffle c1,c2,m), (shuffle t1,t2,m), (shuffle f1,f2,m))
2877bool VectorCombine::foldShuffleOfSelects(Instruction &I) {
2878 ArrayRef<int> Mask;
2879 Value *C1, *T1, *F1, *C2, *T2, *F2;
2880 if (!match(&I, m_Shuffle(m_Select(m_Value(C1), m_Value(T1), m_Value(F1)),
2881 m_Select(m_Value(C2), m_Value(T2), m_Value(F2)),
2882 m_Mask(Mask))))
2883 return false;
2884
2885 auto *Sel1 = cast<Instruction>(I.getOperand(0));
2886 auto *Sel2 = cast<Instruction>(I.getOperand(1));
2887
2888 auto *C1VecTy = dyn_cast<FixedVectorType>(C1->getType());
2889 auto *C2VecTy = dyn_cast<FixedVectorType>(C2->getType());
2890 if (!C1VecTy || !C2VecTy || C1VecTy != C2VecTy)
2891 return false;
2892
2893 auto *SI0FOp = dyn_cast<FPMathOperator>(I.getOperand(0));
2894 auto *SI1FOp = dyn_cast<FPMathOperator>(I.getOperand(1));
2895 // SelectInsts must have the same FMF.
2896 if (((SI0FOp == nullptr) != (SI1FOp == nullptr)) ||
2897 ((SI0FOp != nullptr) &&
2898 (SI0FOp->getFastMathFlags() != SI1FOp->getFastMathFlags())))
2899 return false;
2900
2901 auto *SrcVecTy = cast<FixedVectorType>(T1->getType());
2902 auto *DstVecTy = cast<FixedVectorType>(I.getType());
2904 auto SelOp = Instruction::Select;
2905
2907 SelOp, SrcVecTy, C1VecTy, CmpInst::BAD_ICMP_PREDICATE, CostKind);
2909 SelOp, SrcVecTy, C2VecTy, CmpInst::BAD_ICMP_PREDICATE, CostKind);
2910
2911 InstructionCost OldCost =
2912 CostSel1 + CostSel2 +
2913 TTI.getShuffleCost(SK, DstVecTy, SrcVecTy, Mask, CostKind, 0, nullptr,
2914 {I.getOperand(0), I.getOperand(1)}, &I);
2915
2917 SK, FixedVectorType::get(C1VecTy->getScalarType(), Mask.size()), C1VecTy,
2918 Mask, CostKind, 0, nullptr, {C1, C2});
2919 NewCost += TTI.getShuffleCost(SK, DstVecTy, SrcVecTy, Mask, CostKind, 0,
2920 nullptr, {T1, T2});
2921 NewCost += TTI.getShuffleCost(SK, DstVecTy, SrcVecTy, Mask, CostKind, 0,
2922 nullptr, {F1, F2});
2923 auto *C1C2ShuffledVecTy = FixedVectorType::get(
2924 Type::getInt1Ty(I.getContext()), DstVecTy->getNumElements());
2925 NewCost += TTI.getCmpSelInstrCost(SelOp, DstVecTy, C1C2ShuffledVecTy,
2927
2928 if (!Sel1->hasOneUse())
2929 NewCost += CostSel1;
2930 if (!Sel2->hasOneUse())
2931 NewCost += CostSel2;
2932
2933 LLVM_DEBUG(dbgs() << "Found a shuffle feeding two selects: " << I
2934 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
2935 << "\n");
2936 if (NewCost > OldCost)
2937 return false;
2938
2939 Value *ShuffleCmp = Builder.CreateShuffleVector(C1, C2, Mask);
2940 Value *ShuffleTrue = Builder.CreateShuffleVector(T1, T2, Mask);
2941 Value *ShuffleFalse = Builder.CreateShuffleVector(F1, F2, Mask);
2942 Value *NewSel;
2943 // We presuppose that the SelectInsts have the same FMF.
2944 if (SI0FOp)
2945 NewSel = Builder.CreateSelectFMF(ShuffleCmp, ShuffleTrue, ShuffleFalse,
2946 SI0FOp->getFastMathFlags());
2947 else
2948 NewSel = Builder.CreateSelect(ShuffleCmp, ShuffleTrue, ShuffleFalse);
2949
2950 Worklist.pushValue(ShuffleCmp);
2951 Worklist.pushValue(ShuffleTrue);
2952 Worklist.pushValue(ShuffleFalse);
2953 replaceValue(I, *NewSel);
2954 return true;
2955}
2956
2957/// Try to convert "shuffle (castop), (castop)" with a shared castop operand
2958/// into "castop (shuffle)".
2959bool VectorCombine::foldShuffleOfCastops(Instruction &I) {
2960 Value *V0, *V1;
2961 ArrayRef<int> OldMask;
2962 if (!match(&I, m_Shuffle(m_Value(V0), m_Value(V1), m_Mask(OldMask))))
2963 return false;
2964
2965 // Check whether this is a binary shuffle.
2966 bool IsBinaryShuffle = !isa<UndefValue>(V1);
2967
2968 auto *C0 = dyn_cast<CastInst>(V0);
2969 auto *C1 = dyn_cast<CastInst>(V1);
2970 if (!C0 || (IsBinaryShuffle && !C1))
2971 return false;
2972
2973 Instruction::CastOps Opcode = C0->getOpcode();
2974
2975 // If this is allowed, foldShuffleOfCastops can get stuck in a loop
2976 // with foldBitcastOfShuffle. Reject in favor of foldBitcastOfShuffle.
2977 if (!IsBinaryShuffle && Opcode == Instruction::BitCast)
2978 return false;
2979
2980 if (IsBinaryShuffle) {
2981 if (C0->getSrcTy() != C1->getSrcTy())
2982 return false;
2983 // Handle shuffle(zext_nneg(x), sext(y)) -> sext(shuffle(x,y)) folds.
2984 if (Opcode != C1->getOpcode()) {
2985 if (match(C0, m_SExtLike(m_Value())) && match(C1, m_SExtLike(m_Value())))
2986 Opcode = Instruction::SExt;
2987 else
2988 return false;
2989 }
2990 }
2991
2992 auto *ShuffleDstTy = dyn_cast<FixedVectorType>(I.getType());
2993 auto *CastDstTy = dyn_cast<FixedVectorType>(C0->getDestTy());
2994 auto *CastSrcTy = dyn_cast<FixedVectorType>(C0->getSrcTy());
2995 if (!ShuffleDstTy || !CastDstTy || !CastSrcTy)
2996 return false;
2997
2998 unsigned NumSrcElts = CastSrcTy->getNumElements();
2999 unsigned NumDstElts = CastDstTy->getNumElements();
3000 assert((NumDstElts == NumSrcElts || Opcode == Instruction::BitCast) &&
3001 "Only bitcasts expected to alter src/dst element counts");
3002
3003 // Check for bitcasting of unscalable vector types.
3004 // e.g. <32 x i40> -> <40 x i32>
3005 if (NumDstElts != NumSrcElts && (NumSrcElts % NumDstElts) != 0 &&
3006 (NumDstElts % NumSrcElts) != 0)
3007 return false;
3008
3009 SmallVector<int, 16> NewMask;
3010 if (NumSrcElts >= NumDstElts) {
3011 // The bitcast is from wide to narrow/equal elements. The shuffle mask can
3012 // always be expanded to the equivalent form choosing narrower elements.
3013 assert(NumSrcElts % NumDstElts == 0 && "Unexpected shuffle mask");
3014 unsigned ScaleFactor = NumSrcElts / NumDstElts;
3015 narrowShuffleMaskElts(ScaleFactor, OldMask, NewMask);
3016 } else {
3017 // The bitcast is from narrow elements to wide elements. The shuffle mask
3018 // must choose consecutive elements to allow casting first.
3019 assert(NumDstElts % NumSrcElts == 0 && "Unexpected shuffle mask");
3020 unsigned ScaleFactor = NumDstElts / NumSrcElts;
3021 if (!widenShuffleMaskElts(ScaleFactor, OldMask, NewMask))
3022 return false;
3023 }
3024
3025 auto *NewShuffleDstTy =
3026 FixedVectorType::get(CastSrcTy->getScalarType(), NewMask.size());
3027
3028 // Try to replace a castop with a shuffle if the shuffle is not costly.
3029 InstructionCost CostC0 =
3030 TTI.getCastInstrCost(C0->getOpcode(), CastDstTy, CastSrcTy,
3032
3034 if (IsBinaryShuffle)
3036 else
3038
3039 InstructionCost OldCost = CostC0;
3040 OldCost += TTI.getShuffleCost(ShuffleKind, ShuffleDstTy, CastDstTy, OldMask,
3041 CostKind, 0, nullptr, {}, &I);
3042
3043 InstructionCost NewCost = TTI.getShuffleCost(ShuffleKind, NewShuffleDstTy,
3044 CastSrcTy, NewMask, CostKind);
3045 NewCost += TTI.getCastInstrCost(Opcode, ShuffleDstTy, NewShuffleDstTy,
3047 if (!C0->hasOneUse())
3048 NewCost += CostC0;
3049 if (IsBinaryShuffle) {
3050 InstructionCost CostC1 =
3051 TTI.getCastInstrCost(C1->getOpcode(), CastDstTy, CastSrcTy,
3053 OldCost += CostC1;
3054 if (!C1->hasOneUse())
3055 NewCost += CostC1;
3056 }
3057
3058 LLVM_DEBUG(dbgs() << "Found a shuffle feeding two casts: " << I
3059 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
3060 << "\n");
3061 if (NewCost > OldCost)
3062 return false;
3063
3064 Value *Shuf;
3065 if (IsBinaryShuffle)
3066 Shuf = Builder.CreateShuffleVector(C0->getOperand(0), C1->getOperand(0),
3067 NewMask);
3068 else
3069 Shuf = Builder.CreateShuffleVector(C0->getOperand(0), NewMask);
3070
3071 Value *Cast = Builder.CreateCast(Opcode, Shuf, ShuffleDstTy);
3072
3073 // Intersect flags from the old casts.
3074 if (auto *NewInst = dyn_cast<Instruction>(Cast)) {
3075 NewInst->copyIRFlags(C0);
3076 if (IsBinaryShuffle)
3077 NewInst->andIRFlags(C1);
3078 }
3079
3080 Worklist.pushValue(Shuf);
3081 replaceValue(I, *Cast);
3082 return true;
3083}
3084
3085/// Try to convert any of:
3086/// "shuffle (shuffle x, y), (shuffle y, x)"
3087/// "shuffle (shuffle x, undef), (shuffle y, undef)"
3088/// "shuffle (shuffle x, undef), y"
3089/// "shuffle x, (shuffle y, undef)"
3090/// into "shuffle x, y".
3091bool VectorCombine::foldShuffleOfShuffles(Instruction &I) {
3092 ArrayRef<int> OuterMask;
3093 Value *OuterV0, *OuterV1;
3094 if (!match(&I,
3095 m_Shuffle(m_Value(OuterV0), m_Value(OuterV1), m_Mask(OuterMask))))
3096 return false;
3097
3098 ArrayRef<int> InnerMask0, InnerMask1;
3099 Value *X0, *X1, *Y0, *Y1;
3100 bool Match0 =
3101 match(OuterV0, m_Shuffle(m_Value(X0), m_Value(Y0), m_Mask(InnerMask0)));
3102 bool Match1 =
3103 match(OuterV1, m_Shuffle(m_Value(X1), m_Value(Y1), m_Mask(InnerMask1)));
3104 if (!Match0 && !Match1)
3105 return false;
3106
3107 // If the outer shuffle is a permute, then create a fake inner all-poison
3108 // shuffle. This is easier than accounting for length-changing shuffles below.
3109 SmallVector<int, 16> PoisonMask1;
3110 if (!Match1 && isa<PoisonValue>(OuterV1)) {
3111 X1 = X0;
3112 Y1 = Y0;
3113 PoisonMask1.append(InnerMask0.size(), PoisonMaskElem);
3114 InnerMask1 = PoisonMask1;
3115 Match1 = true; // fake match
3116 }
3117
3118 X0 = Match0 ? X0 : OuterV0;
3119 Y0 = Match0 ? Y0 : OuterV0;
3120 X1 = Match1 ? X1 : OuterV1;
3121 Y1 = Match1 ? Y1 : OuterV1;
3122 auto *ShuffleDstTy = dyn_cast<FixedVectorType>(I.getType());
3123 auto *ShuffleSrcTy = dyn_cast<FixedVectorType>(X0->getType());
3124 auto *ShuffleImmTy = dyn_cast<FixedVectorType>(OuterV0->getType());
3125 if (!ShuffleDstTy || !ShuffleSrcTy || !ShuffleImmTy ||
3126 X0->getType() != X1->getType())
3127 return false;
3128
3129 unsigned NumSrcElts = ShuffleSrcTy->getNumElements();
3130 unsigned NumImmElts = ShuffleImmTy->getNumElements();
3131
3132 // Attempt to merge shuffles, matching upto 2 source operands.
3133 // Replace index to a poison arg with PoisonMaskElem.
3134 // Bail if either inner masks reference an undef arg.
3135 SmallVector<int, 16> NewMask(OuterMask);
3136 Value *NewX = nullptr, *NewY = nullptr;
3137 for (int &M : NewMask) {
3138 Value *Src = nullptr;
3139 if (0 <= M && M < (int)NumImmElts) {
3140 Src = OuterV0;
3141 if (Match0) {
3142 M = InnerMask0[M];
3143 Src = M >= (int)NumSrcElts ? Y0 : X0;
3144 M = M >= (int)NumSrcElts ? (M - NumSrcElts) : M;
3145 }
3146 } else if (M >= (int)NumImmElts) {
3147 Src = OuterV1;
3148 M -= NumImmElts;
3149 if (Match1) {
3150 M = InnerMask1[M];
3151 Src = M >= (int)NumSrcElts ? Y1 : X1;
3152 M = M >= (int)NumSrcElts ? (M - NumSrcElts) : M;
3153 }
3154 }
3155 if (Src && M != PoisonMaskElem) {
3156 assert(0 <= M && M < (int)NumSrcElts && "Unexpected shuffle mask index");
3157 if (isa<UndefValue>(Src)) {
3158 // We've referenced an undef element - if its poison, update the shuffle
3159 // mask, else bail.
3160 if (!isa<PoisonValue>(Src))
3161 return false;
3162 M = PoisonMaskElem;
3163 continue;
3164 }
3165 if (!NewX || NewX == Src) {
3166 NewX = Src;
3167 continue;
3168 }
3169 if (!NewY || NewY == Src) {
3170 M += NumSrcElts;
3171 NewY = Src;
3172 continue;
3173 }
3174 return false;
3175 }
3176 }
3177
3178 if (!NewX) {
3179 replaceValue(I, *PoisonValue::get(ShuffleDstTy));
3180 return true;
3181 }
3182
3183 if (!NewY)
3184 NewY = PoisonValue::get(ShuffleSrcTy);
3185
3186 // Have we folded to an Identity shuffle?
3187 if (ShuffleVectorInst::isIdentityMask(NewMask, NumSrcElts)) {
3188 replaceValue(I, *NewX);
3189 return true;
3190 }
3191
3192 // Try to merge the shuffles if the new shuffle is not costly.
3193 InstructionCost InnerCost0 = 0;
3194 if (Match0)
3195 InnerCost0 = TTI.getInstructionCost(cast<User>(OuterV0), CostKind);
3196
3197 InstructionCost InnerCost1 = 0;
3198 if (Match1)
3199 InnerCost1 = TTI.getInstructionCost(cast<User>(OuterV1), CostKind);
3200
3202
3203 InstructionCost OldCost = InnerCost0 + InnerCost1 + OuterCost;
3204
3205 bool IsUnary = all_of(NewMask, [&](int M) { return M < (int)NumSrcElts; });
3209 InstructionCost NewCost =
3210 TTI.getShuffleCost(SK, ShuffleDstTy, ShuffleSrcTy, NewMask, CostKind, 0,
3211 nullptr, {NewX, NewY});
3212 if (!OuterV0->hasOneUse())
3213 NewCost += InnerCost0;
3214 if (!OuterV1->hasOneUse())
3215 NewCost += InnerCost1;
3216
3217 LLVM_DEBUG(dbgs() << "Found a shuffle feeding two shuffles: " << I
3218 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
3219 << "\n");
3220 if (NewCost > OldCost)
3221 return false;
3222
3223 Value *Shuf = Builder.CreateShuffleVector(NewX, NewY, NewMask);
3224 replaceValue(I, *Shuf);
3225 return true;
3226}
3227
3228/// Try to convert a chain of length-preserving shuffles that are fed by
3229/// length-changing shuffles from the same source, e.g. a chain of length 3:
3230///
3231/// "shuffle (shuffle (shuffle x, (shuffle y, undef)),
3232/// (shuffle y, undef)),
3233// (shuffle y, undef)"
3234///
3235/// into a single shuffle fed by a length-changing shuffle:
3236///
3237/// "shuffle x, (shuffle y, undef)"
3238///
3239/// Such chains arise e.g. from folding extract/insert sequences.
3240bool VectorCombine::foldShufflesOfLengthChangingShuffles(Instruction &I) {
3241 FixedVectorType *TrunkType = dyn_cast<FixedVectorType>(I.getType());
3242 if (!TrunkType)
3243 return false;
3244
3245 unsigned ChainLength = 0;
3246 SmallVector<int> Mask;
3247 SmallVector<int> YMask;
3248 InstructionCost OldCost = 0;
3249 InstructionCost NewCost = 0;
3250 Value *Trunk = &I;
3251 unsigned NumTrunkElts = TrunkType->getNumElements();
3252 Value *Y = nullptr;
3253
3254 for (;;) {
3255 // Match the current trunk against (commutations of) the pattern
3256 // "shuffle trunk', (shuffle y, undef)"
3257 ArrayRef<int> OuterMask;
3258 Value *OuterV0, *OuterV1;
3259 if (ChainLength != 0 && !Trunk->hasOneUse())
3260 break;
3261 if (!match(Trunk, m_Shuffle(m_Value(OuterV0), m_Value(OuterV1),
3262 m_Mask(OuterMask))))
3263 break;
3264 if (OuterV0->getType() != TrunkType) {
3265 // This shuffle is not length-preserving, so it cannot be part of the
3266 // chain.
3267 break;
3268 }
3269
3270 ArrayRef<int> InnerMask0, InnerMask1;
3271 Value *A0, *A1, *B0, *B1;
3272 bool Match0 =
3273 match(OuterV0, m_Shuffle(m_Value(A0), m_Value(B0), m_Mask(InnerMask0)));
3274 bool Match1 =
3275 match(OuterV1, m_Shuffle(m_Value(A1), m_Value(B1), m_Mask(InnerMask1)));
3276 bool Match0Leaf = Match0 && A0->getType() != I.getType();
3277 bool Match1Leaf = Match1 && A1->getType() != I.getType();
3278 if (Match0Leaf == Match1Leaf) {
3279 // Only handle the case of exactly one leaf in each step. The "two leaves"
3280 // case is handled by foldShuffleOfShuffles.
3281 break;
3282 }
3283
3284 SmallVector<int> CommutedOuterMask;
3285 if (Match0Leaf) {
3286 std::swap(OuterV0, OuterV1);
3287 std::swap(InnerMask0, InnerMask1);
3288 std::swap(A0, A1);
3289 std::swap(B0, B1);
3290 llvm::append_range(CommutedOuterMask, OuterMask);
3291 for (int &M : CommutedOuterMask) {
3292 if (M == PoisonMaskElem)
3293 continue;
3294 if (M < (int)NumTrunkElts)
3295 M += NumTrunkElts;
3296 else
3297 M -= NumTrunkElts;
3298 }
3299 OuterMask = CommutedOuterMask;
3300 }
3301 if (!OuterV1->hasOneUse())
3302 break;
3303
3304 if (!isa<UndefValue>(A1)) {
3305 if (!Y)
3306 Y = A1;
3307 else if (Y != A1)
3308 break;
3309 }
3310 if (!isa<UndefValue>(B1)) {
3311 if (!Y)
3312 Y = B1;
3313 else if (Y != B1)
3314 break;
3315 }
3316
3317 auto *YType = cast<FixedVectorType>(A1->getType());
3318 int NumLeafElts = YType->getNumElements();
3319 SmallVector<int> LocalYMask(InnerMask1);
3320 for (int &M : LocalYMask) {
3321 if (M >= NumLeafElts)
3322 M -= NumLeafElts;
3323 }
3324
3325 InstructionCost LocalOldCost =
3328
3329 // Handle the initial (start of chain) case.
3330 if (!ChainLength) {
3331 Mask.assign(OuterMask);
3332 YMask.assign(LocalYMask);
3333 OldCost = NewCost = LocalOldCost;
3334 Trunk = OuterV0;
3335 ChainLength++;
3336 continue;
3337 }
3338
3339 // For the non-root case, first attempt to combine masks.
3340 SmallVector<int> NewYMask(YMask);
3341 bool Valid = true;
3342 for (auto [CombinedM, LeafM] : llvm::zip(NewYMask, LocalYMask)) {
3343 if (LeafM == -1 || CombinedM == LeafM)
3344 continue;
3345 if (CombinedM == -1) {
3346 CombinedM = LeafM;
3347 } else {
3348 Valid = false;
3349 break;
3350 }
3351 }
3352 if (!Valid)
3353 break;
3354
3355 SmallVector<int> NewMask;
3356 NewMask.reserve(NumTrunkElts);
3357 for (int M : Mask) {
3358 if (M < 0 || M >= static_cast<int>(NumTrunkElts))
3359 NewMask.push_back(M);
3360 else
3361 NewMask.push_back(OuterMask[M]);
3362 }
3363
3364 // Break the chain if adding this new step complicates the shuffles such
3365 // that it would increase the new cost by more than the old cost of this
3366 // step.
3367 InstructionCost LocalNewCost =
3369 YType, NewYMask, CostKind) +
3371 TrunkType, NewMask, CostKind);
3372
3373 if (LocalNewCost >= NewCost && LocalOldCost < LocalNewCost - NewCost)
3374 break;
3375
3376 LLVM_DEBUG({
3377 if (ChainLength == 1) {
3378 dbgs() << "Found chain of shuffles fed by length-changing shuffles: "
3379 << I << '\n';
3380 }
3381 dbgs() << " next chain link: " << *Trunk << '\n'
3382 << " old cost: " << (OldCost + LocalOldCost)
3383 << " new cost: " << LocalNewCost << '\n';
3384 });
3385
3386 Mask = NewMask;
3387 YMask = NewYMask;
3388 OldCost += LocalOldCost;
3389 NewCost = LocalNewCost;
3390 Trunk = OuterV0;
3391 ChainLength++;
3392 }
3393 if (ChainLength <= 1)
3394 return false;
3395
3396 // Bail out if all leaves were poison.
3397 if (!Y)
3398 return false;
3399
3400 if (llvm::all_of(Mask, [&](int M) {
3401 return M < 0 || M >= static_cast<int>(NumTrunkElts);
3402 })) {
3403 // Produce a canonical simplified form if all elements are sourced from Y.
3404 for (int &M : Mask) {
3405 if (M >= static_cast<int>(NumTrunkElts))
3406 M = YMask[M - NumTrunkElts];
3407 }
3408 Value *Root =
3409 Builder.CreateShuffleVector(Y, PoisonValue::get(Y->getType()), Mask);
3410 replaceValue(I, *Root);
3411 return true;
3412 }
3413
3414 Value *Leaf =
3415 Builder.CreateShuffleVector(Y, PoisonValue::get(Y->getType()), YMask);
3416 Value *Root = Builder.CreateShuffleVector(Trunk, Leaf, Mask);
3417 replaceValue(I, *Root);
3418 return true;
3419}
3420
3421/// Try to convert
3422/// "shuffle (intrinsic), (intrinsic)" into "intrinsic (shuffle), (shuffle)".
3423bool VectorCombine::foldShuffleOfIntrinsics(Instruction &I) {
3424 Value *V0, *V1;
3425 ArrayRef<int> OldMask;
3426 if (!match(&I, m_Shuffle(m_Value(V0), m_Value(V1), m_Mask(OldMask))))
3427 return false;
3428
3429 auto *II0 = dyn_cast<IntrinsicInst>(V0);
3430 auto *II1 = dyn_cast<IntrinsicInst>(V1);
3431 if (!II0 || !II1)
3432 return false;
3433
3434 Intrinsic::ID IID = II0->getIntrinsicID();
3435 if (IID != II1->getIntrinsicID())
3436 return false;
3437 InstructionCost CostII0 =
3438 TTI.getIntrinsicInstrCost(IntrinsicCostAttributes(IID, *II0), CostKind);
3439 InstructionCost CostII1 =
3440 TTI.getIntrinsicInstrCost(IntrinsicCostAttributes(IID, *II1), CostKind);
3441
3442 auto *ShuffleDstTy = dyn_cast<FixedVectorType>(I.getType());
3443 auto *II0Ty = dyn_cast<FixedVectorType>(II0->getType());
3444 if (!ShuffleDstTy || !II0Ty)
3445 return false;
3446
3447 if (!isTriviallyVectorizable(IID))
3448 return false;
3449
3450 for (unsigned I = 0, E = II0->arg_size(); I != E; ++I) {
3451 Value *Arg0 = II0->getArgOperand(I);
3452 Value *Arg1 = II1->getArgOperand(I);
3454 // Scalar operands must be identical.
3455 if (Arg0 != Arg1)
3456 return false;
3457 } else if (Arg0->getType() != Arg1->getType()) {
3458 // The corresponding vector operands are shuffled together, so they must
3459 // share the same type. For intrinsics overloaded on their operand type
3460 // (e.g. llvm.fptosi.sat), two calls can produce the same result type
3461 // from different operand types; shuffling those would be invalid.
3462 return false;
3463 }
3464 }
3465
3466 InstructionCost OldCost =
3467 CostII0 + CostII1 +
3469 II0Ty, OldMask, CostKind, 0, nullptr, {II0, II1}, &I);
3470
3471 SmallVector<Type *> NewArgsTy;
3472 InstructionCost NewCost = 0;
3473 SmallDenseSet<std::pair<Value *, Value *>> SeenOperandPairs;
3474 for (unsigned I = 0, E = II0->arg_size(); I != E; ++I) {
3476 NewArgsTy.push_back(II0->getArgOperand(I)->getType());
3477 } else {
3478 auto *VecTy = cast<FixedVectorType>(II0->getArgOperand(I)->getType());
3479 auto *ArgTy = FixedVectorType::get(VecTy->getElementType(),
3480 ShuffleDstTy->getNumElements());
3481 NewArgsTy.push_back(ArgTy);
3482 std::pair<Value *, Value *> OperandPair =
3483 std::make_pair(II0->getArgOperand(I), II1->getArgOperand(I));
3484 if (!SeenOperandPairs.insert(OperandPair).second) {
3485 // We've already computed the cost for this operand pair.
3486 continue;
3487 }
3488 NewCost += TTI.getShuffleCost(
3489 TargetTransformInfo::SK_PermuteTwoSrc, ArgTy, VecTy, OldMask,
3490 CostKind, 0, nullptr, {II0->getArgOperand(I), II1->getArgOperand(I)});
3491 }
3492 }
3493 IntrinsicCostAttributes NewAttr(IID, ShuffleDstTy, NewArgsTy);
3494
3495 NewCost += TTI.getIntrinsicInstrCost(NewAttr, CostKind);
3496 if (!II0->hasOneUse())
3497 NewCost += CostII0;
3498 if (II1 != II0 && !II1->hasOneUse())
3499 NewCost += CostII1;
3500
3501 LLVM_DEBUG(dbgs() << "Found a shuffle feeding two intrinsics: " << I
3502 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
3503 << "\n");
3504
3505 if (NewCost > OldCost)
3506 return false;
3507
3508 SmallVector<Value *> NewArgs;
3509 SmallDenseMap<std::pair<Value *, Value *>, Value *> ShuffleCache;
3510 for (unsigned I = 0, E = II0->arg_size(); I != E; ++I)
3512 NewArgs.push_back(II0->getArgOperand(I));
3513 } else {
3514 std::pair<Value *, Value *> OperandPair =
3515 std::make_pair(II0->getArgOperand(I), II1->getArgOperand(I));
3516 auto It = ShuffleCache.find(OperandPair);
3517 if (It != ShuffleCache.end()) {
3518 // Reuse previously created shuffle for this operand pair.
3519 NewArgs.push_back(It->second);
3520 continue;
3521 }
3522 Value *Shuf = Builder.CreateShuffleVector(II0->getArgOperand(I),
3523 II1->getArgOperand(I), OldMask);
3524 ShuffleCache[OperandPair] = Shuf;
3525 NewArgs.push_back(Shuf);
3526 Worklist.pushValue(Shuf);
3527 }
3528 Value *NewIntrinsic = Builder.CreateIntrinsic(ShuffleDstTy, IID, NewArgs);
3529
3530 // Intersect flags from the old intrinsics.
3531 if (auto *NewInst = dyn_cast<Instruction>(NewIntrinsic)) {
3532 NewInst->copyIRFlags(II0);
3533 NewInst->andIRFlags(II1);
3534 }
3535
3536 replaceValue(I, *NewIntrinsic);
3537 return true;
3538}
3539
3540/// Try to convert
3541/// "shuffle (intrinsic), (poison/undef)" into "intrinsic (shuffle)".
3542bool VectorCombine::foldPermuteOfIntrinsic(Instruction &I) {
3543 Value *V0;
3544 ArrayRef<int> Mask;
3545 if (!match(&I, m_Shuffle(m_Value(V0), m_Undef(), m_Mask(Mask))))
3546 return false;
3547
3548 auto *II0 = dyn_cast<IntrinsicInst>(V0);
3549 if (!II0)
3550 return false;
3551
3552 auto *ShuffleDstTy = dyn_cast<FixedVectorType>(I.getType());
3553 auto *IntrinsicSrcTy = dyn_cast<FixedVectorType>(II0->getType());
3554 if (!ShuffleDstTy || !IntrinsicSrcTy)
3555 return false;
3556
3557 // Validate it's a pure permute, mask should only reference the first vector
3558 unsigned NumSrcElts = IntrinsicSrcTy->getNumElements();
3559 if (any_of(Mask, [NumSrcElts](int M) { return M >= (int)NumSrcElts; }))
3560 return false;
3561
3562 Intrinsic::ID IID = II0->getIntrinsicID();
3563 if (!isTriviallyVectorizable(IID))
3564 return false;
3565
3566 // Cost analysis
3568 TTI.getIntrinsicInstrCost(IntrinsicCostAttributes(IID, *II0), CostKind);
3569 InstructionCost OldCost =
3572 IntrinsicSrcTy, Mask, CostKind, 0, nullptr, {V0}, &I);
3573
3574 SmallVector<Type *> NewArgsTy;
3575 InstructionCost NewCost = 0;
3576 for (unsigned I = 0, E = II0->arg_size(); I != E; ++I) {
3578 NewArgsTy.push_back(II0->getArgOperand(I)->getType());
3579 } else {
3580 auto *VecTy = cast<FixedVectorType>(II0->getArgOperand(I)->getType());
3581 auto *ArgTy = FixedVectorType::get(VecTy->getElementType(),
3582 ShuffleDstTy->getNumElements());
3583 NewArgsTy.push_back(ArgTy);
3585 ArgTy, VecTy, Mask, CostKind, 0, nullptr,
3586 {II0->getArgOperand(I)});
3587 }
3588 }
3589 IntrinsicCostAttributes NewAttr(IID, ShuffleDstTy, NewArgsTy);
3590 NewCost += TTI.getIntrinsicInstrCost(NewAttr, CostKind);
3591
3592 // If the intrinsic has multiple uses, we need to account for the cost of
3593 // keeping the original intrinsic around.
3594 if (!II0->hasOneUse())
3595 NewCost += IntrinsicCost;
3596
3597 LLVM_DEBUG(dbgs() << "Found a permute of intrinsic: " << I << "\n OldCost: "
3598 << OldCost << " vs NewCost: " << NewCost << "\n");
3599
3600 if (NewCost > OldCost)
3601 return false;
3602
3603 // Transform
3604 SmallVector<Value *> NewArgs;
3605 for (unsigned I = 0, E = II0->arg_size(); I != E; ++I) {
3607 NewArgs.push_back(II0->getArgOperand(I));
3608 } else {
3609 Value *Shuf = Builder.CreateShuffleVector(II0->getArgOperand(I), Mask);
3610 NewArgs.push_back(Shuf);
3611 Worklist.pushValue(Shuf);
3612 }
3613 }
3614
3615 Value *NewIntrinsic = Builder.CreateIntrinsic(ShuffleDstTy, IID, NewArgs);
3616
3617 if (auto *NewInst = dyn_cast<Instruction>(NewIntrinsic))
3618 NewInst->copyIRFlags(II0);
3619
3620 replaceValue(I, *NewIntrinsic);
3621 return true;
3622}
3623
3624using InstLane = std::pair<Value *, int>;
3625
3626static InstLane lookThroughShuffles(Value *V, int Lane) {
3627 while (auto *SV = dyn_cast<ShuffleVectorInst>(V)) {
3628 unsigned NumElts =
3629 cast<FixedVectorType>(SV->getOperand(0)->getType())->getNumElements();
3630 int M = SV->getMaskValue(Lane);
3631 if (M < 0)
3632 return {nullptr, PoisonMaskElem};
3633 if (static_cast<unsigned>(M) < NumElts) {
3634 V = SV->getOperand(0);
3635 Lane = M;
3636 } else {
3637 V = SV->getOperand(1);
3638 Lane = M - NumElts;
3639 }
3640 }
3641 return InstLane{V, Lane};
3642}
3643
3647 for (InstLane IL : Item) {
3648 auto [U, Lane] = IL;
3649 InstLane OpLane =
3650 U ? lookThroughShuffles(cast<Instruction>(U)->getOperand(Op), Lane)
3651 : InstLane{nullptr, PoisonMaskElem};
3652 NItem.emplace_back(OpLane);
3653 }
3654 return NItem;
3655}
3656
3657/// Detect concat of multiple values into a vector
3659 const TargetTransformInfo &TTI) {
3660 auto *Ty = cast<FixedVectorType>(Item.front().first->getType());
3661 unsigned NumElts = Ty->getNumElements();
3662 if (Item.size() == NumElts || NumElts == 1 || Item.size() % NumElts != 0)
3663 return false;
3664
3665 // Check that the concat is free, usually meaning that the type will be split
3666 // during legalization.
3667 SmallVector<int, 16> ConcatMask(NumElts * 2);
3668 std::iota(ConcatMask.begin(), ConcatMask.end(), 0);
3669 if (TTI.getShuffleCost(TTI::SK_PermuteTwoSrc,
3670 FixedVectorType::get(Ty->getScalarType(), NumElts * 2),
3671 Ty, ConcatMask, CostKind) != 0)
3672 return false;
3673
3674 unsigned NumSlices = Item.size() / NumElts;
3675 // Currently we generate a tree of shuffles for the concats, which limits us
3676 // to a power2.
3677 if (!isPowerOf2_32(NumSlices))
3678 return false;
3679 for (unsigned Slice = 0; Slice < NumSlices; ++Slice) {
3680 Value *SliceV = Item[Slice * NumElts].first;
3681 if (!SliceV || SliceV->getType() != Ty)
3682 return false;
3683 for (unsigned Elt = 0; Elt < NumElts; ++Elt) {
3684 auto [V, Lane] = Item[Slice * NumElts + Elt];
3685 if (Lane != static_cast<int>(Elt) || SliceV != V)
3686 return false;
3687 }
3688 }
3689 return true;
3690}
3691
3692static Value *
3694 const DenseSet<std::pair<Value *, Use *>> &IdentityLeafs,
3695 const DenseSet<std::pair<Value *, Use *>> &SplatLeafs,
3696 const DenseSet<std::pair<Value *, Use *>> &ConcatLeafs,
3697 IRBuilderBase &Builder, InstructionWorklist &WorkList,
3698 const TargetTransformInfo *TTI) {
3699 auto [FrontV, FrontLane] = Item.front();
3700
3701 if (IdentityLeafs.contains(std::make_pair(FrontV, From))) {
3702 return FrontV;
3703 }
3704 if (SplatLeafs.contains(std::make_pair(FrontV, From))) {
3705 SmallVector<int, 16> Mask(Item.size(), FrontLane);
3706 return Builder.CreateShuffleVector(FrontV, Mask);
3707 }
3708 if (ConcatLeafs.contains(std::make_pair(FrontV, From))) {
3709 unsigned NumElts =
3710 cast<FixedVectorType>(FrontV->getType())->getNumElements();
3711 SmallVector<Value *> Values(Item.size() / NumElts, nullptr);
3712 for (unsigned S = 0; S < Values.size(); ++S)
3713 Values[S] = Item[S * NumElts].first;
3714
3715 while (Values.size() > 1) {
3716 NumElts *= 2;
3717 SmallVector<int, 16> Mask(NumElts, 0);
3718 std::iota(Mask.begin(), Mask.end(), 0);
3719 SmallVector<Value *> NewValues(Values.size() / 2, nullptr);
3720 for (unsigned S = 0; S < NewValues.size(); ++S)
3721 NewValues[S] =
3722 Builder.CreateShuffleVector(Values[S * 2], Values[S * 2 + 1], Mask);
3723 Values = NewValues;
3724 }
3725 return Values[0];
3726 }
3727
3728 auto *I = cast<Instruction>(FrontV);
3729
3730 // Handle vector bitcasts that change element count. We cannot use
3731 // generateInstLaneVectorFromOperand for these because the lane indices
3732 // don't map 1:1 through the bitcast.
3733 if (auto *BitCast = dyn_cast<BitCastInst>(I)) {
3734 auto *BCDstTy = dyn_cast<FixedVectorType>(BitCast->getDestTy());
3735 auto *BCSrcTy = dyn_cast<FixedVectorType>(BitCast->getSrcTy());
3736 if (BCDstTy && BCSrcTy &&
3737 BCDstTy->getElementCount() != BCSrcTy->getElementCount()) {
3738 unsigned DstElts = BCDstTy->getNumElements();
3739 unsigned SrcElts = BCSrcTy->getNumElements();
3740 SmallVector<InstLane> NewItem;
3741 if (DstElts > SrcElts) {
3742 // Widening: compress operand Item.
3743 unsigned R = DstElts / SrcElts;
3744 if (Item.size() % R != 0)
3745 return nullptr;
3746 for (unsigned Idx = 0, E = Item.size(); Idx < E; Idx += R) {
3747 auto [V, Lane] = Item[Idx];
3748 if (!V) {
3749 NewItem.push_back({nullptr, PoisonMaskElem});
3750 continue;
3751 }
3752 NewItem.push_back(
3753 lookThroughShuffles(cast<Operator>(V)->getOperand(0), Lane / R));
3754 }
3755 } else {
3756 // Narrowing: expand operand Item.
3757 unsigned R = SrcElts / DstElts;
3758 for (auto [V, Lane] : Item) {
3759 if (!V) {
3760 NewItem.append(R, {nullptr, PoisonMaskElem});
3761 continue;
3762 }
3763 Value *Op = cast<Operator>(V)->getOperand(0);
3764 for (unsigned J = 0; J < R; ++J)
3765 NewItem.push_back(lookThroughShuffles(Op, Lane * R + J));
3766 }
3767 }
3768 Value *Op = generateNewInstTree(NewItem, &BitCast->getOperandUse(0),
3769 IdentityLeafs, SplatLeafs, ConcatLeafs,
3770 Builder, WorkList, TTI);
3771 WorkList.pushValue(Op);
3772 return Builder.CreateBitCast(
3773 Op, FixedVectorType::get(BCDstTy->getScalarType(), Item.size()));
3774 }
3775 }
3776 auto *II = dyn_cast<IntrinsicInst>(I);
3777 unsigned NumOps = I->getNumOperands() - (II ? 1 : 0);
3779 for (unsigned Idx = 0; Idx < NumOps; Idx++) {
3780 if (II &&
3781 isVectorIntrinsicWithScalarOpAtArg(II->getIntrinsicID(), Idx, TTI)) {
3782 Ops[Idx] = II->getOperand(Idx);
3783 continue;
3784 }
3785 Ops[Idx] = generateNewInstTree(
3786 generateInstLaneVectorFromOperand(Item, Idx), &I->getOperandUse(Idx),
3787 IdentityLeafs, SplatLeafs, ConcatLeafs, Builder, WorkList, TTI);
3788 // Don't re-queue the operand of a bitcast we just regenerated. Doing so
3789 // lets foldBitcastShuffle sink the bitcast back into a shuffle(bitcast),
3790 // which foldShuffleToIdentity then re-matches as the same superfluous
3791 // identity - an infinite loop between the two folds.
3792 if (!isa<BitCastInst>(I))
3793 WorkList.pushValue(Ops[Idx]);
3794 }
3795
3796 SmallVector<Value *, 8> ValueList;
3797 for (const auto &Lane : Item)
3798 if (Lane.first)
3799 ValueList.push_back(Lane.first);
3800
3801 Type *DstTy =
3802 FixedVectorType::get(I->getType()->getScalarType(), Item.size());
3803 if (auto *BI = dyn_cast<BinaryOperator>(I)) {
3804 auto *Value = Builder.CreateBinOp((Instruction::BinaryOps)BI->getOpcode(),
3805 Ops[0], Ops[1]);
3806 propagateIRFlags(Value, ValueList);
3807 return Value;
3808 }
3809 if (auto *CI = dyn_cast<CmpInst>(I)) {
3810 auto *Value = Builder.CreateCmp(CI->getPredicate(), Ops[0], Ops[1]);
3811 propagateIRFlags(Value, ValueList);
3812 return Value;
3813 }
3814 if (auto *SI = dyn_cast<SelectInst>(I)) {
3815 auto *Value = Builder.CreateSelect(Ops[0], Ops[1], Ops[2], "", SI);
3816 propagateIRFlags(Value, ValueList);
3817 return Value;
3818 }
3819 if (auto *CI = dyn_cast<CastInst>(I)) {
3820 auto *Value = Builder.CreateCast(CI->getOpcode(), Ops[0], DstTy);
3821 propagateIRFlags(Value, ValueList);
3822 return Value;
3823 }
3824 if (II) {
3825 auto *Value = Builder.CreateIntrinsic(DstTy, II->getIntrinsicID(), Ops);
3826 propagateIRFlags(Value, ValueList);
3827 return Value;
3828 }
3829 assert(isa<UnaryInstruction>(I) && "Unexpected instruction type in Generate");
3830 auto *Value =
3831 Builder.CreateUnOp((Instruction::UnaryOps)I->getOpcode(), Ops[0]);
3832 propagateIRFlags(Value, ValueList);
3833 return Value;
3834}
3835
3836// Starting from a shuffle, look up through operands tracking the shuffled index
3837// of each lane. If we can simplify away the shuffles to identities then
3838// do so.
3839bool VectorCombine::foldShuffleToIdentity(Instruction &I) {
3840 auto *Ty = dyn_cast<FixedVectorType>(I.getType());
3841 if (!Ty || I.use_empty())
3842 return false;
3843
3844 SmallVector<InstLane> Start(Ty->getNumElements());
3845 for (unsigned M = 0, E = Ty->getNumElements(); M < E; ++M)
3846 Start[M] = lookThroughShuffles(&I, M);
3847
3849 Candidates.push_back(std::make_pair(Start, &*I.use_begin()));
3850 DenseSet<std::pair<Value *, Use *>> IdentityLeafs, SplatLeafs, ConcatLeafs;
3851 unsigned NumVisited = 0;
3852 bool TraversedElCountChangingBitcast = false;
3853
3854 while (!Candidates.empty()) {
3855 if (++NumVisited > MaxInstrsToScan)
3856 return false;
3857
3858 auto ItemFrom = Candidates.pop_back_val();
3859 auto Item = ItemFrom.first;
3860 auto From = ItemFrom.second;
3861 auto [FrontV, FrontLane] = Item.front();
3862
3863 // If we found an undef first lane then bail out to keep things simple.
3864 if (!FrontV)
3865 return false;
3866
3867 // Look for an identity value.
3868 if (FrontLane == 0 &&
3869 cast<FixedVectorType>(FrontV->getType())->getNumElements() ==
3870 Item.size() &&
3871 all_of(drop_begin(enumerate(Item)), [Item](const auto &E) {
3872 Value *FrontV = Item.front().first;
3873 return !E.value().first || (isEquivBitcast(E.value().first, FrontV) &&
3874 E.value().second == (int)E.index());
3875 })) {
3876 IdentityLeafs.insert(std::make_pair(FrontV, From));
3877 continue;
3878 }
3879 // Look for constants, for the moment only supporting constant splats.
3880 if (auto *C = dyn_cast<Constant>(FrontV);
3881 C && C->getSplatValue() &&
3882 all_of(drop_begin(Item), [Item](InstLane &IL) {
3883 Value *FrontV = Item.front().first;
3884 Value *V = IL.first;
3885 return !V || (isa<Constant>(V) &&
3886 cast<Constant>(V)->getSplatValue() ==
3887 cast<Constant>(FrontV)->getSplatValue());
3888 })) {
3889 SplatLeafs.insert(std::make_pair(FrontV, From));
3890 continue;
3891 }
3892 // Look for a splat value.
3893 if (all_of(drop_begin(Item), [Item](InstLane &IL) {
3894 auto [FrontV, FrontLane] = Item.front();
3895 auto [V, Lane] = IL;
3896 return !V || (V == FrontV && Lane == FrontLane);
3897 })) {
3898 SplatLeafs.insert(std::make_pair(FrontV, From));
3899 continue;
3900 }
3901
3902 // We need each element to be the same type of value, and check that each
3903 // element has a single use.
3904 auto CheckLaneIsEquivalentToFirst = [Item](InstLane IL) {
3905 Value *FrontV = Item.front().first;
3906 if (!IL.first)
3907 return true;
3908 Value *V = IL.first;
3909 if (auto *I = dyn_cast<Instruction>(V); I && !I->hasOneUser())
3910 return false;
3911 if (V->getValueID() != FrontV->getValueID())
3912 return false;
3913 if (auto *CI = dyn_cast<CmpInst>(V))
3914 if (CI->getPredicate() != cast<CmpInst>(FrontV)->getPredicate())
3915 return false;
3916 if (auto *CI = dyn_cast<CastInst>(V))
3917 if (CI->getSrcTy()->getScalarType() !=
3918 cast<CastInst>(FrontV)->getSrcTy()->getScalarType())
3919 return false;
3920 if (auto *SI = dyn_cast<SelectInst>(V))
3921 if (!isa<VectorType>(SI->getOperand(0)->getType()) ||
3922 SI->getOperand(0)->getType() !=
3923 cast<SelectInst>(FrontV)->getOperand(0)->getType())
3924 return false;
3925 if (isa<CallInst>(V) && !isa<IntrinsicInst>(V))
3926 return false;
3927 auto *II = dyn_cast<IntrinsicInst>(V);
3928 return !II || (isa<IntrinsicInst>(FrontV) &&
3929 II->getIntrinsicID() ==
3930 cast<IntrinsicInst>(FrontV)->getIntrinsicID() &&
3931 !II->hasOperandBundles());
3932 };
3933 if (all_of(drop_begin(Item), CheckLaneIsEquivalentToFirst)) {
3934 // Check the operator is one that we support.
3935 if (isa<BinaryOperator, CmpInst>(FrontV)) {
3936 // We exclude div/rem in case they hit UB from poison lanes.
3937 if (auto *BO = dyn_cast<BinaryOperator>(FrontV);
3938 BO && BO->isIntDivRem())
3939 return false;
3941 &cast<Instruction>(FrontV)->getOperandUse(0));
3943 &cast<Instruction>(FrontV)->getOperandUse(1));
3944 continue;
3945 } else if (isa<UnaryOperator, TruncInst, ZExtInst, SExtInst, FPToSIInst,
3946 FPToUIInst, SIToFPInst, UIToFPInst>(FrontV)) {
3948 &cast<Instruction>(FrontV)->getOperandUse(0));
3949 continue;
3950 } else if (auto *BitCast = dyn_cast<BitCastInst>(FrontV)) {
3951 auto *BCDstTy = dyn_cast<FixedVectorType>(BitCast->getDestTy());
3952 auto *BCSrcTy = dyn_cast<FixedVectorType>(BitCast->getSrcTy());
3953 if (BCDstTy && BCSrcTy) {
3954 ElementCount DstEC = BCDstTy->getElementCount();
3955 ElementCount SrcEC = BCSrcTy->getElementCount();
3956 if (DstEC == SrcEC) {
3957 // Same element count - simple pass-through.
3959 &BitCast->getOperandUse(0));
3960 continue;
3961 }
3962 unsigned DstElts = DstEC.getFixedValue();
3963 unsigned SrcElts = SrcEC.getFixedValue();
3964 if (DstElts > SrcElts && DstElts % SrcElts == 0) {
3965 // Widening bitcast (e.g. <2 x i32> -> <4 x i16>). Compress
3966 // consecutive groups of R destination lanes into one source
3967 // lane.
3968 unsigned R = DstElts / SrcElts;
3970 bool Valid = Item.size() % R == 0;
3971 for (unsigned Idx = 0, E = Item.size(); Valid && Idx < E;
3972 Idx += R) {
3973 auto [V0, L0] = Item[Idx];
3974 if (!V0) {
3975 if (any_of(ArrayRef(Item).slice(Idx + 1, R - 1),
3976 [](InstLane IL) { return IL.first != nullptr; })) {
3977 Valid = false;
3978 break;
3979 }
3980 NItem.push_back({nullptr, PoisonMaskElem});
3981 continue;
3982 }
3983 if (L0 % R != 0) {
3984 Valid = false;
3985 break;
3986 }
3987 for (unsigned J = 1; J < R; ++J) {
3988 auto [VJ, LJ] = Item[Idx + J];
3989 if (!VJ || VJ != V0 || LJ != L0 + (int)J) {
3990 Valid = false;
3991 break;
3992 }
3993 }
3994 if (!Valid)
3995 break;
3997 cast<Operator>(V0)->getOperand(0), L0 / R));
3998 }
3999 if (Valid) {
4000 TraversedElCountChangingBitcast = true;
4001 Candidates.emplace_back(NItem, &BitCast->getOperandUse(0));
4002 continue;
4003 }
4004 } else if (SrcElts > DstElts && SrcElts % DstElts == 0) {
4005 // Narrowing bitcast (e.g. <4 x i16> -> <2 x i32>). Expand
4006 // each destination lane into R source lanes.
4007 unsigned R = SrcElts / DstElts;
4009 for (auto [V, Lane] : Item) {
4010 if (!V) {
4011 NItem.append(R, {nullptr, PoisonMaskElem});
4012 continue;
4013 }
4014 Value *Op = cast<Operator>(V)->getOperand(0);
4015 for (unsigned J = 0; J < R; ++J)
4016 NItem.push_back(lookThroughShuffles(Op, Lane * R + J));
4017 }
4018 TraversedElCountChangingBitcast = true;
4019 Candidates.emplace_back(NItem, &BitCast->getOperandUse(0));
4020 continue;
4021 }
4022 }
4023 } else if (auto *Sel = dyn_cast<SelectInst>(FrontV)) {
4025 &Sel->getOperandUse(0));
4027 &Sel->getOperandUse(1));
4029 &Sel->getOperandUse(2));
4030 continue;
4031 } else if (auto *II = dyn_cast<IntrinsicInst>(FrontV);
4032 II && isTriviallyVectorizable(II->getIntrinsicID()) &&
4033 !II->hasOperandBundles()) {
4034 for (unsigned Op = 0, E = II->getNumOperands() - 1; Op < E; Op++) {
4035 if (isVectorIntrinsicWithScalarOpAtArg(II->getIntrinsicID(), Op,
4036 &TTI)) {
4037 if (!all_of(drop_begin(Item), [Item, Op](InstLane &IL) {
4038 Value *FrontV = Item.front().first;
4039 Value *V = IL.first;
4040 return !V || (cast<Instruction>(V)->getOperand(Op) ==
4041 cast<Instruction>(FrontV)->getOperand(Op));
4042 }))
4043 return false;
4044 continue;
4045 }
4046 Candidates.emplace_back(
4048 &cast<Instruction>(FrontV)->getOperandUse(Op));
4049 }
4050 continue;
4051 }
4052 }
4053
4054 if (isFreeConcat(Item, CostKind, TTI)) {
4055 ConcatLeafs.insert(std::make_pair(FrontV, From));
4056 continue;
4057 }
4058
4059 return false;
4060 }
4061
4062 if (NumVisited <= 1)
4063 return false;
4064
4065 // If the only non-leaf node traversed was a single bitcast that changes
4066 // element count, the fold would just commute the bitcast and shuffle.
4067 // foldBitcastShuffle does the reverse transform, causing an infinite loop.
4068 if (NumVisited == 2 && TraversedElCountChangingBitcast)
4069 return false;
4070
4071 LLVM_DEBUG(dbgs() << "Found a superfluous identity shuffle: " << I << "\n");
4072
4073 // If we got this far, we know the shuffles are superfluous and can be
4074 // removed. Scan through again and generate the new tree of instructions.
4075 Builder.SetInsertPoint(&I);
4076 Value *V =
4077 generateNewInstTree(Start, &*I.use_begin(), IdentityLeafs, SplatLeafs,
4078 ConcatLeafs, Builder, Worklist, &TTI);
4079 replaceValue(I, *V);
4080 return true;
4081}
4082
4083/// Given a commutative reduction, the order of the input lanes does not alter
4084/// the results. We can use this to remove certain shuffles feeding the
4085/// reduction, removing the need to shuffle at all.
4086bool VectorCombine::foldShuffleFromReductions(Instruction &I) {
4087 auto *II = dyn_cast<IntrinsicInst>(&I);
4088 if (!II)
4089 return false;
4090 switch (II->getIntrinsicID()) {
4091 case Intrinsic::vector_reduce_add:
4092 case Intrinsic::vector_reduce_mul:
4093 case Intrinsic::vector_reduce_and:
4094 case Intrinsic::vector_reduce_or:
4095 case Intrinsic::vector_reduce_xor:
4096 case Intrinsic::vector_reduce_smin:
4097 case Intrinsic::vector_reduce_smax:
4098 case Intrinsic::vector_reduce_umin:
4099 case Intrinsic::vector_reduce_umax:
4100 break;
4101 default:
4102 return false;
4103 }
4104
4105 // Find all the inputs when looking through operations that do not alter the
4106 // lane order (binops, for example). Currently we look for a single shuffle,
4107 // and can ignore splat values.
4108 std::queue<Value *> Worklist;
4109 SmallPtrSet<Value *, 4> Visited;
4110 ShuffleVectorInst *Shuffle = nullptr;
4111 if (auto *Op = dyn_cast<Instruction>(I.getOperand(0)))
4112 Worklist.push(Op);
4113
4114 while (!Worklist.empty()) {
4115 Value *CV = Worklist.front();
4116 Worklist.pop();
4117 if (Visited.contains(CV))
4118 continue;
4119
4120 // Splats don't change the order, so can be safely ignored.
4121 if (isSplatValue(CV))
4122 continue;
4123
4124 Visited.insert(CV);
4125
4126 if (auto *CI = dyn_cast<Instruction>(CV)) {
4127 if (CI->isBinaryOp()) {
4128 for (auto *Op : CI->operand_values())
4129 Worklist.push(Op);
4130 continue;
4131 } else if (auto *SV = dyn_cast<ShuffleVectorInst>(CI)) {
4132 if (Shuffle && Shuffle != SV)
4133 return false;
4134 Shuffle = SV;
4135 continue;
4136 }
4137 }
4138
4139 // Anything else is currently an unknown node.
4140 return false;
4141 }
4142
4143 if (!Shuffle)
4144 return false;
4145
4146 // Check all uses of the binary ops and shuffles are also included in the
4147 // lane-invariant operations (Visited should be the list of lanewise
4148 // instructions, including the shuffle that we found).
4149 for (auto *V : Visited)
4150 for (auto *U : V->users())
4151 if (!Visited.contains(U) && U != &I)
4152 return false;
4153
4154 FixedVectorType *VecType =
4155 dyn_cast<FixedVectorType>(II->getOperand(0)->getType());
4156 if (!VecType)
4157 return false;
4158 FixedVectorType *ShuffleInputType =
4160 if (!ShuffleInputType)
4161 return false;
4162 unsigned NumInputElts = ShuffleInputType->getNumElements();
4163
4164 // Find the mask from sorting the lanes into order. This is most likely to
4165 // become a identity or concat mask. Undef elements are pushed to the end.
4166 SmallVector<int> ConcatMask;
4167 Shuffle->getShuffleMask(ConcatMask);
4168 sort(ConcatMask, [](int X, int Y) { return (unsigned)X < (unsigned)Y; });
4169 bool UsesSecondVec =
4170 any_of(ConcatMask, [&](int M) { return M >= (int)NumInputElts; });
4171
4173 UsesSecondVec ? TTI::SK_PermuteTwoSrc : TTI::SK_PermuteSingleSrc, VecType,
4174 ShuffleInputType, Shuffle->getShuffleMask(), CostKind);
4176 UsesSecondVec ? TTI::SK_PermuteTwoSrc : TTI::SK_PermuteSingleSrc, VecType,
4177 ShuffleInputType, ConcatMask, CostKind);
4178
4179 LLVM_DEBUG(dbgs() << "Found a reduction feeding from a shuffle: " << *Shuffle
4180 << "\n");
4181 LLVM_DEBUG(dbgs() << " OldCost: " << OldCost << " vs NewCost: " << NewCost
4182 << "\n");
4183 bool MadeChanges = false;
4184 if (NewCost < OldCost) {
4185 Builder.SetInsertPoint(Shuffle);
4186 Value *NewShuffle = Builder.CreateShuffleVector(
4187 Shuffle->getOperand(0), Shuffle->getOperand(1), ConcatMask);
4188 LLVM_DEBUG(dbgs() << "Created new shuffle: " << *NewShuffle << "\n");
4189 replaceValue(*Shuffle, *NewShuffle);
4190 return true;
4191 }
4192
4193 // See if we can re-use foldSelectShuffle, getting it to reduce the size of
4194 // the shuffle into a nicer order, as it can ignore the order of the shuffles.
4195 MadeChanges |= foldSelectShuffle(*Shuffle, true);
4196 return MadeChanges;
4197}
4198
4199/// Try to fold a chain of shuffles and ops feeding extractelement(..., 0)
4200/// into llvm.vector.reduce.*, by tracking which lanes contribute to the
4201/// extracted lane and reducing the widest vector whose lanes each contribute
4202/// once.
4203///
4204/// For example:
4205///
4206/// %lo = shufflevector <4 x i32> %a, poison, <2 x i32> <i32 0, i32 1>
4207/// %hi = shufflevector <4 x i32> %a, poison, <2 x i32> <i32 2, i32 3>
4208/// %s = add <2 x i32> %lo, %hi
4209/// %sh = shufflevector <2 x i32> %s, poison, <2 x i32> <i32 1, i32 poison>
4210/// %r = add <2 x i32> %s, %sh
4211/// %e = extractelement <2 x i32> %r, i64 0
4212///
4213/// transforms to:
4214///
4215/// %e = call i32 @llvm.vector.reduce.add.v4i32(<4 x i32> %a)
4216bool VectorCombine::foldShuffleChainsToReduce(Instruction &I) {
4217 Value *VecOpEE;
4218 if (!match(&I, m_ExtractElt(m_Value(VecOpEE), m_Zero())))
4219 return false;
4220
4221 auto *FVT = dyn_cast<FixedVectorType>(VecOpEE->getType());
4222 if (!FVT)
4223 return false;
4224
4225 if (FVT->getNumElements() < 2)
4226 return false;
4227
4228 std::optional<Instruction::BinaryOps> CommonBinOp;
4229 std::optional<Intrinsic::ID> CommonCallOp;
4230
4231 if (auto *BO = dyn_cast<BinaryOperator>(VecOpEE)) {
4232 if (!getReductionForBinop(BO->getOpcode()))
4233 return false;
4234 CommonBinOp = BO->getOpcode();
4235 } else if (auto *MMI = dyn_cast<MinMaxIntrinsic>(VecOpEE)) {
4236 CommonCallOp = MMI->getIntrinsicID();
4237 } else {
4238 return false;
4239 }
4240
4241 // For floating-point reductions, track FMF intersection across all binops.
4242 FastMathFlags CommonFMF;
4243 bool IsFloatReduction = false;
4244
4245 // A chain node is one we walk through, either a matching-opcode binop/min-max
4246 // or a single-source shuffle. Anything else is a leaf source.
4247 auto IsChainNode = [&](Value *V) {
4248 if (auto *BO = dyn_cast<BinaryOperator>(V))
4249 return CommonBinOp && BO->getOpcode() == *CommonBinOp;
4250 if (auto *MMI = dyn_cast<MinMaxIntrinsic>(V))
4251 return CommonCallOp && MMI->getIntrinsicID() == *CommonCallOp;
4252 if (auto *SVI = dyn_cast<ShuffleVectorInst>(V))
4253 return isa<PoisonValue>(SVI->getOperand(1));
4254 return false;
4255 };
4256
4257 // Collect the chain, building Nodes in postorder. Bail if the chain is empty
4258 // or exceeds MaxChainNodes.
4259 constexpr unsigned MaxChainNodes = 32;
4260 SmallSetVector<Value *, 16> Nodes;
4261 SmallSetVector<Value *, 4> Sources;
4262 unsigned NumVisited = 0;
4263 auto AddSource = [&](Value *V) {
4264 if (!isa<FixedVectorType>(V->getType()))
4265 return false;
4266 Sources.insert(V);
4267 return true;
4268 };
4269 auto Walk = [&](Value *V, auto &&Walk) -> bool {
4270 if (Nodes.contains(V) || Sources.contains(V))
4271 return true;
4272 if (++NumVisited > MaxChainNodes)
4273 return false;
4274 if (!IsChainNode(V))
4275 return AddSource(V);
4276 // Chain shuffles always have poison as op1, so only op0 matters.
4277 auto *U = cast<Instruction>(V);
4278 unsigned NumOps = isa<ShuffleVectorInst>(U) ? 1 : 2;
4279 for (unsigned I = 0; I != NumOps; ++I)
4280 if (!Walk(U->getOperand(I), Walk))
4281 return false;
4282 if (isa<ShuffleVectorInst>(U) || Nodes.contains(U->getOperand(0)) ||
4283 Nodes.contains(U->getOperand(1))) {
4284 Nodes.insert(V);
4285 return true;
4286 }
4287 // Both operands are leaves so treat this binop as a source rather than
4288 // walking into it.
4289 return AddSource(V);
4290 };
4291 if (!Walk(VecOpEE, Walk) || Nodes.empty())
4292 return false;
4293
4294 bool IsIdempotent =
4295 CommonCallOp || (CommonBinOp && Instruction::isIdempotent(*CommonBinOp));
4296
4297 // For FP reductions, require reassoc on every binop and collect FMF.
4298 for (Value *V : Nodes) {
4299 auto *BinOp = dyn_cast<BinaryOperator>(V);
4300 if (!BinOp || !BinOp->getType()->isFPOrFPVectorTy())
4301 continue;
4302 if (!BinOp->hasAllowReassoc())
4303 return false;
4304 if (!IsFloatReduction) {
4305 CommonFMF = BinOp->getFastMathFlags();
4306 IsFloatReduction = true;
4307 } else {
4308 CommonFMF &= BinOp->getFastMathFlags();
4309 }
4310 }
4311
4312 // Top-down demanded elements. For each chain value, track which lanes feed
4313 // the extracted lane 0 and which feed it more than once. Reverse postorder
4314 // visits every use before its value. A binop forwards its demand to both
4315 // operands and a shuffle follows its mask back to the source lane.
4316 struct Demand {
4317 APInt Lanes;
4318 APInt Duplicates;
4319 };
4320 DenseMap<Value *, Demand> Demands;
4321 auto DemandOf = [&](Value *V) -> Demand & {
4322 unsigned N = cast<FixedVectorType>(V->getType())->getNumElements();
4323 Demand &D = Demands[V];
4324 if (D.Lanes.getBitWidth() != N)
4325 D.Lanes = D.Duplicates = APInt::getZero(N);
4326 return D;
4327 };
4328 DemandOf(VecOpEE).Lanes.setBit(0);
4329 for (Value *V : reverse(Nodes)) {
4330 Demand DV = Demands.lookup(V);
4331 if (DV.Lanes.isZero())
4332 continue;
4333 if (auto *SVI = dyn_cast<ShuffleVectorInst>(V)) {
4334 ArrayRef<int> Mask = SVI->getShuffleMask();
4335 Demand &DS = DemandOf(SVI->getOperand(0));
4336 for (unsigned I = 0, E = Mask.size(); I != E; ++I) {
4337 // Skip lanes that are undemanded or map to poison.
4338 if (!DV.Lanes[I] || Mask[I] < 0 ||
4339 (unsigned)Mask[I] >= DS.Lanes.getBitWidth())
4340 continue;
4341 if (DS.Lanes[Mask[I]] || DV.Duplicates[I])
4342 DS.Duplicates.setBit(Mask[I]);
4343 DS.Lanes.setBit(Mask[I]);
4344 }
4345 } else {
4346 auto *U = cast<User>(V);
4347 for (Value *Op : {U->getOperand(0), U->getOperand(1)}) {
4348 Demand &DOp = DemandOf(Op);
4349 // Lanes demanded through more than one path accumulate in Duplicates.
4350 DOp.Duplicates |= DV.Duplicates | (DOp.Lanes & DV.Lanes);
4351 DOp.Lanes |= DV.Lanes;
4352 }
4353 }
4354 }
4355
4356 // Reducing V replaces the entire chain, so every contribution to the result
4357 // must flow through V. Reject if anything above V reads outside the chain.
4358 auto CoversChain = [&](Value *V) {
4359 SmallVector<Value *, 8> Worklist(1, VecOpEE);
4360 SmallPtrSet<Value *, 8> Seen;
4361 Seen.insert(VecOpEE);
4362 while (!Worklist.empty()) {
4363 auto *U = cast<Instruction>(Worklist.pop_back_val());
4364 unsigned NumOps = isa<ShuffleVectorInst>(U) ? 1 : 2;
4365 for (unsigned I = 0; I != NumOps; ++I) {
4366 Value *Op = U->getOperand(I);
4367 if (Op == V || !Seen.insert(Op).second)
4368 continue;
4369 if (!Nodes.contains(Op))
4370 return false;
4371 Worklist.push_back(Op);
4372 }
4373 }
4374 return true;
4375 };
4376
4377 // Reduce a single cleanly demanded source if there is one, otherwise the
4378 // deepest intermediate that covers the chain.
4379 struct ReductionCut {
4380 Value *Src;
4381 APInt Elts;
4382 };
4383 std::optional<ReductionCut> Cut;
4384 for (Value *S : Sources) {
4385 auto It = Demands.find(S);
4386 if (It == Demands.end() || It->second.Lanes.isZero())
4387 continue;
4388 if (!IsIdempotent && !It->second.Duplicates.isZero()) {
4389 Cut.reset();
4390 break;
4391 }
4392 if (!Cut) {
4393 Cut = ReductionCut{S, It->second.Lanes};
4394 continue;
4395 }
4396 if (!isEquivBitcast(Cut->Src, S)) {
4397 Cut.reset();
4398 break;
4399 }
4400 if (!IsIdempotent && !(Cut->Elts & It->second.Lanes).isZero()) {
4401 Cut.reset();
4402 break;
4403 }
4404 Cut->Elts |= It->second.Lanes;
4405 }
4406 if (!Cut) {
4407 for (Value *V : Nodes) {
4409 continue;
4410 auto It = Demands.find(V);
4411 if (It == Demands.end() || !It->second.Lanes.isAllOnes())
4412 continue;
4413 if (!IsIdempotent && !It->second.Duplicates.isZero())
4414 continue;
4415 if (!CoversChain(V))
4416 continue;
4417 Cut = ReductionCut{V, It->second.Lanes};
4418 break;
4419 }
4420 }
4421 // Reducing one lane is just an extract and can refold forever.
4422 if (!Cut || Cut->Elts.popcount() < 2)
4423 return false;
4424
4425 Intrinsic::ID ReducedOp =
4426 (CommonCallOp ? getMinMaxReductionIntrinsicID(*CommonCallOp)
4427 : getReductionForBinop(*CommonBinOp));
4428 if (!ReducedOp)
4429 return false;
4430
4431 InstructionCost OrigCost = 0;
4432 for (Value *V : Nodes)
4434
4435 auto *SrcVT = cast<FixedVectorType>(Cut->Src->getType());
4436 bool IsPartialReduction = !Cut->Elts.isAllOnes();
4437 FixedVectorType *ReduceVecTy =
4438 IsPartialReduction
4439 ? FixedVectorType::get(FVT->getElementType(), Cut->Elts.popcount())
4440 : SrcVT;
4441
4442 SmallVector<int> ExtractMask;
4443 InstructionCost NewCost = 0;
4444 if (IsPartialReduction) {
4445 for (unsigned I = 0, E = Cut->Elts.getBitWidth(); I != E; ++I)
4446 if (Cut->Elts[I])
4447 ExtractMask.push_back(I);
4448 unsigned SubIdx = 0, SubLen;
4449 auto SK = Cut->Elts.isShiftedMask(SubIdx, SubLen)
4452 NewCost += TTI.getShuffleCost(SK, ReduceVecTy, SrcVT, ExtractMask, CostKind,
4453 SubIdx, ReduceVecTy);
4454 }
4455
4456 IntrinsicCostAttributes ICA(
4457 ReducedOp, ReduceVecTy->getElementType(),
4458 IsFloatReduction
4459 ? SmallVector<Type *, 2>{ReduceVecTy->getElementType(), ReduceVecTy}
4460 : SmallVector<Type *, 2>{ReduceVecTy},
4461 IsFloatReduction ? CommonFMF : FastMathFlags());
4462 NewCost += TTI.getIntrinsicInstrCost(ICA, CostKind);
4463
4464 LLVM_DEBUG(dbgs() << "Found reduction shuffle chain: " << I << "\n OldCost : "
4465 << OrigCost << " vs NewCost: " << NewCost << "\n");
4466
4467 if (!OrigCost.isValid() || !NewCost.isValid())
4468 return false;
4469
4470 if (VecOpEE->hasOneUse() ? (NewCost > OrigCost) : (NewCost >= OrigCost))
4471 return false;
4472
4473 Value *ReduceInput = Cut->Src;
4474 if (IsPartialReduction)
4475 ReduceInput = Builder.CreateShuffleVector(Cut->Src, ExtractMask);
4476
4477 Value *ReducedResult;
4478 if (IsFloatReduction) {
4480 *CommonBinOp, ReduceVecTy->getElementType(), /*AllowRHSConstant=*/false,
4481 CommonFMF.noSignedZeros());
4482 ReducedResult = Builder.CreateIntrinsic(ReducedOp, {ReduceVecTy},
4483 {Identity, ReduceInput}, CommonFMF);
4484 } else {
4485 ReducedResult =
4486 Builder.CreateIntrinsic(ReducedOp, {ReduceVecTy}, {ReduceInput});
4487 }
4488 replaceValue(I, *ReducedResult);
4489
4490 return true;
4491}
4492
4493/// Determine if its more efficient to fold:
4494/// reduce(trunc(x)) -> trunc(reduce(x)).
4495/// reduce(sext(x)) -> sext(reduce(x)).
4496/// reduce(zext(x)) -> zext(reduce(x)).
4497bool VectorCombine::foldCastFromReductions(Instruction &I) {
4498 auto *II = dyn_cast<IntrinsicInst>(&I);
4499 if (!II)
4500 return false;
4501
4502 bool TruncOnly = false;
4503 Intrinsic::ID IID = II->getIntrinsicID();
4504 switch (IID) {
4505 case Intrinsic::vector_reduce_add:
4506 case Intrinsic::vector_reduce_mul:
4507 TruncOnly = true;
4508 break;
4509 case Intrinsic::vector_reduce_and:
4510 case Intrinsic::vector_reduce_or:
4511 case Intrinsic::vector_reduce_xor:
4512 break;
4513 default:
4514 return false;
4515 }
4516
4517 unsigned ReductionOpc = getArithmeticReductionInstruction(IID);
4518 Value *ReductionSrc = I.getOperand(0);
4519
4520 Value *Src;
4521 if (!match(ReductionSrc, m_OneUse(m_Trunc(m_Value(Src)))) &&
4522 (TruncOnly || !match(ReductionSrc, m_OneUse(m_ZExtOrSExt(m_Value(Src))))))
4523 return false;
4524
4525 auto CastOpc =
4526 (Instruction::CastOps)cast<Instruction>(ReductionSrc)->getOpcode();
4527
4528 auto *SrcTy = cast<VectorType>(Src->getType());
4529 auto *ReductionSrcTy = cast<VectorType>(ReductionSrc->getType());
4530 Type *ResultTy = I.getType();
4531
4533 ReductionOpc, ReductionSrcTy, std::nullopt, CostKind);
4534 OldCost += TTI.getCastInstrCost(CastOpc, ReductionSrcTy, SrcTy,
4536 cast<CastInst>(ReductionSrc));
4537 InstructionCost NewCost =
4538 TTI.getArithmeticReductionCost(ReductionOpc, SrcTy, std::nullopt,
4539 CostKind) +
4540 TTI.getCastInstrCost(CastOpc, ResultTy, ReductionSrcTy->getScalarType(),
4542
4543 if (OldCost <= NewCost || !NewCost.isValid())
4544 return false;
4545
4546 Value *NewReduction = Builder.CreateIntrinsic(SrcTy->getScalarType(),
4547 II->getIntrinsicID(), {Src});
4548 Value *NewCast = Builder.CreateCast(CastOpc, NewReduction, ResultTy);
4549 replaceValue(I, *NewCast);
4550 return true;
4551}
4552
4553/// Fold:
4554/// icmp pred (reduce.{add,or,and,umax,umin}(signbit_extract(x))), C
4555/// into:
4556/// icmp sgt/slt (reduce.{or,umax,and,umin}(x)), -1/0
4557///
4558/// Sign-bit reductions produce values with known semantics:
4559/// - reduce.{or,umax}: 0 if no element is negative, 1 if any is
4560/// - reduce.{and,umin}: 1 if all elements are negative, 0 if any isn't
4561/// - reduce.add: count of negative elements (0 to NumElts)
4562///
4563/// Both lshr and ashr are supported:
4564/// - lshr produces 0 or 1, so reduce.add range is [0, N]
4565/// - ashr produces 0 or -1, so reduce.add range is [-N, 0]
4566///
4567/// The fold generalizes to multiple source vectors combined with the same
4568/// operation as the reduction. For example:
4569/// reduce.or(or(shr A, shr B)) conceptually extends the vector
4570/// For reduce.add, this changes the count to M*N where M is the number of
4571/// source vectors.
4572///
4573/// We transform to a direct sign check on the original vector using
4574/// reduce.{or,umax} or reduce.{and,umin}.
4575///
4576/// In spirit, it's similar to foldSignBitCheck in InstCombine.
4577bool VectorCombine::foldSignBitReductionCmp(Instruction &I) {
4578 CmpPredicate Pred;
4579 IntrinsicInst *ReduceOp;
4580 const APInt *CmpVal;
4581 if (!match(&I,
4582 m_ICmp(Pred, m_OneUse(m_AnyIntrinsic(ReduceOp)), m_APInt(CmpVal))))
4583 return false;
4584
4585 Intrinsic::ID OrigIID = ReduceOp->getIntrinsicID();
4586 switch (OrigIID) {
4587 case Intrinsic::vector_reduce_or:
4588 case Intrinsic::vector_reduce_umax:
4589 case Intrinsic::vector_reduce_and:
4590 case Intrinsic::vector_reduce_umin:
4591 case Intrinsic::vector_reduce_add:
4592 break;
4593 default:
4594 return false;
4595 }
4596
4597 Value *ReductionSrc = ReduceOp->getArgOperand(0);
4598 auto *VecTy = dyn_cast<FixedVectorType>(ReductionSrc->getType());
4599 if (!VecTy)
4600 return false;
4601
4602 unsigned BitWidth = VecTy->getScalarSizeInBits();
4603 if (BitWidth == 1)
4604 return false;
4605
4606 unsigned NumElts = VecTy->getNumElements();
4607
4608 // Determine the expected tree opcode for multi-vector patterns.
4609 // The tree opcode must match the reduction's underlying operation.
4610 //
4611 // TODO: for pairs of equivalent operators, we should match both,
4612 // not only the most common.
4613 Instruction::BinaryOps TreeOpcode;
4614 switch (OrigIID) {
4615 case Intrinsic::vector_reduce_or:
4616 case Intrinsic::vector_reduce_umax:
4617 TreeOpcode = Instruction::Or;
4618 break;
4619 case Intrinsic::vector_reduce_and:
4620 case Intrinsic::vector_reduce_umin:
4621 TreeOpcode = Instruction::And;
4622 break;
4623 case Intrinsic::vector_reduce_add:
4624 TreeOpcode = Instruction::Add;
4625 break;
4626 default:
4627 llvm_unreachable("Unexpected intrinsic");
4628 }
4629
4630 // Collect sign-bit extraction leaves from an associative tree of TreeOpcode.
4631 // The tree conceptually extends the vector being reduced.
4632 SmallVector<Value *, 8> Worklist;
4633 SmallVector<Value *, 8> Sources; // Original vectors (X in shr X, BW-1)
4634 Worklist.push_back(ReductionSrc);
4635 std::optional<bool> IsAShr;
4636 constexpr unsigned MaxSources = 8;
4637
4638 // Calculate old cost: all shifts + tree ops + reduction
4639 InstructionCost OldCost = TTI.getInstructionCost(ReduceOp, CostKind);
4640
4641 while (!Worklist.empty() && Worklist.size() <= MaxSources &&
4642 Sources.size() <= MaxSources) {
4643 Value *V = Worklist.pop_back_val();
4644
4645 // Try to match sign-bit extraction: shr X, (bitwidth-1)
4646 Value *X;
4647 if (match(V, m_OneUse(m_Shr(m_Value(X), m_SpecificInt(BitWidth - 1))))) {
4648 auto *Shr = cast<Instruction>(V);
4649
4650 // All shifts must be the same type (all lshr or all ashr)
4651 bool ThisIsAShr = Shr->getOpcode() == Instruction::AShr;
4652 if (!IsAShr)
4653 IsAShr = ThisIsAShr;
4654 else if (*IsAShr != ThisIsAShr)
4655 return false;
4656
4657 Sources.push_back(X);
4658
4659 // As part of the fold, we remove all of the shifts, so we need to keep
4660 // track of their costs.
4661 OldCost += TTI.getInstructionCost(Shr, CostKind);
4662
4663 continue;
4664 }
4665
4666 // Try to extend through a tree node of the expected opcode
4667 Value *A, *B;
4668 if (!match(V, m_OneUse(m_BinOp(TreeOpcode, m_Value(A), m_Value(B)))))
4669 return false;
4670
4671 // We are potentially replacing these operations as well, so we add them
4672 // to the costs.
4674
4675 Worklist.push_back(A);
4676 Worklist.push_back(B);
4677 }
4678
4679 // Must have at least one source and not exceed limit
4680 if (Sources.empty() || Sources.size() > MaxSources ||
4681 Worklist.size() > MaxSources || !IsAShr)
4682 return false;
4683
4684 unsigned NumSources = Sources.size();
4685
4686 // For reduce.add, the total count must fit as a signed integer.
4687 // Range is [0, M*N] for lshr or [-M*N, 0] for ashr.
4688 if (OrigIID == Intrinsic::vector_reduce_add &&
4689 !isIntN(BitWidth, NumSources * NumElts))
4690 return false;
4691
4692 // Compute the boundary value when all elements are negative:
4693 // - Per-element contribution: 1 for lshr, -1 for ashr
4694 // - For add: M*N (total elements across all sources); for others: just 1
4695 unsigned Count =
4696 (OrigIID == Intrinsic::vector_reduce_add) ? NumSources * NumElts : 1;
4697 APInt NegativeVal(CmpVal->getBitWidth(), Count);
4698 if (*IsAShr)
4699 NegativeVal.negate();
4700
4701 // Range is [min(0, AllNegVal), max(0, AllNegVal)]
4702 APInt Zero = APInt::getZero(CmpVal->getBitWidth());
4703 APInt RangeLow = APIntOps::smin(Zero, NegativeVal);
4704 APInt RangeHigh = APIntOps::smax(Zero, NegativeVal);
4705
4706 // Determine comparison semantics:
4707 // - IsEq: true for equality test, false for inequality
4708 // - TestsNegative: true if testing against AllNegVal, false for zero
4709 //
4710 // In addition to EQ/NE against 0 or AllNegVal, we support inequalities
4711 // that fold to boundary tests given the narrow value range:
4712 // < RangeHigh -> != RangeHigh
4713 // > RangeHigh-1 -> == RangeHigh
4714 // > RangeLow -> != RangeLow
4715 // < RangeLow+1 -> == RangeLow
4716 //
4717 // For inequalities, we work with signed predicates only. Unsigned predicates
4718 // are canonicalized to signed when the range is non-negative (where they are
4719 // equivalent). When the range includes negative values, unsigned predicates
4720 // would have different semantics due to wrap-around, so we reject them.
4721 if (!ICmpInst::isEquality(Pred) && !ICmpInst::isSigned(Pred)) {
4722 if (RangeLow.isNegative())
4723 return false;
4724 Pred = ICmpInst::getSignedPredicate(Pred);
4725 }
4726
4727 bool IsEq;
4728 bool TestsNegative;
4729 if (ICmpInst::isEquality(Pred)) {
4730 if (CmpVal->isZero()) {
4731 TestsNegative = false;
4732 } else if (*CmpVal == NegativeVal) {
4733 TestsNegative = true;
4734 } else {
4735 return false;
4736 }
4737 IsEq = Pred == ICmpInst::ICMP_EQ;
4738 } else if (Pred == ICmpInst::ICMP_SLT && *CmpVal == RangeHigh) {
4739 IsEq = false;
4740 TestsNegative = (RangeHigh == NegativeVal);
4741 } else if (Pred == ICmpInst::ICMP_SGT && *CmpVal == RangeHigh - 1) {
4742 IsEq = true;
4743 TestsNegative = (RangeHigh == NegativeVal);
4744 } else if (Pred == ICmpInst::ICMP_SGT && *CmpVal == RangeLow) {
4745 IsEq = false;
4746 TestsNegative = (RangeLow == NegativeVal);
4747 } else if (Pred == ICmpInst::ICMP_SLT && *CmpVal == RangeLow + 1) {
4748 IsEq = true;
4749 TestsNegative = (RangeLow == NegativeVal);
4750 } else {
4751 return false;
4752 }
4753
4754 // For this fold we support four types of checks:
4755 //
4756 // 1. All lanes are negative - AllNeg
4757 // 2. All lanes are non-negative - AllNonNeg
4758 // 3. At least one negative lane - AnyNeg
4759 // 4. At least one non-negative lane - AnyNonNeg
4760 //
4761 // For each case, we can generate the following code:
4762 //
4763 // 1. AllNeg - reduce.and/umin(X) < 0
4764 // 2. AllNonNeg - reduce.or/umax(X) > -1
4765 // 3. AnyNeg - reduce.or/umax(X) < 0
4766 // 4. AnyNonNeg - reduce.and/umin(X) > -1
4767 //
4768 // The table below shows the aggregation of all supported cases
4769 // using these four cases.
4770 //
4771 // Reduction | == 0 | != 0 | == MAX | != MAX
4772 // ------------+-----------+-----------+-----------+-----------
4773 // or/umax | AllNonNeg | AnyNeg | AnyNeg | AllNonNeg
4774 // and/umin | AnyNonNeg | AllNeg | AllNeg | AnyNonNeg
4775 // add | AllNonNeg | AnyNeg | AllNeg | AnyNonNeg
4776 //
4777 // NOTE: MAX = 1 for or/and/umax/umin, and the vector size N for add
4778 //
4779 // For easier codegen and check inversion, we use the following encoding:
4780 //
4781 // 1. Bit-3 === requires or/umax (1) or and/umin (0) check
4782 // 2. Bit-2 === checks < 0 (1) or > -1 (0)
4783 // 3. Bit-1 === universal (1) or existential (0) check
4784 //
4785 // AnyNeg = 0b110: uses or/umax, checks negative, any-check
4786 // AllNonNeg = 0b101: uses or/umax, checks non-neg, all-check
4787 // AnyNonNeg = 0b000: uses and/umin, checks non-neg, any-check
4788 // AllNeg = 0b011: uses and/umin, checks negative, all-check
4789 //
4790 // XOR with 0b011 inverts the check (swaps all/any and neg/non-neg).
4791 //
4792 enum CheckKind : unsigned {
4793 AnyNonNeg = 0b000,
4794 AllNeg = 0b011,
4795 AllNonNeg = 0b101,
4796 AnyNeg = 0b110,
4797 };
4798 // Return true if we fold this check into or/umax and false for and/umin
4799 auto RequiresOr = [](CheckKind C) -> bool { return C & 0b100; };
4800 // Return true if we should check if result is negative and false otherwise
4801 auto IsNegativeCheck = [](CheckKind C) -> bool { return C & 0b010; };
4802 // Logically invert the check
4803 auto Invert = [](CheckKind C) { return CheckKind(C ^ 0b011); };
4804
4805 CheckKind Base;
4806 switch (OrigIID) {
4807 case Intrinsic::vector_reduce_or:
4808 case Intrinsic::vector_reduce_umax:
4809 Base = TestsNegative ? AnyNeg : AllNonNeg;
4810 break;
4811 case Intrinsic::vector_reduce_and:
4812 case Intrinsic::vector_reduce_umin:
4813 Base = TestsNegative ? AllNeg : AnyNonNeg;
4814 break;
4815 case Intrinsic::vector_reduce_add:
4816 Base = TestsNegative ? AllNeg : AllNonNeg;
4817 break;
4818 default:
4819 llvm_unreachable("Unexpected intrinsic");
4820 }
4821
4822 CheckKind Check = IsEq ? Base : Invert(Base);
4823
4824 auto PickCheaper = [&](Intrinsic::ID Arith, Intrinsic::ID MinMax) {
4825 InstructionCost ArithCost =
4827 VecTy, std::nullopt, CostKind);
4828 InstructionCost MinMaxCost =
4830 FastMathFlags(), CostKind);
4831 return ArithCost <= MinMaxCost ? std::make_pair(Arith, ArithCost)
4832 : std::make_pair(MinMax, MinMaxCost);
4833 };
4834
4835 // Choose output reduction based on encoding's MSB
4836 auto [NewIID, NewCost] = RequiresOr(Check)
4837 ? PickCheaper(Intrinsic::vector_reduce_or,
4838 Intrinsic::vector_reduce_umax)
4839 : PickCheaper(Intrinsic::vector_reduce_and,
4840 Intrinsic::vector_reduce_umin);
4841
4842 // Add cost of combining multiple sources with or/and
4843 if (NumSources > 1) {
4844 unsigned CombineOpc =
4845 RequiresOr(Check) ? Instruction::Or : Instruction::And;
4846 NewCost += TTI.getArithmeticInstrCost(CombineOpc, VecTy, CostKind) *
4847 (NumSources - 1);
4848 }
4849
4850 LLVM_DEBUG(dbgs() << "Found sign-bit reduction cmp: " << I << "\n OldCost: "
4851 << OldCost << " vs NewCost: " << NewCost << "\n");
4852
4853 if (NewCost > OldCost)
4854 return false;
4855
4856 // Generate the combined input and reduction
4857 Builder.SetInsertPoint(&I);
4858 Type *ScalarTy = VecTy->getScalarType();
4859
4860 Value *Input;
4861 if (NumSources == 1) {
4862 Input = Sources[0];
4863 } else {
4864 // Combine sources with or/and based on check type
4865 Input = RequiresOr(Check) ? Builder.CreateOr(Sources)
4866 : Builder.CreateAnd(Sources);
4867 }
4868
4869 Value *NewReduce = Builder.CreateIntrinsic(ScalarTy, NewIID, {Input});
4870 Value *NewCmp = IsNegativeCheck(Check) ? Builder.CreateIsNeg(NewReduce)
4871 : Builder.CreateIsNotNeg(NewReduce);
4872 replaceValue(I, *NewCmp);
4873 return true;
4874}
4875
4876/// Fold a zero test of reduce.or or reduce.umax into a boolean reduction.
4877///
4878/// Vectorization may produce IR that compares the result of a scalar reduction
4879/// with zero. Depending on the target, lowering a reduction and a scalar
4880/// comparison separately can cost more than reducing lane-wise comparison
4881/// results. This fold creates the latter form only when it is not costlier.
4882///
4883/// Before:
4884/// %r = call iT @llvm.vector.reduce.or.vNiT(<N x iT> %x)
4885/// %cmp = icmp ne iT %r, 0
4886///
4887/// After:
4888/// %lane.cmp = icmp ne <N x iT> %x, zeroinitializer
4889/// %cmp = call i1 @llvm.vector.reduce.or.vNi1(<N x i1> %lane.cmp)
4890///
4891/// `reduce.or` and `reduce.umax` are non-zero when at least one lane is
4892/// non-zero. Therefore, `icmp ne` uses the existential `reduce.or` test.
4893/// Conversely, `icmp eq` must check that every lane is zero, so it uses the
4894/// universal `reduce.and` test.
4895///
4896/// Before:
4897/// %r = call iT @llvm.vector.reduce.umax.vNiT(<N x iT> %x)
4898/// %cmp = icmp eq iT %r, 0
4899///
4900/// After:
4901/// %lane.cmp = icmp eq <N x iT> %x, zeroinitializer
4902/// %cmp = call i1 @llvm.vector.reduce.and.vNi1(<N x i1> %lane.cmp)
4903bool VectorCombine::foldReductionZeroTest(Instruction &I) {
4904 CmpPredicate Pred;
4905 Value *Op;
4906
4907 if (!match(&I, m_c_ICmp(Pred, m_Value(Op), m_Zero())) ||
4908 !ICmpInst::isEquality(Pred))
4909 return false;
4910
4911 auto *II = dyn_cast<IntrinsicInst>(Op);
4912 if (!II || !II->hasOneUse())
4913 return false;
4914
4915 auto ReduceID = II->getIntrinsicID();
4916 if (ReduceID != Intrinsic::vector_reduce_or &&
4917 ReduceID != Intrinsic::vector_reduce_umax)
4918 return false;
4919
4920 Value *Vec = II->getArgOperand(0);
4921 auto *VecTy = dyn_cast<FixedVectorType>(Vec->getType());
4922 if (!VecTy || !VecTy->getElementType()->isIntegerTy())
4923 return false;
4924
4925 // Map the scalar zero test to an any-lane or all-lane boolean reduction.
4926 Intrinsic::ID NewIID = (Pred == ICmpInst::ICMP_NE)
4927 ? Intrinsic::vector_reduce_or
4928 : Intrinsic::vector_reduce_and;
4929
4930 // This is not an unconditional canonicalization: compare the cost of the
4931 // original scalar reduction and compare with the vector compare and i1
4932 // reduction replacement for both reduce.or and reduce.umax.
4935
4936 auto *CmpTy = cast<VectorType>(CmpInst::makeCmpResultType(VecTy));
4937 InstructionCost NewCost =
4938 TTI.getCmpSelInstrCost(Instruction::ICmp, VecTy, CmpTy, Pred, CostKind);
4940 getArithmeticReductionInstruction(NewIID), CmpTy, std::nullopt, CostKind);
4941
4942 LLVM_DEBUG(dbgs() << "Found a reduction zero test: " << I << "\n OldCost: "
4943 << OldCost << " vs NewCost: " << NewCost << "\n");
4944
4945 if (!OldCost.isValid() || !NewCost.isValid() || NewCost > OldCost)
4946 return false;
4947
4948 Builder.SetInsertPoint(&I);
4949 Value *NewCmp = Builder.CreateICmp(Pred, Vec, Constant::getNullValue(VecTy));
4950 Value *NewReduce = Builder.CreateIntrinsic(NewIID, {CmpTy}, {NewCmp});
4951 replaceValue(I, *NewReduce);
4952 return true;
4953}
4954
4955/// vector.reduce.OP f(X_i) == 0 -> vector.reduce.OP X_i == 0
4956///
4957/// We can prove it for cases when:
4958///
4959/// 1. OP X_i == 0 <=> \forall i \in [1, N] X_i == 0
4960/// 1'. OP X_i == 0 <=> \exists j \in [1, N] X_j == 0
4961/// 2. f(x) == 0 <=> x == 0
4962///
4963/// From 1 and 2 (or 1' and 2), we can infer that
4964///
4965/// OP f(X_i) == 0 <=> OP X_i == 0.
4966///
4967/// (1)
4968/// OP f(X_i) == 0 <=> \forall i \in [1, N] f(X_i) == 0
4969/// (2)
4970/// <=> \forall i \in [1, N] X_i == 0
4971/// (1)
4972/// <=> OP(X_i) == 0
4973///
4974/// For some of the OP's and f's, we need to have domain constraints on X
4975/// to ensure properties 1 (or 1') and 2.
4976bool VectorCombine::foldICmpEqZeroVectorReduce(Instruction &I) {
4977 CmpPredicate Pred;
4978 Value *Op;
4979 if (!match(&I, m_ICmp(Pred, m_Value(Op), m_Zero())) ||
4980 !ICmpInst::isEquality(Pred))
4981 return false;
4982
4983 auto *II = dyn_cast<IntrinsicInst>(Op);
4984 if (!II)
4985 return false;
4986
4987 switch (II->getIntrinsicID()) {
4988 case Intrinsic::vector_reduce_add:
4989 case Intrinsic::vector_reduce_or:
4990 case Intrinsic::vector_reduce_umin:
4991 case Intrinsic::vector_reduce_umax:
4992 case Intrinsic::vector_reduce_smin:
4993 case Intrinsic::vector_reduce_smax:
4994 break;
4995 default:
4996 return false;
4997 }
4998
4999 Value *InnerOp = II->getArgOperand(0);
5000
5001 // TODO: fixed vector type might be too restrictive
5002 if (!II->hasOneUse() || !isa<FixedVectorType>(InnerOp->getType()))
5003 return false;
5004
5005 Value *X = nullptr;
5006
5007 // Check for zero-preserving operations where f(x) = 0 <=> x = 0
5008 //
5009 // 1. f(x) = shl nuw x, y for arbitrary y
5010 // 2. f(x) = mul nuw x, c for defined c != 0
5011 // 3. f(x) = zext x
5012 // 4. f(x) = sext x
5013 // 5. f(x) = neg x
5014 //
5015 if (!(match(InnerOp, m_NUWShl(m_Value(X), m_Value())) || // Case 1
5016 match(InnerOp, m_NUWMul(m_Value(X), m_NonZeroInt())) || // Case 2
5017 match(InnerOp, m_ZExt(m_Value(X))) || // Case 3
5018 match(InnerOp, m_SExt(m_Value(X))) || // Case 4
5019 match(InnerOp, m_Neg(m_Value(X))) // Case 5
5020 ))
5021 return false;
5022
5023 SimplifyQuery S = SQ.getWithInstruction(&I);
5024 auto *XTy = cast<FixedVectorType>(X->getType());
5025
5026 // Check for domain constraints for all supported reductions.
5027 //
5028 // a. OR X_i - has property 1 for every X
5029 // b. UMAX X_i - has property 1 for every X
5030 // c. UMIN X_i - has property 1' for every X
5031 // d. SMAX X_i - has property 1 for X >= 0
5032 // e. SMIN X_i - has property 1' for X >= 0
5033 // f. ADD X_i - has property 1 for X >= 0 && ADD X_i doesn't sign wrap
5034 //
5035 // In order for the proof to work, we need 1 (or 1') to be true for both
5036 // OP f(X_i) and OP X_i and that's why below we check constraints twice.
5037 //
5038 // NOTE: ADD X_i holds property 1 for a mirror case as well, i.e. when
5039 // X <= 0 && ADD X_i doesn't sign wrap. However, due to the nature
5040 // of known bits, we can't reasonably hold knowledge of "either 0
5041 // or negative".
5042 switch (II->getIntrinsicID()) {
5043 case Intrinsic::vector_reduce_add: {
5044 // We need to check that both X_i and f(X_i) have enough leading
5045 // zeros to not overflow.
5046 KnownBits KnownX = computeKnownBits(X, S);
5047 KnownBits KnownFX = computeKnownBits(InnerOp, S);
5048 unsigned NumElems = XTy->getNumElements();
5049 // Adding N elements loses at most ceil(log2(N)) leading bits.
5050 unsigned LostBits = Log2_32_Ceil(NumElems);
5051 unsigned LeadingZerosX = KnownX.countMinLeadingZeros();
5052 unsigned LeadingZerosFX = KnownFX.countMinLeadingZeros();
5053 // Need at least one leading zero left after summation to ensure no overflow
5054 if (LeadingZerosX <= LostBits || LeadingZerosFX <= LostBits)
5055 return false;
5056
5057 // We are not checking whether X or f(X) are positive explicitly because
5058 // we implicitly checked for it when we checked if both cases have enough
5059 // leading zeros to not wrap addition.
5060 break;
5061 }
5062 case Intrinsic::vector_reduce_smin:
5063 case Intrinsic::vector_reduce_smax:
5064 // Check whether X >= 0 and f(X) >= 0
5065 if (!isKnownNonNegative(InnerOp, S) || !isKnownNonNegative(X, S))
5066 return false;
5067
5068 break;
5069 default:
5070 break;
5071 };
5072
5073 LLVM_DEBUG(dbgs() << "Found a reduction to 0 comparison with removable op: "
5074 << *II << "\n");
5075
5076 // For zext/sext, check if the transform is profitable using cost model.
5077 // For other operations (shl, mul, neg), we're removing an instruction
5078 // while keeping the same reduction type, so it's always profitable.
5079 if (isa<ZExtInst>(InnerOp) || isa<SExtInst>(InnerOp)) {
5080 auto *FXTy = cast<FixedVectorType>(InnerOp->getType());
5081 Intrinsic::ID IID = II->getIntrinsicID();
5082
5084 cast<CastInst>(InnerOp)->getOpcode(), FXTy, XTy,
5086
5087 InstructionCost OldReduceCost, NewReduceCost;
5088 switch (IID) {
5089 case Intrinsic::vector_reduce_add:
5090 case Intrinsic::vector_reduce_or:
5091 OldReduceCost = TTI.getArithmeticReductionCost(
5092 getArithmeticReductionInstruction(IID), FXTy, std::nullopt, CostKind);
5093 NewReduceCost = TTI.getArithmeticReductionCost(
5094 getArithmeticReductionInstruction(IID), XTy, std::nullopt, CostKind);
5095 break;
5096 case Intrinsic::vector_reduce_umin:
5097 case Intrinsic::vector_reduce_umax:
5098 case Intrinsic::vector_reduce_smin:
5099 case Intrinsic::vector_reduce_smax:
5100 OldReduceCost = TTI.getMinMaxReductionCost(
5101 getMinMaxReductionIntrinsicOp(IID), FXTy, FastMathFlags(), CostKind);
5102 NewReduceCost = TTI.getMinMaxReductionCost(
5103 getMinMaxReductionIntrinsicOp(IID), XTy, FastMathFlags(), CostKind);
5104 break;
5105 default:
5106 llvm_unreachable("Unexpected reduction");
5107 }
5108
5109 InstructionCost OldCost = OldReduceCost + ExtCost;
5110 InstructionCost NewCost =
5111 NewReduceCost + (InnerOp->hasOneUse() ? 0 : ExtCost);
5112
5113 LLVM_DEBUG(dbgs() << "Found a removable extension before reduction: "
5114 << *InnerOp << "\n OldCost: " << OldCost
5115 << " vs NewCost: " << NewCost << "\n");
5116
5117 // We consider transformation to still be potentially beneficial even
5118 // when the costs are the same because we might remove a use from f(X)
5119 // and unlock other optimizations. Equal costs would just mean that we
5120 // didn't make it worse in the worst case.
5121 if (NewCost > OldCost)
5122 return false;
5123 }
5124
5125 // Since we support zext and sext as f, we might change the scalar type
5126 // of the intrinsic.
5127 Type *Ty = XTy->getScalarType();
5128 Value *NewReduce = Builder.CreateIntrinsic(Ty, II->getIntrinsicID(), {X});
5129 Value *NewCmp =
5130 Builder.CreateICmp(Pred, NewReduce, ConstantInt::getNullValue(Ty));
5131 replaceValue(I, *NewCmp);
5132 return true;
5133}
5134
5135/// Fold comparisons of reduce.or/reduce.and with reduce.umax/reduce.umin
5136/// based on cost, preserving the comparison semantics.
5137///
5138/// We use two fundamental properties for each pair:
5139///
5140/// 1. or(X) == 0 <=> umax(X) == 0
5141/// 2. or(X) == 1 <=> umax(X) == 1
5142/// 3. sign(or(X)) == sign(umax(X))
5143///
5144/// 1. and(X) == -1 <=> umin(X) == -1
5145/// 2. and(X) == -2 <=> umin(X) == -2
5146/// 3. sign(and(X)) == sign(umin(X))
5147///
5148/// From these we can infer the following transformations:
5149/// a. or(X) ==/!= 0 <-> umax(X) ==/!= 0
5150/// b. or(X) s< 0 <-> umax(X) s< 0
5151/// c. or(X) s> -1 <-> umax(X) s> -1
5152/// d. or(X) s< 1 <-> umax(X) s< 1
5153/// e. or(X) ==/!= 1 <-> umax(X) ==/!= 1
5154/// f. or(X) s< 2 <-> umax(X) s< 2
5155/// g. and(X) ==/!= -1 <-> umin(X) ==/!= -1
5156/// h. and(X) s< 0 <-> umin(X) s< 0
5157/// i. and(X) s> -1 <-> umin(X) s> -1
5158/// j. and(X) s> -2 <-> umin(X) s> -2
5159/// k. and(X) ==/!= -2 <-> umin(X) ==/!= -2
5160/// l. and(X) s> -3 <-> umin(X) s> -3
5161///
5162bool VectorCombine::foldEquivalentReductionCmp(Instruction &I) {
5163 CmpPredicate Pred;
5164 Value *ReduceOp;
5165 const APInt *CmpVal;
5166 if (!match(&I, m_ICmp(Pred, m_Value(ReduceOp), m_APInt(CmpVal))))
5167 return false;
5168
5169 auto *II = dyn_cast<IntrinsicInst>(ReduceOp);
5170 if (!II || !II->hasOneUse())
5171 return false;
5172
5173 const auto IsValidOrUmaxCmp = [&]() {
5174 // or === umax for i1
5175 if (CmpVal->getBitWidth() == 1)
5176 return true;
5177
5178 // Cases a and e
5179 bool IsEquality =
5180 (CmpVal->isZero() || CmpVal->isOne()) && ICmpInst::isEquality(Pred);
5181 // Case c
5182 bool IsPositive = CmpVal->isAllOnes() && Pred == ICmpInst::ICMP_SGT;
5183 // Cases b, d, and f
5184 bool IsNegative = (CmpVal->isZero() || CmpVal->isOne() || *CmpVal == 2) &&
5185 Pred == ICmpInst::ICMP_SLT;
5186 return IsEquality || IsPositive || IsNegative;
5187 };
5188
5189 const auto IsValidAndUminCmp = [&]() {
5190 // and === umin for i1
5191 if (CmpVal->getBitWidth() == 1)
5192 return true;
5193
5194 const auto LeadingOnes = CmpVal->countl_one();
5195
5196 // Cases g and k
5197 bool IsEquality =
5198 (CmpVal->isAllOnes() || LeadingOnes + 1 == CmpVal->getBitWidth()) &&
5200 // Case h
5201 bool IsNegative = CmpVal->isZero() && Pred == ICmpInst::ICMP_SLT;
5202 // Cases i, j, and l
5203 bool IsPositive =
5204 // if the number has at least N - 2 leading ones
5205 // and the two LSBs are:
5206 // - 1 x 1 -> -1
5207 // - 1 x 0 -> -2
5208 // - 0 x 1 -> -3
5209 LeadingOnes + 2 >= CmpVal->getBitWidth() &&
5210 ((*CmpVal)[0] || (*CmpVal)[1]) && Pred == ICmpInst::ICMP_SGT;
5211 return IsEquality || IsNegative || IsPositive;
5212 };
5213
5214 Intrinsic::ID OriginalIID = II->getIntrinsicID();
5215 Intrinsic::ID AlternativeIID;
5216
5217 // Check if this is a valid comparison pattern and determine the alternate
5218 // reduction intrinsic.
5219 switch (OriginalIID) {
5220 case Intrinsic::vector_reduce_or:
5221 if (!IsValidOrUmaxCmp())
5222 return false;
5223 AlternativeIID = Intrinsic::vector_reduce_umax;
5224 break;
5225 case Intrinsic::vector_reduce_umax:
5226 if (!IsValidOrUmaxCmp())
5227 return false;
5228 AlternativeIID = Intrinsic::vector_reduce_or;
5229 break;
5230 case Intrinsic::vector_reduce_and:
5231 if (!IsValidAndUminCmp())
5232 return false;
5233 AlternativeIID = Intrinsic::vector_reduce_umin;
5234 break;
5235 case Intrinsic::vector_reduce_umin:
5236 if (!IsValidAndUminCmp())
5237 return false;
5238 AlternativeIID = Intrinsic::vector_reduce_and;
5239 break;
5240 default:
5241 return false;
5242 }
5243
5244 Value *X = II->getArgOperand(0);
5245 auto *VecTy = dyn_cast<FixedVectorType>(X->getType());
5246 if (!VecTy)
5247 return false;
5248
5249 const auto GetReductionCost = [&](Intrinsic::ID IID) -> InstructionCost {
5250 unsigned ReductionOpc = getArithmeticReductionInstruction(IID);
5251 if (ReductionOpc != Instruction::ICmp)
5252 return TTI.getArithmeticReductionCost(ReductionOpc, VecTy, std::nullopt,
5253 CostKind);
5255 FastMathFlags(), CostKind);
5256 };
5257
5258 InstructionCost OrigCost = GetReductionCost(OriginalIID);
5259 InstructionCost AltCost = GetReductionCost(AlternativeIID);
5260
5261 LLVM_DEBUG(dbgs() << "Found equivalent reduction cmp: " << I
5262 << "\n OrigCost: " << OrigCost
5263 << " vs AltCost: " << AltCost << "\n");
5264
5265 if (AltCost >= OrigCost)
5266 return false;
5267
5268 Builder.SetInsertPoint(&I);
5269 Type *ScalarTy = VecTy->getScalarType();
5270 Value *NewReduce = Builder.CreateIntrinsic(ScalarTy, AlternativeIID, {X});
5271 Value *NewCmp =
5272 Builder.CreateICmp(Pred, NewReduce, ConstantInt::get(ScalarTy, *CmpVal));
5273
5274 replaceValue(I, *NewCmp);
5275 return true;
5276}
5277
5278/// Used by foldReduceAddCmpZero to check if we can prove that a value is
5279/// non-positive.
5280/// KnownBits cannot see sext <? x i1> as non-positive: each top bit equals a
5281/// single unknown input bit, which a per-bit lattice cannot track. The fold's
5282/// target shape is popcount-style sums of <N x i1> valid/invalid masks (e.g.
5283/// ray-intersection hits) tested for any-hit.
5284/// Previous attempts to approximate the known bits of such expressions were
5285/// using a fully recursive value tracking approach to infer a constant range
5286/// but ultimately turned to be too expensive in compile time.
5287static bool isKnownNonPositive(const Value *V, const SimplifyQuery &SQ,
5288 unsigned Depth = 0) {
5289 constexpr unsigned MaxLocalDepth = 2;
5290 if (Depth > MaxLocalDepth)
5291 return false;
5292
5293 auto NumSignBits = [&](const Value *X) {
5294 return ComputeNumSignBits(X, SQ.DL, SQ.AC, SQ.CxtI, SQ.DT);
5295 };
5296 if (NumSignBits(V) == V->getType()->getScalarSizeInBits())
5297 return true;
5298
5299 Value *A, *B;
5300 if (match(V, m_Add(m_Value(A), m_Value(B))))
5301 return NumSignBits(A) >= 2 && NumSignBits(B) >= 2 &&
5302 isKnownNonPositive(A, SQ, Depth + 1) &&
5303 isKnownNonPositive(B, SQ, Depth + 1);
5304
5305 return computeKnownBits(V, SQ).isNonPositive();
5306}
5307
5308/// Fold (icmp pred (reduce.add X), 0) to (icmp pred' (reduce.or X), 0) when X
5309/// has lanes known to all be non-negative or all non-positive, so that
5310/// sum == 0 iff every lane is 0. Falls back to reduce.umax if reduce.or is
5311/// more expensive on the target.
5312bool VectorCombine::foldReduceAddCmpZero(Instruction &I) {
5313 CmpPredicate Pred;
5314 Value *Vec;
5315 if (!match(&I, m_ICmp(Pred,
5317 m_Value(Vec))),
5318 m_Zero())))
5319 return false;
5320
5321 auto *VecTy = dyn_cast<FixedVectorType>(Vec->getType());
5322 if (!VecTy || VecTy->getNumElements() < 2)
5323 return false;
5324
5325 SimplifyQuery Q = SQ.getWithInstruction(&I);
5326 bool IsNonNegative = isKnownNonNegative(Vec, Q);
5327 bool IsNonPositive = !IsNonNegative && isKnownNonPositive(Vec, Q);
5328 if (!IsNonNegative && !IsNonPositive)
5329 return false;
5330
5331 // Summing NumElts lanes can consume up to log2(NumElts) sign bits. Require
5332 // strictly more headroom than that so the sum cannot wrap to zero.
5333 unsigned NumElts = VecTy->getNumElements();
5334 unsigned NumSignBits = ComputeNumSignBits(Vec, *DL, SQ.AC, &I, &DT);
5335 if (Log2_32(NumElts) >= NumSignBits)
5336 return false;
5337
5338 ICmpInst::Predicate NewPred;
5339 switch (Pred) {
5340 case ICmpInst::ICMP_EQ:
5341 case ICmpInst::ICMP_ULE:
5342 case ICmpInst::ICMP_SLE:
5343 case ICmpInst::ICMP_SGE:
5344 NewPred = ICmpInst::ICMP_EQ;
5345 break;
5346 case ICmpInst::ICMP_NE:
5347 case ICmpInst::ICMP_UGT:
5348 case ICmpInst::ICMP_SGT:
5349 case ICmpInst::ICMP_SLT:
5350 NewPred = ICmpInst::ICMP_NE;
5351 break;
5352 default:
5353 return false;
5354 }
5355
5356 // SGT and SLE on a non-positive tree, and SLT and SGE on a non-negative
5357 // tree, are tautologies (always true or always false). Leave those to
5358 // InstCombine rather than mapping them here. Remaining signed inequalities
5359 // also need one extra sign bit so the sum cannot flip sign.
5360 if (!IsNonNegative &&
5361 (Pred == ICmpInst::ICMP_SGT || Pred == ICmpInst::ICMP_SLE))
5362 return false;
5363 if (!IsNonPositive &&
5364 (Pred == ICmpInst::ICMP_SLT || Pred == ICmpInst::ICMP_SGE))
5365 return false;
5366 if ((Pred == ICmpInst::ICMP_SGT || Pred == ICmpInst::ICMP_SLE ||
5367 Pred == ICmpInst::ICMP_SLT || Pred == ICmpInst::ICMP_SGE) &&
5368 Log2_32(NumElts) >= NumSignBits - 1)
5369 return false;
5370
5372 Instruction::Add, VecTy, std::nullopt, CostKind);
5374 Instruction::Or, VecTy, std::nullopt, CostKind);
5376 Intrinsic::umax, VecTy, FastMathFlags(), CostKind);
5377 if (!OrCost.isValid() && !UmaxCost.isValid())
5378 return false;
5379 bool UseOr = OrCost.isValid() && (!UmaxCost.isValid() || OrCost <= UmaxCost);
5380 InstructionCost AltCost = UseOr ? OrCost : UmaxCost;
5381 if (AltCost > OrigCost)
5382 return false;
5383
5384 Builder.SetInsertPoint(&I);
5385 Value *NewReduce = UseOr ? Builder.CreateOrReduce(Vec)
5386 : Builder.CreateIntrinsic(
5387 Intrinsic::vector_reduce_umax, {VecTy}, {Vec});
5388 Worklist.pushValue(NewReduce);
5389 Value *NewCmp = Builder.CreateICmp(
5390 NewPred, NewReduce, ConstantInt::getNullValue(VecTy->getScalarType()));
5391 replaceValue(I, *NewCmp);
5392 return true;
5393}
5394
5395/// Returns true if this ShuffleVectorInst eventually feeds into a
5396/// vector reduction intrinsic (e.g., vector_reduce_add) by only following
5397/// chains of shuffles and binary operators (in any combination/order).
5398/// The search does not go deeper than the given Depth.
5400 constexpr unsigned MaxVisited = 32;
5403 bool FoundReduction = false;
5404
5405 WorkList.push_back(SVI);
5406 while (!WorkList.empty()) {
5407 Instruction *I = WorkList.pop_back_val();
5408 for (User *U : I->users()) {
5409 auto *UI = cast<Instruction>(U);
5410 if (!UI || !Visited.insert(UI).second)
5411 continue;
5412 if (Visited.size() > MaxVisited)
5413 return false;
5414 if (auto *II = dyn_cast<IntrinsicInst>(UI)) {
5415 // More than one reduction reached
5416 if (FoundReduction)
5417 return false;
5418 switch (II->getIntrinsicID()) {
5419 case Intrinsic::vector_reduce_add:
5420 case Intrinsic::vector_reduce_mul:
5421 case Intrinsic::vector_reduce_and:
5422 case Intrinsic::vector_reduce_or:
5423 case Intrinsic::vector_reduce_xor:
5424 case Intrinsic::vector_reduce_smin:
5425 case Intrinsic::vector_reduce_smax:
5426 case Intrinsic::vector_reduce_umin:
5427 case Intrinsic::vector_reduce_umax:
5428 FoundReduction = true;
5429 continue;
5430 default:
5431 return false;
5432 }
5433 }
5434
5436 return false;
5437
5438 WorkList.emplace_back(UI);
5439 }
5440 }
5441 return FoundReduction;
5442}
5443
5444/// This method looks for groups of shuffles acting on binops, of the form:
5445/// %x = shuffle ...
5446/// %y = shuffle ...
5447/// %a = binop %x, %y
5448/// %b = binop %x, %y
5449/// shuffle %a, %b, selectmask
5450/// We may, especially if the shuffle is wider than legal, be able to convert
5451/// the shuffle to a form where only parts of a and b need to be computed. On
5452/// architectures with no obvious "select" shuffle, this can reduce the total
5453/// number of operations if the target reports them as cheaper.
5454bool VectorCombine::foldSelectShuffle(Instruction &I, bool FromReduction) {
5455 auto *SVI = cast<ShuffleVectorInst>(&I);
5456 auto *VT = cast<FixedVectorType>(I.getType());
5457 auto *Op0 = dyn_cast<Instruction>(SVI->getOperand(0));
5458 auto *Op1 = dyn_cast<Instruction>(SVI->getOperand(1));
5459 if (!Op0 || !Op1 || Op0 == Op1 || !Op0->isBinaryOp() || !Op1->isBinaryOp() ||
5460 VT != Op0->getType())
5461 return false;
5462
5463 auto *SVI0A = dyn_cast<Instruction>(Op0->getOperand(0));
5464 auto *SVI0B = dyn_cast<Instruction>(Op0->getOperand(1));
5465 auto *SVI1A = dyn_cast<Instruction>(Op1->getOperand(0));
5466 auto *SVI1B = dyn_cast<Instruction>(Op1->getOperand(1));
5467 SmallPtrSet<Instruction *, 4> InputShuffles({SVI0A, SVI0B, SVI1A, SVI1B});
5468 auto checkSVNonOpUses = [&](Instruction *I) {
5469 if (!I || I->getOperand(0)->getType() != VT)
5470 return true;
5471 return any_of(I->users(), [&](User *U) {
5472 return U != Op0 && U != Op1 &&
5473 !(isa<ShuffleVectorInst>(U) &&
5474 (InputShuffles.contains(cast<Instruction>(U)) ||
5475 isInstructionTriviallyDead(cast<Instruction>(U))));
5476 });
5477 };
5478 if (checkSVNonOpUses(SVI0A) || checkSVNonOpUses(SVI0B) ||
5479 checkSVNonOpUses(SVI1A) || checkSVNonOpUses(SVI1B))
5480 return false;
5481
5482 // Collect all the uses that are shuffles that we can transform together. We
5483 // may not have a single shuffle, but a group that can all be transformed
5484 // together profitably.
5486 auto collectShuffles = [&](Instruction *I) {
5487 for (auto *U : I->users()) {
5488 auto *SV = dyn_cast<ShuffleVectorInst>(U);
5489 if (!SV || SV->getType() != VT)
5490 return false;
5491 if ((SV->getOperand(0) != Op0 && SV->getOperand(0) != Op1) ||
5492 (SV->getOperand(1) != Op0 && SV->getOperand(1) != Op1))
5493 return false;
5494 if (!llvm::is_contained(Shuffles, SV))
5495 Shuffles.push_back(SV);
5496 }
5497 return true;
5498 };
5499 if (!collectShuffles(Op0) || !collectShuffles(Op1))
5500 return false;
5501 // From a reduction, we need to be processing a single shuffle, otherwise the
5502 // other uses will not be lane-invariant.
5503 if (FromReduction && Shuffles.size() > 1)
5504 return false;
5505
5506 // Add any shuffle uses for the shuffles we have found, to include them in our
5507 // cost calculations.
5508 if (!FromReduction) {
5509 for (size_t Idx = 0, E = Shuffles.size(); Idx != E; ++Idx) {
5510 for (auto *U : Shuffles[Idx]->users()) {
5511 ShuffleVectorInst *SSV = dyn_cast<ShuffleVectorInst>(U);
5512 if (SSV && isa<UndefValue>(SSV->getOperand(1)) && SSV->getType() == VT)
5513 Shuffles.push_back(SSV);
5514 }
5515 }
5516 }
5517
5518 // For each of the output shuffles, we try to sort all the first vector
5519 // elements to the beginning, followed by the second array elements at the
5520 // end. If the binops are legalized to smaller vectors, this may reduce total
5521 // number of binops. We compute the ReconstructMask mask needed to convert
5522 // back to the original lane order.
5524 SmallVector<SmallVector<int>> OrigReconstructMasks;
5525 int MaxV1Elt = 0, MaxV2Elt = 0;
5526 unsigned NumElts = VT->getNumElements();
5527 for (ShuffleVectorInst *SVN : Shuffles) {
5528 SmallVector<int> Mask;
5529 SVN->getShuffleMask(Mask);
5530
5531 // Check the operands are the same as the original, or reversed (in which
5532 // case we need to commute the mask).
5533 Value *SVOp0 = SVN->getOperand(0);
5534 Value *SVOp1 = SVN->getOperand(1);
5535 if (isa<UndefValue>(SVOp1)) {
5536 auto *SSV = cast<ShuffleVectorInst>(SVOp0);
5537 SVOp0 = SSV->getOperand(0);
5538 SVOp1 = SSV->getOperand(1);
5539 for (int &Elem : Mask) {
5540 if (Elem >= static_cast<int>(SSV->getShuffleMask().size()))
5541 return false;
5542 Elem = Elem < 0 ? Elem : SSV->getMaskValue(Elem);
5543 }
5544 }
5545 if (SVOp0 == Op1 && SVOp1 == Op0) {
5546 std::swap(SVOp0, SVOp1);
5548 }
5549 if (SVOp0 != Op0 || SVOp1 != Op1)
5550 return false;
5551
5552 // Calculate the reconstruction mask for this shuffle, as the mask needed to
5553 // take the packed values from Op0/Op1 and reconstructing to the original
5554 // order.
5555 SmallVector<int> ReconstructMask;
5556 for (unsigned I = 0; I < Mask.size(); I++) {
5557 if (Mask[I] < 0) {
5558 ReconstructMask.push_back(-1);
5559 } else if (Mask[I] < static_cast<int>(NumElts)) {
5560 MaxV1Elt = std::max(MaxV1Elt, Mask[I]);
5561 auto It = find_if(V1, [&](const std::pair<int, int> &A) {
5562 return Mask[I] == A.first;
5563 });
5564 if (It != V1.end())
5565 ReconstructMask.push_back(It - V1.begin());
5566 else {
5567 ReconstructMask.push_back(V1.size());
5568 V1.emplace_back(Mask[I], V1.size());
5569 }
5570 } else {
5571 MaxV2Elt = std::max<int>(MaxV2Elt, Mask[I] - NumElts);
5572 auto It = find_if(V2, [&](const std::pair<int, int> &A) {
5573 return Mask[I] - static_cast<int>(NumElts) == A.first;
5574 });
5575 if (It != V2.end())
5576 ReconstructMask.push_back(NumElts + It - V2.begin());
5577 else {
5578 ReconstructMask.push_back(NumElts + V2.size());
5579 V2.emplace_back(Mask[I] - NumElts, NumElts + V2.size());
5580 }
5581 }
5582 }
5583
5584 // For reductions, we know that the lane ordering out doesn't alter the
5585 // result. In-order can help simplify the shuffle away.
5586 if (FromReduction)
5587 sort(ReconstructMask);
5588 OrigReconstructMasks.push_back(std::move(ReconstructMask));
5589 }
5590
5591 // If the Maximum element used from V1 and V2 are not larger than the new
5592 // vectors, the vectors are already packes and performing the optimization
5593 // again will likely not help any further. This also prevents us from getting
5594 // stuck in a cycle in case the costs do not also rule it out.
5595 if (V1.empty() || V2.empty() ||
5596 (MaxV1Elt == static_cast<int>(V1.size()) - 1 &&
5597 MaxV2Elt == static_cast<int>(V2.size()) - 1))
5598 return false;
5599
5600 // GetBaseMaskValue takes one of the inputs, which may either be a shuffle, a
5601 // shuffle of another shuffle, or not a shuffle (that is treated like a
5602 // identity shuffle).
5603 auto GetBaseMaskValue = [&](Instruction *I, int M) {
5604 auto *SV = dyn_cast<ShuffleVectorInst>(I);
5605 if (!SV)
5606 return M;
5607 if (isa<UndefValue>(SV->getOperand(1)))
5608 if (auto *SSV = dyn_cast<ShuffleVectorInst>(SV->getOperand(0)))
5609 if (InputShuffles.contains(SSV))
5610 return SSV->getMaskValue(SV->getMaskValue(M));
5611 return SV->getMaskValue(M);
5612 };
5613
5614 // Attempt to sort the inputs my ascending mask values to make simpler input
5615 // shuffles and push complex shuffles down to the uses. We sort on the first
5616 // of the two input shuffle orders, to try and get at least one input into a
5617 // nice order.
5618 auto SortBase = [&](Instruction *A, std::pair<int, int> X,
5619 std::pair<int, int> Y) {
5620 int MXA = GetBaseMaskValue(A, X.first);
5621 int MYA = GetBaseMaskValue(A, Y.first);
5622 return MXA < MYA;
5623 };
5624 stable_sort(V1, [&](std::pair<int, int> A, std::pair<int, int> B) {
5625 return SortBase(SVI0A, A, B);
5626 });
5627 stable_sort(V2, [&](std::pair<int, int> A, std::pair<int, int> B) {
5628 return SortBase(SVI1A, A, B);
5629 });
5630 // Calculate our ReconstructMasks from the OrigReconstructMasks and the
5631 // modified order of the input shuffles.
5632 SmallVector<SmallVector<int>> ReconstructMasks;
5633 for (const auto &Mask : OrigReconstructMasks) {
5634 SmallVector<int> ReconstructMask;
5635 for (int M : Mask) {
5636 auto FindIndex = [](const SmallVector<std::pair<int, int>> &V, int M) {
5637 auto It = find_if(V, [M](auto A) { return A.second == M; });
5638 assert(It != V.end() && "Expected all entries in Mask");
5639 return std::distance(V.begin(), It);
5640 };
5641 if (M < 0)
5642 ReconstructMask.push_back(-1);
5643 else if (M < static_cast<int>(NumElts)) {
5644 ReconstructMask.push_back(FindIndex(V1, M));
5645 } else {
5646 ReconstructMask.push_back(NumElts + FindIndex(V2, M));
5647 }
5648 }
5649 ReconstructMasks.push_back(std::move(ReconstructMask));
5650 }
5651
5652 // Calculate the masks needed for the new input shuffles, which get padded
5653 // with undef
5654 SmallVector<int> V1A, V1B, V2A, V2B;
5655 for (unsigned I = 0; I < V1.size(); I++) {
5656 V1A.push_back(GetBaseMaskValue(SVI0A, V1[I].first));
5657 V1B.push_back(GetBaseMaskValue(SVI0B, V1[I].first));
5658 }
5659 for (unsigned I = 0; I < V2.size(); I++) {
5660 V2A.push_back(GetBaseMaskValue(SVI1A, V2[I].first));
5661 V2B.push_back(GetBaseMaskValue(SVI1B, V2[I].first));
5662 }
5663 while (V1A.size() < NumElts) {
5666 }
5667 while (V2A.size() < NumElts) {
5670 }
5671
5672 auto AddShuffleCost = [&](InstructionCost C, Instruction *I) {
5673 auto *SV = dyn_cast<ShuffleVectorInst>(I);
5674 if (!SV)
5675 return C;
5676 return C + TTI.getShuffleCost(isa<UndefValue>(SV->getOperand(1))
5679 VT, VT, SV->getShuffleMask(), CostKind);
5680 };
5681 auto AddShuffleMaskCost = [&](InstructionCost C, ArrayRef<int> Mask) {
5682 return C +
5684 };
5685
5686 unsigned ElementSize = VT->getElementType()->getPrimitiveSizeInBits();
5687 unsigned MaxVectorSize =
5689 unsigned MaxElementsInVector = MaxVectorSize / ElementSize;
5690 if (MaxElementsInVector == 0)
5691 return false;
5692 // When there are multiple shufflevector operations on the same input,
5693 // especially when the vector length is larger than the register size,
5694 // identical shuffle patterns may occur across different groups of elements.
5695 // To avoid overestimating the cost by counting these repeated shuffles more
5696 // than once, we only account for unique shuffle patterns. This adjustment
5697 // prevents inflated costs in the cost model for wide vectors split into
5698 // several register-sized groups.
5699 std::set<SmallVector<int, 4>> UniqueShuffles;
5700 auto AddShuffleMaskAdjustedCost = [&](InstructionCost C, ArrayRef<int> Mask) {
5701 // Compute the cost for performing the shuffle over the full vector.
5702 auto ShuffleCost =
5704 unsigned NumFullVectors = Mask.size() / MaxElementsInVector;
5705 if (NumFullVectors < 2)
5706 return C + ShuffleCost;
5707 SmallVector<int, 4> SubShuffle(MaxElementsInVector);
5708 unsigned NumUniqueGroups = 0;
5709 unsigned NumGroups = Mask.size() / MaxElementsInVector;
5710 // For each group of MaxElementsInVector contiguous elements,
5711 // collect their shuffle pattern and insert into the set of unique patterns.
5712 for (unsigned I = 0; I < NumFullVectors; ++I) {
5713 for (unsigned J = 0; J < MaxElementsInVector; ++J)
5714 SubShuffle[J] = Mask[MaxElementsInVector * I + J];
5715 if (UniqueShuffles.insert(SubShuffle).second)
5716 NumUniqueGroups += 1;
5717 }
5718 return C + ShuffleCost * NumUniqueGroups / NumGroups;
5719 };
5720 auto AddShuffleAdjustedCost = [&](InstructionCost C, Instruction *I) {
5721 auto *SV = dyn_cast<ShuffleVectorInst>(I);
5722 if (!SV)
5723 return C;
5724 SmallVector<int, 16> Mask;
5725 SV->getShuffleMask(Mask);
5726 return AddShuffleMaskAdjustedCost(C, Mask);
5727 };
5728 // Check that input consists of ShuffleVectors applied to the same input
5729 auto AllShufflesHaveSameOperands =
5730 [](SmallPtrSetImpl<Instruction *> &InputShuffles) {
5731 if (InputShuffles.size() < 2)
5732 return false;
5733 ShuffleVectorInst *FirstSV =
5734 dyn_cast<ShuffleVectorInst>(*InputShuffles.begin());
5735 if (!FirstSV)
5736 return false;
5737
5738 Value *In0 = FirstSV->getOperand(0), *In1 = FirstSV->getOperand(1);
5739 return std::all_of(
5740 std::next(InputShuffles.begin()), InputShuffles.end(),
5741 [&](Instruction *I) {
5742 ShuffleVectorInst *SV = dyn_cast<ShuffleVectorInst>(I);
5743 return SV && SV->getOperand(0) == In0 && SV->getOperand(1) == In1;
5744 });
5745 };
5746
5747 // Get the costs of the shuffles + binops before and after with the new
5748 // shuffle masks.
5749 InstructionCost CostBefore =
5750 TTI.getArithmeticInstrCost(Op0->getOpcode(), VT, CostKind) +
5751 TTI.getArithmeticInstrCost(Op1->getOpcode(), VT, CostKind);
5752 CostBefore += std::accumulate(Shuffles.begin(), Shuffles.end(),
5753 InstructionCost(0), AddShuffleCost);
5754 if (AllShufflesHaveSameOperands(InputShuffles)) {
5755 UniqueShuffles.clear();
5756 CostBefore += std::accumulate(InputShuffles.begin(), InputShuffles.end(),
5757 InstructionCost(0), AddShuffleAdjustedCost);
5758 } else {
5759 CostBefore += std::accumulate(InputShuffles.begin(), InputShuffles.end(),
5760 InstructionCost(0), AddShuffleCost);
5761 }
5762
5763 // The new binops will be unused for lanes past the used shuffle lengths.
5764 // These types attempt to get the correct cost for that from the target.
5765 FixedVectorType *Op0SmallVT =
5766 FixedVectorType::get(VT->getScalarType(), V1.size());
5767 FixedVectorType *Op1SmallVT =
5768 FixedVectorType::get(VT->getScalarType(), V2.size());
5769 InstructionCost CostAfter =
5770 TTI.getArithmeticInstrCost(Op0->getOpcode(), Op0SmallVT, CostKind) +
5771 TTI.getArithmeticInstrCost(Op1->getOpcode(), Op1SmallVT, CostKind);
5772 UniqueShuffles.clear();
5773 CostAfter += std::accumulate(ReconstructMasks.begin(), ReconstructMasks.end(),
5774 InstructionCost(0), AddShuffleMaskAdjustedCost);
5775 std::set<SmallVector<int>> OutputShuffleMasks({V1A, V1B, V2A, V2B});
5776 CostAfter +=
5777 std::accumulate(OutputShuffleMasks.begin(), OutputShuffleMasks.end(),
5778 InstructionCost(0), AddShuffleMaskCost);
5779
5780 LLVM_DEBUG(dbgs() << "Found a binop select shuffle pattern: " << I << "\n");
5781 LLVM_DEBUG(dbgs() << " CostBefore: " << CostBefore
5782 << " vs CostAfter: " << CostAfter << "\n");
5783 if (CostBefore < CostAfter ||
5784 (CostBefore == CostAfter && !feedsIntoVectorReduction(SVI)))
5785 return false;
5786
5787 // The cost model has passed, create the new instructions.
5788 auto GetShuffleOperand = [&](Instruction *I, unsigned Op) -> Value * {
5789 auto *SV = dyn_cast<ShuffleVectorInst>(I);
5790 if (!SV)
5791 return I;
5792 if (isa<UndefValue>(SV->getOperand(1)))
5793 if (auto *SSV = dyn_cast<ShuffleVectorInst>(SV->getOperand(0)))
5794 if (InputShuffles.contains(SSV))
5795 return SSV->getOperand(Op);
5796 return SV->getOperand(Op);
5797 };
5798 Builder.SetInsertPoint(*SVI0A->getInsertionPointAfterDef());
5799 Value *NSV0A = Builder.CreateShuffleVector(GetShuffleOperand(SVI0A, 0),
5800 GetShuffleOperand(SVI0A, 1), V1A);
5801 Builder.SetInsertPoint(*SVI0B->getInsertionPointAfterDef());
5802 Value *NSV0B = Builder.CreateShuffleVector(GetShuffleOperand(SVI0B, 0),
5803 GetShuffleOperand(SVI0B, 1), V1B);
5804 Builder.SetInsertPoint(*SVI1A->getInsertionPointAfterDef());
5805 Value *NSV1A = Builder.CreateShuffleVector(GetShuffleOperand(SVI1A, 0),
5806 GetShuffleOperand(SVI1A, 1), V2A);
5807 Builder.SetInsertPoint(*SVI1B->getInsertionPointAfterDef());
5808 Value *NSV1B = Builder.CreateShuffleVector(GetShuffleOperand(SVI1B, 0),
5809 GetShuffleOperand(SVI1B, 1), V2B);
5810 Builder.SetInsertPoint(Op0);
5811 Value *NOp0 = Builder.CreateBinOp((Instruction::BinaryOps)Op0->getOpcode(),
5812 NSV0A, NSV0B);
5813 if (auto *I = dyn_cast<Instruction>(NOp0))
5814 I->copyIRFlags(Op0, true);
5815 Builder.SetInsertPoint(Op1);
5816 Value *NOp1 = Builder.CreateBinOp((Instruction::BinaryOps)Op1->getOpcode(),
5817 NSV1A, NSV1B);
5818 if (auto *I = dyn_cast<Instruction>(NOp1))
5819 I->copyIRFlags(Op1, true);
5820
5821 for (int S = 0, E = ReconstructMasks.size(); S != E; S++) {
5822 Builder.SetInsertPoint(Shuffles[S]);
5823 Value *NSV = Builder.CreateShuffleVector(NOp0, NOp1, ReconstructMasks[S]);
5824 replaceValue(*Shuffles[S], *NSV, false);
5825 }
5826
5827 Worklist.pushValue(NSV0A);
5828 Worklist.pushValue(NSV0B);
5829 Worklist.pushValue(NSV1A);
5830 Worklist.pushValue(NSV1B);
5831 return true;
5832}
5833
5834/// Check if instruction depends on ZExt and this ZExt can be moved after the
5835/// instruction. Move ZExt if it is profitable. For example:
5836/// logic(zext(x),y) -> zext(logic(x,trunc(y)))
5837/// lshr((zext(x),y) -> zext(lshr(x,trunc(y)))
5838/// Cost model calculations takes into account if zext(x) has other users and
5839/// whether it can be propagated through them too.
5840bool VectorCombine::shrinkType(Instruction &I) {
5841 Value *ZExted, *OtherOperand;
5842 if (!match(&I, m_c_BitwiseLogic(m_ZExt(m_Value(ZExted)),
5843 m_Value(OtherOperand))) &&
5844 !match(&I, m_LShr(m_ZExt(m_Value(ZExted)), m_Value(OtherOperand))))
5845 return false;
5846
5847 Value *ZExtOperand = I.getOperand(I.getOperand(0) == OtherOperand ? 1 : 0);
5848
5849 auto *BigTy = cast<FixedVectorType>(I.getType());
5850 auto *SmallTy = cast<FixedVectorType>(ZExted->getType());
5851 unsigned BW = SmallTy->getElementType()->getPrimitiveSizeInBits();
5852
5853 if (I.getOpcode() == Instruction::LShr) {
5854 // Check that the shift amount is less than the number of bits in the
5855 // smaller type. Otherwise, the smaller lshr will return a poison value.
5856 KnownBits ShAmtKB = computeKnownBits(I.getOperand(1), *DL);
5857 if (ShAmtKB.getMaxValue().uge(BW))
5858 return false;
5859 } else {
5860 // Check that the expression overall uses at most the same number of bits as
5861 // ZExted
5862 KnownBits KB = computeKnownBits(&I, *DL);
5863 if (KB.countMaxActiveBits() > BW)
5864 return false;
5865 }
5866
5867 // Calculate costs of leaving current IR as it is and moving ZExt operation
5868 // later, along with adding truncates if needed
5870 Instruction::ZExt, BigTy, SmallTy,
5871 TargetTransformInfo::CastContextHint::None, CostKind);
5872 InstructionCost CurrentCost = ZExtCost;
5873 InstructionCost ShrinkCost = 0;
5874
5875 // Calculate total cost and check that we can propagate through all ZExt users
5876 for (User *U : ZExtOperand->users()) {
5877 auto *UI = cast<Instruction>(U);
5878 if (UI == &I) {
5879 CurrentCost +=
5880 TTI.getArithmeticInstrCost(UI->getOpcode(), BigTy, CostKind);
5881 ShrinkCost +=
5882 TTI.getArithmeticInstrCost(UI->getOpcode(), SmallTy, CostKind);
5883 ShrinkCost += ZExtCost;
5884 continue;
5885 }
5886
5887 if (!Instruction::isBinaryOp(UI->getOpcode()))
5888 return false;
5889
5890 // Check if we can propagate ZExt through its other users
5891 KnownBits KB = computeKnownBits(UI, *DL);
5892 if (KB.countMaxActiveBits() > BW)
5893 return false;
5894
5895 CurrentCost += TTI.getArithmeticInstrCost(UI->getOpcode(), BigTy, CostKind);
5896 ShrinkCost +=
5897 TTI.getArithmeticInstrCost(UI->getOpcode(), SmallTy, CostKind);
5898 ShrinkCost += ZExtCost;
5899 }
5900
5901 // If the other instruction operand is not a constant, we'll need to
5902 // generate a truncate instruction. So we have to adjust cost
5903 if (!isa<Constant>(OtherOperand))
5904 ShrinkCost += TTI.getCastInstrCost(
5905 Instruction::Trunc, SmallTy, BigTy,
5906 TargetTransformInfo::CastContextHint::None, CostKind);
5907
5908 // If the cost of shrinking types and leaving the IR is the same, we'll lean
5909 // towards modifying the IR because shrinking opens opportunities for other
5910 // shrinking optimisations.
5911 if (ShrinkCost > CurrentCost)
5912 return false;
5913
5914 Builder.SetInsertPoint(&I);
5915 Value *Op0 = ZExted;
5916 Value *Op1 = Builder.CreateTrunc(OtherOperand, SmallTy);
5917 // Keep the order of operands the same
5918 if (I.getOperand(0) == OtherOperand)
5919 std::swap(Op0, Op1);
5920 Value *NewBinOp =
5921 Builder.CreateBinOp((Instruction::BinaryOps)I.getOpcode(), Op0, Op1);
5922 cast<Instruction>(NewBinOp)->copyIRFlags(&I);
5923 cast<Instruction>(NewBinOp)->copyMetadata(I);
5924 Value *NewZExtr = Builder.CreateZExt(NewBinOp, BigTy);
5925 replaceValue(I, *NewZExtr);
5926 return true;
5927}
5928
5929/// insert (DstVec, (extract SrcVec, ExtIdx), InsIdx) -->
5930/// shuffle (DstVec, SrcVec, Mask)
5931bool VectorCombine::foldInsExtVectorToShuffle(Instruction &I) {
5932 Value *DstVec, *SrcVec;
5933 uint64_t ExtIdx, InsIdx;
5934 if (!match(&I,
5935 m_InsertElt(m_Value(DstVec),
5936 m_ExtractElt(m_Value(SrcVec), m_ConstantInt(ExtIdx)),
5937 m_ConstantInt(InsIdx))))
5938 return false;
5939
5940 auto *DstVecTy = dyn_cast<FixedVectorType>(I.getType());
5941 auto *SrcVecTy = dyn_cast<FixedVectorType>(SrcVec->getType());
5942 // We can try combining vectors with different element sizes.
5943 if (!DstVecTy || !SrcVecTy ||
5944 SrcVecTy->getElementType() != DstVecTy->getElementType())
5945 return false;
5946
5947 unsigned NumDstElts = DstVecTy->getNumElements();
5948 unsigned NumSrcElts = SrcVecTy->getNumElements();
5949 if (InsIdx >= NumDstElts || ExtIdx >= NumSrcElts || NumDstElts == 1)
5950 return false;
5951
5952 // Insertion into poison is a cheaper single operand shuffle.
5954 SmallVector<int> Mask(NumDstElts, PoisonMaskElem);
5955
5956 bool NeedExpOrNarrow = NumSrcElts != NumDstElts;
5957 bool NeedDstSrcSwap = isa<PoisonValue>(DstVec) && !isa<UndefValue>(SrcVec);
5958 if (NeedDstSrcSwap) {
5960 Mask[InsIdx] = ExtIdx % NumDstElts;
5961 std::swap(DstVec, SrcVec);
5962 } else {
5964 std::iota(Mask.begin(), Mask.end(), 0);
5965 Mask[InsIdx] = (ExtIdx % NumDstElts) + NumDstElts;
5966 }
5967
5968 // Cost
5969 auto *Ins = cast<InsertElementInst>(&I);
5970 auto *Ext = cast<ExtractElementInst>(I.getOperand(1));
5971 InstructionCost InsCost =
5972 TTI.getVectorInstrCost(*Ins, DstVecTy, CostKind, InsIdx);
5973 InstructionCost ExtCost =
5974 TTI.getVectorInstrCost(*Ext, DstVecTy, CostKind, ExtIdx);
5975 InstructionCost OldCost = ExtCost + InsCost;
5976
5977 InstructionCost NewCost = 0;
5978 SmallVector<int> ExtToVecMask;
5979 if (!NeedExpOrNarrow) {
5980 // Ignore 'free' identity insertion shuffle.
5981 // TODO: getShuffleCost should return TCC_Free for Identity shuffles.
5982 if (!ShuffleVectorInst::isIdentityMask(Mask, NumSrcElts))
5983 NewCost += TTI.getShuffleCost(SK, DstVecTy, DstVecTy, Mask, CostKind, 0,
5984 nullptr, {DstVec, SrcVec});
5985 } else {
5986 // When creating a length-changing-vector, always try to keep the relevant
5987 // element in an equivalent position, so that bulk shuffles are more likely
5988 // to be useful.
5989 ExtToVecMask.assign(NumDstElts, PoisonMaskElem);
5990 ExtToVecMask[ExtIdx % NumDstElts] = ExtIdx;
5991 // Add cost for expanding or narrowing
5993 DstVecTy, SrcVecTy, ExtToVecMask, CostKind);
5994 NewCost += TTI.getShuffleCost(SK, DstVecTy, DstVecTy, Mask, CostKind);
5995 }
5996
5997 if (!Ext->hasOneUse())
5998 NewCost += ExtCost;
5999
6000 LLVM_DEBUG(dbgs() << "Found a insert/extract shuffle-like pair: " << I
6001 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
6002 << "\n");
6003
6004 if (OldCost < NewCost)
6005 return false;
6006
6007 if (NeedExpOrNarrow) {
6008 if (!NeedDstSrcSwap)
6009 SrcVec = Builder.CreateShuffleVector(SrcVec, ExtToVecMask);
6010 else
6011 DstVec = Builder.CreateShuffleVector(DstVec, ExtToVecMask);
6012 }
6013
6014 // Canonicalize undef param to RHS to help further folds.
6015 if (isa<UndefValue>(DstVec) && !isa<UndefValue>(SrcVec)) {
6016 ShuffleVectorInst::commuteShuffleMask(Mask, NumDstElts);
6017 std::swap(DstVec, SrcVec);
6018 }
6019
6020 Value *Shuf = Builder.CreateShuffleVector(DstVec, SrcVec, Mask);
6021 replaceValue(I, *Shuf);
6022
6023 return true;
6024}
6025
6026/// Fold away a matched pair of vector.deinterleave/interleave intrinsics
6027/// with a chain of elementwise operations on each between the
6028/// deinterleave and interleave.
6029///
6030/// For example:
6031/// ```
6032/// %d = call { <2 x i16>, <2 x i16> } @deinterleave2.v4i16(<4 x i16> %v)
6033/// %f0 = extractvalue { <2 x i16>, <2 x i16> } %d, 0
6034/// %f1 = extractvalue { <2 x i16>, <2 x i16> } %d, 1
6035///
6036/// %u0 = add <2 x i16> %f0, splat (i16 3)
6037/// %u1 = add <2 x i16> %f1, splat (i16 3)
6038///
6039/// %r = call <4 x i16> @interleave2.v4i16(<2 x i16> %u0, <2 x i16> %u1)
6040/// ```
6041/// Folds to:
6042/// ```
6043/// %r = add <4 x i16> %v, splat (i16 3)
6044/// ```
6045bool VectorCombine::foldDeinterleaveInterleavePair(Instruction &I) {
6047 if (!Deinterleave)
6048 return false;
6049
6050 unsigned Factor =
6052 if (!Factor || Deinterleave->hasOperandBundles() ||
6053 !Deinterleave->hasNUndroppableUses(Factor))
6054 return false;
6055
6056 const Intrinsic::ID ExpectedInterleaveIID =
6058
6059 // Collect one extract for each deinterleaved field.
6060 SmallVector<Use *, 8> CurrentUses(Factor, nullptr);
6061 for (Use &U : Deinterleave->uses()) {
6062 if (U.getUser()->isDroppable())
6063 continue;
6064
6065 auto *Extract = dyn_cast<ExtractValueInst>(U.getUser());
6066 if (!Extract || Extract->getNumIndices() != 1)
6067 return false;
6068
6069 unsigned Index = *Extract->idx_begin();
6070 if (Index >= Factor || CurrentUses[Index])
6071 return false;
6072
6073 CurrentUses[Index] = &U;
6074 }
6075
6076 using ElementwiseStep = SmallVector<Use *, 8>;
6078 IntrinsicInst *Interleave = nullptr;
6079 unsigned NumVisited = 0;
6080
6081 auto GetNumDataOperands = [](Instruction *Inst) {
6082 if (auto *CB = dyn_cast<CallBase>(Inst))
6083 return CB->arg_size(); // Exclude callee operand and bundles.
6084 return Inst->getNumOperands();
6085 };
6086
6087 auto IsSupportedElementwise = [&](Instruction *Inst) {
6088 auto *ResultTy = dyn_cast<VectorType>(Inst->getType());
6089 if (!ResultTy || !isSafeToSpeculativelyExecute(Inst))
6090 return false;
6091
6092 if (auto *II = dyn_cast<IntrinsicInst>(Inst)) {
6093 if (II->hasOperandBundles() ||
6094 !isTriviallyVectorizable(II->getIntrinsicID()))
6095 return false;
6096 } else if (!isa<BinaryOperator, UnaryOperator, CastInst, CmpInst,
6097 SelectInst, FreezeInst>(Inst)) {
6098 return false;
6099 }
6100
6101 // Reject operations that change the element-count.
6102 // E.g., bitcast <vscale x 4 x i16> %v to <vscale x 8 x i8>
6103 for (unsigned Op = 0, E = GetNumDataOperands(Inst); Op != E; ++Op) {
6104 auto *OperandTy = dyn_cast<VectorType>(Inst->getOperand(Op)->getType());
6105 if (OperandTy &&
6106 OperandTy->getElementCount() != ResultTy->getElementCount())
6107 return false;
6108 }
6109
6110 return true;
6111 };
6112
6113 // Traverse the Factor use chains with a breadth-first search.
6114 // At each level, expect every chain to perform the same operation with the
6115 // preceding chain value at the same operand position, until they all reach
6116 // the matching interleave.
6117 while (NumVisited + Factor <= MaxInstrsToScan) {
6118 NumVisited += Factor;
6119
6120 for (Use *&CurrentUse : CurrentUses) {
6121 Use *NextUse = CurrentUse->getUser()->getSingleUndroppableUse();
6122 auto *Next =
6123 NextUse ? dyn_cast<Instruction>(NextUse->getUser()) : nullptr;
6124 if (!Next)
6125 return false;
6126
6127 CurrentUse = NextUse;
6128 }
6129
6130 // Check whether every chain has reached the same interleave.
6131 if (auto *II = dyn_cast<IntrinsicInst>(CurrentUses.front()->getUser());
6132 II && II->getIntrinsicID() == ExpectedInterleaveIID) {
6133 if (II->hasOperandBundles())
6134 return false;
6135
6136 for (unsigned Index = 0; Index != Factor; ++Index)
6137 if (CurrentUses[Index]->getUser() != II ||
6138 CurrentUses[Index]->getOperandNo() != Index)
6139 return false;
6140
6141 Interleave = II;
6142 break;
6143 }
6144
6145 auto *FirstInst = cast<Instruction>(CurrentUses.front()->getUser());
6146 if (!IsSupportedElementwise(FirstInst))
6147 return false;
6148
6149 unsigned ChainOperand = CurrentUses.front()->getOperandNo();
6150 if (any_of(CurrentUses, [&](Use *U) {
6151 auto *Inst = cast<Instruction>(U->getUser());
6152 return Inst != FirstInst && (U->getOperandNo() != ChainOperand ||
6153 !FirstInst->isSameOperationAs(Inst));
6154 }))
6155 return false;
6156
6157 auto GetSplatOrScalar = [](Value *V) {
6158 return isa<VectorType>(V->getType()) ? getSplatValue(V) : V;
6159 };
6160
6161 // Non-chain operands must be either the same scalar or splats of that
6162 // scalar. This intentionally rejects differing poison/undef or non-splat
6163 // vector operands between chains.
6164 for (unsigned Op = 0, E = GetNumDataOperands(FirstInst); Op != E; ++Op) {
6165 if (Op == ChainOperand)
6166 continue;
6167
6168 Value *CommonValue = GetSplatOrScalar(FirstInst->getOperand(Op));
6169 if (!CommonValue || any_of(CurrentUses, [&](Use *U) {
6170 Instruction *Inst = cast<Instruction>(U->getUser());
6171 return Inst != FirstInst &&
6172 GetSplatOrScalar(Inst->getOperand(Op)) != CommonValue;
6173 }))
6174 return false;
6175 }
6176
6177 Steps.push_back(CurrentUses);
6178 }
6179
6180 if (!Interleave)
6181 return false;
6182
6183 // Rebuild the matched elementwise chain at the original vector width.
6184 Value *WideValue = Deinterleave->getArgOperand(0);
6185 ElementCount WideEC =
6186 cast<VectorType>(WideValue->getType())->getElementCount();
6187
6188 auto CreateWideInstruction = [&](Instruction *NarrowInst,
6189 ArrayRef<Value *> NewOperands,
6190 VectorType *WideResultTy) -> Value * {
6191 assert(IsSupportedElementwise(NarrowInst) &&
6192 "Expected supported elementwise");
6193 if (isa<BinaryOperator, UnaryOperator>(NarrowInst))
6194 return Builder.CreateNAryOp(NarrowInst->getOpcode(), NewOperands);
6195 if (auto *Cast = dyn_cast<CastInst>(NarrowInst))
6196 return Builder.CreateCast(Cast->getOpcode(), NewOperands[0],
6197 WideResultTy);
6198 if (auto *Cmp = dyn_cast<CmpInst>(NarrowInst))
6199 return Builder.CreateCmp(Cmp->getPredicate(), NewOperands[0],
6200 NewOperands[1]);
6201 if (isa<SelectInst>(NarrowInst))
6202 return Builder.CreateSelect(
6203 NewOperands[0], NewOperands[1], NewOperands[2], /*Name=*/"",
6204 ProfcheckDisableMetadataFixes ? nullptr : NarrowInst);
6205 if (isa<FreezeInst>(NarrowInst))
6206 return Builder.CreateFreeze(NewOperands[0]);
6207 if (auto *II = dyn_cast<IntrinsicInst>(NarrowInst))
6208 return Builder.CreateIntrinsic(WideResultTy, II->getIntrinsicID(),
6209 NewOperands);
6210 llvm_unreachable("Unsupported instruction");
6211 };
6212
6213 // The BFS has succeeded and collected multiple levels of instructions that
6214 // can be SLP-widened into a chain of wider instructions.
6215 for (const ElementwiseStep &Step : Steps) {
6216 Instruction *NarrowInst = cast<Instruction>(Step.front()->getUser());
6217 unsigned ChainOperand = Step.front()->getOperandNo();
6218
6219 Builder.SetInsertPoint(NarrowInst);
6220 Builder.SetCurrentDebugLocation(NarrowInst->getDebugLoc());
6221
6222 unsigned NumOperands = GetNumDataOperands(NarrowInst);
6223 SmallVector<Value *, 4> NewOperands;
6224 NewOperands.reserve(NumOperands);
6225
6226 for (unsigned Op = 0; Op != NumOperands; ++Op) {
6227 Value *Operand = NarrowInst->getOperand(Op);
6228
6229 if (Op == ChainOperand)
6230 Operand = WideValue;
6231 else if (isa<VectorType>(Operand->getType()))
6232 Operand = Builder.CreateVectorSplat(WideEC, getSplatValue(Operand));
6233 NewOperands.push_back(Operand);
6234 }
6235
6236 auto *WideResultTy =
6237 VectorType::get(NarrowInst->getType()->getScalarType(), WideEC);
6238 Value *NewValue =
6239 CreateWideInstruction(NarrowInst, NewOperands, WideResultTy);
6240
6241 SmallVector<Value *> NarrowInsts =
6242 map_to_vector(Step, [](Use *U) { return cast<Value>(U->getUser()); });
6243 propagateIRFlags(NewValue, NarrowInsts);
6244
6245 if (auto *NewInst = dyn_cast<Instruction>(NewValue))
6246 propagateMetadata(NewInst, NarrowInsts);
6247
6248 WideValue = NewValue;
6249 }
6250
6251 assert(WideValue->getType() == Interleave->getType());
6252 replaceValue(*Interleave, *WideValue);
6253 return true;
6254}
6255
6256/// If we're interleaving 2 constant splats, for instance `<vscale x 8 x i32>
6257/// <splat of 666>` and `<vscale x 8 x i32> <splat of 777>`, we can create a
6258/// larger splat `<vscale x 8 x i64> <splat of ((777 << 32) | 666)>` first
6259/// before casting it back into `<vscale x 16 x i32>`.
6260bool VectorCombine::foldInterleaveIntrinsics(Instruction &I) {
6261 const APInt *SplatVal0, *SplatVal1;
6263 m_APInt(SplatVal0), m_APInt(SplatVal1))))
6264 return false;
6265
6266 LLVM_DEBUG(dbgs() << "VC: Folding interleave2 with two splats: " << I
6267 << "\n");
6268
6269 auto *VTy =
6270 cast<VectorType>(cast<IntrinsicInst>(I).getArgOperand(0)->getType());
6271 auto *ExtVTy = VectorType::getExtendedElementVectorType(VTy);
6272 unsigned Width = VTy->getElementType()->getIntegerBitWidth();
6273
6274 // Just in case the cost of interleave2 intrinsic and bitcast are both
6275 // invalid, in which case we want to bail out, we use <= rather
6276 // than < here. Even they both have valid and equal costs, it's probably
6277 // not a good idea to emit a high-cost constant splat.
6279 TTI.getCastInstrCost(Instruction::BitCast, I.getType(), ExtVTy,
6281 LLVM_DEBUG(dbgs() << "VC: The cost to cast from " << *ExtVTy << " to "
6282 << *I.getType() << " is too high.\n");
6283 return false;
6284 }
6285
6286 APInt NewSplatVal = SplatVal1->zext(Width * 2);
6287 NewSplatVal <<= Width;
6288 NewSplatVal |= SplatVal0->zext(Width * 2);
6289 auto *NewSplat = ConstantVector::getSplat(
6290 ExtVTy->getElementCount(), ConstantInt::get(F.getContext(), NewSplatVal));
6291
6292 IRBuilder<> Builder(&I);
6293 replaceValue(I, *Builder.CreateBitCast(NewSplat, I.getType()));
6294 return true;
6295}
6296
6297/// Given this sequence:
6298/// ```
6299/// %d = llvm.vector.deinterleave2 <vscale x 16 x i32> %v
6300/// %f0 = extractvalue { <vscale x 8 x i32>, <vscale x 8 x i32> } %d, 0
6301/// %f1 = extractvalue { <vscale x 8 x i32>, <vscale x 8 x i32> } %d, 1
6302///
6303/// %low0 = and <vscale x 8 x i32> %f0, splat (i32 65535)
6304/// %low1 = shl <vscale x 8 x i32> %f1, splat (i32 16)
6305/// %merge0 = or disjoint <vscale x 8 x i32> %low0, %low1
6306///
6307/// %high0 = and <vscale x 8 x i32> %f1, splat (i32 -65536)
6308/// %high1 = lshr <vscale x 8 x i32> %f0, splat (i32 16)
6309/// %merge1 = or disjoint <vscale x 8 x i32> %high0, %high1
6310/// ```
6311/// It is actually just de-interleaving a 16-bit vector with double the
6312/// vector length. More generally speaking, it's de-interleaving on a vector
6313/// with half the element width as the original vector.
6314///
6315/// Therefore, we can turn it into:
6316/// ```
6317/// %narrow.v = bitcast <vscale x 16 x i32> %v to <vscale x 32 x i16>
6318/// %d = llvm.vector.deinterleave2 <vscale x 32 x i16> %narrow.v
6319/// %f0 = extractvalue { <vscale x 16 x i16>, <vscale x 16 x i16> } %d, 0
6320/// %f1 = extractvalue { <vscale x 16 x i16>, <vscale x 16 x i16> } %d, 1
6321///
6322/// %merge0 = bitcast <vscale x 16 x i16> %f0 to <vscale x 8 x i32>
6323/// %merge1 = bitcast <vscale x 16 x i16> %f1 to <vscale x 8 x i32>
6324/// ```
6325bool VectorCombine::foldDeinterleaveIntrinsics(Instruction &I) {
6326 if (foldDeinterleaveInterleavePair(I))
6327 return true;
6328
6329 // This pattern involves bitcast that is not compatible with big endian.
6330 if (DL->isBigEndian())
6331 return false;
6332
6333 using namespace PatternMatch;
6334 Value *DeinterleavedVal;
6335 if (!match(&I, m_Deinterleave2(m_Value(DeinterleavedVal))))
6336 return false;
6337
6338 VectorType *VecTy = cast<VectorType>(DeinterleavedVal->getType());
6339 IntegerType *ElementTy = dyn_cast<IntegerType>(VecTy->getElementType());
6340 if (!ElementTy)
6341 return false;
6342 unsigned ElementWidth = ElementTy->getBitWidth();
6343 if (ElementWidth < 2 || !isPowerOf2_32(ElementWidth))
6344 return false;
6345 unsigned HalfElementWidth = ElementWidth / 2;
6346
6347 if (!I.hasNUses(2))
6348 return false;
6349 std::array<ExtractValueInst *, 2> OrigFields{};
6350 for (User *Usr : I.users()) {
6351 auto *E = dyn_cast<ExtractValueInst>(Usr);
6352 // The deinterleave result can only be used by extractions.
6353 if (!E || E->getNumIndices() != 1)
6354 return false;
6355 unsigned Idx = *E->idx_begin();
6356 // A single field cannot be extracted more than once.
6357 if (Idx >= 2 || OrigFields[Idx] || !E->hasNUses(2))
6358 return false;
6359 OrigFields[Idx] = E;
6360 }
6361
6362 // Find the merge instruction (i.e. OR) first.
6363 SmallVector<Instruction *, 2> MergeInsts;
6364 for (auto *FieldUsr : OrigFields[0]->users()) {
6365 if (!FieldUsr->hasOneUse() || !isa<Instruction>(FieldUsr->user_back()))
6366 return false;
6367 MergeInsts.push_back(cast<Instruction>(FieldUsr->user_back()));
6368 }
6369 assert(MergeInsts.size() == 2);
6370
6371 // Pattern match bottom-up from the merge instructions.
6372 auto MatchMerge = [&](void) -> bool {
6373 APInt LoMask = APInt::getLowBitsSet(ElementWidth, HalfElementWidth);
6374 APInt HiMask = APInt::getHighBitsSet(ElementWidth, HalfElementWidth);
6375 return match(MergeInsts[0],
6376 m_c_Or(m_And(m_Specific(OrigFields[0]), m_SpecificInt(LoMask)),
6377 m_Shl(m_Specific(OrigFields[1]),
6378 m_SpecificInt(HalfElementWidth)))) &&
6379 match(MergeInsts[1],
6380 m_c_Or(m_And(m_Specific(OrigFields[1]), m_SpecificInt(HiMask)),
6381 m_LShr(m_Specific(OrigFields[0]),
6382 m_SpecificInt(HalfElementWidth))));
6383 };
6384 if (!MatchMerge()) {
6385 std::swap(MergeInsts[0], MergeInsts[1]);
6386 if (!MatchMerge())
6387 return false;
6388 }
6389
6390 // Profitability check.
6391 InstructionCost OldCost =
6392 TTI.getInstructionCost(MergeInsts[0], CostKind) +
6393 TTI.getInstructionCost(cast<Instruction>(MergeInsts[0]->getOperand(0)),
6394 CostKind) +
6395 TTI.getInstructionCost(cast<Instruction>(MergeInsts[0]->getOperand(1)),
6396 CostKind);
6397 // There are two fields (assuming SHL has the same cost as LSHR).
6398 OldCost *= 2;
6399
6400 auto *NewFieldTy = VecTy->getWithNewBitWidth(HalfElementWidth);
6401 auto *NewVecTy =
6402 VectorType::getDoubleElementsVectorType(cast<VectorType>(NewFieldTy));
6403 InstructionCost NewCost =
6404 TTI.getCastInstrCost(Instruction::BitCast, VecTy, NewVecTy,
6406 TTI.getCastInstrCost(Instruction::BitCast, NewFieldTy,
6407 MergeInsts[0]->getType(), TTI::CastContextHint::None,
6408 CostKind) *
6409 2;
6410 if (OldCost <= NewCost || !NewCost.isValid()) {
6411 LLVM_DEBUG(
6412 dbgs() << "VC: New deinterleave2 sequence cost (" << NewCost << ")"
6413 << " is higher than that of the old one (" << OldCost << ")\n");
6414 return false;
6415 }
6416
6417 // Do the replacement.
6418 IRBuilder<> Builder(&I);
6419 Value *NewVecCast = Builder.CreateBitCast(DeinterleavedVal, NewVecTy);
6420 Value *NewDeinterleave = Builder.CreateIntrinsic(
6421 Intrinsic::vector_deinterleave2, {NewVecTy}, {NewVecCast});
6422 for (auto [Idx, MergeInst] : enumerate(MergeInsts)) {
6423 Value *NewField = Builder.CreateExtractValue(NewDeinterleave, Idx);
6424 NewField = Builder.CreateBitCast(NewField, MergeInst->getType());
6425 replaceValue(*MergeInst, *NewField);
6426 }
6427
6428 return true;
6429}
6430
6431bool VectorCombine::foldBitcastOfVPLoad(Instruction &I) {
6432 const DataLayout &DL = I.getDataLayout();
6433 auto *Cast = dyn_cast<CastInst>(&I);
6434 if (!Cast || !Cast->isNoopCast(DL) || !isa<VectorType>(Cast->getDestTy()))
6435 return false;
6436
6437 // Fold away bit casts of the loaded value by loading the desired type,
6438 // if the mask is all-ones.
6439 Value *EVL;
6440 auto *II = dyn_cast<VPIntrinsic>(I.getOperand(0));
6442 m_Value(), m_AllOnes(), m_Value(EVL)))))
6443 return false;
6444
6445 VectorType *OrigVecTy = cast<VectorType>(II->getType());
6446 Align OrigAlign =
6447 DL.getValueOrABITypeAlignment(II->getPointerAlignment(), OrigVecTy);
6448 ElementCount OrigVecCnt = OrigVecTy->getElementCount();
6449 VectorType *NewVecTy = cast<VectorType>(Cast->getDestTy());
6450 ElementCount NewVecCnt = NewVecTy->getElementCount();
6451
6452 // Right now we only support cases where the NewVec is longer, because for
6453 // cases where it's shorter, we have to be sure that EVL can be exactly
6454 // divided, otherwise it might yield incorrect results or even page faults
6455 // (if we round-up during the division).
6456 if (!(OrigVecCnt.isScalable() == NewVecCnt.isScalable() &&
6457 NewVecCnt.hasKnownScalarFactor(OrigVecCnt)))
6458 return false;
6459
6460 InstructionCost OldCost =
6461 TTI.getMemIntrinsicInstrCost({Intrinsic::vp_load, OrigVecTy,
6462 II->getMemoryPointerParam(), false,
6463 OrigAlign},
6464 CostKind) +
6465 TTI.getCastInstrCost(Instruction::BitCast, Cast->getType(), OrigVecTy,
6468 {Intrinsic::vp_load, NewVecTy, II->getMemoryPointerParam(), false,
6469 OrigAlign},
6470 CostKind);
6471 LLVM_DEBUG(dbgs() << "foldBitcastOfVPLoad: OldCost=" << OldCost
6472 << " NewCost=" << NewCost << "\n");
6473 if (NewCost > OldCost || !NewCost.isValid())
6474 return false;
6475
6476 unsigned Factor = NewVecCnt.getKnownScalarFactor(OrigVecCnt);
6477 Value *NewEVL = Builder.CreateNUWMul(EVL, Builder.getInt32(Factor));
6478 Value *NewMask = Builder.CreateVectorSplat(NewVecCnt, Builder.getTrue());
6479 CallInst *NewVP = Builder.CreateIntrinsicWithoutFolding(
6480 NewVecTy, Intrinsic::vp_load,
6481 {II->getMemoryPointerParam(), NewMask, NewEVL});
6482 // Preserve the original alignment.
6483 NewVP->addParamAttrs(
6484 0, AttrBuilder(II->getContext()).addAlignmentAttr(OrigAlign));
6485 replaceValue(*Cast, *NewVP);
6486 return true;
6487}
6488/// Fold the following cases into a single byte-level bit-reverse operation
6489/// and accepts bswap and bitreverse intrinsics:
6490/// bswap(bitreverse(x)) --> bitcast(bitreverse(bitcast(x)))
6491/// bitreverse(bswap(x)) <--> bitcast(bitreverse(bitcast(x)))
6492/// The direction of the fold is cost-model driven.
6493/// Also supports:
6494/// bitcast(bitreverse(bitcast(x))) --> bitreverse(fshl(x))
6495bool VectorCombine::foldBitOrderReverseAndSwap(Instruction &I) {
6496 Value *X;
6497
6499 Type *Ty = X->getType();
6500 Type *VecTy = I.getOperand(0)->getType();
6501 // Detect the case when bitreversing every octet in X individually. Then we
6502 // can use bswap to reorder the octets before doing a single bitreverse.
6503 bool CanUseBswap =
6504 Ty->isIntegerTy() && Ty == I.getType() && isa<FixedVectorType>(VecTy) &&
6505 cast<FixedVectorType>(VecTy)->getElementType()->isIntegerTy(8) &&
6506 Ty->getIntegerBitWidth() % 16 == 0;
6507 // Detect the case when bitreversing upper and lower half of X
6508 // individually. Then we can use fshl as a rotate operation, to swap the
6509 // halves before doing a single bitreverse.
6510 bool CanUseFshl =
6511 Ty->isIntegerTy() && Ty == I.getType() && isa<FixedVectorType>(VecTy) &&
6512 cast<FixedVectorType>(VecTy)->getElementType()->isIntegerTy() &&
6513 cast<FixedVectorType>(VecTy)->getNumElements() == 2;
6514 if (CanUseBswap || CanUseFshl) {
6515 auto *InnerCall = dyn_cast<Instruction>(I.getOperand(0));
6516 if (!InnerCall)
6517 return false;
6518 auto *InnerBitCast = dyn_cast<BitCastInst>(InnerCall->getOperand(0));
6519 if (!InnerBitCast)
6520 return false;
6521 Constant *HalfBW = ConstantInt::get(Ty, Ty->getIntegerBitWidth() / 2);
6522 InstructionCost OldCost = TTI.getInstructionCost(InnerBitCast, CostKind) +
6523 TTI.getInstructionCost(InnerCall, CostKind) +
6525 IntrinsicCostAttributes ICABSwap(Intrinsic::bswap, Ty, {Ty});
6526 IntrinsicCostAttributes ICABFshl(Intrinsic::fshl, Ty, {X, X, HalfBW},
6527 {Ty, Ty, Ty});
6528 IntrinsicCostAttributes ICABRev(Intrinsic::bitreverse, Ty, {Ty});
6529 InstructionCost NewCost =
6530 TTI.getIntrinsicInstrCost(CanUseBswap ? ICABSwap : ICABFshl,
6531 CostKind) +
6533 if (!InnerCall->hasOneUse())
6534 NewCost += TTI.getInstructionCost(InnerCall, CostKind) +
6535 TTI.getInstructionCost(InnerBitCast, CostKind);
6536 else if (!InnerBitCast->hasOneUse())
6537 NewCost += TTI.getInstructionCost(InnerBitCast, CostKind);
6538 LLVM_DEBUG(dbgs() << "Found bitreverse vector roundtrip: " << I
6539 << "\n OldCost: " << OldCost
6540 << " vs NewCost: " << NewCost << "\n");
6541 if (NewCost.isValid() && NewCost < OldCost) {
6542 Builder.SetInsertPoint(&I);
6543 Value *Swap =
6544 CanUseBswap
6545 ? Builder.CreateUnaryIntrinsic(Intrinsic::bswap, X)
6546 : Builder.CreateIntrinsic(Ty, Intrinsic::fshl, {X, X, HalfBW});
6547 Worklist.pushValue(Swap);
6548 Value *BRev = Builder.CreateUnaryIntrinsic(Intrinsic::bitreverse, Swap);
6549 replaceValue(I, *BRev);
6550 return true;
6551 }
6552 }
6553 }
6554
6555 if (!match(&I, m_BitReverse(m_BSwap(m_Value(X)))) &&
6557 return false;
6558 Type *Ty = I.getType();
6559 Type *I8Ty = Builder.getInt8Ty();
6560 TypeSize ElementSize = DL->getTypeStoreSize(Ty);
6561 ElementCount NewVecCnt = ElementCount::get(ElementSize.getKnownMinValue(),
6562 ElementSize.isScalable());
6563 Type *NewVecTy = VectorType::get(I8Ty, NewVecCnt);
6564 auto *II = cast<IntrinsicInst>(&I);
6565 auto *InnerII = cast<IntrinsicInst>(II->getArgOperand(0));
6566 // OldCost = cost of bitreverse/bswap + cost of bswap/bitreverse
6569 // NewCost = cost of bitcast to byte vector +
6570 // cost of bitreverse/bswap on byte vector +
6571 // cost of bitcast back to original type
6572 InstructionCost CastToVecCost = TTI.getCastInstrCost(
6573 Instruction::BitCast, NewVecTy, Ty, TTI::CastContextHint::None, CostKind);
6574 InstructionCost CastToOrigCost = TTI.getCastInstrCost(
6575 Instruction::BitCast, Ty, NewVecTy, TTI::CastContextHint::None, CostKind);
6576 IntrinsicCostAttributes ICANew(Intrinsic::bitreverse, NewVecTy, {NewVecTy});
6577 InstructionCost NewIntrinsicCost =
6579 InstructionCost NewCost = CastToVecCost + NewIntrinsicCost + CastToOrigCost;
6580 if (!InnerII->hasOneUse())
6581 NewCost += TTI.getInstructionCost(InnerII, CostKind);
6582 LLVM_DEBUG(dbgs() << "Found bitorder reverse and swap: " << I
6583 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
6584 << "\n");
6585 if (!NewCost.isValid() || NewCost >= OldCost)
6586 return false;
6587 // Perform transform: bitcast(arg, <N x i8>), bitreverse, bitcast back
6588 Builder.SetInsertPoint(II);
6589 Value *CastToVec = Builder.CreateBitCast(X, NewVecTy);
6590 Value *NewCall =
6591 Builder.CreateUnaryIntrinsic(Intrinsic::bitreverse, CastToVec);
6592 Value *CastToOrig = Builder.CreateBitCast(NewCall, Ty);
6593 replaceValue(I, *CastToOrig);
6594 return true;
6595}
6596
6597/// Given the maximum shuffle index and load vector type, compute the number of
6598/// elements for the shrunk load, rounding up to the next full vector register
6599/// boundary to avoid scalar remainders that legalize poorly.
6600static unsigned getAlignedNumElements(unsigned MaxIdx, FixedVectorType *LoadTy,
6601 const TargetTransformInfo &TTI,
6602 const DataLayout &DL) {
6603 unsigned RawNumElements = MaxIdx + 1u;
6604 Type *ElemTy = LoadTy->getElementType();
6605 // Skip alignment for illegal element types.
6606 if (!TTI.isTypeLegal(ElemTy))
6607 return RawNumElements;
6608
6609 TypeSize ElemSize = DL.getTypeSizeInBits(ElemTy);
6610 if (ElemSize.isScalable() || ElemSize.isZero())
6611 return RawNumElements;
6612
6615 if (RegSize.isScalable() || RegSize.isZero())
6616 return RawNumElements;
6617
6618 unsigned ElemsPerReg = RegSize.getFixedValue() / ElemSize.getFixedValue();
6619 // If the load already fits in a register, keep the exact size.
6620 // Otherwise round up to the next full register boundary.
6621 if (ElemsPerReg == 0 || RawNumElements <= ElemsPerReg)
6622 return RawNumElements;
6623
6624 return alignTo(RawNumElements, ElemsPerReg);
6625}
6626
6627// Attempt to shrink loads that are only used by shufflevector instructions.
6628bool VectorCombine::shrinkLoadForShuffles(Instruction &I) {
6629 auto *OldLoad = dyn_cast<LoadInst>(&I);
6630 if (!OldLoad || !OldLoad->isSimple())
6631 return false;
6632
6633 auto *OldLoadTy = dyn_cast<FixedVectorType>(OldLoad->getType());
6634 if (!OldLoadTy)
6635 return false;
6636
6637 unsigned const OldNumElements = OldLoadTy->getNumElements();
6638
6639 // Search all uses of load. If all uses are shufflevector instructions, and
6640 // the second operands are all poison values, find the minimum and maximum
6641 // indices of the vector elements referenced by all shuffle masks.
6642 // Otherwise return `std::nullopt`.
6643 using IndexRange = std::pair<int, int>;
6644 auto GetIndexRangeInShuffles = [&]() -> std::optional<IndexRange> {
6645 IndexRange OutputRange = IndexRange(OldNumElements, -1);
6646 for (llvm::Use &Use : I.uses()) {
6647 // Ensure all uses match the required pattern.
6648 User *Shuffle = Use.getUser();
6649 ArrayRef<int> Mask;
6650
6651 if (!match(Shuffle,
6652 m_Shuffle(m_Specific(OldLoad), m_Undef(), m_Mask(Mask))))
6653 return std::nullopt;
6654
6655 // Ignore shufflevector instructions that have no uses.
6656 if (Shuffle->use_empty())
6657 continue;
6658
6659 // Find the min and max indices used by the shufflevector instruction.
6660 for (int Index : Mask) {
6661 if (Index >= 0 && Index < static_cast<int>(OldNumElements)) {
6662 OutputRange.first = std::min(Index, OutputRange.first);
6663 OutputRange.second = std::max(Index, OutputRange.second);
6664 }
6665 }
6666 }
6667
6668 if (OutputRange.second < OutputRange.first)
6669 return std::nullopt;
6670
6671 return OutputRange;
6672 };
6673
6674 // Get the range of vector elements used by shufflevector instructions.
6675 if (std::optional<IndexRange> Indices = GetIndexRangeInShuffles()) {
6676 unsigned const NewNumElements =
6677 getAlignedNumElements(Indices->second, OldLoadTy, TTI, *DL);
6678
6679 // If the range of vector elements is smaller than the full load, attempt
6680 // to create a smaller load.
6681 if (NewNumElements < OldNumElements) {
6682 IRBuilder Builder(&I);
6683 Builder.SetCurrentDebugLocation(I.getDebugLoc());
6684
6685 // Calculate costs of old and new ops.
6686 Type *ElemTy = OldLoadTy->getElementType();
6687 FixedVectorType *NewLoadTy = FixedVectorType::get(ElemTy, NewNumElements);
6688 Value *PtrOp = OldLoad->getPointerOperand();
6689
6691 Instruction::Load, OldLoad->getType(), OldLoad->getAlign(),
6692 OldLoad->getPointerAddressSpace(), CostKind);
6693 InstructionCost NewCost =
6694 TTI.getMemoryOpCost(Instruction::Load, NewLoadTy, OldLoad->getAlign(),
6695 OldLoad->getPointerAddressSpace(), CostKind);
6696
6697 using UseEntry = std::pair<ShuffleVectorInst *, std::vector<int>>;
6699 unsigned const MaxIndex = NewNumElements * 2u;
6700
6701 for (llvm::Use &Use : I.uses()) {
6702 auto *Shuffle = cast<ShuffleVectorInst>(Use.getUser());
6703
6704 // Ignore shufflevector instructions that have no uses.
6705 if (Shuffle->use_empty())
6706 continue;
6707
6708 ArrayRef<int> OldMask = Shuffle->getShuffleMask();
6709
6710 // Create entry for new use.
6711 NewUses.push_back({Shuffle, OldMask});
6712
6713 // Validate mask indices.
6714 for (int Index : OldMask) {
6715 if (Index >= static_cast<int>(MaxIndex))
6716 return false;
6717 }
6718
6719 // Update costs.
6720 OldCost +=
6722 OldLoadTy, OldMask, CostKind);
6723 NewCost +=
6725 NewLoadTy, OldMask, CostKind);
6726 }
6727
6728 LLVM_DEBUG(
6729 dbgs() << "Found a load used only by shufflevector instructions: "
6730 << I << "\n OldCost: " << OldCost
6731 << " vs NewCost: " << NewCost << "\n");
6732
6733 if (OldCost < NewCost || !NewCost.isValid())
6734 return false;
6735
6736 // Create new load of smaller vector.
6737 auto *NewLoad = cast<LoadInst>(
6738 Builder.CreateAlignedLoad(NewLoadTy, PtrOp, OldLoad->getAlign()));
6739 NewLoad->copyMetadata(I);
6740
6741 // Replace all uses.
6742 for (UseEntry &Use : NewUses) {
6743 ShuffleVectorInst *Shuffle = Use.first;
6744 std::vector<int> &NewMask = Use.second;
6745
6746 Builder.SetInsertPoint(Shuffle);
6747 Builder.SetCurrentDebugLocation(Shuffle->getDebugLoc());
6748 Value *NewShuffle = Builder.CreateShuffleVector(
6749 NewLoad, PoisonValue::get(NewLoadTy), NewMask);
6750
6751 replaceValue(*Shuffle, *NewShuffle, false);
6752 }
6753
6754 return true;
6755 }
6756 }
6757 return false;
6758}
6759
6760// Attempt to narrow a phi of shufflevector instructions where the two incoming
6761// values have the same operands but different masks. If the two shuffle masks
6762// are offsets of one another we can use one branch to rotate the incoming
6763// vector and perform one larger shuffle after the phi.
6764bool VectorCombine::shrinkPhiOfShuffles(Instruction &I) {
6765 auto *Phi = dyn_cast<PHINode>(&I);
6766 if (!Phi || Phi->getNumIncomingValues() != 2u)
6767 return false;
6768
6769 Value *Op = nullptr;
6770 ArrayRef<int> Mask0;
6771 ArrayRef<int> Mask1;
6772
6773 if (!match(Phi->getOperand(0u),
6774 m_OneUse(m_Shuffle(m_Value(Op), m_Poison(), m_Mask(Mask0)))) ||
6775 !match(Phi->getOperand(1u),
6776 m_OneUse(m_Shuffle(m_Specific(Op), m_Poison(), m_Mask(Mask1)))))
6777 return false;
6778
6779 auto *Shuf = cast<ShuffleVectorInst>(Phi->getOperand(0u));
6780
6781 // Ensure result vectors are wider than the argument vector.
6782 auto *InputVT = cast<FixedVectorType>(Op->getType());
6783 auto *ResultVT = cast<FixedVectorType>(Shuf->getType());
6784 auto const InputNumElements = InputVT->getNumElements();
6785
6786 if (InputNumElements >= ResultVT->getNumElements())
6787 return false;
6788
6789 // Take the difference of the two shuffle masks at each index. Ignore poison
6790 // values at the same index in both masks.
6791 SmallVector<int, 16> NewMask;
6792 NewMask.reserve(Mask0.size());
6793
6794 for (auto [M0, M1] : zip(Mask0, Mask1)) {
6795 if (M0 >= 0 && M1 >= 0)
6796 NewMask.push_back(M0 - M1);
6797 else if (M0 == -1 && M1 == -1)
6798 continue;
6799 else
6800 return false;
6801 }
6802
6803 // Ensure all elements of the new mask are equal. If the difference between
6804 // the incoming mask elements is the same, the two must be constant offsets
6805 // of one another.
6806 if (NewMask.empty() || !all_equal(NewMask))
6807 return false;
6808
6809 // Create new mask using difference of the two incoming masks.
6810 int MaskOffset = NewMask[0u];
6811 unsigned Index = (InputNumElements + MaskOffset) % InputNumElements;
6812 NewMask.clear();
6813
6814 for (unsigned I = 0u; I < InputNumElements; ++I) {
6815 NewMask.push_back(Index);
6816 Index = (Index + 1u) % InputNumElements;
6817 }
6818
6819 // Calculate costs for worst cases and compare.
6820 auto const Kind = TTI::SK_PermuteSingleSrc;
6821 auto OldCost =
6822 std::max(TTI.getShuffleCost(Kind, ResultVT, InputVT, Mask0, CostKind),
6823 TTI.getShuffleCost(Kind, ResultVT, InputVT, Mask1, CostKind));
6824 auto NewCost = TTI.getShuffleCost(Kind, InputVT, InputVT, NewMask, CostKind) +
6825 TTI.getShuffleCost(Kind, ResultVT, InputVT, Mask1, CostKind);
6826
6827 LLVM_DEBUG(dbgs() << "Found a phi of mergeable shuffles: " << I
6828 << "\n OldCost: " << OldCost << " vs NewCost: " << NewCost
6829 << "\n");
6830
6831 if (NewCost > OldCost)
6832 return false;
6833
6834 // Create new shuffles and narrowed phi.
6835 auto Builder = IRBuilder(Shuf);
6836 Builder.SetCurrentDebugLocation(Shuf->getDebugLoc());
6837 auto *PoisonVal = PoisonValue::get(InputVT);
6838 auto *NewShuf0 = Builder.CreateShuffleVector(Op, PoisonVal, NewMask);
6839 Worklist.push(cast<Instruction>(NewShuf0));
6840
6841 Builder.SetInsertPoint(Phi);
6842 Builder.SetCurrentDebugLocation(Phi->getDebugLoc());
6843 auto *NewPhi = Builder.CreatePHI(NewShuf0->getType(), 2u);
6844 NewPhi->addIncoming(NewShuf0, Phi->getIncomingBlock(0u));
6845 NewPhi->addIncoming(Op, Phi->getIncomingBlock(1u));
6846
6847 Builder.SetInsertPoint(*NewPhi->getInsertionPointAfterDef());
6848 PoisonVal = PoisonValue::get(NewPhi->getType());
6849 auto *NewShuf1 = Builder.CreateShuffleVector(NewPhi, PoisonVal, Mask1);
6850
6851 replaceValue(*Phi, *NewShuf1);
6852 return true;
6853}
6854
6855/// This is the entry point for all transforms. Pass manager differences are
6856/// handled in the callers of this function.
6857bool VectorCombine::run() {
6859 return false;
6860
6861 // Don't attempt vectorization if the target does not support vectors.
6862 if (!TTI.getNumberOfRegisters(TTI.getRegisterClassForType(/*Vector*/ true)))
6863 return false;
6864
6865 LLVM_DEBUG(dbgs() << "\n\nVECTORCOMBINE on " << F.getName() << "\n");
6866
6867 auto FoldInst = [this](Instruction &I) {
6868 Builder.SetInsertPoint(&I);
6869 bool IsVectorType = isa<VectorType>(I.getType());
6870 bool IsFixedVectorType = isa<FixedVectorType>(I.getType());
6871 auto Opcode = I.getOpcode();
6872
6873 LLVM_DEBUG(dbgs() << "VC: Visiting: " << I << '\n');
6874
6875 // These folds should be beneficial regardless of when this pass is run
6876 // in the optimization pipeline.
6877 // The type checking is for run-time efficiency. We can avoid wasting time
6878 // dispatching to folding functions if there's no chance of matching.
6879 if (IsFixedVectorType) {
6880 switch (Opcode) {
6881 case Instruction::InsertElement:
6882 if (vectorizeLoadInsert(I))
6883 return true;
6884 break;
6885 case Instruction::ShuffleVector:
6886 if (widenSubvectorLoad(I))
6887 return true;
6888 break;
6889 default:
6890 break;
6891 }
6892 }
6893
6894 // This transform works with scalable and fixed vectors
6895 // TODO: Identify and allow other scalable transforms
6896 if (IsVectorType) {
6897 if (scalarizeOpOrCmp(I))
6898 return true;
6899 if (scalarizeLoad(I))
6900 return true;
6901 if (scalarizeExtExtract(I))
6902 return true;
6903 if (foldInterleaveIntrinsics(I))
6904 return true;
6905 if (foldBitcastOfVPLoad(I))
6906 return true;
6907 }
6908
6909 if (foldDeinterleaveIntrinsics(I))
6910 return true;
6911
6912 if (Opcode == Instruction::Store)
6913 if (foldInsertElementsToStores(I))
6914 return true;
6915
6916 // If this is an early pipeline invocation of this pass, we are done.
6917 if (TryEarlyFoldsOnly)
6918 return false;
6919
6920 if (Opcode == Instruction::Call)
6921 if (foldBitOrderReverseAndSwap(I))
6922 return true;
6923 if (Opcode == Instruction::BitCast)
6924 if (foldBitOrderReverseAndSwap(I))
6925 return true;
6926
6927 // Otherwise, try folds that improve codegen but may interfere with
6928 // early IR canonicalizations.
6929 // The type checking is for run-time efficiency. We can avoid wasting time
6930 // dispatching to folding functions if there's no chance of matching.
6931 if (IsFixedVectorType) {
6932 switch (Opcode) {
6933 case Instruction::InsertElement:
6934 if (foldInsExtFNeg(I))
6935 return true;
6936 if (foldInsExtBinop(I))
6937 return true;
6938 if (foldInsExtVectorToShuffle(I))
6939 return true;
6940 break;
6941 case Instruction::ShuffleVector:
6942 if (foldPermuteOfBinops(I))
6943 return true;
6944 if (foldShuffleOfBinops(I))
6945 return true;
6946 if (foldShuffleOfSelects(I))
6947 return true;
6948 if (foldShuffleOfCastops(I))
6949 return true;
6950 if (foldShuffleOfShuffles(I))
6951 return true;
6952 if (foldPermuteOfIntrinsic(I))
6953 return true;
6954 if (foldShufflesOfLengthChangingShuffles(I))
6955 return true;
6956 if (foldShuffleOfIntrinsics(I))
6957 return true;
6958 if (foldSelectShuffle(I))
6959 return true;
6960 if (foldShuffleToIdentity(I))
6961 return true;
6962 break;
6963 case Instruction::Load:
6964 if (shrinkLoadForShuffles(I))
6965 return true;
6966 break;
6967 case Instruction::BitCast:
6968 if (foldBitcastShuffle(I))
6969 return true;
6970 if (foldSelectsFromBitcast(I))
6971 return true;
6972 break;
6973 case Instruction::And:
6974 case Instruction::Or:
6975 case Instruction::Xor:
6976 if (foldBitOpOfCastops(I))
6977 return true;
6978 if (foldBitOpOfCastConstant(I))
6979 return true;
6980 break;
6981 case Instruction::PHI:
6982 if (shrinkPhiOfShuffles(I))
6983 return true;
6984 break;
6985 default:
6986 if (shrinkType(I))
6987 return true;
6988 break;
6989 }
6990 } else {
6991 switch (Opcode) {
6992 case Instruction::Call:
6993 if (foldShuffleFromReductions(I))
6994 return true;
6995 if (foldCastFromReductions(I))
6996 return true;
6997 break;
6998 case Instruction::ExtractElement:
6999 if (foldShuffleChainsToReduce(I))
7000 return true;
7001 break;
7002 case Instruction::ICmp:
7003 if (foldSignBitReductionCmp(I))
7004 return true;
7005 if (foldICmpEqZeroVectorReduce(I))
7006 return true;
7007 if (foldReductionZeroTest(I))
7008 return true;
7009 if (foldEquivalentReductionCmp(I))
7010 return true;
7011 if (foldReduceAddCmpZero(I))
7012 return true;
7013 [[fallthrough]];
7014 case Instruction::FCmp:
7015 if (foldExtractExtract(I))
7016 return true;
7017 break;
7018 case Instruction::Or:
7019 if (foldConcatOfBoolMasks(I))
7020 return true;
7021 [[fallthrough]];
7022 default:
7023 if (Instruction::isBinaryOp(Opcode)) {
7024 if (foldExtractExtract(I))
7025 return true;
7026 if (foldExtractedCmps(I))
7027 return true;
7028 if (foldBinopOfReductions(I))
7029 return true;
7030 }
7031 break;
7032 }
7033 }
7034 return false;
7035 };
7036
7037 bool MadeChange = false;
7038 for (BasicBlock &BB : F) {
7039 // Ignore unreachable basic blocks.
7040 if (!DT.isReachableFromEntry(&BB))
7041 continue;
7042 // Use early increment range so that we can erase instructions in loop.
7043 // make_early_inc_range is not applicable here, as the next iterator may
7044 // be invalidated by RecursivelyDeleteTriviallyDeadInstructions.
7045 // We manually maintain the next instruction and update it when it is about
7046 // to be deleted.
7047 Instruction *I = &BB.front();
7048 while (I) {
7049 NextInst = I->getNextNode();
7050 if (!I->isDebugOrPseudoInst())
7051 MadeChange |= FoldInst(*I);
7052 I = NextInst;
7053 }
7054 }
7055
7056 NextInst = nullptr;
7057
7058 while (!Worklist.isEmpty()) {
7059 Instruction *I = Worklist.removeOne();
7060 if (!I)
7061 continue;
7062
7065 continue;
7066 }
7067
7068 MadeChange |= FoldInst(*I);
7069 }
7070
7071 return MadeChange;
7072}
7073
7076 auto &AC = FAM.getResult<AssumptionAnalysis>(F);
7078 DominatorTree &DT = FAM.getResult<DominatorTreeAnalysis>(F);
7079 AAResults &AA = FAM.getResult<AAManager>(F);
7080 const DataLayout *DL = &F.getDataLayout();
7083 VectorCombine Combiner(F, TTI, DT, AA, AC, DL, CostKind, TryEarlyFoldsOnly);
7084 if (!Combiner.run())
7085 return PreservedAnalyses::all();
7088 return PA;
7089}
unsigned RegSize
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
unsigned uint64_t
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
static cl::opt< unsigned > MaxInstrsToScan("aggressive-instcombine-max-scan-instrs", cl::init(64), cl::Hidden, cl::desc("Max number of instructions to scan for aggressive instcombine."))
This is the interface for LLVM's primary stateless and local alias analysis.
#define X(NUM, ENUM, NAME)
Definition ELF.h:856
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< ErlangGC > A("erlang", "erlang-compatible garbage collector")
static GCRegistry::Add< StatepointGC > D("statepoint-example", "an example strategy for statepoint")
static GCRegistry::Add< CoreCLRGC > E("coreclr", "CoreCLR-compatible GC")
static GCRegistry::Add< OcamlGC > B("ocaml", "ocaml 3.10-compatible GC")
static cl::opt< OutputCostKind > CostKind("cost-kind", cl::desc("Target cost kind"), cl::init(OutputCostKind::RecipThroughput), cl::values(clEnumValN(OutputCostKind::RecipThroughput, "throughput", "Reciprocal throughput"), clEnumValN(OutputCostKind::Latency, "latency", "Instruction latency"), clEnumValN(OutputCostKind::CodeSize, "code-size", "Code size"), clEnumValN(OutputCostKind::SizeAndLatency, "size-latency", "Code size and latency"), clEnumValN(OutputCostKind::All, "all", "Print all cost kinds")))
static cl::opt< IntrinsicCostStrategy > IntrinsicCost("intrinsic-cost-strategy", cl::desc("Costing strategy for intrinsic instructions"), cl::init(IntrinsicCostStrategy::InstructionCost), cl::values(clEnumValN(IntrinsicCostStrategy::InstructionCost, "instruction-cost", "Use TargetTransformInfo::getInstructionCost"), clEnumValN(IntrinsicCostStrategy::IntrinsicCost, "intrinsic-cost", "Use TargetTransformInfo::getIntrinsicInstrCost"), clEnumValN(IntrinsicCostStrategy::TypeBasedIntrinsicCost, "type-based-intrinsic-cost", "Calculate the intrinsic cost based only on argument types")))
This file defines the DenseMap class.
#define Check(C,...)
This is the interface for a simple mod/ref and alias analysis over globals.
Hexagon Common GEP
iv users
Definition IVUsers.cpp:48
const size_t AbstractManglingParser< Derived, Alloc >::NumOps
const AbstractManglingParser< Derived, Alloc >::OperatorInfo AbstractManglingParser< Derived, Alloc >::Ops[]
static void eraseInstruction(Instruction &I, ICFLoopSafetyInfo &SafetyInfo, MemorySSAUpdater &MSSAU)
Definition LICM.cpp:1544
#define F(x, y, z)
Definition MD5.cpp:54
#define I(x, y, z)
Definition MD5.cpp:57
#define T1
uint64_t IntrinsicInst * II
FunctionAnalysisManager FAM
if(PassOpts->AAPipeline)
This file contains the declarations for profiling metadata utility functions.
const SmallVectorImpl< MachineOperand > & Cond
Func getContext().diagnose(DiagnosticInfoUnsupported(Func
This file contains some templates that are useful if you are working with the STL at all.
This file defines the scope_exit class, which executes user-defined cleanup logic at scope exit.
This file defines less commonly used SmallVector utilities.
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)
Definition Statistic.h:171
#define LLVM_DEBUG(...)
Definition Debug.h:119
static TableGen::Emitter::Opt Y("gen-skeleton-entry", EmitSkeleton, "Generate example skeleton entry")
static SymbolRef::Type getType(const Symbol *Sym)
Definition TapiFile.cpp:39
This pass exposes codegen information to IR-level passes.
static bool isEquivBitcast(Value *X, Value *Y)
Helper to peek through bitcasts to the same value.
static bool isFreeConcat(ArrayRef< InstLane > Item, TTI::TargetCostKind CostKind, const TargetTransformInfo &TTI)
Detect concat of multiple values into a vector.
static void analyzeCostOfVecReduction(const IntrinsicInst &II, TTI::TargetCostKind CostKind, const TargetTransformInfo &TTI, InstructionCost &CostBeforeReduction, InstructionCost &CostAfterReduction)
static Value * generateNewInstTree(ArrayRef< InstLane > Item, Use *From, const DenseSet< std::pair< Value *, Use * > > &IdentityLeafs, const DenseSet< std::pair< Value *, Use * > > &SplatLeafs, const DenseSet< std::pair< Value *, Use * > > &ConcatLeafs, IRBuilderBase &Builder, InstructionWorklist &WorkList, const TargetTransformInfo *TTI)
static SmallVector< InstLane > generateInstLaneVectorFromOperand(ArrayRef< InstLane > Item, int Op)
static Value * createShiftShuffle(Value *Vec, unsigned OldIndex, unsigned NewIndex, IRBuilderBase &Builder)
Create a shuffle that translates (shifts) 1 element from the input vector to a new element location.
std::pair< Value *, int > InstLane
static bool isKnownNonPositive(const Value *V, const SimplifyQuery &SQ, unsigned Depth=0)
Used by foldReduceAddCmpZero to check if we can prove that a value is non-positive.
static Value * materializeScalarizedGEPIndex(Value *Idx, IntegerType *GEPIndexTy, IRBuilderBase &Builder)
Materialize an index for a scalarized GEP after profitability is known.
static Align computeAlignmentAfterScalarization(Align VectorAlignment, Type *ScalarType, Value *Idx, const DataLayout &DL)
The memory operation on a vector of ScalarType had alignment of VectorAlignment.
static bool feedsIntoVectorReduction(ShuffleVectorInst *SVI)
Returns true if this ShuffleVectorInst eventually feeds into a vector reduction intrinsic (e....
static cl::opt< bool > DisableVectorCombine("disable-vector-combine", cl::init(false), cl::Hidden, cl::desc("Disable all vector combine transforms"))
static bool canWidenLoad(LoadInst *Load, const TargetTransformInfo &TTI)
static const unsigned InvalidIndex
static IntegerType * getScalarizedGEPIndexInfo(VectorType *VecTy, Value *Idx, Type *PtrTy, const DataLayout &DL)
Return the GEP index type if the unsigned vector index Idx can be represented by an inbounds GEP.
static Value * translateExtract(ExtractElementInst *ExtElt, unsigned NewIndex, IRBuilderBase &Builder)
Given an extract element instruction with constant index operand, shuffle the source vector (shift th...
static ScalarizationResult canScalarizeAccess(VectorType *VecTy, Value *Idx, const SimplifyQuery &SQ)
Check if it is legal to scalarize a memory access to VecTy at index Idx.
static cl::opt< unsigned > MaxInstrsToScan("vector-combine-max-scan-instrs", cl::init(30), cl::Hidden, cl::desc("Max number of instructions to scan for vector combining."))
static cl::opt< bool > DisableBinopExtractShuffle("disable-binop-extract-shuffle", cl::init(false), cl::Hidden, cl::desc("Disable binop extract to shuffle transforms"))
static unsigned getAlignedNumElements(unsigned MaxIdx, FixedVectorType *LoadTy, const TargetTransformInfo &TTI, const DataLayout &DL)
Given the maximum shuffle index and load vector type, compute the number of elements for the shrunk l...
static InstLane lookThroughShuffles(Value *V, int Lane)
static bool isMemModifiedBetween(BasicBlock::iterator Begin, BasicBlock::iterator End, const MemoryLocation &Loc, AAResults &AA)
static constexpr int Concat[]
Value * RHS
Value * LHS
A manager for alias analyses.
Class for arbitrary precision integers.
Definition APInt.h:78
LLVM_ABI APInt zext(unsigned width) const
Zero extend to a new width.
Definition APInt.cpp:1056
uint64_t getZExtValue() const
Get zero extended value.
Definition APInt.h:1561
bool isAllOnes() const
Determine if all bits are set. This is true for zero-width values.
Definition APInt.h:368
bool ugt(const APInt &RHS) const
Unsigned greater than comparison.
Definition APInt.h:1187
bool isZero() const
Determine if this value is zero, i.e. all bits are clear.
Definition APInt.h:377
unsigned getBitWidth() const
Return the number of bits in the APInt.
Definition APInt.h:1509
static APInt getSignedMaxValue(unsigned numBits)
Gets maximum signed value of APInt for a specific bit width.
Definition APInt.h:206
bool isNegative() const
Determine sign of this APInt.
Definition APInt.h:326
unsigned countl_one() const
Count the number of leading one bits.
Definition APInt.h:1636
LLVM_ABI APInt sext(unsigned width) const
Sign extend to a new width.
Definition APInt.cpp:1029
static APInt getLowBitsSet(unsigned numBits, unsigned loBitsSet)
Constructs an APInt value that has the bottom loBitsSet bits set.
Definition APInt.h:303
static APInt getHighBitsSet(unsigned numBits, unsigned hiBitsSet)
Constructs an APInt value that has the top hiBitsSet bits set.
Definition APInt.h:293
static APInt getZero(unsigned numBits)
Get the '0' value for the specified bit-width.
Definition APInt.h:197
bool isOne() const
Determine if this is a value of 1.
Definition APInt.h:386
static APInt getOneBitSet(unsigned numBits, unsigned BitNo)
Return an APInt with exactly one bit set in the result.
Definition APInt.h:236
bool uge(const APInt &RHS) const
Unsigned greater or equal comparison.
Definition APInt.h:1226
Represent a constant reference to an array (0 or more elements consecutively in memory),...
Definition ArrayRef.h:40
const T & front() const
Get the first element.
Definition ArrayRef.h:144
size_t size() const
Get the array size.
Definition ArrayRef.h:141
A function analysis which provides an AssumptionCache.
A cache of @llvm.assume calls within a function.
InstListType::iterator iterator
Instruction iterators...
Definition BasicBlock.h:170
BinaryOps getOpcode() const
Definition InstrTypes.h:409
Represents analyses that only rely on functions' control flow.
Definition Analysis.h:73
Value * getArgOperand(unsigned i) const
void addParamAttrs(unsigned ArgNo, const AttrBuilder &B)
Adds attributes to the indicated argument.
static LLVM_ABI CastInst * Create(Instruction::CastOps, Value *S, Type *Ty, const Twine &Name="", InsertPosition InsertBefore=nullptr)
Provides a way to construct any of the CastInst subclasses using an opcode instead of the subclass's ...
static Type * makeCmpResultType(Type *opnd_type)
Create a result type for fcmp/icmp.
Predicate
This enumeration lists the possible predicates for CmpInst subclasses.
Definition InstrTypes.h:740
bool isFPPredicate() const
Definition InstrTypes.h:845
static LLVM_ABI std::optional< CmpPredicate > getMatching(CmpPredicate A, CmpPredicate B)
Compares two CmpPredicates taking samesign into account and returns the canonicalized CmpPredicate if...
Combiner implementation.
Definition Combiner.h:33
static LLVM_ABI Constant * getExtractElement(Constant *Vec, Constant *Idx, Type *OnlyIfReducedTy=nullptr)
static LLVM_ABI Constant * getBinOpIdentity(unsigned Opcode, Type *Ty, bool AllowRHSConstant=false, bool NSZ=false)
Return the identity constant for a binary opcode.
This is the shared class of boolean and integer constants.
Definition Constants.h:87
const APInt & getValue() const
Return the constant as an APInt value reference.
Definition Constants.h:159
This class represents a range of values.
LLVM_ABI ConstantRange urem(const ConstantRange &Other) const
Return a new range representing the possible values resulting from an unsigned remainder operation of...
LLVM_ABI ConstantRange binaryAnd(const ConstantRange &Other) const
Return a new range representing the possible values resulting from a binary-and of a value in this ra...
LLVM_ABI bool contains(const APInt &Val) const
Return true if the specified value is in the set.
static LLVM_ABI Constant * getSplat(ElementCount EC, Constant *Elt)
Return a ConstantVector with the specified constant in each element.
static LLVM_ABI Constant * get(ArrayRef< Constant * > V)
static LLVM_ABI Constant * getNullValue(Type *Ty)
Constructor to create a '0' constant of arbitrary type.
A parsed version of the target data layout string in and methods for querying it.
Definition DataLayout.h:64
ValueT lookup(const_arg_type_t< KeyT > Val) const
Return the entry for the specified key, or a default constructed value if no such entry exists.
Definition DenseMap.h:250
iterator find(const_arg_type_t< KeyT > Val)
Definition DenseMap.h:223
std::pair< iterator, bool > try_emplace(KeyT &&Key, Ts &&...Args)
Definition DenseMap.h:299
bool empty() const
Definition DenseMap.h:171
iterator end()
Definition DenseMap.h:141
Implements a dense probed hash-table based set.
Definition DenseSet.h:281
Analysis pass which computes a DominatorTree.
Definition Dominators.h:241
Concrete subclass of DominatorTreeBase that is used to compute a normal dominator tree.
Definition Dominators.h:122
LLVM_ABI bool isReachableFromEntry(const Use &U) const
Provide an overload for a Use.
LLVM_ABI bool dominates(const BasicBlock *BB, const Use &U) const
Return true if the (end of the) basic block BB dominates the use U.
static constexpr ElementCount get(ScalarTy MinVal, bool Scalable)
Definition TypeSize.h:315
This instruction extracts a single (scalar) element from a VectorType value.
Convenience struct for specifying and reasoning about fast-math flags.
Definition FMF.h:23
bool noSignedZeros() const
Definition FMF.h:67
Class to represent fixed width SIMD vectors.
unsigned getNumElements() const
static FixedVectorType * getDoubleElementsVectorType(FixedVectorType *VTy)
static LLVM_ABI FixedVectorType * get(Type *ElementType, unsigned NumElts)
Definition Type.cpp:867
Predicate getSignedPredicate() const
For example, EQ->EQ, SLE->SLE, UGT->SGT, etc.
bool isEquality() const
Return true if this predicate is either EQ or NE.
Common base class shared among various IRBuilders.
Definition IRBuilder.h:114
LLVM_ABI CallInst * CreateIntrinsicWithoutFolding(Intrinsic::ID ID, ArrayRef< Type * > OverloadTypes, ArrayRef< Value * > Args, FMFSource FMFSource={}, const Twine &Name="", ArrayRef< OperandBundleDef > OpBundles={})
Create a call to intrinsic ID with Args, mangled using OverloadTypes.
Value * CreateNUWMul(Value *LHS, Value *RHS, const Twine &Name="")
Definition IRBuilder.h:1469
Value * CreateInsertElement(Type *VecTy, Value *NewElt, Value *Idx, const Twine &Name="")
Definition IRBuilder.h:2662
Value * CreateExtractElement(Value *Vec, Value *Idx, const Twine &Name="")
Definition IRBuilder.h:2650
LoadInst * CreateAlignedLoad(Type *Ty, Value *Ptr, MaybeAlign Align, const char *Name)
Definition IRBuilder.h:1934
LLVM_ABI Value * CreateSelectFMF(Value *C, Value *True, Value *False, FMFSource FMFSource, const Twine &Name="", Instruction *MDFrom=nullptr)
LLVM_ABI Value * CreateVectorSplat(unsigned NumElts, Value *V, const Twine &Name="")
Return a vector value that contains.
Value * CreateExtractValue(Value *Agg, ArrayRef< unsigned > Idxs, const Twine &Name="")
Definition IRBuilder.h:2709
ConstantInt * getTrue()
Get the constant value for i1 true.
Definition IRBuilder.h:457
LLVM_ABI Value * CreateSelect(Value *C, Value *True, Value *False, const Twine &Name="", Instruction *MDFrom=nullptr)
Value * CreateFreeze(Value *V, const Twine &Name="")
Definition IRBuilder.h:2728
void SetCurrentDebugLocation(const DebugLoc &L)
Set location information used by debugging information.
Definition IRBuilder.h:221
Value * CreateLShr(Value *LHS, Value *RHS, const Twine &Name="", bool isExact=false)
Definition IRBuilder.h:1532
Value * CreateCast(Instruction::CastOps Op, Value *V, Type *DestTy, const Twine &Name="", MDNode *FPMathTag=nullptr, FMFSource FMFSource={})
Definition IRBuilder.h:2277
Value * CreateIsNotNeg(Value *Arg, const Twine &Name="")
Return a boolean value testing if Arg > -1.
Definition IRBuilder.h:2752
Value * CreateInBoundsGEP(Type *Ty, Value *Ptr, ArrayRef< Value * > IdxList, const Twine &Name="")
Definition IRBuilder.h:2019
Value * CreatePointerBitCastOrAddrSpaceCast(Value *V, Type *DestTy, const Twine &Name="")
Definition IRBuilder.h:2302
ConstantInt * getInt64(uint64_t C)
Get a constant 64-bit value.
Definition IRBuilder.h:482
LLVM_ABI Value * CreateOrReduce(Value *Src)
Create a vector int OR reduction intrinsic of the source vector.
ConstantInt * getInt32(uint32_t C)
Get a constant 32-bit value.
Definition IRBuilder.h:477
Value * CreateCmp(CmpInst::Predicate Pred, Value *LHS, Value *RHS, const Twine &Name="", MDNode *FPMathTag=nullptr)
Definition IRBuilder.h:2509
PHINode * CreatePHI(Type *Ty, unsigned NumReservedValues, const Twine &Name="")
Definition IRBuilder.h:2540
InstTy * Insert(InstTy *I, const Twine &Name="") const
Insert and return the specified instruction.
Definition IRBuilder.h:146
Value * CreateIsNeg(Value *Arg, const Twine &Name="")
Return a boolean value testing if Arg < 0.
Definition IRBuilder.h:2747
Value * CreateBitCast(Value *V, Type *DestTy, const Twine &Name="")
Definition IRBuilder.h:2243
LoadInst * CreateLoad(Type *Ty, Value *Ptr, const char *Name)
Provided to resolve 'CreateLoad(Ty, Ptr, "...")' correctly, instead of converting the string to 'bool...
Definition IRBuilder.h:1906
Value * CreateShl(Value *LHS, Value *RHS, const Twine &Name="", bool HasNUW=false, bool HasNSW=false)
Definition IRBuilder.h:1511
LLVM_ABI Value * CreateNAryOp(unsigned Opc, ArrayRef< Value * > Ops, const Twine &Name="", MDNode *FPMathTag=nullptr)
Create either a UnaryOperator or BinaryOperator depending on Opc.
Value * CreateZExt(Value *V, Type *DestTy, const Twine &Name="", bool IsNonNeg=false)
Definition IRBuilder.h:2121
Value * CreateShuffleVector(Value *V1, Value *V2, Value *Mask, const Twine &Name="")
Definition IRBuilder.h:2684
Value * CreateAnd(Value *LHS, Value *RHS, const Twine &Name="")
Definition IRBuilder.h:1570
LLVM_ABI Value * CreateIntrinsic(Intrinsic::ID ID, ArrayRef< Type * > OverloadTypes, ArrayRef< Value * > Args, FMFSource FMFSource={}, const Twine &Name="", ArrayRef< OperandBundleDef > OpBundles={}, function_ref< void(CallInst *)> SetFn=[](CallInst *) {})
Variant to create a possibly constant-folded intrinsic.
StoreInst * CreateStore(Value *Val, Value *Ptr, bool isVolatile=false)
Definition IRBuilder.h:1925
Value * CreateTrunc(Value *V, Type *DestTy, const Twine &Name="", bool IsNUW=false, bool IsNSW=false)
Definition IRBuilder.h:2107
PointerType * getPtrTy(unsigned AddrSpace=0)
Fetch the type representing a pointer.
Definition IRBuilder.h:577
Value * CreateBinOp(Instruction::BinaryOps Opc, Value *LHS, Value *RHS, const Twine &Name="", MDNode *FPMathTag=nullptr)
Definition IRBuilder.h:1731
void SetInsertPoint(BasicBlock *TheBB)
This specifies that created instructions should be appended to the end of the specified block.
Definition IRBuilder.h:181
Value * CreateFNegFMF(Value *V, FMFSource FMFSource, const Twine &Name="", MDNode *FPMathTag=nullptr)
Definition IRBuilder.h:1844
Value * CreateICmp(CmpInst::Predicate P, Value *LHS, Value *RHS, const Twine &Name="")
Definition IRBuilder.h:2485
Value * CreateOr(Value *LHS, Value *RHS, const Twine &Name="", bool IsDisjoint=false)
Definition IRBuilder.h:1592
IntegerType * getInt8Ty()
Fetch the type representing an 8-bit integer.
Definition IRBuilder.h:524
LLVM_ABI Value * CreateUnaryIntrinsic(Intrinsic::ID ID, Value *Op, FMFSource FMFSource={}, const Twine &Name="")
Create a call to intrinsic ID with 1 operand which is mangled on its type.
InstSimplifyFolder - Use InstructionSimplify to fold operations to existing values.
CostType getValue() const
This function is intended to be used as sparingly as possible, since the class provides the full rang...
InstructionWorklist - This is the worklist management logic for InstCombine and other simplification ...
void push(Instruction *I)
Push the instruction onto the worklist stack.
LLVM_ABI void setHasNoUnsignedWrap(bool b=true)
Set or clear the nuw flag on this instruction, which must be an operator which supports this flag.
LLVM_ABI void copyIRFlags(const Value *V, bool IncludeWrapFlags=true)
Convenience method to copy supported exact, fast-math, and (optionally) wrapping flags from V to this...
LLVM_ABI void setHasNoSignedWrap(bool b=true)
Set or clear the nsw flag on this instruction, which must be an operator which supports this flag.
const DebugLoc & getDebugLoc() const
Return the debug location for this node as a DebugLoc.
LLVM_ABI void andIRFlags(const Value *V)
Logical 'and' of any supported wrapping, exact, and fast-math flags of V and this instruction.
bool isBinaryOp() const
LLVM_ABI void setNonNeg(bool b=true)
Set or clear the nneg flag on this instruction, which must be a zext instruction.
LLVM_ABI bool comesBefore(const Instruction *Other) const
Given an instruction Other in the same basic block as this instruction, return true if this instructi...
LLVM_ABI void setMetadata(unsigned KindID, MDNode *Node)
Set the metadata of the specified kind to the specified node.
LLVM_ABI FastMathFlags getFastMathFlags() const LLVM_READONLY
Convenience function for getting all the fast-math flags, which must be an operator which supports th...
LLVM_ABI AAMDNodes getAAMetadata() const
Returns the AA metadata for this instruction.
unsigned getOpcode() const
Returns a member of one of the enums like Instruction::Add.
bool isIdempotent() const
Return true if the instruction is idempotent:
LLVM_ABI void copyMetadata(const Instruction &SrcInst, ArrayRef< unsigned > WL=ArrayRef< unsigned >())
Copy metadata from SrcInst to this instruction.
LLVM_ABI bool hasAllowReassoc() const LLVM_READONLY
Determine whether the allow-reassociation flag is set.
bool isIntDivRem() const
Class to represent integer types.
static LLVM_ABI IntegerType * get(LLVMContext &C, unsigned NumBits)
This static method is the primary way of constructing an IntegerType.
Definition Type.cpp:348
unsigned getBitWidth() const
Get the number of bits in this IntegerType.
A wrapper class for inspecting calls to intrinsic functions.
Intrinsic::ID getIntrinsicID() const
Return the intrinsic ID of this intrinsic.
An instruction for reading from memory.
unsigned getPointerAddressSpace() const
Returns the address space of the pointer operand.
void setAlignment(Align Align)
Type * getPointerOperandType() const
Align getAlign() const
Return the alignment of the access that is being performed.
Representation for a specific memory location.
static LLVM_ABI MemoryLocation get(const LoadInst *LI)
Return a location with information about the memory reference by the given instruction.
void addIncoming(Value *V, BasicBlock *BB)
Add an incoming value to the end of the PHI list.
static LLVM_ABI PoisonValue * get(Type *T)
Static factory methods - Return an 'poison' object of the specified type.
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
PreservedAnalyses & preserveSet()
Mark an analysis set as preserved.
Definition Analysis.h:151
const SDValue & getOperand(unsigned Num) const
bool contains(const_arg_type key) const
Check if the SetVector contains the given key.
Definition SetVector.h:258
bool empty() const
Determine if the SetVector is empty or not.
Definition SetVector.h:100
bool insert(const value_type &X)
Insert a new element into the SetVector.
Definition SetVector.h:157
This instruction constructs a fixed permutation of two input vectors.
int getMaskValue(unsigned Elt) const
Return the shuffle mask value of this instruction for the given element index.
VectorType * getType() const
Overload to return most specific vector type.
static LLVM_ABI void getShuffleMask(const Constant *Mask, SmallVectorImpl< int > &Result)
Convert the input shuffle mask operand to a vector of integers.
static LLVM_ABI bool isIdentityMask(ArrayRef< int > Mask, int NumSrcElts)
Return true if this shuffle mask chooses elements from exactly one source vector without lane crossin...
static void commuteShuffleMask(MutableArrayRef< int > Mask, unsigned InVecNumElts)
Change values in a shuffle permute mask assuming the two vector operands of length InVecNumElts have ...
size_type size() const
Definition SmallPtrSet.h:99
std::pair< iterator, bool > insert(PtrType Ptr)
Inserts Ptr if and only if there is no element in the container equal to Ptr.
bool contains(ConstPtrType Ptr) const
SmallPtrSet - This class implements a set which is optimized for holding SmallSize or less elements.
void assign(size_type NumElts, ValueParamT Elt)
reference emplace_back(ArgTypes &&... Args)
void reserve(size_type N)
void append(ItTy in_start, ItTy in_end)
Add the specified range to the end of the SmallVector.
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
void setAlignment(Align Align)
Analysis pass providing the TargetTransformInfo.
This pass provides access to the codegen interfaces that are needed for IR-level transformations.
static LLVM_ABI CastContextHint getCastContextHint(const Instruction *I)
Calculates a CastContextHint from I.
LLVM_ABI InstructionCost getCmpSelInstrCost(unsigned Opcode, Type *ValTy, Type *CondTy, CmpInst::Predicate VecPred, TTI::TargetCostKind CostKind=TTI::TCK_RecipThroughput, OperandValueInfo Op1Info={OK_AnyValue, OP_None}, OperandValueInfo Op2Info={OK_AnyValue, OP_None}, const Instruction *I=nullptr) const
LLVM_ABI TypeSize getRegisterBitWidth(RegisterKind K) const
static LLVM_ABI OperandValueInfo commonOperandInfo(const Value *X, const Value *Y)
Collect common data between two OperandValueInfo inputs.
LLVM_ABI InstructionCost getMemoryOpCost(unsigned Opcode, Type *Src, Align Alignment, unsigned AddressSpace, TTI::TargetCostKind CostKind=TTI::TCK_RecipThroughput, OperandValueInfo OpdInfo={OK_AnyValue, OP_None}, const Instruction *I=nullptr) const
LLVM_ABI bool allowVectorElementIndexingUsingGEP() const
Returns true if GEP should not be used to index into vectors for this target.
LLVM_ABI InstructionCost getShuffleCost(ShuffleKind Kind, VectorType *DstTy, VectorType *SrcTy, ArrayRef< int > Mask={}, TTI::TargetCostKind CostKind=TTI::TCK_RecipThroughput, int Index=0, VectorType *SubTp=nullptr, ArrayRef< const Value * > Args={}, const Instruction *CxtI=nullptr) const
LLVM_ABI InstructionCost getIntrinsicInstrCost(const IntrinsicCostAttributes &ICA, TTI::TargetCostKind CostKind) const
LLVM_ABI InstructionCost getArithmeticReductionCost(unsigned Opcode, VectorType *Ty, std::optional< FastMathFlags > FMF, TTI::TargetCostKind CostKind=TTI::TCK_RecipThroughput) const
Calculate the cost of vector reduction intrinsics.
LLVM_ABI InstructionCost getCastInstrCost(unsigned Opcode, Type *Dst, Type *Src, TTI::CastContextHint CCH, TTI::TargetCostKind CostKind=TTI::TCK_SizeAndLatency, const Instruction *I=nullptr) const
LLVM_ABI InstructionCost getVectorInstrCost(unsigned Opcode, Type *Val, TTI::TargetCostKind CostKind, unsigned Index=-1, const Value *Op0=nullptr, const Value *Op1=nullptr, TTI::VectorInstrContext VIC=TTI::VectorInstrContext::None) const
LLVM_ABI InstructionCost getGEPCost(Type *PointeeType, const Value *Ptr, ArrayRef< const Value * > Operands, Type *AccessType=nullptr, TargetCostKind CostKind=TCK_SizeAndLatency) const
Estimate the cost of a GEP operation when lowered.
LLVM_ABI unsigned getRegisterClassForType(bool Vector, Type *Ty=nullptr) const
LLVM_ABI InstructionCost getMinMaxReductionCost(Intrinsic::ID IID, VectorType *Ty, FastMathFlags FMF=FastMathFlags(), TTI::TargetCostKind CostKind=TTI::TCK_RecipThroughput) const
TargetCostKind
The kind of cost model.
@ TCK_RecipThroughput
Reciprocal throughput.
@ TCK_CodeSize
Instruction code size.
LLVM_ABI InstructionCost getArithmeticInstrCost(unsigned Opcode, Type *Ty, TTI::TargetCostKind CostKind=TTI::TCK_RecipThroughput, TTI::OperandValueInfo Opd1Info={TTI::OK_AnyValue, TTI::OP_None}, TTI::OperandValueInfo Opd2Info={TTI::OK_AnyValue, TTI::OP_None}, ArrayRef< const Value * > Args={}, const Instruction *CxtI=nullptr, const TargetLibraryInfo *TLibInfo=nullptr) const
This is an approximation of reciprocal throughput of a math/logic op.
LLVM_ABI InstructionCost getMemIntrinsicInstrCost(const MemIntrinsicCostAttributes &MICA, TTI::TargetCostKind CostKind) const
LLVM_ABI unsigned getMinVectorRegisterBitWidth() const
LLVM_ABI InstructionCost getAddressComputationCost(Type *PtrTy, ScalarEvolution *SE, const SCEV *Ptr, TTI::TargetCostKind CostKind) const
LLVM_ABI unsigned getNumberOfRegisters(unsigned ClassID) const
LLVM_ABI InstructionCost getInstructionCost(const User *U, ArrayRef< const Value * > Operands, TargetCostKind CostKind) const
Estimate the cost of a given IR user when lowered.
LLVM_ABI InstructionCost getScalarizationOverhead(VectorType *Ty, const APInt &DemandedElts, bool Insert, bool Extract, TTI::TargetCostKind CostKind, bool ForPoisonSrc=true, ArrayRef< Value * > VL={}, TTI::VectorInstrContext VIC=TTI::VectorInstrContext::None) const
Estimate the overhead of scalarizing an instruction.
ShuffleKind
The various kinds of shuffle patterns for vector queries.
@ SK_PermuteSingleSrc
Shuffle elements of single source vector with any shuffle mask.
@ SK_PermuteTwoSrc
Merge elements from two source vectors into one with any shuffle mask.
@ SK_ExtractSubvector
ExtractSubvector Index indicates start offset.
@ None
The cast is not used with a load/store of any kind.
The instances of the Type class are immutable: once they are created, they are never changed.
Definition Type.h:46
LLVM_ABI unsigned getIntegerBitWidth() const
bool isPointerTy() const
True if this is an instance of PointerType.
Definition Type.h:282
Type * getScalarType() const
If this is a vector type, return the element type, otherwise return 'this'.
Definition Type.h:368
LLVM_ABI TypeSize getPrimitiveSizeInBits() const LLVM_READONLY
Return the basic size of this type if it is a primitive type.
Definition Type.cpp:197
LLVMContext & getContext() const
Return the LLVMContext in which this type was uniqued.
Definition Type.h:130
LLVM_ABI unsigned getScalarSizeInBits() const LLVM_READONLY
If this is a vector type, return the getPrimitiveSizeInBits value for the element type.
Definition Type.cpp:232
bool isFloatingPointTy() const
Return true if this is one of the floating-point types.
Definition Type.h:186
bool isIntegerTy() const
True if this is an instance of IntegerType.
Definition Type.h:257
bool isFPOrFPVectorTy() const
Return true if this is a FP type or a vector of FP.
Definition Type.h:227
A Use represents the edge between a Value definition and its users.
Definition Use.h:35
op_range operands()
Definition User.h:267
Value * getOperand(unsigned i) const
Definition User.h:207
LLVM Value Representation.
Definition Value.h:75
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:255
const Value * stripAndAccumulateInBoundsConstantOffsets(const DataLayout &DL, APInt &Offset) const
This is a wrapper around stripAndAccumulateConstantOffsets with the in-bounds requirement set to fals...
Definition Value.h:727
LLVM_ABI bool hasOneUser() const
Return true if there is exactly one user of this value.
Definition Value.cpp:163
bool hasOneUse() const
Return true if there is exactly one use of this value.
Definition Value.h:439
LLVM_ABI void replaceAllUsesWith(Value *V)
Change all uses of this to point to a new Value.
Definition Value.cpp:553
iterator_range< user_iterator > users()
Definition Value.h:426
LLVM_ABI Align getPointerAlignment(const DataLayout &DL) const
Returns an alignment of the pointer value.
Definition Value.cpp:1002
unsigned getValueID() const
Return an ID for the concrete type of this object.
Definition Value.h:543
LLVM_ABI bool hasNUses(unsigned N) const
Return true if this Value has exactly N uses.
Definition Value.cpp:147
LLVM_ABI const Value * stripPointerCasts() const
Strip off pointer casts, all-zero GEPs and address space casts.
Definition Value.cpp:713
bool use_empty() const
Definition Value.h:346
LLVM_ABI StringRef getName() const
Return a constant reference to the value's name.
Definition Value.cpp:319
bool user_empty() const
Definition Value.h:389
LLVM_ABI PreservedAnalyses run(Function &F, FunctionAnalysisManager &)
static LLVM_ABI VectorType * get(Type *ElementType, ElementCount EC)
This static method is the primary way to construct an VectorType.
Type * getElementType() const
std::pair< iterator, bool > insert(const ValueT &V)
Definition DenseSet.h:209
size_type size() const
Definition DenseSet.h:84
constexpr bool hasKnownScalarFactor(const FixedOrScalableQuantity &RHS) const
Returns true if there exists a value X where RHS.multiplyCoefficientBy(X) will result in a value whos...
Definition TypeSize.h:269
constexpr ScalarTy getFixedValue() const
Definition TypeSize.h:200
constexpr ScalarTy getKnownScalarFactor(const FixedOrScalableQuantity &RHS) const
Returns a value X where RHS.multiplyCoefficientBy(X) will result in a value whose quantity matches ou...
Definition TypeSize.h:277
constexpr bool isScalable() const
Returns whether the quantity is scaled by a runtime quantity (vscale).
Definition TypeSize.h:168
constexpr ScalarTy getKnownMinValue() const
Returns the minimum value this quantity can represent.
Definition TypeSize.h:165
constexpr bool isZero() const
Definition TypeSize.h:153
const ParentTy * getParent() const
Definition ilist_node.h:34
self_iterator getIterator()
Definition ilist_node.h:123
NodeTy * getNextNode()
Get the next node, or nullptr for the list tail.
Definition ilist_node.h:348
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
Abstract Attribute helper functions.
Definition Attributor.h:165
constexpr char Align[]
Key for Kernel::Arg::Metadata::mAlign.
const APInt & smin(const APInt &A, const APInt &B)
Determine the smaller of two APInts considered to be signed.
Definition APInt.h:2275
const APInt & smax(const APInt &A, const APInt &B)
Determine the larger of two APInts considered to be signed.
Definition APInt.h:2280
constexpr std::underlying_type_t< E > Mask()
Get a bitmask with 1s in all places up to the high-order bit of E's largest value.
@ BasicBlock
Various leaf nodes.
Definition ISDOpcodes.h:81
LLVM_ABI Intrinsic::ID getInterleaveIntrinsicID(unsigned Factor)
Returns the corresponding llvm.vector.interleaveN intrinsic for factor N.
SpecificConstantMatch m_ZeroInt()
Convenience matchers for specific integer values.
BinaryOp_match< SpecificConstantMatch, SrcTy, TargetOpcode::G_SUB > m_Neg(const SrcTy &&Src)
Matches a register negated by a G_SUB.
AllOnesConstantMatch m_AllOnes()
OneUse_match< SubPat > m_OneUse(const SubPat &SP)
match_combine_and< Ty... > m_CombineAnd(const Ty &...Ps)
Combine pattern matchers matching all of Ps patterns.
BinaryOp_match< LHS, RHS, Instruction::And > m_And(const LHS &L, const RHS &R)
auto m_BSwap(const Opnd0 &Op0)
auto m_Cmp()
Matches any compare instruction and ignore it.
BinaryOp_match< LHS, RHS, Instruction::Add > m_Add(const LHS &L, const RHS &R)
auto m_BitReverse(const Opnd0 &Op0)
BinaryOp_match< LHS, RHS, Instruction::URem > m_URem(const LHS &L, const RHS &R)
auto m_Poison()
Match an arbitrary poison constant.
ap_match< APInt > m_APInt(const APInt *&Res)
Match a ConstantInt or splatted ConstantVector, binding the specified pointer to the contained APInt.
CastInst_match< OpTy, TruncInst > m_Trunc(const OpTy &Op)
Matches Trunc.
specific_intval< false > m_SpecificInt(const APInt &V)
Match a specific integer value or vector with all elements equal to the value.
bool match(Val *V, const Pattern &P)
match_bind< Instruction > m_Instruction(Instruction *&I)
Match an instruction, capturing it if we match.
specificval_ty m_Specific(const Value *V)
Match if we have a specific specified value.
DisjointOr_match< LHS, RHS > m_DisjointOr(const LHS &L, const RHS &R)
BinOpPred_match< LHS, RHS, is_right_shift_op > m_Shr(const LHS &L, const RHS &R)
Matches logical shift operations.
CmpClass_match< LHS, RHS, ICmpInst, true > m_c_ICmp(CmpPredicate &Pred, const LHS &L, const RHS &R)
Matches an ICmp with a predicate over LHS and RHS in either order.
TwoOps_match< Val_t, Idx_t, Instruction::ExtractElement > m_ExtractElt(const Val_t &Val, const Idx_t &Idx)
Matches ExtractElementInst.
ThreeOps_match< Cond, LHS, RHS, Instruction::Select > m_Select(const Cond &C, const LHS &L, const RHS &R)
Matches SelectInst.
auto m_BinOp()
Match an arbitrary binary operation and ignore it.
auto m_Value()
Match an arbitrary value and ignore it.
BinaryOp_match< LHS, RHS, Instruction::Mul > m_Mul(const LHS &L, const RHS &R)
auto m_Constant()
Match an arbitrary Constant and ignore it.
TwoOps_match< V1_t, V2_t, Instruction::ShuffleVector > m_Shuffle(const V1_t &v1, const V2_t &v2)
Matches ShuffleVectorInst independently of mask value.
cst_pred_ty< is_non_zero_int > m_NonZeroInt()
Match a non-zero integer or a vector with all non-zero elements.
OneOps_match< OpTy, Instruction::Load > m_Load(const OpTy &Op)
Matches LoadInst.
CastInst_match< OpTy, ZExtInst > m_ZExt(const OpTy &Op)
Matches ZExt.
OverflowingBinaryOp_match< LHS, RHS, Instruction::Shl, OverflowingBinaryOperator::NoUnsignedWrap > m_NUWShl(const LHS &L, const RHS &R)
auto m_AnyIntrinsic()
Matches any intrinsic call and ignore it.
OverflowingBinaryOp_match< LHS, RHS, Instruction::Mul, OverflowingBinaryOperator::NoUnsignedWrap > m_NUWMul(const LHS &L, const RHS &R)
BinOpPred_match< LHS, RHS, is_bitwiselogic_op, true > m_c_BitwiseLogic(const LHS &L, const RHS &R)
Matches bitwise logic operations in either order.
CastOperator_match< OpTy, Instruction::BitCast > m_BitCast(const OpTy &Op)
Matches BitCast.
match_combine_or< CastInst_match< OpTy, SExtInst >, NNegZExt_match< OpTy > > m_SExtLike(const OpTy &Op)
Match either "sext" or "zext nneg".
auto m_Intrinsic(const Ts &...Ops)
Match intrinsic calls like this: m_Intrinsic<Intrinsic::fabs>(m_Value(X))
auto m_Deinterleave2(const Opnd &Op)
BinaryOp_match< LHS, RHS, Instruction::LShr > m_LShr(const LHS &L, const RHS &R)
CmpClass_match< LHS, RHS, ICmpInst > m_ICmp(CmpPredicate &Pred, const LHS &L, const RHS &R)
match_combine_or< CastInst_match< OpTy, ZExtInst >, CastInst_match< OpTy, SExtInst > > m_ZExtOrSExt(const OpTy &Op)
FNeg_match< OpTy > m_FNeg(const OpTy &X)
Match 'fneg X' as 'fsub -0.0, X'.
BinaryOp_match< LHS, RHS, Instruction::Shl > m_Shl(const LHS &L, const RHS &R)
auto m_Undef()
Match an arbitrary undef constant.
CastInst_match< OpTy, SExtInst > m_SExt(const OpTy &Op)
Matches SExt.
is_zero m_Zero()
Match any null constant or a vector with all elements equal to 0.
BinaryOp_match< LHS, RHS, Instruction::Or, true > m_c_Or(const LHS &L, const RHS &R)
Matches an Or with LHS and RHS in either order.
ThreeOps_match< Val_t, Elt_t, Idx_t, Instruction::InsertElement > m_InsertElt(const Val_t &Val, const Elt_t &Elt, const Idx_t &Idx)
Matches InsertElementInst.
auto m_ConstantInt()
Match an arbitrary ConstantInt and ignore it.
@ Valid
The data is already valid.
initializer< Ty > init(const Ty &Val)
DXILDebugInfoMap run(Module &M)
@ User
could "use" a pointer
NodeAddr< PhiNode * > Phi
Definition RDFGraph.h:390
NodeAddr< UseNode * > Use
Definition RDFGraph.h:385
friend class Instruction
Iterator for Instructions in a `BasicBlock.
Definition BasicBlock.h:73
unsigned getOpcode(const VPValue *V)
Return the instruction opcode for the recipe defining V or 0 for unsupported recipes and VPValues not...
This is an optimization pass for GlobalISel generic memory operations.
auto drop_begin(T &&RangeOrContainer, size_t N=1)
Return a range covering RangeOrContainer with the first N elements excluded.
Definition STLExtras.h:315
unsigned Log2_32_Ceil(uint32_t Value)
Return the ceil log base 2 of the specified value, 32 if the value is zero.
Definition MathExtras.h:345
LLVM_ABI bool willNotFreeBetween(const Instruction *Assume, const Instruction *CtxI)
Returns true, if no instruction between Assume and CtxI may free (including through synchronization).
@ Offset
Definition DWP.cpp:578
detail::zippy< detail::zip_shortest, T, U, Args... > zip(T &&t, U &&u, Args &&...args)
zip iterator for two or more iteratable types.
Definition STLExtras.h:830
void stable_sort(R &&Range)
Definition STLExtras.h:2116
LLVM_ABI cl::opt< bool > ProfcheckDisableMetadataFixes
Definition LoopInfo.cpp:60
UnaryFunction for_each(R &&Range, UnaryFunction F)
Provide wrappers to std::for_each which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1732
bool all_of(R &&range, UnaryPredicate P)
Provide wrappers to std::all_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1739
LLVM_ABI Intrinsic::ID getMinMaxReductionIntrinsicOp(Intrinsic::ID RdxID)
Returns the min/max intrinsic used when expanding a min/max reduction.
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.
Definition Local.cpp:535
RelativeUniformCounterPtr Values
Definition InstrProf.h:91
LLVM_ABI SDValue peekThroughBitcasts(SDValue V)
Return the non-bitcasted source operand of V if it exists.
auto enumerate(FirstRange &&First, RestRanges &&...Rest)
Given two or more input ranges, returns a new range whose values are tuples (A, B,...
Definition STLExtras.h:2554
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
LLVM_ABI Value * simplifyUnOp(unsigned Opcode, Value *Op, const SimplifyQuery &Q)
Given operand for a UnaryOperator, fold the result or return null.
scope_exit(Callable) -> scope_exit< Callable >
@ Load
The value being inserted comes from a load (InsertElement only).
auto map_to_vector(ContainerTy &&C, FuncTy &&F)
Map a range to a SmallVector with element types deduced from the mapping.
iterator_range< T > make_range(T x, T y)
Convenience function for iterating over sub-ranges.
LLVM_ABI unsigned getArithmeticReductionInstruction(Intrinsic::ID RdxID)
Returns the arithmetic instruction opcode used when expanding a reduction.
void append_range(Container &C, Range &&R)
Wrapper function to append range R to container C.
Definition STLExtras.h:2208
constexpr bool isUIntN(unsigned N, uint64_t x)
Checks if an unsigned integer fits into the given (dynamic) bit width.
Definition MathExtras.h:244
LLVM_ABI Value * simplifyCall(CallBase *Call, Value *Callee, ArrayRef< Value * > Args, const SimplifyQuery &Q)
Given a callsite, callee, and arguments, fold the result or return null.
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...
Definition STLExtras.h:633
LLVM_ABI bool mustSuppressSpeculation(const LoadInst &LI)
Return true if speculation of the given load must be suppressed to avoid ordering or interfering with...
Definition Loads.cpp:452
LLVM_ABI bool widenShuffleMaskElts(int Scale, ArrayRef< int > Mask, SmallVectorImpl< int > &ScaledMask)
Try to transform a shuffle mask by replacing elements with the scaled index for an equivalent mask of...
LLVM_ABI bool isSafeToSpeculativelyExecute(const Instruction *I, const Instruction *CtxI=nullptr, AssumptionCache *AC=nullptr, const DominatorTree *DT=nullptr, const TargetLibraryInfo *TLI=nullptr, bool UseVariableInfo=true, bool IgnoreUBImplyingAttrs=true)
Return true if the instruction does not have any effects besides calculating the result and does not ...
LLVM_ABI Instruction * propagateMetadata(Instruction *I, ArrayRef< Value * > VL)
Specifically, let Kinds = [MD_tbaa, MD_alias_scope, MD_noalias, MD_fpmath, MD_nontemporal,...
LLVM_ABI Value * getSplatValue(const Value *V)
Get splat value if the input is a splat vector or return nullptr.
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Value
Definition InstrProf.h:143
unsigned M1(unsigned Val)
Definition VE.h:377
bool any_of(R &&range, UnaryPredicate P)
Provide wrappers to std::any_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1746
LLVM_ABI bool isInstructionTriviallyDead(Instruction *I, const TargetLibraryInfo *TLI=nullptr)
Return true if the result produced by the instruction is not used, and the instruction will return.
Definition Local.cpp:403
LLVM_ABI bool isSplatValue(const Value *V, int Index=-1, unsigned Depth=0)
Return true if each element of the vector value V is poisoned or equal to every other non-poisoned el...
unsigned Log2_32(uint32_t Value)
Return the floor log base 2 of the specified value, -1 if the value is zero.
Definition MathExtras.h:332
auto reverse(ContainerTy &&C)
Definition STLExtras.h:407
constexpr bool isPowerOf2_32(uint32_t Value)
Return true if the argument is a power of two > 0.
Definition MathExtras.h:280
bool isModSet(const ModRefInfo MRI)
Definition ModRef.h:49
void sort(IteratorTy Start, IteratorTy End)
Definition STLExtras.h:1636
LLVM_ABI void computeKnownBits(const Value *V, KnownBits &Known, const DataLayout &DL, AssumptionCache *AC=nullptr, const Instruction *CxtI=nullptr, const DominatorTree *DT=nullptr, bool UseInstrInfo=true, unsigned Depth=0)
Determine which bits of V are known to be either zero or one and return them in the KnownZero/KnownOn...
LLVM_ABI bool programUndefinedIfPoison(const Instruction *Inst)
LLVM_ABI bool isSafeToLoadUnconditionally(Value *V, Align Alignment, const APInt &Size, const DataLayout &DL, Instruction *ScanFrom, AssumptionCache *AC=nullptr, const DominatorTree *DT=nullptr, const TargetLibraryInfo *TLI=nullptr)
Return true if we know that executing a load from this value cannot trap.
Definition Loads.cpp:456
LLVM_ABI unsigned getDeinterleaveIntrinsicFactor(Intrinsic::ID ID)
Returns the corresponding factor of llvm.vector.deinterleaveN intrinsics.
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
Definition Debug.cpp:209
constexpr uint64_t alignTo(uint64_t Size, Align A)
Returns a multiple of A needed to store Size bytes.
Definition Alignment.h:144
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...
Definition Casting.h:547
LLVM_ABI void propagateIRFlags(Value *I, ArrayRef< Value * > VL, Value *OpValue=nullptr, bool IncludeWrapFlags=true)
Get the intersection (logical and) of all of the potential IR flags of each scalar operation (VL) tha...
MutableArrayRef(T &OneElt) -> MutableArrayRef< T >
constexpr int PoisonMaskElem
@ Other
Any other memory.
Definition ModRef.h:68
TargetTransformInfo TTI
IRBuilder(LLVMContext &, FolderTy, InserterTy, MDNode *, ArrayRef< OperandBundleDef >) -> IRBuilder< FolderTy, InserterTy >
LLVM_ABI Value * simplifyBinOp(unsigned Opcode, Value *LHS, Value *RHS, const SimplifyQuery &Q)
Given operands for a BinaryOperator, fold the result or return null.
LLVM_ABI void narrowShuffleMaskElts(int Scale, ArrayRef< int > Mask, SmallVectorImpl< int > &ScaledMask)
Replace each shuffle mask index with the scaled sequential indices for an equivalent mask of narrowed...
LLVM_ABI Intrinsic::ID getReductionForBinop(Instruction::BinaryOps Opc)
Returns the reduction intrinsic id corresponding to the binary operation.
@ And
Bitwise or logical AND of integers.
LLVM_ABI bool isVectorIntrinsicWithScalarOpAtArg(Intrinsic::ID ID, unsigned ScalarOpdIdx, const TargetTransformInfo *TTI)
Identifies if the vector form of the intrinsic has a scalar operand.
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Count
Definition InstrProf.h:145
DWARFExpression::Operation Op
unsigned M0(unsigned Val)
Definition VE.h:376
ArrayRef(const T &OneElt) -> ArrayRef< T >
LLVM_ABI unsigned ComputeNumSignBits(const Value *Op, const DataLayout &DL, AssumptionCache *AC=nullptr, const Instruction *CxtI=nullptr, const DominatorTree *DT=nullptr, bool UseInstrInfo=true, unsigned Depth=0)
Return the number of times the sign bit of the register is replicated into the other bits.
constexpr unsigned BitWidth
LLVM_ABI bool isGuaranteedToTransferExecutionToSuccessor(const Instruction *I)
Return true if this function can prove that the instruction I will always transfer execution to one o...
LLVM_ABI Constant * getLosslessInvCast(Constant *C, Type *InvCastTo, unsigned CastOp, const DataLayout &DL, PreservedCastFlags *Flags=nullptr)
Try to cast C to InvC losslessly, satisfying CastOp(InvC) equals C, or CastOp(InvC) is a refined valu...
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
auto find_if(R &&Range, UnaryPredicate P)
Provide wrappers to std::find_if which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1772
constexpr bool isIntN(unsigned N, int64_t x)
Checks if an signed integer fits into the given (dynamic) bit width.
Definition MathExtras.h:249
bool is_contained(R &&Range, const E &Element)
Returns true if Element is found in Range.
Definition STLExtras.h:1947
Align commonAlignment(Align A, uint64_t Offset)
Returns the alignment that satisfies both alignments.
Definition Alignment.h:201
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Next
Definition InstrProf.h:147
bool all_equal(std::initializer_list< T > Values)
Returns true if all Values in the initializer lists are equal or the list.
Definition STLExtras.h:2166
LLVM_ABI Value * simplifyCmpInst(CmpPredicate Predicate, Value *LHS, Value *RHS, const SimplifyQuery &Q)
Given operands for a CmpInst, fold the result or return null.
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
LLVM_ABI bool isGuaranteedNotToBePoison(const Value *V, AssumptionCache *AC=nullptr, const Instruction *CtxI=nullptr, const DominatorTree *DT=nullptr, unsigned Depth=0)
Returns true if V cannot be poison, but may be undef.
LLVM_ABI bool isKnownNonNegative(const Value *V, const SimplifyQuery &SQ, unsigned Depth=0)
Returns true if the give value is known to be non-negative.
LLVM_ABI bool isTriviallyVectorizable(Intrinsic::ID ID)
Identify if the intrinsic is trivially vectorizable.
LLVM_ABI Intrinsic::ID getMinMaxReductionIntrinsicID(Intrinsic::ID IID)
Returns the llvm.vector.reduce min/max intrinsic that corresponds to the intrinsic op.
LLVM_ABI ConstantRange computeConstantRange(const Value *V, bool ForSigned, const SimplifyQuery &SQ, unsigned Depth=0)
Determine the possible constant range of an integer or vector of integer value.
void swap(llvm::BitVector &LHS, llvm::BitVector &RHS)
Implement std::swap in terms of BitVector swap.
Definition BitVector.h:880
#define N
LLVM_ABI AAMDNodes adjustForAccess(unsigned AccessSize)
Create a new AAMDNode for accessing AccessSize bytes of this AAMDNode.
This struct is a compact representation of a valid (non-zero power of two) alignment.
Definition Alignment.h:39
unsigned countMaxActiveBits() const
Returns the maximum number of bits needed to represent all possible unsigned values with these known ...
Definition KnownBits.h:310
unsigned countMinLeadingZeros() const
Returns the minimum number of leading zero bits.
Definition KnownBits.h:262
APInt getMaxValue() const
Return the maximal unsigned value possible given these KnownBits.
Definition KnownBits.h:146
const DataLayout & DL
const Instruction * CxtI
const DominatorTree * DT
SimplifyQuery getWithInstruction(const Instruction *I) const
AssumptionCache * AC