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/TargetTransformInfo.h"
19 #include "llvm/CodeGen/TargetLowering.h"
20 #include "llvm/CodeGen/TargetPassConfig.h"
21 #include "llvm/CodeGen/TargetSubtargetInfo.h"
22 #include "llvm/InitializePasses.h"
23 #include "llvm/IR/BasicBlock.h"
24 #include "llvm/IR/Constant.h"
25 #include "llvm/IR/Constants.h"
26 #include "llvm/IR/DerivedTypes.h"
27 #include "llvm/IR/Function.h"
28 #include "llvm/IR/InstrTypes.h"
29 #include "llvm/IR/Instruction.h"
30 #include "llvm/IR/Instructions.h"
31 #include "llvm/IR/IntrinsicInst.h"
32 #include "llvm/IR/Intrinsics.h"
33 #include "llvm/IR/IntrinsicsARM.h"
34 #include "llvm/IR/IRBuilder.h"
35 #include "llvm/IR/PatternMatch.h"
36 #include "llvm/IR/Type.h"
37 #include "llvm/IR/Value.h"
38 #include "llvm/Pass.h"
39 #include "llvm/Support/Casting.h"
40 #include <algorithm>
41 #include <cassert>
42 
43 using namespace llvm;
44 
45 #define DEBUG_TYPE "mve-gather-scatter-lowering"
46 
47 cl::opt<bool> EnableMaskedGatherScatters(
48     "enable-arm-maskedgatscat", cl::Hidden, cl::init(false),
49     cl::desc("Enable the generation of masked gathers and scatters"));
50 
51 namespace {
52 
53 class MVEGatherScatterLowering : public FunctionPass {
54 public:
55   static char ID; // Pass identification, replacement for typeid
56 
57   explicit MVEGatherScatterLowering() : FunctionPass(ID) {
58     initializeMVEGatherScatterLoweringPass(*PassRegistry::getPassRegistry());
59   }
60 
61   bool runOnFunction(Function &F) override;
62 
63   StringRef getPassName() const override {
64     return "MVE gather/scatter lowering";
65   }
66 
67   void getAnalysisUsage(AnalysisUsage &AU) const override {
68     AU.setPreservesCFG();
69     AU.addRequired<TargetPassConfig>();
70     FunctionPass::getAnalysisUsage(AU);
71   }
72 
73 private:
74   // Check this is a valid gather with correct alignment
75   bool isLegalTypeAndAlignment(unsigned NumElements, unsigned ElemSize,
76                                unsigned Alignment);
77   // Check whether Ptr is hidden behind a bitcast and look through it
78   void lookThroughBitcast(Value *&Ptr);
79   // Check for a getelementptr and deduce base and offsets from it, on success
80   // returning the base directly and the offsets indirectly using the Offsets
81   // argument
82   Value *checkGEP(Value *&Offsets, Type *Ty, Value *Ptr, IRBuilder<> &Builder);
83   // Compute the scale of this gather/scatter instruction
84   int computeScale(unsigned GEPElemSize, unsigned MemoryElemSize);
85 
86   bool lowerGather(IntrinsicInst *I);
87   // Create a gather from a base + vector of offsets
88   Value *tryCreateMaskedGatherOffset(IntrinsicInst *I, Value *Ptr,
89                                      Instruction *&Root, IRBuilder<> &Builder);
90   // Create a gather from a vector of pointers
91   Value *tryCreateMaskedGatherBase(IntrinsicInst *I, Value *Ptr,
92                                    IRBuilder<> &Builder);
93 
94   bool lowerScatter(IntrinsicInst *I);
95   // Create a scatter to a base + vector of offsets
96   Value *tryCreateMaskedScatterOffset(IntrinsicInst *I, Value *Ptr,
97                                       IRBuilder<> &Builder);
98   // Create a scatter to a vector of pointers
99   Value *tryCreateMaskedScatterBase(IntrinsicInst *I, Value *Ptr,
100                                     IRBuilder<> &Builder);
101 };
102 
103 } // end anonymous namespace
104 
105 char MVEGatherScatterLowering::ID = 0;
106 
107 INITIALIZE_PASS(MVEGatherScatterLowering, DEBUG_TYPE,
108                 "MVE gather/scattering lowering pass", false, false)
109 
110 Pass *llvm::createMVEGatherScatterLoweringPass() {
111   return new MVEGatherScatterLowering();
112 }
113 
114 bool MVEGatherScatterLowering::isLegalTypeAndAlignment(unsigned NumElements,
115                                                        unsigned ElemSize,
116                                                        unsigned Alignment) {
117   if (((NumElements == 4 &&
118         (ElemSize == 32 || ElemSize == 16 || ElemSize == 8)) ||
119        (NumElements == 8 && (ElemSize == 16 || ElemSize == 8)) ||
120        (NumElements == 16 && ElemSize == 8)) &&
121       ElemSize / 8 <= Alignment)
122     return true;
123   LLVM_DEBUG(dbgs() << "masked gathers/scatters: instruction does not have "
124                     << "valid alignment or vector type \n");
125   return false;
126 }
127 
128 Value *MVEGatherScatterLowering::checkGEP(Value *&Offsets, Type *Ty, Value *Ptr,
129                                           IRBuilder<> &Builder) {
130   GetElementPtrInst *GEP = dyn_cast<GetElementPtrInst>(Ptr);
131   if (!GEP) {
132     LLVM_DEBUG(
133         dbgs() << "masked gathers/scatters: no getelementpointer found\n");
134     return nullptr;
135   }
136   LLVM_DEBUG(dbgs() << "masked gathers/scatters: getelementpointer found."
137                     << " Looking at intrinsic for base + vector of offsets\n");
138   Value *GEPPtr = GEP->getPointerOperand();
139   if (GEPPtr->getType()->isVectorTy()) {
140     return nullptr;
141   }
142   if (GEP->getNumOperands() != 2) {
143     LLVM_DEBUG(dbgs() << "masked gathers/scatters: getelementptr with too many"
144                       << " operands. Expanding.\n");
145     return nullptr;
146   }
147   Offsets = GEP->getOperand(1);
148   // Paranoid check whether the number of parallel lanes is the same
149   assert(Ty->getVectorNumElements() ==
150          Offsets->getType()->getVectorNumElements());
151   // Only <N x i32> offsets can be integrated into an arm gather, any smaller
152   // type would have to be sign extended by the gep - and arm gathers can only
153   // zero extend. Additionally, the offsets do have to originate from a zext of
154   // a vector with element types smaller or equal the type of the gather we're
155   // looking at
156   if (Offsets->getType()->getScalarSizeInBits() != 32)
157     return nullptr;
158   if (ZExtInst *ZextOffs = dyn_cast<ZExtInst>(Offsets))
159     Offsets = ZextOffs->getOperand(0);
160   else if (!(Offsets->getType()->getVectorNumElements() == 4 &&
161              Offsets->getType()->getScalarSizeInBits() == 32))
162     return nullptr;
163 
164   if (Ty != Offsets->getType()) {
165     if ((Ty->getScalarSizeInBits() <
166          Offsets->getType()->getScalarSizeInBits())) {
167       LLVM_DEBUG(dbgs() << "masked gathers/scatters: no correct offset type."
168                         << " Can't create intrinsic.\n");
169       return nullptr;
170     } else {
171       Offsets = Builder.CreateZExt(
172           Offsets, VectorType::getInteger(cast<VectorType>(Ty)));
173     }
174   }
175   // If none of the checks failed, return the gep's base pointer
176   LLVM_DEBUG(dbgs() << "masked gathers/scatters: found correct offsets\n");
177   return GEPPtr;
178 }
179 
180 void MVEGatherScatterLowering::lookThroughBitcast(Value *&Ptr) {
181   // Look through bitcast instruction if #elements is the same
182   if (auto *BitCast = dyn_cast<BitCastInst>(Ptr)) {
183     Type *BCTy = BitCast->getType();
184     Type *BCSrcTy = BitCast->getOperand(0)->getType();
185     if (BCTy->getVectorNumElements() == BCSrcTy->getVectorNumElements()) {
186       LLVM_DEBUG(
187           dbgs() << "masked gathers/scatters: looking through bitcast\n");
188       Ptr = BitCast->getOperand(0);
189     }
190   }
191 }
192 
193 int MVEGatherScatterLowering::computeScale(unsigned GEPElemSize,
194                                            unsigned MemoryElemSize) {
195   // This can be a 32bit load/store scaled by 4, a 16bit load/store scaled by 2,
196   // or a 8bit, 16bit or 32bit load/store scaled by 1
197   if (GEPElemSize == 32 && MemoryElemSize == 32)
198     return 2;
199   else if (GEPElemSize == 16 && MemoryElemSize == 16)
200     return 1;
201   else if (GEPElemSize == 8)
202     return 0;
203   LLVM_DEBUG(dbgs() << "masked gathers/scatters: incorrect scale. Can't "
204                     << "create intrinsic\n");
205   return -1;
206 }
207 
208 bool MVEGatherScatterLowering::lowerGather(IntrinsicInst *I) {
209   using namespace PatternMatch;
210   LLVM_DEBUG(dbgs() << "masked gathers: checking transform preconditions\n");
211 
212   // @llvm.masked.gather.*(Ptrs, alignment, Mask, Src0)
213   // Attempt to turn the masked gather in I into a MVE intrinsic
214   // Potentially optimising the addressing modes as we do so.
215   Type *Ty = I->getType();
216   Value *Ptr = I->getArgOperand(0);
217   unsigned Alignment = cast<ConstantInt>(I->getArgOperand(1))->getZExtValue();
218   Value *Mask = I->getArgOperand(2);
219   Value *PassThru = I->getArgOperand(3);
220 
221   if (!isLegalTypeAndAlignment(Ty->getVectorNumElements(),
222                                Ty->getScalarSizeInBits(), Alignment))
223     return false;
224   lookThroughBitcast(Ptr);
225   assert(Ptr->getType()->isVectorTy() && "Unexpected pointer type");
226 
227   IRBuilder<> Builder(I->getContext());
228   Builder.SetInsertPoint(I);
229   Builder.SetCurrentDebugLocation(I->getDebugLoc());
230 
231   Instruction *Root = I;
232   Value *Load = tryCreateMaskedGatherOffset(I, Ptr, Root, Builder);
233   if (!Load)
234     Load = tryCreateMaskedGatherBase(I, Ptr, Builder);
235   if (!Load)
236     return false;
237 
238   if (!isa<UndefValue>(PassThru) && !match(PassThru, m_Zero())) {
239     LLVM_DEBUG(dbgs() << "masked gathers: found non-trivial passthru - "
240                       << "creating select\n");
241     Load = Builder.CreateSelect(Mask, Load, PassThru);
242   }
243 
244   Root->replaceAllUsesWith(Load);
245   Root->eraseFromParent();
246   if (Root != I)
247     // If this was an extending gather, we need to get rid of the sext/zext
248     // sext/zext as well as of the gather itself
249     I->eraseFromParent();
250   LLVM_DEBUG(dbgs() << "masked gathers: successfully built masked gather\n");
251   return true;
252 }
253 
254 Value *MVEGatherScatterLowering::tryCreateMaskedGatherBase(
255     IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) {
256   using namespace PatternMatch;
257   Type *Ty = I->getType();
258   LLVM_DEBUG(dbgs() << "masked gathers: loading from vector of pointers\n");
259   if (Ty->getVectorNumElements() != 4 || Ty->getScalarSizeInBits() != 32)
260     // Can't build an intrinsic for this
261     return nullptr;
262   Value *Mask = I->getArgOperand(2);
263   if (match(Mask, m_One()))
264     return Builder.CreateIntrinsic(Intrinsic::arm_mve_vldr_gather_base,
265                                    {Ty, Ptr->getType()},
266                                    {Ptr, Builder.getInt32(0)});
267   else
268     return Builder.CreateIntrinsic(
269         Intrinsic::arm_mve_vldr_gather_base_predicated,
270         {Ty, Ptr->getType(), Mask->getType()},
271         {Ptr, Builder.getInt32(0), Mask});
272 }
273 
274 Value *MVEGatherScatterLowering::tryCreateMaskedGatherOffset(
275     IntrinsicInst *I, Value *Ptr, Instruction *&Root, IRBuilder<> &Builder) {
276   using namespace PatternMatch;
277 
278   Type *OriginalTy = I->getType();
279   Type *ResultTy = OriginalTy;
280 
281   unsigned Unsigned = 1;
282   // The size of the gather was already checked in isLegalTypeAndAlignment;
283   // if it was not a full vector width an appropriate extend should follow.
284   auto *Extend = Root;
285   if (OriginalTy->getPrimitiveSizeInBits() < 128) {
286     // Only transform gathers with exactly one use
287     if (!I->hasOneUse())
288       return nullptr;
289 
290     // The correct root to replace is the not the CallInst itself, but the
291     // instruction which extends it
292     Extend = cast<Instruction>(*I->users().begin());
293     if (isa<SExtInst>(Extend)) {
294       Unsigned = 0;
295     } else if (!isa<ZExtInst>(Extend)) {
296       LLVM_DEBUG(dbgs() << "masked gathers: extend needed but not provided. "
297                         << "Expanding\n");
298       return nullptr;
299     }
300     LLVM_DEBUG(dbgs() << "masked gathers: found an extending gather\n");
301     ResultTy = Extend->getType();
302     // The final size of the gather must be a full vector width
303     if (ResultTy->getPrimitiveSizeInBits() != 128) {
304       LLVM_DEBUG(dbgs() << "masked gathers: extending from the wrong type. "
305                         << "Expanding\n");
306       return nullptr;
307     }
308   }
309 
310   Value *Offsets;
311   Value *BasePtr = checkGEP(Offsets, ResultTy, Ptr, Builder);
312   if (!BasePtr)
313     return nullptr;
314 
315   int Scale = computeScale(
316       BasePtr->getType()->getPointerElementType()->getPrimitiveSizeInBits(),
317       OriginalTy->getScalarSizeInBits());
318   if (Scale == -1)
319     return nullptr;
320   Root = Extend;
321 
322   Value *Mask = I->getArgOperand(2);
323   if (!match(Mask, m_One()))
324     return Builder.CreateIntrinsic(
325         Intrinsic::arm_mve_vldr_gather_offset_predicated,
326         {ResultTy, BasePtr->getType(), Offsets->getType(), Mask->getType()},
327         {BasePtr, Offsets, Builder.getInt32(OriginalTy->getScalarSizeInBits()),
328          Builder.getInt32(Scale), Builder.getInt32(Unsigned), Mask});
329   else
330     return Builder.CreateIntrinsic(
331         Intrinsic::arm_mve_vldr_gather_offset,
332         {ResultTy, BasePtr->getType(), Offsets->getType()},
333         {BasePtr, Offsets, Builder.getInt32(OriginalTy->getScalarSizeInBits()),
334          Builder.getInt32(Scale), Builder.getInt32(Unsigned)});
335 }
336 
337 bool MVEGatherScatterLowering::lowerScatter(IntrinsicInst *I) {
338   using namespace PatternMatch;
339   LLVM_DEBUG(dbgs() << "masked scatters: checking transform preconditions\n");
340 
341   // @llvm.masked.scatter.*(data, ptrs, alignment, mask)
342   // Attempt to turn the masked scatter in I into a MVE intrinsic
343   // Potentially optimising the addressing modes as we do so.
344   Value *Input = I->getArgOperand(0);
345   Value *Ptr = I->getArgOperand(1);
346   unsigned Alignment = cast<ConstantInt>(I->getArgOperand(2))->getZExtValue();
347   Type *Ty = Input->getType();
348 
349   if (!isLegalTypeAndAlignment(Ty->getVectorNumElements(),
350                                Ty->getScalarSizeInBits(), Alignment))
351     return false;
352   lookThroughBitcast(Ptr);
353   assert(Ptr->getType()->isVectorTy() && "Unexpected pointer type");
354 
355   IRBuilder<> Builder(I->getContext());
356   Builder.SetInsertPoint(I);
357   Builder.SetCurrentDebugLocation(I->getDebugLoc());
358 
359   Value *Store = tryCreateMaskedScatterOffset(I, Ptr, Builder);
360   if (!Store)
361     Store = tryCreateMaskedScatterBase(I, Ptr, Builder);
362   if (!Store)
363     return false;
364 
365   LLVM_DEBUG(dbgs() << "masked scatters: successfully built masked scatter\n");
366   I->replaceAllUsesWith(Store);
367   I->eraseFromParent();
368   return true;
369 }
370 
371 Value *MVEGatherScatterLowering::tryCreateMaskedScatterBase(
372     IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) {
373   using namespace PatternMatch;
374   Value *Input = I->getArgOperand(0);
375   Value *Mask = I->getArgOperand(3);
376   Type *Ty = Input->getType();
377   // Only QR variants allow truncating
378   if (!(Ty->getVectorNumElements() == 4 && Ty->getScalarSizeInBits() == 32)) {
379     // Can't build an intrinsic for this
380     return nullptr;
381   }
382   //  int_arm_mve_vstr_scatter_base(_predicated) addr, offset, data(, mask)
383   LLVM_DEBUG(dbgs() << "masked scatters: storing to a vector of pointers\n");
384   if (match(Mask, m_One()))
385     return Builder.CreateIntrinsic(Intrinsic::arm_mve_vstr_scatter_base,
386                                    {Ptr->getType(), Input->getType()},
387                                    {Ptr, Builder.getInt32(0), Input});
388   else
389     return Builder.CreateIntrinsic(
390         Intrinsic::arm_mve_vstr_scatter_base_predicated,
391         {Ptr->getType(), Input->getType(), Mask->getType()},
392         {Ptr, Builder.getInt32(0), Input, Mask});
393 }
394 
395 Value *MVEGatherScatterLowering::tryCreateMaskedScatterOffset(
396     IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) {
397   using namespace PatternMatch;
398   Value *Input = I->getArgOperand(0);
399   Value *Mask = I->getArgOperand(3);
400   Type *InputTy = Input->getType();
401   Type *MemoryTy = InputTy;
402   LLVM_DEBUG(dbgs() << "masked scatters: getelementpointer found. Storing"
403                     << " to base + vector of offsets\n");
404   // If the input has been truncated, try to integrate that trunc into the
405   // scatter instruction (we don't care about alignment here)
406   if (TruncInst *Trunc = dyn_cast<TruncInst>(Input)) {
407     Value *PreTrunc = Trunc->getOperand(0);
408     Type *PreTruncTy = PreTrunc->getType();
409     if (PreTruncTy->getPrimitiveSizeInBits() == 128) {
410       Input = PreTrunc;
411       InputTy = PreTruncTy;
412     }
413   }
414   if (InputTy->getPrimitiveSizeInBits() != 128) {
415     LLVM_DEBUG(
416         dbgs() << "masked scatters: cannot create scatters for non-standard"
417                << " input types. Expanding.\n");
418     return nullptr;
419   }
420 
421   Value *Offsets;
422   Value *BasePtr = checkGEP(Offsets, InputTy, Ptr, Builder);
423   if (!BasePtr)
424     return nullptr;
425   int Scale = computeScale(
426       BasePtr->getType()->getPointerElementType()->getPrimitiveSizeInBits(),
427       MemoryTy->getScalarSizeInBits());
428   if (Scale == -1)
429     return nullptr;
430 
431   if (!match(Mask, m_One()))
432     return Builder.CreateIntrinsic(
433         Intrinsic::arm_mve_vstr_scatter_offset_predicated,
434         {BasePtr->getType(), Offsets->getType(), Input->getType(),
435          Mask->getType()},
436         {BasePtr, Offsets, Input,
437          Builder.getInt32(MemoryTy->getScalarSizeInBits()),
438          Builder.getInt32(Scale), Mask});
439   else
440     return Builder.CreateIntrinsic(
441         Intrinsic::arm_mve_vstr_scatter_offset,
442         {BasePtr->getType(), Offsets->getType(), Input->getType()},
443         {BasePtr, Offsets, Input,
444          Builder.getInt32(MemoryTy->getScalarSizeInBits()),
445          Builder.getInt32(Scale)});
446 }
447 
448 bool MVEGatherScatterLowering::runOnFunction(Function &F) {
449   if (!EnableMaskedGatherScatters)
450     return false;
451   auto &TPC = getAnalysis<TargetPassConfig>();
452   auto &TM = TPC.getTM<TargetMachine>();
453   auto *ST = &TM.getSubtarget<ARMSubtarget>(F);
454   if (!ST->hasMVEIntegerOps())
455     return false;
456   SmallVector<IntrinsicInst *, 4> Gathers;
457   SmallVector<IntrinsicInst *, 4> Scatters;
458   for (BasicBlock &BB : F) {
459     for (Instruction &I : BB) {
460       IntrinsicInst *II = dyn_cast<IntrinsicInst>(&I);
461       if (II && II->getIntrinsicID() == Intrinsic::masked_gather)
462         Gathers.push_back(II);
463       else if (II && II->getIntrinsicID() == Intrinsic::masked_scatter)
464         Scatters.push_back(II);
465     }
466   }
467 
468   bool Changed = false;
469   for (IntrinsicInst *I : Gathers)
470     Changed |= lowerGather(I);
471   for (IntrinsicInst *I : Scatters)
472     Changed |= lowerScatter(I);
473 
474   return Changed;
475 }
476