1//===- GCNRegPressure.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
10/// This file implements the GCNRegPressure class.
11///
12//===----------------------------------------------------------------------===//
13
14#include "GCNRegPressure.h"
15#include "AMDGPU.h"
16#include "SIMachineFunctionInfo.h"
17#include "llvm/ADT/SetVector.h"
18#include "llvm/CodeGen/LiveIntervalUnion.h"
19#include "llvm/CodeGen/MachineBasicBlock.h"
20#include "llvm/CodeGen/MachineLoopInfo.h"
21#include "llvm/CodeGen/RegisterPressure.h"
22
23using namespace llvm;
24
25#define DEBUG_TYPE "machine-scheduler"
26
27bool llvm::isEqual(const GCNRPTracker::LiveRegSet &S1,
28 const GCNRPTracker::LiveRegSet &S2) {
29 if (S1.size() != S2.size())
30 return false;
31
32 for (const auto &P : S1) {
33 auto I = S2.find(Val: P.first);
34 if (I == S2.end() || I->second != P.second)
35 return false;
36 }
37 return true;
38}
39
40///////////////////////////////////////////////////////////////////////////////
41// GCNRegPressure
42
43unsigned GCNRegPressure::getRegKind(const TargetRegisterClass *RC,
44 const SIRegisterInfo *STI) {
45 return STI->isSGPRClass(RC)
46 ? SGPR
47 : (STI->isAGPRClass(RC)
48 ? AGPR
49 : (STI->isVectorSuperClass(RC) ? AVGPR : VGPR));
50}
51
52void GCNRegPressure::inc(unsigned Reg,
53 LaneBitmask PrevMask,
54 LaneBitmask NewMask,
55 const MachineRegisterInfo &MRI) {
56 unsigned NewNumCoveredRegs = SIRegisterInfo::getNumCoveredRegs(LM: NewMask);
57 unsigned PrevNumCoveredRegs = SIRegisterInfo::getNumCoveredRegs(LM: PrevMask);
58 if (NewNumCoveredRegs == PrevNumCoveredRegs)
59 return;
60
61 int Sign = 1;
62 if (NewMask < PrevMask) {
63 std::swap(a&: NewMask, b&: PrevMask);
64 std::swap(a&: NewNumCoveredRegs, b&: PrevNumCoveredRegs);
65 Sign = -1;
66 }
67 assert(PrevMask < NewMask && PrevNumCoveredRegs < NewNumCoveredRegs &&
68 "prev mask should always be lesser than new");
69
70 const TargetRegisterClass *RC = MRI.getRegClass(Reg);
71 const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo();
72 const SIRegisterInfo *STI = static_cast<const SIRegisterInfo *>(TRI);
73 unsigned RegKind = getRegKind(RC, STI);
74 if (TRI->getRegSizeInBits(RC: *RC) != 32) {
75 // Reg is from a tuple register class.
76 if (PrevMask.none()) {
77 unsigned TupleIdx = TOTAL_KINDS + RegKind;
78 Value[TupleIdx] += Sign * TRI->getRegClassWeight(RC).RegWeight;
79 }
80 // Pressure scales with number of new registers covered by the new mask.
81 // Note when true16 is enabled, we can no longer safely use the following
82 // approach to calculate the difference in the number of 32-bit registers
83 // between two masks:
84 //
85 // Sign *= SIRegisterInfo::getNumCoveredRegs(~PrevMask & NewMask);
86 //
87 // The issue is that the mask calculation `~PrevMask & NewMask` doesn't
88 // properly account for partial usage of a 32-bit register when dealing with
89 // 16-bit registers.
90 //
91 // Consider this example:
92 // Assume PrevMask = 0b0010 and NewMask = 0b1111. Here, the correct register
93 // usage difference should be 1, because even though PrevMask uses only half
94 // of a 32-bit register, it should still be counted as a full register use.
95 // However, the mask calculation yields `~PrevMask & NewMask = 0b1101`, and
96 // calling `getNumCoveredRegs` returns 2 instead of 1. This incorrect
97 // calculation can lead to integer overflow when Sign = -1.
98 Sign *= NewNumCoveredRegs - PrevNumCoveredRegs;
99 }
100 Value[RegKind] += Sign;
101}
102
103namespace {
104struct RegExcess {
105 unsigned SGPR = 0;
106 unsigned VGPR = 0;
107 unsigned ArchVGPR = 0;
108 unsigned AGPR = 0;
109
110 bool anyExcess() const { return SGPR || VGPR || ArchVGPR || AGPR; }
111 bool hasVectorRegisterExcess() const { return VGPR || ArchVGPR || AGPR; }
112
113 RegExcess(const MachineFunction &MF, const GCNRegPressure &RP)
114 : RegExcess(MF, RP, GCNRPTarget(MF, RP)) {}
115 RegExcess(const MachineFunction &MF, const GCNRegPressure &RP,
116 const GCNRPTarget &Target) {
117 unsigned MaxSGPRs = Target.getMaxSGPRs();
118 unsigned MaxVGPRs = Target.getMaxVGPRs();
119
120 const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>();
121 SGPR = std::max(a: static_cast<int>(RP.getSGPRNum() - MaxSGPRs), b: 0);
122
123 // The number of virtual VGPRs required to handle excess SGPR
124 unsigned WaveSize = ST.getWavefrontSize();
125 unsigned VGPRForSGPRSpills = divideCeil(Numerator: SGPR, Denominator: WaveSize);
126
127 unsigned MaxArchVGPRs = ST.getAddressableNumArchVGPRs();
128
129 // Unified excess pressure conditions, accounting for VGPRs used for SGPR
130 // spills
131 VGPR = std::max(a: static_cast<int>(RP.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) +
132 VGPRForSGPRSpills - MaxVGPRs),
133 b: 0);
134
135 unsigned ArchVGPRLimit = ST.hasGFX90AInsts() ? MaxArchVGPRs : MaxVGPRs;
136 // Arch VGPR excess pressure conditions, accounting for VGPRs used for SGPR
137 // spills
138 ArchVGPR = std::max(a: static_cast<int>(RP.getArchVGPRNum() +
139 VGPRForSGPRSpills - ArchVGPRLimit),
140 b: 0);
141
142 // AGPR excess pressure conditions
143 AGPR = std::max(a: static_cast<int>(RP.getAGPRNum() - ArchVGPRLimit), b: 0);
144 }
145};
146} // namespace
147
148bool GCNRegPressure::less(const MachineFunction &MF, const GCNRegPressure &O,
149 unsigned MaxOccupancy) const {
150 const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>();
151 unsigned DynamicVGPRBlockSize =
152 MF.getInfo<SIMachineFunctionInfo>()->getDynamicVGPRBlockSize();
153
154 const auto SGPROcc = std::min(a: MaxOccupancy,
155 b: ST.getOccupancyWithNumSGPRs(SGPRs: getSGPRNum()));
156 const auto VGPROcc = std::min(
157 a: MaxOccupancy, b: ST.getOccupancyWithNumVGPRs(VGPRs: getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()),
158 DynamicVGPRBlockSize));
159 const auto OtherSGPROcc = std::min(a: MaxOccupancy,
160 b: ST.getOccupancyWithNumSGPRs(SGPRs: O.getSGPRNum()));
161 const auto OtherVGPROcc =
162 std::min(a: MaxOccupancy,
163 b: ST.getOccupancyWithNumVGPRs(VGPRs: O.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()),
164 DynamicVGPRBlockSize));
165
166 const auto Occ = std::min(a: SGPROcc, b: VGPROcc);
167 const auto OtherOcc = std::min(a: OtherSGPROcc, b: OtherVGPROcc);
168
169 // Give first precedence to the better occupancy.
170 if (Occ != OtherOcc)
171 return Occ > OtherOcc;
172
173 unsigned MaxVGPRs = ST.getMaxNumVGPRs(MF);
174
175 RegExcess Excess(MF, *this);
176 RegExcess OtherExcess(MF, O);
177
178 unsigned MaxArchVGPRs = ST.getAddressableNumArchVGPRs();
179
180 bool ExcessRP = Excess.anyExcess();
181 bool OtherExcessRP = OtherExcess.anyExcess();
182
183 // Give second precedence to the reduced number of spills to hold the register
184 // pressure.
185 if (ExcessRP || OtherExcessRP) {
186 // The difference in excess VGPR pressure, after including VGPRs used for
187 // SGPR spills
188 int VGPRDiff =
189 ((OtherExcess.VGPR + OtherExcess.ArchVGPR + OtherExcess.AGPR) -
190 (Excess.VGPR + Excess.ArchVGPR + Excess.AGPR));
191
192 int SGPRDiff = OtherExcess.SGPR - Excess.SGPR;
193
194 if (VGPRDiff != 0)
195 return VGPRDiff > 0;
196 if (SGPRDiff != 0) {
197 unsigned PureExcessVGPR =
198 std::max(a: static_cast<int>(getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) - MaxVGPRs),
199 b: 0) +
200 std::max(a: static_cast<int>(getVGPRNum(UnifiedVGPRFile: false) - MaxArchVGPRs), b: 0);
201 unsigned OtherPureExcessVGPR =
202 std::max(
203 a: static_cast<int>(O.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) - MaxVGPRs),
204 b: 0) +
205 std::max(a: static_cast<int>(O.getVGPRNum(UnifiedVGPRFile: false) - MaxArchVGPRs), b: 0);
206
207 // If we have a special case where there is a tie in excess VGPR, but one
208 // of the pressures has VGPR usage from SGPR spills, prefer the pressure
209 // with SGPR spills.
210 if (PureExcessVGPR != OtherPureExcessVGPR)
211 return SGPRDiff < 0;
212 // If both pressures have the same excess pressure before and after
213 // accounting for SGPR spills, prefer fewer SGPR spills.
214 return SGPRDiff > 0;
215 }
216 }
217
218 bool SGPRImportant = SGPROcc < VGPROcc;
219 const bool OtherSGPRImportant = OtherSGPROcc < OtherVGPROcc;
220
221 // If both pressures disagree on what is more important compare vgprs.
222 if (SGPRImportant != OtherSGPRImportant) {
223 SGPRImportant = false;
224 }
225
226 // Give third precedence to lower register tuple pressure.
227 bool SGPRFirst = SGPRImportant;
228 for (int I = 2; I > 0; --I, SGPRFirst = !SGPRFirst) {
229 if (SGPRFirst) {
230 auto SW = getSGPRTuplesWeight();
231 auto OtherSW = O.getSGPRTuplesWeight();
232 if (SW != OtherSW)
233 return SW < OtherSW;
234 } else {
235 auto VW = getVGPRTuplesWeight();
236 auto OtherVW = O.getVGPRTuplesWeight();
237 if (VW != OtherVW)
238 return VW < OtherVW;
239 }
240 }
241
242 // Give final precedence to lower general RP.
243 return SGPRImportant ? (getSGPRNum() < O.getSGPRNum()):
244 (getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) <
245 O.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()));
246}
247
248Printable llvm::print(const GCNRegPressure &RP, const GCNSubtarget *ST,
249 unsigned DynamicVGPRBlockSize) {
250 return Printable([&RP, ST, DynamicVGPRBlockSize](raw_ostream &OS) {
251 OS << "VGPRs: " << RP.getArchVGPRNum() << ' '
252 << "AGPRs: " << RP.getAGPRNum();
253 if (ST)
254 OS << "(O"
255 << ST->getOccupancyWithNumVGPRs(VGPRs: RP.getVGPRNum(UnifiedVGPRFile: ST->hasGFX90AInsts()),
256 DynamicVGPRBlockSize)
257 << ')';
258 OS << ", SGPRs: " << RP.getSGPRNum();
259 if (ST)
260 OS << "(O" << ST->getOccupancyWithNumSGPRs(SGPRs: RP.getSGPRNum()) << ')';
261 OS << ", LVGPR WT: " << RP.getVGPRTuplesWeight()
262 << ", LSGPR WT: " << RP.getSGPRTuplesWeight();
263 if (ST)
264 OS << " -> Occ: " << RP.getOccupancy(ST: *ST, DynamicVGPRBlockSize);
265 OS << '\n';
266 });
267}
268
269static LaneBitmask getDefRegMask(const MachineOperand &MO,
270 const MachineRegisterInfo &MRI) {
271 assert(MO.isDef() && MO.isReg() && MO.getReg().isVirtual());
272
273 // We don't rely on read-undef flag because in case of tentative schedule
274 // tracking it isn't set correctly yet. This works correctly however since
275 // use mask has been tracked before using LIS.
276 return MO.getSubReg() == 0 ?
277 MRI.getMaxLaneMaskForVReg(Reg: MO.getReg()) :
278 MRI.getTargetRegisterInfo()->getSubRegIndexLaneMask(SubIdx: MO.getSubReg());
279}
280
281static void
282collectVirtualRegUses(SmallVectorImpl<VRegMaskOrUnit> &VRegMaskOrUnits,
283 const MachineInstr &MI, const LiveIntervals &LIS,
284 const MachineRegisterInfo &MRI) {
285
286 auto &TRI = *MRI.getTargetRegisterInfo();
287 for (const auto &MO : MI.operands()) {
288 if (!MO.isReg() || !MO.getReg().isVirtual())
289 continue;
290 if (!MO.isUse() || !MO.readsReg())
291 continue;
292
293 Register Reg = MO.getReg();
294 auto I = llvm::find_if(Range&: VRegMaskOrUnits, P: [Reg](const VRegMaskOrUnit &RM) {
295 return RM.VRegOrUnit.asVirtualReg() == Reg;
296 });
297
298 auto &P = I == VRegMaskOrUnits.end()
299 ? VRegMaskOrUnits.emplace_back(Args: VirtRegOrUnit(Reg),
300 Args: LaneBitmask::getNone())
301 : *I;
302
303 P.LaneMask |= MO.getSubReg() ? TRI.getSubRegIndexLaneMask(SubIdx: MO.getSubReg())
304 : MRI.getMaxLaneMaskForVReg(Reg);
305 }
306
307 SlotIndex InstrSI;
308 for (auto &P : VRegMaskOrUnits) {
309 auto &LI = LIS.getInterval(Reg: P.VRegOrUnit.asVirtualReg());
310 if (!LI.hasSubRanges())
311 continue;
312
313 // For a tentative schedule LIS isn't updated yet but livemask should
314 // remain the same on any schedule. Subreg defs can be reordered but they
315 // all must dominate uses anyway.
316 if (!InstrSI)
317 InstrSI = LIS.getInstructionIndex(Instr: MI).getBaseIndex();
318
319 P.LaneMask = getLiveLaneMask(LI, SI: InstrSI, MRI, LaneMaskFilter: P.LaneMask);
320 }
321}
322
323/// Mostly copy/paste from CodeGen/RegisterPressure.cpp
324static LaneBitmask getLanesWithProperty(
325 const LiveIntervals &LIS, const MachineRegisterInfo &MRI,
326 bool TrackLaneMasks, Register Reg, SlotIndex Pos,
327 function_ref<bool(const LiveRange &LR, SlotIndex Pos)> Property) {
328 assert(Reg.isVirtual());
329 const LiveInterval &LI = LIS.getInterval(Reg);
330 LaneBitmask Result;
331 if (TrackLaneMasks && LI.hasSubRanges()) {
332 for (const LiveInterval::SubRange &SR : LI.subranges()) {
333 if (Property(SR, Pos))
334 Result |= SR.LaneMask;
335 }
336 } else if (Property(LI, Pos)) {
337 Result =
338 TrackLaneMasks ? MRI.getMaxLaneMaskForVReg(Reg) : LaneBitmask::getAll();
339 }
340
341 return Result;
342}
343
344/// Mostly copy/paste from CodeGen/RegisterPressure.cpp
345/// Helper to find a vreg use between two indices {PriorUseIdx, NextUseIdx}.
346/// The query starts with a lane bitmask which gets lanes/bits removed for every
347/// use we find.
348static LaneBitmask findUseBetween(unsigned Reg, LaneBitmask LastUseMask,
349 SlotIndex PriorUseIdx, SlotIndex NextUseIdx,
350 const MachineRegisterInfo &MRI,
351 const SIRegisterInfo *TRI,
352 const LiveIntervals *LIS,
353 bool Upward = false) {
354 for (const MachineOperand &MO : MRI.use_nodbg_operands(Reg)) {
355 if (MO.isUndef())
356 continue;
357 const MachineInstr *MI = MO.getParent();
358 SlotIndex InstSlot = LIS->getInstructionIndex(Instr: *MI).getRegSlot();
359 bool InRange = Upward ? (InstSlot > PriorUseIdx && InstSlot <= NextUseIdx)
360 : (InstSlot >= PriorUseIdx && InstSlot < NextUseIdx);
361 if (!InRange)
362 continue;
363
364 unsigned SubRegIdx = MO.getSubReg();
365 LaneBitmask UseMask = TRI->getSubRegIndexLaneMask(SubIdx: SubRegIdx);
366 LastUseMask &= ~UseMask;
367 if (LastUseMask.none())
368 return LaneBitmask::getNone();
369 }
370 return LastUseMask;
371}
372
373////////////////////////////////////////////////////////////////////////////////
374// GCNRPTarget
375
376GCNRPTarget::GCNRPTarget(const MachineFunction &MF, const GCNRegPressure &RP)
377 : GCNRPTarget(RP, MF) {
378 const Function &F = MF.getFunction();
379 const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>();
380 setTarget(NumSGPRs: ST.getMaxNumSGPRs(F), NumVGPRs: ST.getMaxNumVGPRs(F));
381}
382
383GCNRPTarget::GCNRPTarget(unsigned NumSGPRs, unsigned NumVGPRs,
384 const MachineFunction &MF, const GCNRegPressure &RP)
385 : GCNRPTarget(RP, MF) {
386 setTarget(NumSGPRs, NumVGPRs);
387}
388
389GCNRPTarget::GCNRPTarget(unsigned Occupancy, const MachineFunction &MF,
390 const GCNRegPressure &RP)
391 : GCNRPTarget(RP, MF) {
392 const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>();
393 unsigned DynamicVGPRBlockSize =
394 MF.getInfo<SIMachineFunctionInfo>()->getDynamicVGPRBlockSize();
395 setTarget(NumSGPRs: ST.getMaxNumSGPRs(WavesPerEU: Occupancy, /*Addressable=*/false),
396 NumVGPRs: ST.getMaxNumVGPRs(WavesPerEU: Occupancy, DynamicVGPRBlockSize));
397}
398
399void GCNRPTarget::setTarget(unsigned NumSGPRs, unsigned NumVGPRs) {
400 const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>();
401 MaxSGPRs = std::min(a: ST.getAddressableNumSGPRs(), b: NumSGPRs);
402 MaxVGPRs = std::min(a: ST.getAddressableNumArchVGPRs(), b: NumVGPRs);
403 if (UnifiedRF) {
404 unsigned DynamicVGPRBlockSize =
405 MF.getInfo<SIMachineFunctionInfo>()->getDynamicVGPRBlockSize();
406 MaxUnifiedVGPRs =
407 std::min(a: ST.getAddressableNumVGPRs(DynamicVGPRBlockSize), b: NumVGPRs);
408 } else {
409 MaxUnifiedVGPRs = 0;
410 }
411}
412
413bool GCNRPTarget::isSaveBeneficial(Register Reg) const {
414 const MachineRegisterInfo &MRI = MF.getRegInfo();
415 const TargetRegisterClass *RC = MRI.getRegClass(Reg);
416 const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo();
417 const SIRegisterInfo *SRI = static_cast<const SIRegisterInfo *>(TRI);
418
419 RegExcess Excess(MF, RP, *this);
420
421 if (SRI->isSGPRClass(RC))
422 return Excess.SGPR;
423
424 if (SRI->isAGPRClass(RC))
425 return (UnifiedRF && Excess.VGPR) || Excess.AGPR;
426
427 return (UnifiedRF && Excess.VGPR) || Excess.ArchVGPR;
428}
429
430bool GCNRPTarget::isSaveBeneficial(const GCNRegPressure &SaveRP) const {
431 RegExcess Excess(MF, RP, *this);
432 if (SaveRP.getSGPRNum() != 0 && Excess.SGPR != 0)
433 return true;
434 if (SaveRP.getArchVGPRNum() != 0 && Excess.ArchVGPR != 0)
435 return true;
436 if (SaveRP.getAGPRNum() != 0 && Excess.AGPR != 0)
437 return true;
438 if (UnifiedRF && Excess.VGPR != 0)
439 return SaveRP.getArchVGPRNum() != 0 || SaveRP.getAGPRNum() != 0;
440 return false;
441}
442
443unsigned GCNRPTarget::getNumRegsBenefit(const GCNRegPressure &SaveRP) const {
444 RegExcess Excess(MF, RP, *this);
445 const unsigned NumVGPRAboveAddrLimit =
446 std::min(a: Excess.ArchVGPR, b: SaveRP.getArchVGPRNum()) +
447 std::min(a: Excess.AGPR, b: SaveRP.getAGPRNum());
448 unsigned NumRegsSaved =
449 std::min(a: Excess.SGPR, b: SaveRP.getSGPRNum()) + NumVGPRAboveAddrLimit;
450
451 if (UnifiedRF && Excess.VGPR) {
452 // We have already accounted for excess pressure above addressive limits for
453 // the individual VGPR classes. However for targets with unified RFs there
454 // is also a unified VGPR pressure (ArchVGPR + AGPR combination) limit to
455 // honor that may be more restrictive that the per-VGPR-class limits. We
456 // must also be careful not to double-count VGPR saves that may contribute
457 // to lowering pressure both above the addressable limit in their respective
458 // class as well as in the unified VGPR limit.
459 const unsigned VGPRSave = SaveRP.getArchVGPRNum() + SaveRP.getAGPRNum();
460 if (NumVGPRAboveAddrLimit < VGPRSave)
461 NumRegsSaved += std::min(a: Excess.VGPR, b: VGPRSave - NumVGPRAboveAddrLimit);
462 }
463
464 return NumRegsSaved;
465}
466
467bool GCNRPTarget::satisfied(const GCNRegPressure &TestRP) const {
468 if (TestRP.getSGPRNum() > MaxSGPRs || TestRP.getVGPRNum(UnifiedVGPRFile: false) > MaxVGPRs)
469 return false;
470 if (UnifiedRF && TestRP.getVGPRNum(UnifiedVGPRFile: true) > MaxUnifiedVGPRs)
471 return false;
472 return true;
473}
474
475bool GCNRPTarget::hasVectorRegisterExcess() const {
476 RegExcess Excess(MF, RP, *this);
477 return Excess.hasVectorRegisterExcess();
478}
479
480///////////////////////////////////////////////////////////////////////////////
481// GCNRPTracker
482
483LaneBitmask llvm::getLiveLaneMask(unsigned Reg, SlotIndex SI,
484 const LiveIntervals &LIS,
485 const MachineRegisterInfo &MRI,
486 LaneBitmask LaneMaskFilter) {
487 return getLiveLaneMask(LI: LIS.getInterval(Reg), SI, MRI, LaneMaskFilter);
488}
489
490LaneBitmask llvm::getLiveLaneMask(const LiveInterval &LI, SlotIndex SI,
491 const MachineRegisterInfo &MRI,
492 LaneBitmask LaneMaskFilter) {
493 LaneBitmask LiveMask;
494 if (LI.hasSubRanges()) {
495 for (const auto &S : LI.subranges())
496 if ((S.LaneMask & LaneMaskFilter).any() && S.liveAt(index: SI)) {
497 LiveMask |= S.LaneMask;
498 assert(LiveMask == (LiveMask & MRI.getMaxLaneMaskForVReg(LI.reg())));
499 }
500 } else if (LI.liveAt(index: SI)) {
501 LiveMask = MRI.getMaxLaneMaskForVReg(Reg: LI.reg());
502 }
503 LiveMask &= LaneMaskFilter;
504 return LiveMask;
505}
506
507GCNRPTracker::LiveRegSet llvm::getLiveRegs(SlotIndex SI,
508 const LiveIntervals &LIS,
509 const MachineRegisterInfo &MRI,
510 GCNRegPressure::RegKind RegKind) {
511 GCNRPTracker::LiveRegSet LiveRegs;
512 for (unsigned I = 0, E = MRI.getNumVirtRegs(); I != E; ++I) {
513 auto Reg = Register::index2VirtReg(Index: I);
514 if (RegKind != GCNRegPressure::TOTAL_KINDS &&
515 GCNRegPressure::getRegKind(Reg, MRI) != RegKind)
516 continue;
517 if (!LIS.hasInterval(Reg))
518 continue;
519 auto LiveMask = getLiveLaneMask(Reg, SI, LIS, MRI);
520 if (LiveMask.any())
521 LiveRegs[Reg] = LiveMask;
522 }
523 return LiveRegs;
524}
525
526void GCNRPTracker::reset(const MachineInstr &MI, bool After) {
527 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
528 if (!MI.isDebugInstr()) {
529 SlotIndex SI = LIS.getInstructionIndex(Instr: MI);
530 if (After)
531 SI = SI.getDeadSlot();
532 reset(MRI, SI);
533 return;
534 }
535
536 // Look for the first valid index after the provided debug MI.
537 MachineBasicBlock::const_iterator It = MI.getIterator(),
538 MBBEnd = MI.getParent()->end();
539 MachineBasicBlock::const_iterator NonDbgMI =
540 skipDebugInstructionsForward(It, End: MBBEnd);
541 if (NonDbgMI == MBBEnd) {
542 // There are no non-debug instruction between MI and the end of the
543 // block, so we reset the tracker at the end of the block.
544 reset(MBB: *MI.getParent(), /*End=*/true);
545 return;
546 }
547 // MI is a debug instruction so register pressure before or after it is
548 // identical. Since we moved forward to finding a non-debug instruction
549 // in the block, we reset the tracker before that instruction i.e., at its
550 // base index.
551 reset(MRI, SI: LIS.getInstructionIndex(Instr: *NonDbgMI));
552}
553
554void GCNRPTracker::reset(const MachineBasicBlock &MBB, bool End) {
555 SlotIndex SI = End ? LIS.getSlotIndexes()->getMBBLastIdx(MBB: &MBB)
556 : LIS.getMBBStartIdx(mbb: &MBB);
557 reset(MRI: MBB.getParent()->getRegInfo(), SI);
558}
559
560void GCNRPTracker::reset(const MachineRegisterInfo &MRI, SlotIndex SI) {
561 this->MRI = &MRI;
562 LastTrackedMI = nullptr;
563 LiveRegs = llvm::getLiveRegs(SI, LIS, MRI);
564 MaxPressure = CurPressure = getRegPressure(MRI, LiveRegs);
565}
566
567void GCNRPTracker::reset(const MachineRegisterInfo &MRI,
568 const LiveRegSet &LiveRegs) {
569 this->MRI = &MRI;
570 LastTrackedMI = nullptr;
571 if (&this->LiveRegs != &LiveRegs)
572 this->LiveRegs = LiveRegs;
573 MaxPressure = CurPressure = getRegPressure(MRI, LiveRegs);
574}
575
576/// Mostly copy/paste from CodeGen/RegisterPressure.cpp
577LaneBitmask GCNRPTracker::getLastUsedLanes(Register Reg, SlotIndex Pos) const {
578 return getLanesWithProperty(
579 LIS, MRI: *MRI, TrackLaneMasks: true, Reg, Pos: Pos.getBaseIndex(),
580 Property: [](const LiveRange &LR, SlotIndex Pos) {
581 const LiveRange::Segment *S = LR.getSegmentContaining(Idx: Pos);
582 return S != nullptr && S->end == Pos.getRegSlot();
583 });
584}
585
586////////////////////////////////////////////////////////////////////////////////
587// GCNUpwardRPTracker
588
589void GCNUpwardRPTracker::recede(const MachineInstr &MI) {
590 assert(MRI && "call reset first");
591
592 LastTrackedMI = &MI;
593
594 if (MI.isDebugInstr())
595 return;
596
597 // Kill all defs.
598 GCNRegPressure DefPressure, ECDefPressure;
599 bool HasECDefs = false;
600 for (const MachineOperand &MO : MI.all_defs()) {
601 if (!MO.getReg().isVirtual())
602 continue;
603
604 Register Reg = MO.getReg();
605 LaneBitmask DefMask = getDefRegMask(MO, MRI: *MRI);
606
607 // Treat a def as fully live at the moment of definition: keep a record.
608 if (MO.isEarlyClobber()) {
609 ECDefPressure.inc(Reg, PrevMask: LaneBitmask::getNone(), NewMask: DefMask, MRI: *MRI);
610 HasECDefs = true;
611 } else
612 DefPressure.inc(Reg, PrevMask: LaneBitmask::getNone(), NewMask: DefMask, MRI: *MRI);
613
614 auto I = LiveRegs.find(Val: Reg);
615 if (I == LiveRegs.end())
616 continue;
617
618 LaneBitmask &LiveMask = I->second;
619 LaneBitmask PrevMask = LiveMask;
620 LiveMask &= ~DefMask;
621 CurPressure.inc(Reg, PrevMask, NewMask: LiveMask, MRI: *MRI);
622 if (LiveMask.none())
623 LiveRegs.erase(I);
624 }
625
626 // Update MaxPressure with defs pressure.
627 DefPressure += CurPressure;
628 if (HasECDefs)
629 DefPressure += ECDefPressure;
630 MaxPressure = max(P1: DefPressure, P2: MaxPressure);
631
632 // Make uses alive.
633 SmallVector<VRegMaskOrUnit, 8> RegUses;
634 collectVirtualRegUses(VRegMaskOrUnits&: RegUses, MI, LIS, MRI: *MRI);
635 for (const VRegMaskOrUnit &U : RegUses) {
636 LaneBitmask &LiveMask = LiveRegs[U.VRegOrUnit.asVirtualReg()];
637 LaneBitmask PrevMask = LiveMask;
638 LiveMask |= U.LaneMask;
639 CurPressure.inc(Reg: U.VRegOrUnit.asVirtualReg(), PrevMask, NewMask: LiveMask, MRI: *MRI);
640 }
641
642 // Update MaxPressure with uses plus early-clobber defs pressure.
643 MaxPressure = HasECDefs ? max(P1: CurPressure + ECDefPressure, P2: MaxPressure)
644 : max(P1: CurPressure, P2: MaxPressure);
645
646 assert(CurPressure == getRegPressure(*MRI, LiveRegs));
647}
648
649////////////////////////////////////////////////////////////////////////////////
650// GCNDownwardRPTracker
651
652bool GCNDownwardRPTracker::reset(const MachineInstr &MI,
653 MachineBasicBlock::const_iterator End,
654 const LiveRegSet *LiveRegsCopy) {
655 MBBEnd = MI.getParent()->end();
656 assert((End == MBBEnd || End->getParent()->end() == MBBEnd) &&
657 "end unrelated to MI block");
658 NextMI = &MI;
659 NextMI = skipDebugInstructionsForward(It: NextMI, End);
660
661 // Do not use the MI to compute live registers when a set is provided.
662 // Otherwise the first non-debug instruction after the provided one (or the
663 // end of the block, if no such instruction exists) serves as the basis to
664 // compute a live register set.
665 if (LiveRegsCopy)
666 GCNRPTracker::reset(MRI: MI.getMF()->getRegInfo(), LiveRegs: *LiveRegsCopy);
667 else if (NextMI != MBBEnd)
668 GCNRPTracker::reset(MI: *NextMI, /*After=*/false);
669 else
670 GCNRPTracker::reset(MBB: *MI.getParent(), /*End=*/true);
671 return NextMI != End;
672}
673
674void GCNDownwardRPTracker::retireVirtReg(Register Reg, SlotIndex SI) {
675 const LiveInterval &LI = LIS.getInterval(Reg);
676 if (LI.hasSubRanges()) {
677 auto It = LiveRegs.end();
678 for (const auto &S : LI.subranges()) {
679 if (!S.liveAt(index: SI)) {
680 if (It == LiveRegs.end()) {
681 It = LiveRegs.find(Val: Reg);
682 if (It == LiveRegs.end())
683 llvm_unreachable("register isn't live");
684 }
685 auto PrevMask = It->second;
686 It->second &= ~S.LaneMask;
687 CurPressure.inc(Reg, PrevMask, NewMask: It->second, MRI: *MRI);
688 }
689 }
690 if (It != LiveRegs.end() && It->second.none())
691 LiveRegs.erase(I: It);
692 } else if (!LI.liveAt(index: SI)) {
693 auto It = LiveRegs.find(Val: Reg);
694 if (It == LiveRegs.end())
695 llvm_unreachable("register isn't live");
696 CurPressure.inc(Reg, PrevMask: It->second, NewMask: LaneBitmask::getNone(), MRI: *MRI);
697 LiveRegs.erase(I: It);
698 }
699}
700
701bool GCNDownwardRPTracker::advanceBeforeNext(MachineInstr *MI,
702 bool UseInternalIterator) {
703 assert(MRI && "call reset first");
704 SlotIndex SI;
705 const MachineInstr *CurrMI;
706 if (UseInternalIterator) {
707 if (!LastTrackedMI)
708 return NextMI == MBBEnd;
709
710 assert(NextMI == MBBEnd || !NextMI->isDebugInstr());
711 CurrMI = LastTrackedMI;
712
713 SI = NextMI == MBBEnd
714 ? LIS.getInstructionIndex(Instr: *LastTrackedMI).getDeadSlot()
715 : LIS.getInstructionIndex(Instr: *NextMI).getBaseIndex();
716 } else { //! UseInternalIterator
717 SI = LIS.getInstructionIndex(Instr: *MI).getBaseIndex();
718 CurrMI = MI;
719 }
720
721 assert(SI.isValid());
722
723 // Remove dead registers or mask bits.
724 SmallSet<Register, 8> SeenRegs;
725 for (auto &MO : CurrMI->operands()) {
726 if (!MO.isReg() || !MO.getReg().isVirtual())
727 continue;
728 if (MO.isUse() && CurrMI->getOpcode() == AMDGPU::PHI)
729 break;
730 if (MO.isUse() && !MO.readsReg())
731 continue;
732 if (!UseInternalIterator && MO.isDef())
733 continue;
734 if (!SeenRegs.insert(V: MO.getReg()).second)
735 continue;
736 retireVirtReg(Reg: MO.getReg(), SI);
737 }
738
739 MaxPressure = max(P1: MaxPressure, P2: CurPressure);
740
741 LastTrackedMI = nullptr;
742
743 return UseInternalIterator && (NextMI == MBBEnd);
744}
745
746void GCNDownwardRPTracker::advanceToNext(MachineInstr *MI,
747 bool UseInternalIterator) {
748 if (UseInternalIterator) {
749 LastTrackedMI = &*NextMI++;
750 NextMI = skipDebugInstructionsForward(It: NextMI, End: MBBEnd);
751 } else {
752 LastTrackedMI = MI;
753 }
754
755 const MachineInstr *CurrMI = LastTrackedMI;
756
757 // Add new registers or mask bits.
758 for (const auto &MO : CurrMI->all_defs()) {
759 Register Reg = MO.getReg();
760 if (!Reg.isVirtual())
761 continue;
762 auto &LiveMask = LiveRegs[Reg];
763 auto PrevMask = LiveMask;
764 LiveMask |= getDefRegMask(MO, MRI: *MRI);
765 CurPressure.inc(Reg, PrevMask, NewMask: LiveMask, MRI: *MRI);
766 }
767
768 MaxPressure = max(P1: MaxPressure, P2: CurPressure);
769}
770
771bool GCNDownwardRPTracker::advance(MachineInstr *MI, bool UseInternalIterator) {
772 if (UseInternalIterator && NextMI == MBBEnd)
773 return false;
774
775 advanceBeforeNext(MI, UseInternalIterator);
776 advanceToNext(MI, UseInternalIterator);
777 if (!UseInternalIterator) {
778 const MachineInstr *SavedLastTrackedMI = LastTrackedMI;
779 // We must remove any dead def lanes from the current RP
780 advanceBeforeNext(MI, UseInternalIterator: true);
781 // Restore LastTrackedMI set by advanceToNext, otherwise
782 // speculative queries (bumpDownwardPressure) don't
783 // know the last scheduled instruction and fail to
784 // correctly estimate pressure change.
785 LastTrackedMI = SavedLastTrackedMI;
786 }
787 return true;
788}
789
790bool GCNDownwardRPTracker::advance(MachineBasicBlock::const_iterator End) {
791 bool AnyAdvance = false;
792 while (NextMI != End && advance())
793 AnyAdvance = true;
794 return AnyAdvance;
795}
796
797bool GCNDownwardRPTracker::advance(MachineBasicBlock::const_iterator Begin,
798 MachineBasicBlock::const_iterator End,
799 const LiveRegSet *LiveRegsCopy) {
800 if (!reset(MI: *Begin, End, LiveRegsCopy))
801 return false;
802 return advance(End);
803}
804
805Printable llvm::reportMismatch(const GCNRPTracker::LiveRegSet &LISLR,
806 const GCNRPTracker::LiveRegSet &TrackedLR,
807 const TargetRegisterInfo *TRI, StringRef Pfx) {
808 return Printable([&LISLR, &TrackedLR, TRI, Pfx](raw_ostream &OS) {
809 for (auto const &P : TrackedLR) {
810 auto I = LISLR.find(Val: P.first);
811 if (I == LISLR.end()) {
812 OS << Pfx << printReg(Reg: P.first, TRI) << ":L" << PrintLaneMask(LaneMask: P.second)
813 << " isn't found in LIS reported set\n";
814 } else if (I->second != P.second) {
815 OS << Pfx << printReg(Reg: P.first, TRI)
816 << " masks doesn't match: LIS reported " << PrintLaneMask(LaneMask: I->second)
817 << ", tracked " << PrintLaneMask(LaneMask: P.second) << '\n';
818 }
819 }
820 for (auto const &P : LISLR) {
821 auto I = TrackedLR.find(Val: P.first);
822 if (I == TrackedLR.end()) {
823 OS << Pfx << printReg(Reg: P.first, TRI) << ":L" << PrintLaneMask(LaneMask: P.second)
824 << " isn't found in tracked set\n";
825 }
826 }
827 });
828}
829
830GCNRegPressure
831GCNDownwardRPTracker::bumpDownwardPressure(const MachineInstr *MI,
832 const SIRegisterInfo *TRI) const {
833 assert(!MI->isDebugOrPseudoInstr() && "Expect a nondebug instruction.");
834
835 SlotIndex SlotIdx;
836 SlotIdx = LIS.getInstructionIndex(Instr: *MI).getRegSlot();
837
838 SlotIndex CurrIdx;
839 const MachineBasicBlock *MBB = MI->getParent();
840 MachineBasicBlock::const_iterator StartPos =
841 LastTrackedMI ? std::next(x: LastTrackedMI->getIterator()) : MBB->begin();
842 MachineBasicBlock::const_iterator IdxPos =
843 skipDebugInstructionsForward(It: StartPos, End: MBB->end());
844 if (IdxPos == MBB->end()) {
845 CurrIdx = LIS.getMBBEndIdx(mbb: MBB);
846 } else {
847 CurrIdx = LIS.getInstructionIndex(Instr: *IdxPos).getRegSlot();
848 }
849
850 // Account for register pressure similar to RegPressureTracker::recede().
851 RegisterOperands RegOpers;
852 RegOpers.collect(MI: *MI, TRI: *TRI, MRI: *MRI, TrackLaneMasks: true, /*IgnoreDead=*/false);
853 RegOpers.adjustLaneLiveness(LIS, MRI: *MRI, Pos: SlotIdx);
854 GCNRegPressure TempPressure = CurPressure;
855 // Tracks the live mask reported by the use loop for redefined registers.
856 SmallDenseMap<Register, LaneBitmask, 8> PostUseMask;
857
858 for (const VRegMaskOrUnit &Use : RegOpers.Uses) {
859 if (!Use.VRegOrUnit.isVirtualReg())
860 continue;
861 Register Reg = Use.VRegOrUnit.asVirtualReg();
862 LaneBitmask LastUseMask = getLastUsedLanes(Reg, Pos: SlotIdx);
863 if (LastUseMask.none())
864 continue;
865 // The LastUseMask is queried from the liveness information of instruction
866 // which may be further down the schedule. Some lanes may actually not be
867 // last uses for the current position.
868 // FIXME: allow the caller to pass in the list of vreg uses that remain
869 // to be bottom-scheduled to avoid searching uses at each query.
870 LastUseMask =
871 findUseBetween(Reg, LastUseMask, PriorUseIdx: CurrIdx, NextUseIdx: SlotIdx, MRI: *MRI, TRI, LIS: &LIS);
872 if (LastUseMask.none())
873 continue;
874
875 auto It = LiveRegs.find(Val: Reg);
876 LaneBitmask LiveMask = It != LiveRegs.end() ? It->second : LaneBitmask(0);
877 LaneBitmask NewMask = LiveMask & ~LastUseMask;
878 PostUseMask[Reg] = NewMask;
879 TempPressure.inc(Reg, PrevMask: LiveMask, NewMask, MRI: *MRI);
880 }
881
882 // Generate liveness for defs.
883 for (const VRegMaskOrUnit &Def : RegOpers.Defs) {
884 if (!Def.VRegOrUnit.isVirtualReg())
885 continue;
886 Register Reg = Def.VRegOrUnit.asVirtualReg();
887 auto PostIt = PostUseMask.find(Val: Reg);
888 LaneBitmask LiveMask;
889 if (PostIt != PostUseMask.end()) {
890 LiveMask = PostIt->second;
891 } else {
892 auto It = LiveRegs.find(Val: Reg);
893 LiveMask = It != LiveRegs.end() ? It->second : LaneBitmask(0);
894 }
895
896 LaneBitmask NewMask = LiveMask | Def.LaneMask;
897 TempPressure.inc(Reg, PrevMask: LiveMask, NewMask, MRI: *MRI);
898 }
899
900 return TempPressure;
901}
902
903bool GCNUpwardRPTracker::isValid() const {
904 const auto &SI = LIS.getInstructionIndex(Instr: *LastTrackedMI).getBaseIndex();
905 const auto LISLR = llvm::getLiveRegs(SI, LIS, MRI: *MRI);
906 const auto &TrackedLR = LiveRegs;
907
908 if (!isEqual(S1: LISLR, S2: TrackedLR)) {
909 dbgs() << "\nGCNUpwardRPTracker error: Tracked and"
910 " LIS reported livesets mismatch:\n"
911 << print(LiveRegs: LISLR, MRI: *MRI);
912 reportMismatch(LISLR, TrackedLR, TRI: MRI->getTargetRegisterInfo());
913 return false;
914 }
915
916 auto LISPressure = getRegPressure(MRI: *MRI, LiveRegs: LISLR);
917 if (LISPressure != CurPressure) {
918 dbgs() << "GCNUpwardRPTracker error: Pressure sets different\nTracked: "
919 << print(RP: CurPressure) << "LIS rpt: " << print(RP: LISPressure);
920 return false;
921 }
922 return true;
923}
924
925Printable llvm::print(const GCNRPTracker::LiveRegSet &LiveRegs,
926 const MachineRegisterInfo &MRI) {
927 return Printable([&LiveRegs, &MRI](raw_ostream &OS) {
928 const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo();
929 for (unsigned I = 0, E = MRI.getNumVirtRegs(); I != E; ++I) {
930 Register Reg = Register::index2VirtReg(Index: I);
931 auto It = LiveRegs.find(Val: Reg);
932 if (It != LiveRegs.end() && It->second.any())
933 OS << ' ' << printReg(Reg, TRI) << ':' << PrintLaneMask(LaneMask: It->second);
934 }
935 OS << '\n';
936 });
937}
938
939void GCNRegPressure::dump() const { dbgs() << print(RP: *this); }
940
941static cl::opt<bool> UseDownwardTracker(
942 "amdgpu-print-rp-downward",
943 cl::desc("Use GCNDownwardRPTracker for GCNRegPressurePrinter pass"),
944 cl::init(Val: false), cl::Hidden);
945
946char llvm::GCNRegPressurePrinter::ID = 0;
947char &llvm::GCNRegPressurePrinterID = GCNRegPressurePrinter::ID;
948
949INITIALIZE_PASS(GCNRegPressurePrinter, "amdgpu-print-rp", "", true, true)
950
951// Return lanemask of Reg's subregs that are live-through at [Begin, End] and
952// are fully covered by Mask.
953static LaneBitmask
954getRegLiveThroughMask(const MachineRegisterInfo &MRI, const LiveIntervals &LIS,
955 Register Reg, SlotIndex Begin, SlotIndex End,
956 LaneBitmask Mask = LaneBitmask::getAll()) {
957
958 auto IsInOneSegment = [Begin, End](const LiveRange &LR) -> bool {
959 auto *Segment = LR.getSegmentContaining(Idx: Begin);
960 return Segment && Segment->contains(I: End);
961 };
962
963 LaneBitmask LiveThroughMask;
964 const LiveInterval &LI = LIS.getInterval(Reg);
965 if (LI.hasSubRanges()) {
966 for (auto &SR : LI.subranges()) {
967 if ((SR.LaneMask & Mask) == SR.LaneMask && IsInOneSegment(SR))
968 LiveThroughMask |= SR.LaneMask;
969 }
970 } else {
971 LaneBitmask RegMask = MRI.getMaxLaneMaskForVReg(Reg);
972 if ((RegMask & Mask) == RegMask && IsInOneSegment(LI))
973 LiveThroughMask = RegMask;
974 }
975
976 return LiveThroughMask;
977}
978
979bool GCNRegPressurePrinter::runOnMachineFunction(MachineFunction &MF) {
980 const MachineRegisterInfo &MRI = MF.getRegInfo();
981 const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo();
982 LiveIntervals &LIS = getAnalysis<LiveIntervalsWrapperPass>().getLIS();
983
984 auto &OS = dbgs();
985
986// Leading spaces are important for YAML syntax.
987#define PFX " "
988
989 OS << "---\nname: " << MF.getName() << "\nbody: |\n";
990
991 auto printRP = [](const GCNRegPressure &RP) {
992 return Printable([&RP](raw_ostream &OS) {
993 OS << format(PFX " %-5d", Vals: RP.getSGPRNum())
994 << format(Fmt: " %-5d", Vals: RP.getVGPRNum(UnifiedVGPRFile: false));
995 });
996 };
997
998 auto ReportLISMismatchIfAny = [&](const GCNRPTracker::LiveRegSet &TrackedLR,
999 const GCNRPTracker::LiveRegSet &LISLR) {
1000 if (LISLR != TrackedLR) {
1001 OS << PFX " mis LIS: " << llvm::print(LiveRegs: LISLR, MRI)
1002 << reportMismatch(LISLR, TrackedLR, TRI, PFX " ");
1003 }
1004 };
1005
1006 // Register pressure before and at an instruction (in program order).
1007 SmallVector<std::pair<GCNRegPressure, GCNRegPressure>, 16> RP;
1008
1009 for (auto &MBB : MF) {
1010 RP.clear();
1011 RP.reserve(N: MBB.size());
1012
1013 OS << PFX;
1014 MBB.printName(os&: OS);
1015 OS << ":\n";
1016
1017 SlotIndex MBBStartSlot = LIS.getSlotIndexes()->getMBBStartIdx(mbb: &MBB);
1018 SlotIndex MBBLastSlot = LIS.getSlotIndexes()->getMBBLastIdx(MBB: &MBB);
1019
1020 GCNRPTracker::LiveRegSet LiveIn, LiveOut;
1021 GCNRegPressure RPAtMBBEnd;
1022
1023 if (UseDownwardTracker) {
1024 if (MBB.empty()) {
1025 LiveIn = LiveOut = getLiveRegs(SI: MBBStartSlot, LIS, MRI);
1026 RPAtMBBEnd = getRegPressure(MRI, LiveRegs&: LiveIn);
1027 } else {
1028 GCNDownwardRPTracker RPT(LIS);
1029 RPT.reset(MI: MBB.front(), End: MBB.end());
1030
1031 LiveIn = RPT.getLiveRegs();
1032
1033 while (!RPT.advanceBeforeNext()) {
1034 GCNRegPressure RPBeforeMI = RPT.getPressure();
1035 RPT.resetMaxPressure();
1036 RPT.advanceToNext();
1037 RP.emplace_back(Args&: RPBeforeMI, Args: RPT.getMaxPressure());
1038 }
1039
1040 LiveOut = RPT.getLiveRegs();
1041 RPAtMBBEnd = RPT.getPressure();
1042 }
1043 } else {
1044 GCNUpwardRPTracker RPT(LIS);
1045 RPT.reset(MRI, SI: MBBLastSlot);
1046
1047 LiveOut = RPT.getLiveRegs();
1048 RPAtMBBEnd = RPT.getPressure();
1049
1050 for (auto &MI : reverse(C&: MBB)) {
1051 RPT.resetMaxPressure();
1052 RPT.recede(MI);
1053 if (!MI.isDebugInstr())
1054 RP.emplace_back(Args: RPT.getPressure(), Args: RPT.getMaxPressure());
1055 }
1056
1057 LiveIn = RPT.getLiveRegs();
1058 }
1059
1060 OS << PFX " Live-in: " << llvm::print(LiveRegs: LiveIn, MRI);
1061 if (!UseDownwardTracker)
1062 ReportLISMismatchIfAny(LiveIn, getLiveRegs(SI: MBBStartSlot, LIS, MRI));
1063
1064 OS << PFX " SGPR VGPR\n";
1065 int I = 0;
1066 for (auto &MI : MBB) {
1067 if (!MI.isDebugInstr()) {
1068 auto &[RPBeforeInstr, RPAtInstr] =
1069 RP[UseDownwardTracker ? I : (RP.size() - 1 - I)];
1070 ++I;
1071 OS << printRP(RPBeforeInstr) << '\n' << printRP(RPAtInstr) << " ";
1072 } else
1073 OS << PFX " ";
1074 MI.print(OS);
1075 }
1076 OS << printRP(RPAtMBBEnd) << '\n';
1077
1078 OS << PFX " Live-out:" << llvm::print(LiveRegs: LiveOut, MRI);
1079 if (UseDownwardTracker)
1080 ReportLISMismatchIfAny(LiveOut, getLiveRegs(SI: MBBLastSlot, LIS, MRI));
1081
1082 GCNRPTracker::LiveRegSet LiveThrough;
1083 for (auto [Reg, Mask] : LiveIn) {
1084 LaneBitmask MaskIntersection = Mask & LiveOut.lookup(Val: Reg);
1085 if (MaskIntersection.any()) {
1086 LaneBitmask LTMask = getRegLiveThroughMask(
1087 MRI, LIS, Reg, Begin: MBBStartSlot, End: MBBLastSlot, Mask: MaskIntersection);
1088 if (LTMask.any())
1089 LiveThrough[Reg] = LTMask;
1090 }
1091 }
1092 OS << PFX " Live-thr:" << llvm::print(LiveRegs: LiveThrough, MRI);
1093 OS << printRP(getRegPressure(MRI, LiveRegs&: LiveThrough)) << '\n';
1094 }
1095 OS << "...\n";
1096 return false;
1097
1098#undef PFX
1099}
1100
1101#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
1102LLVM_DUMP_METHOD void llvm::dumpMaxRegPressure(MachineFunction &MF,
1103 GCNRegPressure::RegKind Kind,
1104 LiveIntervals &LIS,
1105 const MachineLoopInfo *MLI) {
1106
1107 const MachineRegisterInfo &MRI = MF.getRegInfo();
1108 const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo();
1109 auto &OS = dbgs();
1110 const char *RegName = GCNRegPressure::getName(Kind);
1111
1112 unsigned MaxNumRegs = 0;
1113 const MachineInstr *MaxPressureMI = nullptr;
1114 GCNUpwardRPTracker RPT(LIS);
1115 for (const MachineBasicBlock &MBB : MF) {
1116 RPT.reset(MRI, LIS.getSlotIndexes()->getMBBEndIdx(&MBB).getPrevSlot());
1117 for (const MachineInstr &MI : reverse(MBB)) {
1118 RPT.recede(MI);
1119 unsigned NumRegs = RPT.getMaxPressure().getNumRegs(Kind);
1120 if (NumRegs > MaxNumRegs) {
1121 MaxNumRegs = NumRegs;
1122 MaxPressureMI = &MI;
1123 }
1124 }
1125 }
1126
1127 SlotIndex MISlot = LIS.getInstructionIndex(*MaxPressureMI);
1128
1129 // Max pressure can occur at either the early-clobber or register slot.
1130 // Choose the maximum liveset between both slots. This is ugly but this is
1131 // diagnostic code.
1132 SlotIndex ECSlot = MISlot.getRegSlot(true);
1133 SlotIndex RSlot = MISlot.getRegSlot(false);
1134 GCNRPTracker::LiveRegSet ECLiveSet = getLiveRegs(ECSlot, LIS, MRI, Kind);
1135 GCNRPTracker::LiveRegSet RLiveSet = getLiveRegs(RSlot, LIS, MRI, Kind);
1136 unsigned ECNumRegs = getRegPressure(MRI, ECLiveSet).getNumRegs(Kind);
1137 unsigned RNumRegs = getRegPressure(MRI, RLiveSet).getNumRegs(Kind);
1138 GCNRPTracker::LiveRegSet *LiveSet =
1139 ECNumRegs > RNumRegs ? &ECLiveSet : &RLiveSet;
1140 SlotIndex MaxPressureSlot = ECNumRegs > RNumRegs ? ECSlot : RSlot;
1141 assert(getRegPressure(MRI, *LiveSet).getNumRegs(Kind) == MaxNumRegs);
1142
1143 // Split live registers into single-def and multi-def sets.
1144 GCNRegPressure SDefPressure, MDefPressure;
1145 SmallVector<Register, 16> SDefRegs, MDefRegs;
1146 for (auto [Reg, LaneMask] : *LiveSet) {
1147 assert(GCNRegPressure::getRegKind(Reg, MRI) == Kind);
1148 LiveInterval &LI = LIS.getInterval(Reg);
1149 if (LI.getNumValNums() == 1 ||
1150 (LI.hasSubRanges() &&
1151 llvm::all_of(LI.subranges(), [](const LiveInterval::SubRange &SR) {
1152 return SR.getNumValNums() == 1;
1153 }))) {
1154 SDefPressure.inc(Reg, LaneBitmask::getNone(), LaneMask, MRI);
1155 SDefRegs.push_back(Reg);
1156 } else {
1157 MDefPressure.inc(Reg, LaneBitmask::getNone(), LaneMask, MRI);
1158 MDefRegs.push_back(Reg);
1159 }
1160 }
1161 unsigned SDefNumRegs = SDefPressure.getNumRegs(Kind);
1162 unsigned MDefNumRegs = MDefPressure.getNumRegs(Kind);
1163 assert(SDefNumRegs + MDefNumRegs == MaxNumRegs);
1164
1165 auto printLoc = [&](const MachineBasicBlock *MBB, SlotIndex SI) {
1166 return Printable([&, MBB, SI](raw_ostream &OS) {
1167 OS << SI << ':' << printMBBReference(*MBB);
1168 if (MLI)
1169 if (const MachineLoop *ML = MLI->getLoopFor(MBB))
1170 OS << " (LoopHdr " << printMBBReference(*ML->getHeader())
1171 << ", Depth " << ML->getLoopDepth() << ")";
1172 });
1173 };
1174
1175 auto PrintRegInfo = [&](Register Reg, LaneBitmask LiveMask) {
1176 GCNRegPressure RegPressure;
1177 RegPressure.inc(Reg, LaneBitmask::getNone(), LiveMask, MRI);
1178 OS << " " << printReg(Reg, TRI) << ':'
1179 << TRI->getRegClassName(MRI.getRegClass(Reg)) << ", LiveMask "
1180 << PrintLaneMask(LiveMask) << " (" << RegPressure.getNumRegs(Kind) << ' '
1181 << RegName << "s)\n";
1182
1183 // Use std::map to sort def/uses by SlotIndex.
1184 std::map<SlotIndex, const MachineInstr *> Instrs;
1185 for (const MachineInstr &MI : MRI.reg_nodbg_instructions(Reg)) {
1186 Instrs[LIS.getInstructionIndex(MI).getRegSlot()] = &MI;
1187 }
1188
1189 for (const auto &[SI, MI] : Instrs) {
1190 OS << " ";
1191 if (MI->definesRegister(Reg, TRI))
1192 OS << "def ";
1193 if (MI->readsRegister(Reg, TRI))
1194 OS << "use ";
1195 OS << printLoc(MI->getParent(), SI) << ": " << *MI;
1196 }
1197 };
1198
1199 OS << "\n*** Register pressure info (" << RegName << "s) for " << MF.getName()
1200 << " ***\n";
1201 OS << "Max pressure is " << MaxNumRegs << ' ' << RegName << "s at "
1202 << printLoc(MaxPressureMI->getParent(), MaxPressureSlot) << ": "
1203 << *MaxPressureMI;
1204
1205 OS << "\nLive registers with single definition (" << SDefNumRegs << ' '
1206 << RegName << "s):\n";
1207
1208 // Sort SDefRegs by number of uses (smallest first)
1209 llvm::sort(SDefRegs, [&](Register A, Register B) {
1210 return std::distance(MRI.use_nodbg_begin(A), MRI.use_nodbg_end()) <
1211 std::distance(MRI.use_nodbg_begin(B), MRI.use_nodbg_end());
1212 });
1213
1214 for (const Register Reg : SDefRegs) {
1215 PrintRegInfo(Reg, LiveSet->lookup(Reg));
1216 }
1217
1218 OS << "\nLive registers with multiple definitions (" << MDefNumRegs << ' '
1219 << RegName << "s):\n";
1220 for (const Register Reg : MDefRegs) {
1221 PrintRegInfo(Reg, LiveSet->lookup(Reg));
1222 }
1223}
1224#endif
1225
1226unsigned llvm::estimateGreedyVGPRPressure(
1227 MachineBasicBlock::const_iterator RegionBegin,
1228 MachineBasicBlock::const_iterator RegionEnd,
1229 const GCNRPTracker::LiveRegSet &LiveIns, const LiveIntervals &LIS,
1230 const MachineRegisterInfo &MRI, const SIRegisterInfo &TRI) {
1231
1232 SetVector<const LiveInterval *> IntervalSet;
1233 IntervalSet.reserve(Size: LiveIns.size());
1234
1235 auto checkAndCollect = [&](Register VReg) {
1236 if (!VReg.isVirtual() || !LIS.hasInterval(Reg: VReg))
1237 return;
1238
1239 const TargetRegisterClass *RC = MRI.getRegClass(Reg: VReg);
1240 if (!TRI.hasVGPRs(RC))
1241 return;
1242
1243 const LiveInterval &LI = LIS.getInterval(Reg: VReg);
1244 IntervalSet.insert(X: &LI);
1245 };
1246
1247 // Collect live-ins.
1248 for (const auto &[RegNum, LaneMask] : LiveIns) {
1249 checkAndCollect(Register(RegNum));
1250 }
1251
1252 // Collect defs in region.
1253 for (MachineBasicBlock::const_iterator I = RegionBegin; I != RegionEnd; ++I) {
1254 for (const MachineOperand &MO : I->operands()) {
1255 if (!MO.isReg() || !MO.isDef())
1256 continue;
1257 checkAndCollect(MO.getReg());
1258 }
1259 }
1260
1261 SmallVector<const LiveInterval *> Intervals = IntervalSet.takeVector();
1262 llvm::sort(C&: Intervals, Comp: [](const LiveInterval *LHS, const LiveInterval *RHS) {
1263 return LHS->beginIndex() < RHS->beginIndex();
1264 });
1265
1266 LiveIntervalUnion::Allocator Alloc;
1267 std::vector<LiveIntervalUnion> RegFile;
1268 unsigned MaxRegsUsed = 0;
1269
1270 // Simulate greedy register allocation, assuming an unlimited number of
1271 // physical registers.
1272 for (const LiveInterval *LI : Intervals) {
1273 const TargetRegisterClass *RC = MRI.getRegClass(Reg: LI->reg());
1274 unsigned Width =
1275 std::max<unsigned>(a: 1, b: TRI.getRegSizeInBits(RC: *RC).getFixedValue() / 32);
1276 unsigned Alignment =
1277 std::max<unsigned>(a: 1, b: TRI.getRegClassAlignmentNumBits(RC) / 32);
1278
1279 unsigned Start = 0;
1280 while (true) {
1281 unsigned End = Start + Width;
1282 if (RegFile.size() < End)
1283 RegFile.resize(new_size: End, x: LiveIntervalUnion(Alloc));
1284
1285 bool Fits = true;
1286 for (unsigned Idx = Start; Idx < End; Idx++) {
1287 LiveIntervalUnion::Query Q(*LI, RegFile[Idx]);
1288 if (Q.checkInterference()) {
1289 Start = alignTo(Value: Idx + 1, Align: Alignment);
1290 Fits = false;
1291 break;
1292 }
1293 }
1294
1295 if (Fits) {
1296 for (unsigned Idx = Start; Idx < End; Idx++)
1297 RegFile[Idx].unify(VirtReg: *LI, Range: *LI);
1298 MaxRegsUsed = std::max(a: MaxRegsUsed, b: End);
1299 break;
1300 }
1301 }
1302 }
1303
1304 return MaxRegsUsed;
1305}
1306