Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions include/rabitqlib/index/hnsw/hnsw.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -777,7 +777,7 @@ inline void HierarchicalNSW::add_point(
inline maxheap<std::pair<float, PID>> HierarchicalNSW::search_base_layer(
PID ep_id, PID cur_c, int layer
) {
HashBasedBooleanSet* vl = visited_list_pool_->get_free_vislist();
VisitedSet* vl = visited_list_pool_->get_free_vislist();

maxheap<std::pair<float, PID>> top_candidates;
minheap<std::pair<float, PID>> candidate_set;
Expand Down Expand Up @@ -1225,7 +1225,7 @@ inline void HierarchicalNSW::searchBaseLayerST_AdaptiveRerankOptDirect(
[[maybe_unused]] const float* query,
BoundedKNN& boundedKNN
) {
HashBasedBooleanSet* vl = visited_list_pool_->get_free_vislist();
VisitedSet* vl = visited_list_pool_->get_free_vislist();

// Use our bounded priority queue instead of the maxheap.
buffer::SearchBuffer<float> candidate_set(ef);
Expand Down
27 changes: 9 additions & 18 deletions include/rabitqlib/index/symqg/qg.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -20,12 +20,12 @@
#include "rabitqlib/quantization/rabitq.hpp"
#include "rabitqlib/utils/array.hpp"
#include "rabitqlib/utils/buffer.hpp"
#include "rabitqlib/utils/hashset.hpp"
#include "rabitqlib/utils/io.hpp"
#include "rabitqlib/utils/memory.hpp"
#include "rabitqlib/utils/rotator.hpp"
#include "rabitqlib/utils/space.hpp"
#include "rabitqlib/utils/visited_pool.hpp"
#include "rabitqlib/utils/visited_set.hpp"

namespace rabitqlib::symqg {

Expand Down Expand Up @@ -94,25 +94,16 @@ class QuantizedGraph {
);
}

void find_candidates(
PID,
size_t,
std::vector<AnnCandidate<T>>&,
HashBasedBooleanSet&,
const std::vector<uint32_t>&
) const;
void
find_candidates(PID, size_t, std::vector<AnnCandidate<T>>&, VisitedSet&, const std::vector<uint32_t>&)
const;

void update_qg(PID, const std::vector<AnnCandidate<T>>&);

void update_results(buffer::SearchBuffer<T>&, HashBasedBooleanSet&, const T*);
void update_results(buffer::SearchBuffer<T>&, VisitedSet&, const T*);

void scan_neighbors(
const BatchQuery<T>&,
PID,
T*,
buffer::SearchBuffer<T>&,
HashBasedBooleanSet&,
size_t
const BatchQuery<T>&, PID, T*, buffer::SearchBuffer<T>&, VisitedSet&, size_t
) const;

public:
Expand Down Expand Up @@ -352,7 +343,7 @@ void QuantizedGraph<T>::scan_neighbors(
PID data_id,
T* est_dist,
buffer::SearchBuffer<T>& search_pool,
HashBasedBooleanSet& vis,
VisitedSet& vis,
size_t cur_degree
) const {
const auto* batch_data = get_batch_data(data_id);
Expand All @@ -378,7 +369,7 @@ void QuantizedGraph<T>::scan_neighbors(

template <typename T>
inline void QuantizedGraph<T>::update_results(
buffer::SearchBuffer<T>& result_pool, HashBasedBooleanSet& vis, const T* query
buffer::SearchBuffer<T>& result_pool, VisitedSet& vis, const T* query
) {
if (result_pool.is_full()) {
return;
Expand Down Expand Up @@ -433,7 +424,7 @@ inline void QuantizedGraph<T>::find_candidates(
PID cur_id,
size_t search_ef,
std::vector<AnnCandidate<T>>& results,
HashBasedBooleanSet& vis,
VisitedSet& vis,
const std::vector<uint32_t>& degrees
) const {
const T* query = get_vector(cur_id);
Expand Down
12 changes: 6 additions & 6 deletions include/rabitqlib/index/symqg/qg_builder.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -11,9 +11,9 @@

#include "rabitqlib/defines.hpp"
#include "rabitqlib/index/symqg/qg.hpp"
#include "rabitqlib/utils/hashset.hpp"
#include "rabitqlib/utils/space.hpp"
#include "rabitqlib/utils/tools.hpp"
#include "rabitqlib/utils/visited_set.hpp"

namespace rabitqlib::symqg {
constexpr size_t kMaxBsIter = 5; // max iter for binary search of pruning bar
Expand All @@ -37,9 +37,9 @@ class QGBuilder {
static constexpr size_t kMaxPrunedSize =
300; // max number of recorded pruned candidates
std::vector<CandidateList> new_neighbors_; // new neighbors for current iteration
std::vector<CandidateList> pruned_neighbors_; // recorded pruned neighbors
std::vector<HashBasedBooleanSet> visited_list_; // list of visited hash set
std::vector<uint32_t> degrees_; // record degree of qg
std::vector<CandidateList> pruned_neighbors_; // recorded pruned neighbors
std::vector<VisitedSet> visited_list_; // per-thread visited sets
std::vector<uint32_t> degrees_; // record degree of qg
void random_init();
void search_new_neighbors(bool refine);
void heuristic_prune(PID, CandidateList&, CandidateList&, bool);
Expand Down Expand Up @@ -67,7 +67,7 @@ class QGBuilder {
, pruned_neighbors_(qg_.num_vertices())
, visited_list_(
num_threads_,
HashBasedBooleanSet(std::min(ef_build_ * ef_build_, num_nodes_ / 10))
VisitedSet(num_nodes_, std::min(ef_build_ * ef_build_, num_nodes_ / 10))
)
, degrees_(qg_.num_vertices(), degree_bound_) {
omp_set_num_threads(static_cast<int>(num_threads_));
Expand Down Expand Up @@ -236,7 +236,7 @@ inline void QGBuilder::search_new_neighbors(bool refine) {
PID cur_id = i;
auto tid = omp_get_thread_num();
CandidateList candidates;
HashBasedBooleanSet& vis = visited_list_[tid];
VisitedSet& vis = visited_list_[tid];
candidates.reserve(2 * kMaxCandidatePoolSize);
vis.clear();
qg_.find_candidates(cur_id, ef_build_, candidates, vis, degrees_);
Expand Down
101 changes: 4 additions & 97 deletions include/rabitqlib/utils/hashset.hpp
Original file line number Diff line number Diff line change
@@ -1,101 +1,8 @@
// This code is modified based on NGT from Yahoo Japan
// https://github.com/yahoojapan/NGT
//
// Copyright (C) 2015 Yahoo Japan Corporation
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//

// This compatibility header preserves the original visited-set API.
#pragma once

#include <climits>
#include <cstring>
#include <iostream>
#include <unordered_set>

#include "rabitqlib/defines.hpp"
#include "rabitqlib/utils/memory.hpp"
#include "rabitqlib/utils/visited_set_hash.hpp"

namespace rabitqlib {
/**
* @brief hash set to record visited vertices
*
*/
class HashBasedBooleanSet {
private:
size_t table_size_ = 0;
PID mask_ = 0;
std::vector<PID, memory::AlignedAllocator<PID>> table_;
std::unordered_set<PID> stl_hash_;

[[nodiscard]] auto hash1(const PID value) const { return value & mask_; }

public:
HashBasedBooleanSet() = default;
~HashBasedBooleanSet() = default;

HashBasedBooleanSet(const HashBasedBooleanSet&) = default;
HashBasedBooleanSet(HashBasedBooleanSet&&) noexcept = default;
HashBasedBooleanSet& operator=(HashBasedBooleanSet&&) noexcept = default;

explicit HashBasedBooleanSet(size_t size) {
size_t bit_size = 0;
size_t bit = size;
while (bit != 0) {
bit_size++;
bit >>= 1;
}
size_t bucket_size = 0x1 << ((bit_size + 4) / 2 + 3);
initialize(bucket_size);
}

void initialize(const size_t table_size) {
table_size_ = table_size;
mask_ = static_cast<PID>(table_size_ - 1);
const PID check_val = hash1(static_cast<PID>(table_size));
if (check_val != 0) {
std::cerr << "[WARN] table size is not 2^N : " << table_size << '\n';
}

table_ = std::vector<PID, memory::AlignedAllocator<PID>>(table_size);
std::fill(table_.begin(), table_.end(), kPidMax);
stl_hash_.clear();
}

void clear() {
std::fill(table_.begin(), table_.end(), kPidMax);
stl_hash_.clear();
}

// get if data_id is in the hashset
[[nodiscard]] bool get(PID data_id) const {
PID val = this->table_[hash1(data_id)];
if (val == data_id) {
return true;
}
return (val != kPidMax && stl_hash_.find(data_id) != stl_hash_.end());
}

void set(PID data_id) {
PID& val = table_[hash1(data_id)];
if (val == data_id) {
return;
}
if (val == kPidMax) {
val = data_id;
} else {
stl_hash_.emplace(data_id);
}
}
};
} // namespace rabitqlib
using HashBasedBooleanSet = HashBasedVisitedSet;
} // namespace rabitqlib
20 changes: 10 additions & 10 deletions include/rabitqlib/utils/visited_pool.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,48 +2,48 @@
#include <deque>
#include <mutex>

#include "rabitqlib/utils/hashset.hpp"
#include "rabitqlib/utils/visited_set.hpp"

namespace rabitqlib {
class VisitedListPool {
std::deque<HashBasedBooleanSet*> pool_;
std::deque<VisitedSet*> pool_;
std::mutex poolguard_;
size_t numelements_;

public:
VisitedListPool(size_t initpoolsize, size_t max_elements) {
numelements_ = max_elements / 10;
numelements_ = max_elements;
for (size_t i = 0; i < initpoolsize; i++) {
pool_.push_front(new HashBasedBooleanSet(numelements_));
pool_.push_front(new VisitedSet(numelements_, numelements_ / 10));
}
}

HashBasedBooleanSet* get_free_vislist() {
HashBasedBooleanSet* rez;
VisitedSet* get_free_vislist() {
VisitedSet* rez;
{
std::unique_lock<std::mutex> lock(poolguard_);
if (pool_.size() > 0) {
rez = pool_.front();
pool_.pop_front();
} else {
rez = new HashBasedBooleanSet(numelements_);
rez = new VisitedSet(numelements_, numelements_ / 10);
}
}
rez->clear();
return rez;
}

void release_vis_list(HashBasedBooleanSet* vl) {
void release_vis_list(VisitedSet* vl) {
std::unique_lock<std::mutex> lock(poolguard_);
pool_.push_front(vl);
}

~VisitedListPool() {
while (pool_.size() > 0) {
HashBasedBooleanSet* rez = pool_.front();
VisitedSet* rez = pool_.front();
pool_.pop_front();
::delete rez;
}
}
};
} // namespace rabitqlib
} // namespace rabitqlib
46 changes: 46 additions & 0 deletions include/rabitqlib/utils/visited_set.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
#pragma once

#include <cstddef>
#include <type_traits>

#include "rabitqlib/defines.hpp"
#include "rabitqlib/utils/visited_set_epoch.hpp"
#include "rabitqlib/utils/visited_set_hash.hpp"

namespace rabitqlib {
namespace detail {
template <typename T>
struct CheckVisitedSet {
static_assert(
std::is_constructible<T, size_t>::value,
"visited set must be constructible from a size_t id space"
);
static_assert(
std::is_constructible<T, size_t, size_t>::value,
"visited set must be constructible from (num_elements, size_hint)"
);
static_assert(
std::is_same<decltype(&T::initialize), void (T::*)(size_t, size_t)>::value,
"visited set must declare: void initialize(size_t num_elements, size_t size_hint)"
);
static_assert(
std::is_same<decltype(&T::clear), void (T::*)()>::value,
"visited set must declare: void clear()"
);
static_assert(
std::is_same<decltype(&T::get), bool (T::*)(PID) const>::value,
"visited set must declare: bool get(PID) const"
);
static_assert(
std::is_same<decltype(&T::set), void (T::*)(PID)>::value,
"visited set must declare: void set(PID)"
);
static constexpr bool kValue = true;
};
} // namespace detail

static_assert(detail::CheckVisitedSet<EpochBasedVisitedSet>::kValue, "");
static_assert(detail::CheckVisitedSet<HashBasedVisitedSet>::kValue, "");

using VisitedSet = HashBasedVisitedSet;
} // namespace rabitqlib
46 changes: 46 additions & 0 deletions include/rabitqlib/utils/visited_set_epoch.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
#pragma once

#include <algorithm>
#include <cstdint>
#include <vector>

#include "rabitqlib/defines.hpp"
#include "rabitqlib/utils/memory.hpp"

namespace rabitqlib {
class EpochBasedVisitedSet {
private:
std::vector<uint16_t, memory::AlignedAllocator<uint16_t>> stamp_;
uint16_t cur_ = 0;

public:
EpochBasedVisitedSet() = default;
~EpochBasedVisitedSet() = default;

EpochBasedVisitedSet(const EpochBasedVisitedSet&) = default;
EpochBasedVisitedSet& operator=(const EpochBasedVisitedSet&) = default;
EpochBasedVisitedSet(EpochBasedVisitedSet&&) noexcept = default;
EpochBasedVisitedSet& operator=(EpochBasedVisitedSet&&) noexcept = default;

explicit EpochBasedVisitedSet(size_t num_elements, size_t = 0) {
initialize(num_elements);
}

void initialize(size_t num_elements, size_t = 0) {
stamp_.assign(num_elements, 0);
cur_ = 0;
clear();
}

void clear() {
if (++cur_ == 0) {
std::fill(stamp_.begin(), stamp_.end(), 0);
cur_ = 1;
}
}

[[nodiscard]] bool get(PID data_id) const { return stamp_[data_id] == cur_; }

void set(PID data_id) { stamp_[data_id] = cur_; }
};
} // namespace rabitqlib
Loading
Loading