1use crate::views::rowmajor::{self, Matrix};
7use diskann_vector::{conversion::CastFromSlice, distance::SquaredL2, PureDistanceFunction};
8use half::f16;
9
10pub 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#[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 0.203688,
180 0.841956,
181 0.855665,
182 0.801917,
183 0.754536,
184 0.312881,
186 0.217382,
187 0.0644115,
188 0.348708,
189 0.999495,
190 0.657741,
192 0.914681,
193 0.555228,
194 0.13253,
195 0.118615,
196 0.356464,
198 0.207449,
199 0.452471,
200 0.925219,
201 0.508498,
202 0.749786,
204 0.90786,
205 0.129618,
206 0.597719,
207 0.000622153,
208 0.569517,
210 0.435447,
211 0.558136,
212 0.480974,
213 0.711425,
214 0.896353,
216 0.275053,
217 0.0427179,
218 0.660916,
219 0.464851,
220 0.558689,
222 0.596543,
223 0.740983,
224 0.122136,
225 0.453822,
226 0.526895,
228 0.492643,
229 0.0951115,
230 0.495487,
231 0.446127,
232 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 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 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 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 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 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 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 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 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, 79, 55, 16, 89, 255, 167, 233, 141, 33, 30, 91, 53, 115, 236, 130, 191, 232, 33, 152, 1, 145, 111, 142, 122, 181, ];
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 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 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 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 let (data, expected) = example_dataset_u8();
341 let m = u8::compute_medoid(data.as_view());
342 assert_eq!(m, expected);
343 }
344
345 fn example_dataset_i8() -> (rowmajor::Owned<i8>, Vec<i8>) {
347 let data: Vec<i8> = vec![
348 -76, 87, 90, 76, 64, -49, -73, -112, -39, 127, 39, 105, 13, -95, -98, -37, -75, -13, 108, 2, -37, -75, -13, 108, 2, 17, -17, 14, -6, 53, ];
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 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 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 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 let (data, expected) = example_dataset_i8();
382 let m = i8::compute_medoid(data.as_view());
383 assert_eq!(m, expected);
384 }
385}