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/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/build_files.h b/src/index/diskann/build_files.h new file mode 100644 index 000000000..9bc8cb091 --- /dev/null +++ b/src/index/diskann/build_files.h @@ -0,0 +1,42 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include + +#include "filemanager/FileManager.h" + +namespace knowhere { +// 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) { + } + DiskANNBuildRegistration(const DiskANNBuildRegistration&) = delete; + DiskANNBuildRegistration& + operator=(const DiskANNBuildRegistration&) = delete; + bool + 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; + return manager_.AddFile(path); + } + + private: + milvus::FileManager& manager_; + std::unordered_set reserved_; +}; +} // namespace knowhere diff --git a/src/index/diskann/diskann.cc b/src/index/diskann/diskann.cc index 936c031b7..e376eba61 100644 --- a/src/index/diskann/diskann.cc +++ b/src/index/diskann/diskann.cc @@ -11,6 +11,9 @@ #include "knowhere/feder/DiskANN.h" +#include + +#include #include #include #include @@ -22,7 +25,9 @@ #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" #include "knowhere/context.h" #include "knowhere/dataset.h" @@ -104,6 +109,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); @@ -112,6 +121,19 @@ 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; + } + // 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; + } + } 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 + "'"; @@ -122,6 +144,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); } @@ -144,7 +169,32 @@ 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) { - return Resource{.memoryCost = file_size_in_bytes / 4, .diskCost = file_size_in_bytes}; + const auto& disk_config = static_cast(config); + 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()); + } } Status @@ -167,7 +217,7 @@ class DiskANNIndexNode : public IndexNode { static std::unique_ptr StaticCreateConfig() { - return std::make_unique(); + return std::make_unique(); } std::unique_ptr @@ -200,7 +250,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 (navigation_store_ != nullptr) { + size += navigation_store_->MemorySize(); + } + return size; } int64_t @@ -283,6 +337,8 @@ class DiskANNIndexNode : public IndexNode { std::atomic_bool is_prepared_; 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_; @@ -359,17 +415,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) { - 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); - 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)); @@ -381,8 +431,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)); @@ -398,7 +448,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)) { @@ -407,8 +457,8 @@ AnyIndexFileExist(const std::string& index_prefix) { } return false; }; - return file_exist(GetNecessaryFilenames(index_prefix, diskann::INNER_PRODUCT, true, true)) || - file_exist(GetOptionalFilenames(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 @@ -431,7 +481,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; @@ -440,7 +494,7 @@ DiskANNIndexNode::Build(const DataSetPtr dataset, std::shared_ptr::Build(const DataSetPtr dataset, std::shared_ptr::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, @@ -493,27 +537,53 @@ 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([&]() { - 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(build_conf, index_prefix_, need_norm, true, true)) { + context->own_output(path); + } + for (const auto& path : GetOptionalFilenames(build_conf, index_prefix_)) { + context->own_output(path); + } + 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); + } })); - // Add file to the file manager - for (auto& filename : GetNecessaryFilenames(index_prefix_, need_norm, true, true)) { - if (!AddFile(filename)) { + // Local generation is complete. Retain formal outputs even if publication + // fails: FileManager may already have uploaded files or updated its state, + // and RemoveFile is not supported by every consumer. Existing local files + // block an automatic same-prefix retry with that partially used manager. + // The context still releases preprocessing and graph/SSD scratch files. + context->commit_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 << "."; return Status::disk_file_error; } } - for (auto& filename : GetOptionalFilenames(index_prefix_)) { - if (file_exists(filename) && !AddFile(filename)) { + 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; } } + count_.store(count); + dim_.store(dim); is_prepared_.store(false); return Status::success; } @@ -530,6 +600,10 @@ DiskANNIndexNode::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 (UsesExternalNavigation(static_cast(*cfg))) { + 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); @@ -586,18 +660,44 @@ DiskANNIndexNode::BuildEmbListIfNeed(const DataSetPtr dataset, std::sh template Status DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr cfg) { - auto prep_conf = static_cast(*cfg); + 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); + } + }); 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; } 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); @@ -613,13 +713,14 @@ 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(), + 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())) { 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 << "."; @@ -630,6 +731,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(); @@ -640,7 +751,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()); + 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); } @@ -658,6 +775,21 @@ DiskANNIndexNode::Deserialize(const BinarySet& binset, std::shared_ptr dim_.store(pq_flash_index_->get_data_dim()); } + if (external_navigation) { + try { + 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 navigation sidecar: " << e.what(); + navigation_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_); @@ -669,19 +801,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), prep_conf.max_degree.value()); - } 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()); - } + 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."; @@ -689,7 +811,11 @@ 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() || 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"; if (TryDiskANNCall([&]() { pq_flash_index_->cache_bfs_levels(num_nodes_to_cache, node_list); }) != Status::success) { @@ -711,10 +837,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 @@ -739,9 +875,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()); })); } @@ -757,6 +894,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; @@ -772,6 +910,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 (UsesExternalNavigation(static_cast(*cfg))) { + 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,12 +975,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 (navigation_store_) { + 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"); @@ -883,7 +1029,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"); } @@ -907,16 +1053,18 @@ 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 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); + bitset_, filter_ratio, navigation.get()); #ifdef NOT_COMPILE_FOR_SWIG knowhere_diskann_search_hops.Observe(stats.n_hops); #endif @@ -927,6 +1075,40 @@ DiskANNIndexNode::Search(const DataSetPtr dataset, std::unique_ptr::Err(Status::diskann_inner_error, "some search failed"); } + if (navigation_store_) { + 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 navigation 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); @@ -1028,6 +1210,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(); @@ -1076,6 +1261,27 @@ 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(); + } + + std::string + Type() const override { + return knowhere::IndexEnum::INDEX_DISKANN_RABITQ; + } +}; + #ifdef KNOWHERE_WITH_CARDINAL KNOWHERE_SIMPLE_REGISTER_DENSE_FLOAT_ALL_GLOBAL(DISKANN_DEPRECATED, DiskANNIndexNode, knowhere::feature::DISK | knowhere::feature::EMB_LIST) @@ -1083,4 +1289,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_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 34a9079f8..d3168d7a3 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; @@ -116,7 +120,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) @@ -126,13 +131,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) @@ -195,5 +202,62 @@ class DiskANNConfig : public BaseConfig { return Status::success; } }; + +// 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; + + KNOWHERE_DECLARE_CONFIG(DiskANNNavigationConfig) { + KNOWHERE_CONFIG_DECLARE_FIELD(navigation_codec) + .description("navigation codec: build defaults to PQ; load detects stored codec unless constrained") + .allow_empty_without_default() + .for_train() + .for_deserialize() + .for_static(); + KNOWHERE_CONFIG_DECLARE_FIELD(rbq_bits) + .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(); + 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(); + } + + 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; + } + return ValidateNavigationConfig(*this, err_msg); + } +}; + +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..41854e34c --- /dev/null +++ b/src/index/diskann/navigation_store.cc @@ -0,0 +1,199 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// 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& +NavigationConfig(const DiskANNConfig& config) { + return static_cast(config); +} + +class RaBitQNavigationBuilder final : public diskann::NavigationBuilder { + public: + explicit RaBitQNavigationBuilder(uint8_t bits) : bits_(bits) { + } + void + build(const diskann::BuildConfig&, const diskann::PreparedBuildContext& context, + const diskann::NavigationTrainingData&) const override { + RaBitQStore::BuildFromFloatBin(context.prepared_source, RaBitQStore::SidecarFilename(context.prefix), bits_); + } + + private: + uint8_t bits_; +}; + +struct NavigationCodec { + bool external; + Status (*validate)(const DiskANNNavigationConfig&, std::string*); + 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&); +}; + +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) -> std::optional { 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) -> 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(bits.value())); + }, + [](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 Codec(config).external; +} + +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); +} + +NavigationFileSet +NavigationFiles(const DiskANNConfig& config, const std::string& prefix) { + return Codec(config).files(prefix); +} + +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) { + return Codec(config).load(prefix); +} +} // namespace knowhere diff --git a/src/index/diskann/navigation_store.h b/src/index/diskann/navigation_store.h new file mode 100644 index 000000000..2cefa9aba --- /dev/null +++ b/src/index/diskann/navigation_store.h @@ -0,0 +1,58 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#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 +// 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); +// 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; + std::vector optional; +}; +NavigationFileSet +NavigationFiles(const DiskANNConfig& config, 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/src/index/diskann/rabitq_store.cc b/src/index/diskann/rabitq_store.cc new file mode 100644 index 000000000..b694b26ea --- /dev/null +++ b/src/index/diskann/rabitq_store.cc @@ -0,0 +1,397 @@ +// 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" +#include "index/diskann/diskann_config.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 RaBitQNavigationDistanceComputer final : public diskann::NavigationDistanceComputer { + public: + RaBitQNavigationDistanceComputer(const faiss::RandomRotationMatrix* rotation, const faiss::IndexRaBitQ* rabitq, + uint8_t query_bits) + : rotation_(rotation), + rabitq_(rabitq), + 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 = 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; + }; + + const auto process_estimate = [&](_u64 i, const uint8_t* code, float estimate) { + if (stats != nullptr) { + ++stats->n_approx_estimates; + } + 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; + } + return; + } + 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(); + } + }; + + // 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]; + 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_; + 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; + if (std::filesystem::exists(temporary_path)) { + throw std::runtime_error("Refusing to overwrite RaBitQ temporary file: " + temporary_path); + } + 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(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 && query_metric != metric::COSINE) { + throw std::invalid_argument("RaBitQ navigation supports L2, IP and COSINE"); + } + 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(static_cast(qb)); +} + +std::unique_ptr +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_, 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); +} + +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 new file mode 100644 index 000000000..b1d8fa9a1 --- /dev/null +++ b/src/index/diskann/rabitq_store.h @@ -0,0 +1,74 @@ +// 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 "index/diskann/navigation_store.h" + +namespace faiss { +struct Index; +struct IndexPreTransform; +struct IndexRaBitQ; +struct RandomRotationMatrix; +} // namespace faiss + +namespace knowhere { + +class RaBitQStore final : public NavigationStore { + 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); + + // 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(); + + RaBitQStore(const RaBitQStore&) = delete; + RaBitQStore& + operator=(const RaBitQStore&) = delete; + + std::unique_ptr + CreateDistanceComputer(uint8_t query_bits = 4) const; + + std::unique_ptr + CreateDistanceComputer(const DiskANNConfig& config) const override; + + int64_t + Count() const override; + + int64_t + Dimension() const override; + + uint8_t + Bits() const; + + size_t + CodeSize() const; + + size_t + MemorySize() const override; + + 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/python/test_diskann_rabitq.py b/tests/python/test_diskann_rabitq.py new file mode 100644 index 000000000..43fcf17c1 --- /dev/null +++ b/tests/python/test_diskann_rabitq.py @@ -0,0 +1,64 @@ +# 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", "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 2b6f9b78e..7b5e22a62 100644 --- a/tests/ut/test_diskann.cc +++ b/tests/ut/test_diskann.cc @@ -11,19 +11,37 @@ #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/impl/io.h" +#include "faiss/index_io.h" +#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" #include "knowhere/comp/knowhere_check.h" #include "knowhere/expected.h" @@ -65,6 +83,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; @@ -80,6 +103,305 @@ constexpr float kL2RangeAp = 0.9; 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.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)); + { + 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. + 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)); + REQUIRE_FALSE(fs::exists(norm)); + REQUIRE_FALSE(fs::exists(config.index_file_path + "_mem.index")); + REQUIRE_FALSE(fs::exists(config.index_file_path + "_build_tmp")); + { + diskann::PreparedBuildContext context(config); + context.create_graph_workspace(); + fs::copy_file(kRawDataPath, context.graph_index_path + "_tempFiles_subshard-0.bin"); + } + REQUIRE_FALSE(fs::exists(config.index_file_path + "_build_tmp")); + // 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 preserves partial publication state", "[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.Reserve("existing")); + REQUIRE(files.Reserve("first")); + REQUIRE(files.Reserve("second")); + REQUIRE(files.Add("first")); + REQUIRE_FALSE(files.Add("second")); + } + REQUIRE(manager.IsExisted("existing").value()); + REQUIRE(manager.IsExisted("first").value()); + REQUIRE(manager.IsExisted("second").value()); + manager.fail = false; + { + knowhere::DiskANNBuildRegistration files(manager); + 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()); +} + +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("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); + WriteRawDataToDisk(kRawDataPath, static_cast(base->GetTensor()), 512, 16); + 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); + 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 { + // 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))) + .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}}; + CAPTURE(codec, fail_first, throw_on_failure); + REQUIRE(index.Build(nullptr, build) != knowhere::Status::success); + REQUIRE(fs::exists(kRawDataPath)); + 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}}; + // 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); + } +} + TEST_CASE("Valid diskann build params test", "[diskann]") { int rows_num = 1000000; auto version = GenTestVersionList(); @@ -474,6 +796,927 @@ 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", "COSINE"}) { + 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 + rows * dim * sizeof(float) / 2); + 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("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 = []() { + 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); + { + 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(kNativeDiskANN, 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; + 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::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::success); + invalid = valid; + invalid["search_cache_budget_gb"] = 0.01; + check_train_config(invalid, knowhere::Status::success); + + // 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); + 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(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); })); + + // 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); +} + +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); + 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); + 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(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(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(9)); + } + } + 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(qb); + auto batch = store.CreateDistanceComputer(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(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"; + 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["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. + deserialize_json["use_bfs_cache"] = false; + } + 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 metadata_reader = std::make_shared(); + diskann::PQFlashIndex metadata_only_index(metadata_reader, diskann::Metric::L2); + 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{}; + std::array distances{}; + 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(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(kNativeDiskANN, version, pack).value(); + auto generic_load = deserialize_json; + // 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; + 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]); + } + } + } + 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); + + 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()); + + 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); + 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); +} + +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("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}) { + 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); + } + } + } + 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; + 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(index_type, 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(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); +} + #ifdef KNOWHERE_WITH_CUVS TEST_CASE("Test DiskANN Build Index", "[diskann]") { auto version = GenTestVersionList(); diff --git a/tests/ut/test_diskann_ssd_pq.cc b/tests/ut/test_diskann_ssd_pq.cc new file mode 100644 index 000000000..09163abf1 --- /dev/null +++ b/tests/ut/test_diskann_ssd_pq.cc @@ -0,0 +1,296 @@ +// 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 "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: + explicit DiskPQReference(diskann::Metric metric) + : PQFlashIndex(std::make_shared(), metric) { + } + + 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); + 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 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")); + // 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 + "_" + 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(); + 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())); + const char* index_type = codec == 2 ? "DISKANN_RABITQ" : kDiskIndexType; + 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", 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") == (disk_pq_dims > 0)); + 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); + 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; + 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); + if (external) + reference.CheckExternalResources(); + if (cached) { + std::vector ids(60); + 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); + 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 < 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)); + } + } + } + // 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(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); +} +#endif diff --git a/thirdparty/DiskANN/include/diskann/aux_utils.h b/thirdparty/DiskANN/include/diskann/aux_utils.h index 078007839..e2bc0f4b6 100644 --- a/thirdparty/DiskANN/include/diskann/aux_utils.h +++ b/thirdparty/DiskANN/include/diskann/aux_utils.h @@ -130,9 +130,52 @@ namespace diskann { bool aisaq_mode = false; uint32_t inline_pq = 0; bool rearrange = false; - int num_entry_points = 0; + int num_entry_points = 0; }; + // One build owns its prepared input and intermediate files. The original + // 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); + ~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 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; + 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; + bool owns_graph_workspace_ = false; + }; + + template + std::unique_ptr prepare_build_context( + const BuildConfig &config); + + class NavigationBuilder; + template + 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/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..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(); @@ -183,18 +187,43 @@ namespace diskann { virtual ~PQDataGetter() {} }; + // 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 ~NavigationDistanceComputer() = 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: + 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); + int load(uint32_t num_threads, const char *index_prefix, + bool load_pq_data = true, + const NavigationMetadata* navigation_metadata = nullptr); 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, @@ -202,7 +231,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 +240,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, + 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); @@ -293,7 +324,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, + 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 @@ -351,6 +383,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; @@ -364,6 +398,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/aux_utils.cpp b/thirdparty/DiskANN/src/aux_utils.cpp index f575cd454..d63a865aa 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 @@ -18,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" @@ -33,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; @@ -1607,8 +1608,73 @@ 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); + 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_); +} + +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) { + 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 (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_); +} + template - int build_disk_index(BuildConfig &config) { +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)) { @@ -1619,33 +1685,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 @@ -1664,7 +1714,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_ @@ -1681,16 +1730,77 @@ 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); + 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, + 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; + 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; + 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 = + 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); + 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); 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."; + 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"; @@ -1727,30 +1837,24 @@ template 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)"; - - size_t num_pq_chunks = - (size_t) (std::floor)(_u64(pq_code_size_limit / points_num)); + 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)"; - 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, 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 (navigation.needs_training_sample() || 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,28 +1875,8 @@ template data_file_to_use.c_str(), 256, (uint32_t) disk_pq_dims, disk_pq_pivots_path, disk_pq_compressed_vectors_path); } - 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) && \ @@ -1867,36 +1951,15 @@ template sample_sampling_rate); if (vamana_index != nullptr) { - 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); - } + 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; LOG_KNOWHERE_INFO_ << "Indexing time: " << diff.count(); - if (config.compare_metric == diskann::Metric::INNER_PRODUCT) { - 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, @@ -1967,6 +2030,20 @@ 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 &, + const NavigationBuilder &); + template int build_disk_index(BuildConfig &, + PreparedBuildContext &, + const NavigationBuilder &); + template int build_disk_index(BuildConfig &, + PreparedBuildContext &, + const NavigationBuilder &); + 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( diff --git a/thirdparty/DiskANN/src/navigation_build.cpp b/thirdparty/DiskANN/src/navigation_build.cpp new file mode 100644 index 000000000..afc3b4ad6 --- /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.graph_index_path) + + 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 diff --git a/thirdparty/DiskANN/src/pq_flash_index.cpp b/thirdparty/DiskANN/src/pq_flash_index.cpp index 1e41b2022..4deeb0f99 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 @@ -166,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++; } @@ -187,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(); } } @@ -243,13 +248,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, @@ -277,10 +284,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 += @@ -519,9 +528,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 +712,10 @@ 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, + 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 = @@ -709,14 +727,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 @@ -726,19 +751,41 @@ 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; - diskann::load_bin<_u8>(pq_compressed_vectors, this->data, 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); + 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); + if (load_pq_data) { + 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) << (load_pq_data ? "Loaded resident PQ navigation" + : "Using external navigation metadata") + << ". #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)) { @@ -746,6 +793,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 * @@ -856,7 +906,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 " @@ -942,13 +992,25 @@ namespace diskann { IOContext &ctx, QueryStats *stats, const knowhere::feder::diskann::FederResultUniq &feder, knowhere::BitsetView bitset_view, - PQDataGetter* pq_data_getter) { + PQDataGetter* pq_data_getter, + 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", + -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 +1037,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]); } @@ -995,8 +1063,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; } @@ -1036,8 +1103,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); @@ -1087,7 +1153,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, + NavigationDistanceComputer* 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__); @@ -1097,15 +1170,34 @@ 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(); + // 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(); - 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( + data.scratch.aligned_query_float); + } size_t bv_cnt = 0; @@ -1129,18 +1221,13 @@ 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; } 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); - this->thread_data.push(data); - this->thread_data.push_notify_all(); - this->reader->put_ctx(ctx); + beam_width, ctx, stats, feder, bitset_view, this, + approx_distance_computer); return; } } @@ -1148,10 +1235,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); - this->thread_data.push(data); - this->thread_data.push_notify_all(); - this->reader->put_ctx(ctx); + beam_width, ctx, stats, feder, bitset_view, this, + approx_distance_computer); return; } @@ -1177,23 +1262,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 +1316,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 +1364,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 +1379,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; @@ -1336,19 +1438,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)); @@ -1362,7 +1453,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 +1498,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 @@ -1513,9 +1613,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(); @@ -1529,6 +1626,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(); @@ -1536,6 +1638,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; @@ -1582,7 +1685,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 @@ -1652,9 +1755,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); } } } @@ -1701,6 +1804,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()) { @@ -1902,20 +2008,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); } @@ -2053,11 +2148,18 @@ 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)); 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; 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); }