[llvm] [AMDGPU] Per-chain MFMA->AGPR conversion (PR #217328)
Matt Arsenault via llvm-commits
llvm-commits at lists.llvm.org
Thu Aug 20 03:34:15 PDT 2026
================
@@ -1399,20 +1457,113 @@ bool RewriteMFMAFormStage::initGCNSchedStage() {
TII = ST.getInstrInfo();
SRI = ST.getRegisterInfo();
- std::vector<std::pair<MachineInstr *, unsigned>> RewriteCands;
- DenseMap<MachineBasicBlock *, std::set<Register>> CopyForUse;
- SmallPtrSet<MachineInstr *, 8> CopyForDef;
+ // Collect all convertible MFMAs.
+ SmallVector<MachineInstr *, 32> AllCands;
+ for (MachineBasicBlock &MBB : MF)
+ for (MachineInstr &MI : MBB)
+ if (isRewriteCandidate(&MI))
+ AllCands.push_back(&MI);
+
+ if (AllCands.empty())
+ return false;
+
+ // Identify accumulator chains and sort by size descending. We operate at
+ // chain granularity rather than individual MFMAs because:
+ // 1. It avoids converting MFMAs from different chains that would each need
+ // separate src2 and dst bridge copies (up to 4 copies for 2 MFMAs from
+ // 2 chains) while only reducing VGPR pressure for 2 instructions.
+ // 2. It reduces the search space from O(N_mfmas) to O(N_chains).
+ //
+ // Longer chains benefit more: each chain requires at most one src2 bridge
+ // copy and one dst bridge copy regardless of length, but reduces VGPR
+ // pressure proportionally to the number of chain members.
+ SmallVector<SmallVector<unsigned, 8>, 16> Chains =
+ identifyAccChains(AllCands, TII);
+ llvm::sort(Chains,
+ [](const auto &A, const auto &B) { return A.size() > B.size(); });
+
+ // Evaluate the cost of converting the first N sorted chains to AGPR form.
+ auto EvaluateProbe = [&](int N) -> int64_t {
+ SmallPtrSet<MachineInstr *, 32> ProbeFilter;
+ for (int I = 0; I < N; ++I)
+ for (unsigned Idx : Chains[I])
+ ProbeFilter.insert(AllCands[Idx]);
+
+ std::vector<std::pair<MachineInstr *, unsigned>> RC;
+ DenseMap<MachineBasicBlock *, std::set<Register>> CU;
+ SmallPtrSet<MachineInstr *, 8> CD;
+ Src2NeedsVGPRCache.clear();
+
+ if (!initHeuristics(RC, CU, CD, ProbeFilter))
+ return std::numeric_limits<int64_t>::max();
+
+ LLVM_DEBUG(dbgs() << "RewriteMFMA probe N=" << N << ":\n");
+ return getRewriteCost(RC, CU, CD);
+ };
+
+ // Start by evaluating all chains.
+ int64_t AllCost = EvaluateProbe(Chains.size());
+
+ LLVM_DEBUG(dbgs() << "RewriteMFMA probe: N=" << Chains.size()
+ << " Cost=" << AllCost << "\n");
+
+ int BestN = AllCost <= 0 ? Chains.size() : 0;
+
+ if (EnableMFMAChainSplitting && Chains.size() > 1) {
+ // Binary search for a good number of chains to convert. Chains are
----------------
arsenm wrote:
Split to helper function?
https://github.com/llvm/llvm-project/pull/217328
More information about the llvm-commits
mailing list