Skip to main content

diskann_utils/sampling/
medoid.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use crate::views::rowmajor::{self, Matrix};
7use diskann_vector::{conversion::CastFromSlice, distance::SquaredL2, PureDistanceFunction};
8use half::f16;
9
10/// Return the row in `data` that is closest to the medoid of all rows.
11pub trait ComputeMedoid: Sized {
12    fn compute_medoid(data: rowmajor::Ref<'_, Self>) -> Vec<Self>;
13}
14
15impl ComputeMedoid for f32 {
16    fn compute_medoid(data: rowmajor::Ref<'_, Self>) -> Vec<Self> {
17        if data.ncols() == 0 {
18            return vec![];
19        }
20
21        let mut sum = vec![0.0f64; data.ncols()];
22        data.rows().for_each(|r| {
23            std::iter::zip(sum.iter_mut(), r.iter()).for_each(|(o, i)| {
24                let i: f64 = (*i).into();
25                *o += i;
26            });
27        });
28
29        let m: Vec<f32> = sum
30            .iter()
31            .map(|s| (s / data.nrows() as f64) as f32)
32            .collect();
33
34        let mut min_dist: f32 = f32::MAX;
35        let mut medoid = None;
36        data.rows().for_each(|r| {
37            let d = SquaredL2::evaluate(m.as_slice(), r);
38            if d < min_dist {
39                min_dist = d;
40                medoid = Some(r);
41            }
42        });
43
44        medoid
45            .map(|x| x.into())
46            .unwrap_or(vec![0.0f32; data.ncols()])
47    }
48}
49
50impl ComputeMedoid for f16 {
51    fn compute_medoid(data: rowmajor::Ref<'_, Self>) -> Vec<Self> {
52        if data.ncols() == 0 {
53            return vec![];
54        }
55
56        let mut sum = vec![0.0f64; data.ncols()];
57        let mut buffer = vec![0.0f32; data.ncols()];
58        data.rows().for_each(|r| {
59            buffer.cast_from_slice(r);
60            std::iter::zip(sum.iter_mut(), buffer.iter()).for_each(|(o, i)| {
61                let i: f64 = (*i).into();
62                *o += i;
63            });
64        });
65
66        std::iter::zip(buffer.iter_mut(), sum.iter()).for_each(|(o, i)| {
67            *o = (*i / data.nrows() as f64) as f32;
68        });
69
70        let mut min_dist: f32 = f32::MAX;
71        let mut medoid = None;
72        data.rows().for_each(|r| {
73            let d = SquaredL2::evaluate(buffer.as_slice(), r);
74            if d < min_dist {
75                min_dist = d;
76                medoid = Some(r);
77            }
78        });
79
80        medoid
81            .map(|x| x.into())
82            .unwrap_or(vec![f16::default(); data.ncols()])
83    }
84}
85
86impl ComputeMedoid for u8 {
87    fn compute_medoid(data: rowmajor::Ref<'_, Self>) -> Vec<Self> {
88        if data.ncols() == 0 {
89            return vec![];
90        }
91
92        let mut sum = vec![0.0f64; data.ncols()];
93        data.rows().for_each(|r| {
94            std::iter::zip(sum.iter_mut(), r.iter()).for_each(|(o, i)| {
95                let i: f64 = (*i).into();
96                *o += i;
97            });
98        });
99
100        let m: Vec<f32> = sum
101            .iter()
102            .map(|s| (s / data.nrows() as f64) as f32)
103            .collect();
104
105        let mut min_dist: f32 = f32::MAX;
106        let mut medoid = None;
107        let mut as_float = vec![0.0f32; data.ncols()];
108        data.rows().for_each(|r| {
109            std::iter::zip(as_float.iter_mut(), r.iter())
110                .for_each(|(dst, src)| *dst = (*src).into());
111            let d = SquaredL2::evaluate(m.as_slice(), &*as_float);
112            if d < min_dist {
113                min_dist = d;
114                medoid = Some(r);
115            }
116        });
117
118        medoid.map(|x| x.into()).unwrap_or(vec![0u8; data.ncols()])
119    }
120}
121
122impl ComputeMedoid for i8 {
123    fn compute_medoid(data: rowmajor::Ref<'_, Self>) -> Vec<Self> {
124        if data.ncols() == 0 {
125            return vec![];
126        }
127
128        let mut sum = vec![0.0f64; data.ncols()];
129        data.rows().for_each(|r| {
130            std::iter::zip(sum.iter_mut(), r.iter()).for_each(|(o, i)| {
131                let i: f64 = (*i).into();
132                *o += i;
133            });
134        });
135
136        let m: Vec<f32> = sum
137            .iter()
138            .map(|s| (s / data.nrows() as f64) as f32)
139            .collect();
140
141        let mut min_dist: f32 = f32::MAX;
142        let mut medoid = None;
143        let mut as_float = vec![0.0f32; data.ncols()];
144        data.rows().for_each(|r| {
145            std::iter::zip(as_float.iter_mut(), r.iter())
146                .for_each(|(dst, src)| *dst = (*src).into());
147            let d = SquaredL2::evaluate(m.as_slice(), &*as_float);
148            if d < min_dist {
149                min_dist = d;
150                medoid = Some(r);
151            }
152        });
153
154        medoid.map(|x| x.into()).unwrap_or(vec![0i8; data.ncols()])
155    }
156}
157
158///////////
159// Tests //
160///////////
161
162#[cfg(not(miri))]
163#[cfg(test)]
164mod tests {
165    use super::*;
166
167    use diskann_wide::cast_f32_to_f16;
168    use rand::{
169        distr::{Distribution, StandardUniform},
170        rngs::StdRng,
171        SeedableRng,
172    };
173
174    use crate::views::rowmajor::MatrixMut;
175
176    fn example_dataset() -> (rowmajor::Owned<f32>, Vec<f32>) {
177        let data: Vec<f32> = vec![
178            // row 0
179            0.203688,
180            0.841956,
181            0.855665,
182            0.801917,
183            0.754536,
184            // row 1
185            0.312881,
186            0.217382,
187            0.0644115,
188            0.348708,
189            0.999495,
190            // row 2
191            0.657741,
192            0.914681,
193            0.555228,
194            0.13253,
195            0.118615,
196            // row 3
197            0.356464,
198            0.207449,
199            0.452471,
200            0.925219,
201            0.508498,
202            // row 4
203            0.749786,
204            0.90786,
205            0.129618,
206            0.597719,
207            0.000622153,
208            // row 5 -- this is the medoid
209            0.569517,
210            0.435447,
211            0.558136,
212            0.480974,
213            0.711425,
214            // row 6
215            0.896353,
216            0.275053,
217            0.0427179,
218            0.660916,
219            0.464851,
220            // row 7
221            0.558689,
222            0.596543,
223            0.740983,
224            0.122136,
225            0.453822,
226            // row 8
227            0.526895,
228            0.492643,
229            0.0951115,
230            0.495487,
231            0.446127,
232            // row 9
233            0.454093,
234            0.160239,
235            0.924585,
236            0.901708,
237            0.329328,
238        ];
239
240        let data = rowmajor::Owned::<f32>::try_from_data(data.into(), 10, 5).unwrap();
241        let expected: Vec<f32> = data.row(5).into();
242        (data, expected)
243    }
244
245    #[test]
246    fn test_f32() {
247        // No Rows
248        let x = rowmajor::Owned::<f32>::from_element(0, 10, 0.0f32);
249        assert_eq!(f32::compute_medoid(x.as_view()), vec![0.0; x.ncols()]);
250
251        // No Cols
252        let x = rowmajor::Owned::<f32>::from_element(10, 0, 0.0f32);
253        assert_eq!(f32::compute_medoid(x.as_view()), Vec::<f32>::new());
254
255        let mut rng = StdRng::seed_from_u64(0xaf2f5fa0b5161acf);
256
257        // One row
258        let dist = StandardUniform;
259        for dim in 1..20 {
260            let x = rowmajor::Owned::<f32>::from_fn(1, dim, |_| dist.sample(&mut rng));
261            assert_eq!(&*f32::compute_medoid(x.as_view()), x.row(0));
262        }
263
264        // Example dataset
265        let (data, expected) = example_dataset();
266        let m = f32::compute_medoid(data.as_view());
267        assert_eq!(m, expected);
268    }
269
270    #[test]
271    fn test_f16() {
272        // No Rows
273        let x = rowmajor::Owned::<f16>::from_element(0, 10, f16::default());
274        assert_eq!(
275            f16::compute_medoid(x.as_view()),
276            vec![f16::default(); x.ncols()]
277        );
278
279        // No Cols
280        let x = rowmajor::Owned::<f16>::from_element(10, 0, f16::default());
281        assert_eq!(f16::compute_medoid(x.as_view()), Vec::<f16>::new());
282
283        let mut rng = StdRng::seed_from_u64(0x88e2f7096fc9b90e);
284
285        // One row
286        let dist = StandardUniform;
287        for dim in 1..20 {
288            let x =
289                rowmajor::Owned::<f16>::from_fn(1, dim, |_| cast_f32_to_f16(dist.sample(&mut rng)));
290            assert_eq!(&*f16::compute_medoid(x.as_view()), x.row(0));
291        }
292
293        // Example dataset
294        let (data, expected) = example_dataset();
295        let mut data_f16 =
296            rowmajor::Owned::<f16>::from_element(data.nrows(), data.ncols(), f16::default());
297        data_f16.as_mut_slice().cast_from_slice(data.as_slice());
298
299        let mut expected_f16 = vec![f16::default(); expected.len()];
300        expected_f16.cast_from_slice(expected.as_slice());
301
302        let m = f16::compute_medoid(data_f16.as_view());
303        assert_eq!(m, expected_f16);
304    }
305
306    fn example_dataset_u8() -> (rowmajor::Owned<u8>, Vec<u8>) {
307        let data: Vec<u8> = vec![
308            52, 215, 218, 204, 192, // row 0
309            79, 55, 16, 89, 255, // row 1
310            167, 233, 141, 33, 30, // row 2
311            91, 53, 115, 236, 130, // row 3
312            191, 232, 33, 152, 1, // row 4
313            145, 111, 142, 122, 181, // row 5 -- this is the medoid
314        ];
315
316        let data = rowmajor::Owned::<u8>::try_from_data(data.into(), 6, 5).unwrap();
317        let expected: Vec<u8> = data.row(5).into();
318        (data, expected)
319    }
320
321    #[test]
322    fn test_u8() {
323        // No Rows
324        let x = rowmajor::Owned::<u8>::from_element(0, 10, 0u8);
325        assert_eq!(u8::compute_medoid(x.as_view()), vec![0u8; x.ncols()]);
326
327        // No Cols
328        let x = rowmajor::Owned::<u8>::from_element(10, 0, 0u8);
329        assert_eq!(u8::compute_medoid(x.as_view()), Vec::<u8>::new());
330        let mut rng = StdRng::seed_from_u64(0x8f2f5fa0b5161acf);
331
332        // One row
333        let dist = StandardUniform;
334        for dim in 1..20 {
335            let x = rowmajor::Owned::<u8>::from_fn(1, dim, |_| dist.sample(&mut rng));
336            assert_eq!(&*u8::compute_medoid(x.as_view()), x.row(0));
337        }
338
339        // Example dataset
340        let (data, expected) = example_dataset_u8();
341        let m = u8::compute_medoid(data.as_view());
342        assert_eq!(m, expected);
343    }
344
345    // This is a test for the i8 medoid function. Each entry is between -128 and 127.
346    fn example_dataset_i8() -> (rowmajor::Owned<i8>, Vec<i8>) {
347        let data: Vec<i8> = vec![
348            -76, 87, 90, 76, 64, // row 0
349            -49, -73, -112, -39, 127, // row 1
350            39, 105, 13, -95, -98, // row 2
351            -37, -75, -13, 108, 2, // row 3
352            -37, -75, -13, 108, 2, // row 4
353            17, -17, 14, -6, 53, // row 5 -- this is the medoid
354        ];
355
356        let data = rowmajor::Owned::<i8>::try_from_data(data.into(), 6, 5).unwrap();
357        let expected: Vec<i8> = data.row(5).into();
358        (data, expected)
359    }
360
361    #[test]
362    fn test_i8() {
363        // No Rows
364        let x = rowmajor::Owned::<i8>::from_element(0, 10, 0i8);
365        assert_eq!(i8::compute_medoid(x.as_view()), vec![0i8; x.ncols()]);
366
367        // No Cols
368        let x = rowmajor::Owned::<i8>::from_element(10, 0, 0i8);
369        assert_eq!(i8::compute_medoid(x.as_view()), Vec::<i8>::new());
370
371        let mut rng = StdRng::seed_from_u64(0x8f2f5fa0b5161acf);
372
373        // One row
374        let dist = StandardUniform;
375        for dim in 1..20 {
376            let x = rowmajor::Owned::<i8>::from_fn(1, dim, |_| dist.sample(&mut rng));
377            assert_eq!(&*i8::compute_medoid(x.as_view()), x.row(0));
378        }
379
380        // Example dataset
381        let (data, expected) = example_dataset_i8();
382        let m = i8::compute_medoid(data.as_view());
383        assert_eq!(m, expected);
384    }
385}