NicheLibrary

This documentation is automatically generated by NotLeonian/competitive-verifier (forked from competitive-verifier/competitive-verifier)

View the Project on GitHub NotLeonian/NicheLibrary

:heavy_check_mark: verify/yukicoder-2512.test.cpp

Depends on

Code

// competitive-verifier: PROBLEM https://yukicoder.me/problems/no/2512

#include <cstdint>
#include <iostream>
#include <vector>

#include "../math/combinatorics/online-binomial-sum.hpp"

class modint998244353 {
  public:
    static constexpr std::uint32_t mod = 998244353;

    modint998244353() : value_(0) {}

    modint998244353(long long value) {
        long long reduced = value % static_cast<long long>(mod);
        if (reduced < 0) {
            reduced += mod;
        }
        value_ = static_cast<std::uint32_t>(reduced);
    }

    std::uint32_t val() const { return value_; }

    modint998244353 inv() const { return pow(*this, mod - 2); }

    modint998244353 &operator+=(const modint998244353 &rhs) {
        std::uint32_t value = value_ + rhs.value_;
        if (value >= mod) {
            value -= mod;
        }
        value_ = value;
        return *this;
    }

    modint998244353 &operator-=(const modint998244353 &rhs) {
        const std::uint32_t value = value_ >= rhs.value_
                                        ? value_ - rhs.value_
                                        : value_ + mod - rhs.value_;
        value_ = value;
        return *this;
    }

    modint998244353 &operator*=(const modint998244353 &rhs) {
        const std::uint64_t value =
            static_cast<std::uint64_t>(value_) * rhs.value_ % mod;
        value_ = static_cast<std::uint32_t>(value);
        return *this;
    }

    modint998244353 &operator/=(const modint998244353 &rhs) {
        return *this *= rhs.inv();
    }

    friend modint998244353 operator+(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs += rhs;
    }

    friend modint998244353 operator-(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs -= rhs;
    }

    friend modint998244353 operator*(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs *= rhs;
    }

    friend modint998244353 operator/(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs /= rhs;
    }

    friend bool operator==(const modint998244353 &lhs,
                           const modint998244353 &rhs) {
        return lhs.value_ == rhs.value_;
    }

  private:
    static modint998244353 pow(modint998244353 base, long long exponent) {
        modint998244353 result(1);
        while (exponent > 0) {
            if (exponent % 2 == 1) {
                result *= base;
            }
            base *= base;
            exponent /= 2;
        }
        return result;
    }

    std::uint32_t value_;
};

int main() {
    std::ios::sync_with_stdio(false);
    std::cin.tie(nullptr);

    constexpr int max_n = 200000;
    constexpr int max_m = 400000;
    static_assert(max_m < static_cast<int>(modint998244353::mod),
                  "max_m must be less than the modulus.");

    OnlineBinomialSum<modint998244353> online_binomial_sum(max_m,
                                                           modint998244353(-2));
    const modint998244353 minus_inv2 = modint998244353(-1) / modint998244353(2);
    std::vector<modint998244353> power_minus_inv2(max_n + 2);
    power_minus_inv2[0] = modint998244353(1);
    for (int i = 0; i <= max_n; ++i) {
        power_minus_inv2[i + 1] = power_minus_inv2[i] * minus_inv2;
    }

    int test_count;
    std::cin >> test_count;
    while (test_count--) {
        int n, m;
        std::cin >> n >> m;
        const modint998244353 sum =
            online_binomial_sum.binom_prefix_sum(n + 1, 2 * m);
        const modint998244353 ans =
            (sum - modint998244353(1)) *
            (modint998244353(0) - power_minus_inv2[n + 1]);
        std::cout << ans.val() << '\n';
    }

    return 0;
}
#line 1 "verify/yukicoder-2512.test.cpp"
// competitive-verifier: PROBLEM https://yukicoder.me/problems/no/2512

#include <cstdint>
#include <iostream>
#include <vector>

#line 1 "math/combinatorics/online-binomial-sum.hpp"



