1pub struct GpuSearchBackend {
13 #[cfg(feature = "gpu")]
15 accelerator: Option<rustyhdf5_gpu::GpuAccelerator>,
16 dim: usize,
18 threshold: usize,
20 num_vectors: usize,
22}
23
24impl GpuSearchBackend {
25 pub fn try_init(
30 vectors: &[Vec<f32>],
31 norms: &[f32],
32 dim: usize,
33 threshold: usize,
34 ) -> Self {
35 #[cfg(feature = "gpu")]
36 {
37 if vectors.len() >= threshold {
38 match rustyhdf5_gpu::GpuAccelerator::new() {
39 Ok(mut accel) => {
40 let flat: Vec<f32> =
41 vectors.iter().flat_map(|v| v.iter().copied()).collect();
42 if accel.upload_vectors(&flat, dim).is_ok()
43 && accel.upload_norms(norms).is_ok()
44 {
45 return Self {
46 accelerator: Some(accel),
47 dim,
48 threshold,
49 num_vectors: vectors.len(),
50 };
51 }
52 }
53 Err(e) => {
54 log_gpu_fallback(&e.to_string());
55 }
56 }
57 }
58
59 Self {
60 accelerator: None,
61 dim,
62 threshold,
63 num_vectors: vectors.len(),
64 }
65 }
66
67 #[cfg(not(feature = "gpu"))]
68 {
69 let _ = (vectors, norms);
70 Self {
71 dim,
72 threshold,
73 num_vectors: 0,
74 }
75 }
76 }
77
78 pub fn is_available(&self) -> bool {
80 #[cfg(feature = "gpu")]
81 {
82 self.accelerator.is_some()
83 }
84 #[cfg(not(feature = "gpu"))]
85 {
86 false
87 }
88 }
89
90 pub fn dim(&self) -> usize {
92 self.dim
93 }
94
95 pub fn threshold(&self) -> usize {
97 self.threshold
98 }
99
100 pub fn re_upload(&mut self, vectors: &[Vec<f32>], norms: &[f32]) {
102 self.num_vectors = vectors.len();
103
104 #[cfg(feature = "gpu")]
105 {
106 if let Some(ref mut accel) = self.accelerator {
108 if vectors.len() >= self.threshold {
109 let flat: Vec<f32> =
110 vectors.iter().flat_map(|v| v.iter().copied()).collect();
111 if accel.upload_vectors(&flat, self.dim).is_err()
112 || accel.upload_norms(norms).is_err()
113 {
114 self.accelerator = None;
115 }
116 } else {
117 self.accelerator = None;
119 }
120 return;
121 }
122
123 if vectors.len() >= self.threshold {
125 if let Ok(mut accel) = rustyhdf5_gpu::GpuAccelerator::new() {
126 let flat: Vec<f32> =
127 vectors.iter().flat_map(|v| v.iter().copied()).collect();
128 if accel.upload_vectors(&flat, self.dim).is_ok()
129 && accel.upload_norms(norms).is_ok()
130 {
131 self.accelerator = Some(accel);
132 }
133 }
134 }
135 }
136
137 #[cfg(not(feature = "gpu"))]
138 {
139 let _ = (vectors, norms);
140 }
141 }
142
143 pub fn search_cosine(
147 &self,
148 query: &[f32],
149 vectors: &[Vec<f32>],
150 norms: &[f32],
151 tombstones: &[u8],
152 k: usize,
153 ) -> Vec<(usize, f32)> {
154 #[cfg(feature = "gpu")]
155 {
156 if let Some(ref accel) = self.accelerator {
157 match accel.cosine_search(query, k.min(self.num_vectors.max(1))) {
158 Ok(mut results) => {
159 results.retain(|(i, _)| {
161 *i < tombstones.len() && tombstones[*i] == 0
162 });
163 results.truncate(k);
164 return results;
165 }
166 Err(_) => {
167 }
169 }
170 }
171 }
172
173 cpu_fallback_cosine(query, vectors, norms, tombstones, k)
174 }
175
176 pub fn search_l2(
178 &self,
179 query: &[f32],
180 vectors: &[Vec<f32>],
181 tombstones: &[u8],
182 k: usize,
183 ) -> Vec<(usize, f32)> {
184 #[cfg(feature = "gpu")]
185 {
186 if let Some(ref accel) = self.accelerator {
187 match accel.l2_search(query, k.min(self.num_vectors.max(1))) {
188 Ok(mut results) => {
189 results.retain(|(i, _)| {
190 *i < tombstones.len() && tombstones[*i] == 0
191 });
192 results.truncate(k);
193 return results;
194 }
195 Err(_) => {
196 }
198 }
199 }
200 }
201
202 cpu_fallback_l2(query, vectors, tombstones, k)
203 }
204
205 pub fn device_info(&self) -> String {
207 #[cfg(feature = "gpu")]
208 {
209 if let Some(ref accel) = self.accelerator {
210 return accel.device_info().to_string();
211 }
212 }
213 "none".to_string()
214 }
215}
216
217pub fn detect_gpu() -> bool {
219 #[cfg(feature = "gpu")]
220 {
221 rustyhdf5_gpu::GpuAccelerator::is_available()
222 }
223 #[cfg(not(feature = "gpu"))]
224 {
225 false
226 }
227}
228
229#[cfg(feature = "gpu")]
230fn log_gpu_fallback(reason: &str) {
231 eprintln!("[edgehdf5-memory] GPU init failed, falling back to CPU: {reason}");
233}
234
235fn cpu_fallback_cosine(
237 query: &[f32],
238 vectors: &[Vec<f32>],
239 norms: &[f32],
240 tombstones: &[u8],
241 k: usize,
242) -> Vec<(usize, f32)> {
243 let query_norm = rustyhdf5_accel::vector_norm(query);
244 if query_norm == 0.0 {
245 return Vec::new();
246 }
247
248 let mut results: Vec<(usize, f32)> = Vec::with_capacity(vectors.len());
249
250 for (i, vec) in vectors.iter().enumerate() {
251 if i < tombstones.len() && tombstones[i] != 0 {
252 continue;
253 }
254 let vec_norm = if i < norms.len() {
255 norms[i]
256 } else {
257 rustyhdf5_accel::vector_norm(vec)
258 };
259 let score = crate::cosine_similarity_prenorm(query, query_norm, vec, vec_norm);
260 results.push((i, score));
261 }
262
263 results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
264 results.truncate(k);
265 results
266}
267
268fn cpu_fallback_l2(
270 query: &[f32],
271 vectors: &[Vec<f32>],
272 tombstones: &[u8],
273 k: usize,
274) -> Vec<(usize, f32)> {
275 let mut results: Vec<(usize, f32)> = Vec::with_capacity(vectors.len());
276
277 for (i, vec) in vectors.iter().enumerate() {
278 if i < tombstones.len() && tombstones[i] != 0 {
279 continue;
280 }
281 let dist = rustyhdf5_accel::l2_distance(query, vec);
282 results.push((i, dist));
283 }
284
285 results.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
287 results.truncate(k);
288 results
289}
290
291#[cfg(test)]
292mod tests {
293 use super::*;
294
295 fn make_vectors(n: usize, dim: usize, seed: u32) -> Vec<Vec<f32>> {
296 let mut s = seed;
297 let mut next = || -> f32 {
298 s = s.wrapping_mul(1103515245).wrapping_add(12345);
299 ((s >> 16) as f32) / 65536.0 - 0.5
300 };
301 (0..n).map(|_| (0..dim).map(|_| next()).collect()).collect()
302 }
303
304 #[test]
305 fn gpu_detect_default_status() {
306 let detected = detect_gpu();
308 let _ = detected;
310 }
311
312 #[test]
313 fn gpu_backend_fallback_when_unavailable() {
314 let vectors = make_vectors(100, 32, 42);
315 let norms: Vec<f32> = vectors
316 .iter()
317 .map(|v| rustyhdf5_accel::vector_norm(v))
318 .collect();
319 let backend = GpuSearchBackend::try_init(&vectors, &norms, 32, 50);
320
321 let tombstones = vec![0u8; 100];
323 let query = vectors[0].clone();
324 let results = backend.search_cosine(&query, &vectors, &norms, &tombstones, 10);
325
326 assert!(!results.is_empty());
327 assert!(results.len() <= 10);
328 assert_eq!(results[0].0, 0);
330 assert!((results[0].1 - 1.0).abs() < 1e-5);
331 }
332
333 #[test]
334 fn gpu_backend_below_threshold() {
335 let vectors = make_vectors(10, 32, 42);
336 let norms: Vec<f32> = vectors
337 .iter()
338 .map(|v| rustyhdf5_accel::vector_norm(v))
339 .collect();
340 let backend = GpuSearchBackend::try_init(&vectors, &norms, 32, 100);
341
342 assert!(!backend.is_available());
343 assert_eq!(backend.dim(), 32);
344 assert_eq!(backend.threshold(), 100);
345 }
346
347 #[test]
348 fn gpu_cosine_cpu_fallback_matches() {
349 let vectors = make_vectors(200, 64, 42);
350 let norms: Vec<f32> = vectors
351 .iter()
352 .map(|v| rustyhdf5_accel::vector_norm(v))
353 .collect();
354 let tombstones = vec![0u8; 200];
355 let query = vectors[5].clone();
356
357 let fallback = cpu_fallback_cosine(&query, &vectors, &norms, &tombstones, 10);
358 let backend = GpuSearchBackend::try_init(&vectors, &norms, 64, 50);
359 let backend_results =
360 backend.search_cosine(&query, &vectors, &norms, &tombstones, 10);
361
362 assert_eq!(fallback.len(), backend_results.len());
363 for (f, b) in fallback.iter().zip(&backend_results) {
364 assert_eq!(f.0, b.0);
365 assert!((f.1 - b.1).abs() < 1e-6);
366 }
367 }
368
369 #[test]
370 fn gpu_l2_search_returns_nearest() {
371 let vectors = vec![
372 vec![0.0, 0.0, 0.0],
373 vec![1.0, 0.0, 0.0],
374 vec![10.0, 10.0, 10.0],
375 ];
376 let tombstones = vec![0u8; 3];
377 let query = vec![0.1, 0.0, 0.0];
378
379 let results = cpu_fallback_l2(&query, &vectors, &tombstones, 3);
380 assert_eq!(results[0].0, 0);
382 assert_eq!(results[1].0, 1);
383 assert_eq!(results[2].0, 2);
384 }
385
386 #[test]
387 fn gpu_search_respects_tombstones() {
388 let vectors = make_vectors(50, 16, 42);
389 let norms: Vec<f32> = vectors
390 .iter()
391 .map(|v| rustyhdf5_accel::vector_norm(v))
392 .collect();
393 let mut tombstones = vec![0u8; 50];
394 tombstones[0] = 1;
395 tombstones[1] = 1;
396
397 let query = vectors[2].clone();
398 let results = cpu_fallback_cosine(&query, &vectors, &norms, &tombstones, 50);
399 assert!(results.iter().all(|r| r.0 != 0 && r.0 != 1));
400 }
401
402 #[test]
403 fn gpu_re_upload_updates_data() {
404 let vectors = make_vectors(10, 16, 42);
405 let norms: Vec<f32> = vectors
406 .iter()
407 .map(|v| rustyhdf5_accel::vector_norm(v))
408 .collect();
409 let mut backend = GpuSearchBackend::try_init(&vectors, &norms, 16, 5);
410
411 let vectors2 = make_vectors(20, 16, 77);
413 let norms2: Vec<f32> = vectors2
414 .iter()
415 .map(|v| rustyhdf5_accel::vector_norm(v))
416 .collect();
417 backend.re_upload(&vectors2, &norms2);
418
419 let tombstones = vec![0u8; 20];
421 let query = vectors2[0].clone();
422 let results = backend.search_cosine(&query, &vectors2, &norms2, &tombstones, 5);
423 assert!(!results.is_empty());
424 }
425
426 #[test]
427 fn gpu_l2_search_respects_tombstones() {
428 let vectors = vec![vec![0.0, 0.0], vec![1.0, 0.0], vec![2.0, 0.0]];
429 let mut tombstones = vec![0u8; 3];
430 tombstones[0] = 1; let query = vec![0.0, 0.0];
433 let results = cpu_fallback_l2(&query, &vectors, &tombstones, 3);
434 assert!(results.iter().all(|r| r.0 != 0));
435 assert_eq!(results[0].0, 1); }
437
438 #[test]
439 fn device_info_returns_string() {
440 let vectors = make_vectors(10, 16, 42);
441 let norms: Vec<f32> = vectors
442 .iter()
443 .map(|v| rustyhdf5_accel::vector_norm(v))
444 .collect();
445 let backend = GpuSearchBackend::try_init(&vectors, &norms, 16, 5);
446 let info = backend.device_info();
447 assert!(!info.is_empty());
448 }
449}