CP-Algorithms Library

This documentation is automatically generated by competitive-verifier/competitive-verifier

View the Project on GitHub cp-algorithms/cp-algorithms-aux

:warning: tests/modint.cpp

Depends on

Code

// dynamic_modint against 128-bit arithmetic and is_prime against a reference, for moduli across
// the signed range: odd moduli above a quarter of the unsigned word and even moduli keep reduced
// residues, the others lazy Montgomery residues.
#include "cp-algo/number_theory/primality.hpp"
#include <bits/stdc++.h>
using namespace cp_algo::math;
using u128 = unsigned __int128;
static uint64_t mulmod(uint64_t a, uint64_t b, uint64_t m) {return uint64_t(u128(a) * b % m);}
static uint64_t powmod(uint64_t a, uint64_t e, uint64_t m) {uint64_t r = 1; for(a %= m; e; e >>= 1, a = mulmod(a, a, m)) if(e & 1) r = mulmod(r, a, m); return r;}
static bool ref_prime(uint64_t n) {
    if(n < 2) return false;
    for(uint64_t p: {2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37}) {if(n % p == 0) return n == p;}
    uint64_t d = n - 1; int s = 0; while(d % 2 == 0) {d /= 2; s++;}
    for(uint64_t a: {2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37}) {
        uint64_t x = powmod(a, d, n); if(x == 1 || x == n - 1) continue;
        bool comp = true; for(int i = 1; i < s && comp; i++) {x = mulmod(x, x, n); if(x == n - 1) comp = false;}
        if(comp) return false;
    }
    return true;
}
template<typename Int> size_t arithmetic(Int m, std::mt19937_64& rng, int iters) {
    using base = dynamic_modint<Int>; size_t bad = 0;
    base::with_mod(m, [&]() {
        for(int it = 0; it < iters; it++) {
            uint64_t x = rng() % uint64_t(m), y = rng() % uint64_t(m), z = rng() % uint64_t(m);
            if(it < 8) {x = it & 1 ? uint64_t(m) - 1 : 0; y = it & 2 ? uint64_t(m) - 1 : 1; z = it & 4 ? uint64_t(m) - 1 : 0;}
            base a, b, c; a.setr(x); b.setr(y); c.setr(z);
            uint64_t um = uint64_t(m);
            bad += (a + b).getr() != (x + y) % um % um || (a + b).getr() != uint64_t((u128(x) + y) % um);
            bad += (a - b).getr() != uint64_t((u128(x) + um - y) % um);
            bad += (-a).getr() != (um - x) % um;
            bad += (a * b).getr() != mulmod(x, y, um);
            bad += ((a * b + c) * a - b * c).getr() != uint64_t((u128(mulmod(uint64_t((u128(mulmod(x, y, um)) + z) % um), x, um)) + um - mulmod(y, z, um)) % um);
            base acc = a; for(int k = 0; k < 5; k++) {acc += acc; acc -= b; acc *= acc;}
            uint64_t r = x; for(int k = 0; k < 5; k++) {r = uint64_t((u128(r) + r) % um); r = uint64_t((u128(r) + um - y) % um); r = mulmod(r, r, um);}
            bad += acc.getr() != r;
        }
        return 0;
    });
    return bad;
}
int main() {
    std::mt19937_64 rng(7);
    size_t bad32 = 0, bad64 = 0, n32 = 0, n64 = 0, pbad32 = 0, pbad64 = 0, np32 = 0, np64 = 0, primes = 0;
    for(int shift: {3, 10, 20, 29, 30, 31}) for(int it = 0; it < 300; it++) {
        uint64_t hi = (uint64_t(1) << shift) - 1, lo = hi / 2 + 1;
        int m = int(lo + rng() % (hi - lo + 1)); if(it < 40) m = int(hi - it); if(m < 3) continue;
        bad32 += arithmetic<int>(m, rng, 60); n32++;
        bool want = ref_prime(uint64_t(m)); primes += want;
        pbad32 += is_prime(m) != want; pbad32 += is_prime(uint32_t(m)) != want; np32++;
    }
    for(int shift: {31, 33, 50, 61, 62, 63}) for(int it = 0; it < 300; it++) {
        uint64_t hi = (uint64_t(1) << shift) - 1, lo = hi / 2 + 1;
        int64_t m = int64_t(lo + rng() % (hi - lo + 1)); if(it < 40) m = int64_t(hi - it);
        bad64 += arithmetic<int64_t>(m, rng, 60); n64++;
        bool want = ref_prime(uint64_t(m)); primes += want;
        pbad64 += is_prime(m) != want; np64++;
    }
    for(int p: {1879048201, 2147395589, 2147395609, 2147483629, 2147483647}) {pbad32 += !is_prime(p); np32++;}
    for(int64_t p: {int64_t(9223372036854775783), int64_t(4611686018427388039), int64_t(9223372036854775643)}) {pbad64 += is_prime(p) != ref_prime(uint64_t(p)); np64++;}
    assert(n32 > 1700 && n64 > 1700 && np32 > n32 && np64 > n64 && primes > 100);
    assert(!bad32 && !bad64 && !pbad32 && !pbad64);
    std::cout << "dynamic_modint and is_prime agree with 128-bit references on " << n32 + n64 << " moduli\n";
}
#line 1 "tests/modint.cpp"
// dynamic_modint against 128-bit arithmetic and is_prime against a reference, for moduli across
// the signed range: odd moduli above a quarter of the unsigned word and even moduli keep reduced
// residues, the others lazy Montgomery residues.
#line 1 "cp-algo/number_theory/primality.hpp"


