Skip to main content

yui_matrix/sparse/
pluq.rs

1//! Sparse PLUQ decomposition and its solvers. `pre_pluq` stops after the
2//! heuristic pivots; `pluq` completes the factorization. `solve_pluq_incr` solves
3//! incrementally, capping the pivots taken per pass.
4
5// Sparse PLUQ decomposition & linear solver.
6// Implemented with the help of Claude Code.
7
8use 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
20/// Result of a sparse PLUQ decomposition.
21///
22/// Satisfies `p * A * q = l * u + s` where `s` is the
23/// `(m - rank) × (n - rank)` Schur complement (bottom-right block).
24pub 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    /// Constructs an `SpPluq` after asserting the shapes are mutually
34    /// consistent: `l.n_cols() == u.n_rows() = r`, `l.n_rows() == p.dim() = m`,
35    /// `u.n_cols() == q.dim() = n`, and `s.shape() == (m - r, n - r)`.
36    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
62/// Converts `a` into the trivial PLUQ whose Schur complement is `a` itself
63impl<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
77/// Computes a partial PLUQ decomposition of `a` under the given pivot-finder
78/// configuration.
79///
80/// Splits the permuted matrix into four blocks `[[a0|a1],[a2|a3]]` at row/col
81/// `r`, then asks Schur to fuse the triangular solve with the Schur update.
82///
83/// Rows: top half `[a0|a1]` is `u` (upper triangular on the left). Schur produces
84///   `l1 = a2·a0⁻¹` and `s = a3 - l1·a1`; final `l = [I_r; l1]`.
85/// Cols: left half `[a0;a2]` is `l` (lower triangular on top). Schur produces
86///   `u1 = a0⁻¹·a1` and `s = a3 - a2·u1`; final `u = [I_r | u1]`.
87pub 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);          // u = [a0 | a1]
107            let l = SpMat::v_stack(SpMat::id(r), l1); // l = [I_r ; l1]
108            (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);            // l = [a0 ; a2]
115            let u = SpMat::h_stack(SpMat::id(r), u1); // u = [I_r | u1]
116            (l, u, s)
117        }
118    };
119
120    SpPluq::new(p, q, l, u, s)
121}
122
123/// Computes a full PLUQ decomposition of `a`.
124///
125/// Iterates `pre_pluq` on the current Schur complement while it keeps finding
126/// sparse pivots; falls back to `dense_pluq_in` once the Schur complement is
127/// non-zero but admits no further sparse pivots.
128pub 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
174// Extracts the compact dense submatrix of `s` using only its non-zero rows/cols.
175// If `transpose` is true, returns S0^T (n0 × m0); otherwise returns S0 (m0 × n0).
176// Also returns the sorted non-zero row/col indices of `s`.
177fn 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(); // S0^T[cj, ri] = S0[ri, cj]
195        } else {
196            mat[(ri, cj)] = v.clone();
197        }
198    }
199
200    (row_idx, col_idx, mat)
201}
202
203// Merges two partial PLUQ decompositions. `pp1` has rank `r1` and shape (m, n);
204// `pp2` is a partial PLUQ of `pp1.s` with rank `r2` and shape (m - r1, n - r1).
205// Returns a partial PLUQ of the same matrix as `pp1` with rank `r1 + r2` and
206// schur complement `pp2.s`.
207fn 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    // MEMO: Even if r2 == 0, there could be non-trivial permutations
219    // when R is not a field.
220
221    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
240/// Solves `a * x = y` over a field using sparse PLUQ.
241///
242/// Returns `Some(x)` if a solution exists, `None` otherwise.
243pub 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())); // y = 0 ⇒ x = 0 solves it — skip the factorization.
251    }
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
266// Solves `L * U * x = y` and returns `x` of length `n = u.n_cols()` with
267// entries beyond `r = l.n_cols()` set to zero (free variables = 0).
268//
269// Requires the top r × r block of L to be unit lower triangular and the top
270// r × r block of U to be invertible upper triangular.
271//
272// Returns `None` when `solve_l(l, y, true)` detects an inconsistent residual.
273// When `l` is square (`l.n_rows() == r`) the residual is empty and the call
274// always succeeds.
275fn 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
286// Solves `l[0..r, 0..r] * z = y[0..r]` by forward substitution, where
287// `r = l.n_cols()`. The top r × r block of L must be lower triangular with
288// non-zero diagonal.
289//
290// If `check_consistency` is true and `r < y.len()`, also verifies the residual
291// `y[r..] - l[r.., :] * z` is zero, returning `None` when it isn't. When `r ==
292// y.len()` the residual is trivially empty so the check is skipped.
293fn 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
315// check l * x == y
316fn 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
339// Solves `u[0..r, 0..r] * x[..r] = y` by back-substitution, where
340// `r = u.n_rows()`, and returns `x` of length `n = u.n_cols()` with entries
341// beyond `r` set to zero. The top r × r block of U must be upper triangular
342// with non-zero diagonal.
343fn 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
360/// Solves `a * x = y` over a field using an incremental sparse PLUQ.
361///
362/// Behaves like [`solve_pluq`] but designed for huge matrices: caps the initial
363/// sparse pre-PLUQ at `max_piv` pivots, then incrementally processes the
364/// remaining Schur complement `chunk` rows at a time. Returns `None` (without
365/// completing the full PLUQ) as soon as a chunk reveals inconsistency.
366pub 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())); // y = 0 ⇒ x = 0 solves it — skip the factorization.
374    }
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        // Apply the chunk's row perm to the tail of yp so it stays in sync with pp.l.
398        let yp_tail = p_next.apply_to(yp[r_old..].to_vec());
399        yp[r_old..].clone_from_slice(&yp_tail);
400
401        // The top `k` rows of pp.s are zero rows (chunk's PLUQ leftover);
402        // they demand `yp[r_new..r_new+k] == L[r_new..r_new+k, :] * z` for
403        // consistency, regardless of future chunks.
404        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
425// Takes the top `min(chunk_size, s.n_rows())` rows of `s`, runs `pluq` on them,
426// and lifts the result to act on all of `s` via `extend_chunk_to_full`.
427// Returns `(pp_chunk_full, r_chunk, c)`.
428fn 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
441// Lifts a PLUQ of the top `c` rows of some matrix `s` (= `pp_chunk`) to a
442// partial PLUQ acting on all rows of `s`, by absorbing the untouched bottom
443// rows `s_rest` (shape (m_s - c, n_s)) into L (below the new pivots) and into
444// the new schur complement.
445//
446// Sparse analog of `dense_pluq_in`, applied to a single chunk.
447fn 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    // Same Schur shape as pre_pluq's Rows branch: u_top (upper triangular) plays
462    // the role of `a`, with `c = s_rest_left`, `b = u_right`, `d = s_rest_right`.
463    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
481// Drops the top `k` zero rows of `pp.s` from `pp.l`, `pp.s`, `pp.p`, and `yp`.
482//
483// Caller is responsible for verifying via `is_consistent_upto` that the
484// `k` rows being removed have zero residual; otherwise the resulting system
485// would silently lose constraints.
486fn 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    // Drop rows [r..r+k] from pp.l: keep [0..r] and [r+k..m], shifted down.
497    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    // Drop the top k rows of pp.s.
508    pp.s = pp.s.submat_rows(k..pp.s.n_rows());
509
510    // Drop yp entries [r..r+k] to stay in sync with pp.l.
511    yp.drain(r..r + k);
512
513    // Drop pp.p entries that map to [r..r+k]; shift later positions down by k.
514    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
527// Composes perm1 with perm2: the first `r` positions stay, the rest are
528// shifted by `r` and remapped by perm2 (where `r = perm1.len() - perm2.len()`).
529fn 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
535// Lifts a compact permutation (acting on compact_idx elements of [0..n]) to the
536// full index space.  compact_idx[k] maps to compact_perm.at(k) (within [0..mr]);
537// all other indices map to consecutive positions starting at mr (in sorted order).
538fn 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    // ---- pre_pluq ----
569
570    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)); // (m-r, n-r) = (4, 5) when r=0
636        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)); // full rank: Schur complement is empty
645    }
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    // ---- pluq ----
664
665    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        // L is unit lower: 1s on diagonal
681        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        // U is unit upper: 1s on diagonal
704        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    // ---- extract_dense ----
712
713    #[test]
714    fn test_extract_dense_no_transpose() {
715        // S has a zero row (row 1) and a zero col (col 1).
716        // Non-zero entries: (0,0)=1, (0,2)=2, (2,0)=3, (2,2)=4.
717        // row_idx=[0,2], col_idx=[0,2].
718        // S0 (2×2) = [[1,2],[3,4]].
719        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        // Same S, but with transpose=true.  S0^T (2×2) = [[1,3],[2,4]].
729        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    // ---- dense_pluq_in ----
737
738    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        // S has a zero row (row 1) and a zero col (col 1).
757        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        // No zero rows/cols: compact_dense gives the full matrix.
777        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    // ---- solve_l ----
802
803    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        // l = [[2, 0], [3, 4]], y = [4, 11]
818        // l[0..2,0..2]*x = [4,11] → x = [2, 5/4]
819        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        // l = [[1,0],[2,1],[3,4]] (3×2 lower triangular with extra row), y = [1,2,3].
828        // Forward sub on top 2×2: z = [1, 2 - 2*1] = [1, 0].
829        // Residual at row 2: 3 - 3*1 - 4*0 = 0 → consistent.
830        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        // Same l as above but y = [1,2,4]. Residual at row 2: 4 - 3 - 0 = 1 ≠ 0 → None.
838        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        // Same inconsistent input as above; with check=false the residual is ignored
846        // and Some(z) is still returned (z is the forward-sub solution on the top).
847        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        // r = 0, y = 0 → returns Some(empty) (third branch with empty submat).
855        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        // r = 0 with non-zero y → residual = y ≠ 0 → None when check=true.
862        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    // ---- is_consistent / is_consistent_upto ----
867
868    #[test]
869    fn test_is_consistent_full() {
870        // l = [[1,0],[2,1],[3,4]], x = [1, 0]:
871        //   y = [1, 2, 3]: residual = [0, 0] → consistent.
872        //   y = [1, 2, 4]: residual at row 2 = 1 ≠ 0 → inconsistent.
873        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        // l = [[1,0],[2,1],[3,4],[5,6]], x = [1, 0], y = [1, 2, 3, 99]:
882        //   row-2 residual = 0; row-3 residual = 99 - 5 = 94.
883        //   k=2: checks rows in [2..2] (none) → trivially true.
884        //   k=3: checks row 2 only → true.
885        //   k=4: checks rows 2 and 3 → false (row 3 fails).
886        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    // ---- solve_u ----
895
896    #[test]
897    fn test_solve_u_square() {
898        // u = [[1, 2], [0, 3]], y = [4, 6].
899        // Back sub: x[1] = 6/3 = 2; x[0] = (4 - 2*2)/1 = 0.
900        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        // u = [[1, 2, 5, 6], [0, 3, 7, 8]], y = [4, 6].
908        // Top 2×2 same as above → x[..2] = [0, 2]; trailing entries are zeros.
909        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        // r = 0, n = 0 → empty input/output.
917        let u: SpMat<R> = SpMat::zero((0, 0));
918        assert_eq!(solve_u(&u, &[]), Vec::<R>::new());
919    }
920
921    // ---- solve_pluq integration tests ----
922
923    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    // ---- extend_chunk_to_full ----
975
976    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    // ---- chunk_pluq ----
1008
1009    #[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    // ---- solve_pluq_incr integration tests ----
1038
1039    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        // exercise different (max_piv, chunk) combinations
1049        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        // Ax = 0 with A=0 has any x as a solution; expect all-zero free vars.
1103        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        // Random sparse system of moderate size; the two solvers should agree.
1116        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        // Pick y that's reachable: y = a * (1, 1, ..., 1) is guaranteed consistent.
1127        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            // Inputs may be inconsistent for this `y`; check both behave the same.
1131            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            // Always-consistent y: both must succeed and produce a solution.
1136            solve_incr_check(&a, &y_consistent, mp, ch);
1137        }
1138    }
1139
1140    // ---- merge_perm ----
1141
1142    #[test]
1143    fn test_merge_perm() {
1144        // perm1 (size 5) = [2, 0, 3, 1, 4]; r = 2; perm2 (size 3) = [1, 2, 0].
1145        // For each i in 0..5, let j = perm1.at(i):
1146        //   i=0: j=2 ≥ r → r + perm2.at(0) = 2 + 1 = 3
1147        //   i=1: j=0 < r → 0
1148        //   i=2: j=3 ≥ r → r + perm2.at(1) = 2 + 2 = 4
1149        //   i=3: j=1 < r → 1
1150        //   i=4: j=4 ≥ r → r + perm2.at(2) = 2 + 0 = 2
1151        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    // ---- extend_perm ----
1160
1161    #[test]
1162    fn test_extend_perm() {
1163        // compact_idx = [1, 3] in full space of size 5.
1164        // compact_perm swaps the two: at(0)=1, at(1)=0.
1165        // Expected:
1166        //   i=1 (compact_idx[0]) -> compact_perm.at(0) = 1
1167        //   i=3 (compact_idx[1]) -> compact_perm.at(1) = 0
1168        //   rest = [0,2,4] -> positions [2,3,4]
1169        //     i=0 -> 2,  i=2 -> 3,  i=4 -> 4
1170        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        // compact_idx = [0, 2, 5] with identity compact_perm.
1183        // extend_perm should equal Perm::forward_indices(7, [0,2,5]).
1184        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}