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 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 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 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 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 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 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 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 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 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}