================ @@ -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
