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