diff --git a/src/tt.cpp b/src/tt.cpp index c298d075..9463f613 100644 --- a/src/tt.cpp +++ b/src/tt.cpp @@ -4,8 +4,8 @@ namespace Clockwork { -static u16 shrink_key(HashKey key) { - return static_cast(key); +static u64 to_fragment(HashKey key) { + return key & TTCluster::FRAGMENT_MASK; } static u64 mulhi64(u64 a, u64 b) { @@ -64,15 +64,13 @@ TT::TT(size_t mb) : resize(mb, 1); } -std::optional TT::probe(const Position& pos, i32 ply) const { - size_t idx = mulhi64(pos.get_hash_key(), m_size); - const auto cluster = this->m_clusters[idx].load(); - const auto key = shrink_key(pos.get_hash_key()); +std::optional TT::probe(const Position& pos, i32 ply) { + size_t cluster_index = mulhi64(pos.get_hash_key(), m_size); + auto& cluster = this->m_clusters[cluster_index]; + auto fragment = to_fragment(pos.get_hash_key()); - for (const auto entry : cluster.entries) { - if (entry.key16 != key) { - continue; - } + if (auto entry_index = cluster.lookup(fragment); entry_index < TTCluster::ENTRY_COUNT) { + auto entry = cluster.load(entry_index); TTData data = {.eval = entry.eval, .move = entry.move, @@ -86,7 +84,7 @@ std::optional TT::probe(const Position& pos, i32 ply) const { return {}; } -TTClusterMemory* TT::addr_key(const u64 key) const { +TTCluster* TT::addr_key(const u64 key) const { size_t idx = mulhi64(key, m_size); return &this->m_clusters[idx]; } @@ -99,32 +97,34 @@ void TT::store(const Position& pos, Depth depth, bool ttpv, Bound bound) { - size_t cluster_index = mulhi64(pos.get_hash_key(), m_size); - auto cluster = this->m_clusters[cluster_index].load(); - const auto key = shrink_key(pos.get_hash_key()); - - auto tte = cluster.entries[0]; - size_t idx = 0; - - if (!(tte.key16 == 0 || tte.key16 == key)) { + size_t cluster_index = mulhi64(pos.get_hash_key(), m_size); + auto& cluster = this->m_clusters[cluster_index]; + auto fragment = to_fragment(pos.get_hash_key()); + + TTEntry tte; + size_t entry_index; + bool fragment_match = false; + + if ((entry_index = cluster.lookup(fragment)) < TTCluster::ENTRY_COUNT) { + fragment_match = true; + tte = cluster.load(entry_index); + } else if ((entry_index = cluster.lookup(0)) < TTCluster::ENTRY_COUNT) { + tte = TTEntry{}; + } else { + tte = cluster.load(0); + entry_index = 0; for (size_t i = 1; i < 3; ++i) { - const auto entry = cluster.entries[i]; - - if (entry.key16 == 0 || entry.key16 == key) { - tte = entry; - idx = i; - break; - } + auto entry = cluster.load(i); if (tte.depth - ((MAX_AGE + m_age - tte.age()) & AGE_MASK) * 4 > entry.depth - ((MAX_AGE + m_age - entry.age()) & AGE_MASK) * 4) { - tte = entry; - idx = i; + tte = entry; + entry_index = i; } } } - if (move == Move::none() && tte.key16 == key) { + if (move == Move::none() && fragment_match) { // if we don't have a best move, and the entry is for the same position, // then we should retain the best move from the previous entry. move = tte.move; @@ -147,9 +147,8 @@ void TT::store(const Position& pos, depth + insert_flag_bonus + (age_differential * age_differential) / 4; //+ i32::from(pv); i32 record_prority = tte.depth + record_flag_bonus; - if (tte.key16 != key || (bound == Bound::Exact && tte.bound() != Bound::Exact) + if (!fragment_match || (bound == Bound::Exact && tte.bound() != Bound::Exact) || insert_priority * 3 >= record_prority * 2) { - tte.key16 = key; tte.move = move; tte.score = score_to_tt(score, ply); tte.eval = static_cast(eval); @@ -157,18 +156,18 @@ void TT::store(const Position& pos, tte.info = make_tt_info(ttpv, bound, m_age); // write back - cluster.entries[idx] = tte; - this->m_clusters[cluster_index].store(cluster); + cluster.store(entry_index, tte); + cluster.set_fragment(entry_index, fragment); } } void TT::resize(size_t mb, usize thread_count) { size_t bytes = mb * 1024 * 1024; - size_t entries = bytes / sizeof(TTClusterMemory); + size_t entries = bytes / sizeof(TTCluster); m_size = entries; - m_clusters = make_unique_for_overwrite_huge_page(m_size); + m_clusters = make_unique_for_overwrite_huge_page(m_size); clear(thread_count); } @@ -201,7 +200,7 @@ void TT::increment_age() { this->m_age = new_age; } -i32 TT::hashfull() const { +i32 TT::hashfull() { if (m_size == 0) { return 0; } @@ -215,8 +214,9 @@ i32 TT::hashfull() const { } for (size_t i = 0; i < num_to_probe; ++i) { - const auto cluster = this->m_clusters[i].load(); - for (const auto& entry : cluster.entries) { + auto& cluster = this->m_clusters[i]; + for (size_t entry_index = 0; entry_index < TTCluster::ENTRY_COUNT; entry_index++) { + auto entry = cluster.load(entry_index); if (entry.age() == m_age && entry.bound() != Bound::None) { occupied_count++; } diff --git a/src/tt.hpp b/src/tt.hpp index f55760f0..b96fa5ea 100644 --- a/src/tt.hpp +++ b/src/tt.hpp @@ -16,7 +16,6 @@ enum Bound : u8 { }; struct TTEntry { - u16 key16; Move move; i16 score; i16 eval; @@ -34,35 +33,55 @@ struct TTEntry { } }; -struct TTCluster { - std::array entries; - std::array padding; -}; +struct alignas(32) TTCluster { +public: + static constexpr usize ENTRY_COUNT = 3; + static constexpr usize FRAGMENT_WIDTH = 21; + static constexpr u64 FRAGMENT_MASK = (1 << FRAGMENT_WIDTH) - 1; + + [[nodiscard]] TTEntry load(usize index) { + u64 raw = std::atomic_ref{entries[index]}.load(std::memory_order_relaxed); + return std::bit_cast(raw); + } -struct TTClusterMemory { - alignas(32) std::array data; + void store(usize index, TTEntry entry) { + u64 raw = std::bit_cast(entry); + std::atomic_ref{entries[index]}.store(raw, std::memory_order_relaxed); + } - [[nodiscard]] auto load() -> TTCluster { - std::array out; - out[0] = std::atomic_ref{this->data[0]}.load(std::memory_order_relaxed); - out[1] = std::atomic_ref{this->data[1]}.load(std::memory_order_relaxed); - out[2] = std::atomic_ref{this->data[2]}.load(std::memory_order_relaxed); - out[3] = std::atomic_ref{this->data[3]}.load(std::memory_order_relaxed); - return std::bit_cast(out); + usize lookup(u64 fragment) { + u64 needle = fragment * FRAGMENTS_LSB; + u64 haystack = std::atomic_ref{fragments}.load(std::memory_order_relaxed); + u64 zeros = needle ^ haystack; + u64 matches = (zeros - FRAGMENTS_LSB) & ~zeros & FRAGMENTS_MSB; + return static_cast(std::countr_zero(matches)) / FRAGMENT_WIDTH; } - auto store(TTCluster cluster) { - std::array mem = std::bit_cast>(cluster); - std::atomic_ref{this->data[0]}.store(mem[0], std::memory_order_relaxed); - std::atomic_ref{this->data[1]}.store(mem[1], std::memory_order_relaxed); - std::atomic_ref{this->data[2]}.store(mem[2], std::memory_order_relaxed); - std::atomic_ref{this->data[3]}.store(mem[3], std::memory_order_relaxed); + u64 get_fragment(usize index) { + u64 f = std::atomic_ref{fragments}.load(std::memory_order_relaxed); + usize shift = FRAGMENT_WIDTH * index; + return (f >> shift) & FRAGMENT_MASK; } + + void set_fragment(usize index, u64 fragment) { + u64 f = std::atomic_ref{fragments}.load(std::memory_order_relaxed); + usize shift = FRAGMENT_WIDTH * index; + f &= ~(FRAGMENT_MASK << shift); + f |= fragment << shift; + std::atomic_ref{fragments}.store(f, std::memory_order_relaxed); + } + +private: + static constexpr u64 FRAGMENTS_LSB = 0x0000'0400'0020'0001; + static constexpr u64 FRAGMENTS_MSB = 0x4000'0200'0010'0000; + + std::array entries; + u64 fragments; }; -static_assert(sizeof(TTEntry) == 10 * sizeof(u8)); +static_assert(sizeof(TTEntry) == 8 * sizeof(u8)); +static_assert(sizeof(TTCluster) == 32 * sizeof(u8)); static_assert(sizeof(TTCluster) == 32 * sizeof(u8)); -static_assert(sizeof(TTClusterMemory) == 32 * sizeof(u8)); struct TTData { Value eval; @@ -92,7 +111,7 @@ class TT { TT(size_t mb = DEFAULT_SIZE_MB); - std::optional probe(const Position& position, i32 ply) const; + std::optional probe(const Position& position, i32 ply); void store(const Position& position, i32 ply, Value eval, @@ -104,14 +123,14 @@ class TT { void resize(size_t mb, usize thread_count); void clear(usize thread_count); void increment_age(); - i32 hashfull() const; - TTClusterMemory* addr_key(const u64 key) const; + i32 hashfull(); + TTCluster* addr_key(const u64 key) const; private: - unique_ptr_huge_page m_clusters; - size_t m_size; - u8 m_age; + unique_ptr_huge_page m_clusters; + size_t m_size; + u8 m_age; }; } // namespace Clockwork