diff --git a/src/index/hnsw/faiss_hnsw.cc b/src/index/hnsw/faiss_hnsw.cc index 57c4e1831..c7296693b 100644 --- a/src/index/hnsw/faiss_hnsw.cc +++ b/src/index/hnsw/faiss_hnsw.cc @@ -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 dists; @@ -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::max() - : 1.0f; + const bool high_filter_ratio = + (bitset_in.count() >= (index->ntotal * HnswSearchThresholds::kHnswSearchKnnBFFilterThreshold)); + workspace.accumulated_alpha = high_filter_ratio ? std::numeric_limits::max() : 1.0f; + workspace.use_brute_force = + (bitset_in.count() >= (index->ntotal * HnswSearchThresholds::kHnswSearchIteratorBFFilterThreshold)); // set up a visitor workspace.graph_visitor = DummyVisitor(); @@ -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 + 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 void next_batch(std::function&)> batch_handler, FilterT& filter) { @@ -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? diff --git a/src/index/hnsw/impl/IndexConditionalWrapper.h b/src/index/hnsw/impl/IndexConditionalWrapper.h index dd6b000ee..53c5e187b 100644 --- a/src/index/hnsw/impl/IndexConditionalWrapper.h +++ b/src/index/hnsw/impl/IndexConditionalWrapper.h @@ -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; }; diff --git a/tests/ut/test_hnsw_rabitq_acceptance.cc b/tests/ut/test_hnsw_rabitq_acceptance.cc index 1f6123ba6..45b6447ad 100644 --- a/tests/ut/test_hnsw_rabitq_acceptance.cc +++ b/tests/ut/test_hnsw_rabitq_acceptance.cc @@ -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 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()); } } diff --git a/tests/ut/test_iterator.cc b/tests/ut/test_iterator.cc index 2021baf85..bbf01d8f4 100644 --- a/tests/ut/test_iterator.cc +++ b/tests/ut/test_iterator.cc @@ -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>( + {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(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(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(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(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 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(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") {