use num_bigint::BigInt;
use num_traits::{One, Signed, ToPrimitive, Zero};
pub fn stirling2(n: impl Into<BigInt>, k: impl Into<BigInt>) -> Option<BigInt> {
let n = n.into();
let k = k.into();
if n.is_negative() || k.is_negative() {
return Some(BigInt::zero());
}
let n: u64 = n.try_into().ok()?;
let k: u64 = k.try_into().ok()?;
Some(stirling2_u64(n, k))
}
fn stirling2_u64(n: u64, k: u64) -> BigInt {
if k > n {
return BigInt::zero();
}
if n == 0 && k == 0 {
return BigInt::one();
}
if k == 0 || n == 0 {
return BigInt::zero();
}
if k == 1 || k == n {
return BigInt::one();
}
if k == 2 {
return (BigInt::one() << (n - 1) as usize) - BigInt::one();
}
if k == n - 1 {
return BigInt::from(n) * BigInt::from(n - 1) / BigInt::from(2);
}
let n = n as usize;
let k = k as usize;
let mut prev = vec![BigInt::zero(); k + 1];
prev[0] = BigInt::one();
for i in 1..=n {
let mut curr = vec![BigInt::zero(); k + 1];
for j in 1..=k.min(i) {
curr[j] = BigInt::from(j as u64) * &prev[j] + &prev[j - 1];
}
prev = curr;
}
prev[k].clone()
}
pub fn stirling1(n: impl Into<BigInt>, k: impl Into<BigInt>) -> Option<BigInt> {
let n = n.into();
let k = k.into();
if n.is_negative() || k.is_negative() {
return Some(BigInt::zero());
}
let n: u64 = n.try_into().ok()?;
let k: u64 = k.try_into().ok()?;
Some(stirling1_u64(n, k))
}
fn stirling1_u64(n: u64, k: u64) -> BigInt {
if k > n {
return BigInt::zero();
}
if n == 0 && k == 0 {
return BigInt::one();
}
if k == 0 || n == 0 {
return BigInt::zero();
}
if k == n {
return BigInt::one();
}
if k == n - 1 {
return -(BigInt::from(n) * BigInt::from(n - 1) / BigInt::from(2));
}
if k == 1 {
let mut fact = BigInt::one();
for i in 1..n {
fact *= BigInt::from(i);
}
if (n - 1).is_multiple_of(2) {
return fact;
} else {
return -fact;
}
}
let n = n as usize;
let k = k as usize;
let mut prev = vec![BigInt::zero(); k + 1];
prev[0] = BigInt::one();
for i in 1..=n {
let mut curr = vec![BigInt::zero(); k + 1];
for j in 1..=k.min(i) {
curr[j] = -BigInt::from((i - 1) as u64) * &prev[j] + &prev[j - 1];
}
prev = curr;
}
prev[k].clone()
}
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))
}
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
}
pub fn partition_count(n: impl Into<BigInt>) -> Option<BigInt> {
let n = n.into();
if n.is_negative() {
return Some(BigInt::zero());
}
if n.is_zero() {
return Some(BigInt::one());
}
let n: usize = n.try_into().ok()?;
Some(partition_count_usize(n))
}
fn partition_count_usize(n: usize) -> BigInt {
let mut table = vec![BigInt::zero(); n + 1];
table[0] = BigInt::one();
for i in 1..=n {
let mut k: i128 = 1;
loop {
let g1 = (k * (3 * k - 1) / 2) as usize;
let g2 = (k * (3 * k + 1) / 2) as usize;
if g1 > i {
break;
}
if k % 2 == 1 {
let term1 = table[i - g1].clone();
table[i] += term1;
if g2 <= i {
let term2 = table[i - g2].clone();
table[i] += term2;
}
} else {
let term1 = table[i - g1].clone();
table[i] -= term1;
if g2 <= i {
let term2 = table[i - g2].clone();
table[i] -= term2;
}
}
k += 1;
}
}
table[n].clone()
}
pub fn npartitions(n: impl Into<BigInt>) -> Option<BigInt> {
partition_count(n)
}
#[derive(Debug, Clone)]
pub struct PartitionIter {
current: Option<Vec<u64>>,
}
impl Iterator for PartitionIter {
type Item = Vec<u64>;
fn next(&mut self) -> Option<Vec<u64>> {
let cur = self.current.take()?;
let mut next = cur.clone();
let ones = next.iter().rev().take_while(|&&p| p == 1).count();
next.truncate(next.len() - ones);
match next.pop() {
None => {
self.current = None;
}
Some(last) => {
let new_part = last - 1;
let mut remainder = ones as u64 + 1; next.push(new_part);
while remainder >= new_part {
next.push(new_part);
remainder -= new_part;
}
if remainder > 0 {
next.push(remainder);
}
self.current = Some(next);
}
}
Some(cur)
}
}
pub fn partitions(n: u64) -> PartitionIter {
PartitionIter {
current: Some(if n == 0 { vec![] } else { vec![n] }),
}
}
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 bell(n: impl Into<BigInt>) -> Option<BigInt> {
let n: BigInt = n.into();
if n.is_negative() {
return Some(BigInt::zero());
}
let n = n.to_usize()?;
if n == 0 {
return Some(BigInt::one());
}
let mut prev = vec![BigInt::one()];
for _ in 1..=n {
let mut row = Vec::with_capacity(prev.len() + 1);
row.push(prev.last().cloned().unwrap_or_else(BigInt::one));
for i in 0..prev.len() {
let next = &row[i] + &prev[i];
row.push(next);
}
prev = row;
}
prev.first().cloned()
}
pub fn catalan(n: impl Into<BigInt>) -> Option<BigInt> {
let n: BigInt = n.into();
if n.is_negative() {
return Some(BigInt::zero());
}
let n = n.to_u64()?;
Some(binomial(2 * n, n) / BigInt::from(n + 1))
}
pub fn derangements(n: impl Into<BigInt>) -> Option<BigInt> {
let n: BigInt = n.into();
if n.is_negative() {
return Some(BigInt::zero());
}
let n = n.to_u64()?;
if n == 0 {
return Some(BigInt::one());
}
let mut a = BigInt::one(); let mut b = BigInt::zero(); for i in 2..=n {
let next = BigInt::from(i - 1) * (&a + &b);
a = b;
b = next;
}
Some(b)
}
#[cfg(test)]
mod tests {
use super::*;
use num_bigint::BigInt;
fn bi(n: i64) -> BigInt {
BigInt::from(n)
}
#[test]
fn stirling2_base_cases() {
assert_eq!(stirling2(0, 0), Some(bi(1)));
assert_eq!(stirling2(1, 0), Some(bi(0)));
assert_eq!(stirling2(0, 1), Some(bi(0)));
assert_eq!(stirling2(1, 1), Some(bi(1)));
}
#[test]
fn stirling2_known_values() {
assert_eq!(stirling2(3, 1), Some(bi(1)));
assert_eq!(stirling2(3, 2), Some(bi(3)));
assert_eq!(stirling2(3, 3), Some(bi(1)));
assert_eq!(stirling2(4, 1), Some(bi(1)));
assert_eq!(stirling2(4, 2), Some(bi(7)));
assert_eq!(stirling2(4, 3), Some(bi(6)));
assert_eq!(stirling2(4, 4), Some(bi(1)));
assert_eq!(stirling2(5, 2), Some(bi(15)));
assert_eq!(stirling2(5, 3), Some(bi(25)));
assert_eq!(stirling2(5, 4), Some(bi(10)));
assert_eq!(stirling2(5, 5), Some(bi(1)));
assert_eq!(stirling2(6, 3), Some(bi(90)));
assert_eq!(stirling2(7, 4), Some(bi(350)));
assert_eq!(stirling2(8, 4), Some(bi(1701)));
assert_eq!(stirling2(10, 5), Some(bi(42525)));
}
#[test]
fn stirling2_k_equals_1_always_1() {
for n in 1..=15u64 {
assert_eq!(stirling2(n, 1u64), Some(bi(1)), "S({n}, 1) should be 1");
}
}
#[test]
fn stirling2_k_equals_n_always_1() {
for n in 0..=15u64 {
assert_eq!(stirling2(n, n), Some(bi(1)), "S({n}, {n}) should be 1");
}
}
#[test]
fn stirling2_k_equals_2() {
for n in 2..=12u64 {
let expected = bi(1 << (n - 1)) - bi(1);
assert_eq!(stirling2(n, 2u64), Some(expected), "S({n}, 2)");
}
}
#[test]
fn stirling2_k_equals_n_minus_1() {
for n in 2..=15u64 {
let expected = bi((n * (n - 1) / 2) as i64);
assert_eq!(stirling2(n, n - 1), Some(expected), "S({n}, {n}-1)");
}
}
#[test]
fn stirling2_k_greater_than_n_is_zero() {
assert_eq!(stirling2(3, 5), Some(bi(0)));
assert_eq!(stirling2(0, 1), Some(bi(0)));
assert_eq!(stirling2(5, 10), Some(bi(0)));
}
#[test]
fn stirling2_negative_inputs() {
assert_eq!(stirling2(-1, 0), Some(bi(0)));
assert_eq!(stirling2(0, -1), Some(bi(0)));
assert_eq!(stirling2(-3, -2), Some(bi(0)));
}
#[test]
fn stirling2_recurrence_identity() {
for n in 2..=10u64 {
for k in 1..=n {
let lhs = stirling2(n, k).unwrap();
let rhs =
bi(k as i64) * stirling2(n - 1, k).unwrap() + stirling2(n - 1, k - 1).unwrap();
assert_eq!(lhs, rhs, "recurrence failed for S({n}, {k})");
}
}
}
#[test]
fn stirling2_row_sum_equals_bell() {
let bells = [1i64, 1, 2, 5, 15, 52, 203, 877, 4140];
for (n, &b) in bells.iter().enumerate() {
let mut sum = BigInt::zero();
for k in 0..=n {
sum += stirling2(n as u64, k as u64).unwrap();
}
assert_eq!(sum, bi(b), "Σ S({n}, k) should equal B({n}) = {b}");
}
}
#[test]
fn stirling2_huge_n_none() {
let huge = BigInt::from(u64::MAX) + BigInt::one();
assert_eq!(stirling2(huge, 2), None);
}
#[test]
fn stirling1_base_cases() {
assert_eq!(stirling1(0, 0), Some(bi(1)));
assert_eq!(stirling1(1, 0), Some(bi(0)));
assert_eq!(stirling1(0, 1), Some(bi(0)));
assert_eq!(stirling1(1, 1), Some(bi(1)));
}
#[test]
fn stirling1_known_values() {
assert_eq!(stirling1(2, 1), Some(bi(-1)));
assert_eq!(stirling1(2, 2), Some(bi(1)));
assert_eq!(stirling1(3, 1), Some(bi(2)));
assert_eq!(stirling1(3, 2), Some(bi(-3)));
assert_eq!(stirling1(3, 3), Some(bi(1)));
assert_eq!(stirling1(4, 1), Some(bi(-6)));
assert_eq!(stirling1(4, 2), Some(bi(11)));
assert_eq!(stirling1(4, 3), Some(bi(-6)));
assert_eq!(stirling1(4, 4), Some(bi(1)));
assert_eq!(stirling1(5, 1), Some(bi(24)));
assert_eq!(stirling1(5, 2), Some(bi(-50)));
assert_eq!(stirling1(5, 3), Some(bi(35)));
assert_eq!(stirling1(5, 4), Some(bi(-10)));
assert_eq!(stirling1(5, 5), Some(bi(1)));
}
#[test]
fn stirling1_k_equals_n_always_1() {
for n in 0..=15u64 {
assert_eq!(stirling1(n, n), Some(bi(1)), "s({n}, {n}) should be 1");
}
}
#[test]
fn stirling1_k_equals_n_minus_1() {
for n in 2..=15u64 {
let expected = -bi((n * (n - 1) / 2) as i64);
assert_eq!(stirling1(n, n - 1), Some(expected), "s({n}, {n}-1)");
}
}
#[test]
fn stirling1_k_equals_1() {
let factorials = [1i64, 1, 2, 6, 24, 120, 720, 5040, 40320];
for n in 1..=9u64 {
let expected = if (n - 1) % 2 == 0 {
bi(factorials[n as usize - 1])
} else {
bi(-factorials[n as usize - 1])
};
assert_eq!(stirling1(n, 1u64), Some(expected), "s({n}, 1)");
}
}
#[test]
fn stirling1_k_greater_than_n_is_zero() {
assert_eq!(stirling1(3, 5), Some(bi(0)));
assert_eq!(stirling1(0, 1), Some(bi(0)));
}
#[test]
fn stirling1_negative_inputs() {
assert_eq!(stirling1(-1, 0), Some(bi(0)));
assert_eq!(stirling1(0, -1), Some(bi(0)));
}
#[test]
fn stirling1_recurrence_identity() {
for n in 2..=10u64 {
for k in 1..=n {
let lhs = stirling1(n, k).unwrap();
let rhs = -bi((n - 1) as i64) * stirling1(n - 1, k).unwrap()
+ stirling1(n - 1, k - 1).unwrap();
assert_eq!(lhs, rhs, "recurrence failed for s({n}, {k})");
}
}
}
#[test]
fn stirling1_row_sum_is_zero_for_n_ge_2() {
for n in 2..=10u64 {
let mut sum = BigInt::zero();
for k in 0..=n {
sum += stirling1(n, k).unwrap();
}
assert_eq!(sum, bi(0), "Σ s({n}, k) should be 0");
}
}
#[test]
fn stirling1_unsigned_row_sum_is_n_factorial() {
let factorials = [1i64, 1, 2, 6, 24, 120, 720, 5040, 40320];
for n in 0..=8u64 {
let mut sum = BigInt::zero();
for k in 0..=n {
let s = stirling1(n, k).unwrap();
sum += if s.is_negative() { -s } else { s };
}
assert_eq!(
sum,
bi(factorials[n as usize]),
"Σ |s({n}, k)| should be {n}!"
);
}
}
#[test]
fn stirling1_huge_n_none() {
let huge = BigInt::from(u64::MAX) + BigInt::one();
assert_eq!(stirling1(huge, 2), None);
}
#[test]
fn stirling_orthogonality() {
for n in 0..=7u64 {
for k in 0..=7u64 {
let mut sum = BigInt::zero();
for j in 0..=n.max(k) {
sum += stirling1(n, j).unwrap() * stirling2(j, k).unwrap();
}
let expected = if n == k { bi(1) } else { bi(0) };
assert_eq!(
sum, expected,
"orthogonality failed: Σ_j s({n},j)*S(j,{k}) = {sum}, expected {expected}"
);
}
}
}
#[test]
fn multinomial_basic() {
assert_eq!(multinomial(6, &[2, 3, 1]), Some(bi(60)));
}
#[test]
fn multinomial_reduces_to_binomial() {
assert_eq!(multinomial(10, &[3, 7]), Some(bi(120)));
assert_eq!(multinomial(6, &[2, 4]), Some(bi(15)));
assert_eq!(multinomial(20, &[10, 10]), Some(bi(184756)));
}
#[test]
fn multinomial_all_ones() {
assert_eq!(multinomial(4, &[1, 1, 1, 1]), Some(bi(24)));
assert_eq!(multinomial(5, &[1, 1, 1, 1, 1]), Some(bi(120)));
}
#[test]
fn multinomial_single_group() {
assert_eq!(multinomial(5, &[5]), Some(bi(1)));
assert_eq!(multinomial(10, &[10]), Some(bi(1)));
}
#[test]
fn multinomial_with_zeros() {
assert_eq!(multinomial(6, &[2, 0, 3, 0, 1]), Some(bi(60)));
}
#[test]
fn multinomial_sum_mismatch_returns_zero() {
assert_eq!(multinomial(5, &[2, 2]), Some(bi(0)));
assert_eq!(multinomial(5, &[3, 3]), Some(bi(0)));
}
#[test]
fn multinomial_negative_k_returns_zero() {
assert_eq!(multinomial(5, &[-1, 6]), Some(bi(0)));
}
#[test]
fn multinomial_zero() {
let empty: &[i64] = &[];
assert_eq!(multinomial(0, empty), Some(bi(1)));
assert_eq!(multinomial(0, &[0]), Some(bi(1)));
}
#[test]
fn multinomial_larger() {
assert_eq!(multinomial(12, &[3, 4, 5]), Some(bi(27720)));
}
#[test]
fn multinomial_huge_k_none() {
let huge = BigInt::from(u64::MAX) + BigInt::one();
assert_eq!(multinomial(huge.clone(), &[huge]), None);
}
#[test]
fn partition_count_small() {
let expected = [1, 1, 2, 3, 5, 7, 11, 15, 22, 30, 42];
for (n, &p) in expected.iter().enumerate() {
assert_eq!(
partition_count(n as u64),
Some(bi(p)),
"p({n}) should be {p}"
);
}
}
#[test]
fn partition_count_medium() {
assert_eq!(partition_count(20u64), Some(bi(627)));
assert_eq!(partition_count(30u64), Some(bi(5604)));
assert_eq!(partition_count(50u64), Some(bi(204226)));
}
#[test]
fn partition_count_100() {
assert_eq!(partition_count(100u64), Some(bi(190569292)));
}
#[test]
fn partition_count_200() {
assert_eq!(
partition_count(200u64),
Some(BigInt::parse_bytes(b"3972999029388", 10).unwrap())
);
}
#[test]
fn partition_count_negative() {
assert_eq!(partition_count(-1), Some(bi(0)));
assert_eq!(partition_count(-100), Some(bi(0)));
}
#[test]
fn partition_count_zero() {
assert_eq!(partition_count(0), Some(bi(1)));
}
}