From 75efb105b643bbf311cf81c64c2b05e37070822a Mon Sep 17 00:00:00 2001 From: Not_Leonian <75620009+NotLeonian@users.noreply.github.com> Date: Sat, 11 Jul 2026 11:09:47 +0900 Subject: [PATCH 1/2] =?UTF-8?q?=E3=80=8C=E4=BA=8C=E9=A0=85=E4=BF=82?= =?UTF-8?q?=E6=95=B0=E3=81=AE=E5=92=8C=EF=BC=88=E3=82=AA=E3=83=B3=E3=83=A9?= =?UTF-8?q?=E3=82=A4=E3=83=B3=EF=BC=89=E3=80=8D=E3=81=AE=E5=AE=9F=E8=A3=85?= =?UTF-8?q?=E3=82=92=E5=A4=89=E6=9B=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../math/combinatorics/online-binomial-sum.md | 63 ++-- math/combinatorics/online-binomial-sum.hpp | 322 ++++++++---------- .../standalone-online-binomial-sum.test.cpp | 177 ++++++++-- verify/yukicoder-2512.test.cpp | 91 ++--- 4 files changed, 377 insertions(+), 276 deletions(-) diff --git a/docs/math/combinatorics/online-binomial-sum.md b/docs/math/combinatorics/online-binomial-sum.md index 43b11a3..a735348 100644 --- a/docs/math/combinatorics/online-binomial-sum.md +++ b/docs/math/combinatorics/online-binomial-sum.md @@ -8,55 +8,68 @@ documentation_of: math/combinatorics/online-binomial-sum.hpp - 整数 $n,m$ と重み $r$ に対する二項係数の prefix sum を $\displaystyle F(n,m)=\sum_{i=0}^{n-1}r^i\binom{m}{i}$ とおく。 - 半開区間の左端、右端をそれぞれ $l,u$ とする。 $\displaystyle \sum_{i=l}^{u-1}r^i\binom{m}{i}$ をオンラインで求める。 - $\binom{m}{i}=0\;(i>m)$ として扱う。 -- クエリでは、前計算を行ったうちの最も近い点から復元する。 +- $n$ と $m$ をそれぞれバケットに分け、バケット境界と各軸の上端をサンプル座標とする。両軸のサンプル座標の直積で $F(n,m)$ を前計算する。 +- クエリでは、マンハッタン距離が最小であるサンプル点から $n$ 方向と $m$ 方向に遷移する。 - バケットサイズを指定できる。 - - 以下、クエリで与えられる $m$ の最大値を $M$ 、バケットサイズを $B$ とする。 - - クエリの時間計算量は $O(B)$ であり、 $B$ が小さいほど各クエリは高速である。 - - 前計算は時間計算量が $O(M^2/B+B^2+M)$ 、空間計算量が $O((M/B)^2+B^2+M)$ である。 $B$ が小さい場合、時空間ともに前計算のコストが大きい。 - - バケットサイズ $B$ を指定しない場合、 $M$ が正であれば $B^2\le M$ を満たす最大の $2$ の冪が選ばれる。 $M=0$ の場合は $B=1$ 。 -- $r=0$ のときは閉形式を用い、 $r=-1$ でも $r+1$ による除算は行わない。 + +以下、クエリで与えられる $m$ の最大値を $M$ 、バケットサイズを $B$ とする。 +- $r=0$ の場合は閉形式を用いる。 +- $r=-1$ の場合は $\displaystyle F(n,m)=(-1)^{n-1}\binom{m-1}{n-1}$ を用いる(除算を用いない)。ただし、 $n=0$ または $m=0$ の場合は別に処理する。 +- $r$ が $0$ でも $-1$ でもない場合、 $n$ 方向の遷移には境界項 $\displaystyle r^n\binom{m}{n}$ を用いる。 $m$ 方向の遷移にはパスカルの三角形の等式から得られる $\displaystyle F(n,m+1)=(r+1)F(n,m)-r^n\binom{m}{n-1}$ を用いる。 +- バケットサイズ $B$ を指定しない場合、 $M$ が正であれば $B^2\le M$ を満たす最大の $2$ の冪が選ばれる。 $M=0$ の場合は $B=1$ である。 ## 使い方 `max_m` を $M$ 、`bucket_size` を $B$ とおく。 -- `OnlineBinomialSum(int max_m, T r, int bucket_size)` +- `OnlineBinomialSum(int max_m, T r, int bucket_size)` - $0\le m\le M$ のクエリに対する前計算を行う。 - `r` は重みである。 - - 引数の `bucket_size` をバケットサイズとして、前計算を行う。 - - 前提: $M\ge 0,\,B>0$ 。 - - 前提: `T` は四則演算と `T()` との等値比較を持つ。 - - 前提: $r \ne 0$ の場合は `r` で除算できる。 - - 前提: `std::numeric_limits::is_integer` が `false` の場合、 $T(1),T(2),\ldots,T(M)$ で除算できる。 - - 備考: 整数型では中間値が `T` の範囲を超えない必要がある。 -- `OnlineBinomialSum(int max_m, T r = T(1))` + - 引数の `bucket_size` をバケットサイズとして前計算を行う。 + - 前提: $M\ge 0,\;B>0$ 。 + - 前提: `T` は素数 $p$ を法とする体の型であり、整数からの構築、四則演算、等値比較を持つ。 + - 前提: $M::is_integer` が `false` の場合、 $T(1),T(2),\ldots,T(M)$ で除算できる。 - - 備考: 整数型では中間値が `T` の範囲を超えない必要がある。 + - 前提: `T` は素数 $p$ を法とする体の型であり、整数からの構築、四則演算、等値比較を持つ。 + - 前提: $Mm+1$ のときも assert 違反にせず、全体の和を返す。 + - 備考: $n>m+1$ の場合も assert 違反にせず、全体の和を返す。 - `T binom_sum(int l, int u, int m) const` - $\displaystyle \sum_{i=l}^{u-1}r^i\binom{m}{i}$ を返す。 - 前提: $0\le l\le u,\;0\le m\le M$ 。 - - 備考: $u>m+1$ や $l>m$ のときも assert 違反にしない。 + - 備考: $u>m+1$ または $l>m$ の場合も assert 違反にしない。 ## 計算量 `max_m` を $M$ 、`bucket_size` を $B$ とおく。 -- コンストラクタ: 時間 $O(M^2/B+B^2+M)$ 、空間 $O((M/B)^2+B^2+M)$ +$r$ が $0$ でも $-1$ でもない場合、 + +- コンストラクタ: 時間 $O(M^2/B+M)$ 、空間 $O((M/B+1)^2+M)$ - `binom_prefix_sum(n, m)`: 時間 $O(B)$ - `binom_sum(l, u, m)`: 時間 $O(B)$ -コンストラクタで $B$ を指定しない場合、 $B=O(\sqrt M)$ である。 -このとき、 +である。各軸の隣接するサンプル座標の差は高々 $B$ であるため、最も近いサンプル点からクエリ点までの遷移回数は高々 $B$ である。 + +$r=-1$ の場合、 + +- コンストラクタ: 時間 $O(M)$ 、空間 $O(M)$ +- `binom_prefix_sum(n, m)`: 時間 $O(1)$ +- `binom_sum(l, u, m)`: 時間 $O(1)$ + +である。 + +$r=0$ の場合、全ての処理は時間 $O(1)$ 、空間 $O(1)$ である。 + +コンストラクタで $B$ を指定せず、 $r$ が $0$ でも $-1$ でもない場合は $B=O(\sqrt M)$ である。このとき、 - コンストラクタ: 時間 $O(M\sqrt M)$ 、空間 $O(M)$ - `binom_prefix_sum(n, m)`: 時間 $O(\sqrt M)$ diff --git a/math/combinatorics/online-binomial-sum.hpp b/math/combinatorics/online-binomial-sum.hpp index 5770c6b..8d2d729 100644 --- a/math/combinatorics/online-binomial-sum.hpp +++ b/math/combinatorics/online-binomial-sum.hpp @@ -3,36 +3,42 @@ // Σ_{i=l}^{u-1} r^i binom(m,i) をオンラインで求める。 // 0 <= l <= u と 0 <= m <= max_m を仮定する。 -// m のバケット境界の累積和を n のバケット境界でサンプルし、 -// バケット内の二項係数を前計算する。 -// T は四則演算と T() との等値比較を持つ型で、 -// r != T() の場合は r で除算できることを前提とする。 -// std::numeric_limits::is_integer が false の場合は、 -// さらに T(1) から T(max_m) で除算できることを前提とする。 -// B をバケットサイズとして、 -// 時間計算量は前計算 O(max_m^2 / B + B^2 + max_m)、クエリ O(B)。 -// 空間計算量は O((max_m / B)^2 + B^2 + max_m)。 +// n と m のバケット境界および上端の直積で累積和をサンプルし、 +// クエリの点に 2 次元的に最も近いサンプル点から復元する。 +// T は素数を法とする体の型で、法 p について max_m < p を仮定する。 +// r = -1 では交代二項和の閉形式を用いる。 +// B をバケットサイズとして、r が 0, -1 でない場合の時間計算量は +// 前計算 O(max_m^2 / B + max_m)、クエリ O(B)。 +// 空間計算量は O((max_m / B + 1)^2 + max_m)。 #include #include #include template struct OnlineBinomialSum { + static_assert(!std::numeric_limits::is_integer, + "T must be a prime-field element type."); + public: int max_m; int bucket_size; T r; + T r_plus_one; bool r_is_zero; - std::vector prefix_sum_offset; - std::vector prefix_sum_table; - std::vector prefix_term_table; - std::vector weighted_binomial_offset; - std::vector weighted_binomial_table; + bool r_is_minus_one; + T r_plus_one_inverse; + std::vector factorial; + std::vector inverse_factorial; std::vector integer_inverse; - T r_inverse; + std::vector power_r; + std::vector sample_n_list; + std::vector sample_m_list; + std::vector sample_sum_table; explicit OnlineBinomialSum(int max_m, T r, int bucket_size) - : max_m(max_m), bucket_size(bucket_size), r(r), r_is_zero(r == T()) { + : 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); @@ -40,80 +46,74 @@ template struct OnlineBinomialSum { return; } - weighted_binomial_offset.assign(bucket_size + 1, 0); - for (int d = 0; d < bucket_size; ++d) { - weighted_binomial_offset[d + 1] = - weighted_binomial_offset[d] + d + 1; + factorial.assign(max_m + 1, T(1)); + for (int i = 1; i <= max_m; ++i) { + factorial[i] = factorial[i - 1] * T(i); } - weighted_binomial_table.assign(weighted_binomial_offset[bucket_size], - T()); - weighted_binomial_table[0] = T(1); - for (int d = 1; d < bucket_size; ++d) { - weighted_binomial_table[weighted_binomial_offset[d]] = T(1); - for (int j = 1; j < d; ++j) { - weighted_binomial_table[weighted_binomial_offset[d] + j] = - weighted_binomial_table[weighted_binomial_offset[d - 1] + - j] + - r * weighted_binomial_table - [weighted_binomial_offset[d - 1] + j - 1]; - } - weighted_binomial_table[weighted_binomial_offset[d] + d] = - r * weighted_binomial_table[weighted_binomial_offset[d - 1] + - d - 1]; + + 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 constexpr (!std::numeric_limits::is_integer) { - r_inverse = T(1) / r; - integer_inverse.assign(max_m + 1, T()); - for (int i = 1; i <= max_m; ++i) { - integer_inverse[i] = T(1) / T(i); - } + if (r_is_minus_one) { + return; } - const int bucket_count = max_m / bucket_size + 1; - prefix_sum_offset.assign(bucket_count + 1, 0); - for (int b = 0; b < bucket_count; ++b) { - prefix_sum_offset[b + 1] = prefix_sum_offset[b] + b + 2; + 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]; } - prefix_sum_table.assign(prefix_sum_offset[bucket_count], T()); - prefix_term_table.assign(prefix_sum_offset[bucket_count], T()); + 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; + } - for (int b = 0; b < bucket_count; ++b) { - const int base = b * bucket_size; - const int offset = prefix_sum_offset[b]; + 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(sample_n_list.size()); + const int sample_m_count = static_cast(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); - for (int i = 0; i <= base; ++i) { - if (i % bucket_size == 0) { - const int q = i / bucket_size; - prefix_sum_table[offset + q] = sum; - prefix_term_table[offset + q] = term; - } - if (i < base) { + 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; - term *= r; - term *= T(base - i); - if constexpr (std::numeric_limits::is_integer) { - term /= T(i + 1); - } else { - term *= integer_inverse[i + 1]; + 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; } - sum += term; - prefix_sum_table[offset + b + 1] = sum; - prefix_term_table[offset + b + 1] = T(); } } explicit OnlineBinomialSum(int max_m, T r = T(1)) - : OnlineBinomialSum(max_m, r, default_bucket_size(max_m)) {} + : OnlineBinomialSum(max_m, r, default_bucket_size(max_m)) {} T binom_prefix_sum(int n, int m) const { assert(n >= 0); assert(m >= 0); assert(m <= max_m); + if (n == 0) { return T(); } @@ -123,140 +123,104 @@ template struct OnlineBinomialSum { if (r_is_zero) { return T(1); } + if (r_is_minus_one) { + if (m == 0) { + return T(1); + } - const int bucket = m / bucket_size; - const int base = bucket * bucket_size; - const int d = m - base; - const int prefix_offset = prefix_sum_offset[bucket]; - const int weight_offset = weighted_binomial_offset[d]; - - int last_j = d; - if (last_j >= n) { - last_j = n - 1; + T ans = binomial(m - 1, n - 1); + if ((n - 1) % 2 == 1) { + ans = T() - ans; + } + return ans; } - const int first_n = n - last_j; - int sample_index = first_n / bucket_size; - if (sample_index > bucket + 1) { - sample_index = bucket + 1; + 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(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) { + sum += power_r[current_n] * binomial(current_m, current_n); + ++current_n; } - int next_sample_index = sample_index + 1; - if (next_sample_index > bucket + 1) { - next_sample_index = bucket + 1; + + while (current_n > n) { + --current_n; + sum -= power_r[current_n] * binomial(current_m, current_n); } - const auto restore_cost = [&](int index) -> int { - const int sample_n = index * bucket_size; - if (sample_n < first_n) { - return n - sample_n; - } - if (sample_n > n) { - return sample_n - first_n; - } - return n - first_n; - }; - if (restore_cost(next_sample_index) < restore_cost(sample_index)) { - sample_index = next_sample_index; + + while (current_m < m) { + sum *= r_plus_one; + sum -= power_r[current_n] * binomial(current_m, current_n - 1); + ++current_m; } - const int sample_n = sample_index * bucket_size; - const int sample_offset = prefix_offset + sample_index; - const T base_term = prefix_term_table[prefix_offset + bucket]; - - const auto move_right = [&](int current_n, T &sum, T &term) -> void { - sum += term; - if (current_n < base) { - term *= r; - term *= T(base - current_n); - if constexpr (std::numeric_limits::is_integer) { - term /= T(current_n + 1); - } else { - term *= integer_inverse[current_n + 1]; - } - } else { - term = T(); - } - }; - const auto move_left = [&](int current_n, T &sum, T &term) -> void { - if (current_n > base + 1) { - term = T(); - return; - } - if (current_n == base + 1) { - term = base_term; - sum -= term; - return; - } - if constexpr (std::numeric_limits::is_integer) { - term /= r; - term *= T(current_n); - term /= T(base - current_n + 1); - } else { - term *= r_inverse; - term *= T(current_n); - term *= integer_inverse[base - current_n + 1]; - } - sum -= term; - }; - - T sum = prefix_sum_table[sample_offset]; - T term = prefix_term_table[sample_offset]; - T ans = T(); - if (sample_n < first_n) { - for (int current_n = sample_n; current_n < first_n; ++current_n) { - move_right(current_n, sum, term); - } - for (int current_n = first_n; current_n <= n; ++current_n) { - const int j = n - current_n; - ans += weighted_binomial_table[weight_offset + j] * sum; - if (current_n < n) { - move_right(current_n, sum, term); - } - } - } else if (sample_n > n) { - for (int current_n = sample_n; current_n > n; --current_n) { - move_left(current_n, sum, term); - } - for (int current_n = n; current_n >= first_n; --current_n) { - const int j = n - current_n; - ans += weighted_binomial_table[weight_offset + j] * sum; - if (current_n > first_n) { - move_left(current_n, sum, term); - } - } - } else { - T right_sum = sum; - T right_term = term; - for (int current_n = sample_n; current_n <= n; ++current_n) { - const int j = n - current_n; - ans += weighted_binomial_table[weight_offset + j] * right_sum; - if (current_n < n) { - move_right(current_n, right_sum, right_term); - } - } - T left_sum = sum; - T left_term = term; - for (int current_n = sample_n; current_n > first_n; --current_n) { - move_left(current_n, left_sum, left_term); - const int j = n - current_n + 1; - ans += weighted_binomial_table[weight_offset + j] * left_sum; - } + while (current_m > m) { + --current_m; + sum += power_r[current_n] * binomial(current_m, current_n - 1); + sum *= r_plus_one_inverse; } - return ans; + + return sum; } T binom_sum(int l, int u, int m) const { assert(l >= 0); assert(l <= u); + return binom_prefix_sum(u, m) - binom_prefix_sum(l, m); } private: + T binomial(int n, int k) const { + if (k < 0 || k > n) { + return T(); + } + + return factorial[n] * inverse_factorial[k] * inverse_factorial[n - k]; + } + + std::vector make_sample_list(int limit) const { + const int full_bucket_count = limit / bucket_size; + std::vector 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 &sample_list, + int value) const { + int index = value / bucket_size; + const int sample_count = static_cast(sample_list.size()); + if (index >= sample_count) { + index = sample_count - 1; + } + 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 (4LL * bucket_size * bucket_size <= max_m) { - bucket_size <<= 1; + bucket_size *= 2; } return bucket_size; diff --git a/verify/standalone-online-binomial-sum.test.cpp b/verify/standalone-online-binomial-sum.test.cpp index 0536b5d..8c217a2 100644 --- a/verify/standalone-online-binomial-sum.test.cpp +++ b/verify/standalone-online-binomial-sum.test.cpp @@ -1,52 +1,136 @@ // competitive-verifier: STANDALONE #include +#include #include #include "../math/combinatorics/online-binomial-sum.hpp" -// 小さい m で、r = -2,-1,0,1,3 と境界をまたぐ n,l,u を愚直解と比較する。 -// n > m + 1、u > m + 1、空区間、l > m の範囲を検証する。 +class modint998244353 { + public: + static constexpr std::uint32_t mod = 998244353; -long long brute_prefix_sum(const std::vector> &binom, - int n, int m, long long r) { - long long ans = 0; - long long pow_r = 1; - for (int i = 0; i <= m; ++i) { - if (i >= n) { - break; + modint998244353() : value_(0) {} + + modint998244353(long long value) { + long long reduced = value % static_cast(mod); + if (reduced < 0) { + reduced += mod; + } + value_ = static_cast(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; } - ans += pow_r * binom[m][i]; - pow_r *= r; + 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(value_) * rhs.value_ % mod; + value_ = static_cast(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> &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; } -long long brute_sum(const std::vector> &binom, int l, - int u, int m, long long r) { - long long ans = 0; - long long pow_r = 1; +mint brute_sum(const std::vector> &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 += pow_r * binom[m][i]; + ans += power_r * binomial[m][i]; } - pow_r *= r; + power_r *= r; } return ans; } -void verify(const std::vector> &binom, - const OnlineBinomialSum &online_binomial_sum, int max_m, - long long r) { +void verify(const std::vector> &binomial, + const OnlineBinomialSum &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(binom, n, m, r)); + 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(binom, l, u, m, r)); + brute_sum(binomial, l, u, m, r)); } } } @@ -54,23 +138,50 @@ void verify(const std::vector> &binom, int main() { constexpr int max_m = 30; - std::vector> binom( - max_m + 1, std::vector(max_m + 1)); + static_assert(max_m < static_cast(mint::mod)); + + std::vector> binomial(max_m + 1, + std::vector(max_m + 1)); for (int n = 0; n <= max_m; ++n) { - binom[n][0] = binom[n][n] = 1; + binomial[n][0] = mint(1); + binomial[n][n] = mint(1); for (int k = 1; k < n; ++k) { - binom[n][k] = binom[n - 1][k - 1] + binom[n - 1][k]; + binomial[n][k] = binomial[n - 1][k - 1] + binomial[n - 1][k]; } } - for (long long r : {-2, -1, 0, 1, 3}) { - OnlineBinomialSum online_binomial_sum(max_m, r); - verify(binom, online_binomial_sum, max_m, r); + const std::vector 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 online_binomial_sum(max_m, r); + verify(binomial, online_binomial_sum, max_m, r); + + for (int bucket_size : bucket_size_list) { + OnlineBinomialSum 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()); + } + } + } - OnlineBinomialSum online_binomial_sum_with_bucket(max_m, r, - 1); - assert(online_binomial_sum_with_bucket.bucket_size == 1); - verify(binom, online_binomial_sum_with_bucket, max_m, r); + for (long long r_value : {-1, 0, 2}) { + OnlineBinomialSum 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; diff --git a/verify/yukicoder-2512.test.cpp b/verify/yukicoder-2512.test.cpp index cb6bbbe..bd57d04 100644 --- a/verify/yukicoder-2512.test.cpp +++ b/verify/yukicoder-2512.test.cpp @@ -10,77 +10,87 @@ class modint998244353 { public: static constexpr std::uint32_t mod = 998244353; - modint998244353() : v_(0) {} - modint998244353(long long x) { - long long y = x % static_cast(mod); - if (y < 0) { - y += mod; + modint998244353() : value_(0) {} + + modint998244353(long long value) { + long long reduced = value % static_cast(mod); + if (reduced < 0) { + reduced += mod; } - v_ = static_cast(y); + value_ = static_cast(reduced); } - std::uint32_t val() const { return v_; } + std::uint32_t val() const { return value_; } + modint998244353 inv() const { return pow(*this, mod - 2); } + modint998244353 &operator+=(const modint998244353 &rhs) { - std::uint32_t x = v_ + rhs.v_; - if (x >= mod) { - x -= mod; + std::uint32_t value = value_ + rhs.value_; + if (value >= mod) { + value -= mod; } - v_ = x; + value_ = value; return *this; } + modint998244353 &operator-=(const modint998244353 &rhs) { - std::uint32_t x = (v_ >= rhs.v_) ? (v_ - rhs.v_) : (v_ + mod - rhs.v_); - v_ = x; + const std::uint32_t value = value_ >= rhs.value_ + ? value_ - rhs.value_ + : value_ + mod - rhs.value_; + value_ = value; return *this; } + modint998244353 &operator*=(const modint998244353 &rhs) { - std::uint64_t x = static_cast(v_) * rhs.v_ % mod; - v_ = static_cast(x); + const std::uint64_t value = + static_cast(value_) * rhs.value_ % mod; + value_ = static_cast(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.v_ == rhs.v_; - } - friend bool operator!=(const modint998244353 &lhs, - const modint998244353 &rhs) { - return lhs.v_ != rhs.v_; + return lhs.value_ == rhs.value_; } private: - static modint998244353 pow(modint998244353 a, long long e) { - modint998244353 r(1); - while (e > 0) { - if (e & 1) { - r *= a; + static modint998244353 pow(modint998244353 base, long long exponent) { + modint998244353 result(1); + while (exponent > 0) { + if (exponent % 2 == 1) { + result *= base; } - a *= a; - e >>= 1; + base *= base; + exponent /= 2; } - return r; + return result; } - std::uint32_t v_; + std::uint32_t value_; }; int main() { @@ -89,24 +99,27 @@ int main() { constexpr int max_n = 200000; constexpr int max_m = 400000; - OnlineBinomialSum binom_sum(max_m, modint998244353(-2)); + static_assert(max_m < static_cast(modint998244353::mod)); + OnlineBinomialSum online_binomial_sum(max_m, + modint998244353(-2)); const modint998244353 minus_inv2 = modint998244353(-1) / modint998244353(2); - std::vector pow_minus_inv2(max_n + 2); - pow_minus_inv2[0] = modint998244353(1); + std::vector power_minus_inv2(max_n + 2); + power_minus_inv2[0] = modint998244353(1); for (int i = 0; i <= max_n; ++i) { - pow_minus_inv2[i + 1] = pow_minus_inv2[i] * minus_inv2; + power_minus_inv2[i + 1] = power_minus_inv2[i] * minus_inv2; } - int t; - std::cin >> t; - while (t--) { + int test_count; + std::cin >> test_count; + while (test_count--) { int n, m; std::cin >> n >> m; - const modint998244353 s = binom_sum.binom_prefix_sum(n + 1, 2 * m); + const modint998244353 sum = + online_binomial_sum.binom_prefix_sum(n + 1, 2 * m); const modint998244353 ans = - (s - modint998244353(1)) * - (modint998244353(0) - pow_minus_inv2[n + 1]); + (sum - modint998244353(1)) * + (modint998244353(0) - power_minus_inv2[n + 1]); std::cout << ans.val() << '\n'; } From 21278afeac4aceb94c47062352fb3066ac9c3ecf Mon Sep 17 00:00:00 2001 From: Not_Leonian <75620009+NotLeonian@users.noreply.github.com> Date: Sat, 11 Jul 2026 11:30:35 +0900 Subject: [PATCH 2/2] =?UTF-8?q?=E5=BE=AE=E4=BF=AE=E6=AD=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/math/combinatorics/online-binomial-sum.md | 10 +++++++--- math/combinatorics/online-binomial-sum.hpp | 11 ++++++++--- verify/standalone-online-binomial-sum.test.cpp | 6 +++++- verify/yukicoder-2512.test.cpp | 3 ++- 4 files changed, 22 insertions(+), 8 deletions(-) diff --git a/docs/math/combinatorics/online-binomial-sum.md b/docs/math/combinatorics/online-binomial-sum.md index a735348..a57cfa1 100644 --- a/docs/math/combinatorics/online-binomial-sum.md +++ b/docs/math/combinatorics/online-binomial-sum.md @@ -14,9 +14,10 @@ documentation_of: math/combinatorics/online-binomial-sum.hpp 以下、クエリで与えられる $m$ の最大値を $M$ 、バケットサイズを $B$ とする。 - $r=0$ の場合は閉形式を用いる。 -- $r=-1$ の場合は $\displaystyle F(n,m)=(-1)^{n-1}\binom{m-1}{n-1}$ を用いる(除算を用いない)。ただし、 $n=0$ または $m=0$ の場合は別に処理する。 +- $r=-1$ の場合は $\displaystyle F(n,m)=(-1)^{n-1}\binom{m-1}{n-1}$ を用い、 $r+1$ による除算を行わない。ただし、 $n=0$ または $m=0$ の場合は別に処理する。 - $r$ が $0$ でも $-1$ でもない場合、 $n$ 方向の遷移には境界項 $\displaystyle r^n\binom{m}{n}$ を用いる。 $m$ 方向の遷移にはパスカルの三角形の等式から得られる $\displaystyle F(n,m+1)=(r+1)F(n,m)-r^n\binom{m}{n-1}$ を用いる。 -- バケットサイズ $B$ を指定しない場合、 $M$ が正であれば $B^2\le M$ を満たす最大の $2$ の冪が選ばれる。 $M=0$ の場合は $B=1$ である。 +- バケットサイズ $B$ を指定しない場合、 $r=0$ であれば $B=1$ とし、既定のバケットサイズを計算しない。 +- バケットサイズ $B$ を指定せず、 $r\ne 0$ の場合、 $M$ が正であれば $B^2\le M$ を満たす最大の $2$ の冪が選ばれる。 $M=0$ の場合は $B=1$ である。 ## 使い方 @@ -28,14 +29,17 @@ documentation_of: math/combinatorics/online-binomial-sum.hpp - 引数の `bucket_size` をバケットサイズとして前計算を行う。 - 前提: $M\ge 0,\;B>0$ 。 - 前提: `T` は素数 $p$ を法とする体の型であり、整数からの構築、四則演算、等値比較を持つ。 + - 前提: `std::numeric_limits::is_integer` は `false` である。 - 前提: $M::is_integer` は `false` である。 - 前提: $M::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)。 @@ -17,7 +21,7 @@ template struct OnlineBinomialSum { static_assert(!std::numeric_limits::is_integer, - "T must be a prime-field element type."); + "std::numeric_limits::is_integer must be false."); public: int max_m; @@ -107,7 +111,8 @@ template struct OnlineBinomialSum { } explicit OnlineBinomialSum(int max_m, T r = T(1)) - : OnlineBinomialSum(max_m, r, default_bucket_size(max_m)) {} + : OnlineBinomialSum(max_m, r, + r == T() ? 1 : default_bucket_size(max_m)) {} T binom_prefix_sum(int n, int m) const { assert(n >= 0); diff --git a/verify/standalone-online-binomial-sum.test.cpp b/verify/standalone-online-binomial-sum.test.cpp index 8c217a2..411c6e9 100644 --- a/verify/standalone-online-binomial-sum.test.cpp +++ b/verify/standalone-online-binomial-sum.test.cpp @@ -138,7 +138,8 @@ void verify(const std::vector> &binomial, int main() { constexpr int max_m = 30; - static_assert(max_m < static_cast(mint::mod)); + static_assert(max_m < static_cast(mint::mod), + "max_m must be less than the modulus."); std::vector> binomial(max_m + 1, std::vector(max_m + 1)); @@ -154,6 +155,9 @@ int main() { for (long long r_value : {-3, -2, -1, 0, 1, 2, 3}) { const mint r(r_value); OnlineBinomialSum 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) { diff --git a/verify/yukicoder-2512.test.cpp b/verify/yukicoder-2512.test.cpp index bd57d04..5d56cef 100644 --- a/verify/yukicoder-2512.test.cpp +++ b/verify/yukicoder-2512.test.cpp @@ -99,7 +99,8 @@ int main() { constexpr int max_n = 200000; constexpr int max_m = 400000; - static_assert(max_m < static_cast(modint998244353::mod)); + static_assert(max_m < static_cast(modint998244353::mod), + "max_m must be less than the modulus."); OnlineBinomialSum online_binomial_sum(max_m, modint998244353(-2));