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

Arseniy Obolenskiy via llvm-commits llvm-commits at lists.llvm.org
Fri Sep 11 06:31:23 PDT 2026


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

>From a99a0e2bb55cf7690ae6c3e232cb9f049cb7b5e2 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 8 Sep 2026 12:57:04 +0200
Subject: [PATCH 1/7] [DAGCombiner] Share byte provider steps with AMDGPU

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
---
 llvm/include/llvm/CodeGen/ByteProvider.h      | 49 +++++++++++--
 llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 62 +++++-----------
 llvm/lib/Target/AMDGPU/SIISelLowering.cpp     | 70 +++++--------------
 3 files changed, 79 insertions(+), 102 deletions(-)

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: {

>From cb24d35bb85968e77526565dcacae0479288116d Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 8 Sep 2026 16:37:28 +0200
Subject: [PATCH 2/7] simplify

---
 llvm/include/llvm/CodeGen/ByteProvider.h | 23 +++++++++++++----------
 1 file changed, 13 insertions(+), 10 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index a3d7bb830fefa..fd73a3d758276 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -16,6 +16,7 @@
 #define LLVM_CODEGEN_BYTEPROVIDER_H
 
 #include "llvm/ADT/STLExtras.h"
+#include "llvm/CodeGen/SelectionDAGNodes.h"
 #include "llvm/Support/DataTypes.h"
 #include <optional>
 #include <type_traits>
@@ -76,15 +77,18 @@ template <typename ISelOp> class ByteProvider {
   }
 };
 
+using SDByteProviderRecurseFn =
+    function_ref<std::optional<ByteProvider<SDValue>>(SDValue, unsigned)>;
+
 /// 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);
