[llvm] [AMDGPU] Per-chain MFMA->AGPR conversion (PR #217328)
Romanov Vlad via llvm-commits
llvm-commits at lists.llvm.org
Mon Sep 7 08:48:32 PDT 2026
================
@@ -1377,6 +1379,122 @@ void RewriteMFMAFormStage::findReachingUses(
}
}
+/// Identify accumulator chains: groups of MFMAs connected via dst->src2.
+static SmallVector<SmallVector<unsigned, 8>, 16>
+identifyAccChains(ArrayRef<MachineInstr *> Cands, const SIInstrInfo *TII) {
+ DenseMap<Register, unsigned> DstToCand;
+ DenseMap<Register, unsigned> Src2ToCand;
+ for (unsigned I = 0; I < Cands.size(); I++) {
+ DstToCand[Cands[I]->getOperand(0).getReg()] = I;
+ MachineOperand *Src2 =
+ TII->getNamedOperand(*Cands[I], AMDGPU::OpName::src2);
+ if (Src2 && Src2->isReg())
+ Src2ToCand[Src2->getReg()] = I;
+ }
+
+ SmallBitVector Visited(Cands.size());
+ auto WalkForward = [&](unsigned Start) {
+ SmallVector<unsigned, 8> Chain;
+ unsigned Cur = Start;
+ while (!Visited[Cur]) {
+ Visited[Cur] = true;
+ Chain.push_back(Cur);
+ auto It = Src2ToCand.find(Cands[Cur]->getOperand(0).getReg());
+ if (It == Src2ToCand.end() || Visited[It->second])
+ break;
+ Cur = It->second;
+ }
+ return Chain;
+ };
+
+ SmallVector<SmallVector<unsigned, 8>, 16> Chains;
+
+ // Collect linear chains from roots(src2 not produced by a candidate).
+ for (unsigned I = 0; I < Cands.size(); I++) {
+ MachineOperand *Src2 =
+ TII->getNamedOperand(*Cands[I], AMDGPU::OpName::src2);
+ if (Src2 && Src2->isReg() && DstToCand.count(Src2->getReg()))
+ continue;
+ Chains.push_back(WalkForward(I));
+ }
+
+ // Remaining chains should by cyclic. Collect them starting at any unvisited
+ // candidate.
+ for (unsigned I = 0; I < Cands.size(); I++) {
+ if (Visited[I])
+ continue;
+ Chains.push_back(WalkForward(I));
+ }
+
+ return Chains;
+}
+
+int64_t RewriteMFMAFormStage::evaluateChainProbe(
+ int N, ArrayRef<SmallVector<unsigned, 8>> Chains,
+ ArrayRef<MachineInstr *> AllCands) {
+ 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);
+}
+
+int RewriteMFMAFormStage::findBestChainCount(
+ ArrayRef<SmallVector<unsigned, 8>> Chains,
+ ArrayRef<MachineInstr *> AllCands) {
+ // Start by evaluating all chains.
+ int64_t AllCost = evaluateChainProbe(Chains.size(), Chains, AllCands);
+
+ LLVM_DEBUG(dbgs() << "RewriteMFMA probe: N=" << Chains.size()
+ << " Cost=" << AllCost << "\n");
+
+ int BestN = AllCost <= 0 ? Chains.size() : 0;
+
+ // Binary search for a good number of chains to convert. Chains are
+ // sorted by length, so we prefer converting the longest ones first as
+ // they provide the most VGPR relief per bridge copy. Converting too
+ // few chains may leave VGPRs over the limit; converting too many may
+ // push AGPRs over the limit. The search tries to find the best
+ // balance. Note: this does not guarantee a globally optimal
+ // solution as that would require evaluating all 2^K subsets of
+ // individual MFMAs. This is an approximation that works well when
+ // longer chains are more profitable. The search tracks the best
+ // cost seen across all probes to handle non-monotonicity.
+ int64_t BestCost = AllCost;
+ int Lo = 1, Hi = (int)Chains.size() - 1;
+
+ while (Lo <= Hi) {
+ int Mid = (Lo + Hi) / 2;
+ int64_t Cost = evaluateChainProbe(Mid, Chains, AllCands);
+
+ LLVM_DEBUG(dbgs() << "RewriteMFMA probe: N=" << Mid << " Cost=" << Cost
+ << "\n");
+
+ if (Cost < BestCost) {
+ BestCost = Cost;
+ BestN = Mid;
+ }
+
+ if (Cost <= 0)
+ Lo = Mid + 1;
+ else
+ Hi = Mid - 1;
+ }
+
+ return BestN;
----------------
romanovvlad wrote:
There is `int BestN = AllCost <= 0 ? Chains.size() : 0;` , so the `BestN` should be 0 when heuristic says that it's non-profitable.
https://github.com/llvm/llvm-project/pull/217328
More information about the llvm-commits
mailing list