[llvm] [DAGCombiner] Share byte provider steps with AMDGPU target (PR #221959)

via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 8 04:06:58 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-amdgpu

Author: Arseniy Obolenskiy (aobolensk)

<details>
<summary>Changes</summary>

DAGCombiner and AMDGPU each had their own copy of the same or/extend logic

Move the shared parts to one place and use one ByteProvider type

discussed in https://github.com/llvm/llvm-project/pull/220169#pullrequestreview-5122750345 

---
Full diff: https://github.com/llvm/llvm-project/pull/221959.diff


3 Files Affected:

- (modified) llvm/include/llvm/CodeGen/ByteProvider.h (+43-6) 
- (modified) llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp (+19-43) 
- (modified) llvm/lib/Target/AMDGPU/SIISelLowering.cpp (+17-53) 


``````````diff
diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index c00335a216458..a3d7bb830fefa 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -8,9 +8,7 @@
 //
 // \file
 // This file implements ByteProvider. The purpose of ByteProvider is to provide
-// a map between a target node's byte (byte position is DestOffset) and the
-// source (and byte position) that provides it (in Src and SrcOffset
-// respectively) See CodeGen/SelectionDAG/DAGCombiner.cpp MatchLoadCombine
+// a map between a byte of a target node and the source that provides it.
 //
 //===----------------------------------------------------------------------===//
 
@@ -49,10 +47,9 @@ template <typename ISelOp> class ByteProvider {
   // For constant zero providers Src is set to nullopt. For actual providers
   // Src represents the node which originally produced the relevant bits.
   std::optional<ISelOp> Src = std::nullopt;
-  // DestOffset is the offset of the byte in the dest we are trying to map for.
+  // DestOffset and SrcOffset are producer defined, see DAGCombiner.cpp and
+  // SIISelLowering.cpp.
   int64_t DestOffset = 0;
-  // SrcOffset is the offset in the ultimate source node that maps to the
-  // DestOffset
   int64_t SrcOffset = 0;
 
   ByteProvider() = default;
@@ -78,6 +75,46 @@ template <typename ISelOp> class ByteProvider {
            Other.SrcOffset == SrcOffset;
   }
 };
+
+/// Visits both operands even once one answers, because \p Recurse may have
+/// side effects (DAGCombiner accumulates an and mask there).
+template <typename ISelOp, typename RecurseT>
+std::optional<ByteProvider<ISelOp>>
+calculateByteProviderForOr(ISelOp Op, unsigned Index, RecurseT Recurse) {
+  std::optional<ByteProvider<ISelOp>> LHS = Recurse(Op.getOperand(0), Index);
+  if (!LHS)
+    return std::nullopt;
+  std::optional<ByteProvider<ISelOp>> RHS = Recurse(Op.getOperand(1), Index);
+  if (!RHS)
+    return std::nullopt;
+
+  // A well formed or has two ByteProviders for each byte, one of which is
+  // constant zero.
+  if (LHS->isConstantZero())
+    return RHS;
+  if (RHS->isConstantZero())
+    return LHS;
+  return std::nullopt;
+}
+
+/// \p NarrowBitWidth is a parameter because it is not always the operand
+/// width, for instance sign_extend_inreg takes it from the VTSDNode.
+template <typename ISelOp, typename RecurseT>
+std::optional<ByteProvider<ISelOp>>
+calculateByteProviderForExtend(ISelOp Op, unsigned Index,
+                               unsigned NarrowBitWidth, bool ZeroFills,
+                               RecurseT Recurse) {
+  if (NarrowBitWidth % 8 != 0)
+    return std::nullopt;
+
+  if (Index >= NarrowBitWidth / 8) {
+    if (!ZeroFills)
+      return std::nullopt;
+    return ByteProvider<ISelOp>::getConstantZero();
+  }
+  return Recurse(Op.getOperand(0), Index);
+}
+
 } // end namespace llvm
 
 #endif // LLVM_CODEGEN_BYTEPROVIDER_H
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 733d0eb9baa40..fcc88929c3549 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -9706,7 +9706,7 @@ SDValue DAGCombiner::MatchRotate(SDValue LHS, SDValue RHS, const SDLoc &DL,
 ///                 LOAD
 ///
 /// *ExtractVectorElement
-using SDByteProvider = ByteProvider<SDNode *>;
+using SDByteProvider = ByteProvider<SDValue>;
 
 static std::optional<SDByteProvider>
 calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
@@ -9736,23 +9736,14 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
   assert(Index < ByteWidth && "invalid index requested");
   (void) ByteWidth;
 
-  switch (Op.getOpcode()) {
-  case ISD::OR: {
-    auto LHS = calculateByteProvider(Op->getOperand(0), Index, Depth + 1,
-                                     VectorIndex, StartingIndex, ByteMask);
-    if (!LHS)
-      return std::nullopt;
-    auto RHS = calculateByteProvider(Op->getOperand(1), Index, Depth + 1,
-                                     VectorIndex, StartingIndex, ByteMask);
-    if (!RHS)
-      return std::nullopt;
+  auto Recurse = [&](SDValue NextOp, unsigned NextIndex) {
+    return calculateByteProvider(NextOp, NextIndex, Depth + 1, VectorIndex,
+                                 StartingIndex, ByteMask);
+  };
 
-    if (LHS->isConstantZero())
-      return RHS;
-    if (RHS->isConstantZero())
-      return LHS;
-    return std::nullopt;
-  }
+  switch (Op.getOpcode()) {
+  case ISD::OR:
+    return calculateByteProviderForOr(Op, Index, Recurse);
   case ISD::SHL: {
     auto ShiftOp = dyn_cast<ConstantSDNode>(Op->getOperand(1));
     if (!ShiftOp)
@@ -9774,25 +9765,12 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
   }
   case ISD::ANY_EXTEND:
   case ISD::SIGN_EXTEND:
-  case ISD::ZERO_EXTEND: {
-    SDValue NarrowOp = Op->getOperand(0);
-    unsigned NarrowBitWidth = NarrowOp.getScalarValueSizeInBits();
-    if (NarrowBitWidth % 8 != 0)
-      return std::nullopt;
-    uint64_t NarrowByteWidth = NarrowBitWidth / 8;
-
-    if (Index >= NarrowByteWidth)
-      return Op.getOpcode() == ISD::ZERO_EXTEND
-                 ? std::optional<SDByteProvider>(
-                       SDByteProvider::getConstantZero())
-                 : std::nullopt;
-    return calculateByteProvider(NarrowOp, Index, Depth + 1, VectorIndex,
-                                 StartingIndex, ByteMask);
-  }
+  case ISD::ZERO_EXTEND:
+    return calculateByteProviderForExtend(
+        Op, Index, Op->getOperand(0).getScalarValueSizeInBits(),
+        Op.getOpcode() == ISD::ZERO_EXTEND, Recurse);
   case ISD::BSWAP:
-    return calculateByteProvider(Op->getOperand(0), ByteWidth - Index - 1,
-                                 Depth + 1, VectorIndex, StartingIndex,
-                                 ByteMask);
+    return Recurse(Op->getOperand(0), ByteWidth - Index - 1);
   case ISD::AND: {
     // Constants are canonicalized to the RHS of AND, so only operand 1 needs
     // to be checked.
@@ -9806,8 +9784,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     if (MaskByte == 0x00)
       return SDByteProvider::getConstantZero();
 
-    auto Result = calculateByteProvider(Op->getOperand(0), Index, Depth + 1,
-                                        VectorIndex, StartingIndex, ByteMask);
+    auto Result = Recurse(Op->getOperand(0), Index);
     if (!Result)
       return std::nullopt;
 
@@ -9847,8 +9824,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     if ((*VectorIndex + 1) * NarrowByteWidth <= StartingIndex)
       return std::nullopt;
 
-    return calculateByteProvider(Op->getOperand(0), Index, Depth + 1,
-                                 VectorIndex, StartingIndex, ByteMask);
+    return Recurse(Op->getOperand(0), Index);
   }
   case ISD::LOAD: {
     auto L = cast<LoadSDNode>(Op.getNode());
@@ -9870,7 +9846,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
                  : std::nullopt;
 
     unsigned BPVectorIndex = VectorIndex.value_or(0U);
-    return SDByteProvider::getSrc(L, Index, BPVectorIndex);
+    return SDByteProvider::getSrc(Op, Index, BPVectorIndex);
   }
   }
 
@@ -10178,7 +10154,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
   bool IsBigEndianTarget = DAG.getDataLayout().isBigEndian();
   auto MemoryByteOffset = [&](SDByteProvider P) {
     assert(P.hasSrc() && "Must be a memory byte provider");
-    auto *Load = cast<LoadSDNode>(P.Src.value());
+    auto *Load = cast<LoadSDNode>(*P.Src);
 
     unsigned LoadBitWidth = Load->getMemoryVT().getScalarSizeInBits();
 
@@ -10216,7 +10192,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
       continue;
     }
     assert(P->hasSrc() && "provenance should either be memory or zero");
-    auto *L = cast<LoadSDNode>(P->Src.value());
+    auto *L = cast<LoadSDNode>(*P->Src);
 
     // All loads must share the same chain
     SDValue LChain = L->getChain();
@@ -10289,7 +10265,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
   // So the combined value can be loaded from the first load address.
   if (MemoryByteOffset(*FirstByteProvider) != 0)
     return SDValue();
-  auto *FirstLoad = cast<LoadSDNode>(FirstByteProvider->Src.value());
+  auto *FirstLoad = cast<LoadSDNode>(*FirstByteProvider->Src);
 
   // Before legalization we allow introducing loads that are wider than legal,
   // which will later be split into legally sized loads. This enables us to
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index 3bad6d330c0ad..969933560c45d 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -15184,30 +15184,16 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
   if (Index > BitWidth / 8 - 1)
     return std::nullopt;
 
+  auto Recurse = [&](SDValue NextOp, unsigned NextIndex) {
+    return calculateByteProvider(NextOp, NextIndex, Depth + 1, StartingIndex);
+  };
+
   bool IsVec = Op.getValueType().isVector();
   switch (Op.getOpcode()) {
-  case ISD::OR: {
+  case ISD::OR:
     if (IsVec)
       return std::nullopt;
-
-    auto RHS = calculateByteProvider(Op.getOperand(1), Index, Depth + 1,
-                                     StartingIndex);
-    if (!RHS)
-      return std::nullopt;
-    auto LHS = calculateByteProvider(Op.getOperand(0), Index, Depth + 1,
-                                     StartingIndex);
-    if (!LHS)
-      return std::nullopt;
-    // A well formed Or will have two ByteProviders for each byte, one of which
-    // is constant zero
-    if (!LHS->isConstantZero() && !RHS->isConstantZero())
-      return std::nullopt;
-    if (!LHS || LHS->isConstantZero())
-      return RHS;
-    if (!RHS || RHS->isConstantZero())
-      return LHS;
-    return std::nullopt;
-  }
+    return calculateByteProviderForOr(Op, Index, Recurse);
 
   case ISD::AND: {
     if (IsVec)
@@ -15238,7 +15224,7 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
 
     // fshr(X,Y,Z): (X << (BW - (Z % BW))) | (Y >> (Z % BW))
     auto *ShiftOp = dyn_cast<ConstantSDNode>(Op->getOperand(2));
-    if (!ShiftOp || Op.getValueType().isVector())
+    if (!ShiftOp)
       return std::nullopt;
 
     uint64_t BitsProvided = Op.getValueSizeInBits();
@@ -15256,7 +15242,7 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     uint64_t BytesProvided = BitsProvided / 8;
     SDValue NextOp = Op.getOperand(NewIndex >= BytesProvided ? 0 : 1);
     NewIndex %= BytesProvided;
-    return calculateByteProvider(NextOp, NewIndex, Depth + 1, StartingIndex);
+    return Recurse(NextOp, NewIndex);
   }
 
   case ISD::SRA:
@@ -15304,10 +15290,8 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     // the index we are trying to provide, then it provides 0s. If not,
     // then this bytes are not definitively 0s, and the corresponding byte
     // of interest is Index - ByteShift of the src
-    return Index < ByteShift
-               ? ByteProvider<SDValue>::getConstantZero()
-               : calculateByteProvider(Op.getOperand(0), Index - ByteShift,
-                                       Depth + 1, StartingIndex);
+    return Index < ByteShift ? ByteProvider<SDValue>::getConstantZero()
+                             : Recurse(Op.getOperand(0), Index - ByteShift);
   }
   case ISD::ANY_EXTEND:
   case ISD::SIGN_EXTEND:
@@ -15318,38 +15302,23 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     if (IsVec)
       return std::nullopt;
 
-    SDValue NarrowOp = Op->getOperand(0);
-    unsigned NarrowBitWidth = NarrowOp.getValueSizeInBits();
+    unsigned NarrowBitWidth = Op->getOperand(0).getValueSizeInBits();
     if (Op->getOpcode() == ISD::SIGN_EXTEND_INREG ||
         Op->getOpcode() == ISD::AssertZext ||
         Op->getOpcode() == ISD::AssertSext) {
       auto *VTSign = cast<VTSDNode>(Op->getOperand(1));
       NarrowBitWidth = VTSign->getVT().getSizeInBits();
     }
-    if (NarrowBitWidth % 8 != 0)
-      return std::nullopt;
-    uint64_t NarrowByteWidth = NarrowBitWidth / 8;
-
-    if (Index >= NarrowByteWidth)
-      return Op.getOpcode() == ISD::ZERO_EXTEND
-                 ? std::optional<ByteProvider<SDValue>>(
-                       ByteProvider<SDValue>::getConstantZero())
-                 : std::nullopt;
-    return calculateByteProvider(NarrowOp, Index, Depth + 1, StartingIndex);
+    return calculateByteProviderForExtend(
+        Op, Index, NarrowBitWidth, Op.getOpcode() == ISD::ZERO_EXTEND, Recurse);
   }
 
   case ISD::TRUNCATE: {
     if (IsVec)
       return std::nullopt;
 
-    uint64_t NarrowByteWidth = BitWidth / 8;
-
-    if (NarrowByteWidth >= Index) {
-      return calculateByteProvider(Op.getOperand(0), Index, Depth + 1,
-                                   StartingIndex);
-    }
-
-    return std::nullopt;
+    // Index is already bounded by BitWidth / 8 above.
+    return Recurse(Op.getOperand(0), Index);
   }
 
   case ISD::CopyFromReg: {
@@ -15377,19 +15346,14 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
                  : std::nullopt;
     }
 
-    if (NarrowByteWidth > Index) {
-      return calculateSrcByte(Op, StartingIndex, Index);
-    }
-
-    return std::nullopt;
+    return calculateSrcByte(Op, StartingIndex, Index);
   }
 
   case ISD::BSWAP: {
     if (IsVec)
       return std::nullopt;
 
-    return calculateByteProvider(Op->getOperand(0), BitWidth / 8 - Index - 1,
-                                 Depth + 1, StartingIndex);
+    return Recurse(Op->getOperand(0), BitWidth / 8 - Index - 1);
   }
 
   case ISD::EXTRACT_VECTOR_ELT: {

``````````

</details>


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


More information about the llvm-commits mailing list