// Σ_{i=l}^{u-1} r^i binom(m,i) をオンラインで求める。
// 0 <= l <= u と 0 <= m <= max_m を仮定する。
// n と m のバケット境界および上端の直積で累積和をサンプルし、
// クエリの点に 2 次元的に最も近いサンプル点から復元する。
// T は素数を法とする体の型である。
// std::numeric_limits<T>::is_integer が false であり、
// 法 p について max_m < p であることを仮定する。
// r = 0 では、バケットサイズを指定しない場合もその計算を行わない。
// 前計算とクエリは時間計算量、空間計算量ともに O(1)。
// r = -1 では交代二項和の閉形式を用いる。
// B をバケットサイズとして、r が 0, -1 でない場合の時間計算量は
// 前計算 O(max_m^2 / B + max_m)、クエリ O(B)。
// 空間計算量は O((max_m / B + 1)^2 + max_m)。

#include <cassert>
#include <limits>
#line 21 "math/combinatorics/online-binomial-sum.hpp"

template <class T> struct OnlineBinomialSum {
    static_assert(!std::numeric_limits<T>::is_integer,
                  "std::numeric_limits<T>::is_integer must be false.");

  public:
    int max_m;
    int bucket_size;
    T r;
    T r_plus_one;
    bool r_is_zero;
    bool r_is_minus_one;
    T r_plus_one_inverse;
    std::vector<T> factorial;
    std::vector<T> inverse_factorial;
    std::vector<T> integer_inverse;
    std::vector<T> power_r;
    std::vector<int> sample_n_list;
    std::vector<int> sample_m_list;
    std::vector<T> sample_sum_table;

    explicit OnlineBinomialSum(int max_m, T r, int bucket_size)
        : max_m(max_m), bucket_size(bucket_size), r(r), r_plus_one(r + T(1)),
          r_is_zero(r == T()), r_is_minus_one(r_plus_one == T()),
          r_plus_one_inverse(T()) {
        assert(max_m >= 0);
        assert(bucket_size > 0);

        if (r_is_zero) {
            return;
        }

        factorial.assign(max_m + 1, T(1));
        for (int i = 1; i <= max_m; ++i) {
            factorial[i] = factorial[i - 1] * T(i);
        }

        inverse_factorial.assign(max_m + 1, T(1));
        inverse_factorial[max_m] = T(1) / factorial[max_m];
        for (int i = max_m; i >= 1; --i) {
            inverse_factorial[i - 1] = inverse_factorial[i] * T(i);
        }

        if (r_is_minus_one) {
            return;
        }

        r_plus_one_inverse = T(1) / r_plus_one;

        integer_inverse.assign(max_m + 1, T());
        for (int i = 1; i <= max_m; ++i) {
            integer_inverse[i] = factorial[i - 1] * inverse_factorial[i];
        }

        power_r.assign(max_m + 2, T());
        power_r[0] = T(1);
        for (int i = 0; i <= max_m; ++i) {
            power_r[i + 1] = power_r[i] * r;
        }

        sample_n_list = make_sample_list(max_m + 1);
        sample_m_list = make_sample_list(max_m);
        sample_sum_table.assign(sample_n_list.size() * sample_m_list.size(),
                                T());

        const int sample_n_count = static_cast<int>(sample_n_list.size());
        const int sample_m_count = static_cast<int>(sample_m_list.size());
        for (int sample_m_index = 0; sample_m_index < sample_m_count;
             ++sample_m_index) {
            const int sample_m = sample_m_list[sample_m_index];
            T sum = T();
            T term = T(1);
            int current_n = 0;

            for (int sample_n_index = 0; sample_n_index < sample_n_count;
                 ++sample_n_index) {
                const int sample_n = sample_n_list[sample_n_index];
                while (current_n < sample_n && current_n <= sample_m) {
                    sum += term;
                    if (current_n < sample_m) {
                        term *= r;
                        term *= T(sample_m - current_n);
                        term *= integer_inverse[current_n + 1];
                    }
                    ++current_n;
                }
                sample_sum_table[sample_m_index * sample_n_count +
                                 sample_n_index] = sum;
            }
        }
    }

    explicit OnlineBinomialSum(int max_m, T r = T(1))
        : OnlineBinomialSum(max_m, r,
                            r == T() ? 1 : default_bucket_size(max_m)) {}

    T binom_prefix_sum(int n, int m) const {
        assert(n >= 0);
        assert(m >= 0);
        assert(m <= max_m);

        return binom_prefix_sum_unchecked(n, m);
    }

    T binom_sum(int l, int u, int m) const {
        assert(l >= 0);
        assert(l <= u);
        assert(m >= 0);
        assert(m <= max_m);

        return binom_prefix_sum_unchecked(u, m) -
               binom_prefix_sum_unchecked(l, m);
    }

  private:
    T binom_prefix_sum_unchecked(int n, int m) const {
        if (n == 0) {
            return T();
        }
        if (n > m) {
            n = m + 1;
        }
        if (r_is_zero) {
            return T(1);
        }
        if (r_is_minus_one) {
            if (m == 0) {
                return T(1);
            }
            if (n > m) {
                return T();
            }

            T ans = binomial(m - 1, n - 1);
            if ((n - 1) % 2 == 1) {
                ans = T() - ans;
            }
            return ans;
        }

        const int sample_n_index = nearest_sample_index(sample_n_list, n);
        const int sample_m_index = nearest_sample_index(sample_m_list, m);
        const int sample_n_count = static_cast<int>(sample_n_list.size());
        int current_n = sample_n_list[sample_n_index];
        int current_m = sample_m_list[sample_m_index];
        T sum =
            sample_sum_table[sample_m_index * sample_n_count + sample_n_index];

        while (current_n < n) {
            if (current_n <= current_m) {
                sum += power_r[current_n] * binomial(current_m, current_n);
            }
            ++current_n;
        }

        while (current_n > n) {
            --current_n;
            if (current_n <= current_m) {
                sum -= power_r[current_n] * binomial(current_m, current_n);
            }
        }

        while (current_m < m) {
            sum *= r_plus_one;
            if (current_n - 1 <= current_m) {
                sum -= power_r[current_n] * binomial(current_m, current_n - 1);
            }
            ++current_m;
        }

        while (current_m > m) {
            --current_m;
            sum += power_r[current_n] * binomial(current_m, current_n - 1);
            sum *= r_plus_one_inverse;
        }

        return sum;
    }

    T binomial(int n, int k) const {
        return factorial[n] * inverse_factorial[k] * inverse_factorial[n - k];
    }

    std::vector<int> make_sample_list(int limit) const {
        const int full_bucket_count = limit / bucket_size;
        std::vector<int> sample_list;
        sample_list.reserve(full_bucket_count + 2);
        for (int index = 0; index <= full_bucket_count; ++index) {
            sample_list.push_back(index * bucket_size);
        }
        if (sample_list.back() != limit) {
            sample_list.push_back(limit);
        }

        return sample_list;
    }

    int nearest_sample_index(const std::vector<int> &sample_list,
                             int value) const {
        int index = value / bucket_size;
        const int sample_count = static_cast<int>(sample_list.size());
        if (index + 1 >= sample_count) {
            return index;
        }
        if (value - sample_list[index] <= sample_list[index + 1] - value) {
            return index;
        }

        return index + 1;
    }

    static int default_bucket_size(int max_m) {
        assert(max_m >= 0);

        int bucket_size = 1;
        while (static_cast<long long>(bucket_size) * bucket_size <= max_m) {
            bucket_size *= 2;
        }

        return bucket_size;
    }
};


