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}