1use urna_format::Int8EmbeddingsView;
8
9use super::search::{greedy_search, layer_search};
10use super::select_neighbors::select_neighbors_heuristic;
11use super::visited::VisitedList;
12use super::{Candidate, HnswIndex, Node, dist_rr};
13use crate::materialize::PackedVectors;
14
15impl HnswIndex {
16 pub fn build(
19 vectors: Vec<f32>,
20 n: usize,
21 dim: usize,
22 m: usize,
23 ef_construction: usize,
24 seed: u64,
25 ) -> Self {
26 assert_eq!(vectors.len(), n * dim);
27 let store = PackedVectors::F32(vectors);
32 let m_max0 = m * 2;
33 let mut rng = LcgRng::new(seed);
34
35 let mut visited = VisitedList::new(n);
38 let mut nodes: Vec<Node> = Vec::with_capacity(n);
39 for _ in 0..n {
40 let level = sample_level(&mut rng, m);
41 let mut neighbors = Vec::with_capacity((level as usize) + 1);
42 for _ in 0..=level {
43 neighbors.push(Vec::new());
44 }
45 nodes.push(Node { level, neighbors });
46 }
47
48 let mut entry_point: u32 = 0;
50 let mut max_level: u32 = 0;
51
52 for i in 0..n {
53 let level = nodes[i].level;
54 if i == 0 {
55 entry_point = 0;
56 max_level = level;
57 continue;
58 }
59
60 let mut q_scratch = store.scratch(dim);
65 let mut sa = store.scratch(dim);
66 let mut sb = store.scratch(dim);
67 let q = store.row(i, dim, &mut q_scratch);
68 let mut curr = entry_point;
69 for layer in (level + 1..=max_level).rev() {
70 curr = greedy_search(curr, q, layer, &nodes, &store, dim);
71 }
72
73 let mut entry = curr;
77 let start_layer = level.min(max_level);
78 for layer in (0..=start_layer).rev() {
79 visited.clear();
80 let candidates = layer_search(
81 &[entry],
82 q,
83 layer,
84 ef_construction,
85 &nodes,
86 &store,
87 dim,
88 i as u32,
89 &mut visited,
90 );
91 let cap_layer = if layer == 0 { m_max0 } else { m };
97 let neighbor_ids =
98 select_neighbors_heuristic(&candidates, m, &store, dim, &mut sa, &mut sb, true);
99 nodes[i].neighbors[layer as usize] = neighbor_ids.clone();
100
101 for &nbr in &neighbor_ids {
104 let nbr_idx = nbr as usize;
105 let layer_idx = layer as usize;
106 if layer_idx >= nodes[nbr_idx].neighbors.len() {
107 continue;
108 }
109 if nodes[nbr_idx].neighbors[layer_idx].len() >= cap_layer {
110 let existing = nodes[nbr_idx].neighbors[layer_idx].clone();
114 let mut all: Vec<Candidate> = Vec::with_capacity(existing.len() + 1);
115 for id in existing {
116 let dist = dist_rr(&store, nbr_idx, id as usize, dim, &mut sa, &mut sb);
117 all.push(Candidate { id, dist });
118 }
119 let dist = dist_rr(&store, nbr_idx, i, dim, &mut sa, &mut sb);
120 all.push(Candidate { id: i as u32, dist });
121 nodes[nbr_idx].neighbors[layer_idx] = select_neighbors_heuristic(
122 &all, cap_layer, &store, dim, &mut sa, &mut sb, true,
123 );
124 } else if !nodes[nbr_idx].neighbors[layer_idx].contains(&(i as u32)) {
125 nodes[nbr_idx].neighbors[layer_idx].push(i as u32);
126 }
127 }
128
129 if let Some(&first) = neighbor_ids.first() {
133 entry = first;
134 } else if !candidates.is_empty() {
135 entry = candidates[0].id;
136 }
137 }
138
139 if level > max_level {
140 max_level = level;
141 entry_point = i as u32;
142 }
143 }
144
145 Self {
146 m,
147 m_max0,
148 ef_construction,
149 entry_point,
150 max_level,
151 nodes,
152 store,
153 dim,
154 n,
155 ef_search: ef_construction,
156 }
157 }
158
159 pub fn build_from_int8(view: &Int8EmbeddingsView<'_>, m: usize, ef: usize, seed: u64) -> Self {
163 let mut vectors = vec![0.0f32; view.n * view.dim];
164 for i in 0..view.n {
165 let scale = view.scale(i);
166 let row = view.row(i);
167 for j in 0..view.dim {
168 vectors[i * view.dim + j] = row[j] as f32 * scale;
169 }
170 }
171 Self::build(vectors, view.n, view.dim, m, ef, seed)
172 }
173
174 pub fn build_from_f16(
176 bytes: &[u8],
177 n: usize,
178 dim: usize,
179 m: usize,
180 ef: usize,
181 seed: u64,
182 ) -> Self {
183 let vectors = urna_format::f16_bytes_to_f32(bytes);
184 Self::build(vectors, n, dim, m, ef, seed)
185 }
186
187 pub fn build_from_f32(
189 bytes: &[u8],
190 n: usize,
191 dim: usize,
192 m: usize,
193 ef: usize,
194 seed: u64,
195 ) -> Self {
196 let mut vectors = vec![0.0f32; n * dim];
197 for (i, slot) in vectors.iter_mut().enumerate() {
198 let off = i * 4;
199 *slot =
200 f32::from_le_bytes([bytes[off], bytes[off + 1], bytes[off + 2], bytes[off + 3]]);
201 }
202 Self::build(vectors, n, dim, m, ef, seed)
203 }
204}
205
206pub(super) fn sample_level(rng: &mut LcgRng, m: usize) -> u32 {
208 let m_l = 1.0 / (m as f64).ln();
209 let r = rng.next_f64();
210 if r <= 0.0 {
211 return 0;
212 }
213 let level = (-(r.ln()) * m_l).floor() as i64;
214 level.clamp(0, 31) as u32 }
216
217pub(super) struct LcgRng {
220 state: u64,
221}
222impl LcgRng {
223 pub(super) fn new(seed: u64) -> Self {
224 Self {
225 state: seed
226 .wrapping_mul(2862933555777941757)
227 .wrapping_add(3037000493),
228 }
229 }
230 pub(super) fn next_u64(&mut self) -> u64 {
231 self.state = self
232 .state
233 .wrapping_mul(6364136223846793005)
234 .wrapping_add(1442695040888963407);
235 self.state
236 }
237 pub(super) fn next_f64(&mut self) -> f64 {
238 ((self.next_u64() >> 11) as f64) * (1.0 / ((1u64 << 53) as f64))
240 }
241}