44#define DEBUG_TYPE "amdgpu-rewrite-agpr-copy-mfma"
47 "Controls which MFMA chains are rewritten to AGPR form");
52 "Number of MFMA instructions rewritten to use AGPR form");
57class AMDGPURewriteAGPRCopyMFMAImpl {
80 TRI(*ST.getRegisterInfo()), MRI(MF.getRegInfo()), VRM(VRM), LRM(LRM),
81 LIS(LIS), LSS(LSS), RegClassInfo(RegClassInfo), MDT(MDT) {}
83 bool isRewriteCandidate(
const MachineInstr &
MI)
const {
92 MCRegister getAssignedAGPR(
Register VReg)
const {
93 MCRegister PhysReg = VRM.getPhys(VReg);
100 return TRI.isAGPRClass(AssignedRC) ? PhysReg : MCRegister();
103 bool tryReassigningMFMAChain(MachineInstr &
MFMA,
Register MFMAHintReg,
116 bool recomputeRegClassExceptRewritable(
117 Register Reg, SmallVectorImpl<MachineInstr *> &RewriteCandidates,
118 SmallSetVector<Register, 4> &RewriteRegs)
const;
120 bool tryFoldCopiesToAGPR(
Register VReg, MCRegister AssignedAGPR)
const;
121 bool tryFoldCopiesFromAGPR(
Register VReg, MCRegister AssignedAGPR)
const;
125 void replaceSpillWithCopyToVReg(MachineInstr &SpillMI,
int SpillFI,
132 SpillReferenceMap &Map)
const;
139 bool isLoadJointlyDominatedByStores(
140 const MachineInstr &LoadMI,
const LiveInterval &SlotLI,
141 const SmallPtrSetImpl<MachineBasicBlock *> &StoreFreeReachable)
const;
145 void eliminateSpillsOfReassignedVGPRs()
const;
150bool AMDGPURewriteAGPRCopyMFMAImpl::recomputeRegClassExceptRewritable(
156 while (!Worklist.
empty()) {
168 MachineInstr *
MI = MO.getParent();
174 if (isRewriteCandidate(*
MI)) {
176 const MCInstrDesc &AGPRDesc =
TII.get(AGPROp);
178 TII.getRegClass(AGPRDesc, MO.getOperandNo());
179 if (!
TRI.hasAGPRs(NewRC))
182 const MachineOperand *VDst =
183 TII.getNamedOperand(*
MI, AMDGPU::OpName::vdst);
184 const MachineOperand *Src2 =
185 TII.getNamedOperand(*
MI, AMDGPU::OpName::src2);
186 for (
const MachineOperand *
Op : {VDst, Src2}) {
194 if (OtherReg !=
Reg && RewriteRegs.
insert(OtherReg))
201 dbgs() <<
"Attempting to replace VGPR MFMA with AGPR version:"
220 unsigned OpNo = &MO - &
MI->getOperand(0);
221 NewRC =
MI->getRegClassConstraintEffect(OpNo, NewRC, &
TII, &
TRI);
222 if (!NewRC || NewRC == OldRC) {
224 <<
" cannot be reassigned to "
225 << (NewRC ?
TRI.getRegClassName(NewRC) :
"NULL")
235bool AMDGPURewriteAGPRCopyMFMAImpl::tryReassigningMFMAChain(
239 SmallVector<MachineInstr *, 4> RewriteCandidates = {&
MFMA};
240 SmallSetVector<Register, 4> RewriteRegs;
244 RewriteRegs.
insert(MFMAHintReg);
255 if (!recomputeRegClassExceptRewritable(MFMAHintReg, RewriteCandidates,
257 LLVM_DEBUG(
dbgs() <<
"Could not recompute the regclass of dst reg "
278 using RecoloringStack =
280 RecoloringStack TentativeReassignments;
282 for (
Register RewriteReg : RewriteRegs) {
284 TentativeReassignments.push_back({&LI, VRM.
getPhys(RewriteReg)});
289 !attemptReassignmentsToAGPR(RewriteRegs, PhysRegHint)) {
291 for (
auto [LI, OldAssign] : TentativeReassignments) {
294 LRM.
assign(*LI, OldAssign);
302 for (
Register InterferingReg : RewriteRegs) {
305 MRI.
setRegClass(InterferingReg, EquivalentAGPRRegClass);
308 for (MachineInstr *RewriteCandidate : RewriteCandidates) {
310 RewriteCandidate->setDesc(
TII.get(NewMFMAOp));
311 ++NumMFMAsRewrittenToAGPR;
320bool AMDGPURewriteAGPRCopyMFMAImpl::attemptReassignmentsToAGPR(
321 SmallSetVector<Register, 4> &InterferingRegs,
MCPhysReg PrefPhysReg)
const {
326 for (
Register InterferingReg : InterferingRegs) {
327 LiveInterval &ReassignLI = LIS.
getInterval(InterferingReg);
331 MCPhysReg Assignable = AMDGPU::NoRegister;
332 if (EquivalentAGPRRegClass->
contains(PrefPhysReg) &&
342 Assignable = PrefPhysReg;
345 RegClassInfo.
getOrder(EquivalentAGPRRegClass);
357 <<
" to a free AGPR\n");
363 LRM.
assign(ReassignLI, Assignable);
375bool AMDGPURewriteAGPRCopyMFMAImpl::tryFoldCopiesToAGPR(
376 Register VReg, MCRegister AssignedAGPR)
const {
377 bool MadeChange =
false;
400 if (isRewriteCandidate(CopySrcDefMI) &&
401 tryReassigningMFMAChain(
402 CopySrcDefMI, CopySrcDefMI.getOperand(0).getReg(), AssignedAGPR))
417bool AMDGPURewriteAGPRCopyMFMAImpl::tryFoldCopiesFromAGPR(
418 Register VReg, MCRegister AssignedAGPR)
const {
419 bool MadeChange =
false;
428 if (!CopyUseMO.readsReg())
431 MachineInstr &CopyUseMI = *CopyUseMO.getParent();
432 if (isRewriteCandidate(CopyUseMI)) {
433 if (tryReassigningMFMAChain(CopyUseMI, CopyDstReg,
443void AMDGPURewriteAGPRCopyMFMAImpl::replaceSpillWithCopyToVReg(
444 MachineInstr &SpillMI,
int SpillFI,
Register VReg)
const {
447 MachineInstr *NewCopy;
461void AMDGPURewriteAGPRCopyMFMAImpl::collectSpillIndexUses(
464 SmallSet<int, 4> NeededFrameIndexes;
465 for (
const LiveInterval *LI : StackIntervals)
468 for (MachineBasicBlock &
MBB : MF) {
469 for (MachineInstr &
MI :
MBB) {
470 for (MachineOperand &MO :
MI.operands()) {
471 if (!MO.isFI() || !NeededFrameIndexes.
count(MO.getIndex()))
474 if (
TII.isVGPRSpill(
MI)) {
475 SmallVector<MachineInstr *, 4> &References =
Map[MO.getIndex()];
484 NeededFrameIndexes.
erase(MO.getIndex());
485 Map.erase(MO.getIndex());
491bool AMDGPURewriteAGPRCopyMFMAImpl::isLoadJointlyDominatedByStores(
492 const MachineInstr &LoadMI,
const LiveInterval &SlotLI,
493 const SmallPtrSetImpl<MachineBasicBlock *> &StoreFreeReachable)
const {
494 const MachineBasicBlock *LoadMBB = LoadMI.
getParent();
499 if (!StoreFreeReachable.
contains(LoadMBB))
512void AMDGPURewriteAGPRCopyMFMAImpl::eliminateSpillsOfReassignedVGPRs()
const {
517 MachineFrameInfo &MFI = MF.getFrameInfo();
520 StackIntervals.
reserve(NumSlots);
522 for (
auto &[Slot, LI] : LSS) {
527 if (
TRI.hasVGPRs(RC))
531 sort(StackIntervals, [](
const LiveInterval *
A,
const LiveInterval *
B) {
534 if (
A->weight() !=
B->weight())
535 return A->weight() >
B->weight();
537 if (
A->getSize() !=
B->getSize())
538 return A->getSize() >
B->getSize();
541 return A->reg().stackSlotIndex() <
B->reg().stackSlotIndex();
556 DenseMap<int, SmallVector<MachineInstr *, 4>> SpillSlotReferences;
557 collectSpillIndexUses(StackIntervals, SpillSlotReferences);
559 for (LiveInterval *LI : StackIntervals) {
561 auto SpillReferences = SpillSlotReferences.find(Slot);
562 if (SpillReferences == SpillSlotReferences.end())
567 SmallPtrSet<MachineBasicBlock *, 4> StoreBlocks;
568 for (MachineInstr *
MI : SpillReferences->second) {
570 StoreBlocks.
insert(
MI->getParent());
573 if (StoreBlocks.
empty()) {
575 <<
": no reachable stores\n");
581 MachineBasicBlock &EntryMBB = MF.front();
582 SmallPtrSet<MachineBasicBlock *, 16> StoreFreeReachable = {&EntryMBB};
585 while (!Worklist.
empty()) {
591 if (StoreFreeReachable.
insert(Succ).second)
597 if (!
llvm::all_of(SpillReferences->second, [&](
const MachineInstr *
MI) {
598 return !MI->mayLoad() ||
599 isLoadJointlyDominatedByStores(*MI, *LI, StoreFreeReachable);
603 <<
": some reachable load not jointly dominated by stores\n");
610 <<
" by reassigning\n");
624 for (MachineInstr *SpillMI : SpillReferences->second)
625 replaceSpillWithCopyToVReg(*SpillMI, Slot, NewVReg);
640 if (!SplitLIs.
empty()) {
641 dbgs() <<
"Split unspilled interval into " << (SplitLIs.
size() + 1)
646 LRM.
assign(NewLI, PhysReg);
647 for (LiveInterval *SplitLI : SplitLIs) {
649 LRM.
assign(*SplitLI, PhysReg);
661 if (!
ST.hasGFX90AInsts())
666 LLVM_DEBUG(
dbgs() <<
"skipping function that did not allocate AGPRs\n");
670 bool MadeChange =
false;
673 Register VReg = Register::index2VirtReg(
I);
674 MCRegister AssignedAGPR = getAssignedAGPR(VReg);
678 if (tryFoldCopiesToAGPR(VReg, AssignedAGPR))
680 if (tryFoldCopiesFromAGPR(VReg, AssignedAGPR))
688 eliminateSpillsOfReassignedVGPRs();
693class AMDGPURewriteAGPRCopyMFMALegacy :
public MachineFunctionPass {
697 AMDGPURewriteAGPRCopyMFMALegacy() : MachineFunctionPass(
ID) {}
701 StringRef getPassName()
const override {
702 return "AMDGPU Rewrite AGPR-Copy-MFMA";
705 void getAnalysisUsage(AnalysisUsage &AU)
const override {
710 AU.
addRequired<MachineRegisterClassInfoWrapperPass>();
728 "AMDGPU Rewrite AGPR-Copy-MFMA",
false,
false)
738char AMDGPURewriteAGPRCopyMFMALegacy::ID = 0;
741 AMDGPURewriteAGPRCopyMFMALegacy::ID;
743bool AMDGPURewriteAGPRCopyMFMALegacy::runOnMachineFunction(
745 if (skipFunction(MF.getFunction()))
748 auto &VRM = getAnalysis<VirtRegMapWrapperLegacy>().getVRM();
749 auto &LRM = getAnalysis<LiveRegMatrixWrapperLegacy>().getLRM();
750 auto &LIS = getAnalysis<LiveIntervalsWrapperPass>().getLIS();
751 auto &LSS = getAnalysis<LiveStacksWrapperLegacy>().getLS();
752 auto &RCI = getAnalysis<MachineRegisterClassInfoWrapperPass>().getRCI();
753 auto &MDT = getAnalysis<MachineDominatorTreeWrapperPass>().getDomTree();
754 AMDGPURewriteAGPRCopyMFMAImpl Impl(MF, VRM, LRM, LIS, LSS, RCI, MDT);
768 AMDGPURewriteAGPRCopyMFMAImpl Impl(MF, VRM, LRM, LIS, LSS, RCI, MDT);
773 .preserve<LiveStacksAnalysis>()
775 .preserve<SlotIndexesAnalysis>()
777 .preserve<LiveRegMatrixAnalysis>()
MachineInstrBuilder & UseMI
AMDGPU Rewrite AGPR Copy MFMA
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
static GCRegistry::Add< ErlangGC > A("erlang", "erlang-compatible garbage collector")
static GCRegistry::Add< CoreCLRGC > E("coreclr", "CoreCLR-compatible GC")
static GCRegistry::Add< OcamlGC > B("ocaml", "ocaml 3.10-compatible GC")
This file provides an implementation of debug counters.
#define DEBUG_COUNTER(VARNAME, COUNTERNAME, DESC)
AMD GCN specific subclass of TargetSubtarget.
const HexagonInstrInfo * TII
Register const TargetRegisterInfo * TRI
Promote Memory to Register
#define INITIALIZE_PASS_DEPENDENCY(depName)
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
Interface definition for SIRegisterInfo.
This file defines the 'Statistic' class, which is designed to be an easy way to expose various metric...
#define STATISTIC(VARNAME, DESC)
PreservedAnalyses run(MachineFunction &MF, MachineFunctionAnalysisManager &MFAM)
PassT::Result & getResult(IRUnitT &IR, ExtraArgTs... ExtraArgs)
Get the result of an analysis pass for a given IR unit.
AnalysisUsage & addRequired()
AnalysisUsage & addPreserved()
Add the specified Pass class to the set of analyses preserved by this pass.
void setPreservesAll()
Set by analyses that do not transform their input at all.
Represents analyses that only rely on functions' control flow.
static bool shouldExecute(CounterInfo &Counter)
bool isReachableFromEntry(const NodeT *A) const
isReachableFromEntry - Return true if A is dominated by the entry block of the function containing it...
SlotIndex getInstructionIndex(const MachineInstr &Instr) const
Returns the base index of the given instruction.
LiveInterval & getInterval(Register Reg)
LLVM_ABI void splitSeparateComponents(LiveInterval &LI, SmallVectorImpl< LiveInterval * > &SplitLIs)
Split separate components in LiveInterval LI into separate intervals.
bool isLiveInToMBB(const LiveRange &LR, const MachineBasicBlock *mbb) const
LiveInterval & createAndComputeVirtRegInterval(Register Reg)
SlotIndex ReplaceMachineInstrInMaps(MachineInstr &MI, MachineInstr &NewMI)
bool liveAt(SlotIndex index) const
LLVM_ABI bool isPhysRegUsed(MCRegister PhysReg) const
Returns true if the given PhysReg has any live intervals assigned.
LLVM_ABI void unassign(const LiveInterval &VirtReg, bool ClearAllReferencingSegments=false)
Unassign VirtReg from its PhysReg.
@ IK_Free
No interference, go ahead and assign.
LLVM_ABI void assign(const LiveInterval &VirtReg, MCRegister PhysReg)
Assign VirtReg to PhysReg.
LLVM_ABI InterferenceKind checkInterference(const LiveInterval &VirtReg, MCRegister PhysReg)
Check for interference before assigning VirtReg to PhysReg.
unsigned getNumIntervals() const
bool contains(MCRegister Reg) const
contains - Return true if the specified register is included in this register class.
iterator_range< succ_iterator > successors()
Analysis pass which computes a MachineDominatorTree.
Analysis pass which computes a MachineDominatorTree.
DominatorTree Class - Concrete subclass of DominatorTreeBase that is used to compute a normal dominat...
bool isSpillSlotObjectIndex(int ObjectIdx) const
Returns true if the specified index corresponds to a spill slot.
void RemoveStackObject(int ObjectIdx)
Remove or mark dead a statically sized stack object.
bool isDeadObjectIndex(int ObjectIdx) const
Returns true if the specified index corresponds to a dead object.
void getAnalysisUsage(AnalysisUsage &AU) const override
getAnalysisUsage - Subclasses that override getAnalysisUsage must call this.
Register getReg(unsigned Idx) const
Get the register for the operand index.
const MachineInstrBuilder & addReg(Register RegNo, RegState Flags={}, unsigned SubReg=0) const
Add a new virtual register operand.
const MachineInstrBuilder & add(const MachineOperand &MO) const
const MachineBasicBlock * getParent() const
bool mayStore(QueryType Type=AnyInBundle) const
Return true if this instruction could possibly modify memory.
const DebugLoc & getDebugLoc() const
Returns the debug location id of this MachineInstr.
const MachineOperand & getOperand(unsigned i) const
LLVM_ABI MachineInstrBundleIterator< MachineInstr > eraseFromParent()
Unlink 'this' from the containing basic block and delete it.
bool isReg() const
isReg - Tests if this is a MO_Register operand.
Register getReg() const
getReg - Returns the register number.
MachineRegisterInfo - Keep track of information for virtual and physical registers,...
const TargetRegisterClass * getRegClass(Register Reg) const
Return the register class of the specified virtual register.
iterator_range< def_instr_iterator > def_instructions(Register Reg) const
LLVM_ABI Register createVirtualRegister(const TargetRegisterClass *RegClass, StringRef Name="")
createVirtualRegister - Create and return a new virtual register in the function with the specified r...
LLVM_ABI void setRegClass(Register Reg, const TargetRegisterClass *RC)
setRegClass - Set the register class of the specified virtual register.
iterator_range< use_instr_iterator > use_instructions(Register Reg) const
iterator_range< reg_nodbg_iterator > reg_nodbg_operands(Register Reg) const
unsigned getNumVirtRegs() const
getNumVirtRegs - Return the number of virtual registers created.
A set of analyses that are preserved following a run of a transformation pass.
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
ArrayRef< MCPhysReg > getOrder(const TargetRegisterClass *RC) const
getOrder - Returns the preferred allocation order for RC.
Wrapper class representing virtual and physical registers.
int stackSlotIndex() const
Compute the frame index from a register value representing a stack slot.
constexpr bool isVirtual() const
Return true if the specified register number is in the virtual register namespace.
constexpr bool isPhysical() const
Return true if the specified register number is in the physical register namespace.
bool insert(const value_type &X)
Insert a new element into the SetVector.
std::pair< iterator, bool > insert(PtrType Ptr)
Inserts Ptr if and only if there is no element in the container equal to Ptr.
bool contains(ConstPtrType Ptr) const
A SetVector that performs no allocations if smaller than a certain size.
size_type count(const T &V) const
count - Return 1 if the element is in the set, 0 otherwise.
std::pair< const_iterator, bool > insert(const T &V)
insert - Insert an element into the set if it isn't already there.
This class consists of common code factored out of the SmallVector class to reduce code duplication b...
void reserve(size_type N)
void push_back(const T &Elt)
MCRegister getPhys(Register virtReg) const
returns the physical register mapped to the specified virtual register
bool hasPhys(Register virtReg) const
returns true if the specified virtual register is mapped to a physical register
LLVM_READONLY int32_t getAGPRFormOp(uint32_t Opcode)
PointerTypeMap run(const Module &M)
Compute the PointerTypeMap for the module M.
This is an optimization pass for GlobalISel generic memory operations.
bool all_of(R &&range, UnaryPredicate P)
Provide wrappers to std::all_of which take ranges instead of having to pass begin/end explicitly.
MachineInstrBuilder BuildMI(MachineFunction &MF, const MIMetadata &MIMD, const MCInstrDesc &MCID)
Builder interface. Specify how to create the initial instruction itself.
AnalysisManager< MachineFunction > MachineFunctionAnalysisManager
LLVM_ABI PreservedAnalyses getMachineFunctionPassPreservedAnalyses()
Returns the minimum set of Analyses that all machine function passes must preserve.
void sort(IteratorTy Start, IteratorTy End)
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
uint16_t MCPhysReg
An unsigned integer type large enough to represent all physical registers, but not necessarily virtua...
DWARFExpression::Operation Op
ArrayRef(const T &OneElt) -> ArrayRef< T >
bool is_contained(R &&Range, const E &Element)
Returns true if Element is found in Range.
char & AMDGPURewriteAGPRCopyMFMALegacyID
LLVM_ABI Printable printReg(Register Reg, const TargetRegisterInfo *TRI=nullptr, unsigned SubIdx=0, const MachineRegisterInfo *MRI=nullptr)
Prints virtual and physical registers with or without a TRI instance.
MCRegisterClass TargetRegisterClass