rig_core/vector_store/
lsh.rs1use 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#[derive(Clone, Default)]
29pub struct LSH {
30 hyperplanes: Vec<Vec<f32>>,
31 num_tables: usize,
32 num_hyperplanes: usize,
33}
34
35impl LSH {
36 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 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#[derive(Clone, Default)]
98pub struct LSHIndex {
99 lsh: LSH,
100 tables: Vec<HashMap<u64, Vec<String>>>, }
102
103impl LSHIndex {
104 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 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 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 pub fn clear(&mut self) {
145 for table in self.tables.iter_mut() {
146 table.clear();
147 }
148 }
149}