1 //===- MVEGatherScatterLowering.cpp - Gather/Scatter lowering -------------===//
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 custom lowers llvm.gather and llvm.scatter instructions to
10 /// arm.mve.gather and arm.mve.scatter intrinsics, optimising the code to
11 /// produce a better final result as we go.
12 //
13 //===----------------------------------------------------------------------===//
14 
15 #include "ARM.h"
16 #include "ARMBaseInstrInfo.h"
17 #include "ARMSubtarget.h"
18 #include "llvm/Analysis/LoopInfo.h"
19 #include "llvm/Analysis/TargetTransformInfo.h"
20 #include "llvm/CodeGen/TargetLowering.h"
21 #include "llvm/CodeGen/TargetPassConfig.h"
22 #include "llvm/CodeGen/TargetSubtargetInfo.h"
23 #include "llvm/InitializePasses.h"
24 #include "llvm/IR/BasicBlock.h"
25 #include "llvm/IR/Constant.h"
26 #include "llvm/IR/Constants.h"
27 #include "llvm/IR/DerivedTypes.h"
28 #include "llvm/IR/Function.h"
29 #include "llvm/IR/InstrTypes.h"
30 #include "llvm/IR/Instruction.h"
31 #include "llvm/IR/Instructions.h"
32 #include "llvm/IR/IntrinsicInst.h"
33 #include "llvm/IR/Intrinsics.h"
34 #include "llvm/IR/IntrinsicsARM.h"
35 #include "llvm/IR/IRBuilder.h"
36 #include "llvm/IR/PatternMatch.h"
37 #include "llvm/IR/Type.h"
38 #include "llvm/IR/Value.h"
39 #include "llvm/Pass.h"
40 #include "llvm/Support/Casting.h"
41 #include "llvm/Transforms/Utils/Local.h"
42 #include <algorithm>
43 #include <cassert>
44 
45 using namespace llvm;
46 
47 #define DEBUG_TYPE "mve-gather-scatter-lowering"
48 
49 cl::opt<bool> EnableMaskedGatherScatters(
50     "enable-arm-maskedgatscat", cl::Hidden, cl::init(false),
51     cl::desc("Enable the generation of masked gathers and scatters"));
52 
53 namespace {
54 
55 class MVEGatherScatterLowering : public FunctionPass {
56 public:
57   static char ID; // Pass identification, replacement for typeid
58 
59   explicit MVEGatherScatterLowering() : FunctionPass(ID) {
60     initializeMVEGatherScatterLoweringPass(*PassRegistry::getPassRegistry());
61   }
62 
63   bool runOnFunction(Function &F) override;
64 
65   StringRef getPassName() const override {
66     return "MVE gather/scatter lowering";
67   }
68 
69   void getAnalysisUsage(AnalysisUsage &AU) const override {
70     AU.setPreservesCFG();
71     AU.addRequired<TargetPassConfig>();
72     AU.addRequired<LoopInfoWrapperPass>();
73     FunctionPass::getAnalysisUsage(AU);
74   }
75 
76 private:
77   LoopInfo *LI = nullptr;
78 
79   // Check this is a valid gather with correct alignment
80   bool isLegalTypeAndAlignment(unsigned NumElements, unsigned ElemSize,
81                                unsigned Alignment);
82   // Check whether Ptr is hidden behind a bitcast and look through it
83   void lookThroughBitcast(Value *&Ptr);
84   // Check for a getelementptr and deduce base and offsets from it, on success
85   // returning the base directly and the offsets indirectly using the Offsets
86   // argument
87   Value *checkGEP(Value *&Offsets, Type *Ty, GetElementPtrInst *GEP,
88                   IRBuilder<> &Builder);
89   // Compute the scale of this gather/scatter instruction
90   int computeScale(unsigned GEPElemSize, unsigned MemoryElemSize);
91   // If the value is a constant, or derived from constants via additions
92   // and multilications, return its numeric value
93   Optional<int64_t> getIfConst(const Value *V);
94   // If Inst is an add instruction, check whether one summand is a
95   // constant. If so, scale this constant and return it together with
96   // the other summand.
97   std::pair<Value *, int64_t> getVarAndConst(Value *Inst, int TypeScale);
98 
99   Value *lowerGather(IntrinsicInst *I);
100   // Create a gather from a base + vector of offsets
101   Value *tryCreateMaskedGatherOffset(IntrinsicInst *I, Value *Ptr,
102                                      Instruction *&Root, IRBuilder<> &Builder);
103   // Create a gather from a vector of pointers
104   Value *tryCreateMaskedGatherBase(IntrinsicInst *I, Value *Ptr,
105                                    IRBuilder<> &Builder,
106                                    unsigned Increment = 0);
107   // Create a gather from a vector of pointers
108   Value *tryCreateMaskedGatherBaseWB(IntrinsicInst *I, Value *Ptr,
109                                      IRBuilder<> &Builder,
110                                      unsigned Increment = 0);
111   // QI gathers can increment their offsets on their own if the increment is
112   // a constant value (digit)
113   Value *tryCreateIncrementingGather(IntrinsicInst *I, Value *BasePtr,
114                                      Value *Ptr, GetElementPtrInst *GEP,
115                                      IRBuilder<> &Builder);
116   // QI gathers can increment their offsets on their own if the increment is
117   // a constant value (digit) - this creates a writeback QI gather
118   Value *tryCreateIncrementingWBGather(IntrinsicInst *I, Value *BasePtr,
119                                        Value *Ptr, unsigned TypeScale,
120                                        IRBuilder<> &Builder);
121 
122   Value *lowerScatter(IntrinsicInst *I);
123   // Create a scatter to a base + vector of offsets
124   Value *tryCreateMaskedScatterOffset(IntrinsicInst *I, Value *Offsets,
125                                       IRBuilder<> &Builder);
126   // Create a scatter to a vector of pointers
127   Value *tryCreateMaskedScatterBase(IntrinsicInst *I, Value *Ptr,
128                                     IRBuilder<> &Builder);
129 
130   // Check whether these offsets could be moved out of the loop they're in
131   bool optimiseOffsets(Value *Offsets, BasicBlock *BB, LoopInfo *LI);
132   // Pushes the given add out of the loop
133   void pushOutAdd(PHINode *&Phi, Value *OffsSecondOperand, unsigned StartIndex);
134   // Pushes the given mul out of the loop
135   void pushOutMul(PHINode *&Phi, Value *IncrementPerRound,
136                   Value *OffsSecondOperand, unsigned LoopIncrement,
137                   IRBuilder<> &Builder);
138 };
139 
140 } // end anonymous namespace
141 
142 char MVEGatherScatterLowering::ID = 0;
143 
144 INITIALIZE_PASS(MVEGatherScatterLowering, DEBUG_TYPE,
145                 "MVE gather/scattering lowering pass", false, false)
146 
147 Pass *llvm::createMVEGatherScatterLoweringPass() {
148   return new MVEGatherScatterLowering();
149 }
150 
151 bool MVEGatherScatterLowering::isLegalTypeAndAlignment(unsigned NumElements,
152                                                        unsigned ElemSize,
153                                                        unsigned Alignment) {
154   if (((NumElements == 4 &&
155         (ElemSize == 32 || ElemSize == 16 || ElemSize == 8)) ||
156        (NumElements == 8 && (ElemSize == 16 || ElemSize == 8)) ||
157        (NumElements == 16 && ElemSize == 8)) &&
158       ElemSize / 8 <= Alignment)
159     return true;
160   LLVM_DEBUG(dbgs() << "masked gathers/scatters: instruction does not have "
161                     << "valid alignment or vector type \n");
162   return false;
163 }
164 
165 Value *MVEGatherScatterLowering::checkGEP(Value *&Offsets, Type *Ty,
166                                           GetElementPtrInst *GEP,
167                                           IRBuilder<> &Builder) {
168   if (!GEP) {
169     LLVM_DEBUG(
170         dbgs() << "masked gathers/scatters: no getelementpointer found\n");
171     return nullptr;
172   }
173   LLVM_DEBUG(dbgs() << "masked gathers/scatters: getelementpointer found."
174                     << " Looking at intrinsic for base + vector of offsets\n");
175   Value *GEPPtr = GEP->getPointerOperand();
176   if (GEPPtr->getType()->isVectorTy()) {
177     return nullptr;
178   }
179   if (GEP->getNumOperands() != 2) {
180     LLVM_DEBUG(dbgs() << "masked gathers/scatters: getelementptr with too many"
181                       << " operands. Expanding.\n");
182     return nullptr;
183   }
184   Offsets = GEP->getOperand(1);
185   // Paranoid check whether the number of parallel lanes is the same
186   assert(cast<VectorType>(Ty)->getNumElements() ==
187          cast<VectorType>(Offsets->getType())->getNumElements());
188   // Only <N x i32> offsets can be integrated into an arm gather, any smaller
189   // type would have to be sign extended by the gep - and arm gathers can only
190   // zero extend. Additionally, the offsets do have to originate from a zext of
191   // a vector with element types smaller or equal the type of the gather we're
192   // looking at
193   if (Offsets->getType()->getScalarSizeInBits() != 32)
194     return nullptr;
195   if (ZExtInst *ZextOffs = dyn_cast<ZExtInst>(Offsets))
196     Offsets = ZextOffs->getOperand(0);
197   else if (!(cast<VectorType>(Offsets->getType())->getNumElements() == 4 &&
198              Offsets->getType()->getScalarSizeInBits() == 32))
199     return nullptr;
200 
201   if (Ty != Offsets->getType()) {
202     if ((Ty->getScalarSizeInBits() <
203          Offsets->getType()->getScalarSizeInBits())) {
204       LLVM_DEBUG(dbgs() << "masked gathers/scatters: no correct offset type."
205                         << " Can't create intrinsic.\n");
206       return nullptr;
207     } else {
208       Offsets = Builder.CreateZExt(
209           Offsets, VectorType::getInteger(cast<VectorType>(Ty)));
210     }
211   }
212   // If none of the checks failed, return the gep's base pointer
213   LLVM_DEBUG(dbgs() << "masked gathers/scatters: found correct offsets\n");
214   return GEPPtr;
215 }
216 
217 void MVEGatherScatterLowering::lookThroughBitcast(Value *&Ptr) {
218   // Look through bitcast instruction if #elements is the same
219   if (auto *BitCast = dyn_cast<BitCastInst>(Ptr)) {
220     auto *BCTy = cast<VectorType>(BitCast->getType());
221     auto *BCSrcTy = cast<VectorType>(BitCast->getOperand(0)->getType());
222     if (BCTy->getNumElements() == BCSrcTy->getNumElements()) {
223       LLVM_DEBUG(
224           dbgs() << "masked gathers/scatters: looking through bitcast\n");
225       Ptr = BitCast->getOperand(0);
226     }
227   }
228 }
229 
230 int MVEGatherScatterLowering::computeScale(unsigned GEPElemSize,
231                                            unsigned MemoryElemSize) {
232   // This can be a 32bit load/store scaled by 4, a 16bit load/store scaled by 2,
233   // or a 8bit, 16bit or 32bit load/store scaled by 1
234   if (GEPElemSize == 32 && MemoryElemSize == 32)
235     return 2;
236   else if (GEPElemSize == 16 && MemoryElemSize == 16)
237     return 1;
238   else if (GEPElemSize == 8)
239     return 0;
240   LLVM_DEBUG(dbgs() << "masked gathers/scatters: incorrect scale. Can't "
241                     << "create intrinsic\n");
242   return -1;
243 }
244 
245 Optional<int64_t> MVEGatherScatterLowering::getIfConst(const Value *V) {
246   const Constant *C = dyn_cast<Constant>(V);
247   if (C != nullptr)
248     return Optional<int64_t>{C->getUniqueInteger().getSExtValue()};
249   if (!isa<Instruction>(V))
250     return Optional<int64_t>{};
251 
252   const Instruction *I = cast<Instruction>(V);
253   if (I->getOpcode() == Instruction::Add ||
254               I->getOpcode() == Instruction::Mul) {
255     Optional<int64_t> Op0 = getIfConst(I->getOperand(0));
256     Optional<int64_t> Op1 = getIfConst(I->getOperand(1));
257     if (!Op0 || !Op1)
258       return Optional<int64_t>{};
259     if (I->getOpcode() == Instruction::Add)
260       return Optional<int64_t>{Op0.getValue() + Op1.getValue()};
261     if (I->getOpcode() == Instruction::Mul)
262       return Optional<int64_t>{Op0.getValue() * Op1.getValue()};
263   }
264   return Optional<int64_t>{};
265 }
266 
267 std::pair<Value *, int64_t>
268 MVEGatherScatterLowering::getVarAndConst(Value *Inst, int TypeScale) {
269   std::pair<Value *, int64_t> ReturnFalse =
270       std::pair<Value *, int64_t>(nullptr, 0);
271   // At this point, the instruction we're looking at must be an add or we
272   // bail out
273   Instruction *Add = dyn_cast<Instruction>(Inst);
274   if (Add == nullptr || Add->getOpcode() != Instruction::Add)
275     return ReturnFalse;
276 
277   Value *Summand;
278   Optional<int64_t> Const;
279   // Find out which operand the value that is increased is
280   if ((Const = getIfConst(Add->getOperand(0))))
281     Summand = Add->getOperand(1);
282   else if ((Const = getIfConst(Add->getOperand(1))))
283     Summand = Add->getOperand(0);
284   else
285     return ReturnFalse;
286 
287   // Check that the constant is small enough for an incrementing gather
288   int64_t Immediate = Const.getValue() << TypeScale;
289   if (Immediate > 512 || Immediate < -512 || Immediate % 4 != 0)
290     return ReturnFalse;
291 
292   return std::pair<Value *, int64_t>(Summand, Immediate);
293 }
294 
295 Value *MVEGatherScatterLowering::lowerGather(IntrinsicInst *I) {
296   using namespace PatternMatch;
297   LLVM_DEBUG(dbgs() << "masked gathers: checking transform preconditions\n");
298 
299   // @llvm.masked.gather.*(Ptrs, alignment, Mask, Src0)
300   // Attempt to turn the masked gather in I into a MVE intrinsic
301   // Potentially optimising the addressing modes as we do so.
302   auto *Ty = cast<VectorType>(I->getType());
303   Value *Ptr = I->getArgOperand(0);
304   unsigned Alignment = cast<ConstantInt>(I->getArgOperand(1))->getZExtValue();
305   Value *Mask = I->getArgOperand(2);
306   Value *PassThru = I->getArgOperand(3);
307 
308   if (!isLegalTypeAndAlignment(Ty->getNumElements(), Ty->getScalarSizeInBits(),
309                                Alignment))
310     return nullptr;
311   lookThroughBitcast(Ptr);
312   assert(Ptr->getType()->isVectorTy() && "Unexpected pointer type");
313 
314   IRBuilder<> Builder(I->getContext());
315   Builder.SetInsertPoint(I);
316   Builder.SetCurrentDebugLocation(I->getDebugLoc());
317 
318   Instruction *Root = I;
319   Value *Load = tryCreateMaskedGatherOffset(I, Ptr, Root, Builder);
320   if (!Load)
321     Load = tryCreateMaskedGatherBase(I, Ptr, Builder);
322   if (!Load)
323     return nullptr;
324 
325   if (!isa<UndefValue>(PassThru) && !match(PassThru, m_Zero())) {
326     LLVM_DEBUG(dbgs() << "masked gathers: found non-trivial passthru - "
327                       << "creating select\n");
328     Load = Builder.CreateSelect(Mask, Load, PassThru);
329   }
330 
331   Root->replaceAllUsesWith(Load);
332   Root->eraseFromParent();
333   if (Root != I)
334     // If this was an extending gather, we need to get rid of the sext/zext
335     // sext/zext as well as of the gather itself
336     I->eraseFromParent();
337 
338   LLVM_DEBUG(dbgs() << "masked gathers: successfully built masked gather\n");
339   return Load;
340 }
341 
342 Value *MVEGatherScatterLowering::tryCreateMaskedGatherBase(IntrinsicInst *I,
343                                                            Value *Ptr,
344                                                            IRBuilder<> &Builder,
345                                                            unsigned Increment) {
346   using namespace PatternMatch;
347   auto *Ty = cast<VectorType>(I->getType());
348   LLVM_DEBUG(dbgs() << "masked gathers: loading from vector of pointers\n");
349   if (Ty->getNumElements() != 4 || Ty->getScalarSizeInBits() != 32)
350     // Can't build an intrinsic for this
351     return nullptr;
352   Value *Mask = I->getArgOperand(2);
353   if (match(Mask, m_One()))
354     return Builder.CreateIntrinsic(Intrinsic::arm_mve_vldr_gather_base,
355                                    {Ty, Ptr->getType()},
356                                    {Ptr, Builder.getInt32(Increment)});
357   else
358     return Builder.CreateIntrinsic(
359         Intrinsic::arm_mve_vldr_gather_base_predicated,
360         {Ty, Ptr->getType(), Mask->getType()},
361         {Ptr, Builder.getInt32(Increment), Mask});
362 }
363 
364 Value *MVEGatherScatterLowering::tryCreateMaskedGatherBaseWB(
365     IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder, unsigned Increment) {
366   using namespace PatternMatch;
367   auto *Ty = cast<VectorType>(I->getType());
368   LLVM_DEBUG(
369       dbgs()
370       << "masked gathers: loading from vector of pointers with writeback\n");
371   if (Ty->getNumElements() != 4 || Ty->getScalarSizeInBits() != 32)
372     // Can't build an intrinsic for this
373     return nullptr;
374   Value *Mask = I->getArgOperand(2);
375   if (match(Mask, m_One()))
376     return Builder.CreateIntrinsic(Intrinsic::arm_mve_vldr_gather_base_wb,
377                                    {Ty, Ptr->getType()},
378                                    {Ptr, Builder.getInt32(Increment)});
379   else
380     return Builder.CreateIntrinsic(
381         Intrinsic::arm_mve_vldr_gather_base_wb_predicated,
382         {Ty, Ptr->getType(), Mask->getType()},
383         {Ptr, Builder.getInt32(Increment), Mask});
384 }
385 
386 Value *MVEGatherScatterLowering::tryCreateMaskedGatherOffset(
387     IntrinsicInst *I, Value *Ptr, Instruction *&Root, IRBuilder<> &Builder) {
388   using namespace PatternMatch;
389 
390   Type *OriginalTy = I->getType();
391   Type *ResultTy = OriginalTy;
392 
393   unsigned Unsigned = 1;
394   // The size of the gather was already checked in isLegalTypeAndAlignment;
395   // if it was not a full vector width an appropriate extend should follow.
396   auto *Extend = Root;
397   if (OriginalTy->getPrimitiveSizeInBits() < 128) {
398     // Only transform gathers with exactly one use
399     if (!I->hasOneUse())
400       return nullptr;
401 
402     // The correct root to replace is not the CallInst itself, but the
403     // instruction which extends it
404     Extend = cast<Instruction>(*I->users().begin());
405     if (isa<SExtInst>(Extend)) {
406       Unsigned = 0;
407     } else if (!isa<ZExtInst>(Extend)) {
408       LLVM_DEBUG(dbgs() << "masked gathers: extend needed but not provided. "
409                         << "Expanding\n");
410       return nullptr;
411     }
412     LLVM_DEBUG(dbgs() << "masked gathers: found an extending gather\n");
413     ResultTy = Extend->getType();
414     // The final size of the gather must be a full vector width
415     if (ResultTy->getPrimitiveSizeInBits() != 128) {
416       LLVM_DEBUG(dbgs() << "masked gathers: extending from the wrong type. "
417                         << "Expanding\n");
418       return nullptr;
419     }
420   }
421 
422   GetElementPtrInst *GEP = dyn_cast<GetElementPtrInst>(Ptr);
423   Value *Offsets;
424   Value *BasePtr = checkGEP(Offsets, ResultTy, GEP, Builder);
425   if (!BasePtr)
426     return nullptr;
427   // Check whether the offset is a constant increment that could be merged into
428   // a QI gather
429   Value *Load =
430       tryCreateIncrementingGather(I, BasePtr, Offsets, GEP, Builder);
431   if (Load)
432     return Load;
433 
434   int Scale = computeScale(
435       BasePtr->getType()->getPointerElementType()->getPrimitiveSizeInBits(),
436       OriginalTy->getScalarSizeInBits());
437   if (Scale == -1)
438     return nullptr;
439   Root = Extend;
440 
441   Value *Mask = I->getArgOperand(2);
442   if (!match(Mask, m_One()))
443     return Builder.CreateIntrinsic(
444         Intrinsic::arm_mve_vldr_gather_offset_predicated,
445         {ResultTy, BasePtr->getType(), Offsets->getType(), Mask->getType()},
446         {BasePtr, Offsets, Builder.getInt32(OriginalTy->getScalarSizeInBits()),
447          Builder.getInt32(Scale), Builder.getInt32(Unsigned), Mask});
448   else
449     return Builder.CreateIntrinsic(
450         Intrinsic::arm_mve_vldr_gather_offset,
451         {ResultTy, BasePtr->getType(), Offsets->getType()},
452         {BasePtr, Offsets, Builder.getInt32(OriginalTy->getScalarSizeInBits()),
453          Builder.getInt32(Scale), Builder.getInt32(Unsigned)});
454 }
455 
456 Value *MVEGatherScatterLowering::tryCreateIncrementingGather(
457     IntrinsicInst *I, Value *BasePtr, Value *Offsets, GetElementPtrInst *GEP,
458     IRBuilder<> &Builder) {
459   auto *Ty = cast<VectorType>(I->getType());
460   // Incrementing gathers only exist for v4i32
461   if (Ty->getNumElements() != 4 || Ty->getScalarSizeInBits() != 32)
462     return nullptr;
463   Loop *L = LI->getLoopFor(I->getParent());
464   if (L == nullptr)
465     // Incrementing gathers are not beneficial outside of a loop
466     return nullptr;
467   LLVM_DEBUG(
468       dbgs() << "masked gathers: trying to build incrementing wb gather\n");
469 
470   // The gep was in charge of making sure the offsets are scaled correctly
471   // - calculate that factor so it can be applied by hand
472   DataLayout DT = I->getParent()->getParent()->getParent()->getDataLayout();
473   int TypeScale =
474       computeScale(DT.getTypeSizeInBits(GEP->getOperand(0)->getType()),
475                    DT.getTypeSizeInBits(GEP->getType()) /
476                        cast<VectorType>(GEP->getType())->getNumElements());
477   if (TypeScale == -1)
478     return nullptr;
479 
480   if (GEP->hasOneUse()) {
481     // Only in this case do we want to build a wb gather, because the wb will
482     // change the phi which does affect other users of the gep (which will still
483     // be using the phi in the old way)
484     Value *Load =
485         tryCreateIncrementingWBGather(I, BasePtr, Offsets, TypeScale, Builder);
486     if (Load != nullptr)
487       return Load;
488   }
489   LLVM_DEBUG(
490       dbgs() << "masked gathers: trying to build incrementing non-wb gather\n");
491 
492   std::pair<Value *, int64_t> Add = getVarAndConst(Offsets, TypeScale);
493   if (Add.first == nullptr)
494     return nullptr;
495   Value *OffsetsIncoming = Add.first;
496   int64_t Immediate = Add.second;
497 
498   // Make sure the offsets are scaled correctly
499   Instruction *ScaledOffsets = BinaryOperator::Create(
500       Instruction::Shl, OffsetsIncoming,
501       Builder.CreateVectorSplat(Ty->getNumElements(), Builder.getInt32(TypeScale)),
502       "ScaledIndex", I);
503   // Add the base to the offsets
504   OffsetsIncoming = BinaryOperator::Create(
505       Instruction::Add, ScaledOffsets,
506       Builder.CreateVectorSplat(
507           Ty->getNumElements(),
508           Builder.CreatePtrToInt(
509               BasePtr,
510               cast<VectorType>(ScaledOffsets->getType())->getElementType())),
511       "StartIndex", I);
512 
513   return cast<IntrinsicInst>(
514       tryCreateMaskedGatherBase(I, OffsetsIncoming, Builder, Immediate));
515 }
516 
517 Value *MVEGatherScatterLowering::tryCreateIncrementingWBGather(
518     IntrinsicInst *I, Value *BasePtr, Value *Offsets, unsigned TypeScale,
519     IRBuilder<> &Builder) {
520   // Check whether this gather's offset is incremented by a constant - if so,
521   // and the load is of the right type, we can merge this into a QI gather
522   Loop *L = LI->getLoopFor(I->getParent());
523   // Offsets that are worth merging into this instruction will be incremented
524   // by a constant, thus we're looking for an add of a phi and a constant
525   PHINode *Phi = dyn_cast<PHINode>(Offsets);
526   if (Phi == nullptr || Phi->getNumIncomingValues() != 2 ||
527       Phi->getParent() != L->getHeader() || Phi->getNumUses() != 2)
528     // No phi means no IV to write back to; if there is a phi, we expect it
529     // to have exactly two incoming values; the only phis we are interested in
530     // will be loop IV's and have exactly two uses, one in their increment and
531     // one in the gather's gep
532     return nullptr;
533 
534   unsigned IncrementIndex =
535       Phi->getIncomingBlock(0) == L->getLoopLatch() ? 0 : 1;
536   // Look through the phi to the phi increment
537   Offsets = Phi->getIncomingValue(IncrementIndex);
538 
539   std::pair<Value *, int64_t> Add = getVarAndConst(Offsets, TypeScale);
540   if (Add.first == nullptr)
541     return nullptr;
542   Value *OffsetsIncoming = Add.first;
543   int64_t Immediate = Add.second;
544   if (OffsetsIncoming != Phi)
545     // Then the increment we are looking at is not an increment of the
546     // induction variable, and we don't want to do a writeback
547     return nullptr;
548 
549   Builder.SetInsertPoint(&Phi->getIncomingBlock(1 - IncrementIndex)->back());
550   unsigned NumElems =
551       cast<VectorType>(OffsetsIncoming->getType())->getNumElements();
552 
553   // Make sure the offsets are scaled correctly
554   Instruction *ScaledOffsets = BinaryOperator::Create(
555       Instruction::Shl, Phi->getIncomingValue(1 - IncrementIndex),
556       Builder.CreateVectorSplat(NumElems, Builder.getInt32(TypeScale)),
557       "ScaledIndex", &Phi->getIncomingBlock(1 - IncrementIndex)->back());
558   // Add the base to the offsets
559   OffsetsIncoming = BinaryOperator::Create(
560       Instruction::Add, ScaledOffsets,
561       Builder.CreateVectorSplat(
562           NumElems,
563           Builder.CreatePtrToInt(
564               BasePtr,
565               cast<VectorType>(ScaledOffsets->getType())->getElementType())),
566       "StartIndex", &Phi->getIncomingBlock(1 - IncrementIndex)->back());
567   // The gather is pre-incrementing
568   OffsetsIncoming = BinaryOperator::Create(
569       Instruction::Sub, OffsetsIncoming,
570       Builder.CreateVectorSplat(NumElems, Builder.getInt32(Immediate)),
571       "PreIncrementStartIndex",
572       &Phi->getIncomingBlock(1 - IncrementIndex)->back());
573   Phi->setIncomingValue(1 - IncrementIndex, OffsetsIncoming);
574 
575   Builder.SetInsertPoint(I);
576 
577   // Build the incrementing gather
578   Value *Load = tryCreateMaskedGatherBaseWB(I, Phi, Builder, Immediate);
579 
580   // One value to be handed to whoever uses the gather, one is the loop
581   // increment
582   Value *ExtractedLoad = Builder.CreateExtractValue(Load, 0, "Gather");
583   Value *Inc = Builder.CreateExtractValue(Load, 1, "GatherIncrement");
584   Instruction *AddInst = cast<Instruction>(Offsets);
585   AddInst->replaceAllUsesWith(Inc);
586   AddInst->eraseFromParent();
587   Phi->setIncomingValue(IncrementIndex, Inc);
588 
589   return ExtractedLoad;
590 }
591 
592 Value *MVEGatherScatterLowering::lowerScatter(IntrinsicInst *I) {
593   using namespace PatternMatch;
594   LLVM_DEBUG(dbgs() << "masked scatters: checking transform preconditions\n");
595 
596   // @llvm.masked.scatter.*(data, ptrs, alignment, mask)
597   // Attempt to turn the masked scatter in I into a MVE intrinsic
598   // Potentially optimising the addressing modes as we do so.
599   Value *Input = I->getArgOperand(0);
600   Value *Ptr = I->getArgOperand(1);
601   unsigned Alignment = cast<ConstantInt>(I->getArgOperand(2))->getZExtValue();
602   auto *Ty = cast<VectorType>(Input->getType());
603 
604   if (!isLegalTypeAndAlignment(Ty->getNumElements(), Ty->getScalarSizeInBits(),
605                                Alignment))
606     return nullptr;
607 
608   lookThroughBitcast(Ptr);
609   assert(Ptr->getType()->isVectorTy() && "Unexpected pointer type");
610 
611   IRBuilder<> Builder(I->getContext());
612   Builder.SetInsertPoint(I);
613   Builder.SetCurrentDebugLocation(I->getDebugLoc());
614 
615   Value *Store = tryCreateMaskedScatterOffset(I, Ptr, Builder);
616   if (!Store)
617     Store = tryCreateMaskedScatterBase(I, Ptr, Builder);
618   if (!Store)
619     return nullptr;
620 
621   LLVM_DEBUG(dbgs() << "masked scatters: successfully built masked scatter\n");
622   I->replaceAllUsesWith(Store);
623   I->eraseFromParent();
624   return Store;
625 }
626 
627 Value *MVEGatherScatterLowering::tryCreateMaskedScatterBase(
628     IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) {
629   using namespace PatternMatch;
630   Value *Input = I->getArgOperand(0);
631   Value *Mask = I->getArgOperand(3);
632   auto *Ty = cast<VectorType>(Input->getType());
633   // Only QR variants allow truncating
634   if (!(Ty->getNumElements() == 4 && Ty->getScalarSizeInBits() == 32)) {
635     // Can't build an intrinsic for this
636     return nullptr;
637   }
638   //  int_arm_mve_vstr_scatter_base(_predicated) addr, offset, data(, mask)
639   LLVM_DEBUG(dbgs() << "masked scatters: storing to a vector of pointers\n");
640   if (match(Mask, m_One()))
641     return Builder.CreateIntrinsic(Intrinsic::arm_mve_vstr_scatter_base,
642                                    {Ptr->getType(), Input->getType()},
643                                    {Ptr, Builder.getInt32(0), Input});
644   else
645     return Builder.CreateIntrinsic(
646         Intrinsic::arm_mve_vstr_scatter_base_predicated,
647         {Ptr->getType(), Input->getType(), Mask->getType()},
648         {Ptr, Builder.getInt32(0), Input, Mask});
649 }
650 
651 Value *MVEGatherScatterLowering::tryCreateMaskedScatterOffset(
652     IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) {
653   using namespace PatternMatch;
654   Value *Input = I->getArgOperand(0);
655   Value *Mask = I->getArgOperand(3);
656   Type *InputTy = Input->getType();
657   Type *MemoryTy = InputTy;
658   LLVM_DEBUG(dbgs() << "masked scatters: getelementpointer found. Storing"
659                     << " to base + vector of offsets\n");
660   // If the input has been truncated, try to integrate that trunc into the
661   // scatter instruction (we don't care about alignment here)
662   if (TruncInst *Trunc = dyn_cast<TruncInst>(Input)) {
663     Value *PreTrunc = Trunc->getOperand(0);
664     Type *PreTruncTy = PreTrunc->getType();
665     if (PreTruncTy->getPrimitiveSizeInBits() == 128) {
666       Input = PreTrunc;
667       InputTy = PreTruncTy;
668     }
669   }
670   if (InputTy->getPrimitiveSizeInBits() != 128) {
671     LLVM_DEBUG(
672         dbgs() << "masked scatters: cannot create scatters for non-standard"
673                << " input types. Expanding.\n");
674     return nullptr;
675   }
676 
677   GetElementPtrInst *GEP = dyn_cast<GetElementPtrInst>(Ptr);
678   Value *Offsets;
679   Value *BasePtr = checkGEP(Offsets, InputTy, GEP, Builder);
680   if (!BasePtr)
681     return nullptr;
682   int Scale = computeScale(
683       BasePtr->getType()->getPointerElementType()->getPrimitiveSizeInBits(),
684       MemoryTy->getScalarSizeInBits());
685   if (Scale == -1)
686     return nullptr;
687 
688   if (!match(Mask, m_One()))
689     return Builder.CreateIntrinsic(
690         Intrinsic::arm_mve_vstr_scatter_offset_predicated,
691         {BasePtr->getType(), Offsets->getType(), Input->getType(),
692          Mask->getType()},
693         {BasePtr, Offsets, Input,
694          Builder.getInt32(MemoryTy->getScalarSizeInBits()),
695          Builder.getInt32(Scale), Mask});
696   else
697     return Builder.CreateIntrinsic(
698         Intrinsic::arm_mve_vstr_scatter_offset,
699         {BasePtr->getType(), Offsets->getType(), Input->getType()},
700         {BasePtr, Offsets, Input,
701          Builder.getInt32(MemoryTy->getScalarSizeInBits()),
702          Builder.getInt32(Scale)});
703 }
704 
705 void MVEGatherScatterLowering::pushOutAdd(PHINode *&Phi,
706                                           Value *OffsSecondOperand,
707                                           unsigned StartIndex) {
708   LLVM_DEBUG(dbgs() << "masked gathers/scatters: optimising add instruction\n");
709   Instruction *InsertionPoint =
710         &cast<Instruction>(Phi->getIncomingBlock(StartIndex)->back());
711   // Initialize the phi with a vector that contains a sum of the constants
712   Instruction *NewIndex = BinaryOperator::Create(
713       Instruction::Add, Phi->getIncomingValue(StartIndex), OffsSecondOperand,
714       "PushedOutAdd", InsertionPoint);
715   unsigned IncrementIndex = StartIndex == 0 ? 1 : 0;
716 
717   // Order such that start index comes first (this reduces mov's)
718   Phi->addIncoming(NewIndex, Phi->getIncomingBlock(StartIndex));
719   Phi->addIncoming(Phi->getIncomingValue(IncrementIndex),
720                    Phi->getIncomingBlock(IncrementIndex));
721   Phi->removeIncomingValue(IncrementIndex);
722   Phi->removeIncomingValue(StartIndex);
723 }
724 
725 void MVEGatherScatterLowering::pushOutMul(PHINode *&Phi,
726                                           Value *IncrementPerRound,
727                                           Value *OffsSecondOperand,
728                                           unsigned LoopIncrement,
729                                           IRBuilder<> &Builder) {
730   LLVM_DEBUG(dbgs() << "masked gathers/scatters: optimising mul instruction\n");
731 
732   // Create a new scalar add outside of the loop and transform it to a splat
733   // by which loop variable can be incremented
734   Instruction *InsertionPoint = &cast<Instruction>(
735         Phi->getIncomingBlock(LoopIncrement == 1 ? 0 : 1)->back());
736 
737   // Create a new index
738   Value *StartIndex = BinaryOperator::Create(
739       Instruction::Mul, Phi->getIncomingValue(LoopIncrement == 1 ? 0 : 1),
740       OffsSecondOperand, "PushedOutMul", InsertionPoint);
741 
742   Instruction *Product =
743       BinaryOperator::Create(Instruction::Mul, IncrementPerRound,
744                              OffsSecondOperand, "Product", InsertionPoint);
745   // Increment NewIndex by Product instead of the multiplication
746   Instruction *NewIncrement = BinaryOperator::Create(
747       Instruction::Add, Phi, Product, "IncrementPushedOutMul",
748       cast<Instruction>(Phi->getIncomingBlock(LoopIncrement)->back())
749           .getPrevNode());
750 
751   Phi->addIncoming(StartIndex,
752                    Phi->getIncomingBlock(LoopIncrement == 1 ? 0 : 1));
753   Phi->addIncoming(NewIncrement, Phi->getIncomingBlock(LoopIncrement));
754   Phi->removeIncomingValue((unsigned)0);
755   Phi->removeIncomingValue((unsigned)0);
756   return;
757 }
758 
759 // Check whether all usages of this instruction are as offsets of
760 // gathers/scatters or simple arithmetics only used by gathers/scatters
761 static bool hasAllGatScatUsers(Instruction *I) {
762   if (I->hasNUses(0)) {
763     return false;
764   }
765   bool Gatscat = true;
766   for (User *U : I->users()) {
767     if (!isa<Instruction>(U))
768       return false;
769     if (isa<GetElementPtrInst>(U) ||
770         isGatherScatter(dyn_cast<IntrinsicInst>(U))) {
771       return Gatscat;
772     } else {
773       unsigned OpCode = cast<Instruction>(U)->getOpcode();
774       if ((OpCode == Instruction::Add || OpCode == Instruction::Mul) &&
775           hasAllGatScatUsers(cast<Instruction>(U))) {
776         continue;
777       }
778       return false;
779     }
780   }
781   return Gatscat;
782 }
783 
784 bool MVEGatherScatterLowering::optimiseOffsets(Value *Offsets, BasicBlock *BB,
785                                                LoopInfo *LI) {
786   LLVM_DEBUG(dbgs() << "masked gathers/scatters: trying to optimize\n");
787   // Optimise the addresses of gathers/scatters by moving invariant
788   // calculations out of the loop
789   if (!isa<Instruction>(Offsets))
790     return false;
791   Instruction *Offs = cast<Instruction>(Offsets);
792   if (Offs->getOpcode() != Instruction::Add &&
793       Offs->getOpcode() != Instruction::Mul)
794     return false;
795   Loop *L = LI->getLoopFor(BB);
796   if (L == nullptr)
797     return false;
798   if (!Offs->hasOneUse()) {
799     if (!hasAllGatScatUsers(Offs))
800       return false;
801   }
802 
803   // Find out which, if any, operand of the instruction
804   // is a phi node
805   PHINode *Phi;
806   int OffsSecondOp;
807   if (isa<PHINode>(Offs->getOperand(0))) {
808     Phi = cast<PHINode>(Offs->getOperand(0));
809     OffsSecondOp = 1;
810   } else if (isa<PHINode>(Offs->getOperand(1))) {
811     Phi = cast<PHINode>(Offs->getOperand(1));
812     OffsSecondOp = 0;
813   } else {
814     bool Changed = true;
815     if (isa<Instruction>(Offs->getOperand(0)) &&
816         L->contains(cast<Instruction>(Offs->getOperand(0))))
817       Changed |= optimiseOffsets(Offs->getOperand(0), BB, LI);
818     if (isa<Instruction>(Offs->getOperand(1)) &&
819         L->contains(cast<Instruction>(Offs->getOperand(1))))
820       Changed |= optimiseOffsets(Offs->getOperand(1), BB, LI);
821     if (!Changed) {
822       return false;
823     } else {
824       if (isa<PHINode>(Offs->getOperand(0))) {
825         Phi = cast<PHINode>(Offs->getOperand(0));
826         OffsSecondOp = 1;
827       } else if (isa<PHINode>(Offs->getOperand(1))) {
828         Phi = cast<PHINode>(Offs->getOperand(1));
829         OffsSecondOp = 0;
830       } else {
831         return false;
832       }
833     }
834   }
835   // A phi node we want to perform this function on should be from the
836   // loop header, and shouldn't have more than 2 incoming values
837   if (Phi->getParent() != L->getHeader() ||
838       Phi->getNumIncomingValues() != 2)
839     return false;
840 
841   // The phi must be an induction variable
842   Instruction *Op;
843   int IncrementingBlock = -1;
844 
845   for (int i = 0; i < 2; i++)
846     if ((Op = dyn_cast<Instruction>(Phi->getIncomingValue(i))) != nullptr)
847       if (Op->getOpcode() == Instruction::Add &&
848           (Op->getOperand(0) == Phi || Op->getOperand(1) == Phi))
849         IncrementingBlock = i;
850   if (IncrementingBlock == -1)
851     return false;
852 
853   Instruction *IncInstruction =
854       cast<Instruction>(Phi->getIncomingValue(IncrementingBlock));
855 
856   // If the phi is not used by anything else, we can just adapt it when
857   // replacing the instruction; if it is, we'll have to duplicate it
858   PHINode *NewPhi;
859   Value *IncrementPerRound = IncInstruction->getOperand(
860       (IncInstruction->getOperand(0) == Phi) ? 1 : 0);
861 
862   // Get the value that is added to/multiplied with the phi
863   Value *OffsSecondOperand = Offs->getOperand(OffsSecondOp);
864 
865   if (IncrementPerRound->getType() != OffsSecondOperand->getType())
866     // Something has gone wrong, abort
867     return false;
868 
869   // Only proceed if the increment per round is a constant or an instruction
870   // which does not originate from within the loop
871   if (!isa<Constant>(IncrementPerRound) &&
872       !(isa<Instruction>(IncrementPerRound) &&
873         !L->contains(cast<Instruction>(IncrementPerRound))))
874     return false;
875 
876   if (Phi->getNumUses() == 2) {
877     // No other users -> reuse existing phi (One user is the instruction
878     // we're looking at, the other is the phi increment)
879     if (IncInstruction->getNumUses() != 1) {
880       // If the incrementing instruction does have more users than
881       // our phi, we need to copy it
882       IncInstruction = BinaryOperator::Create(
883           Instruction::BinaryOps(IncInstruction->getOpcode()), Phi,
884           IncrementPerRound, "LoopIncrement", IncInstruction);
885       Phi->setIncomingValue(IncrementingBlock, IncInstruction);
886     }
887     NewPhi = Phi;
888   } else {
889     // There are other users -> create a new phi
890     NewPhi = PHINode::Create(Phi->getType(), 0, "NewPhi", Phi);
891     std::vector<Value *> Increases;
892     // Copy the incoming values of the old phi
893     NewPhi->addIncoming(Phi->getIncomingValue(IncrementingBlock == 1 ? 0 : 1),
894                         Phi->getIncomingBlock(IncrementingBlock == 1 ? 0 : 1));
895     IncInstruction = BinaryOperator::Create(
896         Instruction::BinaryOps(IncInstruction->getOpcode()), NewPhi,
897         IncrementPerRound, "LoopIncrement", IncInstruction);
898     NewPhi->addIncoming(IncInstruction,
899                         Phi->getIncomingBlock(IncrementingBlock));
900     IncrementingBlock = 1;
901   }
902 
903   IRBuilder<> Builder(BB->getContext());
904   Builder.SetInsertPoint(Phi);
905   Builder.SetCurrentDebugLocation(Offs->getDebugLoc());
906 
907   switch (Offs->getOpcode()) {
908   case Instruction::Add:
909     pushOutAdd(NewPhi, OffsSecondOperand, IncrementingBlock == 1 ? 0 : 1);
910     break;
911   case Instruction::Mul:
912     pushOutMul(NewPhi, IncrementPerRound, OffsSecondOperand, IncrementingBlock,
913                Builder);
914     break;
915   default:
916     return false;
917   }
918   LLVM_DEBUG(
919       dbgs() << "masked gathers/scatters: simplified loop variable add/mul\n");
920 
921   // The instruction has now been "absorbed" into the phi value
922   Offs->replaceAllUsesWith(NewPhi);
923   if (Offs->hasNUses(0))
924     Offs->eraseFromParent();
925   // Clean up the old increment in case it's unused because we built a new
926   // one
927   if (IncInstruction->hasNUses(0))
928     IncInstruction->eraseFromParent();
929 
930   return true;
931 }
932 
933 bool MVEGatherScatterLowering::runOnFunction(Function &F) {
934   if (!EnableMaskedGatherScatters)
935     return false;
936   auto &TPC = getAnalysis<TargetPassConfig>();
937   auto &TM = TPC.getTM<TargetMachine>();
938   auto *ST = &TM.getSubtarget<ARMSubtarget>(F);
939   if (!ST->hasMVEIntegerOps())
940     return false;
941   LI = &getAnalysis<LoopInfoWrapperPass>().getLoopInfo();
942   SmallVector<IntrinsicInst *, 4> Gathers;
943   SmallVector<IntrinsicInst *, 4> Scatters;
944 
945   for (BasicBlock &BB : F) {
946     for (Instruction &I : BB) {
947       IntrinsicInst *II = dyn_cast<IntrinsicInst>(&I);
948       if (II && II->getIntrinsicID() == Intrinsic::masked_gather) {
949         Gathers.push_back(II);
950         if (isa<GetElementPtrInst>(II->getArgOperand(0)))
951           optimiseOffsets(
952               cast<Instruction>(II->getArgOperand(0))->getOperand(1),
953               II->getParent(), LI);
954       } else if (II && II->getIntrinsicID() == Intrinsic::masked_scatter) {
955         Scatters.push_back(II);
956         if (isa<GetElementPtrInst>(II->getArgOperand(1)))
957           optimiseOffsets(
958               cast<Instruction>(II->getArgOperand(1))->getOperand(1),
959               II->getParent(), LI);
960       }
961     }
962   }
963 
964   bool Changed = false;
965   for (unsigned i = 0; i < Gathers.size(); i++) {
966     IntrinsicInst *I = Gathers[i];
967     Value *L = lowerGather(I);
968     if (L == nullptr)
969       continue;
970 
971     // Get rid of any now dead instructions
972     SimplifyInstructionsInBlock(cast<Instruction>(L)->getParent());
973     Changed = true;
974   }
975 
976   for (unsigned i = 0; i < Scatters.size(); i++) {
977     IntrinsicInst *I = Scatters[i];
978     Value *S = lowerScatter(I);
979     if (S == nullptr)
980       continue;
981 
982     // Get rid of any now dead instructions
983     SimplifyInstructionsInBlock(cast<Instruction>(S)->getParent());
984     Changed = true;
985   }
986   return Changed;
987 }
988