1use crate::views::rowmajor::{self, Matrix, MatrixMut};
7use rand::{rngs::StdRng, Rng, SeedableRng};
8
9pub 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 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 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#[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 0.203688,
87 0.841956,
88 0.855665,
89 0.801917,
90 0.754536,
91 0.312881,
93 0.217382,
94 0.0644115,
95 0.348708,
96 0.999495,
97 0.657741,
99 0.914681,
100 0.555228,
101 0.13253,
102 0.118615,
103 0.356464,
105 0.207449,
106 0.452471,
107 0.925219,
108 0.508498,
109 0.749786,
111 0.90786,
112 0.129618,
113 0.597719,
114 0.000622153,
115 0.569517,
117 0.435447,
118 0.558136,
119 0.480974,
120 0.711425,
121 0.896353,
123 0.275053,
124 0.0427179,
125 0.660916,
126 0.464851,
127 0.558689,
129 0.596543,
130 0.740983,
131 0.122136,
132 0.453822,
133 0.526895,
135 0.492643,
136 0.0951115,
137 0.495487,
138 0.446127,
139 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, 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, ];
159
160 rowmajor::Owned::<u8>::try_from_data(data.into(), 6, 5).unwrap()
161 }
162
163 fn example_dataset_i8() -> rowmajor::Owned<i8> {
165 let data: Vec<i8> = vec![
166 -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, ];
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 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 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 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 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}