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

Shilei Tian via llvm-commits llvm-commits at lists.llvm.org
Sun Sep 13 08:34:15 PDT 2026


================
@@ -3302,64 +3370,15 @@ int GCNHazardRecognizer::checkMAIHazards90A(MachineInstr *MI) const {
         else if (ST.hasGFX940Insts() &&
                  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");
-          }
-        }
+        // 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() && !isSameMFMA(Opc, Opc1))
+          NeedWaitStates = std::max(NeedWaitStates,
+                                    getMFMAOverlappedSrcCWaitStates(MI, MI1));
----------------
shiltian wrote:

`{}`

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


More information about the llvm-commits mailing list