From ee5f2318e3ed5e96adf9c5ccd67863e8deec3593 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Wed, 9 Sep 2026 14:17:54 +0000 Subject: [PATCH 01/12] Add DiskANN navigation using shared Faiss RaBitQ storage Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- include/knowhere/comp/index_param.h | 1 + include/knowhere/index/index_table.h | 1 + python/knowhere/knowhere.i | 4 +- src/index/diskann/diskann.cc | 241 +++++++- src/index/diskann/diskann_config.h | 63 +++ src/index/diskann/rabitq_store.cc | 353 ++++++++++++ src/index/diskann/rabitq_store.h | 67 +++ tests/ut/test_diskann.cc | 514 ++++++++++++++++++ .../DiskANN/include/diskann/aux_utils.h | 5 + .../include/diskann/percentile_stats.h | 6 + .../DiskANN/include/diskann/pq_flash_index.h | 26 +- thirdparty/DiskANN/src/aux_utils.cpp | 27 +- thirdparty/DiskANN/src/pq_flash_index.cpp | 188 +++++-- 13 files changed, 1421 insertions(+), 75 deletions(-) create mode 100644 src/index/diskann/rabitq_store.cc create mode 100644 src/index/diskann/rabitq_store.h diff --git a/include/knowhere/comp/index_param.h b/include/knowhere/comp/index_param.h index 3428a2801..5ad93ceab 100644 --- a/include/knowhere/comp/index_param.h +++ b/include/knowhere/comp/index_param.h @@ -61,6 +61,7 @@ constexpr const char* INDEX_HNSW_PRQ = "HNSW_PRQ"; constexpr const char* INDEX_HNSW_RABITQ = "HNSW_RABITQ"; constexpr const char* INDEX_DISKANN = "DISKANN"; +constexpr const char* INDEX_DISKANN_RABITQ = "DISKANN_RABITQ"; constexpr const char* INDEX_AISAQ = "AISAQ"; constexpr const char* INDEX_MINHASH_LSH = "MINHASH_LSH"; diff --git a/include/knowhere/index/index_table.h b/include/knowhere/index/index_table.h index 4d2d97c69..7d3ccedf3 100644 --- a/include/knowhere/index/index_table.h +++ b/include/knowhere/index/index_table.h @@ -113,6 +113,7 @@ static std::set> legal_knowhere_index = { {IndexEnum::INDEX_DISKANN, VecType::VECTOR_FLOAT}, {IndexEnum::INDEX_DISKANN, VecType::VECTOR_FLOAT16}, {IndexEnum::INDEX_DISKANN, VecType::VECTOR_BFLOAT16}, + {IndexEnum::INDEX_DISKANN_RABITQ, VecType::VECTOR_FLOAT}, // aisaq {IndexEnum::INDEX_AISAQ, VecType::VECTOR_FLOAT}, diff --git a/python/knowhere/knowhere.i b/python/knowhere/knowhere.i index 0d2ff6252..8d14b54aa 100644 --- a/python/knowhere/knowhere.i +++ b/python/knowhere/knowhere.i @@ -173,7 +173,9 @@ class IndexWrap { public: IndexWrap(const std::string& name, const int32_t& version) { GILReleaser rel; - if (name == std::string(knowhere::IndexEnum::INDEX_DISKANN) || name == std::string(knowhere::IndexEnum::INDEX_CARDINAL_TIERED)) { + if (name == std::string(knowhere::IndexEnum::INDEX_DISKANN) || + name == std::string(knowhere::IndexEnum::INDEX_DISKANN_RABITQ) || + name == std::string(knowhere::IndexEnum::INDEX_CARDINAL_TIERED)) { std::shared_ptr file_manager = std::make_shared(); auto diskann_pack = knowhere::Pack(file_manager); idx = IndexFactory::Instance().Create(name, version, diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 936c031b7..c7cf01dcf 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -23,6 +23,7 @@ #include "filemanager/FileManager.h" #include "fmt/core.h" #include "index/diskann/diskann_config.h" +#include "index/diskann/rabitq_store.h" #include "knowhere/comp/index_param.h" #include "knowhere/context.h" #include "knowhere/dataset.h" @@ -200,7 +201,11 @@ class DiskANNIndexNode : public IndexNode { LOG_KNOWHERE_ERROR_ << "Diskann not loaded."; return 0; } - return pq_flash_index_->cal_size(); + auto size = pq_flash_index_->cal_size(); + if (rabitq_store_ != nullptr) { + size += rabitq_store_->MemorySize(); + } + return size; } int64_t @@ -225,6 +230,11 @@ class DiskANNIndexNode : public IndexNode { expected GetVectorByStorageIds(const DataSetPtr dataset, milvus::OpContext* op_context) const override; + virtual bool + IsRaBitQ() const { + return false; + } + private: class iterator : public IndexIterator { public: @@ -283,6 +293,7 @@ class DiskANNIndexNode : public IndexNode { std::atomic_bool is_prepared_; std::shared_ptr file_manager_; std::unique_ptr> pq_flash_index_; + std::unique_ptr rabitq_store_; std::atomic_int64_t dim_; std::atomic_int64_t count_; std::shared_ptr search_pool_; @@ -294,6 +305,28 @@ namespace knowhere { namespace { static constexpr float kCacheExpansionRate = 1.2; +bool +HasFilteredBits(const BitsetView& bitset) { + if (bitset.empty()) { + return false; + } + + const auto* data = bitset.data(); + const auto full_bytes = bitset.size() / 8; + for (size_t i = 0; i < full_bytes; ++i) { + if (data[i] != 0) { + return true; + } + } + + const auto remaining_bits = bitset.size() % 8; + if (remaining_bits == 0) { + return false; + } + const auto valid_bits_mask = static_cast((1u << remaining_bits) - 1u); + return (data[full_bytes] & valid_bits_mask) != 0; +} + Status ReadEmbListOffsetFromFile(const std::string& file_path, std::vector& offsets) { std::ifstream in_file(file_path, std::ios::binary); @@ -408,7 +441,7 @@ AnyIndexFileExist(const std::string& index_prefix) { return false; }; return file_exist(GetNecessaryFilenames(index_prefix, diskann::INNER_PRODUCT, true, true)) || - file_exist(GetOptionalFilenames(index_prefix)); + file_exist(GetOptionalFilenames(index_prefix)) || file_exists(RaBitQStore::SidecarFilename(index_prefix)); } inline bool @@ -493,6 +526,8 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr(num_nodes_to_cache), build_conf.shuffle_build.value()}; + diskann_internal_build_config.keep_preprocessed_base = + IsRaBitQ() && diskann_metric == diskann::Metric::INNER_PRODUCT; RETURN_IF_ERROR(TryDiskANNCall([&]() { int res = diskann::build_disk_index(diskann_internal_build_config); if (res != 0) @@ -500,6 +535,34 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr(cfg.get()); + if (rabitq_conf == nullptr) { + LOG_KNOWHERE_ERROR_ << "DISKANN_RABITQ received an unexpected config type"; + return Status::invalid_args; + } + try { + const auto sidecar_path = RaBitQStore::SidecarFilename(index_prefix_); + const auto sidecar_source = diskann_metric == diskann::Metric::INNER_PRODUCT + ? index_prefix_ + "_prepped_base.bin" + : data_path; + LOG_KNOWHERE_INFO_ << "Building DiskANN RaBitQ sidecar: " << sidecar_path; + RaBitQStore::BuildFromFloatBin(sidecar_source, sidecar_path, + static_cast(rabitq_conf->rbq_bits.value())); + if (diskann_metric == diskann::Metric::INNER_PRODUCT) { + std::error_code error; + std::filesystem::remove(sidecar_source, error); + } + } catch (const std::exception& e) { + if (diskann_metric == diskann::Metric::INNER_PRODUCT) { + std::error_code error; + std::filesystem::remove(index_prefix_ + "_prepped_base.bin", error); + } + LOG_KNOWHERE_ERROR_ << "Failed to build DiskANN RaBitQ sidecar: " << e.what(); + return Status::diskann_inner_error; + } + } + // Add file to the file manager for (auto& filename : GetNecessaryFilenames(index_prefix_, need_norm, true, true)) { if (!AddFile(filename)) { @@ -513,6 +576,13 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::BuildEmbListIfNeed(const DataSetPtr dataset, std::sh // If not emb_list metric type, use the default build method return Build(dataset, std::move(cfg), use_knowhere_build_pool); } + if (IsRaBitQ()) { + LOG_KNOWHERE_ERROR_ << "DISKANN_RABITQ does not support embedding-list mode"; + return Status::not_implemented; + } // DiskANN only supports TokenANN strategy auto strategy_type = config.emb_list_strategy.value_or(meta::EMB_LIST_STRATEGY_TOKENANN); @@ -613,7 +687,8 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr // Load file from file manager. for (auto& filename : GetNecessaryFilenames( - index_prefix_, need_norm, prep_conf.search_cache_budget_gb.value() > 0 && !prep_conf.use_bfs_cache.value(), + index_prefix_, need_norm, + prep_conf.search_cache_budget_gb.value() > 0 && !prep_conf.use_bfs_cache.value() && !IsRaBitQ(), prep_conf.warm_up.value())) { if (!LoadFile(filename)) { return Status::disk_file_error; @@ -629,6 +704,13 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr return Status::disk_file_error; } } + if (IsRaBitQ()) { + const auto sidecar_path = RaBitQStore::SidecarFilename(index_prefix_); + if (!LoadFile(sidecar_path)) { + LOG_KNOWHERE_ERROR_ << "Failed to load DiskANN RaBitQ sidecar " << sidecar_path; + return Status::disk_file_error; + } + } // set thread pool search_pool_ = ThreadPool::GetGlobalSearchThreadPool(); @@ -640,7 +722,7 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr pq_flash_index_ = std::make_unique>(reader, diskann_metric); auto disk_ann_call = [&]() { - int res = pq_flash_index_->load(search_pool_->size(), index_prefix_.c_str()); + int res = pq_flash_index_->load(search_pool_->size(), index_prefix_.c_str(), !IsRaBitQ()); if (res != 0) { throw diskann::ANNException("pq_flash_index_->load returned non-zero value: " + std::to_string(res), -1); } @@ -658,6 +740,22 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr dim_.store(pq_flash_index_->get_data_dim()); } + if (IsRaBitQ()) { + try { + rabitq_store_ = std::make_unique(RaBitQStore::SidecarFilename(index_prefix_)); + if (rabitq_store_->Count() != static_cast(pq_flash_index_->get_num_points()) || + rabitq_store_->Dimension() != static_cast(pq_flash_index_->get_data_dim())) { + LOG_KNOWHERE_ERROR_ << "DiskANN graph and RaBitQ sidecar metadata do not match"; + rabitq_store_.reset(); + return Status::invalid_index_error; + } + } catch (const std::exception& e) { + LOG_KNOWHERE_ERROR_ << "Failed to initialize DiskANN RaBitQ sidecar: " << e.what(); + rabitq_store_.reset(); + return Status::invalid_index_error; + } + } + std::string warmup_query_file = diskann::get_sample_data_filename(index_prefix_); // load cache auto cached_nodes_file = diskann::get_cached_nodes_file(index_prefix_); @@ -676,11 +774,11 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr disk_pq_nchunks = prep_conf.disk_pq_dims.value(); } num_nodes_to_cache = GetCachedNodeNum(prep_conf.search_cache_budget_gb.value(), disk_pq_nchunks, - sizeof(_u8), prep_conf.max_degree.value()); + sizeof(_u8), pq_flash_index_->get_max_degree()); } else { num_nodes_to_cache = GetCachedNodeNum(prep_conf.search_cache_budget_gb.value(), pq_flash_index_->get_data_dim(), - sizeof(DataType), prep_conf.max_degree.value()); + sizeof(DataType), pq_flash_index_->get_max_degree()); } if (num_nodes_to_cache > pq_flash_index_->get_num_points() / 3) { LOG_KNOWHERE_ERROR_ << "Failed to generate cache, num_nodes_to_cache(" << num_nodes_to_cache @@ -689,10 +787,16 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr } if (num_nodes_to_cache > 0) { LOG_KNOWHERE_INFO_ << "Caching " << num_nodes_to_cache << " sample nodes around medoid(s)."; - if (prep_conf.use_bfs_cache.value()) { + if (prep_conf.use_bfs_cache.value() || IsRaBitQ()) { + if (IsRaBitQ() && !prep_conf.use_bfs_cache.value()) { + LOG_KNOWHERE_INFO_ << "DISKANN_RABITQ uses BFS cache generation because navigation PQ is not " + "resident"; + } LOG_KNOWHERE_INFO_ << "Use bfs to generate cache list"; - if (TryDiskANNCall([&]() { pq_flash_index_->cache_bfs_levels(num_nodes_to_cache, node_list); }) != - Status::success) { + if (TryDiskANNCall([&]() { + pq_flash_index_->cache_bfs_levels(num_nodes_to_cache, node_list, + prep_conf.bfs_cache_seed.value()); + }) != Status::success) { LOG_KNOWHERE_ERROR_ << "Failed to generate bfs cache for DiskANN."; return Status::diskann_inner_error; } @@ -711,10 +815,20 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr } if (node_list.size() > 0) { + const auto size_before_cache = pq_flash_index_->cal_size(); if (TryDiskANNCall([&]() { pq_flash_index_->load_cache_list(node_list); }) != Status::success) { LOG_KNOWHERE_ERROR_ << "Failed to load cache for DiskANN."; return Status::diskann_inner_error; } + const auto size_after_cache = pq_flash_index_->cal_size(); + uint64_t cache_list_fingerprint = 1469598103934665603ULL; + for (const auto node_id : node_list) { + cache_list_fingerprint ^= node_id; + cache_list_fingerprint *= 1099511628211ULL; + } + LOG_KNOWHERE_INFO_ << "DiskANN node cache: nodes=" << node_list.size() + << ", bytes=" << size_after_cache - size_before_cache + << ", list_fingerprint=" << cache_list_fingerprint; } // warmup @@ -772,6 +886,10 @@ DiskANNIndexNode::DeserializeEmbListIfNeed(const BinarySet& binset, st // If not emb_list metric type, use the default deserialize method return Deserialize(binset, std::move(cfg)); } + if (IsRaBitQ()) { + LOG_KNOWHERE_ERROR_ << "DISKANN_RABITQ does not support embedding-list mode"; + return Status::not_implemented; + } LOG_KNOWHERE_INFO_ << "Deserialize emb_list index and read emb_list offset from file."; @@ -833,6 +951,10 @@ template expected> DiskANNIndexNode::AnnIterator(const DataSetPtr dataset, std::unique_ptr cfg, const BitsetView& bitset, bool use_knowhere_search_pool, milvus::OpContext* op_context) const { + if (IsRaBitQ()) { + return expected>::Err( + Status::not_implemented, "DISKANN_RABITQ does not support iterator search"); + } if (!is_prepared_.load() || !pq_flash_index_) { LOG_KNOWHERE_ERROR_ << "Failed to load diskann."; return expected>::Err(Status::empty_index, "DiskANN not loaded"); @@ -895,6 +1017,22 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptrGetDim(); auto xq = static_cast(dataset->GetTensor()); + bool rbq_probabilistic_refinement = true; + uint8_t rbq_query_bits = 4; + if (IsRaBitQ()) { + if (HasFilteredBits(bitset_)) { + return expected::Err(Status::not_implemented, + "DISKANN_RABITQ does not support bitset search"); + } + const auto* rabitq_conf = dynamic_cast(cfg.get()); + if (rabitq_conf == nullptr || rabitq_store_ == nullptr) { + return expected::Err(Status::invalid_args, + "DISKANN_RABITQ config or sidecar is not initialized"); + } + rbq_probabilistic_refinement = rabitq_conf->rbq_refine_mode.value() == "probabilistic"; + rbq_query_bits = static_cast(rabitq_conf->rbq_bits_query.value_or(4)); + } + feder::diskann::FederResultUniq feder_result; if (search_conf.trace_visit.value()) { if (nq != 1) { @@ -907,16 +1045,20 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptr(k * nq); auto p_dist = std::make_unique(k * nq); + std::vector query_stats(nq); std::vector> futures; futures.reserve(nq); for (int64_t row = 0; row < nq; ++row) { futures.emplace_back(search_pool_->push([&, index = row, p_id_ptr = p_id.get(), p_dist_ptr = p_dist.get()]() { knowhere::checkCancellation(op_context); - diskann::QueryStats stats; + auto& stats = query_stats[index]; + auto approx_distance_computer = IsRaBitQ() ? rabitq_store_->CreateDistanceComputer( + rbq_probabilistic_refinement, rbq_query_bits) + : nullptr; pq_flash_index_->cached_beam_search(xq + (index * dim), k, lsearch, p_id_ptr + (index * k), p_dist_ptr + (index * k), beamwidth, false, &stats, feder_result, - bitset_, filter_ratio); + bitset_, filter_ratio, approx_distance_computer.get()); #ifdef NOT_COMPILE_FOR_SWIG knowhere_diskann_search_hops.Observe(stats.n_hops); #endif @@ -927,6 +1069,40 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptr::Err(Status::diskann_inner_error, "some search failed"); } + if (IsRaBitQ()) { + uint64_t estimates = 0; + uint64_t refinements = 0; + uint64_t pruned = 0; + for (const auto& stats : query_stats) { + estimates += stats.n_approx_estimates; + refinements += stats.n_approx_refinements; + pruned += stats.n_approx_pruned; + } + const double prune_ratio = estimates == 0 ? 0.0 : static_cast(pruned) / estimates; + LOG_KNOWHERE_DEBUG_ << "DiskANN RaBitQ refinement stats: queries=" << nq << ", estimates=" << estimates + << ", full_distances=" << refinements << ", pruned=" << pruned + << ", prune_ratio=" << prune_ratio; + } + + { + double total_us = 0.0; + double cpu_us = 0.0; + double io_us = 0.0; + uint64_t n_ios = 0; + uint64_t n_cache_hits = 0; + for (const auto& stats : query_stats) { + total_us += stats.total_us; + cpu_us += stats.cpu_us; + io_us += stats.io_us; + n_ios += stats.n_ios; + n_cache_hits += stats.n_cache_hits; + } + const double n = static_cast(nq); + LOG_KNOWHERE_DEBUG_ << "DiskANN search stats: queries=" << nq << ", avg_total_us=" << total_us / n + << ", avg_cpu_us=" << cpu_us / n << ", avg_io_us=" << io_us / n + << ", avg_n_ios=" << n_ios / n << ", avg_cache_hits=" << n_cache_hits / n; + } + auto res = GenResultDataSet(nq, k, std::move(p_id), std::move(p_dist)); MapSearchResultIdsToOutIds(res); @@ -1076,6 +1252,47 @@ DiskANNIndexNode::GetCachedNodeNum(const float cache_dram_budget, cons return num_nodes_to_cache; } +template +class DiskANNRaBitQIndexNode : public DiskANNIndexNode { + public: + using DiskANNIndexNode::DiskANNIndexNode; + + static std::unique_ptr + StaticCreateConfig() { + return std::make_unique(); + } + + std::unique_ptr + CreateConfig() const override { + return StaticCreateConfig(); + } + + static Status + StaticConfigCheck(const Config& cfg, PARAM_TYPE param_type, std::string& msg) { + const auto status = DiskANNIndexNode::StaticConfigCheck(cfg, param_type, msg); + if (status != Status::success) { + return status; + } + const auto& base_cfg = static_cast(cfg); + if (base_cfg.emb_list_strategy.has_value() || base_cfg.emb_list_offset_file_path.has_value()) { + msg = "DISKANN_RABITQ does not support embedding-list mode"; + return Status::not_implemented; + } + return Status::success; + } + + std::string + Type() const override { + return knowhere::IndexEnum::INDEX_DISKANN_RABITQ; + } + + protected: + bool + IsRaBitQ() const override { + return true; + } +}; + #ifdef KNOWHERE_WITH_CARDINAL KNOWHERE_SIMPLE_REGISTER_DENSE_FLOAT_ALL_GLOBAL(DISKANN_DEPRECATED, DiskANNIndexNode, knowhere::feature::DISK | knowhere::feature::EMB_LIST) @@ -1083,4 +1300,6 @@ KNOWHERE_SIMPLE_REGISTER_DENSE_FLOAT_ALL_GLOBAL(DISKANN_DEPRECATED, DiskANNIndex KNOWHERE_SIMPLE_REGISTER_DENSE_FLOAT_ALL_GLOBAL(DISKANN, DiskANNIndexNode, knowhere::feature::DISK | knowhere::feature::EMB_LIST) #endif +KNOWHERE_SIMPLE_REGISTER_GLOBAL(DISKANN_RABITQ, DiskANNRaBitQIndexNode, fp32, + knowhere::feature::DISK | knowhere::feature::FLOAT32) } // namespace knowhere diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index 34a9079f8..f791eeac4 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -71,6 +71,9 @@ class DiskANNConfig : public BaseConfig { // cached the nodes on the search paths; 2. do bfs from the entry point and cache them. The first method is suitable // for TopK query heavy circumstances and the second one performed better in range search. CFG_BOOL use_bfs_cache; + // Optional deterministic seed for BFS cache selection. A negative value + // preserves the existing random selection behavior. + CFG_INT bfs_cache_seed; // The beamwidth to be used for search. This is the maximum number of IO requests each query will issue per // iteration of search code. Larger beamwidth will result in fewer IO round-trips per query but might result in // slightly higher total number of IO requests to SSD per query. For the highest query throughput with a fixed SSD @@ -141,6 +144,11 @@ class DiskANNConfig : public BaseConfig { .description("should bfs strategy to cache nodes.") .set_default(false) .for_deserialize(); + KNOWHERE_CONFIG_DECLARE_FIELD(bfs_cache_seed) + .description("seed for deterministic bfs cache selection; -1 uses a random seed.") + .set_default(-1) + .set_range(-1, std::numeric_limits::max()) + .for_deserialize(); KNOWHERE_CONFIG_DECLARE_FIELD(beamwidth) .description("the maximum number of IO requests each query will issue per iteration of search code.") .set_default(diskann::defaults::DEFAULT_DISKANN_BEAMWIDTH) @@ -195,5 +203,60 @@ class DiskANNConfig : public BaseConfig { return Status::success; } }; + +class DiskANNRaBitQConfig : public DiskANNConfig { + public: + CFG_INT rbq_bits; + CFG_INT rbq_bits_query; + CFG_STRING rbq_refine_mode; + + KNOWHERE_DECLARE_CONFIG(DiskANNRaBitQConfig) { + KNOWHERE_CONFIG_DECLARE_FIELD(rbq_bits) + .description("number of RaBitQ bits per database vector dimension") + .set_default(1) + .set_range(1, 9) + .for_train() + .for_static(); + KNOWHERE_CONFIG_DECLARE_FIELD(rbq_bits_query) + .description("query bits for the RaBitQ coarse estimator; 0 uses FP32") + .set_default(4) + .set_range(0, 8) + .for_search(); + KNOWHERE_CONFIG_DECLARE_FIELD(rbq_refine_mode) + .description("RaBitQ multi-bit refinement mode: probabilistic enables error-window pruning; full always " + "computes the complete RaBitQ distance") + .set_default("probabilistic") + .for_search() + .for_range_search() + .for_iterator(); + } + + Status + CheckAndAdjust(PARAM_TYPE param_type, std::string* err_msg) override { + const auto base_status = DiskANNConfig::CheckAndAdjust(param_type, err_msg); + if (base_status != Status::success) { + return base_status; + } + const auto metric = metric_type.value_or(knowhere::metric::L2); + if (metric != knowhere::metric::L2 && metric != knowhere::metric::IP) { + return HandleError(err_msg, "DISKANN_RABITQ supports L2 and IP", Status::invalid_metric_type); + } + const auto database_bits = rbq_bits.value_or(1); + if (database_bits < 1 || database_bits > 9) { + return HandleError(err_msg, "DISKANN_RABITQ supports rbq_bits in [1, 9]", Status::invalid_args); + } + const auto refine_mode = rbq_refine_mode.value_or("probabilistic"); + if (refine_mode != "probabilistic" && refine_mode != "full") { + return HandleError(err_msg, "rbq_refine_mode must be probabilistic or full", Status::invalid_args); + } + if (disk_pq_dims.value_or(0) != 0) { + return HandleError(err_msg, "DISKANN_RABITQ requires disk_pq_dims=0", Status::invalid_args); + } + if (warm_up.value_or(false)) { + return HandleError(err_msg, "DISKANN_RABITQ does not support warm_up", Status::invalid_args); + } + return Status::success; + } +}; } // namespace knowhere #endif /* DISKANN_CONFIG_H */ diff --git a/src/index/diskann/rabitq_store.cc b/src/index/diskann/rabitq_store.cc new file mode 100644 index 000000000..7417baf3c --- /dev/null +++ b/src/index/diskann/rabitq_store.cc @@ -0,0 +1,353 @@ +// Copyright (C) 2019-2026 Zilliz. All rights reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. + +#include "index/diskann/rabitq_store.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +#include "diskann/utils.h" + +namespace knowhere { +namespace { + +constexpr size_t kBuildBlockBytes = 32UL * 1024 * 1024; + +size_t +BlockRows(size_t dim) { + if (dim == 0) { + throw std::invalid_argument("RaBitQ sidecar dimension must be positive"); + } + return std::max(1, kBuildBlockBytes / (dim * sizeof(float))); +} + +template +void +ForEachFloatBinBlock(const std::string& data_path, size_t rows, size_t dim, Fn&& fn) { + std::ifstream input(data_path, std::ios::binary); + if (!input) { + throw std::runtime_error("failed to open RaBitQ source data: " + data_path); + } + + input.seekg(2 * sizeof(uint32_t), std::ios::beg); + const size_t block_rows = BlockRows(dim); + std::vector block(block_rows * dim); + size_t row_offset = 0; + while (row_offset < rows) { + const size_t current_rows = std::min(block_rows, rows - row_offset); + const size_t current_values = current_rows * dim; + input.read(reinterpret_cast(block.data()), current_values * sizeof(float)); + if (!input) { + throw std::runtime_error("short read while building RaBitQ sidecar from: " + data_path); + } + fn(block.data(), current_rows); + row_offset += current_rows; + } +} + +// Single-query rotation application. +// +// RandomRotationMatrix::apply_noalloc() routes through OpenBLAS sgemm_, which +// for a skinny n=1 GEMV spawns an internal worker pool whose launch/packing +// overhead dominates the actual multiply (observed as a serial-p50 regression +// when the search worker thread already saturates the CPU). Since DiskANN +// searches one query at a time, apply the rotation with a plain single-threaded +// GEMV instead. +// +// The stored matrix A is d_in x d_out in column-major order (d_in == d_out == d +// here), and apply_noalloc computes xt = A^T * x, i.e. +// xt[j] = sum_k A[k + j * d] * x[k]. +void +apply_rotation_single_query(const faiss::RandomRotationMatrix* rotation, const float* x, float* xt) { + const int d = rotation->d_in; + const float* a = rotation->A.data(); + for (int j = 0; j < d; ++j) { + const float* col = a + j * d; + float acc = rotation->have_bias ? rotation->b[j] : 0.0f; +#pragma omp simd reduction(+ : acc) + for (int k = 0; k < d; ++k) { + acc += col[k] * x[k]; + } + xt[j] = acc; + } +} + +class RaBitQApproxDistanceComputer final : public diskann::ApproxDistanceComputer { + public: + RaBitQApproxDistanceComputer(const faiss::RandomRotationMatrix* rotation, const faiss::IndexRaBitQ* rabitq, + bool probabilistic_refinement, uint8_t query_bits) + : rotation_(rotation), + rabitq_(rabitq), + probabilistic_refinement_(probabilistic_refinement), + distance_computer_(rabitq->get_quantized_distance_computer(query_bits, false)), + rabitq_distance_computer_(dynamic_cast(distance_computer_.get())) { + if (rabitq_distance_computer_ == nullptr) { + throw std::runtime_error("RaBitQ sidecar returned an incompatible distance computer"); + } + } + + void + set_query(const float* query) override { + const int d = rotation_->d_in; + transformed_query_ = std::make_unique(d); + apply_rotation_single_query(rotation_, query, transformed_query_.get()); + distance_computer_->set_query(transformed_query_.get()); + } + + void + compute_distances(const unsigned* ids, _u64 n_ids, float* distances, float threshold, bool threshold_valid, + diskann::QueryStats* stats) override { + const bool can_prune = probabilistic_refinement_ && threshold_valid && rabitq_->rabitq.nb_bits > 1; + if (!can_prune) { + _u64 i = 0; + for (; i + 4 <= n_ids; i += 4) { + distance_computer_->distances_batch_4(ids[i], ids[i + 1], ids[i + 2], ids[i + 3], distances[i], + distances[i + 1], distances[i + 2], distances[i + 3]); + } + for (; i < n_ids; ++i) { + distances[i] = (*distance_computer_)(ids[i]); + } + if (stats != nullptr && rabitq_->rabitq.nb_bits > 1) { + stats->n_approx_refinements += n_ids; + } + return; + } + + std::array refine_ids{}; + std::array<_u64, 4> refine_positions{}; + size_t pending_refinements = 0; + const auto flush_refinements = [&]() { + distance_computer_->distances_batch_4( + refine_ids[0], refine_ids[1], refine_ids[2], refine_ids[3], + distances[refine_positions[0]], distances[refine_positions[1]], distances[refine_positions[2]], + distances[refine_positions[3]]); + pending_refinements = 0; + }; + + for (_u64 i = 0; i < n_ids; ++i) { + const uint8_t* code = rabitq_->codes.data() + static_cast(ids[i]) * rabitq_->code_size; + const float estimate = rabitq_distance_computer_->distance_to_code_1bit(code); + const size_t code_body_size = (static_cast(rabitq_->d) + 7) / 8; + const auto* factors = reinterpret_cast( + code + code_body_size); + if (stats != nullptr) { + ++stats->n_approx_estimates; + } + if (!faiss::rabitq_utils::should_refine_candidate(estimate, factors->f_error, + rabitq_distance_computer_->g_error, threshold, false)) { + distances[i] = std::numeric_limits::infinity(); + if (stats != nullptr) { + ++stats->n_approx_pruned; + ++stats->n_cmps_saved; + } + continue; + } + refine_ids[pending_refinements] = ids[i]; + refine_positions[pending_refinements] = i; + ++pending_refinements; + if (stats != nullptr) { + ++stats->n_approx_refinements; + } + if (pending_refinements == 4) { + flush_refinements(); + } + } + for (size_t i = 0; i < pending_refinements; ++i) { + const auto id = refine_ids[i]; + const uint8_t* code = rabitq_->codes.data() + static_cast(id) * rabitq_->code_size; + distances[refine_positions[i]] = rabitq_distance_computer_->distance_to_code_full(code); + } + } + + private: + const faiss::RandomRotationMatrix* rotation_; + const faiss::IndexRaBitQ* rabitq_; + const bool probabilistic_refinement_; + std::unique_ptr distance_computer_; + faiss::RaBitQDistanceComputer* rabitq_distance_computer_; + std::unique_ptr transformed_query_; +}; + +} // namespace + +std::string +RaBitQStore::SidecarFilename(const std::string& index_prefix) { + return index_prefix + "_rabitq.index"; +} + +void +RaBitQStore::BuildFromFloatBin(const std::string& data_path, const std::string& sidecar_path, uint8_t rbq_bits) { + if (rbq_bits < 1 || rbq_bits > 9) { + throw std::invalid_argument("RaBitQ database bits must be in [1, 9]"); + } + + size_t rows = 0; + size_t dim = 0; + diskann::get_bin_metadata(data_path, rows, dim); + if (rows == 0 || dim == 0 || dim > static_cast(std::numeric_limits::max())) { + throw std::invalid_argument("invalid RaBitQ source metadata"); + } + + constexpr size_t header_size = 2 * sizeof(uint32_t); + if (dim > (std::numeric_limits::max() - header_size) / sizeof(float) / rows) { + throw std::invalid_argument("RaBitQ source metadata overflows file size"); + } + const auto expected_size = header_size + rows * dim * sizeof(float); + if (std::filesystem::file_size(data_path) != expected_size) { + throw std::runtime_error("RaBitQ source file size does not match float32 metadata"); + } + + auto rotation = std::make_unique(static_cast(dim), static_cast(dim)); + rotation->init(12345); + + std::vector sums(dim, 0.0); + ForEachFloatBinBlock(data_path, rows, dim, [&](const float* block, size_t block_rows) { + for (size_t i = 0; i < block_rows; ++i) { + const float* row = block + i * dim; + for (size_t j = 0; j < dim; ++j) { + sums[j] += row[j]; + } + } + }); + + std::vector mean(dim); + for (size_t j = 0; j < dim; ++j) { + mean[j] = static_cast(sums[j] / static_cast(rows)); + } + std::vector rotated_center(dim); + rotation->apply_noalloc(1, mean.data(), rotated_center.data()); + + auto rabitq = std::make_unique(static_cast(dim), faiss::METRIC_L2, rbq_bits); + rabitq->center = std::move(rotated_center); + rabitq->qb = 4; + rabitq->centered = false; + rabitq->is_trained = true; + auto pretransform = std::make_unique(rotation.get(), rabitq.get()); + pretransform->own_fields = true; + rotation.release(); + rabitq.release(); + + ForEachFloatBinBlock(data_path, rows, dim, [&](const float* block, size_t block_rows) { + // Preserve DiskANN's existing input-block boundaries while sharing + // the same bounded storage population utility as HNSW. + faiss::cppcontrib::knowhere::rabitq_build::add_in_blocks( + *pretransform, static_cast(block_rows), block, + static_cast(BlockRows(dim))); + }); + if (pretransform->ntotal != static_cast(rows)) { + throw std::runtime_error("RaBitQ sidecar point count mismatch after encoding"); + } + + const std::string temporary_path = sidecar_path + ".tmp"; + std::error_code error; + std::filesystem::remove(temporary_path, error); + try { + faiss::cppcontrib::knowhere::write_index(pretransform.get(), temporary_path.c_str()); + std::filesystem::rename(temporary_path, sidecar_path); + } catch (...) { + std::filesystem::remove(temporary_path, error); + throw; + } +} + +RaBitQStore::RaBitQStore(const std::string& sidecar_path) + : index_(faiss::cppcontrib::knowhere::read_index(sidecar_path.c_str())) { + Validate(); +} + +RaBitQStore::~RaBitQStore() = default; + +void +RaBitQStore::Validate() { + pretransform_ = dynamic_cast(index_.get()); + if (pretransform_ == nullptr || pretransform_->chain.size() != 1) { + throw std::runtime_error("DiskANN RaBitQ sidecar must be an IndexPreTransform with one transform"); + } + rotation_ = dynamic_cast(pretransform_->chain[0]); + if (rotation_ == nullptr || !rotation_->is_trained || rotation_->d_in <= 0 || + rotation_->d_in != rotation_->d_out) { + throw std::runtime_error("DiskANN RaBitQ sidecar has an invalid random rotation"); + } + const auto rotation_dim = static_cast(rotation_->d_in); + if (rotation_->A.size() != rotation_dim * rotation_dim || + (rotation_->have_bias ? rotation_->b.size() != rotation_dim : !rotation_->b.empty())) { + throw std::runtime_error("DiskANN RaBitQ sidecar rotation storage is inconsistent"); + } + rabitq_ = dynamic_cast(pretransform_->index); + if (rabitq_ == nullptr || !pretransform_->is_trained || !rabitq_->is_trained) { + throw std::runtime_error("DiskANN RaBitQ sidecar has an invalid RaBitQ leaf"); + } + if (pretransform_->metric_type != faiss::METRIC_L2 || rabitq_->metric_type != faiss::METRIC_L2 || + rabitq_->rabitq.metric_type != faiss::METRIC_L2 || + pretransform_->d != rotation_->d_in || rabitq_->d != rotation_->d_out || + pretransform_->ntotal != rabitq_->ntotal) { + throw std::runtime_error("DiskANN RaBitQ sidecar metadata is inconsistent"); + } + if (rabitq_->ntotal < 0 || rabitq_->rabitq.nb_bits < 1 || rabitq_->rabitq.nb_bits > 9 || + rabitq_->qb > 8 || rabitq_->centered) { + throw std::runtime_error("DiskANN RaBitQ sidecar quantizer metadata is inconsistent"); + } + const auto expected_code_size = rabitq_->rabitq.compute_code_size( + static_cast(rabitq_->d), rabitq_->rabitq.nb_bits); + const auto point_count = static_cast(rabitq_->ntotal); + if (rabitq_->code_size != expected_code_size || rabitq_->rabitq.code_size != expected_code_size || + expected_code_size == 0 || point_count > std::numeric_limits::max() / expected_code_size || + rabitq_->center.size() != static_cast(rabitq_->d) || + rabitq_->codes.size() != point_count * expected_code_size) { + throw std::runtime_error("DiskANN RaBitQ sidecar code storage is inconsistent"); + } +} + +std::unique_ptr +RaBitQStore::CreateDistanceComputer(bool probabilistic_refinement, uint8_t query_bits) const { + if (query_bits > 8) { + throw std::invalid_argument("RaBitQ query bits must be in [0, 8]"); + } + return std::make_unique(rotation_, rabitq_, probabilistic_refinement, query_bits); +} + +int64_t +RaBitQStore::Count() const { + return pretransform_->ntotal; +} + +int64_t +RaBitQStore::Dimension() const { + return pretransform_->d; +} + +uint8_t +RaBitQStore::Bits() const { + return static_cast(rabitq_->rabitq.nb_bits); +} + +size_t +RaBitQStore::CodeSize() const { + return rabitq_->code_size; +} + +size_t +RaBitQStore::MemorySize() const { + return rabitq_->codes.size() * sizeof(uint8_t) + rabitq_->center.size() * sizeof(float) + + rotation_->A.size() * sizeof(float) + rotation_->b.size() * sizeof(float); +} + +} // namespace knowhere diff --git a/src/index/diskann/rabitq_store.h b/src/index/diskann/rabitq_store.h new file mode 100644 index 000000000..6bc223fca --- /dev/null +++ b/src/index/diskann/rabitq_store.h @@ -0,0 +1,67 @@ +// Copyright (C) 2019-2026 Zilliz. All rights reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. + +#pragma once + +#include +#include +#include +#include + +#include "diskann/pq_flash_index.h" + +namespace faiss { +struct Index; +struct IndexPreTransform; +struct IndexRaBitQ; +struct RandomRotationMatrix; +} // namespace faiss + +namespace knowhere { + +class RaBitQStore { + public: + static std::string + SidecarFilename(const std::string& index_prefix); + + static void + BuildFromFloatBin(const std::string& data_path, const std::string& sidecar_path, uint8_t rbq_bits); + + explicit RaBitQStore(const std::string& sidecar_path); + ~RaBitQStore(); + + RaBitQStore(const RaBitQStore&) = delete; + RaBitQStore& + operator=(const RaBitQStore&) = delete; + + std::unique_ptr + CreateDistanceComputer(bool probabilistic_refinement, uint8_t query_bits = 4) const; + + int64_t + Count() const; + + int64_t + Dimension() const; + + uint8_t + Bits() const; + + size_t + CodeSize() const; + + size_t + MemorySize() const; + + private: + void + Validate(); + + std::unique_ptr index_; + const faiss::IndexPreTransform* pretransform_ = nullptr; + const faiss::RandomRotationMatrix* rotation_ = nullptr; + const faiss::IndexRaBitQ* rabitq_ = nullptr; +}; + +} // namespace knowhere diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index 2b6f9b78e..97fe5865f 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -11,19 +11,35 @@ #include +#include +#include #include +#include +#include +#include #include #include +#include "../DiskANN/include/diskann/aux_utils.h" #include "../DiskANN/include/diskann/defaults.h" +#include "../DiskANN/include/diskann/linux_aligned_file_reader.h" +#include "../DiskANN/include/diskann/pq_flash_index.h" #include "catch2/catch_approx.hpp" #include "catch2/catch_test_macros.hpp" #include "catch2/generators/catch_generators.hpp" #include "diskann/diskann_gpu.h" #include "diskann/utils.h" +#include "faiss/IndexPreTransform.h" +#include "faiss/IndexRaBitQ.h" +#include "faiss/VectorTransform.h" +#include "faiss/cppcontrib/knowhere/index_io.h" +#include "faiss/impl/RaBitQUtils.h" +#include "faiss/index_io.h" +#include "faiss/utils/rabitq_simd.h" #include "filemanager/FileManager.h" #include "filemanager/impl/LocalFileManager.h" #include "index/diskann/diskann_config.h" +#include "index/diskann/rabitq_store.h" #include "knowhere/comp/brute_force.h" #include "knowhere/comp/knowhere_check.h" #include "knowhere/expected.h" @@ -80,6 +96,8 @@ constexpr float kL2RangeAp = 0.9; constexpr float kIpRangeAp = 0.9; constexpr float kCosineRangeAp = 0.9; } // namespace + + TEST_CASE("Valid diskann build params test", "[diskann]") { int rows_num = 1000000; auto version = GenTestVersionList(); @@ -127,6 +145,17 @@ TEST_CASE("Valid diskann build params test", "[diskann]") { } } +TEST_CASE("DiskANN navigation PQ can exceed 512 chunks", "[diskann]") { + constexpr size_t rows = 500000; + constexpr size_t dim = 1537; + constexpr size_t matched_code_bytes = 790; + + REQUIRE(diskann::get_num_pq_chunks(static_cast(rows * matched_code_bytes), rows, dim) == + matched_code_bytes); + REQUIRE(diskann::get_num_pq_chunks(static_cast(rows * (dim + 1)), rows, dim) == dim); + REQUIRE(diskann::get_num_pq_chunks(0.0, rows, dim) == 1); +} + TEST_CASE("Invalid diskann params test", "[diskann]") { fs::remove_all(kDir); fs::remove(kDir); @@ -474,6 +503,491 @@ TEST_CASE("Test DiskANN CalcDistByIDs with all vectors cached", "[diskann]") { fs::remove_all(kDir); } +TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { + auto version = GenTestVersionList(); + auto make_pack = []() { + std::shared_ptr file_manager = std::make_shared(); + return knowhere::Pack(file_manager); + }; + + REQUIRE(knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, + make_pack()) + .has_value()); + REQUIRE_FALSE(knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, + make_pack()) + .has_value()); + REQUIRE_FALSE(knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, + make_pack()) + .has_value()); + + auto check_train_config = [&](knowhere::Json json, knowhere::Status expected) { + auto cfg = knowhere::IndexStaticFaced::CreateConfig( + knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); + std::string msg; + REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::TRAIN, &msg) == expected); + }; + + knowhere::Json valid = {{"dim", kDim}, + {"metric_type", knowhere::metric::L2}, + {"index_prefix", kL2IndexPrefix}, + {"data_path", kRawDataPath}, + {"disk_pq_dims", 0}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"rbq_bits", 1}}; + check_train_config(valid, knowhere::Status::success); + valid["rbq_bits"] = 2; + check_train_config(valid, knowhere::Status::success); + valid["rbq_bits"] = 4; + check_train_config(valid, knowhere::Status::success); + valid["rbq_bits"] = 8; + check_train_config(valid, knowhere::Status::success); + valid["rbq_bits"] = 9; + check_train_config(valid, knowhere::Status::success); + valid["rbq_bits"] = 1; + + auto invalid = valid; + invalid["metric_type"] = knowhere::metric::IP; + check_train_config(invalid, knowhere::Status::success); + invalid["metric_type"] = knowhere::metric::COSINE; + check_train_config(invalid, knowhere::Status::invalid_metric_type); + invalid = valid; + invalid["rbq_bits"] = 10; + check_train_config(invalid, knowhere::Status::out_of_range_in_json); + invalid = valid; + invalid["disk_pq_dims"] = 16; + check_train_config(invalid, knowhere::Status::invalid_args); + invalid = valid; + invalid["search_cache_budget_gb"] = 0.01; + check_train_config(invalid, knowhere::Status::success); + + auto check_search_mode = [&](const std::string& mode, knowhere::Status expected) { + auto cfg = knowhere::IndexStaticFaced::CreateConfig( + knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); + knowhere::Json json = {{"dim", kDim}, + {"metric_type", knowhere::metric::L2}, + {"k", kK}, + {"search_list_size", 128}, + {"beamwidth", 8}, + {"rbq_refine_mode", mode}}; + std::string msg; + REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::SEARCH, &msg) == expected); + }; + check_search_mode("probabilistic", knowhere::Status::success); + check_search_mode("full", knowhere::Status::success); + check_search_mode("invalid", knowhere::Status::invalid_args); + for (const int qb : {-1, 0, 4, 8, 9}) { + auto cfg = knowhere::IndexStaticFaced::CreateConfig( + knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); + knowhere::Json json = {{"dim", kDim}, {"metric_type", knowhere::metric::L2}, + {"k", kK}, {"search_list_size", 128}, {"rbq_bits_query", qb}}; + std::string msg; + REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::SEARCH, &msg) == + ((qb >= 0 && qb <= 8) ? knowhere::Status::success : knowhere::Status::out_of_range_in_json)); + } +} + +TEST_CASE("Test DISKANN_RABITQ probabilistic refinement", "[diskann][rabitq]") { + const auto refinement_dir = kDir + "/rabitq_refinement"; + const auto data_path = refinement_dir + "/base.fbin"; + const auto sidecar_path = refinement_dir + "/base_rabitq.index"; + fs::remove_all(refinement_dir); + REQUIRE_NOTHROW(fs::create_directories(refinement_dir)); + + constexpr uint32_t rows = 64; + constexpr uint32_t dim = 128; + auto base_ds = GenDataSet(rows, dim, 30); + const auto* base = static_cast(base_ds->GetTensor()); + WriteRawDataToDisk(data_path, base, rows, dim); + REQUIRE_NOTHROW(knowhere::RaBitQStore::BuildFromFloatBin(data_path, sidecar_path, 4)); + + { + std::unique_ptr index(faiss::cppcontrib::knowhere::read_index(sidecar_path.c_str())); + const auto* pretransform = dynamic_cast(index.get()); + REQUIRE(pretransform != nullptr); + const auto* rabitq = dynamic_cast(pretransform->index); + REQUIRE(rabitq != nullptr); + REQUIRE(rabitq->code_size == rabitq->rabitq.compute_code_size(dim, 4)); + } + + knowhere::RaBitQStore store(sidecar_path); + auto distance_computer = store.CreateDistanceComputer(true, 0); + distance_computer->set_query(base); + std::vector ids(rows); + std::iota(ids.begin(), ids.end(), 0); + std::vector distances(rows); + + diskann::QueryStats full_stats; + distance_computer->compute_distances(ids.data(), rows, distances.data(), + std::numeric_limits::max(), false, &full_stats); + REQUIRE(full_stats.n_approx_estimates == 0); + REQUIRE(full_stats.n_approx_refinements == rows); + REQUIRE(full_stats.n_approx_pruned == 0); + REQUIRE(std::all_of(distances.begin(), distances.end(), [](float distance) { return std::isfinite(distance); })); + + diskann::QueryStats pruned_stats; + distance_computer->compute_distances(ids.data(), rows, distances.data(), 0.0f, true, &pruned_stats); + REQUIRE(pruned_stats.n_approx_estimates == rows); + REQUIRE(pruned_stats.n_approx_refinements == 0); + REQUIRE(pruned_stats.n_approx_pruned == rows); + REQUIRE(pruned_stats.n_cmps_saved == rows); + REQUIRE(std::all_of(distances.begin(), distances.end(), [](float distance) { return std::isinf(distance); })); + + auto full_distance_computer = store.CreateDistanceComputer(false, 0); + full_distance_computer->set_query(base); + diskann::QueryStats explicit_full_stats; + full_distance_computer->compute_distances(ids.data(), rows, distances.data(), 0.0f, true, &explicit_full_stats); + REQUIRE(explicit_full_stats.n_approx_estimates == 0); + REQUIRE(explicit_full_stats.n_approx_refinements == rows); + REQUIRE(explicit_full_stats.n_approx_pruned == 0); + REQUIRE(std::all_of(distances.begin(), distances.end(), [](float distance) { return std::isfinite(distance); })); + + fs::remove_all(refinement_dir); +} + +TEST_CASE("DiskANN RaBitQ shares Faiss codes and request-local query bits", "[diskann][rabitq][sidecar]") { + const auto dir = kDir + "/rabitq_shared_codec"; + REQUIRE_NOTHROW(fs::create_directories(dir)); + for (const uint32_t dim : {33U, 128U}) { + auto data = GenDataSet(17, dim, 73); + const auto* x = static_cast(data->GetTensor()); + const auto raw = dir + "/base.fbin"; + const auto path = dir + "/model.index"; + WriteRawDataToDisk(raw, x, 17, dim); + for (const uint8_t bits : {1, 4, 8, 9}) { + CAPTURE(dim, bits); + knowhere::RaBitQStore::BuildFromFloatBin(raw, path, bits); + knowhere::RaBitQStore store(path); + std::unique_ptr model(faiss::cppcontrib::knowhere::read_index(path.c_str())); + const auto* pt = dynamic_cast(model.get()); + REQUIRE(pt != nullptr); + const auto* rbq = dynamic_cast(pt->index); + REQUIRE(rbq != nullptr); + REQUIRE(store.CodeSize() == rbq->rabitq.compute_code_size(dim, bits)); + std::vector rotated(dim); + std::vector blas_rotated(dim); + pt->chain[0]->apply_noalloc(1, x, blas_rotated.data()); + const auto* rotation = dynamic_cast(pt->chain[0]); + REQUIRE(rotation != nullptr); + // Check the adapter's single-query GEMV formula against BLAS. + // Compiler reduction order can differ across translation units; + // near-zero self distances need a norm-scaled roundoff tolerance. + for (uint32_t j = 0; j < dim; ++j) { + float sum = 0; +#pragma omp simd reduction(+ : sum) + for (uint32_t k = 0; k < dim; ++k) { + sum += rotation->A[k + j * dim] * x[k]; + } + rotated[j] = sum; + REQUIRE(rotated[j] == Catch::Approx(blas_rotated[j]).epsilon(1e-5).margin(1e-4)); + } + const unsigned ids[] = {0, 1, 3, 5, 7, 9, 12}; + std::vector initial(7), after(7); + auto stable = store.CreateDistanceComputer(false, 0); + stable->set_query(x); + stable->compute_distances(ids, 7, initial.data(), 0, false, nullptr); + for (const uint8_t qb : {0, 4, 8}) { + CAPTURE(qb); + auto adapter = store.CreateDistanceComputer(false, qb); + std::unique_ptr native( + rbq->get_quantized_distance_computer(qb, false)); + adapter->set_query(x); + native->set_query(rotated.data()); + std::vector actual(7); + std::vector expected(7); + adapter->compute_distances(ids, 7, actual.data(), 0, false, nullptr); + native->distances_batch_4(ids[0], ids[1], ids[2], ids[3], + expected[0], expected[1], expected[2], expected[3]); + for (size_t i = 4; i < 7; ++i) { + expected[i] = (*native)(ids[i]); + } + for (size_t i = 0; i < 7; ++i) { + double scale = 0; + for (uint32_t j = 0; j < dim; ++j) { + scale += static_cast(x[j]) * x[j] + + static_cast(x[ids[i] * dim + j]) * x[ids[i] * dim + j]; + } + const double roundoff = 32 * std::numeric_limits::epsilon() * scale; + REQUIRE(actual[i] == Catch::Approx(expected[i]).epsilon(1e-5).margin(roundoff)); + } + } + stable->compute_distances(ids, 7, after.data(), 0, false, nullptr); + REQUIRE(initial == after); + REQUIRE(rbq->qb == 4); + REQUIRE_THROWS(store.CreateDistanceComputer(false, 9)); + } + } + fs::remove_all(dir); +} + +TEST_CASE("Test DISKANN_RABITQ rejects inconsistent sidecars", "[diskann][rabitq][sidecar]") { + const auto sidecar_dir = kDir + "/rabitq_sidecar_validation"; + const auto data_path = sidecar_dir + "/base.fbin"; + const auto sidecar_path = sidecar_dir + "/base_rabitq.index"; + fs::remove_all(sidecar_dir); + REQUIRE_NOTHROW(fs::create_directories(sidecar_dir)); + + constexpr uint32_t rows = 32; + constexpr uint32_t dim = 64; + auto base_ds = GenDataSet(rows, dim, 30); + WriteRawDataToDisk(data_path, static_cast(base_ds->GetTensor()), rows, dim); + REQUIRE_NOTHROW(knowhere::RaBitQStore::BuildFromFloatBin(data_path, sidecar_path, 4)); + + const auto corrupt_and_check = [&](const std::string& suffix, const auto& corrupt) { + std::unique_ptr index(faiss::cppcontrib::knowhere::read_index(sidecar_path.c_str())); + auto* pretransform = dynamic_cast(index.get()); + REQUIRE(pretransform != nullptr); + auto* rotation = dynamic_cast(pretransform->chain[0]); + auto* rabitq = dynamic_cast(pretransform->index); + REQUIRE(rotation != nullptr); + REQUIRE(rabitq != nullptr); + corrupt(*pretransform, *rotation, *rabitq); + const auto corrupted_path = sidecar_dir + "/" + suffix + ".index"; + faiss::cppcontrib::knowhere::write_index(index.get(), corrupted_path.c_str()); + REQUIRE_THROWS(knowhere::RaBitQStore(corrupted_path)); + }; + + corrupt_and_check("wrong_outer_metric", [](auto& pretransform, auto&, auto&) { + pretransform.metric_type = faiss::METRIC_INNER_PRODUCT; + }); + + fs::remove_all(sidecar_dir); +} + +TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { + fs::remove_all(kDir); + REQUIRE_NOTHROW(fs::create_directories(kDir)); + + auto version = GenTestVersionList(); + auto base_ds = GenDataSet(kNumRows, kDim, 30); + auto query_ds = GenDataSet(kNumQueries, kDim, 42); + WriteRawDataToDisk(kRawDataPath, static_cast(base_ds->GetTensor()), kNumRows, kDim); + + for (const int rbq_bits : {1, 2, 4, 8, 9}) { + CAPTURE(rbq_bits); + const auto rabitq_dir = kDir + "/rabitq_index_" + std::to_string(rbq_bits); + const auto rabitq_prefix = rabitq_dir + "/l2"; + REQUIRE_NOTHROW(fs::create_directories(rabitq_dir)); + + knowhere::Json build_json = {{"dim", kDim}, + {"metric_type", knowhere::metric::L2}, + {"index_prefix", rabitq_prefix}, + {"data_path", kRawDataPath}, + {"max_degree", defaultMaxDegree}, + {"search_list_size", 128}, + {"pq_code_budget_gb", 0.001}, + {"build_dram_budget_gb", 1.0}, + {"disk_pq_dims", 0}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"rbq_bits", rbq_bits}}; + knowhere::Json deserialize_json = {{"dim", kDim}, + {"metric_type", knowhere::metric::L2}, + {"index_prefix", rabitq_prefix}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"warm_up", false}}; + if (rbq_bits == 4) { + deserialize_json["search_cache_budget_gb"] = 0.00005; + // RaBitQ must safely force BFS even when the default sample-query + // cache mode is requested, because its navigation PQ is not resident. + deserialize_json["use_bfs_cache"] = false; + deserialize_json["bfs_cache_seed"] = 42; + } + knowhere::Json search_json = {{"dim", kDim}, + {"metric_type", knowhere::metric::L2}, + {"k", kK}, + {"search_list_size", 128}, + {"beamwidth", 8}}; + + auto file_manager = std::make_shared(); + auto pack = knowhere::Pack(std::shared_ptr(file_manager)); + knowhere::BinarySet binset; + { + auto index = knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, pack) + .value(); + REQUIRE(index.Build(nullptr, build_json) == knowhere::Status::success); + REQUIRE(fs::exists(rabitq_prefix + "_rabitq.index")); + REQUIRE(index.Serialize(binset) == knowhere::Status::success); + } + + if (rbq_bits == 1) { + auto full_reader = std::make_shared(); + diskann::PQFlashIndex full_pq_index(full_reader, diskann::Metric::L2); + REQUIRE(full_pq_index.load(1, rabitq_prefix.c_str(), true) == 0); + + auto metadata_reader = std::make_shared(); + diskann::PQFlashIndex metadata_only_index(metadata_reader, diskann::Metric::L2); + REQUIRE(metadata_only_index.load(1, rabitq_prefix.c_str(), false) == 0); + REQUIRE(metadata_only_index.get_num_points() == full_pq_index.get_num_points()); + REQUIRE(metadata_only_index.get_data_dim() == full_pq_index.get_data_dim()); + + const auto pq_file_size = fs::file_size(rabitq_prefix + "_pq_compressed.bin"); + const auto pq_code_bytes = pq_file_size - 2 * sizeof(uint32_t); + REQUIRE(full_pq_index.cal_size() - metadata_only_index.cal_size() == pq_code_bytes); + + std::array query{}; + std::array ids{}; + std::array distances{}; + REQUIRE_THROWS(metadata_only_index.cached_beam_search(query.data(), 1, 1, ids.data(), distances.data(), + 1)); + } + + auto index = knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, pack) + .value(); + REQUIRE(index.Deserialize(binset, deserialize_json) == knowhere::Status::success); + REQUIRE(index.Type() == knowhere::IndexEnum::INDEX_DISKANN_RABITQ); + auto result = index.Search(query_ds, search_json, nullptr); + REQUIRE(result.has_value()); + search_json["rbq_refine_mode"] = "full"; + auto full_result = index.Search(query_ds, search_json, nullptr); + REQUIRE(full_result.has_value()); + auto ground_truth = knowhere::BruteForce::Search(base_ds, query_ds, search_json, nullptr); + REQUIRE(ground_truth.has_value()); + REQUIRE(GetKNNRecall(*ground_truth.value(), *result.value()) > 0.5f); + REQUIRE(GetKNNRecall(*ground_truth.value(), *full_result.value()) > 0.5f); + search_json["rbq_refine_mode"] = "probabilistic"; + + std::vector empty_bitset_data((kNumRows + 7) / 8, 0); + auto empty_bitset_result = + index.Search(query_ds, search_json, knowhere::BitsetView(empty_bitset_data.data(), kNumRows)); + REQUIRE(empty_bitset_result.has_value()); + + auto bitset_data = GenerateBitsetWithFirstTbitsSet(kNumRows, 1); + auto bitset_result = index.Search(query_ds, search_json, knowhere::BitsetView(bitset_data.data(), kNumRows)); + REQUIRE_FALSE(bitset_result.has_value()); + REQUIRE(bitset_result.error() == knowhere::Status::not_implemented); + auto iterators = index.AnnIterator(query_ds, search_json, nullptr); + REQUIRE_FALSE(iterators.has_value()); + REQUIRE(iterators.error() == knowhere::Status::not_implemented); + } + + fs::remove_all(kDir); +} + +TEST_CASE("Test DISKANN_RABITQ inner product d+1 sidecar", "[diskann][rabitq][ip]") { + const auto ip_dir = kDir + "/rabitq_ip_index"; + const auto ip_prefix = ip_dir + "/ip"; + fs::remove_all(kDir); + REQUIRE_NOTHROW(fs::create_directories(ip_dir)); + + auto version = GenTestVersionList(); + auto base_ds = GenDataSet(kNumRows, kDim, 30); + auto query_ds = GenDataSet(kNumQueries, kDim, 42); + WriteRawDataToDisk(kRawDataPath, static_cast(base_ds->GetTensor()), kNumRows, kDim); + + knowhere::Json build_json = {{"dim", kDim}, + {"metric_type", knowhere::metric::IP}, + {"index_prefix", ip_prefix}, + {"data_path", kRawDataPath}, + {"max_degree", defaultMaxDegree}, + {"search_list_size", 128}, + {"pq_code_budget_gb", 0.001}, + {"build_dram_budget_gb", 1.0}, + {"disk_pq_dims", 0}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"rbq_bits", 4}}; + knowhere::Json deserialize_json = {{"dim", kDim}, + {"metric_type", knowhere::metric::IP}, + {"index_prefix", ip_prefix}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"warm_up", false}}; + knowhere::Json search_json = {{"dim", kDim}, + {"metric_type", knowhere::metric::IP}, + {"k", kK}, + {"search_list_size", 128}, + {"beamwidth", 8}}; + + auto file_manager = std::make_shared(); + auto pack = knowhere::Pack(std::shared_ptr(file_manager)); + knowhere::BinarySet binset; + { + auto index = knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, pack) + .value(); + REQUIRE(index.Build(nullptr, build_json) == knowhere::Status::success); + REQUIRE(fs::exists(ip_prefix + "_rabitq.index")); + REQUIRE_FALSE(fs::exists(ip_prefix + "_prepped_base.bin")); + REQUIRE(index.Serialize(binset) == knowhere::Status::success); + } + + auto index = knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, pack) + .value(); + REQUIRE(index.Deserialize(binset, deserialize_json) == knowhere::Status::success); + auto result = index.Search(query_ds, search_json, nullptr); + REQUIRE(result.has_value()); + auto ground_truth = knowhere::BruteForce::Search(base_ds, query_ds, search_json, nullptr); + REQUIRE(ground_truth.has_value()); + REQUIRE(GetKNNRecall(*ground_truth.value(), *result.value()) > 0.8f); + + fs::remove_all(kDir); +} + +TEST_CASE("Test AiSAQ clamps navigation PQ for deserialize", "[diskann][aisaq][large_pq]") { + constexpr uint32_t rows = 300; + constexpr uint32_t dim = 513; + const auto test_dir = kDir + "/aisaq_large_pq"; + const auto prefix = test_dir + "/l2"; + const auto data_path = test_dir + "/base.fbin"; + fs::remove_all(test_dir); + REQUIRE_NOTHROW(fs::create_directories(test_dir)); + + auto base_ds = GenDataSet(rows, dim, 30); + WriteRawDataToDisk(data_path, static_cast(base_ds->GetTensor()), rows, dim); + + knowhere::Json build_json = {{"dim", dim}, + {"metric_type", knowhere::metric::L2}, + {"index_prefix", prefix}, + {"data_path", data_path}, + {"max_degree", 16}, + {"search_list_size", 32}, + {"pq_code_budget_gb", + static_cast(rows) * dim / (1024.0 * 1024.0 * 1024.0)}, + {"build_dram_budget_gb", 1.0}, + {"disk_pq_dims", 0}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"inline_pq", -1}, + {"rearrange", false}, + {"num_entry_points", 0}}; + knowhere::Json deserialize_json = {{"dim", dim}, + {"metric_type", knowhere::metric::L2}, + {"index_prefix", prefix}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"use_bfs_cache", true}}; + + auto file_manager = std::make_shared(); + auto pack = knowhere::Pack(std::shared_ptr(file_manager)); + auto version = GenTestVersionList(); + knowhere::BinarySet binset; + { + auto index = knowhere::IndexFactory::Instance().Create("AISAQ", version, pack).value(); + REQUIRE(index.Build(nullptr, build_json) == knowhere::Status::success); + REQUIRE(index.Serialize(binset) == knowhere::Status::success); + } + + std::unique_ptr chunk_offsets; + size_t num_offsets = 0; + size_t offset_dim = 0; + diskann::load_bin(prefix + "_pq_pivots.bin_chunk_offsets.bin", chunk_offsets, num_offsets, offset_dim); + REQUIRE(num_offsets == diskann::defaults::MAX_PQ_CHUNKS + 1); + REQUIRE(offset_dim == 1); + + auto loaded = knowhere::IndexFactory::Instance().Create("AISAQ", version, pack).value(); + REQUIRE(loaded.Deserialize(binset, deserialize_json) == knowhere::Status::success); + + fs::remove_all(test_dir); +} + #ifdef KNOWHERE_WITH_CUVS TEST_CASE("Test DiskANN Build Index", "[diskann]") { auto version = GenTestVersionList(); diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 078007839..0c7e88e17 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -43,6 +43,8 @@ namespace diskann { double get_memory_budget(const std::string &mem_budget_str); double get_memory_budget(double search_ram_budget_in_gb); + size_t get_num_pq_chunks(double pq_code_size_limit, size_t points_num, + size_t dim); void add_new_file_to_single_index(std::string index_file, std::string new_file); @@ -131,6 +133,9 @@ namespace diskann { uint32_t inline_pq = 0; bool rearrange = false; int num_entry_points = 0; + // Keep the temporary MIPS-to-L2 base until the caller builds auxiliary + // indexes that must use the exact same internal d+1 representation. + bool keep_preprocessed_base = false; }; template diff --git a/thirdparty/DiskANN/include/diskann/percentile_stats.h b/thirdparty/DiskANN/include/diskann/percentile_stats.h index 45a7e579e..a276910ab 100644 --- a/thirdparty/DiskANN/include/diskann/percentile_stats.h +++ b/thirdparty/DiskANN/include/diskann/percentile_stats.h @@ -28,6 +28,12 @@ namespace diskann { unsigned n_cache_hits = 0; // # cache_hits unsigned n_hops = 0; // # search hops unsigned n_iters = 0; // # range search iterations + + // Optional approximate-distance refinement statistics. These remain zero + // for ordinary DiskANN PQ navigation. + unsigned n_approx_estimates = 0; + unsigned n_approx_refinements = 0; + unsigned n_approx_pruned = 0; }; template diff --git a/thirdparty/DiskANN/include/diskann/pq_flash_index.h b/thirdparty/DiskANN/include/diskann/pq_flash_index.h index 107e656b3..80c1c1560 100644 --- a/thirdparty/DiskANN/include/diskann/pq_flash_index.h +++ b/thirdparty/DiskANN/include/diskann/pq_flash_index.h @@ -183,6 +183,20 @@ namespace diskann { virtual ~PQDataGetter() {} }; + // Optional query-local navigation distance implementation. Ordinary + // DiskANN leaves this null and uses its resident PQ codes. Knowhere's + // DISKANN_RABITQ supplies one instance per query because the distance + // computer owns transformed-query scratch and is not thread-safe. + class ApproxDistanceComputer { + public: + virtual ~ApproxDistanceComputer() = default; + virtual void set_query(const float* query) = 0; + virtual void compute_distances(const unsigned* ids, _u64 n_ids, + float* distances, float threshold, + bool threshold_valid, + QueryStats* stats) = 0; + }; + template class PQFlashIndex: public PQDataGetter { public: @@ -191,7 +205,8 @@ namespace diskann { ~PQFlashIndex(); // load compressed data, and obtains the handle to the disk-resident index - int load(uint32_t num_threads, const char *index_prefix); + int load(uint32_t num_threads, const char *index_prefix, + bool load_pq_data = true); virtual void load_cache_list(std::vector &node_list); @@ -202,7 +217,8 @@ namespace diskann { _u64 num_nodes_to_cache); virtual void cache_bfs_levels(_u64 num_nodes_to_cache, - std::vector &node_list); + std::vector &node_list, + _s64 bfs_seed = -1); void cached_beam_search( const T *query, const _u64 k_search, const _u64 l_search, _s64 *res_ids, @@ -210,7 +226,8 @@ namespace diskann { const bool use_reorder_data = false, QueryStats *stats = nullptr, const knowhere::feder::diskann::FederResultUniq &feder = nullptr, knowhere::BitsetView bitset_view = nullptr, - const float filter_ratio = -1.0f); + const float filter_ratio = -1.0f, + ApproxDistanceComputer* approx_distance_computer = nullptr); void calc_dist_by_ids(const T *query, const int64_t *ids, const int64_t n, float *const output_dists); @@ -293,7 +310,8 @@ namespace diskann { IOContext &ctx, QueryStats *stats, const knowhere::feder::diskann::FederResultUniq &feder, knowhere::BitsetView bitset_view, - PQDataGetter* pq_data_getter); + PQDataGetter* pq_data_getter, + ApproxDistanceComputer* approx_distance_computer = nullptr); // Assign the index of ids to its corresponding sector and if it is in // cache, write to the output_data diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index f575cd454..1b977d3b2 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -1607,6 +1607,17 @@ void create_aisaq_layout(const std::string base_file, const std::string mem_inde } } +size_t get_num_pq_chunks(double pq_code_size_limit, size_t points_num, + size_t dim) { + if (points_num == 0 || dim == 0) { + return 0; + } + size_t num_pq_chunks = + static_cast(std::floor(pq_code_size_limit / points_num)); + num_pq_chunks = std::max(num_pq_chunks, 1); + return std::min(num_pq_chunks, dim); +} + template int build_disk_index(BuildConfig &config) { if (!knowhere::KnowhereFloatTypeCheck::value && @@ -1733,12 +1744,13 @@ template << " Indexing ram budget: " << indexing_ram_budget << "(GiB)"; - size_t num_pq_chunks = - (size_t) (std::floor)(_u64(pq_code_size_limit / points_num)); - - num_pq_chunks = num_pq_chunks <= 0 ? 1 : num_pq_chunks; - num_pq_chunks = num_pq_chunks > dim ? dim : num_pq_chunks; - num_pq_chunks = num_pq_chunks > diskann::defaults::MAX_PQ_CHUNKS ? diskann::defaults::MAX_PQ_CHUNKS : num_pq_chunks; + // The ordinary PQFlashIndex scratch space is dimension-sized, so an + // in-memory navigation code may safely use up to one chunk per input + // dimension. AiSAQ's on-disk layout and loader retain a 512-chunk limit. + const size_t requested_pq_chunks = get_num_pq_chunks(pq_code_size_limit, points_num, dim); + const size_t num_pq_chunks = config.aisaq_mode + ? std::min(requested_pq_chunks, static_cast(diskann::defaults::MAX_PQ_CHUNKS)) + : requested_pq_chunks; LOG_KNOWHERE_INFO_ << "Compressing " << dim << "-dimensional data into " << num_pq_chunks << " bytes per vector."; @@ -1889,7 +1901,8 @@ template std::chrono::duration diff = e - s; LOG_KNOWHERE_INFO_ << "Indexing time: " << diff.count(); - if (config.compare_metric == diskann::Metric::INNER_PRODUCT) { + if (config.compare_metric == diskann::Metric::INNER_PRODUCT && + !config.keep_preprocessed_base) { std::remove(data_file_to_use.c_str()); } std::remove(mem_index_path.c_str()); diff --git a/thirdparty/DiskANN/src/pq_flash_index.cpp b/thirdparty/DiskANN/src/pq_flash_index.cpp index 1e41b2022..083589ddd 100644 --- a/thirdparty/DiskANN/src/pq_flash_index.cpp +++ b/thirdparty/DiskANN/src/pq_flash_index.cpp @@ -12,6 +12,7 @@ #include #include #include +#include #include #include #include @@ -519,9 +520,15 @@ namespace diskann { template void PQFlashIndex::cache_bfs_levels(_u64 num_nodes_to_cache, - std::vector &node_list) { - std::random_device rng; - std::mt19937 urng(rng()); + std::vector &node_list, + _s64 bfs_seed) { + std::mt19937 urng; + if (bfs_seed < 0) { + std::random_device rng; + urng.seed(rng()); + } else { + urng.seed(static_cast(bfs_seed)); + } node_list.clear(); @@ -697,7 +704,8 @@ namespace diskann { } template - int PQFlashIndex::load(uint32_t num_threads, const char *index_prefix) { + int PQFlashIndex::load(uint32_t num_threads, const char *index_prefix, + bool load_pq_data) { std::string pq_table_bin = get_pq_pivots_filename(std::string(index_prefix)); std::string pq_compressed_vectors = @@ -727,18 +735,37 @@ namespace diskann { this->aligned_dim = ROUND_UP(pq_file_dim, 8); size_t npts_u64, nchunks_u64; - diskann::load_bin<_u8>(pq_compressed_vectors, this->data, npts_u64, - nchunks_u64); + if (load_pq_data) { + diskann::load_bin<_u8>(pq_compressed_vectors, this->data, npts_u64, + nchunks_u64); + } else { + get_bin_metadata(pq_compressed_vectors, npts_u64, nchunks_u64); + const size_t header_size = 2 * sizeof(uint32_t); + if (nchunks_u64 != 0 && + npts_u64 > (std::numeric_limits::max() - header_size) / + nchunks_u64) { + LOG(ERROR) << "PQ compressed vector metadata overflows file size"; + return -1; + } + const size_t expected_size = header_size + npts_u64 * nchunks_u64; + if (get_file_size(pq_compressed_vectors) != expected_size) { + LOG(ERROR) << "PQ compressed vector file size mismatch for " + << pq_compressed_vectors; + return -1; + } + } this->num_points = npts_u64; this->n_chunks = nchunks_u64; pq_table.load_pq_centroid_bin(pq_table_bin.c_str(), nchunks_u64); - LOG(INFO) - << "Loaded PQ centroids and in-memory compressed vectors. #points: " - << num_points << " #dim: " << data_dim - << " #aligned_dim: " << aligned_dim << " #chunks: " << n_chunks; + LOG(INFO) << "Loaded PQ centroids and " + << (load_pq_data ? "in-memory compressed vectors" + : "compressed-vector metadata only") + << ". #points: " << num_points << " #dim: " << data_dim + << " #aligned_dim: " << aligned_dim + << " #chunks: " << n_chunks; std::string disk_pq_pivots_path = this->disk_index_file + "_pq_pivots.bin"; if (file_exists(disk_pq_pivots_path)) { @@ -942,13 +969,22 @@ namespace diskann { IOContext &ctx, QueryStats *stats, const knowhere::feder::diskann::FederResultUniq &feder, knowhere::BitsetView bitset_view, - PQDataGetter* pq_data_getter) { + PQDataGetter* pq_data_getter, + ApproxDistanceComputer* approx_distance_computer) { + if (approx_distance_computer == nullptr && this->data == nullptr) { + throw ANNException( + "resident navigation PQ data is unavailable and no external " + "distance computer was supplied", + -1); + } auto query_scratch = &(data.scratch); const T *query = data.scratch.aligned_query_T; auto beam_width = beam_width_param * kRefineBeamWidthFactor; const float *query_float = data.scratch.aligned_query_float; float *pq_dists = query_scratch->aligned_pqtable_dist_scratch; - pq_table.populate_chunk_distances(query_float, pq_dists); + if (approx_distance_computer == nullptr) { + pq_table.populate_chunk_distances(query_float, pq_dists); + } float *dist_scratch = query_scratch->aligned_dist_scratch; _u8 *pq_coord_scratch = query_scratch->aligned_pq_coord_scratch; constexpr _u32 pq_batch_size = diskann::defaults::MAX_GRAPH_DEGREE; @@ -975,9 +1011,15 @@ namespace diskann { if (pq_batch_ids.size() == pq_batch_size || id == num_points - 1) { const size_t sz = pq_batch_ids.size(); - pq_data_getter->aggregate_pq_coords(pq_batch_ids.data(), sz, this->n_chunks, pq_coord_scratch); - pq_dist_lookup(pq_coord_scratch, sz, this->n_chunks, pq_dists, - dist_scratch); + if (approx_distance_computer != nullptr) { + approx_distance_computer->compute_distances( + pq_batch_ids.data(), sz, dist_scratch, + (std::numeric_limits::max)(), false, stats); + } else { + pq_data_getter->aggregate_pq_coords(pq_batch_ids.data(), sz, this->n_chunks, pq_coord_scratch); + pq_dist_lookup(pq_coord_scratch, sz, this->n_chunks, pq_dists, + dist_scratch); + } for (size_t i = 0; i < sz; ++i) { pq_max_heap.Push(dist_scratch[i], pq_batch_ids[i]); } @@ -1087,7 +1129,14 @@ namespace diskann { const T *query1, const _u64 k_search, const _u64 l_search, _s64 *indices, float *distances, const _u64 beam_width, const bool use_reorder_data, QueryStats *stats, const knowhere::feder::diskann::FederResultUniq &feder, - knowhere::BitsetView bitset_view, const float filter_ratio_in) { + knowhere::BitsetView bitset_view, const float filter_ratio_in, + ApproxDistanceComputer* approx_distance_computer) { + if (approx_distance_computer == nullptr && this->data == nullptr) { + throw ANNException( + "resident navigation PQ data is unavailable and no external " + "distance computer was supplied", + -1); + } if (beam_width > defaults::MAX_N_SECTOR_READS) throw ANNException("Beamwidth can not be higher than MAX_N_SECTOR_READS", -1, __FUNCSIG__, __FILE__, __LINE__); @@ -1107,6 +1156,11 @@ namespace diskann { float query_norm = query_norm_opt.value(); auto ctx = this->reader->get_ctx(); + if (approx_distance_computer != nullptr) { + approx_distance_computer->set_query( + data.scratch.aligned_query_float); + } + size_t bv_cnt = 0; if (!bitset_view.empty()) { @@ -1137,7 +1191,8 @@ namespace diskann { if (bv_cnt >= bitset_view.size() * filter_threshold) { brute_force_beam_search(data, query_norm, k_search, indices, distances, - beam_width, ctx, stats, feder, bitset_view, this); + beam_width, ctx, stats, feder, bitset_view, this, + approx_distance_computer); this->thread_data.push(data); this->thread_data.push_notify_all(); this->reader->put_ctx(ctx); @@ -1148,7 +1203,8 @@ namespace diskann { // Turn to BF is k_search is too large if (k_search > 0.5 * (num_points - bv_cnt)) { brute_force_beam_search(data, query_norm, k_search, indices, distances, - beam_width, ctx, stats, feder, bitset_view, this); + beam_width, ctx, stats, feder, bitset_view, this, + approx_distance_computer); this->thread_data.push(data); this->thread_data.push_notify_all(); this->reader->put_ctx(ctx); @@ -1177,23 +1233,36 @@ namespace diskann { std::vector>> cached_nhoods; cached_nhoods.reserve(2 * beam_width); + std::vector> beam_nhood_order; + beam_nhood_order.reserve(2 * beam_width); // query <-> PQ chunk centers distances float *pq_dists = query_scratch->aligned_pqtable_dist_scratch; - pq_table.populate_chunk_distances(query_float, pq_dists); + if (approx_distance_computer == nullptr) { + pq_table.populate_chunk_distances(query_float, pq_dists); + } // query <-> neighbor list float *dist_scratch = query_scratch->aligned_dist_scratch; _u8 *pq_coord_scratch = query_scratch->aligned_pq_coord_scratch; // lambda to batch compute query<-> node distances in PQ space - auto compute_dists = [this, pq_coord_scratch, pq_dists](const unsigned *ids, - const _u64 n_ids, - float *dists_out) { - aggregate_coords(ids, n_ids, this->data.get(), this->n_chunks, - pq_coord_scratch); - pq_dist_lookup(pq_coord_scratch, n_ids, this->n_chunks, pq_dists, - dists_out); + auto compute_dists = [this, pq_coord_scratch, pq_dists, + approx_distance_computer](const unsigned *ids, + const _u64 n_ids, + float *dists_out, + float threshold, + bool threshold_valid, + QueryStats *stats) { + if (approx_distance_computer != nullptr) { + approx_distance_computer->compute_distances( + ids, n_ids, dists_out, threshold, threshold_valid, stats); + } else { + aggregate_coords(ids, n_ids, this->data.get(), this->n_chunks, + pq_coord_scratch); + pq_dist_lookup(pq_coord_scratch, n_ids, this->n_chunks, pq_dists, + dists_out); + } }; Timer cpu_timer; std::vector retset(l_search + 1); @@ -1218,7 +1287,8 @@ namespace diskann { } } - compute_dists(&best_medoid, 1, dist_scratch); + compute_dists(&best_medoid, 1, dist_scratch, + (std::numeric_limits::max)(), false, stats); retset[0].id = best_medoid; retset[0].flag = true; retset[0].distance = dist_scratch[0]; @@ -1265,6 +1335,7 @@ namespace diskann { frontier_nhoods.clear(); frontier_read_reqs.clear(); cached_nhoods.clear(); + beam_nhood_order.clear(); sector_scratch_idx = 0; // find new beam _u32 marker = k; @@ -1279,11 +1350,13 @@ namespace diskann { if (iter != nhood_cache.end()) { cached_nhoods.push_back( std::make_pair(retset[marker].id, iter->second)); + beam_nhood_order.emplace_back(true, cached_nhoods.size() - 1); if (stats != nullptr) { stats->n_cache_hits++; } } else { frontier.push_back(retset[marker].id); + beam_nhood_order.emplace_back(false, frontier.size() - 1); } } retset[marker].flag = false; @@ -1362,7 +1435,12 @@ namespace diskann { // compute node_nbrs <-> query dists in PQ space cpu_timer.reset(); - compute_dists(node_nbrs, nnbrs, dist_scratch); + const bool threshold_valid = cur_list_size == l_search; + const float threshold = threshold_valid + ? retset[cur_list_size - 1].distance + : (std::numeric_limits::max)(); + compute_dists(node_nbrs, nnbrs, dist_scratch, threshold, + threshold_valid, stats); if (stats != nullptr) { stats->n_cmps += (double) nnbrs; stats->cpu_us += (double) cpu_timer.elapsed(); @@ -1402,30 +1480,34 @@ namespace diskann { } }; - // process cached nhoods - for (auto &cached_nhood : cached_nhoods) { - if (stats != nullptr) { - stats->n_hops++; - } - T *node_fp_coords_copy; - { - std::shared_lock lock(this->cache_mtx); - auto global_cache_iter = coord_cache.find(cached_nhood.first); - node_fp_coords_copy = global_cache_iter->second; + // Preserve the beam order when cache hits and SSD reads are mixed. In + // particular, threshold-aware approximate scorers must observe the same + // candidate sequence with and without a node cache. + for (const auto &[is_cached, nhood_index] : beam_nhood_order) { + if (is_cached) { + auto &cached_nhood = cached_nhoods[nhood_index]; + if (stats != nullptr) { + stats->n_hops++; + } + T *node_fp_coords_copy; + { + std::shared_lock lock(this->cache_mtx); + auto global_cache_iter = coord_cache.find(cached_nhood.first); + node_fp_coords_copy = global_cache_iter->second; + } + process_node(node_fp_coords_copy, cached_nhood.first, + cached_nhood.second.first, cached_nhood.second.second); + } else { + auto &frontier_nhood = frontier_nhoods[nhood_index]; + char *node_disk_buf = + get_offset_to_node(frontier_nhood.second, frontier_nhood.first); + unsigned *node_buf = OFFSET_TO_NODE_NHOOD(node_disk_buf); + T *node_fp_coords = OFFSET_TO_NODE_COORDS(node_disk_buf); + T *node_fp_coords_copy = data_buf; + memcpy(node_fp_coords_copy, node_fp_coords, disk_bytes_per_point); + process_node(node_fp_coords_copy, frontier_nhood.first, *node_buf, + node_buf + 1); } - process_node(node_fp_coords_copy, cached_nhood.first, - cached_nhood.second.first, cached_nhood.second.second); - } - - for (auto &frontier_nhood : frontier_nhoods) { - char *node_disk_buf = - get_offset_to_node(frontier_nhood.second, frontier_nhood.first); - unsigned *node_buf = OFFSET_TO_NODE_NHOOD(node_disk_buf); - T *node_fp_coords = OFFSET_TO_NODE_COORDS(node_disk_buf); - T *node_fp_coords_copy = data_buf; - memcpy(node_fp_coords_copy, node_fp_coords, disk_bytes_per_point); - process_node(node_fp_coords_copy, frontier_nhood.first, *node_buf, - node_buf + 1); } // update best inserted position @@ -2053,7 +2135,9 @@ namespace diskann { index_mem_size += ROUND_UP(num_medoids * aligned_dim * sizeof(float), 32); index_mem_size += num_medoids * aligned_dim * sizeof(uint32_t); // get pq data and pq_table: - index_mem_size += this->num_points * this->n_chunks * sizeof(uint8_t); + if (this->data != nullptr) { + index_mem_size += this->num_points * this->n_chunks * sizeof(uint8_t); + } index_mem_size += this->pq_table.get_total_dims() * 256 * sizeof(float) * 2; index_mem_size += this->pq_table.get_total_dims() * (sizeof(uint32_t) + sizeof(float)); From 87fb14ad48895c2cab1e4b6b63b8a2b3a52d5e68 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Tue, 15 Sep 2026 08:09:36 +0000 Subject: [PATCH 02/12] feat: generalize DiskANN navigation and batch RaBitQ estimates Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 184 ++++++-------- src/index/diskann/diskann_config.h | 47 +++- src/index/diskann/navigation_store.cc | 51 ++++ src/index/diskann/navigation_store.h | 40 +++ src/index/diskann/rabitq_store.cc | 107 ++++++--- src/index/diskann/rabitq_store.h | 15 +- tests/python/test_diskann_rabitq.py | 44 ++++ tests/ut/test_diskann.cc | 227 ++++++++++++++---- .../DiskANN/include/diskann/pq_flash_index.h | 19 +- thirdparty/DiskANN/src/pq_flash_index.cpp | 36 +-- 10 files changed, 544 insertions(+), 226 deletions(-) create mode 100644 src/index/diskann/navigation_store.cc create mode 100644 src/index/diskann/navigation_store.h create mode 100644 tests/python/test_diskann_rabitq.py diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index c7cf01dcf..6d3e726ac 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -11,6 +11,8 @@ #include "knowhere/feder/DiskANN.h" +#include + #include #include #include @@ -23,7 +25,7 @@ #include "filemanager/FileManager.h" #include "fmt/core.h" #include "index/diskann/diskann_config.h" -#include "index/diskann/rabitq_store.h" +#include "index/diskann/navigation_store.h" #include "knowhere/comp/index_param.h" #include "knowhere/context.h" #include "knowhere/dataset.h" @@ -113,6 +115,16 @@ class DiskANNIndexNode : public IndexNode { static Status StaticConfigCheck(const Config& cfg, PARAM_TYPE paramType, std::string& msg) { auto& base_cfg = static_cast(cfg); + if (UsesExternalNavigation(static_cast(cfg))) { + if (!std::is_same_v) { + msg = "external DiskANN navigation currently requires FP32 data"; + return Status::invalid_args; + } + if (base_cfg.emb_list_strategy.has_value() || base_cfg.emb_list_offset_file_path.has_value()) { + msg = "external DiskANN navigation does not support embedding-list mode"; + return Status::not_implemented; + } + } auto strategy = base_cfg.emb_list_strategy.value_or(""); if (strategy == meta::EMB_LIST_STRATEGY_MUVERA || strategy == meta::EMB_LIST_STRATEGY_LEMUR) { msg = "DiskANN only supports TokenANN strategy, got '" + strategy + "'"; @@ -168,7 +180,7 @@ class DiskANNIndexNode : public IndexNode { static std::unique_ptr StaticCreateConfig() { - return std::make_unique(); + return std::make_unique(); } std::unique_ptr @@ -202,8 +214,8 @@ class DiskANNIndexNode : public IndexNode { return 0; } auto size = pq_flash_index_->cal_size(); - if (rabitq_store_ != nullptr) { - size += rabitq_store_->MemorySize(); + if (navigation_store_ != nullptr) { + size += navigation_store_->MemorySize(); } return size; } @@ -230,11 +242,6 @@ class DiskANNIndexNode : public IndexNode { expected GetVectorByStorageIds(const DataSetPtr dataset, milvus::OpContext* op_context) const override; - virtual bool - IsRaBitQ() const { - return false; - } - private: class iterator : public IndexIterator { public: @@ -293,7 +300,7 @@ class DiskANNIndexNode : public IndexNode { std::atomic_bool is_prepared_; std::shared_ptr file_manager_; std::unique_ptr> pq_flash_index_; - std::unique_ptr rabitq_store_; + std::unique_ptr navigation_store_; std::atomic_int64_t dim_; std::atomic_int64_t count_; std::shared_ptr search_pool_; @@ -305,28 +312,6 @@ namespace knowhere { namespace { static constexpr float kCacheExpansionRate = 1.2; -bool -HasFilteredBits(const BitsetView& bitset) { - if (bitset.empty()) { - return false; - } - - const auto* data = bitset.data(); - const auto full_bytes = bitset.size() / 8; - for (size_t i = 0; i < full_bytes; ++i) { - if (data[i] != 0) { - return true; - } - } - - const auto remaining_bits = bitset.size() % 8; - if (remaining_bits == 0) { - return false; - } - const auto valid_bits_mask = static_cast((1u << remaining_bits) - 1u); - return (data[full_bytes] & valid_bits_mask) != 0; -} - Status ReadEmbListOffsetFromFile(const std::string& file_path, std::vector& offsets) { std::ifstream in_file(file_path, std::ios::binary); @@ -431,7 +416,7 @@ GetOptionalFilenames(const std::string& prefix) { } inline bool -AnyIndexFileExist(const std::string& index_prefix) { +AnyIndexFileExist(const std::string& index_prefix, const DiskANNConfig& config) { auto file_exist = [](std::vector filenames) -> bool { for (auto& filename : filenames) { if (file_exists(filename)) { @@ -441,7 +426,7 @@ AnyIndexFileExist(const std::string& index_prefix) { return false; }; return file_exist(GetNecessaryFilenames(index_prefix, diskann::INNER_PRODUCT, true, true)) || - file_exist(GetOptionalFilenames(index_prefix)) || file_exists(RaBitQStore::SidecarFilename(index_prefix)); + file_exist(GetOptionalFilenames(index_prefix)) || file_exist(NavigationFiles(config, index_prefix)); } inline bool @@ -464,7 +449,11 @@ template Status DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr cfg, bool use_knowhere_build_pool) { assert(file_manager_ != nullptr); - auto build_conf = static_cast(*cfg); + const auto& build_conf = static_cast(*cfg); + const bool external_navigation = UsesExternalNavigation(build_conf); + if (external_navigation && !std::is_same_v) { + return Status::invalid_args; + } if (!CheckMetric(build_conf.metric_type.value())) { LOG_KNOWHERE_ERROR_ << "Invalid metric type: " << build_conf.metric_type.value(); return Status::invalid_metric_type; @@ -473,7 +462,7 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::Build(const DataSetPtr dataset, std::shared_ptr(num_nodes_to_cache), build_conf.shuffle_build.value()}; diskann_internal_build_config.keep_preprocessed_base = - IsRaBitQ() && diskann_metric == diskann::Metric::INNER_PRODUCT; + external_navigation && diskann_metric == diskann::Metric::INNER_PRODUCT; RETURN_IF_ERROR(TryDiskANNCall([&]() { int res = diskann::build_disk_index(diskann_internal_build_config); if (res != 0) @@ -535,20 +524,11 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr(cfg.get()); - if (rabitq_conf == nullptr) { - LOG_KNOWHERE_ERROR_ << "DISKANN_RABITQ received an unexpected config type"; - return Status::invalid_args; - } + if (external_navigation) { try { - const auto sidecar_path = RaBitQStore::SidecarFilename(index_prefix_); - const auto sidecar_source = diskann_metric == diskann::Metric::INNER_PRODUCT - ? index_prefix_ + "_prepped_base.bin" - : data_path; - LOG_KNOWHERE_INFO_ << "Building DiskANN RaBitQ sidecar: " << sidecar_path; - RaBitQStore::BuildFromFloatBin(sidecar_source, sidecar_path, - static_cast(rabitq_conf->rbq_bits.value())); + const auto sidecar_source = + diskann_metric == diskann::Metric::INNER_PRODUCT ? index_prefix_ + "_prepped_base.bin" : data_path; + BuildNavigationStore(build_conf, sidecar_source, index_prefix_); if (diskann_metric == diskann::Metric::INNER_PRODUCT) { std::error_code error; std::filesystem::remove(sidecar_source, error); @@ -558,7 +538,7 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::Build(const DataSetPtr dataset, std::shared_ptr::BuildEmbListIfNeed(const DataSetPtr dataset, std::sh // If not emb_list metric type, use the default build method return Build(dataset, std::move(cfg), use_knowhere_build_pool); } - if (IsRaBitQ()) { + if (UsesExternalNavigation(static_cast(*cfg))) { LOG_KNOWHERE_ERROR_ << "DISKANN_RABITQ does not support embedding-list mode"; return Status::not_implemented; } @@ -660,13 +639,25 @@ DiskANNIndexNode::BuildEmbListIfNeed(const DataSetPtr dataset, std::sh template Status DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr cfg) { - auto prep_conf = static_cast(*cfg); + const auto& prep_conf = static_cast(*cfg); + const bool external_navigation = UsesExternalNavigation(prep_conf); + if (external_navigation && !std::is_same_v) { + return Status::invalid_args; + } if (!CheckMetric(prep_conf.metric_type.value())) { return Status::invalid_metric_type; } if (is_prepared_.load()) { return Status::success; } + const auto rollback = folly::makeGuard([this]() { + if (!is_prepared_.load()) { + navigation_store_.reset(); + pq_flash_index_.reset(); + count_.store(-1); + dim_.store(-1); + } + }); if (!(prep_conf.index_prefix.has_value())) { LOG_KNOWHERE_ERROR_ << "DiskANN file path for deserialize is empty." << std::endl; return Status::invalid_param_in_json; @@ -688,7 +679,7 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr // Load file from file manager. for (auto& filename : GetNecessaryFilenames( index_prefix_, need_norm, - prep_conf.search_cache_budget_gb.value() > 0 && !prep_conf.use_bfs_cache.value() && !IsRaBitQ(), + prep_conf.search_cache_budget_gb.value() > 0 && !prep_conf.use_bfs_cache.value() && !external_navigation, prep_conf.warm_up.value())) { if (!LoadFile(filename)) { return Status::disk_file_error; @@ -704,10 +695,9 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr return Status::disk_file_error; } } - if (IsRaBitQ()) { - const auto sidecar_path = RaBitQStore::SidecarFilename(index_prefix_); + for (const auto& sidecar_path : NavigationFiles(prep_conf, index_prefix_)) { if (!LoadFile(sidecar_path)) { - LOG_KNOWHERE_ERROR_ << "Failed to load DiskANN RaBitQ sidecar " << sidecar_path; + LOG_KNOWHERE_ERROR_ << "Failed to load DiskANN navigation sidecar " << sidecar_path; return Status::disk_file_error; } } @@ -722,7 +712,7 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr pq_flash_index_ = std::make_unique>(reader, diskann_metric); auto disk_ann_call = [&]() { - int res = pq_flash_index_->load(search_pool_->size(), index_prefix_.c_str(), !IsRaBitQ()); + int res = pq_flash_index_->load(search_pool_->size(), index_prefix_.c_str(), !external_navigation); if (res != 0) { throw diskann::ANNException("pq_flash_index_->load returned non-zero value: " + std::to_string(res), -1); } @@ -740,18 +730,19 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr dim_.store(pq_flash_index_->get_data_dim()); } - if (IsRaBitQ()) { + navigation_store_.reset(); + if (external_navigation) { try { - rabitq_store_ = std::make_unique(RaBitQStore::SidecarFilename(index_prefix_)); - if (rabitq_store_->Count() != static_cast(pq_flash_index_->get_num_points()) || - rabitq_store_->Dimension() != static_cast(pq_flash_index_->get_data_dim())) { - LOG_KNOWHERE_ERROR_ << "DiskANN graph and RaBitQ sidecar metadata do not match"; - rabitq_store_.reset(); + navigation_store_ = LoadNavigationStore(prep_conf, index_prefix_); + if (navigation_store_->Count() != static_cast(pq_flash_index_->get_num_points()) || + navigation_store_->Dimension() != static_cast(pq_flash_index_->get_data_dim())) { + LOG_KNOWHERE_ERROR_ << "DiskANN graph and navigation sidecar metadata do not match"; + navigation_store_.reset(); return Status::invalid_index_error; } } catch (const std::exception& e) { - LOG_KNOWHERE_ERROR_ << "Failed to initialize DiskANN RaBitQ sidecar: " << e.what(); - rabitq_store_.reset(); + LOG_KNOWHERE_ERROR_ << "Failed to initialize DiskANN navigation sidecar: " << e.what(); + navigation_store_.reset(); return Status::invalid_index_error; } } @@ -787,9 +778,9 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr } if (num_nodes_to_cache > 0) { LOG_KNOWHERE_INFO_ << "Caching " << num_nodes_to_cache << " sample nodes around medoid(s)."; - if (prep_conf.use_bfs_cache.value() || IsRaBitQ()) { - if (IsRaBitQ() && !prep_conf.use_bfs_cache.value()) { - LOG_KNOWHERE_INFO_ << "DISKANN_RABITQ uses BFS cache generation because navigation PQ is not " + if (prep_conf.use_bfs_cache.value() || external_navigation) { + if (external_navigation && !prep_conf.use_bfs_cache.value()) { + LOG_KNOWHERE_INFO_ << "External navigation uses BFS cache generation because navigation PQ is not " "resident"; } LOG_KNOWHERE_INFO_ << "Use bfs to generate cache list"; @@ -853,9 +844,10 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr futures.reserve(warmup_num); for (uint64_t i = 0; i < warmup_num; ++i) { futures.emplace_back(search_pool_->push([&, index = i]() { - pq_flash_index_->cached_beam_search(warmup + (index * warmup_aligned_dim), 1, warmup_L, - warmup_result_ids_64.data() + (index * 1), - warmup_result_dists.data() + (index * 1), 4); + auto navigation = navigation_store_ ? navigation_store_->CreateDistanceComputer(prep_conf) : nullptr; + pq_flash_index_->cached_beam_search( + warmup + (index * warmup_aligned_dim), 1, warmup_L, warmup_result_ids_64.data() + (index * 1), + warmup_result_dists.data() + (index * 1), 4, false, nullptr, nullptr, {}, -1.0f, navigation.get()); })); } @@ -886,7 +878,7 @@ DiskANNIndexNode::DeserializeEmbListIfNeed(const BinarySet& binset, st // If not emb_list metric type, use the default deserialize method return Deserialize(binset, std::move(cfg)); } - if (IsRaBitQ()) { + if (UsesExternalNavigation(static_cast(*cfg))) { LOG_KNOWHERE_ERROR_ << "DISKANN_RABITQ does not support embedding-list mode"; return Status::not_implemented; } @@ -951,16 +943,16 @@ template expected> DiskANNIndexNode::AnnIterator(const DataSetPtr dataset, std::unique_ptr cfg, const BitsetView& bitset, bool use_knowhere_search_pool, milvus::OpContext* op_context) const { - if (IsRaBitQ()) { - return expected>::Err( - Status::not_implemented, "DISKANN_RABITQ does not support iterator search"); + if (navigation_store_ || UsesExternalNavigation(static_cast(*cfg))) { + return expected>::Err(Status::not_implemented, + "DISKANN_RABITQ does not support iterator search"); } if (!is_prepared_.load() || !pq_flash_index_) { LOG_KNOWHERE_ERROR_ << "Failed to load diskann."; return expected>::Err(Status::empty_index, "DiskANN not loaded"); } - auto search_conf = static_cast(*cfg); + const auto& search_conf = static_cast(*cfg); if (!CheckMetric(search_conf.metric_type.value())) { return expected>::Err(Status::invalid_metric_type, "unsupported metric type"); @@ -1005,7 +997,7 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptr::Err(Status::empty_index, "DiskANN not loaded"); } - auto search_conf = static_cast(*cfg); + const auto& search_conf = static_cast(*cfg); if (!CheckMetric(search_conf.metric_type.value())) { return expected::Err(Status::invalid_metric_type, "unsupported metric type"); } @@ -1017,22 +1009,6 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptrGetDim(); auto xq = static_cast(dataset->GetTensor()); - bool rbq_probabilistic_refinement = true; - uint8_t rbq_query_bits = 4; - if (IsRaBitQ()) { - if (HasFilteredBits(bitset_)) { - return expected::Err(Status::not_implemented, - "DISKANN_RABITQ does not support bitset search"); - } - const auto* rabitq_conf = dynamic_cast(cfg.get()); - if (rabitq_conf == nullptr || rabitq_store_ == nullptr) { - return expected::Err(Status::invalid_args, - "DISKANN_RABITQ config or sidecar is not initialized"); - } - rbq_probabilistic_refinement = rabitq_conf->rbq_refine_mode.value() == "probabilistic"; - rbq_query_bits = static_cast(rabitq_conf->rbq_bits_query.value_or(4)); - } - feder::diskann::FederResultUniq feder_result; if (search_conf.trace_visit.value()) { if (nq != 1) { @@ -1053,12 +1029,10 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptrpush([&, index = row, p_id_ptr = p_id.get(), p_dist_ptr = p_dist.get()]() { knowhere::checkCancellation(op_context); auto& stats = query_stats[index]; - auto approx_distance_computer = IsRaBitQ() ? rabitq_store_->CreateDistanceComputer( - rbq_probabilistic_refinement, rbq_query_bits) - : nullptr; + auto navigation = navigation_store_ ? navigation_store_->CreateDistanceComputer(search_conf) : nullptr; pq_flash_index_->cached_beam_search(xq + (index * dim), k, lsearch, p_id_ptr + (index * k), p_dist_ptr + (index * k), beamwidth, false, &stats, feder_result, - bitset_, filter_ratio, approx_distance_computer.get()); + bitset_, filter_ratio, navigation.get()); #ifdef NOT_COMPILE_FOR_SWIG knowhere_diskann_search_hops.Observe(stats.n_hops); #endif @@ -1069,7 +1043,7 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptr::Err(Status::diskann_inner_error, "some search failed"); } - if (IsRaBitQ()) { + if (navigation_store_) { uint64_t estimates = 0; uint64_t refinements = 0; uint64_t pruned = 0; @@ -1079,7 +1053,7 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptr(pruned) / estimates; - LOG_KNOWHERE_DEBUG_ << "DiskANN RaBitQ refinement stats: queries=" << nq << ", estimates=" << estimates + LOG_KNOWHERE_DEBUG_ << "DiskANN navigation refinement stats: queries=" << nq << ", estimates=" << estimates << ", full_distances=" << refinements << ", pruned=" << pruned << ", prune_ratio=" << prune_ratio; } @@ -1285,12 +1259,6 @@ class DiskANNRaBitQIndexNode : public DiskANNIndexNode { Type() const override { return knowhere::IndexEnum::INDEX_DISKANN_RABITQ; } - - protected: - bool - IsRaBitQ() const override { - return true; - } }; #ifdef KNOWHERE_WITH_CARDINAL diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index f791eeac4..2a95b1690 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -204,13 +204,22 @@ class DiskANNConfig : public BaseConfig { } }; -class DiskANNRaBitQConfig : public DiskANNConfig { +// Codec selection and codec-specific knobs are validated at this boundary; +// the disk graph searcher receives only a query-local distance computer. +class DiskANNNavigationConfig : public DiskANNConfig { public: + CFG_STRING navigation_codec; CFG_INT rbq_bits; CFG_INT rbq_bits_query; CFG_STRING rbq_refine_mode; - KNOWHERE_DECLARE_CONFIG(DiskANNRaBitQConfig) { + KNOWHERE_DECLARE_CONFIG(DiskANNNavigationConfig) { + KNOWHERE_CONFIG_DECLARE_FIELD(navigation_codec) + .description("resident navigation codec: PQ or RABITQ") + .set_default("PQ") + .for_train() + .for_deserialize() + .for_static(); KNOWHERE_CONFIG_DECLARE_FIELD(rbq_bits) .description("number of RaBitQ bits per database vector dimension") .set_default(1) @@ -223,8 +232,9 @@ class DiskANNRaBitQConfig : public DiskANNConfig { .set_range(0, 8) .for_search(); KNOWHERE_CONFIG_DECLARE_FIELD(rbq_refine_mode) - .description("RaBitQ multi-bit refinement mode: probabilistic enables error-window pruning; full always " - "computes the complete RaBitQ distance") + .description( + "RaBitQ multi-bit refinement mode: probabilistic enables error-window pruning; full always " + "computes the complete RaBitQ distance") .set_default("probabilistic") .for_search() .for_range_search() @@ -237,6 +247,13 @@ class DiskANNRaBitQConfig : public DiskANNConfig { if (base_status != Status::success) { return base_status; } + const auto codec = navigation_codec.value_or("PQ"); + if (codec == "PQ") { + return Status::success; + } + if (codec != "RABITQ") { + return HandleError(err_msg, "unsupported DiskANN navigation codec", Status::invalid_args); + } const auto metric = metric_type.value_or(knowhere::metric::L2); if (metric != knowhere::metric::L2 && metric != knowhere::metric::IP) { return HandleError(err_msg, "DISKANN_RABITQ supports L2 and IP", Status::invalid_metric_type); @@ -252,11 +269,27 @@ class DiskANNRaBitQConfig : public DiskANNConfig { if (disk_pq_dims.value_or(0) != 0) { return HandleError(err_msg, "DISKANN_RABITQ requires disk_pq_dims=0", Status::invalid_args); } - if (warm_up.value_or(false)) { - return HandleError(err_msg, "DISKANN_RABITQ does not support warm_up", Status::invalid_args); - } return Status::success; } }; + +class DiskANNRaBitQConfig : public DiskANNNavigationConfig { + public: + KNOWHERE_DECLARE_CONFIG(DiskANNRaBitQConfig) { + KNOWHERE_CONFIG_DECLARE_FIELD(navigation_codec) + .description("DISKANN_RABITQ fixes the navigation codec to RABITQ") + .set_default("RABITQ") + .for_train() + .for_deserialize() + .for_static(); + } + Status + CheckAndAdjust(PARAM_TYPE type, std::string* error) override { + if (navigation_codec.value_or("RABITQ") != "RABITQ") { + return HandleError(error, "DISKANN_RABITQ requires navigation_codec=RABITQ", Status::invalid_args); + } + return DiskANNNavigationConfig::CheckAndAdjust(type, error); + } +}; } // namespace knowhere #endif /* DISKANN_CONFIG_H */ diff --git a/src/index/diskann/navigation_store.cc b/src/index/diskann/navigation_store.cc new file mode 100644 index 000000000..5cfc40755 --- /dev/null +++ b/src/index/diskann/navigation_store.cc @@ -0,0 +1,51 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#include "index/diskann/navigation_store.h" + +#include + +#include "index/diskann/diskann_config.h" +#include "index/diskann/rabitq_store.h" + +namespace knowhere { +namespace { +const DiskANNNavigationConfig* +ExternalConfig(const DiskANNConfig& config) { + const auto* navigation = dynamic_cast(&config); + if (!navigation || navigation->navigation_codec.value_or("PQ") == "PQ") { + return nullptr; + } + if (navigation->navigation_codec.value() != "RABITQ") { + throw std::invalid_argument("unsupported DiskANN navigation codec"); + } + return navigation; +} +} // namespace + +bool +UsesExternalNavigation(const DiskANNConfig& config) { + return ExternalConfig(config) != nullptr; +} + +std::vector +NavigationFiles(const DiskANNConfig& config, const std::string& prefix) { + return ExternalConfig(config) ? std::vector{RaBitQStore::SidecarFilename(prefix)} + : std::vector{}; +} + +void +BuildNavigationStore(const DiskANNConfig& config, const std::string& source, const std::string& prefix) { + if (const auto* navigation = ExternalConfig(config)) { + RaBitQStore::BuildFromFloatBin(source, RaBitQStore::SidecarFilename(prefix), + static_cast(navigation->rbq_bits.value_or(1))); + } +} + +std::unique_ptr +LoadNavigationStore(const DiskANNConfig& config, const std::string& prefix) { + if (ExternalConfig(config)) { + return std::make_unique(RaBitQStore::SidecarFilename(prefix)); + } + return nullptr; +} +} // namespace knowhere diff --git a/src/index/diskann/navigation_store.h b/src/index/diskann/navigation_store.h new file mode 100644 index 000000000..a0cbb307a --- /dev/null +++ b/src/index/diskann/navigation_store.h @@ -0,0 +1,40 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#include + +#include "diskann/pq_flash_index.h" + +namespace knowhere { +class DiskANNConfig; + +// Immutable after load. Each request owns its scorer and transformed-query +// scratch; this store must outlive those scorers. PQ remains the native engine +// path (a null external store), without a virtual call per candidate. +class NavigationStore { + public: + virtual ~NavigationStore() = default; + virtual int64_t + Count() const = 0; + virtual int64_t + Dimension() const = 0; + virtual size_t + MemorySize() const = 0; + virtual std::unique_ptr + CreateDistanceComputer(const DiskANNConfig& config) const = 0; +}; + +// Factory boundary: graph/IO code never inspects a concrete quantizer or its +// search parameters. New codecs are added here, not to cached_beam_search. +bool +UsesExternalNavigation(const DiskANNConfig& config); +std::vector +NavigationFiles(const DiskANNConfig& config, const std::string& prefix); +void +BuildNavigationStore(const DiskANNConfig& config, const std::string& prepared_source, const std::string& prefix); +std::unique_ptr +LoadNavigationStore(const DiskANNConfig& config, const std::string& prefix); +} // namespace knowhere diff --git a/src/index/diskann/rabitq_store.cc b/src/index/diskann/rabitq_store.cc index 7417baf3c..fc4308501 100644 --- a/src/index/diskann/rabitq_store.cc +++ b/src/index/diskann/rabitq_store.cc @@ -5,6 +5,14 @@ #include "index/diskann/rabitq_store.h" +#include +#include +#include +#include +#include +#include +#include + #include #include #include @@ -16,15 +24,8 @@ #include #include -#include -#include -#include -#include -#include -#include -#include - #include "diskann/utils.h" +#include "index/diskann/diskann_config.h" namespace knowhere { namespace { @@ -90,10 +91,10 @@ apply_rotation_single_query(const faiss::RandomRotationMatrix* rotation, const f } } -class RaBitQApproxDistanceComputer final : public diskann::ApproxDistanceComputer { +class RaBitQNavigationDistanceComputer final : public diskann::NavigationDistanceComputer { public: - RaBitQApproxDistanceComputer(const faiss::RandomRotationMatrix* rotation, const faiss::IndexRaBitQ* rabitq, - bool probabilistic_refinement, uint8_t query_bits) + RaBitQNavigationDistanceComputer(const faiss::RandomRotationMatrix* rotation, const faiss::IndexRaBitQ* rabitq, + bool probabilistic_refinement, uint8_t query_bits) : rotation_(rotation), rabitq_(rabitq), probabilistic_refinement_(probabilistic_refinement), @@ -135,30 +136,23 @@ class RaBitQApproxDistanceComputer final : public diskann::ApproxDistanceCompute std::array<_u64, 4> refine_positions{}; size_t pending_refinements = 0; const auto flush_refinements = [&]() { - distance_computer_->distances_batch_4( - refine_ids[0], refine_ids[1], refine_ids[2], refine_ids[3], - distances[refine_positions[0]], distances[refine_positions[1]], distances[refine_positions[2]], - distances[refine_positions[3]]); + distance_computer_->distances_batch_4(refine_ids[0], refine_ids[1], refine_ids[2], refine_ids[3], + distances[refine_positions[0]], distances[refine_positions[1]], + distances[refine_positions[2]], distances[refine_positions[3]]); pending_refinements = 0; }; - for (_u64 i = 0; i < n_ids; ++i) { - const uint8_t* code = rabitq_->codes.data() + static_cast(ids[i]) * rabitq_->code_size; - const float estimate = rabitq_distance_computer_->distance_to_code_1bit(code); - const size_t code_body_size = (static_cast(rabitq_->d) + 7) / 8; - const auto* factors = reinterpret_cast( - code + code_body_size); + const auto process_estimate = [&](_u64 i, const uint8_t* code, float estimate) { if (stats != nullptr) { ++stats->n_approx_estimates; } - if (!faiss::rabitq_utils::should_refine_candidate(estimate, factors->f_error, - rabitq_distance_computer_->g_error, threshold, false)) { + if (!rabitq_distance_computer_->should_refine(code, estimate, threshold, false)) { distances[i] = std::numeric_limits::infinity(); if (stats != nullptr) { ++stats->n_approx_pruned; ++stats->n_cmps_saved; } - continue; + return; } refine_ids[pending_refinements] = ids[i]; refine_positions[pending_refinements] = i; @@ -169,6 +163,25 @@ class RaBitQApproxDistanceComputer final : public diskann::ApproxDistanceCompute if (pending_refinements == 4) { flush_refinements(); } + }; + + // Batch independent estimates, then retain the original neighbor order + // and the caller's threshold snapshot when deciding which codes to refine. + _u64 i = 0; + for (; i + 4 <= n_ids; i += 4) { + std::array codes{}; + std::array estimates{}; + for (size_t j = 0; j < 4; ++j) { + codes[j] = rabitq_->codes.data() + static_cast(ids[i + j]) * rabitq_->code_size; + } + rabitq_distance_computer_->distance_to_code_1bit_batch_4(codes.data(), estimates.data()); + for (size_t j = 0; j < 4; ++j) { + process_estimate(i + j, codes[j], estimates[j]); + } + } + for (; i < n_ids; ++i) { + const uint8_t* code = rabitq_->codes.data() + static_cast(ids[i]) * rabitq_->code_size; + process_estimate(i, code, rabitq_distance_computer_->distance_to_code_1bit(code)); } for (size_t i = 0; i < pending_refinements; ++i) { const auto id = refine_ids[i]; @@ -248,9 +261,8 @@ RaBitQStore::BuildFromFloatBin(const std::string& data_path, const std::string& ForEachFloatBinBlock(data_path, rows, dim, [&](const float* block, size_t block_rows) { // Preserve DiskANN's existing input-block boundaries while sharing // the same bounded storage population utility as HNSW. - faiss::cppcontrib::knowhere::rabitq_build::add_in_blocks( - *pretransform, static_cast(block_rows), block, - static_cast(BlockRows(dim))); + faiss::cppcontrib::knowhere::rabitq_build::add_in_blocks(*pretransform, static_cast(block_rows), + block, static_cast(BlockRows(dim))); }); if (pretransform->ntotal != static_cast(rows)) { throw std::runtime_error("RaBitQ sidecar point count mismatch after encoding"); @@ -282,8 +294,7 @@ RaBitQStore::Validate() { throw std::runtime_error("DiskANN RaBitQ sidecar must be an IndexPreTransform with one transform"); } rotation_ = dynamic_cast(pretransform_->chain[0]); - if (rotation_ == nullptr || !rotation_->is_trained || rotation_->d_in <= 0 || - rotation_->d_in != rotation_->d_out) { + if (rotation_ == nullptr || !rotation_->is_trained || rotation_->d_in <= 0 || rotation_->d_in != rotation_->d_out) { throw std::runtime_error("DiskANN RaBitQ sidecar has an invalid random rotation"); } const auto rotation_dim = static_cast(rotation_->d_in); @@ -296,17 +307,16 @@ RaBitQStore::Validate() { throw std::runtime_error("DiskANN RaBitQ sidecar has an invalid RaBitQ leaf"); } if (pretransform_->metric_type != faiss::METRIC_L2 || rabitq_->metric_type != faiss::METRIC_L2 || - rabitq_->rabitq.metric_type != faiss::METRIC_L2 || - pretransform_->d != rotation_->d_in || rabitq_->d != rotation_->d_out || - pretransform_->ntotal != rabitq_->ntotal) { + rabitq_->rabitq.metric_type != faiss::METRIC_L2 || pretransform_->d != rotation_->d_in || + rabitq_->d != rotation_->d_out || pretransform_->ntotal != rabitq_->ntotal) { throw std::runtime_error("DiskANN RaBitQ sidecar metadata is inconsistent"); } - if (rabitq_->ntotal < 0 || rabitq_->rabitq.nb_bits < 1 || rabitq_->rabitq.nb_bits > 9 || - rabitq_->qb > 8 || rabitq_->centered) { + if (rabitq_->ntotal < 0 || rabitq_->rabitq.nb_bits < 1 || rabitq_->rabitq.nb_bits > 9 || rabitq_->qb > 8 || + rabitq_->centered) { throw std::runtime_error("DiskANN RaBitQ sidecar quantizer metadata is inconsistent"); } - const auto expected_code_size = rabitq_->rabitq.compute_code_size( - static_cast(rabitq_->d), rabitq_->rabitq.nb_bits); + const auto expected_code_size = + rabitq_->rabitq.compute_code_size(static_cast(rabitq_->d), rabitq_->rabitq.nb_bits); const auto point_count = static_cast(rabitq_->ntotal); if (rabitq_->code_size != expected_code_size || rabitq_->rabitq.code_size != expected_code_size || expected_code_size == 0 || point_count > std::numeric_limits::max() / expected_code_size || @@ -316,12 +326,33 @@ RaBitQStore::Validate() { } } -std::unique_ptr +std::unique_ptr +RaBitQStore::CreateDistanceComputer(const DiskANNConfig& config) const { + const auto* navigation = dynamic_cast(&config); + if (!navigation) { + throw std::invalid_argument("RaBitQ navigation requires query configuration"); + } + const auto query_metric = config.metric_type.value_or(metric::L2); + if (query_metric != metric::L2 && query_metric != metric::IP) { + throw std::invalid_argument("RaBitQ navigation currently supports L2 and IP"); + } + const auto mode = navigation->rbq_refine_mode.value_or("probabilistic"); + if (mode != "probabilistic" && mode != "full") { + throw std::invalid_argument("invalid RaBitQ refinement mode"); + } + const auto qb = navigation->rbq_bits_query.value_or(4); + if (qb < 0 || qb > 8) { + throw std::invalid_argument("RaBitQ query bits must be in [0, 8]"); + } + return CreateDistanceComputer(mode == "probabilistic", static_cast(qb)); +} + +std::unique_ptr RaBitQStore::CreateDistanceComputer(bool probabilistic_refinement, uint8_t query_bits) const { if (query_bits > 8) { throw std::invalid_argument("RaBitQ query bits must be in [0, 8]"); } - return std::make_unique(rotation_, rabitq_, probabilistic_refinement, query_bits); + return std::make_unique(rotation_, rabitq_, probabilistic_refinement, query_bits); } int64_t diff --git a/src/index/diskann/rabitq_store.h b/src/index/diskann/rabitq_store.h index 6bc223fca..4a4f22653 100644 --- a/src/index/diskann/rabitq_store.h +++ b/src/index/diskann/rabitq_store.h @@ -10,7 +10,7 @@ #include #include -#include "diskann/pq_flash_index.h" +#include "index/diskann/navigation_store.h" namespace faiss { struct Index; @@ -21,7 +21,7 @@ struct RandomRotationMatrix; namespace knowhere { -class RaBitQStore { +class RaBitQStore final : public NavigationStore { public: static std::string SidecarFilename(const std::string& index_prefix); @@ -36,14 +36,17 @@ class RaBitQStore { RaBitQStore& operator=(const RaBitQStore&) = delete; - std::unique_ptr + std::unique_ptr CreateDistanceComputer(bool probabilistic_refinement, uint8_t query_bits = 4) const; + std::unique_ptr + CreateDistanceComputer(const DiskANNConfig& config) const override; + int64_t - Count() const; + Count() const override; int64_t - Dimension() const; + Dimension() const override; uint8_t Bits() const; @@ -52,7 +55,7 @@ class RaBitQStore { CodeSize() const; size_t - MemorySize() const; + MemorySize() const override; private: void diff --git a/tests/python/test_diskann_rabitq.py b/tests/python/test_diskann_rabitq.py new file mode 100644 index 000000000..98d7a58e2 --- /dev/null +++ b/tests/python/test_diskann_rabitq.py @@ -0,0 +1,44 @@ +# Copyright (C) 2026 Zilliz. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import json +import struct + +import knowhere +import numpy as np +import pytest + + +@pytest.mark.parametrize("metric", ["L2", "IP"]) +@pytest.mark.parametrize("kind,codec", [("DISKANN_RABITQ", None), ("DISKANN", "RABITQ"), ("DISKANN", None)]) +def test_navigation_roundtrip(tmp_path, metric, kind, codec): + rng = np.random.default_rng(42) + base = rng.normal(size=(1000, 64)).astype("float32") + query = rng.normal(size=(10, 64)).astype("float32") + source = tmp_path / "base.fbin" + with source.open("wb") as output: + output.write(struct.pack("= 0) and np.all(np.isfinite(distances)) + exact = ((query[:, None, :] - base[None, :, :]) ** 2).sum(2) if metric == "L2" else -(query @ base.T) + truth = np.argsort(exact, axis=1)[:, :10] + recall = sum(len(set(actual) & set(expected)) for actual, expected in zip(ids, truth)) / 100 + assert recall > 0.8 diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index 97fe5865f..e1f40cabe 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -97,7 +97,6 @@ constexpr float kIpRangeAp = 0.9; constexpr float kCosineRangeAp = 0.9; } // namespace - TEST_CASE("Valid diskann build params test", "[diskann]") { int rows_num = 1000000; auto version = GenTestVersionList(); @@ -511,21 +510,18 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { }; REQUIRE(knowhere::IndexFactory::Instance() - .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, - make_pack()) + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, make_pack()) .has_value()); REQUIRE_FALSE(knowhere::IndexFactory::Instance() - .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, - make_pack()) + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, make_pack()) .has_value()); REQUIRE_FALSE(knowhere::IndexFactory::Instance() - .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, - make_pack()) + .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, make_pack()) .has_value()); auto check_train_config = [&](knowhere::Json json, knowhere::Status expected) { - auto cfg = knowhere::IndexStaticFaced::CreateConfig( - knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); + auto cfg = knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, + version); std::string msg; REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::TRAIN, &msg) == expected); }; @@ -539,6 +535,22 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { {"search_cache_budget_gb_ratio", 0}, {"rbq_bits", 1}}; check_train_config(valid, knowhere::Status::success); + { + auto conflict = valid; + conflict["navigation_codec"] = "PQ"; + check_train_config(conflict, knowhere::Status::invalid_args); + for (const auto* codec : {"PQ", "RABITQ", "TQ", "invalid"}) { + auto cfg = + knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN, version); + auto request = valid; + request["navigation_codec"] = codec; + std::string error; + const auto status = knowhere::Config::Load(*cfg, request, knowhere::PARAM_TYPE::TRAIN, &error); + REQUIRE(status == ((std::string(codec) == "PQ" || std::string(codec) == "RABITQ") + ? knowhere::Status::success + : knowhere::Status::invalid_args)); + } + } valid["rbq_bits"] = 2; check_train_config(valid, knowhere::Status::success); valid["rbq_bits"] = 4; @@ -565,14 +577,11 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { check_train_config(invalid, knowhere::Status::success); auto check_search_mode = [&](const std::string& mode, knowhere::Status expected) { - auto cfg = knowhere::IndexStaticFaced::CreateConfig( - knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); - knowhere::Json json = {{"dim", kDim}, - {"metric_type", knowhere::metric::L2}, - {"k", kK}, - {"search_list_size", 128}, - {"beamwidth", 8}, - {"rbq_refine_mode", mode}}; + auto cfg = knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, + version); + knowhere::Json json = {{"dim", kDim}, {"metric_type", knowhere::metric::L2}, + {"k", kK}, {"search_list_size", 128}, + {"beamwidth", 8}, {"rbq_refine_mode", mode}}; std::string msg; REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::SEARCH, &msg) == expected); }; @@ -580,10 +589,13 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { check_search_mode("full", knowhere::Status::success); check_search_mode("invalid", knowhere::Status::invalid_args); for (const int qb : {-1, 0, 4, 8, 9}) { - auto cfg = knowhere::IndexStaticFaced::CreateConfig( - knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); - knowhere::Json json = {{"dim", kDim}, {"metric_type", knowhere::metric::L2}, - {"k", kK}, {"search_list_size", 128}, {"rbq_bits_query", qb}}; + auto cfg = knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, + version); + knowhere::Json json = {{"dim", kDim}, + {"metric_type", knowhere::metric::L2}, + {"k", kK}, + {"search_list_size", 128}, + {"rbq_bits_query", qb}}; std::string msg; REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::SEARCH, &msg) == ((qb >= 0 && qb <= 8) ? knowhere::Status::success : knowhere::Status::out_of_range_in_json)); @@ -621,8 +633,8 @@ TEST_CASE("Test DISKANN_RABITQ probabilistic refinement", "[diskann][rabitq]") { std::vector distances(rows); diskann::QueryStats full_stats; - distance_computer->compute_distances(ids.data(), rows, distances.data(), - std::numeric_limits::max(), false, &full_stats); + distance_computer->compute_distances(ids.data(), rows, distances.data(), std::numeric_limits::max(), false, + &full_stats); REQUIRE(full_stats.n_approx_estimates == 0); REQUIRE(full_stats.n_approx_refinements == rows); REQUIRE(full_stats.n_approx_pruned == 0); @@ -699,8 +711,8 @@ TEST_CASE("DiskANN RaBitQ shares Faiss codes and request-local query bits", "[di std::vector actual(7); std::vector expected(7); adapter->compute_distances(ids, 7, actual.data(), 0, false, nullptr); - native->distances_batch_4(ids[0], ids[1], ids[2], ids[3], - expected[0], expected[1], expected[2], expected[3]); + native->distances_batch_4(ids[0], ids[1], ids[2], ids[3], expected[0], expected[1], expected[2], + expected[3]); for (size_t i = 4; i < 7; ++i) { expected[i] = (*native)(ids[i]); } @@ -723,6 +735,73 @@ TEST_CASE("DiskANN RaBitQ shares Faiss codes and request-local query bits", "[di fs::remove_all(dir); } +TEST_CASE("DiskANN RaBitQ batches preserve scalar decisions", "[diskann][rabitq][batch]") { + const auto dir = kDir + "/rabitq_batch"; + REQUIRE_NOTHROW(fs::create_directories(dir)); + for (const uint32_t dim : {33U, 128U, 769U}) { + auto data = GenDataSet(19, dim, 91); + const auto* x = static_cast(data->GetTensor()); + const auto raw = dir + "/base.fbin"; + const auto path = dir + "/model.index"; + WriteRawDataToDisk(raw, x, 19, dim); + for (uint8_t bits = 1; bits <= 9; ++bits) { + knowhere::RaBitQStore::BuildFromFloatBin(raw, path, bits); + knowhere::RaBitQStore store(path); + std::array ids{}; + std::iota(ids.begin(), ids.end(), 0); + for (uint8_t qb = 0; qb <= 8; ++qb) { + CAPTURE(dim, bits, qb); + auto scalar = store.CreateDistanceComputer(true, qb); + auto batch = store.CreateDistanceComputer(true, qb); + scalar->set_query(x + 18 * dim); + batch->set_query(x + 18 * dim); + std::array full{}; + batch->compute_distances(ids.data(), ids.size(), full.data(), 0, false, nullptr); + auto sorted = full; + std::sort(sorted.begin(), sorted.end()); + for (const float threshold : {0.0f, sorted[4], sorted[8], sorted[16] * 2}) { + for (const size_t count : {0, 1, 3, 4, 5, 7, 16, 17}) { + CAPTURE(threshold, count); + std::array actual{}, expected{}; + diskann::QueryStats scalar_stats{}, batch_stats{}; + for (size_t i = 0; i < count; ++i) { + scalar->compute_distances(ids.data() + i, 1, expected.data() + i, threshold, true, + &scalar_stats); + } + batch->compute_distances(ids.data(), count, actual.data(), threshold, true, &batch_stats); + REQUIRE(scalar_stats.n_approx_estimates == batch_stats.n_approx_estimates); + REQUIRE(scalar_stats.n_approx_pruned == batch_stats.n_approx_pruned); + REQUIRE(scalar_stats.n_approx_refinements == batch_stats.n_approx_refinements); + for (size_t i = 0; i < count; ++i) { + if (std::isinf(expected[i])) { + REQUIRE(actual[i] == expected[i]); + } else { + REQUIRE(actual[i] == Catch::Approx(expected[i]).epsilon(1e-5).margin(1e-3)); + } + } + } + } + } + // Concurrent callers have independent query transforms and query bits. + std::array, 2> sequential, concurrent; + const auto run_query = [&](size_t q, std::vector& result) { + auto dc = store.CreateDistanceComputer(false, q ? 8 : 0); + dc->set_query(x + (17 + q) * dim); + result.resize(ids.size()); + dc->compute_distances(ids.data(), ids.size(), result.data(), 0, false, nullptr); + }; + run_query(0, sequential[0]); + run_query(1, sequential[1]); + std::thread first([&]() { run_query(0, concurrent[0]); }); + std::thread second([&]() { run_query(1, concurrent[1]); }); + first.join(); + second.join(); + REQUIRE(sequential == concurrent); + } + } + fs::remove_all(dir); +} + TEST_CASE("Test DISKANN_RABITQ rejects inconsistent sidecars", "[diskann][rabitq][sidecar]") { const auto sidecar_dir = kDir + "/rabitq_sidecar_validation"; const auto data_path = sidecar_dir + "/base.fbin"; @@ -750,9 +829,8 @@ TEST_CASE("Test DISKANN_RABITQ rejects inconsistent sidecars", "[diskann][rabitq REQUIRE_THROWS(knowhere::RaBitQStore(corrupted_path)); }; - corrupt_and_check("wrong_outer_metric", [](auto& pretransform, auto&, auto&) { - pretransform.metric_type = faiss::METRIC_INNER_PRODUCT; - }); + corrupt_and_check("wrong_outer_metric", + [](auto& pretransform, auto&, auto&) { pretransform.metric_type = faiss::METRIC_INNER_PRODUCT; }); fs::remove_all(sidecar_dir); } @@ -791,6 +869,7 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { {"search_cache_budget_gb_ratio", 0}, {"warm_up", false}}; if (rbq_bits == 4) { + deserialize_json["warm_up"] = true; deserialize_json["search_cache_budget_gb"] = 0.00005; // RaBitQ must safely force BFS even when the default sample-query // cache mode is requested, because its navigation PQ is not resident. @@ -833,17 +912,76 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { std::array query{}; std::array ids{}; std::array distances{}; - REQUIRE_THROWS(metadata_only_index.cached_beam_search(query.data(), 1, 1, ids.data(), distances.data(), - 1)); + REQUIRE_THROWS(metadata_only_index.cached_beam_search(query.data(), 1, 1, ids.data(), distances.data(), 1)); + + struct FailingNavigation : diskann::NavigationDistanceComputer { + bool fail_set = true; + void + set_query(const float*) override { + if (fail_set) + throw std::runtime_error("query preparation failure"); + } + void + compute_distances(const unsigned*, _u64, float*, float, bool, diskann::QueryStats*) override { + throw std::runtime_error("navigation scoring failure"); + } + } failing; + knowhere::RaBitQStore store(rabitq_prefix + "_rabitq.index"); + auto scorer = store.CreateDistanceComputer(false, 0); + std::copy_n(static_cast(query_ds->GetTensor()), kDim, query.data()); + for (bool fail_set : {true, false}) { + failing.fail_set = fail_set; + REQUIRE_THROWS(metadata_only_index.cached_beam_search(query.data(), 1, 16, ids.data(), distances.data(), + 1, false, nullptr, nullptr, {}, -1, &failing)); + // The one-slot scratch pool must still be usable after a throw. + REQUIRE_NOTHROW(metadata_only_index.cached_beam_search(query.data(), 1, 16, ids.data(), + distances.data(), 1, false, nullptr, nullptr, {}, + -1, scorer.get())); + REQUIRE(ids[0] >= 0); + } } auto index = knowhere::IndexFactory::Instance() .Create(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version, pack) .value(); + if (rbq_bits == 4) { + const auto sidecar = rabitq_prefix + "_rabitq.index"; + const auto saved = sidecar + ".saved"; + fs::rename(sidecar, saved); + { + std::ofstream corrupt(sidecar, std::ios::binary); + corrupt << "invalid sidecar"; + } + const auto failed = index.Deserialize(binset, deserialize_json); + fs::rename(saved, sidecar); + REQUIRE(failed == knowhere::Status::invalid_index_error); + REQUIRE(index.Count() == 0); + REQUIRE_FALSE(index.Search(query_ds, search_json, nullptr).has_value()); + } REQUIRE(index.Deserialize(binset, deserialize_json) == knowhere::Status::success); REQUIRE(index.Type() == knowhere::IndexEnum::INDEX_DISKANN_RABITQ); auto result = index.Search(query_ds, search_json, nullptr); REQUIRE(result.has_value()); + if (rbq_bits == 4) { + auto generic = knowhere::IndexFactory::Instance() + .Create(knowhere::IndexEnum::INDEX_DISKANN, version, pack) + .value(); + auto generic_load = deserialize_json; + generic_load["navigation_codec"] = "RABITQ"; + REQUIRE(generic.Deserialize(binset, generic_load) == knowhere::Status::success); + for (const int qb : {0, 4, 8}) { + auto request = search_json; + request["rbq_bits_query"] = qb; + auto expected = index.Search(query_ds, request, nullptr); + auto actual = generic.Search(query_ds, request, nullptr); + REQUIRE(expected.has_value()); + REQUIRE(actual.has_value()); + for (uint32_t i = 0; i < kNumQueries * kK; ++i) { + REQUIRE(actual.value()->GetIds()[i] == expected.value()->GetIds()[i]); + REQUIRE(actual.value()->GetDistance()[i] == expected.value()->GetDistance()[i]); + } + } + } search_json["rbq_refine_mode"] = "full"; auto full_result = index.Search(query_ds, search_json, nullptr); REQUIRE(full_result.has_value()); @@ -858,10 +996,21 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { index.Search(query_ds, search_json, knowhere::BitsetView(empty_bitset_data.data(), kNumRows)); REQUIRE(empty_bitset_result.has_value()); - auto bitset_data = GenerateBitsetWithFirstTbitsSet(kNumRows, 1); - auto bitset_result = index.Search(query_ds, search_json, knowhere::BitsetView(bitset_data.data(), kNumRows)); - REQUIRE_FALSE(bitset_result.has_value()); - REQUIRE(bitset_result.error() == knowhere::Status::not_implemented); + for (const int excluded : {1, 950, 1000}) { + auto bitset_data = GenerateBitsetWithFirstTbitsSet(kNumRows, excluded); + auto filter = knowhere::BitsetView(bitset_data.data(), kNumRows); + auto bitset_result = index.Search(query_ds, search_json, filter); + REQUIRE(bitset_result.has_value()); + for (uint32_t i = 0; i < kNumQueries * kK; ++i) { + const auto id = bitset_result.value()->GetIds()[i]; + if (excluded == 1000) + REQUIRE(id == -1); + else { + REQUIRE(id >= excluded); + REQUIRE(id < kNumRows); + } + } + } auto iterators = index.AnnIterator(query_ds, search_json, nullptr); REQUIRE_FALSE(iterators.has_value()); REQUIRE(iterators.error() == knowhere::Status::not_implemented); @@ -899,11 +1048,8 @@ TEST_CASE("Test DISKANN_RABITQ inner product d+1 sidecar", "[diskann][rabitq][ip {"search_cache_budget_gb", 0}, {"search_cache_budget_gb_ratio", 0}, {"warm_up", false}}; - knowhere::Json search_json = {{"dim", kDim}, - {"metric_type", knowhere::metric::IP}, - {"k", kK}, - {"search_list_size", 128}, - {"beamwidth", 8}}; + knowhere::Json search_json = { + {"dim", kDim}, {"metric_type", knowhere::metric::IP}, {"k", kK}, {"search_list_size", 128}, {"beamwidth", 8}}; auto file_manager = std::make_shared(); auto pack = knowhere::Pack(std::shared_ptr(file_manager)); @@ -949,8 +1095,7 @@ TEST_CASE("Test AiSAQ clamps navigation PQ for deserialize", "[diskann][aisaq][l {"data_path", data_path}, {"max_degree", 16}, {"search_list_size", 32}, - {"pq_code_budget_gb", - static_cast(rows) * dim / (1024.0 * 1024.0 * 1024.0)}, + {"pq_code_budget_gb", static_cast(rows) * dim / (1024.0 * 1024.0 * 1024.0)}, {"build_dram_budget_gb", 1.0}, {"disk_pq_dims", 0}, {"search_cache_budget_gb", 0}, diff --git a/thirdparty/DiskANN/include/diskann/pq_flash_index.h b/thirdparty/DiskANN/include/diskann/pq_flash_index.h index 80c1c1560..d4f885d2a 100644 --- a/thirdparty/DiskANN/include/diskann/pq_flash_index.h +++ b/thirdparty/DiskANN/include/diskann/pq_flash_index.h @@ -183,13 +183,16 @@ namespace diskann { virtual ~PQDataGetter() {} }; - // Optional query-local navigation distance implementation. Ordinary - // DiskANN leaves this null and uses its resident PQ codes. Knowhere's - // DISKANN_RABITQ supplies one instance per query because the distance - // computer owns transformed-query scratch and is not thread-safe. - class ApproxDistanceComputer { + // Optional query-local navigation scorer. Null preserves resident PQ. + // Query inputs and thresholds are in DiskANN prepared space, with smaller + // scores preferred. A valid threshold is a snapshot of the full candidate + // pool, shared by this batch; the scorer must not modify search state. + // A scorer may return +infinity to explicitly reject a candidate under its + // configured approximate policy. That value is not an exact distance. + // Each scorer owns its query scratch and is not shared across queries. + class NavigationDistanceComputer { public: - virtual ~ApproxDistanceComputer() = default; + virtual ~NavigationDistanceComputer() = default; virtual void set_query(const float* query) = 0; virtual void compute_distances(const unsigned* ids, _u64 n_ids, float* distances, float threshold, @@ -227,7 +230,7 @@ namespace diskann { const knowhere::feder::diskann::FederResultUniq &feder = nullptr, knowhere::BitsetView bitset_view = nullptr, const float filter_ratio = -1.0f, - ApproxDistanceComputer* approx_distance_computer = nullptr); + NavigationDistanceComputer* approx_distance_computer = nullptr); void calc_dist_by_ids(const T *query, const int64_t *ids, const int64_t n, float *const output_dists); @@ -311,7 +314,7 @@ namespace diskann { const knowhere::feder::diskann::FederResultUniq &feder, knowhere::BitsetView bitset_view, PQDataGetter* pq_data_getter, - ApproxDistanceComputer* approx_distance_computer = nullptr); + NavigationDistanceComputer* approx_distance_computer = nullptr); // Assign the index of ids to its corresponding sector and if it is in // cache, write to the output_data diff --git a/thirdparty/DiskANN/src/pq_flash_index.cpp b/thirdparty/DiskANN/src/pq_flash_index.cpp index 083589ddd..0a5003692 100644 --- a/thirdparty/DiskANN/src/pq_flash_index.cpp +++ b/thirdparty/DiskANN/src/pq_flash_index.cpp @@ -970,8 +970,11 @@ namespace diskann { const knowhere::feder::diskann::FederResultUniq &feder, knowhere::BitsetView bitset_view, PQDataGetter* pq_data_getter, - ApproxDistanceComputer* approx_distance_computer) { - if (approx_distance_computer == nullptr && this->data == nullptr) { + NavigationDistanceComputer* approx_distance_computer) { + // AiSAQ supplies its own disk-backed PQDataGetter; resident codes are + // required only when this index is itself the PQ provider. + if (approx_distance_computer == nullptr && pq_data_getter == this && + this->data == nullptr) { throw ANNException( "resident navigation PQ data is unavailable and no external " "distance computer was supplied", @@ -1130,7 +1133,7 @@ namespace diskann { float *distances, const _u64 beam_width, const bool use_reorder_data, QueryStats *stats, const knowhere::feder::diskann::FederResultUniq &feder, knowhere::BitsetView bitset_view, const float filter_ratio_in, - ApproxDistanceComputer* approx_distance_computer) { + NavigationDistanceComputer* approx_distance_computer) { if (approx_distance_computer == nullptr && this->data == nullptr) { throw ANNException( "resident navigation PQ data is unavailable and no external " @@ -1146,15 +1149,24 @@ namespace diskann { this->thread_data.wait_for_push_notify(); data = this->thread_data.pop(); } + // Return query resources on every exit, including exceptions from an + // external scorer. Otherwise a failed query permanently consumes a slot. + const auto return_scratch = [this](ThreadData* slot) { + this->thread_data.push(*slot); + this->thread_data.push_notify_all(); + }; + std::unique_ptr, decltype(return_scratch)> scratch_guard(&data, return_scratch); auto query_norm_opt = init_thread_data(data, query1); if (!query_norm_opt.has_value()) { // return an empty answer when calcu a zero point - this->thread_data.push(data); - this->thread_data.push_notify_all(); return; } float query_norm = query_norm_opt.value(); - auto ctx = this->reader->get_ctx(); + auto ctx = this->reader->get_ctx(); + const auto return_context = [this](decltype(ctx)* context) { + this->reader->put_ctx(*context); + }; + std::unique_ptr context_guard(&ctx, return_context); if (approx_distance_computer != nullptr) { approx_distance_computer->set_query( @@ -1183,9 +1195,6 @@ namespace diskann { // like on every other exit from this function. Leaking them here // permanently shrinks the thread_data pool: after max_nthreads such // queries every subsequent search blocks in wait_for_push_notify(). - this->thread_data.push(data); - this->thread_data.push_notify_all(); - this->reader->put_ctx(ctx); return; } @@ -1193,9 +1202,6 @@ namespace diskann { brute_force_beam_search(data, query_norm, k_search, indices, distances, beam_width, ctx, stats, feder, bitset_view, this, approx_distance_computer); - this->thread_data.push(data); - this->thread_data.push_notify_all(); - this->reader->put_ctx(ctx); return; } } @@ -1205,9 +1211,6 @@ namespace diskann { brute_force_beam_search(data, query_norm, k_search, indices, distances, beam_width, ctx, stats, feder, bitset_view, this, approx_distance_computer); - this->thread_data.push(data); - this->thread_data.push_notify_all(); - this->reader->put_ctx(ctx); return; } @@ -1595,9 +1598,6 @@ namespace diskann { } } - this->thread_data.push(data); - this->thread_data.push_notify_all(); - this->reader->put_ctx(ctx); if (stats != nullptr) { stats->total_us = (double) query_timer.elapsed(); From 2c808791394c60a64d3e9ce77f996e0d50416140 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Fri, 18 Sep 2026 09:14:39 +0000 Subject: [PATCH 03/12] fix: harden DiskANN RaBitQ integration for upstream validation Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 24 ++++ src/index/diskann/diskann_config.h | 10 +- src/index/diskann/navigation_store.cc | 14 +++ src/index/diskann/navigation_store.h | 2 + src/index/diskann/rabitq_store.cc | 17 +++ src/index/diskann/rabitq_store.h | 4 + tests/ut/test_diskann.cc | 112 +++++++++++++++--- .../DiskANN/include/diskann/aux_utils.h | 2 - thirdparty/DiskANN/src/aux_utils.cpp | 24 +--- thirdparty/faiss/faiss/impl/io.cpp | 8 ++ 10 files changed, 173 insertions(+), 44 deletions(-) diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 6d3e726ac..722a5aede 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -13,6 +13,7 @@ #include +#include #include #include #include @@ -157,6 +158,29 @@ class DiskANNIndexNode : public IndexNode { static expected StaticEstimateLoadResource(const uint64_t file_size_in_bytes, const int64_t num_rows, const int64_t dim, const knowhere::BaseConfig& config, const IndexVersion& version) { + const auto& disk_config = static_cast(config); + if (UsesExternalNavigation(disk_config)) { + try { + const auto navigation_bytes = EstimateNavigationMemory(disk_config, num_rows, dim); + const long double raw_bytes = static_cast(num_rows) * dim * sizeof(float); + const long double cache_bytes = std::max( + static_cast(disk_config.search_cache_budget_gb.value_or(0)) * (1ULL << 30), + static_cast(disk_config.search_cache_budget_gb_ratio.value_or(0)) * raw_bytes); + // Keep the legacy engine allowance for scratch, PQ tables and + // other loading overhead. Add the codec's persistent storage + // and cache budget explicitly. This is a conservative estimate, + // not an exact RSS prediction or a replacement for Size(). + const long double memory_bytes = + file_size_in_bytes / 4 + static_cast(navigation_bytes) + std::ceil(cache_bytes); + if (!std::isfinite(memory_bytes) || memory_bytes < 0 || + memory_bytes >= static_cast(std::numeric_limits::max())) { + return expected::Err(Status::invalid_args, "DiskANN resource estimate overflows"); + } + return Resource{.memoryCost = static_cast(memory_bytes), .diskCost = file_size_in_bytes}; + } catch (const std::exception& e) { + return expected::Err(Status::invalid_args, e.what()); + } + } return Resource{.memoryCost = file_size_in_bytes / 4, .diskCost = file_size_in_bytes}; } diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index 2a95b1690..58e7277d1 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -129,13 +129,15 @@ class DiskANNConfig : public BaseConfig { .set_default(0) .set_range(0, std::numeric_limits::max()) .for_train() - .for_deserialize(); + .for_deserialize() + .for_static(); KNOWHERE_CONFIG_DECLARE_FIELD(search_cache_budget_gb) .description("the size of cached nodes in GB.") .set_default(0) .set_range(0, std::numeric_limits::max()) .for_train() - .for_deserialize(); + .for_deserialize() + .for_static(); KNOWHERE_CONFIG_DECLARE_FIELD(warm_up) .description("should do warm up before search.") .set_default(false) @@ -236,9 +238,7 @@ class DiskANNNavigationConfig : public DiskANNConfig { "RaBitQ multi-bit refinement mode: probabilistic enables error-window pruning; full always " "computes the complete RaBitQ distance") .set_default("probabilistic") - .for_search() - .for_range_search() - .for_iterator(); + .for_search(); } Status diff --git a/src/index/diskann/navigation_store.cc b/src/index/diskann/navigation_store.cc index 5cfc40755..dbcf30bb0 100644 --- a/src/index/diskann/navigation_store.cc +++ b/src/index/diskann/navigation_store.cc @@ -2,6 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 #include "index/diskann/navigation_store.h" +#include #include #include "index/diskann/diskann_config.h" @@ -27,6 +28,19 @@ UsesExternalNavigation(const DiskANNConfig& config) { return ExternalConfig(config) != nullptr; } +uint64_t +EstimateNavigationMemory(const DiskANNConfig& config, int64_t rows, int64_t dim) { + const auto* navigation = ExternalConfig(config); + if (!navigation) { + return 0; + } + if (dim <= 0 || dim >= std::numeric_limits::max()) { + throw std::invalid_argument("invalid DiskANN navigation dimension"); + } + const auto prepared_dim = dim + (config.metric_type.value_or(metric::L2) == metric::IP ? 1 : 0); + return RaBitQStore::EstimateMemorySize(rows, prepared_dim, static_cast(navigation->rbq_bits.value_or(1))); +} + std::vector NavigationFiles(const DiskANNConfig& config, const std::string& prefix) { return ExternalConfig(config) ? std::vector{RaBitQStore::SidecarFilename(prefix)} diff --git a/src/index/diskann/navigation_store.h b/src/index/diskann/navigation_store.h index a0cbb307a..33c3be1fd 100644 --- a/src/index/diskann/navigation_store.h +++ b/src/index/diskann/navigation_store.h @@ -31,6 +31,8 @@ class NavigationStore { // search parameters. New codecs are added here, not to cached_beam_search. bool UsesExternalNavigation(const DiskANNConfig& config); +uint64_t +EstimateNavigationMemory(const DiskANNConfig& config, int64_t rows, int64_t dim); std::vector NavigationFiles(const DiskANNConfig& config, const std::string& prefix); void diff --git a/src/index/diskann/rabitq_store.cc b/src/index/diskann/rabitq_store.cc index fc4308501..47458c852 100644 --- a/src/index/diskann/rabitq_store.cc +++ b/src/index/diskann/rabitq_store.cc @@ -381,4 +381,21 @@ RaBitQStore::MemorySize() const { rotation_->A.size() * sizeof(float) + rotation_->b.size() * sizeof(float); } +uint64_t +RaBitQStore::EstimateMemorySize(int64_t rows, int64_t prepared_dim, uint8_t rbq_bits) { + if (rows <= 0 || prepared_dim <= 0 || prepared_dim > std::numeric_limits::max() || rbq_bits < 1 || + rbq_bits > 9) { + throw std::invalid_argument("invalid RaBitQ resource estimate dimensions or bits"); + } + const uint64_t d = static_cast(prepared_dim); + const uint64_t code_size = faiss::RaBitQuantizer().compute_code_size(d, rbq_bits); + // d is bounded by the transform's int dimension; d*(d+1)*sizeof(float) + // fits uint64_t. Includes the square rotation and the global centroid. + const uint64_t model_bytes = d * (d + 1) * sizeof(float); + if (static_cast(rows) > (std::numeric_limits::max() - model_bytes) / code_size) { + throw std::overflow_error("RaBitQ resource estimate overflows uint64_t"); + } + return static_cast(rows) * code_size + model_bytes; +} + } // namespace knowhere diff --git a/src/index/diskann/rabitq_store.h b/src/index/diskann/rabitq_store.h index 4a4f22653..bcdc16319 100644 --- a/src/index/diskann/rabitq_store.h +++ b/src/index/diskann/rabitq_store.h @@ -29,6 +29,10 @@ class RaBitQStore final : public NavigationStore { static void BuildFromFloatBin(const std::string& data_path, const std::string& sidecar_path, uint8_t rbq_bits); + // Persistent codes and model only; excludes graph-engine scratch and node cache. + static uint64_t + EstimateMemorySize(int64_t rows, int64_t prepared_dim, uint8_t rbq_bits); + explicit RaBitQStore(const std::string& sidecar_path); ~RaBitQStore(); diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index e1f40cabe..98c55e6da 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -34,6 +34,7 @@ #include "faiss/VectorTransform.h" #include "faiss/cppcontrib/knowhere/index_io.h" #include "faiss/impl/RaBitQUtils.h" +#include "faiss/impl/io.h" #include "faiss/index_io.h" #include "faiss/utils/rabitq_simd.h" #include "filemanager/FileManager.h" @@ -81,6 +82,11 @@ constexpr uint32_t kNumQueries = 10; constexpr uint32_t kDim = 128; constexpr uint32_t kLargeDim = 256; constexpr uint32_t kK = 10; +#ifdef KNOWHERE_WITH_CARDINAL +constexpr const char* kNativeDiskANN = "DISKANN_DEPRECATED"; +#else +constexpr const char* kNativeDiskANN = "DISKANN"; +#endif #ifdef KNOWHERE_WITH_CUVS // This test is expected to run on a GPU with at least 2 GB of RAM constexpr uint32_t defaultMaxDegree = 64; @@ -144,17 +150,6 @@ TEST_CASE("Valid diskann build params test", "[diskann]") { } } -TEST_CASE("DiskANN navigation PQ can exceed 512 chunks", "[diskann]") { - constexpr size_t rows = 500000; - constexpr size_t dim = 1537; - constexpr size_t matched_code_bytes = 790; - - REQUIRE(diskann::get_num_pq_chunks(static_cast(rows * matched_code_bytes), rows, dim) == - matched_code_bytes); - REQUIRE(diskann::get_num_pq_chunks(static_cast(rows * (dim + 1)), rows, dim) == dim); - REQUIRE(diskann::get_num_pq_chunks(0.0, rows, dim) == 1); -} - TEST_CASE("Invalid diskann params test", "[diskann]") { fs::remove_all(kDir); fs::remove(kDir); @@ -502,6 +497,80 @@ TEST_CASE("Test DiskANN CalcDistByIDs with all vectors cached", "[diskann]") { fs::remove_all(kDir); } +TEST_CASE("Faiss sidecar file IO accepts empty arrays", "[diskann][rabitq][sidecar]") { + const auto dir = kDir + "/empty_file_io"; + REQUIRE_NOTHROW(fs::create_directories(dir)); + const auto path = dir + "/empty.bin"; + const uint32_t expected = 123; + { + faiss::FileIOWriter writer(path.c_str()); + REQUIRE(writer(nullptr, sizeof(float), 0) == 0); + REQUIRE(writer(nullptr, 0, 1) == 0); + REQUIRE(writer(&expected, sizeof(expected), 1) == 1); + } + REQUIRE(fs::file_size(path) == sizeof(expected)); + { + faiss::FileIOReader reader(path.c_str()); + REQUIRE(reader(nullptr, sizeof(float), 0) == 0); + REQUIRE(reader(nullptr, 0, 1) == 0); + uint32_t actual = 0; + REQUIRE(reader(&actual, sizeof(actual), 1) == 1); + REQUIRE(actual == expected); + } + fs::remove_all(dir); +} + +TEST_CASE("DiskANN navigation load resource estimate", "[diskann][rabitq][resource]") { + using Static = knowhere::IndexStaticFaced; + const auto version = GenTestVersionList(); + constexpr uint64_t file_bytes = 1ULL << 30; + constexpr int64_t rows = 1000000; + constexpr int64_t dim = 769; + for (const auto* metric : {"L2", "IP"}) { + const int64_t prepared_dim = dim + (std::string(metric) == "IP" ? 1 : 0); + uint64_t previous = 0; + for (int bits = 1; bits <= 9; ++bits) { + knowhere::Json config = {{"metric_type", metric}, {"rbq_bits", bits}}; + auto resource = Static::EstimateLoadResource("DISKANN_RABITQ", version, file_bytes, rows, dim, config); + REQUIRE(resource.has_value()); + const auto code_size = faiss::RaBitQuantizer().compute_code_size(prepared_dim, bits); + const uint64_t model_bytes = rows * code_size + prepared_dim * (prepared_dim + 1) * sizeof(float); + REQUIRE(resource.value().memoryCost == file_bytes / 4 + model_bytes); + REQUIRE(resource.value().diskCost == file_bytes); + REQUIRE(resource.value().memoryCost > previous); + previous = resource.value().memoryCost; + + config["navigation_codec"] = "RABITQ"; + auto generic = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config); + REQUIRE(generic.has_value()); + REQUIRE(generic.value().memoryCost == resource.value().memoryCost); + + config["search_cache_budget_gb"] = 0.25; + auto cached = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config); + REQUIRE(cached.has_value()); + REQUIRE(cached.value().memoryCost == resource.value().memoryCost + (1ULL << 28)); + config["search_cache_budget_gb_ratio"] = 0.5; + auto ratio_cached = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config); + REQUIRE(ratio_cached.has_value()); + REQUIRE(ratio_cached.value().memoryCost == resource.value().memoryCost + rows * dim * sizeof(float) / 2); + + config["navigation_codec"] = "PQ"; + auto pq = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config); + REQUIRE(pq.has_value()); + REQUIRE(pq.value().memoryCost == file_bytes / 4); + REQUIRE(pq.value().diskCost == file_bytes); + } + } + for (const auto& shape : {std::pair{-1, dim}, + {rows, 0}, + {std::numeric_limits::max(), dim}, + {rows, std::numeric_limits::max()}}) { + auto invalid = Static::EstimateLoadResource("DISKANN_RABITQ", version, file_bytes, shape.first, shape.second, + {{"rbq_bits", 8}}); + REQUIRE_FALSE(invalid.has_value()); + } +} + TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { auto version = GenTestVersionList(); auto make_pack = []() { @@ -540,8 +609,7 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { conflict["navigation_codec"] = "PQ"; check_train_config(conflict, knowhere::Status::invalid_args); for (const auto* codec : {"PQ", "RABITQ", "TQ", "invalid"}) { - auto cfg = - knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN, version); + auto cfg = knowhere::IndexStaticFaced::CreateConfig(kNativeDiskANN, version); auto request = valid; request["navigation_codec"] = codec; std::string error; @@ -673,6 +741,7 @@ TEST_CASE("DiskANN RaBitQ shares Faiss codes and request-local query bits", "[di CAPTURE(dim, bits); knowhere::RaBitQStore::BuildFromFloatBin(raw, path, bits); knowhere::RaBitQStore store(path); + REQUIRE(store.MemorySize() == knowhere::RaBitQStore::EstimateMemorySize(17, dim, bits)); std::unique_ptr model(faiss::cppcontrib::knowhere::read_index(path.c_str())); const auto* pt = dynamic_cast(model.get()); REQUIRE(pt != nullptr); @@ -963,9 +1032,8 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { auto result = index.Search(query_ds, search_json, nullptr); REQUIRE(result.has_value()); if (rbq_bits == 4) { - auto generic = knowhere::IndexFactory::Instance() - .Create(knowhere::IndexEnum::INDEX_DISKANN, version, pack) - .value(); + auto generic = + knowhere::IndexFactory::Instance().Create(kNativeDiskANN, version, pack).value(); auto generic_load = deserialize_json; generic_load["navigation_codec"] = "RABITQ"; REQUIRE(generic.Deserialize(binset, generic_load) == knowhere::Status::success); @@ -1014,6 +1082,11 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { auto iterators = index.AnnIterator(query_ds, search_json, nullptr); REQUIRE_FALSE(iterators.has_value()); REQUIRE(iterators.error() == knowhere::Status::not_implemented); + auto range_json = search_json; + range_json["radius"] = 100.0; + auto range = index.RangeSearch(query_ds, range_json, nullptr); + REQUIRE_FALSE(range.has_value()); + REQUIRE(range.error() == knowhere::Status::not_implemented); } fs::remove_all(kDir); @@ -1077,7 +1150,8 @@ TEST_CASE("Test DISKANN_RABITQ inner product d+1 sidecar", "[diskann][rabitq][ip fs::remove_all(kDir); } -TEST_CASE("Test AiSAQ clamps navigation PQ for deserialize", "[diskann][aisaq][large_pq]") { +TEST_CASE("DiskANN and AiSAQ preserve the navigation PQ limit", "[diskann][aisaq][large_pq]") { + const auto* index_type = GENERATE(kNativeDiskANN, "AISAQ"); constexpr uint32_t rows = 300; constexpr uint32_t dim = 513; const auto test_dir = kDir + "/aisaq_large_pq"; @@ -1115,7 +1189,7 @@ TEST_CASE("Test AiSAQ clamps navigation PQ for deserialize", "[diskann][aisaq][l auto version = GenTestVersionList(); knowhere::BinarySet binset; { - auto index = knowhere::IndexFactory::Instance().Create("AISAQ", version, pack).value(); + auto index = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); REQUIRE(index.Build(nullptr, build_json) == knowhere::Status::success); REQUIRE(index.Serialize(binset) == knowhere::Status::success); } @@ -1127,7 +1201,7 @@ TEST_CASE("Test AiSAQ clamps navigation PQ for deserialize", "[diskann][aisaq][l REQUIRE(num_offsets == diskann::defaults::MAX_PQ_CHUNKS + 1); REQUIRE(offset_dim == 1); - auto loaded = knowhere::IndexFactory::Instance().Create("AISAQ", version, pack).value(); + auto loaded = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); REQUIRE(loaded.Deserialize(binset, deserialize_json) == knowhere::Status::success); fs::remove_all(test_dir); diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 0c7e88e17..0032c898c 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -43,8 +43,6 @@ namespace diskann { double get_memory_budget(const std::string &mem_budget_str); double get_memory_budget(double search_ram_budget_in_gb); - size_t get_num_pq_chunks(double pq_code_size_limit, size_t points_num, - size_t dim); void add_new_file_to_single_index(std::string index_file, std::string new_file); diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index 1b977d3b2..b23056d69 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -1607,17 +1607,6 @@ void create_aisaq_layout(const std::string base_file, const std::string mem_inde } } -size_t get_num_pq_chunks(double pq_code_size_limit, size_t points_num, - size_t dim) { - if (points_num == 0 || dim == 0) { - return 0; - } - size_t num_pq_chunks = - static_cast(std::floor(pq_code_size_limit / points_num)); - num_pq_chunks = std::max(num_pq_chunks, 1); - return std::min(num_pq_chunks, dim); -} - template int build_disk_index(BuildConfig &config) { if (!knowhere::KnowhereFloatTypeCheck::value && @@ -1744,13 +1733,12 @@ template << " Indexing ram budget: " << indexing_ram_budget << "(GiB)"; - // The ordinary PQFlashIndex scratch space is dimension-sized, so an - // in-memory navigation code may safely use up to one chunk per input - // dimension. AiSAQ's on-disk layout and loader retain a 512-chunk limit. - const size_t requested_pq_chunks = get_num_pq_chunks(pq_code_size_limit, points_num, dim); - const size_t num_pq_chunks = config.aisaq_mode - ? std::min(requested_pq_chunks, static_cast(diskann::defaults::MAX_PQ_CHUNKS)) - : requested_pq_chunks; + size_t num_pq_chunks = + (size_t) (std::floor)(_u64(pq_code_size_limit / points_num)); + + num_pq_chunks = num_pq_chunks <= 0 ? 1 : num_pq_chunks; + num_pq_chunks = num_pq_chunks > dim ? dim : num_pq_chunks; + num_pq_chunks = num_pq_chunks > diskann::defaults::MAX_PQ_CHUNKS ? diskann::defaults::MAX_PQ_CHUNKS : num_pq_chunks; LOG_KNOWHERE_INFO_ << "Compressing " << dim << "-dimensional data into " << num_pq_chunks << " bytes per vector."; diff --git a/thirdparty/faiss/faiss/impl/io.cpp b/thirdparty/faiss/faiss/impl/io.cpp index cf0018351..7a74e35b8 100644 --- a/thirdparty/faiss/faiss/impl/io.cpp +++ b/thirdparty/faiss/faiss/impl/io.cpp @@ -84,6 +84,10 @@ FileIOReader::~FileIOReader() { } size_t FileIOReader::operator()(void* ptr, size_t size, size_t nitems) { + // Empty serialized vectors may have a null data pointer. + if (size == 0 || nitems == 0) { + return 0; + } return fread(ptr, size, nitems, f); } @@ -119,6 +123,10 @@ FileIOWriter::~FileIOWriter() { } size_t FileIOWriter::operator()(const void* ptr, size_t size, size_t nitems) { + // Avoid passing null to the C stdio API, even for a zero-byte write. + if (size == 0 || nitems == 0) { + return 0; + } return fwrite(ptr, size, nitems, f); } From 4ac08717349b08c48a9ee69e6e365727a2aa46c3 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Mon, 21 Sep 2026 07:07:57 +0000 Subject: [PATCH 04/12] fix: make DiskANN SSD PQ scoring consistent across search paths Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 42 +++--- tests/ut/test_diskann_ssd_pq.cc | 126 ++++++++++++++++++ .../DiskANN/include/diskann/pq_flash_index.h | 20 +++ thirdparty/DiskANN/src/pq_flash_index.cpp | 60 ++++----- 4 files changed, 188 insertions(+), 60 deletions(-) create mode 100644 tests/ut/test_diskann_ssd_pq.cc diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 722a5aede..882f08006 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -108,6 +108,10 @@ class DiskANNIndexNode : public IndexNode { static bool StaticHasRawData(const knowhere::BaseConfig& config, const IndexVersion& version) { + const auto* disk_config = dynamic_cast(&config); + if (disk_config && disk_config->disk_pq_dims.value_or(0) > 0) { + return false; + } knowhere::MetricType metric_type = config.metric_type.has_value() ? config.metric_type.value() : ""; const auto& base_metric = get_sub_metric_type(metric_type).value_or(metric_type); return IsMetricType(base_metric, metric::L2) || IsMetricType(base_metric, metric::COSINE); @@ -136,6 +140,9 @@ class DiskANNIndexNode : public IndexNode { bool HasRawData(const std::string& metric_type) const override { + if (pq_flash_index_ && pq_flash_index_->uses_disk_pq()) { + return false; + } const auto& base_metric = get_sub_metric_type(metric_type).value_or(metric_type); return IsMetricType(base_metric, metric::L2) || IsMetricType(base_metric, metric::COSINE); } @@ -515,18 +522,10 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr 0) { - uint64_t disk_pq_nchunks = dim; - if (std::cmp_less(build_conf.disk_pq_dims.value(), dim)) { - disk_pq_nchunks = build_conf.disk_pq_dims.value(); - } - num_nodes_to_cache = GetCachedNodeNum(build_conf.search_cache_budget_gb.value(), disk_pq_nchunks, sizeof(_u8), - build_conf.max_degree.value()); - } else { - num_nodes_to_cache = GetCachedNodeNum(build_conf.search_cache_budget_gb.value(), dim, sizeof(DataType), - build_conf.max_degree.value()); - } + // The coordinate cache uses aligned T slots even when the SSD payload is PQ. + const auto cache_dim = ROUND_UP(dim + (diskann_metric == diskann::Metric::INNER_PRODUCT ? 1 : 0), 8); + const auto num_nodes_to_cache = GetCachedNodeNum(build_conf.search_cache_budget_gb.value(), cache_dim, + sizeof(DataType), build_conf.max_degree.value()); diskann::BuildConfig diskann_internal_build_config{data_path, index_prefix_, diskann_metric, @@ -782,19 +781,9 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr diskann::load_bin(cached_nodes_file, cached_nodes_ids, num_nodes, nodes_id_dim); node_list.assign(cached_nodes_ids.get(), cached_nodes_ids.get() + num_nodes); } else { - uint64_t num_nodes_to_cache = 0; - if (prep_conf.disk_pq_dims.value() > 0) { - uint64_t disk_pq_nchunks = pq_flash_index_->get_data_dim(); - if (prep_conf.disk_pq_dims.value() < static_cast(pq_flash_index_->get_data_dim())) { - disk_pq_nchunks = prep_conf.disk_pq_dims.value(); - } - num_nodes_to_cache = GetCachedNodeNum(prep_conf.search_cache_budget_gb.value(), disk_pq_nchunks, - sizeof(_u8), pq_flash_index_->get_max_degree()); - } else { - num_nodes_to_cache = - GetCachedNodeNum(prep_conf.search_cache_budget_gb.value(), pq_flash_index_->get_data_dim(), - sizeof(DataType), pq_flash_index_->get_max_degree()); - } + const auto num_nodes_to_cache = + GetCachedNodeNum(prep_conf.search_cache_budget_gb.value(), ROUND_UP(pq_flash_index_->get_data_dim(), 8), + sizeof(DataType), pq_flash_index_->get_max_degree()); if (num_nodes_to_cache > pq_flash_index_->get_num_points() / 3) { LOG_KNOWHERE_ERROR_ << "Failed to generate cache, num_nodes_to_cache(" << num_nodes_to_cache << ") is larger than 1/3 of the total data number."; @@ -1202,6 +1191,9 @@ DiskANNIndexNode::GetVectorByStorageIds(const DataSetPtr dataset, milv LOG_KNOWHERE_ERROR_ << "Failed to load diskann."; return expected::Err(Status::empty_index, "index not loaded"); } + if (pq_flash_index_->uses_disk_pq()) { + return expected::Err(Status::not_implemented, "SSD PQ does not retain original vectors"); + } auto dim = Dim(); auto rows = dataset->GetRows(); auto ids = dataset->GetIds(); diff --git a/tests/ut/test_diskann_ssd_pq.cc b/tests/ut/test_diskann_ssd_pq.cc new file mode 100644 index 000000000..2cb5000c8 --- /dev/null +++ b/tests/ut/test_diskann_ssd_pq.cc @@ -0,0 +1,126 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +#ifdef KNOWHERE_WITH_DISKANN +#include +#include +#include +#include + +#include "catch2/catch_approx.hpp" +#include "catch2/catch_test_macros.hpp" +#include "catch2/generators/catch_generators.hpp" +#include "diskann/linux_aligned_file_reader.h" +#include "diskann/pq_flash_index.h" +#include "filemanager/impl/LocalFileManager.h" +#include "knowhere/index/index_factory.h" +#include "utils.h" + +namespace { +// Decode the actual SSD payload for an independent scalar score reference. +class DiskPQReference : public diskann::PQFlashIndex { + public: + explicit DiskPQReference(diskann::Metric metric) + : PQFlashIndex(std::make_shared(), metric) { + } + + float Score(const std::string& prefix, unsigned id, const float* query, size_t dim) { + std::ifstream file(prefix + "_disk.index", std::ios::binary); + file.seekg(get_node_sector_offset(id) + (long_node ? 0 : (id % nnodes_per_sector) * max_node_len)); + std::vector code(disk_pq_n_chunks); + file.read(reinterpret_cast(code.data()), code.size()); + REQUIRE(file.good()); + std::vector decoded(data_dim); + disk_pq_table.inflate_vector(code.data(), decoded.data()); + double dot = 0, norm = 0, l2 = 0; + for (size_t j = 0; j < dim; ++j) { + dot += double(query[j]) * decoded[j]; + norm += double(query[j]) * query[j]; + l2 += (double(query[j]) - decoded[j]) * (double(query[j]) - decoded[j]); + } + if (metric == diskann::Metric::INNER_PRODUCT) return dot * max_base_norm; + if (metric == diskann::Metric::COSINE) return dot / std::sqrt(norm); + return l2; + } +}; +} // namespace + +TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[diskann][ssd_pq]") { + const auto metric = GENERATE(std::string("L2"), std::string("IP"), std::string("COSINE")); + const auto version = GenTestVersionList(); + constexpr size_t rows = 300, dim = 16, nq = 3, k = 10; + const auto dir = std::filesystem::current_path() / ("ssd_pq_regression_" + metric); + std::filesystem::create_directories(dir); + const auto prefix = (dir / "index").string(); + const auto raw = (dir / "base.bin").string(); + auto base = GenDataSet(rows, dim, 73); + auto queries = GenDataSet(nq, dim, 29); + auto* xb = const_cast(static_cast(base->GetTensor())); + for (size_t i = 0; i < rows; ++i) + for (size_t j = 0; j < dim; ++j) xb[i * dim + j] *= 0.2f + float(i % 7); + { + std::ofstream file(raw, std::ios::binary); + uint32_t header[] = {rows, dim}; + file.write(reinterpret_cast(header), sizeof(header)); + file.write(reinterpret_cast(xb), rows * dim * sizeof(float)); + } + auto pack = knowhere::Pack(std::shared_ptr(std::make_shared())); +#ifdef KNOWHERE_WITH_CARDINAL + const char* index_type = "DISKANN_DEPRECATED"; +#else + const char* index_type = "DISKANN"; +#endif + knowhere::Json config = {{"dim", dim}, {"metric_type", metric}, {"index_prefix", prefix}, {"data_path", raw}, + {"max_degree", 24}, {"search_list_size", 100}, {"pq_code_budget_gb", 0.001}, + {"build_dram_budget_gb", 1.0}, {"disk_pq_dims", 4}, + {"search_cache_budget_gb", 0}, {"search_cache_budget_gb_ratio", 0}}; + auto built = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); + REQUIRE(built.Build(nullptr, config) == knowhere::Status::success); + knowhere::BinarySet binary; + REQUIRE(built.Serialize(binary) == knowhere::Status::success); + const auto dm = metric == "L2" ? diskann::Metric::L2 + : metric == "IP" ? diskann::Metric::INNER_PRODUCT : diskann::Metric::COSINE; + const auto* xq = static_cast(queries->GetTensor()); + for (bool cached : {false, true}) { + auto index = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); + config["search_cache_budget_gb"] = cached ? 0.000005 : 0.0; + config["use_bfs_cache"] = true; + config["warm_up"] = true; + REQUIRE(index.Deserialize(binary, config) == knowhere::Status::success); + REQUIRE_FALSE(index.HasRawData(metric)); + int64_t id = 0; + REQUIRE_FALSE(index.GetVectorByIds(knowhere::GenIdsDataSet(1, &id)).has_value()); + DiskPQReference reference(dm); + REQUIRE(reference.load(1, prefix.c_str()) == 0); + if (cached) { + std::vector ids(60); + std::iota(ids.begin(), ids.end(), 0); + reference.load_cache_list(ids); + } + knowhere::Json search = {{"metric_type", metric}, {"k", k}, {"search_list_size", 128}, {"beamwidth", 4}}; + for (size_t filtered : {size_t(0), size_t(285)}) { + std::vector mask((rows + 7) / 8, 0); + for (size_t i = 0; i < filtered; ++i) mask[i / 8] |= 1u << (i % 8); + auto result = index.Search(queries, search, knowhere::BitsetView(mask.data(), rows)); + REQUIRE(result.has_value()); + for (size_t q = 0; q < nq; ++q) { + for (size_t j = 0; j < k; ++j) { + const auto offset = q * k + j; + const auto label = result.value()->GetIds()[offset]; + REQUIRE(label >= static_cast(filtered)); + REQUIRE(label < rows); + const auto expected = reference.Score(prefix, label, xq + q * dim, dim); + REQUIRE(result.value()->GetDistance()[offset] == Catch::Approx(expected).epsilon(0.0002).margin(0.0002)); + } + } + } + // Exercise both cached and uncached IDs in the independent rerank API. + int64_t ids[] = {1, 12, 61, 299}; + float distances[4]; + reference.calc_dist_by_ids(xq, ids, 4, distances); + for (size_t j = 0; j < 4; ++j) + REQUIRE(distances[j] == Catch::Approx(reference.Score(prefix, ids[j], xq, dim)).epsilon(0.0002).margin(0.0002)); + } + std::filesystem::remove_all(dir); +} +#endif diff --git a/thirdparty/DiskANN/include/diskann/pq_flash_index.h b/thirdparty/DiskANN/include/diskann/pq_flash_index.h index d4f885d2a..af54ee895 100644 --- a/thirdparty/DiskANN/include/diskann/pq_flash_index.h +++ b/thirdparty/DiskANN/include/diskann/pq_flash_index.h @@ -213,6 +213,8 @@ namespace diskann { virtual void load_cache_list(std::vector &node_list); + bool uses_disk_pq() const noexcept { return use_disk_index_pq; } + // asynchronously collect the access frequency of each node in the graph void async_generate_cache_list_from_sample_queries(std::string sample_bin, _u64 l_search, @@ -385,6 +387,24 @@ namespace diskann { } } + // SSD payload scoring, independent of the resident navigation codec. + // Keep the same internal score convention as the uncompressed path. + float disk_distance(const T *query, const float *query_float, + const T *payload, int32_t id) { + if (!use_disk_index_pq) { + return dist_cmp_wrap(query, payload, aligned_dim, id); + } + auto *code = reinterpret_cast<_u8 *>(const_cast(payload)); + if (metric == Metric::INNER_PRODUCT) { + return 2.0f + 2.0f * disk_pq_table.inner_product(query_float, code); + } + if (metric == Metric::COSINE) { + // Disk PQ encodes the normalized base, not the original vectors. + return disk_pq_table.inner_product(query_float, code); + } + return disk_pq_table.l2_distance(query_float, code); + } + float dist_cmp_float_wrap(const float *x, const float *y, size_t d, int32_t u) { if (metric == Metric::COSINE) { diff --git a/thirdparty/DiskANN/src/pq_flash_index.cpp b/thirdparty/DiskANN/src/pq_flash_index.cpp index 0a5003692..38180886d 100644 --- a/thirdparty/DiskANN/src/pq_flash_index.cpp +++ b/thirdparty/DiskANN/src/pq_flash_index.cpp @@ -1040,8 +1040,7 @@ namespace diskann { { std::shared_lock lock(this->cache_mtx); if (coord_cache.find(id) != coord_cache.end()) { - float dist = dist_cmp_wrap(query, coord_cache.at(id), - (size_t) aligned_dim, id); + float dist = disk_distance(query, query_float, coord_cache.at(id), id); max_heap.Push(dist, id); continue; } @@ -1081,8 +1080,7 @@ namespace diskann { char *node_buf = get_offset_to_node(sector_buf, cur_id); memcpy(node_fp_coords_copy, node_buf, disk_bytes_per_point); // Do we really need memcpy here? - float dist = dist_cmp_wrap(query, node_fp_coords_copy, - (size_t) aligned_dim, cur_id); + float dist = disk_distance(query, query_float, node_fp_coords_copy, cur_id); max_heap.Push(dist, cur_id); if (feder != nullptr) { feder->visit_info_.AddTopCandidateInfo(cur_id, dist); @@ -1158,7 +1156,12 @@ namespace diskann { std::unique_ptr, decltype(return_scratch)> scratch_guard(&data, return_scratch); auto query_norm_opt = init_thread_data(data, query1); if (!query_norm_opt.has_value()) { - // return an empty answer when calcu a zero point + // A zero IP/cosine query has no searchable direction. Explicitly mark + // every result missing rather than returning untouched output buffers. + std::fill_n(indices, k_search, -1); + if (distances != nullptr) { + std::fill_n(distances, k_search, -1.0f); + } return; } float query_norm = query_norm_opt.value(); @@ -1412,19 +1415,8 @@ namespace diskann { auto process_node = [&](T *node_fp_coords_copy, auto node_id, auto n_nbr, auto *nbrs) { if (bitset_view.empty() || !bitset_view.test(node_id)) { - float cur_expanded_dist; - if (!use_disk_index_pq) { - cur_expanded_dist = dist_cmp_wrap(query, node_fp_coords_copy, - (size_t) aligned_dim, node_id); - } else { - if (metric == diskann::Metric::INNER_PRODUCT || - metric == diskann::Metric::COSINE) - cur_expanded_dist = disk_pq_table.inner_product( - query_float, (_u8 *) node_fp_coords_copy); - else - cur_expanded_dist = disk_pq_table.l2_distance( - query_float, (_u8 *) node_fp_coords_copy); - } + const float cur_expanded_dist = + disk_distance(query, query_float, node_fp_coords_copy, node_id); full_retset.push_back( Neighbor((unsigned) node_id, cur_expanded_dist, true)); @@ -1611,6 +1603,11 @@ namespace diskann { void PQFlashIndex::calc_dist_by_ids(const T *query_, const int64_t *ids, const int64_t n, float *const output_dists) { + for (int64_t i = 0; i < n; ++i) { + if (ids[i] < 0 || static_cast<_u64>(ids[i]) >= num_points) { + throw ANNException("Invalid storage id", -1); + } + } ThreadData data = this->thread_data.pop(); while (data.scratch.sector_scratch == nullptr) { this->thread_data.wait_for_push_notify(); @@ -1618,6 +1615,7 @@ namespace diskann { } auto query_norm_opt = init_thread_data(data, query_); if (!query_norm_opt.has_value()) { + std::fill_n(output_dists, n, -1.0f); this->thread_data.push(data); this->thread_data.push_notify_all(); return; @@ -1664,7 +1662,7 @@ namespace diskann { if (it != coord_cache.end()) { // Vector is in cache, calculate distance directly output_dists[i] = - dist_cmp_wrap(query, it->second, (size_t) aligned_dim, id); + disk_distance(query, data.scratch.aligned_query_float, it->second, id); } else { // Need to read from disk @@ -1734,9 +1732,9 @@ namespace diskann { char *node_buf = get_offset_to_node(sector_buf, id); T *node_coords = OFFSET_TO_NODE_COORDS(node_buf); - // Calculate raw distance (not PQ distance) + // Score the actual SSD representation; PQ bytes are not float vectors. output_dists[output_idx] = - dist_cmp_wrap(query, node_coords, (size_t) aligned_dim, id); + disk_distance(query, data.scratch.aligned_query_float, node_coords, id); } } } @@ -1783,6 +1781,9 @@ namespace diskann { template void PQFlashIndex::get_vector_by_ids(const int64_t *ids, const int64_t n, T *output_data) { + if (use_disk_index_pq) { + throw ANNException("SSD PQ does not retain original vectors", -1); + } auto sectors_to_visit = get_sectors_layout_and_write_data_from_cache(ids, n, output_data); if (0 == sectors_to_visit.size()) { @@ -1984,20 +1985,9 @@ namespace diskann { auto process_node = [&](T *node_fp_coords_copy, auto node_id, auto n_nbr, auto *nbrs) { if (workspace->bitset.empty() || !workspace->bitset.test(node_id)) { - float cur_expanded_dist; - if (!use_disk_index_pq) { - cur_expanded_dist = - dist_cmp_wrap(workspace->aligned_query_T, node_fp_coords_copy, - (size_t) aligned_dim, node_id); - } else { - if (metric == diskann::Metric::INNER_PRODUCT || - metric == diskann::Metric::COSINE) - cur_expanded_dist = disk_pq_table.inner_product( - workspace->aligned_query_float, (_u8 *) node_fp_coords_copy); - else - cur_expanded_dist = disk_pq_table.l2_distance( - workspace->aligned_query_float, (_u8 *) node_fp_coords_copy); - } + const float cur_expanded_dist = disk_distance( + workspace->aligned_query_T, workspace->aligned_query_float, + node_fp_coords_copy, node_id); workspace->insert_to_full((unsigned) node_id, cur_expanded_dist); } From 02272b90a4a1fb425a1f9808665a76f7220f900c Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Mon, 21 Sep 2026 07:20:00 +0000 Subject: [PATCH 05/12] feat: decouple DiskANN RaBitQ navigation from SSD storage and PQ resources Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 51 ++++--- src/index/diskann/diskann_config.h | 7 +- src/index/diskann/rabitq_store.cc | 4 +- tests/python/test_diskann_rabitq.py | 16 ++- tests/ut/test_diskann.cc | 130 ++++++++++++++++-- tests/ut/test_diskann_ssd_pq.cc | 61 ++++++-- .../DiskANN/include/diskann/aux_utils.h | 6 +- .../DiskANN/include/diskann/pq_flash_index.h | 9 +- thirdparty/DiskANN/src/aux_utils.cpp | 18 ++- thirdparty/DiskANN/src/pq_flash_index.cpp | 54 ++++++-- 10 files changed, 282 insertions(+), 74 deletions(-) diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 882f08006..6828371cc 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -409,16 +409,18 @@ TryDiskANNCall(std::function&& diskann_call) { std::vector GetNecessaryFilenames(const std::string& prefix, const bool need_norm, const bool use_sample_cache, - const bool use_sample_warmup) { + const bool use_sample_warmup, const bool use_pq_navigation = true) { std::vector filenames; auto pq_pivots_filename = diskann::get_pq_pivots_filename(prefix); auto disk_index_filename = diskann::get_disk_index_filename(prefix); - filenames.push_back(pq_pivots_filename); - filenames.push_back(diskann::get_pq_rearrangement_perm_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_chunk_offsets_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_centroid_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_compressed_filename(prefix)); + if (use_pq_navigation) { + filenames.push_back(pq_pivots_filename); + filenames.push_back(diskann::get_pq_rearrangement_perm_filename(pq_pivots_filename)); + filenames.push_back(diskann::get_pq_chunk_offsets_filename(pq_pivots_filename)); + filenames.push_back(diskann::get_pq_centroid_filename(pq_pivots_filename)); + filenames.push_back(diskann::get_pq_compressed_filename(prefix)); + } filenames.push_back(disk_index_filename); if (need_norm) { filenames.push_back(diskann::get_disk_index_max_base_norm_file(disk_index_filename)); @@ -525,7 +527,7 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::Build(const DataSetPtr dataset, std::shared_ptr(num_nodes_to_cache), build_conf.shuffle_build.value()}; - diskann_internal_build_config.keep_preprocessed_base = - external_navigation && diskann_metric == diskann::Metric::INNER_PRODUCT; + const bool navigation_uses_preprocessed_base = external_navigation && need_norm; + diskann_internal_build_config.keep_preprocessed_base = navigation_uses_preprocessed_base; + diskann_internal_build_config.use_pq_navigation = !external_navigation; RETURN_IF_ERROR(TryDiskANNCall([&]() { int res = diskann::build_disk_index(diskann_internal_build_config); if (res != 0) @@ -550,14 +553,14 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::Build(const DataSetPtr dataset, std::shared_ptr::Deserialize(const BinarySet& binset, std::shared_ptr for (auto& filename : GetNecessaryFilenames( index_prefix_, need_norm, prep_conf.search_cache_budget_gb.value() > 0 && !prep_conf.use_bfs_cache.value() && !external_navigation, - prep_conf.warm_up.value())) { + prep_conf.warm_up.value(), !external_navigation)) { if (!LoadFile(filename)) { return Status::disk_file_error; } @@ -725,6 +728,16 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr } } + navigation_store_.reset(); + if (external_navigation) { + try { + navigation_store_ = LoadNavigationStore(prep_conf, index_prefix_); + } catch (const std::exception& e) { + LOG_KNOWHERE_ERROR_ << "Failed to initialize DiskANN navigation sidecar: " << e.what(); + return Status::invalid_index_error; + } + } + // set thread pool search_pool_ = ThreadPool::GetGlobalSearchThreadPool(); @@ -735,7 +748,13 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr pq_flash_index_ = std::make_unique>(reader, diskann_metric); auto disk_ann_call = [&]() { - int res = pq_flash_index_->load(search_pool_->size(), index_prefix_.c_str(), !external_navigation); + typename diskann::PQFlashIndex::NavigationMetadata metadata{}; + if (navigation_store_) { + metadata.count = navigation_store_->Count(); + metadata.dimension = navigation_store_->Dimension(); + } + int res = pq_flash_index_->load(search_pool_->size(), index_prefix_.c_str(), !external_navigation, + navigation_store_ ? &metadata : nullptr); if (res != 0) { throw diskann::ANNException("pq_flash_index_->load returned non-zero value: " + std::to_string(res), -1); } @@ -753,10 +772,8 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr dim_.store(pq_flash_index_->get_data_dim()); } - navigation_store_.reset(); if (external_navigation) { try { - navigation_store_ = LoadNavigationStore(prep_conf, index_prefix_); if (navigation_store_->Count() != static_cast(pq_flash_index_->get_num_points()) || navigation_store_->Dimension() != static_cast(pq_flash_index_->get_data_dim())) { LOG_KNOWHERE_ERROR_ << "DiskANN graph and navigation sidecar metadata do not match"; diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index 58e7277d1..a3208413c 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -255,8 +255,8 @@ class DiskANNNavigationConfig : public DiskANNConfig { return HandleError(err_msg, "unsupported DiskANN navigation codec", Status::invalid_args); } const auto metric = metric_type.value_or(knowhere::metric::L2); - if (metric != knowhere::metric::L2 && metric != knowhere::metric::IP) { - return HandleError(err_msg, "DISKANN_RABITQ supports L2 and IP", Status::invalid_metric_type); + if (metric != knowhere::metric::L2 && metric != knowhere::metric::IP && metric != knowhere::metric::COSINE) { + return HandleError(err_msg, "DISKANN_RABITQ supports L2, IP and COSINE", Status::invalid_metric_type); } const auto database_bits = rbq_bits.value_or(1); if (database_bits < 1 || database_bits > 9) { @@ -266,9 +266,6 @@ class DiskANNNavigationConfig : public DiskANNConfig { if (refine_mode != "probabilistic" && refine_mode != "full") { return HandleError(err_msg, "rbq_refine_mode must be probabilistic or full", Status::invalid_args); } - if (disk_pq_dims.value_or(0) != 0) { - return HandleError(err_msg, "DISKANN_RABITQ requires disk_pq_dims=0", Status::invalid_args); - } return Status::success; } }; diff --git a/src/index/diskann/rabitq_store.cc b/src/index/diskann/rabitq_store.cc index 47458c852..220fbb3f4 100644 --- a/src/index/diskann/rabitq_store.cc +++ b/src/index/diskann/rabitq_store.cc @@ -333,8 +333,8 @@ RaBitQStore::CreateDistanceComputer(const DiskANNConfig& config) const { throw std::invalid_argument("RaBitQ navigation requires query configuration"); } const auto query_metric = config.metric_type.value_or(metric::L2); - if (query_metric != metric::L2 && query_metric != metric::IP) { - throw std::invalid_argument("RaBitQ navigation currently supports L2 and IP"); + if (query_metric != metric::L2 && query_metric != metric::IP && query_metric != metric::COSINE) { + throw std::invalid_argument("RaBitQ navigation supports L2, IP and COSINE"); } const auto mode = navigation->rbq_refine_mode.value_or("probabilistic"); if (mode != "probabilistic" && mode != "full") { diff --git a/tests/python/test_diskann_rabitq.py b/tests/python/test_diskann_rabitq.py index 98d7a58e2..1854b8e10 100644 --- a/tests/python/test_diskann_rabitq.py +++ b/tests/python/test_diskann_rabitq.py @@ -9,12 +9,15 @@ import pytest -@pytest.mark.parametrize("metric", ["L2", "IP"]) +@pytest.mark.parametrize("metric", ["L2", "IP", "COSINE"]) @pytest.mark.parametrize("kind,codec", [("DISKANN_RABITQ", None), ("DISKANN", "RABITQ"), ("DISKANN", None)]) def test_navigation_roundtrip(tmp_path, metric, kind, codec): rng = np.random.default_rng(42) base = rng.normal(size=(1000, 64)).astype("float32") query = rng.normal(size=(10, 64)).astype("float32") + if metric == "COSINE": + base *= rng.uniform(0.1, 10, size=(1000, 1)).astype("float32") + query *= rng.uniform(0.1, 10, size=(10, 1)).astype("float32") source = tmp_path / "base.fbin" with source.open("wb") as output: output.write(struct.pack("= 0) and np.all(np.isfinite(distances)) exact = ((query[:, None, :] - base[None, :, :]) ** 2).sum(2) if metric == "L2" else -(query @ base.T) + if metric == "COSINE": + exact /= np.linalg.norm(query, axis=1)[:, None] * np.linalg.norm(base, axis=1)[None, :] truth = np.argsort(exact, axis=1)[:, :10] recall = sum(len(set(actual) & set(expected)) for actual, expected in zip(ids, truth)) / 100 assert recall > 0.8 + if metric == "COSINE": + expected_scores = -np.take_along_axis(exact, ids, axis=1) + np.testing.assert_allclose(distances, expected_scores, rtol=1e-4, atol=1e-5) + if metric in ("IP", "COSINE"): + zeros = np.zeros((1, 64), dtype="float32") + empty, status = restored.Search(knowhere.ArrayToDataSet(zeros), json.dumps(search), knowhere.GetNullBitSetView()) + assert knowhere.Status(status) == knowhere.Status.success + empty_distances, empty_ids = knowhere.DataSetToArray(empty) + assert np.all(empty_ids == -1) and np.all(empty_distances == -1) diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index 98c55e6da..bca94cf04 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -526,7 +526,7 @@ TEST_CASE("DiskANN navigation load resource estimate", "[diskann][rabitq][resour constexpr uint64_t file_bytes = 1ULL << 30; constexpr int64_t rows = 1000000; constexpr int64_t dim = 769; - for (const auto* metric : {"L2", "IP"}) { + for (const auto* metric : {"L2", "IP", "COSINE"}) { const int64_t prepared_dim = dim + (std::string(metric) == "IP" ? 1 : 0); uint64_t previous = 0; for (int bits = 1; bits <= 9; ++bits) { @@ -633,13 +633,15 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { invalid["metric_type"] = knowhere::metric::IP; check_train_config(invalid, knowhere::Status::success); invalid["metric_type"] = knowhere::metric::COSINE; + check_train_config(invalid, knowhere::Status::success); + invalid["metric_type"] = "HAMMING"; check_train_config(invalid, knowhere::Status::invalid_metric_type); invalid = valid; invalid["rbq_bits"] = 10; check_train_config(invalid, knowhere::Status::out_of_range_in_json); invalid = valid; invalid["disk_pq_dims"] = 16; - check_train_config(invalid, knowhere::Status::invalid_args); + check_train_config(invalid, knowhere::Status::success); invalid = valid; invalid["search_cache_budget_gb"] = 0.01; check_train_config(invalid, knowhere::Status::success); @@ -964,19 +966,14 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { } if (rbq_bits == 1) { - auto full_reader = std::make_shared(); - diskann::PQFlashIndex full_pq_index(full_reader, diskann::Metric::L2); - REQUIRE(full_pq_index.load(1, rabitq_prefix.c_str(), true) == 0); - auto metadata_reader = std::make_shared(); diskann::PQFlashIndex metadata_only_index(metadata_reader, diskann::Metric::L2); - REQUIRE(metadata_only_index.load(1, rabitq_prefix.c_str(), false) == 0); - REQUIRE(metadata_only_index.get_num_points() == full_pq_index.get_num_points()); - REQUIRE(metadata_only_index.get_data_dim() == full_pq_index.get_data_dim()); - - const auto pq_file_size = fs::file_size(rabitq_prefix + "_pq_compressed.bin"); - const auto pq_code_bytes = pq_file_size - 2 * sizeof(uint32_t); - REQUIRE(full_pq_index.cal_size() - metadata_only_index.cal_size() == pq_code_bytes); + const diskann::PQFlashIndex::NavigationMetadata metadata{kNumRows, kDim}; + REQUIRE(metadata_only_index.load(1, rabitq_prefix.c_str(), false, &metadata) == 0); + REQUIRE(metadata_only_index.get_num_points() == kNumRows); + REQUIRE(metadata_only_index.get_data_dim() == kDim); + REQUIRE_FALSE(fs::exists(rabitq_prefix + "_pq_compressed.bin")); + REQUIRE_FALSE(fs::exists(rabitq_prefix + "_pq_pivots.bin")); std::array query{}; std::array ids{}; @@ -1150,6 +1147,104 @@ TEST_CASE("Test DISKANN_RABITQ inner product d+1 sidecar", "[diskann][rabitq][ip fs::remove_all(kDir); } +TEST_CASE("DiskANN RaBitQ cosine uses normalized navigation and original SSD vectors", "[diskann][rabitq][cosine]") { + const auto dir = kDir + "/rabitq_cosine"; + REQUIRE_NOTHROW(fs::create_directories(dir)); + const auto prefix = dir + "/index"; + const auto raw = dir + "/base.fbin"; + auto base = GenDataSet(kNumRows, kDim, 30); + auto query = GenDataSet(kNumQueries, kDim, 42); + auto* xb = const_cast(static_cast(base->GetTensor())); + auto* xq = const_cast(static_cast(query->GetTensor())); + for (uint32_t i = 0; i < kNumRows; ++i) { + const float scale = (i == 0) ? 0.0f : (0.1f + float(i % 100)); + for (uint32_t j = 0; j < kDim; ++j) xb[i * kDim + j] *= scale; + } + for (uint32_t i = 0; i < kNumQueries; ++i) + for (uint32_t j = 0; j < kDim; ++j) xq[i * kDim + j] *= 0.2f + float(i); + WriteRawDataToDisk(raw, xb, kNumRows, kDim); + auto pack = knowhere::Pack(std::shared_ptr(std::make_shared())); + const auto version = GenTestVersionList(); + knowhere::Json config = {{"dim", kDim}, + {"metric_type", "COSINE"}, + {"index_prefix", prefix}, + {"data_path", raw}, + {"max_degree", defaultMaxDegree}, + {"search_list_size", 128}, + {"pq_code_budget_gb", 0.001}, + {"build_dram_budget_gb", 1.0}, + {"disk_pq_dims", 0}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"rbq_bits", 4}}; + knowhere::BinarySet binary; + { + auto built = knowhere::IndexFactory::Instance().Create("DISKANN_RABITQ", version, pack).value(); + REQUIRE(built.Build(nullptr, config) == knowhere::Status::success); + REQUIRE(built.Serialize(binary) == knowhere::Status::success); + REQUIRE_FALSE(fs::exists(prefix + "_prepped_base.bin")); + knowhere::RaBitQStore store(prefix + "_rabitq.index"); + REQUIRE(store.Dimension() == kDim); + } + auto restored = knowhere::IndexFactory::Instance().Create("DISKANN_RABITQ", version, pack).value(); + config["warm_up"] = true; + config["search_cache_budget_gb"] = 0.00005; + config["use_bfs_cache"] = false; + REQUIRE(restored.Deserialize(binary, config) == knowhere::Status::success); + // SSD must still contain original vectors, not the normalized navigation copy. + int64_t raw_id = 1; + auto id_dataset = knowhere::GenIdsDataSet(1, &raw_id); + auto retrieved = restored.GetVectorByIds(id_dataset); + REQUIRE(retrieved.has_value()); + const auto* original = static_cast(retrieved.value()->GetTensor()); + for (uint32_t j = 0; j < kDim; ++j) REQUIRE(original[j] == xb[kDim + j]); + knowhere::Json search = { + {"dim", kDim}, {"metric_type", "COSINE"}, {"k", kK}, {"search_list_size", 128}, {"beamwidth", 8}}; + auto exact = knowhere::BruteForce::Search(base, query, search, nullptr); + REQUIRE(exact.has_value()); + for (const int qb : {0, 4, 8}) { + for (const auto* mode : {"full", "probabilistic"}) { + search["rbq_bits_query"] = qb; + search["rbq_refine_mode"] = mode; + auto result = restored.Search(query, search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(GetKNNRecall(*exact.value(), *result.value()) > 0.8f); + for (uint32_t i = 0; i < kNumQueries; ++i) { + for (uint32_t j = 0; j < kK; ++j) { + const auto offset = i * kK + j; + const auto id = result.value()->GetIds()[offset]; + REQUIRE(id >= 0); + REQUIRE(id < kNumRows); + double dot = 0, norm_x = 0, norm_q = 0; + for (uint32_t k = 0; k < kDim; ++k) { + dot += double(xb[id * kDim + k]) * xq[i * kDim + k]; + norm_x += double(xb[id * kDim + k]) * xb[id * kDim + k]; + norm_q += double(xq[i * kDim + k]) * xq[i * kDim + k]; + } + const double score = norm_x > 0 && norm_q > 0 ? dot / std::sqrt(norm_x * norm_q) : 0; + REQUIRE(std::abs(result.value()->GetDistance()[offset] - score) < 1e-5); + } + } + } + } + for (const int excluded : {1, 950, 1000}) { + auto mask = GenerateBitsetWithFirstTbitsSet(kNumRows, excluded); + auto filtered = restored.Search(query, search, knowhere::BitsetView(mask.data(), kNumRows)); + REQUIRE(filtered.has_value()); + for (uint32_t i = 0; i < kNumQueries * kK; ++i) { + const auto id = filtered.value()->GetIds()[i]; + REQUIRE((id == -1 || (id >= excluded && id < kNumRows))); + if (excluded == 1000) + REQUIRE(id == -1); + } + } + std::vector zero(kDim, 0); + auto zero_result = restored.Search(knowhere::GenDataSet(1, kDim, zero.data()), search, nullptr); + REQUIRE(zero_result.has_value()); + for (uint32_t j = 0; j < kK; ++j) REQUIRE(zero_result.value()->GetIds()[j] == -1); + fs::remove_all(dir); +} + TEST_CASE("DiskANN and AiSAQ preserve the navigation PQ limit", "[diskann][aisaq][large_pq]") { const auto* index_type = GENERATE(kNativeDiskANN, "AISAQ"); constexpr uint32_t rows = 300; @@ -1204,6 +1299,15 @@ TEST_CASE("DiskANN and AiSAQ preserve the navigation PQ limit", "[diskann][aisaq auto loaded = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); REQUIRE(loaded.Deserialize(binset, deserialize_json) == knowhere::Status::success); + auto query = GenDataSet(3, dim, 42); + const knowhere::Json search = { + {"dim", dim}, {"metric_type", "L2"}, {"k", 5}, {"search_list_size", 50}, {"beamwidth", 4}}; + auto result = loaded.Search(query, search, nullptr); + REQUIRE(result.has_value()); + auto exact = knowhere::BruteForce::Search(base_ds, query, search, nullptr); + REQUIRE(exact.has_value()); + REQUIRE(GetKNNRecall(*exact.value(), *result.value()) > 0.8f); + fs::remove_all(test_dir); } diff --git a/tests/ut/test_diskann_ssd_pq.cc b/tests/ut/test_diskann_ssd_pq.cc index 2cb5000c8..8cfc144c2 100644 --- a/tests/ut/test_diskann_ssd_pq.cc +++ b/tests/ut/test_diskann_ssd_pq.cc @@ -24,7 +24,19 @@ class DiskPQReference : public diskann::PQFlashIndex { : PQFlashIndex(std::make_shared(), metric) { } - float Score(const std::string& prefix, unsigned id, const float* query, size_t dim) { + void + CheckExternalResources() { + REQUIRE(data == nullptr); + REQUIRE(pq_table.get_total_dims() == 0); + auto slot = thread_data.pop(); + REQUIRE(slot.scratch.aligned_pq_coord_scratch == nullptr); + REQUIRE(slot.scratch.aligned_pqtable_dist_scratch == nullptr); + REQUIRE(slot.scratch.coord_scratch != nullptr); + thread_data.push(slot); + } + + float + Score(const std::string& prefix, unsigned id, const float* query, size_t dim) { std::ifstream file(prefix + "_disk.index", std::ios::binary); file.seekg(get_node_sector_offset(id) + (long_node ? 0 : (id % nnodes_per_sector) * max_node_len)); std::vector code(disk_pq_n_chunks); @@ -38,8 +50,10 @@ class DiskPQReference : public diskann::PQFlashIndex { norm += double(query[j]) * query[j]; l2 += (double(query[j]) - decoded[j]) * (double(query[j]) - decoded[j]); } - if (metric == diskann::Metric::INNER_PRODUCT) return dot * max_base_norm; - if (metric == diskann::Metric::COSINE) return dot / std::sqrt(norm); + if (metric == diskann::Metric::INNER_PRODUCT) + return dot * max_base_norm; + if (metric == diskann::Metric::COSINE) + return dot / std::sqrt(norm); return l2; } }; @@ -47,9 +61,10 @@ class DiskPQReference : public diskann::PQFlashIndex { TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[diskann][ssd_pq]") { const auto metric = GENERATE(std::string("L2"), std::string("IP"), std::string("COSINE")); + const bool external = GENERATE(false, true); const auto version = GenTestVersionList(); constexpr size_t rows = 300, dim = 16, nq = 3, k = 10; - const auto dir = std::filesystem::current_path() / ("ssd_pq_regression_" + metric); + const auto dir = std::filesystem::current_path() / ("ssd_pq_regression_" + metric + (external ? "_rbq" : "_pq")); std::filesystem::create_directories(dir); const auto prefix = (dir / "index").string(); const auto raw = (dir / "base.bin").string(); @@ -70,16 +85,31 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis #else const char* index_type = "DISKANN"; #endif - knowhere::Json config = {{"dim", dim}, {"metric_type", metric}, {"index_prefix", prefix}, {"data_path", raw}, - {"max_degree", 24}, {"search_list_size", 100}, {"pq_code_budget_gb", 0.001}, - {"build_dram_budget_gb", 1.0}, {"disk_pq_dims", 4}, - {"search_cache_budget_gb", 0}, {"search_cache_budget_gb_ratio", 0}}; + knowhere::Json config = {{"dim", dim}, + {"metric_type", metric}, + {"index_prefix", prefix}, + {"data_path", raw}, + {"max_degree", 24}, + {"search_list_size", 100}, + {"pq_code_budget_gb", 0.001}, + {"build_dram_budget_gb", 1.0}, + {"disk_pq_dims", 4}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}}; + if (external) { + config["navigation_codec"] = "RABITQ"; + config["rbq_bits"] = 4; + } auto built = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); REQUIRE(built.Build(nullptr, config) == knowhere::Status::success); knowhere::BinarySet binary; REQUIRE(built.Serialize(binary) == knowhere::Status::success); - const auto dm = metric == "L2" ? diskann::Metric::L2 - : metric == "IP" ? diskann::Metric::INNER_PRODUCT : diskann::Metric::COSINE; + REQUIRE(std::filesystem::exists(prefix + "_pq_compressed.bin") == !external); + REQUIRE(std::filesystem::exists(prefix + "_pq_pivots.bin") == !external); + REQUIRE(std::filesystem::exists(prefix + "_disk.index_pq_pivots.bin")); + const auto dm = metric == "L2" ? diskann::Metric::L2 + : metric == "IP" ? diskann::Metric::INNER_PRODUCT + : diskann::Metric::COSINE; const auto* xq = static_cast(queries->GetTensor()); for (bool cached : {false, true}) { auto index = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); @@ -91,7 +121,10 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis int64_t id = 0; REQUIRE_FALSE(index.GetVectorByIds(knowhere::GenIdsDataSet(1, &id)).has_value()); DiskPQReference reference(dm); - REQUIRE(reference.load(1, prefix.c_str()) == 0); + const DiskPQReference::NavigationMetadata metadata{rows, dim + (metric == "IP" ? 1 : 0)}; + REQUIRE(reference.load(1, prefix.c_str(), !external, external ? &metadata : nullptr) == 0); + if (external) + reference.CheckExternalResources(); if (cached) { std::vector ids(60); std::iota(ids.begin(), ids.end(), 0); @@ -110,7 +143,8 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis REQUIRE(label >= static_cast(filtered)); REQUIRE(label < rows); const auto expected = reference.Score(prefix, label, xq + q * dim, dim); - REQUIRE(result.value()->GetDistance()[offset] == Catch::Approx(expected).epsilon(0.0002).margin(0.0002)); + REQUIRE(result.value()->GetDistance()[offset] == + Catch::Approx(expected).epsilon(0.0002).margin(0.0002)); } } } @@ -119,7 +153,8 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis float distances[4]; reference.calc_dist_by_ids(xq, ids, 4, distances); for (size_t j = 0; j < 4; ++j) - REQUIRE(distances[j] == Catch::Approx(reference.Score(prefix, ids[j], xq, dim)).epsilon(0.0002).margin(0.0002)); + REQUIRE(distances[j] == + Catch::Approx(reference.Score(prefix, ids[j], xq, dim)).epsilon(0.0002).margin(0.0002)); } std::filesystem::remove_all(dir); } diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 0032c898c..4295bd488 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -131,9 +131,11 @@ namespace diskann { uint32_t inline_pq = 0; bool rearrange = false; int num_entry_points = 0; - // Keep the temporary MIPS-to-L2 base until the caller builds auxiliary - // indexes that must use the exact same internal d+1 representation. + // Keep the temporary MIPS-to-L2 or normalized cosine base until the caller + // builds auxiliary indexes in the same internal representation. bool keep_preprocessed_base = false; + // External navigation owns its codes; SSD PQ remains independently optional. + bool use_pq_navigation = true; }; template diff --git a/thirdparty/DiskANN/include/diskann/pq_flash_index.h b/thirdparty/DiskANN/include/diskann/pq_flash_index.h index af54ee895..62ef2f3eb 100644 --- a/thirdparty/DiskANN/include/diskann/pq_flash_index.h +++ b/thirdparty/DiskANN/include/diskann/pq_flash_index.h @@ -203,13 +203,18 @@ namespace diskann { template class PQFlashIndex: public PQDataGetter { public: + struct NavigationMetadata { + uint64_t count; + uint64_t dimension; + }; PQFlashIndex(std::shared_ptr fileReader, diskann::Metric metric = diskann::Metric::L2); ~PQFlashIndex(); // load compressed data, and obtains the handle to the disk-resident index int load(uint32_t num_threads, const char *index_prefix, - bool load_pq_data = true); + bool load_pq_data = true, + const NavigationMetadata* navigation_metadata = nullptr); virtual void load_cache_list(std::vector &node_list); @@ -374,6 +379,8 @@ namespace diskann { std::unique_ptr<_u8[]> data = nullptr; _u64 n_chunks; FixedChunkPQTable pq_table; + // AiSAQ also needs PQ scratch although its codes are disk resident. + bool use_pq_navigation = true; // distance comparator DISTFUN dist_cmp; diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index b23056d69..1517ab2f7 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -1609,6 +1609,9 @@ void create_aisaq_layout(const std::string base_file, const std::string mem_inde template int build_disk_index(BuildConfig &config) { + if (config.aisaq_mode && !config.use_pq_navigation) { + throw diskann::ANNException("AiSAQ requires PQ navigation", -1); + } if (!knowhere::KnowhereFloatTypeCheck::value && (config.compare_metric == diskann::Metric::INNER_PRODUCT || config.compare_metric == diskann::Metric::COSINE)) { @@ -1743,14 +1746,16 @@ template LOG_KNOWHERE_INFO_ << "Compressing " << dim << "-dimensional data into " << num_pq_chunks << " bytes per vector."; - size_t train_size, train_dim; + size_t train_size = 0, train_dim = 0; std::unique_ptr train_data = nullptr; double p_val = ((double) MAX_PQ_TRAINING_SET_SIZE / (double) points_num); // generates random sample and sets it to train_data and updates // train_size - gen_random_slice(data_file_to_use.c_str(), p_val, train_data, train_size, - train_dim); + if (config.use_pq_navigation || use_disk_pq) { + gen_random_slice(data_file_to_use.c_str(), p_val, train_data, train_size, + train_dim); + } if (use_disk_pq) { if (disk_pq_dims > dim) @@ -1771,6 +1776,7 @@ template data_file_to_use.c_str(), 256, (uint32_t) disk_pq_dims, disk_pq_pivots_path, disk_pq_compressed_vectors_path); } + if (config.use_pq_navigation) { LOG_KNOWHERE_DEBUG_ << "Training data loaded of size " << train_size; // don't translate data to make zero mean for PQ compression. We must not @@ -1793,6 +1799,7 @@ template auto pq_e = std::chrono::high_resolution_clock::now(); std::chrono::duration pq_diff = pq_e - pq_s; LOG_KNOWHERE_INFO_ << "Training PQ codes cost: " << pq_diff.count() << "s"; + } // Gopal. Splitting diskann_dll into separate DLLs for search and build. // This code should only be available in the "build" DLL. #if defined(RELEASE_UNUSED_TCMALLOC_MEMORY_AT_CHECKPOINTS) && \ @@ -1866,7 +1873,7 @@ template gen_random_slice(base_file.c_str(), sample_data_file, sample_sampling_rate); - if (vamana_index != nullptr) { + if (vamana_index != nullptr && config.use_pq_navigation) { auto final_graph = vamana_index->get_graph(); auto entry_point = vamana_index->get_entry_point(); @@ -1889,7 +1896,8 @@ template std::chrono::duration diff = e - s; LOG_KNOWHERE_INFO_ << "Indexing time: " << diff.count(); - if (config.compare_metric == diskann::Metric::INNER_PRODUCT && + if ((config.compare_metric == diskann::Metric::INNER_PRODUCT || + config.compare_metric == diskann::Metric::COSINE) && !config.keep_preprocessed_base) { std::remove(data_file_to_use.c_str()); } diff --git a/thirdparty/DiskANN/src/pq_flash_index.cpp b/thirdparty/DiskANN/src/pq_flash_index.cpp index 38180886d..4725ad461 100644 --- a/thirdparty/DiskANN/src/pq_flash_index.cpp +++ b/thirdparty/DiskANN/src/pq_flash_index.cpp @@ -244,13 +244,15 @@ namespace diskann { diskann::alloc_aligned((void **) &scratch.sector_scratch, (_u64) diskann::defaults::MAX_N_SECTOR_READS * read_len_for_node, diskann::defaults::SECTOR_LEN); - diskann::alloc_aligned( + if (use_pq_navigation) { + diskann::alloc_aligned( (void **) &scratch.aligned_pq_coord_scratch, (_u64) diskann::defaults::MAX_GRAPH_DEGREE * (_u64) this->aligned_dim * sizeof(_u8), 256); diskann::alloc_aligned((void **) &scratch.aligned_pqtable_dist_scratch, 256 * (_u64) this->aligned_dim * sizeof(float), 256); + } diskann::alloc_aligned((void **) &scratch.aligned_dist_scratch, (_u64) diskann::defaults::MAX_GRAPH_DEGREE * sizeof(float), 256); diskann::alloc_aligned((void **) &scratch.aligned_query_T, @@ -278,10 +280,12 @@ namespace diskann { thread_data_size += ROUND_UP(sizeof(T) * this->aligned_dim, 256); thread_data_size += ROUND_UP((_u64) diskann::defaults::MAX_N_SECTOR_READS * read_len_for_node, diskann::defaults::SECTOR_LEN); - thread_data_size += ROUND_UP( + if (use_pq_navigation) { + thread_data_size += ROUND_UP( (_u64) diskann::defaults::MAX_GRAPH_DEGREE * (_u64) this->aligned_dim * sizeof(_u8), 256); thread_data_size += ROUND_UP(256 * (_u64) this->aligned_dim * sizeof(float), 256); + } thread_data_size += ROUND_UP((_u64) diskann::defaults::MAX_GRAPH_DEGREE * sizeof(float), 256); thread_data_size += ROUND_UP(this->aligned_dim * sizeof(T), 8 * sizeof(T)); thread_data_size += @@ -705,7 +709,9 @@ namespace diskann { template int PQFlashIndex::load(uint32_t num_threads, const char *index_prefix, - bool load_pq_data) { + bool load_pq_data, + const NavigationMetadata* navigation_metadata) { + use_pq_navigation = load_pq_data; std::string pq_table_bin = get_pq_pivots_filename(std::string(index_prefix)); std::string pq_compressed_vectors = @@ -717,14 +723,21 @@ namespace diskann { std::string centroids_file = get_disk_index_centroids_filename(std::string(disk_index_file)); - size_t pq_file_dim, pq_file_num_centroids; - get_bin_metadata(pq_table_bin, pq_file_num_centroids, pq_file_dim); + size_t pq_file_dim = 0, pq_file_num_centroids = 0; + if (!load_pq_data && navigation_metadata != nullptr) { + if (navigation_metadata->count == 0 || navigation_metadata->dimension == 0 || + navigation_metadata->dimension > std::numeric_limits::max()) { + throw ANNException("Invalid external navigation metadata", -1); + } + pq_file_dim = navigation_metadata->dimension; + } else { + get_bin_metadata(pq_table_bin, pq_file_num_centroids, pq_file_dim); + if (pq_file_num_centroids != 256) { + throw ANNException("Number of PQ centroids is not 256", -1); + } + } this->disk_index_file = disk_index_file; - if (pq_file_num_centroids != 256) { - LOG(ERROR) << "Error. Number of PQ centroids is not 256. Exitting."; - return -1; - } this->data_dim = pq_file_dim; // will reset later if we use PQ on disk @@ -734,10 +747,12 @@ namespace diskann { this->disk_bytes_per_point = this->data_dim * sizeof(T); this->aligned_dim = ROUND_UP(pq_file_dim, 8); - size_t npts_u64, nchunks_u64; + size_t npts_u64 = 0, nchunks_u64 = 0; if (load_pq_data) { diskann::load_bin<_u8>(pq_compressed_vectors, this->data, npts_u64, nchunks_u64); + } else if (navigation_metadata != nullptr) { + npts_u64 = navigation_metadata->count; } else { get_bin_metadata(pq_compressed_vectors, npts_u64, nchunks_u64); const size_t header_size = 2 * sizeof(uint32_t); @@ -758,11 +773,12 @@ namespace diskann { this->num_points = npts_u64; this->n_chunks = nchunks_u64; - pq_table.load_pq_centroid_bin(pq_table_bin.c_str(), nchunks_u64); + if (load_pq_data) { + pq_table.load_pq_centroid_bin(pq_table_bin.c_str(), nchunks_u64); + } - LOG(INFO) << "Loaded PQ centroids and " - << (load_pq_data ? "in-memory compressed vectors" - : "compressed-vector metadata only") + LOG(INFO) << (load_pq_data ? "Loaded resident PQ navigation" + : "Using external navigation metadata") << ". #points: " << num_points << " #dim: " << data_dim << " #aligned_dim: " << aligned_dim << " #chunks: " << n_chunks; @@ -773,6 +789,9 @@ namespace diskann { // giving 0 chunks to make the pq_table infer from the // chunk_offsets file the correct value disk_pq_table.load_pq_centroid_bin(disk_pq_pivots_path.c_str(), 0); + if (disk_pq_table.get_total_dims() != data_dim) { + throw ANNException("SSD PQ and navigation dimensions do not match", -1); + } disk_pq_n_chunks = disk_pq_table.get_num_chunks(); disk_bytes_per_point = disk_pq_n_chunks * @@ -883,7 +902,7 @@ namespace diskann { diskann::load_aligned_bin(centroids_file, centroid_data, num_centroids, tmp_dim, aligned_tmp_dim); - if (aligned_tmp_dim != aligned_dim || num_centroids != num_medoids) { + if (tmp_dim != data_dim || aligned_tmp_dim != aligned_dim || num_centroids != num_medoids) { std::stringstream stream; stream << "Error loading centroids data file. Expected bin format of " "m times data_dim vector of float, where m is number of " @@ -2132,6 +2151,11 @@ namespace diskann { index_mem_size += this->pq_table.get_total_dims() * (sizeof(uint32_t) + sizeof(float)); index_mem_size += (this->pq_table.get_num_chunks() + 1) * sizeof(uint32_t); + if (use_disk_index_pq) { + index_mem_size += disk_pq_table.get_total_dims() * + (256 * sizeof(float) * 2 + sizeof(uint32_t) + sizeof(float)); + index_mem_size += (disk_pq_table.get_num_chunks() + 1) * sizeof(uint32_t); + } // base norms: if (this->metric == diskann::Metric::COSINE) { index_mem_size += sizeof(float) * this->num_points; From 36f1cd09ebc9620fbef12de583da3f70678b6a03 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Mon, 21 Sep 2026 10:17:18 +0000 Subject: [PATCH 06/12] fix: validate DiskANN navigation config and restore iterator IP scores Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 19 +- src/index/diskann/diskann_config.h | 3 +- tests/python/test_diskann_rabitq.py | 2 + tests/ut/test_diskann_ssd_pq.cc | 167 ++++++++++++++++-- .../DiskANN/include/diskann/pq_flash_index.h | 4 + thirdparty/DiskANN/src/aux_utils.cpp | 30 ++-- thirdparty/DiskANN/src/pq_flash_index.cpp | 22 ++- 7 files changed, 191 insertions(+), 56 deletions(-) diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 6828371cc..1d1a66c26 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -125,7 +125,10 @@ class DiskANNIndexNode : public IndexNode { msg = "external DiskANN navigation currently requires FP32 data"; return Status::invalid_args; } - if (base_cfg.emb_list_strategy.has_value() || base_cfg.emb_list_offset_file_path.has_value()) { + // The strategy has a default even for ordinary vectors. Only + // actual embedding-list inputs select embedding-list mode. + if (get_el_metric_type(base_cfg.metric_type.value_or(metric::L2)).has_value() || + base_cfg.emb_list_offset_file_path.has_value()) { msg = "external DiskANN navigation does not support embedding-list mode"; return Status::not_implemented; } @@ -1274,20 +1277,6 @@ class DiskANNRaBitQIndexNode : public DiskANNIndexNode { return StaticCreateConfig(); } - static Status - StaticConfigCheck(const Config& cfg, PARAM_TYPE param_type, std::string& msg) { - const auto status = DiskANNIndexNode::StaticConfigCheck(cfg, param_type, msg); - if (status != Status::success) { - return status; - } - const auto& base_cfg = static_cast(cfg); - if (base_cfg.emb_list_strategy.has_value() || base_cfg.emb_list_offset_file_path.has_value()) { - msg = "DISKANN_RABITQ does not support embedding-list mode"; - return Status::not_implemented; - } - return Status::success; - } - std::string Type() const override { return knowhere::IndexEnum::INDEX_DISKANN_RABITQ; diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index a3208413c..fb6da77ad 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -119,7 +119,8 @@ class DiskANNConfig : public BaseConfig { KNOWHERE_CONFIG_DECLARE_FIELD(disk_pq_dims) .description("the dimension of compressed vectors stored on the ssd, use 0 to store uncompressed data.") .set_default(0) - .for_train(); + .for_train() + .for_static(); KNOWHERE_CONFIG_DECLARE_FIELD(accelerate_build) .description("a flag to enbale fast build.") .set_default(false) diff --git a/tests/python/test_diskann_rabitq.py b/tests/python/test_diskann_rabitq.py index 1854b8e10..960560fe0 100644 --- a/tests/python/test_diskann_rabitq.py +++ b/tests/python/test_diskann_rabitq.py @@ -27,6 +27,8 @@ def test_navigation_roundtrip(tmp_path, metric, kind, codec): disk_pq_dims=0, search_cache_budget_gb=0, search_cache_budget_gb_ratio=0, rbq_bits=4) if codec is not None: config["navigation_codec"] = codec + if kind == "DISKANN_RABITQ" or codec == "RABITQ": + config.pop("pq_code_budget_gb") version = knowhere.GetCurrentVersion() index = knowhere.CreateIndex(kind, version) assert knowhere.Status(index.Build(knowhere.GetNullDataSet(), json.dumps(config))) == knowhere.Status.success diff --git a/tests/ut/test_diskann_ssd_pq.cc b/tests/ut/test_diskann_ssd_pq.cc index 8cfc144c2..09163abf1 100644 --- a/tests/ut/test_diskann_ssd_pq.cc +++ b/tests/ut/test_diskann_ssd_pq.cc @@ -14,9 +14,16 @@ #include "diskann/pq_flash_index.h" #include "filemanager/impl/LocalFileManager.h" #include "knowhere/index/index_factory.h" +#include "knowhere/index/index_static.h" #include "utils.h" namespace { +#ifdef KNOWHERE_WITH_CARDINAL +constexpr const char* kDiskIndexType = "DISKANN_DEPRECATED"; +#else +constexpr const char* kDiskIndexType = "DISKANN"; +#endif + // Decode the actual SSD payload for an independent scalar score reference. class DiskPQReference : public diskann::PQFlashIndex { public: @@ -59,12 +66,69 @@ class DiskPQReference : public diskann::PQFlashIndex { }; } // namespace -TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[diskann][ssd_pq]") { +TEST_CASE("DiskANN public static configuration and raw data capabilities", "[diskann][ssd_pq][diskann_review]") { + const auto version = GenTestVersionList(); const auto metric = GENERATE(std::string("L2"), std::string("IP"), std::string("COSINE")); - const bool external = GENERATE(false, true); + // Native PQ, codec-selected RBQ, and the fixed RBQ alias. + const int codec = GENERATE(0, 1, 2); + const auto* type = codec == 2 ? "DISKANN_RABITQ" : kDiskIndexType; + knowhere::Json config = {{"dim", 16}, {"metric_type", metric}}; + if (codec == 1) + config["navigation_codec"] = "RABITQ"; + std::string error; + using Static = knowhere::IndexStaticFaced; + REQUIRE(Static::ConfigCheck(type, version, config, error) == knowhere::Status::success); + REQUIRE(Static::HasRawData(type, version, config) == (metric != "IP")); + config["emb_list_strategy"] = "tokenann"; + REQUIRE(Static::ConfigCheck(type, version, config, error) == knowhere::Status::success); + for (int disk_bits : {0, 4}) { + config["disk_pq_dims"] = disk_bits; + REQUIRE(Static::HasRawData(type, version, config) == (disk_bits == 0 && metric != "IP")); + } + if (codec != 0) { + config["emb_list_offset_file_path"] = "unused_offsets.bin"; + REQUIRE(Static::ConfigCheck(type, version, config, error) == knowhere::Status::not_implemented); + config.erase("emb_list_offset_file_path"); + config["metric_type"] = "MAX_SIM_" + metric; + REQUIRE(Static::ConfigCheck(type, version, config, error) != knowhere::Status::success); + } +} + +TEST_CASE("DiskANN iterator uses the same score conversion for normal and final batches", + "[diskann][ssd_pq][diskann_review]") { + const auto metric = GENERATE(diskann::Metric::L2, diskann::Metric::INNER_PRODUCT, diskann::Metric::COSINE); + const float base_norm = GENERATE(0.0f, 1.0f, 7.0f); + const float query[] = {3, 4, 0, 0}; + diskann::IteratorWorkspace workspace(query, metric, 8, metric == diskann::Metric::INNER_PRODUCT ? 5 : 4, 1, + 1, 1, 0, base_norm, knowhere::BitsetView()); + for (float score : {0.4f, 2.8f}) { + CAPTURE(metric, base_norm, score); + const float expected = + metric == diskann::Metric::INNER_PRODUCT ? (score / 2 - 1) * (base_norm != 0 ? base_norm * 5 : 1) : score; + workspace.good_pq_res_count = workspace.next_count + workspace.lsearch; + workspace.insert_to_full(10, score); + workspace.move_full_retset_to_backup(); + REQUIRE(workspace.backup_res.size() == 1); + REQUIRE(workspace.backup_res[0].val == Catch::Approx(expected)); + workspace.backup_res.clear(); + workspace.insert_to_full(10, score); + workspace.move_last_full_retset_to_backup(); + REQUIRE(workspace.backup_res.size() == 1); + REQUIRE(workspace.backup_res[0].val == Catch::Approx(expected)); + REQUIRE(workspace.full_retset.empty()); + workspace.backup_res.clear(); + } +} + +TEST_CASE("DiskANN SSD scores are independent of navigation and cache", "[diskann][ssd_pq][diskann_review]") { + const auto metric = GENERATE(std::string("L2"), std::string("IP"), std::string("COSINE")); + const int codec = GENERATE(0, 1, 2); + const bool external = codec != 0; + const int disk_pq_dims = GENERATE(0, 4); const auto version = GenTestVersionList(); constexpr size_t rows = 300, dim = 16, nq = 3, k = 10; - const auto dir = std::filesystem::current_path() / ("ssd_pq_regression_" + metric + (external ? "_rbq" : "_pq")); + const auto dir = std::filesystem::current_path() / + ("ssd_pq_regression_" + metric + "_" + std::to_string(codec) + "_" + std::to_string(disk_pq_dims)); std::filesystem::create_directories(dir); const auto prefix = (dir / "index").string(); const auto raw = (dir / "base.bin").string(); @@ -80,11 +144,7 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis file.write(reinterpret_cast(xb), rows * dim * sizeof(float)); } auto pack = knowhere::Pack(std::shared_ptr(std::make_shared())); -#ifdef KNOWHERE_WITH_CARDINAL - const char* index_type = "DISKANN_DEPRECATED"; -#else - const char* index_type = "DISKANN"; -#endif + const char* index_type = codec == 2 ? "DISKANN_RABITQ" : kDiskIndexType; knowhere::Json config = {{"dim", dim}, {"metric_type", metric}, {"index_prefix", prefix}, @@ -93,20 +153,37 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis {"search_list_size", 100}, {"pq_code_budget_gb", 0.001}, {"build_dram_budget_gb", 1.0}, - {"disk_pq_dims", 4}, + {"disk_pq_dims", disk_pq_dims}, {"search_cache_budget_gb", 0}, {"search_cache_budget_gb_ratio", 0}}; if (external) { config["navigation_codec"] = "RABITQ"; config["rbq_bits"] = 4; + // Omitted and explicit-zero navigation budgets must both work, + // regardless of whether SSD PQ is enabled. + config.erase("pq_code_budget_gb"); + if (codec == 2) { + config["pq_code_budget_gb"] = 0; + config["pq_code_budget_gb_ratio"] = 0; + } } auto built = knowhere::IndexFactory::Instance().Create(index_type, version, pack).value(); + if (metric == "L2" && disk_pq_dims == 0) { + auto invalid = config; + invalid["build_dram_budget_gb"] = 0; + REQUIRE(built.Build(nullptr, invalid) != knowhere::Status::success); + if (!external) { + invalid = config; + invalid["pq_code_budget_gb"] = 0; + REQUIRE(built.Build(nullptr, invalid) != knowhere::Status::success); + } + } REQUIRE(built.Build(nullptr, config) == knowhere::Status::success); knowhere::BinarySet binary; REQUIRE(built.Serialize(binary) == knowhere::Status::success); REQUIRE(std::filesystem::exists(prefix + "_pq_compressed.bin") == !external); REQUIRE(std::filesystem::exists(prefix + "_pq_pivots.bin") == !external); - REQUIRE(std::filesystem::exists(prefix + "_disk.index_pq_pivots.bin")); + REQUIRE(std::filesystem::exists(prefix + "_disk.index_pq_pivots.bin") == (disk_pq_dims > 0)); const auto dm = metric == "L2" ? diskann::Metric::L2 : metric == "IP" ? diskann::Metric::INNER_PRODUCT : diskann::Metric::COSINE; @@ -117,9 +194,14 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis config["use_bfs_cache"] = true; config["warm_up"] = true; REQUIRE(index.Deserialize(binary, config) == knowhere::Status::success); - REQUIRE_FALSE(index.HasRawData(metric)); + const bool has_raw = disk_pq_dims == 0 && metric != "IP"; + REQUIRE(index.HasRawData(metric) == has_raw); + REQUIRE(knowhere::IndexStaticFaced::HasRawData(index_type, version, config) == has_raw); int64_t id = 0; - REQUIRE_FALSE(index.GetVectorByIds(knowhere::GenIdsDataSet(1, &id)).has_value()); + if (disk_pq_dims > 0) + REQUIRE_FALSE(index.GetVectorByIds(knowhere::GenIdsDataSet(1, &id)).has_value()); + else if (has_raw) + REQUIRE(index.GetVectorByIds(knowhere::GenIdsDataSet(1, &id)).has_value()); DiskPQReference reference(dm); const DiskPQReference::NavigationMetadata metadata{rows, dim + (metric == "IP" ? 1 : 0)}; REQUIRE(reference.load(1, prefix.c_str(), !external, external ? &metadata : nullptr) == 0); @@ -130,6 +212,19 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis std::iota(ids.begin(), ids.end(), 0); reference.load_cache_list(ids); } + const auto expected_score = [&](unsigned id, const float* query) { + if (disk_pq_dims > 0) + return reference.Score(prefix, id, query, dim); + double dot = 0, nq = 0, nb = 0, l2 = 0; + for (size_t j = 0; j < dim; ++j) { + const double b = xb[id * dim + j], q = query[j]; + dot += b * q; + nq += q * q; + nb += b * b; + l2 += (q - b) * (q - b); + } + return float(metric == "L2" ? l2 : metric == "IP" ? dot : dot / std::sqrt(nq * nb)); + }; knowhere::Json search = {{"metric_type", metric}, {"k", k}, {"search_list_size", 128}, {"beamwidth", 4}}; for (size_t filtered : {size_t(0), size_t(285)}) { std::vector mask((rows + 7) / 8, 0); @@ -141,8 +236,8 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis const auto offset = q * k + j; const auto label = result.value()->GetIds()[offset]; REQUIRE(label >= static_cast(filtered)); - REQUIRE(label < rows); - const auto expected = reference.Score(prefix, label, xq + q * dim, dim); + REQUIRE(label < int64_t(rows)); + const auto expected = expected_score(label, xq + q * dim); REQUIRE(result.value()->GetDistance()[offset] == Catch::Approx(expected).epsilon(0.0002).margin(0.0002)); } @@ -153,8 +248,48 @@ TEST_CASE("DiskANN SSD PQ scores are independent of navigation and cache", "[dis float distances[4]; reference.calc_dist_by_ids(xq, ids, 4, distances); for (size_t j = 0; j < 4; ++j) - REQUIRE(distances[j] == - Catch::Approx(reference.Score(prefix, ids[j], xq, dim)).epsilon(0.0002).margin(0.0002)); + REQUIRE(distances[j] == Catch::Approx(expected_score(ids[j], xq)).epsilon(0.0002).margin(0.0002)); + if (!external && metric == "IP") { + // Small L exercises ordinary batches and the tail; L > rows forces + // exhaustion before the ordinary batch threshold is reached. + for (int iterator_l : {2, int(rows + 1)}) { + auto iterator_config = search; + iterator_config["search_list_size"] = iterator_l; + auto result = index.AnnIterator(queries, iterator_config, knowhere::BitsetView(), false); + REQUIRE(result.has_value()); + for (size_t q = 0; q < nq; ++q) { + auto& iterator = result.value()[q]; + std::vector seen(rows, false); + size_t count = 0; + while (iterator->HasNext().value()) { + const auto [id, score] = iterator->Next().value(); + REQUIRE(id >= 0); + REQUIRE(id < int64_t(rows)); + REQUIRE_FALSE(seen[id]); + seen[id] = true; + ++count; + REQUIRE(score == + Catch::Approx(expected_score(id, xq + q * dim)).epsilon(0.0002).margin(0.0002)); + } + // Exhausting an approximate graph traversal need not + // enumerate disconnected/unreachable base vectors. + REQUIRE(count > k); + REQUIRE(count <= rows); + REQUIRE_FALSE(iterator->HasNext().value()); + } + } + // Explicitly observe the low-level exhausted-candidate path, + // rather than relying only on the public iterator's buffering. + auto workspace = reference.getIteratorWorkspace(xq, rows + 1, 4, 0, knowhere::BitsetView()); + reference.getIteratorNextBatch(workspace.get()); + REQUIRE_FALSE(workspace->has_candidates()); + REQUIRE(workspace->full_retset.empty()); + REQUIRE(workspace->next_count == 0); + REQUIRE(workspace->backup_res.size() > k); + for (const auto& item : workspace->backup_res) { + REQUIRE(-item.val == Catch::Approx(expected_score(item.id, xq)).epsilon(0.0002).margin(0.0002)); + } + } } std::filesystem::remove_all(dir); } diff --git a/thirdparty/DiskANN/include/diskann/pq_flash_index.h b/thirdparty/DiskANN/include/diskann/pq_flash_index.h index 62ef2f3eb..c3ca030b3 100644 --- a/thirdparty/DiskANN/include/diskann/pq_flash_index.h +++ b/thirdparty/DiskANN/include/diskann/pq_flash_index.h @@ -126,6 +126,10 @@ namespace diskann { void pop_pq_retset(); + // Convert the internal SSD score to the smaller-is-better iterator + // output. The Knowhere iterator wrapper applies the final IP sign flip. + float output_distance(float distance) const; + void move_full_retset_to_backup(); void move_last_full_retset_to_backup(); diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index 1517ab2f7..d8efc0a24 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -1688,11 +1688,14 @@ template unsigned R = config.max_degree; unsigned L = config.search_list_size; - double pq_code_size_limit = get_memory_budget(config.pq_code_size_gb); - if (pq_code_size_limit <= 0) { - LOG(ERROR) << "Insufficient memory budget (or string was not in right " - "format). Should be > 0."; - return -1; + double pq_code_size_limit = 0; + if (config.use_pq_navigation) { + pq_code_size_limit = get_memory_budget(config.pq_code_size_gb); + if (pq_code_size_limit <= 0) { + LOG(ERROR) << "Insufficient memory budget (or string was not in right " + "format). Should be > 0."; + return -1; + } } double indexing_ram_budget = config.index_mem_gb; if (indexing_ram_budget <= 0) { @@ -1736,16 +1739,6 @@ template << " Indexing ram budget: " << indexing_ram_budget << "(GiB)"; - size_t num_pq_chunks = - (size_t) (std::floor)(_u64(pq_code_size_limit / points_num)); - - num_pq_chunks = num_pq_chunks <= 0 ? 1 : num_pq_chunks; - num_pq_chunks = num_pq_chunks > dim ? dim : num_pq_chunks; - num_pq_chunks = num_pq_chunks > diskann::defaults::MAX_PQ_CHUNKS ? diskann::defaults::MAX_PQ_CHUNKS : num_pq_chunks; - - LOG_KNOWHERE_INFO_ << "Compressing " << dim << "-dimensional data into " - << num_pq_chunks << " bytes per vector."; - size_t train_size = 0, train_dim = 0; std::unique_ptr train_data = nullptr; @@ -1777,6 +1770,13 @@ template disk_pq_pivots_path, disk_pq_compressed_vectors_path); } if (config.use_pq_navigation) { + size_t num_pq_chunks = + (size_t) (std::floor)(_u64(pq_code_size_limit / points_num)); + num_pq_chunks = num_pq_chunks <= 0 ? 1 : num_pq_chunks; + num_pq_chunks = num_pq_chunks > dim ? dim : num_pq_chunks; + num_pq_chunks = num_pq_chunks > diskann::defaults::MAX_PQ_CHUNKS ? diskann::defaults::MAX_PQ_CHUNKS : num_pq_chunks; + LOG_KNOWHERE_INFO_ << "Compressing " << dim << "-dimensional data into " + << num_pq_chunks << " bytes per vector."; LOG_KNOWHERE_DEBUG_ << "Training data loaded of size " << train_size; // don't translate data to make zero mean for PQ compression. We must not diff --git a/thirdparty/DiskANN/src/pq_flash_index.cpp b/thirdparty/DiskANN/src/pq_flash_index.cpp index 4725ad461..4deeb0f99 100644 --- a/thirdparty/DiskANN/src/pq_flash_index.cpp +++ b/thirdparty/DiskANN/src/pq_flash_index.cpp @@ -167,18 +167,22 @@ namespace diskann { } } + template + float IteratorWorkspace::output_distance(float distance) const { + if (metric == diskann::Metric::INNER_PRODUCT) { + distance = distance / 2.0f - 1.0f; + if (max_base_norm != 0) { + distance *= (max_base_norm * query_norm); + } + } + return distance; + } + template void IteratorWorkspace::move_full_retset_to_backup() { if (is_good_pq_enough() && !full_retset.empty()) { auto &nbr = full_retset.top(); - auto dist = nbr.distance; - if (metric == diskann::Metric::INNER_PRODUCT) { - dist = dist / 2.0f - 1.0f; - if (max_base_norm != 0) { - dist *= (max_base_norm * query_norm); - } - } - backup_res.emplace_back(nbr.id, dist); + backup_res.emplace_back(nbr.id, output_distance(nbr.distance)); full_retset.pop(); next_count++; } @@ -188,7 +192,7 @@ namespace diskann { void IteratorWorkspace::move_last_full_retset_to_backup() { while (!full_retset.empty()) { auto &nbr = full_retset.top(); - backup_res.emplace_back(nbr.id, nbr.distance); + backup_res.emplace_back(nbr.id, output_distance(nbr.distance)); full_retset.pop(); } } From 48726b9cad9072a680c9f4d82e24debc43883343 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Tue, 22 Sep 2026 09:20:29 +0000 Subject: [PATCH 07/12] refactor: use probabilistic RaBitQ navigation without experimental switches Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 6 +- src/index/diskann/diskann_config.h | 19 ------ src/index/diskann/rabitq_store.cc | 16 ++--- src/index/diskann/rabitq_store.h | 2 +- tests/ut/test_diskann.cc | 94 +++++++++++++----------------- 5 files changed, 48 insertions(+), 89 deletions(-) diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 1d1a66c26..95a7aa891 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -817,10 +817,8 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr "resident"; } LOG_KNOWHERE_INFO_ << "Use bfs to generate cache list"; - if (TryDiskANNCall([&]() { - pq_flash_index_->cache_bfs_levels(num_nodes_to_cache, node_list, - prep_conf.bfs_cache_seed.value()); - }) != Status::success) { + if (TryDiskANNCall([&]() { pq_flash_index_->cache_bfs_levels(num_nodes_to_cache, node_list); }) != + Status::success) { LOG_KNOWHERE_ERROR_ << "Failed to generate bfs cache for DiskANN."; return Status::diskann_inner_error; } diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index fb6da77ad..0fcd71c42 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -71,9 +71,6 @@ class DiskANNConfig : public BaseConfig { // cached the nodes on the search paths; 2. do bfs from the entry point and cache them. The first method is suitable // for TopK query heavy circumstances and the second one performed better in range search. CFG_BOOL use_bfs_cache; - // Optional deterministic seed for BFS cache selection. A negative value - // preserves the existing random selection behavior. - CFG_INT bfs_cache_seed; // The beamwidth to be used for search. This is the maximum number of IO requests each query will issue per // iteration of search code. Larger beamwidth will result in fewer IO round-trips per query but might result in // slightly higher total number of IO requests to SSD per query. For the highest query throughput with a fixed SSD @@ -147,11 +144,6 @@ class DiskANNConfig : public BaseConfig { .description("should bfs strategy to cache nodes.") .set_default(false) .for_deserialize(); - KNOWHERE_CONFIG_DECLARE_FIELD(bfs_cache_seed) - .description("seed for deterministic bfs cache selection; -1 uses a random seed.") - .set_default(-1) - .set_range(-1, std::numeric_limits::max()) - .for_deserialize(); KNOWHERE_CONFIG_DECLARE_FIELD(beamwidth) .description("the maximum number of IO requests each query will issue per iteration of search code.") .set_default(diskann::defaults::DEFAULT_DISKANN_BEAMWIDTH) @@ -214,7 +206,6 @@ class DiskANNNavigationConfig : public DiskANNConfig { CFG_STRING navigation_codec; CFG_INT rbq_bits; CFG_INT rbq_bits_query; - CFG_STRING rbq_refine_mode; KNOWHERE_DECLARE_CONFIG(DiskANNNavigationConfig) { KNOWHERE_CONFIG_DECLARE_FIELD(navigation_codec) @@ -234,12 +225,6 @@ class DiskANNNavigationConfig : public DiskANNConfig { .set_default(4) .set_range(0, 8) .for_search(); - KNOWHERE_CONFIG_DECLARE_FIELD(rbq_refine_mode) - .description( - "RaBitQ multi-bit refinement mode: probabilistic enables error-window pruning; full always " - "computes the complete RaBitQ distance") - .set_default("probabilistic") - .for_search(); } Status @@ -263,10 +248,6 @@ class DiskANNNavigationConfig : public DiskANNConfig { if (database_bits < 1 || database_bits > 9) { return HandleError(err_msg, "DISKANN_RABITQ supports rbq_bits in [1, 9]", Status::invalid_args); } - const auto refine_mode = rbq_refine_mode.value_or("probabilistic"); - if (refine_mode != "probabilistic" && refine_mode != "full") { - return HandleError(err_msg, "rbq_refine_mode must be probabilistic or full", Status::invalid_args); - } return Status::success; } }; diff --git a/src/index/diskann/rabitq_store.cc b/src/index/diskann/rabitq_store.cc index 220fbb3f4..5027f6e11 100644 --- a/src/index/diskann/rabitq_store.cc +++ b/src/index/diskann/rabitq_store.cc @@ -94,10 +94,9 @@ apply_rotation_single_query(const faiss::RandomRotationMatrix* rotation, const f class RaBitQNavigationDistanceComputer final : public diskann::NavigationDistanceComputer { public: RaBitQNavigationDistanceComputer(const faiss::RandomRotationMatrix* rotation, const faiss::IndexRaBitQ* rabitq, - bool probabilistic_refinement, uint8_t query_bits) + uint8_t query_bits) : rotation_(rotation), rabitq_(rabitq), - probabilistic_refinement_(probabilistic_refinement), distance_computer_(rabitq->get_quantized_distance_computer(query_bits, false)), rabitq_distance_computer_(dynamic_cast(distance_computer_.get())) { if (rabitq_distance_computer_ == nullptr) { @@ -116,7 +115,7 @@ class RaBitQNavigationDistanceComputer final : public diskann::NavigationDistanc void compute_distances(const unsigned* ids, _u64 n_ids, float* distances, float threshold, bool threshold_valid, diskann::QueryStats* stats) override { - const bool can_prune = probabilistic_refinement_ && threshold_valid && rabitq_->rabitq.nb_bits > 1; + const bool can_prune = threshold_valid && rabitq_->rabitq.nb_bits > 1; if (!can_prune) { _u64 i = 0; for (; i + 4 <= n_ids; i += 4) { @@ -193,7 +192,6 @@ class RaBitQNavigationDistanceComputer final : public diskann::NavigationDistanc private: const faiss::RandomRotationMatrix* rotation_; const faiss::IndexRaBitQ* rabitq_; - const bool probabilistic_refinement_; std::unique_ptr distance_computer_; faiss::RaBitQDistanceComputer* rabitq_distance_computer_; std::unique_ptr transformed_query_; @@ -336,23 +334,19 @@ RaBitQStore::CreateDistanceComputer(const DiskANNConfig& config) const { if (query_metric != metric::L2 && query_metric != metric::IP && query_metric != metric::COSINE) { throw std::invalid_argument("RaBitQ navigation supports L2, IP and COSINE"); } - const auto mode = navigation->rbq_refine_mode.value_or("probabilistic"); - if (mode != "probabilistic" && mode != "full") { - throw std::invalid_argument("invalid RaBitQ refinement mode"); - } const auto qb = navigation->rbq_bits_query.value_or(4); if (qb < 0 || qb > 8) { throw std::invalid_argument("RaBitQ query bits must be in [0, 8]"); } - return CreateDistanceComputer(mode == "probabilistic", static_cast(qb)); + return CreateDistanceComputer(static_cast(qb)); } std::unique_ptr -RaBitQStore::CreateDistanceComputer(bool probabilistic_refinement, uint8_t query_bits) const { +RaBitQStore::CreateDistanceComputer(uint8_t query_bits) const { if (query_bits > 8) { throw std::invalid_argument("RaBitQ query bits must be in [0, 8]"); } - return std::make_unique(rotation_, rabitq_, probabilistic_refinement, query_bits); + return std::make_unique(rotation_, rabitq_, query_bits); } int64_t diff --git a/src/index/diskann/rabitq_store.h b/src/index/diskann/rabitq_store.h index bcdc16319..b1d8fa9a1 100644 --- a/src/index/diskann/rabitq_store.h +++ b/src/index/diskann/rabitq_store.h @@ -41,7 +41,7 @@ class RaBitQStore final : public NavigationStore { operator=(const RaBitQStore&) = delete; std::unique_ptr - CreateDistanceComputer(bool probabilistic_refinement, uint8_t query_bits = 4) const; + CreateDistanceComputer(uint8_t query_bits = 4) const; std::unique_ptr CreateDistanceComputer(const DiskANNConfig& config) const override; diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index bca94cf04..7c0371904 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -646,18 +646,12 @@ TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { invalid["search_cache_budget_gb"] = 0.01; check_train_config(invalid, knowhere::Status::success); - auto check_search_mode = [&](const std::string& mode, knowhere::Status expected) { - auto cfg = knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, - version); - knowhere::Json json = {{"dim", kDim}, {"metric_type", knowhere::metric::L2}, - {"k", kK}, {"search_list_size", 128}, - {"beamwidth", 8}, {"rbq_refine_mode", mode}}; - std::string msg; - REQUIRE(knowhere::Config::Load(*cfg, json, knowhere::PARAM_TYPE::SEARCH, &msg) == expected); - }; - check_search_mode("probabilistic", knowhere::Status::success); - check_search_mode("full", knowhere::Status::success); - check_search_mode("invalid", knowhere::Status::invalid_args); + // Experimental controls are not part of either public configuration. + for (const auto* type : {kNativeDiskANN, knowhere::IndexEnum::INDEX_DISKANN_RABITQ}) { + auto cfg = knowhere::IndexStaticFaced::CreateConfig(type, version); + REQUIRE(cfg->__DICT__.count("rbq_refine_mode") == 0); + REQUIRE(cfg->__DICT__.count("bfs_cache_seed") == 0); + } for (const int qb : {-1, 0, 4, 8, 9}) { auto cfg = knowhere::IndexStaticFaced::CreateConfig(knowhere::IndexEnum::INDEX_DISKANN_RABITQ, version); @@ -696,7 +690,7 @@ TEST_CASE("Test DISKANN_RABITQ probabilistic refinement", "[diskann][rabitq]") { } knowhere::RaBitQStore store(sidecar_path); - auto distance_computer = store.CreateDistanceComputer(true, 0); + auto distance_computer = store.CreateDistanceComputer(0); distance_computer->set_query(base); std::vector ids(rows); std::iota(ids.begin(), ids.end(), 0); @@ -718,14 +712,15 @@ TEST_CASE("Test DISKANN_RABITQ probabilistic refinement", "[diskann][rabitq]") { REQUIRE(pruned_stats.n_cmps_saved == rows); REQUIRE(std::all_of(distances.begin(), distances.end(), [](float distance) { return std::isinf(distance); })); - auto full_distance_computer = store.CreateDistanceComputer(false, 0); - full_distance_computer->set_query(base); - diskann::QueryStats explicit_full_stats; - full_distance_computer->compute_distances(ids.data(), rows, distances.data(), 0.0f, true, &explicit_full_stats); - REQUIRE(explicit_full_stats.n_approx_estimates == 0); - REQUIRE(explicit_full_stats.n_approx_refinements == rows); - REQUIRE(explicit_full_stats.n_approx_pruned == 0); - REQUIRE(std::all_of(distances.begin(), distances.end(), [](float distance) { return std::isfinite(distance); })); + // A valid, unbounded threshold refines every candidate; there is no + // separate full-distance search mode. + diskann::QueryStats unbounded_stats; + distance_computer->compute_distances(ids.data(), rows, distances.data(), std::numeric_limits::infinity(), + true, &unbounded_stats); + REQUIRE(unbounded_stats.n_approx_estimates == rows); + REQUIRE(unbounded_stats.n_approx_refinements == rows); + REQUIRE(unbounded_stats.n_approx_pruned == 0); + REQUIRE(std::all_of(distances.begin(), distances.end(), [](float d) { return std::isfinite(d); })); fs::remove_all(refinement_dir); } @@ -769,12 +764,12 @@ TEST_CASE("DiskANN RaBitQ shares Faiss codes and request-local query bits", "[di } const unsigned ids[] = {0, 1, 3, 5, 7, 9, 12}; std::vector initial(7), after(7); - auto stable = store.CreateDistanceComputer(false, 0); + auto stable = store.CreateDistanceComputer(0); stable->set_query(x); stable->compute_distances(ids, 7, initial.data(), 0, false, nullptr); for (const uint8_t qb : {0, 4, 8}) { CAPTURE(qb); - auto adapter = store.CreateDistanceComputer(false, qb); + auto adapter = store.CreateDistanceComputer(qb); std::unique_ptr native( rbq->get_quantized_distance_computer(qb, false)); adapter->set_query(x); @@ -800,7 +795,7 @@ TEST_CASE("DiskANN RaBitQ shares Faiss codes and request-local query bits", "[di stable->compute_distances(ids, 7, after.data(), 0, false, nullptr); REQUIRE(initial == after); REQUIRE(rbq->qb == 4); - REQUIRE_THROWS(store.CreateDistanceComputer(false, 9)); + REQUIRE_THROWS(store.CreateDistanceComputer(9)); } } fs::remove_all(dir); @@ -822,8 +817,8 @@ TEST_CASE("DiskANN RaBitQ batches preserve scalar decisions", "[diskann][rabitq] std::iota(ids.begin(), ids.end(), 0); for (uint8_t qb = 0; qb <= 8; ++qb) { CAPTURE(dim, bits, qb); - auto scalar = store.CreateDistanceComputer(true, qb); - auto batch = store.CreateDistanceComputer(true, qb); + auto scalar = store.CreateDistanceComputer(qb); + auto batch = store.CreateDistanceComputer(qb); scalar->set_query(x + 18 * dim); batch->set_query(x + 18 * dim); std::array full{}; @@ -856,7 +851,7 @@ TEST_CASE("DiskANN RaBitQ batches preserve scalar decisions", "[diskann][rabitq] // Concurrent callers have independent query transforms and query bits. std::array, 2> sequential, concurrent; const auto run_query = [&](size_t q, std::vector& result) { - auto dc = store.CreateDistanceComputer(false, q ? 8 : 0); + auto dc = store.CreateDistanceComputer(q ? 8 : 0); dc->set_query(x + (17 + q) * dim); result.resize(ids.size()); dc->compute_distances(ids.data(), ids.size(), result.data(), 0, false, nullptr); @@ -945,7 +940,6 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { // RaBitQ must safely force BFS even when the default sample-query // cache mode is requested, because its navigation PQ is not resident. deserialize_json["use_bfs_cache"] = false; - deserialize_json["bfs_cache_seed"] = 42; } knowhere::Json search_json = {{"dim", kDim}, {"metric_type", knowhere::metric::L2}, @@ -993,7 +987,7 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { } } failing; knowhere::RaBitQStore store(rabitq_prefix + "_rabitq.index"); - auto scorer = store.CreateDistanceComputer(false, 0); + auto scorer = store.CreateDistanceComputer(0); std::copy_n(static_cast(query_ds->GetTensor()), kDim, query.data()); for (bool fail_set : {true, false}) { failing.fail_set = fail_set; @@ -1047,14 +1041,9 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { } } } - search_json["rbq_refine_mode"] = "full"; - auto full_result = index.Search(query_ds, search_json, nullptr); - REQUIRE(full_result.has_value()); auto ground_truth = knowhere::BruteForce::Search(base_ds, query_ds, search_json, nullptr); REQUIRE(ground_truth.has_value()); REQUIRE(GetKNNRecall(*ground_truth.value(), *result.value()) > 0.5f); - REQUIRE(GetKNNRecall(*ground_truth.value(), *full_result.value()) > 0.5f); - search_json["rbq_refine_mode"] = "probabilistic"; std::vector empty_bitset_data((kNumRows + 7) / 8, 0); auto empty_bitset_result = @@ -1203,27 +1192,24 @@ TEST_CASE("DiskANN RaBitQ cosine uses normalized navigation and original SSD vec auto exact = knowhere::BruteForce::Search(base, query, search, nullptr); REQUIRE(exact.has_value()); for (const int qb : {0, 4, 8}) { - for (const auto* mode : {"full", "probabilistic"}) { - search["rbq_bits_query"] = qb; - search["rbq_refine_mode"] = mode; - auto result = restored.Search(query, search, nullptr); - REQUIRE(result.has_value()); - REQUIRE(GetKNNRecall(*exact.value(), *result.value()) > 0.8f); - for (uint32_t i = 0; i < kNumQueries; ++i) { - for (uint32_t j = 0; j < kK; ++j) { - const auto offset = i * kK + j; - const auto id = result.value()->GetIds()[offset]; - REQUIRE(id >= 0); - REQUIRE(id < kNumRows); - double dot = 0, norm_x = 0, norm_q = 0; - for (uint32_t k = 0; k < kDim; ++k) { - dot += double(xb[id * kDim + k]) * xq[i * kDim + k]; - norm_x += double(xb[id * kDim + k]) * xb[id * kDim + k]; - norm_q += double(xq[i * kDim + k]) * xq[i * kDim + k]; - } - const double score = norm_x > 0 && norm_q > 0 ? dot / std::sqrt(norm_x * norm_q) : 0; - REQUIRE(std::abs(result.value()->GetDistance()[offset] - score) < 1e-5); + search["rbq_bits_query"] = qb; + auto result = restored.Search(query, search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(GetKNNRecall(*exact.value(), *result.value()) > 0.8f); + for (uint32_t i = 0; i < kNumQueries; ++i) { + for (uint32_t j = 0; j < kK; ++j) { + const auto offset = i * kK + j; + const auto id = result.value()->GetIds()[offset]; + REQUIRE(id >= 0); + REQUIRE(id < kNumRows); + double dot = 0, norm_x = 0, norm_q = 0; + for (uint32_t k = 0; k < kDim; ++k) { + dot += double(xb[id * kDim + k]) * xq[i * kDim + k]; + norm_x += double(xb[id * kDim + k]) * xb[id * kDim + k]; + norm_q += double(xq[i * kDim + k]) * xq[i * kDim + k]; } + const double score = norm_x > 0 && norm_q > 0 ? dot / std::sqrt(norm_x * norm_q) : 0; + REQUIRE(std::abs(result.value()->GetDistance()[offset] - score) < 1e-5); } } } From bac91ab0a1c7f9a524dea5dcca9062a1caba268b Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Tue, 22 Sep 2026 09:26:58 +0000 Subject: [PATCH 08/12] refactor: scope DiskANN preprocessing files and build publication Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/build_files.h | 53 ++++++ src/index/diskann/diskann.cc | 52 +++--- tests/ut/test_diskann.cc | 72 +++++++++ .../DiskANN/include/diskann/aux_utils.h | 43 ++++- thirdparty/DiskANN/src/aux_utils.cpp | 151 +++++++++++++----- 5 files changed, 298 insertions(+), 73 deletions(-) create mode 100644 src/index/diskann/build_files.h diff --git a/src/index/diskann/build_files.h b/src/index/diskann/build_files.h new file mode 100644 index 000000000..1843116a9 --- /dev/null +++ b/src/index/diskann/build_files.h @@ -0,0 +1,53 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include + +#include "filemanager/FileManager.h" +#include "knowhere/log.h" + +namespace knowhere { +// Publication is separate from ownership of local build files. Failed uploads +// can have partial effects, so roll back attempted registrations as well. +class DiskANNBuildRegistration { + public: + explicit DiskANNBuildRegistration(milvus::FileManager& manager) : manager_(manager) { + } + DiskANNBuildRegistration(const DiskANNBuildRegistration&) = delete; + DiskANNBuildRegistration& + operator=(const DiskANNBuildRegistration&) = delete; + ~DiskANNBuildRegistration() { + if (!committed_) { + for (auto it = attempted_.rbegin(); it != attempted_.rend(); ++it) { + try { + if (!manager_.RemoveFile(*it)) { + LOG_KNOWHERE_WARNING_ << "Failed to roll back DiskANN registration: " << *it; + } + } catch (const std::exception& e) { + LOG_KNOWHERE_WARNING_ << "Failed to roll back DiskANN registration: " << e.what(); + } + } + } + } + bool + Add(const std::string& path) { + const auto exists = manager_.IsExisted(path); + if (!exists.has_value() || exists.value()) { + return false; + } + attempted_.push_back(path); + return manager_.AddFile(path); + } + void + Commit() noexcept { + committed_ = true; + } + + private: + milvus::FileManager& manager_; + std::vector attempted_; + bool committed_ = false; +}; +} // namespace knowhere diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 95a7aa891..86d06a18c 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -25,6 +25,7 @@ #include "diskann/pq_flash_index.h" #include "filemanager/FileManager.h" #include "fmt/core.h" +#include "index/diskann/build_files.h" #include "index/diskann/diskann_config.h" #include "index/diskann/navigation_store.h" #include "knowhere/comp/index_param.h" @@ -513,8 +514,6 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::Build(const DataSetPtr dataset, std::shared_ptr(num_nodes_to_cache), build_conf.shuffle_build.value()}; - const bool navigation_uses_preprocessed_base = external_navigation && need_norm; - diskann_internal_build_config.keep_preprocessed_base = navigation_uses_preprocessed_base; diskann_internal_build_config.use_pq_navigation = !external_navigation; + std::unique_ptr context; RETURN_IF_ERROR(TryDiskANNCall([&]() { - int res = diskann::build_disk_index(diskann_internal_build_config); - if (res != 0) + context = diskann::prepare_build_context(diskann_internal_build_config); + for (const auto& path : GetNecessaryFilenames(index_prefix_, need_norm, true, true, !external_navigation)) { + context->own_output(path); + } + for (const auto& path : GetOptionalFilenames(index_prefix_)) { + context->own_output(path); + } + for (const auto& path : NavigationFiles(build_conf, index_prefix_)) { + context->own_output(path); + } + const int res = diskann::build_disk_index(diskann_internal_build_config, *context); + if (res != 0) { throw diskann::ANNException("diskann::build_disk_index returned non-zero value: " + std::to_string(res), -1); - })); - - if (external_navigation) { - try { - const auto sidecar_source = - navigation_uses_preprocessed_base ? index_prefix_ + "_prepped_base.bin" : data_path; - BuildNavigationStore(build_conf, sidecar_source, index_prefix_); - if (navigation_uses_preprocessed_base) { - std::error_code error; - std::filesystem::remove(sidecar_source, error); - } - } catch (const std::exception& e) { - if (navigation_uses_preprocessed_base) { - std::error_code error; - std::filesystem::remove(index_prefix_ + "_prepped_base.bin", error); - } - LOG_KNOWHERE_ERROR_ << "Failed to build DiskANN navigation sidecar: " << e.what(); - return Status::diskann_inner_error; } - } + BuildNavigationStore(build_conf, context->prepared_source, index_prefix_); + })); // Add file to the file manager + DiskANNBuildRegistration registration(*file_manager_); for (auto& filename : GetNecessaryFilenames(index_prefix_, need_norm, true, true, !external_navigation)) { - if (!AddFile(filename)) { + if (!registration.Add(filename)) { LOG_KNOWHERE_ERROR_ << "Failed to add file " << filename << "."; return Status::disk_file_error; } } for (auto& filename : GetOptionalFilenames(index_prefix_)) { - if (file_exists(filename) && !AddFile(filename)) { + if (file_exists(filename) && !registration.Add(filename)) { LOG_KNOWHERE_ERROR_ << "Failed to add file " << filename << "."; return Status::disk_file_error; } } for (const auto& sidecar_path : NavigationFiles(build_conf, index_prefix_)) { - if (!AddFile(sidecar_path)) { + if (!registration.Add(sidecar_path)) { LOG_KNOWHERE_ERROR_ << "Failed to add file " << sidecar_path << "."; return Status::disk_file_error; } } + registration.Commit(); + context->commit_outputs(); + count_.store(count); + dim_.store(dim); is_prepared_.store(false); return Status::success; } diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index 7c0371904..f2498ac4c 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -39,6 +39,7 @@ #include "faiss/utils/rabitq_simd.h" #include "filemanager/FileManager.h" #include "filemanager/impl/LocalFileManager.h" +#include "index/diskann/build_files.h" #include "index/diskann/diskann_config.h" #include "index/diskann/rabitq_store.h" #include "knowhere/comp/brute_force.h" @@ -103,6 +104,77 @@ constexpr float kIpRangeAp = 0.9; constexpr float kCosineRangeAp = 0.9; } // namespace +TEST_CASE("DiskANN prepared build files have scoped ownership", "[diskann][build_context]") { + fs::remove_all(kDir); + fs::create_directories(kDir); + constexpr size_t rows = 32, dim = 16; + auto data = GenDataSet(rows, dim, 17); + WriteRawDataToDisk(kRawDataPath, static_cast(data->GetTensor()), rows, dim); + for (const auto metric : {diskann::Metric::L2, diskann::Metric::INNER_PRODUCT, diskann::Metric::COSINE}) { + diskann::BuildConfig config; + config.data_file_path = kRawDataPath; + config.index_file_path = kDir + "/prepared"; + config.compare_metric = metric; + config.use_pq_navigation = false; + const auto temporary = config.index_file_path + "_prepped_base.bin"; + const auto norm = + diskann::get_disk_index_max_base_norm_file(diskann::get_disk_index_filename(config.index_file_path)); + { + auto context = diskann::prepare_build_context(config); + REQUIRE(context->rows == rows); + REQUIRE(context->raw_dim == dim); + REQUIRE(context->prepared_dim == dim + (metric == diskann::Metric::INNER_PRODUCT)); + REQUIRE(context->ssd_source == (metric == diskann::Metric::INNER_PRODUCT ? temporary : kRawDataPath)); + REQUIRE(context->prepared_source == (metric == diskann::Metric::L2 ? kRawDataPath : temporary)); + REQUIRE_THROWS(context->own_temporary(kRawDataPath)); + // Failure after successful preprocessing must release the prepared input. + REQUIRE(diskann::build_disk_index(config, *context) != 0); + } + REQUIRE(fs::exists(kRawDataPath)); + REQUIRE_FALSE(fs::exists(temporary)); + REQUIRE_FALSE(fs::exists(norm)); + REQUIRE_FALSE(fs::exists(config.index_file_path + "_mem.index")); + // A preexisting file is not ours to overwrite or clean up. + if (metric != diskann::Metric::L2) { + fs::copy_file(kRawDataPath, temporary); + REQUIRE_THROWS(diskann::prepare_build_context(config)); + REQUIRE(fs::file_size(temporary) == fs::file_size(kRawDataPath)); + fs::remove(temporary); + } + } +} + +TEST_CASE("DiskANN registration rolls back a partial failure", "[diskann][build_context]") { + class FailingManager : public milvus::LocalFileManager { + public: + bool fail = true; + bool + AddFile(const std::string& path) override { + const auto result = milvus::LocalFileManager::AddFile(path); + return fail && path == "second" ? false : result; + } + } manager; + REQUIRE(manager.AddFile("existing")); + { + knowhere::DiskANNBuildRegistration files(manager); + REQUIRE_FALSE(files.Add("existing")); + REQUIRE(files.Add("first")); + REQUIRE_FALSE(files.Add("second")); + } + REQUIRE(manager.IsExisted("existing").value()); + REQUIRE_FALSE(manager.IsExisted("first").value()); + REQUIRE_FALSE(manager.IsExisted("second").value()); + manager.fail = false; + { + knowhere::DiskANNBuildRegistration files(manager); + REQUIRE(files.Add("first")); + REQUIRE(files.Add("second")); + files.Commit(); + } + REQUIRE(manager.IsExisted("first").value()); + REQUIRE(manager.IsExisted("second").value()); +} + TEST_CASE("Valid diskann build params test", "[diskann]") { int rows_num = 1000000; auto version = GenTestVersionList(); diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 4295bd488..0d468a819 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -130,14 +130,49 @@ namespace diskann { bool aisaq_mode = false; uint32_t inline_pq = 0; bool rearrange = false; - int num_entry_points = 0; - // Keep the temporary MIPS-to-L2 or normalized cosine base until the caller - // builds auxiliary indexes in the same internal representation. - bool keep_preprocessed_base = false; + int num_entry_points = 0; // External navigation owns its codes; SSD PQ remains independently optional. bool use_pq_navigation = true; }; + // One build owns its prepared input and intermediate files. The original + // input is never owned. Outputs are retained only after successful + // publication. + class PreparedBuildContext { + public: + explicit PreparedBuildContext(const BuildConfig &config); + ~PreparedBuildContext(); + PreparedBuildContext(const PreparedBuildContext &) = delete; + PreparedBuildContext &operator=(const PreparedBuildContext &) = delete; + void own_temporary(const std::string &path); + void own_output(const std::string &path); + void commit_outputs() noexcept { + committed_ = true; + } + + const std::string raw_source; + const std::string prefix; + const diskann::Metric metric; + std::string prepared_source; + std::string ssd_source; + size_t rows = 0; + size_t raw_dim = 0; + size_t prepared_dim = 0; + + private: + void own(const std::string &path, std::vector &paths); + std::vector temporaries_; + std::vector outputs_; + bool committed_ = false; + }; + + template + std::unique_ptr prepare_build_context( + const BuildConfig &config); + + template + int build_disk_index(BuildConfig &config, PreparedBuildContext &context); + template int build_disk_index(BuildConfig &config); diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index d8efc0a24..5535d168b 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -5,6 +5,7 @@ #include #include #include +#include #include #include #include @@ -1607,11 +1608,50 @@ void create_aisaq_layout(const std::string base_file, const std::string mem_inde } } -template - int build_disk_index(BuildConfig &config) { - if (config.aisaq_mode && !config.use_pq_navigation) { - throw diskann::ANNException("AiSAQ requires PQ navigation", -1); +PreparedBuildContext::PreparedBuildContext(const BuildConfig &config) + : raw_source(config.data_file_path), prefix(config.index_file_path), + metric(config.compare_metric), prepared_source(raw_source), + ssd_source(raw_source) { + get_bin_metadata(raw_source, rows, raw_dim); + prepared_dim = raw_dim; +} + +void PreparedBuildContext::own(const std::string &path, + std::vector &paths) { + if (std::find(paths.begin(), paths.end(), path) != paths.end()) + return; + if (path == raw_source || std::filesystem::exists(path)) { + throw diskann::ANNException( + "Refusing to overwrite existing build file: " + path, -1); } + paths.push_back(path); +} + +void PreparedBuildContext::own_temporary(const std::string &path) { + own(path, temporaries_); +} +void PreparedBuildContext::own_output(const std::string &path) { + own(path, outputs_); +} + +PreparedBuildContext::~PreparedBuildContext() { + auto remove_owned = [](const std::vector &paths) { + for (auto it = paths.rbegin(); it != paths.rend(); ++it) { + std::error_code error; + std::filesystem::remove(*it, error); + if (error) + LOG_KNOWHERE_WARNING_ << "Could not clean build file " << *it << ": " + << error.message(); + } + }; + remove_owned(temporaries_); + if (!committed_) + remove_owned(outputs_); +} + +template +std::unique_ptr prepare_build_context( + const BuildConfig &config) { if (!knowhere::KnowhereFloatTypeCheck::value && (config.compare_metric == diskann::Metric::INNER_PRODUCT || config.compare_metric == diskann::Metric::COSINE)) { @@ -1622,33 +1662,17 @@ template throw diskann::ANNException(stream.str(), -1); } - _u32 disk_pq_dims = config.disk_pq_dims; - bool use_disk_pq = disk_pq_dims != 0; - - bool reorder_data = config.reorder; - bool ip_prepared = false; - - std::string base_file = config.data_file_path; - std::string data_file_to_use = base_file; - std::string data_file_to_save = base_file; - std::string index_prefix_path = config.index_file_path; - std::string pq_pivots_path = get_pq_pivots_filename(index_prefix_path); - std::string pq_compressed_vectors_path = - get_pq_compressed_filename(index_prefix_path); - std::string mem_index_path = index_prefix_path + "_mem.index"; - std::string disk_index_path = get_disk_index_filename(index_prefix_path); - std::string medoids_path = get_disk_index_medoids_filename(disk_index_path); - std::string centroids_path = - get_disk_index_centroids_filename(disk_index_path); - std::string sample_data_file = get_sample_data_filename(index_prefix_path); - // optional, used if disk index file must store pq data - std::string disk_pq_pivots_path = - index_prefix_path + "_disk.index_pq_pivots.bin"; - // optional, used if disk index must store pq data - std::string disk_pq_compressed_vectors_path = - index_prefix_path + "_disk.index_pq_compressed.bin"; - // optional, used if build mem usage is enough to generate cached nodes - std::string cached_nodes_file = get_cached_nodes_file(index_prefix_path); + auto context = std::make_unique(config); + const auto &base_file = context->raw_source; + const auto &index_prefix_path = context->prefix; + const auto disk_index_path = get_disk_index_filename(index_prefix_path); + auto &data_file_to_use = context->prepared_source; + auto &data_file_to_save = context->ssd_source; + if (config.compare_metric == diskann::Metric::INNER_PRODUCT || + config.compare_metric == diskann::Metric::COSINE) { + context->own_temporary(index_prefix_path + "_prepped_base.bin"); + context->own_output(get_disk_index_max_base_norm_file(disk_index_path)); + } // output a new base file which contains extra dimension with sqrt(1 - // ||x||^2/M^2) for every x, M is max norm of all points. Extra space on @@ -1667,7 +1691,6 @@ template std::string norm_file = get_disk_index_max_base_norm_file(disk_index_path); diskann::save_bin(norm_file, &max_norm_of_base, 1, 1); - ip_prepared = true; } if (config.compare_metric == diskann::Metric::COSINE) { LOG_KNOWHERE_INFO_ @@ -1684,6 +1707,49 @@ template diskann::save_bin(norm_file, norms_of_base.data(), norms_of_base.size(), 1); } + get_bin_metadata(data_file_to_use, context->rows, context->prepared_dim); + return context; +} + +template +int build_disk_index(BuildConfig &config) { + auto context = prepare_build_context(config); + const auto result = build_disk_index(config, *context); + if (result == 0) + context->commit_outputs(); + return result; +} + +template +int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { + if (config.aisaq_mode && !config.use_pq_navigation) { + throw diskann::ANNException("AiSAQ requires PQ navigation", -1); + } + _u32 disk_pq_dims = config.disk_pq_dims; + const bool use_disk_pq = disk_pq_dims != 0; + const bool reorder_data = config.reorder; + const bool ip_prepared = context.metric == diskann::Metric::INNER_PRODUCT; + const auto &base_file = context.raw_source; + const auto &data_file_to_use = context.prepared_source; + const auto &data_file_to_save = context.ssd_source; + const auto &index_prefix_path = context.prefix; + const auto pq_pivots_path = get_pq_pivots_filename(index_prefix_path); + const auto pq_compressed_vectors_path = + get_pq_compressed_filename(index_prefix_path); + const auto mem_index_path = index_prefix_path + "_mem.index"; + const auto disk_index_path = get_disk_index_filename(index_prefix_path); + const auto medoids_path = get_disk_index_medoids_filename(disk_index_path); + const auto centroids_path = + get_disk_index_centroids_filename(disk_index_path); + const auto sample_data_file = get_sample_data_filename(index_prefix_path); + const auto disk_pq_pivots_path = + index_prefix_path + "_disk.index_pq_pivots.bin"; + const auto disk_pq_compressed_vectors_path = + index_prefix_path + "_disk.index_pq_compressed.bin"; + const auto cached_nodes_file = get_cached_nodes_file(index_prefix_path); + context.own_temporary(mem_index_path); + if (use_disk_pq) + context.own_temporary(disk_pq_compressed_vectors_path); unsigned R = config.max_degree; unsigned L = config.search_list_size; @@ -1896,16 +1962,8 @@ template std::chrono::duration diff = e - s; LOG_KNOWHERE_INFO_ << "Indexing time: " << diff.count(); - if ((config.compare_metric == diskann::Metric::INNER_PRODUCT || - config.compare_metric == diskann::Metric::COSINE) && - !config.keep_preprocessed_base) { - std::remove(data_file_to_use.c_str()); - } - std::remove(mem_index_path.c_str()); - if (use_disk_pq) - std::remove(disk_pq_compressed_vectors_path.c_str()); return 0; - } +} template void create_disk_layout(const std::string base_file, const std::string mem_index_file, @@ -1976,6 +2034,17 @@ template template int build_disk_index(BuildConfig &config); template int build_disk_index(BuildConfig &config); template int build_disk_index(BuildConfig &config); + template int build_disk_index(BuildConfig &, PreparedBuildContext &); + template int build_disk_index(BuildConfig &, + PreparedBuildContext &); + template int build_disk_index(BuildConfig &, + PreparedBuildContext &); + template std::unique_ptr prepare_build_context( + const BuildConfig &); + template std::unique_ptr + prepare_build_context(const BuildConfig &); + template std::unique_ptr + prepare_build_context(const BuildConfig &); template std::unique_ptr> build_merged_vamana_index( From 7aba2f88da34dbe31dca52b11c086c76094aad37 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Tue, 22 Sep 2026 09:53:37 +0000 Subject: [PATCH 09/12] refactor: unify DiskANN navigation builders and discover persisted codecs Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- cmake/libs/libdiskann.cmake | 1 + src/index/diskann/diskann.cc | 84 ++++---- src/index/diskann/diskann_aisaq.cc | 11 +- src/index/diskann/diskann_config.h | 25 +-- src/index/diskann/navigation_store.cc | 189 +++++++++++++++--- src/index/diskann/navigation_store.h | 19 +- tests/python/test_diskann_rabitq.py | 8 +- tests/ut/test_diskann.cc | 112 ++++++++++- .../DiskANN/include/diskann/aux_utils.h | 6 +- .../include/diskann/navigation_build.h | 48 +++++ thirdparty/DiskANN/src/aux_utils.cpp | 123 +++++------- thirdparty/DiskANN/src/navigation_build.cpp | 100 +++++++++ 12 files changed, 538 insertions(+), 188 deletions(-) create mode 100644 thirdparty/DiskANN/include/diskann/navigation_build.h create mode 100644 thirdparty/DiskANN/src/navigation_build.cpp diff --git a/cmake/libs/libdiskann.cmake b/cmake/libs/libdiskann.cmake index 9c6ce5eb3..e263bf6f8 100644 --- a/cmake/libs/libdiskann.cmake +++ b/cmake/libs/libdiskann.cmake @@ -14,6 +14,7 @@ include_directories(${double-conversion_INCLUDE_DIRS}) set(DISKANN_SOURCES thirdparty/DiskANN/src/ann_exception.cpp thirdparty/DiskANN/src/aux_utils.cpp + thirdparty/DiskANN/src/navigation_build.cpp thirdparty/DiskANN/src/distance.cpp thirdparty/DiskANN/src/index.cpp thirdparty/DiskANN/src/linux_aligned_file_reader.cpp diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 86d06a18c..199d9a28a 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -336,6 +336,7 @@ class DiskANNIndexNode : public IndexNode { std::shared_ptr file_manager_; std::unique_ptr> pq_flash_index_; std::unique_ptr navigation_store_; + std::string loaded_navigation_codec_; std::atomic_int64_t dim_; std::atomic_int64_t count_; std::shared_ptr search_pool_; @@ -412,19 +413,11 @@ TryDiskANNCall(std::function&& diskann_call) { } std::vector -GetNecessaryFilenames(const std::string& prefix, const bool need_norm, const bool use_sample_cache, - const bool use_sample_warmup, const bool use_pq_navigation = true) { - std::vector filenames; - auto pq_pivots_filename = diskann::get_pq_pivots_filename(prefix); +GetNecessaryFilenames(const DiskANNConfig& config, const std::string& prefix, const bool need_norm, + const bool use_sample_cache, const bool use_sample_warmup) { + auto filenames = NavigationFiles(config, prefix).required; auto disk_index_filename = diskann::get_disk_index_filename(prefix); - if (use_pq_navigation) { - filenames.push_back(pq_pivots_filename); - filenames.push_back(diskann::get_pq_rearrangement_perm_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_chunk_offsets_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_centroid_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_compressed_filename(prefix)); - } filenames.push_back(disk_index_filename); if (need_norm) { filenames.push_back(diskann::get_disk_index_max_base_norm_file(disk_index_filename)); @@ -436,8 +429,8 @@ GetNecessaryFilenames(const std::string& prefix, const bool need_norm, const boo } std::vector -GetOptionalFilenames(const std::string& prefix) { - std::vector filenames; +GetOptionalFilenames(const DiskANNConfig& config, const std::string& prefix) { + auto filenames = NavigationFiles(config, prefix).optional; auto disk_index_filename = diskann::get_disk_index_filename(prefix); auto disk_pq_pivots_file_name = diskann::get_disk_index_pq_pivots_filename(disk_index_filename); filenames.push_back(diskann::get_disk_index_centroids_filename(disk_index_filename)); @@ -462,8 +455,8 @@ AnyIndexFileExist(const std::string& index_prefix, const DiskANNConfig& config) } return false; }; - return file_exist(GetNecessaryFilenames(index_prefix, diskann::INNER_PRODUCT, true, true)) || - file_exist(GetOptionalFilenames(index_prefix)) || file_exist(NavigationFiles(config, index_prefix)); + return file_exist(GetNecessaryFilenames(config, index_prefix, true, true, true)) || + file_exist(GetOptionalFilenames(config, index_prefix)) || file_exist(AllNavigationFiles(index_prefix)); } inline bool @@ -542,47 +535,37 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr(num_nodes_to_cache), build_conf.shuffle_build.value()}; - diskann_internal_build_config.use_pq_navigation = !external_navigation; std::unique_ptr context; RETURN_IF_ERROR(TryDiskANNCall([&]() { context = diskann::prepare_build_context(diskann_internal_build_config); - for (const auto& path : GetNecessaryFilenames(index_prefix_, need_norm, true, true, !external_navigation)) { - context->own_output(path); - } - for (const auto& path : GetOptionalFilenames(index_prefix_)) { + for (const auto& path : GetNecessaryFilenames(build_conf, index_prefix_, need_norm, true, true)) { context->own_output(path); } - for (const auto& path : NavigationFiles(build_conf, index_prefix_)) { + for (const auto& path : GetOptionalFilenames(build_conf, index_prefix_)) { context->own_output(path); } - const int res = diskann::build_disk_index(diskann_internal_build_config, *context); + auto navigation = CreateNavigationBuilder(build_conf, diskann::make_pq_navigation_builder()); + const int res = diskann::build_disk_index(diskann_internal_build_config, *context, *navigation); if (res != 0) { throw diskann::ANNException("diskann::build_disk_index returned non-zero value: " + std::to_string(res), -1); } - BuildNavigationStore(build_conf, context->prepared_source, index_prefix_); })); // Add file to the file manager DiskANNBuildRegistration registration(*file_manager_); - for (auto& filename : GetNecessaryFilenames(index_prefix_, need_norm, true, true, !external_navigation)) { + for (auto& filename : GetNecessaryFilenames(build_conf, index_prefix_, need_norm, true, true)) { if (!registration.Add(filename)) { LOG_KNOWHERE_ERROR_ << "Failed to add file " << filename << "."; return Status::disk_file_error; } } - for (auto& filename : GetOptionalFilenames(index_prefix_)) { + for (auto& filename : GetOptionalFilenames(build_conf, index_prefix_)) { if (file_exists(filename) && !registration.Add(filename)) { LOG_KNOWHERE_ERROR_ << "Failed to add file " << filename << "."; return Status::disk_file_error; } } - for (const auto& sidecar_path : NavigationFiles(build_conf, index_prefix_)) { - if (!registration.Add(sidecar_path)) { - LOG_KNOWHERE_ERROR_ << "Failed to add file " << sidecar_path << "."; - return Status::disk_file_error; - } - } registration.Commit(); context->commit_outputs(); @@ -664,20 +647,22 @@ DiskANNIndexNode::BuildEmbListIfNeed(const DataSetPtr dataset, std::sh template Status DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr cfg) { - const auto& prep_conf = static_cast(*cfg); - const bool external_navigation = UsesExternalNavigation(prep_conf); - if (external_navigation && !std::is_same_v) { - return Status::invalid_args; - } + auto prep_conf = static_cast(*cfg); if (!CheckMetric(prep_conf.metric_type.value())) { return Status::invalid_metric_type; } if (is_prepared_.load()) { + if (prep_conf.index_prefix.value_or("") != index_prefix_ || + (prep_conf.navigation_codec.has_value() && + prep_conf.navigation_codec.value() != loaded_navigation_codec_)) { + return Status::invalid_serialized_index_type; + } return Status::success; } const auto rollback = folly::makeGuard([this]() { if (!is_prepared_.load()) { navigation_store_.reset(); + loaded_navigation_codec_.clear(); pq_flash_index_.reset(); count_.store(-1); dim_.store(-1); @@ -688,6 +673,18 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr return Status::invalid_param_in_json; } index_prefix_ = prep_conf.index_prefix.value(); + const auto detected = DetectNavigationCodec(prep_conf, index_prefix_, *file_manager_); + if (!detected.has_value()) { + LOG_KNOWHERE_ERROR_ << detected.what(); + return detected.error(); + } + prep_conf.navigation_codec = detected.value(); + const bool external_navigation = UsesExternalNavigation(prep_conf); + if (external_navigation && !std::is_same_v) + return Status::invalid_args; + if (external_navigation && (!el_metric_type_.empty() || prep_conf.emb_list_offset_file_path.has_value())) { + return Status::not_implemented; + } bool is_ip = IsMetricType(prep_conf.metric_type.value(), knowhere::metric::IP); bool need_norm = IsMetricType(prep_conf.metric_type.value(), knowhere::metric::IP) || IsMetricType(prep_conf.metric_type.value(), knowhere::metric::COSINE); @@ -703,14 +700,14 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr // Load file from file manager. for (auto& filename : GetNecessaryFilenames( - index_prefix_, need_norm, + prep_conf, index_prefix_, need_norm, prep_conf.search_cache_budget_gb.value() > 0 && !prep_conf.use_bfs_cache.value() && !external_navigation, - prep_conf.warm_up.value(), !external_navigation)) { + prep_conf.warm_up.value())) { if (!LoadFile(filename)) { return Status::disk_file_error; } } - for (auto& filename : GetOptionalFilenames(index_prefix_)) { + for (auto& filename : GetOptionalFilenames(prep_conf, index_prefix_)) { auto is_exist_op = file_manager_->IsExisted(filename); if (!is_exist_op.has_value()) { LOG_KNOWHERE_ERROR_ << "Failed to check existence of file " << filename << "."; @@ -720,12 +717,6 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr return Status::disk_file_error; } } - for (const auto& sidecar_path : NavigationFiles(prep_conf, index_prefix_)) { - if (!LoadFile(sidecar_path)) { - LOG_KNOWHERE_ERROR_ << "Failed to load DiskANN navigation sidecar " << sidecar_path; - return Status::disk_file_error; - } - } navigation_store_.reset(); if (external_navigation) { @@ -890,6 +881,7 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr } } + loaded_navigation_codec_ = detected.value(); is_prepared_.store(true); LOG_KNOWHERE_INFO_ << "End of diskann loading."; return Status::success; @@ -970,7 +962,7 @@ template expected> DiskANNIndexNode::AnnIterator(const DataSetPtr dataset, std::unique_ptr cfg, const BitsetView& bitset, bool use_knowhere_search_pool, milvus::OpContext* op_context) const { - if (navigation_store_ || UsesExternalNavigation(static_cast(*cfg))) { + if (navigation_store_) { return expected>::Err(Status::not_implemented, "DISKANN_RABITQ does not support iterator search"); } diff --git a/src/index/diskann/diskann_aisaq.cc b/src/index/diskann/diskann_aisaq.cc index fb0a236fe..9df18ab58 100644 --- a/src/index/diskann/diskann_aisaq.cc +++ b/src/index/diskann/diskann_aisaq.cc @@ -7,6 +7,7 @@ #include "diskann/aisaq.h" #include "diskann/aux_utils.h" #include "diskann/linux_aligned_file_reader.h" +#include "diskann/navigation_build.h" #include "diskann/pq_flash_aisaq_index.h" #include "diskann/pq_flash_index.h" #include "filemanager/FileManager.h" @@ -213,14 +214,9 @@ TryDiskANNCall(std::function&& diskann_call) { std::vector GetNecessaryFilenames(const std::string& prefix, const bool need_norm, const bool use_sample_cache, const bool use_sample_warmup, const bool rearrange, const bool entry_points) { - std::vector filenames; - auto pq_pivots_filename = diskann::get_pq_pivots_filename(prefix); + auto filenames = diskann::pq_navigation_files(prefix, rearrange); auto disk_index_filename = diskann::get_disk_index_filename(prefix); - filenames.push_back(pq_pivots_filename); - filenames.push_back(diskann::get_pq_rearrangement_perm_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_chunk_offsets_filename(pq_pivots_filename)); - filenames.push_back(diskann::get_pq_centroid_filename(pq_pivots_filename)); filenames.push_back(disk_index_filename); if (need_norm) { filenames.push_back(diskann::get_disk_index_max_base_norm_file(disk_index_filename)); @@ -230,9 +226,6 @@ GetNecessaryFilenames(const std::string& prefix, const bool need_norm, const boo } if (rearrange) { filenames.push_back(diskann::get_index_rearranged_filename(prefix)); - filenames.push_back(diskann::get_pq_compressed_rearranged_filename(prefix)); - } else { - filenames.push_back(diskann::get_pq_compressed_filename(prefix)); } if (entry_points) { filenames.push_back(diskann::get_index_entry_points_filename(prefix)); diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index 0fcd71c42..f3f1cb8ed 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -17,6 +17,10 @@ namespace knowhere { +class DiskANNNavigationConfig; +Status +ValidateNavigationConfig(const DiskANNNavigationConfig& config, std::string* error); + namespace { constexpr const CFG_INT::value_type kSearchListSizeMinValue = 16; @@ -209,8 +213,8 @@ class DiskANNNavigationConfig : public DiskANNConfig { KNOWHERE_DECLARE_CONFIG(DiskANNNavigationConfig) { KNOWHERE_CONFIG_DECLARE_FIELD(navigation_codec) - .description("resident navigation codec: PQ or RABITQ") - .set_default("PQ") + .description("navigation codec: build defaults to PQ; load detects stored codec unless constrained") + .allow_empty_without_default() .for_train() .for_deserialize() .for_static(); @@ -233,22 +237,7 @@ class DiskANNNavigationConfig : public DiskANNConfig { if (base_status != Status::success) { return base_status; } - const auto codec = navigation_codec.value_or("PQ"); - if (codec == "PQ") { - return Status::success; - } - if (codec != "RABITQ") { - return HandleError(err_msg, "unsupported DiskANN navigation codec", Status::invalid_args); - } - const auto metric = metric_type.value_or(knowhere::metric::L2); - if (metric != knowhere::metric::L2 && metric != knowhere::metric::IP && metric != knowhere::metric::COSINE) { - return HandleError(err_msg, "DISKANN_RABITQ supports L2, IP and COSINE", Status::invalid_metric_type); - } - const auto database_bits = rbq_bits.value_or(1); - if (database_bits < 1 || database_bits > 9) { - return HandleError(err_msg, "DISKANN_RABITQ supports rbq_bits in [1, 9]", Status::invalid_args); - } - return Status::success; + return ValidateNavigationConfig(*this, err_msg); } }; diff --git a/src/index/diskann/navigation_store.cc b/src/index/diskann/navigation_store.cc index dbcf30bb0..95e062878 100644 --- a/src/index/diskann/navigation_store.cc +++ b/src/index/diskann/navigation_store.cc @@ -2,64 +2,193 @@ // SPDX-License-Identifier: Apache-2.0 #include "index/diskann/navigation_store.h" +#include +#include #include #include +#include +#include "diskann/aux_utils.h" #include "index/diskann/diskann_config.h" #include "index/diskann/rabitq_store.h" namespace knowhere { namespace { -const DiskANNNavigationConfig* -ExternalConfig(const DiskANNConfig& config) { - const auto* navigation = dynamic_cast(&config); - if (!navigation || navigation->navigation_codec.value_or("PQ") == "PQ") { - return nullptr; +const DiskANNNavigationConfig& +NavigationConfig(const DiskANNConfig& config) { + return static_cast(config); +} + +class RaBitQNavigationBuilder final : public diskann::NavigationBuilder { + public: + explicit RaBitQNavigationBuilder(uint8_t bits) : bits_(bits) { } - if (navigation->navigation_codec.value() != "RABITQ") { - throw std::invalid_argument("unsupported DiskANN navigation codec"); + void + build(const diskann::BuildConfig&, const diskann::PreparedBuildContext& context, + const diskann::NavigationTrainingData&) const override { + RaBitQStore::BuildFromFloatBin(context.prepared_source, RaBitQStore::SidecarFilename(context.prefix), bits_); } - return navigation; + + private: + uint8_t bits_; +}; + +struct NavigationCodec { + bool external; + Status (*validate)(const DiskANNNavigationConfig&, std::string*); + uint64_t (*estimate)(const DiskANNConfig&, int64_t, int64_t); + NavigationFileSet (*files)(const std::string&); + std::unique_ptr (*builder)(const DiskANNConfig&); + std::unique_ptr (*load)(const std::string&); +}; + +const std::unordered_map& +Codecs() { + static const std::unordered_map codecs = { + {"PQ", + {false, [](const DiskANNNavigationConfig&, std::string*) { return Status::success; }, + [](const DiskANNConfig&, int64_t, int64_t) -> uint64_t { return 0; }, + [](const std::string& prefix) { + return NavigationFileSet{diskann::pq_navigation_files(prefix), {}}; + }, + nullptr, [](const std::string&) -> std::unique_ptr { return nullptr; }}}, + {"RABITQ", + {true, + [](const DiskANNNavigationConfig& config, std::string* error) { + const auto metric = config.metric_type.value_or(metric::L2); + if (metric != metric::L2 && metric != metric::IP && metric != metric::COSINE) { + if (error) + *error = "DISKANN_RABITQ supports L2, IP and COSINE"; + return Status::invalid_metric_type; + } + const auto bits = config.rbq_bits.value_or(1); + if (bits < 1 || bits > 9) { + if (error) + *error = "DISKANN_RABITQ supports rbq_bits in [1, 9]"; + return Status::invalid_args; + } + return Status::success; + }, + [](const DiskANNConfig& config, int64_t rows, int64_t dim) { + if (dim <= 0 || dim >= std::numeric_limits::max()) { + throw std::invalid_argument("invalid DiskANN navigation dimension"); + } + const auto prepared_dim = dim + (config.metric_type.value_or(metric::L2) == metric::IP ? 1 : 0); + return RaBitQStore::EstimateMemorySize( + rows, prepared_dim, static_cast(NavigationConfig(config).rbq_bits.value_or(1))); + }, + [](const std::string& prefix) { + return NavigationFileSet{{RaBitQStore::SidecarFilename(prefix)}, {}}; + }, + [](const DiskANNConfig& config) -> std::unique_ptr { + return std::make_unique( + static_cast(NavigationConfig(config).rbq_bits.value_or(1))); + }, + [](const std::string& prefix) -> std::unique_ptr { + return std::make_unique(RaBitQStore::SidecarFilename(prefix)); + }}}}; + return codecs; +} + +const NavigationCodec& +Codec(const DiskANNConfig& config) { + // Older native/AiSAQ callers have a plain DiskANNConfig and use PQ. This + // cast only reads the selector; algorithm dispatch is by registered name. + const auto* navigation = dynamic_cast(&config); + const auto name = navigation ? navigation->navigation_codec.value_or("PQ") : "PQ"; + const auto found = Codecs().find(name); + if (found == Codecs().end()) + throw std::invalid_argument("unsupported DiskANN navigation codec: " + name); + return found->second; } } // namespace +expected +DetectNavigationCodec(const DiskANNNavigationConfig& config, const std::string& prefix, milvus::FileManager& manager) { + std::vector found; + for (const auto& [name, codec] : Codecs()) { + bool present = false; + // Any exclusive required artifact identifies a possible codec, even + // when its primary file is missing. Never disguise an incomplete + // model as absence and silently select another codec. + for (const auto& path : codec.files(prefix).required) { + const auto exists = manager.IsExisted(path); + std::error_code error; + const auto local_exists = std::filesystem::exists(path, error); + if (!exists.has_value() || error) { + return expected::Err(Status::disk_file_error, "Cannot query navigation file: " + path); + } + // A fresh LocalFileManager has not registered already-localized + // files. Remote query failures still remain errors above. + present = present || exists.value() || local_exists; + } + if (present) + found.push_back(name); + } + if (found.empty()) { + return expected::Err(Status::disk_file_error, "No stored DiskANN navigation model found"); + } + if (config.navigation_codec.has_value()) { + const auto& requested = config.navigation_codec.value(); + if (std::find(found.begin(), found.end(), requested) == found.end()) { + return expected::Err(Status::invalid_serialized_index_type, + "Requested navigation codec does not match stored model: " + requested); + } + return requested; + } + if (found.size() != 1) { + return expected::Err( + Status::invalid_serialized_index_type, + "Ambiguous DiskANN navigation files; specify navigation_codec to disambiguate"); + } + return found.front(); +} + +Status +ValidateNavigationConfig(const DiskANNNavigationConfig& config, std::string* error) { + const auto found = Codecs().find(config.navigation_codec.value_or("PQ")); + if (found == Codecs().end()) { + if (error) + *error = "unsupported DiskANN navigation codec"; + return Status::invalid_args; + } + return found->second.validate(config, error); +} + bool UsesExternalNavigation(const DiskANNConfig& config) { - return ExternalConfig(config) != nullptr; + return Codec(config).external; } uint64_t EstimateNavigationMemory(const DiskANNConfig& config, int64_t rows, int64_t dim) { - const auto* navigation = ExternalConfig(config); - if (!navigation) { - return 0; - } - if (dim <= 0 || dim >= std::numeric_limits::max()) { - throw std::invalid_argument("invalid DiskANN navigation dimension"); - } - const auto prepared_dim = dim + (config.metric_type.value_or(metric::L2) == metric::IP ? 1 : 0); - return RaBitQStore::EstimateMemorySize(rows, prepared_dim, static_cast(navigation->rbq_bits.value_or(1))); + return Codec(config).estimate(config, rows, dim); } -std::vector +NavigationFileSet NavigationFiles(const DiskANNConfig& config, const std::string& prefix) { - return ExternalConfig(config) ? std::vector{RaBitQStore::SidecarFilename(prefix)} - : std::vector{}; + return Codec(config).files(prefix); } -void -BuildNavigationStore(const DiskANNConfig& config, const std::string& source, const std::string& prefix) { - if (const auto* navigation = ExternalConfig(config)) { - RaBitQStore::BuildFromFloatBin(source, RaBitQStore::SidecarFilename(prefix), - static_cast(navigation->rbq_bits.value_or(1))); +std::vector +AllNavigationFiles(const std::string& prefix) { + std::vector files; + for (const auto& [name, codec] : Codecs()) { + const auto declared = codec.files(prefix); + files.insert(files.end(), declared.required.begin(), declared.required.end()); + files.insert(files.end(), declared.optional.begin(), declared.optional.end()); } + return files; +} + +std::unique_ptr +CreateNavigationBuilder(const DiskANNConfig& config, std::unique_ptr native_pq) { + const auto& codec = Codec(config); + return codec.builder ? codec.builder(config) : std::move(native_pq); } std::unique_ptr LoadNavigationStore(const DiskANNConfig& config, const std::string& prefix) { - if (ExternalConfig(config)) { - return std::make_unique(RaBitQStore::SidecarFilename(prefix)); - } - return nullptr; + return Codec(config).load(prefix); } } // namespace knowhere diff --git a/src/index/diskann/navigation_store.h b/src/index/diskann/navigation_store.h index 33c3be1fd..a8d100ce2 100644 --- a/src/index/diskann/navigation_store.h +++ b/src/index/diskann/navigation_store.h @@ -6,10 +6,17 @@ #include #include +#include "diskann/navigation_build.h" #include "diskann/pq_flash_index.h" +#include "filemanager/FileManager.h" +#include "knowhere/expected.h" namespace knowhere { class DiskANNConfig; +class DiskANNNavigationConfig; + +expected +DetectNavigationCodec(const DiskANNNavigationConfig& config, const std::string& prefix, milvus::FileManager& manager); // Immutable after load. Each request owns its scorer and transformed-query // scratch; this store must outlive those scorers. PQ remains the native engine @@ -33,10 +40,16 @@ bool UsesExternalNavigation(const DiskANNConfig& config); uint64_t EstimateNavigationMemory(const DiskANNConfig& config, int64_t rows, int64_t dim); -std::vector +struct NavigationFileSet { + std::vector required; + std::vector optional; +}; +NavigationFileSet NavigationFiles(const DiskANNConfig& config, const std::string& prefix); -void -BuildNavigationStore(const DiskANNConfig& config, const std::string& prepared_source, const std::string& prefix); +std::vector +AllNavigationFiles(const std::string& prefix); +std::unique_ptr +CreateNavigationBuilder(const DiskANNConfig& config, std::unique_ptr native_pq); std::unique_ptr LoadNavigationStore(const DiskANNConfig& config, const std::string& prefix); } // namespace knowhere diff --git a/tests/python/test_diskann_rabitq.py b/tests/python/test_diskann_rabitq.py index 960560fe0..43fcf17c1 100644 --- a/tests/python/test_diskann_rabitq.py +++ b/tests/python/test_diskann_rabitq.py @@ -35,8 +35,12 @@ def test_navigation_roundtrip(tmp_path, metric, kind, codec): binary = knowhere.GetBinarySet() assert knowhere.Status(index.Serialize(binary)) == knowhere.Status.success del index - restored = knowhere.CreateIndex(kind, version) - assert knowhere.Status(restored.Deserialize(binary, json.dumps(dict(config, warm_up=True)))) == knowhere.Status.success + # A fresh generic DiskANN node recovers its navigation from stored files, + # without the original codec, bits, PQ budget or other build parameters. + restored = knowhere.CreateIndex("DISKANN", version) + load = dict(metric_type=metric, index_prefix=config["index_prefix"], + search_cache_budget_gb=0, search_cache_budget_gb_ratio=0, warm_up=True) + assert knowhere.Status(restored.Deserialize(binary, json.dumps(load))) == knowhere.Status.success search = dict(dim=64, metric_type=metric, k=10, search_list_size=100, beamwidth=4, rbq_bits_query=4) result, status = restored.Search(knowhere.ArrayToDataSet(query), json.dumps(search), knowhere.GetNullBitSetView()) assert knowhere.Status(status) == knowhere.Status.success diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index f2498ac4c..18272ab6b 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -115,7 +115,7 @@ TEST_CASE("DiskANN prepared build files have scoped ownership", "[diskann][build config.data_file_path = kRawDataPath; config.index_file_path = kDir + "/prepared"; config.compare_metric = metric; - config.use_pq_navigation = false; + config.pq_code_size_gb = 0.001; const auto temporary = config.index_file_path + "_prepped_base.bin"; const auto norm = diskann::get_disk_index_max_base_norm_file(diskann::get_disk_index_filename(config.index_file_path)); @@ -128,7 +128,8 @@ TEST_CASE("DiskANN prepared build files have scoped ownership", "[diskann][build REQUIRE(context->prepared_source == (metric == diskann::Metric::L2 ? kRawDataPath : temporary)); REQUIRE_THROWS(context->own_temporary(kRawDataPath)); // Failure after successful preprocessing must release the prepared input. - REQUIRE(diskann::build_disk_index(config, *context) != 0); + const auto navigation = diskann::make_pq_navigation_builder(); + REQUIRE(diskann::build_disk_index(config, *context, *navigation) != 0); } REQUIRE(fs::exists(kRawDataPath)); REQUIRE_FALSE(fs::exists(temporary)); @@ -175,6 +176,111 @@ TEST_CASE("DiskANN registration rolls back a partial failure", "[diskann][build_ REQUIRE(manager.IsExisted("second").value()); } +TEST_CASE("DiskANN discovers persisted navigation without build parameters", "[diskann][navigation_load]") { + fs::remove_all(kDir); + fs::create_directories(kDir); + constexpr int rows = 512, dim = 16; + auto base = GenDataSet(rows, dim, 19); + auto query = GenDataSet(10, dim, 20); + WriteRawDataToDisk(kRawDataPath, static_cast(base->GetTensor()), rows, dim); + const auto version = GenTestVersionList(); + const auto pq_prefix = kDir + "/pq"; + const auto rbq_prefix = kDir + "/rbq"; + knowhere::BinarySet empty; + auto create = [&](bool rbq = false) { + auto manager = std::make_shared(); + return knowhere::IndexFactory::Instance() + .Create(rbq ? knowhere::IndexEnum::INDEX_DISKANN_RABITQ : kNativeDiskANN, version, + knowhere::Pack(std::shared_ptr(manager))) + .value(); + }; + for (bool rbq : {false, true}) { + auto index = create(); + knowhere::Json build = {{"dim", dim}, + {"metric_type", "L2"}, + {"index_prefix", rbq ? rbq_prefix : pq_prefix}, + {"data_path", kRawDataPath}, + {"max_degree", 16}, + {"search_list_size", 32}, + {"build_dram_budget_gb", 1.0}, + {"pq_code_budget_gb", 0.000004}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"rbq_bits", 4}}; + if (rbq) + build["navigation_codec"] = "RABITQ"; + REQUIRE(index.Build(nullptr, build) == knowhere::Status::success); + } + knowhere::Json load = {{"metric_type", "L2"}, + {"index_prefix", rbq_prefix}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"warm_up", false}}; + knowhere::Json search = {{"metric_type", "L2"}, {"k", 10}, {"search_list_size", 100}, {"beamwidth", 8}}; + auto rbq = create(); + REQUIRE(rbq.Deserialize(empty, load) == knowhere::Status::success); + auto reference = rbq.Search(query, search, nullptr); + REQUIRE(reference.has_value()); + auto alias = create(true); + REQUIRE(alias.Deserialize(empty, load) == knowhere::Status::success); + auto alias_result = alias.Search(query, search, nullptr); + REQUIRE(alias_result.has_value()); + for (int i = 0; i < 100; ++i) { + REQUIRE(reference.value()->GetIds()[i] == alias_result.value()->GetIds()[i]); + REQUIRE(reference.value()->GetDistance()[i] == alias_result.value()->GetDistance()[i]); + } + auto wrong = load; + wrong["navigation_codec"] = "PQ"; + REQUIRE(create().Deserialize(empty, wrong) == knowhere::Status::invalid_serialized_index_type); + REQUIRE(rbq.Deserialize(empty, wrong) == knowhere::Status::invalid_serialized_index_type); + load["index_prefix"] = pq_prefix; + REQUIRE(create().Deserialize(empty, load) == knowhere::Status::success); + REQUIRE(create(true).Deserialize(empty, load) == knowhere::Status::invalid_serialized_index_type); + const auto sidecar = knowhere::RaBitQStore::SidecarFilename(pq_prefix); + fs::copy_file(knowhere::RaBitQStore::SidecarFilename(rbq_prefix), sidecar); + REQUIRE(create().Deserialize(empty, load) == knowhere::Status::invalid_serialized_index_type); + load["navigation_codec"] = "RABITQ"; + auto retry = create(); + REQUIRE(retry.Deserialize(empty, load) == knowhere::Status::success); + fs::resize_file(sidecar, 4); + auto corrupt = create(); + REQUIRE(corrupt.Deserialize(empty, load) != knowhere::Status::success); + REQUIRE_FALSE(corrupt.Search(query, search, nullptr).has_value()); + fs::copy_file(knowhere::RaBitQStore::SidecarFilename(rbq_prefix), sidecar, fs::copy_options::overwrite_existing); + REQUIRE(corrupt.Deserialize(empty, load) == knowhere::Status::success); + load["navigation_codec"] = "PQ"; + REQUIRE(create().Deserialize(empty, load) == knowhere::Status::success); + fs::remove(pq_prefix + "_pq_compressed.bin"); + REQUIRE(create().Deserialize(empty, load) != knowhere::Status::success); +} + +TEST_CASE("DiskANN navigation discovery preserves FileManager errors", "[diskann][navigation_load]") { + class QueryFailureManager : public milvus::LocalFileManager { + public: + bool fail = true; + std::optional + IsExisted(const std::string& path) override { + if (fail) + return std::nullopt; + return milvus::LocalFileManager::IsExisted(path); + } + } manager; + knowhere::DiskANNNavigationConfig config; + const auto prefix = kDir + "/remote_only"; + REQUIRE(manager.AddFile(diskann::get_pq_pivots_filename(prefix))); + auto result = knowhere::DetectNavigationCodec(config, prefix, manager); + REQUIRE_FALSE(result.has_value()); + REQUIRE(result.error() == knowhere::Status::disk_file_error); + manager.fail = false; + result = knowhere::DetectNavigationCodec(config, prefix, manager); + REQUIRE(result.has_value()); + REQUIRE(result.value() == "PQ"); + REQUIRE(manager.AddFile(knowhere::RaBitQStore::SidecarFilename(prefix))); + REQUIRE_FALSE(knowhere::DetectNavigationCodec(config, prefix, manager).has_value()); + config.navigation_codec = "RABITQ"; + REQUIRE(knowhere::DetectNavigationCodec(config, prefix, manager).value() == "RABITQ"); +} + TEST_CASE("Valid diskann build params test", "[diskann]") { int rows_num = 1000000; auto version = GenTestVersionList(); @@ -1098,7 +1204,7 @@ TEST_CASE("Test DISKANN_RABITQ build and search", "[diskann][rabitq]") { auto generic = knowhere::IndexFactory::Instance().Create(kNativeDiskANN, version, pack).value(); auto generic_load = deserialize_json; - generic_load["navigation_codec"] = "RABITQ"; + // Loading the generic node recovers the codec from persisted files. REQUIRE(generic.Deserialize(binset, generic_load) == knowhere::Status::success); for (const int qb : {0, 4, 8}) { auto request = search_json; diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 0d468a819..776f0d135 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -131,8 +131,6 @@ namespace diskann { uint32_t inline_pq = 0; bool rearrange = false; int num_entry_points = 0; - // External navigation owns its codes; SSD PQ remains independently optional. - bool use_pq_navigation = true; }; // One build owns its prepared input and intermediate files. The original @@ -170,8 +168,10 @@ namespace diskann { std::unique_ptr prepare_build_context( const BuildConfig &config); + class NavigationBuilder; template - int build_disk_index(BuildConfig &config, PreparedBuildContext &context); + int build_disk_index(BuildConfig &config, PreparedBuildContext &context, + const NavigationBuilder &navigation); template int build_disk_index(BuildConfig &config); diff --git a/thirdparty/DiskANN/include/diskann/navigation_build.h b/thirdparty/DiskANN/include/diskann/navigation_build.h new file mode 100644 index 000000000..66394c0d8 --- /dev/null +++ b/thirdparty/DiskANN/include/diskann/navigation_build.h @@ -0,0 +1,48 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#include +#include + +namespace diskann { + struct BuildConfig; + class PreparedBuildContext; + + // Build-time adapter only. The graph search hot path is deliberately + // separate. + struct NavigationTrainingData { + const float *data; + size_t rows; + size_t dimension; + }; + + class NavigationBuilder { + public: + virtual ~NavigationBuilder() = default; + virtual bool validate(const BuildConfig &) const { + return true; + } + virtual bool needs_training_sample() const { + return false; + } + virtual bool supports_aisaq() const { + return false; + } + virtual void build(const BuildConfig &, const PreparedBuildContext &, + const NavigationTrainingData &) const = 0; + virtual void build_cache(const BuildConfig &, const PreparedBuildContext &, + const std::vector> &, + unsigned) const { + } + }; + + // Shared file declaration for native PQ build, Knowhere registration/load and + // AiSAQ's rearranged-code layout. The on-disk format is unchanged. + std::vector pq_navigation_files(const std::string &prefix, + bool rearranged = false); + template + std::unique_ptr make_pq_navigation_builder(); +} // namespace diskann diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index 5535d168b..1cce027e6 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -19,6 +19,7 @@ #include "boost/dynamic_bitset.hpp" #include "diskann/aisaq_utils.h" #include "diskann/aux_utils.h" +#include "diskann/navigation_build.h" #include "diskann/cached_io.h" #include "diskann/index.h" #include "diskann/partition_and_pq.h" @@ -34,7 +35,6 @@ namespace diskann { namespace { static constexpr uint32_t kSearchLForCache = 15; - static constexpr float kCacheMemFactor = 1.1; // Currently supported values for graph_degree in cuvs. static const int DEGREE_SIZES[4] = {32, 64, 128, 256}; static bool valid_gpu_params = false; @@ -1714,15 +1714,19 @@ std::unique_ptr prepare_build_context( template int build_disk_index(BuildConfig &config) { auto context = prepare_build_context(config); - const auto result = build_disk_index(config, *context); + auto navigation = make_pq_navigation_builder(); + for (const auto &path : pq_navigation_files(context->prefix)) + context->own_output(path); + const auto result = build_disk_index(config, *context, *navigation); if (result == 0) context->commit_outputs(); return result; } template -int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { - if (config.aisaq_mode && !config.use_pq_navigation) { +int build_disk_index(BuildConfig &config, PreparedBuildContext &context, + const NavigationBuilder &navigation) { + if (config.aisaq_mode && !navigation.supports_aisaq()) { throw diskann::ANNException("AiSAQ requires PQ navigation", -1); } _u32 disk_pq_dims = config.disk_pq_dims; @@ -1733,9 +1737,6 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { const auto &data_file_to_use = context.prepared_source; const auto &data_file_to_save = context.ssd_source; const auto &index_prefix_path = context.prefix; - const auto pq_pivots_path = get_pq_pivots_filename(index_prefix_path); - const auto pq_compressed_vectors_path = - get_pq_compressed_filename(index_prefix_path); const auto mem_index_path = index_prefix_path + "_mem.index"; const auto disk_index_path = get_disk_index_filename(index_prefix_path); const auto medoids_path = get_disk_index_medoids_filename(disk_index_path); @@ -1747,6 +1748,26 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { const auto disk_pq_compressed_vectors_path = index_prefix_path + "_disk.index_pq_compressed.bin"; const auto cached_nodes_file = get_cached_nodes_file(index_prefix_path); + for (const auto &path : {disk_index_path, medoids_path, centroids_path, + sample_data_file, cached_nodes_file}) { + context.own_output(path); + } + if (use_disk_pq) { + for (const auto &path : + {disk_pq_pivots_path, + get_pq_rearrangement_perm_filename(disk_pq_pivots_path), + get_pq_chunk_offsets_filename(disk_pq_pivots_path), + get_pq_centroid_filename(disk_pq_pivots_path)}) { + context.own_output(path); + } + } + if (config.aisaq_mode) { + for (const auto &path : + {get_index_rearranged_filename(index_prefix_path), + get_pq_compressed_rearranged_filename(index_prefix_path), + get_index_entry_points_filename(index_prefix_path)}) + context.own_output(path); + } context.own_temporary(mem_index_path); if (use_disk_pq) context.own_temporary(disk_pq_compressed_vectors_path); @@ -1754,15 +1775,8 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { unsigned R = config.max_degree; unsigned L = config.search_list_size; - double pq_code_size_limit = 0; - if (config.use_pq_navigation) { - pq_code_size_limit = get_memory_budget(config.pq_code_size_gb); - if (pq_code_size_limit <= 0) { - LOG(ERROR) << "Insufficient memory budget (or string was not in right " - "format). Should be > 0."; - return -1; - } - } + if (!navigation.validate(config)) + return -1; double indexing_ram_budget = config.index_mem_gb; if (indexing_ram_budget <= 0) { LOG(ERROR) << "Not building index. Please provide more RAM budget"; @@ -1799,11 +1813,13 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { diskann::get_bin_metadata(data_file_to_use.c_str(), points_num, dim); - LOG_KNOWHERE_INFO_ << "Starting index build for : " << points_num << " vectors with dim: " << dim << " R=" << R << " L=" << L - << " Query RAM budget: " - << pq_code_size_limit / (1024 * 1024 * 1024) << "(GiB)" - << " Indexing ram budget: " << indexing_ram_budget - << "(GiB)"; + LOG_KNOWHERE_INFO_ << "Starting index build for : " << points_num + << " vectors with dim: " << dim << " R=" << R + << " L=" << L + << " Query RAM budget: " << config.pq_code_size_gb + << "(GiB)" + << " Indexing ram budget: " << indexing_ram_budget + << "(GiB)"; size_t train_size = 0, train_dim = 0; std::unique_ptr train_data = nullptr; @@ -1811,7 +1827,7 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { double p_val = ((double) MAX_PQ_TRAINING_SET_SIZE / (double) points_num); // generates random sample and sets it to train_data and updates // train_size - if (config.use_pq_navigation || use_disk_pq) { + if (navigation.needs_training_sample() || use_disk_pq) { gen_random_slice(data_file_to_use.c_str(), p_val, train_data, train_size, train_dim); } @@ -1835,37 +1851,8 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { data_file_to_use.c_str(), 256, (uint32_t) disk_pq_dims, disk_pq_pivots_path, disk_pq_compressed_vectors_path); } - if (config.use_pq_navigation) { - size_t num_pq_chunks = - (size_t) (std::floor)(_u64(pq_code_size_limit / points_num)); - num_pq_chunks = num_pq_chunks <= 0 ? 1 : num_pq_chunks; - num_pq_chunks = num_pq_chunks > dim ? dim : num_pq_chunks; - num_pq_chunks = num_pq_chunks > diskann::defaults::MAX_PQ_CHUNKS ? diskann::defaults::MAX_PQ_CHUNKS : num_pq_chunks; - LOG_KNOWHERE_INFO_ << "Compressing " << dim << "-dimensional data into " - << num_pq_chunks << " bytes per vector."; - LOG_KNOWHERE_DEBUG_ << "Training data loaded of size " << train_size; - - // don't translate data to make zero mean for PQ compression. We must not - // translate for inner product search. - bool make_zero_mean = true; - if (config.compare_metric != diskann::Metric::L2) - make_zero_mean = false; - - auto pq_s = std::chrono::high_resolution_clock::now(); - - LOG_KNOWHERE_INFO_ << "Generating PQ pivots"; - generate_pq_pivots(train_data.get(), train_size, (uint32_t) dim, 256, - (uint32_t) num_pq_chunks, NUM_KMEANS_REPS, - pq_pivots_path, make_zero_mean); - - LOG_KNOWHERE_INFO_ << "Encoding PQ data"; - generate_pq_data_from_pivots(data_file_to_use.c_str(), 256, - (uint32_t) num_pq_chunks, pq_pivots_path, - pq_compressed_vectors_path); - auto pq_e = std::chrono::high_resolution_clock::now(); - std::chrono::duration pq_diff = pq_e - pq_s; - LOG_KNOWHERE_INFO_ << "Training PQ codes cost: " << pq_diff.count() << "s"; - } + navigation.build(config, context, + {train_data.get(), train_size, train_dim}); // Gopal. Splitting diskann_dll into separate DLLs for search and build. // This code should only be available in the "build" DLL. #if defined(RELEASE_UNUSED_TCMALLOC_MEMORY_AT_CHECKPOINTS) && \ @@ -1939,24 +1926,9 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { gen_random_slice(base_file.c_str(), sample_data_file, sample_sampling_rate); - if (vamana_index != nullptr && config.use_pq_navigation) { - auto final_graph = vamana_index->get_graph(); - auto entry_point = vamana_index->get_entry_point(); - - auto generate_cache_mem_usage = - kCacheMemFactor * - (get_file_size(mem_index_path) + get_file_size(sample_data_file) + - get_file_size(pq_compressed_vectors_path) + - get_file_size(pq_pivots_path)) / - (1024 * 1024 * 1024); - - if (config.num_nodes_to_cache > 0 && final_graph->size() != 0 && - generate_cache_mem_usage < config.index_mem_gb) { - generate_cache_list_from_graph_with_pq( - config.num_nodes_to_cache, config.max_degree, config.compare_metric, - sample_data_file, pq_pivots_path, pq_compressed_vectors_path, - entry_point, *final_graph, cached_nodes_file); - } + if (vamana_index != nullptr) { + navigation.build_cache(config, context, *vamana_index->get_graph(), + vamana_index->get_entry_point()); } auto e = std::chrono::high_resolution_clock::now(); std::chrono::duration diff = e - s; @@ -2034,11 +2006,14 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context) { template int build_disk_index(BuildConfig &config); template int build_disk_index(BuildConfig &config); template int build_disk_index(BuildConfig &config); - template int build_disk_index(BuildConfig &, PreparedBuildContext &); + template int build_disk_index(BuildConfig &, PreparedBuildContext &, + const NavigationBuilder &); template int build_disk_index(BuildConfig &, - PreparedBuildContext &); + PreparedBuildContext &, + const NavigationBuilder &); template int build_disk_index(BuildConfig &, - PreparedBuildContext &); + PreparedBuildContext &, + const NavigationBuilder &); template std::unique_ptr prepare_build_context( const BuildConfig &); template std::unique_ptr diff --git a/thirdparty/DiskANN/src/navigation_build.cpp b/thirdparty/DiskANN/src/navigation_build.cpp new file mode 100644 index 000000000..2c3043fef --- /dev/null +++ b/thirdparty/DiskANN/src/navigation_build.cpp @@ -0,0 +1,100 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#include "diskann/navigation_build.h" + +#include +#include +#include + +#include "diskann/aux_utils.h" +#include "diskann/defaults.h" +#include "diskann/partition_and_pq.h" +#include "diskann/utils.h" + +namespace diskann { + std::vector pq_navigation_files(const std::string &prefix, + bool rearranged) { + const auto pivots = get_pq_pivots_filename(prefix); + return {pivots, get_pq_rearrangement_perm_filename(pivots), + get_pq_chunk_offsets_filename(pivots), + get_pq_centroid_filename(pivots), + rearranged ? get_pq_compressed_rearranged_filename(prefix) + : get_pq_compressed_filename(prefix)}; + } + + template + class PQNavigationBuilder final : public NavigationBuilder { + public: + bool validate(const BuildConfig &config) const override { + if (get_memory_budget(config.pq_code_size_gb) <= 0) { + LOG_KNOWHERE_ERROR_ << "PQ navigation requires a positive code budget"; + return false; + } + return true; + } + bool needs_training_sample() const override { + return true; + } + bool supports_aisaq() const override { + return true; + } + void build(const BuildConfig &config, const PreparedBuildContext &context, + const NavigationTrainingData &sample) const override { + const auto start = std::chrono::high_resolution_clock::now(); + const auto budget = get_memory_budget(config.pq_code_size_gb); + size_t chunks = + static_cast(std::floor(_u64(budget / context.rows))); + chunks = std::max(1, chunks); + chunks = std::min(chunks, context.prepared_dim); + chunks = std::min(chunks, diskann::defaults::MAX_PQ_CHUNKS); + const auto pivots = get_pq_pivots_filename(context.prefix); + const auto codes = get_pq_compressed_filename(context.prefix); + LOG_KNOWHERE_INFO_ << "Compressing " << context.prepared_dim + << "-dimensional data into " << chunks + << " bytes per vector."; + generate_pq_pivots(sample.data, sample.rows, + (uint32_t) context.prepared_dim, 256, + (uint32_t) chunks, NUM_KMEANS_REPS, pivots, + context.metric == diskann::Metric::L2); + generate_pq_data_from_pivots(context.prepared_source.c_str(), 256, + (uint32_t) chunks, pivots, codes); + const std::chrono::duration elapsed = + std::chrono::high_resolution_clock::now() - start; + LOG_KNOWHERE_INFO_ << "Training PQ codes cost: " << elapsed.count() + << "s"; + } + void build_cache(const BuildConfig &config, + const PreparedBuildContext &context, + const std::vector> &graph, + unsigned entry_point) const override { + const auto sample_file = get_sample_data_filename(context.prefix); + const auto pivots = get_pq_pivots_filename(context.prefix); + const auto codes = get_pq_compressed_filename(context.prefix); + // Keep the native cache-generation policy and its memory allowance. + constexpr float cache_mem_factor = 1.1; + const auto usage = cache_mem_factor * + (get_file_size(context.prefix + "_mem.index") + + get_file_size(sample_file) + get_file_size(codes) + + get_file_size(pivots)) / + (1024 * 1024 * 1024); + if (config.num_nodes_to_cache > 0 && !graph.empty() && + usage < config.index_mem_gb) { + generate_cache_list_from_graph_with_pq( + config.num_nodes_to_cache, config.max_degree, context.metric, + sample_file, pivots, codes, entry_point, graph, + get_cached_nodes_file(context.prefix)); + } + } + }; + + template + std::unique_ptr make_pq_navigation_builder() { + return std::make_unique>(); + } + template std::unique_ptr + make_pq_navigation_builder(); + template std::unique_ptr + make_pq_navigation_builder(); + template std::unique_ptr + make_pq_navigation_builder(); +} // namespace diskann From 48768d185a90590aba22a9ee22b9457247c246e4 Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Tue, 22 Sep 2026 10:08:16 +0000 Subject: [PATCH 10/12] fix: isolate DiskANN build scratch and reserve publication targets before writes Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/build_files.h | 13 +++- src/index/diskann/diskann.cc | 10 ++- src/index/diskann/rabitq_store.cc | 4 +- tests/ut/test_diskann.cc | 73 ++++++++++++++++++- .../DiskANN/include/diskann/aux_utils.h | 3 + thirdparty/DiskANN/src/aux_utils.cpp | 26 ++++++- thirdparty/DiskANN/src/navigation_build.cpp | 2 +- 7 files changed, 125 insertions(+), 6 deletions(-) diff --git a/src/index/diskann/build_files.h b/src/index/diskann/build_files.h index 1843116a9..124ff871c 100644 --- a/src/index/diskann/build_files.h +++ b/src/index/diskann/build_files.h @@ -3,6 +3,7 @@ #pragma once #include +#include #include #include "filemanager/FileManager.h" @@ -32,11 +33,20 @@ class DiskANNBuildRegistration { } } bool - Add(const std::string& path) { + Reserve(const std::string& path) { + // Check before creating local outputs: FileManager implementations may + // consult the local filesystem as well as their registered objects. const auto exists = manager_.IsExisted(path); if (!exists.has_value() || exists.value()) { return false; } + reserved_.insert(path); + return true; + } + bool + Add(const std::string& path) { + if (!reserved_.count(path)) + return false; attempted_.push_back(path); return manager_.AddFile(path); } @@ -47,6 +57,7 @@ class DiskANNBuildRegistration { private: milvus::FileManager& manager_; + std::unordered_set reserved_; std::vector attempted_; bool committed_ = false; }; diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 199d9a28a..5376eb5a6 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -536,6 +536,15 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr(num_nodes_to_cache), build_conf.shuffle_build.value()}; std::unique_ptr context; + DiskANNBuildRegistration registration(*file_manager_); + for (const auto& path : GetNecessaryFilenames(build_conf, index_prefix_, need_norm, true, true)) { + if (!registration.Reserve(path)) + return Status::disk_file_error; + } + for (const auto& path : GetOptionalFilenames(build_conf, index_prefix_)) { + if (!registration.Reserve(path)) + return Status::disk_file_error; + } RETURN_IF_ERROR(TryDiskANNCall([&]() { context = diskann::prepare_build_context(diskann_internal_build_config); for (const auto& path : GetNecessaryFilenames(build_conf, index_prefix_, need_norm, true, true)) { @@ -553,7 +562,6 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr(kRawDataPath, static_cast(base->GetTensor()), 512, 16); + class FailingManager : public milvus::LocalFileManager { + public: + bool fail = true; + bool + AddFile(const std::string& path) override { + const bool result = milvus::LocalFileManager::AddFile(path); + return fail && path.find("_disk.index") != std::string::npos ? false : result; + } + std::optional + IsExisted(const std::string& path) override { + return milvus::LocalFileManager::IsExisted(path).value() || fs::exists(path); + } + }; + const auto version = GenTestVersionList(); + for (const auto* codec : {"PQ", "RABITQ"}) { + const auto prefix = kDir + "/retry_" + codec; + auto manager = std::make_shared(); + auto index = knowhere::IndexFactory::Instance() + .Create(kNativeDiskANN, version, + knowhere::Pack(std::shared_ptr(manager))) + .value(); + knowhere::Json build = {{"dim", 16}, + {"metric_type", "IP"}, + {"index_prefix", prefix}, + {"data_path", kRawDataPath}, + {"max_degree", 16}, + {"search_list_size", 32}, + {"build_dram_budget_gb", 1.0}, + {"pq_code_budget_gb", 0.000004}, + {"disk_pq_dims", 4}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"rbq_bits", 4}, + {"navigation_codec", codec}}; + REQUIRE(index.Build(nullptr, build) == knowhere::Status::disk_file_error); + for (const auto& entry : fs::directory_iterator(kDir)) { + REQUIRE(entry.path().string().find(prefix) != 0); + } + REQUIRE(fs::exists(kRawDataPath)); + REQUIRE_FALSE(manager->IsExisted(diskann::get_disk_index_filename(prefix)).value()); + manager->fail = false; + REQUIRE(index.Build(nullptr, build) == knowhere::Status::success); + REQUIRE_FALSE(fs::exists(prefix + "_build_tmp")); + REQUIRE_FALSE(fs::exists(prefix + "_prepped_base.bin")); + REQUIRE_FALSE(fs::exists(prefix + "_disk.index_pq_compressed.bin")); + knowhere::BinarySet empty; + knowhere::Json load = {{"metric_type", "IP"}, + {"index_prefix", prefix}, + {"search_cache_budget_gb", 0}, + {"search_cache_budget_gb_ratio", 0}, + {"warm_up", false}}; + REQUIRE(index.Deserialize(empty, load) == knowhere::Status::success); + } +} + TEST_CASE("Valid diskann build params test", "[diskann]") { int rows_num = 1000000; auto version = GenTestVersionList(); diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 776f0d135..0db797035 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -144,12 +144,14 @@ namespace diskann { PreparedBuildContext &operator=(const PreparedBuildContext &) = delete; void own_temporary(const std::string &path); void own_output(const std::string &path); + void create_graph_workspace(); void commit_outputs() noexcept { committed_ = true; } const std::string raw_source; const std::string prefix; + const std::string graph_index_path; const diskann::Metric metric; std::string prepared_source; std::string ssd_source; @@ -162,6 +164,7 @@ namespace diskann { std::vector temporaries_; std::vector outputs_; bool committed_ = false; + bool owns_graph_workspace_ = false; }; template diff --git a/thirdparty/DiskANN/src/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index 1cce027e6..d63a865aa 100644 --- a/thirdparty/DiskANN/src/aux_utils.cpp +++ b/thirdparty/DiskANN/src/aux_utils.cpp @@ -1610,6 +1610,7 @@ void create_aisaq_layout(const std::string base_file, const std::string mem_inde PreparedBuildContext::PreparedBuildContext(const BuildConfig &config) : raw_source(config.data_file_path), prefix(config.index_file_path), + graph_index_path(prefix + "_build_tmp/graph"), metric(config.compare_metric), prepared_source(raw_source), ssd_source(raw_source) { get_bin_metadata(raw_source, rows, raw_dim); @@ -1634,6 +1635,18 @@ void PreparedBuildContext::own_output(const std::string &path) { own(path, outputs_); } +void PreparedBuildContext::create_graph_workspace() { + if (owns_graph_workspace_) + return; + const auto directory = + std::filesystem::path(graph_index_path).parent_path(); + if (!std::filesystem::create_directory(directory)) { + throw diskann::ANNException( + "Build workspace already exists: " + directory.string(), -1); + } + owns_graph_workspace_ = true; +} + PreparedBuildContext::~PreparedBuildContext() { auto remove_owned = [](const std::vector &paths) { for (auto it = paths.rbegin(); it != paths.rend(); ++it) { @@ -1645,6 +1658,16 @@ PreparedBuildContext::~PreparedBuildContext() { } }; remove_owned(temporaries_); + if (owns_graph_workspace_) { + // Only this newly created directory is owned, including any partial + // partition/shard files produced by low-memory graph construction. + std::error_code error; + std::filesystem::remove_all( + std::filesystem::path(graph_index_path).parent_path(), error); + if (error) + LOG_KNOWHERE_WARNING_ << "Could not clean graph workspace: " + << error.message(); + } if (!committed_) remove_owned(outputs_); } @@ -1737,7 +1760,8 @@ int build_disk_index(BuildConfig &config, PreparedBuildContext &context, const auto &data_file_to_use = context.prepared_source; const auto &data_file_to_save = context.ssd_source; const auto &index_prefix_path = context.prefix; - const auto mem_index_path = index_prefix_path + "_mem.index"; + context.create_graph_workspace(); + const auto &mem_index_path = context.graph_index_path; const auto disk_index_path = get_disk_index_filename(index_prefix_path); const auto medoids_path = get_disk_index_medoids_filename(disk_index_path); const auto centroids_path = diff --git a/thirdparty/DiskANN/src/navigation_build.cpp b/thirdparty/DiskANN/src/navigation_build.cpp index 2c3043fef..afc3b4ad6 100644 --- a/thirdparty/DiskANN/src/navigation_build.cpp +++ b/thirdparty/DiskANN/src/navigation_build.cpp @@ -73,7 +73,7 @@ namespace diskann { // Keep the native cache-generation policy and its memory allowance. constexpr float cache_mem_factor = 1.1; const auto usage = cache_mem_factor * - (get_file_size(context.prefix + "_mem.index") + + (get_file_size(context.graph_index_path) + get_file_size(sample_file) + get_file_size(codes) + get_file_size(pivots)) / (1024 * 1024 * 1024); From 4155d20e8d223c0a77a62f133d03a98127658b7d Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Thu, 24 Sep 2026 07:24:43 +0000 Subject: [PATCH 11/12] fix: estimate DiskANN load resources without assuming navigation defaults Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/diskann.cc | 44 ++++----- src/index/diskann/diskann_config.h | 4 +- src/index/diskann/navigation_store.cc | 17 ++-- src/index/diskann/navigation_store.h | 5 +- tests/ut/test_diskann.cc | 123 +++++++++++++++++++++++++- 5 files changed, 162 insertions(+), 31 deletions(-) diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 5376eb5a6..0c220cae9 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -170,29 +170,31 @@ class DiskANNIndexNode : public IndexNode { StaticEstimateLoadResource(const uint64_t file_size_in_bytes, const int64_t num_rows, const int64_t dim, const knowhere::BaseConfig& config, const IndexVersion& version) { const auto& disk_config = static_cast(config); - if (UsesExternalNavigation(disk_config)) { - try { - const auto navigation_bytes = EstimateNavigationMemory(disk_config, num_rows, dim); - const long double raw_bytes = static_cast(num_rows) * dim * sizeof(float); - const long double cache_bytes = std::max( - static_cast(disk_config.search_cache_budget_gb.value_or(0)) * (1ULL << 30), - static_cast(disk_config.search_cache_budget_gb_ratio.value_or(0)) * raw_bytes); - // Keep the legacy engine allowance for scratch, PQ tables and - // other loading overhead. Add the codec's persistent storage - // and cache budget explicitly. This is a conservative estimate, - // not an exact RSS prediction or a replacement for Size(). - const long double memory_bytes = - file_size_in_bytes / 4 + static_cast(navigation_bytes) + std::ceil(cache_bytes); - if (!std::isfinite(memory_bytes) || memory_bytes < 0 || - memory_bytes >= static_cast(std::numeric_limits::max())) { - return expected::Err(Status::invalid_args, "DiskANN resource estimate overflows"); - } - return Resource{.memoryCost = static_cast(memory_bytes), .diskCost = file_size_in_bytes}; - } catch (const std::exception& e) { - return expected::Err(Status::invalid_args, e.what()); + if (num_rows < 0 || dim <= 0) { + return expected::Err(Status::invalid_args, "Invalid DiskANN resource estimate dimensions"); + } + try { + const auto navigation_bytes = EstimateNavigationMemory(disk_config, num_rows, dim); + const long double raw_bytes = static_cast(num_rows) * dim * sizeof(float); + const long double cache_bytes = + std::max(static_cast(disk_config.search_cache_budget_gb.value_or(0)) * (1ULL << 30), + static_cast(disk_config.search_cache_budget_gb_ratio.value_or(0)) * raw_bytes); + // Preserve the legacy engine allowance. Without complete model + // parameters, budget the entire serialized payload as resident + // rather than guessing PQ or applying the RBQ build defaults. + // Cache is additional in every branch, including explicit PQ. + // This estimates admission costs, not exact or peak RSS. + const long double memory_bytes = file_size_in_bytes / 4 + + static_cast(navigation_bytes.value_or(file_size_in_bytes)) + + std::ceil(cache_bytes); + if (!std::isfinite(memory_bytes) || memory_bytes < 0 || + memory_bytes >= static_cast(std::numeric_limits::max())) { + return expected::Err(Status::invalid_args, "DiskANN resource estimate overflows"); } + return Resource{.memoryCost = static_cast(memory_bytes), .diskCost = file_size_in_bytes}; + } catch (const std::exception& e) { + return expected::Err(Status::invalid_args, e.what()); } - return Resource{.memoryCost = file_size_in_bytes / 4, .diskCost = file_size_in_bytes}; } Status diff --git a/src/index/diskann/diskann_config.h b/src/index/diskann/diskann_config.h index f3f1cb8ed..d3168d7a3 100644 --- a/src/index/diskann/diskann_config.h +++ b/src/index/diskann/diskann_config.h @@ -219,8 +219,8 @@ class DiskANNNavigationConfig : public DiskANNConfig { .for_deserialize() .for_static(); KNOWHERE_CONFIG_DECLARE_FIELD(rbq_bits) - .description("number of RaBitQ bits per database vector dimension") - .set_default(1) + .description("number of RaBitQ bits per database vector dimension; build defaults to 1") + .allow_empty_without_default() .set_range(1, 9) .for_train() .for_static(); diff --git a/src/index/diskann/navigation_store.cc b/src/index/diskann/navigation_store.cc index 95e062878..41854e34c 100644 --- a/src/index/diskann/navigation_store.cc +++ b/src/index/diskann/navigation_store.cc @@ -36,7 +36,7 @@ class RaBitQNavigationBuilder final : public diskann::NavigationBuilder { struct NavigationCodec { bool external; Status (*validate)(const DiskANNNavigationConfig&, std::string*); - uint64_t (*estimate)(const DiskANNConfig&, int64_t, int64_t); + std::optional (*estimate)(const DiskANNConfig&, int64_t, int64_t); NavigationFileSet (*files)(const std::string&); std::unique_ptr (*builder)(const DiskANNConfig&); std::unique_ptr (*load)(const std::string&); @@ -47,7 +47,7 @@ Codecs() { static const std::unordered_map codecs = { {"PQ", {false, [](const DiskANNNavigationConfig&, std::string*) { return Status::success; }, - [](const DiskANNConfig&, int64_t, int64_t) -> uint64_t { return 0; }, + [](const DiskANNConfig&, int64_t, int64_t) -> std::optional { return 0; }, [](const std::string& prefix) { return NavigationFileSet{diskann::pq_navigation_files(prefix), {}}; }, @@ -69,13 +69,15 @@ Codecs() { } return Status::success; }, - [](const DiskANNConfig& config, int64_t rows, int64_t dim) { + [](const DiskANNConfig& config, int64_t rows, int64_t dim) -> std::optional { + const auto bits = NavigationConfig(config).rbq_bits; + if (!bits.has_value()) + return std::nullopt; if (dim <= 0 || dim >= std::numeric_limits::max()) { throw std::invalid_argument("invalid DiskANN navigation dimension"); } const auto prepared_dim = dim + (config.metric_type.value_or(metric::L2) == metric::IP ? 1 : 0); - return RaBitQStore::EstimateMemorySize( - rows, prepared_dim, static_cast(NavigationConfig(config).rbq_bits.value_or(1))); + return RaBitQStore::EstimateMemorySize(rows, prepared_dim, static_cast(bits.value())); }, [](const std::string& prefix) { return NavigationFileSet{{RaBitQStore::SidecarFilename(prefix)}, {}}; @@ -160,8 +162,11 @@ UsesExternalNavigation(const DiskANNConfig& config) { return Codec(config).external; } -uint64_t +std::optional EstimateNavigationMemory(const DiskANNConfig& config, int64_t rows, int64_t dim) { + const auto* navigation = dynamic_cast(&config); + if (navigation && !navigation->navigation_codec.has_value()) + return std::nullopt; return Codec(config).estimate(config, rows, dim); } diff --git a/src/index/diskann/navigation_store.h b/src/index/diskann/navigation_store.h index a8d100ce2..2cefa9aba 100644 --- a/src/index/diskann/navigation_store.h +++ b/src/index/diskann/navigation_store.h @@ -3,6 +3,7 @@ #pragma once #include +#include #include #include @@ -38,7 +39,9 @@ class NavigationStore { // search parameters. New codecs are added here, not to cached_beam_search. bool UsesExternalNavigation(const DiskANNConfig& config); -uint64_t +// Missing codec/model parameters require a file-size-based static estimate; +// build defaults must not be mistaken for persisted model metadata. +std::optional EstimateNavigationMemory(const DiskANNConfig& config, int64_t rows, int64_t dim); struct NavigationFileSet { std::vector required; diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index 87fdceb49..71d5ec36d 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -806,7 +806,7 @@ TEST_CASE("DiskANN navigation load resource estimate", "[diskann][rabitq][resour config["navigation_codec"] = "PQ"; auto pq = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config); REQUIRE(pq.has_value()); - REQUIRE(pq.value().memoryCost == file_bytes / 4); + REQUIRE(pq.value().memoryCost == file_bytes / 4 + rows * dim * sizeof(float) / 2); REQUIRE(pq.value().diskCost == file_bytes); } } @@ -820,6 +820,127 @@ TEST_CASE("DiskANN navigation load resource estimate", "[diskann][rabitq][resour } } +TEST_CASE("DiskANN static estimate does not infer persisted navigation from build defaults", + "[diskann][rabitq][resource]") { + using Static = knowhere::IndexStaticFaced; + const auto version = GenTestVersionList(); + constexpr uint64_t file_bytes = 1ULL << 30; + constexpr int64_t rows = 1000000, dim = 1536; + constexpr uint64_t fallback = file_bytes + file_bytes / 4; + for (const auto* metric : {"L2", "IP", "COSINE"}) { + for (const auto* type : {kNativeDiskANN, "DISKANN_RABITQ"}) { + knowhere::Json config = {{"metric_type", metric}, {"disk_pq_dims", 16}}; + auto estimate = [&](const knowhere::Json& params, uint64_t bytes = 1ULL << 30) { + return Static::EstimateLoadResource(type, version, bytes, rows, dim, params); + }; + auto missing = estimate(config); + REQUIRE(missing.has_value()); + REQUIRE(missing.value().memoryCost == fallback); + REQUIRE(missing.value().diskCost == file_bytes); + config["navigation_codec"] = "RABITQ"; + auto missing_bits = estimate(config); + REQUIRE(missing_bits.has_value()); + REQUIRE(missing_bits.value().memoryCost == fallback); + config["search_cache_budget_gb"] = 0.25; + auto cache = estimate(config); + REQUIRE(cache.has_value()); + REQUIRE(cache.value().memoryCost == fallback + (1ULL << 28)); + config["search_cache_budget_gb_ratio"] = 0.5; + auto ratio = estimate(config); + REQUIRE(ratio.has_value()); + REQUIRE(ratio.value().memoryCost == fallback + rows * dim * sizeof(float) / 2); + config["rbq_bits"] = 1; + auto known = estimate(config); + REQUIRE(known.has_value()); + REQUIRE(known.value().memoryCost < ratio.value().memoryCost); + config.erase("rbq_bits"); + REQUIRE_FALSE(estimate(config, std::numeric_limits::max()).has_value()); + } + // A bit count without a codec cannot identify the persisted model. + auto only_bits = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, + {{"metric_type", metric}, {"rbq_bits", 9}}); + REQUIRE(only_bits.has_value()); + REQUIRE(only_bits.value().memoryCost == fallback); + } + for (const auto* codec : {"PQ", "RABITQ"}) { + knowhere::Json config = {{"navigation_codec", codec}, {"search_cache_budget_gb", 0.25}}; + auto result = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config); + REQUIRE(result.has_value()); + REQUIRE(result.value().memoryCost == (std::string(codec) == "PQ" ? file_bytes / 4 : fallback) + (1ULL << 28)); + config["search_cache_budget_gb"] = std::numeric_limits::max(); + REQUIRE_FALSE(Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, config).has_value()); + } + for (const auto& shape : {std::pair{-1, dim}, {rows, 0}, {rows, -1}}) { + REQUIRE_FALSE(Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, shape.first, shape.second, {}) + .has_value()); + } +} + +TEST_CASE("DiskANN automatic RBQ load and resource estimate with SSD PQ", "[diskann][rabitq][resource]") { + using Static = knowhere::IndexStaticFaced; + const auto* metric = GENERATE("L2", "IP", "COSINE"); + const auto bits = GENERATE(1, 8, 9); + CAPTURE(metric, bits); + fs::remove_all(kDir); + fs::create_directories(kDir); + constexpr int rows = 512, dim = 256; + auto base = GenDataSet(rows, dim, 77); + auto query = GenDataSet(4, dim, 78); + WriteRawDataToDisk(kRawDataPath, static_cast(base->GetTensor()), rows, dim); + const auto version = GenTestVersionList(); + const auto prefix = kDir + "/estimate"; + auto create = [&]() { + std::shared_ptr manager = std::make_shared(); + return knowhere::IndexFactory::Instance() + .Create(kNativeDiskANN, version, knowhere::Pack(manager)) + .value(); + }; + knowhere::Json build = {{"dim", dim}, + {"metric_type", metric}, + {"navigation_codec", "RABITQ"}, + {"index_prefix", prefix}, + {"data_path", kRawDataPath}, + {"max_degree", 16}, + {"search_list_size", 32}, + {"build_dram_budget_gb", 1.0}, + {"disk_pq_dims", 4}}; + // Exercise the unchanged one-bit BUILD default, distinct from an unknown + // bit count when estimating a previously persisted model. + if (bits != 1) + build["rbq_bits"] = bits; + REQUIRE(create().Build(nullptr, build) == knowhere::Status::success); + uint64_t file_bytes = 0; + for (const auto& entry : fs::directory_iterator(kDir)) { + if (entry.is_regular_file() && entry.path().string().find(prefix + "_") == 0) + file_bytes += entry.file_size(); + } + knowhere::RaBitQStore stored(knowhere::RaBitQStore::SidecarFilename(prefix)); + const auto prepared_dim = dim + (std::string(metric) == "IP"); + REQUIRE(stored.MemorySize() == knowhere::RaBitQStore::EstimateMemorySize(rows, prepared_dim, bits)); + if (bits >= 8) + REQUIRE(file_bytes / 4 < stored.MemorySize()); + knowhere::Json load = {{"metric_type", metric}, + {"index_prefix", prefix}, + {"search_cache_budget_gb", 0.00001}, + {"use_bfs_cache", true}, + {"warm_up", false}}; + auto estimate = Static::EstimateLoadResource(kNativeDiskANN, version, file_bytes, rows, dim, load); + REQUIRE(estimate.has_value()); + REQUIRE(estimate.value().memoryCost >= file_bytes + file_bytes / 4); + REQUIRE(estimate.value().memoryCost > stored.MemorySize()); + auto loaded = create(); + knowhere::BinarySet empty; + REQUIRE(loaded.Deserialize(empty, load) == knowhere::Status::success); + auto result = loaded.Search( + query, {{"metric_type", metric}, {"k", 10}, {"search_list_size", 128}, {"beamwidth", 4}}, nullptr); + REQUIRE(result.has_value()); + for (int i = 0; i < 40; ++i) { + REQUIRE(result.value()->GetIds()[i] >= 0); + REQUIRE(result.value()->GetIds()[i] < rows); + REQUIRE(std::isfinite(result.value()->GetDistance()[i])); + } +} + TEST_CASE("Test DISKANN_RABITQ constraints", "[diskann][rabitq]") { auto version = GenTestVersionList(); auto make_pack = []() { From 1a12266364b25234d2a3ef21c6b3fc238b5129dc Mon Sep 17 00:00:00 2001 From: ChenLiqing <23721160+CLiqing@users.noreply.github.com> Date: Mon, 28 Sep 2026 03:19:23 +0000 Subject: [PATCH 12/12] fix: retain DiskANN build outputs after publication failure Signed-off-by: ChenLiqing <23721160+CLiqing@users.noreply.github.com> --- src/index/diskann/build_files.h | 26 +----- src/index/diskann/diskann.cc | 9 +- tests/ut/test_diskann.cc | 88 +++++++++++++++---- .../DiskANN/include/diskann/aux_utils.h | 4 +- 4 files changed, 79 insertions(+), 48 deletions(-) diff --git a/src/index/diskann/build_files.h b/src/index/diskann/build_files.h index 124ff871c..9bc8cb091 100644 --- a/src/index/diskann/build_files.h +++ b/src/index/diskann/build_files.h @@ -4,14 +4,12 @@ #include #include -#include #include "filemanager/FileManager.h" -#include "knowhere/log.h" namespace knowhere { -// Publication is separate from ownership of local build files. Failed uploads -// can have partial effects, so roll back attempted registrations as well. +// Publication may have partial effects that FileManager cannot roll back. +// The caller must retain completed local outputs before attempting uploads. class DiskANNBuildRegistration { public: explicit DiskANNBuildRegistration(milvus::FileManager& manager) : manager_(manager) { @@ -19,19 +17,6 @@ class DiskANNBuildRegistration { DiskANNBuildRegistration(const DiskANNBuildRegistration&) = delete; DiskANNBuildRegistration& operator=(const DiskANNBuildRegistration&) = delete; - ~DiskANNBuildRegistration() { - if (!committed_) { - for (auto it = attempted_.rbegin(); it != attempted_.rend(); ++it) { - try { - if (!manager_.RemoveFile(*it)) { - LOG_KNOWHERE_WARNING_ << "Failed to roll back DiskANN registration: " << *it; - } - } catch (const std::exception& e) { - LOG_KNOWHERE_WARNING_ << "Failed to roll back DiskANN registration: " << e.what(); - } - } - } - } bool Reserve(const std::string& path) { // Check before creating local outputs: FileManager implementations may @@ -47,18 +32,11 @@ class DiskANNBuildRegistration { Add(const std::string& path) { if (!reserved_.count(path)) return false; - attempted_.push_back(path); return manager_.AddFile(path); } - void - Commit() noexcept { - committed_ = true; - } private: milvus::FileManager& manager_; std::unordered_set reserved_; - std::vector attempted_; - bool committed_ = false; }; } // namespace knowhere diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 0c220cae9..e376eba61 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -563,7 +563,12 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptrcommit_outputs(); for (auto& filename : GetNecessaryFilenames(build_conf, index_prefix_, need_norm, true, true)) { if (!registration.Add(filename)) { LOG_KNOWHERE_ERROR_ << "Failed to add file " << filename << "."; @@ -577,8 +582,6 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptrcommit_outputs(); count_.store(count); dim_.store(dim); is_prepared_.store(false); diff --git a/tests/ut/test_diskann.cc b/tests/ut/test_diskann.cc index 71d5ec36d..7b5e22a62 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -152,7 +152,7 @@ TEST_CASE("DiskANN prepared build files have scoped ownership", "[diskann][build } } -TEST_CASE("DiskANN registration rolls back a partial failure", "[diskann][build_context]") { +TEST_CASE("DiskANN registration preserves partial publication state", "[diskann][build_context]") { class FailingManager : public milvus::LocalFileManager { public: bool fail = true; @@ -172,16 +172,14 @@ TEST_CASE("DiskANN registration rolls back a partial failure", "[diskann][build_ REQUIRE_FALSE(files.Add("second")); } REQUIRE(manager.IsExisted("existing").value()); - REQUIRE_FALSE(manager.IsExisted("first").value()); - REQUIRE_FALSE(manager.IsExisted("second").value()); + REQUIRE(manager.IsExisted("first").value()); + REQUIRE(manager.IsExisted("second").value()); manager.fail = false; { knowhere::DiskANNBuildRegistration files(manager); - REQUIRE(files.Reserve("first")); - REQUIRE(files.Reserve("second")); - REQUIRE(files.Add("first")); - REQUIRE(files.Add("second")); - files.Commit(); + REQUIRE_FALSE(files.Reserve("first")); + REQUIRE_FALSE(files.Reserve("second")); + REQUIRE_FALSE(files.Add("unreserved")); } REQUIRE(manager.IsExisted("first").value()); REQUIRE(manager.IsExisted("second").value()); @@ -292,7 +290,7 @@ TEST_CASE("DiskANN navigation discovery preserves FileManager errors", "[diskann REQUIRE(knowhere::DetectNavigationCodec(config, prefix, manager).value() == "RABITQ"); } -TEST_CASE("DiskANN failed publication cleans outputs and permits retry", "[diskann][build_context]") { +TEST_CASE("DiskANN failed publication retains outputs and blocks same-prefix retry", "[diskann][build_context]") { fs::remove_all(kDir); fs::create_directories(kDir); auto base = GenDataSet(512, 16, 21); @@ -300,20 +298,48 @@ TEST_CASE("DiskANN failed publication cleans outputs and permits retry", "[diska class FailingManager : public milvus::LocalFileManager { public: bool fail = true; + bool throw_on_failure = false; + std::string fail_path; + size_t add_calls = 0; + size_t remove_calls = 0; + uintmax_t registered_bytes = 0; bool AddFile(const std::string& path) override { + ++add_calls; + registered_bytes += fs::file_size(path); const bool result = milvus::LocalFileManager::AddFile(path); - return fail && path.find("_disk.index") != std::string::npos ? false : result; + if (fail && path == fail_path) { + if (throw_on_failure) + throw std::runtime_error("injected partial upload failure"); + return false; + } + return result; + } + bool + RemoveFile(const std::string&) override { + ++remove_calls; + return false; } std::optional IsExisted(const std::string& path) override { - return milvus::LocalFileManager::IsExisted(path).value() || fs::exists(path); + // Match consumers whose existence check is local but whose upload + // bookkeeping cannot be rolled back by RemoveFile. + return fs::exists(path); } }; + const bool fail_first = GENERATE(false, true); + const bool throw_on_failure = GENERATE(false, true); const auto version = GenTestVersionList(); for (const auto* codec : {"PQ", "RABITQ"}) { const auto prefix = kDir + "/retry_" + codec; + const auto disk_path = diskann::get_disk_index_filename(prefix); + const auto disk_pivots = diskann::get_disk_index_pq_pivots_filename(disk_path); + auto navigation_files = std::string(codec) == "PQ" + ? diskann::pq_navigation_files(prefix) + : std::vector{knowhere::RaBitQStore::SidecarFilename(prefix)}; auto manager = std::make_shared(); + manager->fail_path = fail_first ? navigation_files.front() : disk_pivots; + manager->throw_on_failure = throw_on_failure; auto index = knowhere::IndexFactory::Instance() .Create(kNativeDiskANN, version, knowhere::Pack(std::shared_ptr(manager))) @@ -331,24 +357,48 @@ TEST_CASE("DiskANN failed publication cleans outputs and permits retry", "[diska {"search_cache_budget_gb_ratio", 0}, {"rbq_bits", 4}, {"navigation_codec", codec}}; - REQUIRE(index.Build(nullptr, build) == knowhere::Status::disk_file_error); - for (const auto& entry : fs::directory_iterator(kDir)) { - REQUIRE(entry.path().string().find(prefix) != 0); - } + CAPTURE(codec, fail_first, throw_on_failure); + REQUIRE(index.Build(nullptr, build) != knowhere::Status::success); REQUIRE(fs::exists(kRawDataPath)); - REQUIRE_FALSE(manager->IsExisted(diskann::get_disk_index_filename(prefix)).value()); - manager->fail = false; - REQUIRE(index.Build(nullptr, build) == knowhere::Status::success); + REQUIRE(manager->add_calls >= (fail_first ? 1 : 2)); + REQUIRE(manager->remove_calls == 0); + REQUIRE(manager->milvus::LocalFileManager::IsExisted(manager->fail_path).value()); + navigation_files.insert( + navigation_files.end(), + {disk_path, diskann::get_disk_index_max_base_norm_file(disk_path), + diskann::get_sample_data_filename(prefix), disk_pivots, + diskann::get_pq_rearrangement_perm_filename(disk_pivots), + diskann::get_pq_chunk_offsets_filename(disk_pivots), diskann::get_pq_centroid_filename(disk_pivots)}); + for (const auto& path : navigation_files) { + REQUIRE(fs::exists(path)); + REQUIRE(fs::file_size(path) > 0); + } REQUIRE_FALSE(fs::exists(prefix + "_build_tmp")); REQUIRE_FALSE(fs::exists(prefix + "_prepped_base.bin")); REQUIRE_FALSE(fs::exists(prefix + "_disk.index_pq_compressed.bin")); + const auto calls_before_retry = manager->add_calls; + const auto bytes_before_retry = manager->registered_bytes; + manager->fail = false; + REQUIRE(index.Build(nullptr, build) == knowhere::Status::disk_file_error); + REQUIRE(manager->add_calls == calls_before_retry); + REQUIRE(manager->registered_bytes == bytes_before_retry); + // A fresh node and manager must also refuse to overwrite this prefix. + auto fresh_manager = std::make_shared(); + auto fresh = knowhere::IndexFactory::Instance() + .Create(kNativeDiskANN, version, + knowhere::Pack(std::shared_ptr(fresh_manager))) + .value(); + REQUIRE(fresh.Build(nullptr, build) == knowhere::Status::disk_file_error); + REQUIRE(fresh_manager->add_calls == 0); knowhere::BinarySet empty; knowhere::Json load = {{"metric_type", "IP"}, {"index_prefix", prefix}, {"search_cache_budget_gb", 0}, {"search_cache_budget_gb_ratio", 0}, {"warm_up", false}}; - REQUIRE(index.Deserialize(empty, load) == knowhere::Status::success); + // Local generation completed before upload began; retained files are + // loadable even though publication failed and Build reported failure. + REQUIRE(fresh.Deserialize(empty, load) == knowhere::Status::success); } } diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 0db797035..e2bc0f4b6 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -134,8 +134,8 @@ namespace diskann { }; // One build owns its prepared input and intermediate files. The original - // input is never owned. Outputs are retained only after successful - // publication. + // input is never owned. Commit completed local outputs before publication: + // failed uploads cannot be rolled back by every FileManager implementation. class PreparedBuildContext { public: explicit PreparedBuildContext(const BuildConfig &config);