[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