diff --git a/docs/math/combinatorics/online-binomial-sum.md b/docs/math/combinatorics/online-binomial-sum.md index 0d4a31b..43b11a3 100644 --- a/docs/math/combinatorics/online-binomial-sum.md +++ b/docs/math/combinatorics/online-binomial-sum.md @@ -8,15 +8,32 @@ 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)$ として扱う。 -- クエリでは前計算を行った最も近い点から復元する。 +- クエリでは、前計算を行ったうちの最も近い点から復元する。 +- バケットサイズを指定できる。 + - 以下、クエリで与えられる $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$ による除算は行わない。 ## 使い方 -- `max_m` を $M$ とおく。 +`max_m` を $M$ 、`bucket_size` を $B$ とおく。 + +- `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))` - $0\le m\le M$ のクエリに対する前計算を行う。 - `r` は重みで、省略時は $1$ である。 + - バケットサイズ $B$ は、 $M$ が正であれば $B^2\le M$ を満たす最大の $2$ の冪が選ばれる。 $M=0$ の場合は $B=1$ 。 + - 前提: $M\ge 0$ 。 - 前提: `T` は四則演算と `T()` との等値比較を持つ。 - 前提: $r \ne 0$ の場合は `r` で除算できる。 - 前提: `std::numeric_limits::is_integer` が `false` の場合、 $T(1),T(2),\ldots,T(M)$ で除算できる。 @@ -32,8 +49,17 @@ documentation_of: math/combinatorics/online-binomial-sum.hpp ## 計算量 -`max_m` を $M$ とおく。 +`max_m` を $M$ 、`bucket_size` を $B$ とおく。 + +- コンストラクタ: 時間 $O(M^2/B+B^2+M)$ 、空間 $O((M/B)^2+B^2+M)$ +- `binom_prefix_sum(n, m)`: 時間 $O(B)$ +- `binom_sum(l, u, m)`: 時間 $O(B)$ + +コンストラクタで $B$ を指定しない場合、 $B=O(\sqrt M)$ である。 +このとき、 - コンストラクタ: 時間 $O(M\sqrt M)$ 、空間 $O(M)$ - `binom_prefix_sum(n, m)`: 時間 $O(\sqrt M)$ - `binom_sum(l, u, m)`: 時間 $O(\sqrt M)$ + +である。 diff --git a/math/combinatorics/online-binomial-sum.hpp b/math/combinatorics/online-binomial-sum.hpp index 14ca658..5770c6b 100644 --- a/math/combinatorics/online-binomial-sum.hpp +++ b/math/combinatorics/online-binomial-sum.hpp @@ -9,14 +9,16 @@ // r != T() の場合は r で除算できることを前提とする。 // std::numeric_limits::is_integer が false の場合は、 // さらに T(1) から T(max_m) で除算できることを前提とする。 -// 時間計算量は前計算 O(max_m √max_m)、クエリ O(√max_m)。 -// 空間計算量は O(max_m)。 +// B をバケットサイズとして、 +// 時間計算量は前計算 O(max_m^2 / B + B^2 + max_m)、クエリ O(B)。 +// 空間計算量は O((max_m / B)^2 + B^2 + max_m)。 #include #include #include template struct OnlineBinomialSum { + public: int max_m; int bucket_size; T r; @@ -29,15 +31,14 @@ template struct OnlineBinomialSum { std::vector integer_inverse; T r_inverse; - explicit OnlineBinomialSum(int m, T r = T(1)) - : max_m(m), bucket_size(1), r(r), r_is_zero(r == T()) { + 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()) { assert(max_m >= 0); + assert(bucket_size > 0); + if (r_is_zero) { return; } - while (1LL * bucket_size * bucket_size < max_m + 1) { - bucket_size <<= 1; - } weighted_binomial_offset.assign(bucket_size + 1, 0); for (int d = 0; d < bucket_size; ++d) { @@ -106,6 +107,9 @@ template struct OnlineBinomialSum { } } + explicit OnlineBinomialSum(int max_m, T r = T(1)) + : OnlineBinomialSum(max_m, r, default_bucket_size(max_m)) {} + T binom_prefix_sum(int n, int m) const { assert(n >= 0); assert(m >= 0); @@ -245,6 +249,18 @@ template struct OnlineBinomialSum { assert(l <= u); return binom_prefix_sum(u, m) - binom_prefix_sum(l, m); } + + private: + 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; + } + + return bucket_size; + } }; #endif diff --git a/verify/standalone-online-binomial-sum.test.cpp b/verify/standalone-online-binomial-sum.test.cpp index 525d338..0536b5d 100644 --- a/verify/standalone-online-binomial-sum.test.cpp +++ b/verify/standalone-online-binomial-sum.test.cpp @@ -35,6 +35,23 @@ long long brute_sum(const std::vector> &binom, int l, return ans; } +void verify(const std::vector> &binom, + const OnlineBinomialSum &online_binomial_sum, int max_m, + long long 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)); + } + 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)); + } + } + } +} + int main() { constexpr int max_m = 30; std::vector> binom( @@ -48,18 +65,12 @@ int main() { for (long long r : {-2, -1, 0, 1, 3}) { OnlineBinomialSum online_binomial_sum(max_m, 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)); - } - 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)); - } - } - } + verify(binom, online_binomial_sum, max_m, r); + + 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); } return 0; diff --git a/verify/yukicoder-2512.test.cpp b/verify/yukicoder-2512.test.cpp index 7bd5473..cb6bbbe 100644 --- a/verify/yukicoder-2512.test.cpp +++ b/verify/yukicoder-2512.test.cpp @@ -89,7 +89,7 @@ int main() { constexpr int max_n = 200000; constexpr int max_m = 400000; - OnlineBinomialSum binom(max_m, modint998244353(-2)); + OnlineBinomialSum binom_sum(max_m, modint998244353(-2)); const modint998244353 minus_inv2 = modint998244353(-1) / modint998244353(2); std::vector pow_minus_inv2(max_n + 2); @@ -103,7 +103,7 @@ int main() { while (t--) { int n, m; std::cin >> n >> m; - const modint998244353 s = binom.binom_prefix_sum(n + 1, 2 * m); + const modint998244353 s = binom_sum.binom_prefix_sum(n + 1, 2 * m); const modint998244353 ans = (s - modint998244353(1)) * (modint998244353(0) - pow_minus_inv2[n + 1]);