[llvm] [AMDGPU] Per-chain MFMA->AGPR conversion (PR #217328)

Romanov Vlad via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 7 09:56:25 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:

I see. Will add the check `Cost <= 0`. Thanks!

https://github.com/llvm/llvm-project/pull/217328


More information about the llvm-commits mailing list