Skip to main content

competitive/math/
matrix.rs

1use super::{Field, Invertible, Ring, SemiRing};
2use std::{
3    fmt::{self, Debug},
4    marker::PhantomData,
5    ops::{Add, AddAssign, Index, IndexMut, Mul, MulAssign, Neg, Sub, SubAssign},
6};
7
8pub struct Matrix<R>
9where
10    R: SemiRing,
11{
12    pub shape: (usize, usize),
13    pub data: Vec<Vec<R::T>>,
14    _marker: PhantomData<fn() -> R>,
15}
16
17impl<R> Debug for Matrix<R>
18where
19    R: SemiRing<T: Debug>,
20{
21    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
22        f.debug_struct("Matrix")
23            .field("shape", &self.shape)
24            .field("data", &self.data)
25            .field("_marker", &self._marker)
26            .finish()
27    }
28}
29
30impl<R> Clone for Matrix<R>
31where
32    R: SemiRing,
33{
34    fn clone(&self) -> Self {
35        Self {
36            shape: self.shape,
37            data: self.data.clone(),
38            _marker: self._marker,
39        }
40    }
41}
42
43impl<R> PartialEq for Matrix<R>
44where
45    R: SemiRing<T: PartialEq>,
46{
47    fn eq(&self, other: &Self) -> bool {
48        self.shape == other.shape && self.data == other.data
49    }
50}
51
52impl<R> Eq for Matrix<R> where R: SemiRing<T: Eq> {}
53
54impl<R> Matrix<R>
55where
56    R: SemiRing,
57{
58    pub fn new(shape: (usize, usize), z: R::T) -> Self {
59        Self {
60            shape,
61            data: vec![vec![z; shape.1]; shape.0],
62            _marker: PhantomData,
63        }
64    }
65
66    pub fn from_vec(data: Vec<Vec<R::T>>) -> Self {
67        let shape = (data.len(), data.first().map(Vec::len).unwrap_or_default());
68        assert!(data.iter().all(|r| r.len() == shape.1));
69        Self {
70            shape,
71            data,
72            _marker: PhantomData,
73        }
74    }
75
76    pub fn new_with(shape: (usize, usize), mut f: impl FnMut(usize, usize) -> R::T) -> Self {
77        let data = (0..shape.0)
78            .map(|i| (0..shape.1).map(|j| f(i, j)).collect())
79            .collect();
80        Self {
81            shape,
82            data,
83            _marker: PhantomData,
84        }
85    }
86
87    pub fn zeros(shape: (usize, usize)) -> Self {
88        Self {
89            shape,
90            data: vec![vec![R::zero(); shape.1]; shape.0],
91            _marker: PhantomData,
92        }
93    }
94
95    pub fn eye(shape: (usize, usize)) -> Self {
96        let mut data = vec![vec![R::zero(); shape.1]; shape.0];
97        for (i, d) in data.iter_mut().enumerate().take(shape.1) {
98            d[i] = R::one();
99        }
100        Self {
101            shape,
102            data,
103            _marker: PhantomData,
104        }
105    }
106
107    pub fn transpose(&self) -> Self {
108        Self::new_with((self.shape.1, self.shape.0), |i, j| self[j][i].clone())
109    }
110
111    pub fn map<S, F>(&self, mut f: F) -> Matrix<S>
112    where
113        S: SemiRing,
114        F: FnMut(&R::T) -> S::T,
115    {
116        Matrix::<S>::new_with(self.shape, |i, j| f(&self[i][j]))
117    }
118
119    pub fn add_row_with(&mut self, mut f: impl FnMut(usize, usize) -> R::T) {
120        self.data
121            .push((0..self.shape.1).map(|j| f(self.shape.0, j)).collect());
122        self.shape.0 += 1;
123    }
124
125    pub fn add_col_with(&mut self, mut f: impl FnMut(usize, usize) -> R::T) {
126        for i in 0..self.shape.0 {
127            self.data[i].push(f(i, self.shape.1));
128        }
129        self.shape.1 += 1;
130    }
131
132    pub fn pairwise_assign<F>(&mut self, other: &Self, mut f: F)
133    where
134        F: FnMut(&mut R::T, &R::T),
135    {
136        assert_eq!(self.shape, other.shape);
137        for i in 0..self.shape.0 {
138            for j in 0..self.shape.1 {
139                f(&mut self[i][j], &other[i][j]);
140            }
141        }
142    }
143}
144
145#[derive(Debug)]
146pub struct SystemOfLinearEquationsSolution<R>
147where
148    R: Field<Additive: Invertible, Multiplicative: Invertible>,
149{
150    pub particular: Vec<R::T>,
151    pub basis: Vec<Vec<R::T>>,
152}
153
154impl<R> Matrix<R>
155where
156    R: Field<T: PartialEq, Additive: Invertible, Multiplicative: Invertible>,
157{
158    fn eliminate<const DETERMINANT: bool>(&mut self) -> (usize, R::T) {
159        let (n, m) = self.shape;
160        let mut rank = 0;
161        let mut determinant = R::one();
162        let mut negative = false;
163        for first in (0..m).step_by(64) {
164            if rank == n {
165                break;
166            }
167            let end = (first + 64).min(m);
168            let start = rank;
169            let mut pivots = Vec::new();
170            let mut panel = vec![vec![R::zero(); end - first]; end - first];
171            for col in first..end {
172                if panel[col - first][..col - first]
173                    .iter()
174                    .any(|x| !R::is_zero(x))
175                {
176                    for row in &mut self.data[rank..] {
177                        let value =
178                            R::dot_product(&row[first..col], &panel[col - first][..col - first]);
179                        R::sub_assign(&mut row[col], &value);
180                    }
181                }
182                let Some(pivot) = (rank..n).find(|&i| !R::is_zero(&self[i][col])) else {
183                    continue;
184                };
185                if pivot != rank {
186                    self.data.swap(rank, pivot);
187                    negative = !negative;
188                }
189                if DETERMINANT {
190                    R::mul_assign(&mut determinant, &self[rank][col]);
191                }
192                let inv = R::inv(&self[rank][col]);
193                let row = &mut self.data[rank];
194                for c in col + 1..end {
195                    let value = R::dot_product(&row[first..col], &panel[c - first][..col - first]);
196                    R::sub_assign(&mut row[c], &value);
197                    panel[c - first][col - first] = row[c].clone();
198                }
199                for row in &mut self.data[rank + 1..] {
200                    R::mul_assign(&mut row[col], &inv);
201                }
202                pivots.push(col);
203                rank += 1;
204                if rank == n {
205                    break;
206                }
207            }
208            for i in start..if rank - start < 32 { n } else { rank } {
209                let (upper, lower) = self.data.split_at_mut(i);
210                let row = &mut lower[0];
211                for (j, &col) in pivots[..(i - start).min(pivots.len())].iter().enumerate() {
212                    if R::is_zero(&row[col]) {
213                        continue;
214                    }
215                    let factor = R::neg(&row[col]);
216                    R::add_scaled_assign(&mut row[end..], &upper[start + j][end..], &factor);
217                }
218            }
219            if rank < n
220                && end < m
221                && rank - start >= 32
222                && self.data[rank..]
223                    .iter()
224                    .any(|row| pivots.iter().any(|&col| !R::is_zero(&row[col])))
225            {
226                let lower = Self::new_with((n - rank, rank - start), |i, j| {
227                    R::neg(&self[rank + i][pivots[j]])
228                });
229                let upper = Self::new_with((rank - start, m - end), |i, j| {
230                    self[start + i][end + j].clone()
231                });
232                let update = &lower * &upper;
233                for (row, update) in self.data[rank..].iter_mut().zip(update.data) {
234                    for (x, y) in row[end..].iter_mut().zip(update) {
235                        R::add_assign(x, &y);
236                    }
237                }
238            }
239            for (i, &col) in pivots.iter().enumerate() {
240                for row in &mut self.data[start + i + 1..] {
241                    row[col] = R::zero();
242                }
243            }
244            if DETERMINANT && rank < end {
245                return (rank, R::zero());
246            }
247        }
248        if DETERMINANT && negative {
249            determinant = R::neg(&determinant);
250        }
251        (rank, determinant)
252    }
253
254    /// f: (row, pivot_row, col)
255    pub fn row_reduction_with<F>(&mut self, normalize: bool, mut f: F)
256    where
257        F: FnMut(usize, usize, usize),
258    {
259        let (n, m) = self.shape;
260        let mut c = 0;
261        for r in 0..n {
262            loop {
263                if c >= m {
264                    return;
265                }
266                if let Some(pivot) = (r..n).find(|&p| !R::is_zero(&self[p][c])) {
267                    f(r, pivot, c);
268                    self.data.swap(r, pivot);
269                    break;
270                };
271                c += 1;
272            }
273            let d = R::inv(&self[r][c]);
274            if normalize {
275                for value in &mut self[r][c..m] {
276                    R::mul_assign(value, &d);
277                }
278            }
279            for i in (0..n).filter(|&i| i != r) {
280                let mut e = self[i][c].clone();
281                if !normalize {
282                    R::mul_assign(&mut e, &d);
283                }
284                for j in c..m {
285                    let e = R::mul(&e, &self[r][j]);
286                    R::sub_assign(&mut self[i][j], &e);
287                }
288            }
289            c += 1;
290        }
291    }
292
293    pub fn row_reduction(&mut self, normalize: bool) {
294        self.row_reduction_with(normalize, |_, _, _| {});
295    }
296
297    pub fn rank(&mut self) -> usize {
298        self.eliminate::<false>().0
299    }
300
301    pub fn determinant(&mut self) -> R::T {
302        assert_eq!(self.shape.0, self.shape.1);
303        self.eliminate::<true>().1
304    }
305
306    pub fn solve_system_of_linear_equations(
307        &self,
308        b: &[R::T],
309    ) -> Option<SystemOfLinearEquationsSolution<R>> {
310        assert_eq!(self.shape.0, b.len());
311        let m = self.shape.1;
312        let mut a = Self::new_with((self.shape.0, m + 1), |i, j| {
313            if j == m {
314                b[i].clone()
315            } else {
316                self[i][j].clone()
317            }
318        });
319        let rank = a.eliminate::<false>().0;
320        let mut pivots = Vec::with_capacity(rank);
321        let mut b = Vec::with_capacity(rank);
322        for row in &a.data[..rank] {
323            let c = row.iter().position(|x| !R::is_zero(x)).unwrap();
324            if c == m {
325                return None;
326            }
327            pivots.push(c);
328            b.push(row[m].clone());
329        }
330
331        let mut free = Vec::with_capacity(m - rank);
332        let mut pivot = 0;
333        for c in 0..m {
334            if pivot < rank && pivots[pivot] == c {
335                pivot += 1;
336            } else {
337                free.push(c);
338            }
339        }
340        let mut coefficients: Vec<Vec<_>> = (0..rank)
341            .map(|i| free.iter().map(|&c| a[i][c].clone()).collect())
342            .collect();
343        for k in (0..rank).rev() {
344            let c = pivots[k];
345            let inv = R::inv(&a[k][c]);
346            R::mul_assign(&mut b[k], &inv);
347            let pivot_b = b[k].clone();
348            let (upper, lower) = coefficients.split_at_mut(k);
349            let pivot_coefficients = &mut lower[0];
350            for x in pivot_coefficients.iter_mut() {
351                R::mul_assign(x, &inv);
352            }
353            for ((row, value), coefficients) in a.data[..k].iter_mut().zip(&mut b[..k]).zip(upper) {
354                if R::is_zero(&row[c]) {
355                    continue;
356                }
357                let factor = row[c].clone();
358                row[c] = R::zero();
359                R::sub_assign(value, &R::mul(&factor, &pivot_b));
360                R::add_scaled_assign(coefficients, pivot_coefficients, &R::neg(&factor));
361            }
362        }
363
364        let mut particular = vec![R::zero(); m];
365        for i in 0..rank {
366            particular[pivots[i]] = b[i].clone();
367        }
368        let mut basis = Vec::with_capacity(free.len());
369        for (j, &c) in free.iter().enumerate() {
370            let mut vector = vec![R::zero(); m];
371            vector[c] = R::one();
372            for i in 0..rank {
373                vector[pivots[i]] = R::neg(&coefficients[i][j]);
374            }
375            basis.push(vector);
376        }
377        Some(SystemOfLinearEquationsSolution { particular, basis })
378    }
379
380    pub fn inverse(&self) -> Option<Matrix<R>> {
381        assert_eq!(self.shape.0, self.shape.1);
382        let n = self.shape.0;
383        if n >= 64 {
384            let m = n / 2;
385            let a = Self::new_with((m, m), |i, j| self[i][j].clone());
386            if let Some(mut ai) = a.inverse() {
387                let b = Self::new_with((m, n - m), |i, j| self[i][j + m].clone());
388                let c = Self::new_with((n - m, m), |i, j| self[i + m][j].clone());
389                let mut d = Self::new_with((n - m, n - m), |i, j| self[i + m][j + m].clone());
390                let u = &ai * &b;
391                let v = &c * &ai;
392                d -= &v * &b;
393                let di = d.inverse()?;
394                let r = &u * &di;
395                let t = &di * &v;
396                ai += &r * &v;
397                let mut inverse = Self::zeros((n, n));
398                for i in 0..m {
399                    inverse[i][..m].clone_from_slice(&ai[i]);
400                    for (x, y) in inverse[i][m..].iter_mut().zip(&r[i]) {
401                        *x = R::neg(y);
402                    }
403                }
404                for i in m..n {
405                    for (x, y) in inverse[i][..m].iter_mut().zip(&t[i - m]) {
406                        *x = R::neg(y);
407                    }
408                    inverse[i][m..].clone_from_slice(&di[i - m]);
409                }
410                return Some(inverse);
411            }
412        }
413        let mut a = self.clone();
414        let mut inverse = Self::eye((n, n));
415        let mut ranges: Vec<_> = (0..n).map(|i| (i, i + 1)).collect();
416        for r in 0..n {
417            let pivot = (r..n).find(|&i| !R::is_zero(&a[i][r]))?;
418            a.data.swap(r, pivot);
419            inverse.data.swap(r, pivot);
420            ranges.swap(r, pivot);
421
422            let d = R::inv(&a[r][r]);
423            for x in &mut a[r][r..] {
424                R::mul_assign(x, &d);
425            }
426            let (left, right) = ranges[r];
427            for x in &mut inverse[r][left..right] {
428                R::mul_assign(x, &d);
429            }
430
431            let (a_upper, a_lower) = a.data.split_at_mut(r + 1);
432            let pivot_a = &a_upper[r];
433            let (inverse_upper, inverse_lower) = inverse.data.split_at_mut(r + 1);
434            let pivot_inverse = &inverse_upper[r];
435            let (ranges_upper, ranges_lower) = ranges.split_at_mut(r + 1);
436            let (left, right) = ranges_upper[r];
437            for ((a, inverse), range) in a_lower.iter_mut().zip(inverse_lower).zip(ranges_lower) {
438                if R::is_zero(&a[r]) {
439                    continue;
440                }
441                let e = a[r].clone();
442                a[r] = R::zero();
443                R::add_scaled_assign(&mut a[(r + 1)..], &pivot_a[(r + 1)..], &R::neg(&e));
444                R::add_scaled_assign(
445                    &mut inverse[left..right],
446                    &pivot_inverse[left..right],
447                    &R::neg(&e),
448                );
449                range.0 = range.0.min(left);
450                range.1 = range.1.max(right);
451            }
452        }
453        for r in (0..n).rev() {
454            let (left, right) = ranges[r];
455            let (inverse_upper, inverse_lower) = inverse.data.split_at_mut(r);
456            let pivot_inverse = &inverse_lower[0];
457            let (ranges_upper, _) = ranges.split_at_mut(r);
458            for ((a, inverse), range) in a.data[..r].iter_mut().zip(inverse_upper).zip(ranges_upper)
459            {
460                if R::is_zero(&a[r]) {
461                    continue;
462                }
463                let e = a[r].clone();
464                a[r] = R::zero();
465                R::add_scaled_assign(
466                    &mut inverse[left..right],
467                    &pivot_inverse[left..right],
468                    &R::neg(&e),
469                );
470                range.0 = range.0.min(left);
471                range.1 = range.1.max(right);
472            }
473        }
474        Some(inverse)
475    }
476
477    pub fn characteristic_polynomial(&mut self) -> Vec<R::T> {
478        let n = self.shape.0;
479        if n == 0 {
480            return vec![R::one()];
481        }
482        assert!(self.data.iter().all(|a| a.len() == n));
483        for j in 0..(n - 1) {
484            if let Some(x) = ((j + 1)..n).find(|&x| !R::is_zero(&self[x][j])) {
485                self.data.swap(j + 1, x);
486                self.data.iter_mut().for_each(|a| a.swap(j + 1, x));
487                let inv = R::inv(&self[j + 1][j]);
488                let mut v = vec![];
489                let src = std::mem::take(&mut self[j + 1]);
490                for a in self.data[(j + 2)..].iter_mut() {
491                    let mul = R::mul(&a[j], &inv);
492                    R::add_scaled_assign(&mut a[j..], &src[j..], &R::neg(&mul));
493                    v.push(mul);
494                }
495                self[j + 1] = src;
496                for a in self.data.iter_mut() {
497                    let v = R::dot_product(&a[(j + 2)..], &v);
498                    R::add_assign(&mut a[j + 1], &v);
499                }
500            }
501        }
502        // dp[k][j - k] stores [x^k] det(xI - A[..j, ..j]).
503        let mut dp: Vec<Vec<R::T>> = (0..=n).map(|i| Vec::with_capacity(n + 1 - i)).collect();
504        dp[0].push(R::one());
505        for i in 0..n {
506            let mut c = vec![R::zero(); i + 1];
507            c[i] = R::neg(&self[i][i]);
508            let mut mul = R::one();
509            for j in (0..i).rev() {
510                mul = R::mul(&mul, &self[j + 1][j]);
511                c[j] = R::neg(&R::mul(&mul, &self[j][i]));
512            }
513            for k in (0..=i).rev() {
514                let mut value = R::dot_product(&dp[k], &c[k..]);
515                if k > 0 {
516                    R::add_assign(&mut value, dp[k - 1].last().unwrap());
517                }
518                dp[k].push(value);
519            }
520            dp[i + 1].push(R::one());
521        }
522        dp.into_iter().map(|mut c| c.pop().unwrap()).collect()
523    }
524}
525
526impl<R> Index<usize> for Matrix<R>
527where
528    R: SemiRing,
529{
530    type Output = Vec<R::T>;
531    fn index(&self, index: usize) -> &Self::Output {
532        &self.data[index]
533    }
534}
535
536impl<R> IndexMut<usize> for Matrix<R>
537where
538    R: SemiRing,
539{
540    fn index_mut(&mut self, index: usize) -> &mut Self::Output {
541        &mut self.data[index]
542    }
543}
544
545impl<R> Index<(usize, usize)> for Matrix<R>
546where
547    R: SemiRing,
548{
549    type Output = R::T;
550    fn index(&self, index: (usize, usize)) -> &Self::Output {
551        &self.data[index.0][index.1]
552    }
553}
554
555impl<R> IndexMut<(usize, usize)> for Matrix<R>
556where
557    R: SemiRing,
558{
559    fn index_mut(&mut self, index: (usize, usize)) -> &mut Self::Output {
560        &mut self.data[index.0][index.1]
561    }
562}
563
564macro_rules! impl_matrix_pairwise_binop {
565    ($imp:ident, $method:ident, $imp_assign:ident, $method_assign:ident $(where [$($clauses:tt)*])?) => {
566        impl<R> $imp_assign for Matrix<R>
567        where
568            R: SemiRing,
569            $($($clauses)*)?
570        {
571            fn $method_assign(&mut self, rhs: Self) {
572                self.pairwise_assign(&rhs, |a, b| R::$method_assign(a, b));
573            }
574        }
575        impl<R> $imp_assign<&Matrix<R>> for Matrix<R>
576        where
577            R: SemiRing,
578            $($($clauses)*)?
579        {
580            fn $method_assign(&mut self, rhs: &Self) {
581                self.pairwise_assign(rhs, |a, b| R::$method_assign(a, b));
582            }
583        }
584        impl<R> $imp for Matrix<R>
585        where
586            R: SemiRing,
587            $($($clauses)*)?
588        {
589            type Output = Matrix<R>;
590            fn $method(mut self, rhs: Self) -> Self::Output {
591                self.$method_assign(rhs);
592                self
593            }
594        }
595        impl<R> $imp<&Matrix<R>> for Matrix<R>
596        where
597            R: SemiRing,
598            $($($clauses)*)?
599        {
600            type Output = Matrix<R>;
601            fn $method(mut self, rhs: &Self) -> Self::Output {
602                self.$method_assign(rhs);
603                self
604            }
605        }
606        impl<R> $imp<Matrix<R>> for &Matrix<R>
607        where
608            R: SemiRing,
609            $($($clauses)*)?
610        {
611            type Output = Matrix<R>;
612            fn $method(self, mut rhs: Matrix<R>) -> Self::Output {
613                rhs.pairwise_assign(self, |a, b| *a = R::$method(b, a));
614                rhs
615            }
616        }
617        impl<R> $imp<&Matrix<R>> for &Matrix<R>
618        where
619            R: SemiRing,
620            $($($clauses)*)?
621        {
622            type Output = Matrix<R>;
623            fn $method(self, rhs: &Matrix<R>) -> Self::Output {
624                let mut this = self.clone();
625                this.$method_assign(rhs);
626                this
627            }
628        }
629    };
630}
631
632impl_matrix_pairwise_binop!(Add, add, AddAssign, add_assign);
633impl_matrix_pairwise_binop!(Sub, sub, SubAssign, sub_assign where [R: SemiRing<Additive: Invertible>]);
634
635impl<R> Mul for Matrix<R>
636where
637    R: SemiRing,
638{
639    type Output = Matrix<R>;
640    fn mul(self, rhs: Self) -> Self::Output {
641        (&self).mul(&rhs)
642    }
643}
644impl<R> Mul<&Matrix<R>> for Matrix<R>
645where
646    R: SemiRing,
647{
648    type Output = Matrix<R>;
649    fn mul(self, rhs: &Matrix<R>) -> Self::Output {
650        (&self).mul(rhs)
651    }
652}
653impl<R> Mul<Matrix<R>> for &Matrix<R>
654where
655    R: SemiRing,
656{
657    type Output = Matrix<R>;
658    fn mul(self, rhs: Matrix<R>) -> Self::Output {
659        self.mul(&rhs)
660    }
661}
662impl<R> Mul<&Matrix<R>> for &Matrix<R>
663where
664    R: SemiRing,
665{
666    type Output = Matrix<R>;
667    fn mul(self, rhs: &Matrix<R>) -> Self::Output {
668        assert_eq!(self.shape.1, rhs.shape.0);
669        if let Some(data) = R::try_matrix_product(&self.data, &rhs.data) {
670            return Matrix::from_vec(data);
671        }
672        let rhs = rhs.transpose();
673        Matrix::new_with((self.shape.0, rhs.shape.0), |i, j| {
674            R::dot_product(&self[i], &rhs[j])
675        })
676    }
677}
678
679fn strassen_rec<R: Ring>(
680    a: &[R::T],
681    b: &[R::T],
682    c: &mut [R::T],
683    shape: (usize, usize, usize),
684    stride_a: usize,
685    stride_b: usize,
686) {
687    let (n, m, p) = shape;
688    fn add_block<R: Ring>(
689        a: &[R::T],
690        b: &[R::T],
691        out: &mut [R::T],
692        n: usize,
693        stride_a: usize,
694        stride_b: usize,
695    ) {
696        for ((a, b), c) in a
697            .chunks(stride_a)
698            .zip(b.chunks(stride_b))
699            .zip(out.chunks_exact_mut(n))
700        {
701            for ((a, b), c) in a.iter().zip(b.iter()).zip(c.iter_mut()) {
702                *c = R::add(a, b);
703            }
704        }
705    }
706
707    fn sub_block<R: Ring>(
708        a: &[R::T],
709        b: &[R::T],
710        out: &mut [R::T],
711        n: usize,
712        stride_a: usize,
713        stride_b: usize,
714    ) {
715        for ((a, b), c) in a
716            .chunks(stride_a)
717            .zip(b.chunks(stride_b))
718            .zip(out.chunks_exact_mut(n))
719        {
720            for ((a, b), c) in a.iter().zip(b.iter()).zip(c.iter_mut()) {
721                *c = R::sub(a, b);
722            }
723        }
724    }
725
726    if n.min(m).min(p) <= 128 {
727        let transposed: Vec<_> = (0..p)
728            .flat_map(|j| (0..m).map(move |i| b[i * stride_b + j].clone()))
729            .collect();
730        for (a, c) in a.chunks(stride_a).zip(c.chunks_exact_mut(p)) {
731            for (b, c) in transposed.chunks_exact(m).zip(c) {
732                *c = R::dot_product(&a[..m], b);
733            }
734        }
735        return;
736    }
737    let (h, k, w) = (n / 2, m / 2, p / 2);
738    let a11 = 0;
739    let a12 = k;
740    let a21 = h * stride_a;
741    let a22 = a21 + k;
742    let b11 = 0;
743    let b12 = w;
744    let b21 = k * stride_b;
745    let b22 = b21 + w;
746
747    let block = h * w;
748    let mut buf = vec![R::zero(); h * k + k * w + block * 7];
749    let (s1, rest) = buf.split_at_mut(h * k);
750    let (s2, m_buf) = rest.split_at_mut(k * w);
751    let (m1, rest) = m_buf.split_at_mut(block);
752    let (m2, rest) = rest.split_at_mut(block);
753    let (m3, rest) = rest.split_at_mut(block);
754    let (m4, rest) = rest.split_at_mut(block);
755    let (m5, rest) = rest.split_at_mut(block);
756    let (m6, m7) = rest.split_at_mut(block);
757
758    // (A11 + A22)(B11 + B22)
759    add_block::<R>(&a[a11..], &a[a22..], s1, k, stride_a, stride_a);
760    add_block::<R>(&b[b11..], &b[b22..], s2, w, stride_b, stride_b);
761    strassen_rec::<R>(s1, s2, m1, (h, k, w), k, w);
762
763    // (A21 + A22) B11
764    add_block::<R>(&a[a21..], &a[a22..], s1, k, stride_a, stride_a);
765    strassen_rec::<R>(s1, &b[b11..], m2, (h, k, w), k, stride_b);
766
767    // A11 (B12 - B22)
768    sub_block::<R>(&b[b12..], &b[b22..], s2, w, stride_b, stride_b);
769    strassen_rec::<R>(&a[a11..], s2, m3, (h, k, w), stride_a, w);
770
771    // A22 (B21 - B11)
772    sub_block::<R>(&b[b21..], &b[b11..], s2, w, stride_b, stride_b);
773    strassen_rec::<R>(&a[a22..], s2, m4, (h, k, w), stride_a, w);
774
775    // (A11 + A12) B22
776    add_block::<R>(&a[a11..], &a[a12..], s1, k, stride_a, stride_a);
777    strassen_rec::<R>(s1, &b[b22..], m5, (h, k, w), k, stride_b);
778
779    // (A21 - A11)(B11 + B12)
780    sub_block::<R>(&a[a21..], &a[a11..], s1, k, stride_a, stride_a);
781    add_block::<R>(&b[b11..], &b[b12..], s2, w, stride_b, stride_b);
782    strassen_rec::<R>(s1, s2, m6, (h, k, w), k, w);
783
784    // (A12 - A22)(B21 + B22)
785    sub_block::<R>(&a[a12..], &a[a22..], s1, k, stride_a, stride_a);
786    add_block::<R>(&b[b21..], &b[b22..], s2, w, stride_b, stride_b);
787    strassen_rec::<R>(s1, s2, m7, (h, k, w), k, w);
788
789    let c11 = 0;
790    let c12 = w;
791    let c21 = h * p;
792    let c22 = c21 + w;
793    for ((((m1, m4), m5), m7), c) in m1
794        .iter()
795        .zip(m4.iter())
796        .zip(m5.iter())
797        .zip(m7.iter())
798        .zip(c[c11..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
799    {
800        *c = R::add(m1, m4);
801        R::sub_assign(c, m5);
802        R::add_assign(c, m7);
803    }
804    for ((m3, m5), c) in m3
805        .iter()
806        .zip(m5.iter())
807        .zip(c[c12..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
808    {
809        *c = R::add(m3, m5);
810    }
811    for ((m2, m4), c) in m2
812        .iter()
813        .zip(m4.iter())
814        .zip(c[c21..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
815    {
816        *c = R::add(m2, m4);
817    }
818    for ((((m1, m2), m3), m6), c) in m1
819        .iter()
820        .zip(m2.iter())
821        .zip(m3.iter())
822        .zip(m6.iter())
823        .zip(c[c22..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
824    {
825        *c = R::sub(m1, m2);
826        R::add_assign(c, m3);
827        R::add_assign(c, m6);
828    }
829}
830
831impl<R> Matrix<R>
832where
833    R: Ring,
834{
835    pub fn mul_strassen(&self, rhs: &Matrix<R>) -> Matrix<R> {
836        assert_eq!(self.shape.1, rhs.shape.0);
837        if let Some(data) = R::try_matrix_product(&self.data, &rhs.data) {
838            return Matrix::from_vec(data);
839        }
840        let (n, m) = self.shape;
841        let p = rhs.shape.1;
842        if n == 0 || m == 0 || p == 0 {
843            return Matrix::zeros((n, p));
844        }
845        let split = n.min(m).min(p).div_ceil(128).next_power_of_two();
846        if split <= 2 {
847            return self * rhs;
848        }
849        let rows = n.div_ceil(split) * split;
850        let inner = m.div_ceil(split) * split;
851        let cols = p.div_ceil(split) * split;
852        let mut a = vec![R::zero(); rows * inner];
853        for (a, data) in a.chunks_exact_mut(inner).zip(&self.data) {
854            a[..m].clone_from_slice(data);
855        }
856        let mut b = vec![R::zero(); inner * cols];
857        for (b, data) in b.chunks_exact_mut(cols).zip(&rhs.data) {
858            b[..p].clone_from_slice(data);
859        }
860        let mut c = vec![R::zero(); rows * cols];
861        strassen_rec::<R>(&a, &b, &mut c, (rows, inner, cols), inner, cols);
862        let mut res = Matrix::zeros((n, p));
863        for (data, c) in res.data.iter_mut().zip(c.chunks_exact(cols)) {
864            data.clone_from_slice(&c[..p]);
865        }
866        res
867    }
868}
869
870impl<R> MulAssign<&R::T> for Matrix<R>
871where
872    R: SemiRing,
873{
874    fn mul_assign(&mut self, rhs: &R::T) {
875        for i in 0..self.shape.0 {
876            for j in 0..self.shape.1 {
877                R::mul_assign(&mut self[(i, j)], rhs);
878            }
879        }
880    }
881}
882
883impl<R> Neg for Matrix<R>
884where
885    R: SemiRing<Additive: Invertible>,
886{
887    type Output = Self;
888
889    fn neg(self) -> Self::Output {
890        self.map(|x| R::neg(x))
891    }
892}
893
894impl<R> Neg for &Matrix<R>
895where
896    R: SemiRing<Additive: Invertible>,
897{
898    type Output = Matrix<R>;
899
900    fn neg(self) -> Self::Output {
901        self.map(|x| R::neg(x))
902    }
903}
904
905impl<R> Matrix<R>
906where
907    R: SemiRing,
908{
909    pub fn pow(self, mut n: usize) -> Self {
910        assert_eq!(self.shape.0, self.shape.1);
911        let mut res = Matrix::eye(self.shape);
912        let mut x = self;
913        while n > 0 {
914            if n & 1 == 1 {
915                res = &res * &x;
916            }
917            x = &x * &x;
918            n >>= 1;
919        }
920        res
921    }
922}
923
924impl<R> Matrix<R>
925where
926    R: Ring,
927{
928    pub fn pow_strassen(self, mut n: usize) -> Self {
929        assert_eq!(self.shape.0, self.shape.1);
930        let mut res = Matrix::eye(self.shape);
931        let mut x = self;
932        while n > 0 {
933            if n & 1 == 1 {
934                res = res.mul_strassen(&x);
935            }
936            x = x.mul_strassen(&x);
937            n >>= 1;
938        }
939        res
940    }
941}
942
943#[cfg(test)]
944mod tests {
945    use super::*;
946    use crate::{
947        algebra::AddMulOperation,
948        num::{One, Zero, mint_basic::DynMIntU32},
949        rand, rand_value,
950        tools::Xorshift,
951    };
952
953    type R = AddMulOperation<DynMIntU32>;
954
955    fn random_matrix(rng: &mut Xorshift, shape: (usize, usize)) -> Matrix<R> {
956        if rng.gen_bool(0.5) {
957            Matrix::new_with(shape, |_, _| rng.random(..))
958        } else if rng.gen_bool(0.5) {
959            let r = rng.randf();
960            Matrix::new_with(shape, |_, _| {
961                if rng.gen_bool(r) {
962                    rng.random(..)
963                } else {
964                    DynMIntU32::zero()
965                }
966            })
967        } else if rng.gen_bool(0.5) {
968            let mut mat = Matrix::new_with(shape, |_, _| rng.random(..));
969            let i0 = rng.random(0..shape.0);
970            let i1 = rng.random(0..shape.0);
971            let x: DynMIntU32 = rng.random(..);
972            for j in 0..shape.1 {
973                mat[(i0, j)] = mat[(i1, j)] * x;
974            }
975            mat
976        } else {
977            let mut rows: Vec<_> = (0..shape.0).collect();
978            let mut cols: Vec<_> = (0..shape.1).collect();
979            rng.shuffle(&mut rows);
980            rng.shuffle(&mut cols);
981            let mut mat = Matrix::zeros(shape);
982            for (&i, &j) in rows.iter().zip(&cols) {
983                mat[i][j] = DynMIntU32::one();
984            }
985            mat
986        }
987    }
988
989    #[test]
990    fn test_eye() {
991        for (n, m) in (0..=32).flat_map(|n| (0..=32).map(move |m| (n, m))) {
992            let result = Matrix::<R>::eye((n, m));
993            let expected = Matrix::<R>::new_with((n, m), |i, j| DynMIntU32::from((i == j) as u32));
994            assert_eq!(result, expected);
995        }
996    }
997
998    #[test]
999    fn test_add() {
1000        let mut rng = Xorshift::default();
1001        for _ in 0..100 {
1002            rand!(rng, n: 1..30, m: 1..30);
1003            let a = Matrix::<R>::new_with((n, m), |_, _| rng.random(..));
1004            let b = Matrix::<R>::new_with((n, m), |_, _| rng.random(..));
1005            assert_eq!(&a + &b, a.clone() + b.clone());
1006            assert_eq!(a.clone() + &b, a.clone() + b.clone());
1007            assert_eq!(&a + b.clone(), a.clone() + b.clone());
1008        }
1009    }
1010
1011    #[test]
1012    fn test_sub() {
1013        let mut rng = Xorshift::default();
1014        for _ in 0..100 {
1015            rand!(rng, n: 1..30, m: 1..30);
1016            let a = Matrix::<R>::new_with((n, m), |_, _| rng.random(..));
1017            let b = Matrix::<R>::new_with((n, m), |_, _| rng.random(..));
1018            assert_eq!(&a - &b, a.clone() - b.clone());
1019            assert_eq!(a.clone() - &b, a.clone() - b.clone());
1020            assert_eq!(&a - b.clone(), a.clone() - b.clone());
1021        }
1022    }
1023
1024    #[test]
1025    fn test_mul() {
1026        let mut rng = Xorshift::default();
1027        for _ in 0..100 {
1028            rand!(rng, n: 1..30, m: 1..30, l: 1..30);
1029            let a = Matrix::<R>::new_with((n, m), |_, _| rng.random(..));
1030            let b = Matrix::<R>::new_with((m, l), |_, _| rng.random(..));
1031            assert_eq!(&a * &b, a.clone() * b.clone());
1032            assert_eq!(a.clone() * &b, a.clone() * b.clone());
1033            assert_eq!(&a * b.clone(), a.clone() * b.clone());
1034            assert_eq!(
1035                &a * &b,
1036                Matrix::new_with((n, l), |i, j| (0..m).map(|k| a[i][k] * b[k][j]).sum())
1037            );
1038            let c = rng.random(..);
1039            let mut ac = a.clone();
1040            ac *= &c;
1041            assert_eq!(ac, Matrix::new_with(a.shape, |i, j| a[i][j] * c));
1042        }
1043        for _ in 0..12 {
1044            rand!(rng, n: 257..520, m: 257..520, l: 257..520);
1045            let a = Matrix::<R>::new_with((n, m), |_, _| rng.random(..));
1046            let b = Matrix::<R>::new_with((m, l), |_, _| rng.random(..));
1047            let bt = b.transpose();
1048            let expected = Matrix::new_with((n, l), |i, j| R::dot_product(&a[i], &bt[j]));
1049            assert_eq!(&a * &b, expected);
1050            assert_eq!(a.mul_strassen(&b), expected);
1051        }
1052    }
1053
1054    #[test]
1055    fn test_row_reduction() {
1056        const Q: usize = 1000;
1057        let mut rng = Xorshift::default();
1058        let ps = [2, 3, 1_000_000_007];
1059        for iteration in 0..Q {
1060            let m = ps[rng.random(..ps.len())];
1061            DynMIntU32::set_mod(m);
1062            let n = if iteration < 12 {
1063                rng.random(128..260)
1064            } else {
1065                rng.random(2..=30)
1066            };
1067            let mat = Matrix::<R>::new_with((n, n), |_, _| rng.random(..));
1068            let rank = mat.clone().rank();
1069            let inv = mat.inverse();
1070            assert_eq!(rank == n, inv.is_some());
1071            if let Some(inv) = inv {
1072                assert_eq!(&mat * &inv, Matrix::eye((n, n)));
1073            }
1074        }
1075        for _ in 0..100 {
1076            let m = ps[rng.random(..ps.len())];
1077            DynMIntU32::set_mod(m);
1078            let shape = (rng.random(1..=30), rng.random(1..=30));
1079            let mat = random_matrix(&mut rng, shape);
1080            let mut reduced = mat.clone();
1081            reduced.row_reduction(false);
1082            let expected = reduced
1083                .data
1084                .iter()
1085                .filter(|row| row.iter().any(|x| !R::is_zero(x)))
1086                .count();
1087            assert_eq!(mat.clone().rank(), expected);
1088        }
1089    }
1090
1091    #[test]
1092    fn test_determinant() {
1093        let mut rng = Xorshift::new_with_seed(358224);
1094        let ps = [2, 3, 1_000_000_007];
1095        for iteration in 0..300 {
1096            DynMIntU32::set_mod(ps[rng.random(..ps.len())]);
1097            let n = if iteration < 24 {
1098                rng.random(128..260)
1099            } else {
1100                rng.random(0..32)
1101            };
1102            let mut mat = Matrix::<R>::new_with((n, n), |i, j| {
1103                if i <= j {
1104                    rng.random(..)
1105                } else {
1106                    DynMIntU32::zero()
1107                }
1108            });
1109            if n != 0 && rng.gen_bool(0.5) {
1110                let col = rng.random(..n);
1111                for row in &mut mat.data {
1112                    row[col] = DynMIntU32::zero();
1113                }
1114            }
1115            let mut expected: DynMIntU32 = (0..n).map(|i| mat[i][i]).product();
1116            for _ in 0..3 * n {
1117                let i = rng.random(..n);
1118                let j = rng.random(..n);
1119                if i == j {
1120                    continue;
1121                }
1122                if rng.gen_bool(0.5) {
1123                    mat.data.swap(i, j);
1124                    expected = -expected;
1125                } else {
1126                    let factor: DynMIntU32 = rng.random(..);
1127                    for k in 0..n {
1128                        let x = mat[j][k] * factor;
1129                        mat[i][k] += x;
1130                    }
1131                }
1132            }
1133            let mut reduced = mat.clone();
1134            reduced.row_reduction(false);
1135            let rank = reduced
1136                .data
1137                .iter()
1138                .filter(|row| row.iter().any(|x| !R::is_zero(x)))
1139                .count();
1140            assert_eq!(mat.determinant(), expected);
1141            assert_eq!(mat.rank(), rank);
1142        }
1143    }
1144
1145    #[test]
1146    fn test_system_of_linear_equations() {
1147        let mut rng = Xorshift::new_with_seed(746182);
1148        let ps = [2, 3, 1_000_000_007];
1149        for iteration in 0..300 {
1150            DynMIntU32::set_mod(ps[rng.random(..ps.len())]);
1151            let (n, m): (usize, usize) = if iteration < 24 {
1152                (rng.random(96..200), rng.random(96..200))
1153            } else {
1154                (rng.random(0..32), rng.random(0..32))
1155            };
1156            let r = rng.random(0..=n.min(m));
1157            let mut a = &Matrix::<R>::new_with((n, r), |_, _| rng.random(..))
1158                * &Matrix::new_with((r, m), |_, _| rng.random(..));
1159            for j in 0..m {
1160                if rng.gen_bool(0.1) {
1161                    for row in &mut a.data {
1162                        row[j] = DynMIntU32::zero();
1163                    }
1164                }
1165            }
1166            let b: Vec<DynMIntU32> = if rng.gen_bool(0.5) {
1167                let x: Vec<DynMIntU32> = rand_value!(rng, [..; m]);
1168                a.data.iter().map(|row| R::dot_product(row, &x)).collect()
1169            } else {
1170                rand_value!(rng, [..; n])
1171            };
1172            let mut reduced = a.clone();
1173            reduced.add_col_with(|i, _| b[i]);
1174            reduced.row_reduction(true);
1175            let rank = reduced
1176                .data
1177                .iter()
1178                .filter(|row| row[..m].iter().any(|x| !x.is_zero()))
1179                .count();
1180            let solvable = reduced
1181                .data
1182                .iter()
1183                .all(|row| row[m].is_zero() || row[..m].iter().any(|x| !x.is_zero()));
1184            let solution = a.solve_system_of_linear_equations(&b);
1185            assert_eq!(solution.is_some(), solvable);
1186            if let Some(sol) = solution {
1187                assert_eq!(sol.basis.len(), m - rank);
1188                for (row, expected) in a.data.iter().zip(&b) {
1189                    assert_eq!(R::dot_product(row, &sol.particular), *expected);
1190                    for vector in &sol.basis {
1191                        assert!(R::dot_product(row, vector).is_zero());
1192                    }
1193                }
1194                let mut basis = Matrix::<R>::from_vec(sol.basis);
1195                basis.row_reduction(true);
1196                assert_eq!(
1197                    basis
1198                        .data
1199                        .iter()
1200                        .filter(|row| row.iter().any(|x| !x.is_zero()))
1201                        .count(),
1202                    m - rank
1203                );
1204            }
1205        }
1206        DynMIntU32::set_mod(1_000_000_007);
1207    }
1208}