Skip to main content

diskann_vector/
sparse.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6//! Distance kernels for **sparse** float32/float16 vectors.
7//!
8//! # Storage model
9//!
10//! Each operand is a pair of parallel slices:
11//!
12//! * `idx`: `nnz` dimension indices of a generic integer type (`u16`, `u32`, …), **sorted
13//!   ascending and unique**;
14//! * `val`: `nnz` values parallel to `idx` (`f32`, or `f16` via [`Half`]).
15//!
16//! Indices absent from an operand are implicit zeros.
17//!
18//! # Warning
19//!
20//! The kernels do not verify that `idx` is sorted and unique. Violating this yields
21//! incorrect results, not undefined behavior. Callers that cannot otherwise guarantee this
22//! invariant can check it with [`indices_sorted_unique`].
23//!
24//! # Kernels
25//!
26//! * **Inner product / cosine numerator** — intersection merge over the two sorted index
27//!   arrays; only matching indices contribute.
28//! * **Cosine denominator** — each operand's own L2 norm, computed by reusing
29//!   [`crate::norm::FastL2Norm`] (a sparse vector's norm is just the L2 norm of its value
30//!   array).
31//! * **L2** — direct union merge of squared differences, `sqrt(Σ (x_i − y_i)²)`.
32//!
33//! Accumulation is in `f32`. The merge is scalar. A disjoint-range fast-out skips the
34//! intersection merge when the two index ranges cannot overlap.
35//!
36//! These functions return the **mathematical** value of each metric. Any similarity-score
37//! transform (inner product `x -> -x`, cosine `x -> 1 - x`) is applied by the caller.
38
39use std::cmp::Ordering;
40
41use crate::conversion::CastFromSlice;
42use crate::{norm::FastL2Norm, Half};
43use diskann_wide::arch::dispatch1;
44
45/// Squared-norm floor below which a vector is treated as having zero norm for cosine.
46const NORM_LIMIT: f32 = f32::MIN_POSITIVE;
47
48/// True when two sorted, unique index arrays cannot share any index (either is empty, or their
49/// `[min, max]` ranges don't overlap).
50#[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/// Returns `true` if `idx` is sorted ascending with no duplicates.
59///
60/// # Warning
61///
62/// The kernels in this module require this invariant but do not check it themselves.
63#[inline]
64pub fn indices_sorted_unique<Idx: Ord>(idx: &[Idx]) -> bool {
65    idx.is_sorted_by(|a, b| a < b)
66}
67
68/// Widen both f16 operands into a single f32 buffer (`x` then `y`) using the dispatched SIMD
69/// slice conversion. One allocation; caller splits at `x_val.len()`.
70#[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/// Error returned when a sparse operand's index and value slices differ in length.
80#[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/// Validate that an operand's index and value slices are parallel (equal length).
99#[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/// Zip parallel index/value slices into `(index, value)` pairs for the merge kernels.
112#[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//////////////////////////////
118// Scalar merge kernels     //
119//////////////////////////////
120
121/// Intersection merge: `Σ x·y` over matching indices, accumulated in `f32`. `O(nnz_x + nnz_y)`.
122#[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/// Direct union merge for squared L2: `Σ (x_i − y_i)²`, accumulated in `f32`; an index present
147/// on only one side contributes `v²`. `O(nnz_x + nnz_y)`.
148#[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/// Cosine of the angle from the numerator `dot` and the two operand norms; `0` when either
191/// squared norm underflows [`NORM_LIMIT`], otherwise the ratio clamped to `[-1, 1]`.
192#[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//////////////////////////////
203// f32 kernels              //
204//////////////////////////////
205
206/// `sqrt(Σ (x_i − y_i)²)` for f32 operands.
207#[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/// `Σ x·y` over matching indices for f32 operands.
221#[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/// Cosine similarity `dot / (‖x‖·‖y‖)` for f32 operands, clamped to `[-1, 1]`; `0` when either
237/// norm underflows or the operands are disjoint. Norms reuse [`crate::norm::FastL2Norm`].
238#[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//////////////////////////////
257// f16 kernels              //
258//////////////////////////////
259
260/// `sqrt(Σ (x_i − y_i)²)` for f16 operands; values are pre-widened to f32, then reuse the f32
261/// union merge.
262#[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/// `Σ x·y` over matching indices for f16 operands; a disjoint-range fast-out skips widening
277/// when the operands cannot intersect.
278#[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/// Cosine similarity for f16 operands, clamped to `[-1, 1]`; `0` when either norm underflows or
296/// the operands are disjoint. Values are pre-widened once and reused for the numerator and both
297/// norms.
298#[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//////////////////////////////
316// Tests                    //
317//////////////////////////////
318
319#[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    // Dense f64 reference over the declared dimension, built from the sparse operands.
332    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        // A zero-norm operand has undefined cosine; the kernel reports 0.0 for it.
356        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    // f32 kernels agree with an f64 densified reference for dot, L2, and cosine.
407    #[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    // f16 kernels agree with an f64 densified reference for dot, L2, and cosine.
428    #[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    // Operands with no shared indices have zero dot and zero cosine (f32 and f16).
451    #[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    // An empty (zero-norm) operand yields cosine 0 and L2 equal to the other operand's norm.
477    #[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    // f32 kernels match the f64 densified reference over random, partially-zeroed vectors.
493    #[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    // f16 kernels match the f64 densified reference over random, partially-zeroed vectors.
538    #[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    // Mismatched index/value lengths return a LengthMismatch error rather than panicking.
585    #[test]
586    fn length_mismatch_returns_error() {
587        let xi = [0u16, 1, 2];
588        let xv = [1.0f32, 2.0]; // one value short
589        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    // The kernels are generic over the index type; u32 indices give the same result as u16.
606    #[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}