This documentation is automatically generated by NotLeonian/competitive-verifier (forked from competitive-verifier/competitive-verifier)
// competitive-verifier: STANDALONE
#include <cassert>
#include <cstdint>
#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_;
};
using mint = modint998244353;
mint brute_prefix_sum(const std::vector<std::vector<mint>> &binomial, int n,
int m, mint r) {
mint ans;
mint power_r(1);
for (int i = 0; i <= m && i < n; ++i) {
ans += power_r * binomial[m][i];
power_r *= r;
}
return ans;
}
mint brute_sum(const std::vector<std::vector<mint>> &binomial, int l, int u,
int m, mint r) {
mint ans;
mint power_r(1);
for (int i = 0; i <= m; ++i) {
if (l <= i && i < u) {
ans += power_r * binomial[m][i];
}
power_r *= r;
}
return ans;
}
void verify(const std::vector<std::vector<mint>> &binomial,
const OnlineBinomialSum<mint> &online_binomial_sum, int max_m,
mint r) {
for (int m = 0; m <= max_m; ++m) {
for (int n = 0; n <= max_m + 10; ++n) {
assert(online_binomial_sum.binom_prefix_sum(n, m) ==
brute_prefix_sum(binomial, n, m, r));
}
for (int l = 0; l <= max_m + 5; ++l) {
for (int u = l; u <= max_m + 10; ++u) {
assert(online_binomial_sum.binom_sum(l, u, m) ==
brute_sum(binomial, l, u, m, r));
}
}
}
}
int main() {
constexpr int max_m = 30;
static_assert(max_m < static_cast<int>(mint::mod),
"max_m must be less than the modulus.");
std::vector<std::vector<mint>> binomial(max_m + 1,
std::vector<mint>(max_m + 1));
for (int n = 0; n <= max_m; ++n) {
binomial[n][0] = mint(1);
binomial[n][n] = mint(1);
for (int k = 1; k < n; ++k) {
binomial[n][k] = binomial[n - 1][k - 1] + binomial[n - 1][k];
}
}
const std::vector<int> bucket_size_list = {1, 2, 3, 4, 5, 7, 16, 31, 64};
for (long long r_value : {-3, -2, -1, 0, 1, 2, 3}) {
const mint r(r_value);
OnlineBinomialSum<mint> online_binomial_sum(max_m, r);
if (r == mint(0)) {
assert(online_binomial_sum.bucket_size == 1);
}
verify(binomial, online_binomial_sum, max_m, r);
for (int bucket_size : bucket_size_list) {
OnlineBinomialSum<mint> online_binomial_sum_with_bucket(
max_m, r, bucket_size);
assert(online_binomial_sum_with_bucket.bucket_size == bucket_size);
verify(binomial, online_binomial_sum_with_bucket, max_m, r);
if (r == mint(0)) {
assert(online_binomial_sum_with_bucket.factorial.empty());
assert(
online_binomial_sum_with_bucket.sample_sum_table.empty());
} else if (r == mint(-1)) {
assert(!online_binomial_sum_with_bucket.factorial.empty());
assert(
online_binomial_sum_with_bucket.sample_sum_table.empty());
} else {
assert(
!online_binomial_sum_with_bucket.sample_sum_table.empty());
}
}
}
for (long long r_value : {-1, 0, 2}) {
OnlineBinomialSum<mint> online_binomial_sum(0, mint(r_value), 7);
assert(online_binomial_sum.binom_prefix_sum(0, 0) == mint(0));
assert(online_binomial_sum.binom_prefix_sum(1, 0) == mint(1));
assert(online_binomial_sum.binom_prefix_sum(10, 0) == mint(1));
}
return 0;
}
#line 1 "verify/standalone-online-binomial-sum.test.cpp"
// competitive-verifier: STANDALONE
#include <cassert>
#include <cstdint>
#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)。
#line 19 "math/combinatorics/online-binomial-sum.hpp"
#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/standalone-online-binomial-sum.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_;
};
using mint = modint998244353;
mint brute_prefix_sum(const std::vector<std::vector<mint>> &binomial, int n,
int m, mint r) {
mint ans;
mint power_r(1);
for (int i = 0; i <= m && i < n; ++i) {
ans += power_r * binomial[m][i];
power_r *= r;
}
return ans;
}
mint brute_sum(const std::vector<std::vector<mint>> &binomial, int l, int u,
int m, mint r) {
mint ans;
mint power_r(1);
for (int i = 0; i <= m; ++i) {
if (l <= i && i < u) {
ans += power_r * binomial[m][i];
}
power_r *= r;
}
return ans;
}
void verify(const std::vector<std::vector<mint>> &binomial,
const OnlineBinomialSum<mint> &online_binomial_sum, int max_m,
mint r) {
for (int m = 0; m <= max_m; ++m) {
for (int n = 0; n <= max_m + 10; ++n) {
assert(online_binomial_sum.binom_prefix_sum(n, m) ==
brute_prefix_sum(binomial, n, m, r));
}
for (int l = 0; l <= max_m + 5; ++l) {
for (int u = l; u <= max_m + 10; ++u) {
assert(online_binomial_sum.binom_sum(l, u, m) ==
brute_sum(binomial, l, u, m, r));
}
}
}
}
int main() {
constexpr int max_m = 30;
static_assert(max_m < static_cast<int>(mint::mod),
"max_m must be less than the modulus.");
std::vector<std::vector<mint>> binomial(max_m + 1,
std::vector<mint>(max_m + 1));
for (int n = 0; n <= max_m; ++n) {
binomial[n][0] = mint(1);
binomial[n][n] = mint(1);
for (int k = 1; k < n; ++k) {
binomial[n][k] = binomial[n - 1][k - 1] + binomial[n - 1][k];
}
}
const std::vector<int> bucket_size_list = {1, 2, 3, 4, 5, 7, 16, 31, 64};
for (long long r_value : {-3, -2, -1, 0, 1, 2, 3}) {
const mint r(r_value);
OnlineBinomialSum<mint> online_binomial_sum(max_m, r);
if (r == mint(0)) {
assert(online_binomial_sum.bucket_size == 1);
}
verify(binomial, online_binomial_sum, max_m, r);
for (int bucket_size : bucket_size_list) {
OnlineBinomialSum<mint> online_binomial_sum_with_bucket(
max_m, r, bucket_size);
assert(online_binomial_sum_with_bucket.bucket_size == bucket_size);
verify(binomial, online_binomial_sum_with_bucket, max_m, r);
if (r == mint(0)) {
assert(online_binomial_sum_with_bucket.factorial.empty());
assert(
online_binomial_sum_with_bucket.sample_sum_table.empty());
} else if (r == mint(-1)) {
assert(!online_binomial_sum_with_bucket.factorial.empty());
assert(
online_binomial_sum_with_bucket.sample_sum_table.empty());
} else {
assert(
!online_binomial_sum_with_bucket.sample_sum_table.empty());
}
}
}
for (long long r_value : {-1, 0, 2}) {
OnlineBinomialSum<mint> online_binomial_sum(0, mint(r_value), 7);
assert(online_binomial_sum.binom_prefix_sum(0, 0) == mint(0));
assert(online_binomial_sum.binom_prefix_sum(1, 0) == mint(1));
assert(online_binomial_sum.binom_prefix_sum(10, 0) == mint(1));
}
return 0;
}