use num_bigint::BigInt;
use num_traits::{One, Signed, ToPrimitive, Zero};
pub fn factorial(n: u64) -> BigInt {
if n < 2 {
return BigInt::one();
}
product_range(2, n)
}
fn product_range(lo: u64, hi: u64) -> BigInt {
const LEAF: u64 = 16;
if hi - lo < LEAF {
let mut acc = BigInt::from(lo);
for i in lo + 1..=hi {
acc *= i;
}
return acc;
}
let mid = lo + (hi - lo) / 2;
product_range(lo, mid) * product_range(mid + 1, hi)
}
pub fn binomial(n: impl Into<BigInt>, k: impl Into<BigInt>) -> BigInt {
let n: BigInt = n.into();
let k: BigInt = k.into();
if k.is_negative() {
return BigInt::zero();
}
if !n.is_negative() && k > n {
return BigInt::zero();
}
let k = if !n.is_negative() && &k * 2 > n {
&n - &k
} else {
k
};
let Some(k) = k.to_u64() else {
return BigInt::zero();
};
let mut result = BigInt::one();
for i in 0..k {
result = result * (&n - BigInt::from(i)) / BigInt::from(i + 1);
}
result
}
pub fn multinomial(n: impl Into<BigInt>, ks: &[impl Into<BigInt> + Clone]) -> Option<BigInt> {
let n = n.into();
if n.is_negative() {
return Some(BigInt::zero());
}
let ks_big: Vec<BigInt> = ks.iter().map(|k| k.clone().into()).collect();
let mut sum = BigInt::zero();
for k in &ks_big {
if k.is_negative() {
return Some(BigInt::zero());
}
sum += k;
}
if sum != n {
return Some(BigInt::zero());
}
let ks_u64: Vec<u64> = ks_big
.iter()
.map(|k| k.try_into().ok())
.collect::<Option<Vec<u64>>>()?;
Some(multinomial_u64(&ks_u64))
}
pub fn multinomial_u64(ks: &[u64]) -> BigInt {
let mut result = BigInt::one();
let mut remaining: u64 = ks.iter().sum();
for &k in ks {
if k == 0 {
continue;
}
let mut binom = BigInt::one();
for i in 0..k {
binom *= BigInt::from(remaining - i);
binom /= BigInt::from(i + 1);
}
result *= binom;
remaining -= k;
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn factorial_matches_the_running_product() {
let mut acc = BigInt::one();
assert_eq!(factorial(0), acc);
for n in 1..=200u64 {
acc *= n;
assert_eq!(factorial(n), acc, "{n}!");
}
}
#[test]
fn binomial_symmetry_and_pascal() {
for n in 0..30i64 {
for k in 0..=n {
assert_eq!(binomial(n, k), binomial(n, n - k));
if k > 0 {
assert_eq!(
binomial(n + 1, k),
binomial(n, k - 1) + binomial(n, k),
"Pascal at ({n}, {k})"
);
}
}
}
assert_eq!(binomial(5, -1), BigInt::zero());
assert_eq!(binomial(-1, 3), BigInt::from(-1));
assert_eq!(binomial(-2, 3), BigInt::from(-4));
}
#[test]
fn multinomial_parts() {
assert_eq!(multinomial_u64(&[2, 3, 1]), BigInt::from(60));
assert_eq!(multinomial_u64(&[0, 4]), BigInt::one());
assert_eq!(multinomial(4, &[1, 1, 1, 1]), Some(BigInt::from(24)));
assert_eq!(multinomial(4, &[1, 1, 1]), Some(BigInt::zero()));
assert_eq!(multinomial(-1, &[1]), Some(BigInt::zero()));
assert_eq!(multinomial(3, &[-1, 4]), Some(BigInt::zero()));
}
}