/rust/registry/src/index.crates.io-1949cf8c6b5b557f/rand_distr-0.5.1/src/hypergeometric.rs
Line | Count | Source |
1 | | //! The hypergeometric distribution `Hypergeometric(N, K, n)`. |
2 | | |
3 | | use crate::Distribution; |
4 | | use core::fmt; |
5 | | #[allow(unused_imports)] |
6 | | use num_traits::Float; |
7 | | use rand::distr::uniform::Uniform; |
8 | | use rand::Rng; |
9 | | |
10 | | #[derive(Clone, Copy, Debug, PartialEq)] |
11 | | #[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] |
12 | | enum SamplingMethod { |
13 | | InverseTransform { |
14 | | initial_p: f64, |
15 | | initial_x: i64, |
16 | | }, |
17 | | RejectionAcceptance { |
18 | | m: f64, |
19 | | a: f64, |
20 | | lambda_l: f64, |
21 | | lambda_r: f64, |
22 | | x_l: f64, |
23 | | x_r: f64, |
24 | | p1: f64, |
25 | | p2: f64, |
26 | | p3: f64, |
27 | | }, |
28 | | } |
29 | | |
30 | | /// The [hypergeometric distribution](https://en.wikipedia.org/wiki/Hypergeometric_distribution) `Hypergeometric(N, K, n)`. |
31 | | /// |
32 | | /// This is the distribution of successes in samples of size `n` drawn without |
33 | | /// replacement from a population of size `N` containing `K` success states. |
34 | | /// |
35 | | /// See the [binomial distribution](crate::Binomial) for the analogous distribution |
36 | | /// for sampling with replacement. It is a good approximation when the population |
37 | | /// size is much larger than the sample size. |
38 | | /// |
39 | | /// # Density function |
40 | | /// |
41 | | /// `f(k) = binomial(K, k) * binomial(N-K, n-k) / binomial(N, n)`, |
42 | | /// where `binomial(a, b) = a! / (b! * (a - b)!)`. |
43 | | /// |
44 | | /// # Plot |
45 | | /// |
46 | | /// The following plot of the hypergeometric distribution illustrates the probability of drawing |
47 | | /// `k` successes in `n = 10` draws from a population of `N = 50` items, of which either `K = 12` |
48 | | /// or `K = 35` are successes. |
49 | | /// |
50 | | ///  |
51 | | /// |
52 | | /// # Example |
53 | | /// ``` |
54 | | /// use rand_distr::{Distribution, Hypergeometric}; |
55 | | /// |
56 | | /// let hypergeo = Hypergeometric::new(60, 24, 7).unwrap(); |
57 | | /// let v = hypergeo.sample(&mut rand::rng()); |
58 | | /// println!("{} is from a hypergeometric distribution", v); |
59 | | /// ``` |
60 | | #[derive(Copy, Clone, Debug, PartialEq)] |
61 | | #[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] |
62 | | pub struct Hypergeometric { |
63 | | n1: u64, |
64 | | n2: u64, |
65 | | k: u64, |
66 | | offset_x: i64, |
67 | | sign_x: i64, |
68 | | sampling_method: SamplingMethod, |
69 | | } |
70 | | |
71 | | /// Error type returned from [`Hypergeometric::new`]. |
72 | | #[derive(Clone, Copy, Debug, PartialEq, Eq)] |
73 | | pub enum Error { |
74 | | /// `total_population_size` is too large, causing floating point underflow. |
75 | | PopulationTooLarge, |
76 | | /// `population_with_feature > total_population_size`. |
77 | | ProbabilityTooLarge, |
78 | | /// `sample_size > total_population_size`. |
79 | | SampleSizeTooLarge, |
80 | | } |
81 | | |
82 | | impl fmt::Display for Error { |
83 | 0 | fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { |
84 | 0 | f.write_str(match self { |
85 | | Error::PopulationTooLarge => { |
86 | 0 | "total_population_size is too large causing underflow in geometric distribution" |
87 | | } |
88 | | Error::ProbabilityTooLarge => { |
89 | 0 | "population_with_feature > total_population_size in geometric distribution" |
90 | | } |
91 | | Error::SampleSizeTooLarge => { |
92 | 0 | "sample_size > total_population_size in geometric distribution" |
93 | | } |
94 | | }) |
95 | 0 | } |
96 | | } |
97 | | |
98 | | #[cfg(feature = "std")] |
99 | | impl std::error::Error for Error {} |
100 | | |
101 | | // evaluate fact(numerator.0)*fact(numerator.1) / fact(denominator.0)*fact(denominator.1) |
102 | 0 | fn fraction_of_products_of_factorials(numerator: (u64, u64), denominator: (u64, u64)) -> f64 { |
103 | 0 | let min_top = u64::min(numerator.0, numerator.1); |
104 | 0 | let min_bottom = u64::min(denominator.0, denominator.1); |
105 | | // the factorial of this will cancel out: |
106 | 0 | let min_all = u64::min(min_top, min_bottom); |
107 | | |
108 | 0 | let max_top = u64::max(numerator.0, numerator.1); |
109 | 0 | let max_bottom = u64::max(denominator.0, denominator.1); |
110 | 0 | let max_all = u64::max(max_top, max_bottom); |
111 | | |
112 | 0 | let mut result = 1.0; |
113 | 0 | for i in (min_all + 1)..=max_all { |
114 | 0 | if i <= min_top { |
115 | 0 | result *= i as f64; |
116 | 0 | } |
117 | | |
118 | 0 | if i <= min_bottom { |
119 | 0 | result /= i as f64; |
120 | 0 | } |
121 | | |
122 | 0 | if i <= max_top { |
123 | 0 | result *= i as f64; |
124 | 0 | } |
125 | | |
126 | 0 | if i <= max_bottom { |
127 | 0 | result /= i as f64; |
128 | 0 | } |
129 | | } |
130 | | |
131 | 0 | result |
132 | 0 | } |
133 | | |
134 | | const LOGSQRT2PI: f64 = 0.91893853320467274178; // log(sqrt(2*pi)) |
135 | | |
136 | 0 | fn ln_of_factorial(v: f64) -> f64 { |
137 | | // the paper calls for ln(v!), but also wants to pass in fractions, |
138 | | // so we need to use Stirling's approximation to fill in the gaps: |
139 | | |
140 | | // shift v by 3, because Stirling is bad for small values |
141 | 0 | let v_3 = v + 3.0; |
142 | 0 | let ln_fac = (v_3 + 0.5) * v_3.ln() - v_3 + LOGSQRT2PI + 1.0 / (12.0 * v_3); |
143 | | // make the correction for the shift |
144 | 0 | ln_fac - ((v + 3.0) * (v + 2.0) * (v + 1.0)).ln() |
145 | 0 | } |
146 | | |
147 | | impl Hypergeometric { |
148 | | /// Constructs a new `Hypergeometric` with the shape parameters |
149 | | /// `N = total_population_size`, |
150 | | /// `K = population_with_feature`, |
151 | | /// `n = sample_size`. |
152 | | #[allow(clippy::many_single_char_names)] // Same names as in the reference. |
153 | 0 | pub fn new( |
154 | 0 | total_population_size: u64, |
155 | 0 | population_with_feature: u64, |
156 | 0 | sample_size: u64, |
157 | 0 | ) -> Result<Self, Error> { |
158 | 0 | if population_with_feature > total_population_size { |
159 | 0 | return Err(Error::ProbabilityTooLarge); |
160 | 0 | } |
161 | | |
162 | 0 | if sample_size > total_population_size { |
163 | 0 | return Err(Error::SampleSizeTooLarge); |
164 | 0 | } |
165 | | |
166 | | // set-up constants as function of original parameters |
167 | 0 | let n = total_population_size; |
168 | 0 | let (mut sign_x, mut offset_x) = (1, 0); |
169 | 0 | let (n1, n2) = { |
170 | | // switch around success and failure states if necessary to ensure n1 <= n2 |
171 | 0 | let population_without_feature = n - population_with_feature; |
172 | 0 | if population_with_feature > population_without_feature { |
173 | 0 | sign_x = -1; |
174 | 0 | offset_x = sample_size as i64; |
175 | 0 | (population_without_feature, population_with_feature) |
176 | | } else { |
177 | 0 | (population_with_feature, population_without_feature) |
178 | | } |
179 | | }; |
180 | | // when sampling more than half the total population, take the smaller |
181 | | // group as sampled instead (we can then return n1-x instead). |
182 | | // |
183 | | // Note: the boundary condition given in the paper is `sample_size < n / 2`; |
184 | | // we're deviating here, because when n is even, it doesn't matter whether |
185 | | // we switch here or not, but when n is odd `n/2 < n - n/2`, so switching |
186 | | // when `k == n/2`, we'd actually be taking the _larger_ group as sampled. |
187 | 0 | let k = if sample_size <= n / 2 { |
188 | 0 | sample_size |
189 | | } else { |
190 | 0 | offset_x += n1 as i64 * sign_x; |
191 | 0 | sign_x *= -1; |
192 | 0 | n - sample_size |
193 | | }; |
194 | | |
195 | | // Algorithm H2PE has bounded runtime only if `M - max(0, k-n2) >= 10`, |
196 | | // where `M` is the mode of the distribution. |
197 | | // Use algorithm HIN for the remaining parameter space. |
198 | | // |
199 | | // Voratas Kachitvichyanukul and Bruce W. Schmeiser. 1985. Computer |
200 | | // generation of hypergeometric random variates. |
201 | | // J. Statist. Comput. Simul. Vol.22 (August 1985), 127-145 |
202 | | // https://www.researchgate.net/publication/233212638 |
203 | | const HIN_THRESHOLD: f64 = 10.0; |
204 | 0 | let m = ((k + 1) as f64 * (n1 + 1) as f64 / (n + 2) as f64).floor(); |
205 | 0 | let sampling_method = if m - f64::max(0.0, k as f64 - n2 as f64) < HIN_THRESHOLD { |
206 | 0 | let (initial_p, initial_x) = if k < n2 { |
207 | 0 | ( |
208 | 0 | fraction_of_products_of_factorials((n2, n - k), (n, n2 - k)), |
209 | 0 | 0, |
210 | 0 | ) |
211 | | } else { |
212 | 0 | ( |
213 | 0 | fraction_of_products_of_factorials((n1, k), (n, k - n2)), |
214 | 0 | (k - n2) as i64, |
215 | 0 | ) |
216 | | }; |
217 | | |
218 | 0 | if initial_p <= 0.0 || !initial_p.is_finite() { |
219 | 0 | return Err(Error::PopulationTooLarge); |
220 | 0 | } |
221 | | |
222 | 0 | SamplingMethod::InverseTransform { |
223 | 0 | initial_p, |
224 | 0 | initial_x, |
225 | 0 | } |
226 | | } else { |
227 | 0 | let a = ln_of_factorial(m) |
228 | 0 | + ln_of_factorial(n1 as f64 - m) |
229 | 0 | + ln_of_factorial(k as f64 - m) |
230 | 0 | + ln_of_factorial((n2 - k) as f64 + m); |
231 | | |
232 | 0 | let numerator = (n - k) as f64 * k as f64 * n1 as f64 * n2 as f64; |
233 | 0 | let denominator = (n - 1) as f64 * n as f64 * n as f64; |
234 | 0 | let d = 1.5 * (numerator / denominator).sqrt() + 0.5; |
235 | | |
236 | 0 | let x_l = m - d + 0.5; |
237 | 0 | let x_r = m + d + 0.5; |
238 | | |
239 | 0 | let k_l = f64::exp( |
240 | 0 | a - ln_of_factorial(x_l) |
241 | 0 | - ln_of_factorial(n1 as f64 - x_l) |
242 | 0 | - ln_of_factorial(k as f64 - x_l) |
243 | 0 | - ln_of_factorial((n2 - k) as f64 + x_l), |
244 | | ); |
245 | 0 | let k_r = f64::exp( |
246 | 0 | a - ln_of_factorial(x_r - 1.0) |
247 | 0 | - ln_of_factorial(n1 as f64 - x_r + 1.0) |
248 | 0 | - ln_of_factorial(k as f64 - x_r + 1.0) |
249 | 0 | - ln_of_factorial((n2 - k) as f64 + x_r - 1.0), |
250 | | ); |
251 | | |
252 | 0 | let numerator = x_l * ((n2 - k) as f64 + x_l); |
253 | 0 | let denominator = (n1 as f64 - x_l + 1.0) * (k as f64 - x_l + 1.0); |
254 | 0 | let lambda_l = -((numerator / denominator).ln()); |
255 | | |
256 | 0 | let numerator = (n1 as f64 - x_r + 1.0) * (k as f64 - x_r + 1.0); |
257 | 0 | let denominator = x_r * ((n2 - k) as f64 + x_r); |
258 | 0 | let lambda_r = -((numerator / denominator).ln()); |
259 | | |
260 | | // the paper literally gives `p2 + kL/lambdaL` where it (probably) |
261 | | // should have been `p2 <- p1 + kL/lambdaL`; another print error?! |
262 | 0 | let p1 = 2.0 * d; |
263 | 0 | let p2 = p1 + k_l / lambda_l; |
264 | 0 | let p3 = p2 + k_r / lambda_r; |
265 | | |
266 | 0 | SamplingMethod::RejectionAcceptance { |
267 | 0 | m, |
268 | 0 | a, |
269 | 0 | lambda_l, |
270 | 0 | lambda_r, |
271 | 0 | x_l, |
272 | 0 | x_r, |
273 | 0 | p1, |
274 | 0 | p2, |
275 | 0 | p3, |
276 | 0 | } |
277 | | }; |
278 | | |
279 | 0 | Ok(Hypergeometric { |
280 | 0 | n1, |
281 | 0 | n2, |
282 | 0 | k, |
283 | 0 | offset_x, |
284 | 0 | sign_x, |
285 | 0 | sampling_method, |
286 | 0 | }) |
287 | 0 | } |
288 | | } |
289 | | |
290 | | impl Distribution<u64> for Hypergeometric { |
291 | | #[allow(clippy::many_single_char_names)] // Same names as in the reference. |
292 | 0 | fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> u64 { |
293 | | use SamplingMethod::*; |
294 | | |
295 | | let Hypergeometric { |
296 | 0 | n1, |
297 | 0 | n2, |
298 | 0 | k, |
299 | 0 | sign_x, |
300 | 0 | offset_x, |
301 | 0 | sampling_method, |
302 | 0 | } = *self; |
303 | 0 | let x = match sampling_method { |
304 | | InverseTransform { |
305 | 0 | initial_p: mut p, |
306 | 0 | initial_x: mut x, |
307 | | } => { |
308 | 0 | let mut u = rng.random::<f64>(); |
309 | | |
310 | | // the paper erroneously uses `until n < p`, which doesn't make any sense |
311 | 0 | while u > p && x < k as i64 { |
312 | 0 | u -= p; |
313 | 0 | p *= ((n1 as i64 - x) * (k as i64 - x)) as f64; |
314 | 0 | p /= ((x + 1) * (n2 as i64 - k as i64 + 1 + x)) as f64; |
315 | 0 | x += 1; |
316 | 0 | } |
317 | 0 | x |
318 | | } |
319 | | RejectionAcceptance { |
320 | 0 | m, |
321 | 0 | a, |
322 | 0 | lambda_l, |
323 | 0 | lambda_r, |
324 | 0 | x_l, |
325 | 0 | x_r, |
326 | 0 | p1, |
327 | 0 | p2, |
328 | 0 | p3, |
329 | | } => { |
330 | 0 | let distr_region_select = Uniform::new(0.0, p3).unwrap(); |
331 | | loop { |
332 | 0 | let (y, v) = loop { |
333 | 0 | let u = distr_region_select.sample(rng); |
334 | 0 | let v = rng.random::<f64>(); // for the accept/reject decision |
335 | | |
336 | 0 | if u <= p1 { |
337 | | // Region 1, central bell |
338 | 0 | let y = (x_l + u).floor(); |
339 | 0 | break (y, v); |
340 | 0 | } else if u <= p2 { |
341 | | // Region 2, left exponential tail |
342 | 0 | let y = (x_l + v.ln() / lambda_l).floor(); |
343 | 0 | if y as i64 >= i64::max(0, k as i64 - n2 as i64) { |
344 | 0 | let v = v * (u - p1) * lambda_l; |
345 | 0 | break (y, v); |
346 | 0 | } |
347 | | } else { |
348 | | // Region 3, right exponential tail |
349 | 0 | let y = (x_r - v.ln() / lambda_r).floor(); |
350 | 0 | if y as u64 <= u64::min(n1, k) { |
351 | 0 | let v = v * (u - p2) * lambda_r; |
352 | 0 | break (y, v); |
353 | 0 | } |
354 | | } |
355 | | }; |
356 | | |
357 | | // Step 4: Acceptance/Rejection Comparison |
358 | 0 | if m < 100.0 || y <= 50.0 { |
359 | | // Step 4.1: evaluate f(y) via recursive relationship |
360 | 0 | let mut f = 1.0; |
361 | 0 | if m < y { |
362 | 0 | for i in (m as u64 + 1)..=(y as u64) { |
363 | 0 | f *= (n1 - i + 1) as f64 * (k - i + 1) as f64; |
364 | 0 | f /= i as f64 * (n2 - k + i) as f64; |
365 | 0 | } |
366 | | } else { |
367 | 0 | for i in (y as u64 + 1)..=(m as u64) { |
368 | 0 | f *= i as f64 * (n2 - k + i) as f64; |
369 | 0 | f /= (n1 - i + 1) as f64 * (k - i + 1) as f64; |
370 | 0 | } |
371 | | } |
372 | | |
373 | 0 | if v <= f { |
374 | 0 | break y as i64; |
375 | 0 | } |
376 | | } else { |
377 | | // Step 4.2: Squeezing |
378 | 0 | let y1 = y + 1.0; |
379 | 0 | let ym = y - m; |
380 | 0 | let yn = n1 as f64 - y + 1.0; |
381 | 0 | let yk = k as f64 - y + 1.0; |
382 | 0 | let nk = n2 as f64 - k as f64 + y1; |
383 | 0 | let r = -ym / y1; |
384 | 0 | let s = ym / yn; |
385 | 0 | let t = ym / yk; |
386 | 0 | let e = -ym / nk; |
387 | 0 | let g = yn * yk / (y1 * nk) - 1.0; |
388 | 0 | let dg = if g < 0.0 { 1.0 + g } else { 1.0 }; |
389 | 0 | let gu = g * (1.0 + g * (-0.5 + g / 3.0)); |
390 | 0 | let gl = gu - g.powi(4) / (4.0 * dg); |
391 | 0 | let xm = m + 0.5; |
392 | 0 | let xn = n1 as f64 - m + 0.5; |
393 | 0 | let xk = k as f64 - m + 0.5; |
394 | 0 | let nm = n2 as f64 - k as f64 + xm; |
395 | 0 | let ub = xm * r * (1.0 + r * (-0.5 + r / 3.0)) |
396 | 0 | + xn * s * (1.0 + s * (-0.5 + s / 3.0)) |
397 | 0 | + xk * t * (1.0 + t * (-0.5 + t / 3.0)) |
398 | 0 | + nm * e * (1.0 + e * (-0.5 + e / 3.0)) |
399 | 0 | + y * gu |
400 | 0 | - m * gl |
401 | 0 | + 0.0034; |
402 | 0 | let av = v.ln(); |
403 | 0 | if av > ub { |
404 | 0 | continue; |
405 | 0 | } |
406 | 0 | let dr = if r < 0.0 { |
407 | 0 | xm * r.powi(4) / (1.0 + r) |
408 | | } else { |
409 | 0 | xm * r.powi(4) |
410 | | }; |
411 | 0 | let ds = if s < 0.0 { |
412 | 0 | xn * s.powi(4) / (1.0 + s) |
413 | | } else { |
414 | 0 | xn * s.powi(4) |
415 | | }; |
416 | 0 | let dt = if t < 0.0 { |
417 | 0 | xk * t.powi(4) / (1.0 + t) |
418 | | } else { |
419 | 0 | xk * t.powi(4) |
420 | | }; |
421 | 0 | let de = if e < 0.0 { |
422 | 0 | nm * e.powi(4) / (1.0 + e) |
423 | | } else { |
424 | 0 | nm * e.powi(4) |
425 | | }; |
426 | | |
427 | 0 | if av < ub - 0.25 * (dr + ds + dt + de) + (y + m) * (gl - gu) - 0.0078 { |
428 | 0 | break y as i64; |
429 | 0 | } |
430 | | |
431 | | // Step 4.3: Final Acceptance/Rejection Test |
432 | 0 | let av_critical = a |
433 | 0 | - ln_of_factorial(y) |
434 | 0 | - ln_of_factorial(n1 as f64 - y) |
435 | 0 | - ln_of_factorial(k as f64 - y) |
436 | 0 | - ln_of_factorial((n2 - k) as f64 + y); |
437 | 0 | if v.ln() <= av_critical { |
438 | 0 | break y as i64; |
439 | 0 | } |
440 | | } |
441 | | } |
442 | | } |
443 | | }; |
444 | | |
445 | 0 | (offset_x + sign_x * x) as u64 |
446 | 0 | } |
447 | | } |
448 | | |
449 | | #[cfg(test)] |
450 | | mod test { |
451 | | |
452 | | use super::*; |
453 | | |
454 | | #[test] |
455 | | fn test_hypergeometric_invalid_params() { |
456 | | assert!(Hypergeometric::new(100, 101, 5).is_err()); |
457 | | assert!(Hypergeometric::new(100, 10, 101).is_err()); |
458 | | assert!(Hypergeometric::new(100, 101, 101).is_err()); |
459 | | assert!(Hypergeometric::new(100, 10, 5).is_ok()); |
460 | | } |
461 | | |
462 | | fn test_hypergeometric_mean_and_variance<R: Rng>(n: u64, k: u64, s: u64, rng: &mut R) { |
463 | | let distr = Hypergeometric::new(n, k, s).unwrap(); |
464 | | |
465 | | let expected_mean = s as f64 * k as f64 / n as f64; |
466 | | let expected_variance = { |
467 | | let numerator = (s * k * (n - k) * (n - s)) as f64; |
468 | | let denominator = (n * n * (n - 1)) as f64; |
469 | | numerator / denominator |
470 | | }; |
471 | | |
472 | | let mut results = [0.0; 1000]; |
473 | | for i in results.iter_mut() { |
474 | | *i = distr.sample(rng) as f64; |
475 | | } |
476 | | |
477 | | let mean = results.iter().sum::<f64>() / results.len() as f64; |
478 | | assert!((mean - expected_mean).abs() < expected_mean / 50.0); |
479 | | |
480 | | let variance = |
481 | | results.iter().map(|x| (x - mean) * (x - mean)).sum::<f64>() / results.len() as f64; |
482 | | assert!((variance - expected_variance).abs() < expected_variance / 10.0); |
483 | | } |
484 | | |
485 | | #[test] |
486 | | fn test_hypergeometric() { |
487 | | let mut rng = crate::test::rng(737); |
488 | | |
489 | | // exercise algorithm HIN: |
490 | | test_hypergeometric_mean_and_variance(500, 400, 30, &mut rng); |
491 | | test_hypergeometric_mean_and_variance(250, 200, 230, &mut rng); |
492 | | test_hypergeometric_mean_and_variance(100, 20, 6, &mut rng); |
493 | | test_hypergeometric_mean_and_variance(50, 10, 47, &mut rng); |
494 | | |
495 | | // exercise algorithm H2PE |
496 | | test_hypergeometric_mean_and_variance(5000, 2500, 500, &mut rng); |
497 | | test_hypergeometric_mean_and_variance(10100, 10000, 1000, &mut rng); |
498 | | test_hypergeometric_mean_and_variance(100100, 100, 10000, &mut rng); |
499 | | } |
500 | | |
501 | | #[test] |
502 | | fn hypergeometric_distributions_can_be_compared() { |
503 | | assert_eq!(Hypergeometric::new(1, 2, 3), Hypergeometric::new(1, 2, 3)); |
504 | | } |
505 | | |
506 | | #[test] |
507 | | fn stirling() { |
508 | | let test = [0.5, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]; |
509 | | for &v in test.iter() { |
510 | | let ln_fac = ln_of_factorial(v); |
511 | | assert!((special::Gamma::ln_gamma(v + 1.0).0 - ln_fac).abs() < 1e-4); |
512 | | } |
513 | | } |
514 | | } |