[llvm] [AMDGPU] Add wait states between different MFMAs sharing an accumulator (PR #218363)

Pankaj Dwivedi via llvm-commits llvm-commits at lists.llvm.org
Mon Aug 24 02:50:51 PDT 2026


https://github.com/PankajDwivedi-25 created https://github.com/llvm/llvm-project/pull/218363

A full-register src2/C read after an MFMA write emits no wait states,
relying on accumulator forwarding that only works while the chain stays
on one MFMA. Two different MFMAs sharing an accumulator instead need the
wait states of a partial overlap: on gfx950, v_mfma_f32_16x16x32_f16 then
v_mfma_f32_16x16x16_f16 on the same tuple needs 5 and got none. Compare
canonicalized opcodes so mac and register-bank forms still match.

Assisted-by: Claude Opus 5

>From 95f1eefd60b7baec0e814d4d377256e24279559a Mon Sep 17 00:00:00 2001
From: Pankaj Dwivedi <pankajkumar.divedi at amd.com>
Date: Mon, 24 Aug 2026 09:26:58 +0000
Subject: [PATCH 1/2] [AMDGPU][NFC] Move MFMA overlapped src2/C wait states
 into a helper

---
 .../lib/Target/AMDGPU/GCNHazardRecognizer.cpp | 129 +++++++++---------
 llvm/lib/Target/AMDGPU/GCNHazardRecognizer.h  |   5 +
 2 files changed, 68 insertions(+), 66 deletions(-)

diff --git a/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp b/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp
index 9d01a529eb012..15d4cabe6102f 100644
--- a/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp
+++ b/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp
@@ -3214,6 +3214,68 @@ static int GFX940_XDL_N_PassWritesVGPROverlappedSrcABWaitStates(int NumPasses,
   return NumPasses + 3 + (NumPasses != 2 && IsGFX950);
 }
 
+int GCNHazardRecognizer::getMFMAOverlappedSrcCWaitStates(
+    const MachineInstr *MI, const MachineInstr *MI1) const {
+  const int SMFMA4x4WritesVGPROverlappedSMFMASrcCWaitStates = 2;
+  const int SMFMA16x16WritesVGPROverlappedSMFMASrcCWaitStates = 8;
+  const int SMFMA32x32WritesVGPROverlappedSMFMASrcCWaitStates = 16;
+  const int SMFMA4x4WritesVGPROverlappedDMFMASrcCWaitStates = 3;
+  const int SMFMA16x16WritesVGPROverlappedDMFMASrcCWaitStates = 9;
+  const int SMFMA32x32WritesVGPROverlappedDMFMASrcCWaitStates = 17;
+  const int DMFMA16x16WritesVGPROverlappedSrcCWaitStates = 9;
+  const int GFX950_DMFMA16x16WritesVGPROverlappedSrcCWaitStates = 17;
+  const int DMFMA4x4WritesVGPROverlappedSrcCWaitStates = 4;
+
+  switch (MI1->getOpcode()) {
+  case AMDGPU::V_MFMA_F64_16X16X4F64_e64:
+  case AMDGPU::V_MFMA_F64_16X16X4F64_vgprcd_e64:
+  case AMDGPU::V_MFMA_F64_16X16X4F64_mac_e64:
+  case AMDGPU::V_MFMA_F64_16X16X4F64_mac_vgprcd_e64:
+    if (TII.isXDL(*MI))
+      return 0;
+    return ST.hasGFX950Insts()
+               ? GFX950_DMFMA16x16WritesVGPROverlappedSrcCWaitStates
+               : DMFMA16x16WritesVGPROverlappedSrcCWaitStates;
+  case AMDGPU::V_MFMA_F64_4X4X4F64_e64:
+  case AMDGPU::V_MFMA_F64_4X4X4F64_vgprcd_e64:
+    if (TII.isXDL(*MI))
+      return 0;
+    return DMFMA4x4WritesVGPROverlappedSrcCWaitStates;
+  default:
+    break;
+  }
+
+  int NumPasses = TSchedModel.computeInstrLatency(MI1);
+  if (ST.hasGFX940Insts()) {
+    if (TII.isXDL(*MI) && !TII.isXDL(*MI1))
+      return 0;
+
+    if (!TII.isXDL(*MI1))
+      return GFX940_SMFMA_N_PassWritesVGPROverlappedSMFMASrcCWaitStates(
+          NumPasses);
+    return TII.isXDL(*MI)
+               ? GFX940_XDL_N_PassWritesVGPROverlappedXDLOrSMFMASrcCWaitStates(
+                     NumPasses, ST.hasGFX950Insts())
+               : GFX940_XDL_N_PassWritesVGPROverlappedSGEMMDGEMMSrcCWaitStates(
+                     NumPasses, ST.hasGFX950Insts());
+  }
+
+  bool IsDGEMM = SIInstrInfo::isDGEMM(MI->getOpcode());
+  switch (NumPasses) {
+  case 2:
+    return IsDGEMM ? SMFMA4x4WritesVGPROverlappedDMFMASrcCWaitStates
+                   : SMFMA4x4WritesVGPROverlappedSMFMASrcCWaitStates;
+  case 8:
+    return IsDGEMM ? SMFMA16x16WritesVGPROverlappedDMFMASrcCWaitStates
+                   : SMFMA16x16WritesVGPROverlappedSMFMASrcCWaitStates;
+  case 16:
+    return IsDGEMM ? SMFMA32x32WritesVGPROverlappedDMFMASrcCWaitStates
+                   : SMFMA32x32WritesVGPROverlappedSMFMASrcCWaitStates;
+  default:
+    llvm_unreachable("unexpected number of passes");
+  }
+}
+
 int GCNHazardRecognizer::checkMAIHazards90A(MachineInstr *MI) const {
   int WaitStatesNeeded = 0;
   unsigned Opc = MI->getOpcode();
@@ -3242,15 +3304,6 @@ int GCNHazardRecognizer::checkMAIHazards90A(MachineInstr *MI) const {
   // Loop for both DGEMM and S/HGEMM 2nd instruction.
   for (const MachineOperand &Use : MI->explicit_uses()) {
     const int LegacyVALUNotDotWritesVGPRWaitStates = 2;
-    const int SMFMA4x4WritesVGPROverlappedSMFMASrcCWaitStates = 2;
-    const int SMFMA16x16WritesVGPROverlappedSMFMASrcCWaitStates = 8;
-    const int SMFMA32x32WritesVGPROverlappedSMFMASrcCWaitStates = 16;
-    const int SMFMA4x4WritesVGPROverlappedDMFMASrcCWaitStates = 3;
-    const int SMFMA16x16WritesVGPROverlappedDMFMASrcCWaitStates = 9;
-    const int SMFMA32x32WritesVGPROverlappedDMFMASrcCWaitStates = 17;
-    const int DMFMA16x16WritesVGPROverlappedSrcCWaitStates = 9;
-    const int GFX950_DMFMA16x16WritesVGPROverlappedSrcCWaitStates = 17;
-    const int DMFMA4x4WritesVGPROverlappedSrcCWaitStates = 4;
     const int SMFMA4x4WritesVGPROverlappedSrcABWaitStates = 5;
     const int SMFMA16x16WritesVGPROverlappedSrcABWaitStates = 11;
     const int SMFMA32x32WritesVGPROverlappedSrcABWaitStates = 19;
@@ -3303,63 +3356,7 @@ int GCNHazardRecognizer::checkMAIHazards90A(MachineInstr *MI) const {
                  TSchedModel.computeInstrLatency(MI1) == 2)
           NeedWaitStates = GFX940_SMFMA4x4WritesVGPRFullSrcCWaitStates;
       } else {
-        switch (Opc1) {
-        case AMDGPU::V_MFMA_F64_16X16X4F64_e64:
-        case AMDGPU::V_MFMA_F64_16X16X4F64_vgprcd_e64:
-        case AMDGPU::V_MFMA_F64_16X16X4F64_mac_e64:
-        case AMDGPU::V_MFMA_F64_16X16X4F64_mac_vgprcd_e64:
-          if (!TII.isXDL(*MI))
-            NeedWaitStates =
-                ST.hasGFX950Insts()
-                    ? GFX950_DMFMA16x16WritesVGPROverlappedSrcCWaitStates
-                    : DMFMA16x16WritesVGPROverlappedSrcCWaitStates;
-          break;
-        case AMDGPU::V_MFMA_F64_4X4X4F64_e64:
-        case AMDGPU::V_MFMA_F64_4X4X4F64_vgprcd_e64:
-          if (!TII.isXDL(*MI))
-            NeedWaitStates = DMFMA4x4WritesVGPROverlappedSrcCWaitStates;
-          break;
-        default:
-          int NumPasses = TSchedModel.computeInstrLatency(MI1);
-          if (ST.hasGFX940Insts()) {
-            if (TII.isXDL(*MI) && !TII.isXDL(*MI1))
-              break;
-
-            NeedWaitStates =
-                TII.isXDL(*MI1)
-                    ? (TII.isXDL(*MI)
-                           ? GFX940_XDL_N_PassWritesVGPROverlappedXDLOrSMFMASrcCWaitStates(
-                                 NumPasses, ST.hasGFX950Insts())
-                           : GFX940_XDL_N_PassWritesVGPROverlappedSGEMMDGEMMSrcCWaitStates(
-                                 NumPasses, ST.hasGFX950Insts()))
-                    : GFX940_SMFMA_N_PassWritesVGPROverlappedSMFMASrcCWaitStates(
-                          NumPasses);
-            break;
-          }
-
-          switch (NumPasses) {
-          case 2:
-            NeedWaitStates =
-                SIInstrInfo::isDGEMM(Opc)
-                    ? SMFMA4x4WritesVGPROverlappedDMFMASrcCWaitStates
-                    : SMFMA4x4WritesVGPROverlappedSMFMASrcCWaitStates;
-            break;
-          case 8:
-            NeedWaitStates =
-                SIInstrInfo::isDGEMM(Opc)
-                    ? SMFMA16x16WritesVGPROverlappedDMFMASrcCWaitStates
-                    : SMFMA16x16WritesVGPROverlappedSMFMASrcCWaitStates;
-            break;
-          case 16:
-            NeedWaitStates =
-                SIInstrInfo::isDGEMM(Opc)
-                    ? SMFMA32x32WritesVGPROverlappedDMFMASrcCWaitStates
-                    : SMFMA32x32WritesVGPROverlappedSMFMASrcCWaitStates;
-            break;
-          default:
-            llvm_unreachable("unexpected number of passes");
-          }
-        }
+        NeedWaitStates = getMFMAOverlappedSrcCWaitStates(MI, MI1);
       }
     } else {
       switch (Opc1) {
diff --git a/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.h b/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.h
index 422da557eb47d..8e30449b04fb3 100644
--- a/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.h
+++ b/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.h
@@ -242,6 +242,11 @@ class GCNHazardRecognizer final : public ScheduleHazardRecognizer {
   int checkMAIHazards(MachineInstr *MI) const;
   int checkMAIHazards908(MachineInstr *MI) const;
   int checkMAIHazards90A(MachineInstr *MI) const;
+  /// Wait states needed when \p MI reads as src2/C a register that the earlier
+  /// MFMA \p MI1 wrote, taken from the wait state tables for overlapping
+  /// register tuples.
+  int getMFMAOverlappedSrcCWaitStates(const MachineInstr *MI,
+                                      const MachineInstr *MI1) const;
   /// Pad the latency between neighboring MFMA instructions with s_nops. The
   /// percentage of wait states to fill with s_nops is specified by the command
   /// line option '-amdgpu-mfma-padding-ratio'.

>From cfae73db1499309265e2487c8fe3d40e506a75a0 Mon Sep 17 00:00:00 2001
From: Pankaj Dwivedi <pankajkumar.divedi at amd.com>
Date: Mon, 24 Aug 2026 09:47:10 +0000
Subject: [PATCH 2/2] [AMDGPU] Add wait states between different MFMAs sharing
 an accumulator

---
 .../lib/Target/AMDGPU/GCNHazardRecognizer.cpp | 18 +++++++++++
 .../CodeGen/AMDGPU/mai-hazards-gfx942.mir     | 30 +++++++++++++++++++
 2 files changed, 48 insertions(+)

diff --git a/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp b/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp
index 15d4cabe6102f..ea166d783a945 100644
--- a/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp
+++ b/llvm/lib/Target/AMDGPU/GCNHazardRecognizer.cpp
@@ -3164,6 +3164,16 @@ int GCNHazardRecognizer::checkMAIHazards908(MachineInstr *MI) const {
   return WaitStatesNeeded;
 }
 
+/// Map an MFMA opcode to a form that ignores the vdst/src2 register bank and
+/// the mac form, so two MFMAs can be compared for being the same instruction.
+static unsigned canonicalizeMFMAOpcode(unsigned Opc) {
+  if (int MacOp = AMDGPU::getMFMAEarlyClobberOp(Opc); MacOp != -1)
+    Opc = MacOp;
+  if (int AGPROp = AMDGPU::getAGPRFormOp(Opc); AGPROp != -1)
+    Opc = AGPROp;
+  return Opc;
+}
+
 static int
 GFX940_XDL_N_PassWritesVGPROverlappedXDLOrSMFMASrcCWaitStates(int NumPasses,
                                                               bool IsGFX950) {
@@ -3355,6 +3365,14 @@ int GCNHazardRecognizer::checkMAIHazards90A(MachineInstr *MI) const {
         else if (ST.hasGFX940Insts() &&
                  TSchedModel.computeInstrLatency(MI1) == 2)
           NeedWaitStates = GFX940_SMFMA4x4WritesVGPRFullSrcCWaitStates;
+
+        // The accumulator forwarding path that allows zero wait states is only
+        // available while the chain stays on a single MFMA. Two different MFMAs
+        // sharing an accumulator need the wait states of a partial overlap.
+        if (ST.hasGFX940Insts() &&
+            canonicalizeMFMAOpcode(Opc) != canonicalizeMFMAOpcode(Opc1))
+          NeedWaitStates = std::max(NeedWaitStates,
+                                    getMFMAOverlappedSrcCWaitStates(MI, MI1));
       } else {
         NeedWaitStates = getMFMAOverlappedSrcCWaitStates(MI, MI1);
       }
diff --git a/llvm/test/CodeGen/AMDGPU/mai-hazards-gfx942.mir b/llvm/test/CodeGen/AMDGPU/mai-hazards-gfx942.mir
index 5a08e429283a8..99fd94ff772cb 100644
--- a/llvm/test/CodeGen/AMDGPU/mai-hazards-gfx942.mir
+++ b/llvm/test/CodeGen/AMDGPU/mai-hazards-gfx942.mir
@@ -2512,3 +2512,33 @@ body:             |
     $agpr0_agpr1_agpr2_agpr3_agpr4_agpr5_agpr6_agpr7_agpr8_agpr9_agpr10_agpr11_agpr12_agpr13_agpr14_agpr15_agpr16_agpr17_agpr18_agpr19_agpr20_agpr21_agpr22_agpr23_agpr24_agpr25_agpr26_agpr27_agpr28_agpr29_agpr30_agpr31 = V_MFMA_F32_32X32X4F16_e64 $vgpr126_vgpr127, $vgpr128_vgpr129, $agpr0_agpr1_agpr2_agpr3_agpr4_agpr5_agpr6_agpr7_agpr8_agpr9_agpr10_agpr11_agpr12_agpr13_agpr14_agpr15_agpr16_agpr17_agpr18_agpr19_agpr20_agpr21_agpr22_agpr23_agpr24_agpr25_agpr26_agpr27_agpr28_agpr29_agpr30_agpr31, 0, 0, 0, implicit $mode, implicit $exec
     $agpr2_agpr3_agpr4_agpr5 = V_SMFMAC_F32_16X16X32_F16_e64 $vgpr0_vgpr1, $vgpr2_vgpr3_vgpr4_vgpr5, $vgpr32, 0, 0, $agpr2_agpr3_agpr4_agpr5, implicit $mode, implicit $exec
 ...
+# GCN-LABEL: name: xdl_sgemm16x16_mfma_write_agpr_different_xdl_mfma_read_same_agpr_as_srcc
+# GCN:      V_MFMA_F32_16X16X32_F16
+# GFX942-NEXT: S_NOP 4
+# GFX950-NEXT: S_NOP 5
+# GCN-NEXT: V_MFMA_F32_16X16X16F16
+name:            xdl_sgemm16x16_mfma_write_agpr_different_xdl_mfma_read_same_agpr_as_srcc
+body:             |
+  bb.0:
+    $agpr0_agpr1_agpr2_agpr3 = V_MFMA_F32_16X16X32_F16_e64 $vgpr0_vgpr1_vgpr2_vgpr3, $vgpr4_vgpr5_vgpr6_vgpr7, $agpr0_agpr1_agpr2_agpr3, 0, 0, 0, implicit $mode, implicit $exec
+    $agpr0_agpr1_agpr2_agpr3 = V_MFMA_F32_16X16X16F16_e64 $vgpr0_vgpr1, $vgpr2_vgpr3, $agpr0_agpr1_agpr2_agpr3, 0, 0, 0, implicit $mode, implicit $exec
+...
+# GCN-LABEL: name: nonxdl_sgemm16x16_mfma_write_agpr_different_nonxdl_mfma_read_same_agpr_as_srcc
+# GCN:      V_MFMA_F32_16X16X4F32
+# GCN-NEXT: S_NOP 7
+# GCN-NEXT: V_MFMA_F32_4X4X1F32
+name:            nonxdl_sgemm16x16_mfma_write_agpr_different_nonxdl_mfma_read_same_agpr_as_srcc
+body:             |
+  bb.0:
+    $agpr0_agpr1_agpr2_agpr3 = V_MFMA_F32_16X16X4F32_e64 $vgpr0, $vgpr1, $agpr0_agpr1_agpr2_agpr3, 0, 0, 0, implicit $mode, implicit $exec
+    $agpr0_agpr1_agpr2_agpr3 = V_MFMA_F32_4X4X1F32_e64 $vgpr0, $vgpr1, $agpr0_agpr1_agpr2_agpr3, 0, 0, 0, implicit $mode, implicit $exec
+...
+# GCN-LABEL: name: xdl_sgemm16x16_mac_mfma_write_agpr_nonmac_mfma_read_same_agpr_as_srcc
+# GCN:      V_MFMA_F32_16X16X4F16_mac
+# GCN-NEXT: V_MFMA_F32_16X16X4F16
+name:            xdl_sgemm16x16_mac_mfma_write_agpr_nonmac_mfma_read_same_agpr_as_srcc
+body:             |
+  bb.0:
+    $agpr0_agpr1_agpr2_agpr3_agpr4_agpr5_agpr6_agpr7_agpr8_agpr9_agpr10_agpr11_agpr12_agpr13_agpr14_agpr15 = V_MFMA_F32_16X16X4F16_mac_e64 $vgpr0_vgpr1, $vgpr2_vgpr3, $agpr0_agpr1_agpr2_agpr3_agpr4_agpr5_agpr6_agpr7_agpr8_agpr9_agpr10_agpr11_agpr12_agpr13_agpr14_agpr15, 0, 0, 0, implicit $mode, implicit $exec
+    $agpr16_agpr17_agpr18_agpr19_agpr20_agpr21_agpr22_agpr23_agpr24_agpr25_agpr26_agpr27_agpr28_agpr29_agpr30_agpr31 = V_MFMA_F32_16X16X4F16_e64 $vgpr0_vgpr1, $vgpr2_vgpr3, $agpr0_agpr1_agpr2_agpr3_agpr4_agpr5_agpr6_agpr7_agpr8_agpr9_agpr10_agpr11_agpr12_agpr13_agpr14_agpr15, 0, 0, 0, implicit $mode, implicit $exec
+...



More information about the llvm-commits mailing list