Skip to main content

competitive/math/
black_box_matrix.rs

1use super::{Field, Invertible, Matrix, SemiRing};
2use std::{
3    fmt::{self, Debug},
4    marker::PhantomData,
5};
6
7pub trait BlackBoxMatrix<R>
8where
9    R: SemiRing,
10{
11    fn apply(&self, v: &[R::T]) -> Vec<R::T>;
12
13    fn shape(&self) -> (usize, usize);
14}
15
16impl<R> BlackBoxMatrix<R> for Matrix<R>
17where
18    R: SemiRing,
19{
20    fn apply(&self, v: &[R::T]) -> Vec<R::T> {
21        assert_eq!(self.shape.1, v.len());
22        self.data.iter().map(|row| R::dot_product(row, v)).collect()
23    }
24
25    fn shape(&self) -> (usize, usize) {
26        self.shape
27    }
28}
29
30pub struct SparseMatrix<R>
31where
32    R: SemiRing,
33{
34    shape: (usize, usize),
35    nonzero: Vec<(usize, usize, R::T)>,
36}
37
38impl<R> Debug for SparseMatrix<R>
39where
40    R: SemiRing<T: Debug>,
41{
42    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
43        f.debug_struct("SparseMatrix")
44            .field("shape", &self.shape)
45            .field("nonzero", &self.nonzero)
46            .finish()
47    }
48}
49
50impl<R> Clone for SparseMatrix<R>
51where
52    R: SemiRing,
53{
54    fn clone(&self) -> Self {
55        Self {
56            shape: self.shape,
57            nonzero: self.nonzero.clone(),
58        }
59    }
60}
61
62impl<R> SparseMatrix<R>
63where
64    R: SemiRing,
65{
66    pub fn new(shape: (usize, usize)) -> Self {
67        Self {
68            shape,
69            nonzero: vec![],
70        }
71    }
72    pub fn new_with<F>(shape: (usize, usize), f: F) -> Self
73    where
74        R: SemiRing<T: PartialEq>,
75        F: Fn(usize, usize) -> R::T,
76    {
77        let mut nonzero = vec![];
78        for i in 0..shape.0 {
79            for j in 0..shape.1 {
80                let v = f(i, j);
81                if !R::is_zero(&v) {
82                    nonzero.push((i, j, v));
83                }
84            }
85        }
86        Self { shape, nonzero }
87    }
88    pub fn from_nonzero(shape: (usize, usize), nonzero: Vec<(usize, usize, R::T)>) -> Self {
89        Self { shape, nonzero }
90    }
91}
92
93impl<R> SparseMatrix<R>
94where
95    R: Field<T: PartialEq, Additive: Invertible, Multiplicative: Invertible>,
96{
97    pub fn determinant(&self) -> R::T {
98        assert_eq!(self.shape.0, self.shape.1);
99        let n = self.shape.0;
100        let mut columns = vec![Vec::<(usize, R::T)>::new(); n];
101        for &(i, j, ref value) in &self.nonzero {
102            columns[j].push((i, value.clone()));
103        }
104        let mut degrees = vec![0; n];
105        for column in &mut columns {
106            column.sort_unstable_by_key(|&(i, _)| i);
107            let mut merged: Vec<(usize, R::T)> = Vec::with_capacity(column.len());
108            for (i, value) in column.drain(..) {
109                if let Some((_, x)) = merged.last_mut().filter(|(last, _)| *last == i) {
110                    R::add_assign(x, &value);
111                } else {
112                    merged.push((i, value));
113                }
114            }
115            merged.retain(|(i, value)| {
116                if R::is_zero(value) {
117                    false
118                } else {
119                    degrees[*i] += 1;
120                    true
121                }
122            });
123            *column = merged;
124        }
125        let mut order: Vec<_> = (0..n).collect();
126        order.sort_unstable_by_key(|&j| columns[j].len());
127        let mut lower: Vec<Vec<(usize, R::T)>> = Vec::with_capacity(n);
128        let mut pivots: Vec<Option<usize>> = vec![None; n];
129        let mut x = vec![R::zero(); n];
130        let mut seen = vec![0; n];
131        let mut stack = Vec::new();
132        let mut support = Vec::new();
133        let mut determinant = R::one();
134        for (k, &j) in order.iter().enumerate() {
135            support.clear();
136            for &(i, _) in &columns[j] {
137                if seen[i] == k + 1 {
138                    continue;
139                }
140                seen[i] = k + 1;
141                x[i] = R::zero();
142                stack.push((i, 0));
143                while let Some((i, next)) = stack.last_mut() {
144                    if let Some(pivot) = pivots[*i]
145                        && *next < lower[pivot].len()
146                    {
147                        let row = lower[pivot][*next].0;
148                        *next += 1;
149                        if seen[row] != k + 1 {
150                            seen[row] = k + 1;
151                            x[row] = R::zero();
152                            stack.push((row, 0));
153                        }
154                        continue;
155                    }
156                    support.push(*i);
157                    stack.pop();
158                }
159            }
160            for &(i, ref value) in &columns[j] {
161                x[i] = value.clone();
162            }
163            let mut pivot = None;
164            for &i in support.iter().rev() {
165                if let Some(p) = pivots[i] {
166                    let factor = R::neg(&x[i]);
167                    for &(row, ref value) in &lower[p] {
168                        R::add_assign(&mut x[row], &R::mul(&factor, value));
169                    }
170                } else if !R::is_zero(&x[i]) && pivot.is_none_or(|p| degrees[i] < degrees[p]) {
171                    pivot = Some(i);
172                }
173            }
174            let Some(pivot) = pivot else { return R::zero() };
175            R::mul_assign(&mut determinant, &x[pivot]);
176            let inv = R::inv(&x[pivot]);
177            pivots[pivot] = Some(k);
178            lower.push(
179                support
180                    .iter()
181                    .filter(|&&i| pivots[i].is_none() && !R::is_zero(&x[i]))
182                    .map(|&i| (i, R::mul(&x[i], &inv)))
183                    .collect(),
184            );
185        }
186        for mut permutation in [order, pivots.into_iter().map(Option::unwrap).collect()] {
187            for i in 0..n {
188                while permutation[i] != i {
189                    let j = permutation[i];
190                    permutation.swap(i, j);
191                    determinant = R::neg(&determinant);
192                }
193            }
194        }
195        determinant
196    }
197}
198
199impl<R> From<Matrix<R>> for SparseMatrix<R>
200where
201    R: SemiRing<T: PartialEq>,
202{
203    fn from(mat: Matrix<R>) -> Self {
204        let mut nonzero = vec![];
205        for i in 0..mat.shape.0 {
206            for j in 0..mat.shape.1 {
207                let v = mat[(i, j)].clone();
208                if !R::is_zero(&v) {
209                    nonzero.push((i, j, v));
210                }
211            }
212        }
213        Self {
214            shape: mat.shape,
215            nonzero,
216        }
217    }
218}
219
220impl<R> From<SparseMatrix<R>> for Matrix<R>
221where
222    R: SemiRing,
223{
224    fn from(smat: SparseMatrix<R>) -> Self {
225        let mut mat = Matrix::zeros(smat.shape);
226        for &(i, j, ref v) in &smat.nonzero {
227            R::add_assign(&mut mat[(i, j)], v);
228        }
229        mat
230    }
231}
232
233impl<R> BlackBoxMatrix<R> for SparseMatrix<R>
234where
235    R: SemiRing,
236{
237    fn apply(&self, v: &[R::T]) -> Vec<R::T> {
238        assert_eq!(self.shape.1, v.len());
239        let mut res = vec![R::zero(); self.shape.0];
240        for &(i, j, ref val) in &self.nonzero {
241            R::add_assign(&mut res[i], &R::mul(val, &v[j]));
242        }
243        res
244    }
245
246    fn shape(&self) -> (usize, usize) {
247        self.shape
248    }
249}
250
251pub struct BlackBoxMatrixImpl<R, F> {
252    shape: (usize, usize),
253    apply_fn: F,
254    _marker: PhantomData<fn() -> R>,
255}
256
257impl<R, F> Debug for BlackBoxMatrixImpl<R, F>
258where
259    F: Debug,
260{
261    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
262        f.debug_struct("BlackBoxMatrixImpl")
263            .field("shape", &self.shape)
264            .field("apply_fn", &self.apply_fn)
265            .finish()
266    }
267}
268
269impl<R, F> Clone for BlackBoxMatrixImpl<R, F>
270where
271    F: Clone,
272{
273    fn clone(&self) -> Self {
274        Self {
275            shape: self.shape,
276            apply_fn: self.apply_fn.clone(),
277            _marker: PhantomData,
278        }
279    }
280}
281
282impl<R, F> BlackBoxMatrixImpl<R, F> {
283    pub fn new(shape: (usize, usize), apply_fn: F) -> Self {
284        Self {
285            shape,
286            apply_fn,
287            _marker: PhantomData,
288        }
289    }
290}
291
292impl<R, F> BlackBoxMatrix<R> for BlackBoxMatrixImpl<R, F>
293where
294    R: SemiRing,
295    F: Fn(&[R::T]) -> Vec<R::T>,
296{
297    fn apply(&self, v: &[R::T]) -> Vec<R::T> {
298        assert_eq!(self.shape.1, v.len());
299        (self.apply_fn)(v)
300    }
301
302    fn shape(&self) -> (usize, usize) {
303        self.shape
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310    use crate::{
311        algebra::AddMulOperation,
312        math::{BlackBoxMIntMatrix, Convolve998244353},
313        num::{Zero, montgomery::MInt998244353},
314        rand,
315        tools::Xorshift,
316    };
317
318    type R = AddMulOperation<MInt998244353>;
319
320    fn random_matrix(rng: &mut Xorshift, shape: (usize, usize)) -> Matrix<R> {
321        if rng.gen_bool(0.5) {
322            Matrix::<R>::new_with(shape, |_, _| rng.random(..))
323        } else if rng.gen_bool(0.5) {
324            let r = rng.randf();
325            Matrix::<R>::new_with(shape, |_, _| {
326                if rng.gen_bool(r) {
327                    rng.random(..)
328                } else {
329                    MInt998244353::zero()
330                }
331            })
332        } else {
333            let mut mat = Matrix::<R>::new_with(shape, |_, _| rng.random(..));
334            let i0 = rng.random(0..shape.0);
335            let i1 = rng.random(0..shape.0);
336            let x: MInt998244353 = rng.random(..);
337            for j in 0..shape.1 {
338                mat[(i0, j)] = mat[(i1, j)] * x;
339            }
340            mat
341        }
342    }
343
344    #[test]
345    fn test_apply() {
346        let mut rng = Xorshift::default();
347        for _ in 0..100 {
348            rand!(rng, n: 1..30, m: 1..30);
349            let mat = random_matrix(&mut rng, (n, m));
350            let smat = SparseMatrix::from(mat.clone());
351            let v: Vec<_> = (0..m).map(|_| rng.random(..)).collect();
352            let av = mat.apply(&v);
353            let asv = smat.apply(&v);
354            assert_eq!(av, asv);
355        }
356    }
357
358    #[test]
359    fn test_minimal_polynomial() {
360        let mut rng = Xorshift::default();
361        for _ in 0..100 {
362            rand!(rng, n: 1..30);
363            let a = random_matrix(&mut rng, (n, n));
364            let p = a.minimal_polynomial();
365            assert!(!p.is_empty() && p.len() <= n + 1);
366            assert!(p.iter().any(|x| !x.is_zero()));
367            let mut res = Matrix::<R>::zeros((n, n));
368            let mut pow = Matrix::<R>::eye((n, n));
369            for p in p {
370                for i in 0..n {
371                    for j in 0..n {
372                        res[(i, j)] += p * pow[(i, j)];
373                    }
374                }
375                pow = &pow * &a;
376            }
377            assert_eq!(res, Matrix::<R>::zeros((n, n)));
378        }
379    }
380
381    #[test]
382    fn test_apply_pow() {
383        let mut rng = Xorshift::default();
384        for _ in 0..100 {
385            rand!(rng, n: 1..30, k: 0..1_000_000_000);
386            let a = random_matrix(&mut rng, (n, n));
387            let b: Vec<_> = (0..n).map(|_| rng.random(..)).collect();
388            let expected = a.clone().pow(k).apply(&b);
389            let result = a.apply_pow::<Convolve998244353>(b, k);
390            assert_eq!(result, expected);
391        }
392    }
393
394    #[test]
395    fn test_sparse_determinant() {
396        let mut rng = Xorshift::new_with_seed(94623);
397        for _ in 0..500 {
398            let n = rng.random(0..40);
399            let count = rng.random(0..n * n * 2 + 1);
400            let mut entries = Vec::new();
401            for _ in 0..count {
402                let i = rng.random(0..n);
403                let j = rng.random(0..n);
404                let value: MInt998244353 = rng.random(..);
405                entries.push((i, j, value));
406                if rng.gen_bool(0.25) {
407                    entries.push((i, j, -value));
408                }
409            }
410            let sparse = SparseMatrix::<R>::from_nonzero((n, n), entries);
411            let expected = Matrix::from(sparse.clone()).determinant();
412            assert_eq!(sparse.determinant(), expected);
413        }
414    }
415
416    #[test]
417    fn test_black_box_determinant() {
418        let mut rng = Xorshift::default();
419        for _ in 0..100 {
420            rand!(rng, n: 1..30);
421            let mut a = random_matrix(&mut rng, (n, n));
422            let result = a.black_box_determinant();
423            let expected = a.determinant();
424            assert_eq!(result, expected);
425        }
426    }
427
428    #[test]
429    fn test_black_box_linear_equation() {
430        let mut rng = Xorshift::default();
431        for _ in 0..100 {
432            rand!(rng, n: 1..30);
433            let a = random_matrix(&mut rng, (n, n));
434            let b: Vec<_> = (0..n).map(|_| rng.random(..)).collect();
435            let expected = a
436                .solve_system_of_linear_equations(&b)
437                .map(|sol| sol.particular);
438            let result = a.black_box_linear_equation(b);
439            assert_eq!(result, expected);
440        }
441    }
442}