diff --git a/docs/docs/en/src/indexes/simq.md b/docs/docs/en/src/indexes/simq.md index 8a070111e6..1000ad0702 100644 --- a/docs/docs/en/src/indexes/simq.md +++ b/docs/docs/en/src/indexes/simq.md @@ -175,6 +175,7 @@ auto values = result->GetStatistics( | `simq_coarse_candidate_count` | Unique document candidates produced by coarse search before applying `rerank_k` | | `simq_rerank_candidate_count` | Candidates retained after applying `rerank_k` | | `simq_filtered_candidate_count` | Rerank candidates rejected by the supplied filter | +| `simq_rerank_batch_count` | Number of batched exact-rerank calls issued after filtering | | `simq_result_count` | Results returned after reranking and range/top-k limits | | `simq_limited_size_applied` | Whether `RangeSearch` truncated matches to `limited_size` | diff --git a/docs/docs/zh/src/indexes/simq.md b/docs/docs/zh/src/indexes/simq.md index a1a51f959d..fd298968a2 100644 --- a/docs/docs/zh/src/indexes/simq.md +++ b/docs/docs/zh/src/indexes/simq.md @@ -166,6 +166,7 @@ auto values = result->GetStatistics( | `simq_coarse_candidate_count` | 粗排产生且尚未应用 `rerank_k` 的唯一文档候选数 | | `simq_rerank_candidate_count` | 应用 `rerank_k` 后保留的候选数 | | `simq_filtered_candidate_count` | 被调用方 filter 排除的精排候选数 | +| `simq_rerank_batch_count` | filter 之后执行的批量精确精排调用次数 | | `simq_result_count` | 精排并应用 range/top-k 限制后最终返回的结果数 | | `simq_limited_size_applied` | `RangeSearch` 是否因 `limited_size` 截断结果 | diff --git a/src/algorithm/simq/simq.cpp b/src/algorithm/simq/simq.cpp index d2daf3eb9d..7c9534a235 100644 --- a/src/algorithm/simq/simq.cpp +++ b/src/algorithm/simq/simq.cpp @@ -59,6 +59,7 @@ dump_simq_statistics(const SearchStatistics& stats, uint64_t coarse_candidate_count, uint64_t rerank_candidate_count, uint64_t filtered_candidate_count, + uint64_t rerank_batch_count, uint64_t result_count, bool limited_size_applied) { auto json = JsonType::Parse(stats.Dump()); @@ -67,6 +68,7 @@ dump_simq_statistics(const SearchStatistics& stats, json["simq_coarse_candidate_count"].SetUint64(coarse_candidate_count); json["simq_rerank_candidate_count"].SetUint64(rerank_candidate_count); json["simq_filtered_candidate_count"].SetUint64(filtered_candidate_count); + json["simq_rerank_batch_count"].SetUint64(rerank_batch_count); json["simq_result_count"].SetUint64(result_count); json["simq_limited_size_applied"].SetBool(limited_size_applied); return json.Dump(); @@ -730,7 +732,7 @@ SIMQ::KnnSearch(const DatasetPtr& query, if (total_count_ == 0 || rep_hgraph_ == nullptr) { auto result = Dataset::Make(); - result->Statistics(dump_simq_statistics(stats, 0, 0, 0, 0, 0, 0, false)); + result->Statistics(dump_simq_statistics(stats, 0, 0, 0, 0, 0, 0, 0, false)); return result; } @@ -782,6 +784,8 @@ SIMQ::KnnSearch(const DatasetPtr& query, batch_ids.push_back(doc_id); } + const uint64_t rerank_batch_count = batch_ids.empty() ? 0 : 1; + // Single batched Query call (enables MultiRead in MultiVectorDataCell) if (!batch_ids.empty()) { std::vector batch_dists(batch_ids.size()); @@ -830,6 +834,7 @@ SIMQ::KnnSearch(const DatasetPtr& query, coarse_candidate_count, rerank_candidate_count, filtered_candidate_count, + rerank_batch_count, static_cast(result_ds->GetDim()), false)); return result_ds; @@ -850,7 +855,7 @@ SIMQ::RangeSearch(const DatasetPtr& query, if (total_count_ == 0 || rep_hgraph_ == nullptr) { auto result = Dataset::Make(); - result->Statistics(dump_simq_statistics(stats, 0, 0, 0, 0, 0, 0, false)); + result->Statistics(dump_simq_statistics(stats, 0, 0, 0, 0, 0, 0, 0, false)); return result; } @@ -889,18 +894,34 @@ SIMQ::RangeSearch(const DatasetPtr& query, auto computer = mv_codes_->FactoryComputer(&query_mvs[0]); std::vector> in_range; uint64_t filtered_candidate_count = 0; + + std::vector batch_ids; + batch_ids.reserve(coarse_results.size()); for (auto& [doc_id, _] : coarse_results) { if (filter != nullptr && !filter->CheckValid(this->label_table_->GetLabelById(doc_id))) { ++filtered_candidate_count; continue; } - float dist = 0.0F; + batch_ids.push_back(doc_id); + } + + const uint64_t rerank_batch_count = batch_ids.empty() ? 0 : 1; + if (!batch_ids.empty()) { + in_range.reserve(batch_ids.size()); + std::vector batch_dists(batch_ids.size()); QueryContext query_context{.stats = &stats, .distance_phase = DistanceEvaluationPhase::RERANK}; - mv_codes_->Query(&dist, computer, &doc_id, 1, &query_context); - ++stats.dist_cmp; - if (dist <= radius) { - in_range.emplace_back(dist, doc_id); + mv_codes_->Query(batch_dists.data(), + computer, + batch_ids.data(), + static_cast(batch_ids.size()), + &query_context); + stats.dist_cmp.fetch_add(static_cast(batch_ids.size()), + std::memory_order_relaxed); + for (uint64_t i = 0; i < batch_ids.size(); ++i) { + if (batch_dists[i] <= radius) { + in_range.emplace_back(batch_dists[i], batch_ids[i]); + } } } @@ -929,6 +950,7 @@ SIMQ::RangeSearch(const DatasetPtr& query, coarse_candidate_count, rerank_candidate_count, filtered_candidate_count, + rerank_batch_count, static_cast(in_range.size()), limited_size_applied)); return std::move(result_ds); diff --git a/tests/test_simq.cpp b/tests/test_simq.cpp index ae9f139949..f6df407d70 100644 --- a/tests/test_simq.cpp +++ b/tests/test_simq.cpp @@ -43,6 +43,7 @@ #include #include +#include #include #include #include @@ -275,6 +276,7 @@ struct SimqSearchStats { uint64_t coarse_candidate_count{0}; uint64_t rerank_candidate_count{0}; uint64_t filtered_candidate_count{0}; + uint64_t rerank_batch_count{0}; uint64_t result_count{0}; std::string limited_size_applied; }; @@ -294,9 +296,10 @@ get_simq_search_stats(const vsag::DatasetPtr& result) { "simq_coarse_candidate_count", "simq_rerank_candidate_count", "simq_filtered_candidate_count", + "simq_rerank_batch_count", "simq_result_count", "simq_limited_size_applied"}); - REQUIRE(values.size() == 8); + REQUIRE(values.size() == 9); SimqSearchStats stats; stats.dist_cmp = parse_u64(values[0]); @@ -305,8 +308,9 @@ get_simq_search_stats(const vsag::DatasetPtr& result) { stats.coarse_candidate_count = parse_u64(values[3]); stats.rerank_candidate_count = parse_u64(values[4]); stats.filtered_candidate_count = parse_u64(values[5]); - stats.result_count = parse_u64(values[6]); - stats.limited_size_applied = values[7]; + stats.rerank_batch_count = parse_u64(values[6]); + stats.result_count = parse_u64(values[7]); + stats.limited_size_applied = values[8]; return stats; } @@ -318,9 +322,27 @@ require_simq_search_stats(const vsag::DatasetPtr& result) { REQUIRE(stats.coarse_probe_count > 0); REQUIRE(stats.coarse_candidate_count >= stats.rerank_candidate_count); REQUIRE(stats.rerank_candidate_count >= stats.dist_cmp); + REQUIRE(stats.rerank_batch_count == 1); REQUIRE(stats.result_count == static_cast(result->GetDim())); } +class EvenLabelFilter : public vsag::Filter { +public: + [[nodiscard]] bool + CheckValid(int64_t id) const override { + return id % 2 == 0; + } +}; + +class RejectAllFilter : public vsag::Filter { +public: + [[nodiscard]] bool + CheckValid(int64_t id) const override { + (void)id; + return false; + } +}; + TEST_CASE("SIMQ: one centroid counts nested HGraph entry point", "[simq][statistics]") { TempFile tmp; std::array vector{}; @@ -560,6 +582,65 @@ TEST_CASE("SIMQ: range search", "[simq][range_search]") { auto stats = get_simq_search_stats(rr.value()); REQUIRE(stats.limited_size_applied == "true"); } + + SECTION("matches thresholded KNN with filtering and limited_size") { + auto filter = std::make_shared(); + auto threshold_param = fmt::format( + R"({{"threshold": {}, "simq": {{"coarse_k": 10, "rerank_k": 1000}}}})", radius); + auto expected = index->KnnSearch(one_query, BASE_DOCS, threshold_param, filter); + REQUIRE(expected.has_value()); + + auto actual = index->RangeSearch(one_query, radius, search_param, filter); + REQUIRE(actual.has_value()); + REQUIRE(actual.value()->GetDim() == expected.value()->GetDim()); + for (int64_t i = 0; i < actual.value()->GetDim(); ++i) { + REQUIRE(actual.value()->GetIds()[i] == expected.value()->GetIds()[i]); + REQUIRE(actual.value()->GetDistances()[i] == + Catch::Approx(expected.value()->GetDistances()[i]).margin(1e-6F)); + } + + auto expected_stats = get_simq_search_stats(expected.value()); + auto actual_stats = get_simq_search_stats(actual.value()); + REQUIRE(actual_stats.coarse_candidate_count == expected_stats.coarse_candidate_count); + REQUIRE(actual_stats.rerank_candidate_count == expected_stats.rerank_candidate_count); + REQUIRE(actual_stats.filtered_candidate_count == expected_stats.filtered_candidate_count); + REQUIRE(actual_stats.dist_cmp == expected_stats.dist_cmp); + REQUIRE(actual_stats.rerank_batch_count == 1); + + int64_t limited = std::min(3, expected.value()->GetDim()); + auto limited_result = index->RangeSearch(one_query, radius, search_param, filter, limited); + REQUIRE(limited_result.has_value()); + REQUIRE(limited_result.value()->GetDim() == limited); + for (int64_t i = 0; i < limited; ++i) { + REQUIRE(limited_result.value()->GetIds()[i] == expected.value()->GetIds()[i]); + REQUIRE(limited_result.value()->GetDistances()[i] == + Catch::Approx(expected.value()->GetDistances()[i]).margin(1e-6F)); + } + } + + SECTION("all-filtered candidates avoid rerank IO") { + auto result = index->RangeSearch( + one_query, radius, search_param, std::make_shared()); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == 0); + auto stats = get_simq_search_stats(result.value()); + REQUIRE(stats.rerank_candidate_count > 1); + REQUIRE(stats.filtered_candidate_count == stats.rerank_candidate_count); + REQUIRE(stats.dist_cmp == 0); + REQUIRE(stats.rerank_batch_count == 0); + REQUIRE(stats.result_count == 0); + } + + SECTION("empty radius still reranks in one batch") { + auto result = index->RangeSearch( + one_query, -std::numeric_limits::infinity(), search_param, FilterPtr{}); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == 0); + auto stats = get_simq_search_stats(result.value()); + REQUIRE(stats.dist_cmp > 1); + REQUIRE(stats.rerank_batch_count == 1); + REQUIRE(stats.result_count == 0); + } } TEST_CASE("SIMQ: parameter sweep on coarse_k and rerank_k", "[simq][sweep]") {