diff --git a/docs/docs/en/src/quantization/rabitq_split.md b/docs/docs/en/src/quantization/rabitq_split.md index 8e541b7953..f987a18e96 100644 --- a/docs/docs/en/src/quantization/rabitq_split.md +++ b/docs/docs/en/src/quantization/rabitq_split.md @@ -46,6 +46,7 @@ The relevant parameters are: | `rabitq_bits_per_dim_query` | Must be `32` for split storage. | | `rabitq_error_rate` | Default positive multiplier applied to the lower-bound error term. | | `use_reorder` | Should be `true` so candidates are ranked with the `x+y` distance. | +| `rabitq_fused_datacell` | HGraph only; enables fused graph/code layout. Default: `false`. | The constraints are: @@ -58,6 +59,29 @@ x + y <= 8 If `rabitq_bits_per_dim_precise` is omitted, HGraph and Pyramid use the standard RaBitQ path instead of split storage. +### HGraph fused in-memory layout + +For HGraph only, set `rabitq_fused_datacell` to `true` to store each +bottom-layer node's neighbors, cluster id, label, x-bit code, and y-bit +supplement in one cache-line-aligned record. Pyramid uses ordinary split +storage; `rabitq_fused_datacell` is not a Pyramid parameter. The specialized +HGraph search loop reads the record directly and prefetches graph links and +quantized codes together. The codec uses 16 reproducibly trained residual +clusters. + +The fused layout is opt-in and has stricter constraints than ordinary split +storage: + +- `1 <= x <= 4`, `y >= 1`, and `x + y <= 8`. +- The metric must be L2 or inner product. +- The graph, filter codes, and supplement codes must all use memory IO. +- MCI, `deduplicate_storage`, and force remove must be disabled. +- PCA is not supported in fused v1; omit `rabitq_pca_dim` or set it to `0`. +- The legacy v0.14 serialization format is not supported. + +Indexes created without this option keep their existing layout, behavior, and +serialization format. + Enable the filter/lower-bound search path with: ```json @@ -71,11 +95,25 @@ Enable the filter/lower-bound search path with: } ``` +For Pyramid, put the equivalent search controls under `pyramid`: + +```json +{ + "pyramid": { + "ef_search": 200, + "rabitq_one_bit_search": true, + "rabitq_error_rate": 1.9 + } +} +``` + The external search key is named `rabitq_one_bit_search`, but on a split index it uses all `x` filter bits configured by `rabitq_bits_per_dim_base`. -`hgraph.rabitq_error_rate` overrides the index default for that search. It can -be swept without rebuilding because the stored record contains the geometric -error scale before this multiplier is applied. +`hgraph.rabitq_error_rate` and `pyramid.rabitq_error_rate` override the index +default for their respective searches without requiring a rebuild. Native +fused HGraph records store the geometric error scale before this multiplier is +applied. HNSW-compatible fused `1+7` records retain metadata scaled by the +canonical default and apply an override as a ratio at query time. ## Search pipeline @@ -260,19 +298,29 @@ sum_i q_i * u_i + sum_i q_i * s_i ``` -For L2 with an x-bit lookup filter, HGraph and Pyramid pass the previously computed -filter distance to reorder as a hint. `ComputeDistWithSplitCodeAndFilterDist` -recovers the first term from that hint and computes only the second term from -the y supplement planes: +For `x >= 2`, the canonical HGraph and Pyramid graph-search path carries the +exact x-bit filter inner product from traversal to reorder. Ordinary split +storage exposes it through `QueryWithDistanceLowerBoundAndFilterIP`, and +reorder consumes it through `QueryWithFilterIPHint` and +`ComputeDistWithSplitCodeAndFilterIP`. Fused HGraph uses the same exact hint +semantics while reading the code directly from the node record. These +canonical paths do not recover the inner product from a distance, and full +rerank computes only the second term from the y supplement planes: ```text full contribution = shifted filter contribution + supplement contribution ``` -Thus a `3+5` index reuses the 3-bit filter result and scans only 5 new bit -planes for each reordered candidate. If the hint is unavailable or cannot be -used, the code falls back to `ComputeDistWithSplitCode`, which computes the -same final distance directly from both split records. +Thus ordinary and fused `2+y`, `3+y`, and `4+y` indexes reuse the exact x-bit +filter inner product and scan only the y supplement planes for each reranked +candidate. `QueryWithDistanceHint` and +`ComputeDistWithSplitCodeAndFilterDist` remain compatibility APIs for callers +that only have a filter distance; they are not the canonical graph-search +pipeline. The fused `1+y` traversal uses a four-bit query bit-plane and +popcount approximation; precise reranking recomputes its exact one-bit +contribution because the approximate value is not an exact full-distance +hint. If a usable hint is unavailable, the code computes the same final +distance directly from both split records. ## Memory, disk, and hybrid IO @@ -338,30 +386,46 @@ The split datacell serializes, in order: Create the destination index with parameters compatible with the serialized index, especially `dim`, `metric_type`, x/y bit widths, and query bits. Changing an encoded parameter requires rebuilding the index. Tuning only the -search-time `hgraph.rabitq_error_rate` does not. +search-time `hgraph.rabitq_error_rate` or `pyramid.rabitq_error_rate` does not. + +For a fused index, the codec model is serialized with the split datacell and +the per-node codes are serialized once as part of the bottom-graph slab. +Ordinary and streaming round trips preserve this layout without creating a +second count-scaled copy of the split codes. ## Implementation map | Area | File / entry point | | --- | --- | -| External x/y parameter mapping | `src/algorithm/hgraph/hgraph_param_mapping.cpp` | +| External x/y parameter mapping | `hgraph_param_mapping.cpp`, `pyramid.cpp` | | Split record ownership and IO | `src/datacell/rabitq_split_datacell.h` | | Plane layout and code splitting | `RaBitQuantizer::StoredPlaneIndex`, `SplitCode` | | Filter estimate and lower bound | `ComputeDistWithOneBitLowerBound` | | Direct split distance | `ComputeDistWithSplitCode` | -| Reorder using the filter hint | `ComputeDistWithSplitCodeAndFilterDist` | +| Reorder using the filter hint | `ComputeDistWithSplitCodeAndFilterDist`, `ComputeDistWithSplitCodeAndFilterIP` | | SIMD dispatch | `src/simd/rabitq_simd.cpp` | | AVX2 / AVX512 lookup kernels | `src/simd/avx2.cpp`, `src/simd/avx512.cpp` | | Runnable memory/disk/hybrid example | `examples/cpp/323_index_hgraph_rabitq_split.cpp` | ## Operational notes -- Split storage is currently available on HGraph and Pyramid and requires fp32 query codes. Pyramid enables the one-bit split search path by default for split indexes; pass `rabitq_one_bit_search: false` under `pyramid` to force the standard search path. -- `l2`, `ip`, and `cosine` are supported. The filter-hint reorder shortcut is - currently specialized for L2. +- Split storage is currently available on HGraph and Pyramid and requires fp32 + query codes. Pyramid enables the one-bit split search path by default for + split indexes; pass `rabitq_one_bit_search: false` under `pyramid` to force + the standard search path. +- `l2`, `ip`, and `cosine` are supported. For `x >= 2`, the canonical ordinary + split and fused HGraph paths directly reuse the exact filter inner product + for L2 and inner product. Other cases safely compute the full split distance. +- The fused datacell supports only L2 and inner product and only the in-memory + configuration described above. +- With `support_duplicate: true`, duplicate build probes and alias-expanding + queries use the canonical HGraph searcher; the fused slab remains the code + and graph storage. - Keep `use_reorder: true` unless x-bit traversal accuracy alone has been validated for the dataset. - Changing x, y, metric, or transform parameters requires rebuilding the - index. A search-time `hgraph.rabitq_error_rate` override does not. + index. A search-time `hgraph.rabitq_error_rate` or + `pyramid.rabitq_error_rate` override does not. - Use [RaBitQ](rabitq.md) for the general quantizer description and - [HGraph](../indexes/hgraph.md) and [Pyramid](../indexes/pyramid.md) for the complete index parameter tables. + [HGraph](../indexes/hgraph.md) and [Pyramid](../indexes/pyramid.md) for the + complete index parameter tables. diff --git a/docs/docs/zh/src/quantization/rabitq_split.md b/docs/docs/zh/src/quantization/rabitq_split.md index a0322a0dba..10f13f43bd 100644 --- a/docs/docs/zh/src/quantization/rabitq_split.md +++ b/docs/docs/zh/src/quantization/rabitq_split.md @@ -45,6 +45,7 @@ RaBitQ x+y split 是 HGraph 和 Pyramid 面向低比特底库码的存储与搜 | `rabitq_bits_per_dim_query` | split storage 必须使用 `32`。 | | `rabitq_error_rate` | lower-bound 误差项的默认正数倍率。 | | `use_reorder` | 建议设为 `true`,使用 `x+y` 距离排序候选。 | +| `rabitq_fused_datacell` | 仅用于 HGraph;启用融合布局,默认值为 `false`。 | 参数约束为: @@ -57,6 +58,25 @@ x + y <= 8 如果不配置 `rabitq_bits_per_dim_precise`,HGraph 和 Pyramid 使用 standard RaBitQ 路径, 不会创建 split storage。 +### HGraph 融合内存布局 + +仅对 HGraph,将 `rabitq_fused_datacell` 设为 `true` 后,底层节点的邻居、 +cluster id、label、x-bit code 和 y-bit supplement 会存入同一个 cache-line +对齐的 record。Pyramid 使用普通 split storage;`rabitq_fused_datacell` 不是 +Pyramid 参数。HGraph 专用搜索循环直接读取该 record,并联合预取图邻居和 +量化码。codec 使用固定随机种子可复现训练的 16 个 residual clusters。 + +融合布局是显式启用的,并且比普通 split storage 有更严格的约束: + +- `1 <= x <= 4`、`y >= 1` 且 `x + y <= 8`。 +- metric 必须是 L2 或内积。 +- graph、filter code 和 supplement code 必须全部使用内存 IO。 +- 必须关闭 MCI、`deduplicate_storage` 和 force remove。 +- fused v1 不支持 PCA;请省略 `rabitq_pca_dim` 或将其设为 `0`。 +- 不支持旧版 v0.14 序列化格式。 + +未启用该参数的索引保持原有布局、行为和序列化格式。 + 使用以下搜索参数启用 filter/lower-bound 搜索路径: ```json @@ -70,10 +90,24 @@ x + y <= 8 } ``` +Pyramid 需要把对应搜索参数放在 `pyramid` 下: + +```json +{ + "pyramid": { + "ef_search": 200, + "rabitq_one_bit_search": true, + "rabitq_error_rate": 1.9 + } +} +``` + 外部搜索参数仍命名为 `rabitq_one_bit_search`,但对 split 索引,它会使用 `rabitq_bits_per_dim_base` 配置的全部 `x` 个 filter bits。 -`hgraph.rabitq_error_rate` 可以为单次搜索覆盖索引默认值。record 中保存的是乘倍率前的 -几何误差尺度,因此 sweep 这个搜索参数不需要重建索引。 +`hgraph.rabitq_error_rate` 和 `pyramid.rabitq_error_rate` 可以分别为对应索引的 +单次搜索覆盖默认值,且不需要重建索引。原生 HGraph fused record 保存 +乘倍率前的几何误差尺度;HNSW-compatible fused `1+7` record 保留按规范 +默认值缩放的 metadata,并在查询时按相对该默认值的倍率应用 override。 ## 搜索流程 @@ -252,17 +286,25 @@ sum_i q_i * u_i + sum_i q_i * s_i ``` -对使用 x-bit lookup filter 的 L2 搜索,HGraph 和 Pyramid 会把之前计算的 filter distance -作为 hint 传给 reorder。`ComputeDistWithSplitCodeAndFilterDist` 从 hint 恢复第一项, -只从 y 个 supplement planes 计算第二项: +当 `x >= 2` 时,HGraph 和 Pyramid 的 canonical graph-search 路径会把遍历阶段 +算出的精确 x-bit filter inner product 直接传给 reorder。普通 split storage +通过 `QueryWithDistanceLowerBoundAndFilterIP` 输出该值,reorder 再通过 +`QueryWithFilterIPHint` 和 `ComputeDistWithSplitCodeAndFilterIP` 直接消费。 +HGraph fused 路径从 node record 读取 code,但使用相同的精确 hint 语义。 +这些路径无需从 distance 恢复 inner product,full rerank 只计算 y supplement +planes 对应的第二项: ```text full contribution = shifted filter contribution + supplement contribution ``` -因此 `3+5` 索引会复用 3-bit filter 结果,每个重排候选只扫描 5 个新的 bit-plane。 -如果 hint 不存在或不能使用,代码会回退到 `ComputeDistWithSplitCode`,直接从两个 -split records 计算相同的最终距离。 +因此普通和 fused `2+y`、`3+y`、`4+y` 都会复用精确的 x-bit filter inner +product,每个重排候选只扫描 y 个 supplement planes。 +`QueryWithDistanceHint` 和 `ComputeDistWithSplitCodeAndFilterDist` 仍作为兼容 API, +供只有 filter distance 的调用方使用,但它们不是 canonical graph-search 路径。 +fused `1+y` 的遍历使用 4-bit query bit-plane 与 popcount 近似值;精确重排会 +重新计算它的 1-bit 精确贡献,因为该近似值不能作为精确 full-distance hint。 +如果没有可用 hint,代码会直接从两个 split records 计算相同的最终距离。 ## 内存、磁盘和混合 IO @@ -324,29 +366,40 @@ split datacell 按以下顺序序列化: 创建目标索引时必须使用与序列化索引兼容的参数,尤其是 `dim`、`metric_type`、 x/y bit 数和 query bits。修改编码参数需要重建索引;只调整搜索参数 -`hgraph.rabitq_error_rate` 不需要。 +`hgraph.rabitq_error_rate` 或 `pyramid.rabitq_error_rate` 不需要。 + +对于 fused 索引,codec model 随 split datacell 序列化,每个节点的 code 只在 +bottom-graph slab 中序列化一次。普通和 streaming 往返都会保留该布局, +不会再生成一份随节点数增长的 split code 副本。 ## 实现位置 | 模块 | 文件 / 入口 | | --- | --- | -| 外部 x/y 参数映射 | `src/algorithm/hgraph/hgraph_param_mapping.cpp` | +| 外部 x/y 参数映射 | `hgraph_param_mapping.cpp`、`pyramid.cpp` | | split record 和 IO | `src/datacell/rabitq_split_datacell.h` | | plane 布局和 code 拆分 | `RaBitQuantizer::StoredPlaneIndex`、`SplitCode` | | filter 距离和 lower bound | `ComputeDistWithOneBitLowerBound` | | 直接计算 split distance | `ComputeDistWithSplitCode` | -| 使用 filter hint 的 reorder | `ComputeDistWithSplitCodeAndFilterDist` | +| 使用 filter hint 的 reorder | `ComputeDistWithSplitCodeAndFilterDist`、`ComputeDistWithSplitCodeAndFilterIP` | | SIMD dispatch | `src/simd/rabitq_simd.cpp` | | AVX2 / AVX512 lookup kernel | `src/simd/avx2.cpp`、`src/simd/avx512.cpp` | | 内存/磁盘/混合 IO 示例 | `examples/cpp/323_index_hgraph_rabitq_split.cpp` | ## 使用注意 -- split storage 当前可用于 HGraph 和 Pyramid,并且要求 fp32 query code。Pyramid 的 split 索引默认启用 one-bit split 搜索路径;如需强制使用普通搜索路径,可以在 `pyramid` 搜索参数下传 `rabitq_one_bit_search: false`。 -- 支持 `l2`、`ip` 和 `cosine`;利用 filter hint 的 reorder 快速路径当前针对 L2。 +- split storage 当前可用于 HGraph 和 Pyramid,并且要求 fp32 query code。 + Pyramid 的 split 索引默认启用 one-bit split 搜索路径;如需强制使用普通搜索路径, + 可以在 `pyramid` 搜索参数下传 `rabitq_one_bit_search: false`。 +- 支持 `l2`、`ip` 和 `cosine`。当 `x >= 2` 时,canonical 普通 split 路径和 + HGraph fused 路径会为 L2 和内积直接复用精确 filter inner product;其他情况 + 会安全地计算完整 split distance。 +- fused datacell 只支持 L2、内积以及上文所述的纯内存配置。 +- 启用 `support_duplicate: true` 时,重复向量 build probe 和展开 alias 的查询使用 + HGraph canonical searcher;fused slab 仍负责保存 code 和 graph。 - 除非已经验证仅靠 x-bit 遍历距离能满足召回要求,否则应保持 `use_reorder: true`。 - 修改 x、y、metric 或 transform 参数后必须重建索引;在搜索参数中覆盖 - `hgraph.rabitq_error_rate` 不需要重建。 + `hgraph.rabitq_error_rate` 或 `pyramid.rabitq_error_rate` 不需要重建。 - RaBitQ 通用说明见 [RaBitQ](rabitq.md),完整 HGraph 参数见 - [HGraph 索引](../indexes/hgraph.md)和 [Pyramid 索引](../indexes/pyramid.md)。 + [HGraph 索引](../indexes/hgraph.md) 和 [Pyramid 索引](../indexes/pyramid.md)。 diff --git a/include/vsag/constants.h b/include/vsag/constants.h index 126dc9a643..ff3999080b 100644 --- a/include/vsag/constants.h +++ b/include/vsag/constants.h @@ -237,6 +237,7 @@ extern const char* const HGRAPH_PRECISE_DIRECT_READ; extern const char* const HGRAPH_PARAMETER_EF_RUNTIME; extern const char* const HGRAPH_PARAMETER_HOPS_LIMIT; extern const char* const HGRAPH_PARAMETER_RABITQ_ONE_BIT_SEARCH; +extern const char* const HGRAPH_RABITQ_FUSED_DATACELL; extern const char* const HGRAPH_PARAMETER_BRUTE_FORCE_THRESHOLD; extern const char* const HGRAPH_USE_MCI; extern const char* const HGRAPH_MCI_MCS; diff --git a/src/algorithm/hgraph/hgraph.cpp b/src/algorithm/hgraph/hgraph.cpp index 3b0ffc222f..df35471a1d 100644 --- a/src/algorithm/hgraph/hgraph.cpp +++ b/src/algorithm/hgraph/hgraph.cpp @@ -27,6 +27,8 @@ #include "attr/argparse.h" #include "common.h" #include "datacell/flatten_interface.h" +#include "datacell/hgraph_rabitq_fused_datacell.h" +#include "datacell/rabitq_split_datacell.h" #include "datacell/sparse_graph_datacell.h" #include "dataset_impl.h" #include "impl/filter/filter_headers.h" @@ -36,6 +38,7 @@ #include "impl/pruning_strategy.h" #include "impl/reasoning/search_reasoning.h" #include "impl/reorder/flatten_reorder.h" +#include "impl/searcher/hgraph_rabitq_searcher.h" #include "index/index_impl.h" #include "io/reader_io/reader_io_parameter.h" #include "typing.h" @@ -88,13 +91,32 @@ HGraph::HGraph(const HGraphParameterPtr& hgraph_param, const vsag::IndexCommonPa FlattenInterface::MakeInstance(hgraph_param->precise_codes_param, common_param); } this->searcher_ = std::make_shared(common_param, neighbors_mutex_); + this->rabitq_fused_searcher_ = + std::make_shared(common_param, neighbors_mutex_); this->mci_searcher_ = std::make_shared(common_param); if (this->mci_parameters_.enabled) { this->mci_cliques_ = std::make_shared(common_param.allocator_.get()); } - this->bottom_graph_ = - GraphInterface::MakeInstance(hgraph_param->bottom_graph_param, common_param); + if (hgraph_param->rabitq_fused_datacell) { + auto split_codes = + std::dynamic_pointer_cast(basic_flatten_codes_); + CHECK_ARGUMENT(split_codes != nullptr, + "rabitq_fused_datacell requires in-memory RaBitQ split codes"); + auto graph_param = + std::dynamic_pointer_cast(hgraph_param->bottom_graph_param); + CHECK_ARGUMENT(graph_param != nullptr, "rabitq_fused_datacell requires flat graph storage"); + rabitq_fused_datacell_ = + std::make_shared(graph_param, + split_codes->OneBitCodeSize(), + split_codes->SupplementCodeSize(), + common_param); + split_codes->AttachFusedCodeStorage(rabitq_fused_datacell_.get()); + this->bottom_graph_ = rabitq_fused_datacell_; + } else { + this->bottom_graph_ = + GraphInterface::MakeInstance(hgraph_param->bottom_graph_param, common_param); + } if (this->support_duplicate_) { this->label_table_->SetDuplicateTracker(this->bottom_graph_->GetDuplicateTracker()); } @@ -312,15 +334,19 @@ HGraph::EstimateMemory(uint64_t num_elements) const { static_cast(block_size)); }; - if (this->basic_flatten_codes_->InMemory()) { - auto base_memory = this->basic_flatten_codes_->code_size_ * element_count; - estimate_memory += block_memory_ceil(base_memory, block_size); - } - - if (bottom_graph_->InMemory()) { - auto bottom_graph_memory = - (this->bottom_graph_->maximum_degree_ + 1) * sizeof(InnerIdType) * element_count; - estimate_memory += block_memory_ceil(bottom_graph_memory, block_size); + if (this->rabitq_fused_datacell_ != nullptr) { + const auto fused_memory = this->rabitq_fused_datacell_->RecordSize() * element_count; + estimate_memory += block_memory_ceil(fused_memory, block_size); + } else { + if (this->basic_flatten_codes_->InMemory()) { + auto base_memory = this->basic_flatten_codes_->code_size_ * element_count; + estimate_memory += block_memory_ceil(base_memory, block_size); + } + if (bottom_graph_->InMemory()) { + auto bottom_graph_memory = + (this->bottom_graph_->maximum_degree_ + 1) * sizeof(InnerIdType) * element_count; + estimate_memory += block_memory_ceil(bottom_graph_memory, block_size); + } } if (has_precise_reorder() && this->high_precise_codes_->InMemory() && @@ -377,7 +403,9 @@ HGraph::CalcDistanceById(const float* query, int64_t id, bool calculate_precise_ if (create_new_raw_vector_ && calculate_precise_distance) { flat = this->raw_vector_; } - if (lock.owns_lock() && not this->using_dedup_storage()) { + const bool reads_fused_codes = + this->rabitq_fused_datacell_ != nullptr and flat == this->basic_flatten_codes_; + if (lock.owns_lock() and not this->using_dedup_storage() and not reads_fused_codes) { lock.unlock(); } return InnerIndexInterface::calc_distance_by_id(query, id, flat); @@ -409,7 +437,9 @@ HGraph::CalDistanceById(const float* query, if (create_new_raw_vector_ && calculate_precise_distance) { flat = this->raw_vector_; } - if (lock.owns_lock() && not this->using_dedup_storage()) { + const bool reads_fused_codes = + this->rabitq_fused_datacell_ != nullptr and flat == this->basic_flatten_codes_; + if (lock.owns_lock() and not this->using_dedup_storage() and not reads_fused_codes) { lock.unlock(); } std::vector validity; @@ -443,6 +473,20 @@ InnerIndexPtr HGraph::ExportModel(const IndexCommonParam& param) const { auto index = std::make_shared(this->create_param_ptr_, param); this->basic_flatten_codes_->ExportModel(index->basic_flatten_codes_); + if (this->rabitq_fused_datacell_ != nullptr) { + auto source_split = + std::dynamic_pointer_cast(this->basic_flatten_codes_); + auto target_split = + std::dynamic_pointer_cast(index->basic_flatten_codes_); + CHECK_ARGUMENT(source_split != nullptr and target_split != nullptr and + index->rabitq_fused_datacell_ != nullptr, + "failed to export fused HGraph codec model"); + const auto fused_codec = source_split->ExportFusedCodec(); + if (not fused_codec.empty()) { + target_split->ImportFusedCodec(fused_codec); + index->rabitq_fused_datacell_->SetCodecModel(fused_codec); + } + } if (has_precise_reorder()) { this->high_precise_codes_->ExportModel(index->high_precise_codes_); } @@ -459,11 +503,15 @@ HGraph::GetCodeByInnerId(InnerIdType inner_id, uint8_t* data) const { return; } - if (has_precise_reorder()) { - high_precise_codes_->GetCodesById(inner_id, data); - } else { - basic_flatten_codes_->GetCodesById(inner_id, data); + if (this->has_precise_reorder()) { + this->high_precise_codes_->GetCodesById(inner_id, data); + return; + } + if (this->rabitq_fused_datacell_ != nullptr) { + throw VsagException(ErrorType::UNSUPPORTED_INDEX_OPERATION, + "fused RaBitQ codes do not expose a global merged code"); } + this->basic_flatten_codes_->GetCodesById(inner_id, data); } void @@ -548,6 +596,12 @@ HGraph::GetVectorByInnerId(InnerIdType inner_id, float* data) const { bool release; const auto* buffer = codes->GetCodesById(inner_id, release); if (buffer == nullptr) { + auto split_codes = + std::dynamic_pointer_cast(basic_flatten_codes_); + if (rabitq_fused_datacell_ != nullptr and codes.get() == basic_flatten_codes_.get() and + split_codes != nullptr and split_codes->DecodeFusedById(inner_id, data)) { + return; + } throw VsagException(ErrorType::INTERNAL_ERROR, fmt::format("failed to get vector by inner id {}", inner_id)); } @@ -567,6 +621,7 @@ HGraph::SetImmutable() { auto empty_mutex = std::make_shared(); this->searcher_->SetMutexArray(empty_mutex); this->parallel_searcher_->SetMutexArray(empty_mutex); + this->rabitq_fused_searcher_->SetMutexArray(empty_mutex); this->neighbors_mutex_ = empty_mutex; this->immutable_.store(true, std::memory_order_release); } @@ -646,7 +701,9 @@ void HGraph::init_resize_bit_and_reorder() { if (use_reorder_) { auto reorder_codes = this->get_reorder_codes(); - reorder_ = std::make_shared(reorder_codes, allocator_); + auto fused_graph = + reorder_codes.get() == basic_flatten_codes_.get() ? rabitq_fused_datacell_ : nullptr; + reorder_ = std::make_shared(reorder_codes, allocator_, fused_graph); } } @@ -714,17 +771,29 @@ HGraph::UpdateVector(int64_t id, const DatasetPtr& new_base, bool force_update) void* new_base_vec = nullptr; uint64_t data_size = 0; get_vectors(data_type_, dim_, new_base, &new_base_vec, &data_size); + if (this->rabitq_fused_datacell_ != nullptr) { + CHECK_ARGUMENT(new_base->GetDim() == dim_, + "updated vector dimension must match the index dimension"); + CHECK_ARGUMENT(new_base_vec != nullptr, "updated vector must not be null"); + this->validate_fused_vector_data(static_cast(new_base_vec), 1); + } if (not force_update) { std::shared_lock label_lock(this->label_lookup_mutex_); - // 1. check whether vectors are same - Vector base_data(data_size, allocator_); - GetVectorByInnerId(inner_id, (float*)base_data.data()); - float old_self_dist = this->CalcDistanceById((float*)base_data.data(), id); - float self_dist = this->CalcDistanceById((float*)new_base_vec, id); - if (std::abs(old_self_dist - self_dist) < 1e-3) { - return true; + float self_dist = 0.0F; + if (this->rabitq_fused_datacell_ == nullptr) { + // 1. check whether vectors are same + Vector base_data(data_size, allocator_); + GetVectorByInnerId(inner_id, reinterpret_cast(base_data.data())); + const float old_self_dist = + this->CalcDistanceById(reinterpret_cast(base_data.data()), id); + self_dist = this->CalcDistanceById(static_cast(new_base_vec), id); + if (std::abs(old_self_dist - self_dist) < 1e-3) { + return true; + } + } else { + self_dist = this->CalcDistanceById(static_cast(new_base_vec), id); } // 2. check whether the neighborhood relationship is same @@ -770,12 +839,31 @@ HGraph::UpdateVector(int64_t id, const DatasetPtr& new_base, bool force_update) } std::unique_lock codes_lock(this->persistent_codes_mutex_); bool update_status = basic_flatten_codes_->UpdateVector(new_base_vec, inner_id); + if (update_status and rabitq_fused_datacell_ != nullptr) { + this->sync_fused_node_codes(inner_id, new_base_vec); + } if (has_precise_reorder()) { update_status = update_status && high_precise_codes_->UpdateVector(new_base_vec, inner_id); } return update_status; } +bool +HGraph::UpdateId(int64_t old_id, int64_t new_id) { + if (old_id == new_id) { + return true; + } + if (rabitq_fused_datacell_ == nullptr) { + return InnerIndexInterface::UpdateId(old_id, new_id); + } + std::scoped_lock label_lock(this->label_lookup_mutex_); + auto [found, inner_id] = label_table_->TryGetIdByLabel(old_id, true); + CHECK_ARGUMENT(found, "old label does not exist"); + label_table_->UpdateLabel(old_id, new_id); + rabitq_fused_datacell_->SetLabel(inner_id, new_id); + return true; +} + std::string HGraph::AnalyzeIndexBySearch(const SearchRequest& request) { AnalyzerParam analyzer_param(allocator_); diff --git a/src/algorithm/hgraph/hgraph.h b/src/algorithm/hgraph/hgraph.h index f5c5b78b3b..824d2c0f08 100644 --- a/src/algorithm/hgraph/hgraph.h +++ b/src/algorithm/hgraph/hgraph.h @@ -56,6 +56,8 @@ namespace vsag { class FlattenOptimizedBuildInterface; +class HGraphRaBitQFusedDataCell; +class HGraphRaBitQSearcher; class HGraphOptimizedBuildSession; class IteratorFilterContext; @@ -250,6 +252,9 @@ class HGraph : public InnerIndexInterface { void UpdateAttribute(int64_t id, const AttributeSet& new_attrs) override; + bool + UpdateId(int64_t old_id, int64_t new_id) override; + void UpdateAttribute(int64_t id, const AttributeSet& new_attrs, @@ -337,6 +342,12 @@ class HGraph : public InnerIndexInterface { void insert_persistent_codes_to_slot(const void* data, CodeSlotIdType code_slot_id); + void + sync_fused_node_codes(InnerIdType inner_id, const void* data); + + void + restore_fused_codec(); + /// Ensure physical code storage can hold required_capacity physical slots. void ensure_physical_code_capacity(CodeSlotIdType required_capacity); @@ -363,7 +374,8 @@ class HGraph : public InnerIndexInterface { InnerSearchParam& inner_search_param, const VisitedListPtr& vt, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr, + bool* fused_search_finalized = nullptr) const; /// Overload that accepts an IteratorFilterContext for iterative search. template @@ -375,7 +387,7 @@ class HGraph : public InnerIndexInterface { IteratorFilterContext* iter_ctx, // ctx can be nullptr in adding scenario QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; private: [[nodiscard]] std::shared_lock @@ -468,6 +480,9 @@ class HGraph : public InnerIndexInterface { void validate_add_data(const DatasetPtr& data) const; + void + validate_fused_vector_data(const float* data, uint64_t count) const; + AddContext prepare_add_context(const DatasetPtr& data); @@ -634,7 +649,7 @@ class HGraph : public InnerIndexInterface { int64_t k, IteratorFilterContext* iter_ctx, QueryContext& ctx, - const DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + const RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; /// Run ELP (Edge-Link Pruning) optimizer on the bottom graph. void @@ -837,6 +852,8 @@ class HGraph : public InnerIndexInterface { Vector route_graphs_; // upper-layer route graphs GraphInterfacePtr bottom_graph_{nullptr}; // base-level graph (all vectors) + std::shared_ptr rabitq_fused_datacell_{nullptr}; + std::shared_ptr rabitq_fused_searcher_{nullptr}; SparseGraphDatacellParamPtr hierarchical_datacell_param_{nullptr}; // params for route graphs bool use_elp_optimizer_{false}; // enable ELP edge-link pruning diff --git a/src/algorithm/hgraph/hgraph_add_test.cpp b/src/algorithm/hgraph/hgraph_add_test.cpp index a165a709e5..2f7ce07ba4 100644 --- a/src/algorithm/hgraph/hgraph_add_test.cpp +++ b/src/algorithm/hgraph/hgraph_add_test.cpp @@ -162,6 +162,15 @@ const std::string kBruteForceSearchParams = } // namespace +TEST_CASE("HGraph UpdateId preserves equal-label no-op", "[ut][hgraph][update]") { + auto common_param = MakeCommonParam(8); + auto index = MakeHGraphIndex(MakeFp32HGraphJson(false), common_param); + + auto result = index->UpdateId(404, 404); + REQUIRE(result.has_value()); + REQUIRE(result.value()); +} + TEST_CASE("HGraph exact duplicate fallback supports every dense data type", "[ut][hgraph][duplicate][data_type]") { constexpr int64_t dim = 8; @@ -744,6 +753,95 @@ TEST_CASE("HGraph deduplicate_storage rejects v0.14 serialization", REQUIRE(binary.error().message.find("v0.14") != std::string::npos); } +TEST_CASE("HGraph fused RaBitQ rejects v0.14 serialization", + "[ut][hgraph][rabitq_split][fused][serialize]") { + constexpr int64_t dim = 64; + auto common_param = MakeCommonParam(dim); + common_param.use_old_serial_format_ = true; + auto hgraph_json = vsag::JsonType::Parse(R"({ + "base_quantization_type": "rabitq", + "precise_quantization_type": "rabitq", + "base_io_type": "memory_io", + "base_supplement_io_type": "memory_io", + "rabitq_bits_per_dim_base": 2, + "rabitq_bits_per_dim_precise": 6, + "graph_io_type": "memory_io", + "graph_storage_type": "flat", + "graph_type": "nsw", + "max_degree": 8, + "ef_construction": 32, + "use_reorder": true, + "reorder_source": "base", + "rabitq_fused_datacell": true + })"); + auto index = MakeHGraphIndex(hgraph_json, common_param); + + auto binary = index->Serialize(); + REQUIRE_FALSE(binary.has_value()); + REQUIRE(binary.error().type == vsag::ErrorType::INVALID_ARGUMENT); + REQUIRE(binary.error().message.find("v0.14") != std::string::npos); +} + +TEST_CASE("HGraph fused RaBitQ GetStats decodes vectors from node records", + "[ut][hgraph][rabitq_split][fused][stats]") { + constexpr int64_t dim = 64; + constexpr int64_t count = 32; + auto common_param = MakeCommonParam(dim); + auto hgraph_json = vsag::JsonType::Parse(R"({ + "base_quantization_type": "rabitq", + "precise_quantization_type": "rabitq", + "base_io_type": "memory_io", + "base_supplement_io_type": "memory_io", + "rabitq_bits_per_dim_base": 3, + "rabitq_bits_per_dim_precise": 5, + "rabitq_use_fht": true, + "graph_io_type": "memory_io", + "graph_storage_type": "flat", + "graph_type": "nsw", + "max_degree": 8, + "ef_construction": 32, + "build_thread_count": 1, + "use_reorder": true, + "reorder_source": "base", + "store_raw_vector": false, + "use_mci": false, + "support_duplicate": true, + "rabitq_fused_datacell": true + })"); + auto index = MakeHGraphIndex(hgraph_json, common_param); + + std::vector vectors(static_cast(count) * dim); + std::vector ids(count); + for (int64_t i = 0; i < count; ++i) { + ids[i] = 1000 + i; + for (int64_t d = 0; d < dim; ++d) { + vectors[static_cast(i) * dim + d] = + 0.1F + static_cast((i * 17 + d * 13) % 79) / 100.0F; + } + } + std::copy_n(vectors.data(), dim, vectors.data() + static_cast(count - 1) * dim); + auto base = MakeFloatDataset(vectors, ids, dim, count); + REQUIRE(index->Build(base).has_value()); + + std::string stats; + REQUIRE_NOTHROW(stats = index->GetStats()); + REQUIRE_FALSE(stats.empty()); + const auto parsed_stats = vsag::JsonType::Parse(stats); + REQUIRE(parsed_stats["duplicate_ratio"].GetFloat() > 0.0F); + + const std::vector fetch_ids = {ids.front(), ids.back()}; + auto fetched = index->GetInnerIndex()->GetVectorByIds( + fetch_ids.data(), static_cast(fetch_ids.size()), nullptr); + REQUIRE(fetched != nullptr); + REQUIRE(fetched->GetNumElements() == static_cast(fetch_ids.size())); + REQUIRE(fetched->GetDim() == dim); + const auto* fetched_vectors = fetched->GetFloat32Vectors(); + REQUIRE(fetched_vectors != nullptr); + for (uint64_t i = 0; i < static_cast(fetch_ids.size()) * dim; ++i) { + REQUIRE(std::isfinite(fetched_vectors[i])); + } +} + TEST_CASE("HGraph deduplicate_storage supports precise reorder code path", "[ut][hgraph][duplicate][reorder][add]") { constexpr int64_t dim = 4; diff --git a/src/algorithm/hgraph/hgraph_build.cpp b/src/algorithm/hgraph/hgraph_build.cpp index 1ef4744f4c..977158d851 100644 --- a/src/algorithm/hgraph/hgraph_build.cpp +++ b/src/algorithm/hgraph/hgraph_build.cpp @@ -24,6 +24,8 @@ #include #include "datacell/flatten_datacell_parameter.h" +#include "datacell/hgraph_rabitq_fused_datacell.h" +#include "datacell/rabitq_split_datacell.h" #include "dataset_impl.h" #include "hgraph.h" // IWYU pragma: keep #include "hgraph_fast_build.h" @@ -114,6 +116,16 @@ wait_all_futures(std::vector>& futures) { void HGraph::Train(const DatasetPtr& base) { + if (this->rabitq_fused_datacell_ != nullptr) { + // ODescent may reserve graph IDs before deferred persistent-code training. The split + // datacell count remains zero until those codes are inserted, which distinguishes that + // state (and an empty exported model) from a populated index. + CHECK_ARGUMENT(this->basic_flatten_codes_->TotalCount() == 0, + "cannot retrain a non-empty fused RaBitQ HGraph"); + if (not this->rabitq_fused_datacell_->CodecModel().empty()) { + return; + } + } this->train_codes_with_dataset(this->sample_train_dataset(base)); } @@ -121,6 +133,10 @@ DatasetPtr HGraph::sample_train_dataset(const DatasetPtr& base) const { int64_t total_elements = base->GetNumElements(); int64_t dim = base->GetDim(); + if (rabitq_fused_datacell_ != nullptr) { + return vsag::sample_train_data( + base, total_elements, dim, train_sample_count_, allocator_, 0x52425131U); + } return vsag::sample_train_data(base, total_elements, dim, train_sample_count_, allocator_); } @@ -128,6 +144,14 @@ void HGraph::train_codes_with_dataset(const DatasetPtr& train_data) { const auto* data_ptr = get_data(train_data); this->basic_flatten_codes_->Train(data_ptr, train_data->GetNumElements()); + if (rabitq_fused_datacell_ != nullptr) { + auto split_codes = + std::dynamic_pointer_cast(basic_flatten_codes_); + CHECK_ARGUMENT(split_codes != nullptr, "fused HGraph lost its RaBitQ split codes"); + split_codes->TrainFusedCodec( + static_cast(data_ptr), train_data->GetNumElements(), 16); + rabitq_fused_datacell_->SetCodecModel(split_codes->ExportFusedCodec()); + } if (has_precise_reorder()) { this->high_precise_codes_->Train(data_ptr, train_data->GetNumElements()); } @@ -140,6 +164,9 @@ HGraph::train_codes_with_dataset(const DatasetPtr& train_data) { std::vector HGraph::Build(const DatasetPtr& data) { CHECK_ARGUMENT(GetNumElements() == 0, "index is not empty"); + if (this->rabitq_fused_datacell_ != nullptr) { + this->validate_add_data(data); + } this->build_cache_hit_rate_ = -1.0F; this->build_cache_hit_nodes_ = 0; this->build_cache_missed_nodes_ = 0; @@ -331,6 +358,23 @@ HGraph::Add(const DatasetPtr& data) { return batch.failed_ids; } +void +HGraph::validate_fused_vector_data(const float* data, uint64_t count) const { + if (this->rabitq_fused_datacell_ == nullptr) { + return; + } + CHECK_ARGUMENT(data != nullptr, "fused RaBitQ base vectors must not be null"); + const auto dim = static_cast(this->dim_); + for (uint64_t row = 0; row < count; ++row) { + for (uint64_t d = 0; d < dim; ++d, ++data) { + if (not IsFiniteRaBitQValue(*data)) { + throw VsagException(ErrorType::INVALID_ARGUMENT, + "fused RaBitQ base vectors must contain only finite values"); + } + } + } +} + void HGraph::validate_add_data(const DatasetPtr& data) const { auto base_dim = data->GetDim(); @@ -338,7 +382,10 @@ HGraph::validate_add_data(const DatasetPtr& data) const { CHECK_ARGUMENT(base_dim == dim_, fmt::format("base.dim({}) must be equal to index.dim({})", base_dim, dim_)); } - CHECK_ARGUMENT(get_data(data) != nullptr, "base.float_vector is nullptr"); + const auto* base_data = get_data(data); + CHECK_ARGUMENT(base_data != nullptr, "base.float_vector is nullptr"); + this->validate_fused_vector_data(static_cast(base_data), + static_cast(data->GetNumElements())); } HGraph::AddContext @@ -359,8 +406,14 @@ HGraph::prepare_add_context(const DatasetPtr& data) { { std::scoped_lock lock(this->add_mutex_); if (this->total_count_ == 0) { - context.train_data = this->sample_train_dataset(data); - this->train_codes_with_dataset(context.train_data); + const bool reuse_fused_codec = this->rabitq_fused_datacell_ != nullptr and + not this->rabitq_fused_datacell_->CodecModel().empty(); + if (reuse_fused_codec) { + context.train_data = data; + } else { + context.train_data = this->sample_train_dataset(data); + this->train_codes_with_dataset(context.train_data); + } } } context.first_empty_add = context.train_data != nullptr; @@ -524,6 +577,7 @@ HGraph::insert_persistent_codes(const void* data, InnerIdType inner_id) { void HGraph::insert_persistent_codes_unlocked(const void* data, InnerIdType inner_id) { this->basic_flatten_codes_->InsertVector(data, inner_id); + this->sync_fused_node_codes(inner_id, data); if (has_precise_reorder()) { this->high_precise_codes_->InsertVector(data, inner_id); } @@ -532,6 +586,25 @@ HGraph::insert_persistent_codes_unlocked(const void* data, InnerIdType inner_id) } } +void +HGraph::sync_fused_node_codes(InnerIdType inner_id, const void* data) { + if (rabitq_fused_datacell_ == nullptr) { + return; + } + auto split_codes = + std::dynamic_pointer_cast(basic_flatten_codes_); + CHECK_ARGUMENT(split_codes != nullptr, "fused HGraph lost its RaBitQ split codes"); + const auto label = label_table_->GetLabelById(inner_id); + ByteBuffer one_bit(split_codes->OneBitCodeSize(), allocator_); + ByteBuffer supplement(split_codes->SupplementCodeSize(), allocator_); + uint32_t cluster_id = 0; + CHECK_ARGUMENT(split_codes->EncodeFused( + static_cast(data), one_bit.data, supplement.data, &cluster_id), + "failed to encode fused RaBitQ node codes"); + rabitq_fused_datacell_->SetNodeCodes( + inner_id, label, cluster_id, one_bit.data, supplement.data); +} + void HGraph::insert_persistent_codes_to_slot(const void* data, CodeSlotIdType code_slot_id) { InsertVectorToCodeSlot(this->basic_flatten_codes_, data, code_slot_id); @@ -723,7 +796,9 @@ HGraph::probe_graph_for_add(const void* data, if (this->support_duplicate_) { param.find_duplicate = true; param.duplicate_query_id = - this->using_dedup_storage() ? std::numeric_limits::max() : inner_id; + this->using_dedup_storage() or this->rabitq_fused_datacell_ != nullptr + ? std::numeric_limits::max() + : inner_id; param.duplicate_distance_threshold = this->duplicate_distance_threshold_; } @@ -982,6 +1057,20 @@ HGraph::InitFeatures() { this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_KNN_SEARCH_WITH_EX_FILTER); this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_UPDATE_EXTRA_INFO_CONCURRENT); } + + if (this->rabitq_fused_datacell_ != nullptr) { + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_MERGE_INDEX, false); + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_TUNE, false); + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_ADD_CONCURRENT, false); + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_ADD_SEARCH_CONCURRENT, false); + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_ADD_SEARCH_DELETE_CONCURRENT, + false); + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_UPDATE_VECTOR_CONCURRENT, + false); + if (this->raw_vector_ == nullptr) { + this->index_feature_list_->SetFeature(IndexFeature::SUPPORT_ADD_AFTER_BUILD, false); + } + } } void @@ -1005,14 +1094,16 @@ HGraph::reorder(const void* query, int64_t k, IteratorFilterContext* iter_ctx, QueryContext& ctx, - const DistanceRecordVector* rabitq_lower_bound_candidates) const { + const RaBitQCandidateVector* rabitq_lower_bound_candidates) const { uint64_t size = candidate_heap->Size(); if (k <= 0) { k = static_cast(size); } auto reorder_impl = reorder_; if (reorder_impl == nullptr) { - reorder_impl = std::make_shared(flatten, allocator_); + auto fused_graph = + flatten.get() == basic_flatten_codes_.get() ? rabitq_fused_datacell_ : nullptr; + reorder_impl = std::make_shared(flatten, allocator_, fused_graph); } auto reorder_heap = reorder_impl->Reorder(candidate_heap, static_cast(query), diff --git a/src/algorithm/hgraph/hgraph_fast_build.cpp b/src/algorithm/hgraph/hgraph_fast_build.cpp index 73b262ed3f..afa51b63f0 100644 --- a/src/algorithm/hgraph/hgraph_fast_build.cpp +++ b/src/algorithm/hgraph/hgraph_fast_build.cpp @@ -46,6 +46,9 @@ wait_all_futures(std::vector>& futures) { } // namespace HGraphOptimizedBuildSession::HGraphOptimizedBuildSession(HGraph& hgraph) : hgraph_(&hgraph) { + if (hgraph.rabitq_fused_datacell_ != nullptr) { + return; + } if (hgraph.using_dedup_storage()) { return; } diff --git a/src/algorithm/hgraph/hgraph_param_mapping.cpp b/src/algorithm/hgraph/hgraph_param_mapping.cpp index 7acbdb8e05..7aefc490a6 100644 --- a/src/algorithm/hgraph/hgraph_param_mapping.cpp +++ b/src/algorithm/hgraph/hgraph_param_mapping.cpp @@ -15,6 +15,7 @@ #include #include "common.h" +#include "datacell/graph_datacell_parameter.h" #include "hgraph.h" // IWYU pragma: keep #include "hgraph_parameter.h" #include "quantization/rabitq_quantization/rabitq_quantizer_parameter.h" @@ -131,6 +132,12 @@ HGraph::map_hgraph_param(const JsonType& hgraph_json) { HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY, }, }, + { + HGRAPH_RABITQ_FUSED_DATACELL, + { + HGRAPH_RABITQ_FUSED_DATACELL_KEY, + }, + }, { USE_ATTRIBUTE_FILTER, { @@ -637,6 +644,7 @@ HGraph::map_hgraph_param(const JsonType& hgraph_json) { "{HGRAPH_USE_ENV_OPTIMIZER}": false, "{HGRAPH_IGNORE_REORDER_KEY}": false, "{HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY}": false, + "{HGRAPH_RABITQ_FUSED_DATACELL_KEY}": false, "{RESIZE_INCREASE_COUNT_BIT}": {DEFAULT_RESIZE_INCREASE_COUNT_BIT}, "{HGRAPH_USE_ATTRIBUTE_FILTER_KEY}": false, "{GRAPH_KEY}": { @@ -758,6 +766,19 @@ HGraph::CheckAndMappingExternalParam(const JsonType& external_param, auto hgraph_parameter = std::make_shared(); hgraph_parameter->data_type = common_param.data_type_; hgraph_parameter->FromJson(inner_json); + if (hgraph_parameter->rabitq_fused_datacell) { + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + common_param.metric_ == MetricType::METRIC_TYPE_L2SQR or + common_param.metric_ == MetricType::METRIC_TYPE_IP, + "rabitq_fused_datacell only supports L2 and inner product"); + const auto graph_param = + std::dynamic_pointer_cast(hgraph_parameter->bottom_graph_param); + CHECK_ARGUMENT(graph_param != nullptr, "rabitq_fused_datacell requires flat graph storage"); + const auto graph_io_type = graph_param->io_parameter_->GetTypeName(); + CHECK_ARGUMENT(graph_io_type == IO_TYPE_VALUE_BLOCK_MEMORY_IO or + graph_io_type == IO_TYPE_VALUE_MEMORY_IO, + "rabitq_fused_datacell only supports an in-memory graph"); + } uint64_t max_degree = hgraph_parameter->bottom_graph_param->max_degree_; auto max_degree_threshold = std::max(common_param.dim_, 128); diff --git a/src/algorithm/hgraph/hgraph_parameter.cpp b/src/algorithm/hgraph/hgraph_parameter.cpp index db165b509a..57b8b8810d 100644 --- a/src/algorithm/hgraph/hgraph_parameter.cpp +++ b/src/algorithm/hgraph/hgraph_parameter.cpp @@ -25,6 +25,7 @@ #include "datacell/sparse_vector_datacell_parameter.h" #include "impl/odescent/odescent_graph_parameter.h" #include "inner_string_params.h" +#include "quantization/rabitq_quantization/rabitq_quantizer_parameter.h" #include "utils/param_compat_macros.h" #include "vsag/constants.h" @@ -70,6 +71,9 @@ HGraphParameter::FromJson(const JsonType& json) { if (json.Contains(HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY)) { this->build_by_base = json[HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY].GetBool(); } + if (json.Contains(HGRAPH_RABITQ_FUSED_DATACELL_KEY)) { + this->rabitq_fused_datacell = json[HGRAPH_RABITQ_FUSED_DATACELL_KEY].GetBool(); + } CHECK_ARGUMENT(json.Contains(BASE_CODES_KEY), fmt::format("hgraph parameters must contains {}", BASE_CODES_KEY)); @@ -226,6 +230,40 @@ HGraphParameter::FromJson(const JsonType& json) { CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) not(this->mci_parameters.enabled and this->support_force_remove), "hgraph mci does not support force remove"); + if (this->rabitq_fused_datacell) { + CHECK_ARGUMENT(not this->mci_parameters.enabled, + "rabitq_fused_datacell does not support MCI"); + CHECK_ARGUMENT(not this->deduplicate_storage, + "rabitq_fused_datacell does not support deduplicate_storage"); + CHECK_ARGUMENT(not this->support_force_remove, + "rabitq_fused_datacell does not support force remove"); + CHECK_ARGUMENT(this->base_codes_param->name == RABITQ_SPLIT_DATA_CELL, + "rabitq_fused_datacell requires RaBitQ split codes"); + const auto rabitq_param = std::dynamic_pointer_cast( + this->base_codes_param->quantizer_parameter); + CHECK_ARGUMENT(rabitq_param != nullptr, + "rabitq_fused_datacell requires RaBitQ quantization"); + CHECK_ARGUMENT(rabitq_param->pca_dim_ == 0, + "rabitq_fused_datacell v1 does not support PCA; " + "rabitq_pca_dim must be omitted or set to 0"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + rabitq_param->num_bits_per_dim_filter_ >= 1 and + rabitq_param->num_bits_per_dim_filter_ <= 4 and + rabitq_param->num_bits_per_dim_base_ > rabitq_param->num_bits_per_dim_filter_ and + rabitq_param->num_bits_per_dim_base_ <= 8, + "rabitq_fused_datacell requires split x+y with x in [1, 4], y >= 1, " + "and x+y <= 8"); + const auto io_type = this->base_codes_param->io_parameter->GetTypeName(); + CHECK_ARGUMENT( + io_type == IO_TYPE_VALUE_BLOCK_MEMORY_IO or io_type == IO_TYPE_VALUE_MEMORY_IO, + "rabitq_fused_datacell only supports memory IO"); + const auto supplement_io = this->base_codes_param->supplement_io_parameter; + const auto supplement_io_type = + supplement_io == nullptr ? io_type : supplement_io->GetTypeName(); + CHECK_ARGUMENT(supplement_io_type == IO_TYPE_VALUE_BLOCK_MEMORY_IO or + supplement_io_type == IO_TYPE_VALUE_MEMORY_IO, + "rabitq_fused_datacell only supports in-memory supplement IO"); + } } JsonType @@ -235,6 +273,7 @@ HGraphParameter::ToJson() const { json[HGRAPH_USE_ELP_OPTIMIZER_KEY].SetBool(this->use_elp_optimizer); json[HGRAPH_IGNORE_REORDER_KEY].SetBool(this->ignore_reorder); + json[HGRAPH_RABITQ_FUSED_DATACELL_KEY].SetBool(this->rabitq_fused_datacell); json[REORDER_SOURCE_KEY].SetString(this->reorder_source); json[BASE_CODES_KEY].SetJson(this->base_codes_param->ToJson()); json[GRAPH_KEY].SetJson(this->bottom_graph_param->ToJson()); @@ -293,6 +332,7 @@ HGraphParameter::CheckCompatibility(const ParamPtr& other) const { CHECK_FIELD_EQ(*this, *p, deduplicate_storage); CHECK_FIELD_EQ(*this, *p, duplicate_distance_threshold); CHECK_FIELD_EQ(*this, *p, support_force_remove); + CHECK_FIELD_EQ(*this, *p, rabitq_fused_datacell); CHECK_FIELD_EQ(*this, *p, mci_parameters.enabled); CHECK_FIELD_EQ(*this, *p, mci_parameters.mcs); CHECK_FIELD_EQ(*this, *p, mci_parameters.clique_max); diff --git a/src/algorithm/hgraph/hgraph_parameter.h b/src/algorithm/hgraph/hgraph_parameter.h index fab5e75d74..b99e75f94d 100644 --- a/src/algorithm/hgraph/hgraph_parameter.h +++ b/src/algorithm/hgraph/hgraph_parameter.h @@ -70,6 +70,7 @@ class HGraphParameter : public InnerIndexParameter { bool use_elp_optimizer{false}; bool ignore_reorder{false}; bool build_by_base{false}; + bool rabitq_fused_datacell{false}; uint64_t ef_construction{400}; uint64_t resize_increase_count_bit{DEFAULT_RESIZE_INCREASE_COUNT_BIT}; diff --git a/src/algorithm/hgraph/hgraph_parameter_test.cpp b/src/algorithm/hgraph/hgraph_parameter_test.cpp index b1f9748848..9841f74cff 100644 --- a/src/algorithm/hgraph/hgraph_parameter_test.cpp +++ b/src/algorithm/hgraph/hgraph_parameter_test.cpp @@ -718,6 +718,88 @@ TEST_CASE("HGraph maps RaBitQ x+y split params", "[ut][HGraphParameter]") { REQUIRE(typed_param->reorder_source == std::string("base")); } +TEST_CASE("HGraph maps and validates fused RaBitQ split datacell", "[ut][HGraphParameter]") { + auto make_param = []() { + return vsag::JsonType::Parse(R"({ + "base_quantization_type": "rabitq", + "precise_quantization_type": "rabitq", + "base_io_type": "memory_io", + "base_supplement_io_type": "memory_io", + "rabitq_bits_per_dim_base": 1, + "rabitq_bits_per_dim_precise": 7, + "graph_io_type": "memory_io", + "graph_storage_type": "flat", + "graph_type": "nsw", + "max_degree": 32, + "ef_construction": 200, + "use_reorder": true, + "reorder_source": "base", + "rabitq_fused_datacell": true + })"); + }; + + vsag::IndexCommonParam common_param; + common_param.dim_ = 960; + common_param.metric_ = vsag::MetricType::METRIC_TYPE_L2SQR; + common_param.data_type_ = vsag::DataTypes::DATA_TYPE_FLOAT; + + auto mapped = vsag::HGraph::CheckAndMappingExternalParam(make_param(), common_param); + auto typed_param = std::dynamic_pointer_cast(mapped); + REQUIRE(typed_param != nullptr); + REQUIRE(typed_param->rabitq_fused_datacell); + REQUIRE_FALSE(typed_param->mci_parameters.enabled); + REQUIRE(typed_param->base_codes_param->name == std::string(vsag::RABITQ_SPLIT_DATA_CELL)); + + SECTION("accept filter widths one through four") { + for (int32_t filter_bits = 1; filter_bits <= 4; ++filter_bits) { + auto param = make_param(); + param["rabitq_bits_per_dim_base"].SetInt(filter_bits); + param["rabitq_bits_per_dim_precise"].SetInt(8 - filter_bits); + CAPTURE(filter_bits); + REQUIRE_NOTHROW(vsag::HGraph::CheckAndMappingExternalParam(param, common_param)); + } + } + + SECTION("reject filter widths above four") { + auto param = make_param(); + param["rabitq_bits_per_dim_base"].SetInt(5); + param["rabitq_bits_per_dim_precise"].SetInt(3); + REQUIRE_THROWS(vsag::HGraph::CheckAndMappingExternalParam(param, common_param)); + } + + SECTION("reject MCI") { + auto param = make_param(); + param["use_mci"].SetBool(true); + REQUIRE_THROWS(vsag::HGraph::CheckAndMappingExternalParam(param, common_param)); + } + + SECTION("reject non-memory supplement") { + auto param = make_param(); + param["base_supplement_io_type"].SetString("mmap_io"); + REQUIRE_THROWS(vsag::HGraph::CheckAndMappingExternalParam(param, common_param)); + } + + SECTION("reject cosine") { + common_param.metric_ = vsag::MetricType::METRIC_TYPE_COSINE; + REQUIRE_THROWS(vsag::HGraph::CheckAndMappingExternalParam(make_param(), common_param)); + } + + SECTION("reject PCA with INVALID_ARGUMENT") { + auto param = make_param(); + param[vsag::RABITQ_PCA_DIM].SetInt(480); + bool rejected = false; + try { + static_cast(vsag::HGraph::CheckAndMappingExternalParam(param, common_param)); + } catch (const vsag::VsagException& exception) { + rejected = true; + REQUIRE(exception.error_.type == vsag::ErrorType::INVALID_ARGUMENT); + REQUIRE(std::string(exception.what()).find("does not support PCA") != + std::string::npos); + } + REQUIRE(rejected); + } +} + TEST_CASE("HGraph maps RaBitQ without y bits to standard RaBitQ", "[ut][HGraphParameter]") { auto param = vsag::JsonType::Parse(R"({ "base_quantization_type": "rabitq", diff --git a/src/algorithm/hgraph/hgraph_search.cpp b/src/algorithm/hgraph/hgraph_search.cpp index 7abd62ee8b..f1f8f2581a 100644 --- a/src/algorithm/hgraph/hgraph_search.cpp +++ b/src/algorithm/hgraph/hgraph_search.cpp @@ -17,11 +17,13 @@ #include #include "attr/argparse.h" +#include "datacell/rabitq_split_datacell.h" #include "dataset_impl.h" #include "hgraph.h" // IWYU pragma: keep #include "impl/filter/iterator_filter.h" #include "impl/heap/standard_heap.h" #include "impl/reasoning/search_reasoning.h" +#include "impl/searcher/hgraph_rabitq_searcher.h" #include "utils/util_functions.h" namespace vsag { @@ -78,6 +80,7 @@ HGraph::KnnSearch(const DatasetPtr& query, auto params = HGraphSearchParameters::FromJson(parameters); ctx.rabitq_error_rate = params.rabitq_error_rate; + ctx.enable_rabitq_reorder = params.enable_reorder; CHECK_ARGUMENT( // NOLINT params.ef_search >= 1, fmt::format("ef_search({}) must be at least 1", params.ef_search)); @@ -150,13 +153,15 @@ HGraph::KnnSearch(const DatasetPtr& query, search_param.ef = std::max(params.ef_search, k); search_param.is_inner_id_allowed = ft; search_param.topk = static_cast(search_param.ef); + search_param.rerank_topk = k; search_param.parallel_search_thread_count = params.parallel_search_thread_count; search_param.enable_reorder = params.enable_reorder; search_param.enable_rabitq_one_bit_search = params.rabitq_one_bit_search; + search_param.consider_duplicate = this->support_duplicate_; search_param.skip_ratio = params.skip_ratio; search_param.skip_strategy_type = params.skip_strategy_type; - DistanceRecordVector rabitq_lower_bound_candidates(ctx.alloc); + RaBitQCandidateVector rabitq_lower_bound_candidates(ctx.alloc); auto* rabitq_lower_bound_candidates_ptr = search_param.enable_rabitq_one_bit_search and use_reorder_ and search_param.enable_reorder and reorder_by_base_ @@ -228,7 +233,11 @@ HGraph::search_one_graph(const void* query, InnerSearchParam& inner_search_param, const VisitedListPtr& vt, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates, + bool* fused_search_finalized) const { + if (fused_search_finalized != nullptr) { + *fused_search_finalized = false; + } bool new_visited_list = vt == nullptr; VisitedListPtr visited_list; if (new_visited_list) { @@ -238,7 +247,25 @@ HGraph::search_one_graph(const void* query, visited_list->Reset(); } DistHeapPtr result = nullptr; - if (inner_search_param.parallel_search_thread_count > 1) { + if constexpr (mode == KNN_SEARCH) { + if (not this->support_duplicate_ and rabitq_fused_datacell_ != nullptr and + inner_search_param.distance_batch_func == nullptr and + graph.get() == rabitq_fused_datacell_.get() and + flatten.get() == basic_flatten_codes_.get() and + inner_search_param.parallel_search_thread_count <= 1 and + not inner_search_param.find_duplicate and not inner_search_param.consider_duplicate) { + result = rabitq_fused_searcher_->Search(rabitq_fused_datacell_, + flatten, + visited_list, + query, + inner_search_param, + ctx, + rabitq_lower_bound_candidates, + fused_search_finalized); + } + } + if (result == nullptr and inner_search_param.parallel_search_thread_count > 1 and + this->thread_pool_ != nullptr) { result = this->parallel_searcher_->Search(graph, flatten, visited_list, @@ -247,7 +274,7 @@ HGraph::search_one_graph(const void* query, this->label_table_, ctx, rabitq_lower_bound_candidates); - } else { + } else if (result == nullptr) { result = this->searcher_->Search(graph, flatten, visited_list, @@ -271,7 +298,7 @@ HGraph::search_one_graph(const void* query, InnerSearchParam& inner_search_param, IteratorFilterContext* iter_ctx, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { auto visited_list = this->pool_->TakeOne(); auto result = this->searcher_->Search(graph, flatten, @@ -402,6 +429,7 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { auto params = HGraphSearchParameters::FromJson(request.params_str_); ctx.rabitq_error_rate = params.rabitq_error_rate; + ctx.enable_rabitq_reorder = params.enable_reorder; if (use_custom_distance) { CHECK_ARGUMENT(params.parallel_search_thread_count == 1, @@ -508,10 +536,28 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { auto& vt = vt_guard.visited_list; const auto* raw_query = use_custom_distance ? nullptr : get_data(query); - for (auto i = static_cast(this->route_graphs_.size() - 1); i >= 0; --i) { - auto result = this->search_one_graph( - raw_query, this->route_graphs_[i], this->basic_flatten_codes_, search_param, vt, &ctx); - search_param.ep = result->Top().second; + auto* split_codes = dynamic_cast(basic_flatten_codes_.get()); + if (not use_custom_distance and rabitq_fused_datacell_ != nullptr and split_codes != nullptr) { + search_param.rabitq_fused_computer = split_codes->FactoryFusedComputer(raw_query); + for (auto i = static_cast(this->route_graphs_.size() - 1); i >= 0; --i) { + search_param.ep = + rabitq_fused_searcher_->Route(this->route_graphs_[i], + rabitq_fused_datacell_, + basic_flatten_codes_, + search_param.rabitq_fused_computer, + search_param.ep, + search_param.enable_rabitq_one_bit_search); + } + } else { + for (auto i = static_cast(this->route_graphs_.size() - 1); i >= 0; --i) { + auto result = this->search_one_graph(raw_query, + this->route_graphs_[i], + this->basic_flatten_codes_, + search_param, + vt, + &ctx); + search_param.ep = result->Top().second; + } } FilterPtr ft = this->create_search_filter(request.filter_, params.use_extra_info_filter); @@ -529,7 +575,7 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param.is_inner_id_allowed = ft; search_param.radius = request.radius_; search_param.search_mode = RANGE_SEARCH; - search_param.consider_duplicate = true; + search_param.consider_duplicate = this->support_duplicate_; search_param.range_search_limit_size = static_cast(request.limited_size_); search_param.parallel_search_thread_count = params.parallel_search_thread_count; search_param.enable_reorder = use_custom_distance ? false : params.enable_reorder; @@ -539,13 +585,14 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param.ef = std::max(params.ef_search, k); search_param.is_inner_id_allowed = ft; search_param.topk = static_cast(search_param.ef); + search_param.rerank_topk = k; if (params.topk_factor > 1.0F) { search_param.topk = std::min(search_param.topk, static_cast(static_cast(k) * params.topk_factor)); } search_param.enable_reorder = use_custom_distance ? false : params.enable_reorder; - search_param.consider_duplicate = true; + search_param.consider_duplicate = this->support_duplicate_; search_param.enable_rabitq_one_bit_search = use_custom_distance ? false : params.rabitq_one_bit_search; if (params.enable_time_record) { @@ -571,15 +618,28 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param.skip_ratio = params.skip_ratio; search_param.skip_strategy_type = params.skip_strategy_type; - DistanceRecordVector rabitq_lower_bound_candidates(ctx.alloc); + const bool can_use_fused_direct_search = + not use_custom_distance and not is_range and not this->support_duplicate_ and + rabitq_fused_datacell_ != nullptr and + bottom_graph_.get() == rabitq_fused_datacell_.get() and basic_flatten_codes_ != nullptr and + search_param.parallel_search_thread_count <= 1 and not search_param.find_duplicate and + not search_param.consider_duplicate; + const bool fused_search_can_finalize = + can_use_fused_direct_search and + (not use_reorder_ or not search_param.enable_reorder or reorder_by_base_); + const bool fused_search_needs_candidates = + fused_search_can_finalize and HGraphRaBitQSearcher::ShouldDeferRerank(search_param); + RaBitQCandidateVector rabitq_lower_bound_candidates(ctx.alloc); auto* rabitq_lower_bound_candidates_ptr = - search_param.enable_rabitq_one_bit_search and use_reorder_ and + (not fused_search_can_finalize or fused_search_needs_candidates) and + search_param.enable_rabitq_one_bit_search and use_reorder_ and search_param.enable_reorder and reorder_by_base_ ? &rabitq_lower_bound_candidates : nullptr; DistHeapPtr search_result; bool brute_force_used = false; + bool fused_search_finalized = false; MCIHybridSearchResult mci_result(params, ft); if (not use_custom_distance) { if (params.brute_force_threshold > 0.0F and @@ -604,7 +664,8 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param, vt, &ctx, - rabitq_lower_bound_candidates_ptr); + rabitq_lower_bound_candidates_ptr, + &fused_search_finalized); } } } else { @@ -614,13 +675,18 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param, vt, &ctx, - rabitq_lower_bound_candidates_ptr); + rabitq_lower_bound_candidates_ptr, + &fused_search_finalized); } vt_guard.Release(); + const bool fused_search_already_reranked = + fused_search_finalized and + (not use_reorder_ or not search_param.enable_reorder or reorder_by_base_); + // Reorder - if (mci_result.route != "mci" and not brute_force_used and use_reorder_ and - search_param.enable_reorder) { + if (not fused_search_already_reranked and mci_result.route != "mci" and not brute_force_used and + use_reorder_ and search_param.enable_reorder) { auto limit = is_range ? request.limited_size_ : k; this->reorder(raw_query, this->get_reorder_codes(), @@ -629,8 +695,9 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { nullptr, ctx, rabitq_lower_bound_candidates_ptr); - } else if (mci_result.route != "mci" and not brute_force_used and - search_param.enable_reorder and params.rabitq_one_bit_search) { + } else if (not fused_search_already_reranked and mci_result.route != "mci" and + not brute_force_used and search_param.enable_reorder and + params.rabitq_one_bit_search) { auto limit = is_range ? request.limited_size_ : k; this->reorder(raw_query, this->basic_flatten_codes_, search_result, limit, nullptr, ctx); } diff --git a/src/algorithm/hgraph/hgraph_serialize.cpp b/src/algorithm/hgraph/hgraph_serialize.cpp index d8faf5feee..695521cd21 100644 --- a/src/algorithm/hgraph/hgraph_serialize.cpp +++ b/src/algorithm/hgraph/hgraph_serialize.cpp @@ -19,6 +19,7 @@ #include #include "common.h" +#include "datacell/rabitq_split_datacell.h" #include "datacell/sparse_graph_datacell.h" #include "hgraph.h" // IWYU pragma: keep #include "impl/heap/standard_heap.h" @@ -376,6 +377,11 @@ HGraph::Serialize(StreamWriter& writer) const { "HGraph duplicate code slot mapping does not support v0.14 " "serialization"); } + if (this->rabitq_fused_datacell_ != nullptr) { + throw VsagException(ErrorType::INVALID_ARGUMENT, + "HGraph RaBitQ fused datacell does not support v0.14 " + "serialization"); + } this->serialize_basic_info_v0_14(writer); this->basic_flatten_codes_->Serialize(writer); this->bottom_graph_->Serialize(writer); @@ -897,6 +903,7 @@ HGraph::read_streaming_body(StreamReader& reader, if (this->raw_vector_ != nullptr) { this->has_raw_vector_ = true; } + this->restore_fused_codec(); this->cal_memory_usage(); if (use_elp_optimizer_) { @@ -1069,6 +1076,7 @@ HGraph::Deserialize(StreamReader& reader) { (void)this->code_slot_map_->Resolve(inner_id); } } + this->restore_fused_codec(); this->cal_memory_usage(); // post serialize procedure @@ -1077,6 +1085,33 @@ HGraph::Deserialize(StreamReader& reader) { } } +void +HGraph::restore_fused_codec() { + if (rabitq_fused_datacell_ == nullptr) { + return; + } + auto split_codes = + std::dynamic_pointer_cast(basic_flatten_codes_); + CHECK_ARGUMENT(split_codes != nullptr, "fused HGraph lost its RaBitQ split codes"); + CHECK_ARGUMENT(split_codes->UsesExternalFusedCodeStorage(), + "fused HGraph split codes are not bound to the node slab"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + split_codes->OneBitCodeSize() == rabitq_fused_datacell_->OneBitCodeSize() and + split_codes->SupplementCodeSize() == rabitq_fused_datacell_->FusedSupplementCodeSize(), + "fused HGraph split-code sizes do not match the serialized node layout"); + const auto base_count = basic_flatten_codes_->TotalCount(); + const auto graph_count = rabitq_fused_datacell_->TotalCount(); + CHECK_ARGUMENT(graph_count <= base_count, + "fused HGraph node count exceeds its base-code count"); + CHECK_ARGUMENT(rabitq_fused_datacell_->MaxCapacity() >= base_count, + "fused HGraph node capacity is smaller than its base-code count"); + if (rabitq_fused_datacell_->CodecModel().empty()) { + CHECK_ARGUMENT(base_count == 0, "non-empty fused HGraph is missing its codec model"); + return; + } + split_codes->ImportFusedCodec(rabitq_fused_datacell_->CodecModel()); +} + std::unordered_map HGraph::GetMemoryUsageDetail() const { std::unordered_map memory_usage; diff --git a/src/algorithm/pyramid/pyramid.cpp b/src/algorithm/pyramid/pyramid.cpp index f6bf6d064e..768ac84c0e 100644 --- a/src/algorithm/pyramid/pyramid.cpp +++ b/src/algorithm/pyramid/pyramid.cpp @@ -395,10 +395,10 @@ Pyramid::KnnSearch(const DatasetPtr& query, const bool collect_rabitq_lower_bounds = search_param.enable_rabitq_one_bit_search and use_reorder_ and base_codes_->SupportSplitCodeStorage(); - DistanceRecordVector rabitq_lower_bound_candidates(allocator_); + RaBitQCandidateVector rabitq_lower_bound_candidates(allocator_); std::mutex rabitq_lower_bound_mutex; SearchFunc search_func = [&](const IndexNode* node, const VisitedListPtr& vl) { - DistanceRecordVector local_candidates(allocator_); + RaBitQCandidateVector local_candidates(allocator_); auto* candidates = collect_rabitq_lower_bounds ? &local_candidates : nullptr; auto result = this->search_node(node, vl, @@ -467,10 +467,10 @@ Pyramid::RangeSearch(const DatasetPtr& query, const bool collect_rabitq_lower_bounds = search_param.enable_rabitq_one_bit_search and use_reorder_ and base_codes_->SupportSplitCodeStorage(); - DistanceRecordVector rabitq_lower_bound_candidates(allocator_); + RaBitQCandidateVector rabitq_lower_bound_candidates(allocator_); std::mutex rabitq_lower_bound_mutex; SearchFunc search_func = [&](const IndexNode* node, const VisitedListPtr& vl) { - DistanceRecordVector local_candidates(allocator_); + RaBitQCandidateVector local_candidates(allocator_); auto* candidates = collect_rabitq_lower_bounds ? &local_candidates : nullptr; auto result = this->search_node(node, vl, @@ -507,7 +507,7 @@ Pyramid::search_impl(const DatasetPtr& query, InnerSearchParam& search_param, QueryContext& ctx, const std::string& hierarchy_name, - const DistanceRecordVector* rabitq_lower_bound_candidates) const { + const RaBitQCandidateVector* rabitq_lower_bound_candidates) const { auto h_iter = hierarchies_.find(hierarchy_name); CHECK_ARGUMENT(h_iter != hierarchies_.end(), fmt::format("unknown hierarchy name: '{}'", hierarchy_name)); @@ -1497,7 +1497,7 @@ Pyramid::search_node(const IndexNode* node, const FlattenInterfacePtr& codes, QueryContext& ctx, uint64_t subindex_ef_search, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { std::shared_lock lock(node->mutex_); DistHeapPtr results = nullptr; diff --git a/src/algorithm/pyramid/pyramid.h b/src/algorithm/pyramid/pyramid.h index c8b9289432..8f1afc74f6 100644 --- a/src/algorithm/pyramid/pyramid.h +++ b/src/algorithm/pyramid/pyramid.h @@ -336,7 +336,7 @@ class Pyramid : public InnerIndexInterface { InnerSearchParam& search_param, QueryContext& ctx, const std::string& hierarchy_name, - const DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + const RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; /// Probabilistic check: should total_count trigger a new entry-point update? bool @@ -367,7 +367,7 @@ class Pyramid : public InnerIndexInterface { const FlattenInterfacePtr& codes, QueryContext& ctx, uint64_t subindex_ef_search, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; [[nodiscard]] bool has_precise_reorder() const { diff --git a/src/algorithm/pyramid/pyramid_test.cpp b/src/algorithm/pyramid/pyramid_test.cpp index d325e884ef..dc7adc2523 100644 --- a/src/algorithm/pyramid/pyramid_test.cpp +++ b/src/algorithm/pyramid/pyramid_test.cpp @@ -207,6 +207,7 @@ TEST_CASE("Pyramid promotes flat node at index minimum size", "[ut][pyramid]") { REQUIRE(GetPyramidSubindexCount(index, "graph_subindexes") == 1); REQUIRE(GetPyramidSubindexCount(index, "total_vectors_in_graph") == 3); + bool observed_filter_ip_hint = false; for (int64_t i = 0; i < 3; ++i) { auto query = MakePyramidDataset(vectors.data() + i * PYRAMID_TEST_DIM, nullptr, paths.data() + i, 1); @@ -215,11 +216,18 @@ TEST_CASE("Pyramid promotes flat node at index minimum size", "[ut][pyramid]") { REQUIRE(result->GetDim() == 1); REQUIRE(result->GetIds()[0] == ids[i]); if (split_rabitq) { - auto stats = result->GetStatistics({"reorder_lower_bound_probe_count"}); - REQUIRE(stats.size() == 1); - REQUIRE(std::stoul(stats[0]) > 0); + auto stats = result->GetStatistics({"reorder_lower_bound_probe_count", + "rabitq_filter_count", + "rabitq_reorder_hint_full_count"}); + REQUIRE(stats.size() == 3); + REQUIRE(std::stoul(stats[0]) == 0); + REQUIRE(std::stoul(stats[1]) > 0); + observed_filter_ip_hint |= std::stoul(stats[2]) > 0; } } + if (split_rabitq) { + REQUIRE(observed_filter_ip_hint); + } } TEST_CASE("Pyramid Build stores RaBitQ and SQ8 codes in parallel", "[ut][pyramid]") { diff --git a/src/constants.cpp b/src/constants.cpp index 76b59f85da..c84376030b 100644 --- a/src/constants.cpp +++ b/src/constants.cpp @@ -202,6 +202,7 @@ const char* const HGRAPH_PRECISE_DIRECT_READ = "precise_direct_read"; const char* const HGRAPH_PARAMETER_EF_RUNTIME = "ef_search"; const char* const HGRAPH_PARAMETER_HOPS_LIMIT = "hops_limit"; const char* const HGRAPH_PARAMETER_RABITQ_ONE_BIT_SEARCH = "rabitq_one_bit_search"; +const char* const HGRAPH_RABITQ_FUSED_DATACELL = "rabitq_fused_datacell"; const char* const HGRAPH_PARAMETER_BRUTE_FORCE_THRESHOLD = "brute_force_threshold"; const char* const HGRAPH_USE_MCI = "use_mci"; const char* const HGRAPH_MCI_MCS = "mci_mcs"; diff --git a/src/datacell/flatten_datacell_test.cpp b/src/datacell/flatten_datacell_test.cpp index dbbd9fd81f..a733e2f33f 100644 --- a/src/datacell/flatten_datacell_test.cpp +++ b/src/datacell/flatten_datacell_test.cpp @@ -21,6 +21,7 @@ #include #include #include +#include #include #include #include @@ -32,17 +33,27 @@ #include "flatten_interface_test.h" #include "flatten_optimized_build_interface.h" +#include "hgraph_rabitq_fused_datacell.h" #include "impl/allocator/default_allocator.h" #include "impl/allocator/safe_allocator.h" #include "impl/thread_pool/safe_thread_pool.h" #include "index_common_param.h" +#include "io/memory_io/memory_io_parameter.h" #include "quantization/rabitq_quantization/rabitq_quantizer.h" +#include "rabitq_split_datacell.h" #include "unittest.h" using namespace vsag; namespace { +bool +IsNaNBitPattern(float value) { + uint32_t bits = 0; + std::memcpy(&bits, &value, sizeof(bits)); + return (bits & 0x7FFFFFFFU) > 0x7F800000U; +} + class RejectSecondThreadPool final : public ThreadPool { public: ~RejectSecondThreadPool() override { @@ -214,6 +225,719 @@ TEST_CASE("RaBitQSplitDataCell direct split compute", "[ut][RaBitQSplitDataCell] } } } + +TEST_CASE("RaBitQSplitDataCell fused residual clusters", "[ut][RaBitQSplitDataCell]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + constexpr uint64_t dim = 64; + constexpr uint64_t cluster_count = 16; + constexpr uint64_t vectors_per_cluster = 20; + constexpr uint64_t count = cluster_count * vectors_per_cluster; + Vector vectors(count * dim, 0.0F, allocator.get()); + for (uint64_t cluster = 0; cluster < cluster_count; ++cluster) { + for (uint64_t row = 0; row < vectors_per_cluster; ++row) { + auto* vector = vectors.data() + (cluster * vectors_per_cluster + row) * dim; + vector[cluster] = 100.0F; + vector[(cluster + 17) % dim] = static_cast(row) * 0.001F; + } + } + + auto param_json = JsonType::Parse(R"({ + "codes_type": "rabitq_split", + "io_params": {"type": "memory_io"}, + "quantization_params": { + "type": "rabitq", + "rabitq_version": "split", + "rabitq_bits_per_dim_query": 32, + "rabitq_bits_per_dim_base": 8, + "rabitq_bits_per_dim_filter": 1, + "use_fht": true + } + })"); + auto param = std::make_shared(); + param->FromJson(param_json); + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = dim; + common_param.metric_ = MetricType::METRIC_TYPE_L2SQR; + + auto flatten = FlattenInterface::MakeInstance(param, common_param); + flatten->Train(vectors.data(), count); + auto split = std::dynamic_pointer_cast(flatten); + REQUIRE(split != nullptr); + REQUIRE(split->FusedFilterBits() == 1); + REQUIRE(split->FusedSupplementBits() == 7); + REQUIRE(split->UsesLegacyHnswFusedCodec()); + split->TrainFusedCodec(vectors.data(), count, cluster_count); + const auto codec_model = split->ExportFusedCodec(); + REQUIRE_FALSE(codec_model.empty()); + + auto invalid_centroid_count = codec_model; + const uint64_t oversized_count = std::numeric_limits::max(); + std::memcpy(invalid_centroid_count.data() + sizeof(uint32_t), + &oversized_count, + sizeof(oversized_count)); + REQUIRE_THROWS(split->ImportFusedCodec(invalid_centroid_count)); + + auto trailing_codec_model = codec_model; + trailing_codec_model.push_back('\0'); + REQUIRE_THROWS(split->ImportFusedCodec(trailing_codec_model)); + + Vector one_bit(split->OneBitCodeSize(), allocator.get()); + Vector supplement(split->SupplementCodeSize(), allocator.get()); + const float invalid_values[] = {std::numeric_limits::quiet_NaN(), + std::numeric_limits::infinity(), + -std::numeric_limits::infinity()}; + constexpr uint8_t output_sentinel = 0xA5; + for (const float invalid_value : invalid_values) { + CAPTURE(invalid_value); + std::vector invalid_training(vectors.data(), vectors.data() + vectors.size()); + invalid_training[0] = invalid_value; + REQUIRE_THROWS(split->TrainFusedCodec(invalid_training.data(), count, cluster_count)); + REQUIRE(split->ExportFusedCodec() == codec_model); + + std::vector invalid_vector(vectors.data(), vectors.data() + dim); + invalid_vector[0] = invalid_value; + std::fill(one_bit.begin(), one_bit.end(), output_sentinel); + std::fill(supplement.begin(), supplement.end(), output_sentinel); + uint32_t invalid_cluster_id = cluster_count; + REQUIRE_FALSE(split->EncodeFused( + invalid_vector.data(), one_bit.data(), supplement.data(), &invalid_cluster_id)); + REQUIRE(invalid_cluster_id == cluster_count); + REQUIRE(std::all_of(one_bit.begin(), one_bit.end(), [](uint8_t value) { + return value == output_sentinel; + })); + REQUIRE(std::all_of(supplement.begin(), supplement.end(), [](uint8_t value) { + return value == output_sentinel; + })); + } + + UnorderedSet assigned_clusters(allocator.get()); + for (uint64_t cluster = 0; cluster < cluster_count; ++cluster) { + uint32_t cluster_id = 0; + REQUIRE(split->EncodeFused(vectors.data() + cluster * vectors_per_cluster * dim, + one_bit.data(), + supplement.data(), + &cluster_id)); + assigned_clusters.insert(cluster_id); + } + REQUIRE(assigned_clusters.size() == cluster_count); + + auto computer = split->FactoryFusedComputer(vectors.data()); + uint32_t cluster_id = 0; + REQUIRE(split->EncodeFused(vectors.data(), one_bit.data(), supplement.data(), &cluster_id)); + float coarse_distance = 0.0F; + float lower_bound = 0.0F; + float filter_inner_product = 0.0F; + REQUIRE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + one_bit.data(), + supplement.data(), + &coarse_distance, + &lower_bound, + &filter_inner_product, + nullptr)); + REQUIRE(IsNaNBitPattern(filter_inner_product)); + float full_distance = 0.0F; + REQUIRE_FALSE(split->ComputeFusedFullWithFilterIP( + computer, cluster_id, one_bit.data(), supplement.data(), 0.0F, &full_distance, nullptr)); + REQUIRE(split->ComputeFusedFull( + computer, cluster_id, one_bit.data(), supplement.data(), &full_distance, nullptr)); + REQUIRE(std::isfinite(coarse_distance)); + REQUIRE(std::isfinite(lower_bound)); + REQUIRE(std::isfinite(full_distance)); + + QueryContext narrow_context; + narrow_context.rabitq_error_rate = 0.95F; + QueryContext wide_context; + wide_context.rabitq_error_rate = 3.8F; + float narrow_distance = 0.0F; + float narrow_lower_bound = 0.0F; + float wide_distance = 0.0F; + float wide_lower_bound = 0.0F; + REQUIRE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + one_bit.data(), + supplement.data(), + &narrow_distance, + &narrow_lower_bound, + nullptr, + &narrow_context)); + REQUIRE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + one_bit.data(), + supplement.data(), + &wide_distance, + &wide_lower_bound, + nullptr, + &wide_context)); + REQUIRE(std::abs(coarse_distance - narrow_distance) <= 1e-6F); + REQUIRE(std::abs(coarse_distance - wide_distance) <= 1e-6F); + const float default_gap = coarse_distance - lower_bound; + const float narrow_gap = narrow_distance - narrow_lower_bound; + const float wide_gap = wide_distance - wide_lower_bound; + REQUIRE(default_gap > 1e-6F); + REQUIRE(std::abs(narrow_gap - 0.5F * default_gap) <= 1e-4F * default_gap + 1e-6F); + REQUIRE(std::abs(wide_gap - 2.0F * default_gap) <= 1e-4F * default_gap + 1e-6F); +} + +TEST_CASE("RaBitQSplitDataCell native fused bit splits", "[ut][RaBitQSplitDataCell]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + constexpr uint32_t cluster_count = 16; + constexpr InnerIdType count = 64; + struct FusedSplitCase { + uint32_t filter_bits; + uint32_t supplement_bits; + uint64_t dim; + }; + const FusedSplitCase cases[] = { + {1, 3, 65}, + {2, 6, 960}, + {3, 5, 65}, + {4, 4, 960}, + }; + + constexpr const char* param_template = R"( + {{ + "codes_type": "rabitq_split", + "io_params": {{ + "type": "memory_io" + }}, + "quantization_params": {{ + "type": "rabitq", + "rabitq_version": "split", + "rabitq_bits_per_dim_query": 32, + "rabitq_bits_per_dim_base": {}, + "rabitq_bits_per_dim_filter": {}, + "use_fht": true + }} + }} + )"; + + for (const auto& split_case : cases) { + CAPTURE(split_case.filter_bits, split_case.supplement_bits, split_case.dim); + const uint32_t base_bits = split_case.filter_bits + split_case.supplement_bits; + auto param_json = + JsonType::Parse(fmt::format(param_template, base_bits, split_case.filter_bits)); + auto param = std::make_shared(); + param->FromJson(param_json); + + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = split_case.dim; + common_param.metric_ = MetricType::METRIC_TYPE_L2SQR; + + auto vectors = fixtures::generate_vectors( + count, split_case.dim, false, 31 + static_cast(split_case.filter_bits)); + auto query = fixtures::generate_vectors( + 1, split_case.dim, false, 71 + static_cast(split_case.filter_bits)); + auto encoded_vectors = fixtures::generate_vectors( + 3, split_case.dim, false, 51 + static_cast(split_case.filter_bits)); + auto flatten = FlattenInterface::MakeInstance(param, common_param); + flatten->Train(vectors.data(), count); + auto split = std::dynamic_pointer_cast(flatten); + REQUIRE(split != nullptr); + REQUIRE(split->FusedFilterBits() == split_case.filter_bits); + REQUIRE(split->FusedSupplementBits() == split_case.supplement_bits); + REQUIRE_FALSE(split->UsesLegacyHnswFusedCodec()); + + split->TrainFusedCodec(vectors.data(), count, cluster_count); + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->max_degree_ = 8; + graph_param->init_max_capacity_ = 4; + auto graph = std::make_shared( + graph_param, split->OneBitCodeSize(), split->SupplementCodeSize(), common_param); + split->AttachFusedCodeStorage(graph.get()); + Vector filter_code(split->OneBitCodeSize(), allocator.get()); + Vector supplement_code(split->SupplementCodeSize(), allocator.get()); + auto computer = split->FactoryFusedComputer(query.data()); + REQUIRE(computer != nullptr); + + for (InnerIdType id = 0; id < 3; ++id) { + const auto* encoded_vector = + encoded_vectors.data() + static_cast(id) * split_case.dim; + flatten->InsertVector(encoded_vector, id); + uint32_t cluster_id = cluster_count; + REQUIRE(split->EncodeFused( + encoded_vector, filter_code.data(), supplement_code.data(), &cluster_id)); + REQUIRE(cluster_id < cluster_count); + graph->SetNodeCodes(id, + static_cast(id), + cluster_id, + filter_code.data(), + supplement_code.data()); + + float coarse_distance = 0.0F; + float lower_bound = 0.0F; + float filter_inner_product = 0.0F; + REQUIRE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + filter_code.data(), + supplement_code.data(), + &coarse_distance, + &lower_bound, + &filter_inner_product, + nullptr)); + REQUIRE(std::isfinite(coarse_distance)); + REQUIRE(std::isfinite(lower_bound)); + if (split_case.filter_bits == 1) { + REQUIRE(IsNaNBitPattern(filter_inner_product)); + } else { + REQUIRE(std::isfinite(filter_inner_product)); + } + + float direct_full_distance = 0.0F; + REQUIRE(split->ComputeFusedFull(computer, + cluster_id, + filter_code.data(), + supplement_code.data(), + &direct_full_distance, + nullptr)); + REQUIRE(std::isfinite(direct_full_distance)); + + if (split_case.filter_bits >= 2) { + float hinted_full_distance = 0.0F; + REQUIRE(split->ComputeFusedFullWithFilterIP(computer, + cluster_id, + filter_code.data(), + supplement_code.data(), + filter_inner_product, + &hinted_full_distance, + nullptr)); + REQUIRE(std::isfinite(hinted_full_distance)); + const float tolerance = 2e-4F * std::max({1.0F, + std::abs(direct_full_distance), + std::abs(hinted_full_distance)}); + REQUIRE(std::abs(direct_full_distance - hinted_full_distance) <= tolerance); + REQUIRE_FALSE( + split->ComputeFusedFullWithFilterIP(computer, + cluster_id, + filter_code.data(), + supplement_code.data(), + std::numeric_limits::quiet_NaN(), + &hinted_full_distance, + nullptr)); + + RaBitQFusedTraversalQuery traversal_query; + REQUIRE(split->GetFusedTraversalQuery(computer, &traversal_query)); + std::vector invalid_filter_code(filter_code.begin(), filter_code.end()); + const float invalid_metadata = std::numeric_limits::quiet_NaN(); + std::memcpy(invalid_filter_code.data() + traversal_query.one_bit_metadata_offset, + &invalid_metadata, + sizeof(invalid_metadata)); + REQUIRE_FALSE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + invalid_filter_code.data(), + supplement_code.data(), + &coarse_distance, + &lower_bound, + &filter_inner_product, + nullptr)); + REQUIRE(split->ComputeFusedFull(computer, + cluster_id, + invalid_filter_code.data(), + supplement_code.data(), + &direct_full_distance, + nullptr)); + + auto invalid_bound_code = + std::vector(filter_code.begin(), filter_code.end()); + const float overflowing_error_unit = std::numeric_limits::max(); + std::memcpy(invalid_bound_code.data() + traversal_query.one_bit_metadata_offset + + 2U * sizeof(float), + &overflowing_error_unit, + sizeof(overflowing_error_unit)); + coarse_distance = std::numeric_limits::max(); + REQUIRE_FALSE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + invalid_bound_code.data(), + supplement_code.data(), + &coarse_distance, + &lower_bound, + &filter_inner_product, + nullptr)); + REQUIRE(std::isfinite(coarse_distance)); + REQUIRE(coarse_distance < std::numeric_limits::max()); + + graph->SetNodeCodes(id, + static_cast(id), + cluster_id, + invalid_bound_code.data(), + supplement_code.data()); + SearchStatistics no_reorder_stats; + QueryContext no_reorder_context; + no_reorder_context.stats = &no_reorder_stats; + no_reorder_context.enable_rabitq_reorder = false; + const InnerIdType query_id = id; + float queried_distance = std::numeric_limits::max(); + float queried_lower_bound = std::numeric_limits::max(); + float queried_filter_ip = std::numeric_limits::quiet_NaN(); + flatten->QueryWithDistanceLowerBoundAndFilterIP(&queried_distance, + &queried_lower_bound, + &queried_filter_ip, + computer, + &query_id, + 1, + &no_reorder_context); + REQUIRE(queried_distance == coarse_distance); + REQUIRE(queried_lower_bound == coarse_distance); + REQUIRE(IsNaNBitPattern(queried_filter_ip)); + + float filtered_distance = std::numeric_limits::max(); + flatten->QueryWithDistanceFilter(&filtered_distance, + computer, + &query_id, + 1, + std::numeric_limits::max(), + &no_reorder_context); + REQUIRE(filtered_distance == coarse_distance); + REQUIRE(no_reorder_stats.rabitq_filter_count.load() == 2); + REQUIRE(no_reorder_stats.rabitq_full_count.load() == 0); + REQUIRE(no_reorder_stats.rabitq_filter_fallback_full_count.load() == 0); + graph->SetNodeCodes(id, + static_cast(id), + cluster_id, + filter_code.data(), + supplement_code.data()); + } else { + float hinted_full_distance = 0.0F; + REQUIRE_FALSE(split->ComputeFusedFullWithFilterIP(computer, + cluster_id, + filter_code.data(), + supplement_code.data(), + 0.0F, + &hinted_full_distance, + nullptr)); + } + } + + constexpr InnerIdType alias_id = 3; + flatten->InsertVector(encoded_vectors.data(), alias_id); + uint32_t alias_cluster_id = cluster_count; + REQUIRE(split->EncodeFused( + encoded_vectors.data(), filter_code.data(), supplement_code.data(), &alias_cluster_id)); + graph->SetNodeCodes(alias_id, + static_cast(alias_id), + alias_cluster_id, + filter_code.data(), + supplement_code.data()); + + Vector decoded(split_case.dim, 0.0F, allocator.get()); + Vector decoded_alias(split_case.dim, 0.0F, allocator.get()); + REQUIRE(split->DecodeFusedById(0, decoded.data())); + REQUIRE(split->DecodeFusedById(alias_id, decoded_alias.data())); + float source_norm_sqr = 0.0F; + float decode_error_sqr = 0.0F; + for (uint64_t d = 0; d < split_case.dim; ++d) { + REQUIRE(std::isfinite(decoded[d])); + REQUIRE(std::abs(decoded[d] - decoded_alias[d]) <= 1e-6F); + source_norm_sqr += encoded_vectors[d] * encoded_vectors[d]; + const float error = decoded[d] - encoded_vectors[d]; + decode_error_sqr += error * error; + } + REQUIRE(decode_error_sqr < source_norm_sqr); + REQUIRE_FALSE(split->DecodeFusedById(alias_id + 1, decoded.data())); + REQUIRE_FALSE(split->DecodeFusedById(0, nullptr)); + } +} + +TEST_CASE("RaBitQSplitDataCell fused zero residual metadata", + "[ut][RaBitQSplitDataCell][fused_zero_residual]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + constexpr uint64_t dim = 64; + constexpr InnerIdType count = 32; + constexpr uint32_t cluster_count = 16; + constexpr const char* param_template = R"( + {{ + "codes_type": "rabitq_split", + "io_params": {{ + "type": "memory_io" + }}, + "quantization_params": {{ + "type": "rabitq", + "rabitq_version": "split", + "rabitq_bits_per_dim_query": 32, + "rabitq_bits_per_dim_base": {}, + "rabitq_bits_per_dim_filter": {}, + "use_fht": true + }} + }} + )"; + + Vector vectors(static_cast(count) * dim, 0.0F, allocator.get()); + for (uint64_t d = 0; d < dim; ++d) { + const float value = static_cast(static_cast(d % 13) - 6) * 0.125F; + for (InnerIdType row = 0; row < count; ++row) { + vectors[static_cast(row) * dim + d] = value; + } + } + Vector queries(2 * dim, 0.0F, allocator.get()); + std::copy_n(vectors.data(), dim, queries.data()); + for (uint64_t d = 0; d < dim; ++d) { + queries[dim + d] = vectors[d] + static_cast(static_cast(d % 7) - 3) * 0.05F; + } + + for (const auto metric : {MetricType::METRIC_TYPE_L2SQR, MetricType::METRIC_TYPE_IP}) { + for (uint32_t filter_bits = 1; filter_bits <= 4; ++filter_bits) { + CAPTURE(static_cast(metric), filter_bits); + const uint32_t supplement_bits = filter_bits == 1 ? 3 : 8 - filter_bits; + const uint32_t base_bits = filter_bits + supplement_bits; + auto param_json = JsonType::Parse(fmt::format(param_template, base_bits, filter_bits)); + auto param = std::make_shared(); + param->FromJson(param_json); + + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = dim; + common_param.metric_ = metric; + auto flatten = FlattenInterface::MakeInstance(param, common_param); + flatten->Train(vectors.data(), count); + auto split = std::dynamic_pointer_cast(flatten); + REQUIRE(split != nullptr); + REQUIRE_FALSE(split->UsesLegacyHnswFusedCodec()); + split->TrainFusedCodec(vectors.data(), count, cluster_count); + + Vector one_bit(split->OneBitCodeSize(), allocator.get()); + Vector supplement(split->SupplementCodeSize(), allocator.get()); + uint32_t cluster_id = cluster_count; + REQUIRE( + split->EncodeFused(vectors.data(), one_bit.data(), supplement.data(), &cluster_id)); + REQUIRE(cluster_id < cluster_count); + + auto first_computer = split->FactoryFusedComputer(queries.data()); + RaBitQFusedTraversalQuery traversal_query; + REQUIRE(split->GetFusedTraversalQuery(first_computer, &traversal_query)); + float filter_add = std::numeric_limits::quiet_NaN(); + float filter_rescale = std::numeric_limits::quiet_NaN(); + float filter_error_unit = std::numeric_limits::quiet_NaN(); + const auto* metadata = one_bit.data() + traversal_query.one_bit_metadata_offset; + std::memcpy(&filter_add, metadata, sizeof(filter_add)); + std::memcpy(&filter_rescale, metadata + sizeof(float), sizeof(filter_rescale)); + std::memcpy( + &filter_error_unit, metadata + 2U * sizeof(float), sizeof(filter_error_unit)); + REQUIRE(std::isfinite(filter_add)); + REQUIRE(filter_rescale == 0.0F); + REQUIRE(filter_error_unit == 0.0F); + + for (uint64_t query_id = 0; query_id < 2; ++query_id) { + auto computer = split->FactoryFusedComputer(queries.data() + query_id * dim); + float coarse_distance = std::numeric_limits::max(); + float lower_bound = std::numeric_limits::max(); + float filter_inner_product = std::numeric_limits::quiet_NaN(); + REQUIRE(split->ComputeFusedOneBitWithFilterIP(computer, + cluster_id, + one_bit.data(), + supplement.data(), + &coarse_distance, + &lower_bound, + &filter_inner_product, + nullptr)); + REQUIRE(std::isfinite(coarse_distance)); + REQUIRE(coarse_distance < std::numeric_limits::max()); + REQUIRE(std::isfinite(lower_bound)); + REQUIRE(lower_bound <= coarse_distance + 1e-5F); + double expected_distance = metric == MetricType::METRIC_TYPE_IP ? 1.0 : 0.0; + for (uint64_t d = 0; d < dim; ++d) { + const double base = vectors[d]; + const double query = queries[query_id * dim + d]; + if (metric == MetricType::METRIC_TYPE_IP) { + expected_distance -= base * query; + } else { + const double difference = base - query; + expected_distance += difference * difference; + } + } + const float expected = static_cast(expected_distance); + const float expected_tolerance = 5e-4F * std::max(1.0F, std::fabs(expected)); + REQUIRE(std::fabs(coarse_distance - expected) <= expected_tolerance); + + float full_distance = std::numeric_limits::max(); + REQUIRE(split->ComputeFusedFull(computer, + cluster_id, + one_bit.data(), + supplement.data(), + &full_distance, + nullptr)); + REQUIRE(std::isfinite(full_distance)); + REQUIRE(full_distance < std::numeric_limits::max()); + const float tolerance = + 1e-5F * std::max({1.0F, std::fabs(coarse_distance), std::fabs(full_distance)}); + REQUIRE(std::fabs(full_distance - coarse_distance) <= tolerance); + } + } + } +} + +TEST_CASE("RaBitQSplitDataCell fused model-only serialization", + "[ut][RaBitQSplitDataCell][serialize]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + constexpr uint64_t dim = 64; + constexpr uint32_t cluster_count = 16; + constexpr InnerIdType count = 128; + constexpr uint64_t vectors_per_cluster = count / cluster_count; + Vector vectors(static_cast(count) * dim, 0.0F, allocator.get()); + for (uint32_t cluster = 0; cluster < cluster_count; ++cluster) { + for (uint64_t row = 0; row < vectors_per_cluster; ++row) { + auto* vector = + vectors.data() + (static_cast(cluster) * vectors_per_cluster + row) * dim; + vector[cluster] = 100.0F; + vector[(cluster + 17) % dim] = static_cast(row) * 0.01F; + } + } + + auto param_json = JsonType::Parse(R"({ + "codes_type": "rabitq_split", + "io_params": {"type": "memory_io"}, + "quantization_params": { + "type": "rabitq", + "rabitq_version": "split", + "rabitq_bits_per_dim_query": 32, + "rabitq_bits_per_dim_base": 8, + "rabitq_bits_per_dim_filter": 1, + "use_fht": true + } + })"); + auto param = std::make_shared(); + param->FromJson(param_json); + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = dim; + common_param.metric_ = MetricType::METRIC_TYPE_L2SQR; + + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->max_degree_ = 32; + graph_param->init_max_capacity_ = count; + + auto serialize_flatten = [](const FlattenInterfacePtr& value) { + std::stringstream stream; + IOStreamWriter writer(stream); + value->Serialize(writer); + return stream.str(); + }; + auto serialize_graph = [](const HGraphRaBitQFusedDataCellPtr& value) { + std::stringstream stream; + IOStreamWriter writer(stream); + value->Serialize(writer); + return stream.str(); + }; + auto deserialize_flatten = [](const FlattenInterfacePtr& value, const std::string& payload) { + std::stringstream stream(payload); + IOStreamReader reader(stream); + value->Deserialize(reader); + }; + auto deserialize_graph = [](const HGraphRaBitQFusedDataCellPtr& value, + const std::string& payload) { + std::stringstream stream(payload); + IOStreamReader reader(stream); + value->Deserialize(reader); + }; + + auto flatten = FlattenInterface::MakeInstance(param, common_param); + flatten->Train(vectors.data(), count); + flatten->BatchInsertVector(vectors.data(), count); + auto split = std::dynamic_pointer_cast(flatten); + REQUIRE(split != nullptr); + split->TrainFusedCodec(vectors.data(), count, cluster_count); + + const auto legacy_payload = serialize_flatten(flatten); + const uint64_t memory_with_split_codes = flatten->GetMemoryUsage(); + const uint64_t code_payload_size = + static_cast(count) * (split->OneBitCodeSize() + split->SupplementCodeSize()); + + auto graph = std::make_shared( + graph_param, split->OneBitCodeSize(), split->SupplementCodeSize(), common_param); + Vector one_bit(split->OneBitCodeSize(), allocator.get()); + Vector supplement(split->SupplementCodeSize(), allocator.get()); + Vector empty_neighbors(allocator.get()); + for (InnerIdType id = 0; id < count; ++id) { + uint32_t cluster_id = 0; + REQUIRE(split->EncodeFused(vectors.data() + static_cast(id) * dim, + one_bit.data(), + supplement.data(), + &cluster_id)); + graph->SetNodeCodes( + id, static_cast(id), cluster_id, one_bit.data(), supplement.data()); + graph->InsertNeighborsById(id, empty_neighbors); + } + graph->SetCodecModel(split->ExportFusedCodec()); + split->AttachFusedCodeStorage(graph.get()); + + Vector decoded(dim, 0.0F, allocator.get()); + REQUIRE(split->DecodeFusedById(0, decoded.data())); + for (const float value : decoded) { + REQUIRE(std::isfinite(value)); + } + REQUIRE_FALSE(split->DecodeFusedById(count, decoded.data())); + REQUIRE_FALSE(split->DecodeFusedById(0, nullptr)); + const auto expected_decoded = decoded; + + const auto model_payload = serialize_flatten(flatten); + const auto graph_payload = serialize_graph(graph); + REQUIRE(memory_with_split_codes == flatten->GetMemoryUsage() + code_payload_size); + REQUIRE(static_cast(legacy_payload.size()) + sizeof(uint32_t) == + static_cast(model_payload.size()) + code_payload_size); + + auto make_attached_pair = [&]() { + auto restored_flatten = FlattenInterface::MakeInstance(param, common_param); + auto restored_split = + std::dynamic_pointer_cast(restored_flatten); + REQUIRE(restored_split != nullptr); + auto restored_graph = + std::make_shared(graph_param, + restored_split->OneBitCodeSize(), + restored_split->SupplementCodeSize(), + common_param); + restored_split->AttachFusedCodeStorage(restored_graph.get()); + return std::make_tuple(restored_flatten, restored_split, restored_graph); + }; + + auto [model_flatten, model_split, model_graph] = make_attached_pair(); + deserialize_flatten(model_flatten, model_payload); + deserialize_graph(model_graph, graph_payload); + model_split->ImportFusedCodec(model_graph->CodecModel()); + REQUIRE(model_split->DecodeFusedById(0, decoded.data())); + for (uint64_t d = 0; d < dim; ++d) { + REQUIRE(std::abs(decoded[d] - expected_decoded[d]) <= 1e-6F); + } + + auto query = fixtures::generate_vectors(1, dim, 97); + std::vector ids(count); + std::iota(ids.begin(), ids.end(), 0); + auto query_all = [&](const FlattenInterfacePtr& value) { + auto computer = value->FactoryComputer(query.data()); + std::vector distances(count); + value->Query(distances.data(), computer, ids.data(), count); + return distances; + }; + const auto expected_distances = query_all(flatten); + const auto model_distances = query_all(model_flatten); + for (InnerIdType id = 0; id < count; ++id) { + REQUIRE(std::abs(expected_distances[id] - model_distances[id]) <= 1e-6F); + } + + auto [legacy_flatten, legacy_split, legacy_graph] = make_attached_pair(); + deserialize_flatten(legacy_flatten, legacy_payload); + deserialize_graph(legacy_graph, graph_payload); + legacy_split->ImportFusedCodec(legacy_graph->CodecModel()); + REQUIRE(legacy_split->DecodeFusedById(0, decoded.data())); + for (uint64_t d = 0; d < dim; ++d) { + REQUIRE(std::abs(decoded[d] - expected_decoded[d]) <= 1e-6F); + } + REQUIRE(legacy_split->UsesExternalFusedCodeStorage()); + using MemorySplitDataCell = + RaBitQSplitDataCell; + auto legacy_memory_split = std::dynamic_pointer_cast(legacy_flatten); + REQUIRE(legacy_memory_split != nullptr); + REQUIRE(legacy_memory_split->x_bit_cell_->GetMemoryUsage() == 0); + REQUIRE(legacy_memory_split->supplement_cell_->GetMemoryUsage() == 0); + REQUIRE(legacy_flatten->GetMemoryUsage() == model_flatten->GetMemoryUsage()); + const auto legacy_distances = query_all(legacy_flatten); + for (InnerIdType id = 0; id < count; ++id) { + REQUIRE(std::abs(expected_distances[id] - legacy_distances[id]) <= 1e-6F); + } +} + TEST_CASE("RaBitQSplitDataCell serialize and methods", "[ut][RaBitQSplitDataCell]") { auto allocator = SafeAllocator::FactoryDefaultAllocator(); constexpr uint64_t dim = 64; diff --git a/src/datacell/flatten_interface.h b/src/datacell/flatten_interface.h index 6f1a5376f0..f4b5238eb5 100644 --- a/src/datacell/flatten_interface.h +++ b/src/datacell/flatten_interface.h @@ -81,6 +81,22 @@ class FlattenInterface { } } + virtual void + QueryWithDistanceLowerBoundAndFilterIP(float* result_dists, + float* lower_bounds, + float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* idx, + InnerIdType id_count, + QueryContext* ctx = nullptr) { + this->QueryWithDistanceLowerBound(result_dists, lower_bounds, computer, idx, id_count, ctx); + if (filter_inner_products != nullptr) { + std::fill(filter_inner_products, + filter_inner_products + id_count, + std::numeric_limits::quiet_NaN()); + } + } + virtual void QueryWithDistanceHint(float* result_dists, const float* /*hint_dists*/, @@ -91,6 +107,16 @@ class FlattenInterface { this->Query(result_dists, computer, idx, id_count, ctx); } + virtual void + QueryWithFilterIPHint(float* result_dists, + const float* /*filter_inner_products*/, + const ComputerInterfacePtr& computer, + const InnerIdType* idx, + InnerIdType id_count, + QueryContext* ctx = nullptr) { + this->Query(result_dists, computer, idx, id_count, ctx); + } + virtual ComputerInterfacePtr FactoryComputer(const void* query) = 0; diff --git a/src/datacell/hgraph_rabitq_fused_datacell.cpp b/src/datacell/hgraph_rabitq_fused_datacell.cpp new file mode 100644 index 0000000000..79afc71296 --- /dev/null +++ b/src/datacell/hgraph_rabitq_fused_datacell.cpp @@ -0,0 +1,608 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hgraph_rabitq_fused_datacell.h" + +#include +#include +#include +#include + +#include "rabitq_split_datacell.h" +#include "storage/stream_reader.h" +#include "storage/stream_writer.h" + +namespace vsag { + +namespace { +constexpr uint64_t K_CACHE_LINE_SIZE = 64; +constexpr uint64_t K_COUNT_OFFSET = 0; +constexpr uint64_t K_VERSION_OFFSET = sizeof(uint32_t); +constexpr uint64_t K_HEADER_SIZE = 2 * sizeof(uint32_t); +constexpr uint32_t K_SERIALIZATION_VERSION = 1; +constexpr uint64_t K_FUSED_CLUSTER_COUNT = 16; + +struct fused_wire_layout { + uint64_t record_size{0}; + uint64_t neighbors_offset{0}; + uint64_t cluster_id_offset{0}; + uint64_t label_offset{0}; + uint64_t one_bit_offset{0}; + uint64_t supplement_offset{0}; + uint64_t one_bit_code_size{0}; + uint64_t supplement_code_size{0}; + bool support_remove{false}; + uint32_t remove_flag_bit{0}; +}; + +uint64_t +fused_codec_model_size(int64_t dim) { + CHECK_ARGUMENT(dim > 0, "invalid fused RaBitQ dimension"); + constexpr uint64_t fixed_size = sizeof(uint32_t) + sizeof(uint64_t) + sizeof(uint32_t); + constexpr uint64_t bytes_per_dimension = K_FUSED_CLUSTER_COUNT * sizeof(float); + const auto unsigned_dim = static_cast(dim); + CHECK_ARGUMENT( + unsigned_dim <= (std::numeric_limits::max() - fixed_size) / bytes_per_dimension, + "fused RaBitQ codec size overflow"); + return fixed_size + unsigned_dim * bytes_per_dimension; +} + +uint64_t +remaining_bytes(StreamReader& reader) { + const auto length = reader.Length(); + const auto cursor = reader.GetCursor(); + CHECK_ARGUMENT(cursor <= length, "invalid fused graph reader cursor"); + return length - cursor; +} +} // namespace + +uint64_t +HGraphRaBitQFusedDataCell::AlignUp(uint64_t value, uint64_t alignment) { + return (value + alignment - 1) / alignment * alignment; +} + +HGraphRaBitQFusedDataCell::HGraphRaBitQFusedDataCell(const GraphDataCellParamPtr& graph_param, + uint64_t one_bit_code_size, + uint64_t supplement_code_size, + const IndexCommonParam& common_param) + : storage_(common_param.allocator_.get()), + one_bit_code_size_(one_bit_code_size), + supplement_code_size_(supplement_code_size), + dim_(common_param.dim_) { + CHECK_ARGUMENT(graph_param != nullptr, "fused graph parameter must not be null"); + CHECK_ARGUMENT(graph_param->max_degree_ <= std::numeric_limits::max(), + "fused graph maximum degree exceeds uint32 range"); + CHECK_ARGUMENT(graph_param->init_max_capacity_ <= std::numeric_limits::max(), + "fused graph initial capacity exceeds id range"); + allocator_ = common_param.allocator_.get(); + maximum_degree_ = static_cast(graph_param->max_degree_); + max_capacity_ = static_cast(graph_param->init_max_capacity_); + support_remove_ = graph_param->support_remove_; + remove_flag_bit_ = graph_param->remove_flag_bit_; + constexpr uint32_t id_width = sizeof(InnerIdType) * 8; + CHECK_ARGUMENT(remove_flag_bit_ < id_width, "invalid fused graph remove flag bits"); + if (support_remove_) { + CHECK_ARGUMENT(remove_flag_bit_ > 0, "fused graph removal requires version bits"); + } + id_bit_ = id_width - remove_flag_bit_; + remove_flag_mask_ = id_bit_ == id_width ? std::numeric_limits::max() + : (InnerIdType{1} << id_bit_) - 1U; + if (support_remove_) { + CHECK_ARGUMENT(max_capacity_ <= remove_flag_mask_, + "fused graph initial capacity exceeds remove-id bits"); + } + + neighbors_offset_ = K_HEADER_SIZE; + cluster_id_offset_ = + neighbors_offset_ + static_cast(maximum_degree_) * sizeof(InnerIdType); + label_offset_ = AlignUp(cluster_id_offset_ + sizeof(uint32_t), alignof(LabelType)); + one_bit_offset_ = label_offset_ + sizeof(LabelType); + supplement_offset_ = one_bit_offset_ + one_bit_code_size_; + record_size_ = AlignUp(supplement_offset_ + supplement_code_size_, K_CACHE_LINE_SIZE); + CHECK_ARGUMENT( + max_capacity_ <= + (std::numeric_limits::max() - (K_CACHE_LINE_SIZE - 1)) / record_size_, + "fused graph initial capacity and record stride overflow"); + + if (graph_param->use_reverse_edges_) { + reverse_edges_ = std::make_unique(allocator_); + } + if (graph_param->support_duplicate_) { + InitDuplicateTracker(); + } + Reallocate(max_capacity_); +} + +void +HGraphRaBitQFusedDataCell::Reallocate(InnerIdType new_capacity) { + Vector replacement( + static_cast(new_capacity) * record_size_ + K_CACHE_LINE_SIZE - 1, 0, allocator_); + const auto replacement_address = reinterpret_cast(replacement.data()); + const uint64_t replacement_offset = + (K_CACHE_LINE_SIZE - replacement_address % K_CACHE_LINE_SIZE) % K_CACHE_LINE_SIZE; + if (not storage_.empty()) { + const auto copy_count = std::min(max_capacity_, new_capacity); + std::memcpy(replacement.data() + replacement_offset, + storage_.data() + aligned_offset_, + static_cast(copy_count) * record_size_); + } + storage_ = std::move(replacement); + aligned_offset_ = replacement_offset; + max_capacity_ = new_capacity; + if (duplicate_tracker_ != nullptr) { + duplicate_tracker_->Resize(new_capacity); + } +} + +uint8_t* +HGraphRaBitQFusedDataCell::MutableNodeRecord(InnerIdType id) { + return storage_.data() + aligned_offset_ + static_cast(id) * record_size_; +} + +const InnerIdType* +HGraphRaBitQFusedDataCell::GetNeighborData(const uint8_t* record) const { + return reinterpret_cast(record + neighbors_offset_); +} + +const uint8_t* +HGraphRaBitQFusedDataCell::GetOneBitCode(const uint8_t* record) const { + return record + one_bit_offset_; +} + +const uint8_t* +HGraphRaBitQFusedDataCell::GetSupplementCode(const uint8_t* record) const { + return record + supplement_offset_; +} + +LabelType +HGraphRaBitQFusedDataCell::GetLabel(const uint8_t* record) const { + LabelType label = 0; + std::memcpy(&label, record + label_offset_, sizeof(label)); + return label; +} + +uint32_t +HGraphRaBitQFusedDataCell::GetClusterId(const uint8_t* record) const { + uint32_t cluster_id = 0; + std::memcpy(&cluster_id, record + cluster_id_offset_, sizeof(cluster_id)); + return cluster_id; +} + +uint32_t +HGraphRaBitQFusedDataCell::NodeVersion(const uint8_t* record) { + uint32_t version = 0; + std::memcpy(&version, record + K_VERSION_OFFSET, sizeof(version)); + return version; +} + +void +HGraphRaBitQFusedDataCell::SetNodeVersion(uint8_t* record, uint32_t version) { + std::memcpy(record + K_VERSION_OFFSET, &version, sizeof(version)); +} + +void +HGraphRaBitQFusedDataCell::InsertNeighborsById(InnerIdType id, + const Vector& neighbor_ids) { + CHECK_ARGUMENT(neighbor_ids.size() <= maximum_degree_, + "fused node neighbor count exceeds maximum degree"); + if (id >= max_capacity_) { + Resize(id + 1); + } + + Vector old_neighbors(allocator_); + if (reverse_edges_ != nullptr and id < total_count_) { + GetNeighbors(id, old_neighbors); + } + UpdateReverseEdges(id, old_neighbors, neighbor_ids); + + auto* record = MutableNodeRecord(id); + const auto count = static_cast(neighbor_ids.size()); + std::memcpy(record + K_COUNT_OFFSET, &count, sizeof(count)); + auto* output = reinterpret_cast(record + neighbors_offset_); + for (uint64_t i = 0; i < neighbor_ids.size(); ++i) { + auto neighbor = neighbor_ids[i]; + if (support_remove_) { + const auto version = NodeVersion(GetNodeRecord(neighbor)); + neighbor |= static_cast(version << id_bit_); + } + output[i] = neighbor; + } + auto current = total_count_.load(); + while (current < id + 1 and not total_count_.compare_exchange_weak(current, id + 1)) { + } +} + +void +HGraphRaBitQFusedDataCell::SetNodeCodes(InnerIdType id, + LabelType label, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code) { + SetFusedCodes(id, cluster_id, one_bit_code, supplement_code); + std::memcpy(MutableNodeRecord(id) + label_offset_, &label, sizeof(label)); +} + +bool +HGraphRaBitQFusedDataCell::GetFusedCodeView(InnerIdType id, RaBitQFusedCodeView& view) const { + // Logical duplicate aliases own encoded records but no graph edges, so their IDs may be + // greater than the graph-node high-water mark. Callers validate logical IDs through the + // label table; this layer only enforces the allocated slab boundary. + if (id >= max_capacity_) { + view = {}; + return false; + } + const auto* record = GetNodeRecord(id); + view.one_bit_code = record + one_bit_offset_; + view.supplement_code = record + supplement_offset_; + std::memcpy(&view.cluster_id, record + cluster_id_offset_, sizeof(view.cluster_id)); + return true; +} + +void +HGraphRaBitQFusedDataCell::SetFusedCodes(InnerIdType id, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code) { + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + one_bit_code != nullptr and supplement_code != nullptr, + "fused RaBitQ codes must not be null"); + if (id >= max_capacity_) { + Resize(id + 1); + } + auto* record = MutableNodeRecord(id); + std::memcpy(record + cluster_id_offset_, &cluster_id, sizeof(cluster_id)); + std::memcpy(record + one_bit_offset_, one_bit_code, one_bit_code_size_); + std::memcpy(record + supplement_offset_, supplement_code, supplement_code_size_); +} + +void +HGraphRaBitQFusedDataCell::PrefetchFusedCodes(InnerIdType id, bool include_supplement) const { + const auto* record = GetNodeRecord(id); + const auto prefetch_l1 = [](const uint8_t* begin, uint64_t size) { + const auto begin_address = reinterpret_cast(begin); + const auto begin_offset = begin_address % K_CACHE_LINE_SIZE; + const auto* first = begin - begin_offset; + const auto span = begin_offset + size; + for (uint64_t offset = 0; offset < span; offset += K_CACHE_LINE_SIZE) { + __builtin_prefetch(first + offset, 0, 3); + } + }; + prefetch_l1(record + one_bit_offset_, one_bit_code_size_); + if (include_supplement) { + prefetch_l1(record + supplement_offset_, supplement_code_size_); + } +} + +void +HGraphRaBitQFusedDataCell::SetLabel(InnerIdType id, LabelType label) { + CHECK_ARGUMENT(id < max_capacity_, "fused graph label id does not exist"); + std::memcpy(MutableNodeRecord(id) + label_offset_, &label, sizeof(label)); +} + +bool +HGraphRaBitQFusedDataCell::SyncNodeCodes(InnerIdType id, + LabelType label, + uint32_t cluster_id, + const RaBitQSplitDataCellInterface& split_codes) { + if (id >= max_capacity_) { + Resize(id + 1); + } + auto* record = MutableNodeRecord(id); + if (not split_codes.CopySplitCodes(id, record + one_bit_offset_, record + supplement_offset_)) { + return false; + } + std::memcpy(record + cluster_id_offset_, &cluster_id, sizeof(cluster_id)); + std::memcpy(record + label_offset_, &label, sizeof(label)); + return true; +} + +uint32_t +HGraphRaBitQFusedDataCell::GetNeighborSize(InnerIdType id) const { + uint32_t count = 0; + std::memcpy(&count, GetNodeRecord(id) + K_COUNT_OFFSET, sizeof(count)); + return count; +} + +void +HGraphRaBitQFusedDataCell::GetNeighbors(InnerIdType id, Vector& neighbor_ids) const { + const auto* record = GetNodeRecord(id); + const auto count = GetNeighborSize(id); + if (count > maximum_degree_) { + neighbor_ids.clear(); + return; + } + const auto* input = GetNeighborData(record); + neighbor_ids.clear(); + neighbor_ids.reserve(count); + for (uint32_t i = 0; i < count; ++i) { + InnerIdType neighbor = 0; + if (not ResolveNeighbor(input[i], neighbor)) { + continue; + } + neighbor_ids.push_back(neighbor); + } +} + +bool +HGraphRaBitQFusedDataCell::CheckIdExists(InnerIdType id) const { + return id < total_count_ and id < max_capacity_; +} + +void +HGraphRaBitQFusedDataCell::Resize(InnerIdType new_size) { + std::unique_lock lock(storage_mutex_); + if (new_size <= max_capacity_) { + return; + } + if (support_remove_ and new_size > remove_flag_mask_) { + throw VsagException(ErrorType::INTERNAL_ERROR, "fused graph id capacity exceeded"); + } + Reallocate(new_size); +} + +void +HGraphRaBitQFusedDataCell::Prefetch(InnerIdType id, uint32_t neighbor_i) { + const auto* record = GetNodeRecord(id); + __builtin_prefetch( + record + neighbors_offset_ + static_cast(neighbor_i) * sizeof(InnerIdType), 0, 3); + __builtin_prefetch(record + one_bit_offset_, 0, 3); +} + +void +HGraphRaBitQFusedDataCell::DeleteNeighborsById(InnerIdType id) { + CHECK_ARGUMENT(support_remove_, "remove is disabled for fused graph"); + CHECK_ARGUMENT(id < max_capacity_, "fused graph remove id does not exist"); + auto* record = MutableNodeRecord(id); + const auto version = NodeVersion(record); + CHECK_ARGUMENT(version < (1U << remove_flag_bit_) - 1U, "fused graph node version exhausted"); + SetNodeVersion(record, version + 1); +} + +void +HGraphRaBitQFusedDataCell::RecoverDeleteNeighborsById(InnerIdType id) { + CHECK_ARGUMENT(support_remove_, "remove is disabled for fused graph"); + CHECK_ARGUMENT(id < max_capacity_, "fused graph recover id does not exist"); + auto* record = MutableNodeRecord(id); + const auto version = NodeVersion(record); + CHECK_ARGUMENT(version > 0, "fused graph node has not been removed"); + SetNodeVersion(record, version - 1); +} + +void +HGraphRaBitQFusedDataCell::Move(InnerIdType from, InnerIdType to) { + if (from == to) { + return; + } + if (to >= max_capacity_) { + Resize(to + 1); + } + + Vector source_record(record_size_, allocator_); + std::memcpy(source_record.data(), GetNodeRecord(from), record_size_); + + Vector reverse_neighbors(allocator_); + GetIncomingNeighbors(from, reverse_neighbors); + Vector neighbors(allocator_); + InsertNeighborsById(to, neighbors); + for (const auto reverse_neighbor : reverse_neighbors) { + GetNeighbors(reverse_neighbor, neighbors); + Vector replacement(allocator_); + bool contains_to = false; + for (const auto neighbor : neighbors) { + if (neighbor != from) { + replacement.push_back(neighbor); + } + contains_to = contains_to or neighbor == to; + } + if (not contains_to) { + replacement.push_back(to); + } + InsertNeighborsById(reverse_neighbor, replacement); + } + + Vector from_neighbors(allocator_); + GetNeighbors(from, from_neighbors); + InsertNeighborsById(to, from_neighbors); + from_neighbors.clear(); + InsertNeighborsById(from, from_neighbors); + + auto* target_record = MutableNodeRecord(to); + SetNodeVersion(target_record, NodeVersion(source_record.data())); + std::memcpy(target_record + cluster_id_offset_, + source_record.data() + cluster_id_offset_, + record_size_ - cluster_id_offset_); +} + +void +HGraphRaBitQFusedDataCell::ShrinkToFit(InnerIdType capacity) { + std::unique_lock lock(storage_mutex_); + if (capacity < total_count_) { + capacity = total_count_; + } + Reallocate(capacity); +} + +void +HGraphRaBitQFusedDataCell::Serialize(StreamWriter& writer) { + GraphInterface::Serialize(writer); + StreamWriter::WriteObj(writer, K_SERIALIZATION_VERSION); + StreamWriter::WriteObj(writer, record_size_); + StreamWriter::WriteObj(writer, neighbors_offset_); + StreamWriter::WriteObj(writer, cluster_id_offset_); + StreamWriter::WriteObj(writer, label_offset_); + StreamWriter::WriteObj(writer, one_bit_offset_); + StreamWriter::WriteObj(writer, supplement_offset_); + StreamWriter::WriteObj(writer, one_bit_code_size_); + StreamWriter::WriteObj(writer, supplement_code_size_); + StreamWriter::WriteObj(writer, support_remove_); + StreamWriter::WriteObj(writer, remove_flag_bit_); + StreamWriter::WriteString(writer, codec_model_); + const uint64_t bytes = static_cast(max_capacity_) * record_size_; + StreamWriter::WriteObj(writer, bytes); + writer.Write(reinterpret_cast(storage_.data() + aligned_offset_), bytes); +} + +void +HGraphRaBitQFusedDataCell::Deserialize(StreamReader& reader) { + const auto expected_maximum_degree = maximum_degree_; + const fused_wire_layout expected_layout{record_size_, + neighbors_offset_, + cluster_id_offset_, + label_offset_, + one_bit_offset_, + supplement_offset_, + one_bit_code_size_, + supplement_code_size_, + support_remove_, + remove_flag_bit_}; + + GraphInterface::Deserialize(reader); + uint32_t version = 0; + StreamReader::ReadObj(reader, version); + CHECK_ARGUMENT(version == K_SERIALIZATION_VERSION, + "unsupported fused graph serialization version"); + + fused_wire_layout layout; + StreamReader::ReadObj(reader, layout.record_size); + StreamReader::ReadObj(reader, layout.neighbors_offset); + StreamReader::ReadObj(reader, layout.cluster_id_offset); + StreamReader::ReadObj(reader, layout.label_offset); + StreamReader::ReadObj(reader, layout.one_bit_offset); + StreamReader::ReadObj(reader, layout.supplement_offset); + StreamReader::ReadObj(reader, layout.one_bit_code_size); + StreamReader::ReadObj(reader, layout.supplement_code_size); + uint8_t support_remove = 0; + static_assert(sizeof(support_remove) == sizeof(bool)); + StreamReader::ReadObj(reader, support_remove); + CHECK_ARGUMENT(support_remove <= 1, "invalid fused graph remove flag"); + layout.support_remove = support_remove != 0; + StreamReader::ReadObj(reader, layout.remove_flag_bit); + + constexpr uint32_t id_width = sizeof(InnerIdType) * 8; + CHECK_ARGUMENT(layout.remove_flag_bit < id_width, "invalid fused graph remove flag bits"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + not layout.support_remove or layout.remove_flag_bit > 0, + "fused graph removal requires version bits"); + CHECK_ARGUMENT(total_count_ <= max_capacity_, "invalid fused graph count and capacity"); + CHECK_ARGUMENT(maximum_degree_ == expected_maximum_degree, + "fused graph maximum degree does not match construction parameters"); + + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.record_size > 0 and layout.record_size % K_CACHE_LINE_SIZE == 0, + "invalid fused graph record stride"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.neighbors_offset == K_HEADER_SIZE and + layout.neighbors_offset % alignof(InnerIdType) == 0, + "invalid fused graph neighbor offset"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.cluster_id_offset >= layout.neighbors_offset and + static_cast(maximum_degree_) <= + (layout.cluster_id_offset - layout.neighbors_offset) / sizeof(InnerIdType), + "invalid fused graph cluster offset"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.cluster_id_offset % alignof(uint32_t) == 0 and + layout.label_offset >= layout.cluster_id_offset and + sizeof(uint32_t) <= layout.label_offset - layout.cluster_id_offset, + "invalid fused graph label offset"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.label_offset % alignof(LabelType) == 0 and + layout.one_bit_offset >= layout.label_offset and + sizeof(LabelType) <= layout.one_bit_offset - layout.label_offset, + "invalid fused graph one-bit offset"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.supplement_offset >= layout.one_bit_offset and + layout.one_bit_code_size <= layout.supplement_offset - layout.one_bit_offset, + "invalid fused graph supplement offset"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.record_size >= layout.supplement_offset and + layout.supplement_code_size <= layout.record_size - layout.supplement_offset, + "invalid fused graph code bounds"); + + const uint32_t wire_id_bit = id_width - layout.remove_flag_bit; + const InnerIdType wire_remove_flag_mask = wire_id_bit == id_width + ? std::numeric_limits::max() + : (InnerIdType{1} << wire_id_bit) - 1U; + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + not layout.support_remove or max_capacity_ <= wire_remove_flag_mask, + "fused graph capacity exceeds remove-id bits"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + max_capacity_ <= + (std::numeric_limits::max() - (K_CACHE_LINE_SIZE - 1)) / layout.record_size, + "fused graph capacity and record stride overflow"); + + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.record_size == expected_layout.record_size and + layout.neighbors_offset == expected_layout.neighbors_offset and + layout.cluster_id_offset == expected_layout.cluster_id_offset and + layout.label_offset == expected_layout.label_offset and + layout.one_bit_offset == expected_layout.one_bit_offset and + layout.supplement_offset == expected_layout.supplement_offset, + "fused graph record layout does not match construction parameters"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.one_bit_code_size == expected_layout.one_bit_code_size and + layout.supplement_code_size == expected_layout.supplement_code_size, + "fused graph code sizes do not match construction parameters"); + CHECK_ARGUMENT( // NOLINT(readability-simplify-boolean-expr) + layout.support_remove == expected_layout.support_remove and + layout.remove_flag_bit == expected_layout.remove_flag_bit, + "fused graph remove parameters do not match construction parameters"); + + uint64_t codec_model_size = 0; + StreamReader::ReadObj(reader, codec_model_size); + if (codec_model_size != 0) { + CHECK_ARGUMENT(codec_model_size == fused_codec_model_size(dim_), + "invalid fused RaBitQ codec payload size"); + } + auto available_bytes = remaining_bytes(reader); + CHECK_ARGUMENT(available_bytes >= sizeof(uint64_t), "truncated fused RaBitQ codec payload"); + available_bytes -= sizeof(uint64_t); + CHECK_ARGUMENT(codec_model_size <= available_bytes, "truncated fused RaBitQ codec payload"); + std::string codec_model(codec_model_size, '\0'); + reader.Read(codec_model.data(), codec_model_size); + + uint64_t bytes = 0; + StreamReader::ReadObj(reader, bytes); + const uint64_t expected_bytes = static_cast(max_capacity_) * layout.record_size; + CHECK_ARGUMENT(bytes == expected_bytes, "invalid fused graph payload size"); + CHECK_ARGUMENT(bytes <= remaining_bytes(reader), "truncated fused graph node payload"); + + codec_model_ = std::move(codec_model); + id_bit_ = wire_id_bit; + remove_flag_mask_ = wire_remove_flag_mask; + storage_.clear(); + aligned_offset_ = 0; + Reallocate(max_capacity_); + CHECK_ARGUMENT( + reinterpret_cast(storage_.data() + aligned_offset_) % K_CACHE_LINE_SIZE == 0, + "fused graph node slab is not cache-line aligned"); + if (bytes > 0) { + reader.Read(reinterpret_cast(storage_.data() + aligned_offset_), bytes); + } +} + +uint64_t +HGraphRaBitQFusedDataCell::GetMemoryUsage() const { + uint64_t result = sizeof(*this) + storage_.capacity() + codec_model_.capacity(); + if (reverse_edges_ != nullptr) { + result += reverse_edges_->GetMemoryUsage(); + } + return result; +} + +DuplicateTrackerPtr +HGraphRaBitQFusedDataCell::CreateDuplicateTracker() { + return std::make_shared(allocator_); +} + +} // namespace vsag diff --git a/src/datacell/hgraph_rabitq_fused_datacell.h b/src/datacell/hgraph_rabitq_fused_datacell.h new file mode 100644 index 0000000000..a77420f64b --- /dev/null +++ b/src/datacell/hgraph_rabitq_fused_datacell.h @@ -0,0 +1,300 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include +#include +#include +#include +#include + +#include "dense_duplicate_tracker.h" +#include "graph_datacell_parameter.h" +#include "graph_interface.h" +#include "index_common_param.h" +#include "rabitq_fused_code_storage.h" + +namespace vsag { + +class RaBitQSplitDataCellInterface; + +/** + * Bottom-layer HGraph storage specialized for fused RaBitQ split codes. + * + * The fixed-size, cache-line-aligned record follows the locality strategy used + * by RaBitQ-Library's HNSW implementation: links, cluster id, external label, + * x-bit code and y-bit supplement are addressed from one node pointer. The code sizes are fixed for + * an index and support filter widths from one through four bits. + */ +class HGraphRaBitQFusedDataCell final : public GraphInterface, + public RaBitQFusedCodeStorageInterface { +public: + struct NodeView { + const uint8_t* record; + const InnerIdType* neighbors; + const uint8_t* one_bit_code; + const uint8_t* supplement_code; + uint32_t neighbor_count; + uint32_t cluster_id; + }; + + struct CodeView { + const uint8_t* one_bit_code; + const uint8_t* supplement_code; + uint32_t cluster_id; + }; + + HGraphRaBitQFusedDataCell(const GraphDataCellParamPtr& graph_param, + uint64_t one_bit_code_size, + uint64_t supplement_code_size, + const IndexCommonParam& common_param); + + void + InsertNeighborsById(InnerIdType id, const Vector& neighbor_ids) override; + + void + DeleteNeighborsById(InnerIdType id) override; + + void + RecoverDeleteNeighborsById(InnerIdType id) override; + + [[nodiscard]] uint32_t + GetNeighborSize(InnerIdType id) const override; + + void + GetNeighbors(InnerIdType id, Vector& neighbor_ids) const override; + + [[nodiscard]] bool + CheckIdExists(InnerIdType id) const override; + + void + Resize(InnerIdType new_size) override; + + void + Prefetch(InnerIdType id, uint32_t neighbor_i) override; + + void + Serialize(StreamWriter& writer) override; + + void + Deserialize(StreamReader& reader) override; + + void + Move(InnerIdType from, InnerIdType to) override; + + void + ShrinkToFit(InnerIdType capacity) override; + + [[nodiscard]] uint64_t + GetMemoryUsage() const override; + + DuplicateTrackerPtr + CreateDuplicateTracker() override; + + void + SetNodeCodes(InnerIdType id, + LabelType label, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code); + + [[nodiscard]] bool + GetFusedCodeView(InnerIdType id, RaBitQFusedCodeView& view) const override; + + void + SetFusedCodes(InnerIdType id, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code) override; + + void + PrefetchFusedCodes(InnerIdType id, bool include_supplement) const override; + + void + PrefetchNodeHeader(InnerIdType id) const { + constexpr uint64_t cache_line_size = 64; + const auto* record = GetNodeRecord(id); + for (uint64_t offset = 0; offset < one_bit_offset_; offset += cache_line_size) { + __builtin_prefetch(record + offset, 0, 2); + } + } + + void + PrefetchFusedFilter(InnerIdType id) const { + const auto* record = GetNodeRecord(id); + PrefetchRange(record + cluster_id_offset_, + one_bit_offset_ + one_bit_code_size_ - cluster_id_offset_, + 3); + } + + void + PrefetchFusedSupplement(InnerIdType id) const { + PrefetchRange(GetNodeRecord(id) + supplement_offset_, supplement_code_size_, 2); + } + + [[nodiscard]] uint64_t + FusedOneBitCodeSize() const override { + return one_bit_code_size_; + } + + [[nodiscard]] uint64_t + FusedSupplementCodeSize() const override { + return supplement_code_size_; + } + + void + SetLabel(InnerIdType id, LabelType label); + + bool + SyncNodeCodes(InnerIdType id, + LabelType label, + uint32_t cluster_id, + const RaBitQSplitDataCellInterface& split_codes); + + void + SetCodecModel(std::string codec_model) { + codec_model_ = std::move(codec_model); + } + + [[nodiscard]] const std::string& + CodecModel() const { + return codec_model_; + } + + [[nodiscard]] const uint8_t* + GetNodeRecord(InnerIdType id) const { + return storage_.data() + aligned_offset_ + static_cast(id) * record_size_; + } + + [[nodiscard]] const InnerIdType* + GetNeighborData(const uint8_t* record) const; + + [[nodiscard]] bool + ResolveNeighbor(InnerIdType stored_neighbor, InnerIdType& neighbor) const { + neighbor = stored_neighbor; + if (not support_remove_) { + return neighbor < total_count_; + } + const uint32_t version = neighbor >> id_bit_; + neighbor &= remove_flag_mask_; + return neighbor < total_count_ and NodeVersion(GetNodeRecord(neighbor)) == version; + } + + [[nodiscard]] const uint8_t* + GetOneBitCode(const uint8_t* record) const; + + [[nodiscard]] const uint8_t* + GetSupplementCode(const uint8_t* record) const; + + [[nodiscard]] LabelType + GetLabel(const uint8_t* record) const; + + [[nodiscard]] uint32_t + GetClusterId(const uint8_t* record) const; + + [[nodiscard]] NodeView + GetNodeView(InnerIdType id) const { + const auto* record = + storage_.data() + aligned_offset_ + static_cast(id) * record_size_; + uint32_t neighbor_count = 0; + uint32_t cluster_id = 0; + std::memcpy(&neighbor_count, record, sizeof(neighbor_count)); + std::memcpy(&cluster_id, record + cluster_id_offset_, sizeof(cluster_id)); + return {record, + reinterpret_cast(record + neighbors_offset_), + record + one_bit_offset_, + record + supplement_offset_, + neighbor_count, + cluster_id}; + } + + [[nodiscard]] CodeView + GetCodeView(InnerIdType id) const { + const auto* record = + storage_.data() + aligned_offset_ + static_cast(id) * record_size_; + uint32_t cluster_id = 0; + std::memcpy(&cluster_id, record + cluster_id_offset_, sizeof(cluster_id)); + return {record + one_bit_offset_, record + supplement_offset_, cluster_id}; + } + + [[nodiscard]] uint64_t + RecordSize() const { + return record_size_; + } + + [[nodiscard]] uint64_t + OneBitOffset() const { + return one_bit_offset_; + } + + [[nodiscard]] uint64_t + OneBitCodeSize() const { + return one_bit_code_size_; + } + +private: + static void + PrefetchRange(const uint8_t* begin, uint64_t size, int locality) { + constexpr uintptr_t cache_line_size = 64; + const auto begin_address = reinterpret_cast(begin); + const auto first = begin_address & ~(cache_line_size - 1U); + const auto end = (begin_address + size + cache_line_size - 1U) & ~(cache_line_size - 1U); + for (auto address = first; address < end; address += cache_line_size) { + if (locality == 3) { + __builtin_prefetch(reinterpret_cast(address), 0, 3); + } else { + __builtin_prefetch(reinterpret_cast(address), 0, 2); + } + } + } + + static uint64_t + AlignUp(uint64_t value, uint64_t alignment); + + [[nodiscard]] uint8_t* + MutableNodeRecord(InnerIdType id); + + [[nodiscard]] static uint32_t + NodeVersion(const uint8_t* record); + + static void + SetNodeVersion(uint8_t* record, uint32_t version); + + void + Reallocate(InnerIdType new_capacity); + +private: + Vector storage_; + uint64_t aligned_offset_{0}; + uint64_t record_size_{0}; + uint64_t neighbors_offset_{0}; + uint64_t cluster_id_offset_{0}; + uint64_t label_offset_{0}; + uint64_t one_bit_offset_{0}; + uint64_t supplement_offset_{0}; + uint64_t one_bit_code_size_{0}; + uint64_t supplement_code_size_{0}; + int64_t dim_{0}; + bool support_remove_{false}; + uint32_t remove_flag_bit_{8}; + uint32_t id_bit_{24}; + uint32_t remove_flag_mask_{0x00FFFFFF}; + std::string codec_model_; + mutable std::shared_mutex storage_mutex_; +}; + +DEFINE_POINTER(HGraphRaBitQFusedDataCell); + +} // namespace vsag diff --git a/src/datacell/hgraph_rabitq_fused_datacell_test.cpp b/src/datacell/hgraph_rabitq_fused_datacell_test.cpp new file mode 100644 index 0000000000..0022f1c5d6 --- /dev/null +++ b/src/datacell/hgraph_rabitq_fused_datacell_test.cpp @@ -0,0 +1,261 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hgraph_rabitq_fused_datacell.h" + +#include +#include +#include +#include +#include + +#include "impl/allocator/safe_allocator.h" +#include "io/memory_io/memory_io_parameter.h" +#include "storage/serialization_template_test.h" +#include "unittest.h" + +namespace vsag { + +TEST_CASE("HGraph RaBitQ fused node layout and serialization", "[ut][HGraphRaBitQFusedDataCell]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = 960; + + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->max_degree_ = 32; + graph_param->init_max_capacity_ = 8; + graph_param->support_remove_ = true; + graph_param->remove_flag_bit_ = 8; + + constexpr uint64_t one_bit_size = 136; + constexpr uint64_t supplement_size = 848; + auto graph = std::make_shared( + graph_param, one_bit_size, supplement_size, common_param); + + REQUIRE(reinterpret_cast(graph->GetNodeRecord(0)) % 64 == 0); + REQUIRE(graph->RecordSize() % 64 == 0); + REQUIRE(graph->OneBitOffset() == 152); + + Vector one_bit(one_bit_size, 0x5A, allocator.get()); + Vector supplement(supplement_size, 0xA5, allocator.get()); + graph->SetNodeCodes(0, 42, 7, one_bit.data(), supplement.data()); + Vector neighbors({1, 2, 3}, allocator.get()); + Vector empty_neighbors(allocator.get()); + graph->InsertNeighborsById(0, neighbors); + graph->InsertNeighborsById(1, empty_neighbors); + graph->InsertNeighborsById(2, empty_neighbors); + graph->InsertNeighborsById(3, empty_neighbors); + + const auto* record = graph->GetNodeRecord(0); + REQUIRE(graph->GetLabel(record) == 42); + REQUIRE(graph->GetClusterId(record) == 7); + REQUIRE(std::memcmp(graph->GetOneBitCode(record), one_bit.data(), one_bit_size) == 0); + REQUIRE(std::memcmp(graph->GetSupplementCode(record), supplement.data(), supplement_size) == 0); + + auto restored = std::make_shared( + graph_param, one_bit_size, supplement_size, common_param); + test_serializion(*graph, *restored); + REQUIRE(restored->CodecModel().empty()); + const auto* restored_record = restored->GetNodeRecord(0); + REQUIRE(restored->GetLabel(restored_record) == 42); + REQUIRE(restored->GetClusterId(restored_record) == 7); + REQUIRE(std::memcmp(restored->GetOneBitCode(restored_record), one_bit.data(), one_bit_size) == + 0); + REQUIRE(std::memcmp(restored->GetSupplementCode(restored_record), + supplement.data(), + supplement_size) == 0); + + graph->Move(0, 4); + const auto* moved_record = graph->GetNodeRecord(4); + REQUIRE(graph->GetLabel(moved_record) == 42); + REQUIRE(graph->GetClusterId(moved_record) == 7); +} + +TEST_CASE("HGraph RaBitQ fused delete version invalidates old edges", + "[ut][HGraphRaBitQFusedDataCell]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + IndexCommonParam common_param; + common_param.allocator_ = allocator; + + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->max_degree_ = 4; + graph_param->init_max_capacity_ = 4; + graph_param->support_remove_ = true; + graph_param->remove_flag_bit_ = 8; + + auto graph = std::make_shared(graph_param, 16, 16, common_param); + Vector one_neighbor({1}, allocator.get()); + Vector empty_neighbors(allocator.get()); + graph->InsertNeighborsById(0, one_neighbor); + graph->InsertNeighborsById(1, empty_neighbors); + + Vector neighbors(allocator.get()); + graph->GetNeighbors(0, neighbors); + REQUIRE(neighbors == Vector({1}, allocator.get())); + graph->DeleteNeighborsById(1); + graph->GetNeighbors(0, neighbors); + REQUIRE(neighbors.empty()); + graph->RecoverDeleteNeighborsById(1); + graph->GetNeighbors(0, neighbors); + REQUIRE(neighbors == Vector({1}, allocator.get())); +} + +TEST_CASE("HGraph RaBitQ fused deserialize validates its wire layout", + "[ut][HGraphRaBitQFusedDataCell]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = 8; + + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->max_degree_ = 32; + graph_param->init_max_capacity_ = 8; + graph_param->support_remove_ = true; + graph_param->remove_flag_bit_ = 8; + + constexpr uint64_t one_bit_size = 16; + constexpr uint64_t supplement_size = 16; + auto graph = std::make_shared( + graph_param, one_bit_size, supplement_size, common_param); + const uint64_t expected_codec_model_size = + 16 + 16 * static_cast(common_param.dim_) * sizeof(float); + graph->SetCodecModel(std::string(expected_codec_model_size, '\0')); + Vector empty_neighbors(allocator.get()); + graph->InsertNeighborsById(0, empty_neighbors); + std::stringstream stream; + IOStreamWriter writer(stream); + graph->Serialize(writer); + const auto payload = stream.str(); + + uint64_t cursor = sizeof(std::atomic); + const uint64_t capacity_offset = cursor; + cursor += sizeof(InnerIdType); + cursor += sizeof(uint32_t); // maximum degree + const uint64_t version_offset = cursor; + cursor += sizeof(uint32_t); + const uint64_t record_size_offset = cursor; + cursor += sizeof(uint64_t); + cursor += sizeof(uint64_t); // neighbors offset + cursor += sizeof(uint64_t); // cluster id offset + cursor += sizeof(uint64_t); // label offset + cursor += sizeof(uint64_t); // one-bit offset + const uint64_t supplement_offset_offset = cursor; + cursor += sizeof(uint64_t); + const uint64_t one_bit_size_offset = cursor; + cursor += sizeof(uint64_t); + cursor += sizeof(uint64_t); // supplement code size + const uint64_t support_remove_offset = cursor; + cursor += sizeof(bool); + const uint64_t remove_flag_bit_offset = cursor; + cursor += sizeof(uint32_t); + const uint64_t codec_model_size_offset = cursor; + + uint64_t codec_model_size = 0; + std::memcpy( + &codec_model_size, payload.data() + codec_model_size_offset, sizeof(codec_model_size)); + const uint64_t payload_size_offset = + codec_model_size_offset + sizeof(codec_model_size) + codec_model_size; + + auto overwrite = [](std::string value, uint64_t offset, const auto& replacement) { + std::memcpy(value.data() + offset, &replacement, sizeof(replacement)); + return value; + }; + auto require_rejected = [&](const std::string& value) { + auto restored = std::make_shared( + graph_param, one_bit_size, supplement_size, common_param); + std::stringstream malformed_stream(value); + IOStreamReader reader(malformed_stream); + REQUIRE_THROWS(restored->Deserialize(reader)); + }; + + SECTION("serialization version") { + require_rejected(overwrite(payload, version_offset, uint32_t{2})); + } + + SECTION("remove flag representation") { + require_rejected(overwrite(payload, support_remove_offset, uint8_t{2})); + } + + SECTION("remove flag bit count") { + require_rejected(overwrite( + payload, remove_flag_bit_offset, static_cast(sizeof(InnerIdType) * 8))); + } + + SECTION("non-monotonic code offsets") { + require_rejected(overwrite(payload, supplement_offset_offset, uint64_t{0})); + } + + SECTION("cache-line record stride") { + require_rejected(overwrite(payload, record_size_offset, graph->RecordSize() + 1)); + } + + SECTION("code size") { + require_rejected(overwrite(payload, one_bit_size_offset, one_bit_size + 1)); + } + + SECTION("capacity times stride overflow") { + constexpr uint64_t largest_aligned_stride = std::numeric_limits::max() - 63; + require_rejected(overwrite(payload, record_size_offset, largest_aligned_stride)); + } + + SECTION("payload byte count") { + require_rejected(overwrite(payload, payload_size_offset, uint64_t{0})); + } + + SECTION("codec model length is bounded before allocation") { + require_rejected( + overwrite(payload, codec_model_size_offset, std::numeric_limits::max())); + } + + SECTION("node payload is bounded before allocation") { + constexpr InnerIdType large_capacity = (InnerIdType{1} << 24U) - 1U; + auto malformed = overwrite(payload, capacity_offset, large_capacity); + const uint64_t declared_bytes = static_cast(large_capacity) * graph->RecordSize(); + malformed = overwrite(malformed, payload_size_offset, declared_bytes); + require_rejected(malformed); + } + + SECTION("count exceeds capacity") { + require_rejected(overwrite(payload, capacity_offset, InnerIdType{0})); + } + + SECTION("constructor maximum degree") { + graph_param->max_degree_ = 31; + require_rejected(payload); + } + + SECTION("constructor rejects invalid removal bit count") { + graph_param->remove_flag_bit_ = sizeof(InnerIdType) * 8; + REQUIRE_THROWS(std::make_shared( + graph_param, one_bit_size, supplement_size, common_param)); + } + + SECTION("constructor requires version bits when removal is enabled") { + graph_param->remove_flag_bit_ = 0; + REQUIRE_THROWS(std::make_shared( + graph_param, one_bit_size, supplement_size, common_param)); + } + + SECTION("constructor bounds capacity by removal id bits") { + graph_param->init_max_capacity_ = InnerIdType{1} << 24U; + REQUIRE_THROWS(std::make_shared( + graph_param, one_bit_size, supplement_size, common_param)); + } +} + +} // namespace vsag diff --git a/src/datacell/rabitq_fused_code_storage.h b/src/datacell/rabitq_fused_code_storage.h new file mode 100644 index 0000000000..c6c9955840 --- /dev/null +++ b/src/datacell/rabitq_fused_code_storage.h @@ -0,0 +1,59 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include + +#include "basic_types.h" + +namespace vsag { + +struct RaBitQFusedCodeView { + const uint8_t* one_bit_code{nullptr}; + const uint8_t* supplement_code{nullptr}; + uint32_t cluster_id{0}; +}; + +/** + * Internal non-owning bridge between the RaBitQ model/query processor and a fused node slab. + * + * Fused codes use cluster-residual semantics. The legacy split 1+7 format stores HNSW-compatible + * BinData/ExData records; the x=1..4 native formats retain the ordinary split bit-plane encoding. + * Callers must retain the cluster id and use the fused distance methods selected by the codec. + */ +class RaBitQFusedCodeStorageInterface { +public: + virtual ~RaBitQFusedCodeStorageInterface() = default; + + [[nodiscard]] virtual bool + GetFusedCodeView(InnerIdType id, RaBitQFusedCodeView& view) const = 0; + + virtual void + SetFusedCodes(InnerIdType id, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code) = 0; + + virtual void + PrefetchFusedCodes(InnerIdType id, bool include_supplement) const = 0; + + [[nodiscard]] virtual uint64_t + FusedOneBitCodeSize() const = 0; + + [[nodiscard]] virtual uint64_t + FusedSupplementCodeSize() const = 0; +}; + +} // namespace vsag diff --git a/src/datacell/rabitq_split_datacell.h b/src/datacell/rabitq_split_datacell.h index 283d0cf837..ebdea828dd 100644 --- a/src/datacell/rabitq_split_datacell.h +++ b/src/datacell/rabitq_split_datacell.h @@ -29,6 +29,7 @@ #include "common.h" #include "flatten_interface.h" #include "flatten_optimized_build_interface.h" +#include "impl/cluster/kmeans_cluster.h" #include "impl/thread_pool/safe_thread_pool.h" #include "inner_string_params.h" #include "io/async_io/async_io_parameter.h" @@ -40,6 +41,7 @@ #include "io/mmap_io/mmap_io_parameter.h" #include "quantization/rabitq_quantization/rabitq_quantizer.h" #include "query_context.h" +#include "rabitq_fused_code_storage.h" #include "storage/stream_reader.h" #include "storage/stream_writer.h" #include "type_helpers.h" @@ -50,6 +52,129 @@ namespace vsag { class MMapIO; +struct RaBitQFusedTraversalQuery { + const uint8_t* query_planes{nullptr}; + const float* transformed_query{nullptr}; + const float* cluster_g_add{nullptr}; + const float* cluster_g_error{nullptr}; + uint64_t dim{0}; + uint64_t one_bit_metadata_offset{0}; + uint64_t supplement_metadata_offset{0}; + uint32_t cluster_count{0}; + uint32_t filter_bits{0}; + uint32_t supplement_bits{0}; + float query_delta{0.0F}; + float query_vl{0.0F}; + float query_sum{0.0F}; + float default_rabitq_error_rate{0.0F}; + bool affine{false}; + bool filter_inner_product_is_exact{false}; +}; + +class RaBitQSplitDataCellInterface { +public: + virtual ~RaBitQSplitDataCellInterface() = default; + + [[nodiscard]] virtual uint64_t + OneBitCodeSize() const = 0; + + [[nodiscard]] virtual uint64_t + SupplementCodeSize() const = 0; + + [[nodiscard]] virtual uint32_t + FusedFilterBits() const = 0; + + [[nodiscard]] virtual uint32_t + FusedSupplementBits() const = 0; + + [[nodiscard]] virtual bool + UsesLegacyHnswFusedCodec() const = 0; + + virtual void + AttachFusedCodeStorage(RaBitQFusedCodeStorageInterface* storage) = 0; + + [[nodiscard]] virtual bool + UsesExternalFusedCodeStorage() const = 0; + + virtual bool + CopySplitCodes(InnerIdType id, uint8_t* one_bit_code, uint8_t* supplement_code) const = 0; + + virtual bool + DecodeFusedById(InnerIdType id, float* data) const = 0; + + virtual bool + ComputeOneBitWithFilterIP(const ComputerInterfacePtr& computer, + const uint8_t* one_bit_code, + float* distance, + float* lower_bound, + float* filter_inner_product, + QueryContext* ctx) const = 0; + + virtual bool + ComputeFullWithFilterIP(const ComputerInterfacePtr& computer, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* distance, + QueryContext* ctx) const = 0; + + virtual bool + ComputeFull(const ComputerInterfacePtr& computer, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + QueryContext* ctx) const = 0; + + virtual void + TrainFusedCodec(const float* data, uint64_t count, uint32_t cluster_count) = 0; + + virtual bool + EncodeFused(const float* data, + uint8_t* one_bit_code, + uint8_t* supplement_code, + uint32_t* cluster_id) const = 0; + + virtual ComputerInterfacePtr + FactoryFusedComputer(const void* query) const = 0; + + virtual bool + GetFusedTraversalQuery(const ComputerInterfacePtr& computer, + RaBitQFusedTraversalQuery* query) const = 0; + + virtual bool + ComputeFusedOneBitWithFilterIP(const ComputerInterfacePtr& computer, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + float* lower_bound, + float* filter_inner_product, + QueryContext* ctx) const = 0; + + virtual bool + ComputeFusedFullWithFilterIP(const ComputerInterfacePtr& computer, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* distance, + QueryContext* ctx) const = 0; + + virtual bool + ComputeFusedFull(const ComputerInterfacePtr& computer, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + QueryContext* ctx) const = 0; + + [[nodiscard]] virtual std::string + ExportFusedCodec() const = 0; + + virtual void + ImportFusedCodec(const std::string& serialized) = 0; +}; + template class RaBitQSplitCodeStorage { public: @@ -127,6 +252,16 @@ class RaBitQSplitCodeStorage { io_->Deserialize(reader); } + void + SkipSerialized(StreamReader& reader) { + uint64_t size = 0; + StreamReader::ReadObj(reader, size); + const uint64_t cursor = reader.GetCursor(); + CHECK_ARGUMENT(cursor <= reader.Length() and size <= reader.Length() - cursor, + "serialized RaBitQ split code payload exceeds its stream boundary"); + reader.Seek(cursor + size); + } + [[nodiscard]] uint64_t GetMemoryUsage() const { if constexpr (IOTmpl::InMemory) { @@ -141,7 +276,9 @@ class RaBitQSplitCodeStorage { }; template -class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuildInterface { +class RaBitQSplitDataCell : public FlattenInterface, + public FlattenOptimizedBuildInterface, + public RaBitQSplitDataCellInterface { public: class OptimizedBuildComputer final : public ComputerInterface { public: @@ -153,6 +290,26 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil uint64_t code_sum_{0}; }; + class FusedComputer final : public ComputerInterface { + public: + explicit FusedComputer(Allocator* allocator) + : transformed_query_(allocator), + hnsw_query_planes_(allocator), + hnsw_g_add_(allocator), + hnsw_g_error_(allocator) { + } + + Vector transformed_query_; + Vector hnsw_query_planes_; + float hnsw_query_delta_{0.0F}; + float hnsw_query_vl_{0.0F}; + float hnsw_query_sum_{0.0F}; + Vector hnsw_g_add_; + Vector hnsw_g_error_; + float query_raw_norm_{0.0F}; + typename RaBitQuantizer::norm_type mrq_norm_sqr_{0.0F}; + }; + RaBitQSplitDataCell() = default; explicit RaBitQSplitDataCell(const QuantizerParamPtr& quantization_param, @@ -166,6 +323,8 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil const IOParamPtr& supplement_io_param, const IndexCommonParam& common_param) : common_param_(common_param), allocator_(common_param.allocator_.get()) { + this->quantization_param_ = + std::dynamic_pointer_cast(quantization_param); this->quantizer_ = std::make_shared>(quantization_param, common_param); if (not this->quantizer_->SupportSplitCodeStorage()) { @@ -200,6 +359,10 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil const InnerIdType* idx, InnerIdType id_count, QueryContext* ctx = nullptr) override { + if (fused_code_storage_ != nullptr) { + this->query_fused_full(result_dists, computer, idx, id_count, ctx); + return; + } if (this->optimized_build_active_) { this->query_optimized_build_codes(result_dists, computer, idx, id_count); return; @@ -236,6 +399,10 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil const InnerIdType* idx, InnerIdType id_count, QueryContext* ctx = nullptr) override { + if (fused_code_storage_ != nullptr) { + this->query_fused_full(result_dists, computer, idx, id_count, ctx); + return; + } if (this->optimized_build_active_) { this->query_optimized_build_codes(result_dists, computer, idx, id_count); return; @@ -275,6 +442,11 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil InnerIdType id_count, float threshold, QueryContext* ctx = nullptr) override { + if (fused_code_storage_ != nullptr) { + this->query_fused_with_distance_filter( + result_dists, computer, idx, id_count, threshold, ctx); + return; + } if (this->optimized_build_active_) { this->query_optimized_build_codes(result_dists, computer, idx, id_count); return; @@ -306,7 +478,7 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil throw; } - if (computed and std::isfinite(lower_bound) and lower_bound >= threshold) { + if (computed and IsFiniteRaBitQValue(lower_bound) and lower_bound >= threshold) { this->release_one_bit_code(one_bit_code, one_bit_need_release); result_dists[i] = threshold; continue; @@ -327,6 +499,59 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil } } + void + QueryWithDistanceLowerBoundAndFilterIP(float* result_dists, + float* lower_bounds, + float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* idx, + InnerIdType id_count, + QueryContext* ctx = nullptr) override { + if (fused_code_storage_ != nullptr) { + this->query_fused_lower_bound( + result_dists, lower_bounds, filter_inner_products, computer, idx, id_count, ctx); + return; + } + auto* comp = static_cast>*>(computer.get()); + this->add_filter_count(ctx, id_count); + for (uint32_t i = 0; i < this->prefetch_stride_code_ and i < id_count; ++i) { + this->prefetch_one_bit(idx[i]); + } + for (InnerIdType i = 0; i < id_count; ++i) { + if (i + this->prefetch_stride_code_ < id_count) { + this->prefetch_one_bit(idx[i + this->prefetch_stride_code_]); + } + bool need_release = false; + const auto* one_bit_code = this->get_one_bit_code(idx[i], need_release); + bool computed = false; + try { + computed = this->quantizer_->ComputeDistWithOneBitLowerBoundAndFilterIP( + *comp, + one_bit_code, + result_dists + i, + lower_bounds == nullptr ? nullptr : lower_bounds + i, + filter_inner_products == nullptr ? nullptr : filter_inner_products + i, + this->query_rabitq_error_rate(ctx)); + } catch (...) { + this->release_one_bit_code(one_bit_code, need_release); + throw; + } + if (not computed) { + if (filter_inner_products != nullptr) { + filter_inner_products[i] = std::numeric_limits::quiet_NaN(); + } + this->compute_full_dist_after_one_bit_failure( + idx[i], + one_bit_code, + comp, + result_dists + i, + lower_bounds == nullptr ? nullptr : lower_bounds + i, + ctx); + } + this->release_one_bit_code(one_bit_code, need_release); + } + } + void QueryWithDistanceLowerBound(float* result_dists, float* lower_bounds, @@ -334,6 +559,11 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil const InnerIdType* idx, InnerIdType id_count, QueryContext* ctx = nullptr) override { + if (fused_code_storage_ != nullptr) { + this->query_fused_lower_bound( + result_dists, lower_bounds, nullptr, computer, idx, id_count, ctx); + return; + } if (this->optimized_build_active_) { this->query_optimized_build_codes(result_dists, computer, idx, id_count); if (lower_bounds != nullptr) { @@ -460,8 +690,41 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil } } + void + QueryWithFilterIPHint(float* result_dists, + const float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* idx, + InnerIdType id_count, + QueryContext* ctx = nullptr) override { + if (fused_code_storage_ != nullptr) { + this->query_fused_full_with_filter_ip( + result_dists, filter_inner_products, computer, idx, id_count, ctx); + return; + } + auto* comp = static_cast>*>(computer.get()); + for (uint32_t i = 0; i < this->prefetch_stride_code_ and i < id_count; ++i) { + this->prefetch_full_code(idx[i]); + } + for (InnerIdType i = 0; i < id_count; ++i) { + if (i + this->prefetch_stride_code_ < id_count) { + this->prefetch_full_code(idx[i + this->prefetch_stride_code_]); + } + this->compute_full_dist_with_filter_ip(idx[i], + comp, + result_dists + i, + ctx, + filter_inner_products == nullptr + ? std::numeric_limits::quiet_NaN() + : filter_inner_products[i]); + } + } + ComputerInterfacePtr FactoryComputer(const void* query) override { + if (fused_code_storage_ != nullptr) { + return this->FactoryFusedComputer(query); + } auto computer = this->quantizer_->FactoryComputer(); computer->SetQuery(static_cast(query)); return computer; @@ -491,6 +754,9 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil bool BeginOptimizedBuild(const FlattenOptimizedBuildContext& context) override { + if (this->fused_code_storage_ != nullptr) { + return false; + } if (this->optimized_build_active_ or not this->quantizer_->SupportScalarCodeBuild()) { return false; } @@ -640,6 +906,9 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil "optimized RaBitQ build storage must be resized before inserting vectors"); this->total_count_ = std::max(this->total_count_, idx + 1); } + if (this->fused_code_storage_ != nullptr) { + return; + } this->write_encoded_vector(static_cast(vector), idx); } @@ -649,6 +918,9 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil if (idx >= this->total_count_) { return false; } + if (this->fused_code_storage_ != nullptr) { + return true; + } std::lock_guard lock(this->mutex_); this->write_encoded_vector(static_cast(vector), idx); return true; @@ -665,6 +937,11 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil float ComputePairVectors(InnerIdType id1, InnerIdType id2) override { + if (this->fused_code_storage_ != nullptr) { + throw VsagException( + ErrorType::UNSUPPORTED_INDEX_OPERATION, + "pairwise distance is unavailable for cluster-residual fused RaBitQ codes"); + } if (this->optimized_build_active_) { bool release1 = false; bool release2 = false; @@ -716,6 +993,10 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil if (new_capacity <= this->max_capacity_) { return; } + if (this->fused_code_storage_ != nullptr) { + this->max_capacity_ = new_capacity; + return; + } this->x_bit_cell_->Resize(new_capacity); this->supplement_cell_->Resize(new_capacity); if (this->optimized_build_active_) { @@ -727,6 +1008,10 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void Prefetch(InnerIdType id) override { + if (this->fused_code_storage_ != nullptr) { + this->fused_code_storage_->PrefetchFusedCodes(id, false); + return; + } if (this->optimized_build_active_) { this->optimized_build_scalar_codes_->Prefetch(id, this->optimized_build_record_size_); return; @@ -754,6 +1039,11 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void InitIO(const IOParamPtr& io_param) override { + if (this->fused_code_storage_ != nullptr) { + CHECK_ARGUMENT(OneBitIOTmpl::InMemory and SupplementIOTmpl::InMemory, + "fused RaBitQ code storage requires memory IO"); + return; + } const bool shares_io_param = this->supplement_io_type_.empty(); this->x_bit_cell_->InitIO(SuffixIOParam(io_param, "_onebit", shares_io_param)); // In hybrid mode (one-bit and supplement use different IO backends) @@ -766,6 +1056,11 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void InitIO(const IOParamPtr& one_bit_io_param, const IOParamPtr& supplement_io_param) { + if (this->fused_code_storage_ != nullptr) { + CHECK_ARGUMENT(OneBitIOTmpl::InMemory and SupplementIOTmpl::InMemory, + "fused RaBitQ code storage requires memory IO"); + return; + } const bool shares_io_param = supplement_io_param == nullptr; this->x_bit_cell_->InitIO(SuffixIOParam(one_bit_io_param, "_onebit", shares_io_param)); if (supplement_io_param != nullptr) { @@ -799,6 +1094,499 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil return this->quantizer_->Metric(); } + [[nodiscard]] uint64_t + OneBitCodeSize() const override { + return one_bit_code_size_; + } + + [[nodiscard]] uint64_t + SupplementCodeSize() const override { + return supplement_code_size_; + } + + [[nodiscard]] uint32_t + FusedFilterBits() const override { + return quantizer_->FilterBits(); + } + + [[nodiscard]] uint32_t + FusedSupplementBits() const override { + return quantizer_->ReorderBits(); + } + + [[nodiscard]] bool + UsesLegacyHnswFusedCodec() const override { + return IsLegacyHnswFusedCodec(); + } + + void + AttachFusedCodeStorage(RaBitQFusedCodeStorageInterface* storage) override { + CHECK_ARGUMENT(storage != nullptr, "fused RaBitQ code storage must not be null"); + CHECK_ARGUMENT(not optimized_build_active_, + "cannot attach fused RaBitQ storage during optimized build"); + CHECK_ARGUMENT(storage->FusedOneBitCodeSize() == one_bit_code_size_ and + storage->FusedSupplementCodeSize() == supplement_code_size_, + "fused RaBitQ code sizes do not match the split model"); + fused_code_storage_ = storage; + x_bit_cell_->Shrink(0); + supplement_cell_->Shrink(0); + } + + [[nodiscard]] bool + UsesExternalFusedCodeStorage() const override { + return fused_code_storage_ != nullptr; + } + + bool + CopySplitCodes(InnerIdType id, uint8_t* one_bit_code, uint8_t* supplement_code) const override { + if (this->optimized_build_active_ or one_bit_code == nullptr or + supplement_code == nullptr) { + return false; + } + if (fused_code_storage_ != nullptr) { + RaBitQFusedCodeView view; + if (not fused_code_storage_->GetFusedCodeView(id, view)) { + return false; + } + std::memcpy(one_bit_code, view.one_bit_code, one_bit_code_size_); + std::memcpy(supplement_code, view.supplement_code, supplement_code_size_); + return true; + } + return this->x_bit_cell_->Read(id, one_bit_code) and + this->supplement_cell_->Read(id, supplement_code); + } + + bool + DecodeFusedById(InnerIdType id, float* data) const override { + if (data == nullptr or fused_code_storage_ == nullptr or id >= this->TotalCount() or + fused_quantizers_.empty()) { + return false; + } + RaBitQFusedCodeView view; + if (not fused_code_storage_->GetFusedCodeView(id, view) or + view.cluster_id >= fused_quantizers_.size()) { + return false; + } + return fused_quantizers_[view.cluster_id]->DecodeFusedSplitCode( + view.one_bit_code, view.supplement_code, IsLegacyHnswFusedCodec(), data); + } + + bool + ComputeOneBitWithFilterIP(const ComputerInterfacePtr& computer, + const uint8_t* one_bit_code, + float* distance, + float* lower_bound, + float* filter_inner_product, + QueryContext* ctx) const override { + auto* comp = static_cast>*>(computer.get()); + this->add_filter_count(ctx, 1); + const bool computed = this->quantizer_->ComputeDistWithOneBitLowerBoundAndFilterIP( + *comp, + one_bit_code, + distance, + lower_bound, + filter_inner_product, + this->query_rabitq_error_rate(ctx)); + if (not computed) { + this->add_filter_fallback_full_count(ctx, 1); + } + return computed; + } + + bool + ComputeFullWithFilterIP(const ComputerInterfacePtr& computer, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* distance, + QueryContext* ctx) const override { + if (this->quantizer_->FilterBits() < 2) { + return false; + } + auto* comp = static_cast>*>(computer.get()); + this->add_full_count(ctx, 1); + const bool computed = this->quantizer_->ComputeDistWithSplitCodeAndFilterIP( + *comp, one_bit_code, supplement_code, filter_inner_product, distance); + if (computed) { + this->add_reorder_hint_full_count(ctx, 1); + } else { + this->add_reorder_fallback_full_count(ctx, 1); + } + return computed; + } + + bool + ComputeFull(const ComputerInterfacePtr& computer, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + QueryContext* ctx) const override { + auto* comp = static_cast>*>(computer.get()); + this->add_full_count(ctx, 1); + return this->quantizer_->ComputeDistWithSplitCode( + *comp, one_bit_code, supplement_code, distance); + } + + void + TrainFusedCodec(const float* data, uint64_t count, uint32_t cluster_count) override { + CHECK_ARGUMENT(data != nullptr and count > 0, + "fused RaBitQ training data must not be empty"); + CHECK_ARGUMENT(cluster_count == 16, "fused RaBitQ requires exactly 16 clusters"); + CHECK_ARGUMENT(this->quantization_param_ != nullptr, + "fused RaBitQ quantizer parameter is unavailable"); + CHECK_ARGUMENT(this->AreFusedVectorsFinite(data, count), + "fused RaBitQ training data must contain only finite values"); + + KMeansCluster kmeans( + static_cast(common_param_.dim_), allocator_, common_param_.thread_pool_); + const auto trained_cluster_count = + static_cast(std::min(cluster_count, count)); + kmeans.Run(trained_cluster_count, + data, + count, + 25, + nullptr, + false, + 1e-6F, + KMeansInitMethod::KMEANS_PLUS_PLUS, + 0x52425131U, + true); + fused_centroids_.resize(static_cast(cluster_count) * common_param_.dim_); + for (uint32_t cluster_id = 0; cluster_id < cluster_count; ++cluster_id) { + const auto source_cluster = cluster_id % trained_cluster_count; + std::copy_n( + kmeans.k_centroids_ + static_cast(source_cluster) * common_param_.dim_, + common_param_.dim_, + fused_centroids_.data() + static_cast(cluster_id) * common_param_.dim_); + } + + std::stringstream model_stream; + IOStreamWriter model_writer(model_stream); + quantizer_->Serialize(model_writer); + const auto serialized_model = model_stream.str(); + + fused_quantizers_.clear(); + fused_quantizers_.reserve(cluster_count); + for (uint32_t cluster_id = 0; cluster_id < cluster_count; ++cluster_id) { + auto quantizer = + std::make_shared>(quantization_param_, common_param_); + std::stringstream input(serialized_model); + IOStreamReader model_reader(input); + quantizer->Deserialize(model_reader); + quantizer->SetCentroid(fused_centroids_.data() + + static_cast(cluster_id) * common_param_.dim_); + fused_quantizers_.push_back(std::move(quantizer)); + } + } + + bool + EncodeFused(const float* data, + uint8_t* one_bit_code, + uint8_t* supplement_code, + uint32_t* cluster_id) const override { + if (data == nullptr or one_bit_code == nullptr or supplement_code == nullptr or + cluster_id == nullptr or fused_quantizers_.empty()) { + return false; + } + if (not this->AreFusedVectorsFinite(data, 1)) { + return false; + } + *cluster_id = NearestFusedCluster(data); + ByteBuffer full_code(code_size_, allocator_); + auto& quantizer = fused_quantizers_[*cluster_id]; + if (not quantizer->EncodeOne(data, full_code.data)) { + return false; + } + quantizer->SplitCode(full_code.data, one_bit_code, supplement_code); + if (IsLegacyHnswFusedCodec()) { + quantizer->EncodeHnswOneBitMetadata(data, one_bit_code); + if (not quantizer->EncodeHnswSupplement(data, supplement_code)) { + return false; + } + } else if (not quantizer->EncodeFusedAffineMetadata(data, one_bit_code, supplement_code)) { + return false; + } + return true; + } + + ComputerInterfacePtr + FactoryFusedComputer(const void* query) const override { + auto result = std::make_shared(allocator_); + if (fused_quantizers_.empty()) { + return result; + } + result->transformed_query_.resize(fused_quantizers_.front()->GetDim()); + fused_quantizers_.front()->TransformFusedQuery(static_cast(query), + result->transformed_query_, + result->query_raw_norm_, + result->mrq_norm_sqr_); + if (quantizer_->FilterBits() == 1) { + fused_quantizers_.front()->PrepareHnswFourBitQuery(result->transformed_query_.data(), + result->hnsw_query_planes_, + result->hnsw_query_delta_, + result->hnsw_query_vl_, + result->hnsw_query_sum_); + } else { + double query_sum = 0.0; + for (const float value : result->transformed_query_) { + query_sum += static_cast(value); + } + result->hnsw_query_sum_ = static_cast(query_sum); + } + result->hnsw_g_add_.resize(fused_quantizers_.size()); + result->hnsw_g_error_.resize(fused_quantizers_.size()); + for (uint64_t cluster_id = 0; cluster_id < fused_quantizers_.size(); ++cluster_id) { + fused_quantizers_[cluster_id]->ComputeHnswCentroidTerms( + result->transformed_query_.data(), + result->hnsw_g_add_[cluster_id], + result->hnsw_g_error_[cluster_id]); + } + return result; + } + + bool + GetFusedTraversalQuery(const ComputerInterfacePtr& computer, + RaBitQFusedTraversalQuery* query) const override { + if (query == nullptr) { + return false; + } + *query = {}; + auto* fused_computer = static_cast(computer.get()); + if (fused_computer == nullptr or fused_quantizers_.empty()) { + return false; + } + const auto filter_bits = quantizer_->FilterBits(); + query->query_planes = + filter_bits == 1 ? fused_computer->hnsw_query_planes_.data() : nullptr; + query->transformed_query = fused_computer->transformed_query_.data(); + query->cluster_g_add = fused_computer->hnsw_g_add_.data(); + query->cluster_g_error = fused_computer->hnsw_g_error_.data(); + query->dim = common_param_.dim_; + query->one_bit_metadata_offset = IsLegacyHnswFusedCodec() + ? quantizer_->PlaneBytes() + : quantizer_->OneBitRecordNormOffset(); + query->supplement_metadata_offset = quantizer_->SupplementMetaOffset(); + query->cluster_count = static_cast(fused_quantizers_.size()); + query->filter_bits = filter_bits; + query->supplement_bits = quantizer_->ReorderBits(); + query->query_delta = fused_computer->hnsw_query_delta_; + query->query_vl = fused_computer->hnsw_query_vl_; + query->query_sum = fused_computer->hnsw_query_sum_; + query->default_rabitq_error_rate = quantizer_->DefaultRaBitQErrorRate(); + query->affine = not IsLegacyHnswFusedCodec(); + query->filter_inner_product_is_exact = query->affine and filter_bits >= 2; + return query->transformed_query != nullptr and query->cluster_g_add != nullptr and + query->cluster_g_error != nullptr and + (filter_bits != 1 or query->query_planes != nullptr); + } + + bool + ComputeFusedOneBitWithFilterIP(const ComputerInterfacePtr& computer, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + float* lower_bound, + float* filter_inner_product, + QueryContext* ctx) const override { + auto* fused_computer = static_cast(computer.get()); + if (fused_computer == nullptr or cluster_id >= fused_quantizers_.size()) { + return false; + } + this->add_filter_count(ctx, 1); + if (filter_inner_product != nullptr) { + *filter_inner_product = std::numeric_limits::quiet_NaN(); + } + float local_filter_inner_product = std::numeric_limits::quiet_NaN(); + bool computed = false; + if (IsLegacyHnswFusedCodec()) { + computed = fused_quantizers_[cluster_id]->ComputeHnswOneBit( + fused_computer->hnsw_query_planes_.data(), + fused_computer->hnsw_query_delta_, + fused_computer->hnsw_query_vl_, + fused_computer->hnsw_query_sum_, + fused_computer->hnsw_g_add_[cluster_id], + fused_computer->hnsw_g_error_[cluster_id], + one_bit_code, + supplement_code, + distance, + lower_bound, + &local_filter_inner_product, + this->query_rabitq_error_rate(ctx)); + } else { + RaBitQFusedIPPrecision precision = RaBitQFusedIPPrecision::INVALID; + const auto* query_planes = + quantizer_->FilterBits() == 1 ? fused_computer->hnsw_query_planes_.data() : nullptr; + computed = fused_quantizers_[cluster_id]->ComputeFusedAffineFilter( + fused_computer->transformed_query_.data(), + query_planes, + fused_computer->hnsw_query_delta_, + fused_computer->hnsw_query_vl_, + fused_computer->hnsw_query_sum_, + fused_computer->hnsw_g_add_[cluster_id], + fused_computer->hnsw_g_error_[cluster_id], + one_bit_code, + this->query_rabitq_error_rate(ctx), + distance, + lower_bound, + &local_filter_inner_product, + &precision); + if (computed and quantizer_->FilterBits() >= 2 and + precision == RaBitQFusedIPPrecision::EXACT and filter_inner_product != nullptr) { + *filter_inner_product = local_filter_inner_product; + } + } + if (not computed and (ctx == nullptr or ctx->enable_rabitq_reorder)) { + this->add_filter_fallback_full_count(ctx, 1); + } + return computed; + } + + bool + ComputeFusedFullWithFilterIP(const ComputerInterfacePtr& computer, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* distance, + QueryContext* ctx) const override { + auto* fused_computer = static_cast(computer.get()); + if (fused_computer == nullptr or cluster_id >= fused_quantizers_.size()) { + return false; + } + this->add_full_count(ctx, 1); + bool computed = false; + if (not IsLegacyHnswFusedCodec() and quantizer_->FilterBits() >= 2) { + computed = fused_quantizers_[cluster_id]->ComputeFusedAffineFullWithFilterIP( + fused_computer->transformed_query_.data(), + fused_computer->hnsw_query_sum_, + fused_computer->hnsw_g_add_[cluster_id], + one_bit_code, + supplement_code, + filter_inner_product, + distance); + } + if (computed) { + this->add_reorder_hint_full_count(ctx, 1); + } else { + this->add_reorder_fallback_full_count(ctx, 1); + } + return computed; + } + + bool + ComputeFusedFull(const ComputerInterfacePtr& computer, + uint32_t cluster_id, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + QueryContext* ctx) const override { + auto* fused_computer = static_cast(computer.get()); + if (fused_computer == nullptr or cluster_id >= fused_quantizers_.size()) { + return false; + } + this->add_full_count(ctx, 1); + if (not IsLegacyHnswFusedCodec()) { + return fused_quantizers_[cluster_id]->ComputeFusedAffineFullDirect( + fused_computer->transformed_query_.data(), + fused_computer->hnsw_query_sum_, + fused_computer->hnsw_g_add_[cluster_id], + one_bit_code, + supplement_code, + distance); + } + const float inv_sqrt_dim = 1.0F / std::sqrt(static_cast(common_param_.dim_)); + const float signed_ip = RaBitQFloatBinaryIP(fused_computer->transformed_query_.data(), + one_bit_code, + common_param_.dim_, + inv_sqrt_dim); + const float filter_inner_product = + 0.5F * (signed_ip / inv_sqrt_dim + fused_computer->hnsw_query_sum_); + float lower_bound = 0.0F; + return fused_quantizers_[cluster_id]->ComputeHnswFull( + fused_computer->transformed_query_.data(), + fused_computer->hnsw_query_sum_, + fused_computer->hnsw_g_add_[cluster_id], + fused_computer->hnsw_g_error_[cluster_id], + one_bit_code, + supplement_code, + filter_inner_product, + distance, + &lower_bound); + } + + [[nodiscard]] std::string + ExportFusedCodec() const override { + if (fused_quantizers_.empty()) { + return {}; + } + std::stringstream output; + IOStreamWriter writer(output); + constexpr uint32_t version = 1; + StreamWriter::WriteObj(writer, version); + StreamWriter::WriteVector(writer, fused_centroids_); + const auto cluster_count = static_cast(fused_quantizers_.size()); + StreamWriter::WriteObj(writer, cluster_count); + return output.str(); + } + + void + ImportFusedCodec(const std::string& serialized) override { + CHECK_ARGUMENT(not serialized.empty(), "fused RaBitQ codec payload is empty"); + CHECK_ARGUMENT(this->quantization_param_ != nullptr, + "fused RaBitQ quantizer parameter is unavailable"); + constexpr uint64_t cluster_count = 16; + constexpr uint64_t fixed_payload_size = + sizeof(uint32_t) + sizeof(uint64_t) + sizeof(uint32_t); + constexpr uint64_t bytes_per_dimension = cluster_count * sizeof(float); + CHECK_ARGUMENT(common_param_.dim_ > 0, "invalid fused RaBitQ dimension"); + const auto dim = static_cast(common_param_.dim_); + CHECK_ARGUMENT(dim <= (std::numeric_limits::max() - fixed_payload_size) / + bytes_per_dimension, + "fused RaBitQ codec size overflow"); + const uint64_t expected_centroid_count = cluster_count * dim; + const uint64_t expected_payload_size = + fixed_payload_size + expected_centroid_count * sizeof(float); + CHECK_ARGUMENT(serialized.size() == expected_payload_size, + "invalid fused RaBitQ codec payload size"); + + std::stringstream input(serialized); + IOStreamReader reader(input); + uint32_t version = 0; + StreamReader::ReadObj(reader, version); + CHECK_ARGUMENT(version == 1, "unsupported fused RaBitQ codec version"); + uint64_t centroid_count = 0; + StreamReader::ReadObj(reader, centroid_count); + CHECK_ARGUMENT(centroid_count == expected_centroid_count, + "invalid fused RaBitQ centroid payload"); + fused_centroids_.resize(expected_centroid_count); + reader.Read(reinterpret_cast(fused_centroids_.data()), + expected_centroid_count * sizeof(float)); + uint32_t serialized_cluster_count = 0; + StreamReader::ReadObj(reader, serialized_cluster_count); + CHECK_ARGUMENT(serialized_cluster_count == cluster_count, + "invalid fused RaBitQ cluster count"); + CHECK_ARGUMENT(reader.GetCursor() == reader.Length(), + "trailing fused RaBitQ codec payload"); + + std::stringstream model_output; + IOStreamWriter model_writer(model_output); + quantizer_->Serialize(model_writer); + const auto serialized_model = model_output.str(); + fused_quantizers_.clear(); + fused_quantizers_.reserve(serialized_cluster_count); + for (uint32_t cluster_id = 0; cluster_id < serialized_cluster_count; ++cluster_id) { + auto quantizer = + std::make_shared>(quantization_param_, common_param_); + std::stringstream model_input(serialized_model); + IOStreamReader model_reader(model_input); + quantizer->Deserialize(model_reader); + quantizer->SetCentroid(fused_centroids_.data() + + static_cast(cluster_id) * common_param_.dim_); + fused_quantizers_.push_back(std::move(quantizer)); + } + } + bool Decode(const uint8_t* codes, float* data) override { return this->quantizer_->DecodeOne(codes, data); @@ -809,8 +1597,39 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil return this->quantizer_->EncodeOne(data, codes); } + bool + CompareRawVectorWithId(const void* vector, InnerIdType id) override { + if (this->fused_code_storage_ == nullptr) { + return FlattenInterface::CompareRawVectorWithId(vector, id); + } + if (vector == nullptr) { + return false; + } + ByteBuffer one_bit_code(one_bit_code_size_, allocator_); + ByteBuffer supplement_code(supplement_code_size_, allocator_); + uint32_t cluster_id = 0; + if (not this->EncodeFused(static_cast(vector), + one_bit_code.data, + supplement_code.data, + &cluster_id)) { + return false; + } + RaBitQFusedCodeView stored; + if (not this->fused_code_storage_->GetFusedCodeView(id, stored)) { + return false; + } + return stored.cluster_id == cluster_id and + std::memcmp(stored.one_bit_code, one_bit_code.data, one_bit_code_size_) == 0 and + std::memcmp(stored.supplement_code, supplement_code.data, supplement_code_size_) == + 0; + } + [[nodiscard]] const uint8_t* GetCodesById(InnerIdType id, bool& need_release) const override { + if (this->fused_code_storage_ != nullptr) { + need_release = false; + return nullptr; + } if (this->optimized_build_active_) { auto* codes = static_cast(allocator_->Allocate(this->code_size_)); if (not this->GetCodesById(id, codes)) { @@ -834,6 +1653,9 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil bool GetCodesById(InnerIdType id, uint8_t* codes) const override { + if (this->fused_code_storage_ != nullptr) { + return false; + } if (this->optimized_build_active_) { bool need_release = false; const auto* scalar_code = this->optimized_build_scalar_codes_->Read(id, need_release); @@ -872,6 +1694,14 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil CHECK_ARGUMENT(not this->optimized_build_active_, "cannot serialize RaBitQ split codes during optimized build"); FlattenInterface::Serialize(writer); + if (this->fused_code_storage_ != nullptr) { + StreamWriter::WriteObj(writer, kFusedModelMagic); + StreamWriter::WriteObj(writer, kFusedModelVersion); + StreamWriter::WriteObj(writer, this->one_bit_code_size_); + StreamWriter::WriteObj(writer, this->supplement_code_size_); + this->quantizer_->Serialize(writer); + return; + } StreamWriter::WriteString(writer, this->supplement_io_type_); this->x_bit_cell_->Serialize(writer); this->supplement_cell_->Serialize(writer); @@ -881,6 +1711,43 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void Deserialize(lvalue_or_rvalue reader) override { FlattenInterface::Deserialize(reader); + if (this->fused_code_storage_ != nullptr) { + const uint64_t payload_cursor = reader->GetCursor(); + uint64_t magic = 0; + StreamReader::ReadObj(reader, magic); + if (magic == kFusedModelMagic) { + uint32_t version = 0; + uint64_t serialized_one_bit_size = 0; + uint64_t serialized_supplement_size = 0; + StreamReader::ReadObj(reader, version); + StreamReader::ReadObj(reader, serialized_one_bit_size); + StreamReader::ReadObj(reader, serialized_supplement_size); + CHECK_ARGUMENT(version == kFusedModelVersion, + "unsupported fused RaBitQ model-only serialization version"); + this->quantizer_->Deserialize(reader); + this->refresh_code_sizes(); + CHECK_ARGUMENT(serialized_one_bit_size == this->one_bit_code_size_ and + serialized_supplement_size == this->supplement_code_size_ and + this->fused_code_storage_->FusedOneBitCodeSize() == + this->one_bit_code_size_ and + this->fused_code_storage_->FusedSupplementCodeSize() == + this->supplement_code_size_, + "serialized fused RaBitQ code sizes do not match the node layout"); + return; + } + reader->Seek(payload_cursor); + this->DeserializeSupplementIOType(reader); + this->x_bit_cell_->SkipSerialized(*reader); + this->supplement_cell_->SkipSerialized(*reader); + this->quantizer_->Deserialize(reader); + this->refresh_code_sizes(); + CHECK_ARGUMENT( + this->fused_code_storage_->FusedOneBitCodeSize() == this->one_bit_code_size_ and + this->fused_code_storage_->FusedSupplementCodeSize() == + this->supplement_code_size_, + "legacy split RaBitQ code sizes do not match the fused node layout"); + return; + } this->DeserializeSupplementIOType(reader); this->x_bit_cell_->Deserialize(reader); this->supplement_cell_->Deserialize(reader); @@ -890,6 +1757,10 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void MergeOther(const FlattenInterfacePtr& other, InnerIdType bias) override { + if (this->fused_code_storage_ != nullptr) { + throw VsagException(ErrorType::UNSUPPORTED_INDEX_OPERATION, + "fused RaBitQ code storage does not support MergeOther"); + } auto ptr = std::dynamic_pointer_cast>( other); @@ -912,6 +1783,9 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void Move(InnerIdType from, InnerIdType to) override { + if (this->fused_code_storage_ != nullptr) { + return; + } if (this->optimized_build_active_) { ByteBuffer build_record(this->optimized_build_record_size_, allocator_); this->optimized_build_scalar_codes_->Read(from, build_record.data); @@ -929,6 +1803,10 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil void ShrinkToFit(InnerIdType capacity) override { + if (this->fused_code_storage_ != nullptr) { + this->max_capacity_ = capacity; + return; + } this->x_bit_cell_->Shrink(capacity); this->supplement_cell_->Shrink(capacity); if (this->optimized_build_active_) { @@ -951,12 +1829,17 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil memory += this->optimized_build_code_sums_->capacity() * sizeof(uint64_t); } memory += sizeof(RaBitQuantizer); + memory += fused_quantizers_.size() * sizeof(RaBitQuantizer); + memory += fused_centroids_.capacity() * sizeof(float); return memory; } public: IndexCommonParam common_param_; + RaBitQuantizerParamPtr quantization_param_{nullptr}; std::shared_ptr> quantizer_{nullptr}; + std::vector>> fused_quantizers_; + std::vector fused_centroids_; std::shared_ptr> x_bit_cell_{nullptr}; std::shared_ptr> supplement_cell_{nullptr}; std::shared_ptr> optimized_build_scalar_codes_{nullptr}; @@ -977,8 +1860,50 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil std::string supplement_io_type_{}; bool optimized_build_active_{false}; uint64_t optimized_build_record_size_{0}; + RaBitQFusedCodeStorageInterface* fused_code_storage_{nullptr}; private: + static constexpr uint64_t kFusedModelMagic = 0x524246534D4F444CULL; + static constexpr uint32_t kFusedModelVersion = 1; + + [[nodiscard]] bool + IsLegacyHnswFusedCodec() const { + return quantizer_->FilterBits() == 1 and quantizer_->ReorderBits() == 7; + } + + [[nodiscard]] bool + AreFusedVectorsFinite(const float* data, uint64_t count) const { + if (data == nullptr) { + return false; + } + const auto dim = static_cast(common_param_.dim_); + for (uint64_t row = 0; row < count; ++row) { + for (uint64_t d = 0; d < dim; ++d, ++data) { + if (not IsFiniteRaBitQValue(*data)) { + return false; + } + } + } + return true; + } + + [[nodiscard]] uint32_t + NearestFusedCluster(const float* data) const { + uint32_t nearest = 0; + double nearest_distance = std::numeric_limits::max(); + for (uint32_t cluster_id = 0; cluster_id < fused_quantizers_.size(); ++cluster_id) { + const auto* centroid = + fused_centroids_.data() + static_cast(cluster_id) * common_param_.dim_; + const double distance = + FP32ComputeL2Sqr(data, centroid, static_cast(common_param_.dim_)); + if (distance < nearest_distance) { + nearest_distance = distance; + nearest = cluster_id; + } + } + return nearest; + } + static IOParamPtr SuffixIOParam(const IOParamPtr& io_param, const std::string& suffix, bool split_cache = false) { if (io_param == nullptr) { @@ -1086,6 +2011,180 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil this->supplement_cell_->Write(supplement_code.data, idx); } + [[nodiscard]] RaBitQFusedCodeView + get_fused_code_view(InnerIdType id) const { + RaBitQFusedCodeView view; + if (not fused_code_storage_->GetFusedCodeView(id, view)) { + throw VsagException( + ErrorType::INTERNAL_ERROR, "failed to read fused RaBitQ code for id ", id); + } + return view; + } + + void + prefetch_fused_codes(const InnerIdType* ids, + InnerIdType id_count, + bool include_supplement) const { + const auto count = std::min(prefetch_stride_code_, id_count); + for (InnerIdType i = 0; i < count; ++i) { + fused_code_storage_->PrefetchFusedCodes(ids[i], include_supplement); + } + } + + void + query_fused_full(float* result_dists, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + InnerIdType id_count, + QueryContext* ctx) const { + this->prefetch_fused_codes(ids, id_count, true); + for (InnerIdType i = 0; i < id_count; ++i) { + if (i + prefetch_stride_code_ < id_count) { + fused_code_storage_->PrefetchFusedCodes(ids[i + prefetch_stride_code_], true); + } + const auto view = this->get_fused_code_view(ids[i]); + CHECK_ARGUMENT(this->ComputeFusedFull(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + result_dists + i, + ctx), + "failed to compute fused RaBitQ distance"); + } + } + + void + query_fused_lower_bound(float* result_dists, + float* lower_bounds, + float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + InnerIdType id_count, + QueryContext* ctx) const { + const bool enable_reorder = ctx == nullptr or ctx->enable_rabitq_reorder; + this->prefetch_fused_codes(ids, id_count, false); + for (InnerIdType i = 0; i < id_count; ++i) { + if (i + prefetch_stride_code_ < id_count) { + fused_code_storage_->PrefetchFusedCodes(ids[i + prefetch_stride_code_], false); + } + const auto view = this->get_fused_code_view(ids[i]); + float local_lower_bound = std::numeric_limits::max(); + float local_filter_ip = std::numeric_limits::quiet_NaN(); + const bool computed = this->ComputeFusedOneBitWithFilterIP(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + result_dists + i, + &local_lower_bound, + &local_filter_ip, + ctx); + if (not enable_reorder) { + if (not IsFiniteRaBitQValue(result_dists[i])) { + result_dists[i] = std::numeric_limits::max(); + } + local_lower_bound = result_dists[i]; + local_filter_ip = std::numeric_limits::quiet_NaN(); + } else if (not computed) { + CHECK_ARGUMENT(this->ComputeFusedFull(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + result_dists + i, + ctx), + "failed to compute fused RaBitQ distance"); + local_lower_bound = std::numeric_limits::max(); + local_filter_ip = std::numeric_limits::quiet_NaN(); + } + if (lower_bounds != nullptr) { + lower_bounds[i] = local_lower_bound; + } + if (filter_inner_products != nullptr) { + filter_inner_products[i] = local_filter_ip; + } + } + } + + void + query_fused_with_distance_filter(float* result_dists, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + InnerIdType id_count, + float threshold, + QueryContext* ctx) const { + const bool enable_reorder = ctx == nullptr or ctx->enable_rabitq_reorder; + this->prefetch_fused_codes(ids, id_count, false); + for (InnerIdType i = 0; i < id_count; ++i) { + const auto view = this->get_fused_code_view(ids[i]); + float lower_bound = std::numeric_limits::max(); + float filter_inner_product = std::numeric_limits::quiet_NaN(); + const bool computed = this->ComputeFusedOneBitWithFilterIP(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + result_dists + i, + &lower_bound, + &filter_inner_product, + ctx); + if (not enable_reorder) { + if (not IsFiniteRaBitQValue(result_dists[i])) { + result_dists[i] = std::numeric_limits::max(); + } else if (computed and IsFiniteRaBitQValue(lower_bound) and + lower_bound >= threshold) { + result_dists[i] = threshold; + } + continue; + } + if (computed and IsFiniteRaBitQValue(lower_bound) and lower_bound >= threshold) { + result_dists[i] = threshold; + continue; + } + CHECK_ARGUMENT(this->ComputeFusedFull(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + result_dists + i, + ctx), + "failed to compute fused RaBitQ distance"); + } + } + + void + query_fused_full_with_filter_ip(float* result_dists, + const float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + InnerIdType id_count, + QueryContext* ctx) const { + this->prefetch_fused_codes(ids, id_count, true); + const bool exact_filter_ip_hint = + not IsLegacyHnswFusedCodec() and this->quantizer_->FilterBits() >= 2; + for (InnerIdType i = 0; i < id_count; ++i) { + const auto view = this->get_fused_code_view(ids[i]); + const float filter_ip = filter_inner_products == nullptr + ? std::numeric_limits::quiet_NaN() + : filter_inner_products[i]; + bool computed = false; + if (exact_filter_ip_hint and IsFiniteRaBitQValue(filter_ip)) { + computed = this->ComputeFusedFullWithFilterIP(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + filter_ip, + result_dists + i, + ctx); + } + if (not computed) { + CHECK_ARGUMENT(this->ComputeFusedFull(computer, + view.cluster_id, + view.one_bit_code, + view.supplement_code, + result_dists + i, + ctx), + "failed to compute fused RaBitQ distance"); + } + } + } + void query_optimized_build_codes(float* result_dists, const ComputerInterfacePtr& computer, @@ -1458,8 +2557,9 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil float hint_dist = std::numeric_limits::max()) const { this->add_full_count(ctx, 1); bool computed = false; - const bool has_hint = - std::isfinite(hint_dist) and hint_dist < std::numeric_limits::max(); + const bool has_hint = this->quantizer_->FilterBits() >= 2 and + IsFiniteRaBitQValue(hint_dist) and + hint_dist < std::numeric_limits::max(); if (has_hint) { computed = this->quantizer_->ComputeDistWithSplitCodeAndFilterDist( *computer, one_bit_code, supplement_code, hint_dist, result_dist); @@ -1477,6 +2577,58 @@ class RaBitQSplitDataCell : public FlattenInterface, public FlattenOptimizedBuil } } + void + compute_full_dist_with_filter_ip(const uint8_t* one_bit_code, + const uint8_t* supplement_code, + Computer>* computer, + float* result_dist, + QueryContext* ctx, + float filter_inner_product) const { + this->add_full_count(ctx, 1); + const bool has_hint = + this->quantizer_->FilterBits() >= 2 and IsFiniteRaBitQValue(filter_inner_product); + bool computed = false; + if (has_hint) { + computed = this->quantizer_->ComputeDistWithSplitCodeAndFilterIP( + *computer, one_bit_code, supplement_code, filter_inner_product, result_dist); + } + if (computed) { + this->add_reorder_hint_full_count(ctx, 1); + return; + } + if (has_hint) { + this->add_reorder_fallback_full_count(ctx, 1); + } + if (not this->quantizer_->ComputeDistWithSplitCode( + *computer, one_bit_code, supplement_code, result_dist)) { + ByteBuffer full_code(this->code_size_, allocator_); + this->quantizer_->MergeSplitCode(one_bit_code, supplement_code, full_code.data); + computer->ComputeDist(full_code.data, result_dist); + } + } + + void + compute_full_dist_with_filter_ip(InnerIdType id, + Computer>* computer, + float* result_dist, + QueryContext* ctx, + float filter_inner_product) const { + bool one_bit_need_release = false; + bool supplement_need_release = false; + const auto* one_bit_code = this->get_one_bit_code(id, one_bit_need_release); + const auto* supplement_code = this->get_supplement_code(id, supplement_need_release); + try { + this->compute_full_dist_with_filter_ip( + one_bit_code, supplement_code, computer, result_dist, ctx, filter_inner_product); + } catch (...) { + this->release_one_bit_code(one_bit_code, one_bit_need_release); + this->release_supplement_code(supplement_code, supplement_need_release); + throw; + } + this->release_one_bit_code(one_bit_code, one_bit_need_release); + this->release_supplement_code(supplement_code, supplement_need_release); + } + void compute_full_dist(InnerIdType id, Computer>* computer, diff --git a/src/impl/cluster/kmeans_cluster.cpp b/src/impl/cluster/kmeans_cluster.cpp index a72d5184fd..f6a87dd38a 100644 --- a/src/impl/cluster/kmeans_cluster.cpp +++ b/src/impl/cluster/kmeans_cluster.cpp @@ -52,7 +52,9 @@ KMeansCluster::Run(uint32_t k, double* err, bool use_mse_for_convergence, float threshold, - KMeansInitMethod init_method) { + KMeansInitMethod init_method, + std::optional random_seed, + bool deterministic_reduction) { if (k == 0) { throw VsagException(ErrorType::INVALID_ARGUMENT, "k must be positive"); } @@ -74,7 +76,7 @@ KMeansCluster::Run(uint32_t k, k_centroids_ = static_cast(allocator_->Allocate(size)); std::random_device rd; - std::mt19937 gen(rd()); + std::mt19937 gen(random_seed.has_value() ? *random_seed : rd()); if (init_method == KMeansInitMethod::KMEANS_PLUS_PLUS) { select_initial_centroids_kmeans_plus_plus(datas, count, k, gen); @@ -108,12 +110,20 @@ KMeansCluster::Run(uint32_t k, Vector counts(k, 0, allocator_); Vector new_centroids(static_cast(k) * dim_, 0.0F, allocator_); + const uint64_t block_count = (count + bs - 1) / bs; + const uint64_t centroid_values = static_cast(k) * dim_; + Vector block_counts(allocator_); + Vector block_centroids(allocator_); + if (deterministic_reduction) { + block_counts.resize(block_count * k, 0); + block_centroids.resize(block_count * centroid_values, 0.0F); + } std::mutex merge_mutex; - auto update_centroids_func = [&](uint64_t start, uint64_t end) { + auto update_centroids_func = [&](uint64_t block_id, uint64_t start, uint64_t end) { omp_set_num_threads(1); Vector local_counts(k, 0, allocator_); - Vector local_centroids(static_cast(k) * dim_, 0.0F, allocator_); + Vector local_centroids(centroid_values, 0.0F, allocator_); for (uint64_t i = start; i < end; ++i) { int32_t label = labels[i]; @@ -129,7 +139,13 @@ KMeansCluster::Run(uint32_t k, } } - { + if (deterministic_reduction) { + std::copy( + local_counts.begin(), local_counts.end(), block_counts.data() + block_id * k); + std::copy(local_centroids.begin(), + local_centroids.end(), + block_centroids.data() + block_id * centroid_values); + } else { std::lock_guard lock(merge_mutex); for (uint32_t j = 0; j < k; ++j) { if (local_counts[j] > 0) { @@ -146,13 +162,30 @@ KMeansCluster::Run(uint32_t k, } }; for (uint64_t i = 0; i < count; i += bs) { - futures.emplace_back( - thread_pool_->GeneralEnqueue(update_centroids_func, i, std::min(i + bs, count))); + futures.emplace_back(thread_pool_->GeneralEnqueue( + update_centroids_func, i / bs, i, std::min(i + bs, count))); } for (auto& future : futures) { future.wait(); } futures.clear(); + if (deterministic_reduction) { + for (uint64_t block_id = 0; block_id < block_count; ++block_id) { + for (uint32_t j = 0; j < k; ++j) { + const auto count_offset = block_id * k + j; + if (block_counts[count_offset] > 0) { + counts[j] += block_counts[count_offset]; + BlasFunction::Saxpy(dim_, + 1.0F, + block_centroids.data() + block_id * centroid_values + + static_cast(j) * dim_, + 1, + new_centroids.data() + static_cast(j) * dim_, + 1); + } + } + } + } std::uniform_int_distribution dis(0, count - 1); for (int j = 0; j < k; ++j) { diff --git a/src/impl/cluster/kmeans_cluster.h b/src/impl/cluster/kmeans_cluster.h index 954e6ba922..84bee17dd0 100644 --- a/src/impl/cluster/kmeans_cluster.h +++ b/src/impl/cluster/kmeans_cluster.h @@ -15,6 +15,7 @@ #pragma once +#include #include #include "impl/thread_pool/safe_thread_pool.h" @@ -44,7 +45,9 @@ class KMeansCluster { double* err = nullptr, bool use_mse_for_convergence = false, float threshold = 1e-6F, - KMeansInitMethod init_method = KMeansInitMethod::KMEANS_PLUS_PLUS); + KMeansInitMethod init_method = KMeansInitMethod::KMEANS_PLUS_PLUS, + std::optional random_seed = std::nullopt, + bool deterministic_reduction = false); public: float* k_centroids_{nullptr}; diff --git a/src/impl/cluster/kmeans_cluster_test.cpp b/src/impl/cluster/kmeans_cluster_test.cpp index 4058370571..1733925554 100644 --- a/src/impl/cluster/kmeans_cluster_test.cpp +++ b/src/impl/cluster/kmeans_cluster_test.cpp @@ -102,3 +102,51 @@ TEST_CASE("Kmeans Larger Dim (AMX BF16 path)", "[ut][KMeansCluster]") { } REQUIRE(converged); } + +TEST_CASE("Kmeans seeded fixed-order reduction is reproducible", "[ut][KMeansCluster]") { + constexpr uint32_t k = 8; + constexpr int32_t dim = 17; + constexpr uint64_t count = 4097; + std::vector data(count * dim); + for (uint64_t i = 0; i < count; ++i) { + for (int32_t d = 0; d < dim; ++d) { + data[i * dim + d] = + static_cast(i % k) * 5.0F + + static_cast((i * 31 + static_cast(d) * 17) % 97) * 0.0001F; + } + } + + auto allocator = vsag::SafeAllocator::FactoryDefaultAllocator(); + auto single_thread_pool = vsag::SafeThreadPool::FactoryDefaultThreadPool(); + single_thread_pool->SetPoolSize(1); + auto multi_thread_pool = vsag::SafeThreadPool::FactoryDefaultThreadPool(); + multi_thread_pool->SetPoolSize(4); + vsag::KMeansCluster single_thread(dim, allocator.get(), single_thread_pool); + vsag::KMeansCluster multi_thread(dim, allocator.get(), multi_thread_pool); + + const auto single_thread_labels = single_thread.Run(k, + data.data(), + count, + 6, + nullptr, + false, + 1e-6F, + vsag::KMeansInitMethod::KMEANS_PLUS_PLUS, + 0x52425131U, + true); + const auto multi_thread_labels = multi_thread.Run(k, + data.data(), + count, + 6, + nullptr, + false, + 1e-6F, + vsag::KMeansInitMethod::KMEANS_PLUS_PLUS, + 0x52425131U, + true); + REQUIRE(single_thread_labels == multi_thread_labels); + const uint64_t centroid_values = static_cast(k) * dim; + REQUIRE(std::equal(single_thread.k_centroids_, + single_thread.k_centroids_ + centroid_values, + multi_thread.k_centroids_)); +} diff --git a/src/impl/filter/CMakeLists.txt b/src/impl/filter/CMakeLists.txt index c7d06448d9..c2b2627be2 100644 --- a/src/impl/filter/CMakeLists.txt +++ b/src/impl/filter/CMakeLists.txt @@ -23,6 +23,7 @@ set (FILTER_SRC white_list_filter.h white_list_filter.cpp combined_filter.h + duplicate_group_filter.h iterator_filter.h iterator_filter.cpp ) diff --git a/src/impl/filter/duplicate_group_filter.h b/src/impl/filter/duplicate_group_filter.h new file mode 100644 index 0000000000..23d2b9f6c9 --- /dev/null +++ b/src/impl/filter/duplicate_group_filter.h @@ -0,0 +1,67 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include + +#include "datacell/graph_interface.h" +#include "vsag/filter.h" + +namespace vsag { + +class DuplicateGroupFilter final : public Filter { +public: + DuplicateGroupFilter(FilterPtr filter, GraphInterfacePtr graph) + : filter_(std::move(filter)), graph_(std::move(graph)) { + } + + [[nodiscard]] bool + CheckValid(int64_t id) const override { + const auto group_id = graph_->GetGroupId(static_cast(id)); + if (filter_->CheckValid(group_id)) { + return true; + } + for (const auto duplicate_id : graph_->GetDuplicateIds(group_id)) { + if (filter_->CheckValid(duplicate_id)) { + return true; + } + } + return false; + } + + [[nodiscard]] float + ValidRatio() const override { + // A duplicate group can only increase the effective valid ratio. Returning the source + // ratio is conservative: it preserves the existing skip policy while CheckValid keeps a + // representative traversable when only one of its aliases passes the filter. + return filter_->ValidRatio(); + } + +private: + FilterPtr filter_; + GraphInterfacePtr graph_; +}; + +inline FilterPtr +MakeDuplicateGroupFilter(const FilterPtr& filter, + const GraphInterfacePtr& graph, + bool consider_duplicate) { + if (not consider_duplicate or filter == nullptr) { + return filter; + } + return std::make_shared(filter, graph); +} + +} // namespace vsag diff --git a/src/impl/filter/duplicate_group_filter_test.cpp b/src/impl/filter/duplicate_group_filter_test.cpp new file mode 100644 index 0000000000..7d8fb72c77 --- /dev/null +++ b/src/impl/filter/duplicate_group_filter_test.cpp @@ -0,0 +1,74 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "duplicate_group_filter.h" + +#include "datacell/graph_datacell_parameter.h" +#include "impl/allocator/safe_allocator.h" +#include "index_common_param.h" +#include "io/memory_io/memory_io_parameter.h" +#include "unittest.h" + +namespace vsag { +namespace { + +class SingleIdFilter final : public Filter { +public: + explicit SingleIdFilter(int64_t valid_id) : valid_id_(valid_id) { + } + + [[nodiscard]] bool + CheckValid(int64_t id) const override { + return id == valid_id_; + } + + [[nodiscard]] float + ValidRatio() const override { + return 0.0F; + } + +private: + int64_t valid_id_; +}; + +} // namespace + +TEST_CASE("DuplicateGroupFilter keeps representatives traversable for alias-only filters", + "[ut][DuplicateGroupFilter]") { + IndexCommonParam common_param; + common_param.dim_ = 32; + common_param.allocator_ = SafeAllocator::FactoryDefaultAllocator(); + + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->support_duplicate_ = true; + auto graph = GraphInterface::MakeInstance(graph_param, common_param); + graph->Resize(4); + graph->SetDuplicateId(0, 1); + graph->SetDuplicateId(0, 2); + + FilterPtr alias_only = std::make_shared(2); + const auto group_filter = MakeDuplicateGroupFilter(alias_only, graph, true); + REQUIRE(group_filter != alias_only); + REQUIRE(group_filter->ValidRatio() == 0.0F); + REQUIRE(group_filter->CheckValid(int64_t{0})); + REQUIRE(group_filter->CheckValid(int64_t{1})); + REQUIRE(group_filter->CheckValid(int64_t{2})); + REQUIRE_FALSE(group_filter->CheckValid(int64_t{3})); + + REQUIRE(MakeDuplicateGroupFilter(alias_only, graph, false) == alias_only); + REQUIRE(MakeDuplicateGroupFilter(nullptr, graph, true) == nullptr); +} + +} // namespace vsag diff --git a/src/impl/heap/distance_heap.h b/src/impl/heap/distance_heap.h index 4551618ed6..1adcce166b 100644 --- a/src/impl/heap/distance_heap.h +++ b/src/impl/heap/distance_heap.h @@ -15,6 +15,7 @@ #pragma once +#include #include #include @@ -91,4 +92,14 @@ class DistanceHeap { }; using DistanceRecordVector = Vector; + +struct RaBitQCandidateRecord { + float lower_bound{std::numeric_limits::max()}; + float filter_inner_product{std::numeric_limits::quiet_NaN()}; + InnerIdType id{0}; + // NaN means that an exact full distance is unavailable for the current query. + float full_distance{std::numeric_limits::quiet_NaN()}; +}; + +using RaBitQCandidateVector = Vector; } // namespace vsag diff --git a/src/impl/inner_search_param.h b/src/impl/inner_search_param.h index ecb9ff3b26..e58ebfcd09 100644 --- a/src/impl/inner_search_param.h +++ b/src/impl/inner_search_param.h @@ -18,6 +18,7 @@ #include #include +#include "quantization/computer.h" #include "typing.h" #include "utils/filter_search_skip_strategy.h" #include "utils/pointer_define.h" @@ -39,6 +40,7 @@ class InnerSearchParam { public: int64_t topk{0}; + int64_t rerank_topk{0}; float radius{0.0F}; InnerIdType ep{0}; uint64_t ef{10}; @@ -53,6 +55,7 @@ class InnerSearchParam { bool enable_rabitq_one_bit_search{false}; SearchDistanceBatchFunc distance_batch_func{nullptr}; uint64_t distance_batch_size{1}; + ComputerInterfacePtr rabitq_fused_computer{nullptr}; // for ivf int scan_bucket_size{1}; diff --git a/src/impl/reorder/flatten_reorder.cpp b/src/impl/reorder/flatten_reorder.cpp index 9282e4766f..b853053ea0 100644 --- a/src/impl/reorder/flatten_reorder.cpp +++ b/src/impl/reorder/flatten_reorder.cpp @@ -22,6 +22,7 @@ #include #include "datacell/flatten_interface.h" +#include "datacell/rabitq_split_datacell.h" #include "impl/filter/iterator_filter.h" #include "impl/heap/standard_heap.h" #include "impl/reasoning/search_reasoning.h" @@ -29,6 +30,103 @@ namespace vsag { +void +FlattenReorder::QueryLowerBound(float* distances, + float* lower_bounds, + float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + uint64_t count, + QueryContext* ctx) const { + auto* split_codes = dynamic_cast(flatten_.get()); + if (fused_graph_ == nullptr or split_codes == nullptr) { + flatten_->QueryWithDistanceLowerBoundAndFilterIP( + distances, lower_bounds, filter_inner_products, computer, ids, count, ctx); + return; + } + uint32_t fallback_count = 0; + QueryContext rate_context; + QueryContext* rate_context_ptr = nullptr; + if (ctx != nullptr) { + rate_context.rabitq_error_rate = ctx->rabitq_error_rate; + rate_context_ptr = &rate_context; + } + for (uint64_t i = 0; i < count; ++i) { + const auto node = fused_graph_->GetNodeView(ids[i]); + if (not split_codes->ComputeFusedOneBitWithFilterIP(computer, + node.cluster_id, + node.one_bit_code, + node.supplement_code, + distances + i, + lower_bounds + i, + filter_inner_products + i, + rate_context_ptr)) { + ++fallback_count; + } + } + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->rabitq_filter_count.fetch_add(static_cast(count), + std::memory_order_relaxed); + ctx->stats->rabitq_filter_fallback_full_count.fetch_add(fallback_count, + std::memory_order_relaxed); + } +} + +void +FlattenReorder::QueryFullWithHint(float* distances, + const float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + uint64_t count, + QueryContext* ctx) const { + auto* split_codes = dynamic_cast(flatten_.get()); + if (fused_graph_ == nullptr or split_codes == nullptr) { + flatten_->QueryWithFilterIPHint( + distances, filter_inner_products, computer, ids, count, ctx); + return; + } + uint32_t hint_full_count = 0; + uint32_t fallback_full_count = 0; + uint32_t full_count = 0; + const bool exact_filter_ip_hint = + split_codes->FusedFilterBits() >= 2 and not split_codes->UsesLegacyHnswFusedCodec(); + for (uint64_t i = 0; i < count; ++i) { + const auto* record = fused_graph_->GetNodeRecord(ids[i]); + const auto cluster_id = fused_graph_->GetClusterId(record); + ++full_count; + bool used_hint = false; + if (exact_filter_ip_hint and IsFiniteRaBitQValue(filter_inner_products[i])) { + used_hint = + split_codes->ComputeFusedFullWithFilterIP(computer, + cluster_id, + fused_graph_->GetOneBitCode(record), + fused_graph_->GetSupplementCode(record), + filter_inner_products[i], + distances + i, + nullptr); + } + if (used_hint) { + ++hint_full_count; + } else { + ++fallback_full_count; + CHECK_ARGUMENT(split_codes->ComputeFusedFull(computer, + cluster_id, + fused_graph_->GetOneBitCode(record), + fused_graph_->GetSupplementCode(record), + distances + i, + nullptr), + "failed to compute fused RaBitQ distance"); + } + } + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->rabitq_full_count.fetch_add(full_count, std::memory_order_relaxed); + ctx->stats->rabitq_reorder_hint_full_count.fetch_add(hint_full_count, + std::memory_order_relaxed); + ctx->stats->rabitq_reorder_fallback_full_count.fetch_add(fallback_full_count, + std::memory_order_relaxed); + } +} + namespace { void @@ -55,14 +153,22 @@ FlattenReorder::Reorder(const vsag::DistHeapPtr& input, int64_t topk, QueryContext& ctx, IteratorFilterContext* iter_ctx, - const DistanceRecordVector* rabitq_lower_bound_candidates) { + const RaBitQCandidateVector* rabitq_lower_bound_candidates) { // set query allocator Allocator* query_allocator = select_query_allocator(ctx.alloc, allocator_); const uint64_t heap_candidate_size = input == nullptr ? 0 : input->Size(); + const auto add_iterator_discard = [iter_ctx](float distance, InnerIdType id) { + if (iter_ctx != nullptr) { + iter_ctx->AddDiscardNode(distance, id); + } + }; if (rabitq_lower_bound_candidates == nullptr) { topk = std::min(topk, static_cast(heap_candidate_size)); auto reorder_heap = std::make_shared>(query_allocator, topk); - auto computer = flatten_->FactoryComputer(query); + auto* split_codes = dynamic_cast(flatten_.get()); + auto computer = fused_graph_ != nullptr and split_codes != nullptr + ? split_codes->FactoryFusedComputer(query) + : flatten_->FactoryComputer(query); Vector ids(heap_candidate_size, query_allocator); Vector dists(heap_candidate_size, query_allocator); const auto* candidate_result = input == nullptr ? nullptr : input->GetData(); @@ -70,7 +176,21 @@ FlattenReorder::Reorder(const vsag::DistHeapPtr& input, ids[i] = candidate_result[i].second; } add_reorder_distance_count(ctx, heap_candidate_size); - flatten_->Query(dists.data(), computer, ids.data(), heap_candidate_size, &ctx); + if (fused_graph_ != nullptr and split_codes != nullptr) { + for (uint64_t i = 0; i < heap_candidate_size; ++i) { + const auto* record = fused_graph_->GetNodeRecord(ids[i]); + CHECK_ARGUMENT( + split_codes->ComputeFusedFull(computer, + fused_graph_->GetClusterId(record), + fused_graph_->GetOneBitCode(record), + fused_graph_->GetSupplementCode(record), + dists.data() + i, + &ctx), + "failed to compute fused RaBitQ distance"); + } + } else { + flatten_->Query(dists.data(), computer, ids.data(), heap_candidate_size, &ctx); + } for (uint64_t i = 0; i < heap_candidate_size; ++i) { if (ctx.reasoning_ctx != nullptr) { ctx.reasoning_ctx->RecordReorder( @@ -79,15 +199,15 @@ FlattenReorder::Reorder(const vsag::DistHeapPtr& input, if (reorder_heap->Size() < topk || dists[i] < reorder_heap->Top().first) { reorder_heap->Push(dists[i], candidate_result[i].second); if (reorder_heap->Size() > topk) { - if (iter_ctx != nullptr) { - auto curr = reorder_heap->Top(); - iter_ctx->AddDiscardNode(curr.first, curr.second); - } + const auto curr = reorder_heap->Top(); + add_iterator_discard(curr.first, curr.second); if (ctx.reasoning_ctx != nullptr) { ctx.reasoning_ctx->RecordReorderEviction(reorder_heap->Top().second, 0); } reorder_heap->Pop(); } + } else { + add_iterator_discard(dists[i], candidate_result[i].second); } } return reorder_heap; @@ -99,88 +219,184 @@ FlattenReorder::Reorder(const vsag::DistHeapPtr& input, if (topk <= 0) { topk = static_cast(max_candidate_size); } - auto computer = flatten_->FactoryComputer(query); - if (topk == 0 || max_candidate_size == 0) { + if (topk == 0 or max_candidate_size == 0) { return std::make_shared>(query_allocator, 0); } + auto has_valid_distance = [](float distance) { + return IsFiniteRaBitQValue(distance) and distance < std::numeric_limits::max(); + }; + auto* split_codes = dynamic_cast(flatten_.get()); + const bool accepts_full_distance_hints = fused_graph_ != nullptr and split_codes != nullptr; Vector all_ids(max_candidate_size, query_allocator); Vector lower_bound_probe_dists(max_candidate_size, query_allocator); Vector lower_bounds(max_candidate_size, query_allocator); - UnorderedSet merged_ids(query_allocator); - merged_ids.reserve(max_candidate_size); + Vector filter_inner_products( + max_candidate_size, std::numeric_limits::quiet_NaN(), query_allocator); + Vector full_distances( + max_candidate_size, std::numeric_limits::quiet_NaN(), query_allocator); + UnorderedMap fused_merged_ids(query_allocator); + UnorderedSet legacy_merged_ids(query_allocator); + if (accepts_full_distance_hints) { + fused_merged_ids.reserve(max_candidate_size); + } else { + legacy_merged_ids.reserve(max_candidate_size); + } uint64_t candidate_size = 0; + if (rabitq_lower_bound_candidates != nullptr) { + for (const auto& item : *rabitq_lower_bound_candidates) { + if (not accepts_full_distance_hints) { + if (legacy_merged_ids.insert(item.id).second) { + all_ids[candidate_size] = item.id; + lower_bound_probe_dists[candidate_size] = item.lower_bound; + lower_bounds[candidate_size] = item.lower_bound; + filter_inner_products[candidate_size] = item.filter_inner_product; + ++candidate_size; + } + continue; + } + + const auto [iter, inserted] = fused_merged_ids.emplace(item.id, candidate_size); + if (inserted) { + all_ids[candidate_size] = item.id; + lower_bound_probe_dists[candidate_size] = item.lower_bound; + lower_bounds[candidate_size] = item.lower_bound; + filter_inner_products[candidate_size] = item.filter_inner_product; + full_distances[candidate_size] = item.full_distance; + ++candidate_size; + continue; + } + + const auto idx = iter->second; + if (has_valid_distance(item.lower_bound) and + (not has_valid_distance(lower_bounds[idx]) or + item.lower_bound < lower_bounds[idx])) { + lower_bound_probe_dists[idx] = item.lower_bound; + lower_bounds[idx] = item.lower_bound; + } + if (not IsFiniteRaBitQValue(filter_inner_products[idx]) and + IsFiniteRaBitQValue(item.filter_inner_product)) { + filter_inner_products[idx] = item.filter_inner_product; + } + if (not has_valid_distance(full_distances[idx]) and + has_valid_distance(item.full_distance)) { + full_distances[idx] = item.full_distance; + } + } + } + const uint64_t hinted_candidate_size = candidate_size; + const auto* candidate_result = input == nullptr ? nullptr : input->GetData(); for (uint64_t i = 0; i < heap_candidate_size; ++i) { const auto id = candidate_result[i].second; - if (merged_ids.insert(id).second) { + const bool inserted = accepts_full_distance_hints + ? fused_merged_ids.emplace(id, candidate_size).second + : legacy_merged_ids.insert(id).second; + if (inserted) { all_ids[candidate_size++] = id; } } - const uint64_t heap_unique_size = candidate_size; - if (heap_unique_size > 0) { - add_reorder_lower_bound_probe_count(ctx, heap_unique_size); - flatten_->QueryWithDistanceLowerBound(lower_bound_probe_dists.data(), - lower_bounds.data(), - computer, - all_ids.data(), - heap_unique_size, - &ctx); - } - if (rabitq_lower_bound_candidates != nullptr) { - for (const auto& item : *rabitq_lower_bound_candidates) { - if (merged_ids.insert(item.second).second) { - all_ids[candidate_size] = item.second; - lower_bound_probe_dists[candidate_size] = item.first; - lower_bounds[candidate_size] = item.first; - ++candidate_size; - } + ComputerInterfacePtr computer{nullptr}; + const auto ensure_computer = [&]() -> const ComputerInterfacePtr& { + if (computer == nullptr) { + computer = fused_graph_ != nullptr and split_codes != nullptr + ? split_codes->FactoryFusedComputer(query) + : flatten_->FactoryComputer(query); } + return computer; + }; + + const uint64_t unhinted_candidate_size = candidate_size - hinted_candidate_size; + if (unhinted_candidate_size > 0) { + add_reorder_lower_bound_probe_count(ctx, unhinted_candidate_size); + QueryLowerBound(lower_bound_probe_dists.data() + hinted_candidate_size, + lower_bounds.data() + hinted_candidate_size, + filter_inner_products.data() + hinted_candidate_size, + ensure_computer(), + all_ids.data() + hinted_candidate_size, + unhinted_candidate_size, + &ctx); } topk = std::min(topk, static_cast(candidate_size)); auto reorder_heap = std::make_shared>(query_allocator, topk); - if (topk == 0 || candidate_size == 0) { + if (topk == 0 or candidate_size == 0) { return reorder_heap; } - auto has_valid_lower_bound = [](float lower_bound) { - return std::isfinite(lower_bound) and lower_bound < std::numeric_limits::max(); + const auto push_full_distance = [&](uint64_t idx, float distance) { + if (ctx.reasoning_ctx != nullptr) { + ctx.reasoning_ctx->RecordReorder(all_ids[idx], lower_bound_probe_dists[idx], distance); + } + if (reorder_heap->Size() < topk or distance < reorder_heap->Top().first) { + reorder_heap->Push(distance, all_ids[idx]); + if (reorder_heap->Size() > topk) { + const auto curr = reorder_heap->Top(); + add_iterator_discard(curr.first, curr.second); + if (ctx.reasoning_ctx != nullptr) { + ctx.reasoning_ctx->RecordReorderEviction(curr.second, 0); + } + reorder_heap->Pop(); + } + } else { + add_iterator_discard(distance, all_ids[idx]); + } }; + + bool all_full_distances_available = true; + for (uint64_t i = 0; i < candidate_size; ++i) { + if (not has_valid_distance(full_distances[i])) { + all_full_distances_available = false; + break; + } + } + if (all_full_distances_available) { + for (uint64_t i = 0; i < candidate_size; ++i) { + push_full_distance(i, full_distances[i]); + } + return reorder_heap; + } + bool lower_bounds_available = true; for (uint64_t i = 0; i < candidate_size; ++i) { - if (not has_valid_lower_bound(lower_bounds[i])) { + if (not has_valid_distance(lower_bounds[i])) { lower_bounds_available = false; break; } } if (not lower_bounds_available) { - add_reorder_distance_count(ctx, candidate_size); - flatten_->Query( - lower_bound_probe_dists.data(), computer, all_ids.data(), candidate_size, &ctx); + Vector missing_ids(candidate_size, query_allocator); + Vector missing_hints(candidate_size, query_allocator); + Vector missing_dists(candidate_size, query_allocator); + Vector missing_indices(candidate_size, query_allocator); + uint64_t missing_count = 0; for (uint64_t i = 0; i < candidate_size; ++i) { - if (ctx.reasoning_ctx != nullptr) { - ctx.reasoning_ctx->RecordReorder( - all_ids[i], lower_bounds[i], lower_bound_probe_dists[i]); + if (has_valid_distance(full_distances[i])) { + continue; } - if (reorder_heap->Size() < topk or - lower_bound_probe_dists[i] < reorder_heap->Top().first) { - reorder_heap->Push(lower_bound_probe_dists[i], all_ids[i]); - if (reorder_heap->Size() > topk) { - if (iter_ctx != nullptr) { - auto curr = reorder_heap->Top(); - iter_ctx->AddDiscardNode(curr.first, curr.second); - } - if (ctx.reasoning_ctx != nullptr) { - ctx.reasoning_ctx->RecordReorderEviction(reorder_heap->Top().second, 0); - } - reorder_heap->Pop(); - } + missing_ids[missing_count] = all_ids[i]; + missing_hints[missing_count] = filter_inner_products[i]; + missing_indices[missing_count] = i; + ++missing_count; + } + if (missing_count > 0) { + add_reorder_distance_count(ctx, missing_count); + QueryFullWithHint(missing_dists.data(), + missing_hints.data(), + ensure_computer(), + missing_ids.data(), + missing_count, + &ctx); + for (uint64_t i = 0; i < missing_count; ++i) { + full_distances[missing_indices[i]] = missing_dists[i]; } } + for (uint64_t i = 0; i < candidate_size; ++i) { + push_full_distance(i, full_distances[i]); + } return reorder_heap; } @@ -190,50 +406,82 @@ FlattenReorder::Reorder(const vsag::DistHeapPtr& input, return lower_bounds[lhs] < lower_bounds[rhs]; }); - const uint64_t bootstrap_size = std::min(static_cast(topk), candidate_size); constexpr uint64_t batch_size = 256; - const auto buffer_size = std::max(bootstrap_size, batch_size); - Vector ids(buffer_size, query_allocator); - Vector dists(buffer_size, query_allocator); - Vector hint_dists(buffer_size, query_allocator); - Vector batch_indices(buffer_size, query_allocator); - - for (uint64_t i = 0; i < bootstrap_size; ++i) { - const auto idx = order[i]; - ids[i] = all_ids[idx]; - hint_dists[i] = idx < heap_unique_size ? lower_bound_probe_dists[idx] - : std::numeric_limits::max(); - } - add_reorder_distance_count(ctx, bootstrap_size); - flatten_->QueryWithDistanceHint( - dists.data(), hint_dists.data(), computer, ids.data(), bootstrap_size, &ctx); - for (uint64_t i = 0; i < bootstrap_size; ++i) { - if (ctx.reasoning_ctx != nullptr) { - const auto idx = order[i]; - ctx.reasoning_ctx->RecordReorder(ids[i], lower_bound_probe_dists[idx], dists[i]); + Vector batch_dists(batch_size, query_allocator); + Vector batch_indices(batch_size, query_allocator); + Vector missing_ids(batch_size, query_allocator); + Vector missing_hints(batch_size, query_allocator); + Vector missing_dists(batch_size, query_allocator); + Vector missing_positions(batch_size, query_allocator); + + const auto compute_missing_full_distances = [&](uint64_t batch_count) { + uint64_t missing_count = 0; + for (uint64_t i = 0; i < batch_count; ++i) { + const auto idx = batch_indices[i]; + if (has_valid_distance(full_distances[idx])) { + batch_dists[i] = full_distances[idx]; + continue; + } + missing_ids[missing_count] = all_ids[idx]; + missing_hints[missing_count] = filter_inner_products[idx]; + missing_positions[missing_count] = i; + ++missing_count; + } + if (missing_count == 0) { + return; + } + add_reorder_distance_count(ctx, missing_count); + QueryFullWithHint(missing_dists.data(), + missing_hints.data(), + ensure_computer(), + missing_ids.data(), + missing_count, + &ctx); + for (uint64_t i = 0; i < missing_count; ++i) { + const auto position = missing_positions[i]; + const auto idx = batch_indices[position]; + batch_dists[position] = missing_dists[i]; + full_distances[idx] = missing_dists[i]; + } + }; + + // Seed the exact threshold with every full distance produced during traversal. Missing + // candidates then load their supplement only while their lower bound can still enter top-k. + for (uint64_t i = 0; i < candidate_size; ++i) { + if (has_valid_distance(full_distances[i])) { + push_full_distance(i, full_distances[i]); } - reorder_heap->Push(dists[i], ids[i]); } - uint64_t cursor = bootstrap_size; + uint64_t cursor = 0; while (cursor < candidate_size) { - if (reorder_heap->Size() == topk && - lower_bounds[order[cursor]] >= reorder_heap->Top().first) { + while (cursor < candidate_size and has_valid_distance(full_distances[order[cursor]])) { + ++cursor; + } + if (cursor == candidate_size) { break; } - - const auto threshold = reorder_heap->Top().first; + const bool heap_full = reorder_heap->Size() == static_cast(topk); + const float threshold = + heap_full ? reorder_heap->Top().first : std::numeric_limits::max(); + if (heap_full and lower_bounds[order[cursor]] >= threshold) { + break; + } + const uint64_t batch_limit = + heap_full ? batch_size + : std::min(batch_size, + static_cast(topk) - reorder_heap->Size()); uint64_t batch_count = 0; - while (cursor < candidate_size && batch_count < batch_size) { + while (cursor < candidate_size and batch_count < batch_limit) { const auto idx = order[cursor]; - if (lower_bounds[idx] >= threshold) { + if (has_valid_distance(full_distances[idx])) { + ++cursor; + continue; + } + if (heap_full and lower_bounds[idx] >= threshold) { break; } - ids[batch_count] = all_ids[idx]; - hint_dists[batch_count] = idx < heap_unique_size ? lower_bound_probe_dists[idx] - : std::numeric_limits::max(); - batch_indices[batch_count] = idx; - ++batch_count; + batch_indices[batch_count++] = idx; ++cursor; } @@ -241,26 +489,33 @@ FlattenReorder::Reorder(const vsag::DistHeapPtr& input, break; } - add_reorder_distance_count(ctx, batch_count); - flatten_->QueryWithDistanceHint( - dists.data(), hint_dists.data(), computer, ids.data(), batch_count, &ctx); + compute_missing_full_distances(batch_count); for (uint64_t i = 0; i < batch_count; ++i) { - if (ctx.reasoning_ctx != nullptr) { - ctx.reasoning_ctx->RecordReorder( - ids[i], lower_bound_probe_dists[batch_indices[i]], dists[i]); + push_full_distance(batch_indices[i], batch_dists[i]); + } + } + if (iter_ctx != nullptr) { + // Iterator pages can later expose discard values directly (including on the final drain), + // so every buffered candidate must carry a full distance rather than a lower bound. + while (cursor < candidate_size) { + uint64_t batch_count = 0; + while (cursor < candidate_size and batch_count < batch_size) { + const auto idx = order[cursor++]; + if (not has_valid_distance(full_distances[idx])) { + batch_indices[batch_count++] = idx; + } } - if (dists[i] < reorder_heap->Top().first) { - reorder_heap->Push(dists[i], ids[i]); - if (reorder_heap->Size() > topk) { - if (iter_ctx != nullptr) { - auto curr = reorder_heap->Top(); - iter_ctx->AddDiscardNode(curr.first, curr.second); - } - if (ctx.reasoning_ctx != nullptr) { - ctx.reasoning_ctx->RecordReorderEviction(reorder_heap->Top().second, 0); - } - reorder_heap->Pop(); + if (batch_count == 0) { + continue; + } + compute_missing_full_distances(batch_count); + for (uint64_t i = 0; i < batch_count; ++i) { + const auto idx = batch_indices[i]; + if (ctx.reasoning_ctx != nullptr) { + ctx.reasoning_ctx->RecordReorder( + all_ids[idx], lower_bound_probe_dists[idx], batch_dists[i]); } + add_iterator_discard(batch_dists[i], all_ids[idx]); } } } diff --git a/src/impl/reorder/flatten_reorder.h b/src/impl/reorder/flatten_reorder.h index 00bbb86467..b9473446d3 100644 --- a/src/impl/reorder/flatten_reorder.h +++ b/src/impl/reorder/flatten_reorder.h @@ -18,6 +18,7 @@ #include #include "datacell/flatten_interface.h" +#include "datacell/hgraph_rabitq_fused_datacell.h" #include "impl/heap/distance_heap.h" #include "impl/reorder/reorder.h" #include "utils/pointer_define.h" @@ -25,8 +26,10 @@ namespace vsag { class FlattenReorder : public ReorderInterface { public: - FlattenReorder(const FlattenInterfacePtr& flatten, Allocator* allocator) - : flatten_(flatten), allocator_(allocator) { + FlattenReorder(const FlattenInterfacePtr& flatten, + Allocator* allocator, + HGraphRaBitQFusedDataCellPtr fused_graph = nullptr) + : flatten_(flatten), allocator_(allocator), fused_graph_(std::move(fused_graph)) { } DistHeapPtr @@ -35,10 +38,28 @@ class FlattenReorder : public ReorderInterface { int64_t topk, QueryContext& ctx, IteratorFilterContext* iter_ctx = nullptr, - const DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) override; + const RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) override; private: + void + QueryLowerBound(float* distances, + float* lower_bounds, + float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + uint64_t count, + QueryContext* ctx) const; + + void + QueryFullWithHint(float* distances, + const float* filter_inner_products, + const ComputerInterfacePtr& computer, + const InnerIdType* ids, + uint64_t count, + QueryContext* ctx) const; + const FlattenInterfacePtr flatten_; Allocator* allocator_{nullptr}; + HGraphRaBitQFusedDataCellPtr fused_graph_{nullptr}; }; } // namespace vsag diff --git a/src/impl/reorder/reorder.h b/src/impl/reorder/reorder.h index df9a340b98..ba9bc764b1 100644 --- a/src/impl/reorder/reorder.h +++ b/src/impl/reorder/reorder.h @@ -31,7 +31,7 @@ class ReorderInterface { int64_t topk, QueryContext& ctx, IteratorFilterContext* iter_ctx = nullptr, - const DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) = 0; + const RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) = 0; }; } // namespace vsag diff --git a/src/impl/searcher/CMakeLists.txt b/src/impl/searcher/CMakeLists.txt index b68dafb0f7..d39cc7502b 100644 --- a/src/impl/searcher/CMakeLists.txt +++ b/src/impl/searcher/CMakeLists.txt @@ -17,6 +17,8 @@ set (SEARCHER_SRC basic_searcher.cpp basic_searcher.h + hgraph_rabitq_searcher.cpp + hgraph_rabitq_searcher.h mci_searcher.cpp mci_searcher.h parallel_searcher.cpp diff --git a/src/impl/searcher/basic_searcher.cpp b/src/impl/searcher/basic_searcher.cpp index e9745f3558..e36b1ff5c6 100644 --- a/src/impl/searcher/basic_searcher.cpp +++ b/src/impl/searcher/basic_searcher.cpp @@ -22,6 +22,7 @@ #include #include "datacell/flatten_interface.h" +#include "impl/filter/duplicate_group_filter.h" #include "impl/filter/iterator_filter.h" #include "impl/heap/standard_heap.h" #include "impl/reasoning/search_reasoning.h" @@ -79,7 +80,7 @@ BasicSearcher::Search(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { if (inner_search_param.search_mode == KNN_SEARCH) { return this->search_impl(graph, flatten, @@ -110,7 +111,7 @@ BasicSearcher::SearchWithPresetComputer(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates, + RaBitQCandidateVector* rabitq_lower_bound_candidates, const ComputerInterfacePtr& preset_computer) const { if (inner_search_param.search_mode == KNN_SEARCH) { return this->search_impl(graph, @@ -142,7 +143,7 @@ BasicSearcher::Search(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, IteratorFilterContext* iter_ctx, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { return this->search_impl(graph, flatten, vl, @@ -300,7 +301,7 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, IteratorFilterContext* iter_ctx, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { // set customize query alloctor Allocator* alloc = select_query_allocator(ctx, allocator_); @@ -329,16 +330,85 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, Vector neighbors(graph->MaximumDegree(), alloc); Vector line_dists(graph->MaximumDegree(), alloc); Vector lower_bound_dists(graph->MaximumDegree(), alloc); + Vector filter_inner_products(graph->MaximumDegree(), alloc); + const auto visit_filter = MakeDuplicateGroupFilter( + inner_search_param.is_inner_id_allowed, graph, inner_search_param.consider_duplicate); auto skip_strategy = create_filter_search_skip_strategy( inner_search_param.skip_strategy_type, - inner_search_param.is_inner_id_allowed != nullptr - ? inner_search_param.is_inner_id_allowed->ValidRatio() - : 1.0F, + visit_filter != nullptr ? visit_filter->ValidRatio() : 1.0F, inner_search_param.skip_ratio); if (rabitq_lower_bound_candidates != nullptr) { rabitq_lower_bound_candidates->clear(); } + UnorderedSet expanded_duplicate_groups(alloc); + auto is_result_allowed = [&is_id_allowed](InnerIdType id) { + return is_id_allowed == nullptr or is_id_allowed->CheckValid(id); + }; + auto push_result = [&](InnerIdType id, float distance) { + if (not iter_ctx->CheckPoint(id) or not is_result_allowed(id)) { + return false; + } + top_candidates->Push(distance, id); + return true; + }; + auto push_duplicate_group = [&](InnerIdType id, float distance) { + if (not inner_search_param.consider_duplicate) { + return push_result(id, distance); + } + + const auto group_id = graph->GetGroupId(id); + if (not expanded_duplicate_groups.insert(group_id).second) { + return false; + } + + bool pushed = push_result(group_id, distance); + for (const auto duplicate_id : graph->GetDuplicateIds(group_id)) { + pushed = push_result(duplicate_id, distance) or pushed; + } + return pushed; + }; + auto append_lower_bound_group = [&](InnerIdType id, float bound, float filter_ip) { + if (rabitq_lower_bound_candidates == nullptr) { + return; + } + const auto group_id = inner_search_param.consider_duplicate ? graph->GetGroupId(id) : id; + const auto append = [&](InnerIdType candidate, float candidate_bound, float candidate_ip) { + if (iter_ctx->CheckPoint(candidate) and is_result_allowed(candidate)) { + rabitq_lower_bound_candidates->push_back( + {candidate_bound, candidate_ip, candidate}); + } + }; + append(group_id, bound, filter_ip); + if (inner_search_param.consider_duplicate) { + for (const auto duplicate_id : graph->GetDuplicateIds(group_id)) { + if (not iter_ctx->CheckPoint(duplicate_id) or not is_result_allowed(duplicate_id)) { + continue; + } + float duplicate_distance = 0.0F; + float duplicate_bound = std::numeric_limits::max(); + float duplicate_filter_ip = std::numeric_limits::quiet_NaN(); + flatten->QueryWithDistanceLowerBoundAndFilterIP(&duplicate_distance, + &duplicate_bound, + &duplicate_filter_ip, + computer, + &duplicate_id, + 1, + ctx); + append(duplicate_id, duplicate_bound, duplicate_filter_ip); + } + } + }; + auto trim_top_candidates = [&](uint64_t limit) { + while (top_candidates->Size() > limit) { + const auto candidate = top_candidates->Top(); + if (iter_ctx->CheckPoint(candidate.second)) { + iter_ctx->AddDiscardNode(candidate.first, candidate.second); + } + top_candidates->Pop(); + } + }; + if (!iter_ctx->IsFirstUsed()) { if (iter_ctx->Empty()) { return top_candidates; @@ -346,41 +416,45 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, while (!iter_ctx->Empty()) { uint32_t cur_inner_id = iter_ctx->GetTopID(); float cur_dist = iter_ctx->GetTopDist(); - vl->Set(cur_inner_id); if (iter_ctx->CheckPoint(cur_inner_id)) { - flatten->Query(&cur_dist, computer, &cur_inner_id, 1, ctx); + const auto traversal_id = inner_search_param.consider_duplicate + ? graph->GetGroupId(cur_inner_id) + : cur_inner_id; + const bool group_needs_expansion = + not inner_search_param.consider_duplicate or + expanded_duplicate_groups.find(traversal_id) == expanded_duplicate_groups.end(); + if (group_needs_expansion) { + flatten->Query(&cur_dist, computer, &traversal_id, 1, ctx); + push_duplicate_group(traversal_id, cur_dist); + } // Sign convention: top_candidates stores positive distances (nearest = smallest); - // candidate_set is a max-heap, so distances are negated (nearest = largest, popped first). - top_candidates->Push(cur_dist, cur_inner_id); - candidate_set->Push(-cur_dist, cur_inner_id); - if constexpr (mode == InnerSearchMode::RANGE_SEARCH) { - if (cur_dist > inner_search_param.radius and not top_candidates->Empty()) { - top_candidates->Pop(); - } + // candidate_set is a max-heap, so distances are negated (nearest = largest, + // popped first). + if (not vl->TestAndSet(traversal_id)) { + candidate_set->Push(-cur_dist, traversal_id); } } iter_ctx->PopDiscard(); } if constexpr (mode == InnerSearchMode::KNN_SEARCH) { - while (top_candidates->Size() > ef) { - auto cur_node_pair = top_candidates->Top(); - if (iter_ctx->CheckPoint(cur_node_pair.second)) { - iter_ctx->AddDiscardNode(cur_node_pair.first, cur_node_pair.second); - } - top_candidates->Pop(); - } + trim_top_candidates(ef); } if (not top_candidates->Empty()) { lower_bound = top_candidates->Top().first; } } else { if (inner_search_param.enable_rabitq_one_bit_search) { - flatten->QueryWithDistanceLowerBound(&dist, nullptr, computer, &ep, 1, ctx); + float entry_lower_bound = std::numeric_limits::max(); + float entry_filter_ip = std::numeric_limits::quiet_NaN(); + flatten->QueryWithDistanceLowerBoundAndFilterIP( + &dist, &entry_lower_bound, &entry_filter_ip, computer, &ep, 1, ctx); + append_lower_bound_group(ep, entry_lower_bound, entry_filter_ip); } else { flatten->Query(&dist, computer, &ep, 1, ctx); } - if (not is_id_allowed || is_id_allowed->CheckValid(ep)) { - top_candidates->Push(dist, ep); + push_duplicate_group(ep, dist); + trim_top_candidates(ef); + if (not top_candidates->Empty()) { lower_bound = top_candidates->Top().first; } candidate_set->Push(-dist, ep); @@ -411,7 +485,7 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, count_no_visited = visit(graph, vl, current_node_pair, - inner_search_param.is_inner_id_allowed, + visit_filter, skip_strategy.get(), to_be_visited_id, neighbors); @@ -419,15 +493,16 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, dist_cmp += count_no_visited; bool collect_rabitq_lower_bound = false; - if (inner_search_param.enable_rabitq_one_bit_search and top_candidates->Size() == ef and + if (inner_search_param.enable_rabitq_one_bit_search and rabitq_lower_bound_candidates != nullptr) { collect_rabitq_lower_bound = true; - flatten->QueryWithDistanceLowerBound(line_dists.data(), - lower_bound_dists.data(), - computer, - to_be_visited_id.data(), - count_no_visited, - ctx); + flatten->QueryWithDistanceLowerBoundAndFilterIP(line_dists.data(), + lower_bound_dists.data(), + filter_inner_products.data(), + computer, + to_be_visited_id.data(), + count_no_visited, + ctx); } else if (inner_search_param.enable_rabitq_one_bit_search) { flatten->QueryWithDistanceLowerBound(line_dists.data(), nullptr, @@ -443,32 +518,25 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, for (uint32_t i = 0; i < count_no_visited; i++) { dist = line_dists[i]; const auto cur_id = to_be_visited_id[i]; - const bool id_allowed = not is_id_allowed || is_id_allowed->CheckValid(cur_id); if constexpr (mode == KNN_SEARCH) { - if (collect_rabitq_lower_bound and lower_bound_dists[i] < lower_bound and - id_allowed and iter_ctx->CheckPoint(cur_id)) { - rabitq_lower_bound_candidates->emplace_back(lower_bound_dists[i], cur_id); + if (collect_rabitq_lower_bound and + (top_candidates->Size() < ef or lower_bound_dists[i] < lower_bound)) { + append_lower_bound_group( + cur_id, lower_bound_dists[i], filter_inner_products[i]); } } if (top_candidates->Size() < ef || lower_bound > dist || (mode == RANGE_SEARCH && dist <= inner_search_param.radius)) { - if (!iter_ctx->CheckPoint(cur_id)) { + const bool source_available = iter_ctx->CheckPoint(cur_id); + const bool pushed_group = push_duplicate_group(cur_id, dist); + if (not source_available and not pushed_group) { continue; } candidate_set->Push(-dist, cur_id); flatten->Prefetch(candidate_set->Top().second); - if (id_allowed) { - top_candidates->Push(dist, cur_id); - } if constexpr (mode == KNN_SEARCH) { - if (top_candidates->Size() > ef) { - if (iter_ctx->CheckPoint(top_candidates->Top().second)) { - auto cur_node_pair = top_candidates->Top(); - iter_ctx->AddDiscardNode(cur_node_pair.first, cur_node_pair.second); - } - top_candidates->Pop(); - } + trim_top_candidates(ef); } if (not top_candidates->Empty()) { @@ -479,13 +547,7 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, } if constexpr (mode == KNN_SEARCH) { - while (top_candidates->Size() > inner_search_param.topk) { - auto cur_node_pair = top_candidates->Top(); - if (iter_ctx->CheckPoint(cur_node_pair.second)) { - iter_ctx->AddDiscardNode(cur_node_pair.first, cur_node_pair.second); - } - top_candidates->Pop(); - } + trim_top_candidates(static_cast(inner_search_param.topk)); } return top_candidates; @@ -500,7 +562,7 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates, + RaBitQCandidateVector* rabitq_lower_bound_candidates, const ComputerInterfacePtr& preset_computer) const { // set customize query alloctor Allocator* alloc = select_query_allocator(ctx, allocator_); @@ -539,11 +601,12 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, graph->MaximumDegree())) : 0; Vector custom_labels(custom_batch_capacity, alloc); + Vector filter_inner_products(graph->MaximumDegree(), alloc); + const auto visit_filter = MakeDuplicateGroupFilter( + inner_search_param.is_inner_id_allowed, graph, inner_search_param.consider_duplicate); auto skip_strategy = create_filter_search_skip_strategy( inner_search_param.skip_strategy_type, - inner_search_param.is_inner_id_allowed != nullptr - ? inner_search_param.is_inner_id_allowed->ValidRatio() - : 1.0F, + visit_filter != nullptr ? visit_filter->ValidRatio() : 1.0F, inner_search_param.skip_ratio); if (rabitq_lower_bound_candidates != nullptr) { rabitq_lower_bound_candidates->clear(); @@ -559,6 +622,47 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, return (is_id_allowed == nullptr or is_id_allowed->CheckValid(id)) and (attr_ft == nullptr or attr_ft->CheckValid(id)); }; + auto push_duplicate_candidates = [&](InnerIdType id, float distance) { + if (not inner_search_param.consider_duplicate) { + return; + } + for (const auto duplicate_id : graph->GetDuplicateIds(id)) { + if (check_func(duplicate_id)) { + top_candidates->Push(distance, duplicate_id); + } + } + }; + auto append_lower_bound_candidates = [&](InnerIdType id, float bound, float filter_ip) { + if (rabitq_lower_bound_candidates == nullptr) { + return; + } + const auto group_id = inner_search_param.consider_duplicate ? graph->GetGroupId(id) : id; + const auto append = [&](InnerIdType candidate, float candidate_bound, float candidate_ip) { + if (check_func(candidate)) { + rabitq_lower_bound_candidates->push_back( + {candidate_bound, candidate_ip, candidate}); + } + }; + append(group_id, bound, filter_ip); + if (inner_search_param.consider_duplicate) { + for (const auto duplicate_id : graph->GetDuplicateIds(group_id)) { + if (not check_func(duplicate_id)) { + continue; + } + float duplicate_distance = 0.0F; + float duplicate_bound = std::numeric_limits::max(); + float duplicate_filter_ip = std::numeric_limits::quiet_NaN(); + flatten->QueryWithDistanceLowerBoundAndFilterIP(&duplicate_distance, + &duplicate_bound, + &duplicate_filter_ip, + computer, + &duplicate_id, + 1, + ctx); + append(duplicate_id, duplicate_bound, duplicate_filter_ip); + } + } + }; auto* reasoning = ctx == nullptr ? nullptr : ctx->reasoning_ctx; auto score_ids = [&](const InnerIdType* ids, uint64_t count, float* scores) { @@ -632,20 +736,33 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, if (use_custom_distance) { score_ids(&ep, 1, &dist); } else if (inner_search_param.enable_rabitq_one_bit_search) { - flatten->QueryWithDistanceLowerBound(&dist, nullptr, computer, &ep, 1, ctx); + float entry_lower_bound = std::numeric_limits::max(); + float entry_filter_ip = std::numeric_limits::quiet_NaN(); + flatten->QueryWithDistanceLowerBoundAndFilterIP( + &dist, &entry_lower_bound, &entry_filter_ip, computer, &ep, 1, ctx); + append_lower_bound_candidates(ep, entry_lower_bound, entry_filter_ip); } else { flatten->Query(&dist, computer, &ep, 1, ctx); } ++dist_cmp; if (check_func(ep)) { top_candidates->Push(dist, ep); - lower_bound = top_candidates->Top().first; + } + if (not use_custom_distance) { + push_duplicate_candidates(ep, dist); } if constexpr (mode == InnerSearchMode::RANGE_SEARCH) { - if (dist > inner_search_param.radius and not top_candidates->Empty()) { + while (dist > inner_search_param.radius and not top_candidates->Empty()) { + top_candidates->Pop(); + } + } else if constexpr (mode == InnerSearchMode::KNN_SEARCH) { + while (top_candidates->Size() > ef) { top_candidates->Pop(); } } + if (not top_candidates->Empty()) { + lower_bound = top_candidates->Top().first; + } if (use_custom_distance and inner_search_param.consider_duplicate) { const auto duplicate_ids = graph->GetDuplicateIds(ep); score_duplicates(duplicate_ids, hops); @@ -694,7 +811,7 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, count_no_visited = visit(graph, vl, current_node_pair, - inner_search_param.is_inner_id_allowed, + visit_filter, skip_strategy.get(), to_be_visited_id, neighbors); @@ -703,14 +820,15 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, if (use_custom_distance) { score_ids(to_be_visited_id.data(), count_no_visited, line_dists.data()); } else if (inner_search_param.enable_rabitq_one_bit_search and - top_candidates->Size() == ef and rabitq_lower_bound_candidates != nullptr) { + rabitq_lower_bound_candidates != nullptr) { collect_rabitq_lower_bound = true; - flatten->QueryWithDistanceLowerBound(line_dists.data(), - lower_bound_dists.data(), - computer, - to_be_visited_id.data(), - count_no_visited, - ctx); + flatten->QueryWithDistanceLowerBoundAndFilterIP(line_dists.data(), + lower_bound_dists.data(), + filter_inner_products.data(), + computer, + to_be_visited_id.data(), + count_no_visited, + ctx); } else if (inner_search_param.enable_rabitq_one_bit_search) { flatten->QueryWithDistanceLowerBound(line_dists.data(), nullptr, @@ -735,9 +853,10 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, score_duplicates(duplicate_ids, hops); } if constexpr (mode == KNN_SEARCH) { - if (collect_rabitq_lower_bound and lower_bound_dists[i] < lower_bound and - check_func(cur_id)) { - rabitq_lower_bound_candidates->emplace_back(lower_bound_dists[i], cur_id); + if (collect_rabitq_lower_bound and + (top_candidates->Size() < ef or lower_bound_dists[i] < lower_bound)) { + append_lower_bound_candidates( + cur_id, lower_bound_dists[i], filter_inner_products[i]); } } if (top_candidates->Size() < ef || lower_bound > dist || @@ -749,24 +868,18 @@ BasicSearcher::search_impl(const GraphInterfacePtr& graph, } else if (reasoning != nullptr) { reasoning->RecordFilterReject(cur_id); } - if (inner_search_param.consider_duplicate and not use_custom_distance) { - const auto duplicate_ids = graph->GetDuplicateIds(cur_id); - for (const auto& item : duplicate_ids) { - if (check_func(item)) { - top_candidates->Push(dist, item); - } - } + if (not use_custom_distance) { + push_duplicate_candidates(cur_id, dist); } if constexpr (mode == KNN_SEARCH) { - if (top_candidates->Size() > ef) { + while (top_candidates->Size() > ef) { if (reasoning != nullptr) { reasoning->RecordEviction(top_candidates->Top().second, hops); } top_candidates->Pop(); } } - if (not top_candidates->Empty()) { lower_bound = top_candidates->Top().first; } diff --git a/src/impl/searcher/basic_searcher.h b/src/impl/searcher/basic_searcher.h index 8455f310f8..b0b977ac92 100644 --- a/src/impl/searcher/basic_searcher.h +++ b/src/impl/searcher/basic_searcher.h @@ -54,7 +54,7 @@ class BasicSearcher { const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; DistHeapPtr SearchWithPresetComputer(const GraphInterfacePtr& graph, @@ -64,7 +64,7 @@ class BasicSearcher { const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates, + RaBitQCandidateVector* rabitq_lower_bound_candidates, const ComputerInterfacePtr& preset_computer) const; virtual DistHeapPtr @@ -75,7 +75,7 @@ class BasicSearcher { const InnerSearchParam& inner_search_param, IteratorFilterContext* iter_ctx, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; DistHeapPtr Search(const GraphInterfacePtr& graph, @@ -123,7 +123,7 @@ class BasicSearcher { const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates, + RaBitQCandidateVector* rabitq_lower_bound_candidates, const ComputerInterfacePtr& preset_computer) const; template @@ -135,7 +135,7 @@ class BasicSearcher { const InnerSearchParam& inner_search_param, IteratorFilterContext* iter_ctx, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; template DistHeapPtr diff --git a/src/impl/searcher/hgraph_rabitq_searcher.cpp b/src/impl/searcher/hgraph_rabitq_searcher.cpp new file mode 100644 index 0000000000..5f95b06b95 --- /dev/null +++ b/src/impl/searcher/hgraph_rabitq_searcher.cpp @@ -0,0 +1,1620 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hgraph_rabitq_searcher.h" + +#include +#include +#include +#include +#include +#include + +#include "attr/executor/executor.h" +#include "datacell/rabitq_split_datacell.h" +#include "impl/heap/standard_heap.h" +#include "impl/reasoning/search_reasoning.h" +#include "index_common_param.h" +#include "simd/rabitq_simd.h" +#include "vsag/allocator.h" + +namespace vsag { +namespace { + +class MaybeSharedLock { +public: + MaybeSharedLock(const MutexArrayPtr& mutexes, InnerIdType id) : mutexes_(mutexes), id_(id) { + if (mutexes_ != nullptr) { + mutexes_->SharedLock(id_); + } + } + + ~MaybeSharedLock() { + if (mutexes_ != nullptr) { + mutexes_->SharedUnlock(id_); + } + } + +private: + const MutexArrayPtr& mutexes_; + InnerIdType id_; +}; + +constexpr uint32_t K_FUSED_CLUSTER_COUNT = 16; +constexpr uint32_t K_FULL_DISTANCE_AVAILABLE = 1U << 31U; +constexpr uint32_t K_CANDIDATE_INDEX_MASK = K_FULL_DISTANCE_AVAILABLE - 1U; +constexpr uint32_t K_INVALID_CANDIDATE_INDEX = K_CANDIDATE_INDEX_MASK; + +struct search_buffer_record { + float distance; + InnerIdType id; +}; + +struct bounded_result_record { + float distance; + float filter_inner_product; + InnerIdType id; + uint32_t state; + + [[nodiscard]] bool + FullDistanceAvailable() const { + return (state & K_FULL_DISTANCE_AVAILABLE) != 0; + } + + [[nodiscard]] uint32_t + CandidateIndex() const { + return state & K_CANDIDATE_INDEX_MASK; + } +}; + +static_assert(sizeof(bounded_result_record) == 16); + +/** + * Fixed-capacity sorted result buffer matching RaBitQ-Library's BoundedKNN. + * + * Keeping this concrete avoids virtual heap operations in the per-neighbor hot loop. The exact + * x-bit inner product travels with a deferred candidate, so 2/3/4+y rerank never rescans x. + */ +class BoundedResults { +public: + BoundedResults(uint64_t capacity, Allocator* allocator) + : data_(capacity + 1, allocator), capacity_(capacity) { + } + + void + Insert(InnerIdType id, + float distance, + float filter_inner_product, + bool full_distance_available = false, + uint32_t candidate_index = K_INVALID_CANDIDATE_INDEX) { + if (capacity_ == 0 or (size_ == capacity_ and distance > data_[size_ - 1].distance)) { + return; + } + const auto position = binary_search(distance); + std::memmove(data_.data() + position + 1, + data_.data() + position, + (size_ - position) * sizeof(bounded_result_record)); + data_[position] = { + distance, + filter_inner_product, + id, + candidate_index | (full_distance_available ? K_FULL_DISTANCE_AVAILABLE : 0U)}; + size_ += static_cast(size_ < capacity_); + } + + [[nodiscard]] uint64_t + Size() const { + return size_; + } + + [[nodiscard]] float + WorstDistance() const { + return size_ == capacity_ ? data_[size_ - 1].distance : std::numeric_limits::max(); + } + + [[nodiscard]] const bounded_result_record& + operator[](uint64_t index) const { + return data_[index]; + } + + [[nodiscard]] InnerIdType + WorstId() const { + return data_[size_ - 1].id; + } + +private: + [[nodiscard]] uint64_t + binary_search(float distance) const { + uint64_t lo = 0; + uint64_t length = size_; + while (length > 1) { + const auto half = length >> 1U; + length -= half; + lo += static_cast(data_[lo + half - 1].distance < distance) * half; + } + return lo < size_ and data_[lo].distance < distance ? lo + 1 : lo; + } + +private: + Vector data_; + uint64_t size_{0}; + uint64_t capacity_{0}; +}; + +inline float +read_float(const uint8_t* address) { + float value = 0.0F; + std::memcpy(&value, address, sizeof(value)); + return value; +} + +template +struct AffineFilterIP; + +template <> +struct AffineFilterIP<1> { + static float + Approximate(const RaBitQFusedTraversalQuery& query, const uint8_t* code) { + const uint64_t packed = RaBitQSQ4UBinaryIPWithBaseSum(query.query_planes, code, query.dim); + const auto raw_ip = static_cast(packed); + const auto base_sum = static_cast(packed >> 32U); + return query.query_delta * static_cast(raw_ip) + + query.query_vl * static_cast(base_sum) - 0.5F * query.query_sum; + } + + static float + Exact(const RaBitQFusedTraversalQuery& query, const uint8_t* code) { + return 0.5F * RaBitQFloatBinaryIP(query.transformed_query, code, query.dim, 1.0F); + } + + static void + ApproximateBatch4(const RaBitQFusedTraversalQuery& query, + const uint8_t* code0, + const uint8_t* code1, + const uint8_t* code2, + const uint8_t* code3, + float* results) { + uint64_t packed[4]; + RaBitQSQ4UBinaryIPWithBaseSumBatch4( + query.query_planes, code0, code1, code2, code3, query.dim, packed); + for (uint32_t i = 0; i < 4; ++i) { + const auto raw_ip = static_cast(packed[i]); + const auto base_sum = static_cast(packed[i] >> 32U); + results[i] = query.query_delta * static_cast(raw_ip) + + query.query_vl * static_cast(base_sum) - 0.5F * query.query_sum; + } + } +}; + +template <> +struct AffineFilterIP<2> { + static float + Exact(const RaBitQFusedTraversalQuery& query, const uint8_t* code) { + return RaBitQFloatTwoBitCenteredIP(query.transformed_query, code, query.dim); + } + + static void + ExactBatch4(const RaBitQFusedTraversalQuery& query, + const uint8_t* code0, + const uint8_t* code1, + const uint8_t* code2, + const uint8_t* code3, + float* results) { + RaBitQFloatTwoBitCenteredIPBatch4( + query.transformed_query, code0, code1, code2, code3, query.dim, results); + } +}; + +template <> +struct AffineFilterIP<3> { + static float + Exact(const RaBitQFusedTraversalQuery& query, const uint8_t* code) { + return RaBitQFloatThreeBitCenteredIP(query.transformed_query, code, query.dim); + } + + static void + ExactBatch4(const RaBitQFusedTraversalQuery& query, + const uint8_t* code0, + const uint8_t* code1, + const uint8_t* code2, + const uint8_t* code3, + float* results) { + RaBitQFloatThreeBitCenteredIPBatch4( + query.transformed_query, code0, code1, code2, code3, query.dim, results); + } +}; + +template <> +struct AffineFilterIP<4> { + static float + Exact(const RaBitQFusedTraversalQuery& query, const uint8_t* code) { + return RaBitQFloatFourBitCenteredIP(query.transformed_query, code, query.dim); + } + + static void + ExactBatch4(const RaBitQFusedTraversalQuery& query, + const uint8_t* code0, + const uint8_t* code1, + const uint8_t* code2, + const uint8_t* code3, + float* results) { + RaBitQFloatFourBitCenteredIPBatch4( + query.transformed_query, code0, code1, code2, code3, query.dim, results); + } +}; + +template +class AffineScorer { +public: + static constexpr bool K_EXACT_FILTER_HINT = filter_bits >= 2; + + AffineScorer(const RaBitQFusedTraversalQuery& query, float runtime_error_rate) : query_(query) { + const float error_rate = + IsFiniteRaBitQValue(runtime_error_rate) and runtime_error_rate > 0.0F + ? runtime_error_rate + : query.default_rabitq_error_rate; + const auto count = std::min(query.cluster_count, K_FUSED_CLUSTER_COUNT); + for (uint32_t i = 0; i < count; ++i) { + scaled_cluster_g_error_[i] = error_rate * query.cluster_g_error[i]; + } + } + + [[nodiscard]] bool + Estimate(const HGraphRaBitQFusedDataCell::CodeView& node, + float* distance, + float* lower_bound, + float* filter_inner_product) const { + if (node.cluster_id >= query_.cluster_count or + node.cluster_id >= scaled_cluster_g_error_.size()) { + return false; + } + if constexpr (filter_bits == 1) { + *filter_inner_product = AffineFilterIP<1>::Approximate(query_, node.one_bit_code); + } else { + *filter_inner_product = AffineFilterIP::Exact(query_, node.one_bit_code); + } + const auto* metadata = node.one_bit_code + query_.one_bit_metadata_offset; + const float filter_add = read_float(metadata); + const float filter_rescale = read_float(metadata + sizeof(float)); + const float filter_error_unit = read_float(metadata + 2U * sizeof(float)); + *distance = filter_add + query_.cluster_g_add[node.cluster_id] + + filter_rescale * *filter_inner_product; + const float raw_lower_bound = + *distance - filter_error_unit * scaled_cluster_g_error_[node.cluster_id]; + *lower_bound = raw_lower_bound - 1e-5F * std::max(1.0F, std::fabs(raw_lower_bound)); + return IsFiniteRaBitQValue(*distance) and IsFiniteRaBitQValue(*lower_bound) and + IsFiniteRaBitQValue(*filter_inner_product); + } + + void + EstimateBatch4(const HGraphRaBitQFusedDataCell::CodeView* nodes, + float* distances, + float* lower_bounds, + float* filter_inner_products, + bool* valid) const { + if constexpr (filter_bits == 1) { + AffineFilterIP<1>::ApproximateBatch4(query_, + nodes[0].one_bit_code, + nodes[1].one_bit_code, + nodes[2].one_bit_code, + nodes[3].one_bit_code, + filter_inner_products); + } else { + AffineFilterIP::ExactBatch4(query_, + nodes[0].one_bit_code, + nodes[1].one_bit_code, + nodes[2].one_bit_code, + nodes[3].one_bit_code, + filter_inner_products); + } + for (uint32_t i = 0; i < 4; ++i) { + valid[i] = nodes[i].cluster_id < query_.cluster_count and + nodes[i].cluster_id < scaled_cluster_g_error_.size(); + if (not valid[i]) { + continue; + } + const auto* metadata = nodes[i].one_bit_code + query_.one_bit_metadata_offset; + const float filter_add = read_float(metadata); + const float filter_rescale = read_float(metadata + sizeof(float)); + const float filter_error_unit = read_float(metadata + 2U * sizeof(float)); + distances[i] = filter_add + query_.cluster_g_add[nodes[i].cluster_id] + + filter_rescale * filter_inner_products[i]; + const float raw_lower_bound = + distances[i] - filter_error_unit * scaled_cluster_g_error_[nodes[i].cluster_id]; + lower_bounds[i] = raw_lower_bound - 1e-5F * std::max(1.0F, std::fabs(raw_lower_bound)); + valid[i] = IsFiniteRaBitQValue(distances[i]) and + IsFiniteRaBitQValue(lower_bounds[i]) and + IsFiniteRaBitQValue(filter_inner_products[i]); + } + } + + [[nodiscard]] bool + FullWithHint(const HGraphRaBitQFusedDataCell::CodeView& node, + float exact_filter_inner_product, + float* distance) const { + if constexpr (filter_bits == 1) { + (void)node; + (void)exact_filter_inner_product; + (void)distance; + return false; + } else { + return compute_full(node, exact_filter_inner_product, distance); + } + } + + [[nodiscard]] bool + FullDirect(const HGraphRaBitQFusedDataCell::CodeView& node, float* distance) const { + const float exact_filter_inner_product = + AffineFilterIP::Exact(query_, node.one_bit_code); + return compute_full(node, exact_filter_inner_product, distance); + } + + [[nodiscard]] bool + HasValidFilterDistance(const HGraphRaBitQFusedDataCell::CodeView& node, float distance) const { + return node.cluster_id < query_.cluster_count and + node.cluster_id < scaled_cluster_g_error_.size() and + IsFiniteRaBitQValue(distance) and distance < std::numeric_limits::max(); + } + +private: + [[nodiscard]] bool + compute_full(const HGraphRaBitQFusedDataCell::CodeView& node, + float exact_filter_inner_product, + float* distance) const { + if (node.cluster_id >= query_.cluster_count or + not IsFiniteRaBitQValue(exact_filter_inner_product)) { + return false; + } + const float supplement_ip = RaBitQFloatSupplementCodeIP( + query_.transformed_query, node.supplement_code, query_.dim, query_.supplement_bits); + const auto* metadata = node.supplement_code + query_.supplement_metadata_offset; + const float full_add = read_float(metadata); + const float full_rescale = read_float(metadata + sizeof(float)); + const float supplement_center = + 0.5F * static_cast((1U << query_.supplement_bits) - 1U); + const float full_inner_product = + static_cast(1U << query_.supplement_bits) * exact_filter_inner_product + + supplement_ip - supplement_center * query_.query_sum; + *distance = + full_add + query_.cluster_g_add[node.cluster_id] + full_rescale * full_inner_product; + return IsFiniteRaBitQValue(*distance); + } + +private: + const RaBitQFusedTraversalQuery& query_; + std::array scaled_cluster_g_error_{}; +}; + +class LegacyOneBitScorer { +public: + static constexpr bool K_EXACT_FILTER_HINT = false; + + LegacyOneBitScorer(const RaBitQFusedTraversalQuery& query, float runtime_error_rate) + : query_(query) { + const float error_rate = + IsFiniteRaBitQValue(runtime_error_rate) and runtime_error_rate > 0.0F + ? runtime_error_rate + : query.default_rabitq_error_rate; + const float error_rate_scale = + error_rate / RaBitQuantizerParameter::DEFAULT_RABITQ_ERROR_RATE; + const auto count = std::min(query.cluster_count, K_FUSED_CLUSTER_COUNT); + for (uint32_t i = 0; i < count; ++i) { + scaled_cluster_g_error_[i] = error_rate_scale * query.cluster_g_error[i]; + } + } + + [[nodiscard]] bool + Estimate(const HGraphRaBitQFusedDataCell::CodeView& node, + float* distance, + float* lower_bound, + float* filter_inner_product) const { + if (node.cluster_id >= query_.cluster_count or + node.cluster_id >= scaled_cluster_g_error_.size()) { + return false; + } + const uint64_t packed = + RaBitQSQ4UBinaryIPWithBaseSum(query_.query_planes, node.one_bit_code, query_.dim); + const auto raw_ip = static_cast(packed); + const auto base_sum = static_cast(packed >> 32U); + const float uncentered_ip = query_.query_delta * static_cast(raw_ip) + + query_.query_vl * static_cast(base_sum); + *filter_inner_product = uncentered_ip - 0.5F * query_.query_sum; + const auto* metadata = node.one_bit_code + query_.one_bit_metadata_offset; + const float filter_add = read_float(metadata); + const float filter_rescale = read_float(metadata + sizeof(float)); + const float filter_error = read_float(metadata + 2U * sizeof(float)); + *distance = filter_add + query_.cluster_g_add[node.cluster_id] + + filter_rescale * *filter_inner_product; + *lower_bound = *distance - filter_error * scaled_cluster_g_error_[node.cluster_id]; + return IsFiniteRaBitQValue(*distance) and IsFiniteRaBitQValue(*lower_bound) and + IsFiniteRaBitQValue(*filter_inner_product); + } + + void + EstimateBatch4(const HGraphRaBitQFusedDataCell::CodeView* nodes, + float* distances, + float* lower_bounds, + float* filter_inner_products, + bool* valid) const { + AffineFilterIP<1>::ApproximateBatch4(query_, + nodes[0].one_bit_code, + nodes[1].one_bit_code, + nodes[2].one_bit_code, + nodes[3].one_bit_code, + filter_inner_products); + for (uint32_t i = 0; i < 4; ++i) { + valid[i] = nodes[i].cluster_id < query_.cluster_count and + nodes[i].cluster_id < scaled_cluster_g_error_.size(); + if (not valid[i]) { + continue; + } + const auto* metadata = nodes[i].one_bit_code + query_.one_bit_metadata_offset; + const float filter_add = read_float(metadata); + const float filter_rescale = read_float(metadata + sizeof(float)); + const float filter_error = read_float(metadata + 2U * sizeof(float)); + distances[i] = filter_add + query_.cluster_g_add[nodes[i].cluster_id] + + filter_rescale * filter_inner_products[i]; + lower_bounds[i] = + distances[i] - filter_error * scaled_cluster_g_error_[nodes[i].cluster_id]; + valid[i] = IsFiniteRaBitQValue(distances[i]) and + IsFiniteRaBitQValue(lower_bounds[i]) and + IsFiniteRaBitQValue(filter_inner_products[i]); + } + } + + [[nodiscard]] static bool + FullWithHint(const HGraphRaBitQFusedDataCell::CodeView& /*node*/, + float /*filter_inner_product*/, + float* /*distance*/) { + return false; + } + + [[nodiscard]] bool + FullDirect(const HGraphRaBitQFusedDataCell::CodeView& node, float* distance) const { + if (node.cluster_id >= query_.cluster_count) { + return false; + } + const float exact_centered_ip = AffineFilterIP<1>::Exact(query_, node.one_bit_code); + const float supplement_ip = + (query_.dim & 63U) == 0U + ? RaBitQFloatExCode7IP(query_.transformed_query, node.supplement_code, query_.dim) + : RaBitQFloatSupplementCodeIP( + query_.transformed_query, node.supplement_code, query_.dim, 7); + const auto* metadata = node.supplement_code + query_.supplement_metadata_offset; + const float full_add = read_float(metadata); + const float full_rescale = read_float(metadata + sizeof(float)); + *distance = + full_add + query_.cluster_g_add[node.cluster_id] + + full_rescale * (128.0F * exact_centered_ip + supplement_ip - 63.5F * query_.query_sum); + return IsFiniteRaBitQValue(*distance); + } + + [[nodiscard]] bool + HasValidFilterDistance(const HGraphRaBitQFusedDataCell::CodeView& node, float distance) const { + return node.cluster_id < query_.cluster_count and + node.cluster_id < scaled_cluster_g_error_.size() and + IsFiniteRaBitQValue(distance) and distance < std::numeric_limits::max(); + } + +private: + const RaBitQFusedTraversalQuery& query_; + std::array scaled_cluster_g_error_{}; +}; + +/** + * Sorted linear beam buffer ported from RaBitQ-Library's SearchBuffer + * (Apache-2.0). VSAG keeps the checked flag separate because InnerIdType is + * also used by remove-version encoding. + */ +class SearchBuffer { +public: + SearchBuffer(uint64_t capacity, Allocator* allocator) + : data_(capacity + 1, allocator), capacity_(capacity) { + } + + void + Insert(InnerIdType id, float distance) { + if (IsFull(distance)) { + return; + } + const auto pos = binary_search(distance); + std::memmove(data_.data() + pos + 1, + data_.data() + pos, + (size_ - pos) * sizeof(search_buffer_record)); + data_[pos] = {distance, id}; + size_ += static_cast(size_ < capacity_); + current_ = std::min(current_, pos); + } + + [[nodiscard]] bool + HasNext() const { + return current_ < size_; + } + + InnerIdType + Pop() { + const auto id = data_[current_].id; + data_[current_].id |= K_CHECKED_MASK; + ++current_; + while (current_ < size_ and (data_[current_].id & K_CHECKED_MASK) != 0) { + ++current_; + } + return id; + } + + [[nodiscard]] InnerIdType + NextId() const { + return data_[current_].id; + } + + [[nodiscard]] bool + IsFull(float distance) const { + return size_ == capacity_ and distance > data_[size_ - 1].distance; + } + +private: + static constexpr InnerIdType K_CHECKED_MASK = InnerIdType{1} << (sizeof(InnerIdType) * 8 - 1); + + [[nodiscard]] uint64_t + binary_search(float distance) const { + uint64_t lo = 0; + uint64_t length = size_; + while (length > 1) { + const auto half = length >> 1U; + length -= half; + lo += static_cast(data_[lo + half - 1].distance < distance) * half; + } + return lo < size_ and data_[lo].distance < distance ? lo + 1 : lo; + } + +private: + Vector data_; + uint64_t size_{0}; + uint64_t current_{0}; + uint64_t capacity_{0}; +}; + +template +DistHeapPtr +search_direct_fused(const HGraphRaBitQFusedDataCellPtr& graph, + const VisitedListPtr& visited_list, + const InnerSearchParam& search_param, + QueryContext* ctx, + RaBitQCandidateVector* lower_bound_candidates, + Allocator* allocator, + const MutexArrayPtr& neighbors_mutex, + const Scorer& scorer) { + if (search_param.ef == 0 or not graph->CheckIdExists(search_param.ep)) { + return std::make_shared>(allocator, -1); + } + if (lower_bound_candidates != nullptr) { + lower_bound_candidates->clear(); + const uint64_t reserve_hint = + (static_cast(search_param.ef) + 4U) * graph->MaximumDegree() + 1U; + lower_bound_candidates->reserve(std::min(graph->TotalCount(), reserve_hint)); + } + + const uint64_t rerank_topk = static_cast(std::max( + 1, search_param.rerank_topk > 0 ? search_param.rerank_topk : search_param.topk)); + const bool should_rerank = + search_param.enable_rabitq_one_bit_search and search_param.enable_reorder; + const bool deferred_rerank = HGraphRaBitQSearcher::ShouldDeferRerank(search_param); + const uint64_t deferred_rerank_count = std::max(search_param.ef, 2 * rerank_topk); + BoundedResults results(deferred_rerank ? deferred_rerank_count : rerank_topk, allocator); + SearchBuffer candidate_set(search_param.ef, allocator); + + uint32_t rabitq_filter_count = 0; + uint32_t rabitq_full_count = 0; + uint32_t rabitq_filter_fallback_full_count = 0; + uint32_t rabitq_hint_full_count = 0; + uint32_t rabitq_reorder_fallback_full_count = 0; + uint32_t deferred_finalize_full_count = 0; + Filter* attribute_filter = nullptr; + if (not search_param.executors.empty() and search_param.executors[0] != nullptr) { + search_param.executors[0]->Clear(); + attribute_filter = search_param.executors[0]->Run(); + } + const auto is_allowed = [&search_param, attribute_filter](InnerIdType id) { + return (search_param.is_inner_id_allowed == nullptr or + search_param.is_inner_id_allowed->CheckValid(id)) and + (attribute_filter == nullptr or attribute_filter->CheckValid(id)); + }; + + const auto score_node = [&](const HGraphRaBitQFusedDataCell::CodeView& node, + float* distance, + float* lower_bound, + float* filter_inner_product, + bool* full_distance_available) { + *full_distance_available = false; + if (search_param.enable_rabitq_one_bit_search) { + ++rabitq_filter_count; + if (scorer.Estimate(node, distance, lower_bound, filter_inner_product)) { + return true; + } + if (not should_rerank) { + if (scorer.HasValidFilterDistance(node, *distance)) { + *lower_bound = *distance; + return true; + } + } + ++rabitq_filter_fallback_full_count; + } + ++rabitq_full_count; + *filter_inner_product = std::numeric_limits::quiet_NaN(); + *full_distance_available = scorer.FullDirect(node, distance); + if (not *full_distance_available) { + return false; + } + *lower_bound = *distance; + return true; + }; + const auto refine_node = [&](const HGraphRaBitQFusedDataCell::CodeView& node, + float filter_inner_product, + float* distance) { + ++rabitq_full_count; + if constexpr (Scorer::K_EXACT_FILTER_HINT) { + if (IsFiniteRaBitQValue(filter_inner_product)) { + if (scorer.FullWithHint(node, filter_inner_product, distance)) { + ++rabitq_hint_full_count; + return true; + } + } + } + ++rabitq_reorder_fallback_full_count; + return scorer.FullDirect(node, distance); + }; + const auto record_lower_bound = [&](InnerIdType id, + float lower_bound, + float filter_inner_product, + bool allowed, + float full_distance) { + if (lower_bound_candidates == nullptr or not allowed) { + return K_INVALID_CANDIDATE_INDEX; + } + const auto candidate_index = lower_bound_candidates->size(); + if (candidate_index >= K_INVALID_CANDIDATE_INDEX) { + return K_INVALID_CANDIDATE_INDEX; + } + lower_bound_candidates->push_back({lower_bound, + Scorer::K_EXACT_FILTER_HINT + ? filter_inner_product + : std::numeric_limits::quiet_NaN(), + id, + full_distance}); + return static_cast(candidate_index); + }; + + const auto entry_node = graph->GetCodeView(search_param.ep); + float entry_distance = 0.0F; + float entry_lower_bound = 0.0F; + float entry_filter_ip = std::numeric_limits::quiet_NaN(); + bool entry_full_distance_available = false; + if (not score_node(entry_node, + &entry_distance, + &entry_lower_bound, + &entry_filter_ip, + &entry_full_distance_available)) { + return std::make_shared>(allocator, -1); + } + const bool entry_allowed = is_allowed(search_param.ep); + // RaBitQ-Library inserts the bottom-layer entry point with its full distance. + if (should_rerank and not entry_full_distance_available) { + if (not refine_node(entry_node, entry_filter_ip, &entry_distance)) { + return std::make_shared>(allocator, -1); + } + entry_full_distance_available = true; + } + const auto entry_candidate_index = record_lower_bound( + search_param.ep, + entry_lower_bound, + entry_filter_ip, + entry_allowed, + entry_full_distance_available ? entry_distance : std::numeric_limits::quiet_NaN()); + candidate_set.Insert(search_param.ep, entry_distance); + visited_list->Set(search_param.ep); + if (entry_allowed) { + results.Insert( + search_param.ep, + entry_distance, + Scorer::K_EXACT_FILTER_HINT ? entry_filter_ip : std::numeric_limits::quiet_NaN(), + entry_full_distance_available, + entry_candidate_index); + } + + uint32_t hops = 0; + uint32_t distance_computations = 1; + auto* reasoning = ctx == nullptr ? nullptr : ctx->reasoning_ctx; + InnerIdType last_prefetched_candidate = std::numeric_limits::max(); + const auto prefetch_next_candidate = [&]() { + if (not candidate_set.HasNext()) { + return; + } + const auto next_id = candidate_set.NextId(); + if (next_id != last_prefetched_candidate) { + graph->PrefetchNodeHeader(next_id); + last_prefetched_candidate = next_id; + } + }; + const auto process_scored_node = [&](InnerIdType neighbor, + const HGraphRaBitQFusedDataCell::CodeView& node, + float distance, + float lower_bound, + float filter_inner_product, + bool supplement_prefetched, + bool full_distance_available) { + ++distance_computations; + if (reasoning != nullptr) { + reasoning->RecordVisit(neighbor, distance, hops); + } + + const bool allowed = is_allowed(neighbor); + if (deferred_rerank) { + const auto candidate_index = record_lower_bound( + neighbor, + lower_bound, + filter_inner_product, + allowed, + full_distance_available ? distance : std::numeric_limits::quiet_NaN()); + if (not candidate_set.IsFull(distance)) { + candidate_set.Insert(neighbor, distance); + } + if (allowed) { + results.Insert(neighbor, + distance, + filter_inner_product, + full_distance_available, + candidate_index); + } else if (reasoning != nullptr) { + reasoning->RecordFilterReject(neighbor); + } + prefetch_next_candidate(); + return; + } + + const bool lower_bound_promising = + results.Size() < rerank_topk or + (should_rerank ? lower_bound : distance) < results.WorstDistance(); + const bool promising = + results.Size() < rerank_topk or (should_rerank ? lower_bound < results.WorstDistance() + : distance < results.WorstDistance()); + if (promising and should_rerank and not full_distance_available) { + if (not supplement_prefetched) { + graph->PrefetchFusedSupplement(neighbor); + } + if (not refine_node(node, filter_inner_product, &distance)) { + return; + } + full_distance_available = true; + } + if (lower_bound_promising) { + record_lower_bound( + neighbor, + lower_bound, + filter_inner_product, + allowed, + full_distance_available ? distance : std::numeric_limits::quiet_NaN()); + } + if (not candidate_set.IsFull(distance)) { + candidate_set.Insert(neighbor, distance); + } + if (promising and allowed) { + results.Insert(neighbor, distance, filter_inner_product, full_distance_available); + if (search_param.consider_duplicate) { + for (const auto duplicate : graph->GetDuplicateIds(neighbor)) { + if (is_allowed(duplicate)) { + results.Insert(duplicate, + distance, + std::numeric_limits::quiet_NaN(), + full_distance_available); + } + } + } + } else if (not allowed and reasoning != nullptr) { + reasoning->RecordFilterReject(neighbor); + } + prefetch_next_candidate(); + }; + + while (candidate_set.HasNext()) { + ++hops; + if (hops >= search_param.hops_limit) { + if (reasoning != nullptr) { + reasoning->SetTermination(ReasoningContext::kTerminationHopsLimitReached); + } + break; + } + if (search_param.time_cost != nullptr and search_param.time_cost->CheckOvertime()) { + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->is_timeout.store(true, std::memory_order_relaxed); + } + if (reasoning != nullptr) { + reasoning->SetTermination(ReasoningContext::kTerminationTimeout); + } + break; + } + if (reasoning != nullptr) { + reasoning->AddSearchHop(); + } + + const auto current_id = candidate_set.Pop(); + MaybeSharedLock node_lock(neighbors_mutex, current_id); + const auto current_node = graph->GetNodeView(current_id); + const auto neighbor_count = current_node.neighbor_count; + if (neighbor_count > graph->MaximumDegree()) { + continue; + } + const auto* neighbors = current_node.neighbors; + uint32_t cursor = 0; + while (cursor < neighbor_count) { + InnerIdType batch_ids[4]{}; + uint32_t batch_count = 0; + while (cursor < neighbor_count and batch_count < 4) { + InnerIdType neighbor = 0; + const auto stored_neighbor = neighbors[cursor++]; + if (not graph->ResolveNeighbor(stored_neighbor, neighbor) or + visited_list->TestAndSet(neighbor)) { + continue; + } + batch_ids[batch_count] = neighbor; + graph->PrefetchFusedFilter(neighbor); + ++batch_count; + } + if (batch_count == 0) { + continue; + } + HGraphRaBitQFusedDataCell::CodeView batch_codes[4]{}; + for (uint32_t lane = 0; lane < batch_count; ++lane) { + batch_codes[lane] = graph->GetCodeView(batch_ids[lane]); + } + + float distances[4]{}; + float lower_bounds[4]{}; + float filter_inner_products[4]{std::numeric_limits::quiet_NaN(), + std::numeric_limits::quiet_NaN(), + std::numeric_limits::quiet_NaN(), + std::numeric_limits::quiet_NaN()}; + bool valid[4]{}; + bool full_distance_available[4]{}; + if (batch_count == 4 and search_param.enable_rabitq_one_bit_search) { + rabitq_filter_count += 4; + scorer.EstimateBatch4( + batch_codes, distances, lower_bounds, filter_inner_products, valid); + for (uint32_t lane = 0; lane < 4; ++lane) { + if (not valid[lane]) { + if (not should_rerank) { + if (scorer.HasValidFilterDistance(batch_codes[lane], distances[lane])) { + valid[lane] = true; + lower_bounds[lane] = distances[lane]; + continue; + } + } + ++rabitq_filter_fallback_full_count; + ++rabitq_full_count; + filter_inner_products[lane] = std::numeric_limits::quiet_NaN(); + valid[lane] = scorer.FullDirect(batch_codes[lane], distances + lane); + full_distance_available[lane] = valid[lane]; + lower_bounds[lane] = distances[lane]; + } + } + } else { + for (uint32_t lane = 0; lane < batch_count; ++lane) { + valid[lane] = score_node(batch_codes[lane], + distances + lane, + lower_bounds + lane, + filter_inner_products + lane, + full_distance_available + lane); + } + } + + bool supplement_prefetched[4]{}; + if (not deferred_rerank and should_rerank) { + const bool result_not_full = results.Size() < rerank_topk; + const float bound_snapshot = results.WorstDistance(); + for (uint32_t lane = 0; lane < batch_count; ++lane) { + if (valid[lane] and not full_distance_available[lane] and + (result_not_full or lower_bounds[lane] < bound_snapshot)) { + graph->PrefetchFusedSupplement(batch_ids[lane]); + supplement_prefetched[lane] = true; + } + } + } + for (uint32_t lane = 0; lane < batch_count; ++lane) { + if (valid[lane]) { + process_scored_node(batch_ids[lane], + batch_codes[lane], + distances[lane], + lower_bounds[lane], + filter_inner_products[lane], + supplement_prefetched[lane], + full_distance_available[lane]); + } + } + } + } + + if (deferred_rerank) { + BoundedResults refined(rerank_topk, allocator); + const auto has_valid_full_distance = [](float distance) { + return IsFiniteRaBitQValue(distance) and distance < std::numeric_limits::max(); + }; + constexpr uint32_t k_deferred_prefetch = 4; + const auto initial_prefetch = std::min(k_deferred_prefetch, results.Size()); + for (uint64_t i = 0; i < initial_prefetch; ++i) { + if (not results[i].FullDistanceAvailable()) { + graph->PrefetchFusedSupplement(results[i].id); + } + } + for (uint64_t i = 0; i < results.Size(); ++i) { + const auto lookahead = i + k_deferred_prefetch; + if (lookahead < results.Size() and not results[lookahead].FullDistanceAvailable()) { + graph->PrefetchFusedSupplement(results[lookahead].id); + } + const auto& candidate = results[i]; + float full_distance = candidate.distance; + bool full_distance_available = candidate.FullDistanceAvailable(); + if (not full_distance_available) { + const auto node = graph->GetCodeView(candidate.id); + full_distance_available = + refine_node(node, candidate.filter_inner_product, &full_distance); + } + if (candidate.CandidateIndex() < + (lower_bound_candidates == nullptr ? 0 : lower_bound_candidates->size())) { + (*lower_bound_candidates)[candidate.CandidateIndex()].full_distance = + full_distance_available ? full_distance : std::numeric_limits::max(); + } else if (full_distance_available) { + refined.Insert(candidate.id, + full_distance, + Scorer::K_EXACT_FILTER_HINT + ? candidate.filter_inner_product + : std::numeric_limits::quiet_NaN(), + true); + } + } + + if (lower_bound_candidates != nullptr) { + const auto insert_refined = [&](const RaBitQCandidateRecord& candidate, + float full_distance) { + if (reasoning != nullptr) { + reasoning->RecordReorder(candidate.id, candidate.lower_bound, full_distance); + } + const bool will_evict = + refined.Size() == rerank_topk and full_distance <= refined.WorstDistance(); + const auto evicted_id = will_evict ? refined.WorstId() : 0; + refined.Insert(candidate.id, full_distance, candidate.filter_inner_product, true); + if (will_evict and reasoning != nullptr) { + reasoning->RecordReorderEviction(evicted_id, 0); + } + }; + // Establish an exact kth threshold from every full distance already computed for the + // shortlist or by a filter fallback before considering any missing supplement. + for (const auto& candidate : *lower_bound_candidates) { + if (has_valid_full_distance(candidate.full_distance)) { + insert_refined(candidate, candidate.full_distance); + } + } + for (auto& candidate : *lower_bound_candidates) { + if (has_valid_full_distance(candidate.full_distance) or + candidate.full_distance == std::numeric_limits::max()) { + continue; + } + const bool promising = refined.Size() < rerank_topk or + not IsFiniteRaBitQValue(candidate.lower_bound) or + candidate.lower_bound < refined.WorstDistance(); + if (not promising) { + continue; + } + graph->PrefetchFusedSupplement(candidate.id); + const auto node = graph->GetCodeView(candidate.id); + float full_distance = 0.0F; + ++deferred_finalize_full_count; + if (not refine_node(node, candidate.filter_inner_product, &full_distance)) { + candidate.full_distance = std::numeric_limits::max(); + continue; + } + candidate.full_distance = full_distance; + insert_refined(candidate, full_distance); + } + } + results = std::move(refined); + } + + auto output = std::make_shared>(allocator, -1); + for (uint64_t i = 0; i < results.Size(); ++i) { + output->Push(results[i].distance, results[i].id); + } + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->dist_cmp.fetch_add(distance_computations, std::memory_order_relaxed); + ctx->stats->hops.fetch_add(hops, std::memory_order_relaxed); + ctx->stats->rabitq_filter_count.fetch_add(rabitq_filter_count, std::memory_order_relaxed); + ctx->stats->rabitq_full_count.fetch_add(rabitq_full_count, std::memory_order_relaxed); + ctx->stats->rabitq_filter_fallback_full_count.fetch_add(rabitq_filter_fallback_full_count, + std::memory_order_relaxed); + ctx->stats->rabitq_reorder_hint_full_count.fetch_add(rabitq_hint_full_count, + std::memory_order_relaxed); + ctx->stats->rabitq_reorder_fallback_full_count.fetch_add(rabitq_reorder_fallback_full_count, + std::memory_order_relaxed); + ctx->stats->reorder_distance_count.fetch_add(deferred_finalize_full_count, + std::memory_order_relaxed); + } + return output; +} + +template +InnerIdType +route_direct_fused(const GraphInterfacePtr& route_graph, + const HGraphRaBitQFusedDataCellPtr& fused_graph, + InnerIdType entry_point, + Allocator* allocator, + const MutexArrayPtr& neighbors_mutex, + const Scorer& scorer) { + const auto score = [&fused_graph, &scorer](InnerIdType id, float* distance) { + const auto node = fused_graph->GetCodeView(id); + float lower_bound = 0.0F; + float filter_inner_product = 0.0F; + if (scorer.Estimate(node, distance, &lower_bound, &filter_inner_product) or + scorer.HasValidFilterDistance(node, *distance)) { + return true; + } + return scorer.FullDirect(node, distance); + }; + float current_distance = 0.0F; + if (not score(entry_point, ¤t_distance)) { + return entry_point; + } + Vector neighbors(allocator); + bool changed = true; + while (changed) { + changed = false; + { + MaybeSharedLock node_lock(neighbors_mutex, entry_point); + route_graph->GetNeighbors(entry_point, neighbors); + } + for (const auto neighbor : neighbors) { + float distance = 0.0F; + if (score(neighbor, &distance) and distance < current_distance) { + current_distance = distance; + entry_point = neighbor; + changed = true; + } + } + } + return entry_point; +} + +} // namespace + +HGraphRaBitQSearcher::HGraphRaBitQSearcher(const IndexCommonParam& common_param, + MutexArrayPtr neighbors_mutex) + : allocator_(common_param.allocator_.get()), neighbors_mutex_(std::move(neighbors_mutex)) { +} + +InnerIdType +HGraphRaBitQSearcher::Route(const GraphInterfacePtr& route_graph, + const HGraphRaBitQFusedDataCellPtr& fused_graph, + const FlattenInterfacePtr& flatten, + const ComputerInterfacePtr& computer, + InnerIdType entry_point, + bool enable_one_bit_search) const { + auto* split_codes = dynamic_cast(flatten.get()); + if (route_graph == nullptr or fused_graph == nullptr or split_codes == nullptr or + computer == nullptr or not fused_graph->CheckIdExists(entry_point)) { + return entry_point; + } + RaBitQFusedTraversalQuery traversal_query; + if (enable_one_bit_search and split_codes->GetFusedTraversalQuery(computer, &traversal_query)) { + if (not traversal_query.affine) { + LegacyOneBitScorer scorer(traversal_query, std::numeric_limits::quiet_NaN()); + return route_direct_fused( + route_graph, fused_graph, entry_point, allocator_, neighbors_mutex_, scorer); + } + const float no_runtime_override = std::numeric_limits::quiet_NaN(); + switch (traversal_query.filter_bits) { + case 1: { + AffineScorer<1> scorer(traversal_query, no_runtime_override); + return route_direct_fused( + route_graph, fused_graph, entry_point, allocator_, neighbors_mutex_, scorer); + } + case 2: { + AffineScorer<2> scorer(traversal_query, no_runtime_override); + return route_direct_fused( + route_graph, fused_graph, entry_point, allocator_, neighbors_mutex_, scorer); + } + case 3: { + AffineScorer<3> scorer(traversal_query, no_runtime_override); + return route_direct_fused( + route_graph, fused_graph, entry_point, allocator_, neighbors_mutex_, scorer); + } + case 4: { + AffineScorer<4> scorer(traversal_query, no_runtime_override); + return route_direct_fused( + route_graph, fused_graph, entry_point, allocator_, neighbors_mutex_, scorer); + } + default: + break; + } + } + const auto score = [&](InnerIdType id, float* distance) { + const auto node = fused_graph->GetNodeView(id); + *distance = std::numeric_limits::max(); + if (not enable_one_bit_search) { + return split_codes->ComputeFusedFull(computer, + node.cluster_id, + node.one_bit_code, + node.supplement_code, + distance, + nullptr); + } + float lower_bound = 0.0F; + float filter_inner_product = 0.0F; + if (split_codes->ComputeFusedOneBitWithFilterIP(computer, + node.cluster_id, + node.one_bit_code, + node.supplement_code, + distance, + &lower_bound, + &filter_inner_product, + nullptr) or + (IsFiniteRaBitQValue(*distance) and *distance < std::numeric_limits::max())) { + return true; + } + return split_codes->ComputeFusedFull( + computer, node.cluster_id, node.one_bit_code, node.supplement_code, distance, nullptr); + }; + + float current_distance = 0.0F; + if (not score(entry_point, ¤t_distance)) { + return entry_point; + } + Vector neighbors(allocator_); + bool changed = true; + while (changed) { + changed = false; + { + MaybeSharedLock node_lock(neighbors_mutex_, entry_point); + route_graph->GetNeighbors(entry_point, neighbors); + } + for (const auto neighbor : neighbors) { + float distance = 0.0F; + if (score(neighbor, &distance) and distance < current_distance) { + current_distance = distance; + entry_point = neighbor; + changed = true; + } + } + } + return entry_point; +} + +DistHeapPtr +HGraphRaBitQSearcher::Search(const HGraphRaBitQFusedDataCellPtr& graph, + const FlattenInterfacePtr& flatten, + const VisitedListPtr& visited_list, + const void* query, + const InnerSearchParam& search_param, + QueryContext* ctx, + RaBitQCandidateVector* lower_bound_candidates, + bool* search_finalized) const { + if (search_finalized != nullptr) { + *search_finalized = false; + } + auto* allocator = select_query_allocator(ctx, allocator_); + auto* split_codes = dynamic_cast(flatten.get()); + if (graph == nullptr or split_codes == nullptr or search_param.find_duplicate or + search_param.consider_duplicate) { + return nullptr; + } + if (search_param.ef == 0 or not graph->CheckIdExists(search_param.ep)) { + return std::make_shared>(allocator, -1); + } + + if (lower_bound_candidates != nullptr) { + lower_bound_candidates->clear(); + } + const uint64_t rerank_topk = static_cast(std::max( + 1, search_param.rerank_topk > 0 ? search_param.rerank_topk : search_param.topk)); + const bool should_rerank = + search_param.enable_rabitq_one_bit_search and search_param.enable_reorder; + const bool deferred_rerank = HGraphRaBitQSearcher::ShouldDeferRerank(search_param); + const bool exact_filter_ip_hint = split_codes->FusedFilterBits() >= 2; + const uint64_t deferred_rerank_count = std::max(search_param.ef, 2 * rerank_topk); + auto computer = search_param.rabitq_fused_computer != nullptr + ? search_param.rabitq_fused_computer + : split_codes->FactoryFusedComputer(query); + RaBitQFusedTraversalQuery traversal_query; + const bool has_direct_traversal_query = + split_codes->GetFusedTraversalQuery(computer, &traversal_query); + const float runtime_error_rate = + ctx == nullptr ? std::numeric_limits::quiet_NaN() : ctx->rabitq_error_rate; + if (has_direct_traversal_query) { + const auto run_direct_search = [&](const auto& scorer) { + auto result = search_direct_fused(graph, + visited_list, + search_param, + ctx, + lower_bound_candidates, + allocator, + neighbors_mutex_, + scorer); + if (search_finalized != nullptr) { + *search_finalized = not deferred_rerank or lower_bound_candidates != nullptr; + } + return result; + }; + if (not traversal_query.affine) { + LegacyOneBitScorer scorer(traversal_query, runtime_error_rate); + return run_direct_search(scorer); + } + switch (traversal_query.filter_bits) { + case 1: { + AffineScorer<1> scorer(traversal_query, runtime_error_rate); + return run_direct_search(scorer); + } + case 2: { + AffineScorer<2> scorer(traversal_query, runtime_error_rate); + return run_direct_search(scorer); + } + case 3: { + AffineScorer<3> scorer(traversal_query, runtime_error_rate); + return run_direct_search(scorer); + } + case 4: { + AffineScorer<4> scorer(traversal_query, runtime_error_rate); + return run_direct_search(scorer); + } + default: + break; + } + } + auto result = std::make_shared>(allocator, -1); + SearchBuffer candidate_set(search_param.ef, allocator); + uint32_t rabitq_filter_count = 0; + uint32_t rabitq_full_count = 0; + uint32_t rabitq_filter_fallback_full_count = 0; + uint32_t rabitq_hint_full_count = 0; + uint32_t rabitq_reorder_fallback_full_count = 0; + QueryContext rate_context; + QueryContext* rate_context_ptr = nullptr; + if (ctx != nullptr) { + rate_context.rabitq_error_rate = ctx->rabitq_error_rate; + rate_context_ptr = &rate_context; + } + Filter* attribute_filter = nullptr; + if (not search_param.executors.empty() and search_param.executors[0] != nullptr) { + search_param.executors[0]->Clear(); + attribute_filter = search_param.executors[0]->Run(); + } + const auto is_allowed = [&search_param, attribute_filter](InnerIdType id) { + return (search_param.is_inner_id_allowed == nullptr or + search_param.is_inner_id_allowed->CheckValid(id)) and + (attribute_filter == nullptr or attribute_filter->CheckValid(id)); + }; + + auto score_node = [&](InnerIdType id, + float* distance, + float* lower_bound, + float* filter_inner_product) { + const auto node = graph->GetCodeView(id); + *distance = std::numeric_limits::max(); + *lower_bound = std::numeric_limits::max(); + *filter_inner_product = std::numeric_limits::quiet_NaN(); + if (search_param.enable_rabitq_one_bit_search) { + ++rabitq_filter_count; + if (has_direct_traversal_query and node.cluster_id < traversal_query.cluster_count) { + const uint64_t packed_ip = RaBitQSQ4UBinaryIPWithBaseSum( + traversal_query.query_planes, node.one_bit_code, traversal_query.dim); + const auto raw_ip = static_cast(packed_ip); + const auto base_sum = static_cast(packed_ip >> 32U); + *filter_inner_product = traversal_query.query_delta * static_cast(raw_ip) + + traversal_query.query_vl * static_cast(base_sum); + + float f_add = 0.0F; + float f_rescale = 0.0F; + float f_error = 0.0F; + const auto* metadata = node.one_bit_code + traversal_query.one_bit_metadata_offset; + std::memcpy(&f_add, metadata, sizeof(float)); + std::memcpy(&f_rescale, metadata + sizeof(float), sizeof(float)); + std::memcpy(&f_error, metadata + 2 * sizeof(float), sizeof(float)); + *distance = f_add + traversal_query.cluster_g_add[node.cluster_id] + + f_rescale * (*filter_inner_product - 0.5F * traversal_query.query_sum); + const float effective_error_rate = + IsFiniteRaBitQValue(runtime_error_rate) and runtime_error_rate > 0.0F + ? runtime_error_rate + : traversal_query.default_rabitq_error_rate; + const float error_rate_scale = + effective_error_rate / RaBitQuantizerParameter::DEFAULT_RABITQ_ERROR_RATE; + *lower_bound = *distance - error_rate_scale * f_error * + traversal_query.cluster_g_error[node.cluster_id]; + if (IsFiniteRaBitQValue(*distance) and + (not should_rerank or IsFiniteRaBitQValue(*lower_bound))) { + if (not should_rerank) { + *lower_bound = *distance; + } + return true; + } + } + if (split_codes->ComputeFusedOneBitWithFilterIP(computer, + node.cluster_id, + node.one_bit_code, + node.supplement_code, + distance, + lower_bound, + filter_inner_product, + rate_context_ptr)) { + return true; + } + if (not should_rerank and IsFiniteRaBitQValue(*distance) and + *distance < std::numeric_limits::max()) { + *lower_bound = *distance; + return true; + } + ++rabitq_filter_fallback_full_count; + ++rabitq_full_count; + *filter_inner_product = std::numeric_limits::quiet_NaN(); + const bool computed = split_codes->ComputeFusedFull(computer, + node.cluster_id, + node.one_bit_code, + node.supplement_code, + distance, + nullptr); + *lower_bound = *distance; + return computed; + } + ++rabitq_full_count; + *lower_bound = 0.0F; + *filter_inner_product = std::numeric_limits::quiet_NaN(); + const bool computed = split_codes->ComputeFusedFull( + computer, node.cluster_id, node.one_bit_code, node.supplement_code, distance, nullptr); + *lower_bound = *distance; + return computed; + }; + + const auto refine_node = [&](InnerIdType id, float, float* distance) { + const auto node = graph->GetCodeView(id); + ++rabitq_full_count; + // RaBitQ-Library's HNSW adaptive-rerank path calls + // split_single_fulldist_direct(), which recomputes the x-bit/query inner product + // from the float query. Reusing the 4-bit traversal estimate here changes that + // reference behavior and materially degrades recall on GIST1M. + ++rabitq_reorder_fallback_full_count; + return split_codes->ComputeFusedFull( + computer, node.cluster_id, node.one_bit_code, node.supplement_code, distance, nullptr); + }; + + const auto refine_node_with_hint = [&](InnerIdType id, + float filter_inner_product, + float* distance) { + const auto node = graph->GetCodeView(id); + ++rabitq_full_count; + if (exact_filter_ip_hint and IsFiniteRaBitQValue(filter_inner_product) and + split_codes->ComputeFusedFullWithFilterIP(computer, + node.cluster_id, + node.one_bit_code, + node.supplement_code, + filter_inner_product, + distance, + nullptr)) { + ++rabitq_hint_full_count; + return true; + } + ++rabitq_reorder_fallback_full_count; + return split_codes->ComputeFusedFull( + computer, node.cluster_id, node.one_bit_code, node.supplement_code, distance, nullptr); + }; + + float entry_distance = 0.0F; + float entry_lower_bound = 0.0F; + float entry_filter_ip = 0.0F; + if (not score_node(search_param.ep, &entry_distance, &entry_lower_bound, &entry_filter_ip)) { + return result; + } + if (should_rerank and not deferred_rerank) { + const bool refined = + exact_filter_ip_hint + ? refine_node_with_hint(search_param.ep, entry_filter_ip, &entry_distance) + : refine_node(search_param.ep, entry_filter_ip, &entry_distance); + if (not refined) { + return result; + } + } + candidate_set.Insert(search_param.ep, entry_distance); + visited_list->Set(search_param.ep); + if (is_allowed(search_param.ep)) { + result->Push(entry_distance, search_param.ep); + } + + uint32_t hops = 0; + uint32_t distance_computations = 1; + auto* reasoning = ctx == nullptr ? nullptr : ctx->reasoning_ctx; + const uint32_t prefetch_lookahead = [split_codes]() { + switch (split_codes->FusedFilterBits()) { + case 1: + return 8U; + case 2: + return 5U; + case 3: + return 4U; + default: + return 3U; + } + }(); + const auto prefetch_codes_l1 = [&graph](InnerIdType id) { graph->PrefetchFusedFilter(id); }; + const auto prefetch_graph_l2 = [&graph](InnerIdType id) { graph->PrefetchNodeHeader(id); }; + + while (candidate_set.HasNext()) { + ++hops; + if (hops >= search_param.hops_limit) { + if (reasoning != nullptr) { + reasoning->SetTermination(ReasoningContext::kTerminationHopsLimitReached); + } + break; + } + if (search_param.time_cost != nullptr and search_param.time_cost->CheckOvertime()) { + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->is_timeout.store(true, std::memory_order_relaxed); + } + if (reasoning != nullptr) { + reasoning->SetTermination(ReasoningContext::kTerminationTimeout); + } + break; + } + if (reasoning != nullptr) { + reasoning->AddSearchHop(); + } + + const auto current_id = candidate_set.Pop(); + MaybeSharedLock node_lock(neighbors_mutex_, current_id); + const auto node = graph->GetNodeView(current_id); + const auto neighbor_count = node.neighbor_count; + if (neighbor_count > graph->MaximumDegree()) { + continue; + } + const auto* neighbors = node.neighbors; + const auto initial_prefetch = std::min(prefetch_lookahead, neighbor_count); + for (uint32_t i = 0; i < initial_prefetch; ++i) { + InnerIdType neighbor = 0; + if (graph->ResolveNeighbor(neighbors[i], neighbor)) { + prefetch_codes_l1(neighbor); + } + } + + for (uint32_t i = 0; i < neighbor_count; ++i) { + if (i + prefetch_lookahead < neighbor_count) { + InnerIdType lookahead_neighbor = 0; + if (graph->ResolveNeighbor(neighbors[i + prefetch_lookahead], lookahead_neighbor)) { + prefetch_codes_l1(lookahead_neighbor); + } + } + InnerIdType neighbor = 0; + if (not graph->ResolveNeighbor(neighbors[i], neighbor)) { + continue; + } + if (visited_list->TestAndSet(neighbor)) { + continue; + } + + float distance = 0.0F; + float lower_bound = 0.0F; + float filter_ip = 0.0F; + if (not score_node(neighbor, &distance, &lower_bound, &filter_ip)) { + continue; + } + ++distance_computations; + if (reasoning != nullptr) { + reasoning->RecordVisit(neighbor, distance, hops); + } + + const bool allowed = is_allowed(neighbor); + if (deferred_rerank) { + if (not candidate_set.IsFull(distance)) { + candidate_set.Insert(neighbor, distance); + } + if (allowed) { + result->Push(distance, neighbor); + while (result->Size() > deferred_rerank_count) { + result->Pop(); + } + } else if (reasoning != nullptr) { + reasoning->RecordFilterReject(neighbor); + } + if (candidate_set.HasNext()) { + prefetch_graph_l2(candidate_set.NextId()); + } + continue; + } + const auto result_bound = result->Size() < rerank_topk + ? std::numeric_limits::max() + : result->Top().first; + const bool promising = + result->Size() < rerank_topk or + (should_rerank ? lower_bound < result_bound : distance < result_bound); + if (promising and should_rerank) { + graph->PrefetchFusedSupplement(neighbor); + const bool refined = exact_filter_ip_hint + ? refine_node_with_hint(neighbor, filter_ip, &distance) + : refine_node(neighbor, filter_ip, &distance); + if (not refined) { + continue; + } + } + + if (not candidate_set.IsFull(distance)) { + candidate_set.Insert(neighbor, distance); + } + if (promising and allowed) { + result->Push(distance, neighbor); + if (search_param.consider_duplicate) { + for (const auto duplicate : graph->GetDuplicateIds(neighbor)) { + if (is_allowed(duplicate)) { + result->Push(distance, duplicate); + } + } + } + while (result->Size() > rerank_topk) { + result->Pop(); + } + } else if (not allowed and reasoning != nullptr) { + reasoning->RecordFilterReject(neighbor); + } + if (candidate_set.HasNext()) { + prefetch_graph_l2(candidate_set.NextId()); + } + } + } + + if (deferred_rerank) { + auto refined = std::make_shared>(allocator, -1); + while (not result->Empty()) { + const auto candidate = result->Top(); + result->Pop(); + float full_distance = candidate.first; + float lower_bound = 0.0F; + float filter_inner_product = 0.0F; + const bool scored = + score_node(candidate.second, &full_distance, &lower_bound, &filter_inner_product); + bool refined_ok = false; + if (scored) { + refined_ok = + exact_filter_ip_hint + ? refine_node_with_hint( + candidate.second, filter_inner_product, &full_distance) + : refine_node(candidate.second, filter_inner_product, &full_distance); + } + if (refined_ok) { + refined->Push(full_distance, candidate.second); + while (refined->Size() > rerank_topk) { + refined->Pop(); + } + } + } + result = std::move(refined); + } + while (result->Size() > rerank_topk) { + result->Pop(); + } + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->dist_cmp.fetch_add(distance_computations, std::memory_order_relaxed); + ctx->stats->hops.fetch_add(hops, std::memory_order_relaxed); + ctx->stats->rabitq_filter_count.fetch_add(rabitq_filter_count, std::memory_order_relaxed); + ctx->stats->rabitq_full_count.fetch_add(rabitq_full_count, std::memory_order_relaxed); + ctx->stats->rabitq_filter_fallback_full_count.fetch_add(rabitq_filter_fallback_full_count, + std::memory_order_relaxed); + ctx->stats->rabitq_reorder_hint_full_count.fetch_add(rabitq_hint_full_count, + std::memory_order_relaxed); + ctx->stats->rabitq_reorder_fallback_full_count.fetch_add(rabitq_reorder_fallback_full_count, + std::memory_order_relaxed); + } + return result; +} + +} // namespace vsag diff --git a/src/impl/searcher/hgraph_rabitq_searcher.h b/src/impl/searcher/hgraph_rabitq_searcher.h new file mode 100644 index 0000000000..19696f3ad8 --- /dev/null +++ b/src/impl/searcher/hgraph_rabitq_searcher.h @@ -0,0 +1,69 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "datacell/flatten_interface.h" +#include "datacell/hgraph_rabitq_fused_datacell.h" +#include "impl/heap/distance_heap.h" +#include "impl/inner_search_param.h" +#include "index_common_param_fwd.h" +#include "query_context.h" +#include "utils/lock_strategy.h" +#include "utils/visited_list.h" + +namespace vsag { + +class HGraphRaBitQSearcher { +public: + explicit HGraphRaBitQSearcher(const IndexCommonParam& common_param, + MutexArrayPtr neighbors_mutex); + + static constexpr uint64_t K_DEFERRED_RERANK_MAX_EF = 40; + + [[nodiscard]] static bool + ShouldDeferRerank(const InnerSearchParam& search_param) { + return search_param.enable_rabitq_one_bit_search and search_param.enable_reorder and + search_param.ef <= K_DEFERRED_RERANK_MAX_EF; + } + + void + SetMutexArray(const MutexArrayPtr& neighbors_mutex) { + neighbors_mutex_ = neighbors_mutex; + } + + DistHeapPtr + Search(const HGraphRaBitQFusedDataCellPtr& graph, + const FlattenInterfacePtr& flatten, + const VisitedListPtr& visited_list, + const void* query, + const InnerSearchParam& search_param, + QueryContext* ctx, + RaBitQCandidateVector* lower_bound_candidates, + bool* search_finalized = nullptr) const; + + InnerIdType + Route(const GraphInterfacePtr& route_graph, + const HGraphRaBitQFusedDataCellPtr& fused_graph, + const FlattenInterfacePtr& flatten, + const ComputerInterfacePtr& computer, + InnerIdType entry_point, + bool enable_one_bit_search) const; + +private: + Allocator* allocator_{nullptr}; + MutexArrayPtr neighbors_mutex_{nullptr}; +}; + +} // namespace vsag diff --git a/src/impl/searcher/hgraph_rabitq_searcher_test.cpp b/src/impl/searcher/hgraph_rabitq_searcher_test.cpp new file mode 100644 index 0000000000..1fc847e1f1 --- /dev/null +++ b/src/impl/searcher/hgraph_rabitq_searcher_test.cpp @@ -0,0 +1,597 @@ +// Copyright 2024-present the vsag project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hgraph_rabitq_searcher.h" + +#include +#include +#include +#include + +#include "datacell/flatten_datacell.h" +#include "datacell/flatten_datacell_parameter.h" +#include "datacell/hgraph_rabitq_fused_datacell.h" +#include "datacell/rabitq_split_datacell.h" +#include "impl/allocator/safe_allocator.h" +#include "impl/filter/white_list_filter.h" +#include "impl/heap/standard_heap.h" +#include "impl/reasoning/search_reasoning.h" +#include "impl/reorder/flatten_reorder.h" +#include "index_common_param.h" +#include "io/memory_io/memory_io_parameter.h" +#include "unittest.h" + +namespace vsag { + +TEST_CASE("HGraph RaBitQ route honors one-bit search switch", "[ut][HGraphRaBitQSearcher][route]") { + auto allocator = SafeAllocator::FactoryDefaultAllocator(); + constexpr uint64_t dim = 64; + constexpr InnerIdType count = 64; + constexpr uint32_t cluster_count = 16; + auto vectors = fixtures::generate_vectors(count, dim, false, 73); + + auto param_json = JsonType::Parse(R"({ + "codes_type": "rabitq_split", + "io_params": {"type": "memory_io"}, + "quantization_params": { + "type": "rabitq", + "rabitq_version": "split", + "rabitq_bits_per_dim_query": 32, + "rabitq_bits_per_dim_base": 8, + "rabitq_bits_per_dim_filter": 2, + "use_fht": true + } + })"); + auto param = std::make_shared(); + param->FromJson(param_json); + IndexCommonParam common_param; + common_param.allocator_ = allocator; + common_param.dim_ = dim; + common_param.metric_ = MetricType::METRIC_TYPE_L2SQR; + + auto flatten = FlattenInterface::MakeInstance(param, common_param); + flatten->Train(vectors.data(), count); + auto split = std::dynamic_pointer_cast(flatten); + REQUIRE(split != nullptr); + split->TrainFusedCodec(vectors.data(), count, cluster_count); + auto computer = split->FactoryFusedComputer(vectors.data()); + REQUIRE(computer != nullptr); + + struct EncodedNode { + std::vector filter; + std::vector supplement; + uint32_t cluster_id{0}; + float full_distance{0.0F}; + }; + std::vector encoded; + encoded.reserve(count); + for (InnerIdType id = 0; id < count; ++id) { + EncodedNode node; + node.filter.resize(split->OneBitCodeSize()); + node.supplement.resize(split->SupplementCodeSize()); + REQUIRE(split->EncodeFused(vectors.data() + static_cast(id) * dim, + node.filter.data(), + node.supplement.data(), + &node.cluster_id)); + REQUIRE(split->ComputeFusedFull(computer, + node.cluster_id, + node.filter.data(), + node.supplement.data(), + &node.full_distance, + nullptr)); + encoded.push_back(std::move(node)); + } + const auto [best, worst] = + std::minmax_element(encoded.begin(), encoded.end(), [](const auto& lhs, const auto& rhs) { + return lhs.full_distance < rhs.full_distance; + }); + REQUIRE(best != encoded.end()); + REQUIRE(worst != encoded.end()); + REQUIRE(worst->full_distance > best->full_distance + 1e-4F); + + EncodedNode coarse_best_full_worst = *worst; + EncodedNode coarse_worst_full_best = *best; + RaBitQFusedTraversalQuery traversal_query; + REQUIRE(split->GetFusedTraversalQuery(computer, &traversal_query)); + const auto set_coarse_score = [&](EncodedNode& node, float score) { + const float filter_add = score - traversal_query.cluster_g_add[node.cluster_id]; + constexpr float filter_rescale = 0.0F; + constexpr float filter_error = 0.0F; + auto* metadata = node.filter.data() + traversal_query.one_bit_metadata_offset; + std::memcpy(metadata, &filter_add, sizeof(filter_add)); + std::memcpy(metadata + sizeof(float), &filter_rescale, sizeof(filter_rescale)); + std::memcpy(metadata + 2U * sizeof(float), &filter_error, sizeof(filter_error)); + }; + set_coarse_score(coarse_best_full_worst, -1000.0F); + set_coarse_score(coarse_worst_full_best, 1000.0F); + + auto graph_param = std::make_shared(); + graph_param->io_parameter_ = std::make_shared(); + graph_param->max_degree_ = 4; + graph_param->init_max_capacity_ = 2; + auto graph = std::make_shared( + graph_param, split->OneBitCodeSize(), split->SupplementCodeSize(), common_param); + graph->SetNodeCodes(0, + 0, + coarse_best_full_worst.cluster_id, + coarse_best_full_worst.filter.data(), + coarse_best_full_worst.supplement.data()); + graph->SetNodeCodes(1, + 1, + coarse_worst_full_best.cluster_id, + coarse_worst_full_best.filter.data(), + coarse_worst_full_best.supplement.data()); + Vector neighbor({1}, allocator.get()); + Vector no_neighbors(allocator.get()); + graph->InsertNeighborsById(0, neighbor); + graph->InsertNeighborsById(1, no_neighbors); + + HGraphRaBitQSearcher searcher(common_param, nullptr); + REQUIRE(searcher.Route(graph, graph, flatten, computer, 0, true) == 0); + REQUIRE(searcher.Route(graph, graph, flatten, computer, 0, false) == 1); + + const float invalid_metadata = std::numeric_limits::quiet_NaN(); + std::memcpy(coarse_best_full_worst.filter.data() + traversal_query.one_bit_metadata_offset, + &invalid_metadata, + sizeof(invalid_metadata)); + std::memcpy(coarse_worst_full_best.filter.data() + traversal_query.one_bit_metadata_offset, + &invalid_metadata, + sizeof(invalid_metadata)); + graph->SetNodeCodes(0, + 0, + coarse_best_full_worst.cluster_id, + coarse_best_full_worst.filter.data(), + coarse_best_full_worst.supplement.data()); + graph->SetNodeCodes(1, + 1, + coarse_worst_full_best.cluster_id, + coarse_worst_full_best.filter.data(), + coarse_worst_full_best.supplement.data()); + REQUIRE(searcher.Route(graph, graph, flatten, computer, 0, true) == 1); + + const auto make_search_graph = [&](bool invalidate_filter_add) { + auto search_graph_param = std::make_shared(); + search_graph_param->io_parameter_ = std::make_shared(); + search_graph_param->max_degree_ = 4; + search_graph_param->init_max_capacity_ = 5; + auto search_graph = std::make_shared( + search_graph_param, split->OneBitCodeSize(), split->SupplementCodeSize(), common_param); + for (InnerIdType id = 0; id < 5; ++id) { + auto filter = encoded[id].filter; + auto* metadata = filter.data() + traversal_query.one_bit_metadata_offset; + if (invalidate_filter_add) { + std::memcpy(metadata, &invalid_metadata, sizeof(invalid_metadata)); + } else { + const float overflowing_error = std::numeric_limits::max(); + std::memcpy( + metadata + 2U * sizeof(float), &overflowing_error, sizeof(overflowing_error)); + } + search_graph->SetNodeCodes( + id, id, encoded[id].cluster_id, filter.data(), encoded[id].supplement.data()); + } + Vector four_neighbors({1, 2, 3, 4}, allocator.get()); + search_graph->InsertNeighborsById(0, four_neighbors); + for (InnerIdType id = 1; id < 5; ++id) { + search_graph->InsertNeighborsById(id, no_neighbors); + } + return search_graph; + }; + const auto run_search = [&](const HGraphRaBitQFusedDataCellPtr& search_graph, + uint32_t expected_full_count) { + InnerSearchParam search_param; + search_param.ep = 0; + search_param.ef = 5; + search_param.topk = 5; + search_param.rerank_topk = 5; + search_param.enable_rabitq_one_bit_search = true; + search_param.enable_reorder = false; + search_param.rabitq_fused_computer = computer; + auto visited = std::make_shared(5, allocator.get()); + SearchStatistics statistics; + QueryContext context; + context.stats = &statistics; + auto result = searcher.Search( + search_graph, flatten, visited, vectors.data(), search_param, &context, nullptr); + REQUIRE(result != nullptr); + REQUIRE(result->Size() == 5); + REQUIRE(statistics.rabitq_filter_count.load() == 5); + REQUIRE(statistics.rabitq_full_count.load() == expected_full_count); + REQUIRE(statistics.rabitq_filter_fallback_full_count.load() == expected_full_count); + }; + run_search(make_search_graph(true), 5); + run_search(make_search_graph(false), 0); + + auto full_hint_graph = make_search_graph(false); + InnerSearchParam full_hint_param; + full_hint_param.ep = 0; + full_hint_param.ef = 41; + full_hint_param.topk = 5; + full_hint_param.rerank_topk = 5; + full_hint_param.enable_rabitq_one_bit_search = true; + full_hint_param.enable_reorder = true; + full_hint_param.rabitq_fused_computer = computer; + auto full_hint_visited = std::make_shared(5, allocator.get()); + RaBitQCandidateVector full_hint_candidates(allocator.get()); + SearchStatistics full_hint_search_statistics; + QueryContext full_hint_search_context; + full_hint_search_context.alloc = allocator.get(); + full_hint_search_context.stats = &full_hint_search_statistics; + bool full_hint_search_finalized = false; + auto full_hint_result = searcher.Search(full_hint_graph, + flatten, + full_hint_visited, + vectors.data(), + full_hint_param, + &full_hint_search_context, + &full_hint_candidates, + &full_hint_search_finalized); + REQUIRE(full_hint_result != nullptr); + REQUIRE(full_hint_search_finalized); + REQUIRE(full_hint_result->Size() == 5); + REQUIRE(full_hint_search_statistics.rabitq_full_count.load() == 5); + REQUIRE(full_hint_candidates.size() == 5); + REQUIRE(std::all_of( + full_hint_candidates.begin(), full_hint_candidates.end(), [](const auto& candidate) { + return IsFiniteRaBitQValue(candidate.full_distance); + })); + const auto heap_values_by_id = [](const DistHeapPtr& heap) { + std::vector> values; + values.reserve(heap->Size()); + const auto* data = heap->GetData(); + for (uint64_t i = 0; i < heap->Size(); ++i) { + values.emplace_back(data[i].second, data[i].first); + } + std::sort(values.begin(), values.end()); + return values; + }; + const auto full_hint_values = heap_values_by_id(full_hint_result); + + auto deferred_param = full_hint_param; + deferred_param.ef = 5; + auto deferred_visited = std::make_shared(5, allocator.get()); + RaBitQCandidateVector deferred_candidates(allocator.get()); + SearchStatistics deferred_search_statistics; + QueryContext deferred_search_context; + deferred_search_context.alloc = allocator.get(); + deferred_search_context.stats = &deferred_search_statistics; + bool deferred_search_finalized = false; + auto deferred_result = searcher.Search(full_hint_graph, + flatten, + deferred_visited, + vectors.data(), + deferred_param, + &deferred_search_context, + &deferred_candidates, + &deferred_search_finalized); + REQUIRE(deferred_result != nullptr); + REQUIRE(deferred_search_finalized); + REQUIRE(deferred_result->Size() == 5); + REQUIRE(deferred_search_statistics.rabitq_full_count.load() == 5); + SearchStatistics deferred_reorder_statistics; + QueryContext deferred_reorder_context; + deferred_reorder_context.alloc = allocator.get(); + deferred_reorder_context.stats = &deferred_reorder_statistics; + FlattenReorder deferred_reorder(flatten, allocator.get(), full_hint_graph); + auto deferred_reordered = deferred_reorder.Reorder(deferred_result, + vectors.data(), + 5, + deferred_reorder_context, + nullptr, + &deferred_candidates); + REQUIRE(deferred_reordered != nullptr); + REQUIRE(heap_values_by_id(deferred_reordered) == full_hint_values); + REQUIRE(deferred_reorder_statistics.reorder_distance_count.load() == 0); + REQUIRE(deferred_reorder_statistics.rabitq_full_count.load() == 0); + SECTION("deferred finalize preserves reasoning reorder events") { + ReasoningContext reasoning(allocator.get()); + Vector labels(allocator.get()); + UnorderedMap label_to_inner_id(allocator.get()); + for (InnerIdType id = 0; id < 5; ++id) { + labels.push_back(id); + label_to_inner_id[id] = id; + } + reasoning.InitializeExpectedTargets(labels, label_to_inner_id); + + auto reasoning_param = deferred_param; + reasoning_param.topk = 2; + reasoning_param.rerank_topk = 2; + auto reasoning_visited = std::make_shared(5, allocator.get()); + RaBitQCandidateVector reasoning_candidates(allocator.get()); + QueryContext reasoning_context; + reasoning_context.alloc = allocator.get(); + reasoning_context.reasoning_ctx = &reasoning; + bool reasoning_finalized = false; + auto reasoning_result = searcher.Search(full_hint_graph, + flatten, + reasoning_visited, + vectors.data(), + reasoning_param, + &reasoning_context, + &reasoning_candidates, + &reasoning_finalized); + REQUIRE(reasoning_result != nullptr); + REQUIRE(reasoning_finalized); + REQUIRE(reasoning.reorder_changes_.size() == 5); + + std::vector> exact_topk; + std::vector evicted_ids; + for (const auto& record : reasoning.reorder_changes_) { + REQUIRE(reasoning.expected_traces_.at(record.id).true_distance == record.dist_after); + if (exact_topk.size() == 2 and record.dist_after > exact_topk.back().first) { + continue; + } + if (exact_topk.size() == 2) { + evicted_ids.push_back(exact_topk.back().second); + } + const auto position = std::lower_bound( + exact_topk.begin(), + exact_topk.end(), + std::pair{record.dist_after, 0}, + [](const auto& lhs, const auto& rhs) { return lhs.first < rhs.first; }); + exact_topk.insert(position, {record.dist_after, record.id}); + if (exact_topk.size() > 2) { + exact_topk.pop_back(); + } + } + for (InnerIdType id = 0; id < 5; ++id) { + const bool expected_evicted = + std::find(evicted_ids.begin(), evicted_ids.end(), id) != evicted_ids.end(); + REQUIRE(reasoning.expected_traces_.at(id).reorder_evicted == expected_evicted); + } + } + SECTION("deferred search leaves precise reorder unfinalized without base candidates") { + auto precise_gate_visited = std::make_shared(5, allocator.get()); + bool precise_gate_finalized = true; + auto precise_gate_result = searcher.Search(full_hint_graph, + flatten, + precise_gate_visited, + vectors.data(), + deferred_param, + nullptr, + nullptr, + &precise_gate_finalized); + REQUIRE(precise_gate_result != nullptr); + REQUIRE_FALSE(precise_gate_finalized); + REQUIRE(heap_values_by_id(precise_gate_result) == full_hint_values); + } + + SECTION("deferred finalize handles fully and partially rejected candidates") { + auto reject_all = std::make_shared([](LabelType) { return false; }); + auto all_rejected_param = deferred_param; + all_rejected_param.is_inner_id_allowed = reject_all; + auto all_rejected_visited = std::make_shared(5, allocator.get()); + RaBitQCandidateVector all_rejected_candidates(allocator.get()); + SearchStatistics all_rejected_statistics; + QueryContext all_rejected_context; + all_rejected_context.alloc = allocator.get(); + all_rejected_context.stats = &all_rejected_statistics; + bool all_rejected_finalized = false; + auto all_rejected_result = searcher.Search(full_hint_graph, + flatten, + all_rejected_visited, + vectors.data(), + all_rejected_param, + &all_rejected_context, + &all_rejected_candidates, + &all_rejected_finalized); + REQUIRE(all_rejected_result != nullptr); + REQUIRE(all_rejected_finalized); + REQUIRE(all_rejected_result->Empty()); + REQUIRE(all_rejected_candidates.empty()); + REQUIRE(all_rejected_statistics.reorder_distance_count.load() == 0); + + auto allow_even = + std::make_shared([](LabelType id) { return id % 2 == 0; }); + auto partially_rejected_param = deferred_param; + partially_rejected_param.is_inner_id_allowed = allow_even; + auto partially_rejected_visited = std::make_shared(5, allocator.get()); + RaBitQCandidateVector partially_rejected_candidates(allocator.get()); + bool partially_rejected_finalized = false; + auto partially_rejected_result = searcher.Search(full_hint_graph, + flatten, + partially_rejected_visited, + vectors.data(), + partially_rejected_param, + nullptr, + &partially_rejected_candidates, + &partially_rejected_finalized); + REQUIRE(partially_rejected_result != nullptr); + REQUIRE(partially_rejected_finalized); + REQUIRE(partially_rejected_result->Size() == 3); + REQUIRE(partially_rejected_candidates.size() == 3); + const auto partially_rejected_values = heap_values_by_id(partially_rejected_result); + REQUIRE(std::all_of(partially_rejected_values.begin(), + partially_rejected_values.end(), + [](const auto& value) { return value.first % 2 == 0; })); + REQUIRE(std::all_of(partially_rejected_candidates.begin(), + partially_rejected_candidates.end(), + [](const auto& candidate) { + return candidate.id % 2 == 0 and + IsFiniteRaBitQValue(candidate.full_distance); + })); + } + + SECTION("generic fused fallback does not claim finalized results") { + auto untrained_flatten = FlattenInterface::MakeInstance(param, common_param); + auto untrained_split = + std::dynamic_pointer_cast(untrained_flatten); + REQUIRE(untrained_split != nullptr); + auto generic_fallback_visited = std::make_shared(5, allocator.get()); + RaBitQCandidateVector generic_fallback_candidates(allocator.get()); + bool generic_fallback_finalized = true; + auto generic_fallback_result = searcher.Search(full_hint_graph, + untrained_flatten, + generic_fallback_visited, + vectors.data(), + deferred_param, + nullptr, + &generic_fallback_candidates, + &generic_fallback_finalized); + REQUIRE(generic_fallback_result != nullptr); + REQUIRE(generic_fallback_result->Empty()); + REQUIRE(generic_fallback_candidates.empty()); + REQUIRE_FALSE(generic_fallback_finalized); + } + + for (InnerIdType id = 0; id < 5; ++id) { + auto filter = encoded[id].filter; + const float overflowing_error = std::numeric_limits::max(); + std::memcpy(filter.data() + traversal_query.one_bit_metadata_offset + 2U * sizeof(float), + &overflowing_error, + sizeof(overflowing_error)); + auto invalid_full_hint_supplement = encoded[id].supplement; + std::memcpy( + invalid_full_hint_supplement.data() + traversal_query.supplement_metadata_offset, + &invalid_metadata, + sizeof(invalid_metadata)); + full_hint_graph->SetNodeCodes( + id, id, encoded[id].cluster_id, filter.data(), invalid_full_hint_supplement.data()); + } + + SearchStatistics full_hint_statistics; + QueryContext full_hint_context; + full_hint_context.alloc = allocator.get(); + full_hint_context.stats = &full_hint_statistics; + FlattenReorder full_hint_reorder(flatten, allocator.get(), full_hint_graph); + auto full_hint_reordered = full_hint_reorder.Reorder( + full_hint_result, vectors.data(), 5, full_hint_context, nullptr, &full_hint_candidates); + REQUIRE(full_hint_reordered != nullptr); + REQUIRE(full_hint_reordered->Size() == 5); + REQUIRE(heap_values_by_id(full_hint_reordered) == full_hint_values); + REQUIRE(full_hint_statistics.reorder_distance_count.load() == 0); + REQUIRE(full_hint_statistics.rabitq_full_count.load() == 0); + + graph->SetNodeCodes( + 0, 0, encoded[0].cluster_id, encoded[0].filter.data(), encoded[0].supplement.data()); + RaBitQCandidateVector candidates(allocator.get()); + candidates.push_back({0.0F, std::numeric_limits::max(), static_cast(0)}); + SearchStatistics reorder_statistics; + QueryContext reorder_context; + reorder_context.alloc = allocator.get(); + reorder_context.stats = &reorder_statistics; + FlattenReorder reorder(flatten, allocator.get(), graph); + auto reordered = + reorder.Reorder(nullptr, vectors.data(), 1, reorder_context, nullptr, &candidates); + REQUIRE(reordered != nullptr); + REQUIRE(reordered->Size() == 1); + REQUIRE(reorder_statistics.reorder_distance_count.load() == 1); + REQUIRE(reorder_statistics.rabitq_full_count.load() == 1); + REQUIRE(reorder_statistics.rabitq_reorder_hint_full_count.load() == 0); + REQUIRE(reorder_statistics.rabitq_reorder_fallback_full_count.load() == 1); + + auto invalid_supplement = encoded[0].supplement; + std::memcpy(invalid_supplement.data() + traversal_query.supplement_metadata_offset, + &invalid_metadata, + sizeof(invalid_metadata)); + graph->SetNodeCodes( + 0, 0, encoded[0].cluster_id, encoded[0].filter.data(), invalid_supplement.data()); + + RaBitQCandidateVector reused_candidates(allocator.get()); + reused_candidates.push_back(candidates.front()); + reused_candidates.push_back({encoded[0].full_distance, + std::numeric_limits::quiet_NaN(), + static_cast(0), + encoded[0].full_distance}); + SearchStatistics reused_statistics; + QueryContext reused_context; + reused_context.alloc = allocator.get(); + reused_context.stats = &reused_statistics; + auto reused = + reorder.Reorder(nullptr, vectors.data(), 1, reused_context, nullptr, &reused_candidates); + REQUIRE(reused != nullptr); + REQUIRE(reused->Size() == 1); + REQUIRE(reused->Top().first == encoded[0].full_distance); + REQUIRE(reused_statistics.reorder_distance_count.load() == 0); + REQUIRE(reused_statistics.rabitq_full_count.load() == 0); + + graph->SetNodeCodes( + 1, 1, encoded[1].cluster_id, encoded[1].filter.data(), encoded[1].supplement.data()); + RaBitQCandidateVector mixed_candidates(reused_candidates, allocator.get()); + mixed_candidates.push_back({encoded[1].full_distance, + std::numeric_limits::quiet_NaN(), + static_cast(1)}); + SearchStatistics mixed_statistics; + QueryContext mixed_context; + mixed_context.alloc = allocator.get(); + mixed_context.stats = &mixed_statistics; + auto mixed = + reorder.Reorder(nullptr, vectors.data(), 2, mixed_context, nullptr, &mixed_candidates); + REQUIRE(mixed != nullptr); + REQUIRE(mixed->Size() == 2); + const auto mixed_values = heap_values_by_id(mixed); + const std::vector> expected_mixed_values{ + {static_cast(0), encoded[0].full_distance}, + {static_cast(1), encoded[1].full_distance}}; + REQUIRE(mixed_values == expected_mixed_values); + REQUIRE(mixed_statistics.reorder_distance_count.load() == 1); + REQUIRE(mixed_statistics.rabitq_full_count.load() == 1); + REQUIRE(mixed_statistics.rabitq_reorder_hint_full_count.load() == 0); + REQUIRE(mixed_statistics.rabitq_reorder_fallback_full_count.load() == 1); + REQUIRE_THROWS_AS( + reorder.Reorder(nullptr, vectors.data(), 1, reorder_context, nullptr, &candidates), + VsagException); + + auto input = std::make_shared>(allocator.get(), -1); + input->Push(0.0F, 0); + REQUIRE_THROWS_AS(reorder.Reorder(input, vectors.data(), 1, reorder_context, nullptr, nullptr), + VsagException); + + SECTION("deferred finalize drops a candidate whose full distance fails") { + graph->SetNodeCodes( + 0, 0, encoded[0].cluster_id, encoded[0].filter.data(), encoded[0].supplement.data()); + auto invalid_neighbor_supplement = encoded[1].supplement; + std::memcpy(invalid_neighbor_supplement.data() + traversal_query.supplement_metadata_offset, + &invalid_metadata, + sizeof(invalid_metadata)); + graph->SetNodeCodes(1, + 1, + encoded[1].cluster_id, + encoded[1].filter.data(), + invalid_neighbor_supplement.data()); + + auto failed_full_param = deferred_param; + failed_full_param.ef = 2; + failed_full_param.topk = 2; + failed_full_param.rerank_topk = 2; + auto failed_full_visited = std::make_shared(2, allocator.get()); + RaBitQCandidateVector failed_full_candidates(allocator.get()); + SearchStatistics failed_full_statistics; + QueryContext failed_full_context; + failed_full_context.alloc = allocator.get(); + failed_full_context.stats = &failed_full_statistics; + bool failed_full_finalized = false; + auto failed_full_result = searcher.Search(graph, + flatten, + failed_full_visited, + vectors.data(), + failed_full_param, + &failed_full_context, + &failed_full_candidates, + &failed_full_finalized); + REQUIRE(failed_full_result != nullptr); + REQUIRE(failed_full_finalized); + REQUIRE(failed_full_result->Size() == 1); + REQUIRE(failed_full_result->Top().second == 0); + REQUIRE(failed_full_candidates.size() == 2); + const auto failed_candidate = + std::find_if(failed_full_candidates.begin(), + failed_full_candidates.end(), + [](const auto& candidate) { return candidate.id == 1; }); + REQUIRE(failed_candidate != failed_full_candidates.end()); + REQUIRE(failed_candidate->full_distance == std::numeric_limits::max()); + REQUIRE(failed_full_statistics.rabitq_full_count.load() == 2); + } +} + +} // namespace vsag diff --git a/src/impl/searcher/parallel_searcher.cpp b/src/impl/searcher/parallel_searcher.cpp index 8a12618ca1..49eb3ba8c8 100644 --- a/src/impl/searcher/parallel_searcher.cpp +++ b/src/impl/searcher/parallel_searcher.cpp @@ -23,6 +23,7 @@ #include #include "datacell/flatten_interface.h" +#include "impl/filter/duplicate_group_filter.h" #include "impl/heap/standard_heap.h" #include "utils/filter_search_skip_strategy.h" #include "utils/spsc_queue.h" @@ -85,7 +86,7 @@ ParallelSearcher::Search(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { if (inner_search_param.search_mode == KNN_SEARCH) { return this->search_impl(graph, flatten, @@ -115,7 +116,7 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table, QueryContext* ctx, - DistanceRecordVector* rabitq_lower_bound_candidates) const { + RaBitQCandidateVector* rabitq_lower_bound_candidates) const { // set customize query alloctor Allocator* alloc = select_query_allocator(ctx, allocator_); @@ -147,11 +148,11 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, Vector line_dists(vector_size, alloc); Vector lower_bound_dists(vector_size, alloc); Vector> node_pair(beam, alloc); + const auto visit_filter = MakeDuplicateGroupFilter( + inner_search_param.is_inner_id_allowed, graph, inner_search_param.consider_duplicate); auto skip_strategy = create_filter_search_skip_strategy( inner_search_param.skip_strategy_type, - inner_search_param.is_inner_id_allowed != nullptr - ? inner_search_param.is_inner_id_allowed->ValidRatio() - : 1.0F, + visit_filter != nullptr ? visit_filter->ValidRatio() : 1.0F, inner_search_param.skip_ratio); if (rabitq_lower_bound_candidates != nullptr) { rabitq_lower_bound_candidates->clear(); @@ -167,6 +168,41 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, return (is_id_allowed == nullptr or is_id_allowed->CheckValid(id)) and (attr_ft == nullptr or attr_ft->CheckValid(id)); }; + auto push_duplicate_candidates = [&](InnerIdType id, float distance) { + if (not inner_search_param.consider_duplicate) { + return; + } + for (const auto duplicate_id : graph->GetDuplicateIds(id)) { + if (check_func(duplicate_id)) { + top_candidates->Push(distance, duplicate_id); + } + } + }; + auto append_lower_bound_candidates = [&](InnerIdType id, float bound) { + if (rabitq_lower_bound_candidates == nullptr) { + return; + } + const auto group_id = inner_search_param.consider_duplicate ? graph->GetGroupId(id) : id; + const auto append = [&](InnerIdType candidate, float candidate_bound) { + if (check_func(candidate)) { + rabitq_lower_bound_candidates->push_back( + {candidate_bound, std::numeric_limits::quiet_NaN(), candidate}); + } + }; + append(group_id, bound); + if (inner_search_param.consider_duplicate) { + for (const auto duplicate_id : graph->GetDuplicateIds(group_id)) { + if (not check_func(duplicate_id)) { + continue; + } + float duplicate_distance = 0.0F; + float duplicate_bound = std::numeric_limits::max(); + flatten->QueryWithDistanceLowerBound( + &duplicate_distance, &duplicate_bound, computer, &duplicate_id, 1, ctx); + append(duplicate_id, duplicate_bound); + } + } + }; if (inner_search_param.enable_rabitq_one_bit_search) { flatten->QueryWithDistanceLowerBound(&dist, nullptr, computer, &ep, 1, ctx); @@ -175,13 +211,20 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, } if (check_func(ep)) { top_candidates->Push(dist, ep); - lower_bound = top_candidates->Top().first; } + push_duplicate_candidates(ep, dist); if constexpr (mode == InnerSearchMode::RANGE_SEARCH) { - if (dist > inner_search_param.radius and not top_candidates->Empty()) { + while (dist > inner_search_param.radius and not top_candidates->Empty()) { + top_candidates->Pop(); + } + } else if constexpr (mode == InnerSearchMode::KNN_SEARCH) { + while (top_candidates->Size() > ef) { top_candidates->Pop(); } } + if (not top_candidates->Empty()) { + lower_bound = top_candidates->Top().first; + } if (dist < THRESHOLD_ERROR) { inner_search_param.duplicate_id = ep; } @@ -243,7 +286,7 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, count_no_visited = visit(graph, vl, node_pair, - inner_search_param.is_inner_id_allowed, + visit_filter, skip_strategy.get(), to_be_visited_id, neighbors, @@ -309,9 +352,8 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, dist = line_dists[i]; const auto cur_id = to_be_visited_id[i]; if constexpr (mode == KNN_SEARCH) { - if (collect_rabitq_lower_bound and lower_bound_dists[i] < lower_bound and - check_func(cur_id)) { - rabitq_lower_bound_candidates->emplace_back(lower_bound_dists[i], cur_id); + if (collect_rabitq_lower_bound and lower_bound_dists[i] < lower_bound) { + append_lower_bound_candidates(cur_id, lower_bound_dists[i]); } } if (dist < THRESHOLD_ERROR) { @@ -323,21 +365,13 @@ ParallelSearcher::search_impl(const GraphInterfacePtr& graph, if (check_func(cur_id)) { top_candidates->Push(dist, cur_id); } - if (inner_search_param.consider_duplicate) { - const auto duplicate_ids = graph->GetDuplicateIds(cur_id); - for (const auto& item : duplicate_ids) { - if (check_func(item)) { - top_candidates->Push(dist, item); - } - } - } + push_duplicate_candidates(cur_id, dist); if constexpr (mode == KNN_SEARCH) { - if (top_candidates->Size() > ef) { + while (top_candidates->Size() > ef) { top_candidates->Pop(); } } - if (not top_candidates->Empty()) { lower_bound = top_candidates->Top().first; } diff --git a/src/impl/searcher/parallel_searcher.h b/src/impl/searcher/parallel_searcher.h index 0f4ef39377..7b868d15ff 100644 --- a/src/impl/searcher/parallel_searcher.h +++ b/src/impl/searcher/parallel_searcher.h @@ -40,7 +40,7 @@ class ParallelSearcher { const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table = nullptr, QueryContext* ctx = nullptr, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; void SetMutexArray(MutexArrayPtr new_mutex_array); @@ -67,7 +67,7 @@ class ParallelSearcher { const InnerSearchParam& inner_search_param, const LabelTablePtr& label_table = nullptr, QueryContext* ctx = nullptr, - DistanceRecordVector* rabitq_lower_bound_candidates = nullptr) const; + RaBitQCandidateVector* rabitq_lower_bound_candidates = nullptr) const; private: Allocator* allocator_{nullptr}; diff --git a/src/inner_string_params.h b/src/inner_string_params.h index 6c2115a0f8..313e994d25 100644 --- a/src/inner_string_params.h +++ b/src/inner_string_params.h @@ -47,6 +47,7 @@ const char* const ATTR_PARAMS_KEY = "attr_params"; const char* const HGRAPH_USE_ELP_OPTIMIZER_KEY = "use_elp_optimizer"; const char* const HGRAPH_IGNORE_REORDER_KEY = "ignore_reorder"; const char* const HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY = "build_by_base"; +const char* const HGRAPH_RABITQ_FUSED_DATACELL_KEY = "rabitq_fused_datacell"; const char* const HGRAPH_USE_REVERSE_EDGES_KEY = "use_reverse_edges"; const char* const HGRAPH_PERSIST_SOURCE_ID_KEY = "persist_source_id"; const char* const HGRAPH_MCI_KEY = "mci"; @@ -222,6 +223,7 @@ const std::unordered_map DEFAULT_MAP = { {"HGRAPH_USE_ELP_OPTIMIZER_KEY", HGRAPH_USE_ELP_OPTIMIZER_KEY}, {"HGRAPH_IGNORE_REORDER_KEY", HGRAPH_IGNORE_REORDER_KEY}, {"HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY", HGRAPH_BUILD_BY_BASE_QUANTIZATION_KEY}, + {"HGRAPH_RABITQ_FUSED_DATACELL_KEY", HGRAPH_RABITQ_FUSED_DATACELL_KEY}, {"GRAPH_KEY", GRAPH_KEY}, {"BASE_CODES_KEY", BASE_CODES_KEY}, {"PRECISE_CODES_KEY", PRECISE_CODES_KEY}, diff --git a/src/quantization/computer.h b/src/quantization/computer.h index 34dea07d04..9ab05c37cc 100644 --- a/src/quantization/computer.h +++ b/src/quantization/computer.h @@ -37,7 +37,10 @@ template class Computer : public ComputerInterface { public: explicit Computer(const T* quantizer, Allocator* allocator) - : quantizer_(quantizer), allocator_(allocator), raw_query_(allocator){}; + : quantizer_(quantizer), + allocator_(allocator), + raw_query_(allocator), + auxiliary_codes_(allocator){}; ~Computer() override { if (quantizer_) { @@ -96,6 +99,13 @@ class Computer : public ComputerInterface { const T* quantizer_{nullptr}; uint8_t* buf_{nullptr}; Vector raw_query_; + // Optional query-side representation used by quantizers that keep the + // primary FP32 query for precise scoring while traversing with a compact + // bit-plane representation. + Vector auxiliary_codes_; + float auxiliary_lower_bound_{0.0F}; + float auxiliary_delta_{0.0F}; + float auxiliary_sum_{0.0F}; }; template diff --git a/src/quantization/rabitq_quantization/rabitq_quantizer.cpp b/src/quantization/rabitq_quantization/rabitq_quantizer.cpp index 9d1fbcb823..89f184c7be 100644 --- a/src/quantization/rabitq_quantization/rabitq_quantizer.cpp +++ b/src/quantization/rabitq_quantization/rabitq_quantizer.cpp @@ -31,6 +31,19 @@ namespace vsag { +namespace { + +[[nodiscard]] bool +is_normal_ra_bit_q_value(float value) { + uint32_t bits = 0; + std::memcpy(&bits, &value, sizeof(bits)); + constexpr uint32_t k_exponent_mask = 0x7F800000U; + const uint32_t exponent = bits & k_exponent_mask; + return exponent != 0U and exponent != k_exponent_mask; +} + +} // namespace + template RaBitQuantizer::RaBitQuantizer(int dim, uint64_t pca_dim, @@ -253,6 +266,21 @@ RaBitQuantizer::TrainImpl(const float* data, uint64_t count) { return true; } +template +void +RaBitQuantizer::SetCentroid(const float* centroid) { + CHECK_ARGUMENT(centroid != nullptr, "RaBitQ centroid must not be null"); + Vector pca_centroid(this->original_dim_, 0.0F, this->allocator_); + Vector rotated_centroid(this->dim_, 0.0F, this->allocator_); + if (pca_dim_ != this->original_dim_) { + pca_->Transform(centroid, pca_centroid.data()); + } else { + std::copy(centroid, centroid + original_dim_, pca_centroid.begin()); + } + rom_->Transform(pca_centroid.data(), rotated_centroid.data()); + centroid_.assign(rotated_centroid.begin(), rotated_centroid.end()); +} + inline float ip_obar_q(float ip_yu_q, float q_prime_sum, float y_norm, int B) { // used for recover distance from ip_yu_q @@ -394,10 +422,15 @@ RaBitQuantizer::RaBitQFloatSQIPBySplitCode(const float* query, } const uint64_t plane_bytes = PlaneBytes(); - if (filter_bits == 2 or filter_bits == 3) { - const float centered_filter_ip = - filter_bits == 2 ? RaBitQFloatTwoBitCenteredIP(query, filter_code, this->dim_) - : RaBitQFloatThreeBitCenteredIP(query, filter_code, this->dim_); + if (filter_bits >= 2 and filter_bits <= 4) { + float centered_filter_ip = 0.0F; + if (filter_bits == 2) { + centered_filter_ip = RaBitQFloatTwoBitCenteredIP(query, filter_code, this->dim_); + } else if (filter_bits == 3) { + centered_filter_ip = RaBitQFloatThreeBitCenteredIP(query, filter_code, this->dim_); + } else { + centered_filter_ip = RaBitQFloatFourBitCenteredIP(query, filter_code, this->dim_); + } const float filter_center = 0.5F * static_cast((1U << filter_bits) - 1U); const auto filter_scale = static_cast(1U << supplement_bits); const float supplement_ip = @@ -1265,9 +1298,25 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer float* dists, float* lower_bound, float runtime_rabitq_error_rate) const { + return ComputeDistWithOneBitLowerBoundAndFilterIP( + computer, one_bit_code, dists, lower_bound, nullptr, runtime_rabitq_error_rate); +} + +template +bool +RaBitQuantizer::ComputeDistWithOneBitLowerBoundAndFilterIP( + Computer& computer, + const uint8_t* one_bit_code, + float* dists, + float* lower_bound, + float* filter_inner_product, + float runtime_rabitq_error_rate) const { if (lower_bound != nullptr) { *lower_bound = std::numeric_limits::max(); } + if (filter_inner_product != nullptr) { + *filter_inner_product = std::numeric_limits::quiet_NaN(); + } if (not SupportSplitCodeStorage()) { return false; } @@ -1283,8 +1332,22 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer float filter_ip_yu_q = 0.0F; norm_type base_norm_code = 0.0F; if (FilterBits() == 1) { - filter_ip_estimate = RaBitQFloatBinaryIP( - reinterpret_cast(query), one_bit_code, this->dim_, inv_sqrt_d_); + if (HasFourBitTraversalQuery(computer)) { + const auto packed_ip = RaBitQSQ4UBinaryIPWithBaseSum( + computer.auxiliary_codes_.data(), one_bit_code, this->dim_); + const auto raw_ip = static_cast(packed_ip); + const auto base_sum = static_cast(packed_ip >> 32U); + filter_ip_estimate = recover_dist_between_sq4u_and_fp32(raw_ip, + static_cast(base_sum), + computer.auxiliary_sum_, + computer.auxiliary_lower_bound_, + computer.auxiliary_delta_, + inv_sqrt_d_, + this->dim_); + } else { + filter_ip_estimate = RaBitQFloatBinaryIP( + reinterpret_cast(query), one_bit_code, this->dim_, inv_sqrt_d_); + } } else { sum_type query_raw_sum = *((sum_type*)(query + query_offset_sum_)); memcpy( @@ -1297,6 +1360,10 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer filter_ip_yu_q = RaBitQFloatThreeBitCenteredIP( reinterpret_cast(query), one_bit_code, this->dim_); filter_ip_estimate = base_norm_code <= 0.0F ? 0.0F : filter_ip_yu_q / base_norm_code; + } else if (FilterBits() == 4) { + filter_ip_yu_q = RaBitQFloatFourBitCenteredIP( + reinterpret_cast(query), one_bit_code, this->dim_); + filter_ip_estimate = base_norm_code <= 0.0F ? 0.0F : filter_ip_yu_q / base_norm_code; } else { filter_ip_yu_q = RaBitQFloatSQIPBySplitCode(reinterpret_cast(query), one_bit_code, @@ -1308,6 +1375,11 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer ip_obar_q(filter_ip_yu_q, query_raw_sum, base_norm_code, FilterBits()); } } + // The x=1 traversal value may come from the SQ4 coarse kernel. This API does not carry a + // precision tag, so never expose an x=1 value as a reusable exact inner-product hint. + if (filter_inner_product != nullptr and FilterBits() >= 2) { + *filter_inner_product = filter_ip_estimate; + } float ip_est = filter_ip_estimate / one_bit_error; norm_type query_norm = *((norm_type*)(query + query_offset_norm_)); @@ -1343,7 +1415,7 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer } } - if (not std::isfinite(result)) { + if (not IsFiniteRaBitQValue(result)) { return false; } @@ -1354,7 +1426,7 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer error_type low_bound_error = *((error_type*)(one_bit_code + OneBitRecordLowBoundErrorOffset())); const float effective_error_rate = - std::isfinite(runtime_rabitq_error_rate) and runtime_rabitq_error_rate > 0.0F + IsFiniteRaBitQValue(runtime_rabitq_error_rate) and runtime_rabitq_error_rate > 0.0F ? runtime_rabitq_error_rate : rabitq_error_rate_; float lower_bound_error_term = @@ -1371,7 +1443,7 @@ RaBitQuantizer::ComputeDistWithOneBitLowerBound(Computer } float lower_bound_result = result - lower_bound_error_term; - if (std::isfinite(lower_bound_result)) { + if (IsFiniteRaBitQValue(lower_bound_result)) { *lower_bound = lower_bound_result - 1e-5F * std::max(1.0F, std::fabs(lower_bound_result)); } return true; @@ -1411,7 +1483,18 @@ RaBitQuantizer::ComputeDistsWithOneBitLowerBoundBatch4( return; } - if (FilterBits() != 1 and FilterBits() != 2 and FilterBits() != 3) { + if (FilterBits() < 1 or FilterBits() > 4) { + computed1 = this->ComputeDistWithOneBitLowerBound( + computer, one_bit_code1, &dist1, lower_bound1, runtime_rabitq_error_rate); + computed2 = this->ComputeDistWithOneBitLowerBound( + computer, one_bit_code2, &dist2, lower_bound2, runtime_rabitq_error_rate); + computed3 = this->ComputeDistWithOneBitLowerBound( + computer, one_bit_code3, &dist3, lower_bound3, runtime_rabitq_error_rate); + computed4 = this->ComputeDistWithOneBitLowerBound( + computer, one_bit_code4, &dist4, lower_bound4, runtime_rabitq_error_rate); + return; + } + if (FilterBits() == 1 and HasFourBitTraversalQuery(computer)) { computed1 = this->ComputeDistWithOneBitLowerBound( computer, one_bit_code1, &dist1, lower_bound1, runtime_rabitq_error_rate); computed2 = this->ComputeDistWithOneBitLowerBound( @@ -1471,6 +1554,15 @@ RaBitQuantizer::ComputeDistsWithOneBitLowerBoundBatch4( this->dim_, filter_ip_values); filter_ip_values_are_centered = true; + } else if (FilterBits() == 4) { + RaBitQFloatFourBitCenteredIPBatch4(query_data, + one_bit_code1, + one_bit_code2, + one_bit_code3, + one_bit_code4, + this->dim_, + filter_ip_values); + filter_ip_values_are_centered = true; } auto compute_one = @@ -1523,7 +1615,7 @@ RaBitQuantizer::ComputeDistsWithOneBitLowerBoundBatch4( } } - if (not std::isfinite(result)) { + if (not IsFiniteRaBitQValue(result)) { return false; } @@ -1535,7 +1627,7 @@ RaBitQuantizer::ComputeDistsWithOneBitLowerBoundBatch4( const error_type low_bound_error = *((error_type*)(one_bit_code + OneBitRecordLowBoundErrorOffset())); const float effective_error_rate = - std::isfinite(runtime_rabitq_error_rate) and runtime_rabitq_error_rate > 0.0F + IsFiniteRaBitQValue(runtime_rabitq_error_rate) and runtime_rabitq_error_rate > 0.0F ? runtime_rabitq_error_rate : rabitq_error_rate_; float lower_bound_error_term = 2.0F * base_norm * query_norm * effective_error_rate * @@ -1544,7 +1636,7 @@ RaBitQuantizer::ComputeDistsWithOneBitLowerBoundBatch4( lower_bound_error_term *= 0.5F; } const float lower_bound_result = result - lower_bound_error_term; - if (std::isfinite(lower_bound_result)) { + if (IsFiniteRaBitQValue(lower_bound_result)) { *lower_bound = lower_bound_result - 1e-5F * std::max(1.0F, std::fabs(lower_bound_result)); } @@ -1592,12 +1684,19 @@ RaBitQuantizer::ComputeDistWithSplitCode(Computer& compu } else { sum_type query_raw_sum = *((sum_type*)(query + query_offset_sum_)); float ip_yu_q = 0.0F; - if (FilterBits() == 2 or FilterBits() == 3) { + if (FilterBits() >= 2 and FilterBits() <= 4) { const auto* query_data = reinterpret_cast(query); - const float filter_centered_ip = - FilterBits() == 2 - ? RaBitQFloatTwoBitCenteredIP(query_data, one_bit_code, this->dim_) - : RaBitQFloatThreeBitCenteredIP(query_data, one_bit_code, this->dim_); + float filter_centered_ip = 0.0F; + if (FilterBits() == 2) { + filter_centered_ip = + RaBitQFloatTwoBitCenteredIP(query_data, one_bit_code, this->dim_); + } else if (FilterBits() == 3) { + filter_centered_ip = + RaBitQFloatThreeBitCenteredIP(query_data, one_bit_code, this->dim_); + } else { + filter_centered_ip = + RaBitQFloatFourBitCenteredIP(query_data, one_bit_code, this->dim_); + } const float filter_center = 0.5F * static_cast((1U << FilterBits()) - 1U); const float filter_ip_yu_q = filter_centered_ip + filter_center * query_raw_sum; ip_yu_q = filter_ip_yu_q * static_cast(1U << ReorderBits()); @@ -1679,7 +1778,8 @@ RaBitQuantizer::ComputeDistWithSplitCodeAndFilterDist(Computer::ComputeDistWithSplitCodeAndFilterDist(Computer +bool +RaBitQuantizer::ComputeDistWithSplitCodeAndFilterIP(Computer& computer, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* dists) const { + if constexpr (metric != MetricType::METRIC_TYPE_L2SQR and + metric != MetricType::METRIC_TYPE_IP) { + return false; + } + if (not SupportSplitCodeStorage() or ReorderBits() == 0 or + not IsFiniteRaBitQValue(filter_inner_product) or FilterBits() == 1) { + return false; + } + + const auto* query = computer.buf_; + const auto* split_meta = supplement_code + SupplementMetaOffset(); + const auto code_meta_offset = CodeMetaOffset(); + auto meta_field = [split_meta, code_meta_offset](uint64_t offset) { + return split_meta + (offset - code_meta_offset); + }; + const norm_type query_norm = *((norm_type*)(query + query_offset_norm_)); + const norm_type base_norm = *((norm_type*)(one_bit_code + OneBitRecordNormOffset())); + + norm_type filter_norm_code = 0.0F; + if (FilterBits() == 1) { + if (inv_sqrt_d_ <= 0.0F) { + return false; + } + filter_norm_code = 0.5F / inv_sqrt_d_; + } else { + memcpy(&filter_norm_code, + one_bit_code + OneBitRecordNormCodeOffset(), + sizeof(filter_norm_code)); + } if (filter_norm_code <= 0.0F) { return false; } @@ -1737,7 +1874,8 @@ RaBitQuantizer::ComputeDistWithSplitCodeAndFilterDist(Computer((1U << FilterBits()) - 1U); - const float filter_ip_yu_q = filter_ip_est * filter_norm_code + filter_center * query_raw_sum; + const float filter_ip_yu_q = + filter_inner_product * filter_norm_code + filter_center * query_raw_sum; const float shifted_filter_ip_yu_q = filter_ip_yu_q * static_cast(1U << ReorderBits()); const float supplement_ip_yu_q = RaBitQFloatSupplementCodeIP( reinterpret_cast(query), supplement_code, this->dim_, ReorderBits()); @@ -1761,11 +1899,20 @@ RaBitQuantizer::ComputeDistWithSplitCodeAndFilterDist(Computer::RecoverOrderSQ(const uint8_t* output, uint8_t* input) co template void -RaBitQuantizer::ProcessQueryImpl(const float* query, - Computer& computer) const { - try { - if (computer.buf_ == nullptr) { - computer.buf_ = - reinterpret_cast(this->allocator_->Allocate(this->query_code_size_)); +RaBitQuantizer::PrepareFourBitTraversalQuery(const float* normalized_query, + Computer& computer) const { + constexpr uint32_t k_query_bits = 4; + constexpr uint32_t k_ex_bits = k_query_bits - 1; + constexpr uint32_t k_ex_mask = (1U << k_ex_bits) - 1U; + constexpr float k_center = 7.5F; + const uint64_t plane_bytes = PlaneBytes(); + computer.auxiliary_codes_.assign(k_query_bits * plane_bytes, 0); + + // Port of RaBitQ-Library's 4-bit SplitSingleQuery quantization. The original + // implementation obtains this dimension-dependent constant by averaging 100 + // fixed-seed Gaussian vectors. For a normalized Gaussian vector and 3 + // supplement bits, the fitted factor is 3.21 * sqrt(dim) (99.5 at dim=960). + const float query_scale = 3.21F * std::sqrt(static_cast(this->dim_)); + float reconstruction_ip = 0.0F; + float reconstruction_norm_sqr = 0.0F; + float query_sum = 0.0F; + for (uint64_t i = 0; i < this->dim_; ++i) { + CHECK_ARGUMENT(IsFiniteRaBitQValue(normalized_query[i]), + "RaBitQ query must contain only finite values"); + auto magnitude = + static_cast(query_scale * std::fabs(normalized_query[i]) + 1e-5F); + magnitude = std::min(magnitude, k_ex_mask); + const auto supplement = normalized_query[i] < 0.0F ? k_ex_mask - magnitude : magnitude; + const auto quantized = static_cast( + supplement | (static_cast(normalized_query[i] > 0.0F) << k_ex_bits)); + const float centered = static_cast(quantized) - k_center; + reconstruction_ip += normalized_query[i] * centered; + reconstruction_norm_sqr += centered * centered; + query_sum += static_cast(quantized); + for (uint32_t bit = 0; bit < k_query_bits; ++bit) { + if ((quantized & (1U << bit)) != 0U) { + computer.auxiliary_codes_[static_cast(bit) * plane_bytes + i / 8] |= + static_cast(1U << (i & 7)); + } } - std::fill(computer.buf_, computer.buf_ + this->query_code_size_, 0); + } + const float delta = + reconstruction_norm_sqr > 0.0F ? reconstruction_ip / reconstruction_norm_sqr : 0.0F; + computer.auxiliary_lower_bound_ = -k_center * delta; + computer.auxiliary_delta_ = delta; + computer.auxiliary_sum_ = query_sum; +} - // use residual term in pca, so it's this->original_dim_ - Vector pca_data(this->original_dim_, 0, this->allocator_); - Vector transformed_data(this->dim_, 0, this->allocator_); - Vector normed_data(this->dim_, 0, this->allocator_); +template +bool +RaBitQuantizer::HasFourBitTraversalQuery(const Computer& computer) const { + return FilterBits() == 1 and num_bits_per_dim_query_ == 32 and + computer.auxiliary_codes_.size() == 4 * PlaneBytes(); +} - float query_raw_norm = 0; - if constexpr (metric == MetricType::METRIC_TYPE_IP or - metric == MetricType::METRIC_TYPE_COSINE) { - for (uint64_t d = 0; d < this->dim_; ++d) { - query_raw_norm += query[d] * query[d]; - } +template +void +RaBitQuantizer::TransformFusedQuery(const float* query, + Vector& transformed_query, + float& query_raw_norm, + norm_type& mrq_norm_sqr) const { + Vector pca_data(this->original_dim_, 0, this->allocator_); + query_raw_norm = 0.0F; + mrq_norm_sqr = 0.0F; + if constexpr (metric == MetricType::METRIC_TYPE_IP or + metric == MetricType::METRIC_TYPE_COSINE) { + for (uint64_t d = 0; d < this->dim_; ++d) { + query_raw_norm += query[d] * query[d]; } query_raw_norm = std::sqrt(query_raw_norm); - // 1. pca - if (pca_dim_ != this->original_dim_) { - pca_->Transform(query, pca_data.data()); - if (use_mrq_) { - norm_type mrq_norm_sqr = FP32ComputeIP(pca_data.data() + this->dim_, - pca_data.data() + this->dim_, - this->original_dim_ - this->dim_); + } + if (pca_dim_ != this->original_dim_) { + pca_->Transform(query, pca_data.data()); + if (use_mrq_) { + mrq_norm_sqr = FP32ComputeIP(pca_data.data() + this->dim_, + pca_data.data() + this->dim_, + this->original_dim_ - this->dim_); + } + } else { + pca_data.assign(query, query + original_dim_); + } + rom_->Transform(pca_data.data(), transformed_query.data()); +} + +template +void +RaBitQuantizer::PrepareHnswFourBitQuery(const float* transformed_query, + Vector& query_planes, + float& delta, + float& vl, + float& query_sum) const { + constexpr uint32_t k_query_bits = 4; + constexpr uint32_t k_ex_bits = k_query_bits - 1; + constexpr uint32_t k_ex_mask = (1U << k_ex_bits) - 1U; + constexpr float k_center = 7.5F; + const uint64_t plane_bytes = PlaneBytes(); + query_planes.assign(k_query_bits * plane_bytes, 0); + + double norm_sqr = 0.0; + for (uint64_t i = 0; i < this->dim_; ++i) { + CHECK_ARGUMENT(IsFiniteRaBitQValue(transformed_query[i]), + "RaBitQ query must contain only finite values"); + norm_sqr += static_cast(transformed_query[i]) * transformed_query[i]; + } + const auto norm = static_cast(std::sqrt(norm_sqr)); + const float query_scale = + norm > 0.0F ? 3.21F * std::sqrt(static_cast(this->dim_)) / norm : 0.0F; + double reconstruction_ip = 0.0; + double reconstruction_norm_sqr = 0.0; + query_sum = 0.0F; + for (uint64_t i = 0; i < this->dim_; ++i) { + auto magnitude = + static_cast(query_scale * std::fabs(transformed_query[i]) + 1e-5F); + magnitude = std::min(magnitude, k_ex_mask); + const auto supplement = transformed_query[i] < 0.0F ? k_ex_mask - magnitude : magnitude; + const auto quantized = static_cast( + supplement | (static_cast(transformed_query[i] > 0.0F) << k_ex_bits)); + const float centered = static_cast(quantized) - k_center; + reconstruction_ip += static_cast(transformed_query[i]) * centered; + reconstruction_norm_sqr += static_cast(centered) * centered; + query_sum += transformed_query[i]; + for (uint32_t bit = 0; bit < k_query_bits; ++bit) { + if ((quantized & (1U << bit)) != 0U) { + query_planes[static_cast(bit) * plane_bytes + i / 8] |= + static_cast(1U << (i & 7)); + } + } + } + delta = reconstruction_norm_sqr > 0.0 + ? static_cast(reconstruction_ip / reconstruction_norm_sqr) + : 0.0F; + vl = -k_center * delta; +} + +template +void +RaBitQuantizer::EncodeHnswOneBitMetadata(const float* data, uint8_t* one_bit_code) const { + Vector transformed_data(this->dim_, 0.0F, this->allocator_); + float raw_norm = 0.0F; + norm_type mrq_norm_sqr = 0.0F; + TransformFusedQuery(data, transformed_data, raw_norm, mrq_norm_sqr); + + double residual_norm_sqr = 0.0; + double residual_code_ip = 0.0; + double centroid_code_ip = 0.0; + double code_norm_sqr = 0.0; + for (uint64_t i = 0; i < this->dim_; ++i) { + const float residual = transformed_data[i] - centroid_[i]; + const float centered_code = residual > 0.0F ? 0.5F : -0.5F; + residual_norm_sqr += static_cast(residual) * residual; + residual_code_ip += static_cast(residual) * centered_code; + centroid_code_ip += static_cast(centroid_[i]) * centered_code; + code_norm_sqr += static_cast(centered_code) * centered_code; + } + + const auto l2_sqr = static_cast(residual_norm_sqr); + const float l2_norm = std::sqrt(l2_sqr); + const float safe_ip = std::fabs(residual_code_ip) > 1e-20 + ? static_cast(residual_code_ip) + : std::numeric_limits::infinity(); + float f_add = 0.0F; + float f_rescale = 0.0F; + if constexpr (metric == MetricType::METRIC_TYPE_IP) { + const float residual_centroid_ip = + FP32ComputeIP(transformed_data.data(), centroid_.data(), this->dim_) - + FP32ComputeIP(centroid_.data(), centroid_.data(), this->dim_); + f_add = + 1.0F - residual_centroid_ip + l2_sqr * static_cast(centroid_code_ip) / safe_ip; + f_rescale = -l2_sqr / safe_ip; + } else { + f_add = l2_sqr + 2.0F * l2_sqr * static_cast(centroid_code_ip) / safe_ip; + f_rescale = -2.0F * l2_sqr / safe_ip; + } + const double ratio = residual_norm_sqr * code_norm_sqr / + (static_cast(safe_ip) * static_cast(safe_ip)) - + 1.0; + const float error_scale = metric == MetricType::METRIC_TYPE_IP ? 1.0F : 2.0F; + const float f_error = + error_scale * l2_norm * RaBitQuantizerParameter::DEFAULT_RABITQ_ERROR_RATE * + std::sqrt(static_cast(std::max(0.0, ratio) / std::max(1, this->dim_ - 1))); + std::memcpy(one_bit_code + PlaneBytes(), &f_add, sizeof(float)); + std::memcpy(one_bit_code + PlaneBytes() + sizeof(float), &f_rescale, sizeof(float)); + std::memcpy(one_bit_code + PlaneBytes() + 2 * sizeof(float), &f_error, sizeof(float)); +} + +template +bool +RaBitQuantizer::EncodeHnswSupplement(const float* data, uint8_t* supplement_code) const { + if (data == nullptr or supplement_code == nullptr or centroid_.size() != this->dim_) { + return false; + } + for (uint64_t i = 0; i < this->original_dim_; ++i) { + if (not IsFiniteRaBitQValue(data[i])) { + return false; + } + } + constexpr uint32_t k_ex_bits = 7; + constexpr uint32_t k_ex_mask = (1U << k_ex_bits) - 1U; + constexpr float k_center = 127.5F; + constexpr float k_scale_960 = 1180.0F; + Vector transformed_data(this->dim_, 0.0F, this->allocator_); + Vector residual(this->dim_, 0.0F, this->allocator_); + Vector ex_codes(this->dim_, 0, this->allocator_); + float raw_norm = 0.0F; + norm_type mrq_norm_sqr = 0.0F; + TransformFusedQuery(data, transformed_data, raw_norm, mrq_norm_sqr); + + double residual_norm_sqr = 0.0; + for (uint64_t i = 0; i < this->dim_; ++i) { + if (not IsFiniteRaBitQValue(transformed_data[i]) or not IsFiniteRaBitQValue(centroid_[i])) { + return false; + } + residual[i] = transformed_data[i] - centroid_[i]; + if (not IsFiniteRaBitQValue(residual[i])) { + return false; + } + residual_norm_sqr += static_cast(residual[i]) * residual[i]; + } + const auto residual_norm = static_cast(std::sqrt(residual_norm_sqr)); + if (not IsFiniteRaBitQValue(residual_norm)) { + return false; + } + const float scale = + residual_norm > 0.0F + ? k_scale_960 * std::sqrt(static_cast(this->dim_) / 960.0F) / residual_norm + : 0.0F; + if (not IsFiniteRaBitQValue(scale)) { + return false; + } + double ipnorm = 0.0; + for (uint64_t i = 0; i < this->dim_; ++i) { + const float scaled_magnitude = scale * std::fabs(residual[i]) + 1e-5F; + if (not IsFiniteRaBitQValue(scaled_magnitude)) { + return false; + } + const auto magnitude = + static_cast(std::min(scaled_magnitude, static_cast(k_ex_mask))); + ex_codes[i] = static_cast(residual[i] < 0.0F ? k_ex_mask - magnitude : magnitude); + ipnorm += (static_cast(magnitude) + 0.5) * + std::fabs(static_cast(residual[i])) / + std::max(residual_norm, 1e-30); + } + + std::fill(supplement_code, supplement_code + SupplementPlanesSize(), 0); + if ((this->dim_ & 63U) == 0U) { + // Apache-2.0 RaBitQ-Library ExData packing_7bit_excode layout. + auto* output = supplement_code; + for (uint64_t block = 0; block < this->dim_; block += 64) { + const auto* input = ex_codes.data() + block; + for (uint64_t lane = 0; lane < 16; ++lane) { + output[lane] = static_cast((input[lane] & 0x3FU) | + ((input[48 + lane] & 0x03U) << 6U)); + output[16 + lane] = static_cast((input[16 + lane] & 0x3FU) | + ((input[48 + lane] & 0x0CU) << 4U)); + output[32 + lane] = static_cast((input[32 + lane] & 0x3FU) | + ((input[48 + lane] & 0x30U) << 2U)); + } + uint64_t top_bits = 0; + constexpr uint64_t k_top_mask = 0x0101010101010101ULL; + for (uint64_t lane = 0; lane < 64; lane += 8) { + uint64_t codes = 0; + std::memcpy(&codes, input + lane, sizeof(codes)); + top_bits |= ((codes >> 6U) & k_top_mask) << (lane / 8U); + } + std::memcpy(output + 48, &top_bits, sizeof(top_bits)); + output += 56; + } + } else { + const uint64_t plane_bytes = PlaneBytes(); + for (uint64_t i = 0; i < this->dim_; ++i) { + for (uint32_t bit = 0; bit < k_ex_bits; ++bit) { + if ((ex_codes[i] & (1U << bit)) != 0U) { + supplement_code[static_cast(bit) * plane_bytes + i / 8] |= + static_cast(1U << (i & 7)); + } + } + } + } + + double residual_code_ip = 0.0; + double centroid_code_ip = 0.0; + for (uint64_t i = 0; i < this->dim_; ++i) { + const float total_code = + static_cast(ex_codes[i]) + (residual[i] >= 0.0F ? 128.0F : 0.0F); + const float centered = total_code - k_center; + residual_code_ip += static_cast(residual[i]) * centered; + centroid_code_ip += static_cast(centroid_[i]) * centered; + } + const float safe_ip = std::fabs(residual_code_ip) > 1e-20 + ? static_cast(residual_code_ip) + : std::numeric_limits::infinity(); + const auto inverse_ipnorm = static_cast(1.0 / ipnorm); + const float ipnorm_inv = is_normal_ra_bit_q_value(inverse_ipnorm) ? inverse_ipnorm : 1.0F; + float f_add = 0.0F; + float f_rescale = 0.0F; + if constexpr (metric == MetricType::METRIC_TYPE_IP) { + const float residual_centroid_ip = + FP32ComputeIP(residual.data(), centroid_.data(), this->dim_); + f_add = 1.0F - residual_centroid_ip + + static_cast(residual_norm_sqr * centroid_code_ip / safe_ip); + f_rescale = -ipnorm_inv * residual_norm; + } else { + f_add = static_cast(residual_norm_sqr + + 2.0 * residual_norm_sqr * centroid_code_ip / safe_ip); + f_rescale = -2.0F * ipnorm_inv * residual_norm; + } + if (not IsFiniteRaBitQValue(f_add) or not IsFiniteRaBitQValue(f_rescale)) { + return false; + } + std::memcpy(supplement_code + SupplementMetaOffset(), &f_add, sizeof(float)); + std::memcpy( + supplement_code + SupplementMetaOffset() + sizeof(float), &f_rescale, sizeof(float)); + return true; +} + +template +void +RaBitQuantizer::ComputeHnswCentroidTerms(const float* transformed_query, + float& g_add, + float& g_error) const { + const float centroid_distance_sqr = + FP32ComputeL2Sqr(transformed_query, centroid_.data(), this->dim_); + g_error = std::sqrt(centroid_distance_sqr); + if constexpr (metric == MetricType::METRIC_TYPE_IP) { + g_add = -FP32ComputeIP(transformed_query, centroid_.data(), this->dim_); + } else { + g_add = centroid_distance_sqr; + } +} + +template +bool +RaBitQuantizer::ComputeHnswOneBit(const uint8_t* query_planes, + float query_delta, + float query_vl, + float query_sum, + float g_add, + float g_error, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + float* lower_bound, + float* filter_inner_product, + float runtime_rabitq_error_rate) const { + (void)supplement_code; + const auto packed_ip = RaBitQSQ4UBinaryIPWithBaseSum(query_planes, one_bit_code, this->dim_); + const auto raw_ip = static_cast(packed_ip); + const auto base_sum = static_cast(packed_ip >> 32U); + const float ip_x0_q = + query_delta * static_cast(raw_ip) + query_vl * static_cast(base_sum); + + float f_add = 0.0F; + float f_rescale = 0.0F; + float f_error = 0.0F; + std::memcpy(&f_add, one_bit_code + PlaneBytes(), sizeof(float)); + std::memcpy(&f_rescale, one_bit_code + PlaneBytes() + sizeof(float), sizeof(float)); + std::memcpy(&f_error, one_bit_code + PlaneBytes() + 2 * sizeof(float), sizeof(float)); + + *distance = f_add + g_add + f_rescale * (ip_x0_q - 0.5F * query_sum); + const float effective_error_rate = + IsFiniteRaBitQValue(runtime_rabitq_error_rate) and runtime_rabitq_error_rate > 0.0F + ? runtime_rabitq_error_rate + : rabitq_error_rate_; + const float error_rate_scale = + effective_error_rate / RaBitQuantizerParameter::DEFAULT_RABITQ_ERROR_RATE; + *lower_bound = *distance - error_rate_scale * f_error * g_error; + *filter_inner_product = ip_x0_q; + return IsFiniteRaBitQValue(*distance) and IsFiniteRaBitQValue(*lower_bound); +} + +template +bool +RaBitQuantizer::ComputeHnswFull(const float* transformed_query, + float query_sum, + float g_add, + float g_error, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* distance, + float* lower_bound) const { + if (not IsFiniteRaBitQValue(filter_inner_product)) { + return false; + } + float f_add = 0.0F; + float f_rescale = 0.0F; + float f_error = 0.0F; + std::memcpy(&f_add, supplement_code + SupplementMetaOffset(), sizeof(float)); + std::memcpy( + &f_rescale, supplement_code + SupplementMetaOffset() + sizeof(float), sizeof(float)); + std::memcpy(&f_error, one_bit_code + PlaneBytes() + 2 * sizeof(float), sizeof(float)); + const float supplement_ip = + (this->dim_ & 63U) == 0U + ? RaBitQFloatExCode7IP(transformed_query, supplement_code, this->dim_) + : RaBitQFloatSupplementCodeIP(transformed_query, supplement_code, this->dim_, 7); + *distance = f_add + g_add + + f_rescale * (128.0F * filter_inner_product + supplement_ip - 127.5F * query_sum); + *lower_bound = *distance - f_error * g_error / 128.0F; + return IsFiniteRaBitQValue(*distance) and IsFiniteRaBitQValue(*lower_bound); +} + +template +bool +RaBitQuantizer::EncodeFusedAffineMetadata(const float* data, + uint8_t* one_bit_code, + uint8_t* supplement_code) const { + if constexpr (metric != MetricType::METRIC_TYPE_L2SQR and + metric != MetricType::METRIC_TYPE_IP) { + return false; + } + if (data == nullptr or one_bit_code == nullptr or supplement_code == nullptr or + not SupportSplitCodeStorage() or FilterBits() < 1 or FilterBits() > 4 or + ReorderBits() == 0) { + return false; + } + + const auto read_float = [](const uint8_t* address) { + float value = 0.0F; + std::memcpy(&value, address, sizeof(value)); + return value; + }; + const auto write_float = [](uint8_t* address, float value) { + std::memcpy(address, &value, sizeof(value)); + }; + + const uint32_t filter_bits = FilterBits(); + const uint32_t supplement_bits = ReorderBits(); + const uint32_t base_bits = filter_bits + supplement_bits; + const uint64_t plane_bytes = PlaneBytes(); + const uint64_t filter_meta_offset = OneBitRecordNormOffset(); + const uint64_t supplement_meta_offset = SupplementMetaOffset(); + const auto* supplement_meta = supplement_code + supplement_meta_offset; + const auto full_meta_field = [this, supplement_meta](uint64_t full_code_offset) { + return supplement_meta + (full_code_offset - CodeMetaOffset()); + }; + + const float base_norm = read_float(one_bit_code + OneBitRecordNormOffset()); + const float filter_norm_code = filter_bits == 1 + ? 0.5F / inv_sqrt_d_ + : read_float(one_bit_code + OneBitRecordNormCodeOffset()); + const float filter_error = + std::fabs(read_float(one_bit_code + OneBitRecordOneBitErrorOffset())); + const float filter_epsilon = read_float(one_bit_code + OneBitRecordLowBoundErrorOffset()); + const float full_norm_code = read_float(full_meta_field(offset_norm_code_)); + float full_error = read_float(full_meta_field(offset_error_)); + if (std::fabs(full_error) < 1e-5F) { + full_error = full_error >= 0.0F ? 1.0F : -1.0F; + } + + Vector transformed_data(this->dim_, 0.0F, this->allocator_); + float raw_norm = 0.0F; + norm_type mrq_norm_sqr = 0.0F; + TransformFusedQuery(data, transformed_data, raw_norm, mrq_norm_sqr); + + const float filter_center = 0.5F * static_cast((1U << filter_bits) - 1U); + const float full_center = 0.5F * static_cast((1U << base_bits) - 1U); + double centroid_filter_ip = 0.0; + double centroid_full_ip = 0.0; + double centroid_residual_ip = 0.0; + double residual_norm_sqr = 0.0; + for (uint64_t d = 0; d < this->dim_; ++d) { + const uint64_t byte_idx = d >> 3U; + const auto bit_mask = static_cast(1U << (d & 7U)); + + uint32_t filter_code = 0; + for (uint32_t bit = 0; bit < filter_bits; ++bit) { + const auto* plane = one_bit_code + static_cast(bit) * plane_bytes; + if ((plane[byte_idx] & bit_mask) != 0U) { + filter_code |= 1U << (filter_bits - bit - 1U); + } + } + + uint32_t supplement = 0; + for (uint32_t bit = 0; bit < supplement_bits; ++bit) { + const auto* plane = supplement_code + static_cast(bit) * plane_bytes; + if ((plane[byte_idx] & bit_mask) != 0U) { + supplement |= 1U << bit; + } + } - *(norm_type*)(computer.buf_ + query_offset_mrq_norm_) = mrq_norm_sqr; + const uint32_t full_code = (filter_code << supplement_bits) | supplement; + const auto centroid = static_cast(centroid_[d]); + const double residual = + static_cast(transformed_data[d]) - static_cast(centroid_[d]); + centroid_filter_ip += centroid * (static_cast(filter_code) - filter_center); + centroid_full_ip += centroid * (static_cast(full_code) - full_center); + centroid_residual_ip += centroid * residual; + residual_norm_sqr += residual * residual; + } + + constexpr float metric_scale = metric == MetricType::METRIC_TYPE_IP ? 1.0F : 2.0F; + const float base_norm_sqr = base_norm * base_norm; + float filter_add = std::numeric_limits::quiet_NaN(); + float filter_rescale = std::numeric_limits::quiet_NaN(); + float filter_error_unit = std::numeric_limits::quiet_NaN(); + constexpr double k_normalize_zero_threshold = 1e-5; + const bool degenerate_residual = + std::isfinite(residual_norm_sqr) and residual_norm_sqr < k_normalize_zero_threshold; + if (degenerate_residual) { + const auto residual_norm = static_cast(std::sqrt(std::max(0.0, residual_norm_sqr))); + filter_rescale = 0.0F; + if constexpr (metric == MetricType::METRIC_TYPE_IP) { + filter_add = 1.0F - static_cast(centroid_residual_ip); + } else { + filter_add = static_cast(residual_norm_sqr); + if (pca_dim_ != original_dim_ and use_mrq_) { + filter_add += mrq_norm_sqr; } + } + filter_error_unit = metric_scale * residual_norm; + } else if (IsFiniteRaBitQValue(base_norm) and IsFiniteRaBitQValue(filter_norm_code) and + filter_norm_code > 0.0F and IsFiniteRaBitQValue(filter_error) and + filter_error > 1e-5F and IsFiniteRaBitQValue(filter_epsilon)) { + filter_rescale = -metric_scale * base_norm / (filter_error * filter_norm_code); + if constexpr (metric == MetricType::METRIC_TYPE_IP) { + filter_add = 1.0F - static_cast(centroid_residual_ip) - + filter_rescale * static_cast(centroid_filter_ip); } else { - pca_data.assign(query, query + original_dim_); - } - - // 2. random projection - rom_->Transform(pca_data.data(), transformed_data.data()); - - // 3. norm - float query_norm = NormalizeWithCentroid( - transformed_data.data(), centroid_.data(), normed_data.data(), this->dim_); - - // 4. query quantization - if (num_bits_per_dim_query_ == 4) { - // sq4 quantization - Vector quantized_data(this->dim_, 0, this->allocator_); - float lower_bound = std::numeric_limits::max(); - float upper_bound = std::numeric_limits::lowest(); - float delta = 0.0F; - sum_type query_sum = 0; - EncodeSQ(normed_data.data(), - quantized_data.data(), - upper_bound, - lower_bound, - delta, - query_sum); - ReOrderSQ(quantized_data.data(), reinterpret_cast(computer.buf_)); - // store info - *(float*)(computer.buf_ + query_offset_lb_) = lower_bound; - *(float*)(computer.buf_ + query_offset_delta_) = delta; - *(sum_type*)(computer.buf_ + query_offset_sum_) = query_sum; + filter_add = base_norm_sqr - filter_rescale * static_cast(centroid_filter_ip); + if (pca_dim_ != original_dim_ and use_mrq_) { + filter_add += mrq_norm_sqr; + } + } + filter_error_unit = metric_scale * base_norm * filter_epsilon / filter_error; + } + if (not IsFiniteRaBitQValue(filter_add) or not IsFiniteRaBitQValue(filter_rescale) or + not IsFiniteRaBitQValue(filter_error_unit)) { + filter_add = std::numeric_limits::quiet_NaN(); + filter_rescale = std::numeric_limits::quiet_NaN(); + filter_error_unit = std::numeric_limits::quiet_NaN(); + } + + float full_add = 0.0F; + float full_rescale = 0.0F; + if (degenerate_residual) { + full_add = filter_add; + } else { + if (not IsFiniteRaBitQValue(base_norm) or not IsFiniteRaBitQValue(full_norm_code) or + full_norm_code <= 0.0F or not IsFiniteRaBitQValue(full_error)) { + return false; + } + full_rescale = -metric_scale * base_norm / (full_error * full_norm_code); + if constexpr (metric == MetricType::METRIC_TYPE_IP) { + full_add = 1.0F - static_cast(centroid_residual_ip) - + full_rescale * static_cast(centroid_full_ip); } else { - // store codes - memcpy(computer.buf_, normed_data.data(), normed_data.size() * sizeof(float)); + full_add = base_norm_sqr - full_rescale * static_cast(centroid_full_ip); + if (pca_dim_ != original_dim_ and use_mrq_) { + full_add += mrq_norm_sqr; + } + } + } + if (not IsFiniteRaBitQValue(full_add) or not IsFiniteRaBitQValue(full_rescale)) { + return false; + } + + write_float(one_bit_code + filter_meta_offset, filter_add); + write_float(one_bit_code + filter_meta_offset + sizeof(float), filter_rescale); + write_float(one_bit_code + filter_meta_offset + 2U * sizeof(float), filter_error_unit); + write_float(supplement_code + supplement_meta_offset, full_add); + write_float(supplement_code + supplement_meta_offset + sizeof(float), full_rescale); + return true; +} + +template +bool +RaBitQuantizer::DecodeFusedSplitCode(const uint8_t* one_bit_code, + const uint8_t* supplement_code, + bool legacy_hnsw_codec, + float* data) const { + if constexpr (metric != MetricType::METRIC_TYPE_L2SQR and + metric != MetricType::METRIC_TYPE_IP) { + return false; + } + if (one_bit_code == nullptr or supplement_code == nullptr or data == nullptr or + not SupportSplitCodeStorage() or pca_dim_ != original_dim_ or rom_ == nullptr or + centroid_.size() != this->dim_) { + return false; + } + + const uint32_t filter_bits = FilterBits(); + const uint32_t supplement_bits = ReorderBits(); + const uint32_t base_bits = filter_bits + supplement_bits; + if (filter_bits < 1 or filter_bits > 4 or supplement_bits == 0 or base_bits >= 32 or + (legacy_hnsw_codec and (filter_bits != 1 or supplement_bits != 7))) { + return false; + } + + float full_rescale = 0.0F; + std::memcpy(&full_rescale, + supplement_code + SupplementMetaOffset() + sizeof(float), + sizeof(full_rescale)); + if (not IsFiniteRaBitQValue(full_rescale)) { + return false; + } + + constexpr float metric_scale = metric == MetricType::METRIC_TYPE_IP ? 1.0F : 2.0F; + const float residual_scale = -full_rescale / metric_scale; + if (not IsFiniteRaBitQValue(residual_scale)) { + return false; + } + + const uint64_t plane_bytes = PlaneBytes(); + const float full_center = 0.5F * static_cast((1U << base_bits) - 1U); + Vector transformed_data(this->dim_, 0.0F, this->allocator_); + for (uint64_t d = 0; d < this->dim_; ++d) { + const uint64_t byte_idx = d >> 3U; + const auto bit_mask = static_cast(1U << (d & 7U)); + uint32_t filter_code = 0; + for (uint32_t bit = 0; bit < filter_bits; ++bit) { + const auto* plane = one_bit_code + static_cast(bit) * plane_bytes; + if ((plane[byte_idx] & bit_mask) != 0U) { + filter_code |= 1U << (filter_bits - bit - 1U); + } } - if (num_bits_per_dim_base_ != 1) { - float query_raw_sum = 0; - for (uint32_t d = 0; d < this->dim_; d++) { - query_raw_sum += normed_data[d]; + uint32_t supplement = 0; + if (legacy_hnsw_codec and (this->dim_ & 63U) == 0U) { + constexpr uint64_t legacy_block_size = 56; + constexpr uint64_t legacy_low_dimension_count = 48; + const uint64_t lane = d & 63U; + const auto* block = supplement_code + (d >> 6U) * legacy_block_size; + const uint32_t top = (block[48U + (lane & 7U)] >> (lane >> 3U)) & 1U; + uint32_t low = 0; + if (lane < legacy_low_dimension_count) { + low = block[lane] & 0x3FU; + } else { + const uint64_t packed_lane = lane - legacy_low_dimension_count; + low = ((block[packed_lane] >> 6U) & 0x3U) | + (((block[16U + packed_lane] >> 6U) & 0x3U) << 2U) | + (((block[32U + packed_lane] >> 6U) & 0x3U) << 4U); } - *(sum_type*)(computer.buf_ + query_offset_sum_) = query_raw_sum; + supplement = low | (top << 6U); + } else { + for (uint32_t bit = 0; bit < supplement_bits; ++bit) { + const auto* plane = supplement_code + static_cast(bit) * plane_bytes; + if ((plane[byte_idx] & bit_mask) != 0U) { + supplement |= 1U << bit; + } + } + } + + const uint32_t full_code = (filter_code << supplement_bits) | supplement; + transformed_data[d] = + centroid_[d] + residual_scale * (static_cast(full_code) - full_center); + if (not IsFiniteRaBitQValue(transformed_data[d])) { + return false; + } + } + + rom_->InverseTransform(transformed_data.data(), data); + for (uint64_t d = 0; d < original_dim_; ++d) { + if (not IsFiniteRaBitQValue(data[d])) { + return false; + } + } + return true; +} + +template +bool +RaBitQuantizer::ComputeFusedExactCenteredFilterIP( + const float* transformed_query, + const uint8_t* one_bit_code, + float* centered_filter_inner_product) const { + if (transformed_query == nullptr or one_bit_code == nullptr or + centered_filter_inner_product == nullptr or FilterBits() < 1 or FilterBits() > 4) { + return false; + } + + if (FilterBits() == 1) { + if (inv_sqrt_d_ <= 0.0F) { + return false; } + const float normalized_ip = + RaBitQFloatBinaryIP(transformed_query, one_bit_code, this->dim_, inv_sqrt_d_); + *centered_filter_inner_product = normalized_ip * (0.5F / inv_sqrt_d_); + } else if (FilterBits() == 2) { + *centered_filter_inner_product = + RaBitQFloatTwoBitCenteredIP(transformed_query, one_bit_code, this->dim_); + } else if (FilterBits() == 3) { + *centered_filter_inner_product = + RaBitQFloatThreeBitCenteredIP(transformed_query, one_bit_code, this->dim_); + } else { + *centered_filter_inner_product = + RaBitQFloatFourBitCenteredIP(transformed_query, one_bit_code, this->dim_); + } + return IsFiniteRaBitQValue(*centered_filter_inner_product); +} + +template +bool +RaBitQuantizer::ComputeFusedAffineFilter(const float* transformed_query, + const uint8_t* query_planes, + float query_delta, + float query_vl, + float query_sum, + float g_add, + float g_error, + const uint8_t* one_bit_code, + float runtime_rabitq_error_rate, + float* distance, + float* lower_bound, + float* centered_filter_inner_product, + RaBitQFusedIPPrecision* precision) const { + if (distance != nullptr) { + *distance = std::numeric_limits::max(); + } + if (lower_bound != nullptr) { + *lower_bound = std::numeric_limits::max(); + } + if (centered_filter_inner_product != nullptr) { + *centered_filter_inner_product = std::numeric_limits::quiet_NaN(); + } + if (precision != nullptr) { + *precision = RaBitQFusedIPPrecision::INVALID; + } + if constexpr (metric != MetricType::METRIC_TYPE_L2SQR and + metric != MetricType::METRIC_TYPE_IP) { + return false; + } + if (transformed_query == nullptr or one_bit_code == nullptr or distance == nullptr or + lower_bound == nullptr or centered_filter_inner_product == nullptr or + precision == nullptr or not SupportSplitCodeStorage() or FilterBits() < 1 or + FilterBits() > 4 or not IsFiniteRaBitQValue(query_sum) or not IsFiniteRaBitQValue(g_add) or + not IsFiniteRaBitQValue(g_error)) { + return false; + } + + float filter_ip = 0.0F; + auto filter_precision = RaBitQFusedIPPrecision::EXACT; + if (FilterBits() == 1 and query_planes != nullptr) { + if (not IsFiniteRaBitQValue(query_delta) or not IsFiniteRaBitQValue(query_vl)) { + return false; + } + const uint64_t packed_ip = + RaBitQSQ4UBinaryIPWithBaseSum(query_planes, one_bit_code, this->dim_); + const auto raw_ip = static_cast(packed_ip); + const auto base_sum = static_cast(packed_ip >> 32U); + const float uncentered_filter_ip = + query_delta * static_cast(raw_ip) + query_vl * static_cast(base_sum); + filter_ip = uncentered_filter_ip - 0.5F * query_sum; + filter_precision = RaBitQFusedIPPrecision::APPROXIMATE; + } else if (not ComputeFusedExactCenteredFilterIP(transformed_query, one_bit_code, &filter_ip)) { + return false; + } + + const uint64_t metadata_offset = OneBitRecordNormOffset(); + float filter_add = 0.0F; + float filter_rescale = 0.0F; + float filter_error_unit = 0.0F; + std::memcpy(&filter_add, one_bit_code + metadata_offset, sizeof(filter_add)); + std::memcpy( + &filter_rescale, one_bit_code + metadata_offset + sizeof(float), sizeof(filter_rescale)); + std::memcpy(&filter_error_unit, + one_bit_code + metadata_offset + 2U * sizeof(float), + sizeof(filter_error_unit)); + if (not IsFiniteRaBitQValue(filter_ip) or not IsFiniteRaBitQValue(filter_add) or + not IsFiniteRaBitQValue(filter_rescale) or not IsFiniteRaBitQValue(filter_error_unit)) { + return false; + } + + const float result = filter_add + g_add + filter_rescale * filter_ip; + const float effective_error_rate = + IsFiniteRaBitQValue(runtime_rabitq_error_rate) and runtime_rabitq_error_rate > 0.0F + ? runtime_rabitq_error_rate + : rabitq_error_rate_; + const float lower_bound_result = result - effective_error_rate * filter_error_unit * g_error; + if (not IsFiniteRaBitQValue(result)) { + return false; + } + + // Traversal can still use a valid coarse distance when its error bound overflows and + // supplement reranking is disabled. + *distance = result; + if (not IsFiniteRaBitQValue(lower_bound_result)) { + return false; + } + *lower_bound = lower_bound_result - 1e-5F * std::max(1.0F, std::fabs(lower_bound_result)); + *centered_filter_inner_product = filter_ip; + *precision = filter_precision; + return true; +} + +template +bool +RaBitQuantizer::ComputeFusedAffineFullWithFilterIP( + const float* transformed_query, + float query_sum, + float g_add, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float exact_centered_filter_inner_product, + float* distance) const { + if constexpr (metric != MetricType::METRIC_TYPE_L2SQR and + metric != MetricType::METRIC_TYPE_IP) { + return false; + } + if (transformed_query == nullptr or one_bit_code == nullptr or supplement_code == nullptr or + distance == nullptr or not SupportSplitCodeStorage() or ReorderBits() == 0 or + not IsFiniteRaBitQValue(query_sum) or not IsFiniteRaBitQValue(g_add) or + not IsFiniteRaBitQValue(exact_centered_filter_inner_product)) { + return false; + } + + float full_add = 0.0F; + float full_rescale = 0.0F; + std::memcpy(&full_add, supplement_code + SupplementMetaOffset(), sizeof(full_add)); + std::memcpy(&full_rescale, + supplement_code + SupplementMetaOffset() + sizeof(float), + sizeof(full_rescale)); + if (not IsFiniteRaBitQValue(full_add) or not IsFiniteRaBitQValue(full_rescale)) { + return false; + } + + const uint32_t supplement_bits = ReorderBits(); + const float supplement_ip = RaBitQFloatSupplementCodeIP( + transformed_query, supplement_code, this->dim_, supplement_bits); + const float supplement_center = 0.5F * static_cast((1U << supplement_bits) - 1U); + const float full_centered_ip = + std::ldexp(exact_centered_filter_inner_product, supplement_bits) + supplement_ip - + supplement_center * query_sum; + const float result = full_add + g_add + full_rescale * full_centered_ip; + if (not IsFiniteRaBitQValue(result)) { + return false; + } + *distance = result; + return true; +} + +template +bool +RaBitQuantizer::ComputeFusedAffineFullDirect(const float* transformed_query, + float query_sum, + float g_add, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance) const { + float exact_filter_ip = 0.0F; + if (not ComputeFusedExactCenteredFilterIP(transformed_query, one_bit_code, &exact_filter_ip)) { + return false; + } + return ComputeFusedAffineFullWithFilterIP(transformed_query, + query_sum, + g_add, + one_bit_code, + supplement_code, + exact_filter_ip, + distance); +} + +template +void +RaBitQuantizer::ProcessTransformedFusedQuery(const float* transformed_query, + float query_raw_norm, + norm_type mrq_norm_sqr, + Computer& computer) const { + if (computer.buf_ == nullptr) { + computer.buf_ = + reinterpret_cast(this->allocator_->Allocate(this->query_code_size_)); + } + std::fill(computer.buf_, computer.buf_ + this->query_code_size_, 0); + Vector normed_data(this->dim_, 0, this->allocator_); + const float query_norm = + NormalizeWithCentroid(transformed_query, centroid_.data(), normed_data.data(), this->dim_); - // 5. store norm - *(norm_type*)(computer.buf_ + query_offset_norm_) = query_norm; - if constexpr (metric == MetricType::METRIC_TYPE_IP or - metric == MetricType::METRIC_TYPE_COSINE) { - *(norm_type*)(computer.buf_ + query_offset_raw_norm_) = query_raw_norm; + if (SupportSplitCodeStorage() and FilterBits() == 1 and num_bits_per_dim_query_ == 32) { + PrepareFourBitTraversalQuery(normed_data.data(), computer); + } else { + computer.auxiliary_codes_.clear(); + } + + if (num_bits_per_dim_query_ == 4) { + Vector quantized_data(this->dim_, 0, this->allocator_); + float lower_bound = std::numeric_limits::max(); + float upper_bound = std::numeric_limits::lowest(); + float delta = 0.0F; + sum_type query_sum = 0; + EncodeSQ( + normed_data.data(), quantized_data.data(), upper_bound, lower_bound, delta, query_sum); + ReOrderSQ(quantized_data.data(), reinterpret_cast(computer.buf_)); + *(float*)(computer.buf_ + query_offset_lb_) = lower_bound; + *(float*)(computer.buf_ + query_offset_delta_) = delta; + *(sum_type*)(computer.buf_ + query_offset_sum_) = query_sum; + } else { + memcpy(computer.buf_, normed_data.data(), normed_data.size() * sizeof(float)); + } + + if (num_bits_per_dim_base_ != 1) { + float query_raw_sum = 0; + for (uint32_t d = 0; d < this->dim_; d++) { + query_raw_sum += normed_data[d]; } + *(sum_type*)(computer.buf_ + query_offset_sum_) = query_raw_sum; + } + + *(norm_type*)(computer.buf_ + query_offset_norm_) = query_norm; + if (use_mrq_) { + *(norm_type*)(computer.buf_ + query_offset_mrq_norm_) = mrq_norm_sqr; + } + if constexpr (metric == MetricType::METRIC_TYPE_IP or + metric == MetricType::METRIC_TYPE_COSINE) { + *(norm_type*)(computer.buf_ + query_offset_raw_norm_) = query_raw_norm; + } +} + +template +void +RaBitQuantizer::ProcessQueryImpl(const float* query, + Computer& computer) const { + try { + Vector transformed_data(this->dim_, 0, this->allocator_); + float query_raw_norm = 0.0F; + norm_type mrq_norm_sqr = 0.0F; + TransformFusedQuery(query, transformed_data, query_raw_norm, mrq_norm_sqr); + ProcessTransformedFusedQuery( + transformed_data.data(), query_raw_norm, mrq_norm_sqr, computer); } catch (std::bad_alloc& e) { logger::error("bad alloc when init computer buf"); throw e; diff --git a/src/quantization/rabitq_quantization/rabitq_quantizer.h b/src/quantization/rabitq_quantization/rabitq_quantizer.h index 6df410963b..f8a9673536 100644 --- a/src/quantization/rabitq_quantization/rabitq_quantizer.h +++ b/src/quantization/rabitq_quantization/rabitq_quantizer.h @@ -15,6 +15,8 @@ #pragma once +#include +#include #include #include @@ -26,6 +28,20 @@ namespace vsag { +[[nodiscard]] inline bool +IsFiniteRaBitQValue(float value) { + uint32_t bits = 0; + std::memcpy(&bits, &value, sizeof(bits)); + constexpr uint32_t k_exponent_mask = 0x7F800000U; + return (bits & k_exponent_mask) != k_exponent_mask; +} + +enum class RaBitQFusedIPPrecision : uint8_t { + INVALID = 0, + EXACT = 1, + APPROXIMATE = 2, +}; + /** Implement of RaBitQ Quantization, Integrate MRQ (Minimized Residual Quantization) and Extend-RaBitQ * * RaBitQ: Supports bit-level quantization @@ -71,6 +87,9 @@ class RaBitQuantizer : public Quantizer> { bool TrainImpl(const float* data, uint64_t count); + void + SetCentroid(const float* centroid); + bool EncodeOneImpl(const float* data, uint8_t* codes) const; @@ -92,6 +111,112 @@ class RaBitQuantizer : public Quantizer> { void ProcessQueryImpl(const float* query, Computer& computer) const; + void + TransformFusedQuery(const float* query, + Vector& transformed_query, + float& query_raw_norm, + norm_type& mrq_norm_sqr) const; + + void + ProcessTransformedFusedQuery(const float* transformed_query, + float query_raw_norm, + norm_type mrq_norm_sqr, + Computer& computer) const; + + // RaBitQ-Library HNSW-compatible fused traversal helpers. These intentionally + // quantize the rotated query once, without subtracting a cluster centroid. + void + PrepareHnswFourBitQuery(const float* transformed_query, + Vector& query_planes, + float& delta, + float& vl, + float& query_sum) const; + + void + EncodeHnswOneBitMetadata(const float* data, uint8_t* one_bit_code) const; + + bool + EncodeHnswSupplement(const float* data, uint8_t* supplement_code) const; + + void + ComputeHnswCentroidTerms(const float* transformed_query, float& g_add, float& g_error) const; + + bool + ComputeHnswOneBit( + const uint8_t* query_planes, + float query_delta, + float query_vl, + float query_sum, + float g_add, + float g_error, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance, + float* lower_bound, + float* filter_inner_product, + float runtime_rabitq_error_rate = std::numeric_limits::quiet_NaN()) const; + + bool + ComputeHnswFull(const float* transformed_query, + float query_sum, + float g_add, + float g_error, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* distance, + float* lower_bound) const; + + // Replaces the native split metadata in-place with fused affine coefficients. + // The x/y plane payload and the record sizes are unchanged. + bool + EncodeFusedAffineMetadata(const float* data, + uint8_t* one_bit_code, + uint8_t* supplement_code) const; + + bool + DecodeFusedSplitCode(const uint8_t* one_bit_code, + const uint8_t* supplement_code, + bool legacy_hnsw_codec, + float* data) const; + + bool + ComputeFusedAffineFilter(const float* transformed_query, + const uint8_t* query_planes, + float query_delta, + float query_vl, + float query_sum, + float g_add, + float g_error, + const uint8_t* one_bit_code, + float runtime_rabitq_error_rate, + float* distance, + float* lower_bound, + float* centered_filter_inner_product, + RaBitQFusedIPPrecision* precision) const; + + bool + ComputeFusedAffineFullWithFilterIP(const float* transformed_query, + float query_sum, + float g_add, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float exact_centered_filter_inner_product, + float* distance) const; + + bool + ComputeFusedAffineFullDirect(const float* transformed_query, + float query_sum, + float g_add, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float* distance) const; + + [[nodiscard]] float + DefaultRaBitQErrorRate() const { + return rabitq_error_rate_; + } + void ComputeDistImpl(Computer& computer, const uint8_t* codes, float* dists) const; @@ -235,6 +360,19 @@ class RaBitQuantizer : public Quantizer> { float* lower_bound, float runtime_rabitq_error_rate = std::numeric_limits::quiet_NaN()) const; + // Same coarse x-bit estimate as ComputeDistWithOneBitLowerBound(). When the + // filter-stage inner product is exact, it is returned for direct split + // reranking. The x=1 four-bit traversal estimate is approximate, so that + // mode returns NaN and requires the full path to rescan the x-bit code. + bool + ComputeDistWithOneBitLowerBoundAndFilterIP( + Computer& computer, + const uint8_t* one_bit_code, + float* dists, + float* lower_bound, + float* filter_inner_product, + float runtime_rabitq_error_rate = std::numeric_limits::quiet_NaN()) const; + void ComputeDistsWithOneBitLowerBoundBatch4( Computer& computer, @@ -263,10 +401,11 @@ class RaBitQuantizer : public Quantizer> { float* dists) const; // Computes the full x+y split distance while reusing the filter-stage distance. - // `filter_dist` is the x-bit distance already produced by + // `filter_dist` is the exact x-bit distance already produced by // ComputeDistWithOneBitLowerBound(); it is not a lower bound and not the final // x+y distance. Passing it here lets the reorder path scan only the y-bit - // supplement planes instead of rescanning the x-bit filter planes. + // supplement planes instead of rescanning the x-bit filter planes. The + // x=1 four-bit traversal estimate is rejected because it is approximate. bool ComputeDistWithSplitCodeAndFilterDist(Computer& computer, const uint8_t* one_bit_code, @@ -274,6 +413,13 @@ class RaBitQuantizer : public Quantizer> { float filter_dist, float* dists) const; + bool + ComputeDistWithSplitCodeAndFilterIP(Computer& computer, + const uint8_t* one_bit_code, + const uint8_t* supplement_code, + float filter_inner_product, + float* dists) const; + [[nodiscard]] uint64_t OneBitRecordNormOffset() const; @@ -325,7 +471,19 @@ class RaBitQuantizer : public Quantizer> { [[nodiscard]] uint64_t AlignCodeField(uint64_t size) const; + [[nodiscard]] bool + HasFourBitTraversalQuery(const Computer& computer) const; + private: + bool + ComputeFusedExactCenteredFilterIP(const float* transformed_query, + const uint8_t* one_bit_code, + float* centered_filter_inner_product) const; + + void + PrepareFourBitTraversalQuery(const float* normalized_query, + Computer& computer) const; + bool EncodeOneInternal(const float* data, uint8_t* codes, diff --git a/src/quantization/rabitq_quantization/rabitq_quantizer_test.cpp b/src/quantization/rabitq_quantization/rabitq_quantizer_test.cpp index 5cac37455d..8080c7ea1b 100644 --- a/src/quantization/rabitq_quantization/rabitq_quantizer_test.cpp +++ b/src/quantization/rabitq_quantization/rabitq_quantizer_test.cpp @@ -18,6 +18,7 @@ #include #include #include +#include #include #include #include @@ -34,6 +35,25 @@ using namespace vsag; const auto dims = fixtures::get_common_used_dims(6, 129); const auto counts = {100}; +namespace { + +bool +IsNaNBitPattern(float value) { + uint32_t bits = 0; + std::memcpy(&bits, &value, sizeof(bits)); + return (bits & 0x7FFFFFFFU) > 0x7F800000U; +} + +} // namespace + +TEST_CASE("RaBitQ finite guard survives fast math", "[ut][RaBitQuantizer][rabitq_split]") { + REQUIRE(IsFiniteRaBitQValue(0.0F)); + REQUIRE(IsFiniteRaBitQValue(std::numeric_limits::max())); + REQUIRE_FALSE(IsFiniteRaBitQValue(std::numeric_limits::quiet_NaN())); + REQUIRE_FALSE(IsFiniteRaBitQValue(std::numeric_limits::infinity())); + REQUIRE_FALSE(IsFiniteRaBitQValue(-std::numeric_limits::infinity())); +} + TEST_CASE("RaBitQ Basic Test", "[ut][RaBitQuantizer]") { bool use_fht = GENERATE(true, false); auto num_bits_per_dim_query = GENERATE(4, 32); @@ -520,7 +540,7 @@ TEST_CASE("RaBitQ one-bit split code-code distance", "[ut][RaBitQuantizer]") { TEST_CASE("RaBitQ Split Code Storage", "[ut][RaBitQuantizer]") { auto allocator = SafeAllocator::FactoryDefaultAllocator(); - constexpr auto dim = 64; + const uint64_t dim = GENERATE(64, 960); constexpr auto count = 32; auto vecs = fixtures::generate_vectors(count, dim); @@ -576,11 +596,17 @@ TEST_CASE("RaBitQ Split Code Storage", "[ut][RaBitQuantizer]") { float one_bit_dist = 0.0F; float lower_bound = std::numeric_limits::max(); - REQUIRE(quantizer.ComputeDistWithOneBitLowerBound( - *computer, one_bit_code.data(), &one_bit_dist, &lower_bound)); + float filter_inner_product = 0.0F; + REQUIRE(quantizer.ComputeDistWithOneBitLowerBoundAndFilterIP( + *computer, one_bit_code.data(), &one_bit_dist, &lower_bound, &filter_inner_product)); REQUIRE(std::isfinite(one_bit_dist)); REQUIRE(std::isfinite(lower_bound)); REQUIRE(lower_bound <= one_bit_dist + 1e-5F); + if (filter_bits == 1) { + REQUIRE(IsNaNBitPattern(filter_inner_product)); + } else { + REQUIRE(std::isfinite(filter_inner_product)); + } if (filter_bits > 1) { float stored_filter_norm_code = 0.0F; @@ -604,15 +630,37 @@ TEST_CASE("RaBitQ Split Code Storage", "[ut][RaBitQuantizer]") { } const auto expected_filter_norm_code = static_cast(std::sqrt(filter_norm_sqr)); REQUIRE(std::abs(stored_filter_norm_code - expected_filter_norm_code) <= 1e-5F); + } - if (filter_bits == 2 or filter_bits == 3) { - float hinted_split_dist = 0.0F; + if (quantizer.ReorderBits() > 0 and + (filter_bits == 1 or filter_bits == 2 or filter_bits == 3)) { + float hinted_split_dist = 0.0F; + float direct_ip_split_dist = 0.0F; + if (filter_bits == 1) { + REQUIRE_FALSE( + quantizer.ComputeDistWithSplitCodeAndFilterDist(*computer, + one_bit_code.data(), + supplement_code.data(), + one_bit_dist, + &hinted_split_dist)); + REQUIRE_FALSE(quantizer.ComputeDistWithSplitCodeAndFilterIP(*computer, + one_bit_code.data(), + supplement_code.data(), + 0.0F, + &direct_ip_split_dist)); + } else { REQUIRE(quantizer.ComputeDistWithSplitCodeAndFilterDist(*computer, one_bit_code.data(), supplement_code.data(), one_bit_dist, &hinted_split_dist)); - REQUIRE(std::abs(split_dist - hinted_split_dist) <= 1e-5F); + REQUIRE(quantizer.ComputeDistWithSplitCodeAndFilterIP(*computer, + one_bit_code.data(), + supplement_code.data(), + filter_inner_product, + &direct_ip_split_dist)); + REQUIRE(std::abs(hinted_split_dist - direct_ip_split_dist) <= 1e-5F); + REQUIRE(std::abs(split_dist - direct_ip_split_dist) <= 1e-5F); } } @@ -679,10 +727,10 @@ TEST_CASE("RaBitQ Split Code Storage", "[ut][RaBitQuantizer]") { TEST_CASE("RaBitQ Split IP Batch4 and Reorder Hint", "[ut][RaBitQuantizer]") { auto allocator = SafeAllocator::FactoryDefaultAllocator(); - constexpr uint64_t dim = 64; + const uint64_t dim = GENERATE(64, 960); constexpr uint64_t count = 32; constexpr uint64_t base_bits = 8; - constexpr uint64_t filter_bits = 3; + const uint64_t filter_bits = GENERATE(1, 3); auto vecs = fixtures::generate_vectors(count, dim); RaBitQuantizer quantizer( @@ -713,18 +761,44 @@ TEST_CASE("RaBitQ Split IP Batch4 and Reorder Hint", "[ut][RaBitQuantizer]") { REQUIRE(quantizer.EncodeOne(vecs.data() + i * dim, full_code.data())); quantizer.SplitCode(full_code.data(), one_bit_codes[i].data(), supplement_code.data()); - REQUIRE(quantizer.ComputeDistWithOneBitLowerBound( - *computer, one_bit_codes[i].data(), single_dists + i, single_lower_bounds + i)); + float filter_inner_product = 0.0F; + REQUIRE(quantizer.ComputeDistWithOneBitLowerBoundAndFilterIP(*computer, + one_bit_codes[i].data(), + single_dists + i, + single_lower_bounds + i, + &filter_inner_product)); float split_dist = 0.0F; REQUIRE(quantizer.ComputeDistWithSplitCode( *computer, one_bit_codes[i].data(), supplement_code.data(), &split_dist)); float hinted_split_dist = 0.0F; - REQUIRE(quantizer.ComputeDistWithSplitCodeAndFilterDist(*computer, - one_bit_codes[i].data(), - supplement_code.data(), - single_dists[i], - &hinted_split_dist)); - REQUIRE(std::abs(split_dist - hinted_split_dist) <= 1e-5F); + float direct_ip_split_dist = 0.0F; + if (filter_bits == 1) { + REQUIRE(IsNaNBitPattern(filter_inner_product)); + REQUIRE_FALSE(quantizer.ComputeDistWithSplitCodeAndFilterDist(*computer, + one_bit_codes[i].data(), + supplement_code.data(), + single_dists[i], + &hinted_split_dist)); + REQUIRE_FALSE(quantizer.ComputeDistWithSplitCodeAndFilterIP(*computer, + one_bit_codes[i].data(), + supplement_code.data(), + 0.0F, + &direct_ip_split_dist)); + } else { + REQUIRE(std::isfinite(filter_inner_product)); + REQUIRE(quantizer.ComputeDistWithSplitCodeAndFilterDist(*computer, + one_bit_codes[i].data(), + supplement_code.data(), + single_dists[i], + &hinted_split_dist)); + REQUIRE(quantizer.ComputeDistWithSplitCodeAndFilterIP(*computer, + one_bit_codes[i].data(), + supplement_code.data(), + filter_inner_product, + &direct_ip_split_dist)); + REQUIRE(std::abs(hinted_split_dist - direct_ip_split_dist) <= 1e-5F); + REQUIRE(std::abs(split_dist - direct_ip_split_dist) <= 1e-5F); + } } float batch_dists[4] = {}; diff --git a/src/query_context.h b/src/query_context.h index e2e6b3462a..6f6f481a91 100644 --- a/src/query_context.h +++ b/src/query_context.h @@ -32,6 +32,7 @@ struct QueryContext { SearchStatistics* stats = nullptr; ReasoningContext* reasoning_ctx = nullptr; float rabitq_error_rate = std::numeric_limits::quiet_NaN(); + bool enable_rabitq_reorder = true; }; class SearchStatistics { diff --git a/src/simd/avx2.cpp b/src/simd/avx2.cpp index b20e2a72d9..bb1393673d 100644 --- a/src/simd/avx2.cpp +++ b/src/simd/avx2.cpp @@ -38,6 +38,243 @@ avx2_reduce_add_ps(__m256 a) { namespace vsag::avx2 { +uint64_t +RaBitQSQ4UBinaryIPWithBaseSum(const uint8_t* codes, const uint8_t* bits, uint64_t dim) { +#if defined(ENABLE_AVX2) + const uint64_t num_bytes = (dim + 7) / 8; + const __m256i lookup = _mm256_setr_epi8(0, + 1, + 1, + 2, + 1, + 2, + 2, + 3, + 1, + 2, + 2, + 3, + 2, + 3, + 3, + 4, + 0, + 1, + 1, + 2, + 1, + 2, + 2, + 3, + 1, + 2, + 2, + 3, + 2, + 3, + 3, + 4); + const __m256i low_mask = _mm256_set1_epi8(0x0F); + const auto popcount = [&lookup, &low_mask](__m256i value) { + const auto low = _mm256_and_si256(value, low_mask); + const auto high = _mm256_and_si256(_mm256_srli_epi16(value, 4), low_mask); + const auto counts = + _mm256_add_epi8(_mm256_shuffle_epi8(lookup, low), _mm256_shuffle_epi8(lookup, high)); + return _mm256_sad_epu8(counts, _mm256_setzero_si256()); + }; + + __m256i base_acc = _mm256_setzero_si256(); + __m256i inner_acc[4] = {_mm256_setzero_si256(), + _mm256_setzero_si256(), + _mm256_setzero_si256(), + _mm256_setzero_si256()}; + uint64_t offset = 0; + for (; offset + 32 <= num_bytes; offset += 32) { + const auto base = _mm256_loadu_si256(reinterpret_cast(bits + offset)); + base_acc = _mm256_add_epi64(base_acc, popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + const auto query = _mm256_loadu_si256( + reinterpret_cast(codes + bit * num_bytes + offset)); + inner_acc[bit] = + _mm256_add_epi64(inner_acc[bit], popcount(_mm256_and_si256(query, base))); + } + } + + const uint64_t remaining_words = (num_bytes - offset) / sizeof(int32_t); + if (remaining_words > 0) { + const auto lane_ids = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + const auto mask = + _mm256_cmpgt_epi32(_mm256_set1_epi32(static_cast(remaining_words)), lane_ids); + const auto base = + _mm256_maskload_epi32(reinterpret_cast(bits + offset), mask); + base_acc = _mm256_add_epi64(base_acc, popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + const auto query = _mm256_maskload_epi32( + reinterpret_cast(codes + bit * num_bytes + offset), mask); + inner_acc[bit] = + _mm256_add_epi64(inner_acc[bit], popcount(_mm256_and_si256(query, base))); + } + offset += remaining_words * sizeof(int32_t); + } + + alignas(32) uint64_t lanes[4]; + _mm256_store_si256(reinterpret_cast<__m256i*>(lanes), base_acc); + uint32_t base_sum = static_cast(lanes[0] + lanes[1] + lanes[2] + lanes[3]); + uint32_t inner_product = 0; + for (uint32_t bit = 0; bit < 4; ++bit) { + _mm256_store_si256(reinterpret_cast<__m256i*>(lanes), inner_acc[bit]); + inner_product += static_cast(lanes[0] + lanes[1] + lanes[2] + lanes[3]) << bit; + } + for (; offset < num_bytes; ++offset) { + const auto base = bits[offset]; + base_sum += static_cast(__builtin_popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_product += + static_cast(__builtin_popcount(codes[bit * num_bytes + offset] & base)) + << bit; + } + } + return static_cast(inner_product) | (static_cast(base_sum) << 32U); +#else + return generic::RaBitQSQ4UBinaryIPWithBaseSum(codes, bits, dim); +#endif +} + +void +RaBitQSQ4UBinaryIPWithBaseSumBatch4(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results) { +#if defined(ENABLE_AVX2) + const uint64_t num_bytes = (dim + 7) / 8; + const uint8_t* bases[4] = {bits1, bits2, bits3, bits4}; + const __m256i lookup = _mm256_setr_epi8(0, + 1, + 1, + 2, + 1, + 2, + 2, + 3, + 1, + 2, + 2, + 3, + 2, + 3, + 3, + 4, + 0, + 1, + 1, + 2, + 1, + 2, + 2, + 3, + 1, + 2, + 2, + 3, + 2, + 3, + 3, + 4); + const __m256i low_mask = _mm256_set1_epi8(0x0F); + const auto popcount = [&lookup, &low_mask](__m256i value) { + const auto low = _mm256_and_si256(value, low_mask); + const auto high = _mm256_and_si256(_mm256_srli_epi16(value, 4), low_mask); + const auto counts = + _mm256_add_epi8(_mm256_shuffle_epi8(lookup, low), _mm256_shuffle_epi8(lookup, high)); + return _mm256_sad_epu8(counts, _mm256_setzero_si256()); + }; + + __m256i base_acc[4]; + __m256i inner_acc[4][4]; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + base_acc[base_id] = _mm256_setzero_si256(); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_acc[base_id][bit] = _mm256_setzero_si256(); + } + } + + uint64_t offset = 0; + for (; offset + 32 <= num_bytes; offset += 32) { + __m256i base_values[4]; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + base_values[base_id] = + _mm256_loadu_si256(reinterpret_cast(bases[base_id] + offset)); + base_acc[base_id] = _mm256_add_epi64(base_acc[base_id], popcount(base_values[base_id])); + } + for (uint32_t bit = 0; bit < 4; ++bit) { + const auto query = _mm256_loadu_si256( + reinterpret_cast(codes + bit * num_bytes + offset)); + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + inner_acc[base_id][bit] = + _mm256_add_epi64(inner_acc[base_id][bit], + popcount(_mm256_and_si256(query, base_values[base_id]))); + } + } + } + + const uint64_t remaining_words = (num_bytes - offset) / sizeof(int32_t); + if (remaining_words > 0) { + const auto lane_ids = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + const auto mask = + _mm256_cmpgt_epi32(_mm256_set1_epi32(static_cast(remaining_words)), lane_ids); + __m256i base_values[4]; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + base_values[base_id] = _mm256_maskload_epi32( + reinterpret_cast(bases[base_id] + offset), mask); + base_acc[base_id] = _mm256_add_epi64(base_acc[base_id], popcount(base_values[base_id])); + } + for (uint32_t bit = 0; bit < 4; ++bit) { + const auto query = _mm256_maskload_epi32( + reinterpret_cast(codes + bit * num_bytes + offset), mask); + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + inner_acc[base_id][bit] = + _mm256_add_epi64(inner_acc[base_id][bit], + popcount(_mm256_and_si256(query, base_values[base_id]))); + } + } + offset += remaining_words * sizeof(int32_t); + } + + uint32_t inner_products[4] = {0, 0, 0, 0}; + uint32_t base_sums[4] = {0, 0, 0, 0}; + alignas(32) uint64_t lanes[4]; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + _mm256_store_si256(reinterpret_cast<__m256i*>(lanes), base_acc[base_id]); + base_sums[base_id] = static_cast(lanes[0] + lanes[1] + lanes[2] + lanes[3]); + for (uint32_t bit = 0; bit < 4; ++bit) { + _mm256_store_si256(reinterpret_cast<__m256i*>(lanes), inner_acc[base_id][bit]); + inner_products[base_id] += + static_cast(lanes[0] + lanes[1] + lanes[2] + lanes[3]) << bit; + } + } + for (; offset < num_bytes; ++offset) { + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + const auto base = bases[base_id][offset]; + base_sums[base_id] += static_cast(__builtin_popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_products[base_id] += static_cast(__builtin_popcount( + codes[bit * num_bytes + offset] & base)) + << bit; + } + } + } + for (uint32_t i = 0; i < 4; ++i) { + results[i] = + static_cast(inner_products[i]) | (static_cast(base_sums[i]) << 32U); + } +#else + generic::RaBitQSQ4UBinaryIPWithBaseSumBatch4(codes, bits1, bits2, bits3, bits4, dim, results); +#endif +} + float L2Sqr(const void* pVect1v, const void* pVect2v, const void* qty_ptr) { auto* pVect1 = (float*)pVect1v; @@ -846,6 +1083,39 @@ RaBitQFloatThreeBitCenteredIPBatch4(const float* vector, #endif } +float +RaBitQFloatFourBitCenteredIP(const float* vector, const uint8_t* bits, uint64_t dim) { +#if defined(ENABLE_AVX2) + return simd::RaBitQFloatFourBitCenteredIPImpl>( + vector, bits, dim, &generic::RaBitQFloatFourBitCenteredIP); +#else + return generic::RaBitQFloatFourBitCenteredIP(vector, bits, dim); +#endif +} + +void +RaBitQFloatFourBitCenteredIPBatch4(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results) { +#if defined(ENABLE_AVX2) + simd::RaBitQFloatFourBitCenteredIPBatch4Impl>( + vector, + bits1, + bits2, + bits3, + bits4, + dim, + results, + &generic::RaBitQFloatFourBitCenteredIPBatch4); +#else + generic::RaBitQFloatFourBitCenteredIPBatch4(vector, bits1, bits2, bits3, bits4, dim, results); +#endif +} + float RaBitQFloatThreeBitIPByLookup(const float* lookup, const uint8_t* bits, @@ -1036,6 +1306,68 @@ RaBitQFloatSupplementCodeIP(const float* vector, #endif } +// Ported from RaBitQ-Library's Apache-2.0 ip64_fxu7_avx2 kernel. +float +RaBitQFloatExCode7IP(const float* vector, const uint8_t* compact_code, uint64_t dim) { +#if defined(ENABLE_AVX2) + if ((dim & 63U) != 0U) { + return generic::RaBitQFloatExCode7IP(vector, compact_code, dim); + } + const __m128i mask6 = _mm_set1_epi8(0x3F); + const __m128i mask2 = _mm_set1_epi8(static_cast(0xC0)); + const __m128i top_mask = _mm_set1_epi8(0x40); + __m256 sum = _mm256_setzero_ps(); + + const auto contribute = [&sum](__m128i codes, const float* query) { + __m256 q = _mm256_loadu_ps(query); + __m256 cf = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(codes)); + sum = _mm256_fmadd_ps(q, cf, sum); + q = _mm256_loadu_ps(query + 8); + cf = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(_mm_srli_si128(codes, 8))); + sum = _mm256_fmadd_ps(q, cf, sum); + }; + + for (uint64_t block = 0; block < dim; block += 64) { + const __m128i compact1 = _mm_loadu_si128(reinterpret_cast(compact_code)); + const __m128i compact2 = + _mm_loadu_si128(reinterpret_cast(compact_code + 16)); + const __m128i compact3 = + _mm_loadu_si128(reinterpret_cast(compact_code + 32)); + uint64_t top_bits = 0; + std::memcpy(&top_bits, compact_code + 48, sizeof(top_bits)); + compact_code += 56; + + __m128i code0 = _mm_and_si128(compact1, mask6); + __m128i code1 = _mm_and_si128(compact2, mask6); + __m128i code2 = _mm_and_si128(compact3, mask6); + __m128i code3 = + _mm_or_si128(_mm_or_si128(_mm_srli_epi16(_mm_and_si128(compact1, mask2), 6), + _mm_srli_epi16(_mm_and_si128(compact2, mask2), 4)), + _mm_srli_epi16(_mm_and_si128(compact3, mask2), 2)); + + code0 = _mm_or_si128(code0, + _mm_and_si128(_mm_set_epi64x(top_bits << 5, top_bits << 6), top_mask)); + code1 = _mm_or_si128(code1, + _mm_and_si128(_mm_set_epi64x(top_bits << 3, top_bits << 4), top_mask)); + code2 = _mm_or_si128(code2, + _mm_and_si128(_mm_set_epi64x(top_bits << 1, top_bits << 2), top_mask)); + code3 = + _mm_or_si128(code3, _mm_and_si128(_mm_set_epi64x(top_bits >> 1, top_bits), top_mask)); + + contribute(code0, vector + block); + contribute(code1, vector + block + 16); + contribute(code2, vector + block + 32); + contribute(code3, vector + block + 48); + } + + alignas(32) float lanes[8]; + _mm256_store_ps(lanes, sum); + return lanes[0] + lanes[1] + lanes[2] + lanes[3] + lanes[4] + lanes[5] + lanes[6] + lanes[7]; +#else + return generic::RaBitQFloatExCode7IP(vector, compact_code, dim); +#endif +} + void DivScalar(const float* from, float* to, uint64_t dim, float scalar) { #if defined(ENABLE_AVX2) diff --git a/src/simd/avx512.cpp b/src/simd/avx512.cpp index b2f069c49d..7e4d1b01b6 100644 --- a/src/simd/avx512.cpp +++ b/src/simd/avx512.cpp @@ -801,6 +801,39 @@ RaBitQFloatThreeBitCenteredIPBatch4(const float* vector, #endif } +float +RaBitQFloatFourBitCenteredIP(const float* vector, const uint8_t* bits, uint64_t dim) { +#if defined(ENABLE_AVX512) + return simd::RaBitQFloatFourBitCenteredIPImpl>( + vector, bits, dim, &generic::RaBitQFloatFourBitCenteredIP); +#else + return avx2::RaBitQFloatFourBitCenteredIP(vector, bits, dim); +#endif +} + +void +RaBitQFloatFourBitCenteredIPBatch4(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results) { +#if defined(ENABLE_AVX512) + simd::RaBitQFloatFourBitCenteredIPBatch4Impl>( + vector, + bits1, + bits2, + bits3, + bits4, + dim, + results, + &generic::RaBitQFloatFourBitCenteredIPBatch4); +#else + avx2::RaBitQFloatFourBitCenteredIPBatch4(vector, bits1, bits2, bits3, bits4, dim, results); +#endif +} + float RaBitQFloatThreeBitIPByLookup(const float* lookup, const uint8_t* bits, @@ -1015,7 +1048,12 @@ RaBitQFloatSupplementCodeIP(const float* vector, const uint8_t* supplement_code, uint64_t dim, uint32_t supplement_bits) { +#if defined(ENABLE_AVX512) + return simd::RaBitQFloatSupplementCodeIPImpl>( + vector, supplement_code, dim, supplement_bits); +#else return avx2::RaBitQFloatSupplementCodeIP(vector, supplement_code, dim, supplement_bits); +#endif } void diff --git a/src/simd/avx512vpopcntdq.cpp b/src/simd/avx512vpopcntdq.cpp index 13592c0a84..86c98ad36d 100644 --- a/src/simd/avx512vpopcntdq.cpp +++ b/src/simd/avx512vpopcntdq.cpp @@ -20,6 +20,115 @@ namespace vsag::avx512vpopcntdq { +uint64_t +RaBitQSQ4UBinaryIPWithBaseSum(const uint8_t* codes, const uint8_t* bits, uint64_t dim) { +#if defined(ENABLE_AVX512VPOPCNTDQ) + const uint64_t num_bytes = (dim + 7) / 8; + __m512i base_acc = _mm512_setzero_si512(); + __m512i inner_acc[4] = {_mm512_setzero_si512(), + _mm512_setzero_si512(), + _mm512_setzero_si512(), + _mm512_setzero_si512()}; + uint64_t offset = 0; + for (; offset + 64 <= num_bytes; offset += 64) { + const auto base = _mm512_loadu_si512(reinterpret_cast(bits + offset)); + base_acc = _mm512_add_epi64(base_acc, _mm512_popcnt_epi64(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + const auto query = _mm512_loadu_si512( + reinterpret_cast(codes + bit * num_bytes + offset)); + inner_acc[bit] = _mm512_add_epi64(inner_acc[bit], + _mm512_popcnt_epi64(_mm512_and_si512(query, base))); + } + } + uint32_t base_sum = static_cast(_mm512_reduce_add_epi64(base_acc)); + uint32_t inner_product = 0; + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_product += static_cast(_mm512_reduce_add_epi64(inner_acc[bit])) << bit; + } + for (; offset < num_bytes; ++offset) { + const auto base = bits[offset]; + base_sum += static_cast(__builtin_popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_product += + static_cast(__builtin_popcount(codes[bit * num_bytes + offset] & base)) + << bit; + } + } + return static_cast(inner_product) | (static_cast(base_sum) << 32U); +#else + return avx2::RaBitQSQ4UBinaryIPWithBaseSum(codes, bits, dim); +#endif +} + +void +RaBitQSQ4UBinaryIPWithBaseSumBatch4(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results) { +#if defined(ENABLE_AVX512VPOPCNTDQ) + const uint64_t num_bytes = (dim + 7) / 8; + const uint8_t* bases[4] = {bits1, bits2, bits3, bits4}; + __m512i base_acc[4]; + __m512i inner_acc[4][4]; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + base_acc[base_id] = _mm512_setzero_si512(); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_acc[base_id][bit] = _mm512_setzero_si512(); + } + } + + uint64_t offset = 0; + for (; offset + 64 <= num_bytes; offset += 64) { + __m512i base_values[4]; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + base_values[base_id] = + _mm512_loadu_si512(reinterpret_cast(bases[base_id] + offset)); + base_acc[base_id] = + _mm512_add_epi64(base_acc[base_id], _mm512_popcnt_epi64(base_values[base_id])); + } + for (uint32_t bit = 0; bit < 4; ++bit) { + const auto query = _mm512_loadu_si512( + reinterpret_cast(codes + bit * num_bytes + offset)); + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + inner_acc[base_id][bit] = _mm512_add_epi64( + inner_acc[base_id][bit], + _mm512_popcnt_epi64(_mm512_and_si512(query, base_values[base_id]))); + } + } + } + + uint32_t inner_products[4] = {0, 0, 0, 0}; + uint32_t base_sums[4] = {0, 0, 0, 0}; + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + base_sums[base_id] = static_cast(_mm512_reduce_add_epi64(base_acc[base_id])); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_products[base_id] += + static_cast(_mm512_reduce_add_epi64(inner_acc[base_id][bit])) << bit; + } + } + for (; offset < num_bytes; ++offset) { + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + const auto base = bases[base_id][offset]; + base_sums[base_id] += static_cast(__builtin_popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_products[base_id] += static_cast(__builtin_popcount( + codes[bit * num_bytes + offset] & base)) + << bit; + } + } + } + for (uint32_t i = 0; i < 4; ++i) { + results[i] = + static_cast(inner_products[i]) | (static_cast(base_sums[i]) << 32U); + } +#else + avx2::RaBitQSQ4UBinaryIPWithBaseSumBatch4(codes, bits1, bits2, bits3, bits4, dim, results); +#endif +} + uint32_t RaBitQSQ4UBinaryIP(const uint8_t* codes, const uint8_t* bits, uint64_t dim) { // require dim align with 512 diff --git a/src/simd/generic.cpp b/src/simd/generic.cpp index fa10fc7c77..707035a485 100644 --- a/src/simd/generic.cpp +++ b/src/simd/generic.cpp @@ -649,6 +649,72 @@ RaBitQFloatThreeBitCenteredIPBatch4(const float* vector, } } +float +RaBitQFloatFourBitCenteredIP(const float* vector, const uint8_t* bits, uint64_t dim) { + if (dim == 0) { + return 0.0F; + } + + const uint64_t plane_bytes = (dim + 7) / 8; + const uint8_t* plane0 = bits; + const uint8_t* plane1 = bits + plane_bytes; + const uint8_t* plane2 = bits + 2 * plane_bytes; + const uint8_t* plane3 = bits + 3 * plane_bytes; + float result = 0.0F; + for (uint64_t d = 0; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + float weight = (plane0[byte_idx] & bit_mask) != 0U ? 4.0F : -4.0F; + weight += (plane1[byte_idx] & bit_mask) != 0U ? 2.0F : -2.0F; + weight += (plane2[byte_idx] & bit_mask) != 0U ? 1.0F : -1.0F; + weight += (plane3[byte_idx] & bit_mask) != 0U ? 0.5F : -0.5F; + result += vector[d] * weight; + } + return result; +} + +void +RaBitQFloatFourBitCenteredIPBatch4(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results) { + results[0] = 0.0F; + results[1] = 0.0F; + results[2] = 0.0F; + results[3] = 0.0F; + if (dim == 0) { + return; + } + + const uint64_t plane_bytes = (dim + 7) / 8; + const uint8_t* plane0[4] = {bits1, bits2, bits3, bits4}; + const uint8_t* plane1[4] = { + bits1 + plane_bytes, bits2 + plane_bytes, bits3 + plane_bytes, bits4 + plane_bytes}; + const uint8_t* plane2[4] = {bits1 + 2 * plane_bytes, + bits2 + 2 * plane_bytes, + bits3 + 2 * plane_bytes, + bits4 + 2 * plane_bytes}; + const uint8_t* plane3[4] = {bits1 + 3 * plane_bytes, + bits2 + 3 * plane_bytes, + bits3 + 3 * plane_bytes, + bits4 + 3 * plane_bytes}; + for (uint64_t d = 0; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + const float value = vector[d]; + for (uint32_t i = 0; i < 4; ++i) { + float weight = (plane0[i][byte_idx] & bit_mask) != 0U ? 4.0F : -4.0F; + weight += (plane1[i][byte_idx] & bit_mask) != 0U ? 2.0F : -2.0F; + weight += (plane2[i][byte_idx] & bit_mask) != 0U ? 1.0F : -1.0F; + weight += (plane3[i][byte_idx] & bit_mask) != 0U ? 0.5F : -0.5F; + results[i] += value * weight; + } + } +} + void RaBitQFloatBuildByteIPLookupTable(const float* vector, uint64_t dim, float* lookup) { const uint64_t block_count = (dim + 7) / 8; @@ -796,6 +862,34 @@ RaBitQFloatSupplementCodeIP(const float* vector, return result; } +// Ported from RaBitQ-Library's Apache-2.0 7-bit ExData layout. +float +RaBitQFloatExCode7IP(const float* vector, const uint8_t* compact_code, uint64_t dim) { + if ((dim & 63U) != 0U) { + return 0.0F; + } + float result = 0.0F; + for (uint64_t block = 0; block < dim; block += 64) { + const uint8_t* low_codes = compact_code; + for (uint64_t lane = 0; lane < 48; ++lane) { + const uint32_t top = (compact_code[48 + (lane & 7U)] >> (lane >> 3U)) & 1U; + const uint32_t code = (low_codes[lane] & 0x3FU) | (top << 6U); + result += vector[block + lane] * static_cast(code); + } + for (uint64_t lane = 48; lane < 64; ++lane) { + const uint64_t packed_lane = lane - 48; + const uint32_t low = ((low_codes[packed_lane] >> 6U) & 0x3U) | + (((low_codes[16 + packed_lane] >> 6U) & 0x3U) << 2U) | + (((low_codes[32 + packed_lane] >> 6U) & 0x3U) << 4U); + const uint32_t top = (compact_code[48 + (lane & 7U)] >> (lane >> 3U)) & 1U; + const uint32_t code = low | (top << 6U); + result += vector[block + lane] * static_cast(code); + } + compact_code += 56; + } + return result; +} + float RaBitQFloatSQIP(const float* vector, const uint8_t* codes, uint64_t dim) { if (dim == 0) { @@ -871,6 +965,51 @@ RaBitQSQ4UBinaryIP(const uint8_t* codes, const uint8_t* bits, uint64_t dim) { return result; } +uint64_t +RaBitQSQ4UBinaryIPWithBaseSum(const uint8_t* codes, const uint8_t* bits, uint64_t dim) { + uint32_t inner_product = 0; + uint32_t base_sum = 0; + const uint64_t num_bytes = (dim + 7) / 8; + for (uint64_t i = 0; i < num_bytes; ++i) { + const auto base = bits[i]; + base_sum += static_cast(__builtin_popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_product += + static_cast(__builtin_popcount(codes[bit * num_bytes + i] & base)) << bit; + } + } + return static_cast(inner_product) | (static_cast(base_sum) << 32U); +} + +void +RaBitQSQ4UBinaryIPWithBaseSumBatch4(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results) { + uint32_t inner_products[4] = {0, 0, 0, 0}; + uint32_t base_sums[4] = {0, 0, 0, 0}; + const uint8_t* bases[4] = {bits1, bits2, bits3, bits4}; + const uint64_t num_bytes = (dim + 7) / 8; + for (uint64_t i = 0; i < num_bytes; ++i) { + for (uint32_t base_id = 0; base_id < 4; ++base_id) { + const auto base = bases[base_id][i]; + base_sums[base_id] += static_cast(__builtin_popcount(base)); + for (uint32_t bit = 0; bit < 4; ++bit) { + inner_products[base_id] += + static_cast(__builtin_popcount(codes[bit * num_bytes + i] & base)) + << bit; + } + } + } + for (uint32_t i = 0; i < 4; ++i) { + results[i] = + static_cast(inner_products[i]) | (static_cast(base_sums[i]) << 32U); + } +} + float Normalize(const float* from, float* to, uint64_t dim) { float norm = std::sqrt(FP32ComputeIP(from, from, dim)); diff --git a/src/simd/kernels/rabitq_compute.h b/src/simd/kernels/rabitq_compute.h index d31391ebac..d9216eac19 100644 --- a/src/simd/kernels/rabitq_compute.h +++ b/src/simd/kernels/rabitq_compute.h @@ -357,6 +357,141 @@ RaBitQFloatThreeBitCenteredIPBatch4Impl(const float* vector, } } +template +inline float +RaBitQFloatFourBitCenteredIPImpl(const float* vector, + const uint8_t* bits, + uint64_t dim, + float (*fallback)(const float*, const uint8_t*, uint64_t)) { + if (dim == 0) { + return 0.0F; + } + + constexpr int W = T::Width; + if (dim < static_cast(W)) { + return fallback(vector, bits, dim); + } + + const uint64_t plane_bytes = (dim + 7) / 8; + const uint8_t* plane0 = bits; + const uint8_t* plane1 = bits + plane_bytes; + const uint8_t* plane2 = bits + 2 * plane_bytes; + const uint8_t* plane3 = bits + 3 * plane_bytes; + auto sum = T::zero(); + const auto pos4 = T::set1(4.0F); + const auto neg4 = T::set1(-4.0F); + const auto pos2 = T::set1(2.0F); + const auto neg2 = T::set1(-2.0F); + const auto pos1 = T::set1(1.0F); + const auto neg1 = T::set1(-1.0F); + const auto pos_half = T::set1(0.5F); + const auto neg_half = T::set1(-0.5F); + + uint64_t d = 0; + for (; d + W <= dim; d += W) { + const uint64_t byte_idx = d >> 3; + auto weight = T::bits_to_signed(plane0 + byte_idx, pos4, neg4); + weight = T::add(weight, T::bits_to_signed(plane1 + byte_idx, pos2, neg2)); + weight = T::add(weight, T::bits_to_signed(plane2 + byte_idx, pos1, neg1)); + weight = T::add(weight, T::bits_to_signed(plane3 + byte_idx, pos_half, neg_half)); + + const auto vec = T::load(vector + d); + sum = T::fmadd(weight, vec, sum); + } + + float result = T::reduce_add(sum); + for (; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + const float value = vector[d]; + float weight = (plane0[byte_idx] & bit_mask) != 0U ? 4.0F : -4.0F; + weight += (plane1[byte_idx] & bit_mask) != 0U ? 2.0F : -2.0F; + weight += (plane2[byte_idx] & bit_mask) != 0U ? 1.0F : -1.0F; + weight += (plane3[byte_idx] & bit_mask) != 0U ? 0.5F : -0.5F; + result += value * weight; + } + return result; +} + +template +inline void +RaBitQFloatFourBitCenteredIPBatch4Impl(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results, + void (*fallback)(const float*, + const uint8_t*, + const uint8_t*, + const uint8_t*, + const uint8_t*, + uint64_t, + float*)) { + if (dim == 0) { + results[0] = results[1] = results[2] = results[3] = 0.0F; + return; + } + + constexpr int W = T::Width; + if (dim < static_cast(W)) { + fallback(vector, bits1, bits2, bits3, bits4, dim, results); + return; + } + + const uint64_t plane_bytes = (dim + 7) / 8; + const uint8_t* plane0[4] = {bits1, bits2, bits3, bits4}; + const uint8_t* plane1[4] = { + bits1 + plane_bytes, bits2 + plane_bytes, bits3 + plane_bytes, bits4 + plane_bytes}; + const uint8_t* plane2[4] = {bits1 + 2 * plane_bytes, + bits2 + 2 * plane_bytes, + bits3 + 2 * plane_bytes, + bits4 + 2 * plane_bytes}; + const uint8_t* plane3[4] = {bits1 + 3 * plane_bytes, + bits2 + 3 * plane_bytes, + bits3 + 3 * plane_bytes, + bits4 + 3 * plane_bytes}; + typename T::FloatVec sums[4] = {T::zero(), T::zero(), T::zero(), T::zero()}; + const auto pos4 = T::set1(4.0F); + const auto neg4 = T::set1(-4.0F); + const auto pos2 = T::set1(2.0F); + const auto neg2 = T::set1(-2.0F); + const auto pos1 = T::set1(1.0F); + const auto neg1 = T::set1(-1.0F); + const auto pos_half = T::set1(0.5F); + const auto neg_half = T::set1(-0.5F); + + uint64_t d = 0; + for (; d + W <= dim; d += W) { + const uint64_t byte_idx = d >> 3; + const auto vec = T::load(vector + d); + for (uint32_t i = 0; i < 4; ++i) { + auto weight = T::bits_to_signed(plane0[i] + byte_idx, pos4, neg4); + weight = T::add(weight, T::bits_to_signed(plane1[i] + byte_idx, pos2, neg2)); + weight = T::add(weight, T::bits_to_signed(plane2[i] + byte_idx, pos1, neg1)); + weight = T::add(weight, T::bits_to_signed(plane3[i] + byte_idx, pos_half, neg_half)); + sums[i] = T::fmadd(weight, vec, sums[i]); + } + } + + for (uint32_t i = 0; i < 4; ++i) { + results[i] = T::reduce_add(sums[i]); + } + for (; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + const float value = vector[d]; + for (uint32_t i = 0; i < 4; ++i) { + float weight = (plane0[i][byte_idx] & bit_mask) != 0U ? 4.0F : -4.0F; + weight += (plane1[i][byte_idx] & bit_mask) != 0U ? 2.0F : -2.0F; + weight += (plane2[i][byte_idx] & bit_mask) != 0U ? 1.0F : -1.0F; + weight += (plane3[i][byte_idx] & bit_mask) != 0U ? 0.5F : -0.5F; + results[i] += value * weight; + } + } +} + template inline float RaBitQFloatSplitCodeIPImpl(const float* vector, @@ -409,4 +544,46 @@ RaBitQFloatSplitCodeIPImpl(const float* vector, return result; } +template +inline float +RaBitQFloatSupplementCodeIPImpl(const float* vector, + const uint8_t* supplement_code, + uint64_t dim, + uint32_t supplement_bits) { + if (dim == 0 or supplement_bits == 0) { + return 0.0F; + } + + constexpr int W = T::Width; + const uint64_t plane_bytes = (dim + 7) / 8; + auto sum = T::zero(); + + uint64_t d = 0; + for (; d + W <= dim; d += W) { + const uint64_t byte_idx = d >> 3; + auto code = T::zero(); + for (uint32_t bit = 0; bit < supplement_bits; ++bit) { + const auto* plane = supplement_code + static_cast(bit) * plane_bytes; + const auto weight = T::set1(static_cast(1U << bit)); + code = T::add(code, T::bits_select(plane + byte_idx, weight)); + } + sum = T::fmadd(code, T::load(vector + d), sum); + } + + float result = T::reduce_add(sum); + for (; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + uint32_t code = 0; + for (uint32_t bit = 0; bit < supplement_bits; ++bit) { + const auto* plane = supplement_code + static_cast(bit) * plane_bytes; + if ((plane[byte_idx] & bit_mask) != 0U) { + code += 1U << bit; + } + } + result += vector[d] * static_cast(code); + } + return result; +} + } // namespace vsag::simd diff --git a/src/simd/rabitq_simd.cpp b/src/simd/rabitq_simd.cpp index c85948c7ea..963bd26c52 100644 --- a/src/simd/rabitq_simd.cpp +++ b/src/simd/rabitq_simd.cpp @@ -23,7 +23,49 @@ VSAG_DEFINE_SIMD_DISPATCH(RaBitQFloatBinaryIPBatch4, RaBitQFloatBinaryBatch4Type VSAG_DEFINE_SIMD_DISPATCH(RaBitQFloatThreeBitIPBatch4, RaBitQFloatThreeBitBatch4Type); VSAG_DEFINE_SIMD_DISPATCH(RaBitQFloatSplitCodeIP, RaBitQFloatSplitCodeType); VSAG_DEFINE_SIMD_DISPATCH(RaBitQFloatSupplementCodeIP, RaBitQFloatSupplementCodeType); +static RaBitQFloatExCode7Type +GetRaBitQFloatExCode7IP() { + if (SimdStatus::SupportAVX2()) { +#if defined(ENABLE_AVX2) + return avx2::RaBitQFloatExCode7IP; +#endif + } + return generic::RaBitQFloatExCode7IP; +} +RaBitQFloatExCode7Type RaBitQFloatExCode7IP = GetRaBitQFloatExCode7IP(); VSAG_DEFINE_SIMD_DISPATCH_VPOPCNTDQ(RaBitQSQ4UBinaryIP, RaBitQSQ4UBinaryType); +static RaBitQSQ4UBinaryWithBaseSumType +GetRaBitQSQ4UBinaryIPWithBaseSum() { + if (SimdStatus::SupportAVX512VPOPCNTDQ()) { +#if defined(ENABLE_AVX512VPOPCNTDQ) + return avx512vpopcntdq::RaBitQSQ4UBinaryIPWithBaseSum; +#endif + } + if (SimdStatus::SupportAVX2()) { +#if defined(ENABLE_AVX2) + return avx2::RaBitQSQ4UBinaryIPWithBaseSum; +#endif + } + return generic::RaBitQSQ4UBinaryIPWithBaseSum; +} +RaBitQSQ4UBinaryWithBaseSumType RaBitQSQ4UBinaryIPWithBaseSum = GetRaBitQSQ4UBinaryIPWithBaseSum(); +static RaBitQSQ4UBinaryWithBaseSumBatch4Type +GetRaBitQSQ4UBinaryIPWithBaseSumBatch4() { + if (SimdStatus::SupportAVX512VPOPCNTDQ()) { +#if defined(ENABLE_AVX512VPOPCNTDQ) + return avx512vpopcntdq::RaBitQSQ4UBinaryIPWithBaseSumBatch4; +#endif + } + if (SimdStatus::SupportAVX2()) { +#if defined(ENABLE_AVX2) + return avx2::RaBitQSQ4UBinaryIPWithBaseSumBatch4; +#endif + } + return generic::RaBitQSQ4UBinaryIPWithBaseSumBatch4; +} +RaBitQSQ4UBinaryWithBaseSumBatch4Type RaBitQSQ4UBinaryIPWithBaseSumBatch4 = + GetRaBitQSQ4UBinaryIPWithBaseSumBatch4(); + static RaBitQCodeCodeType GetRaBitQCodeCodeIP() { if (SimdStatus::SupportAVX512()) { @@ -148,6 +190,39 @@ GetRaBitQFloatThreeBitCenteredIPBatch4() { RaBitQFloatThreeBitCenteredBatch4Type RaBitQFloatThreeBitCenteredIPBatch4 = GetRaBitQFloatThreeBitCenteredIPBatch4(); +static RaBitQFloatFourBitCenteredType +GetRaBitQFloatFourBitCenteredIP() { + if (SimdStatus::SupportAVX512()) { +#if defined(ENABLE_AVX512) + return avx512::RaBitQFloatFourBitCenteredIP; +#endif + } + if (SimdStatus::SupportAVX2()) { +#if defined(ENABLE_AVX2) + return avx2::RaBitQFloatFourBitCenteredIP; +#endif + } + return generic::RaBitQFloatFourBitCenteredIP; +} +RaBitQFloatFourBitCenteredType RaBitQFloatFourBitCenteredIP = GetRaBitQFloatFourBitCenteredIP(); + +static RaBitQFloatFourBitCenteredBatch4Type +GetRaBitQFloatFourBitCenteredIPBatch4() { + if (SimdStatus::SupportAVX512()) { +#if defined(ENABLE_AVX512) + return avx512::RaBitQFloatFourBitCenteredIPBatch4; +#endif + } + if (SimdStatus::SupportAVX2()) { +#if defined(ENABLE_AVX2) + return avx2::RaBitQFloatFourBitCenteredIPBatch4; +#endif + } + return generic::RaBitQFloatFourBitCenteredIPBatch4; +} +RaBitQFloatFourBitCenteredBatch4Type RaBitQFloatFourBitCenteredIPBatch4 = + GetRaBitQFloatFourBitCenteredIPBatch4(); + static RaBitQFloatThreeBitByLookupType GetRaBitQFloatThreeBitIPByLookup() { if (SimdStatus::SupportAVX512()) { diff --git a/src/simd/rabitq_simd.h b/src/simd/rabitq_simd.h index 48c8b701de..3253313399 100644 --- a/src/simd/rabitq_simd.h +++ b/src/simd/rabitq_simd.h @@ -24,6 +24,18 @@ namespace avx512vpopcntdq { uint32_t RaBitQSQ4UBinaryIP(const uint8_t* codes, const uint8_t* bits, uint64_t dim); +uint64_t +RaBitQSQ4UBinaryIPWithBaseSum(const uint8_t* codes, const uint8_t* bits, uint64_t dim); + +void +RaBitQSQ4UBinaryIPWithBaseSumBatch4(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results); + } // namespace avx512vpopcntdq namespace avx512 { @@ -77,6 +89,18 @@ RaBitQFloatThreeBitCenteredIPBatch4(const float* vector, uint64_t dim, float* results); +float +RaBitQFloatFourBitCenteredIP(const float* vector, const uint8_t* bits, uint64_t dim); + +void +RaBitQFloatFourBitCenteredIPBatch4(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results); + float RaBitQFloatThreeBitIPByLookup(const float* lookup, const uint8_t* bits, @@ -154,6 +178,18 @@ RotateOp(float* data, int idx, int dim_, int step); } // namespace avx512 namespace avx2 { +uint64_t +RaBitQSQ4UBinaryIPWithBaseSum(const uint8_t* codes, const uint8_t* bits, uint64_t dim); + +void +RaBitQSQ4UBinaryIPWithBaseSumBatch4(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results); + float RaBitQFloatSQIP(const float* vector, const uint8_t* codes, uint64_t dim); @@ -214,6 +250,18 @@ RaBitQFloatThreeBitCenteredIPBatch4(const float* vector, uint64_t dim, float* results); +float +RaBitQFloatFourBitCenteredIP(const float* vector, const uint8_t* bits, uint64_t dim); + +void +RaBitQFloatFourBitCenteredIPBatch4(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results); + float RaBitQFloatThreeBitIPByLookup(const float* lookup, const uint8_t* bits, @@ -261,6 +309,9 @@ RaBitQFloatSupplementCodeIP(const float* vector, uint64_t dim, uint32_t supplement_bits); +float +RaBitQFloatExCode7IP(const float* vector, const uint8_t* compact_code, uint64_t dim); + void FHTRotate(float* data, uint64_t dim_); @@ -425,6 +476,18 @@ RaBitQFloatThreeBitCenteredIPBatch4(const float* vector, uint64_t dim, float* results); +float +RaBitQFloatFourBitCenteredIP(const float* vector, const uint8_t* bits, uint64_t dim); + +void +RaBitQFloatFourBitCenteredIPBatch4(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results); + void RaBitQFloatBuildByteIPLookupTable(const float* vector, uint64_t dim, float* lookup); @@ -475,12 +538,27 @@ RaBitQFloatSupplementCodeIP(const float* vector, uint64_t dim, uint32_t supplement_bits); +float +RaBitQFloatExCode7IP(const float* vector, const uint8_t* compact_code, uint64_t dim); + float RaBitQFloatSQIP(const float* vector, const uint8_t* codes, uint64_t dim); uint32_t RaBitQSQ4UBinaryIP(const uint8_t* codes, const uint8_t* bits, uint64_t dim); +uint64_t +RaBitQSQ4UBinaryIPWithBaseSum(const uint8_t* codes, const uint8_t* bits, uint64_t dim); + +void +RaBitQSQ4UBinaryIPWithBaseSumBatch4(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results); + uint64_t RaBitQCodeCodeIP(const uint8_t* codes1, const uint8_t* codes2, uint64_t dim); void @@ -692,6 +770,18 @@ using RaBitQFloatThreeBitCenteredBatch4Type = void (*)(const float* vector, uint64_t dim, float* results); +using RaBitQFloatFourBitCenteredType = float (*)(const float* vector, + const uint8_t* bits, + uint64_t dim); + +using RaBitQFloatFourBitCenteredBatch4Type = void (*)(const float* vector, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + float* results); + using RaBitQFloatThreeBitByLookupType = float (*)(const float* lookup, const uint8_t* bits, uint64_t dim, @@ -732,8 +822,22 @@ using RaBitQFloatSupplementCodeType = float (*)(const float* vector, const uint8_t* supplement_code, uint64_t dim, uint32_t supplement_bits); +using RaBitQFloatExCode7Type = float (*)(const float* vector, + const uint8_t* compact_code, + uint64_t dim); using RaBitQSQ4UBinaryType = uint32_t (*)(const uint8_t* codes, const uint8_t* bits, uint64_t dim); +using RaBitQSQ4UBinaryWithBaseSumType = uint64_t (*)(const uint8_t* codes, + const uint8_t* bits, + uint64_t dim); + +using RaBitQSQ4UBinaryWithBaseSumBatch4Type = void (*)(const uint8_t* codes, + const uint8_t* bits1, + const uint8_t* bits2, + const uint8_t* bits3, + const uint8_t* bits4, + uint64_t dim, + uint64_t* results); using RaBitQCodeCodeType = uint64_t (*)(const uint8_t* codes1, const uint8_t* codes2, uint64_t dim); using RaBitQPackScalarToSplitPlanesType = void (*)(const uint8_t* scalar_codes, @@ -761,14 +865,19 @@ extern RaBitQFloatTwoBitCenteredType RaBitQFloatTwoBitCenteredIP; extern RaBitQFloatTwoBitCenteredBatch4Type RaBitQFloatTwoBitCenteredIPBatch4; extern RaBitQFloatThreeBitCenteredType RaBitQFloatThreeBitCenteredIP; extern RaBitQFloatThreeBitCenteredBatch4Type RaBitQFloatThreeBitCenteredIPBatch4; +extern RaBitQFloatFourBitCenteredType RaBitQFloatFourBitCenteredIP; +extern RaBitQFloatFourBitCenteredBatch4Type RaBitQFloatFourBitCenteredIPBatch4; extern RaBitQFloatThreeBitByLookupType RaBitQFloatThreeBitIPByLookup; extern RaBitQFloatThreeBitBatch4ByLookupType RaBitQFloatThreeBitIPBatch4ByLookup; extern RaBitQFloatMultiBitByLookupType RaBitQFloatMultiBitIPByLookup; extern RaBitQFloatMultiBitBatch4ByLookupType RaBitQFloatMultiBitIPBatch4ByLookup; extern RaBitQFloatSplitCodeType RaBitQFloatSplitCodeIP; extern RaBitQFloatSupplementCodeType RaBitQFloatSupplementCodeIP; +extern RaBitQFloatExCode7Type RaBitQFloatExCode7IP; extern RaBitQFloatSQType RaBitQFloatSQIP; extern RaBitQSQ4UBinaryType RaBitQSQ4UBinaryIP; +extern RaBitQSQ4UBinaryWithBaseSumType RaBitQSQ4UBinaryIPWithBaseSum; +extern RaBitQSQ4UBinaryWithBaseSumBatch4Type RaBitQSQ4UBinaryIPWithBaseSumBatch4; extern RaBitQCodeCodeType RaBitQCodeCodeIP; extern RaBitQPackScalarToSplitPlanesType RaBitQPackScalarToSplitPlanes; extern FHTRotateType FHTRotate; diff --git a/src/simd/rabitq_simd_test.cpp b/src/simd/rabitq_simd_test.cpp index 70fc697d03..27675b8e14 100644 --- a/src/simd/rabitq_simd_test.cpp +++ b/src/simd/rabitq_simd_test.cpp @@ -16,6 +16,7 @@ #include #include +#include #include "fp32_simd.h" #include "simd_status.h" @@ -130,6 +131,14 @@ TEST_CASE("RaBitQ SQ4U-BQ Compute Codes", "[ut][simd]") { for (auto dim = 0; dim < 17; dim++) { uint32_t result = generic::RaBitQSQ4UBinaryIP(codes.data(), bits.data(), dim); + const auto packed = generic::RaBitQSQ4UBinaryIPWithBaseSum(codes.data(), bits.data(), dim); + REQUIRE(static_cast(packed) == result); + uint32_t expected_base_sum = 0; + for (uint64_t i = 0; i < (static_cast(dim) + 7) / 8; ++i) { + expected_base_sum += static_cast(__builtin_popcount(bits[i])); + } + REQUIRE(static_cast(packed >> 32U) == expected_base_sum); + REQUIRE(avx2::RaBitQSQ4UBinaryIPWithBaseSum(codes.data(), bits.data(), dim) == packed); TEST_ACCURACY_SQ4(RaBitQSQ4UBinaryIP); if (dim == 0) { REQUIRE(result == 0); @@ -143,6 +152,119 @@ TEST_CASE("RaBitQ SQ4U-BQ Compute Codes", "[ut][simd]") { } } +TEST_CASE("RaBitQ SQ4U-BQ with base sum Batch4", "[ut][simd]") { + const std::vector dims = { + 0, 1, 7, 8, 9, 31, 32, 33, 63, 64, 65, 511, 512, 513, 959, 960, 961}; + + for (const auto dim : dims) { + const uint64_t num_bytes = (dim + 7) / 8; + std::vector query_codes(std::max(1, 4 * num_bytes)); + std::vector base_codes(std::max(1, 4 * num_bytes)); + for (uint64_t i = 0; i < query_codes.size(); ++i) { + query_codes[i] = static_cast(53U * i + 17U); + } + for (uint64_t i = 0; i < base_codes.size(); ++i) { + base_codes[i] = static_cast(71U * i + 29U); + } + + if ((dim & 7U) != 0U) { + const auto valid_mask = static_cast((1U << (dim & 7U)) - 1U); + for (uint32_t bit = 0; bit < 4; ++bit) { + query_codes[bit * num_bytes + num_bytes - 1] &= valid_mask; + base_codes[bit * num_bytes + num_bytes - 1] &= valid_mask; + } + } + + const uint8_t* bits1 = base_codes.data(); + const uint8_t* bits2 = bits1 + num_bytes; + const uint8_t* bits3 = bits2 + num_bytes; + const uint8_t* bits4 = bits3 + num_bytes; + const uint8_t* bases[4] = {bits1, bits2, bits3, bits4}; + uint32_t expected_ip[4] = {0, 0, 0, 0}; + uint32_t expected_base_sum[4] = {0, 0, 0, 0}; + for (uint64_t d = 0; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + uint32_t query_code = 0; + for (uint32_t bit = 0; bit < 4; ++bit) { + if ((query_codes[bit * num_bytes + byte_idx] & bit_mask) != 0U) { + query_code += 1U << bit; + } + } + for (uint32_t i = 0; i < 4; ++i) { + if ((bases[i][byte_idx] & bit_mask) != 0U) { + expected_ip[i] += query_code; + ++expected_base_sum[i]; + } + } + } + + uint64_t expected[4]; + for (uint32_t i = 0; i < 4; ++i) { + expected[i] = static_cast(expected_ip[i]) | + (static_cast(expected_base_sum[i]) << 32U); + REQUIRE(generic::RaBitQSQ4UBinaryIPWithBaseSum(query_codes.data(), bases[i], dim) == + expected[i]); + REQUIRE(RaBitQSQ4UBinaryIPWithBaseSum(query_codes.data(), bases[i], dim) == + expected[i]); + if (SimdStatus::SupportAVX2()) { + REQUIRE(avx2::RaBitQSQ4UBinaryIPWithBaseSum(query_codes.data(), bases[i], dim) == + expected[i]); + } + if (SimdStatus::SupportAVX512VPOPCNTDQ()) { + REQUIRE(avx512vpopcntdq::RaBitQSQ4UBinaryIPWithBaseSum( + query_codes.data(), bases[i], dim) == expected[i]); + } + } + + auto check_batch = [&expected](auto func, + const uint8_t* query, + const uint8_t* base1, + const uint8_t* base2, + const uint8_t* base3, + const uint8_t* base4, + uint64_t test_dim) { + uint64_t results[4] = {0, 0, 0, 0}; + func(query, base1, base2, base3, base4, test_dim, results); + for (uint32_t i = 0; i < 4; ++i) { + REQUIRE(results[i] == expected[i]); + } + }; + check_batch(generic::RaBitQSQ4UBinaryIPWithBaseSumBatch4, + query_codes.data(), + bits1, + bits2, + bits3, + bits4, + dim); + check_batch(RaBitQSQ4UBinaryIPWithBaseSumBatch4, + query_codes.data(), + bits1, + bits2, + bits3, + bits4, + dim); + if (SimdStatus::SupportAVX2()) { + check_batch(avx2::RaBitQSQ4UBinaryIPWithBaseSumBatch4, + query_codes.data(), + bits1, + bits2, + bits3, + bits4, + dim); + } + if (SimdStatus::SupportAVX512VPOPCNTDQ()) { + check_batch(avx512vpopcntdq::RaBitQSQ4UBinaryIPWithBaseSumBatch4, + query_codes.data(), + bits1, + bits2, + bits3, + bits4, + dim); + } + } +} + TEST_CASE("RaBitQ scalar-code inner products", "[ut][simd]") { const std::vector dims = {0, 1, 7, 8, 9, 31, 32, 33, 63, 64, 65, 128, 960, 4097}; @@ -562,6 +684,90 @@ TEST_CASE("RaBitQ FP32 three-bit centered SIMD Batch4 Compute Codes", "[ut][simd } } +TEST_CASE("RaBitQ FP32 four-bit centered SIMD Batch4 Compute Codes", "[ut][simd]") { + const std::vector dims = {0, 1, 7, 8, 9, 15, 16, 17, 63, 64, 65, 959, 960, 961}; + + for (const auto dim : dims) { + const uint64_t plane_bytes = (dim + 7) / 8; + std::vector query(dim); + for (uint64_t d = 0; d < dim; ++d) { + query[d] = static_cast(static_cast(d % 29) - 14) * 0.03125F; + } + + std::vector codes(std::max(1, plane_bytes * 4 * 4)); + for (uint64_t i = 0; i < codes.size(); ++i) { + codes[i] = static_cast(43U * i + 5U); + } + + const auto* bits1 = codes.data(); + const auto* bits2 = bits1 + plane_bytes * 4; + const auto* bits3 = bits2 + plane_bytes * 4; + const auto* bits4 = bits3 + plane_bytes * 4; + const uint8_t* all_bits[4] = {bits1, bits2, bits3, bits4}; + + float expected[4] = {0.0F, 0.0F, 0.0F, 0.0F}; + for (uint64_t d = 0; d < dim; ++d) { + const uint64_t byte_idx = d >> 3; + const uint8_t bit_mask = static_cast(1U << (d & 7)); + for (uint32_t i = 0; i < 4; ++i) { + const uint8_t* plane0 = all_bits[i]; + const uint8_t* plane1 = all_bits[i] + plane_bytes; + const uint8_t* plane2 = all_bits[i] + 2 * plane_bytes; + const uint8_t* plane3 = all_bits[i] + 3 * plane_bytes; + float weight = (plane0[byte_idx] & bit_mask) != 0U ? 4.0F : -4.0F; + weight += (plane1[byte_idx] & bit_mask) != 0U ? 2.0F : -2.0F; + weight += (plane2[byte_idx] & bit_mask) != 0U ? 1.0F : -1.0F; + weight += (plane3[byte_idx] & bit_mask) != 0U ? 0.5F : -0.5F; + expected[i] += query[d] * weight; + } + } + + auto check_result = [&expected](const float* result) { + for (uint32_t i = 0; i < 4; ++i) { + REQUIRE(std::abs(expected[i] - result[i]) < 1e-4F); + } + }; + auto check_one = [&](auto func, const uint8_t* bits, float expected_value) { + REQUIRE(std::abs(expected_value - func(query.data(), bits, dim)) < 1e-4F); + }; + + float result[4] = {0.0F, 0.0F, 0.0F, 0.0F}; + generic::RaBitQFloatFourBitCenteredIPBatch4( + query.data(), bits1, bits2, bits3, bits4, dim, result); + check_result(result); + check_one(generic::RaBitQFloatFourBitCenteredIP, bits1, expected[0]); + check_one(generic::RaBitQFloatFourBitCenteredIP, bits2, expected[1]); + check_one(generic::RaBitQFloatFourBitCenteredIP, bits3, expected[2]); + check_one(generic::RaBitQFloatFourBitCenteredIP, bits4, expected[3]); + + RaBitQFloatFourBitCenteredIPBatch4(query.data(), bits1, bits2, bits3, bits4, dim, result); + check_result(result); + check_one(RaBitQFloatFourBitCenteredIP, bits1, expected[0]); + check_one(RaBitQFloatFourBitCenteredIP, bits2, expected[1]); + check_one(RaBitQFloatFourBitCenteredIP, bits3, expected[2]); + check_one(RaBitQFloatFourBitCenteredIP, bits4, expected[3]); + + if (SimdStatus::SupportAVX2()) { + avx2::RaBitQFloatFourBitCenteredIPBatch4( + query.data(), bits1, bits2, bits3, bits4, dim, result); + check_result(result); + check_one(avx2::RaBitQFloatFourBitCenteredIP, bits1, expected[0]); + check_one(avx2::RaBitQFloatFourBitCenteredIP, bits2, expected[1]); + check_one(avx2::RaBitQFloatFourBitCenteredIP, bits3, expected[2]); + check_one(avx2::RaBitQFloatFourBitCenteredIP, bits4, expected[3]); + } + if (SimdStatus::SupportAVX512()) { + avx512::RaBitQFloatFourBitCenteredIPBatch4( + query.data(), bits1, bits2, bits3, bits4, dim, result); + check_result(result); + check_one(avx512::RaBitQFloatFourBitCenteredIP, bits1, expected[0]); + check_one(avx512::RaBitQFloatFourBitCenteredIP, bits2, expected[1]); + check_one(avx512::RaBitQFloatFourBitCenteredIP, bits3, expected[2]); + check_one(avx512::RaBitQFloatFourBitCenteredIP, bits4, expected[3]); + } + } +} + TEST_CASE("RaBitQ FP32 multi-bit lookup SIMD Batch4 Compute Codes", "[ut][simd]") { constexpr float kLookupTolerance = 1e-3F; const std::vector dims = {0, 1, 7, 8, 9, 15, 16, 17, 63, 64, 65, 960}; @@ -833,6 +1039,48 @@ TEST_CASE("RaBitQ FP32 supplement-code SIMD Compute Codes", "[ut][simd]") { } } +TEST_CASE("RaBitQ HNSW 7-bit ExData SIMD", "[ut][simd]") { + constexpr uint64_t dim = 960; + std::vector query(dim); + std::vector codes(dim); + std::vector packed(dim * 7 / 8, 0); + float expected = 0.0F; + for (uint64_t d = 0; d < dim; ++d) { + query[d] = static_cast(static_cast(d % 31) - 15) / 31.0F; + codes[d] = static_cast((d * 73U + 19U) & 127U); + expected += query[d] * static_cast(codes[d]); + } + auto* output = packed.data(); + for (uint64_t block = 0; block < dim; block += 64) { + const auto* input = codes.data() + block; + for (uint64_t lane = 0; lane < 16; ++lane) { + output[lane] = + static_cast((input[lane] & 0x3FU) | ((input[48 + lane] & 0x03U) << 6U)); + output[16 + lane] = static_cast((input[16 + lane] & 0x3FU) | + ((input[48 + lane] & 0x0CU) << 4U)); + output[32 + lane] = static_cast((input[32 + lane] & 0x3FU) | + ((input[48 + lane] & 0x30U) << 2U)); + } + uint64_t top_bits = 0; + constexpr uint64_t k_top_mask = 0x0101010101010101ULL; + for (uint64_t lane = 0; lane < 64; lane += 8) { + uint64_t source = 0; + std::memcpy(&source, input + lane, sizeof(source)); + top_bits |= ((source >> 6U) & k_top_mask) << (lane / 8U); + } + std::memcpy(output + 48, &top_bits, sizeof(top_bits)); + output += 56; + } + + REQUIRE(std::abs(expected - generic::RaBitQFloatExCode7IP(query.data(), packed.data(), dim)) < + 1e-3F); + REQUIRE(std::abs(expected - RaBitQFloatExCode7IP(query.data(), packed.data(), dim)) < 1e-3F); + if (SimdStatus::SupportAVX2()) { + REQUIRE(std::abs(expected - avx2::RaBitQFloatExCode7IP(query.data(), packed.data(), dim)) < + 1e-3F); + } +} + #define BENCHMARK_SIMD_COMPUTE(Simd, Comp) \ BENCHMARK_ADVANCED(#Simd #Comp) { \ for (int i = 0; i < count; ++i) { \ diff --git a/src/utils/util_functions.cpp b/src/utils/util_functions.cpp index 3efcb76ad9..0c5113af26 100644 --- a/src/utils/util_functions.cpp +++ b/src/utils/util_functions.cpp @@ -275,7 +275,8 @@ sample_train_data(const vsag::DatasetPtr& data, int64_t total_elements, int64_t dim, int64_t train_sample_count, - Allocator* allocator) { + Allocator* allocator, + std::optional random_seed) { const int64_t min_train_size = 512; const int64_t max_train_size = 65536; @@ -306,7 +307,7 @@ sample_train_data(const vsag::DatasetPtr& data, sampled_indices.resize(actual_size); std::iota(sampled_indices.begin(), sampled_indices.end(), 0); std::random_device rd; - std::mt19937_64 gen(rd()); + std::mt19937_64 gen(random_seed.has_value() ? *random_seed : rd()); for (int64_t i = sample_count; i < total_elements; ++i) { std::uniform_int_distribution dist(0, i); int64_t j = dist(gen); diff --git a/src/utils/util_functions.h b/src/utils/util_functions.h index 7939933cbb..bd01a66333 100644 --- a/src/utils/util_functions.h +++ b/src/utils/util_functions.h @@ -17,6 +17,7 @@ #pragma once #include +#include #include #include "index_common_param.h" @@ -136,6 +137,7 @@ sample_train_data(const DatasetPtr& data, int64_t total_elements, int64_t dim, int64_t train_sample_count, - Allocator* allocator = nullptr); + Allocator* allocator = nullptr, + std::optional random_seed = std::nullopt); } // namespace vsag diff --git a/src/utils/util_functions_test.cpp b/src/utils/util_functions_test.cpp index 0371d0b9b1..813f7217f5 100644 --- a/src/utils/util_functions_test.cpp +++ b/src/utils/util_functions_test.cpp @@ -17,9 +17,11 @@ #include #include +#include #include #include #include +#include #include "impl/allocator/default_allocator.h" #include "unittest.h" @@ -107,3 +109,30 @@ TEST_CASE("UtilFunctions Basic", "[ut][UtilFunctions]") { REQUIRE_FALSE(t.empty()); } } + +TEST_CASE("sample_train_data honors a fixed random seed", "[ut][UtilFunctions]") { + constexpr int64_t dim = 3; + constexpr int64_t count = 2048; + constexpr int64_t sample_count = 512; + std::vector vectors(count * dim); + std::iota(vectors.begin(), vectors.end(), 0.0F); + std::vector ids(count); + std::iota(ids.begin(), ids.end(), 0); + auto data = Dataset::Make(); + data->NumElements(count) + ->Dim(dim) + ->Ids(ids.data()) + ->Float32Vectors(vectors.data()) + ->Owner(false); + auto allocator = std::make_shared(); + + auto first = sample_train_data(data, count, dim, sample_count, allocator.get(), 0x52425131U); + auto second = sample_train_data(data, count, dim, sample_count, allocator.get(), 0x52425131U); + + REQUIRE(first->GetNumElements() == sample_count); + REQUIRE(second->GetNumElements() == sample_count); + REQUIRE(std::equal(first->GetIds(), first->GetIds() + sample_count, second->GetIds())); + REQUIRE(std::equal(first->GetFloat32Vectors(), + first->GetFloat32Vectors() + sample_count * dim, + second->GetFloat32Vectors())); +} diff --git a/src/utils/visited_list.h b/src/utils/visited_list.h index ef739c4755..647961611b 100644 --- a/src/utils/visited_list.h +++ b/src/utils/visited_list.h @@ -54,6 +54,20 @@ class VisitedList : public ResourceObject { return this->tags_[word_id] == this->tag_ and (this->words_[word_id] & mask) != 0; } + [[nodiscard]] bool + TestAndSet(const InnerIdType& id) { + const auto word_id = static_cast(id) / kBitsPerWord; + const auto mask = WordType{1} << (static_cast(id) % kBitsPerWord); + if (this->tags_[word_id] != this->tag_) { + this->tags_[word_id] = this->tag_; + this->words_[word_id] = mask; + return false; + } + const bool was_set = (this->words_[word_id] & mask) != 0; + this->words_[word_id] |= mask; + return was_set; + } + void Prefetch(const InnerIdType& id) { const auto word_id = static_cast(id) / kBitsPerWord; diff --git a/src/utils/visited_list_test.cpp b/src/utils/visited_list_test.cpp index 2f362d719b..1c907d8d11 100644 --- a/src/utils/visited_list_test.cpp +++ b/src/utils/visited_list_test.cpp @@ -101,6 +101,16 @@ TEST_CASE("VisitedList Basic Test", "[ut][VisitedList]") { REQUIRE_FALSE(vl_ptr->Get(64)); } + SECTION("test and set") { + REQUIRE_FALSE(vl_ptr->TestAndSet(63)); + REQUIRE(vl_ptr->TestAndSet(63)); + REQUIRE_FALSE(vl_ptr->TestAndSet(64)); + REQUIRE(vl_ptr->Get(63)); + REQUIRE(vl_ptr->Get(64)); + vl_ptr->Reset(); + REQUIRE_FALSE(vl_ptr->TestAndSet(63)); + } + SECTION("test memory usage") { const auto word_count = (static_cast(size) + VisitedList::kBitsPerWord - 1) / VisitedList::kBitsPerWord; diff --git a/tests/test_hgraph_rabitq_split.cpp b/tests/test_hgraph_rabitq_split.cpp index f35e14cf58..b13b8608fa 100644 --- a/tests/test_hgraph_rabitq_split.cpp +++ b/tests/test_hgraph_rabitq_split.cpp @@ -14,15 +14,22 @@ #include +#include #include #include #include #include +#include #include +#include +#include #include #include +#include #include "functest.h" +#include "storage/serialization_tags.h" +#include "storage/streaming_serialization_test_utils.h" #include "test_index.h" #include "vsag/engine.h" #include "vsag/resource.h" @@ -185,6 +192,25 @@ class RejectSecondBuildThreadPool final : public vsag::ThreadPool { std::thread worker_{}; }; +class OnlyLabelFilter final : public vsag::Filter { +public: + explicit OnlyLabelFilter(int64_t label) : label_(label) { + } + + bool + CheckValid(int64_t label) const override { + return label == label_; + } + + float + ValidRatio() const override { + return 1.0F; + } + +private: + int64_t label_; +}; + constexpr const char* kSplitSearchParam = R"( { "hgraph": { @@ -264,6 +290,77 @@ TEST_CASE("HGraph RaBitQ Split validates dimension before optimized training", REQUIRE(index->GetNumElements() == base_count); } +TEST_CASE("HGraph fused RaBitQ exported model reuses its codec during build", + "[ft][rabitq_split][hgraph][fused][export_model]") { + using namespace fixtures; + constexpr int64_t dim = 64; + constexpr uint64_t base_count = 64; + constexpr int64_t topk = 10; + const std::string graph_type = GENERATE("odescent", "nsw"); + CAPTURE(graph_type); + + auto param = + HGraphRaBitQSplitTestIndex::GenerateBuildParam("l2", dim, "memory_io", "", 1, 7, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["graph_type"].SetString(graph_type); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["build_thread_count"].SetInt(1); + param = param_json.Dump(); + + auto dataset = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + std::vector target_vectors( + dataset->base_->GetFloat32Vectors(), + dataset->base_->GetFloat32Vectors() + base_count * static_cast(dim)); + for (uint64_t i = 0; i < target_vectors.size(); ++i) { + target_vectors[i] = target_vectors[i] * 1.75F + 3.0F + + static_cast(i % static_cast(dim)) * 0.01F; + } + auto target_base = vsag::Dataset::Make(); + target_base->NumElements(base_count) + ->Dim(dim) + ->Ids(dataset->base_->GetIds()) + ->Float32Vectors(target_vectors.data()) + ->Owner(false); + auto target_query = vsag::Dataset::Make(); + target_query->NumElements(1)->Dim(dim)->Float32Vectors(target_vectors.data())->Owner(false); + + const auto serialized_model = [](const TestIndex::IndexPtr& index) { + auto serialized = index->Serialize(); + REQUIRE(serialized.has_value()); + const auto binary = serialized.value().Get(HGraphRaBitQSplitTestIndex::name); + REQUIRE(binary.data != nullptr); + REQUIRE(binary.size > 0); + return std::string(reinterpret_cast(binary.data.get()), binary.size); + }; + + auto source = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto source_build = source->Build(dataset->base_); + REQUIRE(source_build.has_value()); + + auto model_result = source->ExportModel(); + REQUIRE(model_result.has_value()); + auto model = model_result.value(); + REQUIRE(model->GetNumElements() == 0); + const auto expected_model = serialized_model(model); + auto model_build = model->Build(target_base); + REQUIRE(model_build.has_value()); + REQUIRE(model->GetNumElements() == base_count); + + auto populated_model_result = model->ExportModel(); + REQUIRE(populated_model_result.has_value()); + REQUIRE(serialized_model(populated_model_result.value()) == expected_model); + + auto search_result = model->KnnSearch(target_query, topk, kSplitSearchParam); + REQUIRE(search_result.has_value()); + REQUIRE(search_result.value()->GetDim() == topk); +} + TEST_CASE("HGraph RaBitQ Split drains accepted build tasks after enqueue failure", "[ft][rabitq_split][hgraph]") { using namespace fixtures; @@ -401,3 +498,846 @@ TEST_CASE("HGraph RaBitQ Split Reject Unsupported Hybrid", "[ft][rabitq_split][h auto result = vsag::Factory::CreateIndex(HGraphRaBitQSplitTestIndex::name, bad_param); REQUIRE_FALSE(result.has_value()); } + +TEST_CASE("HGraph RaBitQ split rejects non-finite one-bit queries", + "[ft][rabitq_split][hgraph][fused][validation]") { + using namespace fixtures; + constexpr int64_t dim = 64; + constexpr uint64_t base_count = 64; + constexpr int64_t topk = 10; + const bool fused = GENERATE(false, true); + CAPTURE(fused); + + auto param = + HGraphRaBitQSplitTestIndex::GenerateBuildParam("l2", dim, "memory_io", "", 1, 7, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(fused); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["build_thread_count"].SetInt(1); + + auto source = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param_json.Dump(), true); + auto build_result = index->Build(source->base_); + REQUIRE(build_result.has_value()); + + const float invalid_values[] = {std::numeric_limits::quiet_NaN(), + std::numeric_limits::infinity(), + -std::numeric_limits::infinity()}; + for (const float invalid_value : invalid_values) { + CAPTURE(invalid_value); + std::vector query_vector(source->query_->GetFloat32Vectors(), + source->query_->GetFloat32Vectors() + dim); + query_vector[0] = invalid_value; + auto query = vsag::Dataset::Make(); + query->NumElements(1)->Dim(dim)->Float32Vectors(query_vector.data())->Owner(false); + auto result = index->KnnSearch(query, topk, kSplitSearchParam); + REQUIRE_FALSE(result.has_value()); + REQUIRE(result.error().type == vsag::ErrorType::INVALID_ARGUMENT); + } + + if (fused) { + uint64_t scored_count = 0; + uint64_t largest_batch = 0; + const auto score = [](int64_t label) { return static_cast(label % 1000000); }; + vsag::SearchRequest request; + request.topk_ = topk; + request.params_str_ = fmt::format( + R"({{"hgraph":{{"ef_search":{},"rabitq_one_bit_search":true}}}})", base_count); + request.distance_batch_size_ = 3; + request.distance_batch_func_ = + [&](const int64_t* labels, uint64_t count, float* distances) { + scored_count += count; + largest_batch = std::max(largest_batch, count); + for (uint64_t i = 0; i < count; ++i) { + distances[i] = score(labels[i]); + } + }; + auto result = index->SearchWithRequest(request); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == topk); + REQUIRE(scored_count > 0); + REQUIRE(largest_batch <= request.distance_batch_size_); + for (int64_t i = 0; i < result.value()->GetDim(); ++i) { + REQUIRE(result.value()->GetDistances()[i] == score(result.value()->GetIds()[i])); + } + } +} + +TEST_CASE("HGraph fused RaBitQ split rejects non-finite base vectors", + "[ft][rabitq_split][hgraph][fused][validation]") { + using namespace fixtures; + constexpr int64_t dim = 64; + constexpr uint64_t base_count = 64; + + auto param = + HGraphRaBitQSplitTestIndex::GenerateBuildParam("l2", dim, "memory_io", "", 1, 7, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["graph_type"].SetString("odescent"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(true); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["build_thread_count"].SetInt(1); + param = param_json.Dump(); + + auto source = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + const float invalid_values[] = {std::numeric_limits::quiet_NaN(), + std::numeric_limits::infinity(), + -std::numeric_limits::infinity()}; + for (const float invalid_value : invalid_values) { + CAPTURE(invalid_value); + std::vector invalid_vectors( + source->base_->GetFloat32Vectors(), + source->base_->GetFloat32Vectors() + base_count * static_cast(dim)); + invalid_vectors[0] = invalid_value; + auto invalid_base = vsag::Dataset::Make(); + invalid_base->NumElements(base_count) + ->Dim(dim) + ->Ids(source->base_->GetIds()) + ->Float32Vectors(invalid_vectors.data()) + ->Owner(false); + auto invalid_index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto build_result = invalid_index->Build(invalid_base); + REQUIRE_FALSE(build_result.has_value()); + REQUIRE(build_result.error().type == vsag::ErrorType::INVALID_ARGUMENT); + REQUIRE(invalid_index->GetNumElements() == 0); + } + + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto build_result = index->Build(source->base_); + REQUIRE(build_result.has_value()); + REQUIRE(index->GetNumElements() == base_count); + + const auto* base_ids = source->base_->GetIds(); + const int64_t existing_label = base_ids[0]; + const auto* existing_vector = source->base_->GetFloat32Vectors(); + auto original_distance = index->CalcDistanceById(existing_vector, existing_label); + REQUIRE(original_distance.has_value()); + const int64_t max_label = *std::max_element(base_ids, base_ids + base_count); + uint64_t invalid_value_index = 0; + for (const float invalid_value : invalid_values) { + CAPTURE(invalid_value_index); + std::vector invalid_vector(existing_vector, existing_vector + dim); + invalid_vector[0] = invalid_value; + + const int64_t added_label = max_label + static_cast(invalid_value_index) + 1; + auto invalid_add = vsag::Dataset::Make(); + invalid_add->NumElements(1) + ->Dim(dim) + ->Ids(&added_label) + ->Float32Vectors(invalid_vector.data()) + ->Owner(false); + const uint64_t count_before_add = index->GetNumElements(); + auto add_result = index->Add(invalid_add); + REQUIRE_FALSE(add_result.has_value()); + REQUIRE(add_result.error().type == vsag::ErrorType::INVALID_ARGUMENT); + REQUIRE(index->GetNumElements() == count_before_add); + REQUIRE_FALSE(index->CheckIdExist(added_label)); + + auto invalid_update = vsag::Dataset::Make(); + invalid_update->NumElements(1) + ->Dim(dim) + ->Ids(&existing_label) + ->Float32Vectors(invalid_vector.data()) + ->Owner(false); + auto update_result = index->UpdateVector(existing_label, invalid_update, true); + REQUIRE_FALSE(update_result.has_value()); + REQUIRE(update_result.error().type == vsag::ErrorType::INVALID_ARGUMENT); + REQUIRE(index->GetNumElements() == base_count); + auto unchanged_distance = index->CalcDistanceById(existing_vector, existing_label); + REQUIRE(unchanged_distance.has_value()); + REQUIRE(unchanged_distance.value() == original_distance.value()); + ++invalid_value_index; + } + + std::vector short_vector(existing_vector, existing_vector + dim - 1); + auto wrong_dim_update = vsag::Dataset::Make(); + wrong_dim_update->NumElements(1) + ->Dim(dim - 1) + ->Ids(&existing_label) + ->Float32Vectors(short_vector.data()) + ->Owner(false); + auto wrong_dim_result = index->UpdateVector(existing_label, wrong_dim_update, true); + REQUIRE_FALSE(wrong_dim_result.has_value()); + REQUIRE(wrong_dim_result.error().type == vsag::ErrorType::INVALID_ARGUMENT); + auto unchanged_distance = index->CalcDistanceById(existing_vector, existing_label); + REQUIRE(unchanged_distance.has_value()); + REQUIRE(unchanged_distance.value() == original_distance.value()); +} + +TEST_CASE("HGraph fused RaBitQ split compact round trip", + "[ft][rabitq_split][hgraph][fused][serialize]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 1024; + constexpr int64_t topk = 10; + constexpr uint64_t max_degree = 32; + const uint32_t filter_bits = GENERATE(1U, 2U, 3U, 4U); + const uint32_t supplement_bits = 8U - filter_bits; + INFO(fmt::format("fused split {}+{}", filter_bits, supplement_bits)); + + auto param = HGraphRaBitQSplitTestIndex::GenerateBuildParam( + "l2", dim, "memory_io", "", filter_bits, supplement_bits, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["graph_type"].SetString("odescent"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["build_thread_count"].SetInt(4); + param = param_json.Dump(); + + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto dataset = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + TestIndex::TestBuildIndex(index, dataset, true); + REQUIRE(index->GetNumElements() == base_count); + + const auto* base_ids = dataset->base_->GetIds(); + const int64_t old_label = base_ids[base_count - 1]; + const int64_t new_label = + *std::max_element(base_ids, base_ids + base_count) + static_cast(base_count) + 1; + auto equal_missing = index->UpdateId(new_label, new_label); + REQUIRE(equal_missing.has_value()); + REQUIRE(equal_missing.value()); + auto update = index->UpdateId(old_label, new_label); + REQUIRE(update.has_value()); + REQUIRE(update.value()); + REQUIRE_FALSE(index->CheckIdExist(old_label)); + REQUIRE(index->CheckIdExist(new_label)); + REQUIRE_FALSE(index->UpdateId(new_label, base_ids[0]).has_value()); + REQUIRE_FALSE(index->UpdateId(old_label, new_label + 1).has_value()); + + const float recall = TestIndex::TestKnnSearch(index, dataset, kSplitSearchParam, 0.5F, true); + REQUIRE(recall > 0.5F); + + auto query = get_one_query(dataset->query_, 0); + auto expected = index->KnnSearch(query, topk, kSplitSearchParam); + REQUIRE(expected.has_value()); + REQUIRE(expected.value()->GetDim() == topk); + if (filter_bits == 1) { + const auto stats = expected.value()->GetStatistics({"rabitq_full_count", + "rabitq_reorder_hint_full_count", + "rabitq_reorder_fallback_full_count"}); + REQUIRE(stats.size() == 3); + REQUIRE(std::stoull(stats[0]) > 0); + REQUIRE(std::stoull(stats[1]) == 0); + REQUIRE(std::stoull(stats[2]) == 0); + } + const std::vector expected_ids(expected.value()->GetIds(), + expected.value()->GetIds() + topk); + const std::vector expected_distances(expected.value()->GetDistances(), + expected.value()->GetDistances() + topk); + + auto require_same_result = [&](const TestIndex::IndexPtr& restored) { + REQUIRE(restored->GetNumElements() == base_count); + REQUIRE_FALSE(restored->CheckIdExist(old_label)); + REQUIRE(restored->CheckIdExist(new_label)); + auto actual = restored->KnnSearch(query, topk, kSplitSearchParam); + REQUIRE(actual.has_value()); + REQUIRE(actual.value()->GetDim() == topk); + for (int64_t i = 0; i < topk; ++i) { + REQUIRE(actual.value()->GetIds()[i] == expected_ids[static_cast(i)]); + REQUIRE(std::abs(actual.value()->GetDistances()[i] - + expected_distances[static_cast(i)]) <= 2e-6F); + } + }; + + auto ordinary_result = index->Serialize(); + REQUIRE(ordinary_result.has_value()); + uint64_t ordinary_size = 0; + for (const auto& key : ordinary_result.value().GetKeys()) { + ordinary_size += ordinary_result.value().Get(key).size; + } + auto ordinary_restored = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + REQUIRE(ordinary_restored->Deserialize(ordinary_result.value()).has_value()); + require_same_result(ordinary_restored); + + std::stringstream stream; + REQUIRE(index->SerializeStreaming(stream).has_value()); + const auto bytes = stream.str(); + const auto base_codes = + vsag::test::FindStreamingBlock(bytes, vsag::StreamSerializationTag::BASE_CODES); + const auto bottom_graph = + vsag::test::FindStreamingBlock(bytes, vsag::StreamSerializationTag::BOTTOM_GRAPH); + + // Fused BASE_CODES contains only the codec model. A second per-node split-code copy would + // require at least dim * (x + y) / 8 bytes per vector and violate this bound. + const uint64_t duplicated_split_bytes = base_count * static_cast(dim); + REQUIRE(base_codes.payload_size < duplicated_split_bytes); + + // For dim=128 and M=32, a fused node record is at most 384 bytes: links, metadata, one + // byte/dimension of split codes, and cache-line padding. Allow 64 KiB for graph headers and + // the shared codec model, but no second count-scaled code array. + constexpr uint64_t fused_record_upper_bound = + max_degree * sizeof(vsag::InnerIdType) + static_cast(dim) + 2 * 64; + constexpr uint64_t model_and_header_allowance = 64 * 1024; + constexpr uint64_t container_allowance = 256 * 1024; + REQUIRE(bottom_graph.payload_size <= + base_count * fused_record_upper_bound + model_and_header_allowance); + REQUIRE(ordinary_size <= base_count * fused_record_upper_bound + container_allowance); + REQUIRE(bytes.size() <= base_count * fused_record_upper_bound + container_allowance); + + auto streaming_restored = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + std::stringstream deserialize_stream(bytes); + REQUIRE(streaming_restored->DeserializeStreaming(deserialize_stream).has_value()); + require_same_result(streaming_restored); +} + +TEST_CASE("HGraph fused RaBitQ split trailing duplicate label round trip", + "[ft][rabitq_split][hgraph][fused][duplicate][serialize]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 64; + constexpr int64_t topk = 10; + + auto param = + HGraphRaBitQSplitTestIndex::GenerateBuildParam("l2", dim, "memory_io", "", 1, 7, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["support_duplicate"].SetBool(true); + param_json["index_param"]["build_thread_count"].SetInt(4); + param = param_json.Dump(); + + auto source = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + std::vector labels(source->base_->GetIds(), source->base_->GetIds() + base_count); + const int64_t skipped_label = labels.back(); + const int64_t duplicate_label = labels[base_count - 2]; + labels.back() = duplicate_label; + auto base = vsag::Dataset::Make(); + base->NumElements(base_count) + ->Dim(dim) + ->Ids(labels.data()) + ->Float32Vectors(source->base_->GetFloat32Vectors()) + ->Owner(false); + + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto build_result = index->Build(base); + REQUIRE(build_result.has_value()); + REQUIRE(build_result.value() == std::vector{duplicate_label}); + const uint64_t graph_count = index->GetNumElements(); + REQUIRE(graph_count == base_count - 1); + REQUIRE(graph_count <= base_count); + REQUIRE(index->CheckIdExist(duplicate_label)); + REQUIRE_FALSE(index->CheckIdExist(skipped_label)); + + auto query = vsag::Dataset::Make(); + query->NumElements(1) + ->Dim(dim) + ->Float32Vectors(source->base_->GetFloat32Vectors() + (base_count - 2) * dim) + ->Owner(false); + auto expected = index->KnnSearch(query, topk, kSplitSearchParam); + REQUIRE(expected.has_value()); + REQUIRE(expected.value()->GetDim() == topk); + REQUIRE(expected.value()->GetIds()[0] == duplicate_label); + + auto require_same_semantics = [&](const TestIndex::IndexPtr& restored) { + REQUIRE(restored->GetNumElements() == graph_count); + REQUIRE(restored->CheckIdExist(duplicate_label)); + REQUIRE_FALSE(restored->CheckIdExist(skipped_label)); + auto actual = restored->KnnSearch(query, topk, kSplitSearchParam); + REQUIRE(actual.has_value()); + REQUIRE(actual.value()->GetDim() == topk); + REQUIRE(actual.value()->GetIds()[0] == duplicate_label); + for (int64_t i = 0; i < topk; ++i) { + REQUIRE(actual.value()->GetIds()[i] == expected.value()->GetIds()[i]); + REQUIRE(std::abs(actual.value()->GetDistances()[i] - + expected.value()->GetDistances()[i]) <= 2e-6F); + } + }; + + auto binary = index->Serialize(); + REQUIRE(binary.has_value()); + auto ordinary_restored = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + REQUIRE(ordinary_restored->Deserialize(binary.value()).has_value()); + require_same_semantics(ordinary_restored); + + std::stringstream stream; + REQUIRE(index->SerializeStreaming(stream).has_value()); + auto streaming_restored = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + std::stringstream deserialize_stream(stream.str()); + REQUIRE(streaming_restored->DeserializeStreaming(deserialize_stream).has_value()); + require_same_semantics(streaming_restored); +} + +TEST_CASE("HGraph fused RaBitQ split preserves trailing vector aliases", + "[ft][rabitq_split][hgraph][fused][duplicate][serialize]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 64; + constexpr int64_t duplicate_count = 3; + + auto param = + HGraphRaBitQSplitTestIndex::GenerateBuildParam("l2", dim, "memory_io", "", 1, 7, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["support_duplicate"].SetBool(true); + param_json["index_param"]["build_thread_count"].SetInt(1); + param = param_json.Dump(); + + auto source = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + std::vector vectors(source->base_->GetFloat32Vectors(), + source->base_->GetFloat32Vectors() + base_count * dim); + std::copy_n(vectors.data(), dim, vectors.data() + (base_count - 2) * dim); + std::copy_n(vectors.data(), dim, vectors.data() + (base_count - 1) * dim); + std::vector labels(source->base_->GetIds(), source->base_->GetIds() + base_count); + const std::vector duplicate_labels = { + labels.front(), labels[base_count - 2], labels.back()}; + auto base = vsag::Dataset::Make(); + base->NumElements(base_count) + ->Dim(dim) + ->Ids(labels.data()) + ->Float32Vectors(vectors.data()) + ->Owner(false); + auto query = vsag::Dataset::Make(); + query->NumElements(1)->Dim(dim)->Float32Vectors(vectors.data())->Owner(false); + + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto build_result = index->Build(base); + REQUIRE(build_result.has_value()); + REQUIRE(build_result.value().empty()); + REQUIRE(index->GetNumElements() == base_count); + + auto require_duplicate_semantics = [&](const TestIndex::IndexPtr& restored, + uint32_t ef_search) { + const auto search_param = fmt::format( + R"({{"hgraph":{{"ef_search":{},"rabitq_one_bit_search":true}}}})", ef_search); + auto result = restored->KnnSearch(query, duplicate_count, search_param); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == duplicate_count); + std::vector result_labels(result.value()->GetIds(), + result.value()->GetIds() + duplicate_count); + for (const auto label : duplicate_labels) { + REQUIRE(std::find(result_labels.begin(), result_labels.end(), label) != + result_labels.end()); + } + + auto alias_only = std::make_shared(duplicate_labels.back()); + auto filtered = restored->KnnSearch(query, 1, search_param, alias_only); + REQUIRE(filtered.has_value()); + REQUIRE(filtered.value()->GetDim() == 1); + REQUIRE(filtered.value()->GetIds()[0] == duplicate_labels.back()); + + auto alias_distance = restored->CalcDistanceById(vectors.data(), duplicate_labels.back()); + REQUIRE(alias_distance.has_value()); + REQUIRE(std::isfinite(alias_distance.value())); + const float distance_tolerance = 2e-5F * std::max(1.0F, std::fabs(alias_distance.value())); + for (int64_t i = 0; i < duplicate_count; ++i) { + REQUIRE(std::fabs(result.value()->GetDistances()[i] - alias_distance.value()) <= + distance_tolerance); + } + REQUIRE(std::fabs(filtered.value()->GetDistances()[0] - alias_distance.value()) <= + distance_tolerance); + + vsag::IteratorContext* iterator_context = nullptr; + vsag::FilterPtr iterator_filter = nullptr; + auto iterator_result = restored->KnnSearch( + query, duplicate_count, search_param, iterator_filter, iterator_context, false); + REQUIRE(iterator_result.has_value()); + REQUIRE(iterator_result.value()->GetDim() == duplicate_count); + std::vector iterator_labels(iterator_result.value()->GetIds(), + iterator_result.value()->GetIds() + duplicate_count); + for (const auto label : duplicate_labels) { + REQUIRE(std::find(iterator_labels.begin(), iterator_labels.end(), label) != + iterator_labels.end()); + } + delete iterator_context; + }; + + require_duplicate_semantics(index, 20); + require_duplicate_semantics(index, 80); + + auto binary = index->Serialize(); + REQUIRE(binary.has_value()); + auto ordinary_restored = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + REQUIRE(ordinary_restored->Deserialize(binary.value()).has_value()); + require_duplicate_semantics(ordinary_restored, 20); + + std::stringstream stream; + REQUIRE(index->SerializeStreaming(stream).has_value()); + auto streaming_restored = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + std::stringstream deserialize_stream(stream.str()); + REQUIRE(streaming_restored->DeserializeStreaming(deserialize_stream).has_value()); + require_duplicate_semantics(streaming_restored, 80); +} + +TEST_CASE("HGraph fused RaBitQ split honors disabled reorder", + "[ft][rabitq_split][hgraph][fused][search]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 128; + constexpr int64_t topk = 10; + const uint32_t filter_bits = GENERATE(1U, 2U, 3U, 4U); + const bool support_duplicate = GENERATE(false, true); + + auto param = HGraphRaBitQSplitTestIndex::GenerateBuildParam( + "l2", dim, "memory_io", "", filter_bits, 8U - filter_bits, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["support_duplicate"].SetBool(support_duplicate); + param_json["index_param"]["build_thread_count"].SetInt(1); + + auto source = HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, "l2"); + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param_json.Dump(), true); + auto build_result = index->Build(source->base_); + REQUIRE(build_result.has_value()); + REQUIRE(build_result.value().empty()); + auto query = get_one_query(source->query_, 0); + + for (const uint32_t ef_search : {20U, 80U}) { + for (const uint32_t parallelism : {1U, 2U}) { + CAPTURE(filter_bits, support_duplicate, ef_search, parallelism); + const auto search_param = fmt::format( + R"({{ + "hgraph": {{ + "ef_search": {}, + "rabitq_one_bit_search": true, + "enable_reorder": false, + "parallelism": {} + }} + }})", + ef_search, + parallelism); + auto result = index->KnnSearch(query, topk, search_param); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == topk); + for (int64_t i = 0; i < result.value()->GetDim(); ++i) { + REQUIRE(std::isfinite(result.value()->GetDistances()[i])); + REQUIRE(result.value()->GetDistances()[i] < std::numeric_limits::max()); + } + const auto stats = result.value()->GetStatistics({"rabitq_filter_count", + "rabitq_full_count", + "rabitq_filter_fallback_full_count", + "rabitq_reorder_hint_full_count", + "rabitq_reorder_fallback_full_count", + "reorder_distance_count"}); + REQUIRE(stats.size() == 6); + REQUIRE(std::stoull(stats[0]) > 0); + REQUIRE(std::stoull(stats[1]) == 0); + REQUIRE(std::stoull(stats[2]) == 0); + REQUIRE(std::stoull(stats[3]) == 0); + REQUIRE(std::stoull(stats[4]) == 0); + REQUIRE(std::stoull(stats[5]) == 0); + } + } +} + +TEST_CASE("HGraph fused RaBitQ split reuses full distances across deferred finalize", + "[ft][rabitq_split][hgraph][fused][search][full_hint]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 512; + constexpr int64_t topk = 10; + constexpr uint64_t query_count = 8; + const uint32_t filter_bits = GENERATE(1U, 3U); + const std::string metric_type = GENERATE("l2", "ip"); + const uint32_t supplement_bits = filter_bits == 1U ? 3U : 5U; + + auto param = HGraphRaBitQSplitTestIndex::GenerateBuildParam( + metric_type, dim, "memory_io", "", filter_bits, supplement_bits, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["support_duplicate"].SetBool(false); + param_json["index_param"]["build_thread_count"].SetInt(1); + + auto source = + HGraphRaBitQSplitTestIndex::pool.GetDatasetAndCreate(dim, base_count, metric_type); + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param_json.Dump(), true); + auto build_result = index->Build(source->base_); + REQUIRE(build_result.has_value()); + REQUIRE(build_result.value().empty()); + + for (const uint32_t ef_search : {20U, 40U, 80U}) { + uint64_t full_count = 0; + uint64_t hint_full_count = 0; + uint64_t fallback_full_count = 0; + uint64_t reorder_distance_count = 0; + for (uint64_t query_index = 0; query_index < query_count; ++query_index) { + CAPTURE(filter_bits, supplement_bits, metric_type, ef_search, query_index); + const auto search_param = fmt::format( + R"({{ + "hgraph": {{ + "ef_search": {}, + "rabitq_one_bit_search": true, + "enable_reorder": true, + "parallelism": 1 + }} + }})", + ef_search); + auto query = get_one_query(source->query_, query_index); + auto result = index->KnnSearch(query, topk, search_param); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == topk); + const auto stats = result.value()->GetStatistics({"rabitq_full_count", + "rabitq_reorder_hint_full_count", + "rabitq_reorder_fallback_full_count", + "reorder_distance_count"}); + REQUIRE(stats.size() == 4); + full_count += std::stoull(stats[0]); + hint_full_count += std::stoull(stats[1]); + fallback_full_count += std::stoull(stats[2]); + reorder_distance_count += std::stoull(stats[3]); + } + + REQUIRE(full_count > 0); + REQUIRE(hint_full_count + fallback_full_count == full_count); + if (filter_bits == 1U) { + REQUIRE(hint_full_count == 0); + REQUIRE(fallback_full_count == full_count); + } else { + REQUIRE(hint_full_count == full_count); + REQUIRE(fallback_full_count == 0); + } + if (ef_search <= 40U) { + REQUIRE(reorder_distance_count < full_count); + } else { + REQUIRE(reorder_distance_count == 0); + } + } +} + +TEST_CASE("HGraph fused RaBitQ split keeps zero-residual nodes without reorder", + "[ft][rabitq_split][hgraph][fused][fused_zero_residual]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 64; + constexpr int64_t topk = 10; + const uint32_t filter_bits = GENERATE(1U, 2U, 3U, 4U); + + auto param = HGraphRaBitQSplitTestIndex::GenerateBuildParam( + "l2", dim, "memory_io", "", filter_bits, 3U, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["support_duplicate"].SetBool(false); + param_json["index_param"]["build_thread_count"].SetInt(1); + + std::vector vectors(base_count * static_cast(dim)); + for (int64_t d = 0; d < dim; ++d) { + const float value = static_cast((d % 17) - 8) * 0.125F; + for (uint64_t row = 0; row < base_count; ++row) { + vectors[row * static_cast(dim) + static_cast(d)] = value; + } + } + std::vector labels(base_count); + std::iota(labels.begin(), labels.end(), 0); + auto base = vsag::Dataset::Make(); + base->NumElements(base_count) + ->Dim(dim) + ->Ids(labels.data()) + ->Float32Vectors(vectors.data()) + ->Owner(false); + auto query = vsag::Dataset::Make(); + query->NumElements(1)->Dim(dim)->Float32Vectors(vectors.data())->Owner(false); + + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param_json.Dump(), true); + auto build_result = index->Build(base); + REQUIRE(build_result.has_value()); + REQUIRE(build_result.value().empty()); + + const auto search_param = R"({ + "hgraph": { + "ef_search": 20, + "rabitq_one_bit_search": true, + "enable_reorder": false, + "parallelism": 1 + } + })"; + auto result = index->KnnSearch(query, topk, search_param); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetDim() == topk); + for (int64_t i = 0; i < result.value()->GetDim(); ++i) { + REQUIRE(std::isfinite(result.value()->GetDistances()[i])); + REQUIRE(result.value()->GetDistances()[i] < std::numeric_limits::max()); + } + const auto stats = result.value()->GetStatistics({"rabitq_filter_count", + "rabitq_full_count", + "rabitq_filter_fallback_full_count", + "reorder_distance_count"}); + REQUIRE(stats.size() == 4); + REQUIRE(std::stoull(stats[0]) > 0); + REQUIRE(std::stoull(stats[1]) == 0); + REQUIRE(std::stoull(stats[2]) == 0); + REQUIRE(std::stoull(stats[3]) == 0); +} + +TEST_CASE("HGraph fused RaBitQ split expands a sole representative and its aliases", + "[ft][rabitq_split][hgraph][fused][duplicate][search]") { + using namespace fixtures; + constexpr int64_t dim = 128; + constexpr uint64_t base_count = 3; + + auto param = + HGraphRaBitQSplitTestIndex::GenerateBuildParam("l2", dim, "memory_io", "", 2, 6, true); + auto param_json = vsag::JsonType::Parse(param); + param_json["index_param"]["graph_io_type"].SetString("memory_io"); + param_json["index_param"]["graph_storage_type"].SetString("flat"); + param_json["index_param"]["reorder_source"].SetString("base"); + param_json["index_param"]["rabitq_fused_datacell"].SetBool(true); + param_json["index_param"]["rabitq_use_fht"].SetBool(true); + param_json["index_param"]["store_raw_vector"].SetBool(false); + param_json["index_param"]["use_mci"].SetBool(false); + param_json["index_param"]["support_duplicate"].SetBool(true); + param_json["index_param"]["duplicate_distance_threshold"].SetFloat(1.0F); + param_json["index_param"]["build_thread_count"].SetInt(1); + param = param_json.Dump(); + + std::vector vectors(base_count * dim); + for (int64_t d = 0; d < dim; ++d) { + vectors[d] = static_cast((d * 17) % 29) * 0.125F; + } + for (int64_t d = 0; d < dim; ++d) { + vectors[dim + d] = vectors[d] + static_cast((d % 3) + 1) * 0.005F; + vectors[2 * dim + d] = vectors[d] - static_cast((d % 5) + 1) * 0.008F; + } + std::vector labels = {100, 200, 300}; + auto base = vsag::Dataset::Make(); + base->NumElements(base_count) + ->Dim(dim) + ->Ids(labels.data()) + ->Float32Vectors(vectors.data()) + ->Owner(false); + auto query = vsag::Dataset::Make(); + query->NumElements(1)->Dim(dim)->Float32Vectors(vectors.data())->Owner(false); + + auto index = TestIndex::TestFactory(HGraphRaBitQSplitTestIndex::name, param, true); + auto build_result = index->Build(base); + REQUIRE(build_result.has_value()); + REQUIRE(build_result.value().empty()); + REQUIRE(index->GetNumElements() == base_count); + + const auto stats = vsag::JsonType::Parse(index->GetStats()); + REQUIRE(stats["duplicate_ratio"].GetFloat() > 0.6F); + + std::vector expected_distances; + expected_distances.reserve(base_count); + for (const auto label : labels) { + auto distance = index->CalcDistanceById(vectors.data(), label); + REQUIRE(distance.has_value()); + REQUIRE(std::isfinite(distance.value())); + expected_distances.push_back(distance.value()); + } + REQUIRE(std::fabs(expected_distances.front() - expected_distances.back()) > 1e-4F); + const auto expected_distance = [&](int64_t label) { + const auto found = std::find(labels.begin(), labels.end(), label); + return expected_distances[static_cast(std::distance(labels.begin(), found))]; + }; + + auto require_result = [&](const vsag::DatasetPtr& result, + const std::vector& expected_labels) { + REQUIRE(result != nullptr); + REQUIRE(result->GetDim() == static_cast(expected_labels.size())); + std::vector actual_labels(result->GetIds(), result->GetIds() + result->GetDim()); + std::sort(actual_labels.begin(), actual_labels.end()); + auto sorted_expected = expected_labels; + std::sort(sorted_expected.begin(), sorted_expected.end()); + REQUIRE(actual_labels == sorted_expected); + for (int64_t i = 0; i < result->GetDim(); ++i) { + const auto distance = expected_distance(result->GetIds()[i]); + const float tolerance = 2e-5F * std::max(1.0F, std::fabs(distance)); + REQUIRE(std::fabs(result->GetDistances()[i] - distance) <= tolerance); + } + }; + + const auto single_search_param = R"({"hgraph":{"ef_search":3,"rabitq_one_bit_search":true}})"; + const auto parallel_search_param = R"({ + "hgraph":{ + "ef_search":3, + "rabitq_one_bit_search":true, + "parallelism":2 + } + })"; + auto alias_only = std::make_shared(labels.back()); + + for (const auto* search_param : {single_search_param, parallel_search_param}) { + auto knn = index->KnnSearch(query, base_count, search_param); + REQUIRE(knn.has_value()); + require_result(knn.value(), labels); + + auto filtered_knn = index->KnnSearch(query, 1, search_param, alias_only); + REQUIRE(filtered_knn.has_value()); + require_result(filtered_knn.value(), {labels.back()}); + + const auto max_distance = + *std::max_element(expected_distances.begin(), expected_distances.end()); + const float radius = max_distance + 2e-5F * std::max(1.0F, std::fabs(max_distance)); + auto range = index->RangeSearch(query, radius, search_param); + REQUIRE(range.has_value()); + require_result(range.value(), labels); + + auto filtered_range = index->RangeSearch(query, radius, search_param, alias_only); + REQUIRE(filtered_range.has_value()); + require_result(filtered_range.value(), {labels.back()}); + } + + for (const bool one_bit_search : {false, true}) { + const auto iterator_search_param = fmt::format( + R"({{"hgraph":{{"ef_search":3,"rabitq_one_bit_search":{}}}}})", one_bit_search); + vsag::IteratorContext* iterator_context = nullptr; + vsag::FilterPtr iterator_filter = nullptr; + std::vector iterator_labels; + for (uint64_t i = 0; i < base_count; ++i) { + auto page = index->KnnSearch(query, + 1, + iterator_search_param, + iterator_filter, + iterator_context, + i + 1 == base_count); + REQUIRE(page.has_value()); + REQUIRE(page.value()->GetDim() == 1); + iterator_labels.push_back(page.value()->GetIds()[0]); + const auto distance = expected_distance(page.value()->GetIds()[0]); + const float tolerance = 2e-5F * std::max(1.0F, std::fabs(distance)); + REQUIRE(std::fabs(page.value()->GetDistances()[0] - distance) <= tolerance); + } + std::sort(iterator_labels.begin(), iterator_labels.end()); + REQUIRE(iterator_labels == labels); + delete iterator_context; + + iterator_context = nullptr; + auto filtered_page = + index->KnnSearch(query, 1, iterator_search_param, alias_only, iterator_context, false); + REQUIRE(filtered_page.has_value()); + require_result(filtered_page.value(), {labels.back()}); + delete iterator_context; + } +}