Skip to main content

competitive/math/
binomial_prefix_sum.rs

1use super::{MInt, MIntConvert, MemorizedFactorial, One, Zero};
2use std::{
3    collections::HashMap,
4    fmt::{self, Debug},
5    marker::PhantomData,
6};
7
8pub struct BinomialPrefixSum<M>
9where
10    M: MIntConvert<usize>,
11{
12    query: Vec<(usize, usize)>,
13    _marker: PhantomData<fn() -> M>,
14}
15
16impl<M> Debug for BinomialPrefixSum<M>
17where
18    M: MIntConvert<usize>,
19{
20    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
21        f.debug_struct("BinomialPrefixSum")
22            .field("query", &self.query)
23            .finish()
24    }
25}
26
27impl<M> Default for BinomialPrefixSum<M>
28where
29    M: MIntConvert<usize>,
30{
31    fn default() -> Self {
32        Self {
33            query: Default::default(),
34            _marker: PhantomData,
35        }
36    }
37}
38
39impl<M> BinomialPrefixSum<M>
40where
41    M: MIntConvert<usize>,
42{
43    pub fn new() -> Self {
44        Default::default()
45    }
46
47    pub fn with_capacity(capacity: usize) -> Self {
48        Self {
49            query: Vec::with_capacity(capacity),
50            _marker: PhantomData,
51        }
52    }
53
54    pub fn push(&mut self, n: usize, m: usize) -> usize {
55        let q = self.query.len();
56        self.query.push((m.min(n), n));
57        q
58    }
59
60    pub fn for_each<F>(self, mut f: F)
61    where
62        F: FnMut(usize, MInt<M>),
63    {
64        let query = &self.query;
65        if query.is_empty() {
66            return;
67        }
68        let max_n = query.iter().map(|&(_, n)| n).max().unwrap_or(0);
69        let modulus = M::mod_into();
70        debug_assert!(modulus > 2 && modulus % 2 == 1);
71        debug_assert!(max_n < modulus);
72        debug_assert!(query.iter().all(|&(m, n)| m <= n));
73
74        let fact = MemorizedFactorial::<M>::new(max_n);
75        let inv2 = MInt::<M>::from(2usize).inv();
76        let mut cur = MInt::<M>::one();
77        crate::mo_algorithm!(
78            query,
79            (m, n),
80            |old_m| cur += fact.combination(n, old_m + 1),
81            |new_m| cur -= fact.combination(n, new_m + 1),
82            |old_n| cur = cur + cur - fact.combination(old_n, m),
83            |new_n| cur = (cur + fact.combination(new_n, m)) * inv2,
84            |i| f(i, cur),
85        );
86    }
87
88    pub fn solve(self) -> Vec<MInt<M>> {
89        let mut ans = vec![MInt::zero(); self.query.len()];
90        self.for_each(|i, x| ans[i] = x);
91        ans
92    }
93}
94
95pub struct BinomialPolynomialPrefixSum<M, const K: usize>
96where
97    M: MIntConvert<usize>,
98{
99    query: Vec<(usize, usize, [MInt<M>; K])>,
100}
101
102impl<M, const K: usize> Debug for BinomialPolynomialPrefixSum<M, K>
103where
104    M: MIntConvert<usize>,
105{
106    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
107        f.debug_struct("BinomialPolynomialPrefixSum")
108            .field("query", &self.query)
109            .finish()
110    }
111}
112
113impl<M, const K: usize> Default for BinomialPolynomialPrefixSum<M, K>
114where
115    M: MIntConvert<usize>,
116{
117    fn default() -> Self {
118        Self {
119            query: Default::default(),
120        }
121    }
122}
123
124impl<M, const K: usize> BinomialPolynomialPrefixSum<M, K>
125where
126    M: MIntConvert<usize>,
127{
128    pub fn new() -> Self {
129        Default::default()
130    }
131
132    pub fn with_capacity(capacity: usize) -> Self {
133        Self {
134            query: Vec::with_capacity(capacity),
135        }
136    }
137
138    pub fn push(&mut self, n: usize, m: usize, coef: [MInt<M>; K]) -> usize {
139        let q = self.query.len();
140        self.query.push((n, m, coef));
141        q
142    }
143
144    pub fn for_each_contribution<F>(self, mut f: F)
145    where
146        F: FnMut(usize, MInt<M>),
147    {
148        let mut binom = BinomialPrefixSum::<M>::with_capacity(self.query.len() * K);
149        let mut derived = Vec::with_capacity(self.query.len() * K);
150        let mut stirling = [[MInt::<M>::zero(); K]; K];
151        if K > 0 {
152            stirling[0][0] = MInt::one();
153        }
154        for n in 1..K {
155            for r in 1..=n {
156                stirling[n][r] = stirling[n - 1][r - 1] + MInt::from(r) * stirling[n - 1][r];
157            }
158        }
159        let mut coef_cache = HashMap::with_capacity(self.query.len());
160        for (i, (n, m, coef)) in self.query.iter().enumerate() {
161            let (n, m) = (*n, *m);
162            let coef = *coef_cache.entry(*coef).or_insert_with(|| {
163                let mut falling = [MInt::<M>::zero(); K];
164                for (k, &c) in coef.iter().enumerate() {
165                    if c.is_zero() {
166                        continue;
167                    }
168                    for r in 0..=k {
169                        falling[r] += c * stirling[k][r];
170                    }
171                }
172                falling
173            });
174            let mut falling = MInt::one();
175            for (r, &coef) in coef.iter().enumerate() {
176                if r > n || r > m {
177                    break;
178                }
179                if !coef.is_zero() {
180                    binom.push(n - r, m - r);
181                    derived.push((i, coef * falling));
182                }
183                if r + 1 < K {
184                    falling *= MInt::from(n - r);
185                }
186            }
187        }
188        binom.for_each(|i, x| {
189            let (q, coef) = derived[i];
190            f(q, coef * x);
191        });
192    }
193
194    pub fn solve(self) -> Vec<MInt<M>> {
195        let mut ans = vec![MInt::zero(); self.query.len()];
196        self.for_each_contribution(|i, x| ans[i] += x);
197        ans
198    }
199}
200
201#[cfg(test)]
202mod tests {
203    use super::*;
204    use crate::{
205        num::montgomery::{MInt998244353, Modulo998244353},
206        tools::Xorshift,
207    };
208
209    #[test]
210    fn test_binomial_prefix_sum() {
211        const N: usize = 30;
212        const Q: usize = 300;
213        let mut rng = Xorshift::default();
214        let mut c = vec![vec![MInt998244353::zero(); N + 1]; N + 1];
215        c[0][0] = MInt998244353::one();
216        for i in 1..=N {
217            c[i][0] = MInt998244353::one();
218            c[i][i] = MInt998244353::one();
219            for j in 1..i {
220                c[i][j] = c[i - 1][j - 1] + c[i - 1][j];
221            }
222        }
223        let mut query = vec![];
224        let mut solver = BinomialPrefixSum::<Modulo998244353>::with_capacity(Q);
225        let mut callback_solver = BinomialPrefixSum::<Modulo998244353>::with_capacity(Q);
226        for _ in 0..Q {
227            let (n, m) = rng.random((0..=N, 0..=N * 2));
228            solver.push(n, m);
229            callback_solver.push(n, m);
230            query.push((n, m));
231        }
232        let expected: Vec<MInt998244353> = query
233            .iter()
234            .map(|&(n, m)| (0..=m.min(n)).map(|i| c[n][i]).sum())
235            .collect();
236        let result = solver.solve();
237        assert_eq!(expected, result);
238
239        let mut result = vec![MInt998244353::zero(); query.len()];
240        callback_solver.for_each(|i, x| result[i] = x);
241        assert_eq!(expected, result);
242    }
243
244    #[test]
245    fn test_binomial_polynomial_prefix_sum() {
246        const N: usize = 20;
247        const Q: usize = 300;
248        let mut rng = Xorshift::default();
249        let mut c = vec![vec![MInt998244353::zero(); N + 1]; N + 1];
250        c[0][0] = MInt998244353::one();
251        for i in 1..=N {
252            c[i][0] = MInt998244353::one();
253            c[i][i] = MInt998244353::one();
254            for j in 1..i {
255                c[i][j] = c[i - 1][j - 1] + c[i - 1][j];
256            }
257        }
258
259        macro_rules! check {
260            ($k:expr) => {{
261                const K: usize = $k;
262                let mut query = vec![];
263                let mut solver =
264                    BinomialPolynomialPrefixSum::<Modulo998244353, K>::with_capacity(Q);
265                let mut callback_solver =
266                    BinomialPolynomialPrefixSum::<Modulo998244353, K>::with_capacity(Q);
267                for _ in 0..Q {
268                    let (n, m) = rng.random((0..=N, 0..=N * 2));
269                    let coef = if rng.gen_bool(0.1) {
270                        [MInt998244353::zero(); K]
271                    } else {
272                        std::array::from_fn(|_| rng.random(..))
273                    };
274                    solver.push(n, m, coef);
275                    callback_solver.push(n, m, coef);
276                    query.push((n, m, coef));
277                }
278                let expected: Vec<_> = query
279                    .iter()
280                    .map(|&(n, m, coef)| {
281                        (0..=m.min(n)).fold(MInt998244353::zero(), |s, i| {
282                            let x = MInt998244353::from(i);
283                            let y = coef
284                                .iter()
285                                .rev()
286                                .fold(MInt998244353::zero(), |y, &a| y * x + a);
287                            s + c[n][i] * y
288                        })
289                    })
290                    .collect();
291                let result = solver.solve();
292                assert_eq!(expected, result);
293
294                let mut result = vec![MInt998244353::zero(); query.len()];
295                callback_solver.for_each_contribution(|i, x| result[i] += x);
296                assert_eq!(expected, result);
297            }};
298        }
299        check!(0);
300        check!(1);
301        check!(2);
302        check!(3);
303        check!(4);
304        check!(5);
305    }
306}