m1une's library

This documentation is automatically generated by online-judge-tools/verification-helper

View on GitHub

:heavy_check_mark: Binomial Coefficient Modulo an Arbitrary Modulus
(math/binomial_coefficient_mod.hpp)

Overview

BinomialCoefficientMod computes binomial coefficients modulo one fixed positive integer. The modulus may be prime, a prime power, composite, or 1, and n may be as large as an unsigned 64-bit integer.

The constructor factors the modulus into pairwise coprime prime powers. For each prime power, it precomputes factorial products with factors of that prime removed. A query evaluates the factorials recursively, restores the correct prime exponent, and combines all prime-power residues with the Chinese remainder theorem.

This structure is intended for many queries using the same moderate modulus. Its preprocessing is linear in the sum of the prime-power factors, so it is a good fit for moduli up to roughly a few million. It should not be constructed with a modulus near 10^9 unless that memory and preprocessing cost are acceptable.

ArbitraryModBinomialCoefficient is an alias for BinomialCoefficientMod.

Interface

class BinomialCoefficientMod {
public:
    explicit BinomialCoefficientMod(uint32_t modulus);

    uint32_t modulus() const;
    uint32_t binom(uint64_t n, uint64_t k) const;
    uint32_t operator()(uint64_t n, uint64_t k) const;
};

using ArbitraryModBinomialCoefficient = BinomialCoefficientMod;

Let

\[m = \prod_{i=1}^s p_i^{e_i}\]

be the prime-power factorization of the modulus, and write $q_i = p_i^{e_i}$. The complexities are:

Method Description Complexity
BinomialCoefficientMod(modulus) Factors the modulus and prepares the unit-factorial tables. $O(\sqrt m + \sum_i q_i)$ time and $O(\sum_i q_i)$ memory
modulus() Returns the fixed modulus. $O(1)$
binom(n, k) Returns $\binom{n}{k} \bmod m$; invalid k returns zero. $O(\sum_i \log_{p_i}(n+1))$
operator()(n, k) Alias for binom(n, k). $O(\sum_i \log_{p_i}(n+1))$

Behavioral Notes

Example

#include "math/binomial_coefficient_mod.hpp"

#include <iostream>

int main() {
    m1une::math::BinomialCoefficientMod combinations(12);

    std::cout << combinations.binom(5, 2) << '\n';  // 10
    std::cout << combinations(10, 3) << '\n';       // 0
    std::cout << combinations(1000000000000000000ULL, 1) << '\n';  // 4
}

The object can be reused for any number of queries as long as the modulus remains 12.

Depends on

Required by

Verified with

Code

#ifndef M1UNE_MATH_BINOMIAL_COEFFICIENT_MOD_HPP
#define M1UNE_MATH_BINOMIAL_COEFFICIENT_MOD_HPP 1

#include <cassert>
#include <cstddef>
#include <cstdint>
#include <vector>

#include "number_theory.hpp"

namespace m1une {
namespace math {

// Binomial coefficients modulo a fixed, not necessarily prime, modulus.
class BinomialCoefficientMod {
   private:
    struct PrimePower {
        uint32_t prime;
        int exponent;
        uint32_t modulus;
        uint32_t crt_multiplier;
        std::vector<uint32_t> unit_factorial_prefix;

        uint32_t multiply(uint32_t lhs, uint32_t rhs) const {
            return uint32_t(uint64_t(lhs) * rhs % modulus);
        }

        uint32_t power(uint32_t base, uint64_t exponent_) const {
            uint32_t result = 1 % modulus;
            while (exponent_ > 0) {
                if (exponent_ & 1) result = multiply(result, base);
                base = multiply(base, base);
                exponent_ >>= 1;
            }
            return result;
        }

        uint64_t factorial_valuation(uint64_t n) const {
            uint64_t result = 0;
            while (n > 0) {
                n /= prime;
                result += n;
            }
            return result;
        }

        uint32_t unit_factorial(uint64_t n) const {
            if (n == 0) return 1 % modulus;
            const uint32_t block_product = unit_factorial_prefix.back();
            uint32_t result = power(block_product, n / modulus);
            result = multiply(result, unit_factorial_prefix[std::size_t(n % modulus)]);
            return multiply(result, unit_factorial(n / prime));
        }

        uint32_t binom(uint64_t n, uint64_t k) const {
            if (k > n) return 0;
            const uint64_t valuation = factorial_valuation(n) - factorial_valuation(k) -
                                       factorial_valuation(n - k);
            if (valuation >= uint64_t(exponent)) return 0;

            const uint32_t numerator = unit_factorial(n);
            const uint32_t denominator =
                multiply(unit_factorial(k), unit_factorial(n - k));
            const uint32_t inverse_denominator =
                uint32_t(inv_mod(denominator, modulus));
            uint32_t result = multiply(numerator, inverse_denominator);
            result = multiply(result, power(prime, valuation));
            return result;
        }
    };

