1//===-- AMDGPURewriteAGPRCopyMFMA.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 \brief Try to replace MFMA instructions using VGPRs with MFMA
10/// instructions using AGPRs. We expect MFMAs to be selected using VGPRs, and
11/// only use AGPRs if it helps avoid spilling. In this case, the MFMA will have
12/// copies between AGPRs and VGPRs and the AGPR variant of an MFMA pseudo. This
13/// pass will attempt to delete the cross register bank copy and replace the
14/// MFMA opcode.
15///
16/// TODO:
17/// - Handle rewrites of phis. This must be more careful than normal about the
18/// reassignment. We do not want to introduce an AGPR-to-AGPR copy inside of a
19/// loop, so it depends on the exact assignment of the copy.
20///
21/// - Update LiveIntervals incrementally instead of recomputing from scratch
22///
23//===----------------------------------------------------------------------===//
24
25#include "AMDGPU.h"
26#include "GCNSubtarget.h"
27#include "SIMachineFunctionInfo.h"
28#include "SIRegisterInfo.h"
29#include "llvm/ADT/Statistic.h"
30#include "llvm/CodeGen/LiveIntervals.h"
31#include "llvm/CodeGen/LiveRegMatrix.h"
32#include "llvm/CodeGen/LiveStacks.h"
33#include "llvm/CodeGen/MachineDominators.h"
34#include "llvm/CodeGen/MachineFrameInfo.h"
35#include "llvm/CodeGen/MachineFunctionPass.h"
36#include "llvm/CodeGen/RegisterClassInfo.h"
37#include "llvm/CodeGen/SlotIndexes.h"
38#include "llvm/CodeGen/VirtRegMap.h"
39#include "llvm/InitializePasses.h"
40#include "llvm/Support/DebugCounter.h"
41
42using namespace llvm;
43
44#define DEBUG_TYPE "amdgpu-rewrite-agpr-copy-mfma"
45
46DEBUG_COUNTER(RewriteAGPRCopyMFMACounter, DEBUG_TYPE,
47 "Controls which MFMA chains are rewritten to AGPR form");
48
49namespace {
50
51STATISTIC(NumMFMAsRewrittenToAGPR,
52 "Number of MFMA instructions rewritten to use AGPR form");
53
54/// Map from spill slot frame index to list of instructions which reference it.
55using SpillReferenceMap = DenseMap<int, SmallVector<MachineInstr *, 4>>;
56
57class AMDGPURewriteAGPRCopyMFMAImpl {
58 MachineFunction &MF;
59 const GCNSubtarget &ST;
60 const SIInstrInfo &TII;
61 const SIRegisterInfo &TRI;
62 MachineRegisterInfo &MRI;
63 VirtRegMap &VRM;
64 LiveRegMatrix &LRM;
65 LiveIntervals &LIS;
66 LiveStacks &LSS;
67 const RegisterClassInfo &RegClassInfo;
68 MachineDominatorTree &MDT;
69
70 bool attemptReassignmentsToAGPR(SmallSetVector<Register, 4> &InterferingRegs,
71 MCPhysReg PrefPhysReg) const;
72
73public:
74 AMDGPURewriteAGPRCopyMFMAImpl(MachineFunction &MF, VirtRegMap &VRM,
75 LiveRegMatrix &LRM, LiveIntervals &LIS,
76 LiveStacks &LSS,
77 const RegisterClassInfo &RegClassInfo,
78 MachineDominatorTree &MDT)
79 : MF(MF), ST(MF.getSubtarget<GCNSubtarget>()), TII(*ST.getInstrInfo()),
80 TRI(*ST.getRegisterInfo()), MRI(MF.getRegInfo()), VRM(VRM), LRM(LRM),
81 LIS(LIS), LSS(LSS), RegClassInfo(RegClassInfo), MDT(MDT) {}
82
83 bool isRewriteCandidate(const MachineInstr &MI) const {
84 return TII.isMAI(MI) && AMDGPU::getAGPRFormOp(Opcode: MI.getOpcode()) != -1;
85 }
86
87 /// Find AV_* registers assigned to AGPRs (or virtual registers which were
88 /// already required to be AGPR).
89 ///
90 /// \return the assigned physical register that \p VReg is assigned to if it
91 /// is an AGPR, otherwise MCRegister().
92 MCRegister getAssignedAGPR(Register VReg) const {
93 MCRegister PhysReg = VRM.getPhys(virtReg: VReg);
94 if (!PhysReg)
95 return MCRegister();
96
97 // If this is an AV register, we have to check if the actual assignment is
98 // to an AGPR
99 const TargetRegisterClass *AssignedRC = TRI.getPhysRegBaseClass(Reg: PhysReg);
100 return TRI.isAGPRClass(RC: AssignedRC) ? PhysReg : MCRegister();
101 }
102
103 bool tryReassigningMFMAChain(MachineInstr &MFMA, Register MFMAHintReg,
104 MCPhysReg PhysRegHint) const;
105
106 /// Compute the register class constraints based on the uses of \p Reg,
107 /// excluding MFMA uses from which can be rewritten to change the register
108 /// class constraint. MFMA scale operands need to be constraint checked.
109 /// This should be nearly identical to MachineRegisterInfo::recomputeRegClass.
110
111 /// \p RewriteCandidates will collect the set of MFMA instructions that need
112 /// to have the opcode mutated to perform the replacement.
113 ///
114 /// \p RewriteRegs will accumulate the set of register used by those MFMAs
115 /// that need to have the register classes adjusted.
116 bool recomputeRegClassExceptRewritable(
117 Register Reg, SmallVectorImpl<MachineInstr *> &RewriteCandidates,
118 SmallSetVector<Register, 4> &RewriteRegs) const;
119
120 bool tryFoldCopiesToAGPR(Register VReg, MCRegister AssignedAGPR) const;
121 bool tryFoldCopiesFromAGPR(Register VReg, MCRegister AssignedAGPR) const;
122
123 /// Replace spill instruction \p SpillMI which loads/stores from/to \p SpillFI
124 /// with a COPY to the replacement register value \p VReg.
125 void replaceSpillWithCopyToVReg(MachineInstr &SpillMI, int SpillFI,
126 Register VReg) const;
127
128 /// Create a map from frame index to use instructions for spills. If a use of
129 /// the frame index does not consist only of spill instructions, it will not
130 /// be included in the map.
131 void collectSpillIndexUses(ArrayRef<LiveInterval *> StackIntervals,
132 SpillReferenceMap &Map) const;
133
134 /// Return true if the reload \p LoadMI of the stack slot with live interval
135 /// \p SlotLI is jointly dominated by the slot's spill stores, i.e. every path
136 /// from the entry block to the load passes through a store to the slot before
137 /// the load. \p StoreFreeReachable is the set of blocks reachable from the
138 /// entry block without passing through any store block for the slot.
139 bool isLoadJointlyDominatedByStores(
140 const MachineInstr &LoadMI, const LiveInterval &SlotLI,
141 const SmallPtrSetImpl<MachineBasicBlock *> &StoreFreeReachable) const;
142
143 /// Attempt to unspill VGPRs by finding a free register and replacing the
144 /// spill instructions with copies.
145 void eliminateSpillsOfReassignedVGPRs() const;
146
147 bool run(MachineFunction &MF) const;
148};
149
150bool AMDGPURewriteAGPRCopyMFMAImpl::recomputeRegClassExceptRewritable(
151 Register StartReg, SmallVectorImpl<MachineInstr *> &RewriteCandidates,
152 SmallSetVector<Register, 4> &RewriteRegs) const {
153 SmallVector<Register, 8> Worklist = {StartReg};
154
155 // Recursively visit all transitive MFMA users
156 while (!Worklist.empty()) {
157 Register Reg = Worklist.pop_back_val();
158 const TargetRegisterClass *OldRC = MRI.getRegClass(Reg);
159
160 // Inflate to the equivalent AV_* class.
161 const TargetRegisterClass *NewRC = TRI.getLargestLegalSuperClass(RC: OldRC, MF);
162 if (OldRC == NewRC)
163 return false;
164
165 // Accumulate constraints from all uses.
166 for (MachineOperand &MO : MRI.reg_nodbg_operands(Reg)) {
167 // Apply the effect of the given operand to NewRC.
168 MachineInstr *MI = MO.getParent();
169
170 // We can swap the classes of dst + src2 as a pair to AGPR, so ignore the
171 // effects of rewrite candidates. It just so happens that we can use
172 // either AGPR or VGPR in src0/src1. We still need to check constraint
173 // effects for scale variant, which does not allow AGPR.
174 if (isRewriteCandidate(MI: *MI)) {
175 int AGPROp = AMDGPU::getAGPRFormOp(Opcode: MI->getOpcode());
176 const MCInstrDesc &AGPRDesc = TII.get(Opcode: AGPROp);
177 const TargetRegisterClass *NewRC =
178 TII.getRegClass(MCID: AGPRDesc, OpNum: MO.getOperandNo());
179 if (!TRI.hasAGPRs(RC: NewRC))
180 return false;
181
182 const MachineOperand *VDst =
183 TII.getNamedOperand(MI&: *MI, OperandName: AMDGPU::OpName::vdst);
184 const MachineOperand *Src2 =
185 TII.getNamedOperand(MI&: *MI, OperandName: AMDGPU::OpName::src2);
186 for (const MachineOperand *Op : {VDst, Src2}) {
187 if (!Op->isReg())
188 continue;
189
190 Register OtherReg = Op->getReg();
191 if (OtherReg.isPhysical())
192 return false;
193
194 if (OtherReg != Reg && RewriteRegs.insert(X: OtherReg))
195 Worklist.push_back(Elt: OtherReg);
196 }
197
198 if (!is_contained(Range&: RewriteCandidates, Element: MI)) {
199 LLVM_DEBUG({
200 Register VDstPhysReg = VRM.getPhys(VDst->getReg());
201 dbgs() << "Attempting to replace VGPR MFMA with AGPR version:"
202 << " Dst=[" << printReg(VDst->getReg()) << " => "
203 << printReg(VDstPhysReg, &TRI);
204
205 if (Src2->isReg()) {
206 Register Src2PhysReg = VRM.getPhys(Src2->getReg());
207 dbgs() << "], Src2=[" << printReg(Src2->getReg(), &TRI) << " => "
208 << printReg(Src2PhysReg, &TRI);
209 }
210
211 dbgs() << "]: " << MI;
212 });
213
214 RewriteCandidates.push_back(Elt: MI);
215 }
216
217 continue;
218 }
219
220 unsigned OpNo = &MO - &MI->getOperand(i: 0);
221 NewRC = MI->getRegClassConstraintEffect(OpIdx: OpNo, CurRC: NewRC, TII: &TII, TRI: &TRI);
222 if (!NewRC || NewRC == OldRC) {
223 LLVM_DEBUG(dbgs() << "User of " << printReg(Reg, &TRI)
224 << " cannot be reassigned to "
225 << (NewRC ? TRI.getRegClassName(NewRC) : "NULL")
226 << ": " << *MI);
227 return false;
228 }
229 }
230 }
231
232 return true;
233}
234
235bool AMDGPURewriteAGPRCopyMFMAImpl::tryReassigningMFMAChain(
236 MachineInstr &MFMA, Register MFMAHintReg, MCPhysReg PhysRegHint) const {
237 // src2 and dst have the same physical class constraint; try to preserve
238 // the original src2 subclass if one were to exist.
239 SmallVector<MachineInstr *, 4> RewriteCandidates = {&MFMA};
240 SmallSetVector<Register, 4> RewriteRegs;
241
242 // Make sure we reassign the MFMA we found the copy from first. We want
243 // to ensure dst ends up in the physreg we were originally copying to.
244 RewriteRegs.insert(X: MFMAHintReg);
245
246 // We've found av = COPY (MFMA) (or MFMA (v = COPY av)) and need to verify
247 // that we can trivially rewrite src2 to use the new AGPR. If we can't
248 // trivially replace it, we're going to induce as many copies as we would have
249 // emitted in the first place, as well as need to assign another register, and
250 // need to figure out where to put them. The live range splitting is smarter
251 // than anything we're doing here, so trust it did something reasonable.
252 //
253 // Note recomputeRegClassExceptRewritable will consider the constraints of
254 // this MFMA's src2 as well as the src2/dst of any transitive MFMA users.
255 if (!recomputeRegClassExceptRewritable(StartReg: MFMAHintReg, RewriteCandidates,
256 RewriteRegs)) {
257 LLVM_DEBUG(dbgs() << "Could not recompute the regclass of dst reg "
258 << printReg(MFMAHintReg, &TRI) << '\n');
259 return false;
260 }
261
262 // If src2 and dst are different registers, we need to also reassign the
263 // input to an available AGPR if it is compatible with all other uses.
264 //
265 // If we can't reassign it, we'd need to introduce a different copy
266 // which is likely worse than the copy we'd be saving.
267 //
268 // It's likely that the MFMA is used in sequence with other MFMAs; if we
269 // cannot migrate the full use/def chain of MFMAs, we would need to
270 // introduce intermediate copies somewhere. So we only make the
271 // transform if all the interfering MFMAs can also be migrated. Collect
272 // the set of rewritable MFMAs and check if we can assign an AGPR at
273 // that point.
274 //
275 // If any of the MFMAs aren't reassignable, we give up and rollback to
276 // the original register assignments.
277
278 using RecoloringStack =
279 SmallVector<std::pair<const LiveInterval *, MCRegister>, 8>;
280 RecoloringStack TentativeReassignments;
281
282 for (Register RewriteReg : RewriteRegs) {
283 LiveInterval &LI = LIS.getInterval(Reg: RewriteReg);
284 TentativeReassignments.push_back(Elt: {&LI, VRM.getPhys(virtReg: RewriteReg)});
285 LRM.unassign(VirtReg: LI);
286 }
287
288 if (!DebugCounter::shouldExecute(Counter&: RewriteAGPRCopyMFMACounter) ||
289 !attemptReassignmentsToAGPR(InterferingRegs&: RewriteRegs, PrefPhysReg: PhysRegHint)) {
290 // Roll back the register assignments to the original state.
291 for (auto [LI, OldAssign] : TentativeReassignments) {
292 if (VRM.hasPhys(virtReg: LI->reg()))
293 LRM.unassign(VirtReg: *LI);
294 LRM.assign(VirtReg: *LI, PhysReg: OldAssign);
295 }
296
297 return false;
298 }
299
300 // Fixup the register classes of the virtual registers now that we've
301 // committed to the reassignments.
302 for (Register InterferingReg : RewriteRegs) {
303 const TargetRegisterClass *EquivalentAGPRRegClass =
304 TRI.getEquivalentAGPRClass(SRC: MRI.getRegClass(Reg: InterferingReg));
305 MRI.setRegClass(Reg: InterferingReg, RC: EquivalentAGPRRegClass);
306 }
307
308 for (MachineInstr *RewriteCandidate : RewriteCandidates) {
309 int NewMFMAOp = AMDGPU::getAGPRFormOp(Opcode: RewriteCandidate->getOpcode());
310 RewriteCandidate->setDesc(TII.get(Opcode: NewMFMAOp));
311 ++NumMFMAsRewrittenToAGPR;
312 }
313
314 return true;
315}
316
317/// Attempt to reassign the registers in \p InterferingRegs to be AGPRs, with a
318/// preference to use \p PhysReg first. Returns false if the reassignments
319/// cannot be trivially performed.
320bool AMDGPURewriteAGPRCopyMFMAImpl::attemptReassignmentsToAGPR(
321 SmallSetVector<Register, 4> &InterferingRegs, MCPhysReg PrefPhysReg) const {
322 // FIXME: The ordering may matter here, but we're just taking uselistorder
323 // with the special case of ensuring to process the starting instruction
324 // first. We probably should extract the priority advisor out of greedy and
325 // use that ordering.
326 for (Register InterferingReg : InterferingRegs) {
327 LiveInterval &ReassignLI = LIS.getInterval(Reg: InterferingReg);
328 const TargetRegisterClass *EquivalentAGPRRegClass =
329 TRI.getEquivalentAGPRClass(SRC: MRI.getRegClass(Reg: InterferingReg));
330
331 MCPhysReg Assignable = AMDGPU::NoRegister;
332 if (EquivalentAGPRRegClass->contains(Reg: PrefPhysReg) &&
333 LRM.checkInterference(VirtReg: ReassignLI, PhysReg: PrefPhysReg) ==
334 LiveRegMatrix::IK_Free) {
335 // First try to assign to the AGPR we were already copying to. This
336 // should be the first assignment we attempt. We have to guard
337 // against the use being a subregister (which doesn't have an exact
338 // class match).
339
340 // TODO: If this does happen to be a subregister use, we should
341 // still try to assign to a subregister of the original copy result.
342 Assignable = PrefPhysReg;
343 } else {
344 ArrayRef<MCPhysReg> AllocOrder =
345 RegClassInfo.getOrder(RC: EquivalentAGPRRegClass);
346 for (MCPhysReg Reg : AllocOrder) {
347 if (LRM.checkInterference(VirtReg: ReassignLI, PhysReg: Reg) == LiveRegMatrix::IK_Free) {
348 Assignable = Reg;
349 break;
350 }
351 }
352 }
353
354 if (!Assignable) {
355 LLVM_DEBUG(dbgs() << "Unable to reassign VGPR "
356 << printReg(InterferingReg, &TRI)
357 << " to a free AGPR\n");
358 return false;
359 }
360
361 LLVM_DEBUG(dbgs() << "Reassigning VGPR " << printReg(InterferingReg, &TRI)
362 << " to " << printReg(Assignable, &TRI) << '\n');
363 LRM.assign(VirtReg: ReassignLI, PhysReg: Assignable);
364 }
365
366 return true;
367}
368
369/// Identify copies that look like:
370/// %vdst:vgpr = V_MFMA_.. %src0:av, %src1:av, %src2:vgpr
371/// %agpr = COPY %vgpr
372///
373/// Then try to replace the transitive uses of %src2 and %vdst with the AGPR
374/// versions of the MFMA. This should cover the common case.
375bool AMDGPURewriteAGPRCopyMFMAImpl::tryFoldCopiesToAGPR(
376 Register VReg, MCRegister AssignedAGPR) const {
377 bool MadeChange = false;
378 for (MachineInstr &UseMI : MRI.def_instructions(Reg: VReg)) {
379 if (!UseMI.isCopy())
380 continue;
381
382 Register CopySrcReg = UseMI.getOperand(i: 1).getReg();
383 if (!CopySrcReg.isVirtual())
384 continue;
385
386 // TODO: Handle loop phis copied to AGPR. e.g.
387 //
388 // loop:
389 // %phi:vgpr = COPY %mfma:vgpr
390 // %mfma:vgpr = V_MFMA_xxx_vgprcd_e64 %a, %b, %phi
391 // s_cbranch_vccnz loop
392 //
393 // endloop:
394 // %agpr = mfma
395 //
396 // We need to be sure that %phi is assigned to the same physical register as
397 // %mfma, or else we will just be moving copies into the loop.
398
399 for (MachineInstr &CopySrcDefMI : MRI.def_instructions(Reg: CopySrcReg)) {
400 if (isRewriteCandidate(MI: CopySrcDefMI) &&
401 tryReassigningMFMAChain(
402 MFMA&: CopySrcDefMI, MFMAHintReg: CopySrcDefMI.getOperand(i: 0).getReg(), PhysRegHint: AssignedAGPR))
403 MadeChange = true;
404 }
405 }
406
407 return MadeChange;
408}
409
410/// Identify copies that look like:
411/// %src:vgpr = COPY %src:agpr
412/// %vdst:vgpr = V_MFMA_... %src0:av, %src1:av, %src:vgpr
413///
414/// Then try to replace the transitive uses of %src2 and %vdst with the AGPR
415/// versions of the MFMA. This should cover rarer cases, and will generally be
416/// redundant with tryFoldCopiesToAGPR.
417bool AMDGPURewriteAGPRCopyMFMAImpl::tryFoldCopiesFromAGPR(
418 Register VReg, MCRegister AssignedAGPR) const {
419 bool MadeChange = false;
420 for (MachineInstr &UseMI : MRI.use_instructions(Reg: VReg)) {
421 if (!UseMI.isCopy())
422 continue;
423
424 Register CopyDstReg = UseMI.getOperand(i: 0).getReg();
425 if (!CopyDstReg.isVirtual())
426 continue;
427 for (MachineOperand &CopyUseMO : MRI.reg_nodbg_operands(Reg: CopyDstReg)) {
428 if (!CopyUseMO.readsReg())
429 continue;
430
431 MachineInstr &CopyUseMI = *CopyUseMO.getParent();
432 if (isRewriteCandidate(MI: CopyUseMI)) {
433 if (tryReassigningMFMAChain(MFMA&: CopyUseMI, MFMAHintReg: CopyDstReg,
434 PhysRegHint: VRM.getPhys(virtReg: CopyDstReg)))
435 MadeChange = true;
436 }
437 }
438 }
439
440 return MadeChange;
441}
442
443void AMDGPURewriteAGPRCopyMFMAImpl::replaceSpillWithCopyToVReg(
444 MachineInstr &SpillMI, int SpillFI, Register VReg) const {
445 const DebugLoc &DL = SpillMI.getDebugLoc();
446 MachineBasicBlock &MBB = *SpillMI.getParent();
447 MachineInstr *NewCopy;
448 if (SpillMI.mayStore()) {
449 NewCopy = BuildMI(BB&: MBB, I&: SpillMI, MIMD: DL, MCID: TII.get(Opcode: TargetOpcode::COPY), DestReg: VReg)
450 .add(MO: SpillMI.getOperand(i: 0));
451 } else {
452 NewCopy = BuildMI(BB&: MBB, I&: SpillMI, MIMD: DL, MCID: TII.get(Opcode: TargetOpcode::COPY))
453 .add(MO: SpillMI.getOperand(i: 0))
454 .addReg(RegNo: VReg);
455 }
456
457 LIS.ReplaceMachineInstrInMaps(MI&: SpillMI, NewMI&: *NewCopy);
458 SpillMI.eraseFromParent();
459}
460
461void AMDGPURewriteAGPRCopyMFMAImpl::collectSpillIndexUses(
462 ArrayRef<LiveInterval *> StackIntervals, SpillReferenceMap &Map) const {
463
464 SmallSet<int, 4> NeededFrameIndexes;
465 for (const LiveInterval *LI : StackIntervals)
466 NeededFrameIndexes.insert(V: LI->reg().stackSlotIndex());
467
468 for (MachineBasicBlock &MBB : MF) {
469 for (MachineInstr &MI : MBB) {
470 for (MachineOperand &MO : MI.operands()) {
471 if (!MO.isFI() || !NeededFrameIndexes.count(V: MO.getIndex()))
472 continue;
473
474 if (TII.isVGPRSpill(MI)) {
475 SmallVector<MachineInstr *, 4> &References = Map[MO.getIndex()];
476 References.push_back(Elt: &MI);
477 break;
478 }
479
480 // Verify this was really a spill instruction, if it's not just ignore
481 // all uses.
482
483 // TODO: This should probably be verifier enforced.
484 NeededFrameIndexes.erase(V: MO.getIndex());
485 Map.erase(Val: MO.getIndex());
486 }
487 }
488 }
489}
490
491bool AMDGPURewriteAGPRCopyMFMAImpl::isLoadJointlyDominatedByStores(
492 const MachineInstr &LoadMI, const LiveInterval &SlotLI,
493 const SmallPtrSetImpl<MachineBasicBlock *> &StoreFreeReachable) const {
494 const MachineBasicBlock *LoadMBB = LoadMI.getParent();
495 if (!MDT.isReachableFromEntry(A: LoadMBB))
496 return true;
497
498 // Check if every path passed through a store block.
499 if (!StoreFreeReachable.contains(Ptr: LoadMBB))
500 return true;
501
502 // Otherwise, there exists a path to this block that has not seen any store
503 // yet. We must ensure that within this block there is a store to this slot
504 // before the load. Consult the slot's LiveStacks interval: a store to the
505 // slot before the load means the slot is not live into this block but is
506 // live at the load. If the load reads an undef value, the slot is not live
507 // at the load, failing the joint-dominance check.
508 SlotIndex LoadIdx = LIS.getInstructionIndex(Instr: LoadMI);
509 return SlotLI.liveAt(index: LoadIdx) && !LIS.isLiveInToMBB(LR: SlotLI, mbb: LoadMBB);
510}
511
512void AMDGPURewriteAGPRCopyMFMAImpl::eliminateSpillsOfReassignedVGPRs() const {
513 unsigned NumSlots = LSS.getNumIntervals();
514 if (NumSlots == 0)
515 return;
516
517 MachineFrameInfo &MFI = MF.getFrameInfo();
518
519 SmallVector<LiveInterval *, 32> StackIntervals;
520 StackIntervals.reserve(N: NumSlots);
521
522 for (auto &[Slot, LI] : LSS) {
523 if (!MFI.isSpillSlotObjectIndex(ObjectIdx: Slot) || MFI.isDeadObjectIndex(ObjectIdx: Slot))
524 continue;
525
526 const TargetRegisterClass *RC = LSS.getIntervalRegClass(Slot);
527 if (TRI.hasVGPRs(RC))
528 StackIntervals.push_back(Elt: &LI);
529 }
530
531 sort(C&: StackIntervals, Comp: [](const LiveInterval *A, const LiveInterval *B) {
532 // The ordering has to be strictly weak.
533 /// Sort heaviest intervals first to prioritize their unspilling
534 if (A->weight() != B->weight())
535 return A->weight() > B->weight();
536
537 if (A->getSize() != B->getSize())
538 return A->getSize() > B->getSize();
539
540 // Tie breaker by number to avoid need for stable sort
541 return A->reg().stackSlotIndex() < B->reg().stackSlotIndex();
542 });
543
544 // FIXME: The APIs for dealing with the LiveInterval of a frame index are
545 // cumbersome. LiveStacks owns its LiveIntervals which refer to stack
546 // slots. We cannot use the usual LiveRegMatrix::assign and unassign on these,
547 // and must create a substitute virtual register to do so. This makes
548 // incremental updating here difficult; we need to actually perform the IR
549 // mutation to get the new vreg references in place to compute the register
550 // LiveInterval to perform an assignment to track the new interference
551 // correctly, and we can't simply migrate the LiveInterval we already have.
552 //
553 // To avoid walking through the entire function for each index, pre-collect
554 // all the instructions slot referencess.
555
556 DenseMap<int, SmallVector<MachineInstr *, 4>> SpillSlotReferences;
557 collectSpillIndexUses(StackIntervals, Map&: SpillSlotReferences);
558
559 for (LiveInterval *LI : StackIntervals) {
560 int Slot = LI->reg().stackSlotIndex();
561 auto SpillReferences = SpillSlotReferences.find(Val: Slot);
562 if (SpillReferences == SpillSlotReferences.end())
563 continue;
564
565 // For each spill reload, every path from entry to the reload must pass
566 // through at least one spill store to the same stack slot.
567 SmallPtrSet<MachineBasicBlock *, 4> StoreBlocks;
568 for (MachineInstr *MI : SpillReferences->second) {
569 if (MI->mayStore() && MDT.isReachableFromEntry(A: MI->getParent()))
570 StoreBlocks.insert(Ptr: MI->getParent());
571 }
572
573 if (StoreBlocks.empty()) {
574 LLVM_DEBUG(dbgs() << "Skipping " << printReg(LI->reg(), &TRI)
575 << ": no reachable stores\n");
576 continue;
577 }
578
579 // Compute blocks reachable from entry without passing through a store
580 // block.
581 MachineBasicBlock &EntryMBB = MF.front();
582 SmallPtrSet<MachineBasicBlock *, 16> StoreFreeReachable = {&EntryMBB};
583 SmallVector<MachineBasicBlock *, 16> Worklist = {&EntryMBB};
584
585 while (!Worklist.empty()) {
586 MachineBasicBlock *MBB = Worklist.pop_back_val();
587 if (StoreBlocks.contains(Ptr: MBB))
588 continue;
589
590 for (MachineBasicBlock *Succ : MBB->successors()) {
591 if (StoreFreeReachable.insert(Ptr: Succ).second)
592 Worklist.push_back(Elt: Succ);
593 }
594 }
595
596 // Every reachable reload must be jointly dominated by the slot's stores.
597 if (!llvm::all_of(Range&: SpillReferences->second, P: [&](const MachineInstr *MI) {
598 return !MI->mayLoad() ||
599 isLoadJointlyDominatedByStores(LoadMI: *MI, SlotLI: *LI, StoreFreeReachable);
600 })) {
601 LLVM_DEBUG(
602 dbgs() << "Skipping " << printReg(LI->reg(), &TRI)
603 << ": some reachable load not jointly dominated by stores\n");
604 continue;
605 }
606
607 const TargetRegisterClass *RC = LSS.getIntervalRegClass(Slot);
608
609 LLVM_DEBUG(dbgs() << "Trying to eliminate " << printReg(LI->reg(), &TRI)
610 << " by reassigning\n");
611
612 ArrayRef<MCPhysReg> AllocOrder = RegClassInfo.getOrder(RC);
613
614 // The stack slot's LiveInterval may be discontiguous: a slot can be live
615 // in memory around a spill store and around a much later reload. Once
616 // we unspill the slot into a register, however, the value must reside in
617 // that register continuously from its first reference to its last (modulo
618 // the live range splitting that happens later below). Checking
619 // interference against the slot's discontiguous interval could let us pick
620 // a PhysReg that is busy inside a gap, corrupting it. Instead, check
621 // interference over the range the replacement register will occupy.
622 for (MCPhysReg PhysReg : AllocOrder) {
623 if (LRM.checkInterference(Start: LI->beginIndex(), End: LI->endIndex(), PhysReg))
624 continue;
625
626 LLVM_DEBUG(dbgs() << "Reassigning " << *LI << " to "
627 << printReg(PhysReg, &TRI) << '\n');
628
629 const TargetRegisterClass *RC = LSS.getIntervalRegClass(Slot);
630 Register NewVReg = MRI.createVirtualRegister(RegClass: RC);
631
632 for (MachineInstr *SpillMI : SpillReferences->second)
633 replaceSpillWithCopyToVReg(SpillMI&: *SpillMI, SpillFI: Slot, VReg: NewVReg);
634
635 // TODO: Transferring the stack slot's LiveInterval instead of recomputing
636 // would not be a straight copy: its segments would have to be widened to
637 // cover each store's reaching reloads.
638 LiveInterval &NewLI = LIS.createAndComputeVirtRegInterval(Reg: NewVReg);
639 VRM.grow();
640
641 // A spill slot can be stored to multiple times, so the replacement
642 // vreg may have multiple disconnected live range components. Split
643 // them into separate vregs to maintain the single-component invariant.
644 SmallVector<LiveInterval *, 4> SplitLIs;
645 LIS.splitSeparateComponents(LI&: NewLI, SplitLIs);
646
647 LLVM_DEBUG({
648 if (!SplitLIs.empty()) {
649 dbgs() << "Split unspilled interval into " << (SplitLIs.size() + 1)
650 << " components\n";
651 }
652 });
653
654 LRM.assign(VirtReg: NewLI, PhysReg);
655 for (LiveInterval *SplitLI : SplitLIs) {
656 VRM.grow();
657 LRM.assign(VirtReg: *SplitLI, PhysReg);
658 }
659
660 MFI.RemoveStackObject(ObjectIdx: Slot);
661 break;
662 }
663 }
664}
665
666bool AMDGPURewriteAGPRCopyMFMAImpl::run(MachineFunction &MF) const {
667 // This only applies on subtargets that have a configurable AGPR vs. VGPR
668 // allocation.
669 if (!ST.hasGFX90AInsts())
670 return false;
671
672 // Early exit if no AGPRs were assigned.
673 if (!LRM.isPhysRegUsed(PhysReg: AMDGPU::AGPR0)) {
674 LLVM_DEBUG(dbgs() << "skipping function that did not allocate AGPRs\n");
675 return false;
676 }
677
678 bool MadeChange = false;
679
680 for (unsigned I = 0, E = MRI.getNumVirtRegs(); I != E; ++I) {
681 Register VReg = Register::index2VirtReg(Index: I);
682 MCRegister AssignedAGPR = getAssignedAGPR(VReg);
683 if (!AssignedAGPR)
684 continue;
685
686 if (tryFoldCopiesToAGPR(VReg, AssignedAGPR))
687 MadeChange = true;
688 if (tryFoldCopiesFromAGPR(VReg, AssignedAGPR))
689 MadeChange = true;
690 }
691
692 // If we've successfully rewritten some MFMAs, we've alleviated some VGPR
693 // pressure. See if we can eliminate some spills now that those registers are
694 // more available.
695 if (MadeChange)
696 eliminateSpillsOfReassignedVGPRs();
697
698 return MadeChange;
699}
700
701class AMDGPURewriteAGPRCopyMFMALegacy : public MachineFunctionPass {
702public:
703 static char ID;
704
705 AMDGPURewriteAGPRCopyMFMALegacy() : MachineFunctionPass(ID) {}
706
707 bool runOnMachineFunction(MachineFunction &MF) override;
708
709 StringRef getPassName() const override {
710 return "AMDGPU Rewrite AGPR-Copy-MFMA";
711 }
712
713 void getAnalysisUsage(AnalysisUsage &AU) const override {
714 AU.addRequired<LiveIntervalsWrapperPass>();
715 AU.addRequired<VirtRegMapWrapperLegacy>();
716 AU.addRequired<LiveRegMatrixWrapperLegacy>();
717 AU.addRequired<LiveStacksWrapperLegacy>();
718 AU.addRequired<MachineRegisterClassInfoWrapperPass>();
719 AU.addRequired<MachineDominatorTreeWrapperPass>();
720
721 AU.addPreserved<LiveIntervalsWrapperPass>();
722 AU.addPreserved<VirtRegMapWrapperLegacy>();
723 AU.addPreserved<LiveRegMatrixWrapperLegacy>();
724 AU.addPreserved<LiveStacksWrapperLegacy>();
725 AU.addPreserved<MachineRegisterClassInfoWrapperPass>();
726 AU.addPreserved<MachineDominatorTreeWrapperPass>();
727
728 AU.setPreservesAll();
729 MachineFunctionPass::getAnalysisUsage(AU);
730 }
731};
732
733} // End anonymous namespace.
734
735INITIALIZE_PASS_BEGIN(AMDGPURewriteAGPRCopyMFMALegacy, DEBUG_TYPE,
736 "AMDGPU Rewrite AGPR-Copy-MFMA", false, false)
737INITIALIZE_PASS_DEPENDENCY(LiveIntervalsWrapperPass)
738INITIALIZE_PASS_DEPENDENCY(VirtRegMapWrapperLegacy)
739INITIALIZE_PASS_DEPENDENCY(LiveRegMatrixWrapperLegacy)
740INITIALIZE_PASS_DEPENDENCY(LiveStacksWrapperLegacy)
741INITIALIZE_PASS_DEPENDENCY(MachineRegisterClassInfoWrapperPass)
742INITIALIZE_PASS_DEPENDENCY(MachineDominatorTreeWrapperPass)
743INITIALIZE_PASS_END(AMDGPURewriteAGPRCopyMFMALegacy, DEBUG_TYPE,
744 "AMDGPU Rewrite AGPR-Copy-MFMA", false, false)
745
746char AMDGPURewriteAGPRCopyMFMALegacy::ID = 0;
747
748char &llvm::AMDGPURewriteAGPRCopyMFMALegacyID =
749 AMDGPURewriteAGPRCopyMFMALegacy::ID;
750
751bool AMDGPURewriteAGPRCopyMFMALegacy::runOnMachineFunction(
752 MachineFunction &MF) {
753 if (skipFunction(F: MF.getFunction()))
754 return false;
755
756 auto &VRM = getAnalysis<VirtRegMapWrapperLegacy>().getVRM();
757 auto &LRM = getAnalysis<LiveRegMatrixWrapperLegacy>().getLRM();
758 auto &LIS = getAnalysis<LiveIntervalsWrapperPass>().getLIS();
759 auto &LSS = getAnalysis<LiveStacksWrapperLegacy>().getLS();
760 auto &RCI = getAnalysis<MachineRegisterClassInfoWrapperPass>().getRCI();
761 auto &MDT = getAnalysis<MachineDominatorTreeWrapperPass>().getDomTree();
762 AMDGPURewriteAGPRCopyMFMAImpl Impl(MF, VRM, LRM, LIS, LSS, RCI, MDT);
763 return Impl.run(MF);
764}
765
766PreservedAnalyses
767AMDGPURewriteAGPRCopyMFMAPass::run(MachineFunction &MF,
768 MachineFunctionAnalysisManager &MFAM) {
769 VirtRegMap &VRM = MFAM.getResult<VirtRegMapAnalysis>(IR&: MF);
770 LiveRegMatrix &LRM = MFAM.getResult<LiveRegMatrixAnalysis>(IR&: MF);
771 LiveIntervals &LIS = MFAM.getResult<LiveIntervalsAnalysis>(IR&: MF);
772 LiveStacks &LSS = MFAM.getResult<LiveStacksAnalysis>(IR&: MF);
773 RegisterClassInfo &RCI = MFAM.getResult<MachineRegisterClassAnalysis>(IR&: MF);
774 MachineDominatorTree &MDT = MFAM.getResult<MachineDominatorTreeAnalysis>(IR&: MF);
775
776 AMDGPURewriteAGPRCopyMFMAImpl Impl(MF, VRM, LRM, LIS, LSS, RCI, MDT);
777 if (!Impl.run(MF))
778 return PreservedAnalyses::all();
779 auto PA = getMachineFunctionPassPreservedAnalyses();
780 PA.preserveSet<CFGAnalyses>()
781 .preserve<LiveStacksAnalysis>()
782 .preserve<VirtRegMapAnalysis>()
783 .preserve<SlotIndexesAnalysis>()
784 .preserve<LiveIntervalsAnalysis>()
785 .preserve<LiveRegMatrixAnalysis>()
786 .preserve<MachineRegisterClassAnalysis>();
787 return PA;
788}
789