#line 8 "verify/yukicoder-2512.test.cpp"

class modint998244353 {
  public:
    static constexpr std::uint32_t mod = 998244353;

    modint998244353() : value_(0) {}

    modint998244353(long long value) {
        long long reduced = value % static_cast<long long>(mod);
        if (reduced < 0) {
            reduced += mod;
        }
        value_ = static_cast<std::uint32_t>(reduced);
    }

    std::uint32_t val() const { return value_; }

    modint998244353 inv() const { return pow(*this, mod - 2); }

    modint998244353 &operator+=(const modint998244353 &rhs) {
        std::uint32_t value = value_ + rhs.value_;
        if (value >= mod) {
            value -= mod;
        }
        value_ = value;
        return *this;
    }

    modint998244353 &operator-=(const modint998244353 &rhs) {
        const std::uint32_t value = value_ >= rhs.value_
                                        ? value_ - rhs.value_
                                        : value_ + mod - rhs.value_;
        value_ = value;
        return *this;
    }

    modint998244353 &operator*=(const modint998244353 &rhs) {
        const std::uint64_t value =
            static_cast<std::uint64_t>(value_) * rhs.value_ % mod;
        value_ = static_cast<std::uint32_t>(value);
        return *this;
    }

    modint998244353 &operator/=(const modint998244353 &rhs) {
        return *this *= rhs.inv();
    }

