1//===-- GCNPreRAAntiHints.cpp - MFMA register anti-hints ------------------===//
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/// Insert register allocation anti-hints.
11///
12//===----------------------------------------------------------------------===//
13
14#include "GCNPreRAAntiHints.h"
15#include "GCNSubtarget.h"
16#include "SIInstrInfo.h"
17#include "SIRegisterInfo.h"
18#include "llvm/ADT/STLExtras.h"
19#include "llvm/CodeGen/LiveIntervals.h"
20#include "llvm/CodeGen/MachineRegisterInfo.h"
21#include "llvm/CodeGen/SlotIndexes.h"
22#include "llvm/CodeGen/TargetSchedule.h"
23
24using namespace llvm;
25using namespace llvm::AMDGPU;
26
27#define DEBUG_TYPE "amdgpu-anti-hints"
28
29namespace HC = llvm::AMDGPU::HazardClass;
30
31enum class AntiHintRule {
32 None,
33 MFMAWAW,
34 MFMAWAR,
35 All,
36};
37
38static cl::list<AntiHintRule> AntiHintRuleSelection(
39 "amdgpu-anti-hints-rules", cl::Hidden, cl::CommaSeparated,
40 cl::desc("Anti-hints rules to select."),
41 cl::values(clEnumValN(AntiHintRule::None, "none", "Select no rules"),
42 clEnumValN(AntiHintRule::MFMAWAW, "mfma-waw",
43 "MFMA destination write-after-write"),
44 clEnumValN(AntiHintRule::MFMAWAR, "mfma-war",
45 "XDL MFMA src2 write-after-read"),
46 clEnumValN(AntiHintRule::All, "all",
47 "Select all rules (default)")));
48
49namespace {
50
51// Classify the MI into a HazardClassMask.
52HazardClassMask getInstHazardClass(const MachineInstr &MI,
53 const HazardContext &Ctx) {
54 const SIInstrInfo &TII = *Ctx.TII;
55 HazardClassMask Mask = HC::None;
56
57 if (TII.isLDSDMA(MI))
58 Mask = HC::VALU | HC::VMEM | HC::DS;
59 else if (TII.isWMMA(MI) || SIInstrInfo::isSWMMAC(MI))
60 Mask = HC::WMMA;
61 else if (TII.isMFMA(MI))
62 Mask = HC::MFMA;
63 else if (SIInstrInfo::isTRANS(MI))
64 Mask = HC::TRANS;
65 else if (SIInstrInfo::isVALU(MI, /*AllowLDSDMA=*/true))
66 Mask = HC::VALU;
67 else if (TII.isDS(MI))
68 Mask = HC::DS;
69 else if (TII.isVMEM(MI))
70 Mask = HC::VMEM;
71 else if (TII.isSMRD(MI))
72 Mask = HC::SMEM;
73 else if (TII.isEXP(MI))
74 Mask = HC::EXP;
75 else if (SIInstrInfo::isSALU(MI))
76 Mask = HC::SALU;
77
78 return Mask;
79}
80
81bool isVirtualVGPR(const HazardContext &Ctx, const MachineOperand &MO) {
82 if (!MO.isReg() || !MO.getReg().isVirtual())
83 return false;
84 return Ctx.TRI->hasVGPRs(RC: Ctx.MRI->getRegClass(Reg: MO.getReg()));
85}
86
87void collectOperandRegs(const MachineInstr &MI, HazardOperand Op,
88 const HazardContext &Ctx,
89 SmallVectorImpl<Register> &Out) {
90 const SIInstrInfo &TII = *Ctx.TII;
91 switch (Op) {
92 case HazardOperand::None:
93 break;
94 case HazardOperand::Def:
95 for (const MachineOperand &MO : MI.all_defs()) {
96 if (isVirtualVGPR(Ctx, MO))
97 Out.push_back(Elt: MO.getReg());
98 }
99 break;
100 case HazardOperand::Src0:
101 if (const MachineOperand *MO =
102 TII.getNamedOperand(MI, OperandName: AMDGPU::OpName::src0)) {
103 if (isVirtualVGPR(Ctx, MO: *MO))
104 Out.push_back(Elt: MO->getReg());
105 }
106 break;
107 case HazardOperand::Src1:
108 if (const MachineOperand *MO =
109 TII.getNamedOperand(MI, OperandName: AMDGPU::OpName::src1)) {
110 if (isVirtualVGPR(Ctx, MO: *MO))
111 Out.push_back(Elt: MO->getReg());
112 }
113 break;
114 case HazardOperand::Src2:
115 if (const MachineOperand *MO =
116 TII.getNamedOperand(MI, OperandName: AMDGPU::OpName::src2)) {
117 if (isVirtualVGPR(Ctx, MO: *MO))
118 Out.push_back(Elt: MO->getReg());
119 }
120 break;
121 case HazardOperand::Idx:
122 if (const MachineOperand *MO =
123 TII.getNamedOperand(MI, OperandName: AMDGPU::OpName::idx)) {
124 if (isVirtualVGPR(Ctx, MO: *MO))
125 Out.push_back(Elt: MO->getReg());
126 }
127 break;
128 case HazardOperand::Vaddr:
129 if (const MachineOperand *MO =
130 TII.getNamedOperand(MI, OperandName: AMDGPU::OpName::vaddr)) {
131 if (isVirtualVGPR(Ctx, MO: *MO))
132 Out.push_back(Elt: MO->getReg());
133 }
134 break;
135 case HazardOperand::AnySrc:
136 for (AMDGPU::OpName Name :
137 {AMDGPU::OpName::src0, AMDGPU::OpName::src1, AMDGPU::OpName::src2}) {
138 if (const MachineOperand *MO = TII.getNamedOperand(MI, OperandName: Name)) {
139 if (isVirtualVGPR(Ctx, MO: *MO))
140 Out.push_back(Elt: MO->getReg());
141 }
142 }
143 break;
144 case HazardOperand::AnyUse:
145 for (const MachineOperand &MO : MI.all_uses()) {
146 if (isVirtualVGPR(Ctx, MO))
147 Out.push_back(Elt: MO.getReg());
148 }
149 break;
150 }
151}
152
153enum class MFMAHazardKind { RAW, WAW, WAR };
154
155// MFMA anti-hint wait-state window, mirroring GCNHazardRecognizer.cpp wait
156// states.
157unsigned mfmaWaitStates(const MachineInstr &MFMA, MFMAHazardKind Kind,
158 HazardClassMask ReaderClass, const HazardContext &Ctx) {
159 const SIInstrInfo &TII = *Ctx.TII;
160 const GCNSubtarget &ST = *Ctx.ST;
161 const int NumPasses = Ctx.SchedModel->computeInstrLatency(MI: &MFMA);
162 const bool IsDGEMM = SIInstrInfo::isDGEMM(Opcode: MFMA.getOpcode());
163 const bool Mem = ReaderClass & (HC::VMEM | HC::DS | HC::EXP);
164
165 auto GFX940NPass = [&]() -> unsigned {
166 return TII.isXDL(MI: MFMA)
167 ? NumPasses + 3 + (NumPasses != 2 && ST.hasGFX950Insts())
168 : NumPasses + 2;
169 };
170 auto SMFMANPass = [&]() -> unsigned {
171 switch (NumPasses) {
172 case 2:
173 return 5;
174 case 8:
175 return 11;
176 case 16:
177 return 19;
178 default:
179 return 0;
180 }
181 };
182
183 switch (Kind) {
184 case MFMAHazardKind::RAW:
185 if (IsDGEMM) {
186 switch (NumPasses) {
187 case 4:
188 return Mem ? 9 : 6;
189 case 8:
190 case 16:
191 return Mem ? 18 : (ST.hasGFX950Insts() ? 19 : 11);
192 default:
193 return 0;
194 }
195 }
196 return ST.hasGFX940Insts() ? GFX940NPass() : SMFMANPass();
197
198 case MFMAHazardKind::WAW:
199 if (IsDGEMM) {
200 switch (NumPasses) {
201 case 4:
202 return 6;
203 case 8:
204 case 16:
205 return 11;
206 default:
207 return 0;
208 }
209 }
210 return ST.hasGFX940Insts() ? GFX940NPass() : SMFMANPass();
211
212 case MFMAHazardKind::WAR:
213 switch (NumPasses) {
214 case 2:
215 return 1;
216 case 4:
217 return 3;
218 case 8:
219 return 7;
220 case 16:
221 return 15;
222 default:
223 return 15;
224 }
225 }
226 return 0;
227}
228
229unsigned mfmaWawWindow(const MachineInstr &P, const HazardContext &Ctx) {
230 return mfmaWaitStates(MFMA: P, Kind: MFMAHazardKind::WAW, ReaderClass: HC::None, Ctx);
231}
232unsigned mfmaWarWindow(const MachineInstr &P, const HazardContext &Ctx) {
233 return mfmaWaitStates(MFMA: P, Kind: MFMAHazardKind::WAR, ReaderClass: HC::None, Ctx);
234}
235
236unsigned mfmaReaderRawWindow(const MachineInstr &Producer,
237 HazardClassMask ReaderClass,
238 const HazardContext &Ctx) {
239 return mfmaWaitStates(MFMA: Producer, Kind: MFMAHazardKind::RAW, ReaderClass, Ctx);
240}
241
242bool hasMFMAHazard(const HazardContext &Ctx) {
243 return Ctx.ST->hasGFX90AInsts();
244}
245
246bool ruleSelected(AntiHintRule Rule) {
247 if (AntiHintRuleSelection.empty() ||
248 is_contained(Range&: AntiHintRuleSelection, Element: AntiHintRule::All))
249 return true;
250 return is_contained(Range&: AntiHintRuleSelection, Element: Rule);
251}
252
253bool isMFMAWAWRuleEnabled(const HazardContext &Ctx) {
254 return hasMFMAHazard(Ctx) && ruleSelected(Rule: AntiHintRule::MFMAWAW);
255}
256
257bool isMFMAWARRuleEnabled(const HazardContext &Ctx) {
258 return hasMFMAHazard(Ctx) && ruleSelected(Rule: AntiHintRule::MFMAWAR);
259}
260
261bool isXDLMFMA(const MachineInstr &MI, const HazardContext &Ctx) {
262 return Ctx.TII->isXDL(MI);
263}
264
265unsigned resolveWindow(const ConsumerTarget &CT, const MachineInstr &MI,
266 const HazardContext &Ctx) {
267 // Explicity given window length overrides the computed one.
268 if (CT.Window.OptWindowLength &&
269 CT.Window.OptWindowLength->getNumOccurrences())
270 return *CT.Window.OptWindowLength;
271 if (CT.Window.Fn)
272 return CT.Window.Fn(MI, Ctx);
273 return CT.Window.WindowLength;
274}
275
276// Build the anti-hints rules.
277class HazardRuleSet {
278 SmallVector<HazardAntiHintRule, 0> Rules;
279
280public:
281 class RuleBuilder {
282
283 HazardRuleSet &S;
284 unsigned Idx;
285 HazardAntiHintRule &rule() const { return S.Rules[Idx]; }
286
287 public:
288 RuleBuilder(HazardRuleSet &S, unsigned Idx) : S(S), Idx(Idx) {}
289
290 RuleBuilder &producer(ClassMatch M, HazardOperand Op,
291 InstPredicate Predicate = nullptr) {
292 rule().Producer = {.Match: M, .Op: Op, .Predicate: Predicate};
293 return *this;
294 }
295
296 RuleBuilder &rawCredit(AdvanceForRawWindowFn AdvanceForRawWindow) {
297 rule().AdvanceForRawWindow = AdvanceForRawWindow;
298 return *this;
299 }
300
301 RuleBuilder &consumer(ClassMatch M, HazardOperand Op, WindowSpec Window,
302 HazardClassMask CountMask = 0,
303 ConsumerHint Hint = ConsumerHint::OneDirectional,
304 InstPredicate Predicate = nullptr) {
305 rule().Consumers.push_back(Elt: {.Side: {.Match: M, .Op: Op, .Predicate: Predicate}, .Window: Window, .CounterMask: CountMask, .Hint: Hint});
306 return *this;
307 }
308
309 RuleBuilder &enabledIf(RulePredicate Predicate) {
310 rule().Predicate = Predicate;
311 return *this;
312 }
313 };
314
315 RuleBuilder addRule() {
316 Rules.emplace_back();
317 return RuleBuilder(*this, Rules.size() - 1);
318 }
319
320 SmallVector<HazardAntiHintRule, 0> buildRules() { return std::move(Rules); }
321};
322
323// Here the anti-hints rules are inserted.
324SmallVector<HazardAntiHintRule, 0> buildAntiHintsRules() {
325 using HO = HazardOperand;
326 HazardRuleSet S;
327
328 const WindowSpec MfmaWawWindow{.WindowLength: 0, .OptWindowLength: nullptr, .Fn: mfmaWawWindow};
329 const WindowSpec MfmaWarWindow{.WindowLength: 0, .OptWindowLength: nullptr, .Fn: mfmaWarWindow};
330 const ClassMatch MfmaConsumers = {.AnyOf: HC::DS | HC::VALU | HC::VMEM | HC::TRANS |
331 HC::EXP};
332
333 // MFMA WAW rules
334 S.addRule()
335 .enabledIf(Predicate: isMFMAWAWRuleEnabled)
336 .producer(M: {.AnyOf: HC::MFMA}, Op: HO::Def)
337 .rawCredit(AdvanceForRawWindow: mfmaReaderRawWindow)
338 .consumer(M: MfmaConsumers, Op: HO::Def, Window: MfmaWawWindow, CountMask: HC::None,
339 Hint: ConsumerHint::OneDirectional);
340
341 // MFMA WAR rules
342 S.addRule()
343 .enabledIf(Predicate: isMFMAWARRuleEnabled)
344 .producer(M: {.AnyOf: HC::MFMA}, Op: HO::Src2, Predicate: isXDLMFMA)
345 .rawCredit(AdvanceForRawWindow: mfmaReaderRawWindow)
346 .consumer(M: MfmaConsumers, Op: HO::Def, Window: MfmaWarWindow, CountMask: HC::None,
347 Hint: ConsumerHint::OneDirectional);
348
349 return S.buildRules();
350}
351
352ArrayRef<HazardAntiHintRule> getAntiHintsRules() {
353 static const SmallVector<HazardAntiHintRule, 0> Rules = buildAntiHintsRules();
354 return Rules;
355}
356
357struct AntiHintWindow {
358 SmallVector<Register, 4> Regs;
359 const MachineInstr *Producer = nullptr;
360 unsigned Len = 0;
361 unsigned Elapsed = 0;
362};
363
364using ConsumerTracking = SmallVector<AntiHintWindow, 3>;
365// One per consumer target of a rule.
366using RuleTracking = SmallVector<ConsumerTracking, 3>;
367
368class AntiHintEngine {
369 const HazardContext &Ctx;
370 ArrayRef<HazardAntiHintRule> Rules;
371
372 SmallVector<bool, 8> RuleApplies;
373 bool AnyEnabled = false;
374
375public:
376 AntiHintEngine(const HazardContext &Ctx)
377 : Ctx(Ctx), Rules(getAntiHintsRules()), RuleApplies(Rules.size()) {
378 for (unsigned R = 0; R < Rules.size(); ++R) {
379 const HazardAntiHintRule &Rule = Rules[R];
380 RuleApplies[R] = !Rule.Predicate || Rule.Predicate(Ctx);
381 AnyEnabled |= RuleApplies[R];
382 }
383 }
384
385 void run(MachineFunction &MF) {
386 if (!AnyEnabled)
387 return;
388
389 SmallVector<RuleTracking, 8> Tracking(Rules.size());
390 for (unsigned R = 0; R < Rules.size(); ++R)
391 Tracking[R].resize(N: Rules[R].Consumers.size());
392
393 for (const MachineBasicBlock &MBB : MF) {
394 for (RuleTracking &RT : Tracking) {
395 for (ConsumerTracking &Track : RT)
396 Track.clear();
397 }
398 for (const MachineInstr &MI : MBB) {
399 if (MI.isMetaInstruction())
400 continue;
401 const HazardClassMask C = getInstHazardClass(MI, Ctx);
402 // Wait states this instruction contributes to an open window.
403 unsigned WaitStates = SIInstrInfo::getNumWaitStates(MI);
404 addAntiHintsAndExpire(MI, C, WaitStates, Tracking);
405 addWindows(MI, C, Tracking);
406 }
407 }
408 }
409
410private:
411 // Check if the right class and predicate matches.
412 bool sideMatches(const HazardSide &Side, const MachineInstr &MI,
413 HazardClassMask C) {
414 return Side.Match.matches(Mask: C) &&
415 (!Side.Predicate || Side.Predicate(MI, Ctx));
416 }
417
418 // Producer phase: open a window for each consumer target.
419 void addWindows(const MachineInstr &MI, HazardClassMask C,
420 MutableArrayRef<RuleTracking> Tracking) {
421 for (unsigned R = 0; R < Rules.size(); ++R) {
422 const HazardAntiHintRule &Rule = Rules[R];
423 if (!RuleApplies[R] || !sideMatches(Side: Rule.Producer, MI, C))
424 continue;
425
426 SmallVector<Register, 4> Regs;
427 // Collect the producer regs.
428 collectOperandRegs(MI, Op: Rule.Producer.Op, Ctx, Out&: Regs);
429 if (Regs.empty())
430 continue;
431 for (unsigned ConsumerIdx = 0, E = Rule.Consumers.size();
432 ConsumerIdx != E; ++ConsumerIdx) {
433 unsigned Window = resolveWindow(CT: Rule.Consumers[ConsumerIdx], MI, Ctx);
434 if (!Window)
435 continue;
436 Tracking[R][ConsumerIdx].push_back(Elt: {.Regs: Regs, .Producer: &MI, .Len: Window, .Elapsed: 0});
437 }
438 }
439 }
440
441 void advanceByRawWindow(const MachineInstr &MI, HazardClassMask C,
442 const HazardAntiHintRule &Rule,
443 const ConsumerTarget &CT, ConsumerTracking &Track) {
444
445 if (!Rule.AdvanceForRawWindow || !CT.Side.Match.matches(Mask: C))
446 return;
447 for (AntiHintWindow &Window : Track) {
448 // Determine if def generated by producer is read by the MI consumer.
449 bool ReadsProducerDef = llvm::any_of(
450 Range: Window.Producer->all_defs(), P: [&](const MachineOperand &MO) {
451 return MO.getReg().isVirtual() &&
452 MI.readsVirtualRegister(Reg: MO.getReg());
453 });
454 // Advance by RAW window from the producer if that is larger than current
455 // Window.Elapsed.
456 if (ReadsProducerDef) {
457 Window.Elapsed = std::max(
458 a: Window.Elapsed, b: Rule.AdvanceForRawWindow(*Window.Producer, C, Ctx));
459 }
460 }
461 }
462
463 // Consumer phase: add anti-hints, then charge and expire open windows.
464 void addAntiHintsAndExpire(const MachineInstr &MI, HazardClassMask C,
465 unsigned WaitStates,
466 MutableArrayRef<RuleTracking> Tracking) {
467 for (unsigned R = 0; R < Rules.size(); ++R) {
468 const HazardAntiHintRule &Rule = Rules[R];
469 if (!RuleApplies[R])
470 continue;
471 for (unsigned ConsumerIdx = 0, E = Rule.Consumers.size();
472 ConsumerIdx != E; ++ConsumerIdx) {
473 const ConsumerTarget &CT = Rule.Consumers[ConsumerIdx];
474 ConsumerTracking &Track = Tracking[R][ConsumerIdx];
475 if (Track.empty())
476 continue;
477
478 // Before adding the anti-hints, see if advancing by RAW window will
479 // help remove the window.
480 advanceByRawWindow(MI, C, Rule, CT, Track);
481 llvm::erase_if(C&: Track, P: [](const AntiHintWindow &Window) {
482 return Window.Elapsed >= Window.Len;
483 });
484
485 // Add anti-hints if the consumer matches the instruction.
486 if (sideMatches(Side: CT.Side, MI, C))
487 addAntiHints(CT, MI, Track);
488
489 // This instruction's own wait states count toward the next one.
490 if (!CT.CounterMask || (C & CT.CounterMask)) {
491 for (AntiHintWindow &Window : Track)
492 Window.Elapsed += WaitStates;
493 }
494 }
495 }
496 }
497
498 bool isCopyOf(Register Cand, Register HazardReg) const {
499 const MachineInstr *Def = Ctx.MRI->getUniqueVRegDef(Reg: Cand);
500 return Def && Def->isCopy() && Def->getOperand(i: 1).getReg() == HazardReg;
501 }
502
503 void addAntiHints(const ConsumerTarget &CT, const MachineInstr &MI,
504 const ConsumerTracking &Track) {
505 if (Track.empty())
506 return;
507 SmallVector<Register, 4> ConsumerRegs;
508 // Collect the consumer regs.
509 collectOperandRegs(MI, Op: CT.Side.Op, Ctx, Out&: ConsumerRegs);
510 if (ConsumerRegs.empty())
511 return;
512 SlotIndex Slot = Ctx.LIS->getInstructionIndex(Instr: MI).getRegSlot();
513 auto AntiHint = [&](Register ProducerReg) {
514 if (!Ctx.LIS->hasInterval(Reg: ProducerReg))
515 return;
516 const LiveInterval &ProducerLI = Ctx.LIS->getInterval(Reg: ProducerReg);
517 // Skip a live producer reg.
518 if (ProducerLI.liveAt(index: Slot))
519 return;
520 for (Register ConsumerReg : ConsumerRegs) {
521 if (ConsumerReg == ProducerReg || isCopyOf(Cand: ConsumerReg, HazardReg: ProducerReg))
522 continue;
523
524 Ctx.MRI->addRegAllocationAntiHints(VReg: ConsumerReg, AntiHintVRegs: ProducerReg);
525 if (CT.Hint == ConsumerHint::Symmetric)
526 Ctx.MRI->addRegAllocationAntiHints(VReg: ProducerReg, AntiHintVRegs: ConsumerReg);
527 LLVM_DEBUG(
528 dbgs() << "anti-hint: keep " << printReg(ProducerReg, Ctx.TRI)
529 << (CT.Hint == ConsumerHint::Symmetric ? " <-> " : " <- ")
530 << printReg(ConsumerReg, Ctx.TRI) << " (consumer "
531 << Ctx.TII->getName(MI.getOpcode()) << ")\n");
532 }
533 };
534 for (const AntiHintWindow &Window : Track) {
535 for (Register ProducerReg : Window.Regs)
536 AntiHint(ProducerReg);
537 }
538 }
539};
540
541} // namespace
542
543void AMDGPU::applyAntiHintRules(MachineFunction &MF, const HazardContext &Ctx) {
544 AntiHintEngine(Ctx).run(MF);
545}
546