Coverage Report

Created: 2026-06-30 07:02

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/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
/// ![Hypergeometric distribution](https://raw.githubusercontent.com/rust-random/charts/main/charts/hypergeometric.svg)
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
}