[llvm] [AMDGPU] Per-chain MFMA->AGPR conversion (PR #217328)
Shilei Tian via llvm-commits
llvm-commits at lists.llvm.org
Sun Sep 13 09:21:17 PDT 2026
================
@@ -1413,6 +1415,123 @@ 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
----------------
shiltian wrote:
I'm not sure about this binary search approach. If the first evaluation fails, which means we can't convert all the chains, we're going to skip all the chains before the mid, which means we could skip all the potentially long chains. If those chains are unbalanced, this wouldn't work well, right?
For example, if I have `[128, 96, 12, 8, 4]`, we can't take all of them, but we can take `[96, 12, 8, 4]`. This approach would miss that and potentially use `[12, 8, 4]` instead.
https://github.com/llvm/llvm-project/pull/217328
More information about the llvm-commits
mailing list