1 //===-- MVETPAndVPTOptimisationsPass.cpp ----------------------------------===//
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 /// \file This pass does a few optimisations related to Tail predicated loops
10 /// and MVE VPT blocks before register allocation is performed. For VPT blocks
11 /// the goal is to maximize the sizes of the blocks that will be created by the
12 /// MVE VPT Block Insertion pass (which runs after register allocation). For
13 /// tail predicated loops we transform the loop into something that will
14 /// hopefully make the backend ARMLowOverheadLoops pass's job easier.
15 ///
16 //===----------------------------------------------------------------------===//
17 
18 #include "ARM.h"
19 #include "ARMSubtarget.h"
20 #include "MCTargetDesc/ARMBaseInfo.h"
21 #include "MVETailPredUtils.h"
22 #include "Thumb2InstrInfo.h"
23 #include "llvm/ADT/SmallVector.h"
24 #include "llvm/CodeGen/MachineBasicBlock.h"
25 #include "llvm/CodeGen/MachineDominators.h"
26 #include "llvm/CodeGen/MachineFunction.h"
27 #include "llvm/CodeGen/MachineFunctionPass.h"
28 #include "llvm/CodeGen/MachineInstr.h"
29 #include "llvm/CodeGen/MachineLoopInfo.h"
30 #include "llvm/InitializePasses.h"
31 #include "llvm/Support/Debug.h"
32 #include <cassert>
33 
34 using namespace llvm;
35 
36 #define DEBUG_TYPE "arm-mve-vpt-opts"
37 
38 static cl::opt<bool>
39 MergeEndDec("arm-enable-merge-loopenddec", cl::Hidden,
40     cl::desc("Enable merging Loop End and Dec instructions."),
41     cl::init(true));
42 
43 namespace {
44 class MVETPAndVPTOptimisations : public MachineFunctionPass {
45 public:
46   static char ID;
47   const Thumb2InstrInfo *TII;
48   MachineRegisterInfo *MRI;
49 
50   MVETPAndVPTOptimisations() : MachineFunctionPass(ID) {}
51 
52   bool runOnMachineFunction(MachineFunction &Fn) override;
53 
54   void getAnalysisUsage(AnalysisUsage &AU) const override {
55     AU.addRequired<MachineLoopInfo>();
56     AU.addPreserved<MachineLoopInfo>();
57     AU.addRequired<MachineDominatorTree>();
58     AU.addPreserved<MachineDominatorTree>();
59     MachineFunctionPass::getAnalysisUsage(AU);
60   }
61 
62   StringRef getPassName() const override {
63     return "ARM MVE TailPred and VPT Optimisation Pass";
64   }
65 
66 private:
67   bool LowerWhileLoopStart(MachineLoop *ML);
68   bool MergeLoopEnd(MachineLoop *ML);
69   bool ConvertTailPredLoop(MachineLoop *ML, MachineDominatorTree *DT);
70   MachineInstr &ReplaceRegisterUseWithVPNOT(MachineBasicBlock &MBB,
71                                             MachineInstr &Instr,
72                                             MachineOperand &User,
73                                             Register Target);
74   bool ReduceOldVCCRValueUses(MachineBasicBlock &MBB);
75   bool ReplaceVCMPsByVPNOTs(MachineBasicBlock &MBB);
76   bool ReplaceConstByVPNOTs(MachineBasicBlock &MBB, MachineDominatorTree *DT);
77   bool ConvertVPSEL(MachineBasicBlock &MBB);
78   bool HintDoLoopStartReg(MachineBasicBlock &MBB);
79   MachineInstr *CheckForLRUseInPredecessors(MachineBasicBlock *PreHeader,
80                                             MachineInstr *LoopStart);
81 };
82 
83 char MVETPAndVPTOptimisations::ID = 0;
84 
85 } // end anonymous namespace
86 
87 INITIALIZE_PASS_BEGIN(MVETPAndVPTOptimisations, DEBUG_TYPE,
88                       "ARM MVE TailPred and VPT Optimisations pass", false,
89                       false)
90 INITIALIZE_PASS_DEPENDENCY(MachineLoopInfo)
91 INITIALIZE_PASS_DEPENDENCY(MachineDominatorTree)
92 INITIALIZE_PASS_END(MVETPAndVPTOptimisations, DEBUG_TYPE,
93                     "ARM MVE TailPred and VPT Optimisations pass", false, false)
94 
95 static MachineInstr *LookThroughCOPY(MachineInstr *MI,
96                                      MachineRegisterInfo *MRI) {
97   while (MI && MI->getOpcode() == TargetOpcode::COPY &&
98          MI->getOperand(1).getReg().isVirtual())
99     MI = MRI->getVRegDef(MI->getOperand(1).getReg());
100   return MI;
101 }
102 
103 // Given a loop ML, this attempts to find the t2LoopEnd, t2LoopDec and
104 // corresponding PHI that make up a low overhead loop. Only handles 'do' loops
105 // at the moment, returning a t2DoLoopStart in LoopStart.
106 static bool findLoopComponents(MachineLoop *ML, MachineRegisterInfo *MRI,
107                                MachineInstr *&LoopStart, MachineInstr *&LoopPhi,
108                                MachineInstr *&LoopDec, MachineInstr *&LoopEnd) {
109   MachineBasicBlock *Header = ML->getHeader();
110   MachineBasicBlock *Latch = ML->getLoopLatch();
111   if (!Header || !Latch) {
112     LLVM_DEBUG(dbgs() << "  no Loop Latch or Header\n");
113     return false;
114   }
115 
116   // Find the loop end from the terminators.
117   LoopEnd = nullptr;
118   for (auto &T : Latch->terminators()) {
119     if (T.getOpcode() == ARM::t2LoopEnd && T.getOperand(1).getMBB() == Header) {
120       LoopEnd = &T;
121       break;
122     }
123     if (T.getOpcode() == ARM::t2LoopEndDec &&
124         T.getOperand(2).getMBB() == Header) {
125       LoopEnd = &T;
126       break;
127     }
128   }
129   if (!LoopEnd) {
130     LLVM_DEBUG(dbgs() << "  no LoopEnd\n");
131     return false;
132   }
133   LLVM_DEBUG(dbgs() << "  found loop end: " << *LoopEnd);
134 
135   // Find the dec from the use of the end. There may be copies between
136   // instructions. We expect the loop to loop like:
137   //   $vs = t2DoLoopStart ...
138   // loop:
139   //   $vp = phi [ $vs ], [ $vd ]
140   //   ...
141   //   $vd = t2LoopDec $vp
142   //   ...
143   //   t2LoopEnd $vd, loop
144   if (LoopEnd->getOpcode() == ARM::t2LoopEndDec)
145     LoopDec = LoopEnd;
146   else {
147     LoopDec =
148         LookThroughCOPY(MRI->getVRegDef(LoopEnd->getOperand(0).getReg()), MRI);
149     if (!LoopDec || LoopDec->getOpcode() != ARM::t2LoopDec) {
150       LLVM_DEBUG(dbgs() << "  didn't find LoopDec where we expected!\n");
151       return false;
152     }
153   }
154   LLVM_DEBUG(dbgs() << "  found loop dec: " << *LoopDec);
155 
156   LoopPhi =
157       LookThroughCOPY(MRI->getVRegDef(LoopDec->getOperand(1).getReg()), MRI);
158   if (!LoopPhi || LoopPhi->getOpcode() != TargetOpcode::PHI ||
159       LoopPhi->getNumOperands() != 5 ||
160       (LoopPhi->getOperand(2).getMBB() != Latch &&
161        LoopPhi->getOperand(4).getMBB() != Latch)) {
162     LLVM_DEBUG(dbgs() << "  didn't find PHI where we expected!\n");
163     return false;
164   }
165   LLVM_DEBUG(dbgs() << "  found loop phi: " << *LoopPhi);
166 
167   Register StartReg = LoopPhi->getOperand(2).getMBB() == Latch
168                           ? LoopPhi->getOperand(3).getReg()
169                           : LoopPhi->getOperand(1).getReg();
170   LoopStart = LookThroughCOPY(MRI->getVRegDef(StartReg), MRI);
171   if (!LoopStart || (LoopStart->getOpcode() != ARM::t2DoLoopStart &&
172                      LoopStart->getOpcode() != ARM::t2WhileLoopSetup &&
173                      LoopStart->getOpcode() != ARM::t2WhileLoopStartLR)) {
174     LLVM_DEBUG(dbgs() << "  didn't find Start where we expected!\n");
175     return false;
176   }
177   LLVM_DEBUG(dbgs() << "  found loop start: " << *LoopStart);
178 
179   return true;
180 }
181 
182 static void RevertWhileLoopSetup(MachineInstr *MI, const TargetInstrInfo *TII) {
183   MachineBasicBlock *MBB = MI->getParent();
184   assert(MI->getOpcode() == ARM::t2WhileLoopSetup &&
185          "Only expected a t2WhileLoopSetup in RevertWhileLoopStart!");
186 
187   // Subs
188   MachineInstrBuilder MIB =
189       BuildMI(*MBB, MI, MI->getDebugLoc(), TII->get(ARM::t2SUBri));
190   MIB.add(MI->getOperand(0));
191   MIB.add(MI->getOperand(1));
192   MIB.addImm(0);
193   MIB.addImm(ARMCC::AL);
194   MIB.addReg(ARM::NoRegister);
195   MIB.addReg(ARM::CPSR, RegState::Define);
196 
197   // Attempt to find a t2WhileLoopStart and revert to a t2Bcc.
198   for (MachineInstr &I : MBB->terminators()) {
199     if (I.getOpcode() == ARM::t2WhileLoopStart) {
200       MachineInstrBuilder MIB =
201           BuildMI(*MBB, &I, I.getDebugLoc(), TII->get(ARM::t2Bcc));
202       MIB.add(MI->getOperand(1)); // branch target
203       MIB.addImm(ARMCC::EQ);
204       MIB.addReg(ARM::CPSR);
205       I.eraseFromParent();
206       break;
207     }
208   }
209 
210   MI->eraseFromParent();
211 }
212 
213 // The Hardware Loop insertion and ISel Lowering produce the pseudos for the
214 // start of a while loop:
215 //   %a:gprlr = t2WhileLoopSetup %Cnt
216 //   t2WhileLoopStart %a, %BB
217 // We want to convert those to a single instruction which, like t2LoopEndDec and
218 // t2DoLoopStartTP is both a terminator and produces a value:
219 //   %a:grplr: t2WhileLoopStartLR %Cnt, %BB
220 //
221 // Otherwise if we can't, we revert the loop. t2WhileLoopSetup and
222 // t2WhileLoopStart are not valid past regalloc.
223 bool MVETPAndVPTOptimisations::LowerWhileLoopStart(MachineLoop *ML) {
224   LLVM_DEBUG(dbgs() << "LowerWhileLoopStart on loop "
225                     << ML->getHeader()->getName() << "\n");
226 
227   MachineInstr *LoopEnd, *LoopPhi, *LoopStart, *LoopDec;
228   if (!findLoopComponents(ML, MRI, LoopStart, LoopPhi, LoopDec, LoopEnd))
229     return false;
230 
231   if (LoopStart->getOpcode() != ARM::t2WhileLoopSetup)
232     return false;
233 
234   Register LR = LoopStart->getOperand(0).getReg();
235   auto WLSIt = find_if(MRI->use_nodbg_instructions(LR), [](auto &MI) {
236     return MI.getOpcode() == ARM::t2WhileLoopStart;
237   });
238   if (!MergeEndDec || WLSIt == MRI->use_instr_nodbg_end()) {
239     RevertWhileLoopSetup(LoopStart, TII);
240     RevertLoopDec(LoopStart, TII);
241     RevertLoopEnd(LoopStart, TII);
242     return true;
243   }
244 
245   MachineInstrBuilder MI =
246       BuildMI(*WLSIt->getParent(), *WLSIt, WLSIt->getDebugLoc(),
247               TII->get(ARM::t2WhileLoopStartLR), LR)
248           .add(LoopStart->getOperand(1))
249           .add(WLSIt->getOperand(1));
250   (void)MI;
251   LLVM_DEBUG(dbgs() << "Lowered WhileLoopStart into: " << *MI.getInstr());
252 
253   WLSIt->eraseFromParent();
254   LoopStart->eraseFromParent();
255   return true;
256 }
257 
258 // Return true if this instruction is invalid in a low overhead loop, usually
259 // because it clobbers LR.
260 static bool IsInvalidTPInstruction(MachineInstr &MI) {
261   return MI.isCall() || isLoopStart(MI);
262 }
263 
264 // Starting from PreHeader, search for invalid instructions back until the
265 // LoopStart block is reached. If invalid instructions are found, the loop start
266 // is reverted from a WhileLoopStart to a DoLoopStart on the same loop. Will
267 // return the new DLS LoopStart if updated.
268 MachineInstr *MVETPAndVPTOptimisations::CheckForLRUseInPredecessors(
269     MachineBasicBlock *PreHeader, MachineInstr *LoopStart) {
270   SmallVector<MachineBasicBlock *> Worklist;
271   SmallPtrSet<MachineBasicBlock *, 4> Visited;
272   Worklist.push_back(PreHeader);
273   Visited.insert(LoopStart->getParent());
274 
275   while (!Worklist.empty()) {
276     MachineBasicBlock *MBB = Worklist.pop_back_val();
277     if (Visited.count(MBB))
278       continue;
279 
280     for (MachineInstr &MI : *MBB) {
281       if (!IsInvalidTPInstruction(MI))
282         continue;
283 
284       LLVM_DEBUG(dbgs() << "Found LR use in predecessors, reverting: " << MI);
285 
286       // Create a t2DoLoopStart at the end of the preheader.
287       MachineInstrBuilder MIB =
288           BuildMI(*PreHeader, PreHeader->getFirstTerminator(),
289                   LoopStart->getDebugLoc(), TII->get(ARM::t2DoLoopStart));
290       MIB.add(LoopStart->getOperand(0));
291       MIB.add(LoopStart->getOperand(1));
292 
293       // Make sure to remove the kill flags, to prevent them from being invalid.
294       LoopStart->getOperand(1).setIsKill(false);
295 
296       // Revert the t2WhileLoopStartLR to a CMP and Br.
297       RevertWhileLoopStartLR(LoopStart, TII, ARM::t2Bcc, true);
298       return MIB;
299     }
300 
301     Visited.insert(MBB);
302     for (auto *Pred : MBB->predecessors())
303       Worklist.push_back(Pred);
304   }
305   return LoopStart;
306 }
307 
308 // This function converts loops with t2LoopEnd and t2LoopEnd instructions into
309 // a single t2LoopEndDec instruction. To do that it needs to make sure that LR
310 // will be valid to be used for the low overhead loop, which means nothing else
311 // is using LR (especially calls) and there are no superfluous copies in the
312 // loop. The t2LoopEndDec is a branching terminator that produces a value (the
313 // decrement) around the loop edge, which means we need to be careful that they
314 // will be valid to allocate without any spilling.
315 bool MVETPAndVPTOptimisations::MergeLoopEnd(MachineLoop *ML) {
316   if (!MergeEndDec)
317     return false;
318 
319   LLVM_DEBUG(dbgs() << "MergeLoopEnd on loop " << ML->getHeader()->getName()
320                     << "\n");
321 
322   MachineInstr *LoopEnd, *LoopPhi, *LoopStart, *LoopDec;
323   if (!findLoopComponents(ML, MRI, LoopStart, LoopPhi, LoopDec, LoopEnd))
324     return false;
325 
326   // Check if there is an illegal instruction (a call) in the low overhead loop
327   // and if so revert it now before we get any further. While loops also need to
328   // check the preheaders, but can be reverted to a DLS loop if needed.
329   auto *PreHeader = ML->getLoopPreheader();
330   if (LoopStart->getOpcode() == ARM::t2WhileLoopStartLR && PreHeader)
331     LoopStart = CheckForLRUseInPredecessors(PreHeader, LoopStart);
332 
333   for (MachineBasicBlock *MBB : ML->blocks()) {
334     for (MachineInstr &MI : *MBB) {
335       if (IsInvalidTPInstruction(MI)) {
336         LLVM_DEBUG(dbgs() << "Found LR use in loop, reverting: " << MI);
337         if (LoopStart->getOpcode() == ARM::t2DoLoopStart)
338           RevertDoLoopStart(LoopStart, TII);
339         else
340           RevertWhileLoopStartLR(LoopStart, TII);
341         RevertLoopDec(LoopDec, TII);
342         RevertLoopEnd(LoopEnd, TII);
343         return true;
344       }
345     }
346   }
347 
348   // Remove any copies from the loop, to ensure the phi that remains is both
349   // simpler and contains no extra uses. Because t2LoopEndDec is a terminator
350   // that cannot spill, we need to be careful what remains in the loop.
351   Register PhiReg = LoopPhi->getOperand(0).getReg();
352   Register DecReg = LoopDec->getOperand(0).getReg();
353   Register StartReg = LoopStart->getOperand(0).getReg();
354   // Ensure the uses are expected, and collect any copies we want to remove.
355   SmallVector<MachineInstr *, 4> Copies;
356   auto CheckUsers = [&Copies](Register BaseReg,
357                               ArrayRef<MachineInstr *> ExpectedUsers,
358                               MachineRegisterInfo *MRI) {
359     SmallVector<Register, 4> Worklist;
360     Worklist.push_back(BaseReg);
361     while (!Worklist.empty()) {
362       Register Reg = Worklist.pop_back_val();
363       for (MachineInstr &MI : MRI->use_nodbg_instructions(Reg)) {
364         if (count(ExpectedUsers, &MI))
365           continue;
366         if (MI.getOpcode() != TargetOpcode::COPY ||
367             !MI.getOperand(0).getReg().isVirtual()) {
368           LLVM_DEBUG(dbgs() << "Extra users of register found: " << MI);
369           return false;
370         }
371         Worklist.push_back(MI.getOperand(0).getReg());
372         Copies.push_back(&MI);
373       }
374     }
375     return true;
376   };
377   if (!CheckUsers(PhiReg, {LoopDec}, MRI) ||
378       !CheckUsers(DecReg, {LoopPhi, LoopEnd}, MRI) ||
379       !CheckUsers(StartReg, {LoopPhi}, MRI)) {
380     // Don't leave a t2WhileLoopStartLR without the LoopDecEnd.
381     if (LoopStart->getOpcode() == ARM::t2WhileLoopStartLR) {
382       RevertWhileLoopStartLR(LoopStart, TII);
383       RevertLoopDec(LoopDec, TII);
384       RevertLoopEnd(LoopEnd, TII);
385       return true;
386     }
387     return false;
388   }
389 
390   MRI->constrainRegClass(StartReg, &ARM::GPRlrRegClass);
391   MRI->constrainRegClass(PhiReg, &ARM::GPRlrRegClass);
392   MRI->constrainRegClass(DecReg, &ARM::GPRlrRegClass);
393 
394   if (LoopPhi->getOperand(2).getMBB() == ML->getLoopLatch()) {
395     LoopPhi->getOperand(3).setReg(StartReg);
396     LoopPhi->getOperand(1).setReg(DecReg);
397   } else {
398     LoopPhi->getOperand(1).setReg(StartReg);
399     LoopPhi->getOperand(3).setReg(DecReg);
400   }
401 
402   // Replace the loop dec and loop end as a single instruction.
403   MachineInstrBuilder MI =
404       BuildMI(*LoopEnd->getParent(), *LoopEnd, LoopEnd->getDebugLoc(),
405               TII->get(ARM::t2LoopEndDec), DecReg)
406           .addReg(PhiReg)
407           .add(LoopEnd->getOperand(1));
408   (void)MI;
409   LLVM_DEBUG(dbgs() << "Merged LoopDec and End into: " << *MI.getInstr());
410 
411   LoopDec->eraseFromParent();
412   LoopEnd->eraseFromParent();
413   for (auto *MI : Copies)
414     MI->eraseFromParent();
415   return true;
416 }
417 
418 // Convert t2DoLoopStart to t2DoLoopStartTP if the loop contains VCTP
419 // instructions. This keeps the VCTP count reg operand on the t2DoLoopStartTP
420 // instruction, making the backend ARMLowOverheadLoops passes job of finding the
421 // VCTP operand much simpler.
422 bool MVETPAndVPTOptimisations::ConvertTailPredLoop(MachineLoop *ML,
423                                               MachineDominatorTree *DT) {
424   LLVM_DEBUG(dbgs() << "ConvertTailPredLoop on loop "
425                     << ML->getHeader()->getName() << "\n");
426 
427   // Find some loop components including the LoopEnd/Dec/Start, and any VCTP's
428   // in the loop.
429   MachineInstr *LoopEnd, *LoopPhi, *LoopStart, *LoopDec;
430   if (!findLoopComponents(ML, MRI, LoopStart, LoopPhi, LoopDec, LoopEnd))
431     return false;
432   if (LoopDec != LoopEnd || LoopStart->getOpcode() != ARM::t2DoLoopStart)
433     return false;
434 
435   SmallVector<MachineInstr *, 4> VCTPs;
436   for (MachineBasicBlock *BB : ML->blocks())
437     for (MachineInstr &MI : *BB)
438       if (isVCTP(&MI))
439         VCTPs.push_back(&MI);
440 
441   if (VCTPs.empty()) {
442     LLVM_DEBUG(dbgs() << "  no VCTPs\n");
443     return false;
444   }
445 
446   // Check all VCTPs are the same.
447   MachineInstr *FirstVCTP = *VCTPs.begin();
448   for (MachineInstr *VCTP : VCTPs) {
449     LLVM_DEBUG(dbgs() << "  with VCTP " << *VCTP);
450     if (VCTP->getOpcode() != FirstVCTP->getOpcode() ||
451         VCTP->getOperand(0).getReg() != FirstVCTP->getOperand(0).getReg()) {
452       LLVM_DEBUG(dbgs() << "  VCTP's are not identical\n");
453       return false;
454     }
455   }
456 
457   // Check for the register being used can be setup before the loop. We expect
458   // this to be:
459   //   $vx = ...
460   // loop:
461   //   $vp = PHI [ $vx ], [ $vd ]
462   //   ..
463   //   $vpr = VCTP $vp
464   //   ..
465   //   $vd = t2SUBri $vp, #n
466   //   ..
467   Register CountReg = FirstVCTP->getOperand(1).getReg();
468   if (!CountReg.isVirtual()) {
469     LLVM_DEBUG(dbgs() << "  cannot determine VCTP PHI\n");
470     return false;
471   }
472   MachineInstr *Phi = LookThroughCOPY(MRI->getVRegDef(CountReg), MRI);
473   if (!Phi || Phi->getOpcode() != TargetOpcode::PHI ||
474       Phi->getNumOperands() != 5 ||
475       (Phi->getOperand(2).getMBB() != ML->getLoopLatch() &&
476        Phi->getOperand(4).getMBB() != ML->getLoopLatch())) {
477     LLVM_DEBUG(dbgs() << "  cannot determine VCTP Count\n");
478     return false;
479   }
480   CountReg = Phi->getOperand(2).getMBB() == ML->getLoopLatch()
481                  ? Phi->getOperand(3).getReg()
482                  : Phi->getOperand(1).getReg();
483 
484   // Replace the t2DoLoopStart with the t2DoLoopStartTP, move it to the end of
485   // the preheader and add the new CountReg to it. We attempt to place it late
486   // in the preheader, but may need to move that earlier based on uses.
487   MachineBasicBlock *MBB = LoopStart->getParent();
488   MachineBasicBlock::iterator InsertPt = MBB->getFirstTerminator();
489   for (MachineInstr &Use :
490        MRI->use_instructions(LoopStart->getOperand(0).getReg()))
491     if ((InsertPt != MBB->end() && !DT->dominates(&*InsertPt, &Use)) ||
492         !DT->dominates(ML->getHeader(), Use.getParent())) {
493       LLVM_DEBUG(dbgs() << "  InsertPt could not be a terminator!\n");
494       return false;
495     }
496 
497   MachineInstrBuilder MI = BuildMI(*MBB, InsertPt, LoopStart->getDebugLoc(),
498                                    TII->get(ARM::t2DoLoopStartTP))
499                                .add(LoopStart->getOperand(0))
500                                .add(LoopStart->getOperand(1))
501                                .addReg(CountReg);
502   (void)MI;
503   LLVM_DEBUG(dbgs() << "Replacing " << *LoopStart << "  with "
504                     << *MI.getInstr());
505   MRI->constrainRegClass(CountReg, &ARM::rGPRRegClass);
506   LoopStart->eraseFromParent();
507 
508   return true;
509 }
510 
511 // Returns true if Opcode is any VCMP Opcode.
512 static bool IsVCMP(unsigned Opcode) { return VCMPOpcodeToVPT(Opcode) != 0; }
513 
514 // Returns true if a VCMP with this Opcode can have its operands swapped.
515 // There is 2 kind of VCMP that can't have their operands swapped: Float VCMPs,
516 // and VCMPr instructions (since the r is always on the right).
517 static bool CanHaveSwappedOperands(unsigned Opcode) {
518   switch (Opcode) {
519   default:
520     return true;
521   case ARM::MVE_VCMPf32:
522   case ARM::MVE_VCMPf16:
523   case ARM::MVE_VCMPf32r:
524   case ARM::MVE_VCMPf16r:
525   case ARM::MVE_VCMPi8r:
526   case ARM::MVE_VCMPi16r:
527   case ARM::MVE_VCMPi32r:
528   case ARM::MVE_VCMPu8r:
529   case ARM::MVE_VCMPu16r:
530   case ARM::MVE_VCMPu32r:
531   case ARM::MVE_VCMPs8r:
532   case ARM::MVE_VCMPs16r:
533   case ARM::MVE_VCMPs32r:
534     return false;
535   }
536 }
537 
538 // Returns the CondCode of a VCMP Instruction.
539 static ARMCC::CondCodes GetCondCode(MachineInstr &Instr) {
540   assert(IsVCMP(Instr.getOpcode()) && "Inst must be a VCMP");
541   return ARMCC::CondCodes(Instr.getOperand(3).getImm());
542 }
543 
544 // Returns true if Cond is equivalent to a VPNOT instruction on the result of
545 // Prev. Cond and Prev must be VCMPs.
546 static bool IsVPNOTEquivalent(MachineInstr &Cond, MachineInstr &Prev) {
547   assert(IsVCMP(Cond.getOpcode()) && IsVCMP(Prev.getOpcode()));
548 
549   // Opcodes must match.
550   if (Cond.getOpcode() != Prev.getOpcode())
551     return false;
552 
553   MachineOperand &CondOP1 = Cond.getOperand(1), &CondOP2 = Cond.getOperand(2);
554   MachineOperand &PrevOP1 = Prev.getOperand(1), &PrevOP2 = Prev.getOperand(2);
555 
556   // If the VCMP has the opposite condition with the same operands, we can
557   // replace it with a VPNOT
558   ARMCC::CondCodes ExpectedCode = GetCondCode(Cond);
559   ExpectedCode = ARMCC::getOppositeCondition(ExpectedCode);
560   if (ExpectedCode == GetCondCode(Prev))
561     if (CondOP1.isIdenticalTo(PrevOP1) && CondOP2.isIdenticalTo(PrevOP2))
562       return true;
563   // Check again with operands swapped if possible
564   if (!CanHaveSwappedOperands(Cond.getOpcode()))
565     return false;
566   ExpectedCode = ARMCC::getSwappedCondition(ExpectedCode);
567   return ExpectedCode == GetCondCode(Prev) && CondOP1.isIdenticalTo(PrevOP2) &&
568          CondOP2.isIdenticalTo(PrevOP1);
569 }
570 
571 // Returns true if Instr writes to VCCR.
572 static bool IsWritingToVCCR(MachineInstr &Instr) {
573   if (Instr.getNumOperands() == 0)
574     return false;
575   MachineOperand &Dst = Instr.getOperand(0);
576   if (!Dst.isReg())
577     return false;
578   Register DstReg = Dst.getReg();
579   if (!DstReg.isVirtual())
580     return false;
581   MachineRegisterInfo &RegInfo = Instr.getMF()->getRegInfo();
582   const TargetRegisterClass *RegClass = RegInfo.getRegClassOrNull(DstReg);
583   return RegClass && (RegClass->getID() == ARM::VCCRRegClassID);
584 }
585 
586 // Transforms
587 //    <Instr that uses %A ('User' Operand)>
588 // Into
589 //    %K = VPNOT %Target
590 //    <Instr that uses %K ('User' Operand)>
591 // And returns the newly inserted VPNOT.
592 // This optimization is done in the hopes of preventing spills/reloads of VPR by
593 // reducing the number of VCCR values with overlapping lifetimes.
594 MachineInstr &MVETPAndVPTOptimisations::ReplaceRegisterUseWithVPNOT(
595     MachineBasicBlock &MBB, MachineInstr &Instr, MachineOperand &User,
596     Register Target) {
597   Register NewResult = MRI->createVirtualRegister(MRI->getRegClass(Target));
598 
599   MachineInstrBuilder MIBuilder =
600       BuildMI(MBB, &Instr, Instr.getDebugLoc(), TII->get(ARM::MVE_VPNOT))
601           .addDef(NewResult)
602           .addReg(Target);
603   addUnpredicatedMveVpredNOp(MIBuilder);
604 
605   // Make the user use NewResult instead, and clear its kill flag.
606   User.setReg(NewResult);
607   User.setIsKill(false);
608 
609   LLVM_DEBUG(dbgs() << "  Inserting VPNOT (for spill prevention): ";
610              MIBuilder.getInstr()->dump());
611 
612   return *MIBuilder.getInstr();
613 }
614 
615 // Moves a VPNOT before its first user if an instruction that uses Reg is found
616 // in-between the VPNOT and its user.
617 // Returns true if there is at least one user of the VPNOT in the block.
618 static bool MoveVPNOTBeforeFirstUser(MachineBasicBlock &MBB,
619                                      MachineBasicBlock::iterator Iter,
620                                      Register Reg) {
621   assert(Iter->getOpcode() == ARM::MVE_VPNOT && "Not a VPNOT!");
622   assert(getVPTInstrPredicate(*Iter) == ARMVCC::None &&
623          "The VPNOT cannot be predicated");
624 
625   MachineInstr &VPNOT = *Iter;
626   Register VPNOTResult = VPNOT.getOperand(0).getReg();
627   Register VPNOTOperand = VPNOT.getOperand(1).getReg();
628 
629   // Whether the VPNOT will need to be moved, and whether we found a user of the
630   // VPNOT.
631   bool MustMove = false, HasUser = false;
632   MachineOperand *VPNOTOperandKiller = nullptr;
633   for (; Iter != MBB.end(); ++Iter) {
634     if (MachineOperand *MO =
635             Iter->findRegisterUseOperand(VPNOTOperand, /*isKill*/ true)) {
636       // If we find the operand that kills the VPNOTOperand's result, save it.
637       VPNOTOperandKiller = MO;
638     }
639 
640     if (Iter->findRegisterUseOperandIdx(Reg) != -1) {
641       MustMove = true;
642       continue;
643     }
644 
645     if (Iter->findRegisterUseOperandIdx(VPNOTResult) == -1)
646       continue;
647 
648     HasUser = true;
649     if (!MustMove)
650       break;
651 
652     // Move the VPNOT right before Iter
653     LLVM_DEBUG(dbgs() << "Moving: "; VPNOT.dump(); dbgs() << "  Before: ";
654                Iter->dump());
655     MBB.splice(Iter, &MBB, VPNOT.getIterator());
656     // If we move the instr, and its operand was killed earlier, remove the kill
657     // flag.
658     if (VPNOTOperandKiller)
659       VPNOTOperandKiller->setIsKill(false);
660 
661     break;
662   }
663   return HasUser;
664 }
665 
666 // This optimisation attempts to reduce the number of overlapping lifetimes of
667 // VCCR values by replacing uses of old VCCR values with VPNOTs. For example,
668 // this replaces
669 //    %A:vccr = (something)
670 //    %B:vccr = VPNOT %A
671 //    %Foo = (some op that uses %B)
672 //    %Bar = (some op that uses %A)
673 // With
674 //    %A:vccr = (something)
675 //    %B:vccr = VPNOT %A
676 //    %Foo = (some op that uses %B)
677 //    %TMP2:vccr = VPNOT %B
678 //    %Bar = (some op that uses %A)
679 bool MVETPAndVPTOptimisations::ReduceOldVCCRValueUses(MachineBasicBlock &MBB) {
680   MachineBasicBlock::iterator Iter = MBB.begin(), End = MBB.end();
681   SmallVector<MachineInstr *, 4> DeadInstructions;
682   bool Modified = false;
683 
684   while (Iter != End) {
685     Register VCCRValue, OppositeVCCRValue;
686     // The first loop looks for 2 unpredicated instructions:
687     //    %A:vccr = (instr)     ; A is stored in VCCRValue
688     //    %B:vccr = VPNOT %A    ; B is stored in OppositeVCCRValue
689     for (; Iter != End; ++Iter) {
690       // We're only interested in unpredicated instructions that write to VCCR.
691       if (!IsWritingToVCCR(*Iter) ||
692           getVPTInstrPredicate(*Iter) != ARMVCC::None)
693         continue;
694       Register Dst = Iter->getOperand(0).getReg();
695 
696       // If we already have a VCCRValue, and this is a VPNOT on VCCRValue, we've
697       // found what we were looking for.
698       if (VCCRValue && Iter->getOpcode() == ARM::MVE_VPNOT &&
699           Iter->findRegisterUseOperandIdx(VCCRValue) != -1) {
700         // Move the VPNOT closer to its first user if needed, and ignore if it
701         // has no users.
702         if (!MoveVPNOTBeforeFirstUser(MBB, Iter, VCCRValue))
703           continue;
704 
705         OppositeVCCRValue = Dst;
706         ++Iter;
707         break;
708       }
709 
710       // Else, just set VCCRValue.
711       VCCRValue = Dst;
712     }
713 
714     // If the first inner loop didn't find anything, stop here.
715     if (Iter == End)
716       break;
717 
718     assert(VCCRValue && OppositeVCCRValue &&
719            "VCCRValue and OppositeVCCRValue shouldn't be empty if the loop "
720            "stopped before the end of the block!");
721     assert(VCCRValue != OppositeVCCRValue &&
722            "VCCRValue should not be equal to OppositeVCCRValue!");
723 
724     // LastVPNOTResult always contains the same value as OppositeVCCRValue.
725     Register LastVPNOTResult = OppositeVCCRValue;
726 
727     // This second loop tries to optimize the remaining instructions.
728     for (; Iter != End; ++Iter) {
729       bool IsInteresting = false;
730 
731       if (MachineOperand *MO = Iter->findRegisterUseOperand(VCCRValue)) {
732         IsInteresting = true;
733 
734         // - If the instruction is a VPNOT, it can be removed, and we can just
735         //   replace its uses with LastVPNOTResult.
736         // - Else, insert a new VPNOT on LastVPNOTResult to recompute VCCRValue.
737         if (Iter->getOpcode() == ARM::MVE_VPNOT) {
738           Register Result = Iter->getOperand(0).getReg();
739 
740           MRI->replaceRegWith(Result, LastVPNOTResult);
741           DeadInstructions.push_back(&*Iter);
742           Modified = true;
743 
744           LLVM_DEBUG(dbgs()
745                      << "Replacing all uses of '" << printReg(Result)
746                      << "' with '" << printReg(LastVPNOTResult) << "'\n");
747         } else {
748           MachineInstr &VPNOT =
749               ReplaceRegisterUseWithVPNOT(MBB, *Iter, *MO, LastVPNOTResult);
750           Modified = true;
751 
752           LastVPNOTResult = VPNOT.getOperand(0).getReg();
753           std::swap(VCCRValue, OppositeVCCRValue);
754 
755           LLVM_DEBUG(dbgs() << "Replacing use of '" << printReg(VCCRValue)
756                             << "' with '" << printReg(LastVPNOTResult)
757                             << "' in instr: " << *Iter);
758         }
759       } else {
760         // If the instr uses OppositeVCCRValue, make it use LastVPNOTResult
761         // instead as they contain the same value.
762         if (MachineOperand *MO =
763                 Iter->findRegisterUseOperand(OppositeVCCRValue)) {
764           IsInteresting = true;
765 
766           // This is pointless if LastVPNOTResult == OppositeVCCRValue.
767           if (LastVPNOTResult != OppositeVCCRValue) {
768             LLVM_DEBUG(dbgs() << "Replacing usage of '"
769                               << printReg(OppositeVCCRValue) << "' with '"
770                               << printReg(LastVPNOTResult) << " for instr: ";
771                        Iter->dump());
772             MO->setReg(LastVPNOTResult);
773             Modified = true;
774           }
775 
776           MO->setIsKill(false);
777         }
778 
779         // If this is an unpredicated VPNOT on
780         // LastVPNOTResult/OppositeVCCRValue, we can act like we inserted it.
781         if (Iter->getOpcode() == ARM::MVE_VPNOT &&
782             getVPTInstrPredicate(*Iter) == ARMVCC::None) {
783           Register VPNOTOperand = Iter->getOperand(1).getReg();
784           if (VPNOTOperand == LastVPNOTResult ||
785               VPNOTOperand == OppositeVCCRValue) {
786             IsInteresting = true;
787 
788             std::swap(VCCRValue, OppositeVCCRValue);
789             LastVPNOTResult = Iter->getOperand(0).getReg();
790           }
791         }
792       }
793 
794       // If this instruction was not interesting, and it writes to VCCR, stop.
795       if (!IsInteresting && IsWritingToVCCR(*Iter))
796         break;
797     }
798   }
799 
800   for (MachineInstr *DeadInstruction : DeadInstructions)
801     DeadInstruction->eraseFromParent();
802 
803   return Modified;
804 }
805 
806 // This optimisation replaces VCMPs with VPNOTs when they are equivalent.
807 bool MVETPAndVPTOptimisations::ReplaceVCMPsByVPNOTs(MachineBasicBlock &MBB) {
808   SmallVector<MachineInstr *, 4> DeadInstructions;
809 
810   // The last VCMP that we have seen and that couldn't be replaced.
811   // This is reset when an instruction that writes to VCCR/VPR is found, or when
812   // a VCMP is replaced with a VPNOT.
813   // We'll only replace VCMPs with VPNOTs when this is not null, and when the
814   // current VCMP is the opposite of PrevVCMP.
815   MachineInstr *PrevVCMP = nullptr;
816   // If we find an instruction that kills the result of PrevVCMP, we save the
817   // operand here to remove the kill flag in case we need to use PrevVCMP's
818   // result.
819   MachineOperand *PrevVCMPResultKiller = nullptr;
820 
821   for (MachineInstr &Instr : MBB.instrs()) {
822     if (PrevVCMP) {
823       if (MachineOperand *MO = Instr.findRegisterUseOperand(
824               PrevVCMP->getOperand(0).getReg(), /*isKill*/ true)) {
825         // If we come accross the instr that kills PrevVCMP's result, record it
826         // so we can remove the kill flag later if we need to.
827         PrevVCMPResultKiller = MO;
828       }
829     }
830 
831     // Ignore predicated instructions.
832     if (getVPTInstrPredicate(Instr) != ARMVCC::None)
833       continue;
834 
835     // Only look at VCMPs
836     if (!IsVCMP(Instr.getOpcode())) {
837       // If the instruction writes to VCCR, forget the previous VCMP.
838       if (IsWritingToVCCR(Instr))
839         PrevVCMP = nullptr;
840       continue;
841     }
842 
843     if (!PrevVCMP || !IsVPNOTEquivalent(Instr, *PrevVCMP)) {
844       PrevVCMP = &Instr;
845       continue;
846     }
847 
848     // The register containing the result of the VCMP that we're going to
849     // replace.
850     Register PrevVCMPResultReg = PrevVCMP->getOperand(0).getReg();
851 
852     // Build a VPNOT to replace the VCMP, reusing its operands.
853     MachineInstrBuilder MIBuilder =
854         BuildMI(MBB, &Instr, Instr.getDebugLoc(), TII->get(ARM::MVE_VPNOT))
855             .add(Instr.getOperand(0))
856             .addReg(PrevVCMPResultReg);
857     addUnpredicatedMveVpredNOp(MIBuilder);
858     LLVM_DEBUG(dbgs() << "Inserting VPNOT (to replace VCMP): ";
859                MIBuilder.getInstr()->dump(); dbgs() << "  Removed VCMP: ";
860                Instr.dump());
861 
862     // If we found an instruction that uses, and kills PrevVCMP's result,
863     // remove the kill flag.
864     if (PrevVCMPResultKiller)
865       PrevVCMPResultKiller->setIsKill(false);
866 
867     // Finally, mark the old VCMP for removal and reset
868     // PrevVCMP/PrevVCMPResultKiller.
869     DeadInstructions.push_back(&Instr);
870     PrevVCMP = nullptr;
871     PrevVCMPResultKiller = nullptr;
872   }
873 
874   for (MachineInstr *DeadInstruction : DeadInstructions)
875     DeadInstruction->eraseFromParent();
876 
877   return !DeadInstructions.empty();
878 }
879 
880 bool MVETPAndVPTOptimisations::ReplaceConstByVPNOTs(MachineBasicBlock &MBB,
881                                                MachineDominatorTree *DT) {
882   // Scan through the block, looking for instructions that use constants moves
883   // into VPR that are the negative of one another. These are expected to be
884   // COPY's to VCCRRegClass, from a t2MOVi or t2MOVi16. The last seen constant
885   // mask is kept it or and VPNOT's of it are added or reused as we scan through
886   // the function.
887   unsigned LastVPTImm = 0;
888   Register LastVPTReg = 0;
889   SmallSet<MachineInstr *, 4> DeadInstructions;
890 
891   for (MachineInstr &Instr : MBB.instrs()) {
892     // Look for predicated MVE instructions.
893     int PIdx = llvm::findFirstVPTPredOperandIdx(Instr);
894     if (PIdx == -1)
895       continue;
896     Register VPR = Instr.getOperand(PIdx + 1).getReg();
897     if (!VPR.isVirtual())
898       continue;
899 
900     // From that we are looking for an instruction like %11:vccr = COPY %9:rgpr.
901     MachineInstr *Copy = MRI->getVRegDef(VPR);
902     if (!Copy || Copy->getOpcode() != TargetOpcode::COPY ||
903         !Copy->getOperand(1).getReg().isVirtual() ||
904         MRI->getRegClass(Copy->getOperand(1).getReg()) == &ARM::VCCRRegClass) {
905       LastVPTReg = 0;
906       continue;
907     }
908     Register GPR = Copy->getOperand(1).getReg();
909 
910     // Find the Immediate used by the copy.
911     auto getImm = [&](Register GPR) -> unsigned {
912       MachineInstr *Def = MRI->getVRegDef(GPR);
913       if (Def && (Def->getOpcode() == ARM::t2MOVi ||
914                   Def->getOpcode() == ARM::t2MOVi16))
915         return Def->getOperand(1).getImm();
916       return -1U;
917     };
918     unsigned Imm = getImm(GPR);
919     if (Imm == -1U) {
920       LastVPTReg = 0;
921       continue;
922     }
923 
924     unsigned NotImm = ~Imm & 0xffff;
925     if (LastVPTReg != 0 && LastVPTReg != VPR && LastVPTImm == Imm) {
926       Instr.getOperand(PIdx + 1).setReg(LastVPTReg);
927       if (MRI->use_empty(VPR)) {
928         DeadInstructions.insert(Copy);
929         if (MRI->hasOneUse(GPR))
930           DeadInstructions.insert(MRI->getVRegDef(GPR));
931       }
932       LLVM_DEBUG(dbgs() << "Reusing predicate: in  " << Instr);
933     } else if (LastVPTReg != 0 && LastVPTImm == NotImm) {
934       // We have found the not of a previous constant. Create a VPNot of the
935       // earlier predicate reg and use it instead of the copy.
936       Register NewVPR = MRI->createVirtualRegister(&ARM::VCCRRegClass);
937       auto VPNot = BuildMI(MBB, &Instr, Instr.getDebugLoc(),
938                            TII->get(ARM::MVE_VPNOT), NewVPR)
939                        .addReg(LastVPTReg);
940       addUnpredicatedMveVpredNOp(VPNot);
941 
942       // Use the new register and check if the def is now dead.
943       Instr.getOperand(PIdx + 1).setReg(NewVPR);
944       if (MRI->use_empty(VPR)) {
945         DeadInstructions.insert(Copy);
946         if (MRI->hasOneUse(GPR))
947           DeadInstructions.insert(MRI->getVRegDef(GPR));
948       }
949       LLVM_DEBUG(dbgs() << "Adding VPNot: " << *VPNot << "  to replace use at "
950                         << Instr);
951       VPR = NewVPR;
952     }
953 
954     LastVPTImm = Imm;
955     LastVPTReg = VPR;
956   }
957 
958   for (MachineInstr *DI : DeadInstructions)
959     DI->eraseFromParent();
960 
961   return !DeadInstructions.empty();
962 }
963 
964 // Replace VPSEL with a predicated VMOV in blocks with a VCTP. This is a
965 // somewhat blunt approximation to allow tail predicated with vpsel
966 // instructions. We turn a vselect into a VPSEL in ISEL, but they have slightly
967 // different semantics under tail predication. Until that is modelled we just
968 // convert to a VMOVT (via a predicated VORR) instead.
969 bool MVETPAndVPTOptimisations::ConvertVPSEL(MachineBasicBlock &MBB) {
970   bool HasVCTP = false;
971   SmallVector<MachineInstr *, 4> DeadInstructions;
972 
973   for (MachineInstr &MI : MBB.instrs()) {
974     if (isVCTP(&MI)) {
975       HasVCTP = true;
976       continue;
977     }
978 
979     if (!HasVCTP || MI.getOpcode() != ARM::MVE_VPSEL)
980       continue;
981 
982     MachineInstrBuilder MIBuilder =
983         BuildMI(MBB, &MI, MI.getDebugLoc(), TII->get(ARM::MVE_VORR))
984             .add(MI.getOperand(0))
985             .add(MI.getOperand(1))
986             .add(MI.getOperand(1))
987             .addImm(ARMVCC::Then)
988             .add(MI.getOperand(4))
989             .add(MI.getOperand(2));
990     // Silence unused variable warning in release builds.
991     (void)MIBuilder;
992     LLVM_DEBUG(dbgs() << "Replacing VPSEL: "; MI.dump();
993                dbgs() << "     with VMOVT: "; MIBuilder.getInstr()->dump());
994     DeadInstructions.push_back(&MI);
995   }
996 
997   for (MachineInstr *DeadInstruction : DeadInstructions)
998     DeadInstruction->eraseFromParent();
999 
1000   return !DeadInstructions.empty();
1001 }
1002 
1003 // Add a registry allocation hint for t2DoLoopStart to hint it towards LR, as
1004 // the instruction may be removable as a noop.
1005 bool MVETPAndVPTOptimisations::HintDoLoopStartReg(MachineBasicBlock &MBB) {
1006   bool Changed = false;
1007   for (MachineInstr &MI : MBB.instrs()) {
1008     if (MI.getOpcode() != ARM::t2DoLoopStart)
1009       continue;
1010     Register R = MI.getOperand(1).getReg();
1011     MachineFunction *MF = MI.getParent()->getParent();
1012     MF->getRegInfo().setRegAllocationHint(R, ARMRI::RegLR, 0);
1013     Changed = true;
1014   }
1015   return Changed;
1016 }
1017 
1018 bool MVETPAndVPTOptimisations::runOnMachineFunction(MachineFunction &Fn) {
1019   const ARMSubtarget &STI =
1020       static_cast<const ARMSubtarget &>(Fn.getSubtarget());
1021 
1022   if (!STI.isThumb2() || !STI.hasLOB())
1023     return false;
1024 
1025   TII = static_cast<const Thumb2InstrInfo *>(STI.getInstrInfo());
1026   MRI = &Fn.getRegInfo();
1027   MachineLoopInfo *MLI = &getAnalysis<MachineLoopInfo>();
1028   MachineDominatorTree *DT = &getAnalysis<MachineDominatorTree>();
1029 
1030   LLVM_DEBUG(dbgs() << "********** ARM MVE VPT Optimisations **********\n"
1031                     << "********** Function: " << Fn.getName() << '\n');
1032 
1033   bool Modified = false;
1034   for (MachineLoop *ML : MLI->getBase().getLoopsInPreorder()) {
1035     Modified |= LowerWhileLoopStart(ML);
1036     Modified |= MergeLoopEnd(ML);
1037     Modified |= ConvertTailPredLoop(ML, DT);
1038   }
1039 
1040   for (MachineBasicBlock &MBB : Fn) {
1041     Modified |= HintDoLoopStartReg(MBB);
1042     Modified |= ReplaceConstByVPNOTs(MBB, DT);
1043     Modified |= ReplaceVCMPsByVPNOTs(MBB);
1044     Modified |= ReduceOldVCCRValueUses(MBB);
1045     Modified |= ConvertVPSEL(MBB);
1046   }
1047 
1048   LLVM_DEBUG(dbgs() << "**************************************\n");
1049   return Modified;
1050 }
1051 
1052 /// createMVETPAndVPTOptimisationsPass
1053 FunctionPass *llvm::createMVETPAndVPTOptimisationsPass() {
1054   return new MVETPAndVPTOptimisations();
1055 }
1056