    friend modint998244353 operator+(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs += rhs;
    }

    friend modint998244353 operator-(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs -= rhs;
    }

    friend modint998244353 operator*(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs *= rhs;
    }

    friend modint998244353 operator/(modint998244353 lhs,
                                     const modint998244353 &rhs) {
        return lhs /= rhs;
    }

    friend bool operator==(const modint998244353 &lhs,
                           const modint998244353 &rhs) {
        return lhs.value_ == rhs.value_;
    }

  private:
    static modint998244353 pow(modint998244353 base, long long exponent) {
        modint998244353 result(1);
        while (exponent > 0) {
            if (exponent % 2 == 1) {
                result *= base;
            }
            base *= base;
            exponent /= 2;
        }
        return result;
    }

    std::uint32_t value_;
};

int main() {
    std::ios::sync_with_stdio(false);
    std::cin.tie(nullptr);

    constexpr int max_n = 200000;
    constexpr int max_m = 400000;
    static_assert(max_m < static_cast<int>(modint998244353::mod),
                  "max_m must be less than the modulus.");

    OnlineBinomialSum<modint998244353> online_binomial_sum(max_m,
                                                           modint998244353(-2));
    const modint998244353 minus_inv2 = modint998244353(-1) / modint998244353(2);
    std::vector<modint998244353> power_minus_inv2(max_n + 2);
    power_minus_inv2[0] = modint998244353(1);
    for (int i = 0; i <= max_n; ++i) {
        power_minus_inv2[i + 1] = power_minus_inv2[i] * minus_inv2;
    }

    int test_count;
    std::cin >> test_count;
    while (test_count--) {
        int n, m;
        std::cin >> n >> m;
        const modint998244353 sum =
            online_binomial_sum.binom_prefix_sum(n + 1, 2 * m);
        const modint998244353 ans =
            (sum - modint998244353(1)) *
            (modint998244353(0) - power_minus_inv2[n + 1]);
        std::cout << ans.val() << '\n';
    }

    return 0;
}

Test cases

