1use std::cmp::Ordering;
40
41use crate::conversion::CastFromSlice;
42use crate::{norm::FastL2Norm, Half};
43use diskann_wide::arch::dispatch1;
44
45const NORM_LIMIT: f32 = f32::MIN_POSITIVE;
47
48#[inline]
51fn disjoint_ranges<Idx: Ord>(x_idx: &[Idx], y_idx: &[Idx]) -> bool {
52 x_idx.is_empty()
53 || y_idx.is_empty()
54 || x_idx[x_idx.len() - 1] < y_idx[0]
55 || y_idx[y_idx.len() - 1] < x_idx[0]
56}
57
58#[inline]
64pub fn indices_sorted_unique<Idx: Ord>(idx: &[Idx]) -> bool {
65 idx.is_sorted_by(|a, b| a < b)
66}
67
68#[inline]
71fn widen_pair(x_val: &[Half], y_val: &[Half]) -> Vec<f32> {
72 let mut buf = vec![0.0f32; x_val.len() + y_val.len()];
73 let (xf, yf) = buf.split_at_mut(x_val.len());
74 xf.cast_from_slice(x_val);
75 yf.cast_from_slice(y_val);
76 buf
77}
78
79#[derive(Debug, Clone, Copy, PartialEq, Eq)]
81pub struct LengthMismatch {
82 pub idx_len: usize,
83 pub val_len: usize,
84}
85
86impl std::fmt::Display for LengthMismatch {
87 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88 write!(
89 f,
90 "sparse operand index/value length mismatch: {} indices vs {} values",
91 self.idx_len, self.val_len
92 )
93 }
94}
95
96impl std::error::Error for LengthMismatch {}
97
98#[inline]
100fn check_len<Idx, T>(idx: &[Idx], val: &[T]) -> Result<(), LengthMismatch> {
101 if idx.len() == val.len() {
102 Ok(())
103 } else {
104 Err(LengthMismatch {
105 idx_len: idx.len(),
106 val_len: val.len(),
107 })
108 }
109}
110
111#[inline]
113fn pairs<'a, Idx: Copy>(idx: &'a [Idx], val: &'a [f32]) -> impl Iterator<Item = (Idx, f32)> + 'a {
114 idx.iter().copied().zip(val.iter().copied())
115}
116
117#[inline]
123fn merge_dot<Idx, I, J>(mut x: I, mut y: J) -> f32
124where
125 Idx: Copy + Ord,
126 I: Iterator<Item = (Idx, f32)>,
127 J: Iterator<Item = (Idx, f32)>,
128{
129 let mut acc = 0.0f32;
130 let mut a = x.next();
131 let mut b = y.next();
132 while let (Some((ai, av)), Some((bi, bv))) = (a, b) {
133 match ai.cmp(&bi) {
134 Ordering::Equal => {
135 acc = av.mul_add(bv, acc);
136 a = x.next();
137 b = y.next();
138 }
139 Ordering::Less => a = x.next(),
140 Ordering::Greater => b = y.next(),
141 }
142 }
143 acc
144}
145
146#[inline]
149fn merge_l2_sq<Idx, I, J>(mut x: I, mut y: J) -> f32
150where
151 Idx: Copy + Ord,
152 I: Iterator<Item = (Idx, f32)>,
153 J: Iterator<Item = (Idx, f32)>,
154{
155 let mut acc = 0.0f32;
156 let mut a = x.next();
157 let mut b = y.next();
158 loop {
159 match (a, b) {
160 (Some((ai, av)), Some((bi, bv))) => match ai.cmp(&bi) {
161 Ordering::Equal => {
162 let d = av - bv;
163 acc = d.mul_add(d, acc);
164 a = x.next();
165 b = y.next();
166 }
167 Ordering::Less => {
168 acc = av.mul_add(av, acc);
169 a = x.next();
170 }
171 Ordering::Greater => {
172 acc = bv.mul_add(bv, acc);
173 b = y.next();
174 }
175 },
176 (Some((_, av)), None) => {
177 acc = av.mul_add(av, acc);
178 a = x.next();
179 }
180 (None, Some((_, bv))) => {
181 acc = bv.mul_add(bv, acc);
182 b = y.next();
183 }
184 (None, None) => break,
185 }
186 }
187 acc
188}
189
190#[inline]
193fn cosine_from_parts(dot: f32, nx: f32, ny: f32) -> f32 {
194 if nx * nx < NORM_LIMIT || ny * ny < NORM_LIMIT {
195 0.0
196 } else {
197 let v = dot / (nx * ny);
198 (-1.0f32).max(1.0f32.min(v))
199 }
200}
201
202#[inline]
208pub fn l2_f32<Idx: Copy + Ord>(
209 x_idx: &[Idx],
210 x_val: &[f32],
211 y_idx: &[Idx],
212 y_val: &[f32],
213) -> Result<f32, LengthMismatch> {
214 check_len(x_idx, x_val)?;
215 check_len(y_idx, y_val)?;
216 let d = merge_l2_sq(pairs(x_idx, x_val), pairs(y_idx, y_val));
217 Ok(d.sqrt())
218}
219
220#[inline]
222pub fn inner_product_f32<Idx: Copy + Ord>(
223 x_idx: &[Idx],
224 x_val: &[f32],
225 y_idx: &[Idx],
226 y_val: &[f32],
227) -> Result<f32, LengthMismatch> {
228 check_len(x_idx, x_val)?;
229 check_len(y_idx, y_val)?;
230 if disjoint_ranges(x_idx, y_idx) {
231 return Ok(0.0);
232 }
233 Ok(merge_dot(pairs(x_idx, x_val), pairs(y_idx, y_val)))
234}
235
236#[inline]
239pub fn cosine_f32<Idx: Copy + Ord>(
240 x_idx: &[Idx],
241 x_val: &[f32],
242 y_idx: &[Idx],
243 y_val: &[f32],
244) -> Result<f32, LengthMismatch> {
245 check_len(x_idx, x_val)?;
246 check_len(y_idx, y_val)?;
247 if disjoint_ranges(x_idx, y_idx) {
248 return Ok(0.0);
249 }
250 let dot = merge_dot(pairs(x_idx, x_val), pairs(y_idx, y_val));
251 let nx = dispatch1(FastL2Norm, x_val);
252 let ny = dispatch1(FastL2Norm, y_val);
253 Ok(cosine_from_parts(dot, nx, ny))
254}
255
256#[inline]
263pub fn l2_f16<Idx: Copy + Ord>(
264 x_idx: &[Idx],
265 x_val: &[Half],
266 y_idx: &[Idx],
267 y_val: &[Half],
268) -> Result<f32, LengthMismatch> {
269 check_len(x_idx, x_val)?;
270 check_len(y_idx, y_val)?;
271 let buf = widen_pair(x_val, y_val);
272 let (xf, yf) = buf.split_at(x_val.len());
273 l2_f32(x_idx, xf, y_idx, yf)
274}
275
276#[inline]
279pub fn inner_product_f16<Idx: Copy + Ord>(
280 x_idx: &[Idx],
281 x_val: &[Half],
282 y_idx: &[Idx],
283 y_val: &[Half],
284) -> Result<f32, LengthMismatch> {
285 check_len(x_idx, x_val)?;
286 check_len(y_idx, y_val)?;
287 if disjoint_ranges(x_idx, y_idx) {
288 return Ok(0.0);
289 }
290 let buf = widen_pair(x_val, y_val);
291 let (xf, yf) = buf.split_at(x_val.len());
292 inner_product_f32(x_idx, xf, y_idx, yf)
293}
294
295#[inline]
299pub fn cosine_f16<Idx: Copy + Ord>(
300 x_idx: &[Idx],
301 x_val: &[Half],
302 y_idx: &[Idx],
303 y_val: &[Half],
304) -> Result<f32, LengthMismatch> {
305 check_len(x_idx, x_val)?;
306 check_len(y_idx, y_val)?;
307 if disjoint_ranges(x_idx, y_idx) {
308 return Ok(0.0);
309 }
310 let buf = widen_pair(x_val, y_val);
311 let (xf, yf) = buf.split_at(x_val.len());
312 cosine_f32(x_idx, xf, y_idx, yf)
313}
314
315#[cfg(test)]
320mod test {
321 use super::*;
322
323 use approx::{assert_abs_diff_eq, assert_relative_eq};
324 use diskann_wide::cast_f16_to_f32;
325 use rand::{
326 distr::{Distribution, Uniform},
327 rngs::StdRng,
328 SeedableRng,
329 };
330
331 fn dense(idx: &[u16], val: &[f32], dim: usize) -> Vec<f64> {
333 let mut v = vec![0.0f64; dim];
334 for (&i, &x) in idx.iter().zip(val.iter()) {
335 v[i as usize] = x as f64;
336 }
337 v
338 }
339
340 fn ref_dot(a: &[f64], b: &[f64]) -> f64 {
341 a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
342 }
343
344 fn ref_l2(a: &[f64], b: &[f64]) -> f64 {
345 a.iter()
346 .zip(b.iter())
347 .map(|(x, y)| (x - y) * (x - y))
348 .sum::<f64>()
349 .sqrt()
350 }
351
352 fn ref_cos(a: &[f64], b: &[f64]) -> f64 {
353 let na = a.iter().map(|x| x * x).sum::<f64>().sqrt();
354 let nb = b.iter().map(|x| x * x).sum::<f64>().sqrt();
355 if na == 0.0 || nb == 0.0 {
357 0.0
358 } else {
359 ref_dot(a, b) / (na * nb)
360 }
361 }
362
363 fn to_f16(v: &[f32]) -> Vec<Half> {
364 v.iter().map(|&x| Half::from_f32(x)).collect()
365 }
366
367 fn to_sparse_f32(v: &[f32]) -> (Vec<u16>, Vec<f32>) {
368 let mut idx = Vec::new();
369 let mut val = Vec::new();
370 for (i, &x) in v.iter().enumerate() {
371 if x != 0.0 {
372 idx.push(i as u16);
373 val.push(x);
374 }
375 }
376 (idx, val)
377 }
378
379 fn to_sparse_f16(v: &[Half]) -> (Vec<u16>, Vec<Half>) {
380 let mut idx = Vec::new();
381 let mut val = Vec::new();
382 for (i, &x) in v.iter().enumerate() {
383 if cast_f16_to_f32(x) != 0.0 {
384 idx.push(i as u16);
385 val.push(x);
386 }
387 }
388 (idx, val)
389 }
390
391 fn fill(v: &mut [f32], dist: &Uniform<f32>, rng: &mut StdRng) {
392 for x in v.iter_mut() {
393 *x = dist.sample(rng);
394 }
395 }
396
397 fn drop_half_to_zero(v: &mut [f32], rng: &mut StdRng) {
398 let coin = Uniform::new(0.0f32, 1.0).unwrap();
399 for x in v.iter_mut() {
400 if coin.sample(rng) < 0.5 {
401 *x = 0.0;
402 }
403 }
404 }
405
406 #[test]
408 fn dot_l2_cosine_match_dense_reference_f32() {
409 let dim = 16;
410 let xi = [1u16, 3, 4, 9, 12];
411 let xv = [0.5f32, -1.5, 2.0, 0.25, 3.0];
412 let yi = [0u16, 3, 4, 7, 12, 15];
413 let yv = [1.0f32, 2.0, -0.5, 4.0, 1.25, -2.0];
414
415 let da = dense(&xi, &xv, dim);
416 let db = dense(&yi, &yv, dim);
417
418 let ip = inner_product_f32(&xi, &xv, &yi, &yv).unwrap();
419 let l2 = l2_f32(&xi, &xv, &yi, &yv).unwrap();
420 let cos = cosine_f32(&xi, &xv, &yi, &yv).unwrap();
421
422 assert_abs_diff_eq!(ip as f64, ref_dot(&da, &db), epsilon = 1e-5);
423 assert_abs_diff_eq!(l2 as f64, ref_l2(&da, &db), epsilon = 1e-5);
424 assert_abs_diff_eq!(cos as f64, ref_cos(&da, &db), epsilon = 1e-5);
425 }
426
427 #[test]
429 fn dot_l2_cosine_match_dense_reference_f16() {
430 let dim = 16;
431 let xi = [1u16, 3, 4, 9, 12];
432 let xv = to_f16(&[0.5, -1.5, 2.0, 0.25, 3.0]);
433 let yi = [0u16, 3, 4, 7, 12, 15];
434 let yv = to_f16(&[1.0, 2.0, -0.5, 4.0, 1.25, -2.0]);
435
436 let xvf: Vec<f32> = xv.iter().map(|h| cast_f16_to_f32(*h)).collect();
437 let yvf: Vec<f32> = yv.iter().map(|h| cast_f16_to_f32(*h)).collect();
438 let da = dense(&xi, &xvf, dim);
439 let db = dense(&yi, &yvf, dim);
440
441 let ip = inner_product_f16(&xi, &xv, &yi, &yv).unwrap();
442 let l2 = l2_f16(&xi, &xv, &yi, &yv).unwrap();
443 let cos = cosine_f16(&xi, &xv, &yi, &yv).unwrap();
444
445 assert_abs_diff_eq!(ip as f64, ref_dot(&da, &db), epsilon = 1e-3);
446 assert_abs_diff_eq!(l2 as f64, ref_l2(&da, &db), epsilon = 1e-3);
447 assert_abs_diff_eq!(cos as f64, ref_cos(&da, &db), epsilon = 1e-3);
448 }
449
450 #[test]
452 fn disjoint_ranges_have_zero_dot_and_cosine() {
453 let xi = [1u16, 2, 3];
454 let xv = [1.0f32, 2.0, 3.0];
455 let yi = [10u16, 11, 12];
456 let yv = [1.0f32, 2.0, 3.0];
457 assert_eq!(inner_product_f32(&xi, &xv, &yi, &yv).unwrap(), 0.0);
458 assert_eq!(cosine_f32(&xi, &xv, &yi, &yv).unwrap(), 0.0);
459
460 let xvh = to_f16(&xv);
461 let yvh = to_f16(&yv);
462 assert_eq!(inner_product_f16(&xi, &xvh, &yi, &yvh).unwrap(), 0.0);
463 assert_eq!(cosine_f16(&xi, &xvh, &yi, &yvh).unwrap(), 0.0);
464 }
465
466 #[test]
467 fn indices_sorted_unique_detects_violations() {
468 let empty: [u16; 0] = [];
469 assert!(indices_sorted_unique(&empty));
470 assert!(indices_sorted_unique(&[1u16]));
471 assert!(indices_sorted_unique(&[1u16, 3, 4, 9]));
472 assert!(!indices_sorted_unique(&[1u16, 1]));
473 assert!(!indices_sorted_unique(&[3u16, 1, 4]));
474 }
475
476 #[test]
478 fn empty_operand_cosine_zero_and_l2_is_norm() {
479 let yi = [0u16, 2, 4];
480 let yv = [1.0f32, 2.0, 3.0];
481 let empty_i: [u16; 0] = [];
482 let empty_v: [f32; 0] = [];
483
484 assert_eq!(cosine_f32(&empty_i, &empty_v, &yi, &yv).unwrap(), 0.0);
485 assert_abs_diff_eq!(
486 l2_f32(&empty_i, &empty_v, &yi, &yv).unwrap(),
487 14.0f32.sqrt(),
488 epsilon = 1e-5
489 );
490 }
491
492 #[test]
494 fn matches_dense_reference_over_random_f32() {
495 let mut rng = StdRng::seed_from_u64(0x9e3779b97f4a7c15);
496 let dist = Uniform::new(-100.0f32, 100.0f32).unwrap();
497 for dim in 1..=96usize {
498 for _ in 0..32 {
499 let mut x = vec![0.0f32; dim];
500 let mut y = vec![0.0f32; dim];
501 fill(&mut x, &dist, &mut rng);
502 fill(&mut y, &dist, &mut rng);
503 drop_half_to_zero(&mut x, &mut rng);
504 drop_half_to_zero(&mut y, &mut rng);
505 let (xi, xv) = to_sparse_f32(&x);
506 let (yi, yv) = to_sparse_f32(&y);
507 let da: Vec<f64> = x.iter().map(|&v| v as f64).collect();
508 let db: Vec<f64> = y.iter().map(|&v| v as f64).collect();
509
510 let l2 = l2_f32(&xi, &xv, &yi, &yv).unwrap();
511 assert_relative_eq!(
512 l2 as f64,
513 ref_l2(&da, &db),
514 max_relative = 1e-4,
515 epsilon = 1e-3
516 );
517
518 let ip = inner_product_f32(&xi, &xv, &yi, &yv).unwrap();
519 assert_relative_eq!(
520 ip as f64,
521 ref_dot(&da, &db),
522 max_relative = 1e-4,
523 epsilon = 1e-2
524 );
525
526 let cos = cosine_f32(&xi, &xv, &yi, &yv).unwrap();
527 assert_relative_eq!(
528 cos as f64,
529 ref_cos(&da, &db),
530 max_relative = 1e-4,
531 epsilon = 1e-3
532 );
533 }
534 }
535 }
536
537 #[test]
539 fn matches_dense_reference_over_random_f16() {
540 let mut rng = StdRng::seed_from_u64(0xc2b2ae3d27d4eb4f);
541 let dist = Uniform::new(-10.0f32, 10.0f32).unwrap();
542 for dim in 1..=96usize {
543 for _ in 0..32 {
544 let mut xf = vec![0.0f32; dim];
545 let mut yf = vec![0.0f32; dim];
546 fill(&mut xf, &dist, &mut rng);
547 fill(&mut yf, &dist, &mut rng);
548 drop_half_to_zero(&mut xf, &mut rng);
549 drop_half_to_zero(&mut yf, &mut rng);
550 let x = to_f16(&xf);
551 let y = to_f16(&yf);
552 let (xi, xv) = to_sparse_f16(&x);
553 let (yi, yv) = to_sparse_f16(&y);
554 let da: Vec<f64> = x.iter().map(|v| cast_f16_to_f32(*v) as f64).collect();
555 let db: Vec<f64> = y.iter().map(|v| cast_f16_to_f32(*v) as f64).collect();
556
557 let l2 = l2_f16(&xi, &xv, &yi, &yv).unwrap();
558 assert_relative_eq!(
559 l2 as f64,
560 ref_l2(&da, &db),
561 max_relative = 5e-3,
562 epsilon = 5e-2
563 );
564
565 let ip = inner_product_f16(&xi, &xv, &yi, &yv).unwrap();
566 assert_relative_eq!(
567 ip as f64,
568 ref_dot(&da, &db),
569 max_relative = 5e-3,
570 epsilon = 5e-2
571 );
572
573 let cos = cosine_f16(&xi, &xv, &yi, &yv).unwrap();
574 assert_relative_eq!(
575 cos as f64,
576 ref_cos(&da, &db),
577 max_relative = 5e-3,
578 epsilon = 5e-2
579 );
580 }
581 }
582 }
583
584 #[test]
586 fn length_mismatch_returns_error() {
587 let xi = [0u16, 1, 2];
588 let xv = [1.0f32, 2.0]; let yi = [0u16, 1];
590 let yv = [1.0f32, 2.0];
591
592 assert!(l2_f32(&xi, &xv, &yi, &yv).is_err());
593 assert!(inner_product_f32(&xi, &xv, &yi, &yv).is_err());
594 assert!(cosine_f32(&xi, &xv, &yi, &yv).is_err());
595
596 let xvh = to_f16(&xv);
597 let yvh = to_f16(&yv);
598 assert!(l2_f16(&xi, &xvh, &yi, &yvh).is_err());
599
600 let err = l2_f32(&xi, &xv, &yi, &yv).unwrap_err();
601 assert_eq!(err.idx_len, 3);
602 assert_eq!(err.val_len, 2);
603 }
604
605 #[test]
607 fn generic_over_u32_matches_u16() {
608 let xv = [0.5f32, -1.5, 2.0, 0.25, 3.0];
609 let yv = [1.0f32, 2.0, -0.5, 4.0, 1.25, -2.0];
610 let xi16 = [1u16, 3, 4, 9, 12];
611 let yi16 = [0u16, 3, 4, 7, 12, 15];
612 let xi32 = [1u32, 3, 4, 9, 12];
613 let yi32 = [0u32, 3, 4, 7, 12, 15];
614
615 assert_eq!(
616 l2_f32(&xi16, &xv, &yi16, &yv).unwrap(),
617 l2_f32(&xi32, &xv, &yi32, &yv).unwrap()
618 );
619 assert_eq!(
620 inner_product_f32(&xi16, &xv, &yi16, &yv).unwrap(),
621 inner_product_f32(&xi32, &xv, &yi32, &yv).unwrap()
622 );
623 assert_eq!(
624 cosine_f32(&xi16, &xv, &yi16, &yv).unwrap(),
625 cosine_f32(&xi32, &xv, &yi32, &yv).unwrap()
626 );
627
628 let xvh = to_f16(&xv);
629 let yvh = to_f16(&yv);
630 assert_eq!(
631 l2_f16(&xi16, &xvh, &yi16, &yvh).unwrap(),
632 l2_f16(&xi32, &xvh, &yi32, &yvh).unwrap()
633 );
634 }
635}