1use 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#[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#[derive(Clone, Copy, Debug)]
57pub enum PivotCondition {
58 One, Weight(f64), AnyUnit
59}
60
61#[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
90pub 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 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()); 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 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)>>, target_rows: Vec<Row>, col_wght: Vec<f64>, }
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>>, indices: Vec<Col>,
437 is_piv_row: Vec<bool>, }
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>, queue: VecDeque<Col>,
511 queued: Vec<bool>, touched: Vec<Col>, }
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 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; }
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 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 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}