[llvm] [DAGCombiner] Share byte provider steps with AMDGPU target (PR #221959)
Arseniy Obolenskiy via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 10 07:53:20 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/4] [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/4] 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/4] 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/4] 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)
More information about the llvm-commits
mailing list