[llvm] Refactor LaneBitmask to be a Bitset (PR #191757)
Jiachen Yuan via llvm-commits
llvm-commits at lists.llvm.org
Mon Apr 13 09:51:19 PDT 2026
https://github.com/JiachenYuan updated https://github.com/llvm/llvm-project/pull/191757
>From 99867a0d302d772db2ef2e8e940d98d7c05a3c75 Mon Sep 17 00:00:00 2001
From: Jiachen Yuan <jiacheny at nvidia.com>
Date: Mon, 13 Apr 2026 04:12:06 +0000
Subject: [PATCH] [ADT] Refactor LaneBitmask to be a Bitset
---
llvm/include/llvm/ADT/Bitset.h | 85 ++
llvm/include/llvm/CodeGen/RDFLiveness.h | 2 +-
llvm/include/llvm/CodeGen/RDFRegisters.h | 3 +-
llvm/include/llvm/MC/LaneBitmask.h | 217 ++++-
llvm/lib/CodeGen/MIRParser/MIParser.cpp | 41 +-
llvm/lib/CodeGen/MachineOperand.cpp | 2 +-
llvm/lib/CodeGen/MachineStableHash.cpp | 6 +-
llvm/lib/CodeGen/RDFRegisters.cpp | 5 -
llvm/lib/Target/AMDGPU/SIRegisterInfo.cpp | 8 +-
llvm/lib/Target/AMDGPU/SIRegisterInfo.h | 9 +-
llvm/unittests/CodeGen/CMakeLists.txt | 3 +-
llvm/unittests/CodeGen/LaneBitmaskTest.cpp | 933 ++++++++++++++++++++
llvm/utils/TableGen/RegisterInfoEmitter.cpp | 32 +-
13 files changed, 1256 insertions(+), 90 deletions(-)
create mode 100644 llvm/unittests/CodeGen/LaneBitmaskTest.cpp
diff --git a/llvm/include/llvm/ADT/Bitset.h b/llvm/include/llvm/ADT/Bitset.h
index 9dc0f24b1d9f5..1acf457cf6265 100644
--- a/llvm/include/llvm/ADT/Bitset.h
+++ b/llvm/include/llvm/ADT/Bitset.h
@@ -16,6 +16,7 @@
#ifndef LLVM_ADT_BITSET_H
#define LLVM_ADT_BITSET_H
+#include "llvm/ADT/Hashing.h"
#include "llvm/ADT/bit.h"
#include <array>
#include <climits>
@@ -23,6 +24,10 @@
namespace llvm {
+// Forward declare Bitset and hash_value for friend declarations.
+template <unsigned NumBits> class Bitset;
+template <unsigned NumBits> hash_code hash_value(const Bitset<NumBits> &);
+
/// This is a constexpr reimplementation of a subset of std::bitset. It would be
/// nice to use std::bitset directly, but it doesn't support constant
/// initialization.
@@ -52,6 +57,8 @@ template <unsigned NumBits> class Bitset {
constexpr void maskLastWord() { Bits[getLastWordIndex()] &= RemainderMask; }
protected:
+ constexpr const StorageType &getData() const { return Bits; }
+
constexpr Bitset(const std::array<uint64_t, (NumBits + 63) / 64> &B) {
if constexpr (sizeof(BitWord) == sizeof(uint64_t)) {
for (size_t I = 0; I != B.size(); ++I)
@@ -194,8 +201,86 @@ template <unsigned NumBits> class Bitset {
}
return false;
}
+
+ constexpr Bitset &operator<<=(unsigned N) {
+ if (N == 0)
+ return *this;
+ if (N >= NumBits) {
+ return *this = Bitset();
+ }
+ const unsigned WordShift = N / BitwordBits;
+ const unsigned BitShift = N % BitwordBits;
+ if (BitShift == 0) {
+ for (int I = NumWords - 1; I >= static_cast<int>(WordShift); --I)
+ Bits[I] = Bits[I - WordShift];
+ } else {
+ const unsigned CarryShift = BitwordBits - BitShift;
+ for (int I = NumWords - 1; I > static_cast<int>(WordShift); --I) {
+ Bits[I] = (Bits[I - WordShift] << BitShift) |
+ (Bits[I - WordShift - 1] >> CarryShift);
+ }
+ Bits[WordShift] = Bits[0] << BitShift;
+ }
+ for (unsigned I = 0; I < WordShift; ++I)
+ Bits[I] = 0;
+ maskLastWord();
+ return *this;
+ }
+
+ constexpr Bitset operator<<(unsigned N) const {
+ Bitset Result(*this);
+ Result <<= N;
+ return Result;
+ }
+
+ constexpr Bitset &operator>>=(unsigned N) {
+ if (N == 0)
+ return *this;
+ if (N >= NumBits) {
+ return *this = Bitset();
+ }
+ const unsigned WordShift = N / BitwordBits;
+ const unsigned BitShift = N % BitwordBits;
+ if (BitShift == 0) {
+ for (unsigned I = 0; I < NumWords - WordShift; ++I)
+ Bits[I] = Bits[I + WordShift];
+ } else {
+ const unsigned CarryShift = BitwordBits - BitShift;
+ for (unsigned I = 0; I < NumWords - WordShift - 1; ++I) {
+ Bits[I] = (Bits[I + WordShift] >> BitShift) |
+ (Bits[I + WordShift + 1] << CarryShift);
+ }
+ Bits[NumWords - WordShift - 1] = Bits[NumWords - 1] >> BitShift;
+ }
+ for (unsigned I = NumWords - WordShift; I < NumWords; ++I)
+ Bits[I] = 0;
+ maskLastWord();
+ return *this;
+ }
+
+ constexpr Bitset operator>>(unsigned N) const {
+ Bitset Result(*this);
+ Result >>= N;
+ return Result;
+ }
+
+ friend hash_code hash_value<NumBits>(const Bitset<NumBits> &);
+ friend struct std::hash<Bitset<NumBits>>;
};
+template <unsigned NumBits>
+inline hash_code hash_value(const Bitset<NumBits> &B) {
+ return hash_combine_range(B.Bits.begin(), B.Bits.end());
+}
+
} // end namespace llvm
+namespace std {
+template <unsigned NumBits> struct hash<llvm::Bitset<NumBits>> {
+ size_t operator()(const llvm::Bitset<NumBits> &B) const {
+ return llvm::hash_combine_range(B.Bits.begin(), B.Bits.end());
+ }
+};
+} // end namespace std
+
#endif
diff --git a/llvm/include/llvm/CodeGen/RDFLiveness.h b/llvm/include/llvm/CodeGen/RDFLiveness.h
index fe1034f9b6f8e..bc78b25b36177 100644
--- a/llvm/include/llvm/CodeGen/RDFLiveness.h
+++ b/llvm/include/llvm/CodeGen/RDFLiveness.h
@@ -44,7 +44,7 @@ namespace std {
template <> struct hash<llvm::rdf::detail::NodeRef> {
std::size_t operator()(llvm::rdf::detail::NodeRef R) const {
return std::hash<llvm::rdf::NodeId>{}(R.first) ^
- std::hash<llvm::LaneBitmask::Type>{}(R.second.getAsInteger());
+ std::hash<llvm::LaneBitmask>{}(R.second);
}
};
diff --git a/llvm/include/llvm/CodeGen/RDFRegisters.h b/llvm/include/llvm/CodeGen/RDFRegisters.h
index 48e1e3487f11f..89ca56b32d3e6 100644
--- a/llvm/include/llvm/CodeGen/RDFRegisters.h
+++ b/llvm/include/llvm/CodeGen/RDFRegisters.h
@@ -125,8 +125,7 @@ struct RegisterRef {
}
size_t hash() const {
- return std::hash<RegisterId>{}(Id) ^
- std::hash<LaneBitmask::Type>{}(Mask.getAsInteger());
+ return std::hash<RegisterId>{}(Id) ^ std::hash<LaneBitmask>{}(Mask);
}
static constexpr bool isRegId(RegisterId Id) {
diff --git a/llvm/include/llvm/MC/LaneBitmask.h b/llvm/include/llvm/MC/LaneBitmask.h
index c06ca7dd5b8fc..ba2e285d2f5ae 100644
--- a/llvm/include/llvm/MC/LaneBitmask.h
+++ b/llvm/include/llvm/MC/LaneBitmask.h
@@ -29,72 +29,193 @@
#ifndef LLVM_MC_LANEBITMASK_H
#define LLVM_MC_LANEBITMASK_H
-#include "llvm/Support/Compiler.h"
+#include "llvm/ADT/APInt.h"
+#include "llvm/ADT/Bitset.h"
#include "llvm/Support/Format.h"
+#include "llvm/Support/FormatVariadic.h"
#include "llvm/Support/MathExtras.h"
#include "llvm/Support/Printable.h"
#include "llvm/Support/raw_ostream.h"
+#include <array>
+#include <cassert>
-namespace llvm {
+namespace llvm::detail {
+template <unsigned NumBits> struct LaneBitmaskImpl : public Bitset<NumBits> {
+ static constexpr unsigned BitWidth = NumBits;
- struct LaneBitmask {
- // When changing the underlying type, change the format string as well.
- using Type = uint64_t;
- enum : unsigned { BitWidth = 8*sizeof(Type) };
- constexpr static const char *const FormatStr = "%016llX";
+ constexpr LaneBitmaskImpl() = default;
+ constexpr LaneBitmaskImpl(const LaneBitmaskImpl &) = default;
+ explicit constexpr LaneBitmaskImpl(uint64_t V)
+ : Bitset<NumBits>(std::array<uint64_t, (NumBits + 63) / 64>{V}) {}
+ explicit constexpr LaneBitmaskImpl(
+ const std::array<uint64_t, (NumBits + 63) / 64> &B)
+ : Bitset<NumBits>(B) {}
+ explicit LaneBitmaskImpl(const APInt &N)
+ : Bitset<NumBits>(convertAPIntToArray(N)) {}
+ // Delete the initializer_list constructor to avoid ambiguity with the
+ // std::array constructor.
+ LaneBitmaskImpl(std::initializer_list<unsigned>) = delete;
+ constexpr LaneBitmaskImpl &operator=(const LaneBitmaskImpl &) = default;
- constexpr LaneBitmask() = default;
- explicit constexpr LaneBitmask(Type V) : Mask(V) {}
+ /// Compare as unsigned integers (most-significant word first). This differs
+ /// from Bitset::operator< which compares bit-by-bit from LSB.
+ constexpr bool operator<(const LaneBitmaskImpl &Other) const {
+ const auto &ThisBits = this->getData();
+ const auto &OtherBits = Other.getData();
+ for (int I = ThisBits.size() - 1; I >= 0; --I) {
+ if (ThisBits[I] != OtherBits[I])
+ return ThisBits[I] < OtherBits[I];
+ }
+ return false;
+ }
- constexpr bool operator== (LaneBitmask M) const { return Mask == M.Mask; }
- constexpr bool operator!= (LaneBitmask M) const { return Mask != M.Mask; }
- constexpr bool operator< (LaneBitmask M) const { return Mask < M.Mask; }
- constexpr bool none() const { return Mask == 0; }
- constexpr bool any() const { return Mask != 0; }
- constexpr bool all() const { return ~Mask == 0; }
+ constexpr LaneBitmaskImpl operator~() const {
+ return Bitset<NumBits>::operator~();
+ }
+ constexpr LaneBitmaskImpl operator|(LaneBitmaskImpl M) const {
+ return Bitset<NumBits>::operator|(M);
+ }
+ constexpr LaneBitmaskImpl operator&(LaneBitmaskImpl M) const {
+ return Bitset<NumBits>::operator&(M);
+ }
+ constexpr LaneBitmaskImpl &operator|=(LaneBitmaskImpl M) {
+ Bitset<NumBits>::operator|=(M);
+ return *this;
+ }
+ constexpr LaneBitmaskImpl &operator&=(LaneBitmaskImpl M) {
+ Bitset<NumBits>::operator&=(M);
+ return *this;
+ }
+ constexpr LaneBitmaskImpl operator^(LaneBitmaskImpl M) const {
+ return Bitset<NumBits>::operator^(M);
+ }
+ constexpr LaneBitmaskImpl &operator^=(LaneBitmaskImpl M) {
+ Bitset<NumBits>::operator^=(M);
+ return *this;
+ }
- constexpr LaneBitmask operator~() const {
- return LaneBitmask(~Mask);
- }
- constexpr LaneBitmask operator|(LaneBitmask M) const {
- return LaneBitmask(Mask | M.Mask);
- }
- constexpr LaneBitmask operator&(LaneBitmask M) const {
- return LaneBitmask(Mask & M.Mask);
- }
- LaneBitmask &operator|=(LaneBitmask M) {
- Mask |= M.Mask;
+ constexpr size_t getNumLanes() const { return this->count(); }
+
+ unsigned getHighestLane() const {
+ assert(this->any() && "getHighestLane called on empty mask");
+ const auto &Bits = this->getData();
+ constexpr size_t WordBits = sizeof(decltype(Bits[0])) * 8;
+ for (int I = Bits.size() - 1; I >= 0; --I)
+ if (Bits[I] != 0)
+ return I * WordBits + Log2_64(Bits[I]);
+ llvm_unreachable("should have found a set bit");
+ }
+
+ /// Shift bits left by \p S positions. Zeroes are shifted in from the right.
+ constexpr LaneBitmaskImpl operator<<(unsigned S) const {
+ return Bitset<NumBits>::operator<<(S);
+ }
+
+ /// Shift bits right by \p S positions. Zeroes are shifted in from the left.
+ constexpr LaneBitmaskImpl operator>>(unsigned S) const {
+ return Bitset<NumBits>::operator>>(S);
+ }
+
+ /// Rotate bits left by \p S positions.
+ constexpr LaneBitmaskImpl rotateLeft(unsigned S) const {
+ S = S % NumBits;
+ if (S == 0)
return *this;
- }
- LaneBitmask &operator&=(LaneBitmask M) {
- Mask &= M.Mask;
+ return (*this << S) | (*this >> (NumBits - S));
+ }
+
+ /// Rotate bits right by \p S positions.
+ constexpr LaneBitmaskImpl rotateRight(unsigned S) const {
+ S = S % NumBits;
+ if (S == 0)
return *this;
- }
+ return (*this >> S) | (*this << (NumBits - S));
+ }
- constexpr Type getAsInteger() const { return Mask; }
+ static constexpr LaneBitmaskImpl getNone() { return LaneBitmaskImpl(); }
- unsigned getNumLanes() const { return llvm::popcount(Mask); }
- unsigned getHighestLane() const {
- return Log2_64(Mask);
- }
+ static constexpr LaneBitmaskImpl getAll() {
+ LaneBitmaskImpl Result;
+ Result.set();
+ return Result;
+ }
- static constexpr LaneBitmask getNone() { return LaneBitmask(0); }
- static constexpr LaneBitmask getAll() { return ~LaneBitmask(0); }
- static constexpr LaneBitmask getLane(unsigned Lane) {
- return LaneBitmask(Type(1) << Lane);
- }
+ static constexpr LaneBitmaskImpl getLane(unsigned Lane) {
+ LaneBitmaskImpl Result;
+ Result.set(Lane);
+ return Result;
+ }
+
+private:
+ constexpr LaneBitmaskImpl(const Bitset<NumBits> &B) : Bitset<NumBits>(B) {}
+
+ /// Helper to convert APInt to array format for Bitset constructor.
+ static std::array<uint64_t, (NumBits + 63) / 64>
+ convertAPIntToArray(const APInt &N) {
+ static_assert(std::is_same_v<APInt::WordType, uint64_t>,
+ "APInt::WordType needs to be uint64_t for word-level copy.");
+ assert(N.getBitWidth() <= NumBits &&
+ "Cannot convert to LaneBitmask. The input APInt has "
+ "more bits than LaneBitmask can hold.");
+ std::array<uint64_t, (NumBits + 63) / 64> Result{};
+ const uint64_t *RawData = N.getRawData();
+ const size_t NumWords = N.getNumWords();
+ for (size_t I = 0; I < NumWords && I < Result.size(); ++I)
+ Result[I] = RawData[I];
+ return Result;
+ }
+
+ template <typename, typename> friend struct llvm::format_provider;
+};
- private:
- Type Mask = 0;
- };
+} // end namespace llvm::detail
- /// Create Printable object to print LaneBitmasks on a \ref raw_ostream.
- inline Printable PrintLaneMask(LaneBitmask LaneMask) {
- return Printable([LaneMask](raw_ostream &OS) {
- OS << format(LaneBitmask::FormatStr, LaneMask.getAsInteger());
- });
+namespace llvm {
+using LaneBitmask = detail::LaneBitmaskImpl<64>;
+
+template <unsigned NumBits>
+struct format_provider<detail::LaneBitmaskImpl<NumBits>> {
+ using T = detail::LaneBitmaskImpl<NumBits>;
+ static void format(const T &V, raw_ostream &Stream, StringRef Style) {
+ // Print as hex using platform words from most significant to least.
+ // Only print the first 64 bits if all upper words are zero.
+ const auto &Data = V.getData();
+ constexpr unsigned SizeOfBitword = sizeof(Data[0]);
+ constexpr unsigned HexWidth = SizeOfBitword * 2;
+ constexpr unsigned NumWordsIn64Bits = 8 / SizeOfBitword;
+ T UpperWords = ~T(~0ULL) & V;
+ if (UpperWords.none())
+ for (int I = NumWordsIn64Bits - 1; I >= 0; --I)
+ Stream << format_hex_no_prefix(Data[I], HexWidth, true);
+ else
+ for (int I = Data.size() - 1; I >= 0; --I)
+ Stream << format_hex_no_prefix(Data[I], HexWidth, true);
}
+};
+
+/// Create Printable object to print LaneBitmasks on a \ref raw_ostream.
+template <unsigned NumBits>
+inline Printable PrintLaneMask(detail::LaneBitmaskImpl<NumBits> LaneMask) {
+ return Printable(
+ [LaneMask](raw_ostream &OS) { OS << formatv("{0}", LaneMask); });
+}
+
+template <unsigned NumBits>
+inline hash_code hash_value(const detail::LaneBitmaskImpl<NumBits> &LM) {
+ return hash_value(static_cast<const Bitset<NumBits> &>(LM));
+}
} // end namespace llvm
+namespace std {
+
+template <unsigned NumBits>
+struct hash<llvm::detail::LaneBitmaskImpl<NumBits>> {
+ size_t operator()(const llvm::detail::LaneBitmaskImpl<NumBits> &LM) const {
+ return hash<llvm::Bitset<NumBits>>{}(LM);
+ }
+};
+
+} // end namespace std
+
#endif // LLVM_MC_LANEBITMASK_H
diff --git a/llvm/lib/CodeGen/MIRParser/MIParser.cpp b/llvm/lib/CodeGen/MIRParser/MIParser.cpp
index 84b806ae81f39..39578172862d4 100644
--- a/llvm/lib/CodeGen/MIRParser/MIParser.cpp
+++ b/llvm/lib/CodeGen/MIRParser/MIParser.cpp
@@ -912,12 +912,20 @@ bool MIParser::parseBasicBlockLiveins(MachineBasicBlock &MBB) {
if (Token.isNot(MIToken::IntegerLiteral) &&
Token.isNot(MIToken::HexLiteral))
return error("expected a lane mask");
- static_assert(sizeof(LaneBitmask::Type) == sizeof(uint64_t),
- "Use correct get-function for lane mask");
- LaneBitmask::Type V;
- if (getUint64(V))
- return error("invalid lane mask value");
- Mask = LaneBitmask(V);
+
+ if (Token.is(MIToken::IntegerLiteral)) {
+ // Parse as integer literal (fits in 64 bits).
+ uint64_t V;
+ if (getUint64(V))
+ return error("invalid lane mask value");
+ Mask = LaneBitmask(V);
+ } else {
+ // Parse as hex literal (may be > 64 bits).
+ APInt A;
+ if (getHexUint(A))
+ return error("invalid lane mask value");
+ Mask = LaneBitmask(A);
+ }
lex();
}
MBB.addLiveIn(Reg, Mask);
@@ -3117,12 +3125,21 @@ bool MIParser::parseLaneMaskOperand(MachineOperand &Dest) {
// Parse lanemask.
if (Token.isNot(MIToken::IntegerLiteral) && Token.isNot(MIToken::HexLiteral))
return error("expected a valid lane mask value");
- static_assert(sizeof(LaneBitmask::Type) == sizeof(uint64_t),
- "Use correct get-function for lane mask.");
- LaneBitmask::Type V;
- if (getUint64(V))
- return true;
- LaneBitmask LaneMask(V);
+
+ LaneBitmask LaneMask;
+ if (Token.is(MIToken::IntegerLiteral)) {
+ // Parse as integer literal (fits in 64 bits).
+ uint64_t V;
+ if (getUint64(V))
+ return error("invalid lane mask value");
+ LaneMask = LaneBitmask(V);
+ } else {
+ // Parse as hex literal (may be > 64 bits).
+ APInt A;
+ if (getHexUint(A))
+ return error("invalid lane mask value");
+ LaneMask = LaneBitmask(A);
+ }
lex();
if (expectAndConsume(MIToken::rparen))
diff --git a/llvm/lib/CodeGen/MachineOperand.cpp b/llvm/lib/CodeGen/MachineOperand.cpp
index ac1f201bc8b83..72ff3578dbb2a 100644
--- a/llvm/lib/CodeGen/MachineOperand.cpp
+++ b/llvm/lib/CodeGen/MachineOperand.cpp
@@ -464,7 +464,7 @@ hash_code llvm::hash_value(const MachineOperand &MO) {
return hash_combine(MO.getType(), MO.getTargetFlags(), MO.getShuffleMask());
case MachineOperand::MO_LaneMask:
return hash_combine(MO.getType(), MO.getTargetFlags(),
- MO.getLaneMask().getAsInteger());
+ hash_value(MO.getLaneMask()));
}
llvm_unreachable("Invalid machine operand type");
}
diff --git a/llvm/lib/CodeGen/MachineStableHash.cpp b/llvm/lib/CodeGen/MachineStableHash.cpp
index 2f5f5aeccb2e4..f777d5fe7722b 100644
--- a/llvm/lib/CodeGen/MachineStableHash.cpp
+++ b/llvm/lib/CodeGen/MachineStableHash.cpp
@@ -166,8 +166,12 @@ stable_hash llvm::stableHashValue(const MachineOperand &MO) {
stable_hash_name(SymbolName));
}
case MachineOperand::MO_LaneMask: {
+ // Use the deterministic printed representation for stable hashing.
+ std::string Str;
+ raw_string_ostream OS(Str);
+ OS << PrintLaneMask(MO.getLaneMask());
return stable_hash_combine(MO.getType(), MO.getTargetFlags(),
- MO.getLaneMask().getAsInteger());
+ stable_hash_name(OS.str()));
}
case MachineOperand::MO_CFIIndex:
return stable_hash_combine(MO.getType(), MO.getTargetFlags(),
diff --git a/llvm/lib/CodeGen/RDFRegisters.cpp b/llvm/lib/CodeGen/RDFRegisters.cpp
index ee3e531c6fd5a..55d059ee7092e 100644
--- a/llvm/lib/CodeGen/RDFRegisters.cpp
+++ b/llvm/lib/CodeGen/RDFRegisters.cpp
@@ -407,11 +407,6 @@ raw_ostream &operator<<(raw_ostream &OS, const PrintLaneMaskShort &P) {
if (P.Mask.none())
return OS << ":*none*";
- LaneBitmask::Type Val = P.Mask.getAsInteger();
- if ((Val & 0xffff) == Val)
- return OS << ':' << format("%04llX", Val);
- if ((Val & 0xffffffff) == Val)
- return OS << ':' << format("%08llX", Val);
return OS << ':' << PrintLaneMask(P.Mask);
}
diff --git a/llvm/lib/Target/AMDGPU/SIRegisterInfo.cpp b/llvm/lib/Target/AMDGPU/SIRegisterInfo.cpp
index e00ce4f167c3f..b424858f95ea6 100644
--- a/llvm/lib/Target/AMDGPU/SIRegisterInfo.cpp
+++ b/llvm/lib/Target/AMDGPU/SIRegisterInfo.cpp
@@ -332,11 +332,11 @@ SIRegisterInfo::SIRegisterInfo(const GCNSubtarget &ST)
ST.getHwMode(MCSubtargetInfo::HwMode_RegInfo)),
ST(ST), SpillSGPRToVGPR(EnableSpillSGPRToVGPR), isWave32(ST.isWave32()) {
- assert(getSubRegIndexLaneMask(AMDGPU::sub0).getAsInteger() == 3 &&
- getSubRegIndexLaneMask(AMDGPU::sub31).getAsInteger() == (3ULL << 62) &&
+ assert(getSubRegIndexLaneMask(AMDGPU::sub0) == LaneBitmask(3) &&
+ getSubRegIndexLaneMask(AMDGPU::sub31) == LaneBitmask(3ULL << 62) &&
(getSubRegIndexLaneMask(AMDGPU::lo16) |
- getSubRegIndexLaneMask(AMDGPU::hi16)).getAsInteger() ==
- getSubRegIndexLaneMask(AMDGPU::sub0).getAsInteger() &&
+ getSubRegIndexLaneMask(AMDGPU::hi16)) ==
+ getSubRegIndexLaneMask(AMDGPU::sub0) &&
"getNumCoveredRegs() will not work with generated subreg masks!");
RegPressureIgnoredUnits.resize(getNumRegUnits());
diff --git a/llvm/lib/Target/AMDGPU/SIRegisterInfo.h b/llvm/lib/Target/AMDGPU/SIRegisterInfo.h
index 9d1a9eae75020..e741b14c07b1a 100644
--- a/llvm/lib/Target/AMDGPU/SIRegisterInfo.h
+++ b/llvm/lib/Target/AMDGPU/SIRegisterInfo.h
@@ -404,11 +404,10 @@ class SIRegisterInfo final : public AMDGPUGenRegisterInfo {
static unsigned getNumCoveredRegs(LaneBitmask LM) {
// The assumption is that every lo16 subreg is an even bit and every hi16
// is an adjacent odd bit or vice versa.
- uint64_t Mask = LM.getAsInteger();
- uint64_t Even = Mask & 0xAAAAAAAAAAAAAAAAULL;
- Mask = (Even >> 1) | Mask;
- uint64_t Odd = Mask & 0x5555555555555555ULL;
- return llvm::popcount(Odd);
+ LaneBitmask Even = LM & LaneBitmask(0xAAAAAAAAAAAAAAAAULL);
+ LaneBitmask Mask = (Even >> 1) | LM;
+ LaneBitmask Odd = Mask & LaneBitmask(0x5555555555555555ULL);
+ return Odd.count();
}
// \returns a DWORD offset of a \p SubReg
diff --git a/llvm/unittests/CodeGen/CMakeLists.txt b/llvm/unittests/CodeGen/CMakeLists.txt
index 709017380fa4e..7324c0ba9b5f6 100644
--- a/llvm/unittests/CodeGen/CMakeLists.txt
+++ b/llvm/unittests/CodeGen/CMakeLists.txt
@@ -30,8 +30,9 @@ add_llvm_unittest(CodeGenTests
DwarfStringPoolEntryRefTest.cpp
GCMetadata.cpp
InstrRefLDVTest.cpp
- LowLevelTypeTest.cpp
+ LaneBitmaskTest.cpp
LexicalScopesTest.cpp
+ LowLevelTypeTest.cpp
MachineBasicBlockTest.cpp
MachineDomTreeUpdaterTest.cpp
MachineInstrBundleIteratorTest.cpp
diff --git a/llvm/unittests/CodeGen/LaneBitmaskTest.cpp b/llvm/unittests/CodeGen/LaneBitmaskTest.cpp
new file mode 100644
index 0000000000000..091a3b576b1a7
--- /dev/null
+++ b/llvm/unittests/CodeGen/LaneBitmaskTest.cpp
@@ -0,0 +1,933 @@
+//===- llvm/unittest/CodeGen/LaneBitmaskTest.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/MC/LaneBitmask.h"
+#include "gtest/gtest.h"
+
+using namespace llvm;
+
+namespace {
+
+// Type aliases for clarity.
+using LaneMask64 = detail::LaneBitmaskImpl<64>;
+using LaneMask128 = detail::LaneBitmaskImpl<128>;
+
+TEST(LaneBitmaskTest, ConstructorAndAssignment) {
+ // Test 64-bit version.
+ {
+ LaneMask64 Default;
+ EXPECT_TRUE(Default.none());
+ EXPECT_FALSE(Default.any());
+
+ LaneMask64 FromInt(0x123456789abcdef0);
+ EXPECT_TRUE(FromInt.any());
+ EXPECT_FALSE(FromInt.none());
+
+ std::array<uint64_t, 1> Arr64 = {0x123456789abcdef0};
+ LaneMask64 FromArray(Arr64);
+ EXPECT_TRUE(FromArray.test(4));
+ EXPECT_EQ(FromArray.count(), 32u);
+
+ LaneMask64 None = LaneMask64::getNone();
+ EXPECT_TRUE(None.none());
+
+ LaneMask64 All = LaneMask64::getAll();
+ EXPECT_TRUE(All.all());
+ EXPECT_EQ(All.count(), LaneMask64::BitWidth);
+
+ LaneMask64 Lane5 = LaneMask64::getLane(5);
+ EXPECT_TRUE(Lane5.test(5));
+ EXPECT_EQ(Lane5.getNumLanes(), 1u);
+
+ LaneMask64 Bit0 = LaneMask64::getLane(0);
+ EXPECT_TRUE(Bit0.test(0));
+ EXPECT_FALSE(Bit0.test(1));
+ EXPECT_EQ(Bit0.getHighestLane(), 0u);
+ EXPECT_TRUE((Bit0 >> 1).none());
+ EXPECT_TRUE(Bit0.rotateLeft(1).test(1));
+
+ LaneMask64 BitMax = LaneMask64::getLane(LaneMask64::BitWidth - 1);
+ EXPECT_TRUE(BitMax.test(LaneMask64::BitWidth - 1));
+ EXPECT_FALSE(BitMax.test(LaneMask64::BitWidth - 2));
+ EXPECT_EQ(BitMax.getHighestLane(), LaneMask64::BitWidth - 1);
+ EXPECT_TRUE(BitMax.rotateLeft(1).test(0));
+ EXPECT_TRUE((BitMax << 1).none());
+
+ LaneMask64 Original(0xabcd);
+ LaneMask64 Copied(Original);
+ EXPECT_EQ(Copied, Original);
+
+ LaneMask64 Assigned;
+ Assigned = Original;
+ EXPECT_EQ(Assigned, Original);
+ }
+
+ // Test 128-bit version.
+ {
+ LaneMask128 Default;
+ EXPECT_TRUE(Default.none());
+ EXPECT_FALSE(Default.any());
+
+ LaneMask128 FromInt(0x123456789abcdef0);
+ EXPECT_TRUE(FromInt.any());
+ EXPECT_FALSE(FromInt.none());
+
+ std::array<uint64_t, 2> Arr128 = {0xff, 0xff00};
+ LaneMask128 FromArray(Arr128);
+ EXPECT_TRUE(FromArray.test(0));
+ EXPECT_TRUE(FromArray.test(7));
+ EXPECT_TRUE(FromArray.test(64 + 8));
+ EXPECT_EQ(FromArray.count(), 16u);
+
+ LaneMask128 None = LaneMask128::getNone();
+ EXPECT_TRUE(None.none());
+
+ LaneMask128 All = LaneMask128::getAll();
+ EXPECT_TRUE(All.all());
+ EXPECT_EQ(All.count(), LaneMask128::BitWidth);
+
+ LaneMask128 Lane5 = LaneMask128::getLane(5);
+ EXPECT_TRUE(Lane5.test(5));
+ EXPECT_EQ(Lane5.getNumLanes(), 1u);
+
+ LaneMask128 Lane100 = LaneMask128::getLane(100);
+ EXPECT_TRUE(Lane100.test(100));
+ EXPECT_FALSE(Lane100.test(99));
+ EXPECT_EQ(Lane100.getNumLanes(), 1u);
+
+ LaneMask128 Bit127 = LaneMask128::getLane(127);
+ EXPECT_TRUE(Bit127.test(127));
+ EXPECT_FALSE(Bit127.test(126));
+ EXPECT_EQ(Bit127.getHighestLane(), 127u);
+ EXPECT_TRUE(Bit127.rotateLeft(1).test(0));
+ EXPECT_TRUE((Bit127 << 1).none());
+
+ LaneMask128 Bit63 = LaneMask128::getLane(63);
+ EXPECT_TRUE(Bit63.test(63));
+ EXPECT_FALSE(Bit63.test(62));
+ EXPECT_FALSE(Bit63.test(64));
+ EXPECT_EQ(Bit63.getHighestLane(), 63u);
+
+ LaneMask128 Original = LaneMask128::getLane(100);
+ LaneMask128 Copied(Original);
+ EXPECT_EQ(Copied, Original);
+
+ LaneMask128 Assigned;
+ Assigned = Original;
+ EXPECT_EQ(Assigned, Original);
+ }
+
+ // Constexpr tests for 64-bit.
+ static_assert(LaneMask64::getNone().none(), "getNone() should be empty");
+ static_assert(LaneMask64::getAll().all(), "getAll() should be full");
+ static_assert(LaneMask64::getAll().count() == LaneMask64::BitWidth,
+ "getAll() should have all bits set");
+ static_assert(LaneMask64::getLane(5).test(5), "getLane() should set bit");
+ static_assert(LaneMask64::getLane(0).count() == 1, "getLane() sets one bit");
+ static_assert(LaneMask64::getLane(0).test(0), "Bit 0 is set");
+ static_assert((LaneMask64::getLane(0) >> 1).none(), "Shift clears bit 0");
+ static_assert(
+ LaneMask64::getLane(LaneMask64::BitWidth - 1).rotateLeft(1).test(0),
+ "Rotate highest bit wraps to 0");
+ static_assert((LaneMask64::getLane(LaneMask64::BitWidth - 1) << 1).none(),
+ "Shift clears highest bit");
+
+ // Constexpr tests for 128-bit.
+ static_assert(LaneMask128::getNone().none(), "getNone() should be empty");
+ static_assert(LaneMask128::getAll().all(), "getAll() should be full");
+ static_assert(LaneMask128::getAll().count() == LaneMask128::BitWidth,
+ "getAll() should have all bits set");
+ static_assert(LaneMask128::getLane(100).test(100),
+ "getLane(100) should set bit 100");
+ static_assert(LaneMask128::getLane(127).count() == 1,
+ "getLane() sets one bit");
+ static_assert(LaneMask128::getLane(127).test(127), "Bit 127 is set");
+ static_assert(!LaneMask128::getLane(63).test(64), "Bit 64 not set");
+ static_assert((LaneMask128::getLane(0) >> 1).none(), "Shift clears bit 0");
+ static_assert(
+ LaneMask128::getLane(LaneMask128::BitWidth - 1).rotateLeft(1).test(0),
+ "Rotate highest bit wraps to 0");
+ static_assert((LaneMask128::getLane(LaneMask128::BitWidth - 1) << 1).none(),
+ "Shift clears highest bit");
+ static_assert(
+ []() constexpr {
+ std::array<uint64_t, 1> Arr = {0xff};
+ LaneMask64 M(Arr);
+ return M.count() == 8;
+ }(),
+ "Constexpr array constructor");
+ static_assert(
+ []() constexpr {
+ std::array<uint64_t, 2> Arr = {0xff, 0xff00};
+ LaneMask128 M(Arr);
+ return M.count() == 16;
+ }(),
+ "Constexpr array constructor 128-bit");
+ static_assert(
+ []() constexpr {
+ LaneMask64 Original(0xff);
+ LaneMask64 Copied(Original);
+ return Copied == Original;
+ }(),
+ "Constexpr copy constructor");
+ static_assert(
+ []() constexpr {
+ LaneMask64 Original(0xff);
+ LaneMask64 Assigned;
+ Assigned = Original;
+ return Assigned == Original;
+ }(),
+ "Constexpr copy assignment");
+ static_assert(
+ []() constexpr {
+ LaneMask128 Original = LaneMask128::getLane(100);
+ LaneMask128 Copied(Original);
+ return Copied == Original;
+ }(),
+ "Constexpr copy constructor 128-bit");
+ static_assert(
+ []() constexpr {
+ LaneMask128 Original = LaneMask128::getLane(100);
+ LaneMask128 Assigned;
+ Assigned = Original;
+ return Assigned == Original;
+ }(),
+ "Constexpr copy assignment 128-bit");
+}
+
+TEST(LaneBitmaskTest, ComparisonOperators) {
+ // Test 64-bit version.
+ {
+ LaneMask64 A(0x1234);
+ LaneMask64 B(0x1234);
+ LaneMask64 C(0x5678);
+
+ EXPECT_TRUE(A == B);
+ EXPECT_FALSE(A == C);
+ EXPECT_FALSE(A != B);
+ EXPECT_TRUE(A != C);
+ EXPECT_TRUE(A < C);
+ EXPECT_FALSE(C < A);
+ EXPECT_FALSE(A < B);
+ }
+
+ // Test 128-bit version.
+ {
+ LaneMask128 A(0x1234);
+ LaneMask128 B(0x1234);
+ LaneMask128 C(0x5678);
+
+ EXPECT_TRUE(A == B);
+ EXPECT_FALSE(A == C);
+ EXPECT_FALSE(A != B);
+ EXPECT_TRUE(A != C);
+ EXPECT_TRUE(A < C);
+ EXPECT_FALSE(C < A);
+ EXPECT_FALSE(A < B);
+
+ // Test comparison across word boundaries.
+ LaneMask128 Low(0x1234);
+ LaneMask128 High = LaneMask128::getLane(100);
+ EXPECT_TRUE(Low < High);
+ EXPECT_FALSE(High < Low);
+ }
+
+ // Constexpr comparison operators for 64-bit.
+ static_assert(LaneMask64(0x1234) == LaneMask64(0x1234), "Equality");
+ static_assert(LaneMask64(0x1234) != LaneMask64(0x5678), "Inequality");
+ static_assert(LaneMask64(0x1000) < LaneMask64(0x2000), "Less than");
+
+ // Constexpr comparison operators for 128-bit.
+ static_assert(LaneMask128(0x1234) == LaneMask128(0x1234), "Equality");
+ static_assert(LaneMask128(0x1234) != LaneMask128(0x5678), "Inequality");
+ static_assert(LaneMask128::getLane(64) < LaneMask128::getLane(100),
+ "Cross-word less than");
+}
+
+TEST(LaneBitmaskTest, BitwiseOperators) {
+ // Test 64-bit version.
+ {
+ LaneMask64 A(0xff00);
+ LaneMask64 B(0x0ff0);
+
+ // Test OR.
+ EXPECT_EQ(A | B, LaneMask64(0xfff0));
+
+ // Test AND.
+ EXPECT_EQ(A & B, LaneMask64(0x0f00));
+
+ // Test XOR.
+ EXPECT_EQ(A ^ B, LaneMask64(0xf0f0));
+ EXPECT_EQ((LaneMask64(0xff) ^ LaneMask64(0xff)), LaneMask64::getNone());
+
+ // Test NOT.
+ LaneMask64 NotA = ~A;
+ EXPECT_EQ(NotA.count(), LaneMask64::BitWidth - A.count());
+
+ // Test OR assign.
+ LaneMask64 A2(0xff00);
+ A2 |= B;
+ EXPECT_EQ(A2, LaneMask64(0xfff0));
+
+ // Test AND assign.
+ LaneMask64 A3(0xff00);
+ A3 &= B;
+ EXPECT_EQ(A3, LaneMask64(0x0f00));
+
+ // Test XOR assign.
+ LaneMask64 A4(0xff00);
+ A4 ^= B;
+ EXPECT_EQ(A4, LaneMask64(0xf0f0));
+ }
+
+ // Test 128-bit version.
+ {
+ LaneMask128 A(0xff00);
+ LaneMask128 B(0x0ff0);
+
+ // Test OR.
+ EXPECT_EQ(A | B, LaneMask128(0xfff0));
+
+ // Test AND.
+ EXPECT_EQ(A & B, LaneMask128(0x0f00));
+
+ // Test XOR.
+ EXPECT_EQ(A ^ B, LaneMask128(0xf0f0));
+ EXPECT_EQ((LaneMask128(0xff) ^ LaneMask128(0xff)), LaneMask128::getNone());
+
+ // Test NOT.
+ LaneMask128 NotA = ~A;
+ EXPECT_EQ(NotA.count(), LaneMask128::BitWidth - A.count());
+
+ // Test operations across word boundaries.
+ LaneMask128 Low(0xff);
+ LaneMask128 High = LaneMask128::getLane(100);
+ LaneMask128 Combined = Low | High;
+ EXPECT_TRUE(Combined.test(0));
+ EXPECT_TRUE(Combined.test(100));
+ EXPECT_EQ(Combined.count(), 9u);
+
+ LaneMask128 AndResult = Combined & High;
+ EXPECT_FALSE(AndResult.test(0));
+ EXPECT_TRUE(AndResult.test(100));
+ EXPECT_EQ(AndResult.count(), 1u);
+
+ // Test XOR across word boundaries.
+ EXPECT_TRUE((Combined ^ High).test(0));
+ EXPECT_FALSE((Combined ^ High).test(100));
+ EXPECT_EQ((Combined ^ High).count(), 8u);
+
+ // Test XOR assign.
+ LaneMask128 A4(0xff00);
+ A4 ^= B;
+ EXPECT_EQ(A4, LaneMask128(0xf0f0));
+ }
+
+ // Constexpr bitwise operations for 64-bit.
+ static_assert((LaneMask64(0xff00) | LaneMask64(0x0ff0)) == LaneMask64(0xfff0),
+ "Constexpr OR");
+ static_assert((LaneMask64(0xff00) & LaneMask64(0x0ff0)) == LaneMask64(0x0f00),
+ "Constexpr AND");
+ static_assert((LaneMask64(0xff00) ^ LaneMask64(0x0ff0)) == LaneMask64(0xf0f0),
+ "Constexpr XOR");
+ static_assert((~LaneMask64(0xff)).any(), "Constexpr NOT");
+ static_assert((~LaneMask64::getAll()).none(), "Constexpr NOT of all");
+ static_assert(
+ []() constexpr {
+ LaneMask64 L(0xff00);
+ L |= LaneMask64(0x0ff0);
+ return L == LaneMask64(0xfff0);
+ }(),
+ "Constexpr OR assign");
+ static_assert(
+ []() constexpr {
+ LaneMask64 L(0xff00);
+ L &= LaneMask64(0x0ff0);
+ return L == LaneMask64(0x0f00);
+ }(),
+ "Constexpr AND assign");
+ static_assert(
+ []() constexpr {
+ LaneMask64 L(0xff00);
+ L ^= LaneMask64(0x0ff0);
+ return L == LaneMask64(0xf0f0);
+ }(),
+ "Constexpr XOR assign");
+
+ // Constexpr bitwise operations for 128-bit.
+ static_assert((LaneMask128(0xff00) | LaneMask128(0x0ff0)) ==
+ LaneMask128(0xfff0),
+ "Constexpr OR");
+ static_assert((LaneMask128(0xff00) ^ LaneMask128(0x0ff0)) ==
+ LaneMask128(0xf0f0),
+ "Constexpr XOR");
+ static_assert(
+ (LaneMask128::getLane(64) | LaneMask128::getLane(100)).count() == 2,
+ "Constexpr OR across words");
+ static_assert((LaneMask128::getLane(64) ^ LaneMask128::getLane(64)).none(),
+ "Constexpr XOR with self is empty");
+ static_assert(
+ []() constexpr {
+ LaneMask128 L(0xff00);
+ L |= LaneMask128::getLane(100);
+ return L.count() == 9;
+ }(),
+ "Constexpr OR assign 128-bit");
+ static_assert(
+ []() constexpr {
+ LaneMask128 L(0xff00);
+ L &= LaneMask128(0x0ff0);
+ return L == LaneMask128(0x0f00);
+ }(),
+ "Constexpr AND assign 128-bit");
+ static_assert(
+ []() constexpr {
+ LaneMask128 L(0xff00);
+ L ^= LaneMask128(0x0ff0);
+ return L == LaneMask128(0xf0f0);
+ }(),
+ "Constexpr XOR assign 128-bit");
+}
+
+TEST(LaneBitmaskTest, QueryMethods) {
+ // Test 64-bit version.
+ {
+ LaneMask64 Empty;
+ EXPECT_TRUE(Empty.none());
+ EXPECT_FALSE(Empty.any());
+ EXPECT_FALSE(Empty.all());
+ EXPECT_EQ(Empty.count(), 0u);
+ EXPECT_EQ(Empty.getNumLanes(), 0u);
+
+ LaneMask64 Partial(0x00ff);
+ EXPECT_FALSE(Partial.none());
+ EXPECT_TRUE(Partial.any());
+ EXPECT_FALSE(Partial.all());
+ EXPECT_EQ(Partial.count(), 8u);
+ EXPECT_EQ(Partial.getNumLanes(), 8u);
+
+ LaneMask64 Full = LaneMask64::getAll();
+ EXPECT_FALSE(Full.none());
+ EXPECT_TRUE(Full.any());
+ EXPECT_TRUE(Full.all());
+ EXPECT_EQ(Full.count(), LaneMask64::BitWidth);
+ EXPECT_EQ(Full.getNumLanes(), LaneMask64::BitWidth);
+ EXPECT_EQ(Full.size(), LaneMask64::BitWidth);
+ }
+
+ // Test 128-bit version.
+ {
+ LaneMask128 Empty;
+ EXPECT_TRUE(Empty.none());
+ EXPECT_FALSE(Empty.any());
+ EXPECT_FALSE(Empty.all());
+ EXPECT_EQ(Empty.count(), 0u);
+ EXPECT_EQ(Empty.getNumLanes(), 0u);
+
+ LaneMask128 Partial(0x00ff);
+ EXPECT_FALSE(Partial.none());
+ EXPECT_TRUE(Partial.any());
+ EXPECT_FALSE(Partial.all());
+ EXPECT_EQ(Partial.count(), 8u);
+ EXPECT_EQ(Partial.getNumLanes(), 8u);
+
+ LaneMask128 Full = LaneMask128::getAll();
+ EXPECT_FALSE(Full.none());
+ EXPECT_TRUE(Full.any());
+ EXPECT_TRUE(Full.all());
+ EXPECT_EQ(Full.count(), LaneMask128::BitWidth);
+ EXPECT_EQ(Full.getNumLanes(), LaneMask128::BitWidth);
+
+ LaneMask128 UpperBit = LaneMask128::getLane(127);
+ EXPECT_FALSE(UpperBit.none());
+ EXPECT_TRUE(UpperBit.any());
+ EXPECT_FALSE(UpperBit.all());
+ }
+
+ // Constexpr query methods for 64-bit.
+ static_assert(LaneMask64().none(), "Empty mask");
+ static_assert(!LaneMask64(0x1).none(), "Non-empty mask");
+ static_assert(LaneMask64(0x1).any(), "Any bit set");
+ static_assert(!LaneMask64().any(), "No bits set");
+ static_assert(LaneMask64::getAll().all(), "All bits set");
+ static_assert(!LaneMask64(0xff).all(), "Not all bits set");
+ static_assert(LaneMask64().count() == 0, "Count zero");
+ static_assert(LaneMask64(0x7).count() == 3, "Count three");
+ static_assert(LaneMask64(0x7).getNumLanes() == 3, "Num lanes three");
+ static_assert(LaneMask64().size() == LaneMask64::BitWidth, "Size");
+ static_assert(LaneMask64::getAll().getNumLanes() == LaneMask64::BitWidth,
+ "Full mask num lanes");
+
+ // Constexpr query methods for 128-bit.
+ static_assert(LaneMask128().none(), "Empty mask");
+ static_assert(LaneMask128::getLane(127).any(), "High bit set");
+ static_assert(LaneMask128::getAll().all(), "All bits set");
+ static_assert(LaneMask128().size() == LaneMask128::BitWidth, "Size");
+ static_assert(LaneMask128::getAll().getNumLanes() == 128, "128 lanes");
+}
+
+TEST(LaneBitmaskTest, GetHighestLane) {
+ // Test 64-bit version.
+ EXPECT_EQ(LaneMask64(1).getHighestLane(), 0u);
+ EXPECT_EQ(LaneMask64(1ull << 5).getHighestLane(), 5u);
+ EXPECT_EQ(LaneMask64(1ull << 63).getHighestLane(), 63u);
+ EXPECT_EQ(LaneMask64((1ull << 10) | (1ull << 30)).getHighestLane(), 30u);
+
+ // Test 128-bit version.
+ EXPECT_EQ(LaneMask128(1).getHighestLane(), 0u);
+ EXPECT_EQ(LaneMask128(1ull << 5).getHighestLane(), 5u);
+ EXPECT_EQ(LaneMask128::getLane(100).getHighestLane(), 100u);
+ EXPECT_EQ(LaneMask128::getLane(127).getHighestLane(), 127u);
+ EXPECT_EQ(
+ (LaneMask128::getLane(10) | LaneMask128::getLane(100)).getHighestLane(),
+ 100u);
+}
+
+TEST(LaneBitmaskTest, ShiftAssignOperators) {
+ // Test 64-bit version.
+ {
+ LaneMask64 A1(0xff);
+ A1 <<= 8;
+ EXPECT_EQ(A1, LaneMask64(0xff00));
+
+ LaneMask64 A2(0xff00);
+ A2 >>= 8;
+ EXPECT_EQ(A2, LaneMask64(0xff));
+
+ LaneMask64 A3(0x1);
+ A3 <<= 4;
+ A3 <<= 4;
+ EXPECT_EQ(A3, LaneMask64(0x100));
+ }
+
+ // Test 128-bit version with cross-word shifts.
+ {
+ LaneMask128 A1(0xff);
+ A1 <<= 8;
+ EXPECT_EQ(A1, LaneMask128(0xff00));
+
+ LaneMask128 A2(0xff);
+ A2 <<= 64;
+ EXPECT_FALSE(A2.test(0));
+ EXPECT_TRUE(A2.test(64));
+
+ LaneMask128 A3 = LaneMask128::getLane(100);
+ A3 >>= 50;
+ EXPECT_TRUE(A3.test(50));
+ EXPECT_FALSE(A3.test(100));
+ }
+
+ static_assert(
+ []() constexpr {
+ LaneMask64 L(0xff);
+ L <<= 8;
+ return L == LaneMask64(0xff00);
+ }(),
+ "Constexpr shift left assign");
+
+ static_assert(
+ []() constexpr {
+ LaneMask64 L(0xff00);
+ L >>= 8;
+ return L == LaneMask64(0xff);
+ }(),
+ "Constexpr shift right assign");
+
+ static_assert(
+ []() constexpr {
+ LaneMask128 L(0xff);
+ L <<= 64;
+ return !L.test(0) && L.test(64);
+ }(),
+ "Constexpr shift left assign 128-bit");
+
+ static_assert(
+ []() constexpr {
+ LaneMask128 L = LaneMask128::getLane(100);
+ L >>= 50;
+ return L.test(50) && !L.test(100);
+ }(),
+ "Constexpr shift right assign 128-bit");
+}
+
+TEST(LaneBitmaskTest, ShiftOperators) {
+ // Test 64-bit version.
+ {
+ LaneMask64 A(0xff);
+
+ EXPECT_EQ(A << 0, A);
+ EXPECT_EQ(A << 8, LaneMask64(0xff00));
+ EXPECT_TRUE((A << LaneMask64::BitWidth).none());
+ EXPECT_TRUE((A << (LaneMask64::BitWidth + 10)).none());
+
+ LaneMask64 B(0xff00);
+ EXPECT_EQ(B >> 0, B);
+ EXPECT_EQ(B >> 8, LaneMask64(0xff));
+ EXPECT_TRUE((B >> LaneMask64::BitWidth).none());
+
+ LaneMask64 Bit0 = LaneMask64::getLane(0);
+ LaneMask64 ShiftedLeft = Bit0 << (LaneMask64::BitWidth - 1);
+ EXPECT_FALSE(ShiftedLeft.none());
+ EXPECT_TRUE(ShiftedLeft.test(LaneMask64::BitWidth - 1));
+ EXPECT_TRUE((ShiftedLeft << 1).none());
+ }
+
+ // Test 128-bit version with cross-word shifts.
+ {
+ LaneMask128 A(0xff);
+
+ // Shift left within word.
+ EXPECT_EQ(A << 8, LaneMask128(0xff00));
+
+ // Shift left across word boundary.
+ LaneMask128 B = A << 60;
+ EXPECT_TRUE(B.test(60));
+ EXPECT_TRUE(B.test(67));
+
+ // Shift left completely into upper word.
+ LaneMask128 C = A << 64;
+ EXPECT_FALSE(C.test(0));
+ EXPECT_TRUE(C.test(64));
+ EXPECT_TRUE(C.test(71));
+
+ // Shift right across word boundary.
+ LaneMask128 D = LaneMask128::getLane(100);
+ LaneMask128 E = D >> 50;
+ EXPECT_TRUE(E.test(50));
+ EXPECT_FALSE(E.test(100));
+
+ // Shift by full width.
+ EXPECT_TRUE((A << LaneMask128::BitWidth).none());
+ EXPECT_TRUE((A >> LaneMask128::BitWidth).none());
+ }
+
+ // Constexpr shift operations for 64-bit.
+ static_assert((LaneMask64(0xff) << 8) == LaneMask64(0xff00),
+ "Constexpr shift left");
+ static_assert((LaneMask64(0xff00) >> 8) == LaneMask64(0xff),
+ "Constexpr shift right");
+ static_assert((LaneMask64(0xff) << LaneMask64::BitWidth).none(),
+ "Shift left by BitWidth");
+ static_assert((LaneMask64(0xff) << (LaneMask64::BitWidth + 10)).none(),
+ "Shift left beyond BitWidth");
+ static_assert(
+ []() constexpr {
+ LaneMask64 Bit0 = LaneMask64::getLane(0);
+ LaneMask64 Shifted = Bit0 << (LaneMask64::BitWidth - 1);
+ return !Shifted.none() && Shifted.test(LaneMask64::BitWidth - 1) &&
+ (Shifted << 1).none();
+ }(),
+ "Shift by BitWidth-1 then by 1");
+
+ // Constexpr shift operations for 128-bit.
+ static_assert((LaneMask128(0xff) << 8) == LaneMask128(0xff00),
+ "Constexpr shift left");
+ static_assert((LaneMask128::getLane(64) >> 64).test(0),
+ "Shift across word boundary");
+}
+
+TEST(LaneBitmaskTest, RotateOperators) {
+ // Test 64-bit version.
+ {
+ LaneMask64 A(0xff);
+
+ EXPECT_EQ(A.rotateLeft(0), A);
+ EXPECT_EQ(A.rotateLeft(8), LaneMask64(0xff00));
+ EXPECT_EQ(A.rotateLeft(LaneMask64::BitWidth), A);
+ EXPECT_EQ(A.rotateLeft(LaneMask64::BitWidth * 2), A);
+ EXPECT_EQ(A.rotateLeft(LaneMask64::BitWidth + 8), A.rotateLeft(8));
+
+ LaneMask64 B(0xff00);
+ EXPECT_EQ(B.rotateRight(0), B);
+ EXPECT_EQ(B.rotateRight(8), LaneMask64(0xff));
+ EXPECT_EQ(B.rotateRight(LaneMask64::BitWidth), B);
+ EXPECT_EQ(B.rotateRight(LaneMask64::BitWidth * 3), B);
+ EXPECT_EQ(B.rotateRight(LaneMask64::BitWidth + 8), B.rotateRight(8));
+
+ LaneMask64 C(0x123456789abcdef0);
+ EXPECT_EQ(C.rotateLeft(37).rotateRight(37), C);
+
+ LaneMask64 HighBit = LaneMask64::getLane(LaneMask64::BitWidth - 1);
+ EXPECT_TRUE((HighBit << 1).none());
+ EXPECT_TRUE(HighBit.rotateLeft(1).test(0));
+ }
+
+ // Test 128-bit version with cross-word rotations.
+ {
+ LaneMask128 A(0xff);
+
+ // Rotate left within and across words.
+ EXPECT_EQ(A.rotateLeft(0), A);
+ EXPECT_EQ(A.rotateLeft(8), LaneMask128(0xff00));
+ EXPECT_EQ(A.rotateLeft(LaneMask128::BitWidth), A);
+
+ // Rotate across word boundary.
+ LaneMask128 B = A.rotateLeft(60);
+ EXPECT_TRUE(B.test(60));
+ EXPECT_TRUE(B.test(67));
+
+ // Rotate highest bit wraps to bit 0.
+ LaneMask128 HighBit = LaneMask128::getLane(127);
+ EXPECT_TRUE(HighBit.rotateLeft(1).test(0));
+ EXPECT_FALSE(HighBit.rotateLeft(1).test(127));
+
+ // Rotate from upper word to lower word.
+ LaneMask128 UpperBit = LaneMask128::getLane(100);
+ LaneMask128 Rotated = UpperBit.rotateRight(50);
+ EXPECT_TRUE(Rotated.test(50));
+ EXPECT_FALSE(Rotated.test(100));
+
+ // Verify rotate vs shift difference for 128-bit.
+ EXPECT_TRUE((HighBit << 1).none());
+ EXPECT_TRUE(HighBit.rotateLeft(1).test(0));
+ }
+
+ // Constexpr rotate operations for 64-bit.
+ static_assert(LaneMask64(0xff).rotateLeft(8) == LaneMask64(0xff00),
+ "Constexpr rotate left");
+ static_assert(LaneMask64(0xff00).rotateRight(8) == LaneMask64(0xff),
+ "Constexpr rotate right");
+ static_assert(LaneMask64(0xff).rotateLeft(LaneMask64::BitWidth) ==
+ LaneMask64(0xff),
+ "Rotate by BitWidth is identity");
+ static_assert(LaneMask64(0xff).rotateLeft(LaneMask64::BitWidth * 2) ==
+ LaneMask64(0xff),
+ "Rotate by multiple of BitWidth");
+ static_assert(LaneMask64(0xff).rotateLeft(LaneMask64::BitWidth + 8) ==
+ LaneMask64(0xff).rotateLeft(8),
+ "Rotate by BitWidth + N");
+ static_assert(LaneMask64(0x1234).rotateLeft(37).rotateRight(37) ==
+ LaneMask64(0x1234),
+ "Rotate roundtrip");
+ static_assert(
+ LaneMask64::getLane(LaneMask64::BitWidth - 1).rotateLeft(1).test(0),
+ "Rotate wraps around");
+
+ // Constexpr rotate operations for 128-bit.
+ static_assert(LaneMask128(0xff).rotateLeft(8) == LaneMask128(0xff00),
+ "Constexpr rotate left");
+ static_assert(LaneMask128::getLane(127).rotateLeft(1).test(0),
+ "Rotate highest bit wraps to 0");
+ static_assert(LaneMask128(0xff).rotateLeft(LaneMask128::BitWidth) ==
+ LaneMask128(0xff),
+ "Rotate by 128 is identity");
+}
+
+TEST(LaneBitmaskTest, InheritedBitsetOperations) {
+ // Test 64-bit version.
+ {
+ LaneMask64 M;
+
+ M.set(5);
+ EXPECT_TRUE(M.test(5));
+ EXPECT_TRUE(M[5]);
+ EXPECT_EQ(M.count(), 1u);
+
+ M.set(10);
+ EXPECT_EQ(M.count(), 2u);
+
+ M.reset(5);
+ EXPECT_FALSE(M.test(5));
+ EXPECT_FALSE(M[5]);
+ EXPECT_EQ(M.count(), 1u);
+
+ M.flip(5);
+ EXPECT_TRUE(M.test(5));
+ EXPECT_TRUE(M[5]);
+
+ M.set();
+ EXPECT_TRUE(M.all());
+ EXPECT_EQ(M.getNumLanes(), LaneMask64::BitWidth);
+ }
+
+ // Test 128-bit version.
+ {
+ LaneMask128 M;
+
+ M.set(100);
+ EXPECT_TRUE(M.test(100));
+ EXPECT_TRUE(M[100]);
+ EXPECT_EQ(M.count(), 1u);
+
+ M.set(10);
+ EXPECT_EQ(M.count(), 2u);
+
+ M.reset(100);
+ EXPECT_FALSE(M.test(100));
+ EXPECT_FALSE(M[100]);
+
+ M.flip(127);
+ EXPECT_TRUE(M.test(127));
+ EXPECT_TRUE(M[127]);
+
+ M.set();
+ EXPECT_TRUE(M.all());
+ EXPECT_EQ(M.getNumLanes(), LaneMask128::BitWidth);
+ }
+
+ // Constexpr set/reset/flip/test/operator[] operations for 64-bit.
+ constexpr auto TestSet64 = []() constexpr {
+ LaneMask64 L;
+ L.set(5);
+ return L.test(5) && L[5] && L.count() == 1;
+ };
+ static_assert(TestSet64(), "Constexpr set and test");
+
+ constexpr auto TestReset64 = []() constexpr {
+ LaneMask64 L;
+ L.set(5);
+ L.reset(5);
+ return !L.test(5) && !L[5] && L.none();
+ };
+ static_assert(TestReset64(), "Constexpr reset");
+
+ constexpr auto TestFlip64 = []() constexpr {
+ LaneMask64 L;
+ L.flip(5);
+ return L.test(5) && L[5] && L.count() == 1;
+ };
+ static_assert(TestFlip64(), "Constexpr flip");
+
+ constexpr auto TestSetAll64 = []() constexpr {
+ LaneMask64 L;
+ L.set();
+ return L.all() && L.count() == LaneMask64::BitWidth;
+ };
+ static_assert(TestSetAll64(), "Constexpr set all");
+
+ // Constexpr operations for 128-bit.
+ constexpr auto TestSet128 = []() constexpr {
+ LaneMask128 L;
+ L.set(100);
+ return L.test(100) && L[100] && L.count() == 1;
+ };
+ static_assert(TestSet128(), "Constexpr set bit 100");
+
+ constexpr auto TestReset128 = []() constexpr {
+ LaneMask128 L;
+ L.set(100);
+ L.reset(100);
+ return !L.test(100) && L.none();
+ };
+ static_assert(TestReset128(), "Constexpr reset bit 100");
+
+ constexpr auto TestFlip128 = []() constexpr {
+ LaneMask128 L;
+ L.flip(100);
+ return L.test(100) && L.count() == 1;
+ };
+ static_assert(TestFlip128(), "Constexpr flip bit 100");
+
+ constexpr auto TestSetAll128 = []() constexpr {
+ LaneMask128 L;
+ L.set();
+ return L.all() && L.count() == LaneMask128::BitWidth;
+ };
+ static_assert(TestSetAll128(), "Constexpr set all 128-bit");
+}
+
+TEST(LaneBitmaskTest, APIntConstructor) {
+ APInt Empty(64, 0);
+ LaneMask64 MEmpty(Empty);
+ EXPECT_TRUE(MEmpty.none());
+
+ APInt Full(64, ~0ull);
+ LaneMask64 MFull(Full);
+ EXPECT_EQ(MFull.getNumLanes(), 64u);
+ EXPECT_TRUE(MFull.all());
+
+ APInt A64(64, 0x123456789abcdef0, false);
+ LaneMask64 M64(A64);
+ EXPECT_TRUE(M64.test(4));
+ EXPECT_EQ(M64.getHighestLane(), 60u);
+
+ APInt Small(16, 0xff);
+ LaneMask64 MSmall(Small);
+ EXPECT_EQ(MSmall.getNumLanes(), 8u);
+ EXPECT_TRUE(MSmall.test(0));
+ EXPECT_TRUE(MSmall.test(7));
+ EXPECT_FALSE(MSmall.test(8));
+
+ APInt A128(128, {0xff, 0xff00});
+ LaneMask128 M128(A128);
+ EXPECT_TRUE(M128.test(0));
+ EXPECT_TRUE(M128.test(7));
+ EXPECT_TRUE(M128.test(64 + 8));
+ EXPECT_FALSE(M128.test(64 + 16));
+}
+
+TEST(LaneBitmaskTest, Printing) {
+ std::string Str;
+ raw_string_ostream OS(Str);
+
+ LaneBitmask M(0xABCD);
+ OS << PrintLaneMask(M);
+ OS.flush();
+
+ EXPECT_TRUE(Str.find("ABCD") != std::string::npos);
+}
+
+TEST(LaneBitmaskTest, Hashing) {
+ // Test 64-bit version.
+ {
+ LaneMask64 A(0x1234);
+ LaneMask64 B(0x1234);
+ LaneMask64 C(0x5678);
+
+ EXPECT_EQ(std::hash<LaneMask64>{}(A), std::hash<LaneMask64>{}(B));
+ EXPECT_NE(std::hash<LaneMask64>{}(A), std::hash<LaneMask64>{}(C));
+ }
+
+ // Test 128-bit version.
+ {
+ LaneMask128 A = LaneMask128::getLane(100);
+ LaneMask128 B = LaneMask128::getLane(100);
+ LaneMask128 C = LaneMask128::getLane(50);
+
+ EXPECT_EQ(std::hash<LaneMask128>{}(A), std::hash<LaneMask128>{}(B));
+ EXPECT_NE(std::hash<LaneMask128>{}(A), std::hash<LaneMask128>{}(C));
+ }
+}
+
+TEST(LaneBitmaskTest, MultiWordOperations) {
+ // Test comprehensive operations on 128-bit LaneMask64 across word boundaries.
+ LaneMask128 M;
+ M.set(0);
+ M.set(64);
+ M.set(127);
+
+ EXPECT_EQ(M.getNumLanes(), 3u);
+ EXPECT_EQ(M.getHighestLane(), 127u);
+ EXPECT_TRUE(M.test(0));
+ EXPECT_TRUE(M.test(64));
+ EXPECT_TRUE(M.test(127));
+ EXPECT_FALSE(M.test(63));
+
+ // Test bitwise operations across word boundaries.
+ LaneMask128 N;
+ N.set(64);
+ N.set(100);
+
+ LaneMask128 Or = M | N;
+ EXPECT_EQ(Or.getNumLanes(), 4u);
+
+ LaneMask128 And = M & N;
+ EXPECT_EQ(And.getNumLanes(), 1u);
+ EXPECT_TRUE(And.test(64));
+
+ // Test NOT across multiple words.
+ LaneMask128 NotM = ~M;
+ EXPECT_FALSE(NotM.test(0));
+ EXPECT_FALSE(NotM.test(64));
+ EXPECT_FALSE(NotM.test(127));
+ EXPECT_TRUE(NotM.test(63));
+ EXPECT_TRUE(NotM.test(65));
+ EXPECT_EQ(NotM.getNumLanes(), LaneMask128::BitWidth - 3);
+}
+
+} // namespace
diff --git a/llvm/utils/TableGen/RegisterInfoEmitter.cpp b/llvm/utils/TableGen/RegisterInfoEmitter.cpp
index 0338d0b588352..a0585f8765f9c 100644
--- a/llvm/utils/TableGen/RegisterInfoEmitter.cpp
+++ b/llvm/utils/TableGen/RegisterInfoEmitter.cpp
@@ -671,7 +671,23 @@ static DiffVec &diffEncode(DiffVec &V, unsigned InitVal, Iter Begin, Iter End) {
static void printDiff16(raw_ostream &OS, int16_t Val) { OS << Val; }
static void printMask(raw_ostream &OS, LaneBitmask Val) {
- OS << "LaneBitmask(0x" << PrintLaneMask(Val) << ')';
+ constexpr unsigned NumWords = (LaneBitmask::BitWidth + 63) / 64;
+ // Check if all upper words beyond the first 64 bits are zero.
+ LaneBitmask UpperWords = ~LaneBitmask(~0ULL) & Val;
+ if (UpperWords.none()) {
+ OS << "LaneBitmask(0x" << PrintLaneMask(Val) << ')';
+ } else {
+ // Emit the explicit std::array constructor for multi-word values.
+ // Extract each 64-bit word by shifting and masking.
+ OS << "LaneBitmask(std::array<uint64_t, " << NumWords << ">{";
+ for (unsigned I = 0; I < NumWords; ++I) {
+ if (I > 0)
+ OS << ", ";
+ LaneBitmask Word = (Val >> (I * 64)) & LaneBitmask(~0ULL);
+ OS << "0x" << PrintLaneMask(Word);
+ }
+ OS << "})";
+ }
}
// Try to combine Idx's compose map into Vec if it is compatible.
@@ -881,13 +897,11 @@ void RegisterInfoEmitter::emitComposeSubRegIndexLaneMask(raw_ostream &OS,
" for (const MaskRolOp *Ops =\n"
" &LaneMaskComposeSequences[CompositeSequences[IdxA]];\n"
" Ops->Mask.any(); ++Ops) {\n"
- " LaneBitmask::Type M = LaneMask.getAsInteger() & "
- "Ops->Mask.getAsInteger();\n"
+ " LaneBitmask M = LaneMask & Ops->Mask;\n"
" if (unsigned S = Ops->RotateLeft)\n"
- " Result |= LaneBitmask((M << S) | (M >> (LaneBitmask::BitWidth - "
- "S)));\n"
+ " Result |= M.rotateLeft(S);\n"
" else\n"
- " Result |= LaneBitmask(M);\n"
+ " Result |= M;\n"
" }\n"
" return Result;\n"
"}\n\n";
@@ -903,12 +917,10 @@ void RegisterInfoEmitter::emitComposeSubRegIndexLaneMask(raw_ostream &OS,
" for (const MaskRolOp *Ops =\n"
" &LaneMaskComposeSequences[CompositeSequences[IdxA]];\n"
" Ops->Mask.any(); ++Ops) {\n"
- " LaneBitmask::Type M = LaneMask.getAsInteger();\n"
" if (unsigned S = Ops->RotateLeft)\n"
- " Result |= LaneBitmask((M >> S) | (M << (LaneBitmask::BitWidth - "
- "S)));\n"
+ " Result |= LaneMask.rotateRight(S);\n"
" else\n"
- " Result |= LaneBitmask(M);\n"
+ " Result |= LaneMask;\n"
" }\n"
" return Result;\n"
"}\n\n";
More information about the llvm-commits
mailing list