[llvm] [AMDGPU] Add safe-guard exclusion for MFMA form rewrite (PR #207672)

via llvm-commits llvm-commits at lists.llvm.org
Mon Jul 6 00:31:01 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-amdgpu

Author: xgxanq

<details>
<summary>Changes</summary>

RewriteMFMAFormStage previously used a coarse isRewriteCandidate check that rejected any MFMA whose dst had a non-MFMA/non-COPY user. Replace that with a precise exclusion analysis that screens each candidate for cases where reclassifying to AGPR form would be illegal, and keeps the rest rewritable.

findReachingDefs is reworked into a SubRange-aware implementation backed by collectReachingDefsInRange:
- Subreg uses query the SubRange whose LaneMask fully covers the operand lanes, falling back to all overlapping SubRanges when no single SubRange provides full coverage.
- Full-reg uses collect across all SubRanges and deduplicate via a SmallSet (a full-width def appears in every SubRange but counts once).
- Self-referential defs (src2 == dst MFMA) are skipped in the traversal.

findReachingUses skips implicit operands so that a partial subreg def's implicit full-reg use (RMW lane preservation) is not treated as a real consumer, preventing spurious bridge copies.

New exclusion analysis (computeExclusionSet), run before any rewriting:
- hasSrc2BridgeConflict: detects when the src2 bridge-copy model breaks (a MAI def dominating a non-MAI def, or a src2 use not dominated by any bridge-copy block), with an early-safe return for parallel candidate defs.
- hasDstSubregConflict: detects dst subreg writers that cannot be reclassified to AGPR (non-agnostic writers, or agnostic writers with non-agnostic orphan uses).
- Exclusion propagates forward along the dst->src2 chain and backward to MAI reaching-defs of a conflicted src2.

hasUseRequiringVGPR now takes the exclusion set so that excluded MFMAs, which stay in VGPR form, are correctly treated as VGPR-requiring uses. initHeuristics is split into a candidate-collection pass and a heuristics pass over non-excluded candidates.

Add rewrite-mfma-form-safe-guard.mir covering the exclusion branches and update sched_mfma_rewrite_copies.mir for the new behavior.

---

Patch is 95.31 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/207672.diff


4 Files Affected:

- (modified) llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp (+433-103) 
- (modified) llvm/lib/Target/AMDGPU/GCNSchedStrategy.h (+57-5) 
- (added) llvm/test/CodeGen/AMDGPU/rewrite-mfma-form-safe-guard.mir (+635) 
- (modified) llvm/test/CodeGen/AMDGPU/sched_mfma_rewrite_copies.mir (+107-63) 


``````````diff
diff --git a/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp b/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
index a4f854beaeebe..e893089412115 100644
--- a/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
+++ b/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
@@ -35,6 +35,7 @@
 #include "llvm/CodeGen/MachineBasicBlock.h"
 #include "llvm/CodeGen/MachineBlockFrequencyInfo.h"
 #include "llvm/CodeGen/MachineBranchProbabilityInfo.h"
+#include "llvm/CodeGen/MachineDominators.h"
 #include "llvm/CodeGen/MachineOperand.h"
 #include "llvm/CodeGen/RegisterClassInfo.h"
 #include "llvm/CodeGen/Rematerializer.h"
@@ -1315,48 +1316,108 @@ bool GCNSchedStage::initGCNSchedStage() {
   return true;
 }
 
-void RewriteMFMAFormStage::findReachingDefs(
-    MachineOperand &UseMO, LiveIntervals *LIS,
-    SmallVectorImpl<SlotIndex> &DefIdxs) {
-  MachineInstr *UseMI = UseMO.getParent();
-  LiveInterval &UseLI = LIS->getInterval(UseMO.getReg());
-  VNInfo *VNI = UseLI.getVNInfoAt(LIS->getInstructionIndex(*UseMI));
+static void collectReachingDefsInRange(const LiveRange &QR,
+                                       MachineBasicBlock *UseMBB,
+                                       SlotIndex UseIdx,
+                                       const LiveIntervals *LIS,
+                                       SmallSet<SlotIndex, 8> &ResultSet) {
+  const VNInfo *VNI = QR.getVNInfoAt(UseIdx);
+  if (!VNI)
+    return;
 
-  // If the def is not a PHI, then it must be the only reaching def.
   if (!VNI->isPHIDef()) {
-    DefIdxs.push_back(VNI->def);
+    ResultSet.insert(VNI->def);
     return;
   }
 
-  SmallPtrSet<MachineBasicBlock *, 8> Visited = {UseMI->getParent()};
+  SmallPtrSet<MachineBasicBlock *, 8> Visited;
   SmallVector<MachineBasicBlock *, 8> Worklist;
-
-  // Mark the predecessor blocks for traversal
-  for (MachineBasicBlock *PredMBB : UseMI->getParent()->predecessors()) {
-    Worklist.push_back(PredMBB);
-    Visited.insert(PredMBB);
-  }
+  for (MachineBasicBlock *PredMBB : UseMBB->predecessors())
+    if (Visited.insert(PredMBB).second)
+      Worklist.push_back(PredMBB);
 
   while (!Worklist.empty()) {
     MachineBasicBlock *CurrMBB = Worklist.pop_back_val();
+    SlotIndex CurrMBBEnd = LIS->getMBBEndIdx(CurrMBB).getPrevSlot();
+    const VNInfo *PredVNI = QR.getVNInfoAt(CurrMBBEnd);
+    if (!PredVNI)
+      continue;
 
-    SlotIndex CurrMBBEnd = LIS->getMBBEndIdx(CurrMBB);
-    VNInfo *VNI = UseLI.getVNInfoAt(CurrMBBEnd.getPrevSlot());
-
-    MachineBasicBlock *DefMBB = LIS->getMBBFromIndex(VNI->def);
-
-    // If there is a def in this block, then add it to the list. This is the
-    // reaching def of this path.
-    if (!VNI->isPHIDef()) {
-      DefIdxs.push_back(VNI->def);
+    if (!PredVNI->isPHIDef()) {
+      // Skip self-referential defs (src2 == dst): the MFMA's own result must
+      // not appear as a reaching def of its own src2 operand.
+      if (SlotIndex::isSameInstr(PredVNI->def, UseIdx))
+        continue;
+      ResultSet.insert(PredVNI->def);
       continue;
     }
 
-    for (MachineBasicBlock *PredMBB : DefMBB->predecessors()) {
+    MachineBasicBlock *DefMBB = LIS->getMBBFromIndex(PredVNI->def);
+    for (MachineBasicBlock *PredMBB : DefMBB->predecessors())
       if (Visited.insert(PredMBB).second)
         Worklist.push_back(PredMBB);
+  }
+}
+
+void RewriteMFMAFormStage::findReachingDefs(
+    MachineOperand &UseMO, LiveIntervals *LIS,
+    SmallVectorImpl<SlotIndex> &DefIdxs) {
+  MachineInstr *UseMI = UseMO.getParent();
+  Register UseReg = UseMO.getReg();
+  unsigned UseSubReg = UseMO.getSubReg();
+
+  if (!UseReg.isVirtual() || !LIS->hasInterval(UseReg))
+    return;
+
+  LiveInterval &UseLI = LIS->getInterval(UseReg);
+  SlotIndex UseIdx = LIS->getInstructionIndex(*UseMI);
+
+  // Use a set to deduplicate: a full-reg def appears in every SubRange but
+  // must be inserted into DefIdxs only once.
+  SmallSet<SlotIndex, 8> ResultSet;
+
+  if (UseLI.hasSubRanges()) {
+    const TargetRegisterInfo *TRI = DAG.MRI.getTargetRegisterInfo();
+
+    if (UseSubReg) {
+      // Find the SubRange whose LaneMask fully covers the operand's lanes.
+      // If none does (e.g. register initialised via per-lane subreg writes),
+      // fall back to all overlapping SubRanges to capture every reaching def.
+      LaneBitmask UseLanes = TRI->getSubRegIndexLaneMask(UseSubReg);
+      bool FoundFullCoverage = false;
+      for (LiveInterval::SubRange &SR : UseLI.subranges()) {
+        if ((SR.LaneMask & UseLanes) == UseLanes) {
+          collectReachingDefsInRange(SR, UseMI->getParent(), UseIdx, LIS,
+                                     ResultSet);
+          FoundFullCoverage = true;
+          break;
+        }
+      }
+      if (!FoundFullCoverage) {
+        for (LiveInterval::SubRange &SR : UseLI.subranges())
+          if ((SR.LaneMask & UseLanes).any())
+            collectReachingDefsInRange(SR, UseMI->getParent(), UseIdx, LIS,
+                                       ResultSet);
+      }
+    } else {
+      // Full-reg use: query every subrange so that partial (subreg) defs on
+      // different lanes are all captured.
+      for (LiveInterval::SubRange &SR : UseLI.subranges())
+        collectReachingDefsInRange(SR, UseMI->getParent(), UseIdx, LIS,
+                                   ResultSet);
+      // If the register has SubRanges but none covers the use point,
+      // LiveIntervals is malformed: a full-width def must appear in at least
+      // one SubRange.
+      assert(!ResultSet.empty() &&
+             "hasSubRanges() but no SubRange live at full-reg use: "
+             "LiveInterval construction is inconsistent");
     }
+  } else {
+    collectReachingDefsInRange(UseLI, UseMI->getParent(), UseIdx, LIS,
+                               ResultSet);
   }
+
+  DefIdxs.append(ResultSet.begin(), ResultSet.end());
 }
 
 void RewriteMFMAFormStage::findReachingUses(
@@ -1365,6 +1426,12 @@ void RewriteMFMAFormStage::findReachingUses(
   SlotIndex DefIdx = LIS->getInstructionIndex(*DefMI);
   for (MachineOperand &UseMO :
        DAG.MRI.use_nodbg_operands(DefMI->getOperand(0).getReg())) {
+    // Skip implicit operands: partial subreg defs carry an implicit use of the
+    // full register for RMW lane preservation, not as a real consumer of the
+    // MFMA dst value. Treating them as reaching uses inserts a spurious bridge
+    // copy before the partial def, corrupting MappedReg's live range.
+    if (UseMO.isImplicit())
+      continue;
     SmallVector<SlotIndex, 8> ReachingDefIndexes;
     findReachingDefs(UseMO, LIS, ReachingDefIndexes);
 
@@ -2290,12 +2357,29 @@ void GCNSchedStage::modifyRegionSchedule(unsigned RegionIdx,
   DAG.Regions[RegionIdx].first = MIOrder.front();
 }
 
+static unsigned getDefSubReg(const MachineInstr &MI, Register Reg) {
+  for (const MachineOperand &MO : MI.operands())
+    if (MO.isReg() && MO.isDef() && MO.getReg() == Reg)
+      return MO.getSubReg();
+  return AMDGPU::NoSubRegister;
+}
+
+/// Returns true if \p MI is a class-agnostic subreg writer (COPY or AV_MOV).
+/// These lower to v_accvgpr_write after AGPR reclassification and are legal.
+static bool isAgnosticSubregWriter(const MachineInstr *MI) {
+  if (MI->isCopy())
+    return true;
+  unsigned Opc = MI->getOpcode();
+  return Opc == AMDGPU::AV_MOV_B32_IMM_PSEUDO ||
+         Opc == AMDGPU::AV_MOV_B64_IMM_PSEUDO;
+}
+
 /// Returns true if reaching def \p RD will be in AGPR form after the rewrite
 /// and so needs no bridge copy: a candidate MFMA in \p RewriteSet, an
 /// AV_MOV_*_IMM_PSEUDO, or a copy from a candidate src2 reg in \p CandSrc2Regs.
 /// A non-candidate MFMA stays in VGPR form and still needs a bridge.
 static bool isReachingDefAGPRForm(
-    MachineInstr *RD, const SmallPtrSetImpl<MachineInstr *> &RewriteSet,
+    MachineInstr *RD, const SmallSetVector<MachineInstr *, 16> &RewriteSet,
     const DenseSet<Register> &CandSrc2Regs, const SIInstrInfo &TII) {
   if (TII.isMAI(*RD))
     return RewriteSet.contains(RD);
@@ -2309,16 +2393,27 @@ static bool isReachingDefAGPRForm(
 
 bool RewriteMFMAFormStage::hasUseRequiringVGPR(
     ArrayRef<SlotIndex> Src2ReachingDefs,
-    const SmallPtrSetImpl<MachineInstr *> &RewriteSet) {
+    const SmallSetVector<MachineInstr *, 16> &RewriteSet,
+    const SmallPtrSetImpl<MachineInstr *> &ExcludedMFMAs) {
   for (SlotIndex RDIdx : Src2ReachingDefs) {
     const MachineInstr *RD = DAG.LIS->getInstructionFromIndex(RDIdx);
+    // If RD is a non-excluded MFMA candidate, its dst will be reclassified to
+    // AGPR and Case 2 will insert vreg bridge copies for all non-MAI uses of
+    // its dst. Those uses therefore do not impose a VGPR constraint on src2.
+    // The !ExcludedMFMAs.count(RD) guard is defensive:
+    // propagateExclusionForward ensures that if an excluded MFMA's dst is used
+    // as src2, the consumer MFMA is also excluded and hasUseRequiringVGPR is
+    // never called for it.
+    if (TII->isMAI(*RD) && RewriteSet.contains(RD) && !ExcludedMFMAs.count(RD))
+      continue;
     SmallVector<MachineOperand *, 8> ReachingUses;
     findReachingUses(RD, DAG.LIS, ReachingUses);
     for (const MachineOperand *UseMO : ReachingUses) {
       const MachineInstr *UseMI = UseMO->getParent();
       if (UseMI->isCopy())
         continue;
-      if (TII->isMAI(*UseMI) && RewriteSet.contains(UseMI))
+      if (TII->isMAI(*UseMI) && RewriteSet.contains(UseMI) &&
+          !ExcludedMFMAs.count(UseMI))
         continue;
       return true;
     }
@@ -2354,24 +2449,247 @@ bool RewriteMFMAFormStage::isRewriteCandidate(MachineInstr *MI) const {
     return false;
   if (AMDGPU::getMFMASrcCVDstAGPROp(MI->getOpcode()) == -1)
     return false;
-  // Reject candidates whose users force an unavoidable bridge copy.
-  Register DstReg = MI->getOperand(0).getReg();
-  for (const MachineOperand &Use : DAG.MRI.use_nodbg_operands(DstReg)) {
-    if (!TII->isMAI(*Use.getParent()) && !Use.getParent()->isCopy())
-      return false;
-  }
   return true;
 }
 
+bool RewriteMFMAFormStage::hasSrc2BridgeConflict(ArrayRef<SlotIndex> DefIdxs,
+                                                 Register Src2Reg) const {
+  auto &MDT = DAG.LIS->getDomTree();
+
+  SmallVector<MachineInstr *, 8> MAIMIs;
+  SmallPtrSet<MachineBasicBlock *, 8> BridgeCopyBlocks; // non-MAI def blocks
+
+  for (SlotIndex SI : DefIdxs) {
+    MachineInstr *MI = DAG.LIS->getInstructionFromIndex(SI);
+    if (TII->isMFMA(*MI))
+      MAIMIs.push_back(MI);
+    else
+      BridgeCopyBlocks.insert(MI->getParent());
+  }
+
+  if (BridgeCopyBlocks.empty())
+    return false; // All defs are MAI; no bridge copies needed.
+
+  // Check 1: MAI def dominates non-MAI def (partial subreg overwrite).
+  // The MFMA first writes all lanes (AGPR), then a non-MAI instruction
+  // partially overwrites some lanes.  A bridge copy at the non-MAI def
+  // would read an already-AGPR register and must partially update it —
+  // a read-modify-write that the bridge-copy mechanism cannot implement.
+  if (!MAIMIs.empty()) {
+    SmallVector<MachineInstr *, 8> NonMAIMIs;
+    for (SlotIndex SI : DefIdxs) {
+      MachineInstr *MI = DAG.LIS->getInstructionFromIndex(SI);
+      if (!TII->isMFMA(*MI))
+        NonMAIMIs.push_back(MI);
+    }
+
+    for (MachineInstr *M : MAIMIs)
+      for (MachineInstr *N : NonMAIMIs)
+        if (MDT.dominates(M, N))
+          return true;
+
+    // Early-safe return: if every (MAI, non-MAI) pair is parallel and every
+    // MAI def is itself a rewrite candidate, the rewrite is safe without
+    // invoking Check 2.
+    if (!NonMAIMIs.empty()) {
+      bool AllParallelAndCandidates = true;
+      for (MachineInstr *M : MAIMIs) {
+        if (!isRewriteCandidate(M)) {
+          AllParallelAndCandidates = false;
+          break;
+        }
+        for (MachineInstr *N : NonMAIMIs) {
+          if (MDT.dominates(N, M)) {
+            AllParallelAndCandidates = false;
+            break;
+          }
+        }
+        if (!AllParallelAndCandidates)
+          break;
+      }
+      if (AllParallelAndCandidates)
+        return false;
+    }
+  }
+
+  // Check 2: every use of Src2Reg must be dominated by at least one
+  // bridge-copy block; otherwise %MappedReg would be undefined on some path.
+  //
+  // use_nodbg_operands is used here (not findReachingUses) because rewrite()
+  // replaces the src2 operand of the MFMA with %MappedReg uniformly — it does
+  // not distinguish which reaching def flows to which use.  %MappedReg is
+  // defined only in bridge-copy blocks (after non-MAI defs); if any use of
+  // Src2Reg is in a block not dominated by any bridge-copy block, %MappedReg
+  // would be undefined on that path regardless of whether the use is reached
+  // by a MAI or non-MAI def.  findReachingUses(non-MAI RD) would miss uses
+  // that arrive only via MAI defs or other defs (e.g. IMPLICIT_DEF on a
+  // bypass path), causing a false negative and silent undefined-read.
+  for (const MachineOperand &UseMO : DAG.MRI.use_nodbg_operands(Src2Reg)) {
+    const MachineBasicBlock *UseBlock = UseMO.getParent()->getParent();
+    bool Covered = any_of(BridgeCopyBlocks, [&](const MachineBasicBlock *B) {
+      return MDT.dominates(B, UseBlock);
+    });
+    if (!Covered)
+      return true;
+  }
+
+  return false;
+}
+
+void RewriteMFMAFormStage::propagateExclusionForward(
+    MachineInstr *Root, SmallPtrSetImpl<MachineInstr *> &ExcludedMFMAs) {
+  SmallVector<MachineInstr *, 8> Worklist = {Root};
+  while (!Worklist.empty()) {
+    MachineInstr *ExclMI = Worklist.pop_back_val();
+    MachineOperand &DstMO = ExclMI->getOperand(0);
+    if (!DstMO.isReg() || !DstMO.getReg().isVirtual())
+      continue;
+    Register DstReg = DstMO.getReg();
+    for (MachineOperand &UseMO : DAG.MRI.use_nodbg_operands(DstReg)) {
+      MachineInstr *UserMI = UseMO.getParent();
+      if (!isRewriteCandidate(UserMI))
+        continue;
+      MachineOperand *UserSrc2 =
+          TII->getNamedOperand(*UserMI, AMDGPU::OpName::src2);
+      if (!UserSrc2 || !UserSrc2->isReg() || UserSrc2->getReg() != DstReg)
+        continue;
+      if (ExcludedMFMAs.insert(UserMI).second) {
+        LLVM_DEBUG(dbgs() << "[initHeuristics] exclude downstream MFMA "
+                             "(src2 = excluded MFMA dst): "
+                          << *UserMI);
+        Worklist.push_back(UserMI);
+      }
+    }
+  }
+}
+
+// rewrite()'s design principle: reclassify DstReg from VGPR to AGPR class so
+// that the MFMA emits its result directly into AGPR, eliminating the need for
+// a post-MFMA VGPR→AGPR copy.  For this reclassification to be legal, every
+// def and every use of DstReg throughout its live range must support
+// AGPR-class operands.  hasDstSubregConflict checks this before rewriting:
+//
+bool RewriteMFMAFormStage::hasDstSubregConflict(Register DstReg,
+                                                MachineInstr *MFMA) {
+  SmallVector<MachineOperand *, 8> DstReachingUses;
+  findReachingUses(MFMA, DAG.LIS, DstReachingUses);
+  SmallPtrSet<MachineInstr *, 8> CheckedDefs;
+  SmallVector<MachineInstr *, 4> SafeSubregDefs;
+
+  // Phase 1: classify subreg reaching defs.
+  // Non-agnostic subreg def → immediate conflict.
+  // Agnostic (COPY/AV_MOV) subreg def → defer to orphan-use check.
+  for (MachineOperand *RUOp : DstReachingUses) {
+    SmallVector<SlotIndex, 8> ReachingDefs;
+    findReachingDefs(*RUOp, DAG.LIS, ReachingDefs);
+    for (SlotIndex RDIdx : ReachingDefs) {
+      MachineInstr *RD = DAG.LIS->getInstructionFromIndex(RDIdx);
+      if (!CheckedDefs.insert(RD).second)
+        continue;
+      if (getDefSubReg(*RD, DstReg) == AMDGPU::NoSubRegister || TII->isMAI(*RD))
+        continue;
+      if (!isAgnosticSubregWriter(RD))
+        return true;
+      SafeSubregDefs.push_back(RD);
+    }
+  }
+
+  // Phase 2: agnostic subreg defs lower to v_accvgpr_write after reclassify,
+  // so their orphan uses read an AGPR sub-register.  Only COPY and AV_MOV can
+  // legally source an AGPR sub-register; anything else is a conflict.
+  for (MachineInstr *RD : SafeSubregDefs) {
+    SlotIndex RDIdx = DAG.LIS->getInstructionIndex(*RD);
+    for (MachineOperand &UseMO : DAG.MRI.use_nodbg_operands(DstReg)) {
+      if (UseMO.isImplicit() || TII->isMAI(*UseMO.getParent()))
+        continue;
+      SmallVector<SlotIndex, 8> UseReachingDefs;
+      findReachingDefs(UseMO, DAG.LIS, UseReachingDefs);
+      if (any_of(UseReachingDefs,
+                 [RDIdx](SlotIndex SI) {
+                   return SlotIndex::isSameInstr(SI, RDIdx);
+                 }) &&
+          !isAgnosticSubregWriter(UseMO.getParent()))
+        return true;
+    }
+  }
+  return false;
+}
+
+SmallPtrSet<MachineInstr *, 16> RewriteMFMAFormStage::computeExclusionSet(
+    const SmallSetVector<MachineInstr *, 16> &RewriteSet) {
+  // Per-MI checks run cheapest-first:
+  //   1. hasSrc2BridgeConflict: bridge COPY after non-MAI src2 def would be
+  //      a read-modify-write on AGPR (MAI def dominates non-MAI def), or
+  //      absent on some CFG path to a src2 use.
+  //   2. hasDstSubregConflict: DstReg has non-MAI subreg writers that cannot
+  //      be reclassified to AGPR.  Skipped when check 1 already forces
+  //      exclusion (!HasConflict &&).
+  // Exclusion propagates forward (dst→src2 chain via propagateExclusionForward)
+  // and backward (MAI reaching-defs of a conflicted src2).
+  SmallPtrSet<MachineInstr *, 16> ExcludedMFMAs;
+  for (MachineInstr *MI : RewriteSet) {
+    MachineOperand *Src2 = TII->getNamedOperand(*MI, AMDGPU::OpName::src2);
+    Register DstReg = MI->getOperand(0).getReg();
+
+    bool HasConflict = false;
+    SmallVector<SlotIndex, 8> Src2Defs;
+    if (Src2->isReg()) {
+      findReachingDefs(*Src2, DAG.LIS, Src2Defs);
+      LLVM_DEBUG({
+        dbgs() << "[computeExclusionSet] candidate: " << *MI;
+        dbgs() << "  src2 reaching defs (" << Src2Defs.size() << "):\n";
+        for (SlotIndex SI : Src2Defs) {
+          MachineInstr *D = DAG.LIS->getInstructionFromIndex(SI);
+          dbgs() << "    " << SI << " opcode=" << (D ? (int)D->getOpcode() : -1)
+                 << "\n";
+          if (D)
+            dbgs() << "    " << *D;
+        }
+      });
+      HasConflict = hasSrc2BridgeConflict(Src2Defs, Src2->getReg());
+    }
+    bool HasDstSubregDef = !HasConflict && hasDstSubregConflict(DstReg, MI);
+
+    if (!HasConflict && !HasDstSubregDef)
+      continue;
+
+    if (ExcludedMFMAs.insert(MI).second) {
+      LLVM_DEBUG(
+          dbgs() << "[computeExclusionSet] exclude MFMA ("
+                 << (HasConflict ? "src2 dominance conflict" : "")
+                 << (HasConflict && HasDstSubregDef ? " + " : "")
+                 << (HasDstSubregDef ? "dst non-MAI subreg overwrite" : "")
+                 << "): " << *MI);
+      propagateExclusionForward(MI, ExcludedMFMAs);
+    }
+
+    // Backward: exclude MAI reaching-defs that are themselves candidates.
+    if (HasConflict) {
+      for (SlotIndex SI : Src2Defs) {
+        MachineInstr *DefMI = DAG.LIS->getInstructionFromIndex(SI);
+        if (TII->isMFMA(*DefMI) && isRewriteCandidate(DefMI) &&
+            ExcludedMFMAs.insert(DefMI).second) {
+          LLVM_DEBUG(dbgs() << "[computeExclusionSet] exclude MAI def "
+                               "(backward from src2 dominance conflict): "
+                            << *DefMI);
+          propagateExclusionForward(DefMI, ExcludedMFMAs);
+        }
+      }
+    }
+  }
+  return ExcludedMFMAs;
+}
+
 bool RewriteMFMAFormStage::initHeuristics(
     std::vector<std::pair<MachineInstr *, unsigned>> &RewriteCands,
     DenseMap<MachineBasicBlock *, std::set<Register>> &CopyForUse,
     SmallPtrSetImpl<MachineInstr *> &CopyForDef) {
   bool Changed = false;
 
-  // Collect the candidate group, its members share AGPR-form operands
-  // post-rewrite, so reaching defs feeding any member don't need bridge copy.
-  SmallPtrSet<MachineInstr *, 16> RewriteSet;
+  // Pass 1: collect candidate group.
+  // RewriteSet/CandSrc2Regs are needed by isReachingDefAGPRForm and
+  // hasUseRequiringVGPR; collect them before any setDesc/setRegClass changes.
+  SmallSetVector<MachineInstr *, 16> RewriteSet;
   DenseSet<Register> CandSrc2Regs;
   for (MachineBasicBlock &MBB : MF) {
     for (MachineInstr &MI : MBB) {
@@ -2384,86 +2702,98 @@ bool RewriteMFMAFormStage::initHeuristics(
     }
   }
 
-  // Prepare for the heuristics
-  for (MachineBasicBlock &MBB : MF) {
-   ...
[truncated]

``````````

</details>


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


More information about the llvm-commits mailing list