#line 1 "cp-algo/number_theory/modint.hpp"


#line 1 "cp-algo/math/common.hpp"


#include <functional>
#include <cstdint>
#include <cassert>
#include <bit>
#include <vector>
#include <algorithm>
namespace cp_algo::math {
#ifdef CP_ALGO_MAXN
    const int maxn = CP_ALGO_MAXN;
#else
    const int maxn = 1 << 19;
#endif
    const int magic = 64; // threshold for sizes to run the naive algo

    // Nonnegative 64-bit exponents, with an associative operation and its identity.
    // Windows >1 precompute odd powers only when that saves operations.
    template<int window = 1>
    auto bpow(auto const& x, auto n, auto const& one, auto op) {
        static_assert(window >= 1 && window <= 6);
        if constexpr(window > 1) {
            if(n == 0) {return one;}
            int bits = std::bit_width(uint64_t(n));
            auto low_bit = [&](int high) {
                int low = std::max(0, high - window + 1);
                while(!((n >> low) & 1)) {low++;}
                return low;
            };
            int first = low_bit(bits - 1);
            int cost = (1 << (window - 1)) + first;
            for(int j = first - 1; j >= 0;) {
                if(!((n >> j) & 1)) {j--;}
                else {cost++; j = low_bit(j) - 1;}
            }
            // Do not pay for the table when binary powering uses fewer operations.
            if(cost >= bits + std::popcount(uint64_t(n)) - 2) {return bpow<1>(x, n, one, op);}
            using T = std::decay_t<decltype(x)>;
            std::vector<T> odd;
            odd.reserve(1 << (window - 1));
            odd.push_back(x);
            auto square = op(x, x);
            while(odd.size() < size_t(1 << (window - 1))) {odd.push_back(op(odd.back(), square));}
            auto ans = odd[(n >> first) / 2];
            for(int j = first - 1; j >= 0;) {
                if(!((n >> j) & 1)) {ans = op(ans, ans); j--;}
                else {
                    int low = low_bit(j), length = j - low + 1;
                    auto digit = (n >> low) & ((1u << length) - 1);
                    for(int i = 0; i < length; i++) {ans = op(ans, ans);}
                    ans = op(ans, odd[digit / 2]);
                    j = low - 1;
                }
            }
            return ans;
        } else {
            if (n == 0) {
                return one;
            }
            auto ans = x;
            for(int j = std::bit_width<uint64_t>(n) - 2; ~j; j--) {
                ans = op(ans, ans);
                if((n >> j) & 1) {
                    ans = op(ans, x);
                }
            }
            return ans;
        }
    }
    template<int window = 1>
    auto bpow(auto x, auto n, auto ans) {
        return bpow<window>(x, n, ans, std::multiplies{});
    }
    template<typename T>
    T bpow(T const& x, auto n) {
        return bpow(x, n, T(1));
    }
    inline constexpr auto inv2(auto x) {
        assert(x % 2);
        std::make_unsigned_t<decltype(x)> y = 1;
        while(y * x != 1) {
            y *= 2 - x * y;
        }
        return y;
    }
}

#line 4 "cp-algo/number_theory/modint.hpp"
#include <iostream>
#line 6 "cp-algo/number_theory/modint.hpp"
namespace cp_algo::math {

