================
@@ -0,0 +1,510 @@
+//===-- GCNPreRAAntiHints.cpp - MFMA register anti-hints 
------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM 
Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+/// \file
+/// Insert register allocation anti-hints.
+///
+//===----------------------------------------------------------------------===//
+
+#include "GCNPreRAAntiHints.h"
+#include "GCNSubtarget.h"
+#include "SIInstrInfo.h"
+#include "SIRegisterInfo.h"
+#include "llvm/CodeGen/LiveIntervals.h"
+#include "llvm/CodeGen/MachineRegisterInfo.h"
+#include "llvm/CodeGen/SlotIndexes.h"
+#include "llvm/CodeGen/TargetSchedule.h"
+
+using namespace llvm;
+using namespace llvm::AMDGPU;
+
+#define DEBUG_TYPE "amdgpu-anti-hints"
+
+namespace HC = llvm::AMDGPU::HazardClass;
+
+static cl::opt<std::string> AntiHintRuleSelection(
+    "amdgpu-anti-hints-rules", cl::Hidden,
+    cl::desc("Comma-separated anti-hints rules (waw, war), or all or none."),
+    cl::init("all"));
+
+namespace {
+
+// Classify the MI into a HazardClassMask.
+HazardClassMask getInstHazardClass(const MachineInstr &MI,
+                                   const HazardContext &Ctx) {
+  const SIInstrInfo &TII = *Ctx.TII;
+  HazardClassMask Mask = HC::None;
+
+  if (TII.isLDSDMA(MI))
+    Mask = HC::VALU | HC::VMEM | HC::DS;
+  else if (TII.isWMMA(MI) || SIInstrInfo::isSWMMAC(MI))
+    Mask = HC::WMMA;
+  else if (TII.isMFMA(MI))
+    Mask = HC::MFMA;
+  else if (SIInstrInfo::isTRANS(MI))
+    Mask = HC::TRANS;
+  else if (SIInstrInfo::isVALU(MI, /*AllowLDSDMA=*/true))
+    Mask = HC::VALU;
+  else if (TII.isDS(MI))
+    Mask = HC::DS;
+  else if (TII.isVMEM(MI))
+    Mask = HC::VMEM;
+  else if (TII.isSMRD(MI))
+    Mask = HC::SMEM;
+  else if (TII.isEXP(MI))
+    Mask = HC::EXP;
+  else if (SIInstrInfo::isSALU(MI))
+    Mask = HC::SALU;
+
+  return Mask;
+}
+
+void collectOperandRegs(const MachineInstr &MI, HazardOperand Op,
+                        const HazardContext &Ctx,
+                        SmallVectorImpl<Register> &Out) {
+  const SIInstrInfo &TII = *Ctx.TII;
+  auto Add = [&](const MachineOperand *MO) {
+    if (MO && MO->isReg() && MO->getReg().isVirtual() &&
+        Ctx.TRI->hasVGPRs(Ctx.MRI->getRegClass(MO->getReg())))
+      Out.push_back(MO->getReg());
+  };
+  auto Named = [&](AMDGPU::OpName N) { Add(TII.getNamedOperand(MI, N)); };
+  switch (Op) {
+  case HazardOperand::None:
+    break;
+  case HazardOperand::Def:
+    for (const MachineOperand &MO : MI.operands())
+      if (MO.isReg() && MO.isDef())
+        Add(&MO);
+    break;
+  case HazardOperand::Src0:
+    Named(AMDGPU::OpName::src0);
+    break;
+  case HazardOperand::Src1:
+    Named(AMDGPU::OpName::src1);
+    break;
+  case HazardOperand::Src2:
+    Named(AMDGPU::OpName::src2);
+    break;
+  case HazardOperand::Idx:
+    Named(AMDGPU::OpName::idx);
+    break;
+  case HazardOperand::Vaddr:
+    Named(AMDGPU::OpName::vaddr);
+    break;
+  case HazardOperand::AnySrc:
+    Named(AMDGPU::OpName::src0);
+    Named(AMDGPU::OpName::src1);
+    Named(AMDGPU::OpName::src2);
+    break;
+  case HazardOperand::AnyUse:
+    for (const MachineOperand &MO : MI.uses())
+      if (MO.isReg() && MO.isUse())
+        Add(&MO);
+    break;
+  }
+}
+
+enum class MFMAHazardKind { RAW, WAW, WAR };
+
+// MFMA anti-hint wait-state window, mirroring GCNHazardRecognizer.cpp wait
+// states.
+unsigned mfmaWaitStates(const MachineInstr &MFMA, MFMAHazardKind Kind,
+                        HazardClassMask ReaderClass, const HazardContext &Ctx) 
{
+  const SIInstrInfo &TII = *Ctx.TII;
+  const GCNSubtarget &ST = *Ctx.ST;
+  const int NumPasses = Ctx.SchedModel->computeInstrLatency(&MFMA);
+  const bool IsDGEMM = SIInstrInfo::isDGEMM(MFMA.getOpcode());
+  const bool Mem = ReaderClass & (HC::VMEM | HC::DS | HC::EXP);
+
+  auto GFX940NPass = [&]() -> unsigned {
+    return TII.isXDL(MFMA)
+               ? NumPasses + 3 + (NumPasses != 2 && ST.hasGFX950Insts())
+               : NumPasses + 2;
+  };
+  auto SMFMANPass = [&]() -> unsigned {
+    switch (NumPasses) {
+    case 2:
+      return 5;
+    case 8:
+      return 11;
+    case 16:
+      return 19;
+    default:
+      return 0;
+    }
+  };
+
+  switch (Kind) {
+  case MFMAHazardKind::RAW:
+    if (IsDGEMM) {
+      switch (NumPasses) {
+      case 4:
+        return Mem ? 9 : 6;
+      case 8:
+      case 16:
+        return Mem ? 18 : (ST.hasGFX950Insts() ? 19 : 11);
+      default:
+        return 0;
+      }
+    }
+    return ST.hasGFX940Insts() ? GFX940NPass() : SMFMANPass();
+
+  case MFMAHazardKind::WAW:
+    if (IsDGEMM) {
+      switch (NumPasses) {
+      case 4:
+        return 6;
+      case 8:
+      case 16:
+        return 11;
+      default:
+        return 0;
+      }
+    }
+    return ST.hasGFX940Insts() ? GFX940NPass() : SMFMANPass();
+
+  case MFMAHazardKind::WAR:
+    switch (NumPasses) {
+    case 2:
+      return 1;
+    case 4:
+      return 3;
+    case 8:
+      return 7;
+    case 16:
+      return 15;
+    default:
+      return 15;
+    }
+  }
+  return 0;
+}
+
+unsigned mfmaWawWindow(const MachineInstr &P, const HazardContext &Ctx) {
+  return mfmaWaitStates(P, MFMAHazardKind::WAW, HC::None, Ctx);
+}
+unsigned mfmaWarWindow(const MachineInstr &P, const HazardContext &Ctx) {
+  return mfmaWaitStates(P, MFMAHazardKind::WAR, HC::None, Ctx);
+}
+
+unsigned mfmaReaderRawWindow(const MachineInstr &Producer,
+                             HazardClassMask ReaderClass,
+                             const HazardContext &Ctx) {
+  return mfmaWaitStates(Producer, MFMAHazardKind::RAW, ReaderClass, Ctx);
+}
+
+bool hasMFMAHazard(const HazardContext &Ctx) {
+  return Ctx.ST->hasGFX90AInsts();
+}
+
+bool ruleSelected(StringRef Name) {
+  StringRef Selection(AntiHintRuleSelection);
+  if (Selection.equals_insensitive("all"))
+    return true;
+  if (Selection.equals_insensitive("none"))
+    return false;
+  SmallVector<StringRef, 3> Selected;
+  Selection.split(Selected, ',', /*MaxSplit=*/-1, /*KeepEmpty=*/false);
+  return llvm::any_of(Selected, [Name](StringRef S) {
+    return S.trim().equals_insensitive(Name);
+  });
+}
----------------
qcolombet wrote:

Instead of doing the parsing, etc., you should be able to use `cl::bits` or 
something like that.
Look at the full command line doc if you need details 
`llvm/docs/CommandLine.rst`.

https://github.com/llvm/llvm-project/pull/218075
_______________________________________________
llvm-branch-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/llvm-branch-commits

Reply via email to