Skip to main content

kernel/
nav.rs

1//! 2k: the vector navigation tier (FACT-02's graph stage).
2//!
3//! A newcomer's map. The scan tier reads EVERY fingerprint per query --
4//! honest but linear. This tier gives each vector a handful of links to
5//! its most-similar vectors, forming a small-world web (the Vamana family,
6//! as in DiskANN); a query then WALKS: start at a central vector (the
7//! medoid), repeatedly move toward whichever known neighbour looks closest
8//! to the query, keep the best `ef` candidates, stop when no frontier
9//! candidate can beat the worst kept result. Hops ~ logarithmic; a million
10//! vectors answer in a few hundred row reads instead of a million.
11//!
12//! Everything is ordinary rows (D4): (0x11, field, id) -> [norm][code]
13//! [links]. One web per FIELD -- links are bare ids, so a shared keyspace
14//! would wire a 384-dim column's vectors into an 8-dim column's
15//! neighbourhoods and both walks would read the other's rows as garbage.
16//! The 2-bit fingerprint rides IN the nav row, so one read per visited
17//! node yields both topology and ranking data. The exact tier (rescore)
18//! still ranks the survivors: approximation can miss, never misrank.
19//!
20//! Freshness without an LSM: the fold records a WATERMARK as a fast path for
21//! the ordinary increasing-id tail. Vectors written at-or-below it get an
22//! explicit pending marker in the catalog. Queries merge both pending sets
23//! with the walk, and folds consume both, so arbitrary caller ids remain
24//! visible without an O(N) sweep. Deletes leave dangling links that walks
25//! skip (missing row = dead), healed at fold.
26
27use crate::graph::{Graph, Metric, CAT_NAV_PENDING};
28use crate::keys;
29use crate::vecquant::{self, AffineQuery};
30use crate::{Error, Result};
31
32/// Max neighbours kept per node. DiskANN's sweet spot region; the row
33/// stays ~0.8KB at 2-bit/2048-pad codes.
34pub const NAV_R: usize = 32;
35/// Build-time beam width (bigger = better graph, slower fold).
36pub const NAV_L_BUILD: usize = 64;
37/// Pruning slack: a candidate is dropped if some kept neighbour is more
38/// than ALPHA closer to it than the candidate is to the node (Vamana's
39/// robust prune -- keeps links spread out instead of clustered).
40/// Applied as ALPHA^2 because we compare SQUARED L2 -- alpha on squared
41/// distances is sqrt(alpha) on true ones, which quietly weakened the
42/// spread rule to 1.095 (recall 0.535 measured before the fix).
43pub const NAV_ALPHA: f32 = 1.2;
44/// The fold commits AND checkpoints every this-many inserts. A fold that
45/// commits once holds every neighbour-row rewrite of the whole batch in
46/// the WAL (~33 row images per insert: 15GB measured at 1M) -- the WAL
47/// only truncates at a checkpoint, so the bound must checkpoint. Crash
48/// mid-fold is already safe at ANY boundary: the watermark rides each
49/// insert, so a reopened store simply resumes the fold above it.
50pub const NAV_FOLD_CHECKPOINT_EVERY: u64 = 25_000;
51
52/// Backlink lists may grow to this before being pruned back to NAV_R.
53/// Pruning on EVERY overflow re-read ~33 full vectors per neighbour per
54/// insert; letting lists run to 2R amortises that ~32x (Vamana batch
55/// builds do the same). The row decoder already accepts n <= 2R.
56pub const NAV_R_SLACK: usize = NAV_R * 2;
57
58/// Map an f32 to order-preserving u64 bits (positive floats sort by raw
59/// bits with the sign bit set; negatives flip ALL bits). The width
60/// matters: f32::to_bits is u32, and a u64-width flip sets the upper 32
61/// bits, ranking every NEGATIVE estimate above every positive one --
62/// backwards. Negative L2 estimates (norm^2 - 2*dot omits |q|^2, so near
63/// neighbours go negative) are exactly the BEST candidates; the inverted
64/// order made recall FALL as the beam widened (0.495 -> 0.285 at 1M,
65/// measured) because wider beams found, then evicted, more of them.
66pub(crate) fn sortable(d: f32) -> u64 {
67    let b = d.to_bits();
68    (if d >= 0.0 { b | 0x8000_0000 } else { !b }) as u64
69}
70
71struct NavRow {
72    norm: f32,
73    code_off: usize,
74    code_len: usize,
75    neighbors: Vec<u64>,
76}
77
78fn decode_row(v: &[u8], code_len: usize) -> Result<NavRow> {
79    let count_at = 4usize.checked_add(code_len).ok_or(Error::Corrupt {
80        page_no: 0,
81        why: "navigation row code length overflows",
82    })?;
83    let header_end = count_at.checked_add(2).ok_or(Error::Corrupt {
84        page_no: 0,
85        why: "navigation row header length overflows",
86    })?;
87    let count_bytes = v.get(count_at..header_end).ok_or(Error::Corrupt {
88        page_no: 0,
89        why: "navigation row is shorter than its code and count",
90    })?;
91    let norm = f32::from_le_bytes(v.get(..4).ok_or(Error::Corrupt {
92        page_no: 0,
93        why: "navigation row has no norm",
94    })?.try_into().unwrap());
95    let n = u16::from_le_bytes(count_bytes.try_into().unwrap()) as usize;
96    if n > NAV_R_SLACK {
97        return Err(Error::Corrupt { page_no: 0, why: "navigation row has too many neighbours" });
98    }
99    let expected = n.checked_mul(8).and_then(|bytes| header_end.checked_add(bytes))
100        .ok_or(Error::Corrupt { page_no: 0, why: "navigation row length overflows" })?;
101    if v.len() != expected {
102        return Err(Error::Corrupt { page_no: 0, why: "navigation row length is invalid" });
103    }
104    let mut neighbors = Vec::with_capacity(n);
105    for i in 0..n {
106        let o = header_end + i * 8;
107        neighbors.push(u64::from_le_bytes(v[o..o + 8].try_into().unwrap()));
108    }
109    Ok(NavRow { norm, code_off: 4, code_len, neighbors })
110}
111
112fn encode_row(norm: f32, code: &[u8], neighbors: &[u64]) -> Vec<u8> {
113    let mut v = Vec::with_capacity(4 + code.len() + 2 + neighbors.len() * 8);
114    v.extend_from_slice(&norm.to_le_bytes());
115    v.extend_from_slice(code);
116    v.extend_from_slice(&(neighbors.len() as u16).to_le_bytes());
117    for n in neighbors { v.extend_from_slice(&n.to_le_bytes()); }
118    v
119}
120
121impl Graph {
122    /// Append ids explicitly known to have been written behind the watermark.
123    /// Empty pending sets cost one catalog seek, never a vector-range sweep.
124    fn nav_pending_ids(&self, field: u64, out: &mut Vec<u64>) -> Result<()> {
125        let prefix = keys::catalog_field(CAT_NAV_PENDING, field);
126        self.store_ref().scan(&prefix)?.for_each_ref(|key, _| {
127            if !key.starts_with(&prefix) { return false; }
128            if key.len() == 25 { out.push(keys::u64_at(key, 17)); }
129            true
130        })
131    }
132
133    fn nav_code_len(&self, field: u64) -> usize {
134        vecquant::code_len(vecquant::pad_dim(self.vec_dim(field) as usize),
135                           self.vec_bits_pub(field))
136    }
137
138    /// Diagnostic walk: like the query path, but also returns every id the
139    /// beam VISITED (estimated), not only the ef it kept. Separates "the
140    /// walk never reached the region" from "reached it but ranked it out".
141    pub fn nav_walk_diag(&self, field: u64, q: &[f32], ef: usize)
142        -> Result<(Vec<u64>, std::collections::HashSet<u64>)>
143    {
144        let medoid = match self.vec_meta(field) {
145            Some(m) if m.medoid != 0 => m.medoid,
146            _ => return Ok((Vec::new(), Default::default())),
147        };
148        let Some(enc) = self.encoder(field) else { return Ok((Vec::new(), Default::default())) };
149        let aq = enc.affine_query(q);
150        let mut visited = std::collections::HashSet::new();
151        let kept = self.nav_beam_inner(field, &aq, Metric::L2, ef, medoid, Some(&mut visited))?;
152        Ok((kept.into_iter().map(|(_, id)| id).collect(), visited))
153    }
154
155    /// Beam walk over the nav rows: returns up to `ef` candidate ids by
156    /// estimated distance. Reads ONE row per visited node.
157    fn nav_beam(&self, field: u64, aq: &AffineQuery, metric: Metric, ef: usize, start: u64)
158        -> Result<Vec<(f32, u64)>>
159    {
160        self.nav_beam_inner(field, aq, metric, ef, start, None)
161    }
162
163    fn nav_beam_inner(&self, field: u64, aq: &AffineQuery, metric: Metric, ef: usize, start: u64,
164                      mut diag: Option<&mut std::collections::HashSet<u64>>)
165        -> Result<Vec<(f32, u64)>>
166    {
167        let code_len = self.nav_code_len(field);
168        let est = |row: &NavRow, v: &[u8]| -> f32 {
169            let dot = vecquant::dot_est_affine2(row.norm, &v[row.code_off..row.code_off + row.code_len], aq);
170            match metric {
171                Metric::L2 | Metric::L1 => row.norm * row.norm - 2.0 * dot,
172                Metric::Dot => -dot,
173                Metric::Cosine => if row.norm > 0.0 { -dot / row.norm } else { 0.0 },
174            }
175        };
176        let mut visited: std::collections::HashSet<u64> = std::collections::HashSet::new();
177        // frontier: nearest-first (Reverse); kept: worst-first, bounded ef
178        let mut frontier2: std::collections::BinaryHeap<std::cmp::Reverse<(u64, u64)>> =
179            std::collections::BinaryHeap::new();
180        let mut kept: std::collections::BinaryHeap<(u64, u64)> = std::collections::BinaryHeap::new();
181        let seed = match self.store_ref().get(&keys::nav_key(field, start))? {
182            Some(v) => v,
183            None => return Ok(Vec::new()),
184        };
185        let row = decode_row(&seed, code_len)?;
186        let d0 = est(&row, &seed);
187        visited.insert(start);
188        frontier2.push(std::cmp::Reverse((sortable(d0), start)));
189        kept.push((sortable(d0), start));
190        let mut dists: std::collections::HashMap<u64, f32> = std::collections::HashMap::new();
191        dists.insert(start, d0);
192
193        while let Some(std::cmp::Reverse((ds, id))) = frontier2.pop() {
194            // stop when the closest frontier item cannot beat the worst kept
195            if kept.len() >= ef {
196                if let Some(&(worst, _)) = kept.peek() {
197                    if ds > worst { break; }
198                }
199            }
200            let Some(v) = self.store_ref().get(&keys::nav_key(field, id))? else { continue };
201            let row = decode_row(&v, code_len)?;
202            for &nb in &row.neighbors {
203                if !visited.insert(nb) { continue; }
204                if let Some(d) = diag.as_deref_mut() { d.insert(nb); }
205                let Some(nv) = self.store_ref().get(&keys::nav_key(field, nb))? else { continue };
206                let nrow = decode_row(&nv, code_len)?;
207                let nd = est(&nrow, &nv);
208                let nds = sortable(nd);
209                let admit = kept.len() < ef || kept.peek().map(|&(w, _)| nds < w).unwrap_or(true);
210                if admit {
211                    kept.push((nds, nb));
212                    if kept.len() > ef { kept.pop(); }
213                    frontier2.push(std::cmp::Reverse((nds, nb)));
214                    dists.insert(nb, nd);
215                }
216            }
217        }
218        let mut out: Vec<(f32, u64)> = kept.into_iter()
219            .map(|(_, id)| (dists.get(&id).copied().unwrap_or(f32::MAX), id))
220            .collect();
221        out.sort_by(|a, b| a.0.total_cmp(&b.0));
222        Ok(out)
223    }
224
225    /// Wire `id` into the graph (Vamana insert): beam-search its
226    /// neighbourhood, robust-prune to NAV_R links, write the row, add
227    /// pruned backlinks. Cost ∝ beam + degree, never ∝ store (Law 2).
228    fn nav_insert(&mut self, field: u64, id: u64, medoid: u64, refresh: bool) -> Result<()> {
229        let code_len = self.nav_code_len(field);
230        let Some(vrow) = self.store_ref().get(&keys::vcode_key(field, id))? else { return Ok(()) };
231        if vrow.len() != 4usize.saturating_add(code_len) {
232            return Err(Error::Corrupt { page_no: 0, why: "vector code row length is invalid" });
233        }
234        let norm = f32::from_le_bytes(vrow[0..4].try_into().unwrap());
235        let code = vrow[4..].to_vec();
236        // the query IS this vector: reuse its exact f32 form for the walk
237        let Some(v) = self.get_vec(field, id)? else { return Ok(()) };
238        let Some(enc) = self.encoder(field) else { return Ok(()) };
239        let aq = enc.affine_query(&v);
240
241        let mut cands = if id == medoid && !refresh {
242            Vec::new()
243        } else {
244            self.nav_beam(field, &aq, Metric::L2, NAV_L_BUILD, medoid)?
245        };
246        // A behind-watermark write may replace a folded vector. Its old row
247        // can be reached by the beam, but it cannot be its own neighbour.
248        cands.retain(|&(_, candidate)| candidate != id);
249        // re-price candidates EXACTLY before pruning (estimates found them;
250        // truth ranks them), all through one per-insert vector cache
251        let mut vcache: std::collections::HashMap<u64, Vec<f32>> =
252            std::collections::HashMap::new();
253        vcache.insert(id, v.clone());
254        for c in cands.iter_mut() {
255            if let Some(d) = self.pair_dist_cached(field, &mut vcache, id, c.1)? { c.0 = d; }
256        }
257        let pruned = self.robust_prune(field, &mut vcache, id, &cands, code_len)?;
258        self.store().put(&keys::nav_key(field, id), &encode_row(norm, &code, &pruned))?;
259        // backlinks, pruned per neighbour
260        for &nb in &pruned {
261            let Some(nv) = self.store_ref().get(&keys::nav_key(field, nb))? else { continue };
262            let mut nrow = decode_row(&nv, code_len)?;
263            if nrow.neighbors.contains(&id) { continue; }
264            nrow.neighbors.push(id);
265            let links = if nrow.neighbors.len() > NAV_R_SLACK {
266                let mut cds = Vec::with_capacity(nrow.neighbors.len());
267                for &x in &nrow.neighbors {
268                    cds.push((self.pair_dist_cached(field, &mut vcache, nb, x)?.unwrap_or(f32::MAX), x));
269                }
270                self.robust_prune(field, &mut vcache, nb, &cds, code_len)?
271            } else { nrow.neighbors.clone() };
272            let ncode = nv[4..4 + code_len].to_vec();
273            self.store().put(&keys::nav_key(field, nb), &encode_row(nrow.norm, &ncode, &links))?;
274        }
275        Ok(())
276    }
277
278    /// EXACT L2^2 between two stored vectors -- build-time only, through a
279    /// per-insert cache (the uncached form re-read vectors O(kept x cands)
280    /// times: 12ms per insert, measured). Estimates find candidates;
281    /// exact distances rank and prune them (the DiskANN posture).
282    fn pair_dist_cached(&self, field: u64, cache: &mut std::collections::HashMap<u64, Vec<f32>>,
283                        a: u64, b: u64) -> Result<Option<f32>> {
284        for id in [a, b] {
285            if !cache.contains_key(&id) {
286                let Some(vector) = self.get_vec(field, id)? else { return Ok(None) };
287                cache.insert(id, vector);
288            }
289        }
290        let (Some(va), Some(vb)) = (cache.get(&a), cache.get(&b)) else { return Ok(None) };
291        Ok(Some(va.iter().zip(vb).map(|(x, y)| (x - y) * (x - y)).sum()))
292    }
293
294    /// Vamana's robust prune: keep the closest candidate, drop any other
295    /// candidate that some kept neighbour dominates (ALPHA * d(kept, cand)
296    /// < d(node, cand)); repeat to NAV_R.
297    fn robust_prune(&self, field: u64, vcache: &mut std::collections::HashMap<u64, Vec<f32>>,
298                    _node: u64, cands: &[(f32, u64)], code_len: usize)
299        -> Result<Vec<u64>>
300    {
301        let mut sorted: Vec<(f32, u64)> = cands.to_vec();
302        sorted.sort_by(|a, b| a.0.total_cmp(&b.0));
303        sorted.dedup_by_key(|c| c.1);
304        let mut kept: Vec<(f32, u64)> = Vec::new();
305        let _ = code_len;
306        'cand: for &(d, c) in &sorted {
307            if kept.len() >= NAV_R { break; }
308            for &(_, k) in &kept {
309                let Some(dkc) = self.pair_dist_cached(field, vcache, k, c)? else { continue };
310                if NAV_ALPHA * NAV_ALPHA * dkc < d { continue 'cand; }
311            }
312            kept.push((d, c));
313        }
314        Ok(kept.into_iter().map(|(_, id)| id).collect())
315    }
316
317    /// Fold: wire the increasing-id tail plus explicit out-of-order pending
318    /// ids into the graph, in id order. Cost ∝ new vectors, not the field.
319    pub fn fold_nav(&mut self, field: u64) -> Result<u64> {
320        let Some(mut meta) = self.vec_meta(field) else { return Ok(0) };
321        let mut stragglers: Vec<u64> = Vec::new();
322        self.nav_pending_ids(field, &mut stragglers)?;
323        stragglers.sort_unstable();
324        stragglers.dedup();
325        let mut pending = stragglers.clone();
326        // The contiguous tail remains the fast path and preserves existing
327        // dock/bulk-loaded stores, which have no per-vector pending markers.
328        if let Some(next) = meta.watermark.checked_add(1) {
329            let from = keys::vcode_key(field, next);
330            let it = self.store_ref().scan(&from)?;
331            it.for_each_ref(|key, _| {
332                if key.len() != 17 || key[0] != keys::TAG_VCODE
333                    || keys::u64_at(key, 1) != field { return false; }
334                pending.push(keys::u64_at(key, 9));
335                true
336            })?;
337        }
338        pending.sort_unstable();
339        pending.dedup();
340        let n = pending.len() as u64;
341        let mut since = 0u64;
342        for id in pending {
343            let m = match meta.medoid {
344                0 => {
345                    meta.medoid = id;
346                    self.set_vec_meta(field, meta)?;
347                    id
348                }
349                m => m,
350            };
351            let refresh = stragglers.binary_search(&id).is_ok();
352            self.nav_insert(field, id, m, refresh)?;
353            if refresh {
354                self.store().delete(
355                    &keys::catalog_field_item(CAT_NAV_PENDING, field, id))?;
356            }
357            meta.watermark = meta.watermark.max(id);
358            self.set_vec_meta(field, meta)?;
359            since += 1;
360            if since >= self.nav_fold_every {
361                self.commit()?;
362                self.checkpoint()?; // truncates the WAL: disk high-water ∝ interval
363                since = 0;
364            }
365        }
366        self.commit()?;
367        Ok(n)
368    }
369
370    /// Graph-accelerated nearest: beam walk over the graph, scan tier for the
371    /// increasing-id head and explicit out-of-order pending set, then exact
372    /// rescore over the union. Falls back to pure scan when no graph exists.
373    pub fn nearest_nav(&self, field: u64, q: &[f32], k: usize, metric: Metric, oversample: usize)
374        -> Result<Vec<(u64, f32)>>
375    {
376        let meta = match self.vec_meta(field) {
377            Some(m) if m.medoid != 0 => m,
378            _ => return self.nearest(field, q, k, metric, oversample),
379        };
380        let Some(enc) = self.encoder(field) else {
381            return self.nearest(field, q, k, metric, oversample);
382        };
383        let aq = enc.affine_query(q);
384        let ef = (k * oversample).max(k).max(64);
385        let mut cands: Vec<u64> = self.nav_beam(field, &aq, metric, ef, meta.medoid)?
386            .into_iter().map(|(_, id)| id).collect();
387        let mut stragglers = Vec::new();
388        self.nav_pending_ids(field, &mut stragglers)?;
389        for id in stragglers {
390            // An overwrite may still be reachable through its old nav row.
391            // Avoid returning the same id twice after exact rescoring.
392            if !cands.contains(&id) { cands.push(id); }
393        }
394        // Head: codes above the watermark, contiguous within this field.
395        if let Some(next) = meta.watermark.checked_add(1) {
396            let head_from = keys::vcode_key(field, next);
397            let it = self.store_ref().scan(&head_from)?;
398            it.for_each_ref(|key, _| {
399                if key.len() != 17 || key[0] != keys::TAG_VCODE
400                    || keys::u64_at(key, 1) != field { return false; }
401                cands.push(keys::u64_at(key, 9));
402                true
403            })?;
404        }
405        self.rescore(field, &cands, q, metric, k)
406    }
407}
408
409#[cfg(test)]
410mod tests {
411    use super::*;
412    use crate::store::{Config, Store};
413
414    /// The beam orders candidates by these keys: more-negative estimate =
415    /// better = smaller key. The u64-width flip bug ranked all negatives
416    /// worst; this pins strict monotonicity across the sign boundary.
417    #[test]
418    fn sortable_orders_floats_across_the_sign_boundary() {
419        let seq = [-3.5f32, -1.0, -0.25, 0.0, 0.25, 1.0, 3.5];
420        for w in seq.windows(2) {
421            assert!(sortable(w[0]) < sortable(w[1]),
422                    "{} must sort below {}", w[0], w[1]);
423        }
424    }
425
426    #[test]
427    fn malformed_navigation_rows_are_errors_not_missing_nodes() {
428        assert!(matches!(decode_row(&[0; 3], 8), Err(Error::Corrupt { .. })));
429
430        let mut trailing = encode_row(1.0, &[0; 8], &[7]);
431        trailing.push(0);
432        assert!(matches!(decode_row(&trailing, 8), Err(Error::Corrupt { .. })));
433    }
434
435    /// Pins the alpha semantics: distances here are SQUARED L2, so the
436    /// spread rule must use ALPHA^2. Geometry chosen so that a candidate
437    /// sits in the band between sqrt(1.2)-slack and 1.2-slack: the correct
438    /// rule keeps it, the squared-alpha-as-is bug drops it.
439    #[test]
440    fn robust_prune_alpha_acts_on_true_distances() {
441        let d = tempfile::TempDir::new().unwrap();
442        let g = Graph::new(Store::create(d.path(), Config::default()).unwrap()).unwrap();
443        let mut vc: std::collections::HashMap<u64, Vec<f32>> = std::collections::HashMap::new();
444        vc.insert(10, vec![1.0, 0.0]);          // c1: nearest, always kept
445        vc.insert(20, vec![0.68375, 1.10567]);  // c2: d(p,.)=1.3, d(c1,.)=1.15
446        vc.insert(30, vec![1.05, 0.0]);         // c3: dominated by c1, must drop
447        let cands = vec![
448            (1.0f32, 10u64),      // d(p,c1)^2
449            (1.69, 20),           // 1.3^2
450            (1.1025, 30),         // 1.05^2
451        ];
452        let kept = g.robust_prune(0, &mut vc, 0, &cands, 0).unwrap();
453        // 1.2 * d(c1,c2) = 1.38 > 1.3         -> c2 survives the spread rule
454        // 1.2 * d(c1,c3) = 0.06 < 1.05        -> c3 dropped
455        assert_eq!(kept, vec![10, 20], "alpha must apply to true distances, not squared");
456    }
457}