    template<typename modint, typename _Int>
    struct modint_base {
        using Int = _Int;
        using UInt = std::make_unsigned_t<Int>;
        static constexpr size_t bits = sizeof(Int) * 8;
        using Int2 = std::conditional_t<bits <= 32, int64_t, __int128_t>;
        using UInt2 = std::conditional_t<bits <= 32, uint64_t, __uint128_t>;
        constexpr static Int mod() {
            return modint::mod();
        }
        constexpr static UInt remod() {
            return modint::remod();
        }
        constexpr static UInt2 modmod() {
            return UInt2(mod()) * mod();
        }
        constexpr modint_base() = default;
        constexpr modint_base(Int2 rr) {
            to_modint().setr(UInt((rr + modmod()) % mod()));
        }
        constexpr modint inv() const {
            return bpow(to_modint(), mod() - 2);
        }
        modint operator - () const {
            modint neg;
            neg.r = std::min(-r, remod() - r);
            return neg;
        }
        modint& operator /= (const modint &t) {
            return to_modint() *= t.inv();
        }
        modint& operator *= (const modint &t) {
            r = UInt(UInt2(r) * t.r % mod());
            return to_modint();
        }
        modint& operator += (const modint &t) {
            r += t.r; r = std::min(r, r - remod());
            return to_modint();
        }
        modint& operator -= (const modint &t) {
            r -= t.r; r = std::min(r, r + remod());
            return to_modint();
        }
        modint operator + (const modint &t) const {return modint(to_modint()) += t;}
        modint operator - (const modint &t) const {return modint(to_modint()) -= t;}
        modint operator * (const modint &t) const {return modint(to_modint()) *= t;}
        modint operator / (const modint &t) const {return modint(to_modint()) /= t;}
        // Why <=> doesn't work?..
        auto operator == (const modint &t) const {return to_modint().getr() == t.getr();}
        auto operator != (const modint &t) const {return to_modint().getr() != t.getr();}
        auto operator <= (const modint &t) const {return to_modint().getr() <= t.getr();}
        auto operator >= (const modint &t) const {return to_modint().getr() >= t.getr();}
        auto operator < (const modint &t) const {return to_modint().getr() < t.getr();}
        auto operator > (const modint &t) const {return to_modint().getr() > t.getr();}
        Int rem() const {
            UInt R = to_modint().getr();
            return R - (R > (UInt)mod() / 2) * mod();
        }
        constexpr void setr(UInt rr) {
            r = rr;
        }
        constexpr UInt getr() const {
            return r;
        }

        // Only use these if you really know what you're doing!
        static uint64_t modmod8() {return uint64_t(8 * modmod());}
        void add_unsafe(UInt t) {r += t;}
        void pseudonormalize() {r = std::min(r, r - modmod8());}
        modint const& normalize() {
            if(r >= (UInt)mod()) {
                r %= mod();
            }
            return to_modint();
        }
        void setr_direct(UInt rr) {r = rr;}
        UInt getr_direct() const {return r;}
    protected:
        UInt r;
    private:
        constexpr modint& to_modint() {return static_cast<modint&>(*this);}
        constexpr modint const& to_modint() const {return static_cast<modint const&>(*this);}
    };
    template<typename modint>
    concept modint_type = std::is_base_of_v<modint_base<modint, typename modint::Int>, modint>;
    template<modint_type modint>
    decltype(std::cin)& operator >> (decltype(std::cin) &in, modint &x) {
        typename modint::UInt r;
        auto &res = in >> r;
        x.setr(r);
        return res;
    }
    template<modint_type modint>
    decltype(std::cout)& operator << (decltype(std::cout) &out, modint const& x) {
        return out << x.getr();
    }

    template<auto m>
    struct modint: modint_base<modint<m>, decltype(m)> {
        using Base = modint_base<modint<m>, decltype(m)>;
        using Base::Base;
        static constexpr Base::Int mod() {return m;}
        static constexpr Base::UInt remod() {return m;}
        auto getr() const {return Base::r;}
    };

    // Odd moduli up to a quarter of the unsigned word keep Montgomery residues lazily in
    // [0, 2 mod): remod() = 2 mod, and both 4 mod and ab + q * mod fit. Any other modulus keeps
    // fully reduced residues, remod() = mod, which is what the sums of modint_base need then:
    // an even one without the Montgomery form, a wide odd one with a reduction that subtracts
    // high words instead of adding double words that would overflow.
    template<typename Int = int>
    struct dynamic_modint: modint_base<dynamic_modint<Int>, Int> {
        using Base = modint_base<dynamic_modint<Int>, Int>;
        using Base::Base;

