Skip to main content

diskann_utils/sampling/
random.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use rand::{rngs::StdRng, Rng};
7use rand_distr::StandardNormal;
8
9pub trait RoundFromf32 {
10    fn round_from_f32(x: f32) -> Self;
11}
12
13impl RoundFromf32 for f32 {
14    fn round_from_f32(x: f32) -> Self {
15        x
16    }
17}
18impl RoundFromf32 for i8 {
19    fn round_from_f32(x: f32) -> Self {
20        x.round() as i8
21    }
22}
23impl RoundFromf32 for u8 {
24    fn round_from_f32(x: f32) -> Self {
25        x.round() as u8
26    }
27}
28impl RoundFromf32 for half::f16 {
29    fn round_from_f32(x: f32) -> Self {
30        half::f16::from_f32(x)
31    }
32}
33
34pub trait WithApproximateNorm: Sized {
35    fn with_approximate_norm(dim: usize, norm: f32, rng: &mut StdRng) -> Vec<Self>;
36}
37
38impl WithApproximateNorm for f32 {
39    fn with_approximate_norm(dim: usize, norm: f32, rng: &mut StdRng) -> Vec<Self> {
40        generate_random_vector_with_norm_signed(dim, norm, true, rng, |x: f32| x)
41    }
42}
43
44impl WithApproximateNorm for half::f16 {
45    fn with_approximate_norm(dim: usize, norm: f32, rng: &mut StdRng) -> Vec<Self> {
46        // Small QOL improvement, `diskann_wide::cast_f32_to_f16` works under `Miri` while `half::f16::from_f32`
47        // does not.
48        generate_random_vector_with_norm_signed(dim, norm, true, rng, diskann_wide::cast_f32_to_f16)
49    }
50}
51
52impl WithApproximateNorm for u8 {
53    fn with_approximate_norm(dim: usize, norm: f32, rng: &mut StdRng) -> Vec<Self> {
54        generate_random_vector_with_norm_signed(dim, norm, false, rng, |x| x as u8)
55    }
56}
57
58impl WithApproximateNorm for i8 {
59    fn with_approximate_norm(dim: usize, norm: f32, rng: &mut StdRng) -> Vec<Self> {
60        generate_random_vector_with_norm_signed(dim, norm, true, rng, |x| x as i8)
61    }
62}
63
64// This function uses StandardNormal distribution. StandardNormal creates uniformly
65// distributed points on sphere surface, making the graph easier to navigate.
66fn generate_random_vector_with_norm_signed<T, F>(
67    dim: usize,
68    norm: f32,
69    signed: bool,
70    rng: &mut StdRng,
71    f: F,
72) -> Vec<T>
73where
74    F: Fn(f32) -> T,
75{
76    let mut vec: Vec<f32> = (0..dim).map(|_| rng.sample(StandardNormal)).collect();
77    let current_norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
78    let scale = norm / current_norm;
79    if signed {
80        vec.iter_mut().for_each(|x| *x *= scale);
81    } else {
82        vec.iter_mut().for_each(|x| *x = (*x * scale).abs());
83    };
84    vec.into_iter().map(f).collect()
85}
86
87#[cfg(not(miri))]
88#[cfg(test)]
89mod tests {
90    use super::*;
91    use rand::SeedableRng;
92    use rstest::rstest;
93
94    #[rstest]
95    #[case(1, 0.01)]
96    #[case(100, 0.01)]
97    #[case(171, 5.0)]
98    #[case(1024, 100.7)]
99    fn test_generate_random_vector_with_norm_f32(#[case] dim: usize, #[case] norm: f32) {
100        let seed = 42;
101        let mut rng = StdRng::seed_from_u64(seed);
102        let vec: Vec<f32> = WithApproximateNorm::with_approximate_norm(dim, norm, &mut rng);
103        let computed_norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
104        let tolerance = 1e-5;
105        assert!((computed_norm - norm).abs() / norm < tolerance);
106    }
107
108    #[rstest]
109    #[case(1, 0.01)]
110    #[case(100, 0.01)]
111    #[case(171, 5.0)]
112    #[case(1024, 100.7)]
113    fn test_generate_random_vector_with_norm_half_f16(#[case] dim: usize, #[case] norm: f32) {
114        let seed = 42;
115        let mut rng = StdRng::seed_from_u64(seed);
116        let vec: Vec<half::f16> = WithApproximateNorm::with_approximate_norm(dim, norm, &mut rng);
117        let computed_norm: f32 = vec
118            .iter()
119            .map(|x| {
120                let val: f32 = x.to_f32();
121                val * val
122            })
123            .sum::<f32>()
124            .sqrt();
125        let tolerance = 1e-2; // half precision
126        assert!((computed_norm - norm).abs() / norm < tolerance);
127    }
128
129    #[rstest]
130    #[case(17, 50.0)]
131    #[case(1024, 1007.0)]
132    fn test_generate_random_vector_with_norm_u8(#[case] dim: usize, #[case] norm: f32) {
133        let seed = 42;
134        let mut rng = StdRng::seed_from_u64(seed);
135        let vec: Vec<u8> = WithApproximateNorm::with_approximate_norm(dim, norm, &mut rng);
136        let computed_norm: f32 = vec
137            .iter()
138            .map(|&x| {
139                let val: f32 = x as f32;
140                val * val
141            })
142            .sum::<f32>()
143            .sqrt();
144        let tolerance = 1e-1; // due to quantization
145        assert!((computed_norm - norm).abs() / norm < tolerance);
146    }
147
148    #[rstest]
149    #[case(17, 50.0)]
150    #[case(1024, 1007.0)]
151    fn test_generate_random_vector_with_norm_i8(#[case] dim: usize, #[case] norm: f32) {
152        let seed = 42;
153        let mut rng = StdRng::seed_from_u64(seed);
154        let vec: Vec<i8> = WithApproximateNorm::with_approximate_norm(dim, norm, &mut rng);
155        let computed_norm: f32 = vec
156            .iter()
157            .map(|&x| {
158                let val: f32 = x as f32;
159                val * val
160            })
161            .sum::<f32>()
162            .sqrt();
163        let tolerance = 1e-1; // due to quantization
164        assert!((computed_norm - norm).abs() / norm < tolerance);
165    }
166
167    #[rstest]
168    #[case(3.6f32, 4i8)]
169    #[case(2.3f32, 2i8)]
170    #[case(-1.5f32, -2i8)]
171    fn test_round_f32_to_i8(#[case] input: f32, #[case] expected: i8) {
172        let result: i8 = RoundFromf32::round_from_f32(input);
173        assert_eq!(result, expected);
174    }
175
176    #[rstest]
177    #[case(3.6f32, 4u8)]
178    #[case(2.3f32, 2u8)]
179    #[case(-1.5f32, 0u8)]
180    fn test_round_f32_to_u8(#[case] input: f32, #[case] expected: u8) {
181        let result: u8 = RoundFromf32::round_from_f32(input);
182        assert_eq!(result, expected);
183    }
184
185    #[rstest]
186    #[case(3.6f32, half::f16::from_f32(3.6f32))]
187    #[case(2.3f32, half::f16::from_f32(2.3f32))]
188    #[case(-1.5f32, half::f16::from_f32(-1.5f32))]
189    fn test_round_f32_to_f16(#[case] input: f32, #[case] expected: half::f16) {
190        let result: half::f16 = RoundFromf32::round_from_f32(input);
191        assert_eq!(result, expected);
192    }
193
194    #[rstest]
195    #[case(3.6f32, 3.6f32)]
196    #[case(2.3f32, 2.3f32)]
197    #[case(-1.5f32, -1.5f32)]
198    fn test_round_f32_to_f32(#[case] input: f32, #[case] expected: f32) {
199        let result: f32 = RoundFromf32::round_from_f32(input);
200        assert_eq!(result, expected);
201    }
202
203    /// Test that generated points are evenly distributed on a circle.
204    ///
205    /// **Testing methodology:**
206    /// 1. Split the circle into 36 buckets (signed) or 9 buckets (unsigned), each covering 10 degrees
207    /// 2. Generate points and count how many fall into each angular bucket
208    /// 3. Check that each bucket's count is within `tolerance_sigmas × σ` of the expected count,
209    ///    where σ = sqrt(expected) is the statistical noise for random sampling
210    /// 4. Fail if any bucket deviates too much (indicates clustering instead of uniform distribution)
211    ///
212    /// **Tolerance levels:**
213    ///   - tolerance_sigmas = 1.0 → Very strict, only allows ±1σ deviation (about 68% of buckets would naturally fall within this)
214    ///   - tolerance_sigmas = 3.0 → Moderate, allows ±3σ deviation (99.7% would naturally fall within this)
215    ///   - tolerance_sigmas = 6.0 → Very lenient, allows ±6σ deviation (99.9997% would naturally fall within this)
216    #[rstest]
217    #[case(true, 500, 3.0, 42)]
218    #[case(true, 500, 3.0, 43)]
219    #[case(true, 500, 3.0, 44)]
220    #[case(false, 500, 3.0, 42)]
221    #[case(false, 500, 3.0, 43)]
222    #[case(false, 500, 3.0, 44)]
223    fn test_generate_random_vector_with_norm_signed_produces_uniform_distribution_on_circle(
224        #[case] signed: bool,
225        #[case] expected_per_bucket: usize,
226        #[case] tolerance_sigmas: f32,
227        #[case] seed: u64,
228    ) {
229        let dim = 2;
230        let norm = 1.0;
231        let mut rng = StdRng::seed_from_u64(seed);
232
233        // Step 1: Pick number of buckets and calculate samples
234        let num_buckets = if signed { 36 } else { 9 };
235        let num_samples = num_buckets * expected_per_bucket;
236
237        // Generate samples
238        let samples: Vec<Vec<f32>> = (0..num_samples)
239            .map(|_| generate_random_vector_with_norm_signed(dim, norm, signed, &mut rng, |x| x))
240            .collect();
241
242        // Step 2: Count hits per bucket
243        let mut counts = vec![0usize; num_buckets];
244
245        for sample in &samples {
246            let theta = sample[1].atan2(sample[0]); // atan2(y, x) returns [-π, π]
247
248            // Map to bucket: floor(θ / 2π × buckets)
249            let bucket = if signed {
250                // Full circle [0, 2π) → [0, 36)
251                let normalized_theta = if theta < 0.0 {
252                    theta + 2.0 * std::f32::consts::PI
253                } else {
254                    theta
255                };
256                ((normalized_theta / (2.0 * std::f32::consts::PI)) * num_buckets as f32).floor()
257                    as usize
258                    % num_buckets
259            } else {
260                // First quadrant [0, π/2) → [0, 9)
261                ((theta / (std::f32::consts::PI / 2.0)) * num_buckets as f32).floor() as usize
262            };
263
264            counts[bucket] += 1;
265        }
266
267        // Step 3: Check each bucket is within tolerance_sigmas × σ
268        // Noise per bucket: σ ≈ sqrt(expected)
269        // Threshold: |observed - expected| / expected > tolerance_sigmas / sqrt(expected)
270        let sigma = (expected_per_bucket as f32).sqrt();
271        let threshold = tolerance_sigmas / sigma;
272
273        let failed_count = counts
274            .iter()
275            .filter(|&&observed| {
276                let deviation = (observed as f32 - expected_per_bucket as f32).abs()
277                    / expected_per_bucket as f32;
278                deviation > threshold
279            })
280            .count();
281
282        assert_eq!(
283            failed_count,
284            0,
285            "Distribution not uniform: {} out of {} bucket(s) had point counts that deviated more than {}σ from expected. \
286             This indicates the generator is producing clustered points instead of evenly distributed points on the circle surface.",
287            failed_count,
288            num_buckets,
289            tolerance_sigmas
290        );
291    }
292}