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}