    uint32_t _modulus;
    std::vector<PrimePower> _prime_powers;

   public:
    explicit BinomialCoefficientMod(uint32_t modulus) : _modulus(modulus) {
        assert(modulus >= 1);
        uint32_t remaining = modulus;
        for (uint32_t prime = 2; uint64_t(prime) * prime <= remaining; prime++) {
            if (remaining % prime != 0) continue;
            int exponent = 0;
            uint32_t prime_power = 1;
            do {
                remaining /= prime;
                prime_power *= prime;
                exponent++;
            } while (remaining % prime == 0);
            _prime_powers.push_back(
                PrimePower{prime, exponent, prime_power, 0, {}});
        }
        if (remaining > 1) {
            _prime_powers.push_back(PrimePower{remaining, 1, remaining, 0, {}});
        }

        for (PrimePower& component : _prime_powers) {
            component.unit_factorial_prefix.resize(std::size_t(component.modulus));
            component.unit_factorial_prefix[0] = 1;
            for (uint32_t value = 1; value < component.modulus; value++) {
                component.unit_factorial_prefix[value] =
                    component.unit_factorial_prefix[value - 1];
                if (value % component.prime != 0) {
                    component.unit_factorial_prefix[value] = component.multiply(
                        component.unit_factorial_prefix[value], value);
                }
            }

            const uint32_t other = modulus / component.modulus;
            const uint32_t inverse =
                uint32_t(inv_mod(other, component.modulus));
            component.crt_multiplier =
                uint32_t(uint64_t(other) * inverse % modulus);
        }
    }

    uint32_t modulus() const {
        return _modulus;
    }

    uint32_t binom(uint64_t n, uint64_t k) const {
        if (k > n || _modulus == 1) return 0;
        uint64_t result = 0;
        for (const PrimePower& component : _prime_powers) {
            const uint32_t residue = component.binom(n, k);
            result += uint64_t(residue) * component.crt_multiplier % _modulus;
            result %= _modulus;
        }
        return uint32_t(result);
    }

    uint32_t operator()(uint64_t n, uint64_t k) const {
        return binom(n, k);
    }
};

using ArbitraryModBinomialCoefficient = BinomialCoefficientMod;

}  // namespace math
}  // namespace m1une

#endif  // M1UNE_MATH_BINOMIAL_COEFFICIENT_MOD_HPP
#line 1 "math/binomial_coefficient_mod.hpp"



#include <cassert>
#include <cstddef>
#include <cstdint>
#include <vector>

#line 1 "math/number_theory.hpp"



#line 6 "math/number_theory.hpp"
#include <limits>
#include <tuple>
#include <utility>
#line 10 "math/number_theory.hpp"

