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
62 changes: 55 additions & 7 deletions src/index/hnsw/faiss_hnsw.cc
Original file line number Diff line number Diff line change
Expand Up @@ -867,6 +867,12 @@ struct FaissHnswIteratorWorkspace {
// hnsw layer.
bool initial_search_done = false;

// whether to scan the unfiltered points instead of traversing the graph.
// At high filter ratios the graph traversal visits almost every node while
// looking for the few unfiltered ones, so a sequential scan over the
// unfiltered points is much cheaper (kHnswSearchIteratorBFFilterThreshold).
bool use_brute_force = false;

// accumulated elements
std::vector<DistId> dists;

Expand All @@ -888,10 +894,11 @@ class FaissHnswIterator : public IndexIterator {
labels{labels_in},
label_to_internal_offset(label_to_internal_offset_in),
mv_base_offset(mv_base_offset_in) {
workspace.accumulated_alpha =
(bitset_in.count() >= (index->ntotal * HnswSearchThresholds::kHnswSearchKnnBFFilterThreshold))
? std::numeric_limits<float>::max()
: 1.0f;
const bool high_filter_ratio =
(bitset_in.count() >= (index->ntotal * HnswSearchThresholds::kHnswSearchKnnBFFilterThreshold));
workspace.accumulated_alpha = high_filter_ratio ? std::numeric_limits<float>::max() : 1.0f;
workspace.use_brute_force =
(bitset_in.count() >= (index->ntotal * HnswSearchThresholds::kHnswSearchIteratorBFFilterThreshold));

// set up a visitor
workspace.graph_visitor = DummyVisitor();
Expand Down Expand Up @@ -996,6 +1003,41 @@ class FaissHnswIterator : public IndexIterator {
}

protected:
// Computes distances to every unfiltered point and stores them into workspace.dists.
// The stored values follow the sign convention of workspace.qdis (negated for similarity
// metrics), so the post-processing in next_batch() applies unchanged. If a refine index is
// available, its exact distances are used, the same as the brute-force kNN Search does.
template <typename FilterT>
void
brute_force_scan(FilterT& filter) {
// workspace.qdis is already sign-wrapped; workspace.qdis_refine is not.
const bool use_refine = (workspace.qdis_refine != nullptr);
auto& qdis = use_refine ? *workspace.qdis_refine : *workspace.qdis;
const float sign =
(use_refine && faiss::cppcontrib::knowhere::is_similarity_metric(index->metric_type)) ? -1.0f : 1.0f;
const faiss::idx_t ntotal = index->ntotal;

faiss::idx_t ids[4];
size_t n_ids = 0;
for (faiss::idx_t i = 0; i < ntotal; i++) {
if (!filter.is_member(i)) {
continue;
}
ids[n_ids++] = i;
if (n_ids == 4) {
float dis[4];
qdis.distances_batch_4(ids[0], ids[1], ids[2], ids[3], dis[0], dis[1], dis[2], dis[3]);
for (size_t j = 0; j < 4; j++) {
workspace.dists.emplace_back(ids[j], sign * dis[j]);
}
n_ids = 0;
}
}
for (size_t j = 0; j < n_ids; j++) {
workspace.dists.emplace_back(ids[j], sign * qdis(ids[j]));
}
}

template <typename FilterT>
void
next_batch(std::function<void(const std::vector<DistId>&)> batch_handler, FilterT& filter) {
Expand All @@ -1013,9 +1055,15 @@ class FaissHnswIterator : public IndexIterator {
// whether to track hnsw stats
constexpr bool track_hnsw_stats = true;

// accumulate elements for a new batch?
if (!workspace.initial_search_done) {
// yes
if (workspace.use_brute_force) {
// emit every unfiltered point in a single batch; the base class keeps
// them in a heap and returns them in order.
if (!workspace.initial_search_done) {
brute_force_scan(filter);
workspace.initial_search_done = true;
}
} else if (!workspace.initial_search_done) {
// accumulate elements for a new batch
faiss::cppcontrib::knowhere::HNSWStats stats;

// is the graph empty?
Expand Down
4 changes: 4 additions & 0 deletions src/index/hnsw/impl/IndexConditionalWrapper.h
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,10 @@ struct SearchParametersHNSWWrapper;
struct HnswSearchThresholds {
static constexpr float kHnswSearchKnnBFFilterThreshold = 0.93f;
static constexpr float kHnswSearchRangeBFFilterThreshold = 0.97f;
// AnnIterator switches to a scan over the unfiltered points at this filter ratio. It is higher
// than the kNN threshold because, at typical ef values, the graph traversal of an iterator
// stays cheaper than the brute force up to ~97%.
static constexpr float kHnswSearchIteratorBFFilterThreshold = 0.97f;
static constexpr float kHnswSearchBFTopkThreshold = 0.5f;
};

Expand Down
13 changes: 10 additions & 3 deletions tests/ut/test_hnsw_rabitq_acceptance.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1147,9 +1147,16 @@ TEST_CASE("RaBitQ filtered results and exhausted iterators match full-code refer
REQUIRE(distance == Catch::Approx(reference[id]).margin(1e-4));
REQUIRE(iter_ids.size() <= size_t(n - excluded));
}
auto expected_reachable = reachable;
for (int i = 0; i < excluded; ++i) expected_reachable.erase(i);
REQUIRE(iter_ids == expected_reachable);
// Below the iterator brute-force threshold the iterator traverses the graph and returns
// the reachable points; at or above it, it scans and returns every unfiltered point.
std::set<int64_t> expected_iter_ids;
if (excluded >= n * knowhere::HnswSearchThresholds::kHnswSearchIteratorBFFilterThreshold) {
for (int i = excluded; i < n; ++i) expected_iter_ids.insert(i);
} else {
expected_iter_ids = reachable;
for (int i = 0; i < excluded; ++i) expected_iter_ids.erase(i);
}
REQUIRE(iter_ids == expected_iter_ids);
REQUIRE_FALSE(it->HasNext().value());
}
}
Expand Down
61 changes: 61 additions & 0 deletions tests/ut/test_iterator.cc
Original file line number Diff line number Diff line change
Expand Up @@ -366,6 +366,67 @@ TEST_CASE("Test Iterator Mem Index With Float Vector", "[float metrics]") {
}
}

// At high filter ratios the HNSW iterator scans the unfiltered points instead of traversing the graph,
// so it must return every unfiltered point, in the same order as a brute-force kNN search.
SECTION("Test HNSW iterator falls back to brute force at high filter ratio") {
using std::make_tuple;
auto [name, gen] = GENERATE_REF(table<std::string, std::function<knowhere::Json()>>(
{make_tuple(knowhere::IndexEnum::INDEX_HNSW, hnsw_gen),
make_tuple(knowhere::IndexEnum::INDEX_HNSW_SQ, hnsw_sq_gen),
make_tuple(knowhere::IndexEnum::INDEX_HNSW_SQ, hnsw_sq_refine_flat_gen),
make_tuple(knowhere::IndexEnum::INDEX_HNSW_PQ, hnsw_pq_gen),
make_tuple(knowhere::IndexEnum::INDEX_HNSW_PQ, hnsw_pq_refine_flat_gen)}));
auto idx = knowhere::IndexFactory::Instance().Create<knowhere::fp32>(name, version).value();
auto cfg_json = gen().dump();
CAPTURE(name, cfg_json);
knowhere::Json json = knowhere::Json::parse(cfg_json);
REQUIRE(idx.Type() == name);
REQUIRE(idx.Build(train_ds, json) == knowhere::Status::success);

// HNSW (flat) and refined indexes scan with exact distances.
const bool is_exact = (name == knowhere::IndexEnum::INDEX_HNSW) || json.value("refine", false);
std::vector<std::function<std::vector<uint8_t>(size_t, size_t)>> gen_bitset_funcs = {
GenerateBitsetWithFirstTbitsSet, GenerateBitsetWithRandomTbitsSet};
const auto bitset_percentages = {0.98f, 0.99f};
for (const float percentage : bitset_percentages) {
for (const auto& gen_func : gen_bitset_funcs) {
const size_t n_filtered = percentage * nb;
auto bitset_data = gen_func(nb, n_filtered);
knowhere::BitsetView bitset(bitset_data.data(), nb);
const size_t n_unfiltered = nb - n_filtered;

// reference: exact brute force for the exact indexes, and the index's own kNN Search for the
// quantized ones (which brute-forces over the same quantized storage at this filter ratio).
auto ref = is_exact ? knowhere::BruteForce::Search<knowhere::fp32>(train_ds, query_ds, json, bitset)
: idx.Search(query_ds, json, bitset);
REQUIRE(ref.has_value());
const auto ref_ids = ref.value()->GetIds();
const size_t expected_k = std::min<size_t>(topk, n_unfiltered);

auto its = idx.AnnIterator(query_ds, json, bitset);
REQUIRE(its.has_value());
size_t n_hit = 0;
for (int64_t i = 0; i < nq; ++i) {
auto& it = its.value()[i];
std::unordered_set<int64_t> ref_topk(ref_ids + i * topk, ref_ids + i * topk + expected_k);
// every unfiltered point is reachable through the iterator.
size_t n_returned = 0;
while (it->HasNext().value()) {
auto [id, dist] = it->Next().value();
REQUIRE(!bitset.test(id));
if (n_returned < expected_k && ref_topk.count(id) > 0) {
n_hit++;
}
n_returned++;
}
REQUIRE(n_returned == n_unfiltered);
}
const float recall = static_cast<float>(n_hit) / (nq * expected_k);
REQUIRE(recall > 0.99f);
}
}
}

// certain unit tests are disabled, because they are way too slow at this moment
// todo: re-enable later
SECTION("Test Search with Bitset using iterator insufficient results") {
Expand Down
Loading