Env Name Status Elapsed Memory
g++ 06_hand_1 :heavy_check_mark: AC 5749 ms 11 MB
g++ 06_hand_2 :heavy_check_mark: AC 4010 ms 11 MB
g++ 06_hand_3 :heavy_check_mark: AC 2516 ms 11 MB
g++ 06_hand_4 :heavy_check_mark: AC 3871 ms 11 MB
g++ 3_small_1 :heavy_check_mark: AC 1940 ms 11 MB
g++ 3_small_10 :heavy_check_mark: AC 1937 ms 11 MB
g++ 3_small_2 :heavy_check_mark: AC 1867 ms 11 MB
g++ 3_small_3 :heavy_check_mark: AC 1940 ms 11 MB
g++ 3_small_4 :heavy_check_mark: AC 1937 ms 11 MB
g++ 3_small_5 :heavy_check_mark: AC 1928 ms 11 MB
g++ 3_small_6 :heavy_check_mark: AC 1938 ms 11 MB
g++ 3_small_7 :heavy_check_mark: AC 1940 ms 11 MB
g++ 3_small_8 :heavy_check_mark: AC 1935 ms 11 MB
g++ 3_small_9 :heavy_check_mark: AC 1955 ms 11 MB
g++ 4_max_1 :heavy_check_mark: AC 5484 ms 11 MB
g++ 4_max_10 :heavy_check_mark: AC 5481 ms 11 MB
g++ 4_max_2 :heavy_check_mark: AC 5422 ms 11 MB
g++ 4_max_3 :heavy_check_mark: AC 5496 ms 11 MB
g++ 4_max_4 :heavy_check_mark: AC 5413 ms 11 MB
g++ 4_max_5 :heavy_check_mark: AC 5496 ms 11 MB
g++ 4_max_6 :heavy_check_mark: AC 5484 ms 11 MB
g++ 4_max_7 :heavy_check_mark: AC 5483 ms 11 MB
g++ 4_max_8 :heavy_check_mark: AC 5434 ms 11 MB
g++ 4_max_9 :heavy_check_mark: AC 5486 ms 11 MB
g++ 5_large_case_1 :heavy_check_mark: AC 5423 ms 11 MB
g++ 5_large_case_2 :heavy_check_mark: AC 5407 ms 11 MB
g++ 5_large_case_3 :heavy_check_mark: AC 5418 ms 11 MB
g++ 5_large_case_4 :heavy_check_mark: AC 5439 ms 11 MB
g++ 5_large_case_5 :heavy_check_mark: AC 5444 ms 11 MB
clang++ 06_hand_1 :heavy_check_mark: AC 7495 ms 11 MB
clang++ 06_hand_2 :heavy_check_mark: AC 5023 ms 11 MB
clang++ 06_hand_3 :heavy_check_mark: AC 3107 ms 11 MB
clang++ 06_hand_4 :heavy_check_mark: AC 4623 ms 11 MB
clang++ 3_small_1 :heavy_check_mark: AC 2144 ms 11 MB
clang++ 3_small_10 :heavy_check_mark: AC 2139 ms 11 MB
clang++ 3_small_2 :heavy_check_mark: AC 2139 ms 11 MB
clang++ 3_small_3 :heavy_check_mark: AC 2138 ms 11 MB
clang++ 3_small_4 :heavy_check_mark: AC 2140 ms 11 MB
clang++ 3_small_5 :heavy_check_mark: AC 2138 ms 11 MB
clang++ 3_small_6 :heavy_check_mark: AC 2137 ms 11 MB
clang++ 3_small_7 :heavy_check_mark: AC 2137 ms 11 MB
clang++ 3_small_8 :heavy_check_mark: AC 2140 ms 11 MB
clang++ 3_small_9 :heavy_check_mark: AC 2139 ms 11 MB
clang++ 4_max_1 :heavy_check_mark: AC 6732 ms 11 MB
clang++ 4_max_10 :heavy_check_mark: AC 6725 ms 11 MB
clang++ 4_max_2 :heavy_check_mark: AC 6722 ms 11 MB
clang++ 4_max_3 :heavy_check_mark: AC 6750 ms 11 MB
clang++ 4_max_4 :heavy_check_mark: AC 6724 ms 11 MB
clang++ 4_max_5 :heavy_check_mark: AC 6725 ms 11 MB
clang++ 4_max_6 :heavy_check_mark: AC 6748 ms 11 MB
clang++ 4_max_7 :heavy_check_mark: AC 6729 ms 11 MB
clang++ 4_max_8 :heavy_check_mark: AC 6724 ms 11 MB
clang++ 4_max_9 :heavy_check_mark: AC 6730 ms 11 MB
clang++ 5_large_case_1 :heavy_check_mark: AC 6657 ms 11 MB
clang++ 5_large_case_2 :heavy_check_mark: AC 6674 ms 11 MB
clang++ 5_large_case_3 :heavy_check_mark: AC 6669 ms 11 MB
clang++ 5_large_case_4 :heavy_check_mark: AC 6678 ms 11 MB
clang++ 5_large_case_5 :heavy_check_mark: AC 6665 ms 11 MB
Back to top page