Skip to main content

diskann_utils/sampling/
latin_hypercube.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use crate::views::rowmajor::{self, Matrix, MatrixMut};
7use rand::{rngs::StdRng, Rng, SeedableRng};
8
9/// Return multiple rows sampled using Latin Hypercube Sampling in `data` that aproximetely uniformly distributed.
10/// This makes the assumtion that the data is uniformly distributed.
11pub trait SampleLatinHyperCube: Sized + Copy + Default {
12    fn sample_latin_hypercube(
13        data: rowmajor::Ref<'_, Self>,
14        num_samples: usize,
15        seed: Option<u64>,
16    ) -> rowmajor::Owned<Self>;
17}
18
19impl<T: Sized + Copy + Default> SampleLatinHyperCube for T {
20    fn sample_latin_hypercube(
21        data: rowmajor::Ref<'_, Self>,
22        num_samples: usize,
23        seed: Option<u64>,
24    ) -> rowmajor::Owned<Self> {
25        let nrows = data.nrows();
26        let ncols = data.ncols();
27        if ncols == 0 || nrows == 0 {
28            return rowmajor::Owned::from_element(num_samples, ncols, T::default());
29        }
30
31        let seed = seed.unwrap_or(0xaf2f5fa0b5161acf);
32        let mut rng = StdRng::seed_from_u64(seed);
33        let mut result = rowmajor::Owned::from_element(num_samples, ncols, T::default());
34
35        // sample a random partitions down the diagonal
36        for (s, res) in result.rows_mut().enumerate() {
37            for (idx, val) in res.iter_mut().enumerate() {
38                let step = nrows / num_samples;
39                let value = data
40                    .get_row(rng.random_range(s * step..(s + 1) * step))
41                    .unwrap()
42                    .get(idx)
43                    .unwrap();
44                *val = *value;
45            }
46        }
47
48        // shuffle the dimensions between the vectors for random sampling
49        for start_idx in 0..num_samples {
50            for dim_idx in 0..ncols {
51                let swap_idx = rng.random_range(0..num_samples);
52                let swap = *result.element(start_idx, dim_idx);
53                *result.element_mut(start_idx, dim_idx) = *result.element(swap_idx, dim_idx);
54                *result.element_mut(swap_idx, dim_idx) = swap;
55            }
56        }
57
58        result
59    }
60}
61
62///////////
63// Tests //
64///////////
65
66#[cfg(not(miri))]
67#[cfg(test)]
68mod tests {
69    use std::fmt::Display;
70
71    use crate::assert_contains;
72
73    use diskann_vector::conversion::CastFromSlice;
74    use half::f16;
75    use rand::{
76        distr::{Distribution, StandardUniform},
77        rngs::StdRng,
78        SeedableRng,
79    };
80
81    use super::*;
82
83    fn example_dataset() -> rowmajor::Owned<f32> {
84        let data: Vec<f32> = vec![
85            // row 0
86            0.203688,
87            0.841956,
88            0.855665,
89            0.801917,
90            0.754536,
91            // row 1
92            0.312881,
93            0.217382,
94            0.0644115,
95            0.348708,
96            0.999495,
97            // row 2
98            0.657741,
99            0.914681,
100            0.555228,
101            0.13253,
102            0.118615,
103            // row 3
104            0.356464,
105            0.207449,
106            0.452471,
107            0.925219,
108            0.508498,
109            // row 4
110            0.749786,
111            0.90786,
112            0.129618,
113            0.597719,
114            0.000622153,
115            // row 5 -- this is the medoid
116            0.569517,
117            0.435447,
118            0.558136,
119            0.480974,
120            0.711425,
121            // row 6
122            0.896353,
123            0.275053,
124            0.0427179,
125            0.660916,
126            0.464851,
127            // row 7
128            0.558689,
129            0.596543,
130            0.740983,
131            0.122136,
132            0.453822,
133            // row 8
134            0.526895,
135            0.492643,
136            0.0951115,
137            0.495487,
138            0.446127,
139            // row 9
140            0.454093,
141            0.160239,
142            0.924585,
143            0.901708,
144            0.329328,
145        ];
146
147        rowmajor::Owned::<f32>::try_from_data(data.into(), 10, 5).unwrap()
148    }
149
150    fn example_dataset_u8() -> rowmajor::Owned<u8> {
151        let data: Vec<u8> = vec![
152            52, 215, 218, 204, 192, // row 0
153            79, 55, 16, 89, 255, // row 1
154            167, 233, 141, 33, 30, // row 2
155            91, 53, 115, 236, 130, // row 3
156            191, 232, 33, 152, 1, // row 4
157            145, 111, 142, 122, 181, // row 5 -- this is the medoid
158        ];
159
160        rowmajor::Owned::<u8>::try_from_data(data.into(), 6, 5).unwrap()
161    }
162
163    // This is a test for the i8 function. Each entry is between -128 and 127.
164    fn example_dataset_i8() -> rowmajor::Owned<i8> {
165        let data: Vec<i8> = vec![
166            -76, 87, 90, 76, 64, // row 0
167            -49, -73, -112, -39, 127, // row 1
168            39, 105, 13, -95, -98, // row 2
169            -37, -75, -13, 108, 2, // row 3
170            -37, -75, -13, 108, 2, // row 4
171            17, -17, 14, -6, 53, // row 5 -- this is the medoid
172        ];
173
174        rowmajor::Owned::<i8>::try_from_data(data.into(), 6, 5).unwrap()
175    }
176
177    fn test_for_type<T>(data: rowmajor::Owned<T>)
178    where
179        T: SampleLatinHyperCube + PartialEq + std::fmt::Debug + Display,
180        StandardUniform: Distribution<T>,
181    {
182        // No Rows
183        let x = rowmajor::Owned::<T>::from_element(0, 10, T::default());
184        assert_eq!(
185            T::sample_latin_hypercube(x.as_view(), 1, None),
186            rowmajor::Owned::<T>::from_element(1, x.ncols(), T::default())
187        );
188
189        // No Cols0
190        let x = rowmajor::Owned::<T>::from_element(1, 0, T::default());
191        assert_eq!(
192            T::sample_latin_hypercube(x.as_view(), 1, None),
193            rowmajor::Owned::<T>::from_element(1, x.ncols(), T::default())
194        );
195
196        let mut rng: StdRng = StdRng::seed_from_u64(0xaf2f5fa0b5161acf);
197
198        // One row
199        let dist = StandardUniform;
200        for dim in 1..20 {
201            let x = rowmajor::Owned::<T>::from_fn(1, dim, |_| dist.sample(&mut rng));
202            assert_eq!(
203                T::sample_latin_hypercube(x.as_view(), 1, None),
204                rowmajor::Owned::<T>::try_from_data(x.row(0).to_vec().into_boxed_slice(), 1, dim)
205                    .unwrap()
206            );
207        }
208
209        // Example dataset
210        let starts = T::sample_latin_hypercube(data.as_view(), 2, None);
211        for s in starts.rows() {
212            for (col, &val) in s.iter().enumerate() {
213                let col_vals: Vec<T> = (0..data.nrows())
214                    .map(|row| {
215                        *data
216                            .get_row(row)
217                            .expect("Row must exist")
218                            .get(col)
219                            .expect("Column must exist")
220                    })
221                    .collect();
222                assert_contains!(
223                    col_vals,
224                    val,
225                    "Value {} in column {} not found in data",
226                    val,
227                    col
228                );
229            }
230        }
231    }
232
233    #[test]
234    fn test_f32() {
235        test_for_type(example_dataset())
236    }
237
238    #[test]
239    fn test_f16() {
240        let data = example_dataset();
241        let mut data_f16 =
242            rowmajor::Owned::<f16>::from_element(data.nrows(), data.ncols(), f16::default());
243        data_f16.as_mut_slice().cast_from_slice(data.as_slice());
244        test_for_type(data_f16);
245    }
246
247    #[test]
248    fn test_u8() {
249        test_for_type(example_dataset_u8());
250    }
251
252    #[test]
253    fn test_i8() {
254        test_for_type(example_dataset_i8());
255    }
256}