        // Out of line, so that the hot path stays as small as it was.
        [[gnu::noinline, gnu::cold]] static Base::UInt m_reduce_reduced(Base::UInt2 ab) {
            if(mod() % 2 == 0) {return typename Base::UInt(ab % mod());}
            // q * mod has the low word of ab, so the difference of the high words is exact.
            typename Base::UInt q = -(typename Base::UInt(ab) * inverse);
            auto high = typename Base::UInt(ab >> Base::bits);
            auto low = typename Base::UInt(typename Base::UInt2(q) * typename Base::UInt(mod()) >> Base::bits);
            return high >= low ? high - low : high - low + mod();
        }
        static Base::UInt m_reduce(Base::UInt2 ab) {
            if(imod() == 0) [[unlikely]] {return m_reduce_reduced(ab);}
            typename Base::UInt2 m = typename Base::UInt(ab) * imod();
            return typename Base::UInt((ab + m * mod()) >> Base::bits);
        }
        static Base::UInt m_transform(Base::UInt a) {
            if(mod() % 2 == 0) [[unlikely]] {
                return a;
            } else {
                return m_reduce(a * pw128());
            }
        }
        dynamic_modint& operator *= (const dynamic_modint &t) {
            Base::r = m_reduce(typename Base::UInt2(Base::r) * t.r);
            return *this;
        }
        void setr(Base::UInt rr) {
            Base::r = m_transform(rr);
        }
        Base::UInt getr() const {
            typename Base::UInt res = m_reduce(Base::r);
            return std::min(res, res - mod());
        }
        static Int mod() {return m;}
        static Base::UInt remod() {return rm;}
        static Base::UInt imod() {return im;}
        static Base::UInt2 pw128() {return r2;}
        static void switch_mod(Int nm) {
            m = nm;
            bool lazy = m % 2 && typename Base::UInt(m) <= typename Base::UInt(-1) / 4;
            rm = typename Base::UInt(m) * (lazy ? 2 : 1);
            inverse = m % 2 ? inv2(-m) : 0;
            im = lazy ? inverse : 0;
            r2 = static_cast<Base::UInt>(static_cast<Base::UInt2>(-1) % m + 1);
        }

        // Wrapper for temp switching
        auto static with_mod(Int tmp, auto callback) {
            struct scoped {
                Int prev = mod();
                ~scoped() {switch_mod(prev);}
            } _;
            switch_mod(tmp);
            return callback();
        }
    private:
        static thread_local Int m;
        // im: -1 / mod modulo 2^bits for lazy residues and 0 for reduced ones; inverse: the same
        // for every odd mod; rm: the value of remod().
        static thread_local Base::UInt im, r2, inverse, rm;
    };
    template<typename Int>
    Int thread_local dynamic_modint<Int>::m = 1;
    template<typename Int>
    dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::im = -1;
    template<typename Int>
    dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::r2 = 0;
    template<typename Int>
    dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::inverse = -1;
    template<typename Int>
    dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::rm = 2;
}

#line 6 "cp-algo/number_theory/primality.hpp"
namespace cp_algo::math {
    // https://en.wikipedia.org/wiki/Miller–Rabin_primality_test
    template<typename _Int>
    bool is_prime(_Int m) {
        using Int = std::make_signed_t<_Int>;
        using UInt = std::make_unsigned_t<Int>;
        if(m == 1 || m % 2 == 0) {
            return m == 2;
        }
        // m - 1 = 2^s * d
        int s = std::countr_zero(UInt(m - 1));
        auto d = (m - 1) >> s;
        using base = dynamic_modint<Int>;
        auto test = [&](base x) {
            x = bpow(x, d);
            if(std::abs(x.rem()) <= 1) {
                return true;
            }
            for(int i = 1; i < s && x != -1; i++) {
                x *= x;
            }
            return x == -1;
        };
        return base::with_mod(m, [&]() {
#ifdef CP_ALGO_NUMBER_THEORY_PRIMALITY_BASES_HPP
            uint16_t base2 = 7, base3 = 61;
            if (m != uint32_t(m)) {
                base2 = base_table1[uint32_t(m * 0xAD625B89) >> 18];
                base3 = base_table2[base2 >> 13];
            }
            return test(2) && test(base2) && test(base3);
#else
            return std::ranges::all_of(std::array{2, 325, 9375, 28178, 450775, 9780504, 1795265022}, test);
#endif
        });
    }
}

