Skip to content
Open
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
1 change: 1 addition & 0 deletions docs/docs/en/src/indexes/simq.md
Original file line number Diff line number Diff line change
Expand Up @@ -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` |

Expand Down
1 change: 1 addition & 0 deletions docs/docs/zh/src/indexes/simq.md
Original file line number Diff line number Diff line change
Expand Up @@ -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` 截断结果 |

Expand Down
36 changes: 29 additions & 7 deletions src/algorithm/simq/simq.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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());
Expand All @@ -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();
Expand Down Expand Up @@ -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;
}

Expand Down Expand Up @@ -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<float> batch_dists(batch_ids.size());
Expand Down Expand Up @@ -830,6 +834,7 @@ SIMQ::KnnSearch(const DatasetPtr& query,
coarse_candidate_count,
rerank_candidate_count,
filtered_candidate_count,
rerank_batch_count,
static_cast<uint64_t>(result_ds->GetDim()),
false));
return result_ds;
Expand All @@ -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;
}

Expand Down Expand Up @@ -889,18 +894,34 @@ SIMQ::RangeSearch(const DatasetPtr& query,
auto computer = mv_codes_->FactoryComputer(&query_mvs[0]);
std::vector<std::pair<float, InnerIdType>> in_range;
uint64_t filtered_candidate_count = 0;

std::vector<InnerIdType> 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<float> 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<InnerIdType>(batch_ids.size()),
Comment thread
CharlesXu-HQ marked this conversation as resolved.
Comment thread
CharlesXu-HQ marked this conversation as resolved.
&query_context);
stats.dist_cmp.fetch_add(static_cast<uint32_t>(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]);
}
}
}

Expand Down Expand Up @@ -929,6 +950,7 @@ SIMQ::RangeSearch(const DatasetPtr& query,
coarse_candidate_count,
rerank_candidate_count,
filtered_candidate_count,
rerank_batch_count,
static_cast<uint64_t>(in_range.size()),
limited_size_applied));
return std::move(result_ds);
Expand Down
87 changes: 84 additions & 3 deletions tests/test_simq.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
#include <unistd.h>

#include <algorithm>
#include <catch2/catch_approx.hpp>
#include <catch2/catch_test_macros.hpp>
#include <catch2/generators/catch_generators.hpp>
#include <cmath>
Expand Down Expand Up @@ -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;
};
Expand All @@ -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]);
Expand All @@ -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;
}

Expand All @@ -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<uint64_t>(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<float, SIMQ_DIM> vector{};
Expand Down Expand Up @@ -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<EvenLabelFilter>();
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<int64_t>(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<RejectAllFilter>());
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<float>::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]") {
Expand Down