namespace m1une {
namespace math {

namespace internal {

inline long long safe_mod(long long x, long long mod) {
    x %= mod;
    if (x < 0) x += mod;
    return x;
}

inline unsigned __int128 floor_sum_unsigned(unsigned long long n, unsigned long long mod, unsigned long long a,
                                            unsigned long long b) {
    unsigned __int128 answer = 0;
    while (true) {
        if (a >= mod) {
            answer += static_cast<unsigned __int128>(n) * (n - 1) / 2 * (a / mod);
            a %= mod;
        }
        if (b >= mod) {
            answer += static_cast<unsigned __int128>(n) * (b / mod);
            b %= mod;
        }

        const unsigned __int128 y_max = static_cast<unsigned __int128>(a) * n + b;
        if (y_max < mod) break;
        n = static_cast<unsigned long long>(y_max / mod);
        b = static_cast<unsigned long long>(y_max % mod);
        unsigned long long tmp = mod;
        mod = a;
        a = tmp;
    }
    return answer;
}

}  // namespace internal

// Returns (g, x, y), where g = gcd(a, b) is nonnegative and
// a * x + b * y = g. Returns (0, 0, 0) when a = b = 0.
inline std::tuple<long long, long long, long long> extended_gcd(long long a,
                                                               long long b) {
    using i128 = __int128;
    if (a == 0 && b == 0) return {0, 0, 0};

    i128 old_remainder = a;
    i128 remainder = b;
    if (old_remainder < 0) old_remainder = -old_remainder;
    if (remainder < 0) remainder = -remainder;
    i128 old_x = 1;
    i128 x = 0;
    i128 old_y = 0;
    i128 y = 1;

    while (remainder != 0) {
        i128 quotient = old_remainder / remainder;

        i128 next = old_remainder - quotient * remainder;
        old_remainder = remainder;
        remainder = next;

        next = old_x - quotient * x;
        old_x = x;
        x = next;

        next = old_y - quotient * y;
        old_y = y;
        y = next;
    }

    if (a < 0) old_x = -old_x;
    if (b < 0) old_y = -old_y;

#ifndef NDEBUG
    const i128 minimum = std::numeric_limits<long long>::min();
    const i128 maximum = std::numeric_limits<long long>::max();
    assert(old_remainder <= maximum);
    assert(minimum <= old_x && old_x <= maximum);
    assert(minimum <= old_y && old_y <= maximum);
#endif
    return {static_cast<long long>(old_remainder), static_cast<long long>(old_x),
            static_cast<long long>(old_y)};
}

inline long long pow_mod(long long x, unsigned long long exponent, long long mod) {
    assert(mod >= 1);
    if (mod == 1) return 0;

    unsigned long long base = static_cast<unsigned long long>(internal::safe_mod(x, mod));
    unsigned long long result = 1;
    const unsigned long long unsigned_mod = static_cast<unsigned long long>(mod);
    while (exponent > 0) {
        if (exponent & 1) {
            result = static_cast<unsigned long long>(static_cast<unsigned __int128>(result) * base % unsigned_mod);
        }
        base = static_cast<unsigned long long>(static_cast<unsigned __int128>(base) * base % unsigned_mod);
        exponent >>= 1;
    }
    return static_cast<long long>(result);
}

// Returns gcd(a, mod) and x such that a * x is congruent to gcd(a, mod)
// modulo mod. The returned x is in [0, mod / gcd(a, mod)).
inline std::pair<long long, long long> inv_gcd(long long a, long long mod) {
    assert(mod >= 1);
    a = internal::safe_mod(a, mod);
    if (a == 0) return {mod, 0};

    long long s = mod;
    long long t = a;
    long long m0 = 0;
    long long m1 = 1;
    while (t > 0) {
        const long long quotient = s / t;
        s -= t * quotient;
        m0 -= m1 * quotient;

        long long tmp = s;
        s = t;
        t = tmp;
        tmp = m0;
        m0 = m1;
        m1 = tmp;
    }
    if (m0 < 0) m0 += mod / s;
    return {s, m0};
}

inline long long inv_mod(long long x, long long mod) {
    const auto result = inv_gcd(x, mod);
    assert(result.first == 1);
    return result.second;
}

// Returns the smallest nonnegative solution and the least common multiple of
// the moduli. Returns {0, 0} when the system is inconsistent.
inline std::pair<long long, long long> crt(const std::vector<long long>& remainders,
                                           const std::vector<long long>& moduli) {
    assert(remainders.size() == moduli.size());

    long long r0 = 0;
    long long m0 = 1;
    for (int i = 0; i < int(remainders.size()); i++) {
        assert(moduli[i] >= 1);
        long long r1 = internal::safe_mod(remainders[i], moduli[i]);
        long long m1 = moduli[i];

        if (m0 < m1) {
            long long tmp = r0;
            r0 = r1;
            r1 = tmp;
            tmp = m0;
            m0 = m1;
            m1 = tmp;
        }
        if (m0 % m1 == 0) {
            if (r0 % m1 != r1) return {0, 0};
            continue;
        }

        const auto inverse = inv_gcd(m0, m1);
        const long long gcd = inverse.first;
        const long long reduced_modulus = m1 / gcd;
        const __int128 difference = static_cast<__int128>(r1) - r0;
        if (difference % gcd != 0) return {0, 0};

        __int128 multiplier = difference / gcd % reduced_modulus;
        multiplier = multiplier * inverse.second % reduced_modulus;
        if (multiplier < 0) multiplier += reduced_modulus;

        const __int128 new_modulus = static_cast<__int128>(m0) * reduced_modulus;
        assert(new_modulus <= std::numeric_limits<long long>::max());
        __int128 new_remainder = static_cast<__int128>(r0) + multiplier * m0;
        new_remainder %= new_modulus;
        if (new_remainder < 0) new_remainder += new_modulus;
        r0 = static_cast<long long>(new_remainder);
        m0 = static_cast<long long>(new_modulus);
    }
    return {r0, m0};
}

// Returns sum_{i=0}^{n-1} floor((a * i + b) / mod).
inline long long floor_sum(long long n, long long mod, long long a, long long b) {
    assert(n >= 0);
    assert(mod >= 1);

    const long long normalized_a = internal::safe_mod(a, mod);
    const long long normalized_b = internal::safe_mod(b, mod);
    __int128 answer = (static_cast<__int128>(a) - normalized_a) / mod * n * (n - 1) / 2;
    answer += (static_cast<__int128>(b) - normalized_b) / mod * n;
    answer += internal::floor_sum_unsigned(static_cast<unsigned long long>(n), static_cast<unsigned long long>(mod),
                                           static_cast<unsigned long long>(normalized_a),
                                           static_cast<unsigned long long>(normalized_b));

    assert(answer >= std::numeric_limits<long long>::min());
    assert(answer <= std::numeric_limits<long long>::max());
    return static_cast<long long>(answer);
}

}  // namespace math
}  // namespace m1une


