Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 29 additions & 3 deletions docs/math/combinatorics/online-binomial-sum.md
Original file line number Diff line number Diff line change
Expand Up @@ -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<T>(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<T>::is_integer` が `false` の場合、 $T(1),T(2),\ldots,T(M)$ で除算できる。
- 備考: 整数型では中間値が `T` の範囲を超えない必要がある。
- `OnlineBinomialSum<T>(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<T>::is_integer` が `false` の場合、 $T(1),T(2),\ldots,T(M)$ で除算できる。
Expand All @@ -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)$

である。
30 changes: 23 additions & 7 deletions math/combinatorics/online-binomial-sum.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,14 +9,16 @@
// r != T() の場合は r で除算できることを前提とする。
// std::numeric_limits<T>::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 <cassert>
#include <limits>
#include <vector>

template <class T> struct OnlineBinomialSum {
public:
int max_m;
int bucket_size;
T r;
Expand All @@ -29,15 +31,14 @@ template <class T> struct OnlineBinomialSum {
std::vector<T> 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) {
Expand Down Expand Up @@ -106,6 +107,9 @@ template <class T> struct OnlineBinomialSum {
}
}

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

T binom_prefix_sum(int n, int m) const {
assert(n >= 0);
assert(m >= 0);
Expand Down Expand Up @@ -245,6 +249,18 @@ template <class T> 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
35 changes: 23 additions & 12 deletions verify/standalone-online-binomial-sum.test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,23 @@ long long brute_sum(const std::vector<std::vector<long long>> &binom, int l,
return ans;
}

void verify(const std::vector<std::vector<long long>> &binom,
const OnlineBinomialSum<long long> &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<std::vector<long long>> binom(
Expand All @@ -48,18 +65,12 @@ int main() {

for (long long r : {-2, -1, 0, 1, 3}) {
OnlineBinomialSum<long long> 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<long long> 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;
Expand Down
4 changes: 2 additions & 2 deletions verify/yukicoder-2512.test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ int main() {

constexpr int max_n = 200000;
constexpr int max_m = 400000;
OnlineBinomialSum<modint998244353> binom(max_m, modint998244353(-2));
OnlineBinomialSum<modint998244353> binom_sum(max_m, modint998244353(-2));

const modint998244353 minus_inv2 = modint998244353(-1) / modint998244353(2);
std::vector<modint998244353> pow_minus_inv2(max_n + 2);
Expand All @@ -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]);
Expand Down
Loading