Skip to main content

yui_matrix/sparse/
pivot.rs

1//! Heuristic pivot finder for sparse PLUQ.
2//!
3//! Implementation based on:
4//!
5//! - "Parallel Sparse PLUQ Factorization modulo p", Charles Bouillaguet,
6//!   Claire Delaplace, Marie-Emilie Voge.
7//!   <https://hal.inria.fr/hal-01646133/document>
8//! - see also: SpaSM (Sparse direct Solver Modulo p),
9//!   <https://github.com/cbouilla/spasm>.
10
11use std::cmp::Ordering;
12use std::collections::VecDeque;
13use itertools::Itertools;
14use log::*;
15
16use yui_core::abst::{Ring, RingOps};
17use yui_core::algo::TopSort;
18use yui_core::util::log::log_step_crossed;
19use crate::Perm;
20use super::*;
21
22cfg_if::cfg_if! {
23    if #[cfg(feature = "multithread")] {
24        use std::cell::RefCell;
25        use std::sync::RwLock;
26        use thread_local::ThreadLocal;
27        use rayon::prelude::*;
28        use yui_core::util::sync::SyncCounter;
29    }
30}
31
32const LOG_THRESHOLD: usize = 10_000;
33const DEFAULT_MAX_PIVOT: usize = usize::MAX;
34
35/// Whether pivots are picked along rows (each pivot eliminates a row's
36/// other entries) or along columns. Affects how the resulting `L`/`U`
37/// blocks are oriented after reduction.
38#[derive(Clone, Copy, PartialEq, Eq, Debug)]
39pub enum PivotType {
40    Rows, Cols
41}
42
43impl PivotType {
44    fn str(&self) -> &'static str {
45        match self {
46            PivotType::Rows => "row",
47            PivotType::Cols => "col"
48        }
49    }
50}
51
52/// Which entries are eligible as pivots:
53/// - `One` — only `±1`.
54/// - `Weight(w)` — any unit whose `c_weight()` is at most `w`.
55/// - `AnyUnit` — any unit of the ring.
56#[derive(Clone, Copy, Debug)]
57pub enum PivotCondition {
58    One, Weight(f64), AnyUnit
59}
60
61/// Knobs for [`find_pivots`].
62#[derive(Clone, Copy, Debug)]
63pub struct PivotFinderConfig {
64    pub piv_type: PivotType,
65    pub piv_cond: PivotCondition,
66    pub max_pivots: usize,
67}
68
69impl Default for PivotFinderConfig {
70    fn default() -> Self {
71        Self {
72            piv_type: PivotType::Rows,
73            piv_cond: PivotCondition::One,
74            max_pivots: DEFAULT_MAX_PIVOT,
75        }
76    }
77}
78
79impl PivotCondition {
80    fn is_cand<R>(&self, r: &R) -> bool
81    where R: Ring, for<'x> &'x R: RingOps<R> {
82        match self {
83            PivotCondition::One       => r.is_pm_one(),
84            PivotCondition::Weight(w) => r.is_unit() && r.c_weight() <= *w,
85            PivotCondition::AnyUnit   => r.is_unit(),
86        }
87    }
88}
89
90/// Searches for pivots in `a` and returns row/column permutations `(p, q)`
91/// that bring the chosen pivots to the leading r×r block of `p·a·q⁻¹`, along
92/// with the rank `r` (number of pivots found).
93pub fn find_pivots<R>(a: &SpMat<R>, config: PivotFinderConfig) -> (Perm, Perm, usize)
94where R: Ring, for<'x> &'x R: RingOps<R> {
95    let (m, n) = a.shape();
96
97    if a.is_zero() {
98        return (Perm::id(m), Perm::id(n), 0);
99    }
100
101    debug!("find {} pivots: {:?}", config.piv_type.str(), a.shape());
102
103    let mut pf = PivotFinder::new(a, &config);
104    pf.find_pivots();
105    let pivs = pf.result();
106
107    debug!("  found {} {} pivots", pivs.len(), config.piv_type.str());
108
109    let p = Perm::forward_indices(m, pivs.iter().map(|(i, _)| *i));
110    let q = Perm::forward_indices(n, pivs.iter().map(|(_, j)| *j));
111    let r = pivs.len();
112
113    (p, q, r)
114}
115
116type Row = usize;
117type Col = usize;
118
119pub struct PivotFinder {
120    str: MatrixStr,
121    pivots: PivotData,
122    piv_type: PivotType,
123    max_pivots: usize,
124}
125
126impl PivotFinder {
127    pub fn new<R>(a: &SpMat<R>, config: &PivotFinderConfig) -> Self
128    where R: Ring, for<'x> &'x R: RingOps<R> {
129        let str = MatrixStr::new(a, config.piv_type, config.piv_cond);
130        let pivots = PivotData::new(a, config.piv_type);
131        PivotFinder { str, pivots, piv_type: config.piv_type, max_pivots: config.max_pivots }
132    }
133
134    pub fn find_pivots(&mut self) {
135        trace!("pivots: {:?} ..", self.str.shape());
136
137        self.find_fl_pivots();
138        if self.pivots.count() < self.max_pivots {
139            self.find_fl_col_pivots();
140        }
141        if self.pivots.count() < self.max_pivots {
142            self.find_cycle_free_pivots();
143        }
144
145        trace!("pivots: {:?} => {}.", self.str.shape(), self.pivots.count());
146    }
147
148    pub fn result(&self) -> Vec<(usize, usize)> {
149        let mut ts = TopSort::new();
150        for (i, j) in self.pivots.iter() {
151            ts.add_node(j);
152            for j2 in self.str.cols_in(i).filter(|&j2| j != j2 && self.pivots.has_col(j2)) {
153                ts.add_edge(j, j2);
154            }
155        }
156        let sorted = ts.into_sorted().unwrap();
157        let is_row_type = self.piv_type == PivotType::Rows;
158
159        sorted.into_iter().map(|j| {
160            let i = self.pivots.row_for(j).unwrap();
161            if is_row_type { (i, j) } else { (j, i) }
162        }).collect_vec()
163    }
164
165    #[cfg(test)]
166    fn rows(&self) -> Row {
167        self.str.shape.0
168    }
169
170    fn cols(&self) -> Col {
171        self.str.shape.1
172    }
173
174    fn remain_rows(&self) -> impl Iterator<Item = Row> + '_ {
175        self.str.target_rows.iter().copied()
176            .filter(|&i| !self.pivots.is_piv_row(i))
177    }
178
179    // occupancy bitmap over columns: the finder already holds O(n_cols) state (col_wght, RowWorker),
180    // so the bitmap is free, and lookups are one per matrix entry — they must be cheaper than hashing.
181    fn occupied_cols(&self) -> Vec<bool> {
182        let mut occ = vec![false; self.cols()];
183        for (i, _) in self.pivots.iter() {
184            for j in self.str.cols_in(i) {
185                occ[j] = true;
186            }
187        }
188        occ
189    }
190
191    fn find_fl_pivots(&mut self) {
192        let remain_rows: Vec<_> = self.remain_rows().collect();
193
194        for i in remain_rows {
195            if self.pivots.count() >= self.max_pivots { break; }
196
197            let Some((j, is_cand)) = self.str.head(i) else { continue };
198
199            if is_cand && !self.pivots.has_col(j) {
200                self.pivots.set(i, j);
201            }
202        }
203
204        let piv_count = self.pivots.count();
205
206        trace!("  fl-pivots: +{}.", piv_count);
207    }
208
209    fn find_fl_col_pivots(&mut self) {
210        let before_piv_count = self.pivots.count();
211
212        let remain_rows: Vec<_> = self.remain_rows().collect();
213        let mut occ_cols = self.occupied_cols();
214
215        for i in remain_rows {
216            if self.pivots.count() >= self.max_pivots { break; }
217
218            let mut cands = vec![];
219
220            for (j, is_cand) in self.str.entries_in(i) {
221                if is_cand && !occ_cols[j] {
222                    cands.push(j);
223                }
224            }
225
226            let Some(j) = cands.into_iter().min_by(|&j1, &j2|
227                self.str.cmp_cols(j1, j2)
228            ) else { continue };
229
230            self.pivots.set(i, j);
231
232            for j in self.str.cols_in(i) {
233                occ_cols[j] = true;
234            }
235        }
236
237        let piv_count = self.pivots.count();
238
239        trace!("  fl-col-pivots: +{}, total: {}.", piv_count - before_piv_count, piv_count);
240    }
241
242    fn find_cycle_free_pivots(&mut self) {
243        let before_piv_count = self.pivots.count();
244
245        cfg_if::cfg_if! {
246            if #[cfg(feature = "multithread")] {
247                self.find_cycle_free_pivots_m();
248            } else {
249                self.find_cycle_free_pivots_s();
250            }
251        }
252
253        let piv_count = self.pivots.count();
254
255        trace!("  cycle-free-pivots: +{}, total: {}.", piv_count - before_piv_count, piv_count);
256    }
257
258    #[allow(unused)]
259    fn find_cycle_free_pivots_s(&mut self) {
260        let remain_rows: Vec<_> = self.remain_rows().collect();
261        let total_rows = remain_rows.len();
262
263        trace!("  start find-cycle-free-pivots: {total_rows} rows");
264
265        let n = self.cols();
266        let mut w = RowWorker::new(n);
267        let mut row_count = 0;
268
269        for i in remain_rows {
270            if self.pivots.count() >= self.max_pivots { break; }
271
272            if let Some(j) = w.find_cycle_free_pivots(i, &self.str, &self.pivots) {
273                self.pivots.set(i, j);
274            }
275
276            row_count += 1;
277            if log_step_crossed(row_count, row_count - 1, total_rows, LOG_THRESHOLD) {
278                let c = self.pivots.count();
279                trace!("    [{row_count}/{total_rows}], {c} pivots.");
280            }
281        }
282    }
283
284     #[cfg(feature = "multithread")]
285     fn find_cycle_free_pivots_m(&mut self) {
286        let remain_rows = self.remain_rows().collect_vec();
287        let total_rows = remain_rows.len();
288
289        trace!("  start find-cycle-free-pivots: {total_rows} rows");
290
291        let n = self.cols();
292        let count = SyncCounter::new(self.pivots.count()); // lock-free mirror of `pivots.count()`
293        let pivots = RwLock::new(
294            std::mem::take(&mut self.pivots)
295        );
296        let loc_pivots_tls = ThreadLocal::new();
297        let loc_worker_tls = ThreadLocal::new();
298
299        let row_counter = SyncCounter::new(0);
300
301        remain_rows.par_iter().for_each(|&i| {
302            if count.count() >= self.max_pivots { return; }
303
304            let mut loc_pivots = init_tls(&loc_pivots_tls, ||
305                pivots.read().unwrap().clone()
306            ).borrow_mut();
307
308            let mut w = init_tls(&loc_worker_tls, ||
309                RowWorker::new(n)
310            ).borrow_mut();
311
312            loc_pivots.update_from(&pivots.read().unwrap());
313            w.init(i, &self.str, &loc_pivots);
314
315            self.find_cycle_free_pivots_in(&pivots, &count, &mut loc_pivots, &mut w);
316
317            let row_count = row_counter.incr();
318            if log_step_crossed(row_count, row_count - 1, total_rows, LOG_THRESHOLD) {
319                let c = loc_pivots.count();
320                trace!("    [{row_count}/{total_rows}], {c} pivots.");
321            }
322        });
323
324        self.pivots = pivots.into_inner().unwrap();
325     }
326
327     #[cfg(feature = "multithread")]
328     fn find_cycle_free_pivots_in(&self, pivots: &RwLock<PivotData>, count: &SyncCounter, loc_pivots: &mut PivotData, w: &mut RowWorker) {
329        loop {
330            w.traverse(&self.str, loc_pivots);
331
332            let Some(j) = w.choose_candidate(&self.str) else {
333                break
334            };
335
336            // If changes are made in other threads, update `loc_pivots` and retry.
337            // Otherwise, modify `pivots` and exit.
338
339            let mut pivots = pivots.write().unwrap();
340            w.update_diff(loc_pivots, &pivots);
341
342            if w.should_retry() {
343                loc_pivots.update_from(&pivots);
344                continue
345            } else {
346                if pivots.count() < self.max_pivots {
347                    pivots.set(w.row, j);
348                    count.set(pivots.count());
349                }
350                break
351            }
352        }
353     }
354
355}
356
357#[cfg(feature = "multithread")]
358fn init_tls<T, F>(tl: &ThreadLocal<RefCell<T>>, f: F) -> &RefCell<T>
359where T: Send, F: FnOnce() -> T {
360    tl.get_or(|| RefCell::new( f() ) )
361}
362
363struct MatrixStr {
364    shape: (usize, usize),
365    entries: Vec<Vec<(Col, bool)>>,    // [row -> sorted [(col, is_cand)]]
366    target_rows: Vec<Row>,             // non-empty rows, sorted by (row_wght, row index)
367    col_wght: Vec<f64>,                // [col -> weight]
368}
369
370impl MatrixStr {
371    fn new<R>(a: &SpMat<R>, piv_type: PivotType, pivot_cond: PivotCondition) -> Self
372    where R: Ring, for<'x> &'x R: RingOps<R> {
373        let shape = match piv_type {
374            PivotType::Rows => a.shape(),
375            PivotType::Cols => (a.n_cols(), a.n_rows())
376        };
377        let t = match piv_type {
378            PivotType::Rows => |i: usize, j: usize| (i, j),
379            PivotType::Cols => |i, j| (j, i)
380        };
381
382        let (m, n) = shape;
383        let mut entries = vec![vec![]; m];
384        let mut row_wght = vec![0.0; m];
385        let mut col_wght = vec![0.0; n];
386
387        for (i, j, r) in a.iter() {
388            if r.is_zero() { continue }
389
390            let (i, j) = t(i, j);
391            entries[i].push((j, pivot_cond.is_cand(r)));
392
393            let w = r.c_weight();
394            row_wght[i] += w;
395            col_wght[j] += w;
396        }
397
398        let mut target_rows: Vec<Row> = (0..m).filter(|&i| !entries[i].is_empty()).collect();
399        target_rows.sort_unstable_by(|&i1, &i2| {
400            row_wght[i1].partial_cmp(&row_wght[i2])
401                .unwrap_or(Ordering::Equal)
402                .then(i1.cmp(&i2))
403        });
404
405        Self { shape, entries, col_wght, target_rows }
406    }
407
408    fn shape(&self) -> (usize, usize) {
409        self.shape
410    }
411
412    fn head(&self, i: Row) -> Option<(Col, bool)> {
413        self.entries[i].first().copied()
414    }
415
416    fn entries_in(&self, i: Row) -> impl Iterator<Item = (Col, bool)> + '_ {
417        self.entries[i].iter().copied()
418    }
419
420    fn cols_in(&self, i: Row) -> impl Iterator<Item = Col> + '_ {
421        self.entries[i].iter().map(|(c, _)| *c)
422    }
423
424    fn cmp_cols(&self, j1: Col, j2: Col) -> Ordering {
425        if let Some(o) = self.col_wght[j1].partial_cmp(&self.col_wght[j2]) {
426            o.then(Ord::cmp(&j1, &j2))
427        } else {
428            Ordering::Equal
429        }
430    }
431}
432
433#[derive(Clone, Default)]
434struct PivotData {
435    data: Vec<Option<Row>>,   // col -> row
436    indices: Vec<Col>,
437    is_piv_row: Vec<bool>,    // row -> is this a pivot row?
438}
439
440impl PivotData {
441    fn new<R>(a: &SpMat<R>, piv_type: PivotType) -> Self
442    where R: Ring, for<'x> &'x R: RingOps<R> {
443        let (m, n) = match piv_type {
444            PivotType::Rows => (a.n_rows(), a.n_cols()),
445            PivotType::Cols => (a.n_cols(), a.n_rows()),
446        };
447        let data = vec![None; n];
448        let indices = vec![];
449        let is_piv_row = vec![false; m];
450        Self { data, indices, is_piv_row }
451    }
452
453    fn count(&self) -> usize {
454        self.indices.len()
455    }
456
457    fn has_col(&self, j: Col) -> bool {
458        self.data[j].is_some()
459    }
460
461    fn is_piv_row(&self, i: Row) -> bool {
462        self.is_piv_row[i]
463    }
464
465    fn row_for(&self, j: Col) -> Option<Row> {
466        self.data[j]
467    }
468
469    fn set(&mut self, i: Row, j: Col) {
470        assert!(!self.has_col(j));
471        self.data[j] = Some(i);
472        self.indices.push(j);
473        self.is_piv_row[i] = true;
474    }
475
476    fn iter(&self) -> impl Iterator<Item = (Row, Col)> + '_ {
477        self.indices.iter().map(|&j| {
478            let i = self.data[j].unwrap();
479            (i, j)
480        })
481    }
482
483    #[allow(unused)]
484    fn pivot_at(&self, k: usize) -> (Row, Col) {
485        let j = self.indices[k];
486        let i = self.data[j].unwrap();
487        (i, j)
488    }
489
490    fn update_from(&mut self, from: &Self) {
491        debug_assert!(self.count() <= from.count());
492        for k in self.count() .. from.count() {
493            let (i, j) = from.pivot_at(k);
494            self.set(i, j);
495        }
496    }
497}
498
499#[repr(u8)]
500#[derive(Clone, Copy, PartialEq, Eq, Debug)]
501enum EntryStatus {
502    None, Candidate, Occupied
503}
504
505struct RowWorker {
506    row: usize,
507    status: Vec<EntryStatus>,
508    ncand: usize,
509    candidates: Vec<Col>,    // columns ever marked Candidate (may include stale entries reclassified to Occupied)
510    queue: VecDeque<Col>,
511    queued: Vec<bool>,       // queued[j] = true iff j has ever been pushed to `queue`
512    touched: Vec<Col>,       // columns whose status/queued changed — `clear` resets only these, O(touched) not O(n)
513}
514
515impl RowWorker {
516    fn new(size: usize) -> Self {
517        let status = vec![EntryStatus::None; size];
518        let candidates = Vec::new();
519        let queue = VecDeque::new();
520        let queued = vec![false; size];
521        let touched = Vec::new();
522        RowWorker { row: 0, status, ncand: 0, candidates, queue, queued, touched }
523    }
524
525    fn clear(&mut self) {
526        self.row = 0;
527        self.ncand = 0;
528        for &j in self.touched.iter() {
529            self.status[j] = EntryStatus::None;
530            self.queued[j] = false;
531        }
532        self.touched.clear();
533        self.candidates.clear();
534        self.queue.clear();
535    }
536
537    //  i [  o       #     # ]     [  o   x   x      # ]     [  o   x   x   x  # ]
538    //    [  |               ]     [  |   :   :        ]     [  |   :   :   :    ]
539    //    [  *   .   .       ] ~~> [  * - o - .        ] ~~> [  * - o - .   :    ]
540    //    [                  ]     [      |            ]     [      |       :    ]
541    //    [      *       .   ]     [      *       .    ]     [      *-------.    ]
542    //
543    //  o: queued, #: candidate, x: occupied
544
545    fn find_cycle_free_pivots(&mut self, i: usize, str: &MatrixStr, pivots: &PivotData) -> Option<Col> {
546        self.init(i, str, pivots);
547        self.traverse(str, pivots);
548        self.choose_candidate(str)
549    }
550
551    fn init(&mut self, i: usize, str: &MatrixStr, pivots: &PivotData) {
552        self.clear();
553        self.row = i;
554
555        for (j, is_cand) in str.entries_in(i) {
556            if pivots.has_col(j) {
557                self.enqueue(j);
558                self.set_occupied(j);
559            } else if is_cand {
560                self.set_candidate(j);
561            } else {
562                self.set_occupied(j);
563            }
564        }
565    }
566
567    fn traverse(&mut self, str: &MatrixStr, pivots: &PivotData) {
568        if !self.has_candidate() {
569            return
570        }
571
572        while let Some(j) = self.dequeue() {
573            let i2 = pivots.row_for(j).unwrap();
574
575            for j2 in str.cols_in(i2) {
576                if pivots.has_col(j2) && !self.is_queued(j2) {
577                    self.enqueue(j2);
578                }
579
580                self.set_occupied(j2);
581
582                if !self.has_candidate() {
583                    break
584                }
585            }
586        }
587    }
588
589    fn choose_candidate(&self, str: &MatrixStr) -> Option<Col> {
590        if self.ncand == 0 { return None; }
591        self.candidates.iter().copied()
592            .filter(|&j| self.is_candidate(j))
593            .min_by(|&j1, &j2|
594                str.cmp_cols(j1, j2)
595            )
596    }
597
598    #[allow(dead_code)]
599    fn update_diff(&mut self, loc_pivots: &PivotData, pivots: &PivotData) {
600        debug_assert!(loc_pivots.count() <= pivots.count());
601        for k in loc_pivots.count()..pivots.count() {
602            let j = pivots.indices[k];
603            if self.is_candidate(j) || self.is_occupied(j) {
604                self.enqueue(j);
605                self.set_occupied(j);
606            }
607        }
608    }
609
610    fn should_retry(&self) -> bool {
611        !self.queue.is_empty()
612    }
613
614    fn has_candidate(&self) -> bool {
615        self.ncand > 0
616    }
617
618    fn is_candidate(&self, i: usize) -> bool {
619        self.status[i] == EntryStatus::Candidate
620    }
621
622    fn set_candidate(&mut self, i: usize) {
623        assert_eq!(self.status[i], EntryStatus::None);
624        self.status[i] = EntryStatus::Candidate;
625        self.candidates.push(i);
626        self.touched.push(i);
627        self.ncand += 1;
628    }
629
630    fn is_occupied(&self, i: usize) -> bool {
631        self.status[i] == EntryStatus::Occupied
632    }
633
634    fn set_occupied(&mut self, i: usize) {
635        match self.status[i] {
636            EntryStatus::Candidate => {
637                self.ncand -= 1; // already in `touched` via set_candidate
638            }
639            EntryStatus::None => {
640                self.touched.push(i);
641            }
642            EntryStatus::Occupied => {
643                return;
644            }
645        }
646        self.status[i] = EntryStatus::Occupied;
647    }
648
649    fn enqueue(&mut self, i: Col) {
650        self.queue.push_back(i);
651        if !self.queued[i] {
652            self.queued[i] = true;
653            self.touched.push(i);
654        }
655    }
656
657    fn dequeue(&mut self) -> Option<Col> {
658        self.queue.pop_front()
659    }
660
661    fn is_queued(&self, i: Col) -> bool {
662        self.queued[i]
663    }
664}
665
666#[cfg(test)]
667mod tests {
668    use super::*;
669    use num_traits::{Zero, One};
670
671    #[test]
672    fn str_init() {
673        let a = SpMat::from_row_major((6, 9), [
674            1, 0, 1, 0, 0, 1, 1, 0, 1,
675            0, 1, 1, 1, 0, 1, 0, 2, 0,
676            0, 0, 1, 1, 0, 0, 0, 1, 1,
677            0, 1, 1, 0, 3, 0, 0, 0, 0,
678            0, 1, 0, 1, 0, 0, 1, 0, 1,
679            1, 0, 1, 0, 1, 1, 0, 1, 1
680        ]);
681        let str = MatrixStr::new(&a, PivotType::Rows, PivotCondition::One);
682
683        assert_eq!(str.entries, vec![
684            vec![(0,true), (2,true), (5,true), (6,true), (8,true)],
685            vec![(1,true), (2,true), (3,true), (5,true), (7,false)],
686            vec![(2,true), (3,true), (7,true), (8,true)],
687            vec![(1,true), (2,true), (4,false)],
688            vec![(1,true), (3,true), (6,true), (8,true)],
689            vec![(0,true), (2,true), (4,true), (5,true), (7,true), (8,true)],
690        ]);
691        assert_eq!(str.col_wght, vec![2.0, 3.0, 5.0, 3.0, 4.0, 3.0, 2.0, 4.0, 4.0]);
692        assert_eq!(str.target_rows, vec![2, 4, 0, 3, 1, 5]);
693    }
694
695    #[test]
696    fn str_row_head() {
697        let a = SpMat::from_row_major((4, 4), [
698            1, 0, 1, 0,
699            0, 1, 1, 1,
700            0, 0, 0, 0,
701            0, 0, 1, 1,
702        ]);
703        let str = MatrixStr::new(&a, PivotType::Rows, PivotCondition::One);
704
705        assert_eq!(str.head(0), Some((0, true)));
706        assert_eq!(str.head(1), Some((1, true)));
707        assert_eq!(str.head(2), None);
708        assert_eq!(str.head(3), Some((2, true)));
709    }
710
711    #[test]
712    fn rows_cols() {
713        let a = SpMat::<i32>::from_row_major((4, 3), []);
714        let pf = PivotFinder::new(&a, &Default::default());
715        assert_eq!(pf.rows(), 4);
716        assert_eq!(pf.cols(), 3);
717    }
718
719    #[test]
720    fn pivot_data() {
721        let a = SpMat::from_row_major((2, 4), [
722            1, 0, 1, 0,
723            0, 0, 1, 1,
724        ]);
725        let mut piv = PivotData::new(&a, PivotType::Rows);
726
727        assert_eq!(piv.count(), 0);
728        assert!(!piv.has_col(0));
729        assert_eq!(piv.row_for(0), None);
730
731        piv.set(1, 2);
732
733        assert_eq!(piv.count(), 1);
734        assert!(piv.has_col(2));
735        assert_eq!(piv.row_for(2), Some(1));
736    }
737
738    #[test]
739    fn remain_rows() {
740        let a = SpMat::from_row_major((4, 4), [
741            1, 0, 1, 0,
742            0, 1, 1, 1,
743            0, 0, 0, 0,
744            0, 0, 1, 1,
745        ]);
746        let mut pf = PivotFinder::new(&a, &Default::default());
747
748        assert_eq!(pf.remain_rows().collect_vec(), vec![0,3,1]);
749
750        pf.pivots.set(0, 0);
751
752        assert_eq!(pf.remain_rows().collect_vec(), vec![3,1]);
753
754        pf.pivots.set(1, 1);
755
756        assert_eq!(pf.remain_rows().collect_vec(), vec![3]);
757    }
758
759    #[test]
760    fn pivots() {
761        let a = SpMat::from_row_major((4, 4), [
762            1, 0, 1, 0,
763            0, 1, 1, 1,
764            0, 0, 0, 0,
765            0, 0, 1, 1,
766        ]);
767        let mut pf = PivotFinder::new(&a, &Default::default());
768
769        let occ = |pf: &PivotFinder| pf.occupied_cols().iter().positions(|&b| b).collect_vec();
770
771        assert!(occ(&pf).is_empty());
772
773        pf.pivots.set(0, 0);
774
775        assert_eq!(occ(&pf), vec![0, 2]);
776
777        pf.pivots.set(1, 1);
778
779        assert_eq!(occ(&pf), vec![0, 1, 2, 3]);
780    }
781
782    #[test]
783    fn find_fl_pivots() {
784        let a = SpMat::from_row_major((6, 9), [
785            1, 0, 1, 0, 0, 1, 1, 0, 1,
786            0, 1, 1, 1, 0, 1, 0, 1, 0,
787            0, 0, 1, 1, 0, 0, 0, 1, 1,
788            0, 1, 1, 0, 1, 0, 0, 0, 0,
789            0, 0, 1, 1, 0, 0, 0, 0, 0,
790            0, 0, 0, 0, 0, 1, 0, 1, 1
791        ]);
792        let mut pf = PivotFinder::new(&a, &Default::default());
793
794        pf.find_fl_pivots();
795
796        assert_eq!(pf.pivots.iter().collect_vec(), vec![(4, 2), (3, 1), (5, 5), (0, 0)]);
797    }
798
799    #[test]
800    fn find_fl_col_pivots() {
801        let a = SpMat::from_row_major((6, 9), [
802            1, 0, 0, 0, 0, 1, 0, 0, 1,
803            0, 1, 1, 1, 0, 1, 0, 1, 0,
804            0, 0, 1, 1, 0, 0, 0, 1, 1,
805            0, 1, 0, 0, 1, 0, 0, 0, 0,
806            0, 0, 1, 0, 0, 0, 0, 0, 0,
807            0, 1, 0, 0, 0, 1, 0, 1, 0
808        ]);
809        let mut pf = PivotFinder::new(&a, &Default::default());
810
811        pf.find_fl_col_pivots();
812
813        assert_eq!(pf.pivots.iter().collect_vec(), vec![(4, 2), (3, 4), (0, 0), (5, 7), (2, 3)]);
814    }
815
816    #[test]
817    fn find_fl_row_col_pivots() {
818        let a = SpMat::from_row_major((6, 9), [
819            1, 0, 0, 0, 0, 1, 0, 0, 1,
820            0, 1, 1, 1, 0, 1, 0, 1, 0,
821            0, 0, 1, 1, 0, 0, 0, 1, 1,
822            0, 1, 0, 0, 1, 0, 0, 0, 0,
823            0, 0, 1, 0, 0, 0, 0, 0, 0,
824            0, 1, 0, 0, 0, 1, 0, 1, 0
825        ]);
826        let mut pf = PivotFinder::new(&a, &Default::default());
827
828        pf.find_fl_pivots();
829
830        assert_eq!(pf.pivots.iter().collect_vec(), vec![(4, 2), (3, 1), (0, 0)]);
831
832        pf.find_fl_col_pivots();
833
834        assert_eq!(pf.pivots.iter().collect_vec(), vec![(4, 2), (3, 1), (0, 0), (5, 7), (2, 3)]);
835    }
836
837    #[test]
838    fn find_cycle_free_pivots_s() {
839        let a = SpMat::from_row_major((6, 9), [
840            1, 0, 0, 0, 0, 1, 0, 0, 1,
841            0, 1, 1, 1, 0, 1, 0, 1, 0,
842            0, 0, 1, 1, 0, 0, 0, 1, 1,
843            0, 1, 0, 0, 1, 0, 0, 0, 0,
844            0, 0, 1, 0, 0, 0, 0, 0, 0,
845            0, 1, 0, 0, 0, 1, 0, 1, 0
846        ]);
847        let mut pf = PivotFinder::new(&a, &Default::default());
848
849        pf.find_cycle_free_pivots_s();
850
851        assert_eq!(pf.pivots.iter().collect_vec(), vec![(4, 2), (3, 4), (0, 0), (5, 1), (2, 3)]);
852    }
853
854    #[cfg(feature = "multithread")]
855    #[test]
856    fn find_cycle_free_pivots_m() {
857        let a = SpMat::from_row_major((6, 9), [
858            1, 0, 0, 0, 0, 1, 0, 0, 1,
859            0, 1, 1, 1, 0, 1, 0, 1, 0,
860            0, 0, 1, 1, 0, 0, 0, 1, 1,
861            0, 1, 0, 0, 1, 0, 0, 0, 0,
862            0, 0, 1, 0, 0, 0, 0, 0, 0,
863            0, 1, 0, 0, 0, 1, 0, 1, 0
864        ]);
865        let mut pf = PivotFinder::new(&a, &Default::default());
866
867        pf.find_cycle_free_pivots_m();
868
869        assert!(pf.pivots.count() >= 5);
870    }
871
872    #[test]
873    fn zero() {
874        let a = SpMat::from_row_major((1, 1), [0]);
875        let (p, q, r) = find_pivots(&a, Default::default());
876        assert_eq!(r, 0);
877        assert_eq!(p.len(), 1);
878        assert_eq!(q.len(), 1);
879        assert!(p.is_id());
880        assert!(q.is_id());
881    }
882
883    #[test]
884    fn id_1() {
885        let a = SpMat::from_row_major((1, 1), [1]);
886        let (p, q, r) = find_pivots(&a, Default::default());
887        assert_eq!(r, 1);
888        assert_eq!(p.len(), 1);
889        assert_eq!(q.len(), 1);
890        assert!(p.is_id());
891        assert!(q.is_id());
892    }
893
894    #[test]
895    fn id_2() {
896        let a = SpMat::from_row_major((2, 2), [
897            1, 0, 0, 1
898        ]);
899        let (p, q, r) = find_pivots(&a, Default::default());
900        assert_eq!(r, 2);
901        assert_eq!(p.len(), 2);
902        assert_eq!(q.len(), 2);
903    }
904
905    #[test]
906    fn result() {
907        let a = SpMat::from_row_major((6, 9), [
908            1, 0, 0, 0, 0, 1, 0, 0, 1,
909            0, 1, 1, 1, 0, 1, 0, 1, 0,
910            0, 0, 1, 1, 0, 0, 0, 1, 1,
911            0, 1, 0, 0, 1, 0, 0, 0, 0,
912            0, 0, 1, 0, 0, 0, 0, 0, 0,
913            0, 1, 0, 0, 0, 1, 0, 1, 0
914        ]);
915        let (p, q, r) = find_pivots(&a, Default::default());
916        assert_eq!(r, 5);
917        assert_eq!(p.len(), 6);
918        assert_eq!(q.len(), 9);
919
920        let b = a.permute(&p, &q).into_dense();
921
922        assert!((0..r).all(|i| b[(i, i)].is_one()));
923        assert!((0..r).all(|j| {
924            (j+1..r).all(|i| b[(i, j)].is_zero())
925        }));
926    }
927
928    #[test]
929    fn result_cols() {
930        let a = SpMat::from_row_major((6, 9), [
931            1, 0, 0, 0, 0, 1, 0, 0, 1,
932            0, 1, 1, 1, 0, 1, 0, 1, 0,
933            0, 0, 1, 1, 0, 0, 0, 1, 1,
934            0, 1, 0, 0, 1, 0, 0, 0, 0,
935            0, 0, 1, 0, 0, 0, 0, 0, 0,
936            0, 1, 0, 0, 0, 1, 0, 1, 0
937        ]);
938        let config = PivotFinderConfig { piv_type: PivotType::Cols, ..Default::default() };
939        let (p, q, r) = find_pivots(&a, config);
940        assert_eq!(r, 6);
941        assert_eq!(p.len(), 6);
942        assert_eq!(q.len(), 9);
943
944        let b = a.permute(&p, &q).into_dense();
945
946        assert!((0..r).all(|i| b[(i, i)].is_one()));
947        assert!((0..r).all(|i| {
948            (i+1..r).all(|j| b[(i, j)].is_zero())
949        }));
950    }
951
952    #[test]
953    fn rand() {
954        let d = 0.1;
955        let shape = (60, 80);
956        let a = SpMat::<i32>::rand(shape, d);
957
958        let (p, q, r) = find_pivots(&a, Default::default());
959        assert!(r > 10);
960        assert_eq!(p.len(), shape.0);
961        assert_eq!(q.len(), shape.1);
962
963        let b = a.permute(&p, &q).into_dense();
964
965        assert!((0..r).all(|i| b[(i, i)].is_one()));
966        assert!((0..r).all(|j| {
967            (j+1..r).all(|i| b[(i, j)].is_zero())
968        }))
969    }
970
971    #[test]
972    fn max_pivots_zero() {
973        // the cap is tested before a pivot is taken, so zero really means zero.
974        let a: SpMat<i64> = SpMat::id(3);
975        let config = PivotFinderConfig { max_pivots: 0, ..Default::default() };
976        let (_, _, r) = find_pivots(&a, config);
977        assert_eq!(r, 0);
978    }
979
980    #[test]
981    fn max_pivots() {
982        // Full rank of this matrix is 5; limiting to 3 must return ≤ 3 pivots.
983        let a = SpMat::from_row_major((6, 9), [
984            1, 0, 0, 0, 0, 1, 0, 0, 1,
985            0, 1, 1, 1, 0, 1, 0, 1, 0,
986            0, 0, 1, 1, 0, 0, 0, 1, 1,
987            0, 1, 0, 0, 1, 0, 0, 0, 0,
988            0, 0, 1, 0, 0, 0, 0, 0, 0,
989            0, 1, 0, 0, 0, 1, 0, 1, 0
990        ]);
991        let config = PivotFinderConfig { max_pivots: 3, ..Default::default() };
992        let (p, q, r) = find_pivots(&a, config);
993        assert!(r <= 3);
994        assert_eq!(p.len(), 6);
995        assert_eq!(q.len(), 9);
996
997        let b = a.permute(&p, &q).into_dense();
998        assert!((0..r).all(|i| b[(i, i)].is_one()));
999        assert!((0..r).all(|j| {
1000            (j+1..r).all(|i| b[(i, j)].is_zero())
1001        }));
1002    }
1003
1004}