#line 5 "tests/modint.cpp"
#include <bits/stdc++.h>
using namespace cp_algo::math;
using u128 = unsigned __int128;
static uint64_t mulmod(uint64_t a, uint64_t b, uint64_t m) {return uint64_t(u128(a) * b % m);}
static uint64_t powmod(uint64_t a, uint64_t e, uint64_t m) {uint64_t r = 1; for(a %= m; e; e >>= 1, a = mulmod(a, a, m)) if(e & 1) r = mulmod(r, a, m); return r;}
static bool ref_prime(uint64_t n) {
    if(n < 2) return false;
    for(uint64_t p: {2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37}) {if(n % p == 0) return n == p;}
    uint64_t d = n - 1; int s = 0; while(d % 2 == 0) {d /= 2; s++;}
    for(uint64_t a: {2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37}) {
        uint64_t x = powmod(a, d, n); if(x == 1 || x == n - 1) continue;
        bool comp = true; for(int i = 1; i < s && comp; i++) {x = mulmod(x, x, n); if(x == n - 1) comp = false;}
        if(comp) return false;
    }
    return true;
}
template<typename Int> size_t arithmetic(Int m, std::mt19937_64& rng, int iters) {
    using base = dynamic_modint<Int>; size_t bad = 0;
    base::with_mod(m, [&]() {
        for(int it = 0; it < iters; it++) {
            uint64_t x = rng() % uint64_t(m), y = rng() % uint64_t(m), z = rng() % uint64_t(m);
            if(it < 8) {x = it & 1 ? uint64_t(m) - 1 : 0; y = it & 2 ? uint64_t(m) - 1 : 1; z = it & 4 ? uint64_t(m) - 1 : 0;}
            base a, b, c; a.setr(x); b.setr(y); c.setr(z);
            uint64_t um = uint64_t(m);
            bad += (a + b).getr() != (x + y) % um % um || (a + b).getr() != uint64_t((u128(x) + y) % um);
            bad += (a - b).getr() != uint64_t((u128(x) + um - y) % um);
            bad += (-a).getr() != (um - x) % um;
            bad += (a * b).getr() != mulmod(x, y, um);
            bad += ((a * b + c) * a - b * c).getr() != uint64_t((u128(mulmod(uint64_t((u128(mulmod(x, y, um)) + z) % um), x, um)) + um - mulmod(y, z, um)) % um);
            base acc = a; for(int k = 0; k < 5; k++) {acc += acc; acc -= b; acc *= acc;}
            uint64_t r = x; for(int k = 0; k < 5; k++) {r = uint64_t((u128(r) + r) % um); r = uint64_t((u128(r) + um - y) % um); r = mulmod(r, r, um);}
            bad += acc.getr() != r;
        }
        return 0;
    });
    return bad;
}
int main() {
    std::mt19937_64 rng(7);
    size_t bad32 = 0, bad64 = 0, n32 = 0, n64 = 0, pbad32 = 0, pbad64 = 0, np32 = 0, np64 = 0, primes = 0;
    for(int shift: {3, 10, 20, 29, 30, 31}) for(int it = 0; it < 300; it++) {
        uint64_t hi = (uint64_t(1) << shift) - 1, lo = hi / 2 + 1;
        int m = int(lo + rng() % (hi - lo + 1)); if(it < 40) m = int(hi - it); if(m < 3) continue;
        bad32 += arithmetic<int>(m, rng, 60); n32++;
        bool want = ref_prime(uint64_t(m)); primes += want;
        pbad32 += is_prime(m) != want; pbad32 += is_prime(uint32_t(m)) != want; np32++;
    }
    for(int shift: {31, 33, 50, 61, 62, 63}) for(int it = 0; it < 300; it++) {
        uint64_t hi = (uint64_t(1) << shift) - 1, lo = hi / 2 + 1;
        int64_t m = int64_t(lo + rng() % (hi - lo + 1)); if(it < 40) m = int64_t(hi - it);
        bad64 += arithmetic<int64_t>(m, rng, 60); n64++;
        bool want = ref_prime(uint64_t(m)); primes += want;
        pbad64 += is_prime(m) != want; np64++;
    }
    for(int p: {1879048201, 2147395589, 2147395609, 2147483629, 2147483647}) {pbad32 += !is_prime(p); np32++;}
    for(int64_t p: {int64_t(9223372036854775783), int64_t(4611686018427388039), int64_t(9223372036854775643)}) {pbad64 += is_prime(p) != ref_prime(uint64_t(p)); np64++;}
    assert(n32 > 1700 && n64 > 1700 && np32 > n32 && np64 > n64 && primes > 100);
    assert(!bad32 && !bad64 && !pbad32 && !pbad64);
    std::cout << "dynamic_modint and is_prime agree with 128-bit references on " << n32 + n64 << " moduli\n";
}
#line 1 "tests/modint.cpp"
#line 1 "cp-algo/number_theory/primality.hpp"
#line 1 "cp-algo/number_theory/modint.hpp"
#line 1 "cp-algo/math/common.hpp"
#include <functional>
#include <cstdint>
#include <cassert>
#include <bit>
#include <vector>
#include <algorithm>
namespace cp_algo::math{
#ifdef CP_ALGO_MAXN
const int maxn=CP_ALGO_MAXN;
#else
const int maxn=1<<19;
#endif
const int magic=64;template<int window=1>auto bpow(auto const&x,auto n,auto const&one,auto op){static_assert(window>=1&&window<=6);if constexpr(window>1){if(n==0){return one;}int bits=std::bit_width(uint64_t(n));auto low_bit=[&](int high){int low=std::max(0,high-window+1);while(!((n>>low)&1)){low++;}return low;};int first=low_bit(bits-1);int cost=(1<<(window-1))+first;for(int j=first-1;j>=0;){if(!((n>>j)&1)){j--;}else{cost++;j=low_bit(j)-1;}}if(cost>=bits+std::popcount(uint64_t(n))-2){return bpow<1>(x,n,one,op);}using T=std::decay_t<decltype(x)>;std::vector<T>odd;odd.reserve(1<<(window-1));odd.push_back(x);auto square=op(x,x);while(odd.size()<size_t(1<<(window-1))){odd.push_back(op(odd.back(),square));}auto ans=odd[(n>>first)/2];for(int j=first-1;j>=0;){if(!((n>>j)&1)){ans=op(ans,ans);j--;}else{int low=low_bit(j),length=j-low+1;auto digit=(n>>low)&((1u<<length)-1);for(int i=0;i<length;i++){ans=op(ans,ans);}ans=op(ans,odd[digit/2]);j=low-1;}}return ans;}else{if(n==0){return one;}auto ans=x;for(int j=std::bit_width<uint64_t>(n)-2;~j;j--){ans=op(ans,ans);if((n>>j)&1){ans=op(ans,x);}}return ans;}}template<int window=1>auto bpow(auto x,auto n,auto ans){return bpow<window>(x,n,ans,std::multiplies{});}template<typename T>T bpow(T const&x,auto n){return bpow(x,n,T(1));}inline constexpr auto inv2(auto x){assert(x%2);std::make_unsigned_t<decltype(x)>y=1;while(y*x!=1){y*=2-x*y;}return y;}}
#line 4 "cp-algo/number_theory/modint.hpp"
#include <iostream>
#line 6 "cp-algo/number_theory/modint.hpp"
namespace cp_algo::math{template<typename modint,typename _Int>struct modint_base{using Int=_Int;using UInt=std::make_unsigned_t<Int>;static constexpr size_t bits=sizeof(Int)*8;using Int2=std::conditional_t<bits<=32,int64_t,__int128_t>;using UInt2=std::conditional_t<bits<=32,uint64_t,__uint128_t>;constexpr static Int mod(){return modint::mod();}constexpr static UInt remod(){return modint::remod();}constexpr static UInt2 modmod(){return UInt2(mod())*mod();}constexpr modint_base()=default;constexpr modint_base(Int2 rr){to_modint().setr(UInt((rr+modmod())%mod()));}constexpr modint inv()const{return bpow(to_modint(),mod()-2);}modint operator-()const{modint neg;neg.r=std::min(-r,remod()-r);return neg;}modint&operator/=(const modint&t){return to_modint()*=t.inv();}modint&operator*=(const modint&t){r=UInt(UInt2(r)*t.r%mod());return to_modint();}modint&operator+=(const modint&t){r+=t.r;r=std::min(r,r-remod());return to_modint();}modint&operator-=(const modint&t){r-=t.r;r=std::min(r,r+remod());return to_modint();}modint operator+(const modint&t)const{return modint(to_modint())+=t;}modint operator-(const modint&t)const{return modint(to_modint())-=t;}modint operator*(const modint&t)const{return modint(to_modint())*=t;}modint operator/(const modint&t)const{return modint(to_modint())/=t;}auto operator==(const modint&t)const{return to_modint().getr()==t.getr();}auto operator!=(const modint&t)const{return to_modint().getr()!=t.getr();}auto operator<=(const modint&t)const{return to_modint().getr()<=t.getr();}auto operator>=(const modint&t)const{return to_modint().getr()>=t.getr();}auto operator<(const modint&t)const{return to_modint().getr()<t.getr();}auto operator>(const modint&t)const{return to_modint().getr()>t.getr();}Int rem()const{UInt R=to_modint().getr();return R-(R>(UInt)mod()/2)*mod();}constexpr void setr(UInt rr){r=rr;}constexpr UInt getr()const{return r;}static uint64_t modmod8(){return uint64_t(8*modmod());}void add_unsafe(UInt t){r+=t;}void pseudonormalize(){r=std::min(r,r-modmod8());}modint const&normalize(){if(r>=(UInt)mod()){r%=mod();}return to_modint();}void setr_direct(UInt rr){r=rr;}UInt getr_direct()const{return r;}protected:UInt r;private:constexpr modint&to_modint(){return static_cast<modint&>(*this);}constexpr modint const&to_modint()const{return static_cast<modint const&>(*this);}};template<typename modint>concept modint_type=std::is_base_of_v<modint_base<modint,typename modint::Int>,modint>;template<modint_type modint>decltype(std::cin)&operator>>(decltype(std::cin)&in,modint&x){typename modint::UInt r;auto&res=in>>r;x.setr(r);return res;}template<modint_type modint>decltype(std::cout)&operator<<(decltype(std::cout)&out,modint const&x){return out<<x.getr();}template<auto m>struct modint:modint_base<modint<m>,decltype(m)>{using Base=modint_base<modint<m>,decltype(m)>;using Base::Base;static constexpr Base::Int mod(){return m;}static constexpr Base::UInt remod(){return m;}auto getr()const{return Base::r;}};template<typename Int=int>struct dynamic_modint:modint_base<dynamic_modint<Int>,Int>{using Base=modint_base<dynamic_modint<Int>,Int>;using Base::Base;[[gnu::noinline,gnu::cold]]static Base::UInt m_reduce_reduced(Base::UInt2 ab){if(mod()%2==0){return typename Base::UInt(ab%mod());}typename Base::UInt q=-(typename Base::UInt(ab)*inverse);auto high=typename Base::UInt(ab>>Base::bits);auto low=typename Base::UInt(typename Base::UInt2(q)*typename Base::UInt(mod())>>Base::bits);return high>=low?high-low:high-low+mod();}static Base::UInt m_reduce(Base::UInt2 ab){if(imod()==0)[[unlikely]]{return m_reduce_reduced(ab);}typename Base::UInt2 m=typename Base::UInt(ab)*imod();return typename Base::UInt((ab+m*mod())>>Base::bits);}static Base::UInt m_transform(Base::UInt a){if(mod()%2==0)[[unlikely]]{return a;}else{return m_reduce(a*pw128());}}dynamic_modint&operator*=(const dynamic_modint&t){Base::r=m_reduce(typename Base::UInt2(Base::r)*t.r);return*this;}void setr(Base::UInt rr){Base::r=m_transform(rr);}Base::UInt getr()const{typename Base::UInt res=m_reduce(Base::r);return std::min(res,res-mod());}static Int mod(){return m;}static Base::UInt remod(){return rm;}static Base::UInt imod(){return im;}static Base::UInt2 pw128(){return r2;}static void switch_mod(Int nm){m=nm;bool lazy=m%2&&typename Base::UInt(m)<=typename Base::UInt(-1)/4;rm=typename Base::UInt(m)*(lazy?2:1);inverse=m%2?inv2(-m):0;im=lazy?inverse:0;r2=static_cast<Base::UInt>(static_cast<Base::UInt2>(-1)%m+1);}auto static with_mod(Int tmp,auto callback){struct scoped{Int prev=mod();~scoped(){switch_mod(prev);}}_;switch_mod(tmp);return callback();}private:static thread_local Int m;static thread_local Base::UInt im,r2,inverse,rm;};template<typename Int>Int thread_local dynamic_modint<Int>::m=1;template<typename Int>dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::im=-1;template<typename Int>dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::r2=0;template<typename Int>dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::inverse=-1;template<typename Int>dynamic_modint<Int>::Base::UInt thread_local dynamic_modint<Int>::rm=2;}
#line 6 "cp-algo/number_theory/primality.hpp"
namespace cp_algo::math{template<typename _Int>bool is_prime(_Int m){using Int=std::make_signed_t<_Int>;using UInt=std::make_unsigned_t<Int>;if(m==1||m%2==0){return m==2;}int s=std::countr_zero(UInt(m-1));auto d=(m-1)>>s;using base=dynamic_modint<Int>;auto test=[&](base x){x=bpow(x,d);if(std::abs(x.rem())<=1){return true;}for(int i=1;i<s&&x!=-1;i++){x*=x;}return x==-1;};return base::with_mod(m,[&](){
#ifdef CP_ALGO_NUMBER_THEORY_PRIMALITY_BASES_HPP
uint16_t base2=7,base3=61;if(m!=uint32_t(m)){base2=base_table1[uint32_t(m*0xAD625B89)>>18];base3=base_table2[base2>>13];}return test(2)&&test(base2)&&test(base3);
#else
return std::ranges::all_of(std::array{2,325,9375,28178,450775,9780504,1795265022},test);
#endif
});}}
#line 5 "tests/modint.cpp"
#include <bits/stdc++.h>
using namespace cp_algo::math;using u128=unsigned __int128;static uint64_t mulmod(uint64_t a,uint64_t b,uint64_t m){return uint64_t(u128(a)*b%m);}static uint64_t powmod(uint64_t a,uint64_t e,uint64_t m){uint64_t r=1;for(a%=m;e;e>>=1,a=mulmod(a,a,m))if(e&1)r=mulmod(r,a,m);return r;}static bool ref_prime(uint64_t n){if(n<2)return false;for(uint64_t p:{2,3,5,7,11,13,17,19,23,29,31,37}){if(n%p==0)return n==p;}uint64_t d=n-1;int s=0;while(d%2==0){d/=2;s++;}for(uint64_t a:{2,3,5,7,11,13,17,19,23,29,31,37}){uint64_t x=powmod(a,d,n);if(x==1||x==n-1)continue;bool comp=true;for(int i=1;i<s&&comp;i++){x=mulmod(x,x,n);if(x==n-1)comp=false;}if(comp)return false;}return true;}template<typename Int>size_t arithmetic(Int m,std::mt19937_64&rng,int iters){using base=dynamic_modint<Int>;size_t bad=0;base::with_mod(m,[&](){for(int it=0;it<iters;it++){uint64_t x=rng()%uint64_t(m),y=rng()%uint64_t(m),z=rng()%uint64_t(m);if(it<8){x=it&1?uint64_t(m)-1:0;y=it&2?uint64_t(m)-1:1;z=it&4?uint64_t(m)-1:0;}base a,b,c;a.setr(x);b.setr(y);c.setr(z);uint64_t um=uint64_t(m);bad+=(a+b).getr()!=(x+y)%um%um||(a+b).getr()!=uint64_t((u128(x)+y)%um);bad+=(a-b).getr()!=uint64_t((u128(x)+um-y)%um);bad+=(-a).getr()!=(um-x)%um;bad+=(a*b).getr()!=mulmod(x,y,um);bad+=((a*b+c)*a-b*c).getr()!=uint64_t((u128(mulmod(uint64_t((u128(mulmod(x,y,um))+z)%um),x,um))+um-mulmod(y,z,um))%um);base acc=a;for(int k=0;k<5;k++){acc+=acc;acc-=b;acc*=acc;}uint64_t r=x;for(int k=0;k<5;k++){r=uint64_t((u128(r)+r)%um);r=uint64_t((u128(r)+um-y)%um);r=mulmod(r,r,um);}bad+=acc.getr()!=r;}return 0;});return bad;}int main(){std::mt19937_64 rng(7);size_t bad32=0,bad64=0,n32=0,n64=0,pbad32=0,pbad64=0,np32=0,np64=0,primes=0;for(int shift:{3,10,20,29,30,31})for(int it=0;it<300;it++){uint64_t hi=(uint64_t(1)<<shift)-1,lo=hi/2+1;int m=int(lo+rng()%(hi-lo+1));if(it<40)m=int(hi-it);if(m<3)continue;bad32+=arithmetic<int>(m,rng,60);n32++;bool want=ref_prime(uint64_t(m));primes+=want;pbad32+=is_prime(m)!=want;pbad32+=is_prime(uint32_t(m))!=want;np32++;}for(int shift:{31,33,50,61,62,63})for(int it=0;it<300;it++){uint64_t hi=(uint64_t(1)<<shift)-1,lo=hi/2+1;int64_t m=int64_t(lo+rng()%(hi-lo+1));if(it<40)m=int64_t(hi-it);bad64+=arithmetic<int64_t>(m,rng,60);n64++;bool want=ref_prime(uint64_t(m));primes+=want;pbad64+=is_prime(m)!=want;np64++;}for(int p:{1879048201,2147395589,2147395609,2147483629,2147483647}){pbad32+=!is_prime(p);np32++;}for(int64_t p:{int64_t(9223372036854775783),int64_t(4611686018427388039),int64_t(9223372036854775643)}){pbad64+=is_prime(p)!=ref_prime(uint64_t(p));np64++;}assert(n32>1700&&n64>1700&&np32>n32&&np64>n64&&primes>100);assert(!bad32&&!bad64&&!pbad32&&!pbad64);std::cout<<"dynamic_modint and is_prime agree with 128-bit references on "<<n32+n64<<" moduli\n";}
Back to top page