The HNSW index keeps its own copy of every vector, which at 100K x 384 f32 is ~146 MiB — the largest single item in the 2.43x footprint now that the agent stores embeddings once. `Storage::Int8` cuts that copy to a quarter by scaling each row to i8. The scale is per row, not global. Unit-length rows in d dimensions have components around 1/sqrt(d), so a fixed [-1, 1] scale spends fewer than 12 of the 255 levels on a 128-dimensional vector; measured against an exact ranking that gives 0.35 top-10 overlap. Scaling each row by its own largest component uses the full range and brings it to 0.99. Quantised distances still cost recall on their own, and `ef` does not buy it back because the loss is in the distances rather than the graph: at N=100K recall@10 tops out at 0.967 against f32's 0.9995. Re-scoring a wider candidate pool against the exact vectors removes the gap (0.9940 vs 0.9945 at ef=64) for ~13% of query throughput and ~16% of build time. That is the intended use, so it is what the test asserts — against ground truth, not against the f32 index, whose own mistakes a re-scored search is entitled to get right. Default is unchanged: `Storage::Float32`, chosen by every existing constructor. Serialized indexes carry f32 vectors and no storage tag, so a quantised index is rebuilt rather than loaded; `compact()` keeps the storage it was given. The harness grows `--int8` and `--rerank` axes, and reports the storage in each table header. Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>
2445 lines
87 KiB
Rust
2445 lines
87 KiB
Rust
//! HNSW index implementation with HDF5 serialization.
|
|
|
|
use std::collections::BinaryHeap;
|
|
|
|
use clawhdf5_format::attribute::extract_attributes_full;
|
|
use clawhdf5_format::data_layout::DataLayout;
|
|
use clawhdf5_format::data_read::{read_as_f32, read_as_i32, read_raw_data_full};
|
|
use clawhdf5_format::dataspace::Dataspace;
|
|
use clawhdf5_format::datatype::Datatype;
|
|
use clawhdf5_format::error::FormatError;
|
|
use clawhdf5_format::file_writer::{AttrValue, FileWriter as FmtWriter};
|
|
use clawhdf5_format::filter_pipeline::FilterPipeline;
|
|
use clawhdf5_format::group_v2::resolve_path_any;
|
|
use clawhdf5_format::message_type::MessageType;
|
|
use clawhdf5_format::object_header::ObjectHeader;
|
|
use clawhdf5_format::signature::find_signature;
|
|
use clawhdf5_format::superblock::Superblock;
|
|
use clawhdf5_io::FileWriter as IoFileWriter;
|
|
|
|
/// Distance metric for the HNSW index.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
pub enum DistanceMetric {
|
|
/// L2 (Euclidean) distance.
|
|
L2,
|
|
/// Cosine distance (1 - cosine_similarity).
|
|
Cosine,
|
|
}
|
|
|
|
impl DistanceMetric {
|
|
fn as_str(self) -> &'static str {
|
|
match self {
|
|
DistanceMetric::L2 => "l2",
|
|
DistanceMetric::Cosine => "cosine",
|
|
}
|
|
}
|
|
|
|
fn from_str(s: &str) -> Option<Self> {
|
|
match s {
|
|
"l2" => Some(DistanceMetric::L2),
|
|
"cosine" => Some(DistanceMetric::Cosine),
|
|
_ => None,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Compute distance between two vectors using the given metric.
|
|
///
|
|
/// Delegates to `clawhdf5-accel`'s runtime-dispatched SIMD kernels (AVX2 on
|
|
/// x86_64, NEON on aarch64, portable scalar fallback elsewhere) — this is
|
|
/// the hottest loop in both HNSW build and every `hybrid_search` query.
|
|
fn compute_distance(a: &[f32], b: &[f32], metric: DistanceMetric) -> f32 {
|
|
match metric {
|
|
DistanceMetric::L2 => clawhdf5_accel::l2_distance(a, b),
|
|
// Both sides are unit length (see `prepare`), so cosine similarity is
|
|
// the plain dot product. Computing it as dot / (|a| * |b|) re-derived
|
|
// both norms on every call — three reductions instead of one, in the
|
|
// innermost loop of both build and search.
|
|
DistanceMetric::Cosine => 1.0 - clawhdf5_accel::dot_product(a, b),
|
|
}
|
|
}
|
|
|
|
/// Put a vector in the form the index stores and compares: unit length for the
|
|
/// cosine metric, unchanged for L2. A zero vector stays zero, giving distance 1
|
|
/// to everything — what the cosine kernel reports for a degenerate input.
|
|
fn prepare(mut v: Vec<f32>, metric: DistanceMetric) -> Vec<f32> {
|
|
if metric == DistanceMetric::Cosine {
|
|
let norm = clawhdf5_accel::vector_norm(&v);
|
|
if norm > f32::EPSILON {
|
|
let inv = 1.0 / norm;
|
|
v.iter_mut().for_each(|x| *x *= inv);
|
|
} else {
|
|
v.iter_mut().for_each(|x| *x = 0.0);
|
|
}
|
|
}
|
|
v
|
|
}
|
|
|
|
/// Assign a random level to a new node based on the HNSW probability distribution.
|
|
///
|
|
/// Uses a deterministic approach based on the node index for reproducibility.
|
|
fn assign_level(node_id: usize, m: usize) -> usize {
|
|
let ml = 1.0 / (m as f64).ln();
|
|
// Use a simple hash-based pseudo-random for reproducibility
|
|
let hash = splitmix64(node_id as u64);
|
|
let uniform = (hash >> 11) as f64 / (1u64 << 53) as f64;
|
|
(-uniform.ln() * ml).floor() as usize
|
|
}
|
|
|
|
/// Simple splitmix64 hash for deterministic level assignment.
|
|
fn splitmix64(mut x: u64) -> u64 {
|
|
x = x.wrapping_add(0x9e3779b97f4a7c15);
|
|
x = (x ^ (x >> 30)).wrapping_mul(0xbf58476d1ce4e5b9);
|
|
x = (x ^ (x >> 27)).wrapping_mul(0x94d049bb133111eb);
|
|
x ^ (x >> 31)
|
|
}
|
|
|
|
/// Candidate neighbor for priority queue operations.
|
|
#[derive(Debug, Clone)]
|
|
struct Candidate {
|
|
id: usize,
|
|
distance: f32,
|
|
}
|
|
|
|
impl PartialEq for Candidate {
|
|
fn eq(&self, other: &Self) -> bool {
|
|
self.distance.to_bits() == other.distance.to_bits() && self.id == other.id
|
|
}
|
|
}
|
|
|
|
impl Eq for Candidate {}
|
|
|
|
impl PartialOrd for Candidate {
|
|
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
|
|
Some(self.cmp(other))
|
|
}
|
|
}
|
|
|
|
impl Ord for Candidate {
|
|
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
|
|
// Reverse ordering for min-heap behavior
|
|
other
|
|
.distance
|
|
.partial_cmp(&self.distance)
|
|
.unwrap_or(std::cmp::Ordering::Equal)
|
|
}
|
|
}
|
|
|
|
/// Max-heap candidate (furthest first).
|
|
#[derive(Debug, Clone)]
|
|
struct FarCandidate {
|
|
id: usize,
|
|
distance: f32,
|
|
}
|
|
|
|
impl PartialEq for FarCandidate {
|
|
fn eq(&self, other: &Self) -> bool {
|
|
self.distance.to_bits() == other.distance.to_bits() && self.id == other.id
|
|
}
|
|
}
|
|
|
|
impl Eq for FarCandidate {}
|
|
|
|
impl PartialOrd for FarCandidate {
|
|
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
|
|
Some(self.cmp(other))
|
|
}
|
|
}
|
|
|
|
impl Ord for FarCandidate {
|
|
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
|
|
self.distance
|
|
.partial_cmp(&other.distance)
|
|
.unwrap_or(std::cmp::Ordering::Equal)
|
|
}
|
|
}
|
|
|
|
/// How the index keeps its copy of the vectors.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
|
|
pub enum Storage {
|
|
/// Exactly as given: `dim * 4` bytes per vector.
|
|
#[default]
|
|
Float32,
|
|
/// Each component scaled to an `i8`: `dim` bytes per vector, a quarter of
|
|
/// the space, at some cost in precision.
|
|
///
|
|
/// Only meaningful for [`DistanceMetric::Cosine`]: rows are stored
|
|
/// unit-length, so a quantised dot product reconstructs the similarity
|
|
/// directly. Requesting it for `L2` keeps `Float32`, because an L2
|
|
/// distance cannot be recovered from a dot product alone.
|
|
Int8,
|
|
}
|
|
|
|
/// The index's copy of the vectors, flat and row-major.
|
|
#[derive(Debug, Clone)]
|
|
enum Vectors {
|
|
F32 {
|
|
dim: usize,
|
|
flat: Vec<f32>,
|
|
},
|
|
/// `flat[i * dim + j]` is component `j` of vector `i` divided by
|
|
/// `scales[i]`; multiplying back recovers it.
|
|
///
|
|
/// The scale is per row rather than global. A unit-length row in `d`
|
|
/// dimensions has components around `1/sqrt(d)`, so a fixed `[-1, 1]`
|
|
/// scale spends fewer than 12 of the 255 levels on a 128-dimensional
|
|
/// vector and the reconstruction error swamps the gaps between near
|
|
/// neighbours — measured at 0.35 top-10 overlap with the exact ranking.
|
|
/// Scaling each row by its own largest component uses the full range.
|
|
Int8 {
|
|
dim: usize,
|
|
flat: Vec<i8>,
|
|
scales: Vec<f32>,
|
|
},
|
|
}
|
|
|
|
/// Levels either side of zero. 127, not 128, so the range is symmetric.
|
|
const INT8_LEVELS: f32 = 127.0;
|
|
|
|
/// Quantise one row, returning the codes and the scale that inverts them.
|
|
fn quantise_row(v: &[f32], out: &mut Vec<i8>) -> f32 {
|
|
let max_abs = v.iter().fold(0.0f32, |m, x| m.max(x.abs()));
|
|
if max_abs <= f32::MIN_POSITIVE {
|
|
out.extend(core::iter::repeat_n(0i8, v.len()));
|
|
return 0.0;
|
|
}
|
|
let inv = INT8_LEVELS / max_abs;
|
|
out.extend(
|
|
v.iter()
|
|
.map(|x| (x * inv).round().clamp(-INT8_LEVELS, INT8_LEVELS) as i8),
|
|
);
|
|
max_abs / INT8_LEVELS
|
|
}
|
|
|
|
impl Vectors {
|
|
fn new(dim: usize, storage: Storage, metric: DistanceMetric) -> Self {
|
|
match storage {
|
|
Storage::Int8 if metric == DistanceMetric::Cosine => Vectors::Int8 {
|
|
dim,
|
|
flat: Vec::new(),
|
|
scales: Vec::new(),
|
|
},
|
|
_ => Vectors::F32 {
|
|
dim,
|
|
flat: Vec::new(),
|
|
},
|
|
}
|
|
}
|
|
|
|
fn dim(&self) -> usize {
|
|
match self {
|
|
Vectors::F32 { dim, .. } | Vectors::Int8 { dim, .. } => *dim,
|
|
}
|
|
}
|
|
|
|
fn storage(&self) -> Storage {
|
|
match self {
|
|
Vectors::F32 { .. } => Storage::Float32,
|
|
Vectors::Int8 { .. } => Storage::Int8,
|
|
}
|
|
}
|
|
|
|
fn len(&self) -> usize {
|
|
let dim = self.dim();
|
|
if dim == 0 {
|
|
return 0;
|
|
}
|
|
match self {
|
|
Vectors::F32 { flat, .. } => flat.len() / dim,
|
|
Vectors::Int8 { flat, .. } => flat.len() / dim,
|
|
}
|
|
}
|
|
|
|
/// Set the row width, for a store seeded empty by `new`.
|
|
fn set_dim(&mut self, new_dim: usize) {
|
|
match self {
|
|
Vectors::F32 { dim, .. } | Vectors::Int8 { dim, .. } => *dim = new_dim,
|
|
}
|
|
}
|
|
|
|
fn push(&mut self, vector: &[f32]) {
|
|
match self {
|
|
Vectors::F32 { flat, .. } => flat.extend_from_slice(vector),
|
|
Vectors::Int8 { flat, scales, .. } => scales.push(quantise_row(vector, flat)),
|
|
}
|
|
}
|
|
|
|
/// Row `i` as `f32`, for callers that need the values back (serialization,
|
|
/// and the f32 fast paths). Quantised rows are reconstructed, so this is
|
|
/// lossy in exactly the way the storage is.
|
|
fn row(&self, i: usize) -> Vec<f32> {
|
|
let dim = self.dim();
|
|
let start = i * dim;
|
|
match self {
|
|
Vectors::F32 { flat, .. } => flat[start..start + dim].to_vec(),
|
|
Vectors::Int8 { flat, scales, .. } => flat[start..start + dim]
|
|
.iter()
|
|
.map(|&q| f32::from(q) * scales[i])
|
|
.collect(),
|
|
}
|
|
}
|
|
|
|
/// Distance between two stored vectors.
|
|
fn dist(&self, a: usize, b: usize, metric: DistanceMetric) -> f32 {
|
|
let dim = self.dim();
|
|
match self {
|
|
Vectors::F32 { flat, .. } => {
|
|
let (x, y) = (a * dim, b * dim);
|
|
compute_distance(&flat[x..x + dim], &flat[y..y + dim], metric)
|
|
}
|
|
Vectors::Int8 { flat, scales, .. } => {
|
|
let (x, y) = (a * dim, b * dim);
|
|
let dot = dot_i8(&flat[x..x + dim], &flat[y..y + dim]);
|
|
1.0 - dot as f32 * scales[a] * scales[b]
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Distance from a prepared query to stored vector `i`.
|
|
fn dist_query(&self, query: &Query, i: usize, metric: DistanceMetric) -> f32 {
|
|
let dim = self.dim();
|
|
let start = i * dim;
|
|
match (self, query) {
|
|
(Vectors::F32 { flat, .. }, Query::F32(q)) => {
|
|
compute_distance(q, &flat[start..start + dim], metric)
|
|
}
|
|
(Vectors::Int8 { flat, scales, .. }, Query::Int8(q, q_scale)) => {
|
|
let dot = dot_i8(q, &flat[start..start + dim]);
|
|
1.0 - dot as f32 * q_scale * scales[i]
|
|
}
|
|
// Mixed forms cannot occur: `Query` is built from the same storage.
|
|
_ => f32::MAX,
|
|
}
|
|
}
|
|
|
|
/// Build a store from prepared rows.
|
|
fn from_rows(rows: &[Vec<f32>], storage: Storage, metric: DistanceMetric) -> Self {
|
|
let dim = rows.first().map_or(0, Vec::len);
|
|
let mut out = Vectors::new(dim, storage, metric);
|
|
for row in rows {
|
|
out.push(row);
|
|
}
|
|
out
|
|
}
|
|
|
|
/// Prepare `query` for comparison against this store.
|
|
fn query(&self, query: Vec<f32>) -> Query {
|
|
match self {
|
|
Vectors::F32 { .. } => Query::F32(query),
|
|
Vectors::Int8 { .. } => {
|
|
let mut codes = Vec::with_capacity(query.len());
|
|
let scale = quantise_row(&query, &mut codes);
|
|
Query::Int8(codes, scale)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
/// What a layer search is measuring distance *to*: an incoming query, or a
|
|
/// node already in the index (which is what insertion compares against).
|
|
enum Target<'a> {
|
|
Query(&'a Query),
|
|
Node(usize),
|
|
}
|
|
|
|
impl Vectors {
|
|
fn dist_to(&self, target: &Target<'_>, i: usize, metric: DistanceMetric) -> f32 {
|
|
match target {
|
|
Target::Query(q) => self.dist_query(q, i, metric),
|
|
Target::Node(n) => self.dist(*n, i, metric),
|
|
}
|
|
}
|
|
}
|
|
|
|
/// A search query in whichever form the store compares against.
|
|
enum Query {
|
|
F32(Vec<f32>),
|
|
/// Codes and the scale that inverts them, as in [`Vectors::Int8`].
|
|
Int8(Vec<i8>, f32),
|
|
}
|
|
|
|
/// Sum of products, widened so it cannot overflow: `dim` terms of at most
|
|
/// `127 * 127`, so `i32` suffices for any realistic dimension.
|
|
fn dot_i8(a: &[i8], b: &[i8]) -> i32 {
|
|
// Four independent accumulators over 32-lane blocks: the widening product
|
|
// has to sit in a fixed-length chunk for the vectoriser to see it, and the
|
|
// separate accumulators keep it off one dependency chain.
|
|
const LANE: usize = 8;
|
|
let (a_blocks, a_tail) = a.as_chunks::<{ LANE * 4 }>();
|
|
let (b_blocks, b_tail) = b.as_chunks::<{ LANE * 4 }>();
|
|
let mut acc = [0i32; 4];
|
|
for (x, y) in a_blocks.iter().zip(b_blocks) {
|
|
for (lane, slot) in acc.iter_mut().enumerate() {
|
|
let mut sum = 0i32;
|
|
for k in 0..LANE {
|
|
sum += i32::from(x[lane * LANE + k]) * i32::from(y[lane * LANE + k]);
|
|
}
|
|
*slot += sum;
|
|
}
|
|
}
|
|
let tail: i32 = a_tail
|
|
.iter()
|
|
.zip(b_tail)
|
|
.map(|(&x, &y)| i32::from(x) * i32::from(y))
|
|
.sum();
|
|
acc[0] + acc[1] + acc[2] + acc[3] + tail
|
|
}
|
|
|
|
/// Magic for [`HnswIndex::graph_to_bytes`].
|
|
const GRAPH_MAGIC: &[u8; 4] = b"CHG1";
|
|
|
|
/// On-disk format version for the serialized HNSW index.
|
|
///
|
|
/// - Version 1: original layout (`vectors`, `graph_layer_*`, `config`), no
|
|
/// deletion support and no explicit version tag.
|
|
/// - Version 2: adds a `format_version` attribute and a `deleted` bitset dataset
|
|
/// so live insert/delete state survives a save/load round-trip.
|
|
///
|
|
/// Files written before this constant existed are treated as version 1 on load.
|
|
pub const HNSW_FORMAT_VERSION: i64 = 2;
|
|
|
|
/// HNSW (Hierarchical Navigable Small World) approximate nearest neighbor index.
|
|
///
|
|
/// Supports building an index from vectors, incremental insertion and soft
|
|
/// deletion, searching for nearest neighbors, and serializing/deserializing to
|
|
/// HDF5 format.
|
|
#[derive(Debug, Clone)]
|
|
pub struct HnswIndex {
|
|
/// All vectors in the index, flat and row-major.
|
|
vectors: Vectors,
|
|
/// Adjacency lists per layer. `graph[layer][node]` = list of neighbor IDs.
|
|
graph: Vec<Vec<Vec<usize>>>,
|
|
/// Soft-deletion flags, one per node. Deleted nodes remain in the graph for
|
|
/// connectivity but are never returned from [`HnswIndex::search`].
|
|
deleted: Vec<bool>,
|
|
/// Entry point node ID.
|
|
entry_point: usize,
|
|
/// Maximum number of connections per node (per layer).
|
|
m: usize,
|
|
/// Maximum connections for layer 0 (typically 2*m).
|
|
m_max0: usize,
|
|
/// ef parameter used during construction.
|
|
ef_construction: usize,
|
|
/// Maximum layer assigned to each node.
|
|
node_levels: Vec<usize>,
|
|
/// Distance metric.
|
|
metric: DistanceMetric,
|
|
}
|
|
|
|
impl HnswIndex {
|
|
/// Build an HNSW index from a set of vectors.
|
|
///
|
|
/// # Parameters
|
|
/// - `vectors`: The vectors to index. All must have the same dimension.
|
|
/// - `m`: Maximum number of connections per node (higher = more accurate, more memory).
|
|
/// - `ef_construction`: Size of the dynamic candidate list during construction.
|
|
///
|
|
/// Uses L2 distance by default. Use [`build_with_metric`] to specify the metric.
|
|
pub fn build(vectors: &[Vec<f32>], m: usize, ef_construction: usize) -> Self {
|
|
Self::build_with_metric(vectors, m, ef_construction, DistanceMetric::L2)
|
|
}
|
|
|
|
/// Build an HNSW index with a specific distance metric.
|
|
pub fn build_with_metric(
|
|
vectors: &[Vec<f32>],
|
|
m: usize,
|
|
ef_construction: usize,
|
|
metric: DistanceMetric,
|
|
) -> Self {
|
|
Self::build_with(vectors, m, ef_construction, metric, Storage::default())
|
|
}
|
|
|
|
/// Build an index, choosing how the vectors are stored.
|
|
///
|
|
/// [`Storage::Int8`] keeps them at a quarter of the size; see its docs for
|
|
/// what that costs and when it applies.
|
|
pub fn build_with(
|
|
vectors: &[Vec<f32>],
|
|
m: usize,
|
|
ef_construction: usize,
|
|
metric: DistanceMetric,
|
|
storage: Storage,
|
|
) -> Self {
|
|
assert!(!vectors.is_empty(), "cannot build index from empty vectors");
|
|
assert!(m >= 2, "m must be at least 2");
|
|
let dim = vectors[0].len();
|
|
for v in vectors {
|
|
assert_eq!(v.len(), dim, "all vectors must have the same dimension");
|
|
}
|
|
|
|
let m_max0 = m * 2;
|
|
let n = vectors.len();
|
|
let mut prepared = Vectors::new(dim, storage, metric);
|
|
for v in vectors {
|
|
prepared.push(&prepare(v.clone(), metric));
|
|
}
|
|
let vectors = &prepared;
|
|
|
|
// Assign levels to all nodes
|
|
let mut node_levels = Vec::with_capacity(n);
|
|
let mut max_level = 0;
|
|
for i in 0..n {
|
|
let level = assign_level(i, m);
|
|
if level > max_level {
|
|
max_level = level;
|
|
}
|
|
node_levels.push(level);
|
|
}
|
|
|
|
// Initialize graph layers
|
|
let num_layers = max_level + 1;
|
|
let mut graph: Vec<Vec<Vec<usize>>> = Vec::with_capacity(num_layers);
|
|
for _ in 0..num_layers {
|
|
graph.push(vec![Vec::new(); n]);
|
|
}
|
|
|
|
let mut entry_point = 0;
|
|
let mut ep_level = node_levels[0];
|
|
|
|
// Insert in batches. Each batch is planned against the graph as it
|
|
// stood when the batch began (read-only, so the plans are independent
|
|
// and run in parallel with the `parallel` feature), then linked, then
|
|
// every neighbour list that overflowed is pruned once. Pruning is ~90%
|
|
// of a build's distance evaluations, and a node that gains several
|
|
// back-links in one batch is pruned once instead of once per link.
|
|
//
|
|
// Nodes in the same batch cannot see each other while planning, so
|
|
// batches start at one node and grow only as the graph does — a batch
|
|
// is never more than a small fraction of what is already linked. The
|
|
// result is deterministic and identical with or without `parallel`.
|
|
let mut next = 1;
|
|
while next < n {
|
|
let mut end = (next + batch_len(next)).min(n);
|
|
// A node that raises the top layer becomes the new entry point and
|
|
// changes how every later node descends: give it a batch alone.
|
|
if let Some(tall) = (next..end).find(|&i| node_levels[i] > ep_level) {
|
|
end = if tall == next { next + 1 } else { tall };
|
|
}
|
|
|
|
let plans = plan_batch(
|
|
vectors,
|
|
&graph,
|
|
&node_levels,
|
|
next..end,
|
|
entry_point,
|
|
ep_level,
|
|
(m, m_max0, ef_construction),
|
|
metric,
|
|
);
|
|
|
|
let mut overflowed: Vec<(usize, usize)> = Vec::new();
|
|
for (offset, plan) in plans.into_iter().enumerate() {
|
|
let node = next + offset;
|
|
for (layer, selected) in plan {
|
|
let max_conn = if layer == 0 { m_max0 } else { m };
|
|
for &neighbor in &selected {
|
|
let list = &mut graph[layer][neighbor];
|
|
list.push(node);
|
|
if list.len() == max_conn + 1 {
|
|
overflowed.push((layer, neighbor));
|
|
}
|
|
}
|
|
graph[layer][node] = selected;
|
|
}
|
|
}
|
|
prune_overflowed(vectors, &mut graph, overflowed, (m, m_max0), metric);
|
|
|
|
for (i, &level) in node_levels.iter().enumerate().take(end).skip(next) {
|
|
if level > ep_level {
|
|
entry_point = i;
|
|
ep_level = level;
|
|
}
|
|
}
|
|
next = end;
|
|
}
|
|
|
|
Self {
|
|
vectors: prepared,
|
|
graph,
|
|
deleted: vec![false; n],
|
|
entry_point,
|
|
m,
|
|
m_max0,
|
|
ef_construction,
|
|
node_levels,
|
|
metric,
|
|
}
|
|
}
|
|
|
|
/// Create an empty index with the given parameters. Used as the starting
|
|
/// point for incremental [`HnswIndex::insert`] and as the result of
|
|
/// [`HnswIndex::compact`] when every vector has been deleted.
|
|
pub fn new(m: usize, ef_construction: usize, metric: DistanceMetric) -> Self {
|
|
Self::new_with(m, ef_construction, metric, Storage::default())
|
|
}
|
|
|
|
/// [`HnswIndex::new`], choosing how the vectors are stored.
|
|
pub fn new_with(
|
|
m: usize,
|
|
ef_construction: usize,
|
|
metric: DistanceMetric,
|
|
storage: Storage,
|
|
) -> Self {
|
|
assert!(m >= 2, "m must be at least 2");
|
|
Self {
|
|
// The dimension is set by the first insert.
|
|
vectors: Vectors::new(0, storage, metric),
|
|
graph: Vec::new(),
|
|
deleted: Vec::new(),
|
|
entry_point: 0,
|
|
m,
|
|
m_max0: m * 2,
|
|
ef_construction,
|
|
node_levels: Vec::new(),
|
|
metric,
|
|
}
|
|
}
|
|
|
|
/// Insert a single vector into the index incrementally and return its id.
|
|
///
|
|
/// The id is the vector's position in insertion order and is stable for the
|
|
/// life of the index (until [`HnswIndex::compact`] renumbers survivors).
|
|
/// Inserting into an empty index seeds the entry point.
|
|
///
|
|
/// # Panics
|
|
/// Panics if `vector`'s dimension does not match the existing vectors.
|
|
pub fn insert(&mut self, vector: Vec<f32>) -> usize {
|
|
let vector = prepare(vector, self.metric);
|
|
let id = self.vectors.len();
|
|
|
|
// Seed an empty index.
|
|
if id == 0 {
|
|
let node_level = assign_level(0, self.m);
|
|
self.vectors.set_dim(vector.len());
|
|
self.vectors.push(&vector);
|
|
self.deleted.push(false);
|
|
self.node_levels.push(node_level);
|
|
self.graph = (0..=node_level).map(|_| vec![Vec::new(); 1]).collect();
|
|
self.entry_point = 0;
|
|
return 0;
|
|
}
|
|
|
|
assert_eq!(
|
|
vector.len(),
|
|
self.vectors.dim(),
|
|
"insert dimension mismatch"
|
|
);
|
|
|
|
let node_level = assign_level(id, self.m);
|
|
self.vectors.push(&vector);
|
|
self.deleted.push(false);
|
|
self.node_levels.push(node_level);
|
|
|
|
// Grow every existing layer with an empty adjacency slot for `id`, and
|
|
// add any brand-new top layers this node introduces.
|
|
for layer in self.graph.iter_mut() {
|
|
layer.push(Vec::new());
|
|
}
|
|
while self.graph.len() <= node_level {
|
|
self.graph.push(vec![Vec::new(); id + 1]);
|
|
}
|
|
|
|
let ep_level = self.node_levels[self.entry_point];
|
|
let mut ep = self.entry_point;
|
|
|
|
// Phase 1: greedy descent from the top down to node_level + 1.
|
|
for layer in (node_level + 1..=ep_level).rev() {
|
|
ep = greedy_closest(
|
|
&self.vectors,
|
|
&self.graph[layer],
|
|
&Target::Node(id),
|
|
ep,
|
|
self.metric,
|
|
);
|
|
}
|
|
|
|
// Phase 2: search and connect from min(node_level, ep_level) down to 0.
|
|
let bottom = node_level.min(ep_level);
|
|
for layer in (0..=bottom).rev() {
|
|
let max_conn = if layer == 0 { self.m_max0 } else { self.m };
|
|
let neighbors = search_layer(
|
|
&self.vectors,
|
|
&self.graph[layer],
|
|
&Target::Node(id),
|
|
ep,
|
|
self.ef_construction,
|
|
self.metric,
|
|
None,
|
|
);
|
|
let scored: Vec<(usize, f32)> = neighbors.iter().map(|c| (c.id, c.distance)).collect();
|
|
let selected = select_neighbors(&self.vectors, &scored, max_conn, self.metric);
|
|
self.graph[layer][id] = selected.clone();
|
|
link_back(
|
|
&self.vectors,
|
|
&mut self.graph[layer],
|
|
id,
|
|
&selected,
|
|
max_conn,
|
|
self.metric,
|
|
);
|
|
if !selected.is_empty() {
|
|
ep = selected[0];
|
|
}
|
|
}
|
|
|
|
// Promote the entry point if this node sits on a higher layer.
|
|
if node_level > ep_level {
|
|
self.entry_point = id;
|
|
}
|
|
|
|
id
|
|
}
|
|
|
|
/// Soft-delete the vector with the given id. The node stays in the graph so
|
|
/// traversal/connectivity is preserved, but it will never be returned from
|
|
/// [`HnswIndex::search`]. Idempotent; out-of-range ids are ignored.
|
|
///
|
|
/// Returns `true` if the id existed and was not already deleted.
|
|
pub fn mark_deleted(&mut self, id: usize) -> bool {
|
|
if id >= self.deleted.len() || self.deleted[id] {
|
|
return false;
|
|
}
|
|
self.deleted[id] = true;
|
|
true
|
|
}
|
|
|
|
/// Returns whether the vector with the given id is soft-deleted.
|
|
pub fn is_deleted(&self, id: usize) -> bool {
|
|
self.deleted.get(id).copied().unwrap_or(false)
|
|
}
|
|
|
|
/// Number of soft-deleted vectors still occupying the index.
|
|
pub fn deleted_count(&self) -> usize {
|
|
self.deleted.iter().filter(|&&d| d).count()
|
|
}
|
|
|
|
/// Number of live (non-deleted) vectors.
|
|
pub fn active_len(&self) -> usize {
|
|
self.vectors.len() - self.deleted_count()
|
|
}
|
|
|
|
/// Rebuild the index from scratch, dropping all soft-deleted vectors and
|
|
/// renumbering the survivors into a compact `0..active_len` id space.
|
|
///
|
|
/// Returns a mapping from old id to new id (`None` for dropped vectors) so
|
|
/// callers can rewrite any external id references they keep.
|
|
pub fn compact(&mut self) -> Vec<Option<usize>> {
|
|
let mut mapping = vec![None; self.vectors.len()];
|
|
let mut surviving: Vec<Vec<f32>> = Vec::with_capacity(self.active_len());
|
|
for (old, slot) in mapping.iter_mut().enumerate() {
|
|
if !self.deleted[old] {
|
|
*slot = Some(surviving.len());
|
|
surviving.push(self.vectors.row(old));
|
|
}
|
|
}
|
|
// Rebuilding must keep the storage the caller chose; a compaction is
|
|
// not the place to silently quadruple the index's memory.
|
|
let storage = self.vectors.storage();
|
|
*self = if surviving.is_empty() {
|
|
Self::new_with(self.m, self.ef_construction, self.metric, storage)
|
|
} else {
|
|
Self::build_with(
|
|
&surviving,
|
|
self.m,
|
|
self.ef_construction,
|
|
self.metric,
|
|
storage,
|
|
)
|
|
};
|
|
mapping
|
|
}
|
|
|
|
/// Search the index for the `k` nearest neighbors to the query vector.
|
|
///
|
|
/// # Parameters
|
|
/// - `query`: The query vector.
|
|
/// - `k`: Number of nearest neighbors to return.
|
|
/// - `ef`: Size of the dynamic candidate list during search (must be >= k).
|
|
///
|
|
/// # Returns
|
|
/// A vector of `(id, distance)` pairs sorted by distance (closest first).
|
|
pub fn search(&self, query: &[f32], k: usize, ef: usize) -> Vec<(usize, f32)> {
|
|
if self.vectors.len() == 0 {
|
|
return Vec::new();
|
|
}
|
|
assert_eq!(query.len(), self.vectors.dim(), "query dimension mismatch");
|
|
let ef = ef.max(k);
|
|
// Prepared and, for a quantised store, quantised once per search
|
|
// rather than once per comparison.
|
|
let prepared = self.vectors.query(prepare(query.to_vec(), self.metric));
|
|
let target = Target::Query(&prepared);
|
|
|
|
let mut ep = self.entry_point;
|
|
let top_layer = self.graph.len().saturating_sub(1);
|
|
|
|
// Greedy search from top layer down to layer 1
|
|
for layer in (1..=top_layer).rev() {
|
|
ep = greedy_closest(&self.vectors, &self.graph[layer], &target, ep, self.metric);
|
|
}
|
|
|
|
// Search layer 0 for the ef nearest *live* nodes. Deleted nodes are
|
|
// still traversed (they remain valid graph waypoints) but take no
|
|
// result slot, so deletions near the query don't shrink the answer.
|
|
let candidates = search_layer(
|
|
&self.vectors,
|
|
&self.graph[0],
|
|
&target,
|
|
ep,
|
|
ef,
|
|
self.metric,
|
|
Some(&self.deleted),
|
|
);
|
|
|
|
candidates
|
|
.into_iter()
|
|
.take(k)
|
|
.map(|c| (c.id, c.distance))
|
|
.collect()
|
|
}
|
|
|
|
/// Save the index to an HDF5 file via the given writer.
|
|
pub fn save_to_hdf5(&self, writer: &mut IoFileWriter) -> Result<(), FormatError> {
|
|
let bytes = self.to_hdf5_bytes()?;
|
|
writer
|
|
.write_bytes_owned(bytes)
|
|
.map_err(|e| FormatError::SerializationError(e.to_string()))?;
|
|
Ok(())
|
|
}
|
|
|
|
/// Serialize the index to HDF5 bytes.
|
|
pub fn to_hdf5_bytes(&self) -> Result<Vec<u8>, FormatError> {
|
|
let mut fw = FmtWriter::new();
|
|
let n = self.vectors.len();
|
|
let dim = self.vectors.dim();
|
|
|
|
// Flatten vectors into a 1D array for storage
|
|
let mut flat_vectors: Vec<f32> = Vec::with_capacity(n * dim);
|
|
for i in 0..n {
|
|
flat_vectors.extend_from_slice(&self.vectors.row(i));
|
|
}
|
|
|
|
let mut group = fw.create_group("ann");
|
|
|
|
group
|
|
.create_dataset("vectors")
|
|
.with_f32_data(&flat_vectors)
|
|
.with_shape(&[n as u64, dim as u64])
|
|
.set_attr("rows", AttrValue::I64(n as i64))
|
|
.set_attr("cols", AttrValue::I64(dim as i64));
|
|
|
|
// Serialize graph layers: store as flat i32 arrays with metadata
|
|
let num_layers = self.graph.len();
|
|
for (layer_idx, layer) in self.graph.iter().enumerate() {
|
|
// Flatten: for each node, store [count, neighbor1, neighbor2, ...]
|
|
let mut flat: Vec<i32> = Vec::new();
|
|
for neighbors in layer {
|
|
flat.push(neighbors.len() as i32);
|
|
for &n_id in neighbors {
|
|
flat.push(n_id as i32);
|
|
}
|
|
}
|
|
|
|
let ds_name = format!("graph_layer_{layer_idx}");
|
|
group
|
|
.create_dataset(&ds_name)
|
|
.with_i32_data(&flat)
|
|
.set_attr("layer", AttrValue::I64(layer_idx as i64));
|
|
}
|
|
|
|
// Soft-deletion bitset (format version 2+): 0 = live, 1 = deleted.
|
|
let deleted_i32: Vec<i32> = self.deleted.iter().map(|&d| d as i32).collect();
|
|
group.create_dataset("deleted").with_i32_data(&deleted_i32);
|
|
|
|
// Store config as attributes on a small dataset
|
|
let node_levels_i32: Vec<i32> = self.node_levels.iter().map(|&l| l as i32).collect();
|
|
group
|
|
.create_dataset("config")
|
|
.with_i32_data(&node_levels_i32)
|
|
.set_attr("format_version", AttrValue::I64(HNSW_FORMAT_VERSION))
|
|
.set_attr("m", AttrValue::I64(self.m as i64))
|
|
.set_attr(
|
|
"ef_construction",
|
|
AttrValue::I64(self.ef_construction as i64),
|
|
)
|
|
.set_attr("entry_point", AttrValue::I64(self.entry_point as i64))
|
|
.set_attr("num_layers", AttrValue::I64(num_layers as i64))
|
|
.set_attr(
|
|
"metric",
|
|
AttrValue::String(self.metric.as_str().to_string()),
|
|
)
|
|
.set_attr("num_vectors", AttrValue::I64(n as i64))
|
|
.set_attr("dimension", AttrValue::I64(dim as i64));
|
|
|
|
let finished = group.finish();
|
|
fw.add_group(finished);
|
|
fw.finish()
|
|
}
|
|
|
|
/// Load an HNSW index from HDF5 bytes.
|
|
///
|
|
/// The HDF5 data must contain the `/ann/vectors`, `/ann/graph_layer_*`,
|
|
/// and `/ann/config` datasets as produced by [`to_hdf5_bytes`].
|
|
pub fn load_from_hdf5(data: &[u8]) -> Result<Self, FormatError> {
|
|
let sig_offset = find_signature(data)?;
|
|
let sb = Superblock::parse(data, sig_offset)?;
|
|
|
|
// Read config dataset and its attributes
|
|
let config_attrs = read_dataset_attrs(data, &sb, "ann/config")?;
|
|
let config_raw = read_dataset_raw(data, &sb, "ann/config")?;
|
|
let config_dt = read_dataset_datatype(data, &sb, "ann/config")?;
|
|
let node_levels_i32 = read_as_i32(&config_raw, &config_dt)?;
|
|
|
|
// Files written before format version 2 have no version attribute; treat
|
|
// them as version 1. Reject anything newer than we understand.
|
|
let format_version = get_attr_i64_opt(&config_attrs, "format_version").unwrap_or(1);
|
|
if format_version > HNSW_FORMAT_VERSION {
|
|
return Err(FormatError::SerializationError(format!(
|
|
"unsupported HNSW format version {format_version} (this build understands up to {HNSW_FORMAT_VERSION})"
|
|
)));
|
|
}
|
|
|
|
let m = get_attr_i64(&config_attrs, "m")? as usize;
|
|
let ef_construction = get_attr_i64(&config_attrs, "ef_construction")? as usize;
|
|
let entry_point = get_attr_i64(&config_attrs, "entry_point")? as usize;
|
|
let num_layers = get_attr_i64(&config_attrs, "num_layers")? as usize;
|
|
let n = get_attr_i64(&config_attrs, "num_vectors")? as usize;
|
|
let dim = get_attr_i64(&config_attrs, "dimension")? as usize;
|
|
let metric_str = get_attr_string(&config_attrs, "metric")?;
|
|
let metric = DistanceMetric::from_str(&metric_str).ok_or_else(|| {
|
|
FormatError::SerializationError(format!("unknown metric: {metric_str}"))
|
|
})?;
|
|
|
|
let node_levels: Vec<usize> = node_levels_i32.iter().map(|&l| l as usize).collect();
|
|
|
|
// Read vectors
|
|
let vectors_raw = read_dataset_raw(data, &sb, "ann/vectors")?;
|
|
let vectors_dt = read_dataset_datatype(data, &sb, "ann/vectors")?;
|
|
let flat_vectors = read_as_f32(&vectors_raw, &vectors_dt)?;
|
|
let mut vectors = Vec::with_capacity(n);
|
|
for i in 0..n {
|
|
let start = i * dim;
|
|
let end = start + dim;
|
|
if end > flat_vectors.len() {
|
|
return Err(FormatError::DataSizeMismatch {
|
|
expected: end,
|
|
actual: flat_vectors.len(),
|
|
});
|
|
}
|
|
// Files written before vectors were stored unit-length hold the
|
|
// raw ones; preparing is idempotent, so this handles both.
|
|
vectors.push(prepare(flat_vectors[start..end].to_vec(), metric));
|
|
}
|
|
|
|
// Read graph layers
|
|
let mut graph = Vec::with_capacity(num_layers);
|
|
for layer_idx in 0..num_layers {
|
|
let ds_name = format!("ann/graph_layer_{layer_idx}");
|
|
let layer_raw = read_dataset_raw(data, &sb, &ds_name)?;
|
|
let layer_dt = read_dataset_datatype(data, &sb, &ds_name)?;
|
|
let flat = read_as_i32(&layer_raw, &layer_dt)?;
|
|
|
|
let mut layer_graph = Vec::with_capacity(n);
|
|
let mut pos = 0;
|
|
while pos < flat.len() {
|
|
let count = flat[pos] as usize;
|
|
pos += 1;
|
|
let mut neighbors = Vec::with_capacity(count);
|
|
for _ in 0..count {
|
|
if pos >= flat.len() {
|
|
return Err(FormatError::SerializationError(
|
|
"truncated graph data".into(),
|
|
));
|
|
}
|
|
neighbors.push(flat[pos] as usize);
|
|
pos += 1;
|
|
}
|
|
layer_graph.push(neighbors);
|
|
}
|
|
// Pad with empty if needed (nodes not present at this layer)
|
|
while layer_graph.len() < n {
|
|
layer_graph.push(Vec::new());
|
|
}
|
|
graph.push(layer_graph);
|
|
}
|
|
|
|
// Deleted bitset (version 2+). Older files default every node to live.
|
|
let deleted = if format_version >= 2 {
|
|
let deleted_raw = read_dataset_raw(data, &sb, "ann/deleted")?;
|
|
let deleted_dt = read_dataset_datatype(data, &sb, "ann/deleted")?;
|
|
let deleted_i32 = read_as_i32(&deleted_raw, &deleted_dt)?;
|
|
let mut deleted: Vec<bool> = deleted_i32.iter().map(|&d| d != 0).collect();
|
|
deleted.resize(n, false);
|
|
deleted
|
|
} else {
|
|
vec![false; n]
|
|
};
|
|
|
|
Ok(Self {
|
|
// Serialized files carry f32 vectors and no storage tag: a
|
|
// quantised index is rebuilt, not loaded.
|
|
vectors: Vectors::from_rows(&vectors, Storage::Float32, metric),
|
|
graph,
|
|
deleted,
|
|
entry_point,
|
|
m,
|
|
m_max0: m * 2,
|
|
ef_construction,
|
|
node_levels,
|
|
metric,
|
|
})
|
|
}
|
|
|
|
/// Serialize the **graph only** — levels, tombstones and adjacency, not the
|
|
/// vectors — for a caller that already stores the vectors elsewhere (the
|
|
/// agent's record cache). [`HnswIndex::to_hdf5_bytes`] writes a complete,
|
|
/// self-contained index including a full copy of every vector, which would
|
|
/// double such a store's size. Reattach with
|
|
/// [`HnswIndex::from_graph_bytes`].
|
|
///
|
|
/// Layout (little endian): magic `CHG1`, then u32 fields `n`, `m`,
|
|
/// `m_max0`, `ef_construction`, `entry_point`, `num_layers`, `metric`;
|
|
/// `n` level bytes; `n` tombstone bytes; per layer, per node that exists on
|
|
/// that layer: u32 neighbour count + u32 ids; trailing CRC32 of all of it.
|
|
pub fn graph_to_bytes(&self) -> Vec<u8> {
|
|
let n = self.vectors.len();
|
|
let mut out = Vec::with_capacity(32 + n * 2 + n * self.m_max0 * 4);
|
|
out.extend_from_slice(GRAPH_MAGIC);
|
|
for field in [
|
|
n,
|
|
self.m,
|
|
self.m_max0,
|
|
self.ef_construction,
|
|
self.entry_point,
|
|
self.graph.len(),
|
|
match self.metric {
|
|
DistanceMetric::L2 => 0,
|
|
DistanceMetric::Cosine => 1,
|
|
},
|
|
] {
|
|
out.extend_from_slice(&(field as u32).to_le_bytes());
|
|
}
|
|
out.extend(self.node_levels.iter().map(|&l| l.min(255) as u8));
|
|
out.extend(self.deleted.iter().map(|&d| u8::from(d)));
|
|
for (layer, adjacency) in self.graph.iter().enumerate() {
|
|
for (node, neighbors) in adjacency.iter().enumerate() {
|
|
if self.node_levels[node] < layer {
|
|
continue; // node does not exist on this layer
|
|
}
|
|
out.extend_from_slice(&(neighbors.len() as u32).to_le_bytes());
|
|
for &id in neighbors {
|
|
out.extend_from_slice(&(id as u32).to_le_bytes());
|
|
}
|
|
}
|
|
}
|
|
let crc = clawhdf5_format::checksum::crc32(&out);
|
|
out.extend_from_slice(&crc.to_le_bytes());
|
|
out
|
|
}
|
|
|
|
/// Rebuild an index from [`HnswIndex::graph_to_bytes`] output and the
|
|
/// vectors it was built over (same order). Every structural claim in
|
|
/// `bytes` is validated — a corrupt or mismatched graph is an error, never
|
|
/// an index that panics or walks out of bounds during a search.
|
|
pub fn from_graph_bytes(bytes: &[u8], vectors: Vec<Vec<f32>>) -> Result<Self, FormatError> {
|
|
Self::from_graph_bytes_with(bytes, vectors, Storage::default())
|
|
}
|
|
|
|
/// As [`from_graph_bytes`](Self::from_graph_bytes), choosing how the
|
|
/// rehydrated vectors are stored.
|
|
pub fn from_graph_bytes_with(
|
|
bytes: &[u8],
|
|
vectors: Vec<Vec<f32>>,
|
|
storage: Storage,
|
|
) -> Result<Self, FormatError> {
|
|
let bad = |what: &str| FormatError::SerializationError(format!("HNSW graph: {what}"));
|
|
let body_len = bytes
|
|
.len()
|
|
.checked_sub(4)
|
|
.filter(|&l| l >= GRAPH_MAGIC.len() + 7 * 4)
|
|
.ok_or_else(|| bad("truncated"))?;
|
|
let (body, crc_bytes) = bytes.split_at(body_len);
|
|
if &body[..4] != GRAPH_MAGIC {
|
|
return Err(bad("bad magic"));
|
|
}
|
|
let stored_crc =
|
|
u32::from_le_bytes([crc_bytes[0], crc_bytes[1], crc_bytes[2], crc_bytes[3]]);
|
|
if clawhdf5_format::checksum::crc32(body) != stored_crc {
|
|
return Err(bad("checksum mismatch"));
|
|
}
|
|
|
|
let mut pos = 4;
|
|
let next_u32 = |pos: &mut usize| -> Result<usize, FormatError> {
|
|
let b = body.get(*pos..*pos + 4).ok_or_else(|| bad("truncated"))?;
|
|
*pos += 4;
|
|
Ok(u32::from_le_bytes([b[0], b[1], b[2], b[3]]) as usize)
|
|
};
|
|
let n = next_u32(&mut pos)?;
|
|
let m = next_u32(&mut pos)?;
|
|
let m_max0 = next_u32(&mut pos)?;
|
|
let ef_construction = next_u32(&mut pos)?;
|
|
let entry_point = next_u32(&mut pos)?;
|
|
let num_layers = next_u32(&mut pos)?;
|
|
let metric = match next_u32(&mut pos)? {
|
|
0 => DistanceMetric::L2,
|
|
1 => DistanceMetric::Cosine,
|
|
_ => return Err(bad("unknown metric")),
|
|
};
|
|
if n != vectors.len() {
|
|
return Err(bad("vector count does not match the graph"));
|
|
}
|
|
if n == 0 || entry_point >= n || m < 2 || num_layers == 0 || num_layers > 256 {
|
|
return Err(bad("invalid header"));
|
|
}
|
|
let dim = vectors[0].len();
|
|
if vectors.iter().any(|v| v.len() != dim) {
|
|
return Err(bad("vectors have mixed dimensions"));
|
|
}
|
|
|
|
let levels = body.get(pos..pos + n).ok_or_else(|| bad("truncated"))?;
|
|
pos += n;
|
|
let node_levels: Vec<usize> = levels.iter().map(|&l| l as usize).collect();
|
|
if node_levels.iter().any(|&l| l >= num_layers)
|
|
|| node_levels[entry_point] + 1 != num_layers
|
|
{
|
|
return Err(bad("levels inconsistent with layer count"));
|
|
}
|
|
let deleted: Vec<bool> = body
|
|
.get(pos..pos + n)
|
|
.ok_or_else(|| bad("truncated"))?
|
|
.iter()
|
|
.map(|&d| d != 0)
|
|
.collect();
|
|
pos += n;
|
|
|
|
let mut graph: Vec<Vec<Vec<usize>>> = Vec::with_capacity(num_layers);
|
|
for layer in 0..num_layers {
|
|
let max_conn = if layer == 0 { m_max0 } else { m };
|
|
let mut adjacency = vec![Vec::new(); n];
|
|
for (node, slot) in adjacency.iter_mut().enumerate() {
|
|
if node_levels[node] < layer {
|
|
continue;
|
|
}
|
|
let count = next_u32(&mut pos)?;
|
|
if count > max_conn {
|
|
return Err(bad("neighbour list exceeds the connection limit"));
|
|
}
|
|
let mut neighbors = Vec::with_capacity(count);
|
|
for _ in 0..count {
|
|
let id = next_u32(&mut pos)?;
|
|
// A neighbour must exist, and exist on this layer.
|
|
if id >= n || node_levels[id] < layer {
|
|
return Err(bad("neighbour id out of range for its layer"));
|
|
}
|
|
neighbors.push(id);
|
|
}
|
|
*slot = neighbors;
|
|
}
|
|
graph.push(adjacency);
|
|
}
|
|
if pos != body.len() {
|
|
return Err(bad("trailing bytes"));
|
|
}
|
|
|
|
Ok(Self {
|
|
vectors: Vectors::from_rows(
|
|
&vectors
|
|
.into_iter()
|
|
.map(|v| prepare(v, metric))
|
|
.collect::<Vec<_>>(),
|
|
storage,
|
|
metric,
|
|
),
|
|
graph,
|
|
deleted,
|
|
entry_point,
|
|
m,
|
|
m_max0,
|
|
ef_construction,
|
|
node_levels,
|
|
metric,
|
|
})
|
|
}
|
|
|
|
/// Returns the number of vectors in the index.
|
|
pub fn len(&self) -> usize {
|
|
self.vectors.len()
|
|
}
|
|
|
|
/// Returns true if the index is empty.
|
|
pub fn is_empty(&self) -> bool {
|
|
self.vectors.len() == 0
|
|
}
|
|
|
|
/// How this index stores its copy of the vectors.
|
|
pub fn storage(&self) -> Storage {
|
|
self.vectors.storage()
|
|
}
|
|
|
|
/// Returns the dimension of vectors in the index.
|
|
pub fn dimension(&self) -> usize {
|
|
self.vectors.dim()
|
|
}
|
|
|
|
/// Returns the number of layers in the graph.
|
|
pub fn num_layers(&self) -> usize {
|
|
self.graph.len()
|
|
}
|
|
|
|
/// Returns the distance metric used by this index.
|
|
pub fn metric(&self) -> DistanceMetric {
|
|
self.metric
|
|
}
|
|
|
|
/// Returns the maximum number of connections at layer 0.
|
|
pub fn m_max0(&self) -> usize {
|
|
self.m_max0
|
|
}
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Internal HNSW algorithms
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Greedy search: find the single closest node to `query` starting from `ep`.
|
|
fn greedy_closest(
|
|
vectors: &Vectors,
|
|
layer: &[Vec<usize>],
|
|
target: &Target<'_>,
|
|
mut ep: usize,
|
|
metric: DistanceMetric,
|
|
) -> usize {
|
|
let mut best_dist = vectors.dist_to(target, ep, metric);
|
|
loop {
|
|
let mut changed = false;
|
|
for &neighbor in &layer[ep] {
|
|
let d = vectors.dist_to(target, neighbor, metric);
|
|
if d < best_dist {
|
|
best_dist = d;
|
|
ep = neighbor;
|
|
changed = true;
|
|
}
|
|
}
|
|
if !changed {
|
|
break;
|
|
}
|
|
}
|
|
ep
|
|
}
|
|
|
|
/// Search a single layer for the ef closest nodes to `query`.
|
|
/// Best-first search of one layer, returning up to `ef` nodes by ascending
|
|
/// distance.
|
|
///
|
|
/// `skip` marks nodes that must not be *returned* (soft-deleted ones). They
|
|
/// are still traversed — a tombstone is a perfectly good waypoint — but they
|
|
/// never occupy one of the `ef` result slots. Filtering them out afterwards
|
|
/// instead meant a query whose neighbourhood had been deleted got back fewer
|
|
/// than `k` results, or none, however many live records were nearby.
|
|
fn search_layer(
|
|
vectors: &Vectors,
|
|
layer: &[Vec<usize>],
|
|
target: &Target<'_>,
|
|
ep: usize,
|
|
ef: usize,
|
|
metric: DistanceMetric,
|
|
skip: Option<&[bool]>,
|
|
) -> Vec<Candidate> {
|
|
let ep_dist = vectors.dist_to(target, ep, metric);
|
|
|
|
// Min-heap of candidates to explore
|
|
let mut candidates = BinaryHeap::new();
|
|
candidates.push(Candidate {
|
|
id: ep,
|
|
distance: ep_dist,
|
|
});
|
|
|
|
// Max-heap of current results (furthest first)
|
|
let mut results = BinaryHeap::new();
|
|
if !skip.is_some_and(|s| s[ep]) {
|
|
results.push(FarCandidate {
|
|
id: ep,
|
|
distance: ep_dist,
|
|
});
|
|
}
|
|
|
|
VISITED.with_borrow_mut(|visited| {
|
|
visited.begin(vectors.len());
|
|
visited.insert(ep);
|
|
search_layer_visit(
|
|
vectors, layer, target, ef, metric, skip, visited, candidates, results,
|
|
)
|
|
})
|
|
}
|
|
|
|
/// Which nodes a layer search has already seen. A `HashSet` allocated per call
|
|
/// was the hottest non-arithmetic cost in both build and query; this is one
|
|
/// `u32` stamp per node, reused across calls: a node is visited iff its stamp
|
|
/// equals the current epoch, so "clearing" is just bumping the epoch.
|
|
#[derive(Default)]
|
|
struct Visited {
|
|
stamps: Vec<u32>,
|
|
epoch: u32,
|
|
}
|
|
|
|
impl Visited {
|
|
fn begin(&mut self, n: usize) {
|
|
if self.stamps.len() < n {
|
|
self.stamps.resize(n, 0);
|
|
}
|
|
self.epoch = self.epoch.wrapping_add(1);
|
|
if self.epoch == 0 {
|
|
// Wrapped: stale stamps could collide with the new epoch.
|
|
self.stamps.iter_mut().for_each(|s| *s = 0);
|
|
self.epoch = 1;
|
|
}
|
|
}
|
|
|
|
/// Mark `id` visited; `true` if it was not already.
|
|
fn insert(&mut self, id: usize) -> bool {
|
|
let seen = self.stamps[id] == self.epoch;
|
|
self.stamps[id] = self.epoch;
|
|
!seen
|
|
}
|
|
}
|
|
|
|
thread_local! {
|
|
/// Per-thread scratch, so `search(&self)` stays shareable across threads.
|
|
static VISITED: std::cell::RefCell<Visited> = std::cell::RefCell::new(Visited::default());
|
|
}
|
|
|
|
#[allow(clippy::too_many_arguments)]
|
|
fn search_layer_visit(
|
|
vectors: &Vectors,
|
|
layer: &[Vec<usize>],
|
|
target: &Target<'_>,
|
|
ef: usize,
|
|
metric: DistanceMetric,
|
|
skip: Option<&[bool]>,
|
|
visited: &mut Visited,
|
|
mut candidates: BinaryHeap<Candidate>,
|
|
mut results: BinaryHeap<FarCandidate>,
|
|
) -> Vec<Candidate> {
|
|
while let Some(closest) = candidates.pop() {
|
|
let furthest_dist = results.peek().map_or(f32::MAX, |f| f.distance);
|
|
if closest.distance > furthest_dist && results.len() >= ef {
|
|
break;
|
|
}
|
|
|
|
for &neighbor in &layer[closest.id] {
|
|
if !visited.insert(neighbor) {
|
|
continue;
|
|
}
|
|
|
|
let d = vectors.dist_to(target, neighbor, metric);
|
|
let furthest_dist = results.peek().map_or(f32::MAX, |f| f.distance);
|
|
|
|
if d < furthest_dist || results.len() < ef {
|
|
candidates.push(Candidate {
|
|
id: neighbor,
|
|
distance: d,
|
|
});
|
|
if skip.is_some_and(|s| s[neighbor]) {
|
|
continue; // explore through it, but never return it
|
|
}
|
|
results.push(FarCandidate {
|
|
id: neighbor,
|
|
distance: d,
|
|
});
|
|
if results.len() > ef {
|
|
results.pop();
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// Convert to sorted vec
|
|
let mut result: Vec<Candidate> = results
|
|
.into_iter()
|
|
.map(|f| Candidate {
|
|
id: f.id,
|
|
distance: f.distance,
|
|
})
|
|
.collect();
|
|
result.sort_by(|a, b| {
|
|
a.distance
|
|
.partial_cmp(&b.distance)
|
|
.unwrap_or(std::cmp::Ordering::Equal)
|
|
});
|
|
result
|
|
}
|
|
|
|
/// Choose up to `max_conn` neighbours for a node from `candidates` (sorted by
|
|
/// ascending distance to that node) — the HNSW paper's Algorithm 4 with
|
|
/// `keepPrunedConnections`.
|
|
///
|
|
/// Taking the plain `max_conn` closest is what breaks the graph on clustered
|
|
/// data: every link of a node inside a tight cluster goes to that same cluster,
|
|
/// so clusters become islands that a search entering elsewhere can never
|
|
/// reach, however large `ef` is. Instead a candidate is accepted only if it is
|
|
/// closer to the node than to every neighbour already accepted, which spreads
|
|
/// links across directions and keeps the long edges that join clusters. Any
|
|
/// remaining slots are then filled with the closest rejected candidates, so a
|
|
/// node is never left under-connected.
|
|
fn select_neighbors(
|
|
vectors: &Vectors,
|
|
candidates: &[(usize, f32)],
|
|
max_conn: usize,
|
|
metric: DistanceMetric,
|
|
) -> Vec<usize> {
|
|
if candidates.len() <= max_conn {
|
|
return candidates.iter().map(|&(id, _)| id).collect();
|
|
}
|
|
let mut selected: Vec<usize> = Vec::with_capacity(max_conn);
|
|
let mut rejected: Vec<usize> = Vec::new();
|
|
for &(id, dist_to_node) in candidates {
|
|
if selected.len() >= max_conn {
|
|
break;
|
|
}
|
|
let diverse = selected
|
|
.iter()
|
|
.all(|&s| vectors.dist(id, s, metric) > dist_to_node);
|
|
if diverse {
|
|
selected.push(id);
|
|
} else {
|
|
rejected.push(id);
|
|
}
|
|
}
|
|
for id in rejected {
|
|
if selected.len() >= max_conn {
|
|
break;
|
|
}
|
|
selected.push(id);
|
|
}
|
|
selected
|
|
}
|
|
|
|
/// How many nodes to plan together once `linked` nodes are in the graph.
|
|
fn batch_len(linked: usize) -> usize {
|
|
(linked / 16).clamp(1, 512)
|
|
}
|
|
|
|
/// For each node in `batch`: the neighbours to link it to on each of its
|
|
/// layers, found by searching the graph as it currently stands.
|
|
#[allow(clippy::too_many_arguments)]
|
|
fn plan_batch(
|
|
vectors: &Vectors,
|
|
graph: &[Vec<Vec<usize>>],
|
|
node_levels: &[usize],
|
|
batch: std::ops::Range<usize>,
|
|
entry_point: usize,
|
|
ep_level: usize,
|
|
(m, m_max0, ef_construction): (usize, usize, usize),
|
|
metric: DistanceMetric,
|
|
) -> Vec<Vec<(usize, Vec<usize>)>> {
|
|
let plan_one = |i: usize| -> Vec<(usize, Vec<usize>)> {
|
|
let node_level = node_levels[i];
|
|
let mut ep = entry_point;
|
|
// Phase 1: greedy descent from the top layer down to node_level + 1.
|
|
for layer in (node_level + 1..=ep_level).rev() {
|
|
ep = greedy_closest(vectors, &graph[layer], &Target::Node(i), ep, metric);
|
|
}
|
|
// Phase 2: search and select on every layer the node lives on.
|
|
let mut plan = Vec::with_capacity(node_level.min(ep_level) + 1);
|
|
for layer in (0..=node_level.min(ep_level)).rev() {
|
|
let max_conn = if layer == 0 { m_max0 } else { m };
|
|
let neighbors = search_layer(
|
|
vectors,
|
|
&graph[layer],
|
|
&Target::Node(i),
|
|
ep,
|
|
ef_construction,
|
|
metric,
|
|
None,
|
|
);
|
|
let scored: Vec<(usize, f32)> = neighbors.iter().map(|c| (c.id, c.distance)).collect();
|
|
let selected = select_neighbors(vectors, &scored, max_conn, metric);
|
|
if let Some(&closest) = selected.first() {
|
|
ep = closest;
|
|
}
|
|
plan.push((layer, selected));
|
|
}
|
|
plan
|
|
};
|
|
|
|
#[cfg(feature = "parallel")]
|
|
if batch.len() >= PARALLEL_MIN {
|
|
use rayon::prelude::*;
|
|
return batch.into_par_iter().map(plan_one).collect();
|
|
}
|
|
batch.map(plan_one).collect()
|
|
}
|
|
|
|
/// Prune every `(layer, node)` neighbour list in `overflowed` back to its
|
|
/// limit. Each list belongs to a different node, so they are independent.
|
|
fn prune_overflowed(
|
|
vectors: &Vectors,
|
|
graph: &mut [Vec<Vec<usize>>],
|
|
overflowed: Vec<(usize, usize)>,
|
|
(m, m_max0): (usize, usize),
|
|
metric: DistanceMetric,
|
|
) {
|
|
let limit = |layer: usize| if layer == 0 { m_max0 } else { m };
|
|
|
|
#[cfg(feature = "parallel")]
|
|
if overflowed.len() >= PARALLEL_MIN {
|
|
use rayon::prelude::*;
|
|
let mut work: Vec<(usize, usize, Vec<usize>)> = overflowed
|
|
.into_iter()
|
|
.map(|(layer, node)| (layer, node, std::mem::take(&mut graph[layer][node])))
|
|
.collect();
|
|
work.par_iter_mut().for_each(|(layer, node, list)| {
|
|
prune_connections(vectors, list, *node, limit(*layer), metric);
|
|
});
|
|
for (layer, node, list) in work {
|
|
graph[layer][node] = list;
|
|
}
|
|
return;
|
|
}
|
|
|
|
for (layer, node) in overflowed {
|
|
prune_connections(vectors, &mut graph[layer][node], node, limit(layer), metric);
|
|
}
|
|
}
|
|
|
|
/// Fewest independent tasks worth handing to the thread pool.
|
|
#[cfg(feature = "parallel")]
|
|
const PARALLEL_MIN: usize = 8;
|
|
|
|
/// Add the back-link `neighbor -> new_id` for every selected neighbour, pruning
|
|
/// each list that overflows.
|
|
///
|
|
/// Used by incremental [`HnswIndex::insert`]. (A single insert's handful of
|
|
/// prunes is too fine-grained to parallelise profitably — measured 1.45x on 16
|
|
/// cores; bulk builds batch their pruning instead, see `prune_overflowed`.)
|
|
fn link_back(
|
|
vectors: &Vectors,
|
|
layer: &mut [Vec<usize>],
|
|
new_id: usize,
|
|
selected: &[usize],
|
|
max_conn: usize,
|
|
metric: DistanceMetric,
|
|
) {
|
|
let mut overflowed: Vec<usize> = Vec::new();
|
|
for &neighbor in selected {
|
|
layer[neighbor].push(new_id);
|
|
if layer[neighbor].len() > max_conn {
|
|
overflowed.push(neighbor);
|
|
}
|
|
}
|
|
|
|
for node in overflowed {
|
|
prune_connections(vectors, &mut layer[node], node, max_conn, metric);
|
|
}
|
|
}
|
|
|
|
/// Trim `node`'s neighbour list back to `max_conn` with [`select_neighbors`].
|
|
fn prune_connections(
|
|
vectors: &Vectors,
|
|
neighbors: &mut Vec<usize>,
|
|
node: usize,
|
|
max_conn: usize,
|
|
metric: DistanceMetric,
|
|
) {
|
|
if neighbors.len() <= max_conn {
|
|
return;
|
|
}
|
|
let mut scored: Vec<(usize, f32)> = neighbors
|
|
.iter()
|
|
.map(|&n| (n, vectors.dist(node, n, metric)))
|
|
.collect();
|
|
scored.sort_by(|a, b| a.1.total_cmp(&b.1).then(a.0.cmp(&b.0)));
|
|
*neighbors = select_neighbors(vectors, &scored, max_conn, metric);
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// HDF5 reading helpers
|
|
// ---------------------------------------------------------------------------
|
|
|
|
fn read_dataset_raw(data: &[u8], sb: &Superblock, path: &str) -> Result<Vec<u8>, FormatError> {
|
|
let addr = resolve_path_any(data, sb, path)?;
|
|
let header = ObjectHeader::parse(data, addr as usize, sb.offset_size, sb.length_size)?;
|
|
|
|
let dt_msg = header
|
|
.messages
|
|
.iter()
|
|
.find(|m| m.msg_type == MessageType::Datatype)
|
|
.ok_or(FormatError::DatasetMissingData)?;
|
|
let (datatype, _) = Datatype::parse(&dt_msg.data)?;
|
|
|
|
let ds_msg = header
|
|
.messages
|
|
.iter()
|
|
.find(|m| m.msg_type == MessageType::Dataspace)
|
|
.ok_or(FormatError::DatasetMissingShape)?;
|
|
let dataspace = Dataspace::parse(&ds_msg.data, sb.length_size)?;
|
|
|
|
let dl_msg = header
|
|
.messages
|
|
.iter()
|
|
.find(|m| m.msg_type == MessageType::DataLayout)
|
|
.ok_or(FormatError::DatasetMissingData)?;
|
|
let layout = DataLayout::parse(&dl_msg.data, sb.offset_size, sb.length_size)?;
|
|
|
|
let pipeline = header
|
|
.messages
|
|
.iter()
|
|
.find(|m| m.msg_type == MessageType::FilterPipeline)
|
|
.and_then(|msg| FilterPipeline::parse(&msg.data).ok());
|
|
|
|
read_raw_data_full(
|
|
data,
|
|
&layout,
|
|
&dataspace,
|
|
&datatype,
|
|
pipeline.as_ref(),
|
|
sb.offset_size,
|
|
sb.length_size,
|
|
)
|
|
}
|
|
|
|
fn read_dataset_datatype(
|
|
data: &[u8],
|
|
sb: &Superblock,
|
|
path: &str,
|
|
) -> Result<Datatype, FormatError> {
|
|
let addr = resolve_path_any(data, sb, path)?;
|
|
let header = ObjectHeader::parse(data, addr as usize, sb.offset_size, sb.length_size)?;
|
|
let dt_msg = header
|
|
.messages
|
|
.iter()
|
|
.find(|m| m.msg_type == MessageType::Datatype)
|
|
.ok_or(FormatError::DatasetMissingData)?;
|
|
let (datatype, _) = Datatype::parse(&dt_msg.data)?;
|
|
Ok(datatype)
|
|
}
|
|
|
|
fn read_dataset_attrs(
|
|
data: &[u8],
|
|
sb: &Superblock,
|
|
path: &str,
|
|
) -> Result<Vec<(String, AttrValue)>, FormatError> {
|
|
let addr = resolve_path_any(data, sb, path)?;
|
|
let header = ObjectHeader::parse(data, addr as usize, sb.offset_size, sb.length_size)?;
|
|
let attr_msgs = extract_attributes_full(data, &header, sb.offset_size, sb.length_size)?;
|
|
|
|
let mut result = Vec::new();
|
|
for attr in &attr_msgs {
|
|
let name = attr.name.clone();
|
|
if let Some(val) = decode_simple_attr(attr) {
|
|
result.push((name, val));
|
|
}
|
|
}
|
|
Ok(result)
|
|
}
|
|
|
|
fn decode_simple_attr(attr: &clawhdf5_format::attribute::AttributeMessage) -> Option<AttrValue> {
|
|
let raw = &attr.raw_data;
|
|
match &attr.datatype {
|
|
Datatype::FixedPoint { size, signed, .. } => {
|
|
if *signed {
|
|
match size {
|
|
8 => {
|
|
if raw.len() >= 8 {
|
|
let val = i64::from_le_bytes(raw[..8].try_into().ok()?);
|
|
Some(AttrValue::I64(val))
|
|
} else {
|
|
None
|
|
}
|
|
}
|
|
4 => {
|
|
if raw.len() >= 4 {
|
|
let val = i32::from_le_bytes(raw[..4].try_into().ok()?) as i64;
|
|
Some(AttrValue::I64(val))
|
|
} else {
|
|
None
|
|
}
|
|
}
|
|
_ => None,
|
|
}
|
|
} else if raw.len() >= *size as usize {
|
|
let val = u64::from_le_bytes({
|
|
let mut buf = [0u8; 8];
|
|
buf[..*size as usize].copy_from_slice(&raw[..*size as usize]);
|
|
buf
|
|
});
|
|
Some(AttrValue::U64(val))
|
|
} else {
|
|
None
|
|
}
|
|
}
|
|
Datatype::FloatingPoint { size: 8, .. } => {
|
|
if raw.len() >= 8 {
|
|
let val = f64::from_le_bytes(raw[..8].try_into().ok()?);
|
|
Some(AttrValue::F64(val))
|
|
} else {
|
|
None
|
|
}
|
|
}
|
|
Datatype::String { size, .. } => {
|
|
let s = std::str::from_utf8(&raw[..*size as usize]).ok()?;
|
|
Some(AttrValue::String(s.trim_end_matches('\0').to_string()))
|
|
}
|
|
_ => None,
|
|
}
|
|
}
|
|
|
|
fn get_attr_i64(attrs: &[(String, AttrValue)], name: &str) -> Result<i64, FormatError> {
|
|
for (n, v) in attrs {
|
|
if n == name {
|
|
return match v {
|
|
AttrValue::I64(val) => Ok(*val),
|
|
AttrValue::U64(val) => Ok(*val as i64),
|
|
_ => Err(FormatError::SerializationError(format!(
|
|
"attribute {name} is not an integer"
|
|
))),
|
|
};
|
|
}
|
|
}
|
|
Err(FormatError::SerializationError(format!(
|
|
"missing attribute: {name}"
|
|
)))
|
|
}
|
|
|
|
/// Like [`get_attr_i64`] but returns `None` when the attribute is absent or not
|
|
/// an integer, instead of erroring. Used for optional/back-compat attributes.
|
|
fn get_attr_i64_opt(attrs: &[(String, AttrValue)], name: &str) -> Option<i64> {
|
|
attrs
|
|
.iter()
|
|
.find(|(n, _)| n == name)
|
|
.and_then(|(_, v)| match v {
|
|
AttrValue::I64(val) => Some(*val),
|
|
AttrValue::U64(val) => Some(*val as i64),
|
|
_ => None,
|
|
})
|
|
}
|
|
|
|
fn get_attr_string(attrs: &[(String, AttrValue)], name: &str) -> Result<String, FormatError> {
|
|
for (n, v) in attrs {
|
|
if n == name {
|
|
return match v {
|
|
AttrValue::String(s) => Ok(s.clone()),
|
|
_ => Err(FormatError::SerializationError(format!(
|
|
"attribute {name} is not a string"
|
|
))),
|
|
};
|
|
}
|
|
}
|
|
Err(FormatError::SerializationError(format!(
|
|
"missing attribute: {name}"
|
|
)))
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Tests
|
|
// ---------------------------------------------------------------------------
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use std::collections::HashSet;
|
|
|
|
/// Tight, well-separated clusters — the shape real embeddings have, and
|
|
/// the case plain closest-M neighbour selection fails on: each cluster
|
|
/// becomes an island, so recall is capped no matter how large `ef` is.
|
|
fn clustered(n: usize, dim: usize, clusters: usize, seed: u64) -> Vec<Vec<f32>> {
|
|
let mut state = seed;
|
|
let mut next = move || {
|
|
state = state.wrapping_add(0x9E37_79B9_7F4A_7C15);
|
|
let mut z = state;
|
|
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
|
|
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
|
|
((z ^ (z >> 31)) >> 40) as f32 / (1u64 << 24) as f32 - 0.5
|
|
};
|
|
let centres: Vec<Vec<f32>> = (0..clusters)
|
|
.map(|_| (0..dim).map(|_| next() * 10.0).collect())
|
|
.collect();
|
|
(0..n)
|
|
.map(|i| {
|
|
centres[i % clusters]
|
|
.iter()
|
|
.map(|c| c + next() * 0.5)
|
|
.collect()
|
|
})
|
|
.collect()
|
|
}
|
|
|
|
fn recall_at_10(
|
|
index: &HnswIndex,
|
|
vectors: &[Vec<f32>],
|
|
queries: &[Vec<f32>],
|
|
ef: usize,
|
|
) -> f64 {
|
|
let mut hits = 0;
|
|
for q in queries {
|
|
let mut exact: Vec<(usize, f32)> = vectors
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, v)| (i, compute_distance(q, v, DistanceMetric::L2)))
|
|
.collect();
|
|
exact.sort_by(|a, b| a.1.total_cmp(&b.1));
|
|
let want: Vec<usize> = exact[..10].iter().map(|e| e.0).collect();
|
|
hits += index
|
|
.search(q, 10, ef)
|
|
.iter()
|
|
.filter(|(id, _)| want.contains(id))
|
|
.count();
|
|
}
|
|
hits as f64 / (10 * queries.len()) as f64
|
|
}
|
|
|
|
#[test]
|
|
fn clustered_data_keeps_high_recall() {
|
|
// Data and queries come from the same clusters: one draw, split.
|
|
let mut vectors = clustered(3060, 24, 30, 1);
|
|
let queries = vectors.split_off(3000);
|
|
let built = HnswIndex::build_with_metric(&vectors, 8, 40, DistanceMetric::L2);
|
|
let recall = recall_at_10(&built, &vectors, &queries, 64);
|
|
assert!(recall >= 0.95, "bulk build recall@10 = {recall}");
|
|
|
|
// Incremental inserts go through the same neighbour selection.
|
|
let mut incremental = HnswIndex::new(8, 40, DistanceMetric::L2);
|
|
for v in &vectors {
|
|
incremental.insert(v.clone());
|
|
}
|
|
let recall = recall_at_10(&incremental, &vectors, &queries, 64);
|
|
assert!(recall >= 0.95, "incremental recall@10 = {recall}");
|
|
}
|
|
|
|
#[test]
|
|
fn int8_storage_needs_an_exact_re_score_to_match_f32() {
|
|
// Cosine only: rows are unit-length, so a quantised dot product
|
|
// reconstructs the similarity directly.
|
|
//
|
|
// Not the `clustered` generator: its clusters are far tighter than any
|
|
// real embedding, so neighbours sit closer together than the
|
|
// quantisation error and top-10 identity there is noise — that would
|
|
// measure the fixture, not the storage.
|
|
let mut vectors = make_random_vectors(3060, 128, 5);
|
|
let queries = vectors.split_off(3000);
|
|
let f32_index =
|
|
HnswIndex::build_with(&vectors, 8, 40, DistanceMetric::Cosine, Storage::Float32);
|
|
let quantised =
|
|
HnswIndex::build_with(&vectors, 8, 40, DistanceMetric::Cosine, Storage::Int8);
|
|
assert_eq!(quantised.storage(), Storage::Int8);
|
|
|
|
// Ground truth, not the f32 index's answers: re-scoring can beat that
|
|
// index, and measuring against it would score being right as drift.
|
|
let truth: Vec<Vec<usize>> = queries
|
|
.iter()
|
|
.map(|q| {
|
|
let mut d: Vec<(usize, f32)> = vectors
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, v)| (i, compute_distance(q, v, DistanceMetric::Cosine)))
|
|
.collect();
|
|
d.sort_by(|a, b| a.1.total_cmp(&b.1));
|
|
d[..10].iter().map(|x| x.0).collect()
|
|
})
|
|
.collect();
|
|
|
|
let recall = |got: &dyn Fn(&[f32]) -> Vec<usize>| -> f64 {
|
|
let mut hits = 0;
|
|
for (q, want) in queries.iter().zip(&truth) {
|
|
hits += got(q).iter().filter(|id| want.contains(id)).count();
|
|
}
|
|
hits as f64 / (10 * queries.len()) as f64
|
|
};
|
|
|
|
let exact_recall = recall(&|q| f32_index.search(q, 10, 64).iter().map(|r| r.0).collect());
|
|
let raw_recall = recall(&|q| quantised.search(q, 10, 64).iter().map(|r| r.0).collect());
|
|
// Quantised distances alone cost recall, and `ef` cannot buy it back:
|
|
// the loss is in the distances, not in the graph.
|
|
assert!(
|
|
raw_recall < exact_recall,
|
|
"int8 alone should cost recall: {raw_recall} vs {exact_recall}"
|
|
);
|
|
|
|
// Re-scoring a wider candidate pool against the exact vectors — what a
|
|
// caller holding them (the agent's embedding cache) does — puts it
|
|
// back, because only the *ordering* was approximate.
|
|
let rescored_recall = recall(&|q| {
|
|
let mut pool: Vec<(usize, f32)> = quantised
|
|
.search(q, 40, 64)
|
|
.into_iter()
|
|
.map(|(id, _)| {
|
|
(
|
|
id,
|
|
compute_distance(q, &vectors[id], DistanceMetric::Cosine),
|
|
)
|
|
})
|
|
.collect();
|
|
pool.sort_by(|a, b| a.1.total_cmp(&b.1));
|
|
pool.truncate(10);
|
|
pool.into_iter().map(|p| p.0).collect()
|
|
});
|
|
assert!(
|
|
rescored_recall >= exact_recall - 0.01,
|
|
"int8 + exact re-score should match f32: {rescored_recall} vs {exact_recall} (raw {raw_recall})"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn int8_storage_falls_back_to_f32_for_non_cosine_metrics() {
|
|
// L2 distance is not recoverable from a quantised dot product, so the
|
|
// store silently stays f32 rather than returning wrong distances.
|
|
let vectors = clustered(100, 8, 5, 3);
|
|
let index = HnswIndex::build_with(&vectors, 8, 40, DistanceMetric::L2, Storage::Int8);
|
|
assert_eq!(index.storage(), Storage::Float32);
|
|
}
|
|
|
|
#[test]
|
|
fn deletions_near_the_query_do_not_shrink_or_degrade_results() {
|
|
let mut vectors = clustered(2040, 16, 20, 11);
|
|
let queries = vectors.split_off(2000);
|
|
let mut index = HnswIndex::build_with_metric(&vectors, 8, 40, DistanceMetric::L2);
|
|
|
|
let mut short = 0;
|
|
let mut hits = 0;
|
|
for q in &queries {
|
|
// Delete this query's 40 nearest neighbours: more than ef, so every
|
|
// candidate a plain search collects is a tombstone.
|
|
let mut exact: Vec<(usize, f32)> = vectors
|
|
.iter()
|
|
.enumerate()
|
|
.filter(|(i, _)| !index.is_deleted(*i))
|
|
.map(|(i, v)| (i, compute_distance(q, v, DistanceMetric::L2)))
|
|
.collect();
|
|
exact.sort_by(|a, b| a.1.total_cmp(&b.1));
|
|
for &(id, _) in &exact[..40] {
|
|
index.mark_deleted(id);
|
|
}
|
|
let want: Vec<usize> = exact[40..50].iter().map(|e| e.0).collect();
|
|
|
|
let got = index.search(q, 10, 32);
|
|
assert!(got.iter().all(|(id, _)| !index.is_deleted(*id)));
|
|
short += usize::from(got.len() < 10);
|
|
hits += got.iter().filter(|(id, _)| want.contains(id)).count();
|
|
}
|
|
assert_eq!(short, 0, "searches returned fewer than k live results");
|
|
let recall = hits as f64 / (10 * queries.len()) as f64;
|
|
assert!(recall >= 0.9, "recall@10 among live records = {recall}");
|
|
}
|
|
|
|
#[test]
|
|
fn bulk_build_is_deterministic() {
|
|
// Batched planning runs on a thread pool with the `parallel` feature;
|
|
// the graph must not depend on scheduling. (It is also the same graph
|
|
// with and without the feature: both take this exact code path.)
|
|
let vectors = clustered(2500, 16, 20, 21);
|
|
let a = HnswIndex::build_with_metric(&vectors, 8, 40, DistanceMetric::Cosine);
|
|
let b = HnswIndex::build_with_metric(&vectors, 8, 40, DistanceMetric::Cosine);
|
|
assert_eq!(a.graph_to_bytes(), b.graph_to_bytes());
|
|
}
|
|
|
|
#[test]
|
|
fn batches_stay_a_small_fraction_of_the_graph() {
|
|
assert_eq!(batch_len(1), 1);
|
|
assert_eq!(batch_len(15), 1);
|
|
assert_eq!(batch_len(160), 10);
|
|
assert_eq!(batch_len(1_000_000), 512);
|
|
}
|
|
|
|
#[test]
|
|
fn graph_bytes_round_trip_gives_identical_searches() {
|
|
let mut vectors = clustered(1260, 16, 12, 9);
|
|
let queries = vectors.split_off(1200);
|
|
let mut index = HnswIndex::build_with_metric(&vectors, 8, 40, DistanceMetric::L2);
|
|
index.mark_deleted(3);
|
|
index.mark_deleted(700);
|
|
|
|
let bytes = index.graph_to_bytes();
|
|
// The graph is a small fraction of the vectors it indexes... not
|
|
// necessarily at dim 16, but it must not embed them.
|
|
assert!(bytes.len() < 1200 * (16 * 2 + 2) * 4);
|
|
let restored = HnswIndex::from_graph_bytes(&bytes, vectors.clone()).unwrap();
|
|
assert_eq!(restored.deleted_count(), 2);
|
|
for q in &queries {
|
|
assert_eq!(restored.search(q, 10, 50), index.search(q, 10, 50));
|
|
}
|
|
// A restored index keeps working incrementally.
|
|
let mut restored = restored;
|
|
let id = restored.insert(queries[0].clone());
|
|
assert_eq!(restored.search(&queries[0], 1, 50)[0].0, id);
|
|
}
|
|
|
|
#[test]
|
|
fn damaged_or_mismatched_graph_bytes_are_errors() {
|
|
let vectors = clustered(300, 8, 6, 4);
|
|
let index = HnswIndex::build_with_metric(&vectors, 6, 30, DistanceMetric::Cosine);
|
|
let bytes = index.graph_to_bytes();
|
|
|
|
// Wrong vector set.
|
|
assert!(HnswIndex::from_graph_bytes(&bytes, vectors[..299].to_vec()).is_err());
|
|
// Every truncation.
|
|
for len in 0..bytes.len() {
|
|
assert!(
|
|
HnswIndex::from_graph_bytes(&bytes[..len], vectors.clone()).is_err(),
|
|
"truncated to {len}"
|
|
);
|
|
}
|
|
// A flipped bit anywhere.
|
|
for i in (0..bytes.len()).step_by(7) {
|
|
let mut damaged = bytes.clone();
|
|
damaged[i] ^= 0x10;
|
|
assert!(
|
|
HnswIndex::from_graph_bytes(&damaged, vectors.clone()).is_err(),
|
|
"bit flip at {i}"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn structurally_invalid_graph_with_a_valid_checksum_is_rejected() {
|
|
// The CRC only proves the bytes are what was written; a hostile or
|
|
// buggy writer can checksum nonsense. Out-of-range neighbour ids must
|
|
// still be caught, or search would index out of bounds.
|
|
let vectors = clustered(50, 4, 3, 5);
|
|
let index = HnswIndex::build_with_metric(&vectors, 4, 20, DistanceMetric::L2);
|
|
let mut bytes = index.graph_to_bytes();
|
|
let body_len = bytes.len() - 4;
|
|
// First neighbour id of node 0 on layer 0 sits right after the header,
|
|
// levels, tombstones and node 0's count.
|
|
let at = 4 + 7 * 4 + 50 + 50 + 4;
|
|
bytes[at..at + 4].copy_from_slice(&9999u32.to_le_bytes());
|
|
let crc = clawhdf5_format::checksum::crc32(&bytes[..body_len]);
|
|
bytes[body_len..].copy_from_slice(&crc.to_le_bytes());
|
|
assert!(HnswIndex::from_graph_bytes(&bytes, vectors).is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn select_neighbors_prefers_diverse_directions_and_fills_up() {
|
|
// Node at the origin. Three candidates bunched together on the right,
|
|
// one on the left. With room for two, plain closest-M would take two
|
|
// from the bunch and lose the only link leftwards.
|
|
let vectors = vec![
|
|
vec![0.0, 0.0], // 0: the node
|
|
vec![1.0, 0.0], // 1
|
|
vec![1.1, 0.0], // 2
|
|
vec![1.2, 0.0], // 3
|
|
vec![-2.0, 0.0], // 4
|
|
];
|
|
let store = Vectors::from_rows(&vectors, Storage::Float32, DistanceMetric::L2);
|
|
let scored: Vec<(usize, f32)> = (1..5)
|
|
.map(|i| {
|
|
(
|
|
i,
|
|
compute_distance(&vectors[0], &vectors[i], DistanceMetric::L2),
|
|
)
|
|
})
|
|
.collect();
|
|
assert_eq!(
|
|
select_neighbors(&store, &scored, 2, DistanceMetric::L2),
|
|
[1, 4]
|
|
);
|
|
// Spare capacity is filled with the closest rejected candidates.
|
|
assert_eq!(
|
|
select_neighbors(&store, &scored, 3, DistanceMetric::L2),
|
|
[1, 4, 2]
|
|
);
|
|
}
|
|
|
|
fn make_random_vectors(n: usize, dim: usize, seed: u64) -> Vec<Vec<f32>> {
|
|
let mut vectors = Vec::with_capacity(n);
|
|
let mut state = seed;
|
|
for _ in 0..n {
|
|
let mut v = Vec::with_capacity(dim);
|
|
for _ in 0..dim {
|
|
state = splitmix64(state);
|
|
let val = (state >> 40) as f32 / 16777216.0 - 0.5;
|
|
v.push(val);
|
|
}
|
|
vectors.push(v);
|
|
}
|
|
vectors
|
|
}
|
|
|
|
#[test]
|
|
fn build_small_index() {
|
|
let vectors = vec![
|
|
vec![1.0, 0.0, 0.0],
|
|
vec![0.0, 1.0, 0.0],
|
|
vec![0.0, 0.0, 1.0],
|
|
vec![1.0, 1.0, 0.0],
|
|
];
|
|
let index = HnswIndex::build(&vectors, 4, 16);
|
|
assert_eq!(index.len(), 4);
|
|
assert_eq!(index.dimension(), 3);
|
|
assert!(!index.is_empty());
|
|
}
|
|
|
|
#[test]
|
|
fn search_exact_match() {
|
|
let vectors = vec![
|
|
vec![1.0, 0.0],
|
|
vec![0.0, 1.0],
|
|
vec![1.0, 1.0],
|
|
vec![0.5, 0.5],
|
|
];
|
|
let index = HnswIndex::build(&vectors, 4, 16);
|
|
let results = index.search(&[1.0, 0.0], 1, 16);
|
|
assert_eq!(results.len(), 1);
|
|
assert_eq!(results[0].0, 0); // exact match
|
|
assert!(results[0].1 < 1e-6); // distance ~0
|
|
}
|
|
|
|
#[test]
|
|
fn search_k_neighbors() {
|
|
let vectors = make_random_vectors(50, 8, 42);
|
|
let index = HnswIndex::build(&vectors, 8, 32);
|
|
let results = index.search(&vectors[0], 5, 32);
|
|
assert_eq!(results.len(), 5);
|
|
// First result should be the query itself
|
|
assert_eq!(results[0].0, 0);
|
|
assert!(results[0].1 < 1e-6);
|
|
// Results should be sorted by distance
|
|
for i in 1..results.len() {
|
|
assert!(results[i].1 >= results[i - 1].1);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn search_returns_correct_count() {
|
|
let vectors = make_random_vectors(20, 4, 123);
|
|
let index = HnswIndex::build(&vectors, 4, 16);
|
|
let results = index.search(&vectors[5], 3, 16);
|
|
assert_eq!(results.len(), 3);
|
|
}
|
|
|
|
#[test]
|
|
fn cosine_distance() {
|
|
let a = vec![1.0, 0.0];
|
|
let b = vec![0.0, 1.0];
|
|
let d = compute_distance(&a, &b, DistanceMetric::Cosine);
|
|
assert!((d - 1.0).abs() < 1e-6); // orthogonal vectors = cosine distance 1
|
|
|
|
let c = vec![1.0, 0.0];
|
|
let d_same = compute_distance(&a, &c, DistanceMetric::Cosine);
|
|
assert!(d_same < 1e-6); // same direction = cosine distance 0
|
|
}
|
|
|
|
#[test]
|
|
fn l2_distance() {
|
|
let a = vec![0.0, 0.0];
|
|
let b = vec![3.0, 4.0];
|
|
let d = compute_distance(&a, &b, DistanceMetric::L2);
|
|
assert!((d - 5.0).abs() < 1e-6); // 3-4-5 triangle
|
|
}
|
|
|
|
#[test]
|
|
fn cosine_index_build_and_search() {
|
|
let vectors = vec![
|
|
vec![1.0, 0.0, 0.0],
|
|
vec![0.9, 0.1, 0.0],
|
|
vec![0.0, 1.0, 0.0],
|
|
vec![0.0, 0.0, 1.0],
|
|
];
|
|
let index = HnswIndex::build_with_metric(&vectors, 4, 16, DistanceMetric::Cosine);
|
|
let results = index.search(&[1.0, 0.0, 0.0], 2, 16);
|
|
assert_eq!(results.len(), 2);
|
|
// The closest should be vector 0 (exact) or vector 1 (very similar)
|
|
assert!(results[0].0 == 0 || results[0].0 == 1);
|
|
}
|
|
|
|
#[test]
|
|
fn save_and_load_roundtrip() {
|
|
let vectors = make_random_vectors(30, 4, 999);
|
|
let index = HnswIndex::build(&vectors, 4, 16);
|
|
|
|
let bytes = index.to_hdf5_bytes().unwrap();
|
|
assert!(!bytes.is_empty());
|
|
assert_eq!(&bytes[..8], b"\x89HDF\r\n\x1a\n");
|
|
|
|
let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap();
|
|
assert_eq!(loaded.len(), index.len());
|
|
assert_eq!(loaded.dimension(), index.dimension());
|
|
assert_eq!(loaded.metric(), DistanceMetric::L2);
|
|
assert_eq!(loaded.m, index.m);
|
|
assert_eq!(loaded.ef_construction, index.ef_construction);
|
|
assert_eq!(loaded.entry_point, index.entry_point);
|
|
|
|
// Verify vectors match
|
|
for i in 0..loaded.len() {
|
|
assert_eq!(loaded.vectors.row(i), index.vectors.row(i));
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn save_and_load_preserves_search_results() {
|
|
let vectors = make_random_vectors(40, 6, 777);
|
|
let index = HnswIndex::build(&vectors, 6, 24);
|
|
|
|
let query = &vectors[10];
|
|
let original_results = index.search(query, 5, 24);
|
|
|
|
let bytes = index.to_hdf5_bytes().unwrap();
|
|
let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap();
|
|
let loaded_results = loaded.search(query, 5, 24);
|
|
|
|
assert_eq!(original_results.len(), loaded_results.len());
|
|
for (orig, load) in original_results.iter().zip(loaded_results.iter()) {
|
|
assert_eq!(orig.0, load.0);
|
|
assert!((orig.1 - load.1).abs() < 1e-6);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn cosine_roundtrip() {
|
|
let vectors = make_random_vectors(20, 3, 555);
|
|
let index = HnswIndex::build_with_metric(&vectors, 4, 16, DistanceMetric::Cosine);
|
|
|
|
let bytes = index.to_hdf5_bytes().unwrap();
|
|
let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap();
|
|
assert_eq!(loaded.metric(), DistanceMetric::Cosine);
|
|
}
|
|
|
|
#[test]
|
|
fn index_metadata() {
|
|
let vectors = make_random_vectors(10, 5, 111);
|
|
let index = HnswIndex::build(&vectors, 3, 12);
|
|
assert_eq!(index.len(), 10);
|
|
assert_eq!(index.dimension(), 5);
|
|
assert!(index.num_layers() >= 1);
|
|
assert_eq!(index.metric(), DistanceMetric::L2);
|
|
}
|
|
|
|
#[test]
|
|
fn save_to_file_writer() {
|
|
let vectors = make_random_vectors(15, 3, 333);
|
|
let index = HnswIndex::build(&vectors, 4, 16);
|
|
|
|
let dir = std::env::temp_dir();
|
|
let path = dir.join("clawhdf5_ann_test_save.h5");
|
|
let mut writer = IoFileWriter::create(&path).unwrap();
|
|
index.save_to_hdf5(&mut writer).unwrap();
|
|
|
|
// Verify file exists and has HDF5 signature
|
|
let data = std::fs::read(&path).unwrap();
|
|
assert_eq!(&data[..8], b"\x89HDF\r\n\x1a\n");
|
|
|
|
// Load it back
|
|
let loaded = HnswIndex::load_from_hdf5(&data).unwrap();
|
|
assert_eq!(loaded.len(), 15);
|
|
|
|
std::fs::remove_file(&path).ok();
|
|
}
|
|
|
|
#[test]
|
|
fn search_accuracy_l2() {
|
|
// Build a small index and verify that brute-force search agrees
|
|
let vectors = make_random_vectors(100, 8, 42);
|
|
let index = HnswIndex::build(&vectors, 16, 64);
|
|
|
|
let query = &vectors[50];
|
|
let k = 10;
|
|
let results = index.search(query, k, 64);
|
|
|
|
// Brute-force nearest neighbors
|
|
let mut brute: Vec<(usize, f32)> = vectors
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, v)| (i, compute_distance(query, v, DistanceMetric::L2)))
|
|
.collect();
|
|
brute.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
|
|
brute.truncate(k);
|
|
|
|
// With high ef and m, HNSW should find at least 80% of true neighbors
|
|
let hnsw_ids: HashSet<usize> = results.iter().map(|r| r.0).collect();
|
|
let brute_ids: HashSet<usize> = brute.iter().map(|r| r.0).collect();
|
|
let overlap = hnsw_ids.intersection(&brute_ids).count();
|
|
assert!(overlap >= k * 8 / 10, "HNSW recall too low: {overlap}/{k}");
|
|
}
|
|
|
|
#[test]
|
|
fn search_accuracy_cosine() {
|
|
let vectors = make_random_vectors(100, 8, 99);
|
|
let index = HnswIndex::build_with_metric(&vectors, 16, 64, DistanceMetric::Cosine);
|
|
|
|
let query = &vectors[25];
|
|
let k = 10;
|
|
let results = index.search(query, k, 64);
|
|
|
|
let mut brute: Vec<(usize, f32)> = vectors
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, v)| (i, compute_distance(query, v, DistanceMetric::Cosine)))
|
|
.collect();
|
|
brute.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
|
|
brute.truncate(k);
|
|
|
|
let hnsw_ids: HashSet<usize> = results.iter().map(|r| r.0).collect();
|
|
let brute_ids: HashSet<usize> = brute.iter().map(|r| r.0).collect();
|
|
let overlap = hnsw_ids.intersection(&brute_ids).count();
|
|
assert!(
|
|
overlap >= k * 8 / 10,
|
|
"HNSW cosine recall too low: {overlap}/{k}"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn distance_metric_str_roundtrip() {
|
|
assert_eq!(DistanceMetric::from_str("l2"), Some(DistanceMetric::L2));
|
|
assert_eq!(
|
|
DistanceMetric::from_str("cosine"),
|
|
Some(DistanceMetric::Cosine)
|
|
);
|
|
assert_eq!(DistanceMetric::from_str("unknown"), None);
|
|
assert_eq!(DistanceMetric::L2.as_str(), "l2");
|
|
assert_eq!(DistanceMetric::Cosine.as_str(), "cosine");
|
|
}
|
|
|
|
#[test]
|
|
fn cosine_zero_vector() {
|
|
let a = vec![0.0, 0.0];
|
|
let b = vec![1.0, 0.0];
|
|
let d = compute_distance(&a, &b, DistanceMetric::Cosine);
|
|
assert!((d - 1.0).abs() < 1e-6); // zero vector -> distance 1
|
|
}
|
|
|
|
#[test]
|
|
fn cosine_near_zero_vector() {
|
|
// Tiny-but-nonzero, identical-direction vectors: denom is well
|
|
// below f32::EPSILON but not exactly 0.0. Must still be treated
|
|
// as a degenerate/unreliable direction (distance 1, "maximally
|
|
// dissimilar"), not as an exact match (distance 0).
|
|
let a = vec![1e-4, 1e-4];
|
|
let b = vec![1e-4, 1e-4];
|
|
let d = compute_distance(&a, &b, DistanceMetric::Cosine);
|
|
assert!((d - 1.0).abs() < 1e-6);
|
|
}
|
|
|
|
#[test]
|
|
fn insert_into_empty_index() {
|
|
let mut index = HnswIndex::new(4, 16, DistanceMetric::L2);
|
|
assert!(index.is_empty());
|
|
let id = index.insert(vec![1.0, 0.0, 0.0]);
|
|
assert_eq!(id, 0);
|
|
assert_eq!(index.len(), 1);
|
|
let results = index.search(&[1.0, 0.0, 0.0], 1, 16);
|
|
assert_eq!(results, vec![(0, 0.0)]);
|
|
}
|
|
|
|
#[test]
|
|
fn incremental_insert_matches_batch_recall() {
|
|
// Build one index incrementally and one in batch from the same vectors,
|
|
// then confirm the incremental index has acceptable recall vs brute force.
|
|
let vectors = make_random_vectors(120, 8, 2024);
|
|
let mut incremental = HnswIndex::new(16, 64, DistanceMetric::L2);
|
|
for v in &vectors {
|
|
incremental.insert(v.clone());
|
|
}
|
|
assert_eq!(incremental.len(), vectors.len());
|
|
|
|
let query = &vectors[60];
|
|
let k = 10;
|
|
let results = incremental.search(query, k, 64);
|
|
|
|
let mut brute: Vec<(usize, f32)> = vectors
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, v)| (i, compute_distance(query, v, DistanceMetric::L2)))
|
|
.collect();
|
|
brute.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
|
|
brute.truncate(k);
|
|
|
|
let hnsw_ids: HashSet<usize> = results.iter().map(|r| r.0).collect();
|
|
let brute_ids: HashSet<usize> = brute.iter().map(|r| r.0).collect();
|
|
let overlap = hnsw_ids.intersection(&brute_ids).count();
|
|
assert!(
|
|
overlap >= k * 8 / 10,
|
|
"incremental HNSW recall too low: {overlap}/{k}"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn mark_deleted_excludes_from_search() {
|
|
let vectors = vec![
|
|
vec![1.0, 0.0],
|
|
vec![0.9, 0.1],
|
|
vec![0.0, 1.0],
|
|
vec![0.1, 0.9],
|
|
];
|
|
let mut index = HnswIndex::build(&vectors, 4, 16);
|
|
|
|
// Exact match on vector 0 before deletion.
|
|
let before = index.search(&[1.0, 0.0], 1, 16);
|
|
assert_eq!(before[0].0, 0);
|
|
|
|
assert!(index.mark_deleted(0));
|
|
assert!(index.is_deleted(0));
|
|
assert!(!index.mark_deleted(0)); // idempotent
|
|
assert_eq!(index.deleted_count(), 1);
|
|
assert_eq!(index.active_len(), 3);
|
|
|
|
// Vector 0 must no longer be returned; nearest is now vector 1.
|
|
let after = index.search(&[1.0, 0.0], 2, 16);
|
|
assert!(after.iter().all(|(id, _)| *id != 0));
|
|
assert_eq!(after[0].0, 1);
|
|
}
|
|
|
|
#[test]
|
|
fn compact_drops_deleted_and_renumbers() {
|
|
let vectors = make_random_vectors(20, 4, 4242);
|
|
let mut index = HnswIndex::build(&vectors, 8, 32);
|
|
index.mark_deleted(3);
|
|
index.mark_deleted(7);
|
|
index.mark_deleted(11);
|
|
|
|
let mapping = index.compact();
|
|
assert_eq!(mapping.len(), 20);
|
|
assert_eq!(index.len(), 17);
|
|
assert_eq!(index.deleted_count(), 0);
|
|
// Deleted ids map to None; survivors map to a dense 0..17 range.
|
|
assert!(mapping[3].is_none() && mapping[7].is_none() && mapping[11].is_none());
|
|
let mut new_ids: Vec<usize> = mapping.iter().filter_map(|m| *m).collect();
|
|
new_ids.sort_unstable();
|
|
assert_eq!(new_ids, (0..17).collect::<Vec<_>>());
|
|
}
|
|
|
|
#[test]
|
|
fn compact_all_deleted_yields_empty_index() {
|
|
let vectors = make_random_vectors(5, 3, 1);
|
|
let mut index = HnswIndex::build(&vectors, 4, 16);
|
|
for i in 0..5 {
|
|
index.mark_deleted(i);
|
|
}
|
|
index.compact();
|
|
assert!(index.is_empty());
|
|
assert!(index.search(&[0.0, 0.0, 0.0], 3, 16).is_empty());
|
|
}
|
|
|
|
#[test]
|
|
fn versioned_roundtrip_preserves_deletions() {
|
|
let vectors = make_random_vectors(30, 4, 8080);
|
|
let mut index = HnswIndex::build(&vectors, 6, 24);
|
|
index.mark_deleted(5);
|
|
index.mark_deleted(12);
|
|
|
|
let bytes = index.to_hdf5_bytes().unwrap();
|
|
let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap();
|
|
|
|
assert_eq!(loaded.len(), index.len());
|
|
assert!(loaded.is_deleted(5));
|
|
assert!(loaded.is_deleted(12));
|
|
assert_eq!(loaded.deleted_count(), 2);
|
|
// A deleted vector is not returned even after a round-trip.
|
|
let results = loaded.search(&vectors[5], 5, 24);
|
|
assert!(results.iter().all(|(id, _)| *id != 5));
|
|
}
|
|
|
|
#[test]
|
|
fn insert_then_save_load_search() {
|
|
let mut index = HnswIndex::new(8, 32, DistanceMetric::Cosine);
|
|
let vectors = make_random_vectors(25, 5, 31337);
|
|
for v in &vectors {
|
|
index.insert(v.clone());
|
|
}
|
|
let bytes = index.to_hdf5_bytes().unwrap();
|
|
let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap();
|
|
assert_eq!(loaded.len(), 25);
|
|
assert_eq!(loaded.metric(), DistanceMetric::Cosine);
|
|
let results = loaded.search(&vectors[0], 3, 32);
|
|
assert_eq!(results.len(), 3);
|
|
assert_eq!(results[0].0, 0);
|
|
}
|
|
}
|