1use log::debug;
9use yui_core::abst::{Ring, RingOps, Field, FieldOps};
10
11use crate::{MatTrait, Perm};
12use crate::dense::Mat;
13use crate::dense::pluq::pluq as dense_pluq;
14use super::SpMat;
15use super::SpVec;
16use super::pivot::{PivotFinderConfig, PivotType, find_pivots};
17use super::schur::Schur;
18use super::triang::{TriangularType, solve_triangular_vec};
19
20pub struct SpPluq<R> {
25 pub p: Perm,
26 pub q: Perm,
27 pub l: SpMat<R>,
28 pub u: SpMat<R>,
29 pub s: SpMat<R>,
30}
31
32impl<R> SpPluq<R> {
33 pub fn new(p: Perm, q: Perm, l: SpMat<R>, u: SpMat<R>, s: SpMat<R>) -> Self {
37 let r = l.n_cols();
38 let m = l.n_rows();
39 let n = u.n_cols();
40 assert_eq!(r, u.n_rows(), "l.n_cols() must match u.n_rows()");
41 assert_eq!(m, p.len(), "l.n_rows() must match p.len()");
42 assert_eq!(n, q.len(), "u.n_cols() must match q.len()");
43 assert_eq!(s.shape(), (m - r, n - r), "s shape must be (m - r, n - r)");
44 Self { p, q, l, u, s }
45 }
46
47 pub fn rank(&self) -> usize { self.l.n_cols() }
48
49 pub fn take_l(&mut self) -> SpMat<R> {
50 std::mem::take(&mut self.l)
51 }
52
53 pub fn take_u(&mut self) -> SpMat<R> {
54 std::mem::take(&mut self.u)
55 }
56
57 pub fn take_s(&mut self) -> SpMat<R> {
58 std::mem::take(&mut self.s)
59 }
60}
61
62impl<R> From<SpMat<R>> for SpPluq<R>
64where R: Ring, for<'x> &'x R: RingOps<R> {
65 fn from(a: SpMat<R>) -> Self {
66 let (m, n) = a.shape();
67 Self::new(
68 Perm::id(m),
69 Perm::id(n),
70 SpMat::zero((m, 0)),
71 SpMat::zero((0, n)),
72 a,
73 )
74 }
75}
76
77pub fn pre_pluq<R>(a: &SpMat<R>, config: PivotFinderConfig) -> SpPluq<R>
88where R: Ring, for<'x> &'x R: RingOps<R> {
89 debug!("compute sparse pre-pluq: {:?}", a.shape());
90
91 let (m, n) = a.shape();
92 let piv_type = config.piv_type;
93 let (p, q, r) = find_pivots(a, config);
94
95 if r == 0 {
96 return SpPluq::new(Perm::id(m), Perm::id(n), SpMat::zero((m, 0)), SpMat::zero((0, n)), a.clone());
97 }
98
99 let [a0, a1, a2, a3] = a.permute_and_split(&p, &q, r);
100
101 let (l, u, s) = match piv_type {
102 PivotType::Rows => {
103 let sch = Schur::from_blocks(TriangularType::Upper, [&a0, &a1, &a2, &a3], false, true);
104 let (s, _, row_mult) = sch.disassemble();
105 let l1 = row_mult.unwrap();
106 let u = SpMat::h_stack(a0, a1); let l = SpMat::v_stack(SpMat::id(r), l1); (l, u, s)
109 },
110 PivotType::Cols => {
111 let sch = Schur::from_blocks(TriangularType::Lower, [&a0, &a1, &a2, &a3], true, false);
112 let (s, col_mult, _) = sch.disassemble();
113 let u1 = col_mult.unwrap();
114 let l = SpMat::v_stack(a0, a2); let u = SpMat::h_stack(SpMat::id(r), u1); (l, u, s)
117 }
118 };
119
120 SpPluq::new(p, q, l, u, s)
121}
122
123pub fn pluq<R>(a: &SpMat<R>, config: PivotFinderConfig) -> SpPluq<R>
129where R: Ring, for<'x> &'x R: RingOps<R> {
130 debug!("compute sparse pluq: {:?}", a.shape());
131
132 let piv_type = config.piv_type;
133 let mut pp = SpPluq::from(a.clone());
134
135 while !pp.s.is_zero() {
136 let pp_next = pre_pluq(&pp.s, config);
137 if pp_next.rank() == 0 { break; }
138
139 merge_pluq(&mut pp, pp_next);
140 }
141
142 if pp.s.is_zero() { return pp; }
143
144 let pp_dense = dense_pluq_in(&pp.s, piv_type);
145 merge_pluq(&mut pp, pp_dense);
146 pp
147}
148
149fn dense_pluq_in<R>(s: &SpMat<R>, piv_type: PivotType) -> SpPluq<R>
150where R: Ring, for<'x> &'x R: RingOps<R> {
151 let transpose = piv_type == PivotType::Rows;
152 let (ms, ns) = s.shape();
153 let (row_idx, col_idx, mat) = extract_dense(s, transpose);
154 let (m0, n0) = (row_idx.len(), col_idx.len());
155
156 let raw = dense_pluq(&mat);
157 let dp = if transpose { raw.transpose() } else { raw };
158
159 let p2 = extend_perm(ms, &row_idx, dp.p);
160 let q2 = extend_perm(ns, &col_idx, dp.q);
161
162 let mut l2 = SpMat::from(dp.l);
163 l2.extend_by_zero(ms - m0, 0);
164
165 let mut u2 = SpMat::from(dp.u);
166 u2.extend_by_zero(0, ns - n0);
167
168 let mut s2 = SpMat::from(dp.s);
169 s2.extend_by_zero(ms - m0, ns - n0);
170
171 SpPluq::new(p2, q2, l2, u2, s2)
172}
173
174fn extract_dense<R>(s: &SpMat<R>, transpose: bool) -> (Vec<usize>, Vec<usize>, Mat<R>)
178where R: Ring, for<'x> &'x R: RingOps<R> {
179 use std::collections::BTreeSet;
180
181 let row_idx: Vec<usize> = s.iter_nz().map(|(i, _, _)| i).collect::<BTreeSet<_>>().into_iter().collect();
182 let col_idx: Vec<usize> = s.iter_nz().map(|(_, j, _)| j).collect::<BTreeSet<_>>().into_iter().collect();
183 let (m0, n0) = (row_idx.len(), col_idx.len());
184
185 let row_perm = Perm::forward_indices(s.n_rows(), row_idx.iter().copied());
186 let col_perm = Perm::forward_indices(s.n_cols(), col_idx.iter().copied());
187
188 let shape = if transpose { (n0, m0) } else { (m0, n0) };
189 let mut mat = Mat::zero(shape);
190
191 for (i, j, v) in s.iter_nz() {
192 let (ri, cj) = (row_perm.at(i), col_perm.at(j));
193 if transpose {
194 mat[(cj, ri)] = v.clone(); } else {
196 mat[(ri, cj)] = v.clone();
197 }
198 }
199
200 (row_idx, col_idx, mat)
201}
202
203fn merge_pluq<R>(pp1: &mut SpPluq<R>, pp2: SpPluq<R>)
208where R: Ring, for<'x> &'x R: RingOps<R> {
209 debug!("merge pluq: {} + {}", pp1.rank(), pp2.rank());
210
211 let (m, n) = (pp1.l.n_rows(), pp1.u.n_cols());
212 let r1 = pp1.rank();
213 let r2 = pp2.rank();
214
215 assert_eq!(pp2.l.n_rows(), m - r1);
216 assert_eq!(pp2.u.n_cols(), n - r1);
217
218 pp1.l = {
222 let [l0, l1] = pp1.take_l().v_split(r1);
223 let l1 = l1.permute_rows(&pp2.p);
224 let zero_tr = SpMat::zero((r1, r2));
225 SpMat::block_combine([l0, zero_tr, l1, pp2.l])
226 };
227
228 pp1.u = {
229 let [u0, u1] = pp1.take_u().h_split(r1);
230 let u1 = u1.permute_cols(&pp2.q);
231 let zero_bl = SpMat::zero((r2, r1));
232 SpMat::block_combine([u0, u1, zero_bl, pp2.u])
233 };
234
235 pp1.s = pp2.s;
236 pp1.p = merge_perm(&pp1.p, pp2.p);
237 pp1.q = merge_perm(&pp1.q, pp2.q);
238}
239
240pub fn solve_pluq<R>(a: &SpMat<R>, y: &SpVec<R>) -> Option<SpVec<R>>
244where R: Field, for<'x> &'x R: FieldOps<R> {
245 debug!("solve pluq, a: {:?}", a.shape());
246
247 assert_eq!(y.dim(), a.n_rows());
248
249 if y.is_zero() {
250 return Some(SpVec::zero(a.n_cols())); }
252
253 let pp = pluq(a, PivotFinderConfig {
254 piv_type: PivotType::Rows,
255 ..Default::default()
256 });
257
258 let y_dense = y.clone().into_dense();
259 let yp = pp.p.apply_to(y_dense);
260 let xq = solve_lu(&pp.l, &pp.u, &yp)?;
261 let x = pp.q.apply_inv_to(xq);
262
263 Some(SpVec::from(x))
264}
265
266fn solve_lu<R>(l: &SpMat<R>, u: &SpMat<R>, y: &[R]) -> Option<Vec<R>>
276where R: Field, for<'x> &'x R: FieldOps<R> {
277 assert_eq!(l.n_cols(), u.n_rows());
278 assert_eq!(y.len(), l.n_rows());
279
280 let z = solve_l(l, y, true)?;
281 let x = solve_u(u, &z);
282
283 Some(x)
284}
285
286fn solve_l<R>(l: &SpMat<R>, y: &[R], check_consistency: bool) -> Option<Vec<R>>
294where R: Field, for<'x> &'x R: FieldOps<R> {
295 assert_eq!(l.n_rows(), y.len());
296 let r = l.n_cols();
297
298 let x = if r == y.len() {
299 let y = SpVec::from(y.to_vec());
300 solve_triangular_vec(TriangularType::Lower, l, &y).into_dense()
301 } else {
302 let l0 = l.submat(0..r, 0..r);
303 let y0 = SpVec::from(y[..r].to_vec());
304 let x = solve_triangular_vec(TriangularType::Lower, &l0, &y0).into_dense();
305
306 if check_consistency && !is_consistent(l, y, &x) {
307 return None;
308 }
309 x
310 };
311
312 Some(x)
313}
314
315fn is_consistent<R>(l: &SpMat<R>, y: &[R], x: &[R]) -> bool
317where R: Ring, for<'x> &'x R: RingOps<R> {
318 is_consistent_upto(l, y, x, y.len())
319}
320
321fn is_consistent_upto<R>(l: &SpMat<R>, y: &[R], x: &[R], k: usize) -> bool
322where R: Ring, for<'x> &'x R: RingOps<R> {
323 assert_eq!(l.n_rows(), y.len());
324 assert_eq!(l.n_cols(), x.len());
325 assert!(x.len() <= k && k <= y.len());
326
327 let r = x.len();
328 let mut res = y[r..k].to_vec();
329
330 for (i, j, v) in l.iter_nz() {
331 if r <= i && i < k {
332 res[i - r] -= v * &x[j];
333 }
334 }
335
336 res.iter().all(|v| v.is_zero())
337}
338
339fn solve_u<R>(u: &SpMat<R>, y: &[R]) -> Vec<R>
344where R: Field, for<'x> &'x R: FieldOps<R> {
345 let (r, n) = u.shape();
346 assert_eq!(y.len(), r);
347 assert!(n >= r);
348
349 let mut x = if n == r {
350 solve_triangular_vec(TriangularType::Upper, u, &SpVec::from(y.to_vec())).into_dense()
351 } else {
352 let u0 = u.submat(0..r, 0..r);
353 solve_triangular_vec(TriangularType::Upper, &u0, &SpVec::from(y.to_vec())).into_dense()
354 };
355
356 x.resize(n, R::zero());
357 x
358}
359
360pub fn solve_pluq_incr<R>(a: &SpMat<R>, y: &SpVec<R>, max_piv: usize, chunk: usize) -> Option<SpVec<R>>
367where R: Field, for<'x> &'x R: FieldOps<R> {
368 debug!("solve pluq (incremental), a: {:?}", a.shape());
369
370 assert_eq!(y.dim(), a.n_rows());
371
372 if y.is_zero() {
373 return Some(SpVec::zero(a.n_cols())); }
375
376 let mut pp = pre_pluq(a, PivotFinderConfig {
377 piv_type: PivotType::Rows,
378 max_pivots: max_piv,
379 ..Default::default()
380 });
381 let y_dense = y.clone().into_dense();
382 let mut yp = pp.p.apply_to(y_dense);
383
384 let mut step = 1;
385 let total_step = (a.n_rows() - pp.rank()) / chunk + 1;
386
387 while pp.s.n_rows() > 0 {
388 debug!("(step {}/{})", step, total_step);
389 debug!(" current rank: {}", pp.rank());
390
391 let r_old = pp.rank();
392 let (pp_next, r_next, c) = chunk_pluq(pp.take_s(), chunk);
393 let p_next = pp_next.p.clone();
394
395 merge_pluq(&mut pp, pp_next);
396
397 let yp_tail = p_next.apply_to(yp[r_old..].to_vec());
399 yp[r_old..].clone_from_slice(&yp_tail);
400
401 let k = c - r_next;
405 let z = solve_l(&pp.l, &yp, false).unwrap();
406
407 if !is_consistent_upto(&pp.l, &yp, &z, z.len() + k) {
408 debug!("found inconsistency at step {}/{}.", step, total_step);
409 return None;
410 }
411
412 trim_zero_rows(&mut pp, &mut yp, k);
413 step += 1;
414 }
415
416 debug!("pluq complete, rank: {}", pp.rank());
417 debug!("solve pluq..");
418
419 let xq = solve_lu(&pp.l, &pp.u, &yp)?;
420 let x = pp.q.apply_inv_to(xq);
421
422 Some(SpVec::from(x))
423}
424
425fn chunk_pluq<R>(s: SpMat<R>, chunk_size: usize) -> (SpPluq<R>, usize, usize)
429where R: Ring, for<'x> &'x R: RingOps<R> {
430 let c = chunk_size.min(s.n_rows());
431 let [s_chunk, s_rest] = s.v_split(c);
432 let pp_chunk = pluq(&s_chunk, PivotFinderConfig {
433 piv_type: PivotType::Rows,
434 ..Default::default()
435 });
436 let r_chunk = pp_chunk.rank();
437 let pp_full = extend_chunk_to_full(pp_chunk, s_rest);
438 (pp_full, r_chunk, c)
439}
440
441fn extend_chunk_to_full<R>(pp_chunk: SpPluq<R>, s_rest: SpMat<R>) -> SpPluq<R>
448where R: Ring, for<'x> &'x R: RingOps<R> {
449 let (c, n_s) = (pp_chunk.l.n_rows(), pp_chunk.u.n_cols());
450 let r_chunk = pp_chunk.rank();
451 let m_rest = s_rest.n_rows();
452 let m_s = c + m_rest;
453
454 assert_eq!(s_rest.n_cols(), n_s);
455 assert_eq!(pp_chunk.s.shape(), (c - r_chunk, n_s - r_chunk));
456
457 let s_rest_q = s_rest.permute_cols(&pp_chunk.q);
458 let [s_rest_left, s_rest_right] = s_rest_q.h_split(r_chunk);
459 let [u_top, u_right] = pp_chunk.u.clone().h_split(r_chunk);
460
461 let sch = Schur::from_blocks(
464 TriangularType::Upper,
465 [&u_top, &u_right, &s_rest_left, &s_rest_right],
466 false, true
467 );
468 let (s_ext, _, row_mult) = sch.disassemble();
469 let l_ext = row_mult.unwrap();
470
471 let chunk_idx: Vec<usize> = (0..c).collect();
472 let p = extend_perm(m_s, &chunk_idx, pp_chunk.p);
473 let q = pp_chunk.q;
474 let l = SpMat::v_stack(pp_chunk.l, l_ext);
475 let u = pp_chunk.u;
476 let s = SpMat::v_stack(pp_chunk.s, s_ext);
477
478 SpPluq::new(p, q, l, u, s)
479}
480
481fn trim_zero_rows<R>(pp: &mut SpPluq<R>, yp: &mut Vec<R>, k: usize)
487where R: Ring, for<'x> &'x R: RingOps<R> {
488 if k == 0 { return; }
489
490 let r = pp.rank();
491 let m = pp.l.n_rows();
492 assert!(r + k <= m);
493 assert_eq!(yp.len(), m);
494 assert_eq!(pp.p.len(), m);
495
496 pp.l = pp.l.extract((m - k, r), |i, j| {
498 if i < r {
499 Some((i, j))
500 } else if i < r + k {
501 None
502 } else {
503 Some((i - k, j))
504 }
505 });
506
507 pp.s = pp.s.submat_rows(k..pp.s.n_rows());
509
510 yp.drain(r..r + k);
512
513 let new_p_at: Vec<usize> = (0..m).filter_map(|i| {
515 let pos = pp.p.at(i);
516 if pos < r {
517 Some(pos)
518 } else if pos < r + k {
519 None
520 } else {
521 Some(pos - k)
522 }
523 }).collect();
524 pp.p = Perm::new(new_p_at);
525}
526
527fn merge_perm(perm1: &Perm, perm2: Perm) -> Perm {
530 assert!(perm1.len() >= perm2.len());
531 let r = perm1.len() - perm2.len();
532 perm2.shift(r) * perm1
533}
534
535fn extend_perm(n: usize, compact_idx: &[usize], compact_perm: Perm) -> Perm {
539 let c = compact_idx.len();
540
541 assert!(n >= c);
542 assert_eq!(compact_perm.len(), c);
543
544 let front = Perm::forward_indices(n, compact_idx.iter().copied());
545 compact_perm.extend(n - c) * front
546}
547
548#[cfg(test)]
549mod tests {
550 use super::*;
551 use num_traits::One;
552
553 fn cfg(piv_type: PivotType) -> PivotFinderConfig {
554 PivotFinderConfig { piv_type, ..Default::default() }
555 }
556
557 fn sample() -> SpMat<i32> {
558 SpMat::from_row_major((6, 9), [
559 1, 0, 0, 0, 0, 1, 0, 0, 1,
560 0, 1, 1, 1, 0, 1, 0, 1, 0,
561 0, 0, 1, 1, 0, 0, 0, 1, 1,
562 0, 1, 0, 0, 1, 0, 0, 0, 0,
563 0, 0, 1, 0, 0, 0, 0, 0, 0,
564 0, 1, 0, 0, 0, 1, 0, 1, 0,
565 ])
566 }
567
568 fn check_pre_pluq_rows(a: &SpMat<i32>) {
571 let pp = pre_pluq(a, cfg(PivotType::Rows));
572 let (m, n) = a.shape();
573 let r = pp.rank();
574
575 assert_eq!(pp.l.shape(), (m, r));
576 assert_eq!(pp.u.shape(), (r, n));
577 assert_eq!(pp.s.shape(), (m - r, n - r));
578
579 let paq = a.permute(&pp.p, &pp.q);
580 let rem_full = SpMat::from_entries((m, n),
581 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
582 );
583 assert_eq!(paq, &pp.l * &pp.u + &rem_full);
584
585 let b = pp.u.clone().into_dense();
586 for k in 0..r {
587 assert!(b[(k, k)].is_one(), "u[{k},{k}] should be a pivot (=1)");
588 }
589 for j in 0..r {
590 for i in j + 1..r {
591 assert_eq!(b[(i, j)], 0, "u[{i},{j}] should be zero (below diagonal)");
592 }
593 }
594 }
595
596 fn check_pre_pluq_cols(a: &SpMat<i32>) {
597 let pp = pre_pluq(a, cfg(PivotType::Cols));
598 let (m, n) = a.shape();
599 let r = pp.rank();
600
601 assert_eq!(pp.l.shape(), (m, r));
602 assert_eq!(pp.u.shape(), (r, n));
603 assert_eq!(pp.s.shape(), (m - r, n - r));
604
605 let paq = a.permute(&pp.p, &pp.q);
606 let rem_full = SpMat::from_entries((m, n),
607 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
608 );
609 assert_eq!(paq, &pp.l * &pp.u + &rem_full);
610
611 let b = pp.l.clone().into_dense();
612 for k in 0..r {
613 assert!(b[(k, k)].is_one(), "l[{k},{k}] should be a pivot (=1)");
614 }
615 for i in 0..r {
616 for j in i + 1..r {
617 assert_eq!(b[(i, j)], 0, "l[{i},{j}] should be zero (above diagonal)");
618 }
619 }
620 }
621
622 #[test]
623 fn test_pre_pluq_rows() { check_pre_pluq_rows(&sample()); }
624
625 #[test]
626 fn test_pre_pluq_cols() { check_pre_pluq_cols(&sample()); }
627
628 #[test]
629 fn test_pre_pluq_zero() {
630 let a = SpMat::<i32>::zero((4, 5));
631 let pp = pre_pluq(&a, cfg(PivotType::Rows));
632 assert_eq!(pp.rank(), 0);
633 assert_eq!(pp.l.shape(), (4, 0));
634 assert_eq!(pp.u.shape(), (0, 5));
635 assert_eq!(pp.s.shape(), (4, 5)); assert_eq!(pp.s, a.permute(&pp.p, &pp.q));
637 }
638
639 #[test]
640 fn test_pre_pluq_square_full_rank() {
641 let a = SpMat::from_row_major((3, 3), [1, 0, 0, 0, 1, 0, 0, 0, 1]);
642 let pp = pre_pluq(&a, cfg(PivotType::Rows));
643 assert_eq!(pp.rank(), 3);
644 assert_eq!(pp.s.shape(), (0, 0)); }
646
647 #[test]
648 fn test_pre_pluq_rank_rows() {
649 assert_eq!(pre_pluq(&sample(), cfg(PivotType::Rows)).rank(), 5);
650 }
651
652 #[test]
653 fn test_pre_pluq_rank_cols() {
654 assert_eq!(pre_pluq(&sample(), cfg(PivotType::Cols)).rank(), 6);
655 }
656
657 #[test]
658 fn test_pre_pluq_rand_rows() { check_pre_pluq_rows(&SpMat::<i32>::rand((40, 60), 0.1)); }
659
660 #[test]
661 fn test_pre_pluq_rand_cols() { check_pre_pluq_cols(&SpMat::<i32>::rand((40, 60), 0.1)); }
662
663 fn check_pluq_rows(a: &SpMat<i32>) {
666 let pp = pluq(a, cfg(PivotType::Rows));
667 let (m, n) = a.shape();
668 let r = pp.rank();
669
670 assert_eq!(pp.l.shape(), (m, r));
671 assert_eq!(pp.u.shape(), (r, n));
672 assert_eq!(pp.s.shape(), (m - r, n - r));
673
674 let paq = a.permute(&pp.p, &pp.q);
675 let rem = SpMat::from_entries((m, n),
676 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
677 );
678 assert_eq!(paq, &pp.l * &pp.u + &rem, "p*A*q != l*u + rest");
679
680 let lb = pp.l.clone().into_dense();
682 for k in 0..r {
683 assert!(lb[(k, k)].is_one(), "l[{k},{k}] should be 1");
684 for j in k + 1..r { assert_eq!(lb[(k, j)], 0, "l[{k},{j}] above diag"); }
685 }
686 }
687
688 fn check_pluq_cols(a: &SpMat<i32>) {
689 let pp = pluq(a, cfg(PivotType::Cols));
690 let (m, n) = a.shape();
691 let r = pp.rank();
692
693 assert_eq!(pp.l.shape(), (m, r));
694 assert_eq!(pp.u.shape(), (r, n));
695 assert_eq!(pp.s.shape(), (m - r, n - r));
696
697 let paq = a.permute(&pp.p, &pp.q);
698 let rem = SpMat::from_entries((m, n),
699 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
700 );
701 assert_eq!(paq, &pp.l * &pp.u + &rem, "p*A*q != l*u + rest");
702
703 let ub = pp.u.clone().into_dense();
705 for k in 0..r {
706 assert!(ub[(k, k)].is_one(), "u[{k},{k}] should be 1");
707 for i in k + 1..r { assert_eq!(ub[(i, k)], 0, "u[{i},{k}] below diag"); }
708 }
709 }
710
711 #[test]
714 fn test_extract_dense_no_transpose() {
715 let s = SpMat::from_row_major((3, 3), [1i32, 0, 2, 0, 0, 0, 3, 0, 4]);
720 let (row_idx, col_idx, mat) = extract_dense(&s, false);
721 assert_eq!(row_idx, vec![0usize, 2]);
722 assert_eq!(col_idx, vec![0usize, 2]);
723 assert_eq!(mat, crate::dense::Mat::from_row_major((2, 2), [1i32, 2, 3, 4]));
724 }
725
726 #[test]
727 fn test_extract_dense_transpose() {
728 let s = SpMat::from_row_major((3, 3), [1i32, 0, 2, 0, 0, 0, 3, 0, 4]);
730 let (row_idx, col_idx, mat) = extract_dense(&s, true);
731 assert_eq!(row_idx, vec![0usize, 2]);
732 assert_eq!(col_idx, vec![0usize, 2]);
733 assert_eq!(mat, crate::dense::Mat::from_row_major((2, 2), [1i32, 3, 2, 4]));
734 }
735
736 fn check_dense_pluq_in(s: &SpMat<i32>, piv_type: PivotType) {
739 let pp = dense_pluq_in(s, piv_type);
740 let (ms, ns) = s.shape();
741 let r = pp.rank();
742
743 assert_eq!(pp.l.shape(), (ms, r));
744 assert_eq!(pp.u.shape(), (r, ns));
745 assert_eq!(pp.s.shape(), (ms - r, ns - r));
746
747 let psq = s.permute(&pp.p, &pp.q);
748 let rem = SpMat::from_entries((ms, ns),
749 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
750 );
751 assert_eq!(psq, &pp.l * &pp.u + &rem, "p*s*q != l*u + rest");
752 }
753
754 #[test]
755 fn test_dense_pluq_in_cols_with_zero_row_and_col() {
756 let s = SpMat::from_row_major((3, 3), [1i32, 0, 2, 0, 0, 0, 3, 0, 4]);
758 check_dense_pluq_in(&s, PivotType::Cols);
759 }
760
761 #[test]
762 fn test_dense_pluq_in_rows_with_zero_row_and_col() {
763 let s = SpMat::from_row_major((3, 3), [1i32, 0, 2, 0, 0, 0, 3, 0, 4]);
764 check_dense_pluq_in(&s, PivotType::Rows);
765 }
766
767 #[test]
768 fn test_dense_pluq_in_all_zero() {
769 let s = SpMat::<i32>::zero((4, 5));
770 check_dense_pluq_in(&s, PivotType::Cols);
771 check_dense_pluq_in(&s, PivotType::Rows);
772 }
773
774 #[test]
775 fn test_dense_pluq_in_no_zero_rows_or_cols() {
776 let s = SpMat::from_row_major((3, 3), [1i32,2,3,4,5,6,7,8,9]);
778 check_dense_pluq_in(&s, PivotType::Cols);
779 check_dense_pluq_in(&s, PivotType::Rows);
780 }
781
782 #[test]
783 fn test_pluq_rows() { check_pluq_rows(&sample()); }
784
785 #[test]
786 fn test_pluq_cols() { check_pluq_cols(&sample()); }
787
788 #[test]
789 fn test_pluq_zero() {
790 let a = SpMat::<i32>::zero((4, 5));
791 check_pluq_rows(&a);
792 check_pluq_cols(&a);
793 }
794
795 #[test]
796 fn test_pluq_rand_rows() { check_pluq_rows(&SpMat::<i32>::rand((40, 60), 0.1)); }
797
798 #[test]
799 fn test_pluq_rand_cols() { check_pluq_cols(&SpMat::<i32>::rand((40, 60), 0.1)); }
800
801 use yui_core::num::Ratio;
804 type R = Ratio<i64>;
805 fn r(n: i64) -> R { R::from(n) }
806
807 fn sp_mat(shape: (usize, usize), data: impl IntoIterator<Item = R>) -> SpMat<R> {
808 SpMat::from_row_major(shape, data)
809 }
810
811 fn sp_vec(data: impl IntoIterator<Item = R>) -> SpVec<R> {
812 SpVec::from(data.into_iter().collect::<Vec<_>>())
813 }
814
815 #[test]
816 fn test_solve_l_square() {
817 let l = sp_mat((2, 2), [r(2), r(0), r(3), r(4)]);
820 let y = vec![r(4), r(11)];
821 let x = solve_l(&l, &y, true);
822 assert_eq!(x, Some(vec![r(2), r(5)/r(4)]));
823 }
824
825 #[test]
826 fn test_solve_l_rectangular_consistent() {
827 let l = sp_mat((3, 2), [r(1), r(0), r(2), r(1), r(3), r(4)]);
831 let y = vec![r(1), r(2), r(3)];
832 assert_eq!(solve_l(&l, &y, true), Some(vec![r(1), r(0)]));
833 }
834
835 #[test]
836 fn test_solve_l_rectangular_inconsistent() {
837 let l = sp_mat((3, 2), [r(1), r(0), r(2), r(1), r(3), r(4)]);
839 let y = vec![r(1), r(2), r(4)];
840 assert_eq!(solve_l(&l, &y, true), None);
841 }
842
843 #[test]
844 fn test_solve_l_no_check() {
845 let l = sp_mat((3, 2), [r(1), r(0), r(2), r(1), r(3), r(4)]);
848 let y = vec![r(1), r(2), r(4)];
849 assert_eq!(solve_l(&l, &y, false), Some(vec![r(1), r(0)]));
850 }
851
852 #[test]
853 fn test_solve_l_zero_cols_consistent() {
854 let l: SpMat<R> = SpMat::zero((3, 0));
856 assert_eq!(solve_l(&l, &[r(0); 3], true), Some(vec![]));
857 }
858
859 #[test]
860 fn test_solve_l_zero_cols_inconsistent() {
861 let l: SpMat<R> = SpMat::zero((3, 0));
863 assert_eq!(solve_l(&l, &[r(1), r(0), r(0)], true), None);
864 }
865
866 #[test]
869 fn test_is_consistent_full() {
870 let l = sp_mat((3, 2), [r(1), r(0), r(2), r(1), r(3), r(4)]);
874 let x = vec![r(1), r(0)];
875 assert!( is_consistent(&l, &[r(1), r(2), r(3)], &x));
876 assert!(!is_consistent(&l, &[r(1), r(2), r(4)], &x));
877 }
878
879 #[test]
880 fn test_is_consistent_upto_partial() {
881 let l = sp_mat((4, 2), [r(1), r(0), r(2), r(1), r(3), r(4), r(5), r(6)]);
887 let y = vec![r(1), r(2), r(3), r(99)];
888 let x = vec![r(1), r(0)];
889 assert!( is_consistent_upto(&l, &y, &x, 2));
890 assert!( is_consistent_upto(&l, &y, &x, 3));
891 assert!(!is_consistent_upto(&l, &y, &x, 4));
892 }
893
894 #[test]
897 fn test_solve_u_square() {
898 let u = sp_mat((2, 2), [r(1), r(2), r(0), r(3)]);
901 let y = vec![r(4), r(6)];
902 assert_eq!(solve_u(&u, &y), vec![r(0), r(2)]);
903 }
904
905 #[test]
906 fn test_solve_u_rectangular() {
907 let u = sp_mat((2, 4), [r(1), r(2), r(5), r(6), r(0), r(3), r(7), r(8)]);
910 let y = vec![r(4), r(6)];
911 assert_eq!(solve_u(&u, &y), vec![r(0), r(2), r(0), r(0)]);
912 }
913
914 #[test]
915 fn test_solve_u_empty() {
916 let u: SpMat<R> = SpMat::zero((0, 0));
918 assert_eq!(solve_u(&u, &[]), Vec::<R>::new());
919 }
920
921 fn solve_check(a: &SpMat<R>, y: &SpVec<R>) -> SpVec<R> {
924 let x = solve_pluq(a, y).expect("expected a solution");
925 assert_eq!(&(a * &x), y, "A*x != y");
926 x
927 }
928
929 #[test]
930 fn test_solve_square() {
931 let a = sp_mat((2, 2), [r(1), r(2), r(3), r(4)]);
932 solve_check(&a, &sp_vec([r(5), r(6)]));
933 }
934
935 #[test]
936 fn test_solve_overdetermined_consistent() {
937 let a = sp_mat((3, 2), [r(1), r(0), r(0), r(1), r(1), r(1)]);
938 solve_check(&a, &sp_vec([r(2), r(3), r(5)]));
939 }
940
941 #[test]
942 fn test_solve_overdetermined_inconsistent() {
943 let a = sp_mat((3, 2), [r(1), r(0), r(0), r(1), r(1), r(1)]);
944 assert!(solve_pluq(&a, &sp_vec([r(1), r(1), r(0)])).is_none());
945 }
946
947 #[test]
948 fn test_solve_underdetermined() {
949 let a = sp_mat((2, 3), [r(1), r(0), r(2), r(0), r(1), r(3)]);
950 solve_check(&a, &sp_vec([r(4), r(5)]));
951 }
952
953 #[test]
954 fn test_solve_zero_rhs() {
955 let a = sp_mat((2, 2), [r(1), r(2), r(3), r(4)]);
956 let x = solve_check(&a, &sp_vec([r(0), r(0)]));
957 assert_eq!(x, sp_vec([r(0), r(0)]));
958 }
959
960 #[test]
961 fn test_solve_no_solution() {
962 let a = sp_mat((2, 2), [r(1), r(2), r(2), r(4)]);
963 assert!(solve_pluq(&a, &sp_vec([r(1), r(0)])).is_none());
964 }
965
966 #[test]
967 fn test_solve_identity() {
968 let a: SpMat<R> = SpMat::from_entries((4, 4), (0..4).map(|k| (k, k, r(1))));
969 let y = sp_vec([r(1), r(2), r(3), r(4)]);
970 let x = solve_check(&a, &y);
971 assert_eq!(x, y);
972 }
973
974 fn check_extend_chunk(s: &SpMat<i32>, c: usize) {
977 let (m, n) = s.shape();
978 assert!(c <= m);
979 let [s_top, s_rest] = s.clone().v_split(c);
980 let pp_chunk = pluq(&s_top, cfg(PivotType::Rows));
981 let pp = extend_chunk_to_full(pp_chunk, s_rest);
982 let r = pp.rank();
983
984 assert_eq!(pp.l.shape(), (m, r));
985 assert_eq!(pp.u.shape(), (r, n));
986 assert_eq!(pp.s.shape(), (m - r, n - r));
987
988 let psq = s.permute(&pp.p, &pp.q);
989 let rem = SpMat::from_entries((m, n),
990 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
991 );
992 assert_eq!(psq, &pp.l * &pp.u + &rem, "p*s*q != l*u + rest (c = {c})");
993 }
994
995 #[test]
996 fn test_extend_chunk_top() { check_extend_chunk(&sample(), 3); }
997
998 #[test]
999 fn test_extend_chunk_full() { check_extend_chunk(&sample(), 6); }
1000
1001 #[test]
1002 fn test_extend_chunk_empty() { check_extend_chunk(&sample(), 0); }
1003
1004 #[test]
1005 fn test_extend_chunk_rand() { check_extend_chunk(&SpMat::<i32>::rand((40, 60), 0.1), 17); }
1006
1007 #[test]
1010 fn test_chunk_pluq() {
1011 let s = sample();
1012 let (m, n) = s.shape();
1013 let (pp, r_chunk, c) = chunk_pluq(s.clone(), 3);
1014 let r = pp.rank();
1015
1016 assert_eq!(c, 3);
1017 assert_eq!(r, r_chunk);
1018 assert_eq!(pp.l.shape(), (m, r));
1019 assert_eq!(pp.u.shape(), (r, n));
1020 assert_eq!(pp.s.shape(), (m - r, n - r));
1021
1022 let psq = s.permute(&pp.p, &pp.q);
1023 let rem = SpMat::from_entries((m, n),
1024 pp.s.iter_nz().map(|(i, j, v)| (i + r, j + r, *v))
1025 );
1026 assert_eq!(psq, &pp.l * &pp.u + &rem);
1027 }
1028
1029 #[test]
1030 fn test_chunk_pluq_oversize() {
1031 let s = sample();
1032 let m = s.n_rows();
1033 let (_, _, c) = chunk_pluq(s, 100);
1034 assert_eq!(c, m);
1035 }
1036
1037 fn solve_incr_check(a: &SpMat<R>, y: &SpVec<R>, max_piv: usize, chunk: usize) -> SpVec<R> {
1040 let x = solve_pluq_incr(a, y, max_piv, chunk).expect("expected a solution");
1041 assert_eq!(&(a * &x), y, "A*x != y (max_piv={max_piv}, chunk={chunk})");
1042 x
1043 }
1044
1045 #[test]
1046 fn test_solve_incr_square() {
1047 let a = sp_mat((2, 2), [r(1), r(2), r(3), r(4)]);
1048 for (mp, ch) in [(0, 1), (0, 2), (1, 1), (usize::MAX, 1), (usize::MAX, 100)] {
1050 solve_incr_check(&a, &sp_vec([r(5), r(6)]), mp, ch);
1051 }
1052 }
1053
1054 #[test]
1055 fn test_solve_incr_overdetermined_consistent() {
1056 let a = sp_mat((3, 2), [r(1), r(0), r(0), r(1), r(1), r(1)]);
1057 solve_incr_check(&a, &sp_vec([r(2), r(3), r(5)]), 0, 2);
1058 solve_incr_check(&a, &sp_vec([r(2), r(3), r(5)]), 1, 1);
1059 }
1060
1061 #[test]
1062 fn test_solve_incr_overdetermined_inconsistent() {
1063 let a = sp_mat((3, 2), [r(1), r(0), r(0), r(1), r(1), r(1)]);
1064 for (mp, ch) in [(0, 1), (0, 3), (usize::MAX, 1)] {
1065 assert!(solve_pluq_incr(&a, &sp_vec([r(1), r(1), r(0)]), mp, ch).is_none());
1066 }
1067 }
1068
1069 #[test]
1070 fn test_solve_incr_underdetermined() {
1071 let a = sp_mat((2, 3), [r(1), r(0), r(2), r(0), r(1), r(3)]);
1072 solve_incr_check(&a, &sp_vec([r(4), r(5)]), 0, 1);
1073 solve_incr_check(&a, &sp_vec([r(4), r(5)]), 1, 1);
1074 }
1075
1076 #[test]
1077 fn test_solve_incr_zero_rhs() {
1078 let a = sp_mat((2, 2), [r(1), r(2), r(3), r(4)]);
1079 let x = solve_incr_check(&a, &sp_vec([r(0), r(0)]), 0, 1);
1080 assert_eq!(x, sp_vec([r(0), r(0)]));
1081 }
1082
1083 #[test]
1084 fn test_solve_incr_no_solution() {
1085 let a = sp_mat((2, 2), [r(1), r(2), r(2), r(4)]);
1086 for (mp, ch) in [(0, 1), (0, 2), (usize::MAX, 1)] {
1087 assert!(solve_pluq_incr(&a, &sp_vec([r(1), r(0)]), mp, ch).is_none());
1088 }
1089 }
1090
1091 #[test]
1092 fn test_solve_incr_identity() {
1093 let a: SpMat<R> = SpMat::from_entries((4, 4), (0..4).map(|k| (k, k, r(1))));
1094 let y = sp_vec([r(1), r(2), r(3), r(4)]);
1095 let x = solve_incr_check(&a, &y, 0, 2);
1096 assert_eq!(x, y);
1097 }
1098
1099 #[test]
1100 fn test_solve_incr_zero_matrix_zero_rhs() {
1101 let a = SpMat::<R>::zero((3, 4));
1102 let x = solve_pluq_incr(&a, &sp_vec([r(0); 3]), 0, 1).expect("zero rhs is consistent");
1104 assert_eq!(x, sp_vec([r(0); 4]));
1105 }
1106
1107 #[test]
1108 fn test_solve_incr_zero_matrix_nonzero_rhs() {
1109 let a = SpMat::<R>::zero((3, 4));
1110 assert!(solve_pluq_incr(&a, &sp_vec([r(1), r(0), r(0)]), 0, 1).is_none());
1111 }
1112
1113 #[test]
1114 fn test_solve_incr_matches_solve_pluq() {
1115 let a: SpMat<R> = sp_mat((6, 9), [
1117 r(1), r(0), r(0), r(0), r(0), r(1), r(0), r(0), r(1),
1118 r(0), r(1), r(1), r(1), r(0), r(1), r(0), r(1), r(0),
1119 r(0), r(0), r(1), r(1), r(0), r(0), r(0), r(1), r(1),
1120 r(0), r(1), r(0), r(0), r(1), r(0), r(0), r(0), r(0),
1121 r(0), r(0), r(1), r(0), r(0), r(0), r(0), r(0), r(0),
1122 r(0), r(1), r(0), r(0), r(0), r(1), r(0), r(1), r(0),
1123 ]);
1124 let y = sp_vec([r(1), r(2), r(3), r(0), r(1), r(0)]);
1125
1126 let y_consistent = &a * &sp_vec(vec![r(1); 9]);
1128
1129 for (mp, ch) in [(0, 1), (0, 3), (2, 2), (usize::MAX, 2)] {
1130 let x_full = solve_pluq(&a, &y);
1132 let x_incr = solve_pluq_incr(&a, &y, mp, ch);
1133 assert_eq!(x_full.is_some(), x_incr.is_some(), "mp={mp}, ch={ch}");
1134
1135 solve_incr_check(&a, &y_consistent, mp, ch);
1137 }
1138 }
1139
1140 #[test]
1143 fn test_merge_perm() {
1144 let perm1 = Perm::from_indices([2, 0, 3, 1, 4]);
1152 let perm2 = Perm::from_indices([1, 2, 0]);
1153 let p = merge_perm(&perm1, perm2);
1154 for (i, expected) in [3, 0, 4, 1, 2].iter().enumerate() {
1155 assert_eq!(p.at(i), *expected, "mismatch at i={i}");
1156 }
1157 }
1158
1159 #[test]
1162 fn test_extend_perm() {
1163 let cp = Perm::from_indices([1, 0]);
1171 let idx = vec![1usize, 3];
1172 let p = extend_perm(5, &idx, cp);
1173 assert_eq!(p.at(0), 2);
1174 assert_eq!(p.at(1), 1);
1175 assert_eq!(p.at(2), 3);
1176 assert_eq!(p.at(3), 0);
1177 assert_eq!(p.at(4), 4);
1178 }
1179
1180 #[test]
1181 fn test_extend_perm_identity() {
1182 let cp = Perm::id(3);
1185 let idx = vec![0usize, 2, 5];
1186 let p = extend_perm(7, &idx, cp);
1187 let expected = Perm::forward_indices(7, idx.iter().copied());
1188 for i in 0..7 {
1189 assert_eq!(p.at(i), expected.at(i), "mismatch at i={i}");
1190 }
1191 }
1192
1193}