#line 10 "math/binomial_coefficient_mod.hpp"

namespace m1une {
namespace math {

// Binomial coefficients modulo a fixed, not necessarily prime, modulus.
class BinomialCoefficientMod {
   private:
    struct PrimePower {
        uint32_t prime;
        int exponent;
        uint32_t modulus;
        uint32_t crt_multiplier;
        std::vector<uint32_t> unit_factorial_prefix;

        uint32_t multiply(uint32_t lhs, uint32_t rhs) const {
            return uint32_t(uint64_t(lhs) * rhs % modulus);
        }

        uint32_t power(uint32_t base, uint64_t exponent_) const {
            uint32_t result = 1 % modulus;
            while (exponent_ > 0) {
                if (exponent_ & 1) result = multiply(result, base);
                base = multiply(base, base);
                exponent_ >>= 1;
            }
            return result;
        }

        uint64_t factorial_valuation(uint64_t n) const {
            uint64_t result = 0;
            while (n > 0) {
                n /= prime;
                result += n;
            }
            return result;
        }

        uint32_t unit_factorial(uint64_t n) const {
            if (n == 0) return 1 % modulus;
            const uint32_t block_product = unit_factorial_prefix.back();
            uint32_t result = power(block_product, n / modulus);
            result = multiply(result, unit_factorial_prefix[std::size_t(n % modulus)]);
            return multiply(result, unit_factorial(n / prime));
        }

        uint32_t binom(uint64_t n, uint64_t k) const {
            if (k > n) return 0;
            const uint64_t valuation = factorial_valuation(n) - factorial_valuation(k) -
                                       factorial_valuation(n - k);
            if (valuation >= uint64_t(exponent)) return 0;

            const uint32_t numerator = unit_factorial(n);
            const uint32_t denominator =
                multiply(unit_factorial(k), unit_factorial(n - k));
            const uint32_t inverse_denominator =
                uint32_t(inv_mod(denominator, modulus));
            uint32_t result = multiply(numerator, inverse_denominator);
            result = multiply(result, power(prime, valuation));
            return result;
        }
    };

    uint32_t _modulus;
    std::vector<PrimePower> _prime_powers;

   public:
    explicit BinomialCoefficientMod(uint32_t modulus) : _modulus(modulus) {
        assert(modulus >= 1);
        uint32_t remaining = modulus;
        for (uint32_t prime = 2; uint64_t(prime) * prime <= remaining; prime++) {
            if (remaining % prime != 0) continue;
            int exponent = 0;
            uint32_t prime_power = 1;
            do {
                remaining /= prime;
                prime_power *= prime;
                exponent++;
            } while (remaining % prime == 0);
            _prime_powers.push_back(
                PrimePower{prime, exponent, prime_power, 0, {}});
        }
        if (remaining > 1) {
            _prime_powers.push_back(PrimePower{remaining, 1, remaining, 0, {}});
        }

        for (PrimePower& component : _prime_powers) {
            component.unit_factorial_prefix.resize(std::size_t(component.modulus));
            component.unit_factorial_prefix[0] = 1;
            for (uint32_t value = 1; value < component.modulus; value++) {
                component.unit_factorial_prefix[value] =
                    component.unit_factorial_prefix[value - 1];
                if (value % component.prime != 0) {
                    component.unit_factorial_prefix[value] = component.multiply(
                        component.unit_factorial_prefix[value], value);
                }
            }

            const uint32_t other = modulus / component.modulus;
            const uint32_t inverse =
                uint32_t(inv_mod(other, component.modulus));
            component.crt_multiplier =
                uint32_t(uint64_t(other) * inverse % modulus);
        }
    }

    uint32_t modulus() const {
        return _modulus;
    }

    uint32_t binom(uint64_t n, uint64_t k) const {
        if (k > n || _modulus == 1) return 0;
        uint64_t result = 0;
        for (const PrimePower& component : _prime_powers) {
            const uint32_t residue = component.binom(n, k);
            result += uint64_t(residue) * component.crt_multiplier % _modulus;
            result %= _modulus;
        }
        return uint32_t(result);
    }

    uint32_t operator()(uint64_t n, uint64_t k) const {
        return binom(n, k);
    }
};

using ArbitraryModBinomialCoefficient = BinomialCoefficientMod;

}  // namespace math
}  // namespace m1une
Back to top page