+inline std::optional<ByteProvider<SDValue>>
+calculateByteProviderForOr(SDValue Op, unsigned Index,
+                           SDByteProviderRecurseFn Recurse) {
+  std::optional<ByteProvider<SDValue>> LHS = Recurse(Op.getOperand(0), Index);
   if (!LHS)
     return std::nullopt;
-  std::optional<ByteProvider<ISelOp>> RHS = Recurse(Op.getOperand(1), Index);
+  std::optional<ByteProvider<SDValue>> RHS = Recurse(Op.getOperand(1), Index);
   if (!RHS)
     return std::nullopt;
 
@@ -99,18 +103,17 @@ calculateByteProviderForOr(ISelOp Op, unsigned Index, RecurseT Recurse) {
 
 /// \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,
+inline std::optional<ByteProvider<SDValue>>
+calculateByteProviderForExtend(SDValue Op, unsigned Index,
                                unsigned NarrowBitWidth, bool ZeroFills,
-                               RecurseT Recurse) {
+                               SDByteProviderRecurseFn Recurse) {
   if (NarrowBitWidth % 8 != 0)
     return std::nullopt;
 
   if (Index >= NarrowBitWidth / 8) {
     if (!ZeroFills)
       return std::nullopt;
-    return ByteProvider<ISelOp>::getConstantZero();
+    return ByteProvider<SDValue>::getConstantZero();
   }
   return Recurse(Op.getOperand(0), Index);
 }

>From dba46a772743aff05f235c7939bcf70c9cc4b615 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 8 Sep 2026 19:36:12 +0200
Subject: [PATCH 3/7] rm comment

---
 llvm/include/llvm/CodeGen/ByteProvider.h | 3 +--
 1 file changed, 1 insertion(+), 2 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index fd73a3d758276..5d6788d5be160 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -48,8 +48,7 @@ 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 and SrcOffset are producer defined, see DAGCombiner.cpp and
-  // SIISelLowering.cpp.
+  // DestOffset and SrcOffset are producer defined, see DAGCombiner.cpp.
   int64_t DestOffset = 0;
   int64_t SrcOffset = 0;
 

>From 5652e2589dfe459db1c38ed2fab6cbaa1d95937d Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Thu, 10 Sep 2026 16:53:07 +0200
Subject: [PATCH 4/7] address comments

---
 llvm/include/llvm/CodeGen/ByteProvider.h      | 40 +++++-----------
 llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 17 ++++---
 llvm/lib/Target/AMDGPU/SIISelLowering.cpp     | 47 +++++++++----------
 3 files changed, 43 insertions(+), 61 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index 5d6788d5be160..3159e8aa72c84 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -15,11 +15,9 @@
 #ifndef LLVM_CODEGEN_BYTEPROVIDER_H
 #define LLVM_CODEGEN_BYTEPROVIDER_H
 
-#include "llvm/ADT/STLExtras.h"
+#include "llvm/ADT/STLFunctionalExtras.h"
 #include "llvm/CodeGen/SelectionDAGNodes.h"
-#include "llvm/Support/DataTypes.h"
 #include <optional>
-#include <type_traits>
 
 namespace llvm {
 
@@ -28,41 +26,29 @@ namespace llvm {
 /// some other productive instruction (e.g. arithmetic instructions).
 /// Bit manipulation instructions like shifts are not ByteProviders, rather
 /// are used to extract Bytes.
-template <typename ISelOp> class ByteProvider {
+class ByteProvider {
 private:
-  ByteProvider(std::optional<ISelOp> Src, int64_t DestOffset, int64_t SrcOffset)
+  ByteProvider(std::optional<SDValue> Src, int64_t DestOffset,
+               int64_t SrcOffset)
       : Src(Src), DestOffset(DestOffset), SrcOffset(SrcOffset) {}
 
-  // TODO -- use constraint in c++20
-  // Does this type correspond with an operation in selection DAG
-  // Only allow classes with member function getOpcode
-  template <typename U>
-  using check_has_getOpcode =
-      decltype(std::declval<std::remove_pointer_t<U> &>().getOpcode());
-
-  template <typename U>
-  static constexpr bool has_getOpcode =
-      is_detected<check_has_getOpcode, U>::value;
-
 public:
   // 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;
+  std::optional<SDValue> Src = std::nullopt;
   // DestOffset and SrcOffset are producer defined, see DAGCombiner.cpp.
   int64_t DestOffset = 0;
   int64_t SrcOffset = 0;
 
   ByteProvider() = default;
 
-  static ByteProvider getSrc(std::optional<ISelOp> Val, int64_t ByteOffset,
+  static ByteProvider getSrc(std::optional<SDValue> Val, int64_t ByteOffset,
                              int64_t VectorOffset) {
-    static_assert(has_getOpcode<ISelOp>,
-                  "ByteProviders must contain an operation in selection DAG.");
     return ByteProvider(Val, ByteOffset, VectorOffset);
   }
 
   static ByteProvider getConstantZero() {
-    return ByteProvider<ISelOp>(std::nullopt, 0, 0);
+    return ByteProvider(std::nullopt, 0, 0);
   }
   bool isConstantZero() const { return !Src; }
 
@@ -77,17 +63,17 @@ template <typename ISelOp> class ByteProvider {
 };
 
 using SDByteProviderRecurseFn =
-    function_ref<std::optional<ByteProvider<SDValue>>(SDValue, unsigned)>;
+    function_ref<std::optional<ByteProvider>(SDValue, unsigned)>;
 
 /// Visits both operands even once one answers, because \p Recurse may have
 /// side effects (DAGCombiner accumulates an and mask there).
-inline std::optional<ByteProvider<SDValue>>
+inline std::optional<ByteProvider>
 calculateByteProviderForOr(SDValue Op, unsigned Index,
                            SDByteProviderRecurseFn Recurse) {
-  std::optional<ByteProvider<SDValue>> LHS = Recurse(Op.getOperand(0), Index);
+  std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
   if (!LHS)
     return std::nullopt;
-  std::optional<ByteProvider<SDValue>> RHS = Recurse(Op.getOperand(1), Index);
+  std::optional<ByteProvider> RHS = Recurse(Op.getOperand(1), Index);
   if (!RHS)
     return std::nullopt;
 
@@ -102,7 +88,7 @@ calculateByteProviderForOr(SDValue Op, unsigned Index,
 
 /// \p NarrowBitWidth is a parameter because it is not always the operand
 /// width, for instance sign_extend_inreg takes it from the VTSDNode.
-inline std::optional<ByteProvider<SDValue>>
+inline std::optional<ByteProvider>
 calculateByteProviderForExtend(SDValue Op, unsigned Index,
                                unsigned NarrowBitWidth, bool ZeroFills,
                                SDByteProviderRecurseFn Recurse) {
@@ -112,7 +98,7 @@ calculateByteProviderForExtend(SDValue Op, unsigned Index,
   if (Index >= NarrowBitWidth / 8) {
     if (!ZeroFills)
       return std::nullopt;
-    return ByteProvider<SDValue>::getConstantZero();
+    return ByteProvider::getConstantZero();
   }
   return Recurse(Op.getOperand(0), Index);
 }
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index fcc88929c3549..61a33c8431009 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -9706,9 +9706,8 @@ SDValue DAGCombiner::MatchRotate(SDValue LHS, SDValue RHS, const SDLoc &DL,
 ///                 LOAD
 ///
 /// *ExtractVectorElement
-using SDByteProvider = ByteProvider<SDValue>;
 
-static std::optional<SDByteProvider>
+static std::optional<ByteProvider>
 calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
                       std::optional<uint64_t> VectorIndex,
                       unsigned StartingIndex = 0,
@@ -9759,7 +9758,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     // provide, then do not provide anything. Otherwise, subtract the index by
     // the amount we shifted by.
     return Index < ByteShift
-               ? SDByteProvider::getConstantZero()
+               ? ByteProvider::getConstantZero()
                : calculateByteProvider(Op->getOperand(0), Index - ByteShift,
                                        Depth + 1, VectorIndex, Index, ByteMask);
   }
@@ -9782,7 +9781,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
         MaskOp->getAPIntValue().extractBitsAsZExtValue(8, Index * 8);
 
     if (MaskByte == 0x00)
-      return SDByteProvider::getConstantZero();
+      return ByteProvider::getConstantZero();
 
     auto Result = Recurse(Op->getOperand(0), Index);
     if (!Result)
@@ -9841,12 +9840,12 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     // question
     if (Index >= NarrowByteWidth)
       return L->getExtensionType() == ISD::ZEXTLOAD
-                 ? std::optional<SDByteProvider>(
-                       SDByteProvider::getConstantZero())
+                 ? std::optional<ByteProvider>(
+                       ByteProvider::getConstantZero())
                  : std::nullopt;
 
     unsigned BPVectorIndex = VectorIndex.value_or(0U);
-    return SDByteProvider::getSrc(Op, Index, BPVectorIndex);
+    return ByteProvider::getSrc(Op, Index, BPVectorIndex);
   }
   }
 
@@ -10152,7 +10151,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
   unsigned ByteWidth = VT.getSizeInBits() / 8;
 
   bool IsBigEndianTarget = DAG.getDataLayout().isBigEndian();
-  auto MemoryByteOffset = [&](SDByteProvider P) {
+  auto MemoryByteOffset = [&](ByteProvider P) {
     assert(P.hasSrc() && "Must be a memory byte provider");
     auto *Load = cast<LoadSDNode>(*P.Src);
 
@@ -10169,7 +10168,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
   SDValue Chain;
 
   SmallPtrSet<LoadSDNode *, 8> Loads;
-  std::optional<SDByteProvider> FirstByteProvider;
+  std::optional<ByteProvider> FirstByteProvider;
   int64_t FirstOffset = INT64_MAX;
 
   // Check if all the bytes of the OR we are looking at are loaded from the same
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index 969933560c45d..7e1643e4dc3d4 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -15101,9 +15101,10 @@ SDValue SITargetLowering::performAndCombine(SDNode *N,
 // ultimately provides. \p SrcIndex is the byte of the src that maps to this
 // dest of the or byte. \p Depth tracks how many recursive iterations we have
 // performed.
-static const std::optional<ByteProvider<SDValue>>
-calculateSrcByte(const SDValue Op, uint64_t DestByte, uint64_t SrcIndex = 0,
-                 unsigned Depth = 0) {
+static const std::optional<ByteProvider> calculateSrcByte(const SDValue Op,
+                                                          uint64_t DestByte,
+                                                          uint64_t SrcIndex = 0,
+                                                          unsigned Depth = 0) {
   // We may need to recursively traverse a series of SRLs
   if (Depth >= 6)
     return std::nullopt;
@@ -15112,7 +15113,7 @@ calculateSrcByte(const SDValue Op, uint64_t DestByte, uint64_t SrcIndex = 0,
     return std::nullopt;
 
   if (Op.getValueType().isVector())
-    return ByteProvider<SDValue>::getSrc(Op, DestByte, SrcIndex);
+    return ByteProvider::getSrc(Op, DestByte, SrcIndex);
 
   switch (Op->getOpcode()) {
   case ISD::TRUNCATE: {
@@ -15158,7 +15159,7 @@ calculateSrcByte(const SDValue Op, uint64_t DestByte, uint64_t SrcIndex = 0,
   }
 
   default: {
-    return ByteProvider<SDValue>::getSrc(Op, DestByte, SrcIndex);
+    return ByteProvider::getSrc(Op, DestByte, SrcIndex);
   }
   }
   llvm_unreachable("fully handled switch");
@@ -15170,7 +15171,7 @@ calculateSrcByte(const SDValue Op, uint64_t DestByte, uint64_t SrcIndex = 0,
 // the byte position of the Op that corresponds with the originally requested
 // byte of the Or \p Depth tracks how many recursive iterations we have
 // performed. \p StartingIndex is the originally requested byte of the Or
-static const std::optional<ByteProvider<SDValue>>
+static const std::optional<ByteProvider>
 calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
                       unsigned StartingIndex = 0) {
   // Finding Src tree of RHS of or typically requires at least 1 additional
@@ -15212,7 +15213,7 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
       // is not well formatted
       if (IndexMask & BitMask)
         return std::nullopt;
-      return ByteProvider<SDValue>::getConstantZero();
+      return ByteProvider::getConstantZero();
     }
 
     return calculateSrcByte(Op->getOperand(0), StartingIndex, Index);
@@ -15270,7 +15271,7 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     // SRA's out-of-range bytes are sign bits, not constant zero.
     if (Op.getOpcode() == ISD::SRA)
       return std::nullopt;
-    return ByteProvider<SDValue>::getConstantZero();
+    return ByteProvider::getConstantZero();
   }
 
   case ISD::SHL: {
@@ -15290,7 +15291,7 @@ 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()
+    return Index < ByteShift ? ByteProvider::getConstantZero()
                              : Recurse(Op.getOperand(0), Index - ByteShift);
   }
   case ISD::ANY_EXTEND:
@@ -15341,8 +15342,7 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     // question
     if (Index >= NarrowByteWidth) {
       return L->getExtensionType() == ISD::ZEXTLOAD
-                 ? std::optional<ByteProvider<SDValue>>(
-                       ByteProvider<SDValue>::getConstantZero())
+                 ? std::optional<ByteProvider>(ByteProvider::getConstantZero())
                  : std::nullopt;
     }
 
@@ -15385,8 +15385,7 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     auto NextIndex = IdxMask > 0x03 ? IdxMask % 4 : IdxMask;
 
     return IdxMask != 0x0c ? calculateSrcByte(NextOp, StartingIndex, NextIndex)
-                           : ByteProvider<SDValue>(
-                                 ByteProvider<SDValue>::getConstantZero());
+                           : ByteProvider::getConstantZero();
   }
 
   default: {
@@ -15531,13 +15530,13 @@ static SDValue getDWordFromOffset(SelectionDAG &DAG, SDLoc SL, SDValue Src,
 static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
   SelectionDAG &DAG = DCI.DAG;
   [[maybe_unused]] EVT VT = N->getValueType(0);
-  SmallVector<ByteProvider<SDValue>, 8> PermNodes;
+  SmallVector<ByteProvider, 8> PermNodes;
 
   // VT is known to be MVT::i32, so we need to provide 4 bytes.
   assert(VT == MVT::i32);
   for (int i = 0; i < 4; i++) {
     // Find the ByteProvider that provides the ith byte of the result of OR
-    std::optional<ByteProvider<SDValue>> P =
+    std::optional<ByteProvider> P =
         calculateByteProvider(SDValue(N, 0), i, 0, /*StartingIndex = */ i);
     // TODO support constantZero
     if (!P || P->isConstantZero())
@@ -15905,13 +15904,13 @@ SITargetLowering::performZeroOrAnyExtendCombine(SDNode *N,
   // possible we're missing out on some combine opportunities, but we'd need to
   // weigh the cost of extracting the byte from the upper dwords.
 
-  std::optional<ByteProvider<SDValue>> BP0 =
+  std::optional<ByteProvider> BP0 =
       calculateByteProvider(SDValue(N, 0), 0, 0, 0);
   if (!BP0 || BP0->SrcOffset >= 4 || !BP0->Src)
     return SDValue();
   SDValue V0 = *BP0->Src;
 
-  std::optional<ByteProvider<SDValue>> BP1 =
+  std::optional<ByteProvider> BP1 =
       calculateByteProvider(SDValue(N, 0), 1, 0, 1);
   if (!BP1 || BP1->SrcOffset >= 4 || !BP1->Src)
     return SDValue();
@@ -17454,8 +17453,7 @@ SITargetLowering::foldAddSub64WithZeroLowBitsTo32(SDNode *N,
 
 // Collect the ultimate src of each of the mul node's operands, and confirm
 // each operand is 8 bytes.
-static std::optional<ByteProvider<SDValue>>
-handleMulOperand(const SDValue &MulOperand) {
+static std::optional<ByteProvider> handleMulOperand(const SDValue &MulOperand) {
   auto Byte0 = calculateByteProvider(MulOperand, 0, 0);
   if (!Byte0 || Byte0->isConstantZero()) {
     return std::nullopt;
@@ -17487,8 +17485,7 @@ struct DotSrc {
   int64_t DWordOffset;
 };
 
-static void placeSources(ByteProvider<SDValue> &Src0,
-                         ByteProvider<SDValue> &Src1,
+static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
                          SmallVectorImpl<DotSrc> &Src0s,
                          SmallVectorImpl<DotSrc> &Src1s, int Step) {
 
@@ -17503,7 +17500,7 @@ static void placeSources(ByteProvider<SDValue> &Src0,
   }
 
   for (int BPI = 0; BPI < 2; BPI++) {
-    std::pair<ByteProvider<SDValue>, ByteProvider<SDValue>> BPP = {Src0, Src1};
+    std::pair<ByteProvider, ByteProvider> BPP = {Src0, Src1};
     if (BPI == 1) {
       BPP = {Src1, Src0};
     }
@@ -17647,9 +17644,9 @@ static bool isMul(const SDValue Op) {
 }
 
 static std::optional<bool>
-checkDot4MulSignedness(const SDValue &N, ByteProvider<SDValue> &Src0,
-                       ByteProvider<SDValue> &Src1, const SDValue &S0Op,
-                       const SDValue &S1Op, const SelectionDAG &DAG) {
+checkDot4MulSignedness(const SDValue &N, ByteProvider &Src0, ByteProvider &Src1,
+                       const SDValue &S0Op, const SDValue &S1Op,
+                       const SelectionDAG &DAG) {
   // If we both ops are i8s (pre legalize-dag), then the signedness semantics
   // of the dot4 is irrelevant.
   if (S0Op.getValueSizeInBits() == 8 && S1Op.getValueSizeInBits() == 8)

>From 6761bcdaf5f5316b0ecef7218313c1877cf39e20 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Thu, 10 Sep 2026 17:43:50 +0200
Subject: [PATCH 5/7] comments

---
 llvm/include/llvm/CodeGen/ByteProvider.h      | 72 +++++++------------
 llvm/lib/CodeGen/ByteProvider.cpp             | 53 ++++++++++++++
 llvm/lib/CodeGen/CMakeLists.txt               |  1 +
 llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp |  9 ++-
 llvm/lib/Target/AMDGPU/SIISelLowering.cpp     | 31 ++++----
 5 files changed, 100 insertions(+), 66 deletions(-)
 create mode 100644 llvm/lib/CodeGen/ByteProvider.cpp

diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index 3159e8aa72c84..e6ac9cb8d73b8 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -16,11 +16,13 @@
 #define LLVM_CODEGEN_BYTEPROVIDER_H
 
 #include "llvm/ADT/STLFunctionalExtras.h"
-#include "llvm/CodeGen/SelectionDAGNodes.h"
 #include <optional>
 
 namespace llvm {
 
+class SDNode;
+class SDValue;
+
 /// Represents known origin of an individual byte in combine pattern. The
 /// value of the byte is either constant zero, or comes from memory /
 /// some other productive instruction (e.g. arithmetic instructions).
@@ -28,36 +30,39 @@ namespace llvm {
 /// are used to extract Bytes.
 class ByteProvider {
 private:
-  ByteProvider(std::optional<SDValue> Src, int64_t DestOffset,
+  ByteProvider(SDNode *Node, unsigned ResNo, int64_t DestOffset,
                int64_t SrcOffset)
-      : Src(Src), DestOffset(DestOffset), SrcOffset(SrcOffset) {}
+      : Node(Node), ResNo(ResNo), DestOffset(DestOffset), SrcOffset(SrcOffset) {
+  }
 
 public:
-  // For constant zero providers Src is set to nullopt. For actual providers
-  // Src represents the node which originally produced the relevant bits.
-  std::optional<SDValue> Src = std::nullopt;
+  // For constant zero providers Node is null. For actual providers Node and
+  // ResNo represent the SDValue which originally produced the relevant bits.
+  SDNode *Node = nullptr;
+  unsigned ResNo = 0;
   // DestOffset and SrcOffset are producer defined, see DAGCombiner.cpp.
   int64_t DestOffset = 0;
   int64_t SrcOffset = 0;
 
   ByteProvider() = default;
 
-  static ByteProvider getSrc(std::optional<SDValue> Val, int64_t ByteOffset,
-                             int64_t VectorOffset) {
-    return ByteProvider(Val, ByteOffset, VectorOffset);
-  }
+  static ByteProvider getSrc(SDValue Val, int64_t ByteOffset,
+                             int64_t VectorOffset);
 
-  static ByteProvider getConstantZero() {
-    return ByteProvider(std::nullopt, 0, 0);
-  }
-  bool isConstantZero() const { return !Src; }
+  static ByteProvider getConstantZero() { return ByteProvider(); }
+  bool isConstantZero() const { return !Node; }
 
-  bool hasSrc() const { return Src.has_value(); }
+  bool hasSrc() const { return Node != nullptr; }
 
-  bool hasSameSrc(const ByteProvider &Other) const { return Other.Src == Src; }
+  /// Returns the SDValue this byte comes from. Only valid if hasSrc().
+  SDValue getSrc() const;
+
+  bool hasSameSrc(const ByteProvider &Other) const {
+    return Other.Node == Node && Other.ResNo == ResNo;
+  }
 
   bool operator==(const ByteProvider &Other) const {
-    return Other.Src == Src && Other.DestOffset == DestOffset &&
+    return hasSameSrc(Other) && Other.DestOffset == DestOffset &&
            Other.SrcOffset == SrcOffset;
   }
 };
@@ -67,41 +72,16 @@ using SDByteProviderRecurseFn =
 
 /// Visits both operands even once one answers, because \p Recurse may have
 /// side effects (DAGCombiner accumulates an and mask there).
-inline std::optional<ByteProvider>
+std::optional<ByteProvider>
 calculateByteProviderForOr(SDValue Op, unsigned Index,
-                           SDByteProviderRecurseFn Recurse) {
-  std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
-  if (!LHS)
-    return std::nullopt;
-  std::optional<ByteProvider> 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;
-}
+                           SDByteProviderRecurseFn Recurse);
 
 /// \p NarrowBitWidth is a parameter because it is not always the operand
 /// width, for instance sign_extend_inreg takes it from the VTSDNode.
-inline std::optional<ByteProvider>
+std::optional<ByteProvider>
 calculateByteProviderForExtend(SDValue Op, unsigned Index,
                                unsigned NarrowBitWidth, bool ZeroFills,
-                               SDByteProviderRecurseFn Recurse) {
-  if (NarrowBitWidth % 8 != 0)
-    return std::nullopt;
-
-  if (Index >= NarrowBitWidth / 8) {
-    if (!ZeroFills)
-      return std::nullopt;
-    return ByteProvider::getConstantZero();
-  }
-  return Recurse(Op.getOperand(0), Index);
-}
+                               SDByteProviderRecurseFn Recurse);
 
 } // end namespace llvm
 
diff --git a/llvm/lib/CodeGen/ByteProvider.cpp b/llvm/lib/CodeGen/ByteProvider.cpp
new file mode 100644
index 0000000000000..383fca0bce0e2
--- /dev/null
+++ b/llvm/lib/CodeGen/ByteProvider.cpp
@@ -0,0 +1,53 @@
+//===-- ByteProvider.cpp -------------------------------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "llvm/CodeGen/ByteProvider.h"
+#include "llvm/CodeGen/SelectionDAGNodes.h"
+
+using namespace llvm;
+
+ByteProvider ByteProvider::getSrc(SDValue Val, int64_t ByteOffset,
+                                  int64_t VectorOffset) {
+  return ByteProvider(Val.getNode(), Val.getResNo(), ByteOffset, VectorOffset);
+}
+
+SDValue ByteProvider::getSrc() const { return SDValue(Node, ResNo); }
+
+std::optional<ByteProvider>
+llvm::calculateByteProviderForOr(SDValue Op, unsigned Index,
+                                 SDByteProviderRecurseFn Recurse) {
+  std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
+  if (!LHS)
+    return std::nullopt;
+  std::optional<ByteProvider> 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;
+}
+
+std::optional<ByteProvider>
+llvm::calculateByteProviderForExtend(SDValue Op, unsigned Index,
+                                     unsigned NarrowBitWidth, bool ZeroFills,
+                                     SDByteProviderRecurseFn Recurse) {
+  if (NarrowBitWidth % 8 != 0)
+    return std::nullopt;
+
+  if (Index >= NarrowBitWidth / 8) {
+    if (!ZeroFills)
+      return std::nullopt;
+    return ByteProvider::getConstantZero();
+  }
+  return Recurse(Op.getOperand(0), Index);
+}
diff --git a/llvm/lib/CodeGen/CMakeLists.txt b/llvm/lib/CodeGen/CMakeLists.txt
index 99dfb4bb09df7..c068d7d1878e2 100644
--- a/llvm/lib/CodeGen/CMakeLists.txt
+++ b/llvm/lib/CodeGen/CMakeLists.txt
@@ -43,6 +43,7 @@ add_llvm_component_library(LLVMCodeGen
   BranchFolding.cpp
   BranchRelaxation.cpp
   BreakFalseDeps.cpp
+  ByteProvider.cpp
   BasicBlockSections.cpp
   BasicBlockPathCloning.cpp
   BasicBlockSectionsProfileReader.cpp
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 61a33c8431009..94fc4dbf43a5a 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -9840,8 +9840,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     // question
     if (Index >= NarrowByteWidth)
       return L->getExtensionType() == ISD::ZEXTLOAD
-                 ? std::optional<ByteProvider>(
-                       ByteProvider::getConstantZero())
+                 ? std::optional<ByteProvider>(ByteProvider::getConstantZero())
                  : std::nullopt;
 
     unsigned BPVectorIndex = VectorIndex.value_or(0U);
@@ -10153,7 +10152,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
   bool IsBigEndianTarget = DAG.getDataLayout().isBigEndian();
   auto MemoryByteOffset = [&](ByteProvider P) {
     assert(P.hasSrc() && "Must be a memory byte provider");
-    auto *Load = cast<LoadSDNode>(*P.Src);
+    auto *Load = cast<LoadSDNode>(P.getSrc());
 
     unsigned LoadBitWidth = Load->getMemoryVT().getScalarSizeInBits();
 
@@ -10191,7 +10190,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
       continue;
     }
     assert(P->hasSrc() && "provenance should either be memory or zero");
-    auto *L = cast<LoadSDNode>(*P->Src);
+    auto *L = cast<LoadSDNode>(P->getSrc());
 
     // All loads must share the same chain
     SDValue LChain = L->getChain();
@@ -10264,7 +10263,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);
+  auto *FirstLoad = cast<LoadSDNode>(FirstByteProvider->getSrc());
 
   // 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 7e1643e4dc3d4..97464c061f4cb 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -15567,7 +15567,7 @@ static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
 
       // Set the index of the second distinct Src node
       SecondSrc = {i, PermNodes[i].SrcOffset / 4};
-      assert(!(PermNodes[SecondSrc->first].Src->getValueSizeInBits() % 8));
+      assert(!(PermNodes[SecondSrc->first].getSrc().getValueSizeInBits() % 8));
       SrcByteAdjust = 0;
     }
     assert((PermOp.SrcOffset % 4) + SrcByteAdjust < 8);
@@ -15575,7 +15575,7 @@ static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
     PermMask |= ((PermOp.SrcOffset % 4) + SrcByteAdjust) << (i * 8);
   }
   SDLoc DL(N);
-  SDValue Op = *PermNodes[FirstSrc.first].Src;
+  SDValue Op = PermNodes[FirstSrc.first].getSrc();
   Op = getDWordFromOffset(DAG, DL, Op, FirstSrc.second);
   assert(Op.getValueSizeInBits() == 32);
 
@@ -15592,7 +15592,7 @@ static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
       return DAG.getBitcast(MVT::getIntegerVT(32), Op);
   }
 
-  SDValue OtherOp = SecondSrc ? *PermNodes[SecondSrc->first].Src : Op;
+  SDValue OtherOp = SecondSrc ? PermNodes[SecondSrc->first].getSrc() : Op;
 
   if (SecondSrc) {
     OtherOp = getDWordFromOffset(DAG, DL, OtherOp, SecondSrc->second);
@@ -15906,16 +15906,16 @@ SITargetLowering::performZeroOrAnyExtendCombine(SDNode *N,
 
   std::optional<ByteProvider> BP0 =
       calculateByteProvider(SDValue(N, 0), 0, 0, 0);
-  if (!BP0 || BP0->SrcOffset >= 4 || !BP0->Src)
+  if (!BP0 || BP0->SrcOffset >= 4 || !BP0->hasSrc())
     return SDValue();
-  SDValue V0 = *BP0->Src;
+  SDValue V0 = BP0->getSrc();
 
   std::optional<ByteProvider> BP1 =
       calculateByteProvider(SDValue(N, 0), 1, 0, 1);
-  if (!BP1 || BP1->SrcOffset >= 4 || !BP1->Src)
+  if (!BP1 || BP1->SrcOffset >= 4 || !BP1->hasSrc())
     return SDValue();
 
-  SDValue V1 = *BP1->Src;
+  SDValue V1 = BP1->getSrc();
 
   if (V0 == V1)
     return SDValue();
@@ -17489,12 +17489,12 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
                          SmallVectorImpl<DotSrc> &Src0s,
                          SmallVectorImpl<DotSrc> &Src1s, int Step) {
 
-  assert(Src0.Src.has_value() && Src1.Src.has_value());
+  assert(Src0.hasSrc() && Src1.hasSrc());
   // Src0s and Src1s are empty, just place arbitrarily.
   if (Step == 0) {
-    Src0s.push_back({*Src0.Src, ((Src0.SrcOffset % 4) << 24) + 0x0c0c0c,
+    Src0s.push_back({Src0.getSrc(), ((Src0.SrcOffset % 4) << 24) + 0x0c0c0c,
                      Src0.SrcOffset / 4});
-    Src1s.push_back({*Src1.Src, ((Src1.SrcOffset % 4) << 24) + 0x0c0c0c,
+    Src1s.push_back({Src1.getSrc(), ((Src1.SrcOffset % 4) << 24) + 0x0c0c0c,
                      Src1.SrcOffset / 4});
     return;
   }
@@ -17518,7 +17518,7 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
     for (int I = 0; I < 2; I++) {
       SmallVectorImpl<DotSrc> &Srcs = I == 0 ? Src0s : Src1s;
       auto MatchesFirst = [&BPP](DotSrc &IterElt) {
-        return IterElt.SrcOp == *BPP.first.Src &&
+        return IterElt.SrcOp == BPP.first.getSrc() &&
                (IterElt.DWordOffset == (BPP.first.SrcOffset / 4));
       };
 
@@ -17532,14 +17532,15 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
     if (FirstGroup != -1) {
       SmallVectorImpl<DotSrc> &Srcs = FirstGroup == 1 ? Src0s : Src1s;
       auto MatchesSecond = [&BPP](DotSrc &IterElt) {
-        return IterElt.SrcOp == *BPP.second.Src &&
+        return IterElt.SrcOp == BPP.second.getSrc() &&
                (IterElt.DWordOffset == (BPP.second.SrcOffset / 4));
       };
       auto *Match = llvm::find_if(Srcs, MatchesSecond);
       if (Match != Srcs.end()) {
         Match->PermMask = addPermMasks(SecondMask, Match->PermMask);
       } else
-        Srcs.push_back({*BPP.second.Src, SecondMask, BPP.second.SrcOffset / 4});
+        Srcs.push_back(
+            {BPP.second.getSrc(), SecondMask, BPP.second.SrcOffset / 4});
       return;
     }
   }
@@ -17551,11 +17552,11 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
   unsigned FMask = 0xFF << (8 * (3 - Step));
 
   Src0s.push_back(
-      {*Src0.Src,
+      {Src0.getSrc(),
        ((Src0.SrcOffset % 4) << (8 * (3 - Step)) | (ZeroMask & ~FMask)),
        Src0.SrcOffset / 4});
   Src1s.push_back(
-      {*Src1.Src,
+      {Src1.getSrc(),
        ((Src1.SrcOffset % 4) << (8 * (3 - Step)) | (ZeroMask & ~FMask)),
        Src1.SrcOffset / 4});
 }

>From df93ab2864c505d0b4b41c6776c8b403e05187c6 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Fri, 11 Sep 2026 13:37:43 +0200
Subject: [PATCH 6/7] address comments

---
 llvm/include/llvm/CodeGen/ByteProvider.h      | 32 ++++-------
 llvm/lib/CodeGen/ByteProvider.cpp             | 14 +----
 llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp |  6 +-
 llvm/lib/Target/AMDGPU/SIISelLowering.cpp     | 55 +++++++++++--------
 4 files changed, 49 insertions(+), 58 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index e6ac9cb8d73b8..0492d07f2a03d 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -16,13 +16,11 @@
 #define LLVM_CODEGEN_BYTEPROVIDER_H
 
 #include "llvm/ADT/STLFunctionalExtras.h"
+#include "llvm/CodeGen/SelectionDAGNodes.h"
 #include <optional>
 
 namespace llvm {
 
-class SDNode;
-class SDValue;
-
 /// Represents known origin of an individual byte in combine pattern. The
 /// value of the byte is either constant zero, or comes from memory /
 /// some other productive instruction (e.g. arithmetic instructions).
@@ -30,16 +28,13 @@ class SDValue;
 /// are used to extract Bytes.
 class ByteProvider {
 private:
-  ByteProvider(SDNode *Node, unsigned ResNo, int64_t DestOffset,
-               int64_t SrcOffset)
-      : Node(Node), ResNo(ResNo), DestOffset(DestOffset), SrcOffset(SrcOffset) {
-  }
+  ByteProvider(SDValue Src, int64_t DestOffset, int64_t SrcOffset)
+      : Src(Src), DestOffset(DestOffset), SrcOffset(SrcOffset) {}
 
 public:
-  // For constant zero providers Node is null. For actual providers Node and
-  // ResNo represent the SDValue which originally produced the relevant bits.
-  SDNode *Node = nullptr;
-  unsigned ResNo = 0;
+  // For constant zero providers Src is null. For actual providers Src is the
+  // value which originally produced the relevant bits.
+  SDValue Src;
   // DestOffset and SrcOffset are producer defined, see DAGCombiner.cpp.
   int64_t DestOffset = 0;
   int64_t SrcOffset = 0;
@@ -47,19 +42,16 @@ class ByteProvider {
   ByteProvider() = default;
 
   static ByteProvider getSrc(SDValue Val, int64_t ByteOffset,
-                             int64_t VectorOffset);
+                             int64_t VectorOffset) {
+    return ByteProvider(Val, ByteOffset, VectorOffset);
+  }
 
   static ByteProvider getConstantZero() { return ByteProvider(); }
-  bool isConstantZero() const { return !Node; }
+  bool isConstantZero() const { return !Src; }
 
-  bool hasSrc() const { return Node != nullptr; }
+  bool hasSrc() const { return static_cast<bool>(Src); }
 
-  /// Returns the SDValue this byte comes from. Only valid if hasSrc().
-  SDValue getSrc() const;
-
-  bool hasSameSrc(const ByteProvider &Other) const {
-    return Other.Node == Node && Other.ResNo == ResNo;
-  }
+  bool hasSameSrc(const ByteProvider &Other) const { return Other.Src == Src; }
 
   bool operator==(const ByteProvider &Other) const {
     return hasSameSrc(Other) && Other.DestOffset == DestOffset &&
diff --git a/llvm/lib/CodeGen/ByteProvider.cpp b/llvm/lib/CodeGen/ByteProvider.cpp
index 383fca0bce0e2..2a1b951378021 100644
--- a/llvm/lib/CodeGen/ByteProvider.cpp
+++ b/llvm/lib/CodeGen/ByteProvider.cpp
@@ -7,26 +7,18 @@
 //===----------------------------------------------------------------------===//
 
 #include "llvm/CodeGen/ByteProvider.h"
-#include "llvm/CodeGen/SelectionDAGNodes.h"
 
 using namespace llvm;
 
-ByteProvider ByteProvider::getSrc(SDValue Val, int64_t ByteOffset,
-                                  int64_t VectorOffset) {
-  return ByteProvider(Val.getNode(), Val.getResNo(), ByteOffset, VectorOffset);
-}
-
-SDValue ByteProvider::getSrc() const { return SDValue(Node, ResNo); }
-
 std::optional<ByteProvider>
 llvm::calculateByteProviderForOr(SDValue Op, unsigned Index,
                                  SDByteProviderRecurseFn Recurse) {
-  std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
-  if (!LHS)
-    return std::nullopt;
   std::optional<ByteProvider> RHS = Recurse(Op.getOperand(1), Index);
   if (!RHS)
     return std::nullopt;
+  std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
+  if (!LHS)
+    return std::nullopt;
 
   // A well formed or has two ByteProviders for each byte, one of which is
   // constant zero.
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 94fc4dbf43a5a..b5d5ddfcc47e7 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -10152,7 +10152,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
   bool IsBigEndianTarget = DAG.getDataLayout().isBigEndian();
   auto MemoryByteOffset = [&](ByteProvider P) {
     assert(P.hasSrc() && "Must be a memory byte provider");
-    auto *Load = cast<LoadSDNode>(P.getSrc());
+    auto *Load = cast<LoadSDNode>(P.Src);
 
     unsigned LoadBitWidth = Load->getMemoryVT().getScalarSizeInBits();
 
@@ -10190,7 +10190,7 @@ SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
       continue;
     }
     assert(P->hasSrc() && "provenance should either be memory or zero");
-    auto *L = cast<LoadSDNode>(P->getSrc());
+    auto *L = cast<LoadSDNode>(P->Src);
 
     // All loads must share the same chain
     SDValue LChain = L->getChain();
@@ -10263,7 +10263,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->getSrc());
+  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 97464c061f4cb..cfd372c81f3bd 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -15185,16 +15185,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:
     if (IsVec)
       return std::nullopt;
-    return calculateByteProviderForOr(Op, Index, Recurse);
+    return calculateByteProviderForOr(
+        Op, Index, [&](SDValue NextOp, unsigned NextIndex) {
+          return calculateByteProvider(NextOp, NextIndex, Depth + 1,
+                                       StartingIndex);
+        });
 
   case ISD::AND: {
     if (IsVec)
@@ -15243,7 +15243,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 Recurse(NextOp, NewIndex);
+    return calculateByteProvider(NextOp, NewIndex, Depth + 1, StartingIndex);
   }
 
   case ISD::SRA:
@@ -15291,8 +15291,10 @@ 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::getConstantZero()
-                             : Recurse(Op.getOperand(0), Index - ByteShift);
+    if (Index < ByteShift)
+      return ByteProvider::getConstantZero();
+    return calculateByteProvider(Op.getOperand(0), Index - ByteShift, Depth + 1,
+                                 StartingIndex);
   }
   case ISD::ANY_EXTEND:
   case ISD::SIGN_EXTEND:
@@ -15311,7 +15313,11 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
       NarrowBitWidth = VTSign->getVT().getSizeInBits();
     }
     return calculateByteProviderForExtend(
-        Op, Index, NarrowBitWidth, Op.getOpcode() == ISD::ZERO_EXTEND, Recurse);
+        Op, Index, NarrowBitWidth, Op.getOpcode() == ISD::ZERO_EXTEND,
+        [&](SDValue NextOp, unsigned NextIndex) {
+          return calculateByteProvider(NextOp, NextIndex, Depth + 1,
+                                       StartingIndex);
+        });
   }
 
   case ISD::TRUNCATE: {
@@ -15319,7 +15325,8 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
       return std::nullopt;
 
     // Index is already bounded by BitWidth / 8 above.
-    return Recurse(Op.getOperand(0), Index);
+    return calculateByteProvider(Op.getOperand(0), Index, Depth + 1,
+                                 StartingIndex);
   }
 
   case ISD::CopyFromReg: {
@@ -15353,7 +15360,8 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     if (IsVec)
       return std::nullopt;
 
-    return Recurse(Op->getOperand(0), BitWidth / 8 - Index - 1);
+    return calculateByteProvider(Op->getOperand(0), BitWidth / 8 - Index - 1,
+                                 Depth + 1, StartingIndex);
   }
 
   case ISD::EXTRACT_VECTOR_ELT: {
@@ -15567,7 +15575,7 @@ static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
 
       // Set the index of the second distinct Src node
       SecondSrc = {i, PermNodes[i].SrcOffset / 4};
-      assert(!(PermNodes[SecondSrc->first].getSrc().getValueSizeInBits() % 8));
+      assert(!(PermNodes[SecondSrc->first].Src.getValueSizeInBits() % 8));
       SrcByteAdjust = 0;
     }
     assert((PermOp.SrcOffset % 4) + SrcByteAdjust < 8);
@@ -15575,7 +15583,7 @@ static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
     PermMask |= ((PermOp.SrcOffset % 4) + SrcByteAdjust) << (i * 8);
   }
   SDLoc DL(N);
-  SDValue Op = PermNodes[FirstSrc.first].getSrc();
+  SDValue Op = PermNodes[FirstSrc.first].Src;
   Op = getDWordFromOffset(DAG, DL, Op, FirstSrc.second);
   assert(Op.getValueSizeInBits() == 32);
 
@@ -15592,7 +15600,7 @@ static SDValue matchPERM(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
       return DAG.getBitcast(MVT::getIntegerVT(32), Op);
   }
 
-  SDValue OtherOp = SecondSrc ? PermNodes[SecondSrc->first].getSrc() : Op;
+  SDValue OtherOp = SecondSrc ? PermNodes[SecondSrc->first].Src : Op;
 
   if (SecondSrc) {
     OtherOp = getDWordFromOffset(DAG, DL, OtherOp, SecondSrc->second);
@@ -15908,14 +15916,14 @@ SITargetLowering::performZeroOrAnyExtendCombine(SDNode *N,
       calculateByteProvider(SDValue(N, 0), 0, 0, 0);
   if (!BP0 || BP0->SrcOffset >= 4 || !BP0->hasSrc())
     return SDValue();
-  SDValue V0 = BP0->getSrc();
+  SDValue V0 = BP0->Src;
 
   std::optional<ByteProvider> BP1 =
       calculateByteProvider(SDValue(N, 0), 1, 0, 1);
   if (!BP1 || BP1->SrcOffset >= 4 || !BP1->hasSrc())
     return SDValue();
 
-  SDValue V1 = BP1->getSrc();
+  SDValue V1 = BP1->Src;
 
   if (V0 == V1)
     return SDValue();
@@ -17492,9 +17500,9 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
   assert(Src0.hasSrc() && Src1.hasSrc());
   // Src0s and Src1s are empty, just place arbitrarily.
   if (Step == 0) {
-    Src0s.push_back({Src0.getSrc(), ((Src0.SrcOffset % 4) << 24) + 0x0c0c0c,
+    Src0s.push_back({Src0.Src, ((Src0.SrcOffset % 4) << 24) + 0x0c0c0c,
                      Src0.SrcOffset / 4});
-    Src1s.push_back({Src1.getSrc(), ((Src1.SrcOffset % 4) << 24) + 0x0c0c0c,
+    Src1s.push_back({Src1.Src, ((Src1.SrcOffset % 4) << 24) + 0x0c0c0c,
                      Src1.SrcOffset / 4});
     return;
   }
@@ -17518,7 +17526,7 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
     for (int I = 0; I < 2; I++) {
       SmallVectorImpl<DotSrc> &Srcs = I == 0 ? Src0s : Src1s;
       auto MatchesFirst = [&BPP](DotSrc &IterElt) {
-        return IterElt.SrcOp == BPP.first.getSrc() &&
+        return IterElt.SrcOp == BPP.first.Src &&
                (IterElt.DWordOffset == (BPP.first.SrcOffset / 4));
       };
 
@@ -17532,15 +17540,14 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
     if (FirstGroup != -1) {
       SmallVectorImpl<DotSrc> &Srcs = FirstGroup == 1 ? Src0s : Src1s;
       auto MatchesSecond = [&BPP](DotSrc &IterElt) {
-        return IterElt.SrcOp == BPP.second.getSrc() &&
+        return IterElt.SrcOp == BPP.second.Src &&
                (IterElt.DWordOffset == (BPP.second.SrcOffset / 4));
       };
       auto *Match = llvm::find_if(Srcs, MatchesSecond);
       if (Match != Srcs.end()) {
         Match->PermMask = addPermMasks(SecondMask, Match->PermMask);
       } else
-        Srcs.push_back(
-            {BPP.second.getSrc(), SecondMask, BPP.second.SrcOffset / 4});
+        Srcs.push_back({BPP.second.Src, SecondMask, BPP.second.SrcOffset / 4});
       return;
     }
   }
@@ -17552,11 +17559,11 @@ static void placeSources(ByteProvider &Src0, ByteProvider &Src1,
   unsigned FMask = 0xFF << (8 * (3 - Step));
 
   Src0s.push_back(
-      {Src0.getSrc(),
+      {Src0.Src,
        ((Src0.SrcOffset % 4) << (8 * (3 - Step)) | (ZeroMask & ~FMask)),
        Src0.SrcOffset / 4});
   Src1s.push_back(
-      {Src1.getSrc(),
+      {Src1.Src,
        ((Src1.SrcOffset % 4) << (8 * (3 - Step)) | (ZeroMask & ~FMask)),
        Src1.SrcOffset / 4});
 }

>From 6d7399003a8e0f71eb8ea77b38e71af353bb94fe Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Fri, 11 Sep 2026 15:31:08 +0200
Subject: [PATCH 7/7] rework  byte provider callbacks

---
 llvm/include/llvm/CodeGen/ByteProvider.h      | 28 ++++----
 llvm/lib/CodeGen/ByteProvider.cpp             | 32 +++------
 llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 50 ++++++++------
 llvm/lib/Target/AMDGPU/SIISelLowering.cpp     | 66 ++++++++++---------
 4 files changed, 85 insertions(+), 91 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/ByteProvider.h b/llvm/include/llvm/CodeGen/ByteProvider.h
index 0492d07f2a03d..c0d66b1652cd8 100644
--- a/llvm/include/llvm/CodeGen/ByteProvider.h
+++ b/llvm/include/llvm/CodeGen/ByteProvider.h
@@ -15,7 +15,6 @@
 #ifndef LLVM_CODEGEN_BYTEPROVIDER_H
 #define LLVM_CODEGEN_BYTEPROVIDER_H
 
-#include "llvm/ADT/STLFunctionalExtras.h"
 #include "llvm/CodeGen/SelectionDAGNodes.h"
 #include <optional>
 
@@ -35,9 +34,8 @@ class ByteProvider {
   // For constant zero providers Src is null. For actual providers Src is the
   // value which originally produced the relevant bits.
   SDValue Src;
-  // DestOffset and SrcOffset are producer defined, see DAGCombiner.cpp.
-  int64_t DestOffset = 0;
-  int64_t SrcOffset = 0;
+  int64_t DestOffset = 0; // Load byte in DAGCombiner, unused in AMDGPU.
+  int64_t SrcOffset = 0;  // Vector lane in DAGCombiner, byte in Src in AMDGPU.
 
   ByteProvider() = default;
 
@@ -59,21 +57,17 @@ class ByteProvider {
   }
 };
 
-using SDByteProviderRecurseFn =
-    function_ref<std::optional<ByteProvider>(SDValue, unsigned)>;
-
-/// Visits both operands even once one answers, because \p Recurse may have
-/// side effects (DAGCombiner accumulates an and mask there).
+/// In a well formed or, one of the two byte providers is constant zero.
 std::optional<ByteProvider>
-calculateByteProviderForOr(SDValue Op, unsigned Index,
-                           SDByteProviderRecurseFn Recurse);
+selectOrByteProvider(const std::optional<ByteProvider> &LHS,
+                     const std::optional<ByteProvider> &RHS);
 
-/// \p NarrowBitWidth is a parameter because it is not always the operand
-/// width, for instance sign_extend_inreg takes it from the VTSDNode.
-std::optional<ByteProvider>
-calculateByteProviderForExtend(SDValue Op, unsigned Index,
-                               unsigned NarrowBitWidth, bool ZeroFills,
-                               SDByteProviderRecurseFn Recurse);
+enum class NarrowByteAction { Unknown, ConstantZero, FromNarrow };
+
+/// FromNarrow keeps \p Index. \p NarrowBitWidth is not always the operand
+/// width, sign_extend_inreg takes it from the VTSDNode.
+NarrowByteAction classifyNarrowByte(unsigned Index, unsigned NarrowBitWidth,
+                                    bool ZeroFills);
 
 } // end namespace llvm
 
diff --git a/llvm/lib/CodeGen/ByteProvider.cpp b/llvm/lib/CodeGen/ByteProvider.cpp
index 2a1b951378021..83b9e5fe8f25f 100644
--- a/llvm/lib/CodeGen/ByteProvider.cpp
+++ b/llvm/lib/CodeGen/ByteProvider.cpp
@@ -11,17 +11,10 @@
 using namespace llvm;
 
 std::optional<ByteProvider>
-llvm::calculateByteProviderForOr(SDValue Op, unsigned Index,
-                                 SDByteProviderRecurseFn Recurse) {
-  std::optional<ByteProvider> RHS = Recurse(Op.getOperand(1), Index);
-  if (!RHS)
+llvm::selectOrByteProvider(const std::optional<ByteProvider> &LHS,
+                           const std::optional<ByteProvider> &RHS) {
+  if (!LHS || !RHS)
     return std::nullopt;
-  std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
-  if (!LHS)
-    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())
@@ -29,17 +22,12 @@ llvm::calculateByteProviderForOr(SDValue Op, unsigned Index,
   return std::nullopt;
 }
 
-std::optional<ByteProvider>
-llvm::calculateByteProviderForExtend(SDValue Op, unsigned Index,
-                                     unsigned NarrowBitWidth, bool ZeroFills,
-                                     SDByteProviderRecurseFn Recurse) {
+NarrowByteAction llvm::classifyNarrowByte(unsigned Index,
+                                          unsigned NarrowBitWidth,
+                                          bool ZeroFills) {
   if (NarrowBitWidth % 8 != 0)
-    return std::nullopt;
-
-  if (Index >= NarrowBitWidth / 8) {
-    if (!ZeroFills)
-      return std::nullopt;
-    return ByteProvider::getConstantZero();
-  }
-  return Recurse(Op.getOperand(0), Index);
+    return NarrowByteAction::Unknown;
+  if (Index < NarrowBitWidth / 8)
+    return NarrowByteAction::FromNarrow;
+  return ZeroFills ? NarrowByteAction::ConstantZero : NarrowByteAction::Unknown;
 }
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index b5d5ddfcc47e7..56f0811659963 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -9733,7 +9733,7 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     return std::nullopt;
   unsigned ByteWidth = BitWidth / 8;
   assert(Index < ByteWidth && "invalid index requested");
-  (void) ByteWidth;
+  (void)ByteWidth;
 
   auto Recurse = [&](SDValue NextOp, unsigned NextIndex) {
     return calculateByteProvider(NextOp, NextIndex, Depth + 1, VectorIndex,
@@ -9741,8 +9741,12 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
   };
 
   switch (Op.getOpcode()) {
-  case ISD::OR:
-    return calculateByteProviderForOr(Op, Index, Recurse);
+  case ISD::OR: {
+    std::optional<ByteProvider> LHS = Recurse(Op.getOperand(0), Index);
+    if (!LHS)
+      return std::nullopt;
+    return selectOrByteProvider(LHS, Recurse(Op.getOperand(1), Index));
+  }
   case ISD::SHL: {
     auto ShiftOp = dyn_cast<ConstantSDNode>(Op->getOperand(1));
     if (!ShiftOp)
@@ -9764,10 +9768,19 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
   }
   case ISD::ANY_EXTEND:
   case ISD::SIGN_EXTEND:
-  case ISD::ZERO_EXTEND:
-    return calculateByteProviderForExtend(
-        Op, Index, Op->getOperand(0).getScalarValueSizeInBits(),
-        Op.getOpcode() == ISD::ZERO_EXTEND, Recurse);
+  case ISD::ZERO_EXTEND: {
+    SDValue NarrowOp = Op->getOperand(0);
+    switch (classifyNarrowByte(Index, NarrowOp.getScalarValueSizeInBits(),
+                               Op.getOpcode() == ISD::ZERO_EXTEND)) {
+    case NarrowByteAction::Unknown:
+      return std::nullopt;
+    case NarrowByteAction::ConstantZero:
+      return ByteProvider::getConstantZero();
+    case NarrowByteAction::FromNarrow:
+      return Recurse(NarrowOp, Index);
+    }
+    llvm_unreachable("fully handled switch");
+  }
   case ISD::BSWAP:
     return Recurse(Op->getOperand(0), ByteWidth - Index - 1);
   case ISD::AND: {
@@ -9830,21 +9843,16 @@ calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
     if (!L->isSimple() || L->isIndexed())
       return std::nullopt;
 
-    unsigned NarrowBitWidth = L->getMemoryVT().getScalarSizeInBits();
-    if (NarrowBitWidth % 8 != 0)
+    switch (classifyNarrowByte(Index, L->getMemoryVT().getScalarSizeInBits(),
+                               L->getExtensionType() == ISD::ZEXTLOAD)) {
+    case NarrowByteAction::Unknown:
       return std::nullopt;
-    uint64_t NarrowByteWidth = NarrowBitWidth / 8;
-
-    // If the width of the load does not reach byte we are trying to provide for
-    // and it is not a ZEXTLOAD, then the load does not provide for the byte in
-    // question
-    if (Index >= NarrowByteWidth)
-      return L->getExtensionType() == ISD::ZEXTLOAD
-                 ? std::optional<ByteProvider>(ByteProvider::getConstantZero())
-                 : std::nullopt;
-
-    unsigned BPVectorIndex = VectorIndex.value_or(0U);
-    return ByteProvider::getSrc(Op, Index, BPVectorIndex);
+    case NarrowByteAction::ConstantZero:
+      return ByteProvider::getConstantZero();
+    case NarrowByteAction::FromNarrow:
+      return ByteProvider::getSrc(Op, Index, VectorIndex.value_or(0U));
+    }
+    llvm_unreachable("fully handled switch");
   }
   }
 
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index cfd372c81f3bd..af8c559ccb800 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -15101,10 +15101,10 @@ SDValue SITargetLowering::performAndCombine(SDNode *N,
 // ultimately provides. \p SrcIndex is the byte of the src that maps to this
 // dest of the or byte. \p Depth tracks how many recursive iterations we have
 // performed.
-static const std::optional<ByteProvider> calculateSrcByte(const SDValue Op,
-                                                          uint64_t DestByte,
-                                                          uint64_t SrcIndex = 0,
-                                                          unsigned Depth = 0) {
+static std::optional<ByteProvider> calculateSrcByte(const SDValue Op,
+                                                    uint64_t DestByte,
+                                                    uint64_t SrcIndex = 0,
+                                                    unsigned Depth = 0) {
   // We may need to recursively traverse a series of SRLs
   if (Depth >= 6)
     return std::nullopt;
@@ -15171,7 +15171,7 @@ static const std::optional<ByteProvider> calculateSrcByte(const SDValue Op,
 // the byte position of the Op that corresponds with the originally requested
 // byte of the Or \p Depth tracks how many recursive iterations we have
 // performed. \p StartingIndex is the originally requested byte of the Or
-static const std::optional<ByteProvider>
+static std::optional<ByteProvider>
 calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
                       unsigned StartingIndex = 0) {
   // Finding Src tree of RHS of or typically requires at least 1 additional
@@ -15187,14 +15187,18 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
 
   bool IsVec = Op.getValueType().isVector();
   switch (Op.getOpcode()) {
-  case ISD::OR:
+  case ISD::OR: {
     if (IsVec)
       return std::nullopt;
-    return calculateByteProviderForOr(
-        Op, Index, [&](SDValue NextOp, unsigned NextIndex) {
-          return calculateByteProvider(NextOp, NextIndex, Depth + 1,
-                                       StartingIndex);
-        });
+
+    std::optional<ByteProvider> RHS = calculateByteProvider(
+        Op.getOperand(1), Index, Depth + 1, StartingIndex);
+    if (!RHS)
+      return std::nullopt;
+    return selectOrByteProvider(calculateByteProvider(Op.getOperand(0), Index,
+                                                      Depth + 1, StartingIndex),
+                                RHS);
+  }
 
   case ISD::AND: {
     if (IsVec)
@@ -15305,19 +15309,24 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
     if (IsVec)
       return std::nullopt;
 
-    unsigned NarrowBitWidth = Op->getOperand(0).getValueSizeInBits();
+    SDValue NarrowOp = Op->getOperand(0);
+    unsigned NarrowBitWidth = NarrowOp.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();
     }
-    return calculateByteProviderForExtend(
-        Op, Index, NarrowBitWidth, Op.getOpcode() == ISD::ZERO_EXTEND,
-        [&](SDValue NextOp, unsigned NextIndex) {
-          return calculateByteProvider(NextOp, NextIndex, Depth + 1,
-                                       StartingIndex);
-        });
+    switch (classifyNarrowByte(Index, NarrowBitWidth,
+                               Op.getOpcode() == ISD::ZERO_EXTEND)) {
+    case NarrowByteAction::Unknown:
+      return std::nullopt;
+    case NarrowByteAction::ConstantZero:
+      return ByteProvider::getConstantZero();
+    case NarrowByteAction::FromNarrow:
+      return calculateByteProvider(NarrowOp, Index, Depth + 1, StartingIndex);
+    }
+    llvm_unreachable("fully handled switch");
   }
 
   case ISD::TRUNCATE: {
@@ -15339,21 +15348,16 @@ calculateByteProvider(const SDValue &Op, unsigned Index, unsigned Depth,
   case ISD::LOAD: {
     auto *L = cast<LoadSDNode>(Op.getNode());
 
-    unsigned NarrowBitWidth = L->getMemoryVT().getSizeInBits();
-    if (NarrowBitWidth % 8 != 0)
+    switch (classifyNarrowByte(Index, L->getMemoryVT().getSizeInBits(),
+                               L->getExtensionType() == ISD::ZEXTLOAD)) {
+    case NarrowByteAction::Unknown:
       return std::nullopt;
-    uint64_t NarrowByteWidth = NarrowBitWidth / 8;
-
-    // If the width of the load does not reach byte we are trying to provide for
-    // and it is not a ZEXTLOAD, then the load does not provide for the byte in
-    // question
-    if (Index >= NarrowByteWidth) {
-      return L->getExtensionType() == ISD::ZEXTLOAD
-                 ? std::optional<ByteProvider>(ByteProvider::getConstantZero())
-                 : std::nullopt;
+    case NarrowByteAction::ConstantZero:
+      return ByteProvider::getConstantZero();
+    case NarrowByteAction::FromNarrow:
+      return calculateSrcByte(Op, StartingIndex, Index);
     }
-
-    return calculateSrcByte(Op, StartingIndex, Index);
+    llvm_unreachable("fully handled switch");
   }
 
   case ISD::BSWAP: {



More information about the llvm-commits mailing list