[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