1//===-- GCNSchedStrategy.cpp - GCN Scheduler Strategy ---------------------===//
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 contains a MachineSchedStrategy implementation for maximizing wave
11/// occupancy on GCN hardware.
12///
13/// This pass will apply multiple scheduling stages to the same function.
14/// Regions are first recorded in GCNScheduleDAGMILive::schedule. The actual
15/// entry point for the scheduling of those regions is
16/// GCNScheduleDAGMILive::runSchedStages.
17
18/// Generally, the reason for having multiple scheduling stages is to account
19/// for the kernel-wide effect of register usage on occupancy. Usually, only a
20/// few scheduling regions will have register pressure high enough to limit
21/// occupancy for the kernel, so constraints can be relaxed to improve ILP in
22/// other regions.
23///
24//===----------------------------------------------------------------------===//
25
26#include "GCNSchedStrategy.h"
27#include "AMDGPUIGroupLP.h"
28#include "GCNHazardRecognizer.h"
29#include "GCNRegPressure.h"
30#include "SIMachineFunctionInfo.h"
31#include "Utils/AMDGPUBaseInfo.h"
32#include "llvm/ADT/BitVector.h"
33#include "llvm/ADT/STLExtras.h"
34#include "llvm/CodeGen/CalcSpillWeights.h"
35#include "llvm/CodeGen/MachineBasicBlock.h"
36#include "llvm/CodeGen/MachineBlockFrequencyInfo.h"
37#include "llvm/CodeGen/MachineBranchProbabilityInfo.h"
38#include "llvm/CodeGen/MachineCycleAnalysis.h"
39#include "llvm/CodeGen/MachineOperand.h"
40#include "llvm/CodeGen/RegisterPressure.h"
41#include "llvm/CodeGen/Rematerializer.h"
42#include "llvm/MC/LaneBitmask.h"
43#include "llvm/MC/MCSchedule.h"
44#include "llvm/MC/TargetRegistry.h"
45#include "llvm/Support/ErrorHandling.h"
46
47#define DEBUG_TYPE "machine-scheduler"
48
49using namespace llvm;
50
51static cl::opt<bool> DisableUnclusterHighRP(
52 "amdgpu-disable-unclustered-high-rp-reschedule", cl::Hidden,
53 cl::desc("Disable unclustered high register pressure "
54 "reduction scheduling stage."),
55 cl::init(Val: false));
56
57static cl::opt<bool> DisableClusteredLowOccupancy(
58 "amdgpu-disable-clustered-low-occupancy-reschedule", cl::Hidden,
59 cl::desc("Disable clustered low occupancy "
60 "rescheduling for ILP scheduling stage."),
61 cl::init(Val: false));
62
63static cl::opt<unsigned> ScheduleMetricBias(
64 "amdgpu-schedule-metric-bias", cl::Hidden,
65 cl::desc(
66 "Sets the bias which adds weight to occupancy vs latency. Set it to "
67 "100 to chase the occupancy only."),
68 cl::init(Val: 10));
69
70static cl::opt<bool>
71 RelaxedOcc("amdgpu-schedule-relaxed-occupancy", cl::Hidden,
72 cl::desc("Relax occupancy targets for kernels which are memory "
73 "bound (amdgpu-membound-threshold), or "
74 "Wave Limited (amdgpu-limit-wave-threshold)."),
75 cl::init(Val: false));
76
77static cl::opt<bool> GCNTrackers(
78 "amdgpu-use-amdgpu-trackers", cl::Hidden,
79 cl::desc("Use the AMDGPU specific RPTrackers during scheduling"),
80 cl::init(Val: false));
81
82static cl::opt<unsigned> PendingQueueLimit(
83 "amdgpu-scheduler-pending-queue-limit", cl::Hidden,
84 cl::desc(
85 "Max (Available+Pending) size to inspect pending queue (0 disables)"),
86 cl::init(Val: 256));
87
88#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
89#define DUMP_MAX_REG_PRESSURE
90static cl::opt<bool> PrintMaxRPRegUsageBeforeScheduler(
91 "amdgpu-print-max-reg-pressure-regusage-before-scheduler", cl::Hidden,
92 cl::desc("Print a list of live registers along with their def/uses at the "
93 "point of maximum register pressure before scheduling."),
94 cl::init(false));
95
96static cl::opt<bool> PrintMaxRPRegUsageAfterScheduler(
97 "amdgpu-print-max-reg-pressure-regusage-after-scheduler", cl::Hidden,
98 cl::desc("Print a list of live registers along with their def/uses at the "
99 "point of maximum register pressure after scheduling."),
100 cl::init(false));
101#endif
102
103static cl::opt<bool> DisableRewriteMFMAFormSchedStage(
104 "amdgpu-disable-rewrite-mfma-form-sched-stage", cl::Hidden,
105 cl::desc("Disable rewrite mfma rewrite scheduling stage"), cl::init(Val: true));
106
107bool VGPRThresholdParser::parse(cl::Option &O, StringRef ArgName, StringRef Arg,
108 unsigned &Value) {
109 if (Arg.getAsInteger(Radix: 0, Result&: Value))
110 return O.error(Message: "'" + Arg + "' value invalid for uint argument!");
111
112 if (Value > 100)
113 return O.error(Message: "'" + Arg + "' value must be in the range [0, 100]!");
114
115 return false;
116}
117
118cl::opt<unsigned, false, VGPRThresholdParser> llvm::VGPRThresholdPercentOpt(
119 "amdgpu-vgpr-threshold-percent", cl::Hidden,
120 cl::desc("Percent of VGPR limits that we should use as RP threshold "
121 "during scheduling. We have two limits relevant to scheduling: "
122 "Critical (avoid decreasing occupancy), Excess (avoid spilling). "
123 "This flag scales both limits back by an equal percent: (0 = use "
124 " default calculation, 1-100 = use percentage), default: 0"),
125 cl::init(Val: 0));
126
127const unsigned ScheduleMetrics::ScaleFactor = 100;
128
129GCNSchedStrategy::GCNSchedStrategy(const MachineSchedContext *C)
130 : GenericScheduler(C), TargetOccupancy(0), MF(nullptr),
131 DownwardTracker(*C->LIS), UpwardTracker(*C->LIS), HasHighPressure(false) {
132 if (GCNTrackers.getNumOccurrences() > 0)
133 GCNTrackersOverride = GCNTrackers;
134 VGPRThresholdPercent = VGPRThresholdPercentOpt;
135}
136
137void GCNSchedStrategy::initialize(ScheduleDAGMI *DAG) {
138 GenericScheduler::initialize(dag: DAG);
139
140 MF = &DAG->MF;
141
142 const GCNSubtarget &ST = MF->getSubtarget<GCNSubtarget>();
143
144 SGPRExcessLimit =
145 Context->RegClassInfo->getNumAllocatableRegs(RC: &AMDGPU::SGPR_32RegClass);
146 VGPRExcessLimit =
147 Context->RegClassInfo->getNumAllocatableRegs(RC: &AMDGPU::VGPR_32RegClass);
148 AGPRExcessLimit =
149 Context->RegClassInfo->getNumAllocatableRegs(RC: &AMDGPU::AGPR_32RegClass);
150
151 SIMachineFunctionInfo &MFI = *MF->getInfo<SIMachineFunctionInfo>();
152 // Set the initial TargetOccupnacy to the maximum occupancy that we can
153 // achieve for this function. This effectively sets a lower bound on the
154 // 'Critical' register limits in the scheduler.
155 // Allow for lower occupancy targets if kernel is wave limited or memory
156 // bound, and using the relaxed occupancy feature.
157 TargetOccupancy =
158 RelaxedOcc ? MFI.getMinAllowedOccupancy() : MFI.getOccupancy();
159 SGPRCriticalLimit =
160 std::min(a: ST.getMaxNumSGPRs(WavesPerEU: TargetOccupancy, Addressable: true), b: SGPRExcessLimit);
161
162 if (!KnownExcessRP) {
163 VGPRCriticalLimit = std::min(
164 a: ST.getMaxNumVGPRs(WavesPerEU: TargetOccupancy, DynamicVGPRBlockSize: MFI.getDynamicVGPRBlockSize()),
165 b: VGPRExcessLimit);
166 } else {
167 // This is similar to ST.getMaxNumVGPRs(TargetOccupancy) result except
168 // returns a reasonably small number for targets with lots of VGPRs, such
169 // as GFX10 and GFX11.
170 LLVM_DEBUG(dbgs() << "Region is known to spill, use alternative "
171 "VGPRCriticalLimit calculation method.\n");
172 unsigned DynamicVGPRBlockSize = MFI.getDynamicVGPRBlockSize();
173 unsigned Granule =
174 AMDGPU::IsaInfo::getVGPRAllocGranule(STI: ST, DynamicVGPRBlockSize);
175 unsigned Addressable =
176 AMDGPU::IsaInfo::getAddressableNumVGPRs(STI: ST, DynamicVGPRBlockSize);
177 unsigned VGPRBudget = alignDown(Value: Addressable / TargetOccupancy, Align: Granule);
178 VGPRBudget = std::max(a: VGPRBudget, b: Granule);
179 VGPRCriticalLimit = std::min(a: VGPRBudget, b: VGPRExcessLimit);
180 }
181
182 // Reuse VGPR critical limit
183 AGPRCriticalLimit = std::min(a: VGPRCriticalLimit, b: AGPRExcessLimit);
184
185 // Apply VGPR excess threshold percentage if specified.
186 if (VGPRThresholdPercent > 0) {
187 [[maybe_unused]] unsigned OriginalVGPRExcessLimit = VGPRExcessLimit;
188 [[maybe_unused]] unsigned OriginalVGPRCriticalLimit = VGPRCriticalLimit;
189 VGPRExcessLimit = (VGPRThresholdPercent * VGPRExcessLimit + 99) / 100;
190 VGPRCriticalLimit = (VGPRThresholdPercent * VGPRCriticalLimit + 99) / 100;
191 LLVM_DEBUG(dbgs() << "Applied VGPR excess threshold "
192 << VGPRThresholdPercent << "%, VGPRExcessLimit: "
193 << OriginalVGPRExcessLimit << " -> " << VGPRExcessLimit
194 << ". VGPRCriticalLimit: " << OriginalVGPRCriticalLimit
195 << " -> " << VGPRCriticalLimit << '\n');
196 } else {
197 VGPRExcessLimit -= std::min(a: VGPRLimitBias + ErrorMargin, b: VGPRExcessLimit);
198 VGPRCriticalLimit -=
199 std::min(a: VGPRLimitBias + ErrorMargin, b: VGPRCriticalLimit);
200 }
201
202 // Subtract error margin and bias from register limits and avoid overflow.
203 SGPRCriticalLimit -= std::min(a: SGPRLimitBias + ErrorMargin, b: SGPRCriticalLimit);
204 SGPRExcessLimit -= std::min(a: SGPRLimitBias + ErrorMargin, b: SGPRExcessLimit);
205
206 AGPRExcessLimit -= std::min(a: VGPRLimitBias + ErrorMargin, b: AGPRExcessLimit);
207 AGPRCriticalLimit -= std::min(a: VGPRLimitBias + ErrorMargin, b: AGPRCriticalLimit);
208
209 LLVM_DEBUG(dbgs() << "VGPRCriticalLimit = " << VGPRCriticalLimit
210 << ", VGPRExcessLimit = " << VGPRExcessLimit
211 << ", AGPRCriticalLimit = " << AGPRCriticalLimit
212 << ", AGPRExcessLimit = " << AGPRExcessLimit
213 << ", SGPRCriticalLimit = " << SGPRCriticalLimit
214 << ", SGPRExcessLimit = " << SGPRExcessLimit << "\n\n");
215}
216
217/// Checks whether \p SU can use the cached DAG pressure diffs to compute the
218/// current register pressure.
219///
220/// This works for the common case, but it has a few exceptions that have been
221/// observed through trial and error:
222/// - Explicit physical register operands
223/// - Subregister definitions
224///
225/// In both of those cases, PressureDiff doesn't represent the actual pressure,
226/// and querying LiveIntervals through the RegPressureTracker is needed to get
227/// an accurate value.
228///
229/// We should eventually only use PressureDiff for maximum performance, but this
230/// already allows 80% of SUs to take the fast path without changing scheduling
231/// at all. Further changes would either change scheduling, or require a lot
232/// more logic to recover an accurate pressure estimate from the PressureDiffs.
233static bool canUsePressureDiffs(const SUnit &SU) {
234 if (!SU.isInstr())
235 return false;
236
237 // Cannot use pressure diffs for subregister defs or with physregs, it's
238 // imprecise in both cases. For a bundle, check the instructions inside it:
239 // the BUNDLE header only has implicit operands.
240 for (const auto &Op : const_mi_bundle_ops(MI: *SU.getInstr())) {
241 if (!Op.isReg() || Op.isImplicit())
242 continue;
243 if (Op.getReg().isPhysical() ||
244 (Op.isDef() && Op.getSubReg() != AMDGPU::NoSubRegister))
245 return false;
246 }
247 return true;
248}
249
250void GCNSchedStrategy::getRegisterPressures(
251 bool AtTop, const RegPressureTracker &RPTracker, SUnit *SU,
252 std::vector<unsigned> &Pressure, std::vector<unsigned> &MaxPressure,
253 GCNDownwardRPTracker &DownwardTracker, GCNUpwardRPTracker &UpwardTracker,
254 ScheduleDAGMI *DAG, const SIRegisterInfo *SRI) {
255 // getDownwardPressure() and getUpwardPressure() make temporary changes to
256 // the tracker, so we need to pass those function a non-const copy.
257 RegPressureTracker &TempTracker = const_cast<RegPressureTracker &>(RPTracker);
258 if (!useGCNTrackers()) {
259 AtTop
260 ? TempTracker.getDownwardPressure(MI: SU->getInstr(), PressureResult&: Pressure, MaxPressureResult&: MaxPressure)
261 : TempTracker.getUpwardPressure(MI: SU->getInstr(), PressureResult&: Pressure, MaxPressureResult&: MaxPressure);
262
263 return;
264 }
265
266 // GCNTrackers
267 Pressure.resize(new_size: 4, x: 0);
268 MachineInstr *MI = SU->getInstr();
269 GCNRegPressure NewPressure;
270 if (AtTop) {
271 GCNDownwardRPTracker TempDownwardTracker(DownwardTracker);
272 NewPressure = TempDownwardTracker.bumpDownwardPressure(MI, TRI: SRI);
273 } else {
274 GCNUpwardRPTracker TempUpwardTracker(UpwardTracker);
275 TempUpwardTracker.recede(MI: *MI);
276 NewPressure = TempUpwardTracker.getPressure();
277 }
278 Pressure[AMDGPU::RegisterPressureSets::SReg_32] = NewPressure.getSGPRNum();
279 Pressure[AMDGPU::RegisterPressureSets::VGPR_32] =
280 NewPressure.getArchVGPRNum();
281 Pressure[AMDGPU::RegisterPressureSets::AGPR_32] = NewPressure.getAGPRNum();
282}
283
284void GCNSchedStrategy::initCandidate(SchedCandidate &Cand, SUnit *SU,
285 bool AtTop,
286 const RegPressureTracker &RPTracker,
287 const SIRegisterInfo *SRI,
288 unsigned SGPRPressure,
289 unsigned VGPRPressure,
290 unsigned AGPRPressure, bool IsBottomUp) {
291 Cand.SU = SU;
292 Cand.AtTop = AtTop;
293
294 if (!DAG->isTrackingPressure())
295 return;
296
297 Pressure.clear();
298 MaxPressure.clear();
299
300 // We try to use the cached PressureDiffs in the ScheduleDAG whenever
301 // possible over querying the RegPressureTracker.
302 //
303 // RegPressureTracker will make a lot of LIS queries which are very
304 // expensive, it is considered a slow function in this context.
305 //
306 // PressureDiffs are precomputed and cached, and getPressureDiff is just a
307 // trivial lookup into an array. It is pretty much free.
308 //
309 // In EXPENSIVE_CHECKS, we always query RPTracker to verify the results of
310 // PressureDiffs.
311 if (AtTop || !canUsePressureDiffs(SU: *SU) || useGCNTrackers()) {
312 getRegisterPressures(AtTop, RPTracker, SU, Pressure, MaxPressure,
313 DownwardTracker, UpwardTracker, DAG, SRI);
314 } else {
315 // Reserve 4 slots.
316 Pressure.resize(new_size: 4, x: 0);
317 Pressure[AMDGPU::RegisterPressureSets::SReg_32] = SGPRPressure;
318 Pressure[AMDGPU::RegisterPressureSets::VGPR_32] = VGPRPressure;
319 Pressure[AMDGPU::RegisterPressureSets::AGPR_32] = AGPRPressure;
320
321 for (const auto &Diff : DAG->getPressureDiff(SU)) {
322 if (!Diff.isValid())
323 continue;
324 // PressureDiffs is always bottom-up so if we're working top-down we need
325 // to invert its sign.
326 Pressure[Diff.getPSet()] +=
327 (IsBottomUp ? Diff.getUnitInc() : -Diff.getUnitInc());
328 }
329
330#ifdef EXPENSIVE_CHECKS
331 std::vector<unsigned> CheckPressure, CheckMaxPressure;
332 getRegisterPressures(AtTop, RPTracker, SU, CheckPressure, CheckMaxPressure,
333 DownwardTracker, UpwardTracker, DAG, SRI);
334 if (Pressure[AMDGPU::RegisterPressureSets::SReg_32] !=
335 CheckPressure[AMDGPU::RegisterPressureSets::SReg_32] ||
336 Pressure[AMDGPU::RegisterPressureSets::VGPR_32] !=
337 CheckPressure[AMDGPU::RegisterPressureSets::VGPR_32] ||
338 Pressure[AMDGPU::RegisterPressureSets::AGPR_32] !=
339 CheckPressure[AMDGPU::RegisterPressureSets::AGPR_32]) {
340 errs() << "Register Pressure is inaccurate when calculated through "
341 "PressureDiff\n"
342 << "SGPR got " << Pressure[AMDGPU::RegisterPressureSets::SReg_32]
343 << ", expected "
344 << CheckPressure[AMDGPU::RegisterPressureSets::SReg_32] << "\n"
345 << "VGPR got " << Pressure[AMDGPU::RegisterPressureSets::VGPR_32]
346 << ", expected "
347 << CheckPressure[AMDGPU::RegisterPressureSets::VGPR_32] << "\n"
348 << "AGPR got " << Pressure[AMDGPU::RegisterPressureSets::AGPR_32]
349 << ", expected "
350 << CheckPressure[AMDGPU::RegisterPressureSets::AGPR_32] << "\n";
351 report_fatal_error("inaccurate register pressure calculation");
352 }
353#endif
354 }
355
356 unsigned NewAGPRPressure = Pressure[AMDGPU::RegisterPressureSets::AGPR_32];
357 unsigned NewSGPRPressure = Pressure[AMDGPU::RegisterPressureSets::SReg_32];
358 unsigned NewVGPRPressure = Pressure[AMDGPU::RegisterPressureSets::VGPR_32];
359
360 // If two instructions increase the pressure of different register sets
361 // by the same amount, the generic scheduler will prefer to schedule the
362 // instruction that increases the set with the least amount of registers,
363 // which in our case would be SGPRs. This is rarely what we want, so
364 // when we report excess/critical register pressure, we do it either
365 // only for VGPRs, AGPRs or SGPRs. Priority: VGPR > AGPR > SGPR.
366
367 // FIXME: Better heuristics to determine whether to prefer SGPRs or VGPRs.
368 const unsigned MaxVGPRPressureInc = 16;
369 bool ShouldTrackVGPRs = VGPRPressure + MaxVGPRPressureInc >= VGPRExcessLimit;
370 bool ShouldTrackAGPRs = AGPRExcessLimit > 0 && !ShouldTrackVGPRs &&
371 AGPRPressure + MaxVGPRPressureInc >= AGPRExcessLimit;
372 bool ShouldTrackSGPRs =
373 !ShouldTrackVGPRs && !ShouldTrackAGPRs && SGPRPressure >= SGPRExcessLimit;
374 // FIXME: We have to enter REG-EXCESS before we reach the actual threshold
375 // to increase the likelihood we don't go over the limits. We should improve
376 // the analysis to look through dependencies to find the path with the least
377 // register pressure.
378 // We only need to update the RPDelta for instructions that increase register
379 // pressure. Instructions that decrease or keep reg pressure the same will be
380 // marked as RegExcess in tryCandidate() when they are compared with
381 // instructions that increase the register pressure.
382 if (ShouldTrackVGPRs && NewVGPRPressure >= VGPRExcessLimit) {
383 HasHighPressure = true;
384 Cand.RPDelta.Excess = PressureChange(AMDGPU::RegisterPressureSets::VGPR_32);
385 Cand.RPDelta.Excess.setUnitInc(NewVGPRPressure - VGPRExcessLimit);
386 }
387
388 if (ShouldTrackAGPRs && NewAGPRPressure >= AGPRExcessLimit) {
389 HasHighPressure = true;
390 Cand.RPDelta.Excess = PressureChange(AMDGPU::RegisterPressureSets::AGPR_32);
391 Cand.RPDelta.Excess.setUnitInc(NewAGPRPressure - AGPRExcessLimit);
392 }
393
394 if (ShouldTrackSGPRs && NewSGPRPressure >= SGPRExcessLimit) {
395 HasHighPressure = true;
396 Cand.RPDelta.Excess = PressureChange(AMDGPU::RegisterPressureSets::SReg_32);
397 Cand.RPDelta.Excess.setUnitInc(NewSGPRPressure - SGPRExcessLimit);
398 }
399
400 // Register pressure is considered 'CRITICAL' if it is approaching a value
401 // that would reduce the wave occupancy for the execution unit. When
402 // register pressure is 'CRITICAL', increasing SGPR, VGPR, and AGPR
403 // pressure all has the same cost, so we pick the most critical type.
404
405 int SGPRDelta = NewSGPRPressure - SGPRCriticalLimit;
406 int VGPRDelta = NewVGPRPressure - VGPRCriticalLimit;
407 int AGPRDelta = AGPRExcessLimit > 0 ? NewAGPRPressure - AGPRCriticalLimit
408 : std::numeric_limits<int>::min();
409
410 if (SGPRDelta >= 0 || VGPRDelta >= 0 || AGPRDelta >= 0) {
411 HasHighPressure = true;
412 // Pick the most critical type.
413 if (VGPRDelta >= SGPRDelta && VGPRDelta >= AGPRDelta) {
414 Cand.RPDelta.CriticalMax =
415 PressureChange(AMDGPU::RegisterPressureSets::VGPR_32);
416 Cand.RPDelta.CriticalMax.setUnitInc(VGPRDelta);
417 } else if (AGPRDelta >= SGPRDelta) {
418 Cand.RPDelta.CriticalMax =
419 PressureChange(AMDGPU::RegisterPressureSets::AGPR_32);
420 Cand.RPDelta.CriticalMax.setUnitInc(AGPRDelta);
421 } else {
422 Cand.RPDelta.CriticalMax =
423 PressureChange(AMDGPU::RegisterPressureSets::SReg_32);
424 Cand.RPDelta.CriticalMax.setUnitInc(SGPRDelta);
425 }
426 }
427}
428
429static bool shouldCheckPending(SchedBoundary &Zone,
430 const TargetSchedModel *SchedModel) {
431 bool HasBufferedModel =
432 SchedModel->hasInstrSchedModel() && SchedModel->getMicroOpBufferSize();
433 unsigned Combined = Zone.Available.size() + Zone.Pending.size();
434 return Combined <= PendingQueueLimit && HasBufferedModel;
435}
436
437static SUnit *pickOnlyChoice(SchedBoundary &Zone,
438 const TargetSchedModel *SchedModel) {
439 // pickOnlyChoice() releases pending instructions and checks for new hazards.
440 SUnit *OnlyChoice = Zone.pickOnlyChoice();
441 if (!shouldCheckPending(Zone, SchedModel) || Zone.Pending.empty())
442 return OnlyChoice;
443
444 return nullptr;
445}
446
447void GCNSchedStrategy::printCandidateDecision(const SchedCandidate &Current,
448 const SchedCandidate &Preferred) {
449 LLVM_DEBUG({
450 dbgs() << "Prefer:\t\t";
451 DAG->dumpNode(*Preferred.SU);
452
453 if (Current.SU) {
454 dbgs() << "Not:\t";
455 DAG->dumpNode(*Current.SU);
456 }
457
458 dbgs() << "Reason:\t\t";
459 traceCandidate(Preferred);
460 });
461}
462
463// This function is mostly cut and pasted from
464// GenericScheduler::pickNodeFromQueue()
465void GCNSchedStrategy::pickNodeFromQueue(SchedBoundary &Zone,
466 const CandPolicy &ZonePolicy,
467 const RegPressureTracker &RPTracker,
468 SchedCandidate &Cand, bool &IsPending,
469 bool IsBottomUp) {
470 const SIRegisterInfo *SRI = static_cast<const SIRegisterInfo *>(TRI);
471 ArrayRef<unsigned> Pressure = RPTracker.getRegSetPressureAtPos();
472 unsigned SGPRPressure = 0;
473 unsigned VGPRPressure = 0;
474 unsigned AGPRPressure = 0;
475 IsPending = false;
476 if (DAG->isTrackingPressure()) {
477 if (!useGCNTrackers()) {
478 SGPRPressure = Pressure[AMDGPU::RegisterPressureSets::SReg_32];
479 VGPRPressure = Pressure[AMDGPU::RegisterPressureSets::VGPR_32];
480 AGPRPressure = Pressure[AMDGPU::RegisterPressureSets::AGPR_32];
481 } else {
482 GCNRPTracker *T = IsBottomUp
483 ? static_cast<GCNRPTracker *>(&UpwardTracker)
484 : static_cast<GCNRPTracker *>(&DownwardTracker);
485 SGPRPressure = T->getPressure().getSGPRNum();
486 VGPRPressure = T->getPressure().getArchVGPRNum();
487 AGPRPressure = T->getPressure().getAGPRNum();
488 }
489 }
490 LLVM_DEBUG(dbgs() << "Available Q:\n");
491 ReadyQueue &AQ = Zone.Available;
492 for (SUnit *SU : AQ) {
493
494 SchedCandidate TryCand(ZonePolicy);
495 initCandidate(Cand&: TryCand, SU, AtTop: Zone.isTop(), RPTracker, SRI, SGPRPressure,
496 VGPRPressure, AGPRPressure, IsBottomUp);
497 // Pass SchedBoundary only when comparing nodes from the same boundary.
498 SchedBoundary *ZoneArg = Cand.AtTop == TryCand.AtTop ? &Zone : nullptr;
499 tryCandidate(Cand, TryCand, Zone: ZoneArg);
500 if (TryCand.Reason != NoCand) {
501 // Initialize resource delta if needed in case future heuristics query it.
502 if (TryCand.ResDelta == SchedResourceDelta())
503 TryCand.initResourceDelta(DAG: Zone.DAG, SchedModel);
504 LLVM_DEBUG(printCandidateDecision(Cand, TryCand));
505 Cand.setBest(TryCand);
506 } else {
507 printCandidateDecision(Current: TryCand, Preferred: Cand);
508 }
509 }
510
511 if (!shouldCheckPending(Zone, SchedModel))
512 return;
513
514 LLVM_DEBUG(dbgs() << "Pending Q:\n");
515 ReadyQueue &PQ = Zone.Pending;
516 for (SUnit *SU : PQ) {
517
518 SchedCandidate TryCand(ZonePolicy);
519 initCandidate(Cand&: TryCand, SU, AtTop: Zone.isTop(), RPTracker, SRI, SGPRPressure,
520 VGPRPressure, AGPRPressure, IsBottomUp);
521 // Pass SchedBoundary only when comparing nodes from the same boundary.
522 SchedBoundary *ZoneArg = Cand.AtTop == TryCand.AtTop ? &Zone : nullptr;
523 tryPendingCandidate(Cand, TryCand, Zone: ZoneArg);
524 if (TryCand.Reason != NoCand) {
525 // Initialize resource delta if needed in case future heuristics query it.
526 if (TryCand.ResDelta == SchedResourceDelta())
527 TryCand.initResourceDelta(DAG: Zone.DAG, SchedModel);
528 LLVM_DEBUG(printCandidateDecision(Cand, TryCand));
529 IsPending = true;
530 Cand.setBest(TryCand);
531 } else {
532 printCandidateDecision(Current: TryCand, Preferred: Cand);
533 }
534 }
535}
536
537// This function is mostly cut and pasted from
538// GenericScheduler::pickNodeBidirectional()
539SUnit *GCNSchedStrategy::pickNodeBidirectional(bool &IsTopNode,
540 bool &PickedPending) {
541 // Schedule as far as possible in the direction of no choice. This is most
542 // efficient, but also provides the best heuristics for CriticalPSets.
543 if (SUnit *SU = pickOnlyChoice(Zone&: Bot, SchedModel)) {
544 IsTopNode = false;
545 return SU;
546 }
547 if (SUnit *SU = pickOnlyChoice(Zone&: Top, SchedModel)) {
548 IsTopNode = true;
549 return SU;
550 }
551 // Set the bottom-up policy based on the state of the current bottom zone
552 // and the instructions outside the zone, including the top zone.
553 CandPolicy BotPolicy;
554 setPolicy(Policy&: BotPolicy, /*IsPostRA=*/false, CurrZone&: Bot, OtherZone: &Top);
555 // Set the top-down policy based on the state of the current top zone and
556 // the instructions outside the zone, including the bottom zone.
557 CandPolicy TopPolicy;
558 setPolicy(Policy&: TopPolicy, /*IsPostRA=*/false, CurrZone&: Top, OtherZone: &Bot);
559
560 bool BotPending = false;
561 // See if BotCand is still valid (because we previously scheduled from Top).
562 LLVM_DEBUG(dbgs() << "Picking from Bot:\n");
563 if (!BotCand.isValid() || BotCand.SU->isScheduled ||
564 BotCand.Policy != BotPolicy) {
565 BotCand.reset(NewPolicy: CandPolicy());
566 pickNodeFromQueue(Zone&: Bot, ZonePolicy: BotPolicy, RPTracker: DAG->getBotRPTracker(), Cand&: BotCand,
567 IsPending&: BotPending,
568 /*IsBottomUp=*/true);
569 assert(BotCand.Reason != NoCand && "failed to find the first candidate");
570 } else {
571 LLVM_DEBUG(traceCandidate(BotCand));
572#ifndef NDEBUG
573 if (shouldVerifyScheduling()) {
574 SchedCandidate TCand;
575 TCand.reset(CandPolicy());
576 pickNodeFromQueue(Bot, BotPolicy, DAG->getBotRPTracker(), TCand,
577 BotPending,
578 /*IsBottomUp=*/true);
579 assert(TCand.SU == BotCand.SU &&
580 "Last pick result should correspond to re-picking right now");
581 }
582#endif
583 }
584
585 bool TopPending = false;
586 // Check if the top Q has a better candidate.
587 LLVM_DEBUG(dbgs() << "Picking from Top:\n");
588 if (!TopCand.isValid() || TopCand.SU->isScheduled ||
589 TopCand.Policy != TopPolicy) {
590 TopCand.reset(NewPolicy: CandPolicy());
591 pickNodeFromQueue(Zone&: Top, ZonePolicy: TopPolicy, RPTracker: DAG->getTopRPTracker(), Cand&: TopCand,
592 IsPending&: TopPending,
593 /*IsBottomUp=*/false);
594 assert(TopCand.Reason != NoCand && "failed to find the first candidate");
595 } else {
596 LLVM_DEBUG(traceCandidate(TopCand));
597#ifndef NDEBUG
598 if (shouldVerifyScheduling()) {
599 SchedCandidate TCand;
600 TCand.reset(CandPolicy());
601 pickNodeFromQueue(Top, TopPolicy, DAG->getTopRPTracker(), TCand,
602 TopPending,
603 /*IsBottomUp=*/false);
604 assert(TCand.SU == TopCand.SU &&
605 "Last pick result should correspond to re-picking right now");
606 }
607#endif
608 }
609
610 // Pick best from BotCand and TopCand.
611 LLVM_DEBUG(dbgs() << "Top Cand: "; traceCandidate(TopCand);
612 dbgs() << "Bot Cand: "; traceCandidate(BotCand););
613 SchedCandidate Cand = BotPending ? TopCand : BotCand;
614 SchedCandidate TryCand = BotPending ? BotCand : TopCand;
615 PickedPending = BotPending && TopPending;
616
617 TryCand.Reason = NoCand;
618 if (BotPending || TopPending) {
619 PickedPending |= tryPendingCandidate(Cand, TryCand&: TopCand, Zone: nullptr);
620 } else {
621 tryCandidate(Cand, TryCand, Zone: nullptr);
622 }
623
624 if (TryCand.Reason != NoCand) {
625 Cand.setBest(TryCand);
626 }
627
628 LLVM_DEBUG(dbgs() << "Picking: "; traceCandidate(Cand););
629
630 IsTopNode = Cand.AtTop;
631 return Cand.SU;
632}
633
634// This function is mostly cut and pasted from
635// GenericScheduler::pickNode()
636SUnit *GCNSchedStrategy::pickNode(bool &IsTopNode) {
637 if (DAG->top() == DAG->bottom()) {
638 assert(Top.Available.empty() && Top.Pending.empty() &&
639 Bot.Available.empty() && Bot.Pending.empty() && "ReadyQ garbage");
640 return nullptr;
641 }
642 bool PickedPending;
643 SUnit *SU;
644 do {
645 PickedPending = false;
646 if (RegionPolicy.OnlyTopDown) {
647 SU = pickOnlyChoice(Zone&: Top, SchedModel);
648 if (!SU) {
649 CandPolicy NoPolicy;
650 TopCand.reset(NewPolicy: NoPolicy);
651 pickNodeFromQueue(Zone&: Top, ZonePolicy: NoPolicy, RPTracker: DAG->getTopRPTracker(), Cand&: TopCand,
652 IsPending&: PickedPending,
653 /*IsBottomUp=*/false);
654 assert(TopCand.Reason != NoCand && "failed to find a candidate");
655 SU = TopCand.SU;
656 }
657 IsTopNode = true;
658 } else if (RegionPolicy.OnlyBottomUp) {
659 SU = pickOnlyChoice(Zone&: Bot, SchedModel);
660 if (!SU) {
661 CandPolicy NoPolicy;
662 BotCand.reset(NewPolicy: NoPolicy);
663 pickNodeFromQueue(Zone&: Bot, ZonePolicy: NoPolicy, RPTracker: DAG->getBotRPTracker(), Cand&: BotCand,
664 IsPending&: PickedPending,
665 /*IsBottomUp=*/true);
666 assert(BotCand.Reason != NoCand && "failed to find a candidate");
667 SU = BotCand.SU;
668 }
669 IsTopNode = false;
670 } else {
671 SU = pickNodeBidirectional(IsTopNode, PickedPending);
672 }
673 } while (SU->isScheduled);
674
675 if (PickedPending) {
676 unsigned ReadyCycle = IsTopNode ? SU->TopReadyCycle : SU->BotReadyCycle;
677 SchedBoundary &Zone = IsTopNode ? Top : Bot;
678 unsigned CurrentCycle = Zone.getCurrCycle();
679 if (ReadyCycle > CurrentCycle)
680 Zone.bumpCycle(NextCycle: ReadyCycle);
681
682 // FIXME: checkHazard() doesn't give information about which cycle the
683 // hazard will resolve so just keep bumping the cycle by 1. This could be
684 // made more efficient if checkHazard() returned more details.
685 while (Zone.checkHazard(SU))
686 Zone.bumpCycle(NextCycle: Zone.getCurrCycle() + 1);
687
688 Zone.releasePending();
689 }
690
691 if (SU->isTopReady())
692 Top.removeReady(SU);
693 if (SU->isBottomReady())
694 Bot.removeReady(SU);
695
696 LLVM_DEBUG(dbgs() << "Scheduling " << *SU << " " << *SU->getInstr());
697 return SU;
698}
699
700void GCNSchedStrategy::schedNode(SUnit *SU, bool IsTopNode) {
701 if (useGCNTrackers()) {
702 MachineInstr *MI = SU->getInstr();
703 IsTopNode ? (void)DownwardTracker.advance(MI, UseInternalIterator: false)
704 : UpwardTracker.recede(MI: *MI);
705 }
706
707 return GenericScheduler::schedNode(SU, IsTopNode);
708}
709
710GCNSchedStageID GCNSchedStrategy::getCurrentStage() {
711 assert(CurrentStage && CurrentStage != SchedStages.end());
712 return *CurrentStage;
713}
714
715bool GCNSchedStrategy::advanceStage() {
716 assert(CurrentStage != SchedStages.end());
717 if (!CurrentStage)
718 CurrentStage = SchedStages.begin();
719 else
720 CurrentStage++;
721
722 return CurrentStage != SchedStages.end();
723}
724
725bool GCNSchedStrategy::hasNextStage() const {
726 assert(CurrentStage);
727 return std::next(x: CurrentStage) != SchedStages.end();
728}
729
730bool GCNSchedStrategy::tryPendingCandidate(SchedCandidate &Cand,
731 SchedCandidate &TryCand,
732 SchedBoundary *Zone) const {
733 // Initialize the candidate if needed.
734 if (!Cand.isValid()) {
735 TryCand.Reason = NodeOrder;
736 return true;
737 }
738
739 // Bias PhysReg Defs and copies to their uses and defined respectively.
740 if (tryGreater(TryVal: biasPhysReg(SU: TryCand.SU, isTop: TryCand.AtTop),
741 CandVal: biasPhysReg(SU: Cand.SU, isTop: Cand.AtTop), TryCand, Cand, Reason: PhysReg))
742 return TryCand.Reason != NoCand;
743
744 // Avoid exceeding the target's limit.
745 if (DAG->isTrackingPressure() &&
746 tryPressure(TryP: TryCand.RPDelta.Excess, CandP: Cand.RPDelta.Excess, TryCand, Cand,
747 Reason: RegExcess, TRI, MF: DAG->MF))
748 return TryCand.Reason != NoCand;
749
750 // Avoid increasing the max critical pressure in the scheduled region.
751 if (DAG->isTrackingPressure() &&
752 tryPressure(TryP: TryCand.RPDelta.CriticalMax, CandP: Cand.RPDelta.CriticalMax,
753 TryCand, Cand, Reason: RegCritical, TRI, MF: DAG->MF))
754 return TryCand.Reason != NoCand;
755
756 bool SameBoundary = Zone != nullptr;
757 if (SameBoundary) {
758 TryCand.initResourceDelta(DAG, SchedModel);
759 if (tryLess(TryVal: TryCand.ResDelta.CritResources, CandVal: Cand.ResDelta.CritResources,
760 TryCand, Cand, Reason: ResourceReduce))
761 return TryCand.Reason != NoCand;
762 if (tryGreater(TryVal: TryCand.ResDelta.DemandedResources,
763 CandVal: Cand.ResDelta.DemandedResources, TryCand, Cand,
764 Reason: ResourceDemand))
765 return TryCand.Reason != NoCand;
766 }
767
768 return false;
769}
770
771GCNMaxOccupancySchedStrategy::GCNMaxOccupancySchedStrategy(
772 const MachineSchedContext *C, bool IsLegacyScheduler)
773 : GCNSchedStrategy(C) {
774 SchedStages.push_back(Elt: GCNSchedStageID::OccInitialSchedule);
775 if (!DisableRewriteMFMAFormSchedStage)
776 SchedStages.push_back(Elt: GCNSchedStageID::RewriteMFMAForm);
777 SchedStages.push_back(Elt: GCNSchedStageID::UnclusteredHighRPReschedule);
778 SchedStages.push_back(Elt: GCNSchedStageID::ClusteredLowOccupancyReschedule);
779 SchedStages.push_back(Elt: GCNSchedStageID::PreRARematerialize);
780 if (IsLegacyScheduler)
781 GCNTrackersOverride = std::nullopt;
782}
783
784GCNMaxILPSchedStrategy::GCNMaxILPSchedStrategy(const MachineSchedContext *C)
785 : GCNSchedStrategy(C) {
786 SchedStages.push_back(Elt: GCNSchedStageID::ILPInitialSchedule);
787}
788
789bool GCNMaxILPSchedStrategy::tryCandidate(SchedCandidate &Cand,
790 SchedCandidate &TryCand,
791 SchedBoundary *Zone) const {
792 // Initialize the candidate if needed.
793 if (!Cand.isValid()) {
794 TryCand.Reason = NodeOrder;
795 return true;
796 }
797
798 // Avoid spilling by exceeding the register limit.
799 if (DAG->isTrackingPressure() &&
800 tryPressure(TryP: TryCand.RPDelta.Excess, CandP: Cand.RPDelta.Excess, TryCand, Cand,
801 Reason: RegExcess, TRI, MF: DAG->MF))
802 return TryCand.Reason != NoCand;
803
804 // Bias PhysReg Defs and copies to their uses and defined respectively.
805 if (tryGreater(TryVal: biasPhysReg(SU: TryCand.SU, isTop: TryCand.AtTop),
806 CandVal: biasPhysReg(SU: Cand.SU, isTop: Cand.AtTop), TryCand, Cand, Reason: PhysReg))
807 return TryCand.Reason != NoCand;
808
809 bool SameBoundary = Zone != nullptr;
810 if (SameBoundary) {
811 // Prioritize instructions that read unbuffered resources by stall cycles.
812 if (tryLess(TryVal: Zone->getLatencyStallCycles(SU: TryCand.SU),
813 CandVal: Zone->getLatencyStallCycles(SU: Cand.SU), TryCand, Cand, Reason: Stall))
814 return TryCand.Reason != NoCand;
815
816 // Avoid critical resource consumption and balance the schedule.
817 TryCand.initResourceDelta(DAG, SchedModel);
818 if (tryLess(TryVal: TryCand.ResDelta.CritResources, CandVal: Cand.ResDelta.CritResources,
819 TryCand, Cand, Reason: ResourceReduce))
820 return TryCand.Reason != NoCand;
821 if (tryGreater(TryVal: TryCand.ResDelta.DemandedResources,
822 CandVal: Cand.ResDelta.DemandedResources, TryCand, Cand,
823 Reason: ResourceDemand))
824 return TryCand.Reason != NoCand;
825
826 // Unconditionally try to reduce latency.
827 if (tryLatency(TryCand, Cand, Zone&: *Zone))
828 return TryCand.Reason != NoCand;
829
830 // Weak edges are for clustering and other constraints.
831 if (tryLess(TryVal: getWeakLeft(SU: TryCand.SU, isTop: TryCand.AtTop),
832 CandVal: getWeakLeft(SU: Cand.SU, isTop: Cand.AtTop), TryCand, Cand, Reason: Weak))
833 return TryCand.Reason != NoCand;
834 }
835
836 // Keep clustered nodes together to encourage downstream peephole
837 // optimizations which may reduce resource requirements.
838 //
839 // This is a best effort to set things up for a post-RA pass. Optimizations
840 // like generating loads of multiple registers should ideally be done within
841 // the scheduler pass by combining the loads during DAG postprocessing.
842 unsigned CandZoneCluster = Cand.AtTop ? TopClusterID : BotClusterID;
843 unsigned TryCandZoneCluster = TryCand.AtTop ? TopClusterID : BotClusterID;
844 bool CandIsClusterSucc =
845 isTheSameCluster(A: CandZoneCluster, B: Cand.SU->ParentClusterIdx);
846 bool TryCandIsClusterSucc =
847 isTheSameCluster(A: TryCandZoneCluster, B: TryCand.SU->ParentClusterIdx);
848 if (tryGreater(TryVal: TryCandIsClusterSucc, CandVal: CandIsClusterSucc, TryCand, Cand,
849 Reason: Cluster))
850 return TryCand.Reason != NoCand;
851
852 // Avoid increasing the max critical pressure in the scheduled region.
853 if (DAG->isTrackingPressure() &&
854 tryPressure(TryP: TryCand.RPDelta.CriticalMax, CandP: Cand.RPDelta.CriticalMax,
855 TryCand, Cand, Reason: RegCritical, TRI, MF: DAG->MF))
856 return TryCand.Reason != NoCand;
857
858 // Avoid increasing the max pressure of the entire region.
859 if (DAG->isTrackingPressure() &&
860 tryPressure(TryP: TryCand.RPDelta.CurrentMax, CandP: Cand.RPDelta.CurrentMax, TryCand,
861 Cand, Reason: RegMax, TRI, MF: DAG->MF))
862 return TryCand.Reason != NoCand;
863
864 if (SameBoundary) {
865 // Fall through to original instruction order.
866 if ((Zone->isTop() && TryCand.SU->NodeNum < Cand.SU->NodeNum) ||
867 (!Zone->isTop() && TryCand.SU->NodeNum > Cand.SU->NodeNum)) {
868 TryCand.Reason = NodeOrder;
869 return true;
870 }
871 }
872 return false;
873}
874
875GCNMaxMemoryClauseSchedStrategy::GCNMaxMemoryClauseSchedStrategy(
876 const MachineSchedContext *C)
877 : GCNSchedStrategy(C) {
878 SchedStages.push_back(Elt: GCNSchedStageID::MemoryClauseInitialSchedule);
879}
880
881/// GCNMaxMemoryClauseSchedStrategy tries best to clause memory instructions as
882/// much as possible. This is achieved by:
883// 1. Prioritize clustered operations before stall latency heuristic.
884// 2. Prioritize long-latency-load before stall latency heuristic.
885///
886/// \param Cand provides the policy and current best candidate.
887/// \param TryCand refers to the next SUnit candidate, otherwise uninitialized.
888/// \param Zone describes the scheduled zone that we are extending, or nullptr
889/// if Cand is from a different zone than TryCand.
890/// \return \c true if TryCand is better than Cand (Reason is NOT NoCand)
891bool GCNMaxMemoryClauseSchedStrategy::tryCandidate(SchedCandidate &Cand,
892 SchedCandidate &TryCand,
893 SchedBoundary *Zone) const {
894 // Initialize the candidate if needed.
895 if (!Cand.isValid()) {
896 TryCand.Reason = NodeOrder;
897 return true;
898 }
899
900 // Bias PhysReg Defs and copies to their uses and defined respectively.
901 if (tryGreater(TryVal: biasPhysReg(SU: TryCand.SU, isTop: TryCand.AtTop),
902 CandVal: biasPhysReg(SU: Cand.SU, isTop: Cand.AtTop), TryCand, Cand, Reason: PhysReg))
903 return TryCand.Reason != NoCand;
904
905 if (DAG->isTrackingPressure()) {
906 // Avoid exceeding the target's limit.
907 if (tryPressure(TryP: TryCand.RPDelta.Excess, CandP: Cand.RPDelta.Excess, TryCand, Cand,
908 Reason: RegExcess, TRI, MF: DAG->MF))
909 return TryCand.Reason != NoCand;
910
911 // Avoid increasing the max critical pressure in the scheduled region.
912 if (tryPressure(TryP: TryCand.RPDelta.CriticalMax, CandP: Cand.RPDelta.CriticalMax,
913 TryCand, Cand, Reason: RegCritical, TRI, MF: DAG->MF))
914 return TryCand.Reason != NoCand;
915 }
916
917 // MaxMemoryClause-specific: We prioritize clustered instructions as we would
918 // get more benefit from clausing these memory instructions.
919 unsigned CandZoneCluster = Cand.AtTop ? TopClusterID : BotClusterID;
920 unsigned TryCandZoneCluster = TryCand.AtTop ? TopClusterID : BotClusterID;
921 bool CandIsClusterSucc =
922 isTheSameCluster(A: CandZoneCluster, B: Cand.SU->ParentClusterIdx);
923 bool TryCandIsClusterSucc =
924 isTheSameCluster(A: TryCandZoneCluster, B: TryCand.SU->ParentClusterIdx);
925 if (tryGreater(TryVal: TryCandIsClusterSucc, CandVal: CandIsClusterSucc, TryCand, Cand,
926 Reason: Cluster))
927 return TryCand.Reason != NoCand;
928
929 // We only compare a subset of features when comparing nodes between
930 // Top and Bottom boundary. Some properties are simply incomparable, in many
931 // other instances we should only override the other boundary if something
932 // is a clear good pick on one boundary. Skip heuristics that are more
933 // "tie-breaking" in nature.
934 bool SameBoundary = Zone != nullptr;
935 if (SameBoundary) {
936 // For loops that are acyclic path limited, aggressively schedule for
937 // latency. Within an single cycle, whenever CurrMOps > 0, allow normal
938 // heuristics to take precedence.
939 if (Rem.IsAcyclicLatencyLimited && !Zone->getCurrMOps() &&
940 tryLatency(TryCand, Cand, Zone&: *Zone))
941 return TryCand.Reason != NoCand;
942
943 // MaxMemoryClause-specific: Prioritize long latency memory load
944 // instructions in top-bottom order to hide more latency. The mayLoad check
945 // is used to exclude store-like instructions, which we do not want to
946 // scheduler them too early.
947 bool TryMayLoad =
948 TryCand.SU->isInstr() && TryCand.SU->getInstr()->mayLoad();
949 bool CandMayLoad = Cand.SU->isInstr() && Cand.SU->getInstr()->mayLoad();
950
951 if (TryMayLoad || CandMayLoad) {
952 bool TryLongLatency =
953 TryCand.SU->Latency > 10 * Cand.SU->Latency && TryMayLoad;
954 bool CandLongLatency =
955 10 * TryCand.SU->Latency < Cand.SU->Latency && CandMayLoad;
956
957 if (tryGreater(TryVal: Zone->isTop() ? TryLongLatency : CandLongLatency,
958 CandVal: Zone->isTop() ? CandLongLatency : TryLongLatency, TryCand,
959 Cand, Reason: Stall))
960 return TryCand.Reason != NoCand;
961 }
962 // Prioritize instructions that read unbuffered resources by stall cycles.
963 if (tryLess(TryVal: Zone->getLatencyStallCycles(SU: TryCand.SU),
964 CandVal: Zone->getLatencyStallCycles(SU: Cand.SU), TryCand, Cand, Reason: Stall))
965 return TryCand.Reason != NoCand;
966 }
967
968 if (SameBoundary) {
969 // Weak edges are for clustering and other constraints.
970 if (tryLess(TryVal: getWeakLeft(SU: TryCand.SU, isTop: TryCand.AtTop),
971 CandVal: getWeakLeft(SU: Cand.SU, isTop: Cand.AtTop), TryCand, Cand, Reason: Weak))
972 return TryCand.Reason != NoCand;
973 }
974
975 // Avoid increasing the max pressure of the entire region.
976 if (DAG->isTrackingPressure() &&
977 tryPressure(TryP: TryCand.RPDelta.CurrentMax, CandP: Cand.RPDelta.CurrentMax, TryCand,
978 Cand, Reason: RegMax, TRI, MF: DAG->MF))
979 return TryCand.Reason != NoCand;
980
981 if (SameBoundary) {
982 // Avoid critical resource consumption and balance the schedule.
983 TryCand.initResourceDelta(DAG, SchedModel);
984 if (tryLess(TryVal: TryCand.ResDelta.CritResources, CandVal: Cand.ResDelta.CritResources,
985 TryCand, Cand, Reason: ResourceReduce))
986 return TryCand.Reason != NoCand;
987 if (tryGreater(TryVal: TryCand.ResDelta.DemandedResources,
988 CandVal: Cand.ResDelta.DemandedResources, TryCand, Cand,
989 Reason: ResourceDemand))
990 return TryCand.Reason != NoCand;
991
992 // Avoid serializing long latency dependence chains.
993 // For acyclic path limited loops, latency was already checked above.
994 if (!RegionPolicy.DisableLatencyHeuristic && TryCand.Policy.ReduceLatency &&
995 !Rem.IsAcyclicLatencyLimited && tryLatency(TryCand, Cand, Zone&: *Zone))
996 return TryCand.Reason != NoCand;
997
998 // Fall through to original instruction order.
999 if (Zone->isTop() == (TryCand.SU->NodeNum < Cand.SU->NodeNum)) {
1000 assert(TryCand.SU->NodeNum != Cand.SU->NodeNum);
1001 TryCand.Reason = NodeOrder;
1002 return true;
1003 }
1004 }
1005
1006 return false;
1007}
1008
1009GCNScheduleDAGMILive::GCNScheduleDAGMILive(
1010 MachineSchedContext *C, std::unique_ptr<MachineSchedStrategy> S)
1011 : ScheduleDAGMILive(C, std::move(S)), ST(MF.getSubtarget<GCNSubtarget>()),
1012 MFI(*MF.getInfo<SIMachineFunctionInfo>()),
1013 StartingOccupancy(MFI.getOccupancy()), MinOccupancy(StartingOccupancy),
1014 RegionLiveOuts(this, /*IsLiveOut=*/true) {
1015
1016 // We want regions with a single MI to be scheduled so that we can reason
1017 // about them correctly during scheduling stages that move MIs between regions
1018 // (e.g., rematerialization).
1019 ScheduleSingleMIRegions = true;
1020 LLVM_DEBUG(dbgs() << "Starting occupancy is " << StartingOccupancy << ".\n");
1021 if (RelaxedOcc) {
1022 MinOccupancy = std::min(a: MFI.getMinAllowedOccupancy(), b: StartingOccupancy);
1023 if (MinOccupancy != StartingOccupancy)
1024 LLVM_DEBUG(dbgs() << "Allowing Occupancy drops to " << MinOccupancy
1025 << ".\n");
1026 }
1027}
1028
1029std::unique_ptr<GCNSchedStage>
1030GCNScheduleDAGMILive::createSchedStage(GCNSchedStageID SchedStageID) {
1031 switch (SchedStageID) {
1032 case GCNSchedStageID::OccInitialSchedule:
1033 return std::make_unique<OccInitialScheduleStage>(args&: SchedStageID, args&: *this);
1034 case GCNSchedStageID::RewriteMFMAForm:
1035 return std::make_unique<RewriteMFMAFormStage>(args&: SchedStageID, args&: *this);
1036 case GCNSchedStageID::UnclusteredHighRPReschedule:
1037 return std::make_unique<UnclusteredHighRPStage>(args&: SchedStageID, args&: *this);
1038 case GCNSchedStageID::ClusteredLowOccupancyReschedule:
1039 return std::make_unique<ClusteredLowOccStage>(args&: SchedStageID, args&: *this);
1040 case GCNSchedStageID::PreRARematerialize:
1041 return std::make_unique<PreRARematStage>(args&: SchedStageID, args&: *this);
1042 case GCNSchedStageID::ILPInitialSchedule:
1043 return std::make_unique<ILPInitialScheduleStage>(args&: SchedStageID, args&: *this);
1044 case GCNSchedStageID::MemoryClauseInitialSchedule:
1045 return std::make_unique<MemoryClauseInitialScheduleStage>(args&: SchedStageID,
1046 args&: *this);
1047 case GCNSchedStageID::LiveIntervalRPReschedule:
1048 return std::make_unique<LiveIntervalRPStage>(args&: SchedStageID, args&: *this);
1049 }
1050
1051 llvm_unreachable("Unknown SchedStageID.");
1052}
1053
1054void GCNScheduleDAGMILive::schedule() {
1055 // Collect all scheduling regions. The actual scheduling is performed in
1056 // GCNScheduleDAGMILive::finalizeSchedule.
1057 Regions.push_back(Elt: std::pair(RegionBegin, RegionEnd));
1058}
1059
1060GCNRegPressure
1061GCNScheduleDAGMILive::getRealRegPressure(unsigned RegionIdx) const {
1062 if (Regions[RegionIdx].first == Regions[RegionIdx].second)
1063 return llvm::getRegPressure(MRI, LiveRegs: LiveIns[RegionIdx]);
1064 GCNDownwardRPTracker RPTracker(*LIS);
1065 RPTracker.advance(Begin: Regions[RegionIdx].first, End: Regions[RegionIdx].second,
1066 LiveRegsCopy: &LiveIns[RegionIdx]);
1067 return RPTracker.moveMaxPressure();
1068}
1069
1070static MachineInstr *getLastMIForRegion(MachineBasicBlock::iterator RegionBegin,
1071 MachineBasicBlock::iterator RegionEnd) {
1072 assert(RegionBegin != RegionEnd && "Region must not be empty");
1073 return &*skipDebugInstructionsBackward(It: std::prev(x: RegionEnd), Begin: RegionBegin);
1074}
1075
1076void GCNScheduleDAGMILive::computeBlockPressure(unsigned RegionIdx,
1077 const MachineBasicBlock *MBB) {
1078 GCNDownwardRPTracker RPTracker(*LIS);
1079
1080 // If the block has the only successor then live-ins of that successor are
1081 // live-outs of the current block. We can reuse calculated live set if the
1082 // successor will be sent to scheduling past current block.
1083
1084 // However, due to the bug in LiveInterval analysis it may happen that two
1085 // predecessors of the same successor block have different lane bitmasks for
1086 // a live-out register. Workaround that by sticking to one-to-one relationship
1087 // i.e. one predecessor with one successor block.
1088 const MachineBasicBlock *OnlySucc = nullptr;
1089 if (MBB->succ_size() == 1) {
1090 auto *Candidate = *MBB->succ_begin();
1091 if (!Candidate->empty() && Candidate->pred_size() == 1) {
1092 SlotIndexes *Ind = LIS->getSlotIndexes();
1093 if (Ind->getMBBStartIdx(mbb: MBB) < Ind->getMBBStartIdx(mbb: Candidate))
1094 OnlySucc = Candidate;
1095 }
1096 }
1097
1098 // Scheduler sends regions from the end of the block upwards.
1099 size_t CurRegion = RegionIdx;
1100 for (size_t E = Regions.size(); CurRegion != E; ++CurRegion)
1101 if (Regions[CurRegion].first->getParent() != MBB)
1102 break;
1103 --CurRegion;
1104
1105 auto I = MBB->begin();
1106 auto LiveInIt = MBBLiveIns.find(Val: MBB);
1107 auto &Rgn = Regions[CurRegion];
1108 auto *NonDbgMI = &*skipDebugInstructionsForward(It: Rgn.first, End: Rgn.second);
1109 if (LiveInIt != MBBLiveIns.end()) {
1110 auto LiveIn = std::move(LiveInIt->second);
1111 RPTracker.reset(MI: *MBB->begin(), End: MBB->end(), LiveRegs: &LiveIn);
1112 MBBLiveIns.erase(I: LiveInIt);
1113 } else {
1114 I = Rgn.first;
1115 auto LRS = BBLiveInMap.lookup(Val: NonDbgMI);
1116#ifdef EXPENSIVE_CHECKS
1117 assert(isEqual(getLiveRegsBefore(*NonDbgMI, *LIS), LRS));
1118#endif
1119 RPTracker.reset(MI: *I, End: I->getParent()->end(), LiveRegs: &LRS);
1120 }
1121
1122 for (;;) {
1123 I = RPTracker.getNext();
1124
1125 if (Regions[CurRegion].first == I || NonDbgMI == I) {
1126 LiveIns[CurRegion] = RPTracker.getLiveRegs();
1127 RPTracker.clearMaxPressure();
1128 }
1129
1130 if (Regions[CurRegion].second == I) {
1131 Pressure[CurRegion] = RPTracker.moveMaxPressure();
1132 if (CurRegion-- == RegionIdx)
1133 break;
1134 auto &Rgn = Regions[CurRegion];
1135 NonDbgMI = &*skipDebugInstructionsForward(It: Rgn.first, End: Rgn.second);
1136 }
1137 RPTracker.advanceBeforeNext();
1138 RPTracker.advanceToNext();
1139 }
1140
1141 if (OnlySucc) {
1142 if (I != MBB->end()) {
1143 RPTracker.advanceBeforeNext();
1144 RPTracker.advanceToNext();
1145 RPTracker.advance(End: MBB->end());
1146 }
1147 MBBLiveIns[OnlySucc] = RPTracker.moveLiveRegs();
1148 }
1149}
1150
1151DenseMap<MachineInstr *, GCNRPTracker::LiveRegSet>
1152GCNScheduleDAGMILive::getRegionLiveInMap() const {
1153 assert(!Regions.empty());
1154 std::vector<MachineInstr *> RegionFirstMIs;
1155 RegionFirstMIs.reserve(n: Regions.size());
1156 for (auto &[RegionBegin, RegionEnd] : reverse(C: Regions))
1157 RegionFirstMIs.push_back(
1158 x: &*skipDebugInstructionsForward(It: RegionBegin, End: RegionEnd));
1159
1160 return getLiveRegMap(R&: RegionFirstMIs, /*After=*/false, LIS&: *LIS);
1161}
1162
1163DenseMap<MachineInstr *, GCNRPTracker::LiveRegSet>
1164GCNScheduleDAGMILive::getRegionLiveOutMap() const {
1165 assert(!Regions.empty());
1166 std::vector<MachineInstr *> RegionLastMIs;
1167 RegionLastMIs.reserve(n: Regions.size());
1168 for (auto &[RegionBegin, RegionEnd] : reverse(C: Regions)) {
1169 // Skip empty regions.
1170 if (RegionBegin == RegionEnd)
1171 continue;
1172 RegionLastMIs.push_back(x: getLastMIForRegion(RegionBegin, RegionEnd));
1173 }
1174 return getLiveRegMap(R&: RegionLastMIs, /*After=*/true, LIS&: *LIS);
1175}
1176
1177void RegionPressureMap::buildLiveRegMap() {
1178 IdxToInstruction.clear();
1179
1180 RegionLiveRegMap =
1181 IsLiveOut ? DAG->getRegionLiveOutMap() : DAG->getRegionLiveInMap();
1182 for (unsigned I = 0; I < DAG->Regions.size(); I++) {
1183 auto &[RegionBegin, RegionEnd] = DAG->Regions[I];
1184 // Skip empty regions.
1185 if (RegionBegin == RegionEnd)
1186 continue;
1187 MachineInstr *RegionKey =
1188 IsLiveOut ? getLastMIForRegion(RegionBegin, RegionEnd) : &*RegionBegin;
1189 IdxToInstruction[I] = RegionKey;
1190 }
1191}
1192
1193void GCNScheduleDAGMILive::finalizeSchedule() {
1194 // Start actual scheduling here. This function is called by the base
1195 // MachineScheduler after all regions have been recorded by
1196 // GCNScheduleDAGMILive::schedule().
1197 LiveIns.resize(N: Regions.size());
1198 Pressure.resize(N: Regions.size());
1199 RegionsWithHighRP.resize(N: Regions.size());
1200 RegionsWithExcessRP.resize(N: Regions.size());
1201 RegionsWithIGLPInstrs.resize(N: Regions.size());
1202 RegionsWithHighRP.reset();
1203 RegionsWithExcessRP.reset();
1204 RegionsWithIGLPInstrs.reset();
1205
1206 runSchedStages();
1207}
1208
1209void GCNScheduleDAGMILive::runSchedStages() {
1210 LLVM_DEBUG(dbgs() << "All regions recorded, starting actual scheduling.\n");
1211
1212 GCNSchedStrategy &S = static_cast<GCNSchedStrategy &>(*SchedImpl);
1213 if (!Regions.empty()) {
1214 BBLiveInMap = getRegionLiveInMap();
1215 if (S.useGCNTrackers())
1216 RegionLiveOuts.buildLiveRegMap();
1217 }
1218
1219#ifdef DUMP_MAX_REG_PRESSURE
1220 if (PrintMaxRPRegUsageBeforeScheduler) {
1221 dumpMaxRegPressure(MF, GCNRegPressure::VGPR, *LIS, MLI);
1222 dumpMaxRegPressure(MF, GCNRegPressure::SGPR, *LIS, MLI);
1223 LIS->dump();
1224 }
1225#endif
1226
1227 while (S.advanceStage()) {
1228 auto Stage = createSchedStage(SchedStageID: S.getCurrentStage());
1229 if (!Stage->initGCNSchedStage())
1230 continue;
1231
1232 for (auto Region : Regions) {
1233 RegionBegin = Region.first;
1234 RegionEnd = Region.second;
1235 // Setup for scheduling the region and check whether it should be skipped.
1236 if (!Stage->initGCNRegion()) {
1237 Stage->advanceRegion();
1238 exitRegion();
1239 continue;
1240 }
1241
1242 if (S.useGCNTrackers()) {
1243 const unsigned RegionIdx = Stage->getRegionIdx();
1244 S.getDownwardTracker()->reset(MRI, LiveRegs: LiveIns[RegionIdx]);
1245 S.getUpwardTracker()->reset(
1246 MRI, LiveRegs: RegionLiveOuts.getLiveRegsForRegionIdx(RegionIdx));
1247 }
1248
1249 ScheduleDAGMILive::schedule();
1250 Stage->finalizeGCNRegion();
1251 Stage->advanceRegion();
1252 exitRegion();
1253 }
1254
1255 Stage->finalizeGCNSchedStage();
1256 }
1257
1258#ifdef DUMP_MAX_REG_PRESSURE
1259 if (PrintMaxRPRegUsageAfterScheduler) {
1260 dumpMaxRegPressure(MF, GCNRegPressure::VGPR, *LIS, MLI);
1261 dumpMaxRegPressure(MF, GCNRegPressure::SGPR, *LIS, MLI);
1262 LIS->dump();
1263 }
1264#endif
1265}
1266
1267#ifndef NDEBUG
1268raw_ostream &llvm::operator<<(raw_ostream &OS, const GCNSchedStageID &StageID) {
1269 switch (StageID) {
1270 case GCNSchedStageID::OccInitialSchedule:
1271 OS << "Max Occupancy Initial Schedule";
1272 break;
1273 case GCNSchedStageID::RewriteMFMAForm:
1274 OS << "Instruction Rewriting Reschedule";
1275 break;
1276 case GCNSchedStageID::UnclusteredHighRPReschedule:
1277 OS << "Unclustered High Register Pressure Reschedule";
1278 break;
1279 case GCNSchedStageID::ClusteredLowOccupancyReschedule:
1280 OS << "Clustered Low Occupancy Reschedule";
1281 break;
1282 case GCNSchedStageID::PreRARematerialize:
1283 OS << "Pre-RA Rematerialize";
1284 break;
1285 case GCNSchedStageID::ILPInitialSchedule:
1286 OS << "Max ILP Initial Schedule";
1287 break;
1288 case GCNSchedStageID::MemoryClauseInitialSchedule:
1289 OS << "Max memory clause Initial Schedule";
1290 break;
1291 case GCNSchedStageID::LiveIntervalRPReschedule:
1292 OS << "Live Interval RP Reschedule";
1293 break;
1294 }
1295
1296 return OS;
1297}
1298#endif
1299
1300GCNSchedStage::GCNSchedStage(GCNSchedStageID StageID, GCNScheduleDAGMILive &DAG)
1301 : DAG(DAG), S(static_cast<GCNSchedStrategy &>(*DAG.SchedImpl)), MF(DAG.MF),
1302 MFI(DAG.MFI), ST(DAG.ST), StageID(StageID) {}
1303
1304bool GCNSchedStage::initGCNSchedStage() {
1305 if (!DAG.LIS)
1306 return false;
1307
1308 LLVM_DEBUG(dbgs() << "Starting scheduling stage: " << StageID << "\n");
1309 return true;
1310}
1311
1312void RewriteMFMAFormStage::findReachingDefs(
1313 MachineOperand &UseMO, LiveIntervals *LIS,
1314 SmallVectorImpl<SlotIndex> &DefIdxs) {
1315 MachineInstr *UseMI = UseMO.getParent();
1316 LiveInterval &UseLI = LIS->getInterval(Reg: UseMO.getReg());
1317 VNInfo *VNI = UseLI.getVNInfoAt(Idx: LIS->getInstructionIndex(Instr: *UseMI));
1318
1319 // If the def is not a PHI, then it must be the only reaching def.
1320 if (!VNI->isPHIDef()) {
1321 DefIdxs.push_back(Elt: VNI->def);
1322 return;
1323 }
1324
1325 SmallPtrSet<MachineBasicBlock *, 8> Visited = {UseMI->getParent()};
1326 SmallVector<MachineBasicBlock *, 8> Worklist;
1327
1328 // Mark the predecessor blocks for traversal
1329 for (MachineBasicBlock *PredMBB : UseMI->getParent()->predecessors()) {
1330 Worklist.push_back(Elt: PredMBB);
1331 Visited.insert(Ptr: PredMBB);
1332 }
1333
1334 while (!Worklist.empty()) {
1335 MachineBasicBlock *CurrMBB = Worklist.pop_back_val();
1336
1337 SlotIndex CurrMBBEnd = LIS->getMBBEndIdx(mbb: CurrMBB);
1338 VNInfo *VNI = UseLI.getVNInfoAt(Idx: CurrMBBEnd.getPrevSlot());
1339
1340 MachineBasicBlock *DefMBB = LIS->getMBBFromIndex(index: VNI->def);
1341
1342 // If there is a def in this block, then add it to the list. This is the
1343 // reaching def of this path.
1344 if (!VNI->isPHIDef()) {
1345 DefIdxs.push_back(Elt: VNI->def);
1346 continue;
1347 }
1348
1349 for (MachineBasicBlock *PredMBB : DefMBB->predecessors()) {
1350 if (Visited.insert(Ptr: PredMBB).second)
1351 Worklist.push_back(Elt: PredMBB);
1352 }
1353 }
1354}
1355
1356void RewriteMFMAFormStage::findReachingUses(
1357 const MachineInstr *DefMI, LiveIntervals *LIS,
1358 SmallVectorImpl<MachineOperand *> &ReachingUses) {
1359 SlotIndex DefIdx = LIS->getInstructionIndex(Instr: *DefMI);
1360 for (MachineOperand &UseMO :
1361 DAG.MRI.use_nodbg_operands(Reg: DefMI->getOperand(i: 0).getReg())) {
1362 SmallVector<SlotIndex, 8> ReachingDefIndexes;
1363 findReachingDefs(UseMO, LIS, DefIdxs&: ReachingDefIndexes);
1364
1365 // If we find a use that contains this DefMI in its reachingDefs, then it is
1366 // a reaching use.
1367 if (any_of(Range&: ReachingDefIndexes, P: [DefIdx](SlotIndex RDIdx) {
1368 return SlotIndex::isSameInstr(A: RDIdx, B: DefIdx);
1369 }))
1370 ReachingUses.push_back(Elt: &UseMO);
1371 }
1372}
1373
1374bool RewriteMFMAFormStage::initGCNSchedStage() {
1375 // We only need to run this pass if the architecture supports AGPRs.
1376 // Additionally, we don't use AGPRs at occupancy levels above 1 so there
1377 // is no need for this pass in that case, either.
1378 const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>();
1379 if (!ST.hasGFX90AInsts() || MFI.getMinWavesPerEU() > 1)
1380 return false;
1381
1382 RegionsWithExcessArchVGPR.resize(N: DAG.Regions.size());
1383 RegionsWithExcessArchVGPR.reset();
1384 for (unsigned Region = 0; Region < DAG.Regions.size(); Region++) {
1385 GCNRegPressure PressureBefore = DAG.Pressure[Region];
1386 if (PressureBefore.getArchVGPRNum() > ST.getAddressableNumArchVGPRs())
1387 RegionsWithExcessArchVGPR[Region] = true;
1388 }
1389
1390 if (RegionsWithExcessArchVGPR.none())
1391 return false;
1392
1393 TII = ST.getInstrInfo();
1394 SRI = ST.getRegisterInfo();
1395
1396 std::vector<std::pair<MachineInstr *, unsigned>> RewriteCands;
1397 DenseMap<MachineBasicBlock *, std::set<Register>> CopyForUse;
1398 SmallPtrSet<MachineInstr *, 8> CopyForDef;
1399
1400 if (!initHeuristics(RewriteCands, CopyForUse, CopyForDef))
1401 return false;
1402
1403 int64_t Cost = getRewriteCost(RewriteCands, CopyForUse, CopyForDef);
1404
1405 // If we haven't found the beneficial conditions, prefer the VGPR form which
1406 // may result in less cross RC copies.
1407 if (Cost > 0)
1408 return false;
1409
1410 return rewrite(RewriteCands);
1411}
1412
1413bool UnclusteredHighRPStage::initGCNSchedStage() {
1414 if (DisableUnclusterHighRP)
1415 return false;
1416
1417 if (!GCNSchedStage::initGCNSchedStage())
1418 return false;
1419
1420 if (DAG.RegionsWithHighRP.none() && DAG.RegionsWithExcessRP.none())
1421 return false;
1422
1423 SavedMutations.swap(x&: DAG.Mutations);
1424 DAG.addMutation(
1425 Mutation: createIGroupLPDAGMutation(Phase: AMDGPU::SchedulingPhase::PreRAReentry));
1426
1427 InitialOccupancy = DAG.MinOccupancy;
1428 // Aggressively try to reduce register pressure in the unclustered high RP
1429 // stage. Temporarily increase occupancy target in the region.
1430 TempTargetOccupancy = MFI.getMaxWavesPerEU() > DAG.MinOccupancy
1431 ? InitialOccupancy + 1
1432 : InitialOccupancy;
1433 IsAnyRegionScheduled = false;
1434 S.SGPRLimitBias = S.HighRPSGPRBias;
1435 S.VGPRLimitBias = S.HighRPVGPRBias;
1436
1437 LLVM_DEBUG(
1438 dbgs()
1439 << "Retrying function scheduling without clustering. "
1440 "Aggressively try to reduce register pressure to achieve occupancy "
1441 << TempTargetOccupancy << ".\n");
1442
1443 return true;
1444}
1445
1446bool ClusteredLowOccStage::initGCNSchedStage() {
1447 if (DisableClusteredLowOccupancy)
1448 return false;
1449
1450 if (!GCNSchedStage::initGCNSchedStage())
1451 return false;
1452
1453 // Don't bother trying to improve ILP in lower RP regions if occupancy has not
1454 // been dropped. All regions will have already been scheduled with the ideal
1455 // occupancy targets.
1456 if (DAG.StartingOccupancy <= DAG.MinOccupancy)
1457 return false;
1458
1459 LLVM_DEBUG(
1460 dbgs() << "Retrying function scheduling with lowest recorded occupancy "
1461 << DAG.MinOccupancy << ".\n");
1462 return true;
1463}
1464
1465/// Allows to easily filter for this stage's debug output.
1466#define REMAT_PREFIX "[PreRARemat] "
1467#define REMAT_DEBUG(X) LLVM_DEBUG(dbgs() << REMAT_PREFIX; X;)
1468
1469#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
1470Printable PreRARematStage::ScoredRemat::print() const {
1471 return Printable([&](raw_ostream &OS) {
1472 OS << '(' << MaxFreq << ", " << FreqDiff << ", " << RegionImpact << ')';
1473 });
1474}
1475#endif
1476
1477bool PreRARematStage::initGCNSchedStage() {
1478 // FIXME: This pass will invalidate cached BBLiveInMap and MBBLiveIns for
1479 // regions inbetween the defs and region we sinked the def to. Will need to be
1480 // fixed if there is another pass after this pass.
1481 assert(!S.hasNextStage());
1482
1483 if (!GCNSchedStage::initGCNSchedStage() || DAG.Regions.size() <= 1)
1484 return false;
1485
1486#ifndef NDEBUG
1487 auto PrintTargetRegions = [&]() -> void {
1488 if (TargetRegions.none()) {
1489 dbgs() << REMAT_PREFIX << "No target regions\n";
1490 return;
1491 }
1492 dbgs() << REMAT_PREFIX << "Target regions:\n";
1493 for (unsigned I : TargetRegions.set_bits())
1494 dbgs() << REMAT_PREFIX << " [" << I << "] " << RPTargets[I] << '\n';
1495 };
1496#endif
1497
1498 // Set an objective for the stage based on current RP in each region.
1499 REMAT_DEBUG({
1500 dbgs() << "Analyzing ";
1501 MF.getFunction().printAsOperand(dbgs(), false);
1502 dbgs() << ": ";
1503 });
1504 if (!setObjective()) {
1505 LLVM_DEBUG(dbgs() << "no objective to achieve, occupancy is maximal at "
1506 << MFI.getMaxWavesPerEU() << '\n');
1507 return false;
1508 }
1509 LLVM_DEBUG({
1510 if (TargetOcc) {
1511 dbgs() << "increase occupancy from " << *TargetOcc - 1 << '\n';
1512 } else {
1513 dbgs() << "reduce spilling (minimum target occupancy is "
1514 << MFI.getMinWavesPerEU() << ")\n";
1515 }
1516 PrintTargetRegions();
1517 });
1518
1519 // We need up-to-date live-out info. to query live-out register masks in
1520 // regions containing rematerializable instructions.
1521 DAG.RegionLiveOuts.buildLiveRegMap();
1522
1523 if (!Remater.analyze()) {
1524 REMAT_DEBUG(dbgs() << "No rematerializable registers\n");
1525 return false;
1526 }
1527 const ScoredRemat::FreqInfo FreqInfo(MF, DAG);
1528
1529 // Set of registers already marked for potential remterialization; used to
1530 // avoid rematerialization chains.
1531 SmallSet<Register, 4> MarkedRegs;
1532
1533 // Collect candidates. We have more restrictions on what we can track here
1534 // compared to the rematerializer.
1535 SmallVector<ScoredRemat, 8> Candidates;
1536 // Map registers to candidate indices. Use ~0u as null value
1537 // since 0 is a valid index.
1538 IndexedMap<unsigned, VirtReg2IndexFunctor> DefRegToCandIdx(~0u);
1539 DefRegToCandIdx.resize(S: DAG.MRI.getNumVirtRegs());
1540 const unsigned NumRegions = DAG.Regions.size();
1541
1542 for (unsigned RegIdx = 0, E = Remater.getNumRegs(); RegIdx < E; ++RegIdx) {
1543 const Rematerializer::Reg &CandReg = Remater.getReg(RegIdx);
1544
1545 // All users must be in a single region.
1546 if (CandReg.Uses.size() != 1)
1547 continue;
1548 const auto [UseRegion, Users] = *CandReg.Uses.begin();
1549
1550 // Rematerialization moves the defining instruction into the region of its
1551 // use, which may sit under different control dependencies (e.g., across a
1552 // change of EXEC). Convergent operations must not be made control-dependent
1553 // on additional values, so they cannot be safely relocated this way. This
1554 // mirrors the check MachineSink performs before sinking an instruction.
1555 if (any_of(Range: CandReg.Defs,
1556 P: [](const MachineInstr *DefMI) { return DefMI->isConvergent(); }))
1557 continue;
1558
1559 // We further filter the registers that we can rematerialize based on our
1560 // current tracking capabilities in the stage. Users cannot themselves be
1561 // marked rematerializable, and no register operand of the defining MI can
1562 // be marked rematerializable. We also do not rematerialize an instruction
1563 // if it uses registers that aren't available at its use. This ensures that
1564 // we are not extending any live range while rematerializing.
1565 if (llvm::any_of(Range: Users, P: [&MarkedRegs](const MachineInstr *UserMI) {
1566 assert(UserMI->getNumOperands() > 0 &&
1567 "user must have at least one operand");
1568 const MachineOperand &UseMO = UserMI->getOperand(i: 0);
1569 return UseMO.isReg() && MarkedRegs.contains(V: UseMO.getReg());
1570 }))
1571 continue;
1572 MachineInstr *FirstUseMI =
1573 CandReg.getRegionUseBounds(UseRegion, LIS: *DAG.LIS).first;
1574 assert(FirstUseMI && "there must be a user in the region");
1575 SlotIndex FirstUseIdx =
1576 DAG.LIS->getInstructionIndex(Instr: *FirstUseMI).getRegSlot(EC: true);
1577 SlotIndex RefIdx =
1578 DAG.LIS->getInstructionIndex(Instr: *CandReg.getLastDef()).getRegSlot(EC: true);
1579 if (llvm::any_of(Range: CandReg.Dependencies, P: [&](RegisterIdx DepRegIdx) {
1580 const Rematerializer::Reg &DepReg = Remater.getReg(RegIdx: DepRegIdx);
1581 Register DepDefReg = DepReg.getDefReg();
1582 return MarkedRegs.contains(V: DepDefReg) ||
1583 !Remater.isRegIdenticalAtUses(Reg: DepDefReg, Mask: DepReg.Mask, RefSlot: RefIdx,
1584 Uses: {FirstUseIdx});
1585 }))
1586 continue;
1587 if (llvm::any_of(Range: Remater.getUnrematableDeps(RegIdx),
1588 P: [&](const std::pair<Register, LaneBitmask> &RegAndMask) {
1589 const auto &[Reg, Mask] = RegAndMask;
1590 return !Remater.isRegIdenticalAtUses(Reg, Mask, RefSlot: RefIdx,
1591 Uses: {FirstUseIdx});
1592 }))
1593 continue;
1594
1595 Register DefReg = CandReg.getDefReg();
1596 MarkedRegs.insert(V: DefReg);
1597 DefRegToCandIdx[DefReg] = Candidates.size();
1598 Candidates.emplace_back(Args&: RegIdx, Args: NumRegions);
1599 }
1600
1601 // Initialize the LiveIn and LiveOut sets of all candidates.
1602 // Iterating all regions and their live regs once is considerably
1603 // more efficient than querying those structures for each candidate
1604 // separately in ScoredRemat::init.
1605 for (unsigned I = 0; I < NumRegions; ++I) {
1606 for (const auto &[Reg, Mask] : DAG.LiveIns[I]) {
1607 if (!Register::isVirtualRegister(Reg))
1608 continue;
1609 unsigned CandIdx = DefRegToCandIdx[Reg];
1610 if (CandIdx != ~0u)
1611 Candidates[CandIdx].LiveIn.set(I);
1612 }
1613 for (const auto &[Reg, Mask] :
1614 DAG.RegionLiveOuts.getLiveRegsForRegionIdx(RegionIdx: I)) {
1615 if (!Register::isVirtualRegister(Reg))
1616 continue;
1617 unsigned CandIdx = DefRegToCandIdx[Reg];
1618 if (CandIdx != ~0u)
1619 Candidates[CandIdx].LiveOut.set(I);
1620 }
1621 }
1622
1623 // Finish initializing candidates.
1624 SmallVector<unsigned> CandidateOrder;
1625 for (auto [CandIdx, Cand] : enumerate(First&: Candidates)) {
1626 Cand.init(Freq: FreqInfo, Remater, DAG);
1627 Cand.update(TargetRegions, RPTargets, Freq: FreqInfo, ReduceSpill: !TargetOcc);
1628 if (!Cand.hasNullScore())
1629 CandidateOrder.push_back(Elt: CandIdx);
1630 }
1631
1632 if (TargetOcc) {
1633 // Every rematerialization we do here is likely to move the instruction
1634 // into a higher frequency region, increasing the total sum latency of the
1635 // instruction itself. This is acceptable if we are eliminating a spill in
1636 // the process, but when the goal is increasing occupancy we get nothing
1637 // out of rematerialization if occupancy is not increased in the end; in
1638 // such cases we want to roll back the rematerialization.
1639 Rollback = std::make_unique<RollbackSupport>(args&: Remater);
1640 }
1641
1642 // Rematerialize registers in successive rounds until all RP targets are
1643 // satisifed or until we run out of rematerialization candidates.
1644 BitVector RecomputeRP(DAG.Regions.size());
1645 for (;;) {
1646 RecomputeRP.reset();
1647
1648 // Sort candidates in increasing score order.
1649 sort(C&: CandidateOrder, Comp: [&](unsigned LHSIndex, unsigned RHSIndex) {
1650 return Candidates[LHSIndex] < Candidates[RHSIndex];
1651 });
1652
1653 REMAT_DEBUG({
1654 dbgs() << "==== NEW REMAT ROUND ====\n"
1655 << REMAT_PREFIX
1656 << "Candidates with non-null score, in rematerialization order:\n";
1657 for (const ScoredRemat &Cand : reverse(Candidates)) {
1658 dbgs() << REMAT_PREFIX << " " << Cand.print() << " | "
1659 << Remater.printRematReg(Cand.RegIdx) << '\n';
1660 }
1661 PrintTargetRegions();
1662 });
1663
1664 // Rematerialize registers in decreasing score order until we estimate
1665 // that all RP targets are satisfied or until rematerialization candidates
1666 // are no longer useful to decrease RP.
1667 while (!CandidateOrder.empty()) {
1668 const ScoredRemat &Cand = Candidates[CandidateOrder.back()];
1669 const Rematerializer::Reg &Reg = Remater.getReg(RegIdx: Cand.RegIdx);
1670
1671 // When previous rematerializations in this round have already satisfied
1672 // RP targets in all regions this rematerialization can impact, we have a
1673 // good indication that our scores have diverged significantly from
1674 // reality, in which case we interrupt this round and re-score. This also
1675 // ensures that every rematerialization we perform is possibly impactful
1676 // in at least one target region.
1677 if (!Cand.maybeBeneficial(TargetRegions, RPTargets)) {
1678 REMAT_DEBUG(dbgs() << "Interrupt round on stale score for "
1679 << Cand.print() << " | "
1680 << Remater.printRematReg(Cand.RegIdx));
1681 break;
1682 }
1683 CandidateOrder.pop_back();
1684
1685#ifdef EXPENSIVE_CHECKS
1686 // All uses are known to be available / live at the remat point. Thus,
1687 // the uses should already be live in to the using region.
1688 for (const MachineInstr *DefMI : Reg.Defs) {
1689 for (const MachineOperand &MO : DefMI->operands()) {
1690 // Exclude the defined register. We are rematerializing all
1691 // instructions defining it so we don't care that its value is
1692 // available at the remat point.
1693 if (!MO.isReg() || !MO.getReg() || !MO.readsReg() || MO.isDef())
1694 continue;
1695
1696 Register UseReg = MO.getReg();
1697 if (!UseReg.isVirtual())
1698 continue;
1699
1700 LiveInterval &LI = DAG.LIS->getInterval(UseReg);
1701 LaneBitmask LM = DAG.MRI.getMaxLaneMaskForVReg(MO.getReg());
1702 if (LI.hasSubRanges() && MO.getSubReg())
1703 LM = DAG.TRI->getSubRegIndexLaneMask(MO.getSubReg());
1704
1705 const unsigned UseRegion = Reg.Uses.begin()->first;
1706 LaneBitmask LiveInMask = DAG.LiveIns[UseRegion].at(UseReg);
1707 LaneBitmask UncoveredLanes = LM & ~(LiveInMask & LM);
1708 // If this register has lanes not covered by the LiveIns, be sure they
1709 // do not map to any subrange. ref:
1710 // machine-scheduler-sink-trivial-remats.mir::omitted_subrange
1711 if (UncoveredLanes.any()) {
1712 assert(LI.hasSubRanges());
1713 for (LiveInterval::SubRange &SR : LI.subranges())
1714 assert((SR.LaneMask & UncoveredLanes).none());
1715 }
1716 }
1717 }
1718#endif
1719
1720 // Remove the register from all regions where it is a live-in or live-out,
1721 // then rematerialize the register.
1722 REMAT_DEBUG(dbgs() << "** REMAT " << Remater.printRematReg(Cand.RegIdx)
1723 << '\n');
1724 removeFromLiveMaps(Reg: Reg.getDefReg(), LiveIn: Cand.LiveIn, LiveOut: Cand.LiveOut);
1725 if (Rollback) {
1726 Rollback->LiveMapUpdates.emplace_back(Args: Cand.RegIdx, Args: Cand.LiveIn,
1727 Args: Cand.LiveOut);
1728 }
1729 Cand.rematerialize(Remater);
1730
1731 // Adjust RP targets. The save is guaranteed in regions in which the
1732 // register is live-through and unused but optimistic in all other regions
1733 // where the register is live.
1734 updateRPTargets(Regions: Cand.Live, RPSave: Cand.RPSave);
1735 RecomputeRP |= Cand.UnpredictableRPSave;
1736 RescheduleRegions |= Cand.Live;
1737 if (!TargetRegions.any()) {
1738 REMAT_DEBUG(dbgs() << "All targets cleared, verifying...\n");
1739 break;
1740 }
1741 }
1742
1743 if (!updateAndVerifyRPTargets(Regions: RecomputeRP) && !TargetRegions.any()) {
1744 REMAT_DEBUG(dbgs() << "Objectives achieved!\n");
1745 break;
1746 }
1747
1748 // Update the score of remaining candidates and filter out those that have
1749 // become useless from the vector. Candidates never become useful after
1750 // having been useless for a round, so we can freely drop them without
1751 // losing any future rematerialization opportunity.
1752 unsigned NumUsefulCandidates = 0;
1753 for (unsigned CandIdx : CandidateOrder) {
1754 ScoredRemat &Candidate = Candidates[CandIdx];
1755 Candidate.update(TargetRegions, RPTargets, Freq: FreqInfo, ReduceSpill: !TargetOcc);
1756 if (!Candidate.hasNullScore())
1757 CandidateOrder[NumUsefulCandidates++] = CandIdx;
1758 }
1759 if (NumUsefulCandidates == 0) {
1760 REMAT_DEBUG(dbgs() << "Stop on exhausted rematerialization candidates\n");
1761 break;
1762 }
1763 CandidateOrder.truncate(N: NumUsefulCandidates);
1764 }
1765
1766 if (RescheduleRegions.none())
1767 return false;
1768
1769 // Commit all pressure changes to the DAG and compute minimum achieved
1770 // occupancy in impacted regions.
1771 REMAT_DEBUG(dbgs() << "==== REMAT RESULTS ====\n");
1772 unsigned DynamicVGPRBlockSize = MFI.getDynamicVGPRBlockSize();
1773 for (unsigned I : RescheduleRegions.set_bits()) {
1774 DAG.Pressure[I] = RPTargets[I].getCurrentRP();
1775 REMAT_DEBUG(dbgs() << '[' << I << "] Achieved occupancy "
1776 << DAG.Pressure[I].getOccupancy(ST, DynamicVGPRBlockSize)
1777 << " (" << RPTargets[I] << ")\n");
1778 }
1779 AchievedOcc = MFI.getMaxWavesPerEU();
1780 for (const GCNRegPressure &RP : DAG.Pressure) {
1781 AchievedOcc =
1782 std::min(a: AchievedOcc, b: RP.getOccupancy(ST, DynamicVGPRBlockSize));
1783 }
1784
1785 REMAT_DEBUG({
1786 dbgs() << "Retrying function scheduling with new min. occupancy of "
1787 << AchievedOcc << " from rematerializing (original was "
1788 << DAG.MinOccupancy;
1789 if (TargetOcc)
1790 dbgs() << ", target was " << *TargetOcc;
1791 dbgs() << ")\n";
1792 });
1793
1794 DAG.setTargetOccupancy(getStageTargetOccupancy());
1795 return true;
1796}
1797
1798void GCNSchedStage::finalizeGCNSchedStage() {
1799 DAG.finishBlock();
1800 LLVM_DEBUG(dbgs() << "Ending scheduling stage: " << StageID << "\n");
1801}
1802
1803void UnclusteredHighRPStage::finalizeGCNSchedStage() {
1804 SavedMutations.swap(x&: DAG.Mutations);
1805 S.SGPRLimitBias = S.VGPRLimitBias = 0;
1806 if (DAG.MinOccupancy > InitialOccupancy) {
1807 assert(IsAnyRegionScheduled);
1808 LLVM_DEBUG(dbgs() << StageID
1809 << " stage successfully increased occupancy to "
1810 << DAG.MinOccupancy << '\n');
1811 } else if (!IsAnyRegionScheduled) {
1812 assert(DAG.MinOccupancy == InitialOccupancy);
1813 LLVM_DEBUG(dbgs() << StageID
1814 << ": No regions scheduled, min occupancy stays at "
1815 << DAG.MinOccupancy << ", MFI occupancy stays at "
1816 << MFI.getOccupancy() << ".\n");
1817 }
1818
1819 GCNSchedStage::finalizeGCNSchedStage();
1820}
1821
1822bool GCNSchedStage::initGCNRegion() {
1823 // Skip empty scheduling region.
1824 if (DAG.begin() == DAG.end())
1825 return false;
1826
1827 // Check whether this new region is also a new block.
1828 if (DAG.RegionBegin->getParent() != CurrentMBB)
1829 setupNewBlock();
1830
1831 unsigned NumRegionInstrs = std::distance(first: DAG.begin(), last: DAG.end());
1832 DAG.enterRegion(bb: CurrentMBB, begin: DAG.begin(), end: DAG.end(), regioninstrs: NumRegionInstrs);
1833
1834 // Skip regions with 1 schedulable instruction.
1835 if (DAG.begin() == std::prev(x: DAG.end()))
1836 return false;
1837
1838 LLVM_DEBUG(dbgs() << "********** MI Scheduling **********\n");
1839 LLVM_DEBUG(dbgs() << MF.getName() << ":" << printMBBReference(*CurrentMBB)
1840 << " " << CurrentMBB->getName()
1841 << "\n From: " << *DAG.begin() << " To: ";
1842 if (DAG.RegionEnd != CurrentMBB->end()) dbgs() << *DAG.RegionEnd;
1843 else dbgs() << "End";
1844 dbgs() << " RegionInstrs: " << NumRegionInstrs << '\n');
1845
1846 // Save original instruction order before scheduling for possible revert.
1847 Unsched.clear();
1848 Unsched.reserve(n: DAG.NumRegionInstrs);
1849 if (StageID == GCNSchedStageID::OccInitialSchedule ||
1850 StageID == GCNSchedStageID::ILPInitialSchedule) {
1851 const SIInstrInfo *SII = static_cast<const SIInstrInfo *>(DAG.TII);
1852 for (auto &I : DAG) {
1853 Unsched.push_back(x: &I);
1854 if (SII->isIGLPMutationOnly(Opcode: I.getOpcode()))
1855 DAG.RegionsWithIGLPInstrs[RegionIdx] = true;
1856 }
1857 } else {
1858 for (auto &I : DAG)
1859 Unsched.push_back(x: &I);
1860 }
1861
1862 PressureBefore = DAG.Pressure[RegionIdx];
1863
1864 LLVM_DEBUG(
1865 dbgs() << "Pressure before scheduling:\nRegion live-ins:"
1866 << print(DAG.LiveIns[RegionIdx], DAG.MRI)
1867 << "Region live-in pressure: "
1868 << print(llvm::getRegPressure(DAG.MRI, DAG.LiveIns[RegionIdx]))
1869 << "Region register pressure: " << print(PressureBefore));
1870
1871 S.HasHighPressure = false;
1872 S.KnownExcessRP = isRegionWithExcessRP();
1873
1874 if (DAG.RegionsWithIGLPInstrs[RegionIdx] &&
1875 StageID != GCNSchedStageID::UnclusteredHighRPReschedule) {
1876 SavedMutations.clear();
1877 SavedMutations.swap(x&: DAG.Mutations);
1878 bool IsInitialStage = StageID == GCNSchedStageID::OccInitialSchedule ||
1879 StageID == GCNSchedStageID::ILPInitialSchedule;
1880 DAG.addMutation(Mutation: createIGroupLPDAGMutation(
1881 Phase: IsInitialStage ? AMDGPU::SchedulingPhase::Initial
1882 : AMDGPU::SchedulingPhase::PreRAReentry));
1883 }
1884
1885 return true;
1886}
1887
1888bool UnclusteredHighRPStage::initGCNRegion() {
1889 // Only reschedule regions that have excess register pressure (i.e. spilling)
1890 // or had minimum occupancy at the beginning of the stage (as long as
1891 // rescheduling of previous regions did not make occupancy drop back down to
1892 // the initial minimum).
1893 unsigned DynamicVGPRBlockSize = DAG.MFI.getDynamicVGPRBlockSize();
1894 // If no region has been scheduled yet, the DAG has not yet been updated with
1895 // the occupancy target. So retrieve it from the temporary.
1896 unsigned CurrentTargetOccupancy =
1897 IsAnyRegionScheduled ? DAG.MinOccupancy : TempTargetOccupancy;
1898 if (!DAG.RegionsWithExcessRP[RegionIdx] &&
1899 (CurrentTargetOccupancy <= InitialOccupancy ||
1900 DAG.Pressure[RegionIdx].getOccupancy(ST, DynamicVGPRBlockSize) !=
1901 InitialOccupancy))
1902 return false;
1903
1904 bool IsSchedulingThisRegion = GCNSchedStage::initGCNRegion();
1905 // If this is the first region scheduled during this stage, make the target
1906 // occupancy changes in the DAG and MFI.
1907 if (!IsAnyRegionScheduled && IsSchedulingThisRegion) {
1908 IsAnyRegionScheduled = true;
1909 if (MFI.getMaxWavesPerEU() > DAG.MinOccupancy)
1910 DAG.setTargetOccupancy(TempTargetOccupancy);
1911 }
1912 return IsSchedulingThisRegion;
1913}
1914
1915bool ClusteredLowOccStage::initGCNRegion() {
1916 // We may need to reschedule this region if it wasn't rescheduled in the last
1917 // stage, or if we found it was testing critical register pressure limits in
1918 // the unclustered reschedule stage. The later is because we may not have been
1919 // able to raise the min occupancy in the previous stage so the region may be
1920 // overly constrained even if it was already rescheduled.
1921 if (!DAG.RegionsWithHighRP[RegionIdx])
1922 return false;
1923
1924 return GCNSchedStage::initGCNRegion();
1925}
1926
1927bool PreRARematStage::initGCNRegion() {
1928 return !RevertAllRegions && RescheduleRegions[RegionIdx] &&
1929 GCNSchedStage::initGCNRegion();
1930}
1931
1932void GCNSchedStage::setupNewBlock() {
1933 if (CurrentMBB)
1934 DAG.finishBlock();
1935
1936 CurrentMBB = DAG.RegionBegin->getParent();
1937 DAG.startBlock(bb: CurrentMBB);
1938 // Get real RP for the region if it hasn't be calculated before. After the
1939 // initial schedule stage real RP will be collected after scheduling.
1940 if (StageID == GCNSchedStageID::OccInitialSchedule ||
1941 StageID == GCNSchedStageID::ILPInitialSchedule ||
1942 StageID == GCNSchedStageID::MemoryClauseInitialSchedule)
1943 DAG.computeBlockPressure(RegionIdx, MBB: CurrentMBB);
1944}
1945
1946void GCNSchedStage::finalizeGCNRegion() {
1947 DAG.Regions[RegionIdx] = std::pair(DAG.RegionBegin, DAG.RegionEnd);
1948 if (S.HasHighPressure)
1949 DAG.RegionsWithHighRP[RegionIdx] = true;
1950
1951 // Revert scheduling if we have dropped occupancy or there is some other
1952 // reason that the original schedule is better.
1953 checkScheduling();
1954
1955 if (DAG.RegionsWithIGLPInstrs[RegionIdx] &&
1956 StageID != GCNSchedStageID::UnclusteredHighRPReschedule)
1957 SavedMutations.swap(x&: DAG.Mutations);
1958}
1959
1960void PreRARematStage::finalizeGCNRegion() {
1961 GCNSchedStage::finalizeGCNRegion();
1962 // When the goal is to increase occupancy, all regions must reach the target
1963 // occupancy for rematerializations to be possibly useful, otherwise we will
1964 // just hurt latency for no benefit. If minimum occupancy drops below the
1965 // target there is no point in trying to re-schedule further regions.
1966 if (!TargetOcc)
1967 return;
1968 RegionReverts.emplace_back(Args&: RegionIdx, Args&: Unsched, Args&: PressureBefore);
1969 if (DAG.MinOccupancy < *TargetOcc) {
1970 REMAT_DEBUG(dbgs() << "Region " << RegionIdx
1971 << " cannot meet occupancy target, interrupting "
1972 "re-scheduling in all regions\n");
1973 RevertAllRegions = true;
1974 }
1975}
1976
1977void GCNSchedStage::checkScheduling() {
1978 // Check the results of scheduling.
1979 PressureAfter = DAG.getRealRegPressure(RegionIdx);
1980
1981 LLVM_DEBUG(dbgs() << "Pressure after scheduling: " << print(PressureAfter));
1982 LLVM_DEBUG(dbgs() << "Region: " << RegionIdx << ".\n");
1983
1984 unsigned DynamicVGPRBlockSize = DAG.MFI.getDynamicVGPRBlockSize();
1985
1986 if (PressureAfter.getSGPRNum() <= S.SGPRCriticalLimit &&
1987 PressureAfter.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) <= S.VGPRCriticalLimit) {
1988 DAG.Pressure[RegionIdx] = PressureAfter;
1989
1990 // Early out if we have achieved the occupancy target.
1991 LLVM_DEBUG(dbgs() << "Pressure in desired limits, done.\n");
1992 return;
1993 }
1994
1995 unsigned TargetOccupancy = std::min(
1996 a: S.getTargetOccupancy(), b: ST.getOccupancyWithWorkGroupSizes(MF).second);
1997 unsigned WavesAfter = std::min(
1998 a: TargetOccupancy, b: PressureAfter.getOccupancy(ST, DynamicVGPRBlockSize));
1999 unsigned WavesBefore = std::min(
2000 a: TargetOccupancy, b: PressureBefore.getOccupancy(ST, DynamicVGPRBlockSize));
2001 LLVM_DEBUG(dbgs() << "Occupancy before scheduling: " << WavesBefore
2002 << ", after " << WavesAfter << ".\n");
2003
2004 // We may not be able to keep the current target occupancy because of the just
2005 // scheduled region. We might still be able to revert scheduling if the
2006 // occupancy before was higher, or if the current schedule has register
2007 // pressure higher than the excess limits which could lead to more spilling.
2008 unsigned NewOccupancy = std::max(a: WavesAfter, b: WavesBefore);
2009
2010 // Allow memory bound functions to drop to 4 waves if not limited by an
2011 // attribute.
2012 if (WavesAfter < WavesBefore && WavesAfter < DAG.MinOccupancy &&
2013 WavesAfter >= MFI.getMinAllowedOccupancy()) {
2014 LLVM_DEBUG(dbgs() << "Function is memory bound, allow occupancy drop up to "
2015 << MFI.getMinAllowedOccupancy() << " waves\n");
2016 NewOccupancy = WavesAfter;
2017 }
2018
2019 if (NewOccupancy < DAG.MinOccupancy) {
2020 DAG.MinOccupancy = NewOccupancy;
2021 MFI.limitOccupancy(Limit: DAG.MinOccupancy);
2022 LLVM_DEBUG(dbgs() << "Occupancy lowered for the function to "
2023 << DAG.MinOccupancy << ".\n");
2024 }
2025 // The maximum number of arch VGPR on non-unified register file, or the
2026 // maximum VGPR + AGPR in the unified register file case.
2027 unsigned MaxVGPRs = ST.getMaxNumVGPRs(MF);
2028 // The maximum number of arch VGPR for both unified and non-unified register
2029 // file.
2030 unsigned MaxArchVGPRs = std::min(a: MaxVGPRs, b: ST.getAddressableNumArchVGPRs());
2031 unsigned MaxSGPRs = ST.getMaxNumSGPRs(MF);
2032
2033 if (PressureAfter.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) > MaxVGPRs ||
2034 PressureAfter.getArchVGPRNum() > MaxArchVGPRs ||
2035 PressureAfter.getAGPRNum() > MaxArchVGPRs ||
2036 PressureAfter.getSGPRNum() > MaxSGPRs) {
2037 DAG.RegionsWithHighRP[RegionIdx] = true;
2038 DAG.RegionsWithExcessRP[RegionIdx] = true;
2039 }
2040
2041 // Revert if this region's schedule would cause a drop in occupancy or
2042 // spilling.
2043 if (shouldRevertScheduling(WavesAfter)) {
2044 modifyRegionSchedule(RegionIdx, MIOrder: Unsched);
2045 std::tie(args&: DAG.RegionBegin, args&: DAG.RegionEnd) = DAG.Regions[RegionIdx];
2046 } else {
2047 DAG.Pressure[RegionIdx] = PressureAfter;
2048 }
2049}
2050
2051unsigned
2052GCNSchedStage::computeSUnitReadyCycle(const SUnit &SU, unsigned CurrCycle,
2053 DenseMap<unsigned, unsigned> &ReadyCycles,
2054 const TargetSchedModel &SM) {
2055 unsigned ReadyCycle = CurrCycle;
2056 for (auto &D : SU.Preds) {
2057 if (D.isAssignedRegDep()) {
2058 MachineInstr *DefMI = D.getSUnit()->getInstr();
2059 unsigned Latency = SM.computeInstrLatency(MI: DefMI);
2060 unsigned DefReady = ReadyCycles[DAG.getSUnit(MI: DefMI)->NodeNum];
2061 ReadyCycle = std::max(a: ReadyCycle, b: DefReady + Latency);
2062 }
2063 }
2064 ReadyCycles[SU.NodeNum] = ReadyCycle;
2065 return ReadyCycle;
2066}
2067
2068#ifndef NDEBUG
2069struct EarlierIssuingCycle {
2070 bool operator()(std::pair<MachineInstr *, unsigned> A,
2071 std::pair<MachineInstr *, unsigned> B) const {
2072 return A.second < B.second;
2073 }
2074};
2075
2076static void printScheduleModel(std::set<std::pair<MachineInstr *, unsigned>,
2077 EarlierIssuingCycle> &ReadyCycles) {
2078 if (ReadyCycles.empty())
2079 return;
2080 unsigned BBNum = ReadyCycles.begin()->first->getParent()->getNumber();
2081 dbgs() << "\n################## Schedule time ReadyCycles for MBB : " << BBNum
2082 << " ##################\n# Cycle #\t\t\tInstruction "
2083 " "
2084 " \n";
2085 unsigned IPrev = 1;
2086 for (auto &I : ReadyCycles) {
2087 if (I.second > IPrev + 1)
2088 dbgs() << "****************************** BUBBLE OF " << I.second - IPrev
2089 << " CYCLES DETECTED ******************************\n\n";
2090 dbgs() << "[ " << I.second << " ] : " << *I.first << "\n";
2091 IPrev = I.second;
2092 }
2093}
2094#endif
2095
2096ScheduleMetrics
2097GCNSchedStage::getScheduleMetrics(const std::vector<SUnit> &InputSchedule) {
2098#ifndef NDEBUG
2099 std::set<std::pair<MachineInstr *, unsigned>, EarlierIssuingCycle>
2100 ReadyCyclesSorted;
2101#endif
2102 const TargetSchedModel &SM = ST.getInstrInfo()->getSchedModel();
2103 unsigned SumBubbles = 0;
2104 DenseMap<unsigned, unsigned> ReadyCycles;
2105 unsigned CurrCycle = 0;
2106 for (auto &SU : InputSchedule) {
2107 unsigned ReadyCycle =
2108 computeSUnitReadyCycle(SU, CurrCycle, ReadyCycles, SM);
2109 SumBubbles += ReadyCycle - CurrCycle;
2110#ifndef NDEBUG
2111 ReadyCyclesSorted.insert(std::make_pair(SU.getInstr(), ReadyCycle));
2112#endif
2113 CurrCycle = ++ReadyCycle;
2114 }
2115#ifndef NDEBUG
2116 LLVM_DEBUG(
2117 printScheduleModel(ReadyCyclesSorted);
2118 dbgs() << "\n\t"
2119 << "Metric: "
2120 << (SumBubbles
2121 ? (SumBubbles * ScheduleMetrics::ScaleFactor) / CurrCycle
2122 : 1)
2123 << "\n\n");
2124#endif
2125
2126 return ScheduleMetrics(CurrCycle, SumBubbles);
2127}
2128
2129ScheduleMetrics
2130GCNSchedStage::getScheduleMetrics(const GCNScheduleDAGMILive &DAG) {
2131#ifndef NDEBUG
2132 std::set<std::pair<MachineInstr *, unsigned>, EarlierIssuingCycle>
2133 ReadyCyclesSorted;
2134#endif
2135 const TargetSchedModel &SM = ST.getInstrInfo()->getSchedModel();
2136 unsigned SumBubbles = 0;
2137 DenseMap<unsigned, unsigned> ReadyCycles;
2138 unsigned CurrCycle = 0;
2139 for (auto &MI : DAG) {
2140 SUnit *SU = DAG.getSUnit(MI: &MI);
2141 if (!SU)
2142 continue;
2143 unsigned ReadyCycle =
2144 computeSUnitReadyCycle(SU: *SU, CurrCycle, ReadyCycles, SM);
2145 SumBubbles += ReadyCycle - CurrCycle;
2146#ifndef NDEBUG
2147 ReadyCyclesSorted.insert(std::make_pair(SU->getInstr(), ReadyCycle));
2148#endif
2149 CurrCycle = ++ReadyCycle;
2150 }
2151#ifndef NDEBUG
2152 LLVM_DEBUG(
2153 printScheduleModel(ReadyCyclesSorted);
2154 dbgs() << "\n\t"
2155 << "Metric: "
2156 << (SumBubbles
2157 ? (SumBubbles * ScheduleMetrics::ScaleFactor) / CurrCycle
2158 : 1)
2159 << "\n\n");
2160#endif
2161
2162 return ScheduleMetrics(CurrCycle, SumBubbles);
2163}
2164
2165bool GCNSchedStage::shouldRevertScheduling(unsigned WavesAfter) {
2166 if (WavesAfter < DAG.MinOccupancy)
2167 return true;
2168
2169 // For dynamic VGPR mode, we don't want to waste any VGPR blocks.
2170 if (DAG.MFI.isDynamicVGPREnabled()) {
2171 unsigned BlocksBefore = AMDGPU::IsaInfo::getAllocatedNumVGPRBlocks(
2172 STI: ST, NumVGPRs: PressureBefore.getVGPRNum(UnifiedVGPRFile: false),
2173 DynamicVGPRBlockSize: DAG.MFI.getDynamicVGPRBlockSize());
2174 unsigned BlocksAfter = AMDGPU::IsaInfo::getAllocatedNumVGPRBlocks(
2175 STI: ST, NumVGPRs: PressureAfter.getVGPRNum(UnifiedVGPRFile: false), DynamicVGPRBlockSize: DAG.MFI.getDynamicVGPRBlockSize());
2176 if (BlocksAfter > BlocksBefore)
2177 return true;
2178 }
2179
2180 return false;
2181}
2182
2183bool OccInitialScheduleStage::shouldRevertScheduling(unsigned WavesAfter) {
2184 if (PressureAfter == PressureBefore)
2185 return false;
2186
2187 if (GCNSchedStage::shouldRevertScheduling(WavesAfter))
2188 return true;
2189
2190 if (mayCauseSpilling(WavesAfter))
2191 return true;
2192
2193 return false;
2194}
2195
2196bool UnclusteredHighRPStage::shouldRevertScheduling(unsigned WavesAfter) {
2197 // If RP is not reduced in the unclustered reschedule stage, revert to the
2198 // old schedule.
2199 if ((WavesAfter <=
2200 PressureBefore.getOccupancy(ST, DynamicVGPRBlockSize: DAG.MFI.getDynamicVGPRBlockSize()) &&
2201 mayCauseSpilling(WavesAfter)) ||
2202 GCNSchedStage::shouldRevertScheduling(WavesAfter)) {
2203 LLVM_DEBUG(dbgs() << "Unclustered reschedule did not help.\n");
2204 return true;
2205 }
2206
2207 // Do not attempt to relax schedule even more if we are already spilling.
2208 if (isRegionWithExcessRP())
2209 return false;
2210
2211 LLVM_DEBUG(
2212 dbgs()
2213 << "\n\t *** In shouldRevertScheduling ***\n"
2214 << " *********** BEFORE UnclusteredHighRPStage ***********\n");
2215 ScheduleMetrics MBefore = getScheduleMetrics(InputSchedule: DAG.SUnits);
2216 LLVM_DEBUG(
2217 dbgs()
2218 << "\n *********** AFTER UnclusteredHighRPStage ***********\n");
2219 ScheduleMetrics MAfter = getScheduleMetrics(DAG);
2220 unsigned OldMetric = MBefore.getMetric();
2221 unsigned NewMetric = MAfter.getMetric();
2222 unsigned WavesBefore = std::min(
2223 a: S.getTargetOccupancy(),
2224 b: PressureBefore.getOccupancy(ST, DynamicVGPRBlockSize: DAG.MFI.getDynamicVGPRBlockSize()));
2225 unsigned Profit =
2226 ((WavesAfter * ScheduleMetrics::ScaleFactor) / WavesBefore *
2227 ((OldMetric + ScheduleMetricBias) * ScheduleMetrics::ScaleFactor) /
2228 NewMetric) /
2229 ScheduleMetrics::ScaleFactor;
2230 LLVM_DEBUG(dbgs() << "\tMetric before " << MBefore << "\tMetric after "
2231 << MAfter << "Profit: " << Profit << "\n");
2232 return Profit < ScheduleMetrics::ScaleFactor;
2233}
2234
2235bool ClusteredLowOccStage::shouldRevertScheduling(unsigned WavesAfter) {
2236 if (PressureAfter == PressureBefore)
2237 return false;
2238
2239 if (GCNSchedStage::shouldRevertScheduling(WavesAfter))
2240 return true;
2241
2242 if (mayCauseSpilling(WavesAfter))
2243 return true;
2244
2245 return false;
2246}
2247
2248bool PreRARematStage::shouldRevertScheduling(unsigned WavesAfter) {
2249 // When trying to increase occupancy (TargetOcc == true) the stage manages
2250 // region reverts globally (all or none), so we always return false here.
2251 return !TargetOcc && mayCauseSpilling(WavesAfter);
2252}
2253
2254bool ILPInitialScheduleStage::shouldRevertScheduling(unsigned WavesAfter) {
2255 if (mayCauseSpilling(WavesAfter))
2256 return true;
2257
2258 return false;
2259}
2260
2261bool MemoryClauseInitialScheduleStage::shouldRevertScheduling(
2262 unsigned WavesAfter) {
2263 return mayCauseSpilling(WavesAfter);
2264}
2265
2266static cl::opt<bool> EnableLiveIntervalRPReschedule(
2267 "amdgpu-lirp-reschedule", cl::Hidden,
2268 cl::desc("Enable live interval RP reschedule stage"), cl::init(Val: true));
2269
2270static cl::opt<unsigned> LiveIntervalRPThreshold(
2271 "amdgpu-lirp-threshold", cl::Hidden,
2272 cl::desc("Percent increase of live interval RP over instant pressure to "
2273 "trigger rescheduling"),
2274 cl::init(Val: 10));
2275
2276static cl::opt<unsigned> LiveIntervalRPVGPRReduction(
2277 "amdgpu-lirp-vgpr-reduction", cl::Hidden,
2278 cl::desc(
2279 "Reduction factor (percent) for VGPR threshold during live interval RP "
2280 "reschedule stage"),
2281 cl::init(Val: 90));
2282
2283static cl::opt<unsigned> LiveIntervalRPInstantLowerBound(
2284 "amdgpu-lirp-instant-lower-bound", cl::Hidden,
2285 cl::desc("Lower bound (percent of the VGPR excess limit) on instant RP, "
2286 "below which a region is skipped"),
2287 cl::init(Val: 10));
2288
2289bool LiveIntervalRPStage::initGCNSchedStage() {
2290 if (!EnableLiveIntervalRPReschedule)
2291 return false;
2292
2293 if (!GCNSchedStage::initGCNSchedStage())
2294 return false;
2295
2296 if (!S.VGPRThresholdPercent) {
2297 LLVM_DEBUG(dbgs() << "LIRP: expected VGPRThresholdPercent to be enabled, "
2298 "not using live interval RP reschedule stage\n");
2299 return false;
2300 }
2301
2302 return true;
2303}
2304
2305bool LiveIntervalRPStage::initGCNRegion() {
2306 unsigned InstantRP = DAG.Pressure[RegionIdx].getArchVGPRNum();
2307 auto [RegionBegin, RegionEnd] = DAG.Regions[RegionIdx];
2308 if (RegionBegin == RegionEnd)
2309 return false;
2310
2311 unsigned LIRP = estimateGreedyVGPRPressure(
2312 RegionBegin, RegionEnd, LiveIns: DAG.LiveIns[RegionIdx], LIS: *DAG.getLIS(),
2313 MRI: DAG.MF.getRegInfo(), TRI: static_cast<const SIRegisterInfo &>(*DAG.TRI));
2314
2315 unsigned NewVGPRThresholdPercent =
2316 (S.VGPRThresholdPercent * LiveIntervalRPVGPRReduction + 99) / 100;
2317
2318 LLVM_DEBUG(dbgs() << "LIRP: Region " << RegionIdx
2319 << ", VGPRThresholdPercent: " << S.VGPRThresholdPercent
2320 << " -> " << NewVGPRThresholdPercent
2321 << ", VGPRExcessLimit=" << S.VGPRExcessLimit
2322 << ", VGPRCriticalLimit=" << S.VGPRCriticalLimit
2323 << ", InstantRP=" << InstantRP << ", LIRP=" << LIRP);
2324
2325 bool DoRescheduling = false;
2326 // Lower bound on InstantRP to skip over tiny regions.
2327 unsigned InstantRPLowerBound =
2328 S.VGPRExcessLimit * LiveIntervalRPInstantLowerBound / 100;
2329 if (LIRP > S.VGPRExcessLimit) {
2330 LLVM_DEBUG(dbgs() << " [LIRP exceeds the limit (" << S.VGPRExcessLimit
2331 << "), rescheduling]");
2332 DoRescheduling = true;
2333 } else if (LIRP > InstantRP && InstantRP > InstantRPLowerBound) {
2334 unsigned IncreasePercent = ((LIRP - InstantRP) * 100) / InstantRP;
2335 if (IncreasePercent > LiveIntervalRPThreshold) {
2336 LLVM_DEBUG(dbgs() << " [" << IncreasePercent << "% > "
2337 << LiveIntervalRPThreshold << "%, rescheduling]");
2338 DoRescheduling = true;
2339 }
2340 }
2341 LLVM_DEBUG(dbgs() << '\n');
2342
2343 if (DoRescheduling && GCNSchedStage::initGCNRegion()) {
2344 SavedVGPRExcessLimit = S.VGPRExcessLimit;
2345 SavedVGPRCriticalLimit = S.VGPRCriticalLimit;
2346 SavedVGPRThresholdPercent = S.VGPRThresholdPercent;
2347 S.VGPRThresholdPercent = NewVGPRThresholdPercent;
2348 return true;
2349 }
2350
2351 return false;
2352}
2353
2354void LiveIntervalRPStage::finalizeGCNRegion() {
2355 S.VGPRExcessLimit = SavedVGPRExcessLimit;
2356 S.VGPRCriticalLimit = SavedVGPRCriticalLimit;
2357 S.VGPRThresholdPercent = SavedVGPRThresholdPercent;
2358 GCNSchedStage::finalizeGCNRegion();
2359}
2360
2361bool GCNSchedStage::mayCauseSpilling(unsigned WavesAfter) {
2362 if (WavesAfter <= MFI.getMinWavesPerEU() && isRegionWithExcessRP() &&
2363 !PressureAfter.less(MF, O: PressureBefore)) {
2364 LLVM_DEBUG(dbgs() << "New pressure will result in more spilling.\n");
2365 return true;
2366 }
2367
2368 return false;
2369}
2370
2371void GCNSchedStage::modifyRegionSchedule(unsigned RegionIdx,
2372 ArrayRef<MachineInstr *> MIOrder) {
2373 assert(static_cast<size_t>(std::distance(DAG.Regions[RegionIdx].first,
2374 DAG.Regions[RegionIdx].second)) ==
2375 MIOrder.size() &&
2376 "instruction number mismatch");
2377 if (MIOrder.empty())
2378 return;
2379
2380 LLVM_DEBUG(dbgs() << "Reverting scheduling for region " << RegionIdx << '\n');
2381
2382 // Reconstruct MI sequence by moving instructions in desired order before
2383 // the current region's start.
2384 MachineBasicBlock::iterator RegionEnd = DAG.Regions[RegionIdx].first;
2385 MachineBasicBlock *MBB = MIOrder.front()->getParent();
2386 for (MachineInstr *MI : MIOrder) {
2387 // Either move the next MI in order before the end of the region or move the
2388 // region end past the MI if it is at the correct position.
2389 MachineBasicBlock::iterator MII = MI->getIterator();
2390 if (MII != RegionEnd) {
2391 // Will subsequent splice move MI up past a non-debug instruction?
2392 bool NonDebugReordered =
2393 !MI->isDebugInstr() &&
2394 skipDebugInstructionsForward(It: RegionEnd, End: MII) != MII;
2395 MBB->splice(Where: RegionEnd, Other: MBB, From: MI);
2396 // Only update LiveIntervals information if non-debug instructions are
2397 // reordered. Otherwise debug instructions could cause code generation to
2398 // change.
2399 if (NonDebugReordered)
2400 DAG.LIS->handleMove(MI&: *MI, UpdateFlags: true);
2401 } else {
2402 // MI is already at the expected position. However, earlier splices in
2403 // this loop may have changed neighboring slot indices, so this MI's
2404 // slot index can become non-monotonic w.r.t. the physical MBB order.
2405 // Only re-seat when monotonicity is actually violated to avoid
2406 // unnecessary LiveInterval changes that could perturb scheduling.
2407 if (!MI->isDebugInstr()) {
2408 SlotIndex MIIdx = DAG.LIS->getInstructionIndex(Instr: *MI);
2409 SlotIndex PrevIdx = DAG.LIS->getSlotIndexes()->getIndexBefore(MI: *MI);
2410 if (PrevIdx >= MIIdx)
2411 DAG.LIS->handleMove(MI&: *MI, UpdateFlags: true);
2412 }
2413 ++RegionEnd;
2414 }
2415 if (MI->isDebugInstr()) {
2416 LLVM_DEBUG(dbgs() << "Scheduling " << *MI);
2417 continue;
2418 }
2419
2420 // Reset read-undef flags and update them later.
2421 RegisterOperands::restoreLivenessFlags(MI&: *MI, TRI: *DAG.TRI, MRI: DAG.MRI, LIS&: *DAG.LIS,
2422 TrackLaneMasks: DAG.ShouldTrackLaneMasks);
2423 LLVM_DEBUG(dbgs() << "Scheduling " << *MI);
2424 }
2425
2426 // The region end doesn't change throughout scheduling since it itself is
2427 // outside the region (whether that is a MBB end or a terminator MI).
2428 assert(RegionEnd == DAG.Regions[RegionIdx].second && "region end mismatch");
2429 DAG.Regions[RegionIdx].first = MIOrder.front();
2430}
2431
2432/// Returns true if reaching def \p RD will be in AGPR form after the rewrite
2433/// and so needs no bridge copy: a candidate MFMA in \p RewriteSet, an
2434/// AV_MOV_*_IMM_PSEUDO, or a copy from a candidate src2 reg in \p CandSrc2Regs.
2435/// A non-candidate MFMA stays in VGPR form and still needs a bridge.
2436static bool isReachingDefAGPRForm(
2437 MachineInstr *RD, const SmallPtrSetImpl<MachineInstr *> &RewriteSet,
2438 const DenseSet<Register> &CandSrc2Regs, const SIInstrInfo &TII) {
2439 if (TII.isMAI(MI: *RD))
2440 return RewriteSet.contains(Ptr: RD);
2441 if (RD->getOpcode() == AMDGPU::AV_MOV_B32_IMM_PSEUDO ||
2442 RD->getOpcode() == AMDGPU::AV_MOV_B64_IMM_PSEUDO)
2443 return true;
2444 if (RD->isCopy() && CandSrc2Regs.contains(V: RD->getOperand(i: 1).getReg()))
2445 return true;
2446 return false;
2447}
2448
2449bool RewriteMFMAFormStage::hasUseRequiringVGPR(
2450 ArrayRef<SlotIndex> Src2ReachingDefs,
2451 const SmallPtrSetImpl<MachineInstr *> &RewriteSet) {
2452 for (SlotIndex RDIdx : Src2ReachingDefs) {
2453 const MachineInstr *RD = DAG.LIS->getInstructionFromIndex(index: RDIdx);
2454 SmallVector<MachineOperand *, 8> ReachingUses;
2455 findReachingUses(DefMI: RD, LIS: DAG.LIS, ReachingUses);
2456 for (const MachineOperand *UseMO : ReachingUses) {
2457 const MachineInstr *UseMI = UseMO->getParent();
2458 if (UseMI->isCopy())
2459 continue;
2460 if (TII->isMAI(MI: *UseMI) && RewriteSet.contains(Ptr: UseMI))
2461 continue;
2462 return true;
2463 }
2464 }
2465 return false;
2466}
2467
2468void RewriteMFMAFormStage::resetRewriteCandsToVGPR(
2469 ArrayRef<std::pair<MachineInstr *, unsigned>> RewriteCands) {
2470 for (auto [MI, OriginalOpcode] : RewriteCands) {
2471 assert(TII->isMAI(*MI));
2472 const TargetRegisterClass *ADefRC =
2473 DAG.MRI.getRegClass(Reg: MI->getOperand(i: 0).getReg());
2474 const TargetRegisterClass *VDefRC = SRI->getEquivalentVGPRClass(SRC: ADefRC);
2475 DAG.MRI.setRegClass(Reg: MI->getOperand(i: 0).getReg(), RC: VDefRC);
2476 MI->setDesc(TII->get(Opcode: OriginalOpcode));
2477
2478 MachineOperand *Src2 = TII->getNamedOperand(MI&: *MI, OperandName: AMDGPU::OpName::src2);
2479 if (!Src2->isReg())
2480 continue;
2481
2482 // Have to get src types separately since subregs may cause C and D
2483 // registers to be different types even though the actual operand is
2484 // the same size.
2485 const TargetRegisterClass *AUseRC = DAG.MRI.getRegClass(Reg: Src2->getReg());
2486 const TargetRegisterClass *VUseRC = SRI->getEquivalentVGPRClass(SRC: AUseRC);
2487 DAG.MRI.setRegClass(Reg: Src2->getReg(), RC: VUseRC);
2488 }
2489}
2490
2491bool RewriteMFMAFormStage::isRewriteCandidate(MachineInstr *MI) const {
2492 if (!static_cast<const SIInstrInfo *>(DAG.TII)->isMAI(MI: *MI))
2493 return false;
2494 if (AMDGPU::getAGPRFormOp(Opcode: MI->getOpcode()) == -1)
2495 return false;
2496 // Reject candidates whose users force an unavoidable bridge copy.
2497 Register DstReg = MI->getOperand(i: 0).getReg();
2498 for (const MachineInstr &UseMI : DAG.MRI.use_nodbg_instructions(Reg: DstReg)) {
2499 if (!TII->isMAI(MI: UseMI) && !UseMI.isCopy())
2500 return false;
2501 }
2502 return true;
2503}
2504
2505bool RewriteMFMAFormStage::initHeuristics(
2506 std::vector<std::pair<MachineInstr *, unsigned>> &RewriteCands,
2507 DenseMap<MachineBasicBlock *, std::set<Register>> &CopyForUse,
2508 SmallPtrSetImpl<MachineInstr *> &CopyForDef) {
2509 bool Changed = false;
2510
2511 // Collect the candidate group, its members share AGPR-form operands
2512 // post-rewrite, so reaching defs feeding any member don't need bridge copy.
2513 SmallPtrSet<MachineInstr *, 16> RewriteSet;
2514 DenseSet<Register> CandSrc2Regs;
2515 for (MachineBasicBlock &MBB : MF) {
2516 for (MachineInstr &MI : MBB) {
2517 if (!isRewriteCandidate(MI: &MI))
2518 continue;
2519 RewriteSet.insert(Ptr: &MI);
2520 MachineOperand *Src2 = TII->getNamedOperand(MI, OperandName: AMDGPU::OpName::src2);
2521 if (Src2 && Src2->isReg())
2522 CandSrc2Regs.insert(V: Src2->getReg());
2523 }
2524 }
2525
2526 // Prepare for the heuristics
2527 for (MachineBasicBlock &MBB : MF) {
2528 for (MachineInstr &MI : MBB) {
2529 if (!isRewriteCandidate(MI: &MI))
2530 continue;
2531
2532 int ReplacementOp = AMDGPU::getAGPRFormOp(Opcode: MI.getOpcode());
2533 assert(ReplacementOp != -1);
2534
2535 RewriteCands.push_back(x: {&MI, MI.getOpcode()});
2536 MI.setDesc(TII->get(Opcode: ReplacementOp));
2537
2538 MachineOperand *Src2 = TII->getNamedOperand(MI, OperandName: AMDGPU::OpName::src2);
2539 if (Src2->isReg()) {
2540 SmallVector<SlotIndex, 8> Src2ReachingDefs;
2541 findReachingDefs(UseMO&: *Src2, LIS: DAG.LIS, DefIdxs&: Src2ReachingDefs);
2542
2543 // If src2 has a use that must remain VGPR, it cannot be reclassified to
2544 // AGPR.
2545 bool Src2NeedsVGPR = hasUseRequiringVGPR(Src2ReachingDefs, RewriteSet);
2546 Src2NeedsVGPRCache[&MI] = Src2NeedsVGPR;
2547
2548 for (SlotIndex RDIdx : Src2ReachingDefs) {
2549 MachineInstr *RD = DAG.LIS->getInstructionFromIndex(index: RDIdx);
2550 if (!Src2NeedsVGPR &&
2551 isReachingDefAGPRForm(RD, RewriteSet, CandSrc2Regs, TII: *TII))
2552 continue;
2553 CopyForDef.insert(Ptr: RD);
2554 }
2555 }
2556
2557 MachineOperand &Dst = MI.getOperand(i: 0);
2558 SmallVector<MachineOperand *, 8> DstReachingUses;
2559
2560 findReachingUses(DefMI: &MI, LIS: DAG.LIS, ReachingUses&: DstReachingUses);
2561
2562 for (MachineOperand *RUOp : DstReachingUses) {
2563 MachineInstr *UserMI = RUOp->getParent();
2564 // Group members read the AGPR result directly.
2565 if (TII->isMAI(MI: *UserMI) && RewriteSet.contains(Ptr: UserMI))
2566 continue;
2567
2568 // For any user of the result of the MFMA which is not an MFMA, we
2569 // insert a copy. For a given register, we will only insert one copy
2570 // per user block.
2571 CopyForUse[UserMI->getParent()].insert(x: RUOp->getReg());
2572
2573 if (TII->isMAI(MI: *UserMI))
2574 continue;
2575
2576 SmallVector<SlotIndex, 8> DstUsesReachingDefs;
2577 findReachingDefs(UseMO&: *RUOp, LIS: DAG.LIS, DefIdxs&: DstUsesReachingDefs);
2578
2579 for (SlotIndex RDIndex : DstUsesReachingDefs) {
2580 MachineInstr *RD = DAG.LIS->getInstructionFromIndex(index: RDIndex);
2581 if (TII->isMAI(MI: *RD))
2582 continue;
2583
2584 // For any definition of the user of the MFMA which is not an MFMA,
2585 // we insert a copy. We do this to transform all the reaching defs
2586 // of this use to AGPR. By doing this, we can insert a copy from
2587 // AGPR to VGPR at the user rather than after the MFMA.
2588 CopyForDef.insert(Ptr: RD);
2589 }
2590 }
2591
2592 // Do the rewrite to allow for updated RP calculation.
2593 const TargetRegisterClass *VDefRC = DAG.MRI.getRegClass(Reg: Dst.getReg());
2594 const TargetRegisterClass *ADefRC = SRI->getEquivalentAGPRClass(SRC: VDefRC);
2595 DAG.MRI.setRegClass(Reg: Dst.getReg(), RC: ADefRC);
2596 if (Src2->isReg()) {
2597 // Have to get src types separately since subregs may cause C and D
2598 // registers to be different types even though the actual operand is
2599 // the same size.
2600 const TargetRegisterClass *VUseRC = DAG.MRI.getRegClass(Reg: Src2->getReg());
2601 const TargetRegisterClass *AUseRC = SRI->getEquivalentAGPRClass(SRC: VUseRC);
2602 DAG.MRI.setRegClass(Reg: Src2->getReg(), RC: AUseRC);
2603 }
2604 Changed = true;
2605 }
2606 }
2607
2608 return Changed;
2609}
2610
2611int64_t RewriteMFMAFormStage::getRewriteCost(
2612 ArrayRef<std::pair<MachineInstr *, unsigned>> RewriteCands,
2613 const DenseMap<MachineBasicBlock *, std::set<Register>> &CopyForUse,
2614 const SmallPtrSetImpl<MachineInstr *> &CopyForDef) {
2615 MachineBlockFrequencyInfo *MBFI = DAG.MBFI;
2616
2617 int64_t BestSpillCost = 0;
2618 int64_t Cost = 0;
2619 uint64_t EntryFreq = MBFI->getEntryFreq().getFrequency();
2620
2621 std::pair<unsigned, unsigned> MaxVectorRegs =
2622 ST.getMaxNumVectorRegs(F: MF.getFunction());
2623 unsigned ArchVGPRThreshold = MaxVectorRegs.first;
2624 unsigned AGPRThreshold = MaxVectorRegs.second;
2625 unsigned CombinedThreshold = ST.getMaxNumVGPRs(MF);
2626
2627 for (unsigned Region = 0; Region < DAG.Regions.size(); Region++) {
2628 if (!RegionsWithExcessArchVGPR[Region])
2629 continue;
2630
2631 GCNRegPressure &PressureBefore = DAG.Pressure[Region];
2632 unsigned SpillCostBefore = PressureBefore.getVGPRSpills(
2633 MF, ArchVGPRThreshold, AGPRThreshold, CombinedThreshold);
2634
2635 // For the cases we care about (i.e. ArchVGPR usage is greater than the
2636 // addressable limit), rewriting alone should bring pressure to manageable
2637 // level. If we find any such region, then the rewrite is potentially
2638 // beneficial.
2639 GCNRegPressure PressureAfter = DAG.getRealRegPressure(RegionIdx: Region);
2640 unsigned SpillCostAfter = PressureAfter.getVGPRSpills(
2641 MF, ArchVGPRThreshold, AGPRThreshold, CombinedThreshold);
2642
2643 uint64_t BlockFreq =
2644 MBFI->getBlockFreq(MBB: DAG.Regions[Region].first->getParent())
2645 .getFrequency();
2646
2647 bool RelativeFreqIsDenom = EntryFreq > BlockFreq;
2648 uint64_t RelativeFreq = EntryFreq && BlockFreq
2649 ? (RelativeFreqIsDenom ? EntryFreq / BlockFreq
2650 : BlockFreq / EntryFreq)
2651 : 1;
2652
2653 // This assumes perfect spilling / splitting -- using one spill / copy
2654 // instruction and one restoreFrom / copy for each excess register,
2655 int64_t SpillCost = ((int)SpillCostAfter - (int)SpillCostBefore) * 2;
2656
2657 // Also account for the block frequency.
2658 if (RelativeFreqIsDenom)
2659 SpillCost /= (int64_t)RelativeFreq;
2660 else
2661 SpillCost *= (int64_t)RelativeFreq;
2662
2663 // If we have increased spilling in any block, just bail.
2664 if (SpillCost > 0) {
2665 resetRewriteCandsToVGPR(RewriteCands);
2666 return SpillCost;
2667 }
2668
2669 if (SpillCost < BestSpillCost)
2670 BestSpillCost = SpillCost;
2671 }
2672
2673 // Set the cost to the largest decrease in spill cost in order to not double
2674 // count spill reductions.
2675 Cost = BestSpillCost;
2676 assert(Cost <= 0);
2677
2678 unsigned CopyCost = 0;
2679
2680 // For each CopyForDef, increase the cost by the register size while
2681 // accounting for block frequency.
2682 for (MachineInstr *DefMI : CopyForDef) {
2683 Register DefReg = DefMI->getOperand(i: 0).getReg();
2684 uint64_t DefFreq =
2685 EntryFreq
2686 ? MBFI->getBlockFreq(MBB: DefMI->getParent()).getFrequency() / EntryFreq
2687 : 1;
2688
2689 const TargetRegisterClass *RC = DAG.MRI.getRegClass(Reg: DefReg);
2690 CopyCost += RC->getCopyCost() * DefFreq;
2691 }
2692
2693 // Account for CopyForUse copies in each block that the register is used.
2694 for (auto &[UseBlock, UseRegs] : CopyForUse) {
2695 uint64_t UseFreq =
2696 EntryFreq ? MBFI->getBlockFreq(MBB: UseBlock).getFrequency() / EntryFreq : 1;
2697
2698 for (Register UseReg : UseRegs) {
2699 const TargetRegisterClass *RC = DAG.MRI.getRegClass(Reg: UseReg);
2700 CopyCost += RC->getCopyCost() * UseFreq;
2701 }
2702 }
2703
2704 // Reset the classes that were changed to AGPR for better register bank
2705 // analysis. We must do rewriting after copy-insertion, as some defs of the
2706 // register may require VGPR. Additionally, if we bail out and don't perform
2707 // the rewrite then these need to be restored anyway.
2708 resetRewriteCandsToVGPR(RewriteCands);
2709
2710 return Cost + CopyCost;
2711}
2712
2713bool RewriteMFMAFormStage::rewrite(
2714 ArrayRef<std::pair<MachineInstr *, unsigned>> RewriteCands) {
2715 DenseMap<MachineInstr *, unsigned> FirstMIToRegion;
2716 DenseMap<MachineInstr *, unsigned> LastMIToRegion;
2717
2718 for (unsigned Region = 0; Region < DAG.Regions.size(); Region++) {
2719 RegionBoundaries Entry = DAG.Regions[Region];
2720 if (Entry.first == Entry.second)
2721 continue;
2722
2723 FirstMIToRegion[&*Entry.first] = Region;
2724 if (Entry.second != Entry.first->getParent()->end())
2725 LastMIToRegion[&*Entry.second] = Region;
2726 }
2727
2728 // Rewrite the MFMAs to AGPR, and insert any copies as needed.
2729 // The general assumption of the algorithm (and the previous cost calculation)
2730 // is that it is better to insert the copies in the MBB of the def of the src2
2731 // operands, and in the MBB of the user of the dest operands. This is based on
2732 // the assumption that the MFMAs are likely to appear in loop bodies, while
2733 // the src2 and dest operands are live-in / live-out of the loop. Due to this
2734 // design, the algorithm for finding copy insertion points is more
2735 // complicated.
2736 //
2737 // There are three main cases to handle: 1. the reaching defs of the src2
2738 // operands, 2. the reaching uses of the dst operands, and 3. the reaching
2739 // defs of the reaching uses of the dst operand.
2740 //
2741 // In the first case, we simply insert copies after each of the reaching
2742 // definitions. In the second case, we collect all the uses of a given dest
2743 // and organize them by MBB. Then, we insert 1 copy for each MBB before the
2744 // earliest use. Since the use may have multiple reaching defs, and since we
2745 // want to replace the register it is using with the result of the copy, we
2746 // must handle case 3. In the third case, we simply insert a copy after each
2747 // of the reaching defs to connect to the copy of the reaching uses of the dst
2748 // reg. This allows us to avoid inserting copies next to the MFMAs.
2749 //
2750 // While inserting the copies, we maintain a map of operands which will use
2751 // different regs (i.e. the result of the copies). For example, a case 1 src2
2752 // operand will use the register result of the copies after the reaching defs,
2753 // as opposed to the original register. Now that we have completed our copy
2754 // analysis and placement, we can bulk update the registers. We do this
2755 // separately as to avoid complicating the reachingDef and reachingUse
2756 // queries.
2757 //
2758 // While inserting the copies, we also maintain a list or registers which we
2759 // will want to reclassify as AGPR. After doing the copy insertion and the
2760 // register replacement, we can finally do the reclassification. This uses the
2761 // redef map, as the registers we are interested in reclassifying may be
2762 // replaced by the result of a copy. We must do this after the copy analysis
2763 // and placement as we must have an accurate redef map -- otherwise we may end
2764 // up creating illegal instructions.
2765
2766 // The original registers of the MFMA that need to be reclassified as AGPR.
2767 DenseSet<Register> RewriteRegs;
2768 // The map of an original register in the MFMA to a new register (result of a
2769 // copy) that it should be replaced with.
2770 DenseMap<Register, Register> RedefMap;
2771 // The map of the original MFMA registers to the relevant MFMA operands.
2772 DenseMap<Register, DenseSet<MachineOperand *>> ReplaceMap;
2773 // The map of reaching defs for a given register -- to avoid duplicate copies.
2774 DenseMap<Register, SmallPtrSet<MachineInstr *, 8>> ReachingDefCopyMap;
2775 // The map of reaching uses for a given register by basic block -- to avoid
2776 // duplicate copies and to calculate per MBB insert pts.
2777 DenseMap<unsigned, DenseMap<Register, SmallPtrSet<MachineOperand *, 8>>>
2778 ReachingUseTracker;
2779
2780 // Collect the candidate group; its members share AGPR-form operands
2781 // post-rewrite, so reaching defs feeding any member need no bridge copy.
2782 SmallPtrSet<MachineInstr *, 16> RewriteCandsSet;
2783 DenseSet<Register> RewriteSrc2Regs;
2784 for (auto &[MI, OriginalOpcode] : RewriteCands) {
2785 RewriteCandsSet.insert(Ptr: MI);
2786 MachineOperand *Src2 = TII->getNamedOperand(MI&: *MI, OperandName: AMDGPU::OpName::src2);
2787 if (Src2 && Src2->isReg())
2788 RewriteSrc2Regs.insert(V: Src2->getReg());
2789 }
2790
2791 for (auto &[MI, OriginalOpcode] : RewriteCands) {
2792 int ReplacementOp = AMDGPU::getAGPRFormOp(Opcode: MI->getOpcode());
2793 if (ReplacementOp == -1)
2794 continue;
2795 MI->setDesc(TII->get(Opcode: ReplacementOp));
2796
2797 // Case 1: insert copies for the reaching defs of the Src2Reg.
2798 MachineOperand *Src2 = TII->getNamedOperand(MI&: *MI, OperandName: AMDGPU::OpName::src2);
2799 if (Src2->isReg()) {
2800 Register Src2Reg = Src2->getReg();
2801 if (!Src2Reg.isVirtual())
2802 return false;
2803
2804 Register MappedReg = Src2->getReg();
2805 SmallVector<SlotIndex, 8> Src2ReachingDefs;
2806 findReachingDefs(UseMO&: *Src2, LIS: DAG.LIS, DefIdxs&: Src2ReachingDefs);
2807 SmallSetVector<MachineInstr *, 8> Src2DefsReplace;
2808
2809 // If src2 has a use that must remain VGPR, it cannot be reclassified to
2810 // AGPR.
2811 bool Src2NeedsVGPR = Src2NeedsVGPRCache.lookup(Val: MI);
2812
2813 for (SlotIndex RDIndex : Src2ReachingDefs) {
2814 MachineInstr *RD = DAG.LIS->getInstructionFromIndex(index: RDIndex);
2815 if (!Src2NeedsVGPR &&
2816 isReachingDefAGPRForm(RD, RewriteSet: RewriteCandsSet, CandSrc2Regs: RewriteSrc2Regs, TII: *TII))
2817 continue;
2818
2819 Src2DefsReplace.insert(X: RD);
2820 }
2821
2822 if (!Src2DefsReplace.empty()) {
2823 auto RI = RedefMap.find(Val: Src2Reg);
2824 if (RI != RedefMap.end()) {
2825 MappedReg = RI->second;
2826 } else {
2827 assert(!ReachingDefCopyMap.contains(Src2Reg));
2828 const TargetRegisterClass *Src2RC = DAG.MRI.getRegClass(Reg: Src2Reg);
2829 const TargetRegisterClass *VGPRRC =
2830 SRI->getEquivalentVGPRClass(SRC: Src2RC);
2831
2832 // Track the mapping of the original register to the new register.
2833 MappedReg = DAG.MRI.createVirtualRegister(RegClass: VGPRRC);
2834 RedefMap[Src2Reg] = MappedReg;
2835 }
2836
2837 // If none exists, create a copy from this reaching def.
2838 // We may have inserted a copy already in an earlier iteration.
2839 for (MachineInstr *RD : Src2DefsReplace) {
2840 // Do not create redundant copies.
2841 if (ReachingDefCopyMap[Src2Reg].insert(Ptr: RD).second) {
2842 MachineInstrBuilder VGPRCopy =
2843 BuildMI(BB&: *RD->getParent(), I: std::next(x: RD->getIterator()),
2844 MIMD: RD->getDebugLoc(), MCID: TII->get(Opcode: TargetOpcode::COPY))
2845 .addDef(RegNo: MappedReg, Flags: {}, SubReg: 0)
2846 .addUse(RegNo: Src2Reg, Flags: {}, SubReg: 0);
2847 DAG.LIS->InsertMachineInstrInMaps(MI&: *VGPRCopy);
2848
2849 // If this reaching def was the last MI in the region, update the
2850 // region boundaries.
2851 if (LastMIToRegion.contains(Val: RD)) {
2852 unsigned UpdateRegion = LastMIToRegion[RD];
2853 DAG.Regions[UpdateRegion].second = VGPRCopy;
2854 LastMIToRegion.erase(Val: RD);
2855 }
2856 }
2857 }
2858 }
2859
2860 // Track the register for reclassification
2861 RewriteRegs.insert(V: Src2Reg);
2862
2863 // Always insert the operand for replacement. If this corresponds with a
2864 // chain of tied-def we may not see the VGPR requirement until later.
2865 ReplaceMap[Src2Reg].insert(V: Src2);
2866 }
2867
2868 // Case 2 and Case 3: insert copies before the reaching uses of the dsts,
2869 // and after the reaching defs of the reaching uses of the dsts.
2870
2871 MachineOperand *Dst = &MI->getOperand(i: 0);
2872 Register DstReg = Dst->getReg();
2873 if (!DstReg.isVirtual())
2874 return false;
2875
2876 Register MappedReg = DstReg;
2877 SmallVector<MachineOperand *, 8> DstReachingUses;
2878
2879 SmallVector<MachineOperand *, 8> DstReachingUseCopies;
2880 SmallVector<MachineInstr *, 8> DstUseDefsReplace;
2881
2882 findReachingUses(DefMI: MI, LIS: DAG.LIS, ReachingUses&: DstReachingUses);
2883
2884 for (MachineOperand *RUOp : DstReachingUses) {
2885 MachineInstr *UserMI = RUOp->getParent();
2886 // Group members read the AGPR result directly.
2887 if (TII->isMAI(MI: *UserMI) && RewriteCandsSet.contains(Ptr: UserMI))
2888 continue;
2889
2890 // If there is a non mai reaching use, then we need a copy.
2891 if (find(Range&: DstReachingUseCopies, Val: RUOp) == DstReachingUseCopies.end())
2892 DstReachingUseCopies.push_back(Elt: RUOp);
2893
2894 // Non-rewritten MAI: its defs aren't being reclassified.
2895 if (TII->isMAI(MI: *UserMI))
2896 continue;
2897
2898 SmallVector<SlotIndex, 8> DstUsesReachingDefs;
2899 findReachingDefs(UseMO&: *RUOp, LIS: DAG.LIS, DefIdxs&: DstUsesReachingDefs);
2900
2901 for (SlotIndex RDIndex : DstUsesReachingDefs) {
2902 MachineInstr *RD = DAG.LIS->getInstructionFromIndex(index: RDIndex);
2903 if (TII->isMAI(MI: *RD))
2904 continue;
2905
2906 // If there is a non mai reaching def of this reaching use, then we will
2907 // need a copy.
2908 if (find(Range&: DstUseDefsReplace, Val: RD) == DstUseDefsReplace.end())
2909 DstUseDefsReplace.push_back(Elt: RD);
2910 }
2911 }
2912
2913 if (!DstUseDefsReplace.empty()) {
2914 auto RI = RedefMap.find(Val: DstReg);
2915 if (RI != RedefMap.end()) {
2916 MappedReg = RI->second;
2917 } else {
2918 assert(!ReachingDefCopyMap.contains(DstReg));
2919 const TargetRegisterClass *DstRC = DAG.MRI.getRegClass(Reg: DstReg);
2920 const TargetRegisterClass *VGPRRC = SRI->getEquivalentVGPRClass(SRC: DstRC);
2921
2922 // Track the mapping of the original register to the new register.
2923 MappedReg = DAG.MRI.createVirtualRegister(RegClass: VGPRRC);
2924 RedefMap[DstReg] = MappedReg;
2925 }
2926
2927 // If none exists, create a copy from this reaching def.
2928 // We may have inserted a copy already in an earlier iteration.
2929 for (MachineInstr *RD : DstUseDefsReplace) {
2930 // Do not create reundant copies.
2931 if (ReachingDefCopyMap[DstReg].insert(Ptr: RD).second) {
2932 MachineInstrBuilder VGPRCopy =
2933 BuildMI(BB&: *RD->getParent(), I: std::next(x: RD->getIterator()),
2934 MIMD: RD->getDebugLoc(), MCID: TII->get(Opcode: TargetOpcode::COPY))
2935 .addDef(RegNo: MappedReg, Flags: {}, SubReg: 0)
2936 .addUse(RegNo: DstReg, Flags: {}, SubReg: 0);
2937 DAG.LIS->InsertMachineInstrInMaps(MI&: *VGPRCopy);
2938
2939 // If this reaching def was the last MI in the region, update the
2940 // region boundaries.
2941 auto LMI = LastMIToRegion.find(Val: RD);
2942 if (LMI != LastMIToRegion.end()) {
2943 unsigned UpdateRegion = LMI->second;
2944 DAG.Regions[UpdateRegion].second = VGPRCopy;
2945 LastMIToRegion.erase(Val: RD);
2946 }
2947 }
2948 }
2949 }
2950
2951 DenseSet<MachineOperand *> &DstRegSet = ReplaceMap[DstReg];
2952 // One AGPR→VGPR copy per dst register, shared by all same-block uses.
2953 Register SameBlockCopyReg;
2954 MachineInstr *EarliestSameBlockUse = nullptr;
2955 for (MachineOperand *RU : DstReachingUseCopies) {
2956 MachineBasicBlock *RUBlock = RU->getParent()->getParent();
2957 // Just keep track of the reaching use of this register by block. After we
2958 // have scanned all the MFMAs we can find optimal insert pts.
2959 if (RUBlock != MI->getParent()) {
2960 ReachingUseTracker[RUBlock->getNumber()][DstReg].insert(Ptr: RU);
2961 continue;
2962 }
2963
2964 // Lazily create the copy register on first same-block use.
2965 if (!SameBlockCopyReg.isValid()) {
2966 const TargetRegisterClass *DstRC = DAG.MRI.getRegClass(Reg: DstReg);
2967 const TargetRegisterClass *VGPRRC = SRI->getEquivalentVGPRClass(SRC: DstRC);
2968 SameBlockCopyReg = DAG.MRI.createVirtualRegister(RegClass: VGPRRC);
2969 }
2970
2971 // Track the earliest use for copy insertion point.
2972 MachineInstr *UseInst = RU->getParent();
2973 if (!EarliestSameBlockUse ||
2974 SlotIndex::isEarlierInstr(
2975 A: DAG.LIS->getInstructionIndex(Instr: *UseInst),
2976 B: DAG.LIS->getInstructionIndex(Instr: *EarliestSameBlockUse)))
2977 EarliestSameBlockUse = UseInst;
2978 RU->setReg(SameBlockCopyReg);
2979 }
2980
2981 // Insert the copy before the earliest same-block use.
2982 if (SameBlockCopyReg.isValid()) {
2983 MachineInstrBuilder VGPRCopy =
2984 BuildMI(BB&: *EarliestSameBlockUse->getParent(),
2985 I: EarliestSameBlockUse->getIterator(), MIMD: DebugLoc(),
2986 MCID: TII->get(Opcode: TargetOpcode::COPY), DestReg: SameBlockCopyReg)
2987 .addUse(RegNo: DstReg, Flags: {}, SubReg: 0);
2988 DAG.LIS->InsertMachineInstrInMaps(MI&: *VGPRCopy);
2989 DstRegSet.insert(V: &VGPRCopy->getOperand(i: 1));
2990 }
2991
2992 // Track the register for reclassification
2993 RewriteRegs.insert(V: DstReg);
2994
2995 // Insert the dst operand for replacement. If this dst is in a chain of
2996 // tied-def MFMAs, and the first src2 needs to be replaced with a new reg,
2997 // all the correspond operands need to be replaced.
2998 DstRegSet.insert(V: Dst);
2999 }
3000
3001 // Handle the copies for dst uses.
3002 using RUBType =
3003 std::pair<unsigned, DenseMap<Register, SmallPtrSet<MachineOperand *, 8>>>;
3004 for (RUBType RUBlockEntry : ReachingUseTracker) {
3005 using RUDType = std::pair<Register, SmallPtrSet<MachineOperand *, 8>>;
3006 for (RUDType RUDst : RUBlockEntry.second) {
3007 MachineOperand *OpBegin = *RUDst.second.begin();
3008 SlotIndex InstPt = DAG.LIS->getInstructionIndex(Instr: *OpBegin->getParent());
3009
3010 // Find the earliest use in this block.
3011 for (MachineOperand *User : RUDst.second) {
3012 SlotIndex NewInstPt = DAG.LIS->getInstructionIndex(Instr: *User->getParent());
3013 if (SlotIndex::isEarlierInstr(A: NewInstPt, B: InstPt))
3014 InstPt = NewInstPt;
3015 }
3016
3017 const TargetRegisterClass *DstRC = DAG.MRI.getRegClass(Reg: RUDst.first);
3018 const TargetRegisterClass *VGPRRC = SRI->getEquivalentVGPRClass(SRC: DstRC);
3019 Register NewUseReg = DAG.MRI.createVirtualRegister(RegClass: VGPRRC);
3020 MachineInstr *UseInst = DAG.LIS->getInstructionFromIndex(index: InstPt);
3021
3022 MachineInstrBuilder VGPRCopy =
3023 BuildMI(BB&: *UseInst->getParent(), I: UseInst->getIterator(),
3024 MIMD: UseInst->getDebugLoc(), MCID: TII->get(Opcode: TargetOpcode::COPY))
3025 .addDef(RegNo: NewUseReg, Flags: {}, SubReg: 0)
3026 .addUse(RegNo: RUDst.first, Flags: {}, SubReg: 0);
3027 DAG.LIS->InsertMachineInstrInMaps(MI&: *VGPRCopy);
3028
3029 // If this UseInst was the first MI in the region, update the region
3030 // boundaries.
3031 auto FI = FirstMIToRegion.find(Val: UseInst);
3032 if (FI != FirstMIToRegion.end()) {
3033 unsigned UpdateRegion = FI->second;
3034 DAG.Regions[UpdateRegion].first = VGPRCopy;
3035 FirstMIToRegion.erase(Val: UseInst);
3036 }
3037
3038 // Replace the operand for all users.
3039 for (MachineOperand *User : RUDst.second) {
3040 User->setReg(NewUseReg);
3041 }
3042
3043 // Track the copy source operand for replacement.
3044 ReplaceMap[RUDst.first].insert(V: &VGPRCopy->getOperand(i: 1));
3045 }
3046 }
3047
3048 // We may have needed to insert copies after the reaching defs of the MFMAs.
3049 // Replace the original register with the result of the copy for all relevant
3050 // operands.
3051 for (std::pair<Register, Register> NewDef : RedefMap) {
3052 Register OldReg = NewDef.first;
3053 Register NewReg = NewDef.second;
3054
3055 // Replace the register for any associated operand in the MFMA chain.
3056 for (MachineOperand *ReplaceOp : ReplaceMap[OldReg])
3057 ReplaceOp->setReg(NewReg);
3058 }
3059
3060 // Finally, do the reclassification of the MFMA registers.
3061 for (Register RewriteReg : RewriteRegs) {
3062 Register RegToRewrite = RewriteReg;
3063
3064 // Be sure to update the replacement register and not the original.
3065 auto RI = RedefMap.find(Val: RewriteReg);
3066 if (RI != RedefMap.end())
3067 RegToRewrite = RI->second;
3068
3069 const TargetRegisterClass *CurrRC = DAG.MRI.getRegClass(Reg: RegToRewrite);
3070 const TargetRegisterClass *AGPRRC = SRI->getEquivalentAGPRClass(SRC: CurrRC);
3071
3072 DAG.MRI.setRegClass(Reg: RegToRewrite, RC: AGPRRC);
3073 }
3074
3075 // Bulk update the LIS.
3076 DAG.LIS->reanalyze(MF&: DAG.MF);
3077 // Liveins may have been modified for cross RC copies
3078 RegionPressureMap LiveInUpdater(&DAG, false);
3079 LiveInUpdater.buildLiveRegMap();
3080
3081 for (unsigned Region = 0; Region < DAG.Regions.size(); Region++)
3082 DAG.LiveIns[Region] = LiveInUpdater.getLiveRegsForRegionIdx(RegionIdx: Region);
3083
3084 DAG.Pressure[RegionIdx] = DAG.getRealRegPressure(RegionIdx);
3085
3086 return true;
3087}
3088
3089unsigned PreRARematStage::getStageTargetOccupancy() const {
3090 return TargetOcc ? *TargetOcc : MFI.getMinWavesPerEU();
3091}
3092
3093bool PreRARematStage::setObjective() {
3094 const Function &F = MF.getFunction();
3095
3096 // Set up "spilling targets" for all regions.
3097 unsigned MaxSGPRs = ST.getMaxNumSGPRs(F);
3098 unsigned MaxVGPRs = ST.getMaxNumVGPRs(F);
3099 bool HasVectorRegisterExcess = false;
3100 for (unsigned I = 0, E = DAG.Regions.size(); I != E; ++I) {
3101 const GCNRegPressure &RP = DAG.Pressure[I];
3102 GCNRPTarget &Target = RPTargets.emplace_back(Args&: MaxSGPRs, Args&: MaxVGPRs, Args&: MF, Args: RP);
3103 if (!Target.satisfied())
3104 TargetRegions.set(I);
3105 HasVectorRegisterExcess |= Target.hasVectorRegisterExcess();
3106 }
3107
3108 if (HasVectorRegisterExcess || DAG.MinOccupancy >= MFI.getMaxWavesPerEU()) {
3109 // In addition to register usage being above addressable limits, occupancy
3110 // below the minimum is considered like "spilling" as well.
3111 TargetOcc = std::nullopt;
3112 } else {
3113 // There is no spilling and room to improve occupancy; set up "increased
3114 // occupancy targets" for all regions.
3115 TargetOcc = DAG.MinOccupancy + 1;
3116 const unsigned VGPRBlockSize = MFI.getDynamicVGPRBlockSize();
3117 MaxSGPRs = ST.getMaxNumSGPRs(WavesPerEU: *TargetOcc, Addressable: false);
3118 MaxVGPRs = ST.getMaxNumVGPRs(WavesPerEU: *TargetOcc, DynamicVGPRBlockSize: VGPRBlockSize);
3119 for (auto [I, Target] : enumerate(First&: RPTargets)) {
3120 Target.setTarget(NumSGPRs: MaxSGPRs, NumVGPRs: MaxVGPRs);
3121 if (!Target.satisfied())
3122 TargetRegions.set(I);
3123 }
3124 }
3125
3126 return TargetRegions.any();
3127}
3128
3129bool PreRARematStage::ScoredRemat::maybeBeneficial(
3130 const BitVector &TargetRegions, ArrayRef<GCNRPTarget> RPTargets) const {
3131 for (unsigned I : TargetRegions.set_bits()) {
3132 if (Live[I] && RPTargets[I].isSaveBeneficial(SaveRP: RPSave))
3133 return true;
3134 }
3135 return false;
3136}
3137
3138PreRARematStage::ScoredRemat::FreqInfo::FreqInfo(
3139 MachineFunction &MF, const GCNScheduleDAGMILive &DAG) {
3140 MachineBranchProbabilityInfo MBPI;
3141 MachineCycleInfo MCI;
3142 MCI.compute(F&: MF);
3143 MachineBlockFrequencyInfo MBFI(MF, MBPI, MCI);
3144
3145 const unsigned NumRegions = DAG.Regions.size();
3146 MinFreq = MBFI.getEntryFreq().getFrequency();
3147 MaxFreq = 0;
3148 Regions.reserve(N: NumRegions);
3149 for (unsigned I = 0; I < NumRegions; ++I) {
3150 MachineBasicBlock *MBB = DAG.Regions[I].first->getParent();
3151 uint64_t BlockFreq = MBFI.getBlockFreq(MBB).getFrequency();
3152 Regions.push_back(Elt: BlockFreq);
3153 if (BlockFreq && BlockFreq < MinFreq)
3154 MinFreq = BlockFreq;
3155 else if (BlockFreq > MaxFreq)
3156 MaxFreq = BlockFreq;
3157 }
3158 if (!MinFreq)
3159 return;
3160
3161 // Scale everything down if frequencies are high.
3162 if (MinFreq >= ScaleFactor * ScaleFactor) {
3163 for (uint64_t &Freq : Regions)
3164 Freq /= ScaleFactor;
3165 MinFreq /= ScaleFactor;
3166 MaxFreq /= ScaleFactor;
3167 }
3168}
3169
3170void PreRARematStage::ScoredRemat::init(const FreqInfo &Freq,
3171 const Rematerializer &Remater,
3172 GCNScheduleDAGMILive &DAG) {
3173 const Rematerializer::Reg &Reg = Remater.getReg(RegIdx);
3174 Register DefReg = Reg.getDefReg();
3175 assert(Reg.Uses.size() == 1 && "expected users in single region");
3176 const unsigned UseRegion = Reg.Uses.begin()->first;
3177
3178 Live |= LiveIn;
3179 Live |= LiveOut;
3180
3181 for (unsigned I : Live.set_bits()) {
3182 // If the register is both unused and live-through in the region, the
3183 // latter's RP is guaranteed to decrease.
3184 if (!LiveIn[I] || !LiveOut[I] || I == UseRegion)
3185 UnpredictableRPSave.set(I);
3186 }
3187 RPSave.inc(Reg: DefReg, PrevMask: LaneBitmask::getNone(), NewMask: Reg.Mask, MRI: DAG.MRI);
3188
3189 // Get frequencies of defining and using regions. A rematerialization from the
3190 // least frequent region to the most frequent region will yield the greatest
3191 // in order to penalize rematerializations from or into regions whose
3192 int64_t DefOrMin = std::max(a: Freq.Regions[Reg.DefRegion], b: Freq.MinFreq);
3193 int64_t UseOrMax = Freq.Regions[UseRegion];
3194 if (!UseOrMax)
3195 UseOrMax = Freq.MaxFreq;
3196 FreqDiff = DefOrMin - UseOrMax;
3197}
3198
3199void PreRARematStage::ScoredRemat::update(const BitVector &TargetRegions,
3200 ArrayRef<GCNRPTarget> RPTargets,
3201 const FreqInfo &FreqInfo,
3202 bool ReduceSpill) {
3203 MaxFreq = 0;
3204 RegionImpact = 0;
3205 for (unsigned I : TargetRegions.set_bits()) {
3206 if (!Live[I])
3207 continue;
3208
3209 // The rematerialization must contribute positively in at least one
3210 // register class with usage above the RP target for this region to
3211 // contribute to the score.
3212 const GCNRPTarget &RegionTarget = RPTargets[I];
3213 const unsigned NumRegsBenefit = RegionTarget.getNumRegsBenefit(SaveRP: RPSave);
3214 if (!NumRegsBenefit)
3215 continue;
3216
3217 // Regions in which RP is guaranteed to decrease have more weight.
3218 RegionImpact += (UnpredictableRPSave[I] ? 1 : 2) * NumRegsBenefit;
3219
3220 if (ReduceSpill) {
3221 uint64_t Freq = FreqInfo.Regions[I];
3222 if (UnpredictableRPSave[I]) {
3223 // Apply a frequency penalty in regions in which we are not sure that RP
3224 // will decrease.
3225 Freq /= 2;
3226 }
3227 MaxFreq = std::max(a: MaxFreq, b: Freq);
3228 }
3229 }
3230}
3231
3232void PreRARematStage::ScoredRemat::rematerialize(
3233 Rematerializer &Remater) const {
3234 const Rematerializer::Reg &Reg = Remater.getReg(RegIdx);
3235 Rematerializer::DependencyReuseInfo DRI;
3236 for (RegisterIdx DepRegIdx : Reg.Dependencies)
3237 DRI.reuse(DepIdx: DepRegIdx);
3238 unsigned UseRegion = Reg.Uses.begin()->first;
3239 Remater.rematerializeToRegion(RootIdx: RegIdx, UseRegion, DRI);
3240}
3241
3242void PreRARematStage::updateRPTargets(const BitVector &Regions,
3243 const GCNRegPressure &RPSave) {
3244 for (unsigned I : Regions.set_bits()) {
3245 RPTargets[I].saveRP(SaveRP: RPSave);
3246 if (TargetRegions[I] && RPTargets[I].satisfied()) {
3247 REMAT_DEBUG(dbgs() << " [" << I << "] Target reached!\n");
3248 TargetRegions.reset(Idx: I);
3249 }
3250 }
3251}
3252
3253bool PreRARematStage::updateAndVerifyRPTargets(const BitVector &Regions) {
3254 bool TooOptimistic = false;
3255 for (unsigned I : Regions.set_bits()) {
3256 GCNRPTarget &Target = RPTargets[I];
3257 Target.setRP(DAG.getRealRegPressure(RegionIdx: I));
3258
3259 // Since we were optimistic in assessing RP decreases in these regions, we
3260 // may need to remark the target as a target region if RP didn't decrease
3261 // as expected.
3262 if (!TargetRegions[I] && !Target.satisfied()) {
3263 REMAT_DEBUG(dbgs() << " [" << I << "] Incorrect RP estimation\n");
3264 TooOptimistic = true;
3265 TargetRegions.set(I);
3266 }
3267 }
3268 return TooOptimistic;
3269}
3270
3271void PreRARematStage::removeFromLiveMaps(Register Reg, const BitVector &LiveIn,
3272 const BitVector &LiveOut) {
3273 assert(LiveIn.size() == DAG.Regions.size() &&
3274 LiveOut.size() == DAG.Regions.size() && "region num mismatch");
3275 for (unsigned I : LiveIn.set_bits())
3276 DAG.LiveIns[I].erase(Val: Reg);
3277 for (unsigned I : LiveOut.set_bits())
3278 DAG.RegionLiveOuts.getLiveRegsForRegionIdx(RegionIdx: I).erase(Val: Reg);
3279}
3280
3281void PreRARematStage::addToLiveMaps(Register Reg, LaneBitmask Mask,
3282 const BitVector &LiveIn,
3283 const BitVector &LiveOut) {
3284 assert(LiveIn.size() == DAG.Regions.size() &&
3285 LiveOut.size() == DAG.Regions.size() && "region num mismatch");
3286 std::pair<Register, LaneBitmask> LiveReg(Reg, Mask);
3287 for (unsigned I : LiveIn.set_bits())
3288 DAG.LiveIns[I].insert(KV: LiveReg);
3289 for (unsigned I : LiveOut.set_bits())
3290 DAG.RegionLiveOuts.getLiveRegsForRegionIdx(RegionIdx: I).insert(KV: LiveReg);
3291}
3292
3293void PreRARematStage::finalizeGCNSchedStage() {
3294 // We consider that reducing spilling is always beneficial so we never
3295 // rollback rematerializations or revert scheduling in such cases.
3296 if (!TargetOcc)
3297 return;
3298
3299 // When increasing occupancy, it is possible that re-scheduling is not able to
3300 // achieve the target occupancy in all regions, in which case re-scheduling in
3301 // all regions should be reverted.
3302 if (DAG.MinOccupancy >= *TargetOcc)
3303 return;
3304
3305 // Revert re-scheduling in all affected regions.
3306 for (const auto &[RegionIdx, OrigMIOrder, MaxPressure] : RegionReverts) {
3307 REMAT_DEBUG(dbgs() << "Reverting re-scheduling in region " << RegionIdx
3308 << '\n');
3309 DAG.Pressure[RegionIdx] = MaxPressure;
3310 modifyRegionSchedule(RegionIdx, MIOrder: OrigMIOrder);
3311 }
3312
3313 // It is possible that re-scheduling lowers occupancy over the one achieved
3314 // just through rematerializations, in which case we revert re-scheduling in
3315 // all regions but do not roll back rematerializations.
3316 if (AchievedOcc >= *TargetOcc) {
3317 DAG.setTargetOccupancy(AchievedOcc);
3318 return;
3319 }
3320
3321 // Reset the target occupancy to what it was pre-rematerialization.
3322 DAG.setTargetOccupancy(*TargetOcc - 1);
3323
3324 // Roll back changes made by the stage, then recompute pressure in all
3325 // affected regions.
3326 REMAT_DEBUG(dbgs() << "==== ROLLBACK ====\n");
3327 assert(Rollback && "rollbacker should be defined");
3328 Rollback->Listener.rollback(Remater);
3329 for (const auto &[RegIdx, LiveIn, LiveOut] : Rollback->LiveMapUpdates) {
3330 const Rematerializer::Reg &Reg = Remater.getReg(RegIdx);
3331 addToLiveMaps(Reg: Reg.getDefReg(), Mask: Reg.Mask, LiveIn, LiveOut);
3332 }
3333
3334#ifdef EXPENSIVE_CHECKS
3335 // In particular, we want to check for coherent MI/slot order in regions in
3336 // which reverts and/or rollbacks may have happened.
3337 MF.verify();
3338#endif
3339 for (unsigned I : RescheduleRegions.set_bits())
3340 DAG.Pressure[I] = DAG.getRealRegPressure(RegionIdx: I);
3341
3342 GCNSchedStage::finalizeGCNSchedStage();
3343}
3344
3345void GCNScheduleDAGMILive::setTargetOccupancy(unsigned TargetOccupancy) {
3346 MinOccupancy = TargetOccupancy;
3347 if (MFI.getOccupancy() < TargetOccupancy)
3348 MFI.increaseOccupancy(MF, Limit: MinOccupancy);
3349 else
3350 MFI.limitOccupancy(Limit: MinOccupancy);
3351}
3352
3353static bool hasIGLPInstrs(ScheduleDAGInstrs *DAG) {
3354 const SIInstrInfo *SII = static_cast<const SIInstrInfo *>(DAG->TII);
3355 return any_of(Range&: *DAG, P: [SII](MachineBasicBlock::iterator MI) {
3356 return SII->isIGLPMutationOnly(Opcode: MI->getOpcode());
3357 });
3358}
3359
3360GCNPostScheduleDAGMILive::GCNPostScheduleDAGMILive(
3361 MachineSchedContext *C, std::unique_ptr<MachineSchedStrategy> S,
3362 bool RemoveKillFlags)
3363 : ScheduleDAGMI(C, std::move(S), RemoveKillFlags) {}
3364
3365void GCNPostScheduleDAGMILive::schedule() {
3366 HasIGLPInstrs = hasIGLPInstrs(DAG: this);
3367 if (HasIGLPInstrs) {
3368 SavedMutations.clear();
3369 SavedMutations.swap(x&: Mutations);
3370 addMutation(Mutation: createIGroupLPDAGMutation(Phase: AMDGPU::SchedulingPhase::PostRA));
3371 }
3372
3373 ScheduleDAGMI::schedule();
3374}
3375
3376void GCNPostScheduleDAGMILive::finalizeSchedule() {
3377 if (HasIGLPInstrs)
3378 SavedMutations.swap(x&: Mutations);
3379
3380 ScheduleDAGMI::finalizeSchedule();
3381}
3382