1use 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 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
64fn 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; 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; 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; 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 #[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 let num_buckets = if signed { 36 } else { 9 };
235 let num_samples = num_buckets * expected_per_bucket;
236
237 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 let mut counts = vec![0usize; num_buckets];
244
245 for sample in &samples {
246 let theta = sample[1].atan2(sample[0]); let bucket = if signed {
250 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 ((theta / (std::f32::consts::PI / 2.0)) * num_buckets as f32).floor() as usize
262 };
263
264 counts[bucket] += 1;
265 }
266
267 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}