Skip to main content

rig_core/vector_store/
lsh.rs

1//! Random-projection hashing for approximate vector-search candidates.
2//!
3//! ```
4//! use rig_core::vector_store::lsh::LSHIndex;
5//!
6//! let mut index = LSHIndex::new(2, 4, 8);
7//! index.insert("document", &[1.0, 0.0]);
8//! assert_eq!(index.query(&[1.0, 0.0]), vec!["document"]);
9//! ```
10
11use fastrand::Rng;
12use std::collections::HashMap;
13
14#[cfg(test)]
15fn lsh_rng() -> Rng {
16    Rng::with_seed(0x5eed_fade_cafe_beef)
17}
18
19#[cfg(not(test))]
20fn lsh_rng() -> Rng {
21    Rng::new()
22}
23
24/// Locality Sensitive Hashing (LSH) with random projection.
25/// Uses random hyperplanes to hash similar vectors into the same buckets for efficient
26/// approximate nearest neighbor search. See <https://www.pinecone.io/learn/series/faiss/locality-sensitive-hashing-random-projection/>
27/// for details on how LSH works.
28#[derive(Clone, Default)]
29pub struct LSH {
30    hyperplanes: Vec<Vec<f32>>,
31    num_tables: usize,
32    num_hyperplanes: usize,
33}
34
35impl LSH {
36    /// Creates random normalized projection vectors. Use at most 64 hyperplanes
37    /// per table and ensure `num_tables * num_hyperplanes` fits in `usize`.
38    pub fn new(dim: usize, num_tables: usize, num_hyperplanes: usize) -> Self {
39        let mut rng = lsh_rng();
40        let mut hyperplanes = Vec::new();
41
42        for _ in 0..(num_tables * num_hyperplanes) {
43            let mut plane = vec![0.0; dim];
44
45            for val in plane.iter_mut() {
46                *val = rng.f32() * 2.0 - 1.0;
47            }
48
49            let norm: f32 = plane.iter().map(|x| x * x).sum::<f32>().sqrt();
50            if norm > 0.0 {
51                for val in plane.iter_mut() {
52                    *val /= norm;
53                }
54            }
55
56            hyperplanes.push(plane);
57        }
58
59        Self {
60            hyperplanes,
61            num_tables,
62            num_hyperplanes,
63        }
64    }
65
66    /// Computes sign bits using f32 projections. Supply a valid table index
67    /// and a vector matching the configured dimension; lengths are not checked.
68    pub fn hash(&self, vector: &[f64], table_idx: usize) -> u64 {
69        let mut hash = 0u64;
70        let start = table_idx * self.num_hyperplanes;
71
72        for (i, hyperplane) in self
73            .hyperplanes
74            .get(start..start + self.num_hyperplanes)
75            .unwrap_or(&[])
76            .iter()
77            .enumerate()
78        {
79            let dot: f32 = vector
80                .iter()
81                .zip(hyperplane.iter())
82                .map(|(v, h)| (*v as f32) * h)
83                .sum();
84
85            if dot >= 0.0 {
86                hash |= 1u64 << i;
87            }
88        }
89
90        hash
91    }
92}
93
94/// LSH Index for document IDs.
95/// Stores document IDs in a hashmap of hash values to document IDs.
96/// This allows for efficient lookup of document IDs by hash value.
97#[derive(Clone, Default)]
98pub struct LSHIndex {
99    lsh: LSH,
100    tables: Vec<HashMap<u64, Vec<String>>>, // Hash -> document IDs
101}
102
103impl LSHIndex {
104    /// Create a new LSHIndex.
105    pub fn new(dim: usize, num_tables: usize, num_hyperplanes: usize) -> Self {
106        let lsh = LSH::new(dim, num_tables, num_hyperplanes);
107        let tables = vec![HashMap::new(); num_tables];
108
109        Self { lsh, tables }
110    }
111
112    /// Insert a document ID with its embedding
113    pub fn insert(&mut self, id: &str, embedding: &[f64]) {
114        for table_idx in 0..self.lsh.num_tables {
115            let hash = self.lsh.hash(embedding, table_idx);
116            if let Some(table) = self.tables.get_mut(table_idx) {
117                table.entry(hash).or_default().push(id.to_owned());
118            }
119        }
120    }
121
122    /// Returns unique IDs sharing a bucket in any table, in unspecified order.
123    pub fn query(&self, embedding: &[f64]) -> Vec<String> {
124        use std::collections::HashSet;
125
126        let mut candidates = HashSet::new();
127
128        for table_idx in 0..self.lsh.num_tables {
129            let hash = self.lsh.hash(embedding, table_idx);
130
131            if let Some(ids) = self
132                .tables
133                .get(table_idx)
134                .and_then(|table| table.get(&hash))
135            {
136                candidates.extend(ids.iter().cloned());
137            }
138        }
139
140        candidates.into_iter().collect()
141    }
142
143    /// Clear all tables
144    pub fn clear(&mut self) {
145        for table in self.tables.iter_mut() {
146            table.clear();
147        }
148    }
149}