[llvm] [AMDGPU] RewriteMFMAFormStage: pre-check dst/src2 eligibility in initHeuristics() (PR #199825)

via llvm-commits llvm-commits at lists.llvm.org
Wed Jun 3 08:47:33 PDT 2026


https://github.com/xgxanq updated https://github.com/llvm/llvm-project/pull/199825

>From b3b426c82b8c2b1af76f100b89a34f862f16c644 Mon Sep 17 00:00:00 2001
From: anqfu <anqfu at amd.com>
Date: Wed, 27 May 2026 03:54:59 +0000
Subject: [PATCH] [AMDGPU] RewriteMFMAFormStage: pre-check dst/src2 eligibility
 in initHeuristics()

Pre-check dst/src2 eligibility in initHeuristics(): non-virtual dst or
src2 registers could slip into RewriteCands and trigger downstream
asserts. Add a virtual-register pre-validation gate before the
candidate is added and its opcode/class is changed.
---
 llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp   |  9 ++-
 .../rewrite-mfma-form-check-half-rewrite.mir  | 79 +++++++++++++++++++
 2 files changed, 86 insertions(+), 2 deletions(-)
 create mode 100644 llvm/test/CodeGen/AMDGPU/rewrite-mfma-form-check-half-rewrite.mir

diff --git a/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp b/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
index 11b783a1024da..ce01269bad370 100644
--- a/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
+++ b/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
@@ -2307,10 +2307,16 @@ bool RewriteMFMAFormStage::initHeuristics(
       int ReplacementOp = AMDGPU::getMFMASrcCVDstAGPROp(MI.getOpcode());
       assert(ReplacementOp != -1);
 
+      MachineOperand *Src2 = TII->getNamedOperand(MI, AMDGPU::OpName::src2);
+      MachineOperand &Dst = MI.getOperand(0);
+      // Pre-validate: both dst and src2 (if a register) must be virtual.
+      if (!Dst.getReg().isVirtual() ||
+          (Src2->isReg() && !Src2->getReg().isVirtual()))
+        continue;
+
       RewriteCands.push_back({&MI, MI.getOpcode()});
       MI.setDesc(TII->get(ReplacementOp));
 
-      MachineOperand *Src2 = TII->getNamedOperand(MI, AMDGPU::OpName::src2);
       if (Src2->isReg()) {
         SmallVector<SlotIndex, 8> Src2ReachingDefs;
         findReachingDefs(*Src2, DAG.LIS, Src2ReachingDefs);
@@ -2323,7 +2329,6 @@ bool RewriteMFMAFormStage::initHeuristics(
         }
       }
 
-      MachineOperand &Dst = MI.getOperand(0);
       SmallVector<MachineOperand *, 8> DstReachingUses;
 
       findReachingUses(&MI, DAG.LIS, DstReachingUses);
diff --git a/llvm/test/CodeGen/AMDGPU/rewrite-mfma-form-check-half-rewrite.mir b/llvm/test/CodeGen/AMDGPU/rewrite-mfma-form-check-half-rewrite.mir
new file mode 100644
index 0000000000000..4522a1ec2db37
--- /dev/null
+++ b/llvm/test/CodeGen/AMDGPU/rewrite-mfma-form-check-half-rewrite.mir
@@ -0,0 +1,79 @@
+# RUN: llc -mtriple=amdgcn-amd-amdhsa -mcpu=gfx950 \
+# RUN:     -run-pass=machine-scheduler \
+# RUN:     -amdgpu-disable-rewrite-mfma-form-sched-stage=false \
+# RUN:     -o - %s | FileCheck %s
+#
+# Test: RewriteMFMAFormStage — physical register in MFMA src2/dst must not crash.
+#
+# Without the pre-check in initHeuristics():
+#   MFMA_B (physical src2): passes isRewriteCandidate(), then initHeuristics()
+#   calls MRI.getRegClass(Src2=physreg) which invokes virtRegIndex() and
+#   asserts isVirtual() — crash.
+#   MFMA_C (physical dst): passes isRewriteCandidate() via MRI physical-reg
+#   use tracking, then initHeuristics() calls MRI.setRegClass(Dst=physreg)
+#   which asserts isVirtual() — same crash class.
+#
+# Fix: in initHeuristics(), skip any MFMA where dst or src2 is non-virtual:
+#   if (!Dst.getReg().isVirtual() || (Src2->isReg() && !Src2->getReg().isVirtual()))
+#     continue;
+#
+# CHECK-LABEL: name: test_bug2_physreg_src2
+# CHECK-DAG:      %res_a:vreg_128_align2 = {{.*}} V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64
+# CHECK-DAG:      %copy_a:vreg_128_align2 = COPY %res_a
+# CHECK-DAG:      %res_b:vreg_128_align2 = {{.*}} V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64
+# CHECK-DAG:      %copy_b:vreg_128_align2 = COPY %res_b
+# CHECK-DAG:      $vgpr8_vgpr9_vgpr10_vgpr11 = {{.*}} V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64
+# CHECK-DAG:      %copy_c:vreg_128_align2 = COPY $vgpr8_vgpr9_vgpr10_vgpr11
+
+--- |
+  define void @test_bug2_physreg_src2() #0 {
+  entry:
+    unreachable
+  }
+
+  attributes #0 = { "amdgpu-waves-per-eu"="1,1" "amdgpu-flat-work-group-size"="64,64" }
+...
+
+---
+name:            test_bug2_physreg_src2
+tracksRegLiveness: true
+body: |
+  bb.0:
+    liveins: $vgpr4, $vgpr5, $vgpr6, $vgpr7
+
+    ; 9 x vreg_1024 = 288 VGPRs > 256: triggers RegionsWithExcessArchVGPR.
+    %0:vreg_1024 = IMPLICIT_DEF
+    %1:vreg_1024 = IMPLICIT_DEF
+    %2:vreg_1024 = IMPLICIT_DEF
+    %3:vreg_1024 = IMPLICIT_DEF
+    %4:vreg_1024 = IMPLICIT_DEF
+    %5:vreg_1024 = IMPLICIT_DEF
+    %6:vreg_1024 = IMPLICIT_DEF
+    %7:vreg_1024 = IMPLICIT_DEF
+    %8:vreg_1024 = IMPLICIT_DEF
+    %9:av_128_align2 = IMPLICIT_DEF
+    %10:av_128_align2 = IMPLICIT_DEF
+    %11:vreg_64_align2 = IMPLICIT_DEF
+    %12:vgpr_32 = IMPLICIT_DEF
+
+    SCHED_BARRIER 0
+
+    ; MFMA_A: all-virtual operands.
+    %acc_a:vreg_128_align2 = IMPLICIT_DEF
+    %res_a:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 %9:av_128_align2, %10:av_128_align2, %acc_a:vreg_128_align2, 4, 4, %11.sub0:vreg_64_align2, %12:vgpr_32, 0, 0, implicit $mode, implicit $exec
+
+    ; MFMA_B: physical src2 — exercises the src2 pre-check.
+    %res_b:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 %9:av_128_align2, %10:av_128_align2, $vgpr4_vgpr5_vgpr6_vgpr7, 4, 4, %11.sub0:vreg_64_align2, %12:vgpr_32, 0, 0, implicit $mode, implicit $exec
+
+    ; MFMA_C: physical dst — exercises the dst pre-check.
+    %acc_c:vreg_128_align2 = IMPLICIT_DEF
+    $vgpr8_vgpr9_vgpr10_vgpr11 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 %9:av_128_align2, %10:av_128_align2, %acc_c:vreg_128_align2, 4, 4, %11.sub0:vreg_64_align2, %12:vgpr_32, 0, 0, implicit $mode, implicit $exec
+
+    KILL %0, %1, %2, %3, %4, %5, %6, %7, %8
+
+    %copy_a:vreg_128_align2 = COPY %res_a
+    %copy_b:vreg_128_align2 = COPY %res_b
+    %copy_c:vreg_128_align2 = COPY $vgpr8_vgpr9_vgpr10_vgpr11
+    KILL %copy_a, %copy_b, %copy_c
+
+    S_ENDPGM 0



More information about the llvm-commits mailing list