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