[llvm] [AMDGPU] implement no-sdwa prediction in td and support globalisel (PR #193518)

via llvm-commits llvm-commits at lists.llvm.org
Wed Apr 22 09:27:09 PDT 2026


https://github.com/xiongzile updated https://github.com/llvm/llvm-project/pull/193518

>From 8eac21d53703d8e265d926864b0534f31d3cbfa8 Mon Sep 17 00:00:00 2001
From: Zile Xiong <xiongzile at bytedance.com>
Date: Wed, 22 Apr 2026 20:41:21 +0800
Subject: [PATCH] [AMDGPU] implement no-sdwa prediction in td

---
 llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.cpp | 20 ----------
 llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.h   |  1 -
 llvm/lib/Target/AMDGPU/AMDGPUInstructions.td  | 39 ++++++++++++++++++-
 3 files changed, 38 insertions(+), 22 deletions(-)

diff --git a/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.cpp b/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.cpp
index c2322bd922f31..598ab6f7354f5 100644
--- a/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.cpp
+++ b/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.cpp
@@ -897,26 +897,6 @@ void AMDGPUDAGToDAGISel::Select(SDNode *N) {
   SelectCode(N);
 }
 
-bool AMDGPUDAGToDAGISel::isSDWAOperand(const SDNode *N) const {
-  if (!Subtarget->hasSDWA())
-    return false;
-
-  if (N->getOpcode() == ISD::SIGN_EXTEND_INREG) {
-    EVT VT = cast<VTSDNode>(N->getOperand(1))->getVT();
-    return VT.getScalarSizeInBits() == 8 || VT.getScalarSizeInBits() == 16;
-  }
-
-  if (N->getOpcode() == ISD::AND)
-    if (auto *RHS = dyn_cast<ConstantSDNode>(N->getOperand(1)))
-      return RHS->getZExtValue() == 0xFF || RHS->getZExtValue() == 0xFFFF;
-
-  if (N->getOpcode() == ISD::SRA || N->getOpcode() == ISD::SRL)
-    if (auto *RHS = dyn_cast<ConstantSDNode>(N->getOperand(1)))
-      return (RHS->getZExtValue() % 8) == 0;
-
-  return false;
-}
-
 bool AMDGPUDAGToDAGISel::isUniformBr(const SDNode *N) const {
   const BasicBlock *BB = FuncInfo->MBB->getBasicBlock();
   const Instruction *Term = BB->getTerminator();
diff --git a/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.h b/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.h
index 0a3631e4dbf59..8138a89f9b3a9 100644
--- a/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.h
+++ b/llvm/lib/Target/AMDGPU/AMDGPUISelDAGToDAG.h
@@ -74,7 +74,6 @@ class AMDGPUDAGToDAGISel : public SelectionDAGISel {
 protected:
   void SelectBuildVector(SDNode *N, unsigned RegClassID);
   void SelectVectorShuffle(SDNode *N);
-  bool isSDWAOperand(const SDNode *N) const;
 
 private:
   std::pair<SDValue, SDValue> foldFrameIndex(SDValue N) const;
diff --git a/llvm/lib/Target/AMDGPU/AMDGPUInstructions.td b/llvm/lib/Target/AMDGPU/AMDGPUInstructions.td
index 529b2990f9b3a..4d8bdaf1b8287 100644
--- a/llvm/lib/Target/AMDGPU/AMDGPUInstructions.td
+++ b/llvm/lib/Target/AMDGPU/AMDGPUInstructions.td
@@ -230,11 +230,48 @@ class BinOp_no_sdwa<SDPatternOperator binop> : PatFrag<
   (ops node:$lhs, node:$rhs),
   (binop node:$lhs, node:$rhs),
   [{
+    auto isSDWAOperand = [&](const SDNode *N) -> bool {
+      if (!Subtarget->hasSDWA())
+        return false;
+
+      if (N->getOpcode() == ISD::SIGN_EXTEND_INREG) {
+        EVT VT = cast<VTSDNode>(N->getOperand(1))->getVT();
+        return VT.getScalarSizeInBits() == 8 ||
+               VT.getScalarSizeInBits() == 16;
+      }
+
+      if (N->getOpcode() == ISD::AND) {
+        if (const auto *RHS = dyn_cast<ConstantSDNode>(N->getOperand(1)))
+          return RHS->getZExtValue() == 0xFF ||
+                 RHS->getZExtValue() == 0xFFFF;
+      }
+
+      if (N->getOpcode() == ISD::SRA || N->getOpcode() == ISD::SRL) {
+        if (const auto *RHS = dyn_cast<ConstantSDNode>(N->getOperand(1)))
+          return (RHS->getZExtValue() % 8) == 0;
+      }
+
+      return false;
+    };
+
     return !isSDWAOperand(Op.getOperand(0).getNode()) &&
            !isSDWAOperand(Op.getOperand(1).getNode());
   }]> {
   let GISelPredicateCode = [{
-    return true; // TODO
+    auto isSDWAOperand = [&](const MachineInstr& mi) -> bool {
+        if (!Subtarget->hasSDWA())
+                return false;
+        if (mi->getOpcode() == AMDGPU::G_SEXT_INREG) {
+            const MachineOperand &Op = mi.getOperand(1);
+            if (!Op.isImm())
+                return false;
+
+            unsigned BitWidth = Op.getImm();
+            return BitWidth == 8 || BitWidth == 16;
+        }
+        //TODO:  by @Elio
+    }
+    return isSDWAOperand(MI);
   }];
 }
 



More information about the llvm-commits mailing list