1//===- Scalarizer.cpp - Scalarize 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 converts vector operations into scalar operations (or, optionally,
10// operations on smaller vector widths), in order to expose optimization
11// opportunities on the individual scalar operations.
12// It is mainly intended for targets that do not have vector units, but it
13// may also be useful for revectorizing code to different vector widths.
14//
15//===----------------------------------------------------------------------===//
16
17#include "llvm/Transforms/Scalar/Scalarizer.h"
18#include "llvm/ADT/PostOrderIterator.h"
19#include "llvm/ADT/SmallVector.h"
20#include "llvm/ADT/Twine.h"
21#include "llvm/Analysis/TargetTransformInfo.h"
22#include "llvm/Analysis/VectorUtils.h"
23#include "llvm/IR/Argument.h"
24#include "llvm/IR/BasicBlock.h"
25#include "llvm/IR/Constants.h"
26#include "llvm/IR/DataLayout.h"
27#include "llvm/IR/DerivedTypes.h"
28#include "llvm/IR/Dominators.h"
29#include "llvm/IR/Function.h"
30#include "llvm/IR/IRBuilder.h"
31#include "llvm/IR/InstVisitor.h"
32#include "llvm/IR/InstrTypes.h"
33#include "llvm/IR/Instruction.h"
34#include "llvm/IR/Instructions.h"
35#include "llvm/IR/Intrinsics.h"
36#include "llvm/IR/LLVMContext.h"
37#include "llvm/IR/Module.h"
38#include "llvm/IR/Type.h"
39#include "llvm/IR/Value.h"
40#include "llvm/InitializePasses.h"
41#include "llvm/Support/Casting.h"
42#include "llvm/Transforms/Utils/Local.h"
43#include <cassert>
44#include <cstdint>
45#include <iterator>
46#include <map>
47#include <utility>
48
49using namespace llvm;
50
51#define DEBUG_TYPE "scalarizer"
52
53static BasicBlock::iterator skipPastPhiNodesAndDbg(BasicBlock::iterator Itr) {
54 BasicBlock *BB = Itr->getParent();
55 if (isa<PHINode>(Val: Itr))
56 Itr = BB->getFirstInsertionPt();
57 if (Itr != BB->end())
58 Itr = skipDebugIntrinsics(It: Itr);
59 return Itr;
60}
61
62// Used to store the scattered form of a vector.
63using ValueVector = SmallVector<Value *, 8>;
64
65// Used to map a vector Value and associated type to its scattered form.
66// The associated type is only non-null for pointer values that are "scattered"
67// when used as pointer operands to load or store.
68//
69// We use std::map because we want iterators to persist across insertion and
70// because the values are relatively large.
71using ScatterMap = std::map<std::pair<Value *, Type *>, ValueVector>;
72
73// Lists Instructions that have been replaced with scalar implementations,
74// along with a pointer to their scattered forms.
75using GatherList = SmallVector<std::pair<Instruction *, ValueVector *>, 16>;
76
77namespace {
78
79struct VectorSplit {
80 // The type of the vector.
81 FixedVectorType *VecTy = nullptr;
82
83 // The number of elements packed in a fragment (other than the remainder).
84 unsigned NumPacked = 0;
85
86 // The number of fragments (scalars or smaller vectors) into which the vector
87 // shall be split.
88 unsigned NumFragments = 0;
89
90 // The type of each complete fragment.
91 Type *SplitTy = nullptr;
92
93 // The type of the remainder (last) fragment; null if all fragments are
94 // complete.
95 Type *RemainderTy = nullptr;
96
97 Type *getFragmentType(unsigned I) const {
98 return RemainderTy && I == NumFragments - 1 ? RemainderTy : SplitTy;
99 }
100};
101
102// Provides a very limited vector-like interface for lazily accessing one
103// component of a scattered vector or vector pointer.
104class Scatterer {
105public:
106 Scatterer() = default;
107
108 // Scatter V into Size components. If new instructions are needed,
109 // insert them before BBI. If Cache is nonnull, use it to cache
110 // the results.
111 Scatterer(BasicBlock::iterator bbi, Value *v, const VectorSplit &VS,
112 ValueVector *cachePtr = nullptr);
113
114 // Return component I, creating a new Value for it if necessary.
115 Value *operator[](unsigned I);
116
117 // Return the number of components.
118 unsigned size() const { return VS.NumFragments; }
119
120private:
121 BasicBlock::iterator BBI;
122 Value *V;
123 VectorSplit VS;
124 bool IsPointer;
125 ValueVector *CachePtr;
126 ValueVector Tmp;
127};
128
129// FCmpSplitter(FCI)(Builder, X, Y, Name) uses Builder to create an FCmp
130// called Name that compares X and Y in the same way as FCI.
131struct FCmpSplitter {
132 FCmpSplitter(FCmpInst &fci) : FCI(fci) {}
133
134 Value *operator()(IRBuilder<> &Builder, Value *Op0, Value *Op1,
135 const Twine &Name) const {
136 return Builder.CreateFCmp(P: FCI.getPredicate(), LHS: Op0, RHS: Op1, Name);
137 }
138
139 FCmpInst &FCI;
140};
141
142// ICmpSplitter(ICI)(Builder, X, Y, Name) uses Builder to create an ICmp
143// called Name that compares X and Y in the same way as ICI.
144struct ICmpSplitter {
145 ICmpSplitter(ICmpInst &ici) : ICI(ici) {}
146
147 Value *operator()(IRBuilder<> &Builder, Value *Op0, Value *Op1,
148 const Twine &Name) const {
149 return Builder.CreateICmp(P: ICI.getPredicate(), LHS: Op0, RHS: Op1, Name);
150 }
151
152 ICmpInst &ICI;
153};
154
155// UnarySplitter(UO)(Builder, X, Name) uses Builder to create
156// a unary operator like UO called Name with operand X.
157struct UnarySplitter {
158 UnarySplitter(UnaryOperator &uo) : UO(uo) {}
159
160 Value *operator()(IRBuilder<> &Builder, Value *Op, const Twine &Name) const {
161 return Builder.CreateUnOp(Opc: UO.getOpcode(), V: Op, Name);
162 }
163
164 UnaryOperator &UO;
165};
166
167// BinarySplitter(BO)(Builder, X, Y, Name) uses Builder to create
168// a binary operator like BO called Name with operands X and Y.
169struct BinarySplitter {
170 BinarySplitter(BinaryOperator &bo) : BO(bo) {}
171
172 Value *operator()(IRBuilder<> &Builder, Value *Op0, Value *Op1,
173 const Twine &Name) const {
174 return Builder.CreateBinOp(Opc: BO.getOpcode(), LHS: Op0, RHS: Op1, Name);
175 }
176
177 BinaryOperator &BO;
178};
179
180// Information about a load or store that we're scalarizing.
181struct VectorLayout {
182 VectorLayout() = default;
183
184 // Return the alignment of fragment Frag.
185 Align getFragmentAlign(unsigned Frag) {
186 return commonAlignment(A: VecAlign, Offset: Frag * SplitSize);
187 }
188
189 // The split of the underlying vector type.
190 VectorSplit VS;
191
192 // The alignment of the vector.
193 Align VecAlign;
194
195 // The size of each (non-remainder) fragment in bytes.
196 uint64_t SplitSize = 0;
197};
198} // namespace
199
200static bool isStructOfMatchingFixedVectors(Type *Ty) {
201 if (!isa<StructType>(Val: Ty))
202 return false;
203 unsigned StructSize = Ty->getNumContainedTypes();
204 if (StructSize < 1)
205 return false;
206 FixedVectorType *VecTy = dyn_cast<FixedVectorType>(Val: Ty->getContainedType(i: 0));
207 if (!VecTy)
208 return false;
209 unsigned VecSize = VecTy->getNumElements();
210 for (unsigned I = 1; I < StructSize; I++) {
211 VecTy = dyn_cast<FixedVectorType>(Val: Ty->getContainedType(i: I));
212 if (!VecTy || VecSize != VecTy->getNumElements())
213 return false;
214 }
215 return true;
216}
217
218/// Concatenate the given fragments to a single vector value of the type
219/// described in @p VS.
220static Value *concatenate(IRBuilder<> &Builder, ArrayRef<Value *> Fragments,
221 const VectorSplit &VS, Twine Name) {
222 unsigned NumElements = VS.VecTy->getNumElements();
223 SmallVector<int> ExtendMask;
224 SmallVector<int> InsertMask;
225
226 if (VS.NumPacked > 1) {
227 // Prepare the shufflevector masks once and re-use them for all
228 // fragments.
229 ExtendMask.resize(N: NumElements, NV: -1);
230 for (unsigned I = 0; I < VS.NumPacked; ++I)
231 ExtendMask[I] = I;
232
233 InsertMask.resize(N: NumElements);
234 for (unsigned I = 0; I < NumElements; ++I)
235 InsertMask[I] = I;
236 }
237
238 Value *Res = PoisonValue::get(T: VS.VecTy);
239 for (unsigned I = 0; I < VS.NumFragments; ++I) {
240 Value *Fragment = Fragments[I];
241
242 unsigned NumPacked = VS.NumPacked;
243 if (I == VS.NumFragments - 1 && VS.RemainderTy) {
244 if (auto *RemVecTy = dyn_cast<FixedVectorType>(Val: VS.RemainderTy))
245 NumPacked = RemVecTy->getNumElements();
246 else
247 NumPacked = 1;
248 }
249
250 if (NumPacked == 1) {
251 Res = Builder.CreateInsertElement(Vec: Res, NewElt: Fragment, Idx: I * VS.NumPacked,
252 Name: Name + ".upto" + Twine(I));
253 } else {
254 if (NumPacked < VS.NumPacked) {
255 // If last pack of remained bits not match current ExtendMask size.
256 ExtendMask.truncate(N: NumPacked);
257 ExtendMask.resize(N: NumElements, NV: -1);
258 }
259
260 Fragment = Builder.CreateShuffleVector(
261 V1: Fragment, V2: PoisonValue::get(T: Fragment->getType()), Mask: ExtendMask);
262 if (I == 0) {
263 Res = Fragment;
264 } else {
265 for (unsigned J = 0; J < NumPacked; ++J)
266 InsertMask[I * VS.NumPacked + J] = NumElements + J;
267 Res = Builder.CreateShuffleVector(V1: Res, V2: Fragment, Mask: InsertMask,
268 Name: Name + ".upto" + Twine(I));
269 for (unsigned J = 0; J < NumPacked; ++J)
270 InsertMask[I * VS.NumPacked + J] = I * VS.NumPacked + J;
271 }
272 }
273 }
274
275 return Res;
276}
277
278namespace {
279class ScalarizerVisitor : public InstVisitor<ScalarizerVisitor, bool> {
280public:
281 ScalarizerVisitor(DominatorTree *DT, const TargetTransformInfo *TTI,
282 ScalarizerPassOptions Options)
283 : DT(DT), TTI(TTI),
284 ScalarizeVariableInsertExtract(Options.ScalarizeVariableInsertExtract),
285 ScalarizeLoadStore(Options.ScalarizeLoadStore),
286 ScalarizeMinBits(Options.ScalarizeMinBits) {}
287
288 bool visit(Function &F);
289
290 // InstVisitor methods. They return true if the instruction was scalarized,
291 // false if nothing changed.
292 bool visitInstruction(Instruction &I) { return false; }
293 bool visitSelectInst(SelectInst &SI);
294 bool visitICmpInst(ICmpInst &ICI);
295 bool visitFCmpInst(FCmpInst &FCI);
296 bool visitUnaryOperator(UnaryOperator &UO);
297 bool visitBinaryOperator(BinaryOperator &BO);
298 bool visitGetElementPtrInst(GetElementPtrInst &GEPI);
299 bool visitCastInst(CastInst &CI);
300 bool visitBitCastInst(BitCastInst &BCI);
301 bool visitInsertElementInst(InsertElementInst &IEI);
302 bool visitExtractElementInst(ExtractElementInst &EEI);
303 bool visitExtractValueInst(ExtractValueInst &EVI);
304 bool visitShuffleVectorInst(ShuffleVectorInst &SVI);
305 bool visitPHINode(PHINode &PHI);
306 bool visitLoadInst(LoadInst &LI);
307 bool visitStoreInst(StoreInst &SI);
308 bool visitCallInst(CallInst &ICI);
309 bool visitFreezeInst(FreezeInst &FI);
310
311private:
312 Scatterer scatter(Instruction *Point, Value *V, const VectorSplit &VS);
313 void gather(Instruction *Op, const ValueVector &CV, const VectorSplit &VS);
314 void replaceUses(Instruction *Op, Value *CV);
315 bool canTransferMetadata(unsigned Kind);
316 void transferMetadataAndIRFlags(Instruction *Op, const ValueVector &CV);
317 std::optional<VectorSplit> getVectorSplit(Type *Ty);
318 std::optional<VectorLayout> getVectorLayout(Type *Ty, Align Alignment,
319 const DataLayout &DL);
320 bool finish();
321
322 template<typename T> bool splitUnary(Instruction &, const T &);
323 template<typename T> bool splitBinary(Instruction &, const T &);
324
325 bool splitCall(CallInst &CI);
326
327 ScatterMap Scattered;
328 GatherList Gathered;
329 bool Scalarized;
330
331 SmallVector<WeakTrackingVH, 32> PotentiallyDeadInstrs;
332
333 DominatorTree *DT;
334 const TargetTransformInfo *TTI;
335
336 const bool ScalarizeVariableInsertExtract;
337 const bool ScalarizeLoadStore;
338 const unsigned ScalarizeMinBits;
339};
340
341class ScalarizerLegacyPass : public FunctionPass {
342public:
343 static char ID;
344 ScalarizerPassOptions Options;
345 ScalarizerLegacyPass() : FunctionPass(ID), Options() {}
346 ScalarizerLegacyPass(const ScalarizerPassOptions &Options);
347 bool runOnFunction(Function &F) override;
348 void getAnalysisUsage(AnalysisUsage &AU) const override;
349};
350
351} // end anonymous namespace
352
353ScalarizerLegacyPass::ScalarizerLegacyPass(const ScalarizerPassOptions &Options)
354 : FunctionPass(ID), Options(Options) {}
355
356void ScalarizerLegacyPass::getAnalysisUsage(AnalysisUsage &AU) const {
357 AU.addRequired<DominatorTreeWrapperPass>();
358 AU.addRequired<TargetTransformInfoWrapperPass>();
359 AU.addPreserved<DominatorTreeWrapperPass>();
360}
361
362char ScalarizerLegacyPass::ID = 0;
363INITIALIZE_PASS_BEGIN(ScalarizerLegacyPass, "scalarizer",
364 "Scalarize vector operations", false, false)
365INITIALIZE_PASS_DEPENDENCY(DominatorTreeWrapperPass)
366INITIALIZE_PASS_DEPENDENCY(TargetTransformInfoWrapperPass)
367INITIALIZE_PASS_END(ScalarizerLegacyPass, "scalarizer",
368 "Scalarize vector operations", false, false)
369
370Scatterer::Scatterer(BasicBlock::iterator bbi, Value *v, const VectorSplit &VS,
371 ValueVector *cachePtr)
372 : BBI(bbi), V(v), VS(VS), CachePtr(cachePtr) {
373 IsPointer = V->getType()->isPointerTy();
374 if (!CachePtr) {
375 Tmp.resize(N: VS.NumFragments, NV: nullptr);
376 } else {
377 assert((CachePtr->empty() || VS.NumFragments == CachePtr->size() ||
378 IsPointer) &&
379 "Inconsistent vector sizes");
380 if (VS.NumFragments > CachePtr->size())
381 CachePtr->resize(N: VS.NumFragments, NV: nullptr);
382 }
383}
384
385// Return fragment Frag, creating a new Value for it if necessary.
386Value *Scatterer::operator[](unsigned Frag) {
387 ValueVector &CV = CachePtr ? *CachePtr : Tmp;
388 // Try to reuse a previous value.
389 if (CV[Frag])
390 return CV[Frag];
391 IRBuilder<> Builder(BBI);
392 if (IsPointer) {
393 if (Frag == 0)
394 CV[Frag] = V;
395 else
396 CV[Frag] = Builder.CreateConstGEP1_32(Ty: VS.SplitTy, Ptr: V, Idx0: Frag,
397 Name: V->getName() + ".i" + Twine(Frag));
398 return CV[Frag];
399 }
400
401 Type *FragmentTy = VS.getFragmentType(I: Frag);
402
403 if (auto *VecTy = dyn_cast<FixedVectorType>(Val: FragmentTy)) {
404 SmallVector<int> Mask;
405 for (unsigned J = 0; J < VecTy->getNumElements(); ++J)
406 Mask.push_back(Elt: Frag * VS.NumPacked + J);
407 CV[Frag] =
408 Builder.CreateShuffleVector(V1: V, V2: PoisonValue::get(T: V->getType()), Mask,
409 Name: V->getName() + ".i" + Twine(Frag));
410 } else {
411 // Search through a chain of InsertElementInsts looking for element Frag.
412 // Record other elements in the cache. The new V is still suitable
413 // for all uncached indices.
414 while (true) {
415 InsertElementInst *Insert = dyn_cast<InsertElementInst>(Val: V);
416 if (!Insert)
417 break;
418 ConstantInt *Idx = dyn_cast<ConstantInt>(Val: Insert->getOperand(i_nocapture: 2));
419 if (!Idx)
420 break;
421 unsigned J = Idx->getZExtValue();
422 V = Insert->getOperand(i_nocapture: 0);
423 if (Frag * VS.NumPacked == J) {
424 CV[Frag] = Insert->getOperand(i_nocapture: 1);
425 return CV[Frag];
426 }
427
428 if (VS.NumPacked == 1 && !CV[J]) {
429 // Only cache the first entry we find for each index we're not actively
430 // searching for. This prevents us from going too far up the chain and
431 // caching incorrect entries.
432 CV[J] = Insert->getOperand(i_nocapture: 1);
433 }
434 }
435 CV[Frag] = Builder.CreateExtractElement(Vec: V, Idx: Frag * VS.NumPacked,
436 Name: V->getName() + ".i" + Twine(Frag));
437 }
438
439 return CV[Frag];
440}
441
442bool ScalarizerLegacyPass::runOnFunction(Function &F) {
443 if (skipFunction(F))
444 return false;
445
446 DominatorTree *DT = &getAnalysis<DominatorTreeWrapperPass>().getDomTree();
447 const TargetTransformInfo *TTI =
448 &getAnalysis<TargetTransformInfoWrapperPass>().getTTI(F);
449 ScalarizerVisitor Impl(DT, TTI, Options);
450 return Impl.visit(F);
451}
452
453FunctionPass *llvm::createScalarizerPass(const ScalarizerPassOptions &Options) {
454 return new ScalarizerLegacyPass(Options);
455}
456
457bool ScalarizerVisitor::visit(Function &F) {
458 assert(Gathered.empty() && Scattered.empty());
459
460 Scalarized = false;
461
462 // To ensure we replace gathered components correctly we need to do an ordered
463 // traversal of the basic blocks in the function.
464 ReversePostOrderTraversal<BasicBlock *> RPOT(&F.getEntryBlock());
465 for (BasicBlock *BB : RPOT) {
466 for (BasicBlock::iterator II = BB->begin(), IE = BB->end(); II != IE;) {
467 Instruction *I = &*II;
468 bool Done = InstVisitor::visit(I);
469 ++II;
470 if (Done && I->getType()->isVoidTy()) {
471 I->eraseFromParent();
472 Scalarized = true;
473 }
474 }
475 }
476 return finish();
477}
478
479// Return a scattered form of V that can be accessed by Point. V must be a
480// vector or a pointer to a vector.
481Scatterer ScalarizerVisitor::scatter(Instruction *Point, Value *V,
482 const VectorSplit &VS) {
483 if (Argument *VArg = dyn_cast<Argument>(Val: V)) {
484 // Put the scattered form of arguments in the entry block,
485 // so that it can be used everywhere.
486 Function *F = VArg->getParent();
487 BasicBlock *BB = &F->getEntryBlock();
488 return Scatterer(BB->begin(), V, VS, &Scattered[{V, VS.SplitTy}]);
489 }
490 if (Instruction *VOp = dyn_cast<Instruction>(Val: V)) {
491 // When scalarizing PHI nodes we might try to examine/rewrite InsertElement
492 // nodes in predecessors. If those predecessors are unreachable from entry,
493 // then the IR in those blocks could have unexpected properties resulting in
494 // infinite loops in Scatterer::operator[]. By simply treating values
495 // originating from instructions in unreachable blocks as undef we do not
496 // need to analyse them further.
497 if (!DT->isReachableFromEntry(A: VOp->getParent()))
498 return Scatterer(Point->getIterator(), PoisonValue::get(T: V->getType()),
499 VS);
500 // Put the scattered form of an instruction directly after the
501 // instruction, skipping over PHI nodes and debug intrinsics.
502 return Scatterer(
503 skipPastPhiNodesAndDbg(Itr: std::next(x: BasicBlock::iterator(VOp))), V, VS,
504 &Scattered[{V, VS.SplitTy}]);
505 }
506 // In the fallback case, just put the scattered before Point and
507 // keep the result local to Point.
508 return Scatterer(Point->getIterator(), V, VS);
509}
510
511// Replace Op with the gathered form of the components in CV. Defer the
512// deletion of Op and creation of the gathered form to the end of the pass,
513// so that we can avoid creating the gathered form if all uses of Op are
514// replaced with uses of CV.
515void ScalarizerVisitor::gather(Instruction *Op, const ValueVector &CV,
516 const VectorSplit &VS) {
517 transferMetadataAndIRFlags(Op, CV);
518
519 // If we already have a scattered form of Op (created from ExtractElements
520 // of Op itself), replace them with the new form.
521 ValueVector &SV = Scattered[{Op, VS.SplitTy}];
522 if (!SV.empty()) {
523 for (unsigned I = 0, E = SV.size(); I != E; ++I) {
524 Value *V = SV[I];
525 if (V == nullptr || SV[I] == CV[I])
526 continue;
527
528 Instruction *Old = cast<Instruction>(Val: V);
529 if (isa<Instruction>(Val: CV[I]))
530 CV[I]->takeName(V: Old);
531 Old->replaceAllUsesWith(V: CV[I]);
532 PotentiallyDeadInstrs.emplace_back(Args&: Old);
533 }
534 }
535 SV = CV;
536 Gathered.push_back(Elt: GatherList::value_type(Op, &SV));
537}
538
539// Replace Op with CV and collect Op has a potentially dead instruction.
540void ScalarizerVisitor::replaceUses(Instruction *Op, Value *CV) {
541 if (CV != Op) {
542 Op->replaceAllUsesWith(V: CV);
543 PotentiallyDeadInstrs.emplace_back(Args&: Op);
544 Scalarized = true;
545 }
546}
547
548// Return true if it is safe to transfer the given metadata tag from
549// vector to scalar instructions.
550bool ScalarizerVisitor::canTransferMetadata(unsigned Tag) {
551 return (Tag == LLVMContext::MD_tbaa
552 || Tag == LLVMContext::MD_fpmath
553 || Tag == LLVMContext::MD_tbaa_struct
554 || Tag == LLVMContext::MD_invariant_load
555 || Tag == LLVMContext::MD_alias_scope
556 || Tag == LLVMContext::MD_noalias
557 || Tag == LLVMContext::MD_mem_parallel_loop_access
558 || Tag == LLVMContext::MD_access_group);
559}
560
561// Transfer metadata from Op to the instructions in CV if it is known
562// to be safe to do so.
563void ScalarizerVisitor::transferMetadataAndIRFlags(Instruction *Op,
564 const ValueVector &CV) {
565 SmallVector<std::pair<unsigned, MDNode *>, 4> MDs;
566 Op->getAllMetadataOtherThanDebugLoc(MDs);
567 for (Value *V : CV) {
568 if (Instruction *New = dyn_cast<Instruction>(Val: V)) {
569 for (const auto &MD : MDs)
570 if (canTransferMetadata(Tag: MD.first))
571 New->setMetadata(KindID: MD.first, Node: MD.second);
572 New->copyIRFlags(V: Op);
573 if (Op->getDebugLoc() && !New->getDebugLoc())
574 New->setDebugLoc(Op->getDebugLoc());
575 }
576 }
577}
578
579// Determine how Ty is split, if at all.
580std::optional<VectorSplit> ScalarizerVisitor::getVectorSplit(Type *Ty) {
581 VectorSplit Split;
582 Split.VecTy = dyn_cast<FixedVectorType>(Val: Ty);
583 if (!Split.VecTy)
584 return {};
585
586 unsigned NumElems = Split.VecTy->getNumElements();
587 Type *ElemTy = Split.VecTy->getElementType();
588
589 if (NumElems == 1 || ElemTy->isPointerTy() ||
590 2 * ElemTy->getScalarSizeInBits() > ScalarizeMinBits) {
591 Split.NumPacked = 1;
592 Split.NumFragments = NumElems;
593 Split.SplitTy = ElemTy;
594 } else {
595 Split.NumPacked = ScalarizeMinBits / ElemTy->getScalarSizeInBits();
596 if (Split.NumPacked >= NumElems)
597 return {};
598
599 Split.NumFragments = divideCeil(Numerator: NumElems, Denominator: Split.NumPacked);
600 Split.SplitTy = FixedVectorType::get(ElementType: ElemTy, NumElts: Split.NumPacked);
601
602 unsigned RemainderElems = NumElems % Split.NumPacked;
603 if (RemainderElems > 1)
604 Split.RemainderTy = FixedVectorType::get(ElementType: ElemTy, NumElts: RemainderElems);
605 else if (RemainderElems == 1)
606 Split.RemainderTy = ElemTy;
607 }
608
609 return Split;
610}
611
612// Try to fill in Layout from Ty, returning true on success. Alignment is
613// the alignment of the vector, or std::nullopt if the ABI default should be
614// used.
615std::optional<VectorLayout>
616ScalarizerVisitor::getVectorLayout(Type *Ty, Align Alignment,
617 const DataLayout &DL) {
618 std::optional<VectorSplit> VS = getVectorSplit(Ty);
619 if (!VS)
620 return {};
621
622 VectorLayout Layout;
623 Layout.VS = *VS;
624 // Check that we're dealing with full-byte fragments.
625 if (!DL.typeSizeEqualsStoreSize(Ty: VS->SplitTy) ||
626 (VS->RemainderTy && !DL.typeSizeEqualsStoreSize(Ty: VS->RemainderTy)))
627 return {};
628 Layout.VecAlign = Alignment;
629 Layout.SplitSize = DL.getTypeStoreSize(Ty: VS->SplitTy);
630 return Layout;
631}
632
633// Scalarize one-operand instruction I, using Split(Builder, X, Name)
634// to create an instruction like I with operand X and name Name.
635template<typename Splitter>
636bool ScalarizerVisitor::splitUnary(Instruction &I, const Splitter &Split) {
637 std::optional<VectorSplit> VS = getVectorSplit(Ty: I.getType());
638 if (!VS)
639 return false;
640
641 std::optional<VectorSplit> OpVS;
642 if (I.getOperand(i: 0)->getType() == I.getType()) {
643 OpVS = VS;
644 } else {
645 OpVS = getVectorSplit(Ty: I.getOperand(i: 0)->getType());
646 if (!OpVS || VS->NumPacked != OpVS->NumPacked)
647 return false;
648 }
649
650 IRBuilder<> Builder(&I);
651 Scatterer Op = scatter(Point: &I, V: I.getOperand(i: 0), VS: *OpVS);
652 assert(Op.size() == VS->NumFragments && "Mismatched unary operation");
653 ValueVector Res;
654 Res.resize(N: VS->NumFragments);
655 for (unsigned Frag = 0; Frag < VS->NumFragments; ++Frag)
656 Res[Frag] = Split(Builder, Op[Frag], I.getName() + ".i" + Twine(Frag));
657 gather(Op: &I, CV: Res, VS: *VS);
658 return true;
659}
660
661// Scalarize two-operand instruction I, using Split(Builder, X, Y, Name)
662// to create an instruction like I with operands X and Y and name Name.
663template<typename Splitter>
664bool ScalarizerVisitor::splitBinary(Instruction &I, const Splitter &Split) {
665 std::optional<VectorSplit> VS = getVectorSplit(Ty: I.getType());
666 if (!VS)
667 return false;
668
669 std::optional<VectorSplit> OpVS;
670 if (I.getOperand(i: 0)->getType() == I.getType()) {
671 OpVS = VS;
672 } else {
673 OpVS = getVectorSplit(Ty: I.getOperand(i: 0)->getType());
674 if (!OpVS || VS->NumPacked != OpVS->NumPacked)
675 return false;
676 }
677
678 IRBuilder<> Builder(&I);
679 Scatterer VOp0 = scatter(Point: &I, V: I.getOperand(i: 0), VS: *OpVS);
680 Scatterer VOp1 = scatter(Point: &I, V: I.getOperand(i: 1), VS: *OpVS);
681 assert(VOp0.size() == VS->NumFragments && "Mismatched binary operation");
682 assert(VOp1.size() == VS->NumFragments && "Mismatched binary operation");
683 ValueVector Res;
684 Res.resize(N: VS->NumFragments);
685 for (unsigned Frag = 0; Frag < VS->NumFragments; ++Frag) {
686 Value *Op0 = VOp0[Frag];
687 Value *Op1 = VOp1[Frag];
688 Res[Frag] = Split(Builder, Op0, Op1, I.getName() + ".i" + Twine(Frag));
689 }
690 gather(Op: &I, CV: Res, VS: *VS);
691 return true;
692}
693
694/// If a call to a vector typed intrinsic function, split into a scalar call per
695/// element if possible for the intrinsic.
696bool ScalarizerVisitor::splitCall(CallInst &CI) {
697 Type *CallType = CI.getType();
698 bool AreAllVectorsOfMatchingSize = isStructOfMatchingFixedVectors(Ty: CallType);
699 std::optional<VectorSplit> VS;
700 if (AreAllVectorsOfMatchingSize)
701 VS = getVectorSplit(Ty: CallType->getContainedType(i: 0));
702 else
703 VS = getVectorSplit(Ty: CallType);
704 if (!VS)
705 return false;
706
707 Function *F = CI.getCalledFunction();
708 if (!F)
709 return false;
710
711 Intrinsic::ID ID = F->getIntrinsicID();
712
713 if (ID == Intrinsic::not_intrinsic || !isTriviallyScalarizable(ID))
714 return false;
715
716 // unsigned NumElems = VT->getNumElements();
717 unsigned NumArgs = CI.arg_size();
718
719 ValueVector ScalarOperands(NumArgs);
720 SmallVector<Scatterer, 8> Scattered(NumArgs);
721 SmallVector<int> OverloadIdx(NumArgs, -1);
722
723 SmallVector<llvm::Type *, 3> Tys;
724 // Add return type if intrinsic is overloaded on it.
725 if (isVectorIntrinsicWithOverloadTypeAtArg(ID, OpdIdx: -1, TTI))
726 Tys.push_back(Elt: VS->SplitTy);
727
728 if (AreAllVectorsOfMatchingSize) {
729 for (unsigned I = 1; I < CallType->getNumContainedTypes(); I++) {
730 std::optional<VectorSplit> CurrVS =
731 getVectorSplit(Ty: cast<FixedVectorType>(Val: CallType->getContainedType(i: I)));
732 // It is possible for VectorSplit.NumPacked >= NumElems. If that happens a
733 // VectorSplit is not returned and we will bailout of handling this call.
734 // The secondary bailout case is if NumPacked does not match. This can
735 // happen if ScalarizeMinBits is not set to the default. This means with
736 // certain ScalarizeMinBits intrinsics like frexp will only scalarize when
737 // the struct elements have the same bitness.
738 if (!CurrVS || CurrVS->NumPacked != VS->NumPacked)
739 return false;
740 if (isVectorIntrinsicWithStructReturnOverloadAtField(ID, RetIdx: I, TTI))
741 Tys.push_back(Elt: CurrVS->SplitTy);
742 }
743 }
744 // Assumes that any vector type has the same number of elements as the return
745 // vector type, which is true for all current intrinsics.
746 for (unsigned I = 0; I != NumArgs; ++I) {
747 Value *OpI = CI.getOperand(i_nocapture: I);
748 if ([[maybe_unused]] auto *OpVecTy =
749 dyn_cast<FixedVectorType>(Val: OpI->getType())) {
750 assert(OpVecTy->getNumElements() == VS->VecTy->getNumElements());
751 std::optional<VectorSplit> OpVS = getVectorSplit(Ty: OpI->getType());
752 if (!OpVS || OpVS->NumPacked != VS->NumPacked) {
753 // The natural split of the operand doesn't match the result. This could
754 // happen if the vector elements are different and the ScalarizeMinBits
755 // option is used.
756 //
757 // We could in principle handle this case as well, at the cost of
758 // complicating the scattering machinery to support multiple scattering
759 // granularities for a single value.
760 return false;
761 }
762
763 Scattered[I] = scatter(Point: &CI, V: OpI, VS: *OpVS);
764 if (isVectorIntrinsicWithOverloadTypeAtArg(ID, OpdIdx: I, TTI)) {
765 OverloadIdx[I] = Tys.size();
766 Tys.push_back(Elt: OpVS->SplitTy);
767 }
768 } else {
769 ScalarOperands[I] = OpI;
770 if (isVectorIntrinsicWithOverloadTypeAtArg(ID, OpdIdx: I, TTI))
771 Tys.push_back(Elt: OpI->getType());
772 }
773 }
774
775 ValueVector Res(VS->NumFragments);
776 ValueVector ScalarCallOps(NumArgs);
777
778 Function *NewIntrin =
779 Intrinsic::getOrInsertDeclaration(M: F->getParent(), id: ID, OverloadTys: Tys);
780 IRBuilder<> Builder(&CI);
781
782 // Perform actual scalarization, taking care to preserve any scalar operands.
783 for (unsigned I = 0; I < VS->NumFragments; ++I) {
784 bool IsRemainder = I == VS->NumFragments - 1 && VS->RemainderTy;
785 ScalarCallOps.clear();
786
787 if (IsRemainder)
788 Tys[0] = VS->RemainderTy;
789
790 for (unsigned J = 0; J != NumArgs; ++J) {
791 if (isVectorIntrinsicWithScalarOpAtArg(ID, ScalarOpdIdx: J, TTI)) {
792 ScalarCallOps.push_back(Elt: ScalarOperands[J]);
793 } else {
794 ScalarCallOps.push_back(Elt: Scattered[J][I]);
795 if (IsRemainder && OverloadIdx[J] >= 0)
796 Tys[OverloadIdx[J]] = Scattered[J][I]->getType();
797 }
798 }
799
800 if (IsRemainder)
801 NewIntrin = Intrinsic::getOrInsertDeclaration(M: F->getParent(), id: ID, OverloadTys: Tys);
802
803 Res[I] = Builder.CreateCall(Callee: NewIntrin, Args: ScalarCallOps,
804 Name: CI.getName() + ".i" + Twine(I));
805 }
806
807 gather(Op: &CI, CV: Res, VS: *VS);
808 return true;
809}
810
811bool ScalarizerVisitor::visitSelectInst(SelectInst &SI) {
812 std::optional<VectorSplit> VS = getVectorSplit(Ty: SI.getType());
813 if (!VS)
814 return false;
815
816 std::optional<VectorSplit> CondVS;
817 if (isa<FixedVectorType>(Val: SI.getCondition()->getType())) {
818 CondVS = getVectorSplit(Ty: SI.getCondition()->getType());
819 if (!CondVS || CondVS->NumPacked != VS->NumPacked) {
820 // This happens when ScalarizeMinBits is used.
821 return false;
822 }
823 }
824
825 IRBuilder<> Builder(&SI);
826 Scatterer VOp1 = scatter(Point: &SI, V: SI.getOperand(i_nocapture: 1), VS: *VS);
827 Scatterer VOp2 = scatter(Point: &SI, V: SI.getOperand(i_nocapture: 2), VS: *VS);
828 assert(VOp1.size() == VS->NumFragments && "Mismatched select");
829 assert(VOp2.size() == VS->NumFragments && "Mismatched select");
830 ValueVector Res;
831 Res.resize(N: VS->NumFragments);
832
833 if (CondVS) {
834 Scatterer VOp0 = scatter(Point: &SI, V: SI.getOperand(i_nocapture: 0), VS: *CondVS);
835 assert(VOp0.size() == CondVS->NumFragments && "Mismatched select");
836 for (unsigned I = 0; I < VS->NumFragments; ++I) {
837 Value *Op0 = VOp0[I];
838 Value *Op1 = VOp1[I];
839 Value *Op2 = VOp2[I];
840 Res[I] = Builder.CreateSelect(C: Op0, True: Op1, False: Op2,
841 Name: SI.getName() + ".i" + Twine(I));
842 }
843 } else {
844 Value *Op0 = SI.getOperand(i_nocapture: 0);
845 for (unsigned I = 0; I < VS->NumFragments; ++I) {
846 Value *Op1 = VOp1[I];
847 Value *Op2 = VOp2[I];
848 Res[I] = Builder.CreateSelect(C: Op0, True: Op1, False: Op2,
849 Name: SI.getName() + ".i" + Twine(I));
850 }
851 }
852 gather(Op: &SI, CV: Res, VS: *VS);
853 return true;
854}
855
856bool ScalarizerVisitor::visitICmpInst(ICmpInst &ICI) {
857 return splitBinary(I&: ICI, Split: ICmpSplitter(ICI));
858}
859
860bool ScalarizerVisitor::visitFCmpInst(FCmpInst &FCI) {
861 return splitBinary(I&: FCI, Split: FCmpSplitter(FCI));
862}
863
864bool ScalarizerVisitor::visitUnaryOperator(UnaryOperator &UO) {
865 return splitUnary(I&: UO, Split: UnarySplitter(UO));
866}
867
868bool ScalarizerVisitor::visitBinaryOperator(BinaryOperator &BO) {
869 return splitBinary(I&: BO, Split: BinarySplitter(BO));
870}
871
872bool ScalarizerVisitor::visitGetElementPtrInst(GetElementPtrInst &GEPI) {
873 std::optional<VectorSplit> VS = getVectorSplit(Ty: GEPI.getType());
874 if (!VS)
875 return false;
876
877 IRBuilder<> Builder(&GEPI);
878 unsigned NumIndices = GEPI.getNumIndices();
879
880 // The base pointer and indices might be scalar even if it's a vector GEP.
881 SmallVector<Value *, 8> ScalarOps{1 + NumIndices};
882 SmallVector<Scatterer, 8> ScatterOps{1 + NumIndices};
883
884 for (unsigned I = 0; I < 1 + NumIndices; ++I) {
885 if (auto *VecTy =
886 dyn_cast<FixedVectorType>(Val: GEPI.getOperand(i_nocapture: I)->getType())) {
887 std::optional<VectorSplit> OpVS = getVectorSplit(Ty: VecTy);
888 if (!OpVS || OpVS->NumPacked != VS->NumPacked) {
889 // This can happen when ScalarizeMinBits is used.
890 return false;
891 }
892 ScatterOps[I] = scatter(Point: &GEPI, V: GEPI.getOperand(i_nocapture: I), VS: *OpVS);
893 } else {
894 ScalarOps[I] = GEPI.getOperand(i_nocapture: I);
895 }
896 }
897
898 ValueVector Res;
899 Res.resize(N: VS->NumFragments);
900 for (unsigned I = 0; I < VS->NumFragments; ++I) {
901 SmallVector<Value *, 8> SplitOps;
902 SplitOps.resize(N: 1 + NumIndices);
903 for (unsigned J = 0; J < 1 + NumIndices; ++J) {
904 if (ScalarOps[J])
905 SplitOps[J] = ScalarOps[J];
906 else
907 SplitOps[J] = ScatterOps[J][I];
908 }
909 Res[I] = Builder.CreateGEP(Ty: GEPI.getSourceElementType(), Ptr: SplitOps[0],
910 IdxList: ArrayRef(SplitOps).drop_front(),
911 Name: GEPI.getName() + ".i" + Twine(I));
912 if (GEPI.isInBounds())
913 if (GetElementPtrInst *NewGEPI = dyn_cast<GetElementPtrInst>(Val: Res[I]))
914 NewGEPI->setIsInBounds();
915 }
916 gather(Op: &GEPI, CV: Res, VS: *VS);
917 return true;
918}
919
920bool ScalarizerVisitor::visitCastInst(CastInst &CI) {
921 std::optional<VectorSplit> DestVS = getVectorSplit(Ty: CI.getDestTy());
922 if (!DestVS)
923 return false;
924
925 std::optional<VectorSplit> SrcVS = getVectorSplit(Ty: CI.getSrcTy());
926 if (!SrcVS || SrcVS->NumPacked != DestVS->NumPacked)
927 return false;
928
929 IRBuilder<> Builder(&CI);
930 Scatterer Op0 = scatter(Point: &CI, V: CI.getOperand(i_nocapture: 0), VS: *SrcVS);
931 assert(Op0.size() == SrcVS->NumFragments && "Mismatched cast");
932 ValueVector Res;
933 Res.resize(N: DestVS->NumFragments);
934 for (unsigned I = 0; I < DestVS->NumFragments; ++I)
935 Res[I] =
936 Builder.CreateCast(Op: CI.getOpcode(), V: Op0[I], DestTy: DestVS->getFragmentType(I),
937 Name: CI.getName() + ".i" + Twine(I));
938 gather(Op: &CI, CV: Res, VS: *DestVS);
939 return true;
940}
941
942bool ScalarizerVisitor::visitBitCastInst(BitCastInst &BCI) {
943 std::optional<VectorSplit> DstVS = getVectorSplit(Ty: BCI.getDestTy());
944 std::optional<VectorSplit> SrcVS = getVectorSplit(Ty: BCI.getSrcTy());
945
946 if (DstVS && !SrcVS && BCI.getSrcTy()->isIntegerTy() && !DstVS->RemainderTy &&
947 DstVS->NumPacked == 1 && DstVS->SplitTy->isIntegerTy()) {
948 IRBuilder<> Builder(&BCI);
949 Builder.SetCurrentDebugLocation(BCI.getDebugLoc());
950 ValueVector Res(DstVS->NumFragments);
951 unsigned FragmentBits = DstVS->SplitTy->getPrimitiveSizeInBits();
952 bool IsBigEndian = BCI.getDataLayout().isBigEndian();
953 for (unsigned I = 0; I < DstVS->NumFragments; ++I) {
954 unsigned FragmentIndex = IsBigEndian ? DstVS->NumFragments - I - 1 : I;
955 Value *Fragment = BCI.getOperand(i_nocapture: 0);
956 if (FragmentIndex)
957 Fragment = Builder.CreateLShr(LHS: Fragment, RHS: FragmentIndex * FragmentBits);
958 Res[I] = Builder.CreateTruncOrBitCast(V: Fragment, DestTy: DstVS->getFragmentType(I),
959 Name: BCI.getName() + ".i" + Twine(I));
960 }
961 gather(Op: &BCI, CV: Res, VS: *DstVS);
962 return true;
963 }
964
965 if (!DstVS && SrcVS && BCI.getDestTy()->isIntegerTy() &&
966 !SrcVS->RemainderTy && SrcVS->NumPacked == 1 &&
967 SrcVS->SplitTy->isIntegerTy()) {
968 IRBuilder<> Builder(&BCI);
969 Builder.SetCurrentDebugLocation(BCI.getDebugLoc());
970 Scatterer Op0 = scatter(Point: &BCI, V: BCI.getOperand(i_nocapture: 0), VS: *SrcVS);
971 Value *Result = nullptr;
972 unsigned FragmentBits = SrcVS->SplitTy->getPrimitiveSizeInBits();
973 bool IsBigEndian = BCI.getDataLayout().isBigEndian();
974 for (unsigned I = 0; I < SrcVS->NumFragments; ++I) {
975 unsigned FragmentIndex = IsBigEndian ? SrcVS->NumFragments - I - 1 : I;
976 Value *Fragment = Builder.CreateZExtOrTrunc(V: Op0[I], DestTy: BCI.getDestTy());
977 if (FragmentIndex)
978 Fragment = Builder.CreateShl(LHS: Fragment, RHS: FragmentIndex * FragmentBits);
979 Result = Result ? Builder.CreateOr(LHS: Result, RHS: Fragment) : Fragment;
980 }
981 replaceUses(Op: &BCI, CV: Result);
982 return true;
983 }
984
985 if (!DstVS || !SrcVS || DstVS->RemainderTy || SrcVS->RemainderTy)
986 return false;
987
988 const bool isPointerTy = DstVS->VecTy->getElementType()->isPointerTy();
989
990 // Vectors of pointers are always fully scalarized.
991 assert(!isPointerTy || (DstVS->NumPacked == 1 && SrcVS->NumPacked == 1));
992
993 IRBuilder<> Builder(&BCI);
994 Scatterer Op0 = scatter(Point: &BCI, V: BCI.getOperand(i_nocapture: 0), VS: *SrcVS);
995 ValueVector Res;
996 Res.resize(N: DstVS->NumFragments);
997
998 unsigned DstSplitBits = DstVS->SplitTy->getPrimitiveSizeInBits();
999 unsigned SrcSplitBits = SrcVS->SplitTy->getPrimitiveSizeInBits();
1000
1001 if (isPointerTy || DstSplitBits == SrcSplitBits) {
1002 assert(DstVS->NumFragments == SrcVS->NumFragments);
1003 for (unsigned I = 0; I < DstVS->NumFragments; ++I) {
1004 Res[I] = Builder.CreateBitCast(V: Op0[I], DestTy: DstVS->getFragmentType(I),
1005 Name: BCI.getName() + ".i" + Twine(I));
1006 }
1007 } else if (SrcSplitBits % DstSplitBits == 0) {
1008 // Convert each source fragment to the same-sized destination vector and
1009 // then scatter the result to the destination.
1010 VectorSplit MidVS;
1011 MidVS.NumPacked = DstVS->NumPacked;
1012 MidVS.NumFragments = SrcSplitBits / DstSplitBits;
1013 MidVS.VecTy = FixedVectorType::get(ElementType: DstVS->VecTy->getElementType(),
1014 NumElts: MidVS.NumPacked * MidVS.NumFragments);
1015 MidVS.SplitTy = DstVS->SplitTy;
1016
1017 unsigned ResI = 0;
1018 for (unsigned I = 0; I < SrcVS->NumFragments; ++I) {
1019 Value *V = Op0[I];
1020
1021 // Look through any existing bitcasts before converting to <N x t2>.
1022 // In the best case, the resulting conversion might be a no-op.
1023 Instruction *VI;
1024 while ((VI = dyn_cast<Instruction>(Val: V)) &&
1025 VI->getOpcode() == Instruction::BitCast)
1026 V = VI->getOperand(i: 0);
1027
1028 V = Builder.CreateBitCast(V, DestTy: MidVS.VecTy, Name: V->getName() + ".cast");
1029
1030 Scatterer Mid = scatter(Point: &BCI, V, VS: MidVS);
1031 for (unsigned J = 0; J < MidVS.NumFragments; ++J)
1032 Res[ResI++] = Mid[J];
1033 }
1034 } else if (DstSplitBits % SrcSplitBits == 0) {
1035 // Gather enough source fragments to make up a destination fragment and
1036 // then convert to the destination type.
1037 VectorSplit MidVS;
1038 MidVS.NumFragments = DstSplitBits / SrcSplitBits;
1039 MidVS.NumPacked = SrcVS->NumPacked;
1040 MidVS.VecTy = FixedVectorType::get(ElementType: SrcVS->VecTy->getElementType(),
1041 NumElts: MidVS.NumPacked * MidVS.NumFragments);
1042 MidVS.SplitTy = SrcVS->SplitTy;
1043
1044 unsigned SrcI = 0;
1045 SmallVector<Value *, 8> ConcatOps;
1046 ConcatOps.resize(N: MidVS.NumFragments);
1047 for (unsigned I = 0; I < DstVS->NumFragments; ++I) {
1048 for (unsigned J = 0; J < MidVS.NumFragments; ++J)
1049 ConcatOps[J] = Op0[SrcI++];
1050 Value *V = concatenate(Builder, Fragments: ConcatOps, VS: MidVS,
1051 Name: BCI.getName() + ".i" + Twine(I));
1052 Res[I] = Builder.CreateBitCast(V, DestTy: DstVS->getFragmentType(I),
1053 Name: BCI.getName() + ".i" + Twine(I));
1054 }
1055 } else {
1056 return false;
1057 }
1058
1059 gather(Op: &BCI, CV: Res, VS: *DstVS);
1060 return true;
1061}
1062
1063bool ScalarizerVisitor::visitInsertElementInst(InsertElementInst &IEI) {
1064 std::optional<VectorSplit> VS = getVectorSplit(Ty: IEI.getType());
1065 if (!VS)
1066 return false;
1067
1068 IRBuilder<> Builder(&IEI);
1069 Scatterer Op0 = scatter(Point: &IEI, V: IEI.getOperand(i_nocapture: 0), VS: *VS);
1070 Value *NewElt = IEI.getOperand(i_nocapture: 1);
1071 Value *InsIdx = IEI.getOperand(i_nocapture: 2);
1072
1073 ValueVector Res;
1074 Res.resize(N: VS->NumFragments);
1075
1076 if (auto *CI = dyn_cast<ConstantInt>(Val: InsIdx)) {
1077 unsigned Idx = CI->getZExtValue();
1078 unsigned Fragment = Idx / VS->NumPacked;
1079 for (unsigned I = 0; I < VS->NumFragments; ++I) {
1080 if (I == Fragment) {
1081 bool IsPacked = VS->NumPacked > 1;
1082 if (Fragment == VS->NumFragments - 1 && VS->RemainderTy &&
1083 !VS->RemainderTy->isVectorTy())
1084 IsPacked = false;
1085 if (IsPacked) {
1086 Res[I] =
1087 Builder.CreateInsertElement(Vec: Op0[I], NewElt, Idx: Idx % VS->NumPacked);
1088 } else {
1089 Res[I] = NewElt;
1090 }
1091 } else {
1092 Res[I] = Op0[I];
1093 }
1094 }
1095 } else {
1096 // Never split a variable insertelement that isn't fully scalarized.
1097 if (!ScalarizeVariableInsertExtract || VS->NumPacked > 1)
1098 return false;
1099
1100 for (unsigned I = 0; I < VS->NumFragments; ++I) {
1101 Value *ShouldReplace =
1102 Builder.CreateICmpEQ(LHS: InsIdx, RHS: ConstantInt::get(Ty: InsIdx->getType(), V: I),
1103 Name: InsIdx->getName() + ".is." + Twine(I));
1104 Value *OldElt = Op0[I];
1105 Res[I] = Builder.CreateSelect(C: ShouldReplace, True: NewElt, False: OldElt,
1106 Name: IEI.getName() + ".i" + Twine(I));
1107 }
1108 }
1109
1110 gather(Op: &IEI, CV: Res, VS: *VS);
1111 return true;
1112}
1113
1114bool ScalarizerVisitor::visitExtractValueInst(ExtractValueInst &EVI) {
1115 Value *Op = EVI.getOperand(i_nocapture: 0);
1116 Type *OpTy = Op->getType();
1117 ValueVector Res;
1118 if (!isStructOfMatchingFixedVectors(Ty: OpTy))
1119 return false;
1120 if (CallInst *CI = dyn_cast<CallInst>(Val: Op)) {
1121 Function *F = CI->getCalledFunction();
1122 if (!F)
1123 return false;
1124 Intrinsic::ID ID = F->getIntrinsicID();
1125 if (ID == Intrinsic::not_intrinsic || !isTriviallyScalarizable(ID))
1126 return false;
1127 // Note: Fall through means Operand is a`CallInst` and it is defined in
1128 // `isTriviallyScalarizable`.
1129 } else
1130 return false;
1131 Type *VecType = cast<FixedVectorType>(Val: OpTy->getContainedType(i: 0));
1132 std::optional<VectorSplit> VS = getVectorSplit(Ty: VecType);
1133 if (!VS)
1134 return false;
1135 for (unsigned I = 1; I < OpTy->getNumContainedTypes(); I++) {
1136 std::optional<VectorSplit> CurrVS =
1137 getVectorSplit(Ty: cast<FixedVectorType>(Val: OpTy->getContainedType(i: I)));
1138 // It is possible for VectorSplit.NumPacked >= NumElems. If that happens a
1139 // VectorSplit is not returned and we will bailout of handling this call.
1140 // The secondary bailout case is if NumPacked does not match. This can
1141 // happen if ScalarizeMinBits is not set to the default. This means with
1142 // certain ScalarizeMinBits intrinsics like frexp will only scalarize when
1143 // the struct elements have the same bitness.
1144 if (!CurrVS || CurrVS->NumPacked != VS->NumPacked)
1145 return false;
1146 }
1147 IRBuilder<> Builder(&EVI);
1148 Scatterer Op0 = scatter(Point: &EVI, V: Op, VS: *VS);
1149 assert(!EVI.getIndices().empty() && "Make sure an index exists");
1150 // Note for our use case we only care about the top level index.
1151 unsigned Index = EVI.getIndices()[0];
1152 for (unsigned OpIdx = 0; OpIdx < Op0.size(); ++OpIdx) {
1153 Value *ResElem = Builder.CreateExtractValue(
1154 Agg: Op0[OpIdx], Idxs: Index, Name: EVI.getName() + ".elem" + Twine(Index));
1155 Res.push_back(Elt: ResElem);
1156 }
1157
1158 Type *ActualVecType = cast<FixedVectorType>(Val: OpTy->getContainedType(i: Index));
1159 std::optional<VectorSplit> AVS = getVectorSplit(Ty: ActualVecType);
1160 gather(Op: &EVI, CV: Res, VS: *AVS);
1161 return true;
1162}
1163
1164bool ScalarizerVisitor::visitExtractElementInst(ExtractElementInst &EEI) {
1165 std::optional<VectorSplit> VS = getVectorSplit(Ty: EEI.getOperand(i_nocapture: 0)->getType());
1166 if (!VS)
1167 return false;
1168
1169 IRBuilder<> Builder(&EEI);
1170 Scatterer Op0 = scatter(Point: &EEI, V: EEI.getOperand(i_nocapture: 0), VS: *VS);
1171 Value *ExtIdx = EEI.getOperand(i_nocapture: 1);
1172
1173 if (auto *CI = dyn_cast<ConstantInt>(Val: ExtIdx)) {
1174 unsigned Idx = CI->getZExtValue();
1175 if (Idx >= VS->VecTy->getNumElements())
1176 return false;
1177 unsigned Fragment = Idx / VS->NumPacked;
1178 Value *Res = Op0[Fragment];
1179 bool IsPacked = VS->NumPacked > 1;
1180 if (Fragment == VS->NumFragments - 1 && VS->RemainderTy &&
1181 !VS->RemainderTy->isVectorTy())
1182 IsPacked = false;
1183 if (IsPacked)
1184 Res = Builder.CreateExtractElement(Vec: Res, Idx: Idx % VS->NumPacked);
1185 replaceUses(Op: &EEI, CV: Res);
1186 return true;
1187 }
1188
1189 // Never split a variable extractelement that isn't fully scalarized.
1190 if (!ScalarizeVariableInsertExtract || VS->NumPacked > 1)
1191 return false;
1192
1193 Value *Res = PoisonValue::get(T: VS->VecTy->getElementType());
1194 for (unsigned I = 0; I < VS->NumFragments; ++I) {
1195 Value *ShouldExtract =
1196 Builder.CreateICmpEQ(LHS: ExtIdx, RHS: ConstantInt::get(Ty: ExtIdx->getType(), V: I),
1197 Name: ExtIdx->getName() + ".is." + Twine(I));
1198 Value *Elt = Op0[I];
1199 Res = Builder.CreateSelect(C: ShouldExtract, True: Elt, False: Res,
1200 Name: EEI.getName() + ".upto" + Twine(I));
1201 }
1202 replaceUses(Op: &EEI, CV: Res);
1203 return true;
1204}
1205
1206bool ScalarizerVisitor::visitShuffleVectorInst(ShuffleVectorInst &SVI) {
1207 std::optional<VectorSplit> VS = getVectorSplit(Ty: SVI.getType());
1208 std::optional<VectorSplit> VSOp =
1209 getVectorSplit(Ty: SVI.getOperand(i_nocapture: 0)->getType());
1210 if (!VS || !VSOp || VS->NumPacked > 1 || VSOp->NumPacked > 1)
1211 return false;
1212
1213 Scatterer Op0 = scatter(Point: &SVI, V: SVI.getOperand(i_nocapture: 0), VS: *VSOp);
1214 Scatterer Op1 = scatter(Point: &SVI, V: SVI.getOperand(i_nocapture: 1), VS: *VSOp);
1215 ValueVector Res;
1216 Res.resize(N: VS->NumFragments);
1217
1218 for (unsigned I = 0; I < VS->NumFragments; ++I) {
1219 int Selector = SVI.getMaskValue(Elt: I);
1220 if (Selector < 0)
1221 Res[I] = PoisonValue::get(T: VS->VecTy->getElementType());
1222 else if (unsigned(Selector) < Op0.size())
1223 Res[I] = Op0[Selector];
1224 else
1225 Res[I] = Op1[Selector - Op0.size()];
1226 }
1227 gather(Op: &SVI, CV: Res, VS: *VS);
1228 return true;
1229}
1230
1231bool ScalarizerVisitor::visitPHINode(PHINode &PHI) {
1232 std::optional<VectorSplit> VS = getVectorSplit(Ty: PHI.getType());
1233 if (!VS)
1234 return false;
1235
1236 IRBuilder<> Builder(&PHI);
1237 ValueVector Res;
1238 Res.resize(N: VS->NumFragments);
1239
1240 unsigned NumOps = PHI.getNumOperands();
1241 for (unsigned I = 0; I < VS->NumFragments; ++I) {
1242 Res[I] = Builder.CreatePHI(Ty: VS->getFragmentType(I), NumReservedValues: NumOps,
1243 Name: PHI.getName() + ".i" + Twine(I));
1244 }
1245
1246 for (unsigned I = 0; I < NumOps; ++I) {
1247 Scatterer Op = scatter(Point: &PHI, V: PHI.getIncomingValue(i: I), VS: *VS);
1248 BasicBlock *IncomingBlock = PHI.getIncomingBlock(i: I);
1249 for (unsigned J = 0; J < VS->NumFragments; ++J)
1250 cast<PHINode>(Val: Res[J])->addIncoming(V: Op[J], BB: IncomingBlock);
1251 }
1252 gather(Op: &PHI, CV: Res, VS: *VS);
1253 return true;
1254}
1255
1256bool ScalarizerVisitor::visitLoadInst(LoadInst &LI) {
1257 if (!ScalarizeLoadStore)
1258 return false;
1259 if (!LI.isSimple())
1260 return false;
1261
1262 std::optional<VectorLayout> Layout = getVectorLayout(
1263 Ty: LI.getType(), Alignment: LI.getAlign(), DL: LI.getDataLayout());
1264 if (!Layout)
1265 return false;
1266
1267 IRBuilder<> Builder(&LI);
1268 Scatterer Ptr = scatter(Point: &LI, V: LI.getPointerOperand(), VS: Layout->VS);
1269 ValueVector Res;
1270 Res.resize(N: Layout->VS.NumFragments);
1271
1272 for (unsigned I = 0; I < Layout->VS.NumFragments; ++I) {
1273 Res[I] = Builder.CreateAlignedLoad(Ty: Layout->VS.getFragmentType(I), Ptr: Ptr[I],
1274 Align: Align(Layout->getFragmentAlign(Frag: I)),
1275 Name: LI.getName() + ".i" + Twine(I));
1276 }
1277 gather(Op: &LI, CV: Res, VS: Layout->VS);
1278 return true;
1279}
1280
1281bool ScalarizerVisitor::visitStoreInst(StoreInst &SI) {
1282 if (!ScalarizeLoadStore)
1283 return false;
1284 if (!SI.isSimple())
1285 return false;
1286
1287 Value *FullValue = SI.getValueOperand();
1288 std::optional<VectorLayout> Layout = getVectorLayout(
1289 Ty: FullValue->getType(), Alignment: SI.getAlign(), DL: SI.getDataLayout());
1290 if (!Layout)
1291 return false;
1292
1293 IRBuilder<> Builder(&SI);
1294 Scatterer VPtr = scatter(Point: &SI, V: SI.getPointerOperand(), VS: Layout->VS);
1295 Scatterer VVal = scatter(Point: &SI, V: FullValue, VS: Layout->VS);
1296
1297 ValueVector Stores;
1298 Stores.resize(N: Layout->VS.NumFragments);
1299 for (unsigned I = 0; I < Layout->VS.NumFragments; ++I) {
1300 Value *Val = VVal[I];
1301 Value *Ptr = VPtr[I];
1302 Stores[I] =
1303 Builder.CreateAlignedStore(Val, Ptr, Align: Layout->getFragmentAlign(Frag: I));
1304 }
1305 transferMetadataAndIRFlags(Op: &SI, CV: Stores);
1306 return true;
1307}
1308
1309bool ScalarizerVisitor::visitCallInst(CallInst &CI) {
1310 return splitCall(CI);
1311}
1312
1313bool ScalarizerVisitor::visitFreezeInst(FreezeInst &FI) {
1314 return splitUnary(I&: FI, Split: [](IRBuilder<> &Builder, Value *Op, const Twine &Name) {
1315 return Builder.CreateFreeze(V: Op, Name);
1316 });
1317}
1318
1319// Delete the instructions that we scalarized. If a full vector result
1320// is still needed, recreate it using InsertElements.
1321bool ScalarizerVisitor::finish() {
1322 // The presence of data in Gathered or Scattered indicates changes
1323 // made to the Function.
1324 if (Gathered.empty() && Scattered.empty() && !Scalarized)
1325 return false;
1326 for (const auto &GMI : Gathered) {
1327 Instruction *Op = GMI.first;
1328 ValueVector &CV = *GMI.second;
1329 if (!Op->use_empty()) {
1330 // The value is still needed, so recreate it using a series of
1331 // insertelements and/or shufflevectors.
1332 Value *Res;
1333 if (auto *Ty = dyn_cast<FixedVectorType>(Val: Op->getType())) {
1334 BasicBlock *BB = Op->getParent();
1335 IRBuilder<> Builder(Op);
1336 if (isa<PHINode>(Val: Op))
1337 Builder.SetInsertPoint(BB->getFirstInsertionPt());
1338
1339 VectorSplit VS = *getVectorSplit(Ty);
1340 assert(VS.NumFragments == CV.size());
1341
1342 Res = concatenate(Builder, Fragments: CV, VS, Name: Op->getName());
1343
1344 Res->takeName(V: Op);
1345 } else if (auto *Ty = dyn_cast<StructType>(Val: Op->getType())) {
1346 BasicBlock *BB = Op->getParent();
1347 IRBuilder<> Builder(Op);
1348 if (isa<PHINode>(Val: Op))
1349 Builder.SetInsertPoint(BB->getFirstInsertionPt());
1350
1351 // Iterate over each element in the struct
1352 unsigned NumOfStructElements = Ty->getNumElements();
1353 SmallVector<ValueVector, 4> ElemCV(NumOfStructElements);
1354 for (unsigned I = 0; I < NumOfStructElements; ++I) {
1355 for (auto *CVelem : CV) {
1356 Value *Elem = Builder.CreateExtractValue(
1357 Agg: CVelem, Idxs: I, Name: Op->getName() + ".elem" + Twine(I));
1358 ElemCV[I].push_back(Elt: Elem);
1359 }
1360 }
1361 Res = PoisonValue::get(T: Ty);
1362 for (unsigned I = 0; I < NumOfStructElements; ++I) {
1363 Type *ElemTy = Ty->getElementType(N: I);
1364 assert(isa<FixedVectorType>(ElemTy) &&
1365 "Only Structs of all FixedVectorType supported");
1366 VectorSplit VS = *getVectorSplit(Ty: ElemTy);
1367 assert(VS.NumFragments == CV.size());
1368
1369 Value *ConcatenatedVector =
1370 concatenate(Builder, Fragments: ElemCV[I], VS, Name: Op->getName());
1371 Res = Builder.CreateInsertValue(Agg: Res, Val: ConcatenatedVector, Idxs: I,
1372 Name: Op->getName() + ".insert");
1373 }
1374 } else {
1375 assert(CV.size() == 1 && Op->getType() == CV[0]->getType());
1376 Res = CV[0];
1377 if (Op == Res)
1378 continue;
1379 }
1380 Op->replaceAllUsesWith(V: Res);
1381 }
1382 PotentiallyDeadInstrs.emplace_back(Args&: Op);
1383 }
1384 Gathered.clear();
1385 Scattered.clear();
1386 Scalarized = false;
1387
1388 RecursivelyDeleteTriviallyDeadInstructionsPermissive(DeadInsts&: PotentiallyDeadInstrs);
1389
1390 return true;
1391}
1392
1393PreservedAnalyses ScalarizerPass::run(Function &F, FunctionAnalysisManager &AM) {
1394 DominatorTree *DT = &AM.getResult<DominatorTreeAnalysis>(IR&: F);
1395 const TargetTransformInfo *TTI = &AM.getResult<TargetIRAnalysis>(IR&: F);
1396 ScalarizerVisitor Impl(DT, TTI, Options);
1397 bool Changed = Impl.visit(F);
1398 PreservedAnalyses PA;
1399 PA.preserve<DominatorTreeAnalysis>();
1400 return Changed ? PA : PreservedAnalyses::all();
1401}
1402