Skip to main content

competitive/math/formal_power_series/
berlekamp_massey.rs

1use super::{FormalPowerSeries, FormalPowerSeriesCoefficient, NttReuse, One, Zero};
2use std::mem::{replace, swap};
3
4#[derive(Clone)]
5struct FpsMatrix<T, C> {
6    a00: FormalPowerSeries<T, C>,
7    a01: FormalPowerSeries<T, C>,
8    a10: FormalPowerSeries<T, C>,
9    a11: FormalPowerSeries<T, C>,
10}
11
12struct FrequencyMatrix<C>
13where
14    C: NttReuse,
15{
16    a00: C::F,
17    a01: C::F,
18    a10: C::F,
19    a11: C::F,
20}
21
22impl<C> Clone for FrequencyMatrix<C>
23where
24    C: NttReuse,
25    C::F: Clone,
26{
27    fn clone(&self) -> Self {
28        Self {
29            a00: self.a00.clone(),
30            a01: self.a01.clone(),
31            a10: self.a10.clone(),
32            a11: self.a11.clone(),
33        }
34    }
35}
36
37impl<T, C> FpsMatrix<T, C>
38where
39    T: FormalPowerSeriesCoefficient,
40    C: NttReuse<T = Vec<T>>,
41    C::F: Clone,
42{
43    fn identity() -> Self {
44        Self {
45            a00: FormalPowerSeries::one(),
46            a01: FormalPowerSeries::zero(),
47            a10: FormalPowerSeries::zero(),
48            a11: FormalPowerSeries::one(),
49        }
50    }
51
52    fn multiply_vector(
53        &self,
54        p: &FormalPowerSeries<T, C>,
55        q: &FormalPowerSeries<T, C>,
56    ) -> (FormalPowerSeries<T, C>, FormalPowerSeries<T, C>) {
57        (
58            add(&self.a00 * p, &self.a01 * q),
59            add(&self.a10 * p, &self.a11 * q),
60        )
61    }
62
63    fn left_multiply_step(&mut self, quotient: &[T]) {
64        swap(&mut self.a00, &mut self.a10);
65        swap(&mut self.a01, &mut self.a11);
66        let quotient: FormalPowerSeries<T, C> = FormalPowerSeries::from_vec(quotient.to_vec());
67        self.a10 = add(
68            replace(&mut self.a10, FormalPowerSeries::zero()),
69            quotient.clone() * &self.a00,
70        );
71        self.a11 = add(
72            replace(&mut self.a11, FormalPowerSeries::zero()),
73            quotient * &self.a01,
74        );
75    }
76
77    fn transform(&self, length: usize) -> FrequencyMatrix<C> {
78        FrequencyMatrix {
79            a00: reduced_transform(&self.a00, length),
80            a01: reduced_transform(&self.a01, length),
81            a10: reduced_transform(&self.a10, length),
82            a11: reduced_transform(&self.a11, length),
83        }
84    }
85
86    fn extend_transform(&self, frequency: FrequencyMatrix<C>, length: usize) -> FrequencyMatrix<C> {
87        fn extend<T, C>(fps: &FormalPowerSeries<T, C>, frequency: C::F, length: usize) -> C::F
88        where
89            T: FormalPowerSeriesCoefficient,
90            C: NttReuse<T = Vec<T>>,
91            C::F: Clone,
92        {
93            if fps.length() <= length / 2 {
94                C::ntt_doubling(frequency, false)
95            } else {
96                reduced_transform(fps, length)
97            }
98        }
99
100        FrequencyMatrix {
101            a00: extend(&self.a00, frequency.a00, length),
102            a01: extend(&self.a01, frequency.a01, length),
103            a10: extend(&self.a10, frequency.a10, length),
104            a11: extend(&self.a11, frequency.a11, length),
105        }
106    }
107}
108
109impl<T, C> FrequencyMatrix<C>
110where
111    T: FormalPowerSeriesCoefficient,
112    C: NttReuse<T = Vec<T>>,
113    C::F: Clone,
114{
115    fn product_sum(left_a: &C::F, right_a: &C::F, left_b: &C::F, right_b: &C::F) -> C::F {
116        let mut result = left_a.clone();
117        C::multiply_prefix(&mut result, right_a);
118        C::multiply_add(&mut result, left_b, right_b);
119        result
120    }
121
122    fn multiply(&self, right: &Self) -> Self {
123        Self {
124            a00: Self::product_sum(&self.a00, &right.a00, &self.a01, &right.a10),
125            a01: Self::product_sum(&self.a00, &right.a01, &self.a01, &right.a11),
126            a10: Self::product_sum(&self.a10, &right.a00, &self.a11, &right.a10),
127            a11: Self::product_sum(&self.a10, &right.a01, &self.a11, &right.a11),
128        }
129    }
130
131    fn apply(&self, p: &C::F, q: &C::F, length: usize) -> (Vec<T>, Vec<T>) {
132        (
133            C::inverse_transform_ntt(Self::product_sum(p, &self.a00, q, &self.a01), length),
134            C::inverse_transform_ntt(Self::product_sum(p, &self.a10, q, &self.a11), length),
135        )
136    }
137
138    fn left_multiply_step(self, quotient: &FormalPowerSeries<T, C>, length: usize) -> Self {
139        let negative_quotient = reduced_transform(&(-quotient), length);
140        let mut a10 = self.a00;
141        C::multiply_add(&mut a10, &negative_quotient, &self.a10);
142        let mut a11 = self.a01;
143        C::multiply_add(&mut a11, &negative_quotient, &self.a11);
144        let result = Self {
145            a00: self.a10,
146            a01: self.a11,
147            a10,
148            a11,
149        };
150        if C::MULTIPLE {
151            result.inverse_transform(length).transform(length)
152        } else {
153            result
154        }
155    }
156
157    fn inverse_transform(self, length: usize) -> FpsMatrix<T, C> {
158        FpsMatrix {
159            a00: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a00, length)),
160            a01: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a01, length)),
161            a10: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a10, length)),
162            a11: FormalPowerSeries::from_vec(C::inverse_transform_ntt(self.a11, length)),
163        }
164    }
165}
166
167fn berlekamp_massey_naive<T>(a: &[T], max_work: usize) -> Option<Vec<T>>
168where
169    T: FormalPowerSeriesCoefficient,
170{
171    let n = a.len();
172    let mut b = Vec::with_capacity(n + 1);
173    let mut c = Vec::with_capacity(n + 1);
174    let mut temporary = Vec::with_capacity(n + 1);
175    b.push(T::one());
176    c.push(T::one());
177    let mut y = T::one();
178    let mut work = 0usize;
179    for k in 1..=n {
180        let c_len = c.len();
181        work = work.saturating_add(c_len);
182        if work > max_work {
183            return None;
184        }
185        let mut x = T::zero();
186        for (c, a) in c.iter().zip(&a[k - c_len..]) {
187            x += c.clone() * a.clone();
188        }
189        b.push(T::zero());
190        let b_len = b.len();
191        if x.is_zero() {
192            continue;
193        }
194        let frequency = x.clone() / y.clone();
195        if c_len < b_len {
196            swap(&mut c, &mut temporary);
197            c.clear();
198            c.resize_with(b_len - c_len, T::zero);
199            c.extend(temporary.iter().cloned());
200            for (c, b) in c.iter_mut().rev().zip(b.iter().rev()) {
201                *c -= frequency.clone() * b.clone();
202            }
203            swap(&mut b, &mut temporary);
204            y = x;
205        } else {
206            for (c, b) in c.iter_mut().rev().zip(b.iter().rev()) {
207                *c -= frequency.clone() * b.clone();
208            }
209        }
210    }
211    c.reverse();
212    Some(c)
213}
214
215impl<T, C> FormalPowerSeries<T, C>
216where
217    T: FormalPowerSeriesCoefficient,
218    C: NttReuse<T = Vec<T>>,
219    C::F: Clone,
220{
221    pub fn berlekamp_massey(input: &[T]) -> Self {
222        if input.last().is_none_or(|value| value.is_zero())
223            && input.iter().all(|value| value.is_zero())
224        {
225            return Self::one();
226        }
227        let max_work = if input.len() <= 1536 {
228            usize::MAX
229        } else {
230            input.len().saturating_mul(2)
231        };
232        if let Some(recurrence) = berlekamp_massey_naive(input, max_work) {
233            return Self::from_vec(recurrence);
234        }
235        let n = input.len();
236        let leading_zeros = input.iter().take_while(|value| value.is_zero()).count();
237        let sequence = Self::from_vec(input.to_vec()).trimed();
238        let mut modulus = Self::zeros(n + 1);
239        modulus[n] = T::one();
240        let (matrix, _) = half_gcd(&modulus, &sequence, n / 2, n.max(1).next_power_of_two());
241        let (x, y) = matrix.multiply_vector(&modulus, &sequence);
242        let mut recurrence = if y.length() == 0 {
243            matrix.a01.clone()
244        } else {
245            matrix.a11.clone()
246        };
247        let recurrence_leading_zeros = recurrence
248            .iter()
249            .take_while(|value| value.is_zero())
250            .count();
251        if recurrence_leading_zeros > 0 {
252            let (division, _) = x.div_rem(y.clone());
253            recurrence = add(recurrence * division, matrix.a01);
254        }
255        let inverse = T::one() / &recurrence[0];
256        for value in recurrence.iter_mut() {
257            *value *= &inverse;
258        }
259        let minimum_length = (leading_zeros + 2).max(y.length() + 1);
260        if recurrence.length() < minimum_length {
261            recurrence.resize(minimum_length);
262        }
263        recurrence
264    }
265}
266
267fn degree<T, C>(fps: &FormalPowerSeries<T, C>) -> isize {
268    fps.length() as isize - 1
269}
270
271fn add<T, C>(
272    left: FormalPowerSeries<T, C>,
273    right: FormalPowerSeries<T, C>,
274) -> FormalPowerSeries<T, C>
275where
276    T: FormalPowerSeriesCoefficient,
277{
278    (left + right).trimed()
279}
280
281fn tail<T, C>(fps: &FormalPowerSeries<T, C>, start: isize) -> FormalPowerSeries<T, C>
282where
283    T: FormalPowerSeriesCoefficient,
284{
285    let start = start.max(0) as usize;
286    if start >= fps.length() {
287        FormalPowerSeries::zero()
288    } else {
289        FormalPowerSeries::from_vec(fps.data[start..].to_vec())
290    }
291}
292
293fn coefficient<T, C>(fps: &FormalPowerSeries<T, C>, index: isize) -> T
294where
295    T: FormalPowerSeriesCoefficient,
296{
297    if index < 0 {
298        T::zero()
299    } else {
300        fps.coeff(index as usize)
301    }
302}
303
304fn brute_force<T, C>(
305    mut p: FormalPowerSeries<T, C>,
306    mut q: FormalPowerSeries<T, C>,
307    k: usize,
308) -> FpsMatrix<T, C>
309where
310    T: FormalPowerSeriesCoefficient,
311    C: NttReuse<T = Vec<T>>,
312    C::F: Clone,
313{
314    let threshold = degree(&p) - k as isize;
315    let mut matrix = FpsMatrix::identity();
316    while q.length() as isize > threshold {
317        let q_degree = q.length() - 1;
318        let mut negative_quotient = vec![T::zero(); p.length() - q.length() + 1];
319        let inverse = -T::one() / &q[q_degree];
320        for i in (0..negative_quotient.len()).rev() {
321            negative_quotient[i] = p[i + q_degree].clone() * &inverse;
322            p[i + q_degree] = T::zero();
323            for j in 0..q_degree {
324                let value = negative_quotient[i].clone() * &q[j];
325                p[i + j] += &value;
326            }
327        }
328        matrix.left_multiply_step(&negative_quotient);
329        p.truncate(q_degree);
330        p.trim_tail_zeros();
331        swap(&mut p, &mut q);
332    }
333    matrix
334}
335
336fn reduced_transform<T, C>(fps: &FormalPowerSeries<T, C>, length: usize) -> C::F
337where
338    T: FormalPowerSeriesCoefficient,
339    C: NttReuse<T = Vec<T>>,
340{
341    let mut coefficients = vec![T::zero(); length];
342    for (i, value) in fps.iter().enumerate() {
343        coefficients[i & (length - 1)] += value;
344    }
345    C::transform_ntt(coefficients, length)
346}
347
348fn transform_window<T, C>(fps: &FormalPowerSeries<T, C>, end: isize, length: usize) -> C::F
349where
350    T: FormalPowerSeriesCoefficient,
351    C: NttReuse<T = Vec<T>>,
352{
353    let start = end - length as isize;
354    let coefficients = (start..end).map(|index| coefficient(fps, index)).collect();
355    C::transform_ntt(coefficients, length)
356}
357
358fn half_gcd<T, C>(
359    p: &FormalPowerSeries<T, C>,
360    q: &FormalPowerSeries<T, C>,
361    k: usize,
362    length: usize,
363) -> (FpsMatrix<T, C>, FrequencyMatrix<C>)
364where
365    T: FormalPowerSeriesCoefficient,
366    C: NttReuse<T = Vec<T>>,
367    C::F: Clone,
368{
369    let d = degree(p);
370    if degree(q) < d - k as isize {
371        let matrix = FpsMatrix::identity();
372        let frequency = matrix.transform(length);
373        return (matrix, frequency);
374    }
375    if k == 1 {
376        let matrix = FpsMatrix {
377            a00: FormalPowerSeries::zero(),
378            a01: FormalPowerSeries::one(),
379            a10: FormalPowerSeries::one(),
380            a11: -(tail(p, d - 2) / tail(q, d - 2)),
381        };
382        let frequency = matrix.transform(length);
383        return (matrix, frequency);
384    }
385    if p.length().min(q.length()) <= 32 {
386        let matrix = brute_force(p.clone(), q.clone(), k);
387        let frequency = matrix.transform(length);
388        return (matrix, frequency);
389    }
390
391    let half = length / 2;
392    if k <= half {
393        let (matrix, frequency) = half_gcd(p, q, k, half);
394        let frequency = matrix.extend_transform(frequency, length);
395        return (matrix, frequency);
396    }
397
398    let (matrix, mut matrix_frequency) = half_gcd(
399        &tail(p, d - 2 * half as isize),
400        &tail(q, d - 2 * half as isize),
401        half,
402        length,
403    );
404    let degeneracy = half as isize - degree(&matrix.a11);
405
406    let (p0, q0) = matrix_frequency.apply(
407        &transform_window(p, d - half as isize + degeneracy, length),
408        &transform_window(q, d - half as isize + degeneracy, length),
409        length,
410    );
411    let (p1, q1) = matrix_frequency.apply(
412        &transform_window(p, d - 2 * half as isize, length),
413        &transform_window(q, d - 2 * half as isize, length),
414        length,
415    );
416    let part_length = (half as isize + degeneracy) as usize;
417    let mut p_reduced = p1[length - part_length..].to_vec();
418    p_reduced.extend_from_slice(&p0[length - part_length..]);
419    let mut q_reduced = q1[length - part_length..].to_vec();
420    q_reduced.extend_from_slice(&q0[length - part_length..]);
421    let mut q_reduced: FormalPowerSeries<T, C> = FormalPowerSeries::from_vec(q_reduced).trimed();
422
423    let position = d - half as isize + degeneracy;
424    let mut leading = T::zero();
425    for i in 0..=position {
426        leading += coefficient(p, i) * coefficient(&matrix.a00, position - i)
427            + coefficient(q, i) * coefficient(&matrix.a01, position - i);
428    }
429    p_reduced.push(leading);
430    let mut p_reduced: FormalPowerSeries<T, C> = FormalPowerSeries::from_vec(p_reduced);
431    if degree(&q_reduced) < 3 * half as isize + degeneracy - k as isize {
432        return (matrix, matrix_frequency);
433    }
434
435    let mut remaining = k as isize - degree(&matrix.a11);
436    let mut top_product = matrix.a11.data.last().unwrap().clone();
437    let mut product_degree = degree(&matrix.a11);
438    if degeneracy > 0 {
439        let skip = (2 * half as isize + 2 * degeneracy - (d - half as isize + degeneracy)).max(0);
440        let (division, remainder) = tail(&p_reduced, skip).div_rem(tail(&q_reduced, skip));
441        remaining -= degree(&division);
442        top_product *= -division.data.last().unwrap().clone();
443        product_degree += degree(&division);
444        matrix_frequency = matrix_frequency.left_multiply_step(&division, length);
445        swap(&mut p_reduced, &mut q_reduced);
446        q_reduced = FormalPowerSeries::zeros(skip as usize);
447        q_reduced.data.extend(remainder.data);
448    }
449
450    let start = 3 * half as isize + degeneracy - k as isize - remaining;
451    let (right_matrix, right_frequency) = half_gcd(
452        &tail(&p_reduced, start),
453        &tail(&q_reduced, start),
454        remaining as usize,
455        length,
456    );
457    let product_frequency = right_frequency.multiply(&matrix_frequency);
458    let mut product = product_frequency.clone().inverse_transform(length);
459    product.a00.truncate(k);
460    product.a00.trim_tail_zeros();
461    product.a01.truncate(k);
462    product.a01.trim_tail_zeros();
463    product.a10.truncate(k);
464    product.a10.trim_tail_zeros();
465    product_degree += degree(&right_matrix.a11);
466    if product_degree == length as isize {
467        product.a11.resize(k + 1);
468        let highest = top_product * right_matrix.a11.data.last().unwrap();
469        product.a11[k] = highest.clone();
470        product.a11[0] -= highest;
471    }
472    product.a11.trim_tail_zeros();
473    let product_frequency = if C::MULTIPLE {
474        product.transform(length)
475    } else {
476        product_frequency
477    };
478    (product, product_frequency)
479}
480
481#[cfg(test)]
482mod tests {
483    use crate::math::formal_power_series::berlekamp_massey::berlekamp_massey_naive;
484    use crate::{
485        math::{FormalPowerSeries, FormalPowerSeriesCoefficient, Fps, Fps998244353, NttReuse},
486        num::{MInt, One, Zero, mint_basic::Modulo1000000009, montgomery::MInt998244353},
487        tools::Xorshift,
488    };
489
490    fn verify<T, C>(sequence: &[T], _: &FormalPowerSeries<T, C>)
491    where
492        T: Copy + FormalPowerSeriesCoefficient + std::fmt::Debug,
493        C: NttReuse<T = Vec<T>>,
494        C::F: Clone,
495    {
496        let expected = berlekamp_massey_naive(sequence, usize::MAX).unwrap();
497        let actual: FormalPowerSeries<T, C> = FormalPowerSeries::berlekamp_massey(sequence);
498        assert_eq!(actual.length(), expected.len());
499        for i in actual.length() - 1..sequence.len() {
500            let value = actual
501                .iter()
502                .enumerate()
503                .fold(T::zero(), |sum, (j, &coefficient)| {
504                    sum + coefficient * sequence[i - j]
505                });
506            assert!(value.is_zero());
507        }
508    }
509
510    #[test]
511    fn berlekamp_massey_random() {
512        let mut rng = Xorshift::default();
513        let direct_marker: Fps998244353 = FormalPowerSeries::zero();
514        let arbitrary_marker: Fps<Modulo1000000009> = FormalPowerSeries::zero();
515        for iteration in 0..20 {
516            let n = if iteration < 5 {
517                rng.random(0..=256)
518            } else if iteration == 15 {
519                1537
520            } else {
521                rng.random(257..=600)
522            };
523            let mut direct: Vec<MInt998244353> = rng.random_iter(..).take(n).collect();
524            if iteration >= 5 && iteration % 5 == 0 {
525                let degree = rng.random(n * 5 / 8..n * 7 / 8);
526                let ratio: MInt998244353 = rng.random(2u32..998244351).into();
527                let mut recurrence: Vec<MInt998244353> =
528                    rng.random_iter(..).take(degree + 1).collect();
529                if recurrence[degree].is_zero() {
530                    recurrence[degree] = MInt998244353::one();
531                }
532                direct[0] = MInt998244353::one();
533                for i in 1..degree {
534                    direct[i] = direct[i - 1] * ratio;
535                }
536                for i in degree..n {
537                    direct[i] = (1..=degree).fold(MInt998244353::zero(), |sum, j| {
538                        sum + recurrence[j] * direct[i - j]
539                    });
540                }
541            } else if iteration % 3 == 0 {
542                direct
543                    .iter_mut()
544                    .skip(n * 3 / 4)
545                    .for_each(|x| *x = Zero::zero());
546            }
547            verify(&direct, &direct_marker);
548
549            let mut arbitrary: Vec<MInt<Modulo1000000009>> = rng.random_iter(..).take(n).collect();
550            if iteration % 4 == 0 {
551                arbitrary
552                    .iter_mut()
553                    .take(n / 4)
554                    .for_each(|x| *x = Zero::zero());
555            }
556            verify(&arbitrary, &arbitrary_marker);
557        }
558    }
559}