Skip to main content

conspire/math/sparse/factor/lu/
mod.rs

1#[cfg(target_arch = "x86_64")]
2mod avx;
3#[cfg(test)]
4mod test;
5
6use super::super::{SparseError, matrix::CscMatrix};
7use super::gemm::{CHUNK, NONE, axpy, etree, gemm_wide, max_below, reach_sorted, supernodes};
8use crate::{
9    ABS_TOL,
10    math::{Scalar, Vector},
11};
12
13/// Threshold for preferring the diagonal pivot, which preserves the
14/// fill-reducing ordering when the matrix has a symmetric pattern.
15const PIVOT_TOL: Scalar = 0.001;
16
17/// A sparse LU factorization, PAQ = LU, with L stored as dense supernodal panels.
18pub struct CscLu {
19    pub(super) fill: usize,
20    pub(super) sn_of: Vec<usize>,
21    pub(super) sn_start: Vec<usize>,
22    pub(super) sn_rows_ptr: Vec<usize>,
23    pub(super) sn_rows: Vec<usize>,
24    pub(super) sn_panel_ptr: Vec<usize>,
25    pub(super) sn_values: Vec<Scalar>,
26    pub(super) u_col_ptr: Vec<usize>,
27    pub(super) u_row_idx: Vec<usize>,
28    pub(super) u_values: Vec<Scalar>,
29    pub(super) pinv: Vec<usize>,
30    pub(super) q: Vec<usize>,
31}
32
33impl CscMatrix {
34    /// Factors PA = LU using the Gilbert-Peierls method with partial pivoting.
35    pub fn lu(&self) -> Result<CscLu, SparseError> {
36        self.factor((0..self.height()).collect())
37    }
38    /// Factors PAQ = LU using the Gilbert-Peierls method with partial pivoting,
39    /// with a fill-reducing approximate minimum degree column ordering.
40    pub fn lu_amd(&self) -> Result<CscLu, SparseError> {
41        self.factor(self.amd())
42    }
43    /// Builds the factorization pattern symbolically for the fill-reducing
44    /// ordering assuming pivots from a maximum transversal (the diagonal when it
45    /// is structurally full), which is the exact fill pattern for a symmetric
46    /// pattern and a superset otherwise; all values are zero until a
47    /// refactorization supplies them.
48    pub fn lu_symbolic(&self) -> Result<CscLu, SparseError> {
49        let n = self.height();
50        assert_eq!(n, self.width());
51        let matching = self.maxtrans().ok_or(SparseError::Singular)?;
52        let q = if matching.iter().enumerate().all(|(c, &r)| r == c) {
53            self.amd()
54        } else {
55            let mut rinv = vec![0; n];
56            matching.iter().enumerate().for_each(|(c, &r)| rinv[r] = c);
57            let mut permuted = Vec::with_capacity(self.nonzeros());
58            (0..n).for_each(|c| {
59                self.column(c)
60                    .for_each(|(r, _)| permuted.push((rinv[r], c)))
61            });
62            CscMatrix::from_pattern(n, n, permuted).amd()
63        };
64        let mut pinv = vec![NONE; n];
65        q.iter()
66            .enumerate()
67            .for_each(|(j, &q_j)| pinv[matching[q_j]] = j);
68        Ok(self.symbolic(q, pinv, &vec![false; n]))
69    }
70    pub(super) fn symbolic(&self, q: Vec<usize>, pinv: Vec<usize>, locked: &[bool]) -> CscLu {
71        let n = self.height();
72        let mut qinv = vec![NONE; n];
73        q.iter().enumerate().for_each(|(j, &q_j)| qinv[q_j] = j);
74        let (parent, adj_ptr, adj) = etree(self, n, n, |c| qinv[c], |r| pinv[r]);
75        let mut mark = vec![NONE; n];
76        let mut row = Vec::new();
77        let mut l_count = vec![1_usize; n];
78        let mut u_col_ptr = Vec::with_capacity(n + 1);
79        let mut u_row_idx = Vec::new();
80        u_col_ptr.push(0);
81        (0..n).for_each(|k| {
82            reach_sorted(k, &adj_ptr, &adj, &parent, &mut mark, &mut row);
83            row.iter().for_each(|&j| l_count[j] += 1);
84            u_row_idx.extend_from_slice(&row);
85            u_row_idx.push(k);
86            u_col_ptr.push(u_row_idx.len());
87        });
88        let mut l_col_ptr = Vec::with_capacity(n + 1);
89        l_col_ptr.push(0);
90        l_count.iter().for_each(|count| {
91            l_col_ptr.push(l_col_ptr.last().unwrap() + count);
92        });
93        let mut l_row_idx = vec![0; l_col_ptr[n]];
94        let mut next = vec![0; n];
95        (0..n).for_each(|j| {
96            l_row_idx[l_col_ptr[j]] = j;
97            next[j] = l_col_ptr[j] + 1;
98        });
99        (0..n).for_each(|k| {
100            u_row_idx[u_col_ptr[k]..u_col_ptr[k + 1] - 1]
101                .iter()
102                .for_each(|&j| {
103                    l_row_idx[next[j]] = k;
104                    next[j] += 1;
105                });
106        });
107        let fill = l_row_idx.len() + u_row_idx.len();
108        let l_values = vec![0.0; l_row_idx.len()];
109        let (sn_of, sn_start, sn_rows_ptr, sn_rows, sn_panel_ptr, sn_values) =
110            supernodes(&l_col_ptr, &l_row_idx, &l_values, n, locked);
111        CscLu {
112            fill,
113            sn_of,
114            sn_start,
115            sn_rows_ptr,
116            sn_rows,
117            sn_panel_ptr,
118            sn_values,
119            u_col_ptr,
120            u_row_idx,
121            u_values: vec![0.0; fill - l_values.len()],
122            pinv,
123            q,
124        }
125    }
126    fn factor(&self, q: Vec<usize>) -> Result<CscLu, SparseError> {
127        let n = self.height();
128        assert_eq!(n, self.width());
129        let mut pinv = vec![NONE; n];
130        let mut l_cols = Vec::<Vec<(usize, Scalar)>>::with_capacity(n);
131        let mut u_cols = Vec::<Vec<(usize, Scalar)>>::with_capacity(n);
132        let mut x = vec![0.0; n];
133        let mut mark = vec![0; n];
134        let mut order = vec![0; n];
135        let mut stack = vec![0; n];
136        let mut pstack = vec![0; n];
137        let mut top;
138        for (j, &q_j) in q.iter().enumerate() {
139            top = reach(
140                self.column(q_j).map(|(i, _)| i),
141                &l_cols,
142                &pinv,
143                j + 1,
144                &mut mark,
145                &mut order,
146                &mut stack,
147                &mut pstack,
148            );
149            self.column(q_j).for_each(|(i, value)| x[i] = *value);
150            order[top..n].iter().for_each(|&i| {
151                if pinv[i] != NONE {
152                    let x_i = x[i];
153                    l_cols[pinv[i]]
154                        .iter()
155                        .skip(1)
156                        .for_each(|&(row, value)| x[row] -= value * x_i);
157                }
158            });
159            let mut pivot_row = NONE;
160            let mut pivot_abs = 0.0;
161            order[top..n].iter().for_each(|&i| {
162                if pinv[i] == NONE && x[i].abs() > pivot_abs {
163                    pivot_abs = x[i].abs();
164                    pivot_row = i;
165                }
166            });
167            if pivot_row == NONE || pivot_abs < ABS_TOL {
168                return Err(SparseError::Singular);
169            }
170            if pinv[q_j] == NONE && x[q_j].abs() >= PIVOT_TOL * pivot_abs {
171                pivot_row = q_j;
172            }
173            let pivot = x[pivot_row];
174            pinv[pivot_row] = j;
175            let mut l_col = vec![(pivot_row, 1.0)];
176            let mut u_col = Vec::new();
177            order[top..n].iter().for_each(|&i| {
178                if pinv[i] == NONE {
179                    l_col.push((i, x[i] / pivot));
180                } else if i != pivot_row {
181                    u_col.push((pinv[i], x[i]));
182                }
183                x[i] = 0.0;
184            });
185            u_col.push((j, pivot));
186            u_col.sort_unstable_by_key(|&(row, _)| row);
187            l_cols.push(l_col);
188            u_cols.push(u_col);
189        }
190        let (l_col_ptr, l_row_idx, l_values) = compress(l_cols, &pinv);
191        let (u_col_ptr, u_row_idx, u_values) = compress(u_cols, &(0..n).collect::<Vec<usize>>());
192        let fill = l_values.len() + u_values.len();
193        let (sn_of, sn_start, sn_rows_ptr, sn_rows, sn_panel_ptr, sn_values) =
194            supernodes(&l_col_ptr, &l_row_idx, &l_values, n, &vec![false; n]);
195        Ok(CscLu {
196            fill,
197            sn_of,
198            sn_start,
199            sn_rows_ptr,
200            sn_rows,
201            sn_panel_ptr,
202            sn_values,
203            u_col_ptr,
204            u_row_idx,
205            u_values,
206            pinv,
207            q,
208        })
209    }
210}
211
212impl CscLu {
213    /// Solve a system of linear equations using the factorization.
214    pub fn solve(&self, b: &Vector) -> Vector {
215        let n = self.pinv.len();
216        let mut x = vec![0.0; n];
217        let mut below = vec![0.0; self.max_below()];
218        self.pinv
219            .iter()
220            .enumerate()
221            .for_each(|(i, &p_i)| x[p_i] = b[i]);
222        (0..self.sn_start.len() - 1).for_each(|s| {
223            let t1 = self.sn_start[s];
224            let t2 = self.sn_start[s + 1];
225            let width = t2 - t1;
226            let rows = &self.sn_rows[self.sn_rows_ptr[s]..self.sn_rows_ptr[s + 1]];
227            let m = rows.len();
228            let panel = &self.sn_values[self.sn_panel_ptr[s]..self.sn_panel_ptr[s + 1]];
229            (0..width).for_each(|c| {
230                let x_c = x[t1 + c];
231                if x_c != 0.0 {
232                    let column = &panel[c * m..(c + 1) * m];
233                    x[t1 + c + 1..t2]
234                        .iter_mut()
235                        .zip(column[c + 1..width].iter())
236                        .for_each(|(x_r, value)| *x_r -= value * x_c);
237                    below[..m - width]
238                        .iter_mut()
239                        .zip(column[width..].iter())
240                        .for_each(|(below_r, value)| *below_r += value * x_c);
241                }
242            });
243            rows[width..]
244                .iter()
245                .zip(below[..m - width].iter_mut())
246                .for_each(|(&row, below_r)| {
247                    x[row] -= *below_r;
248                    *below_r = 0.0;
249                });
250        });
251        (0..n).rev().for_each(|j| {
252            let end = self.u_col_ptr[j + 1];
253            x[j] /= self.u_values[end - 1];
254            let x_j = x[j];
255            if x_j != 0.0 {
256                self.u_row_idx[self.u_col_ptr[j]..end - 1]
257                    .iter()
258                    .zip(self.u_values[self.u_col_ptr[j]..end - 1].iter())
259                    .for_each(|(&row, value)| x[row] -= value * x_j);
260            }
261        });
262        let mut solution = Vector::zero(n);
263        self.q
264            .iter()
265            .enumerate()
266            .for_each(|(j, &q_j)| solution[q_j] = x[j]);
267        solution
268    }
269    /// The number of nonzero entries in the factors.
270    pub fn nonzeros(&self) -> usize {
271        self.fill
272    }
273    /// Recomputes the factorization for new values in the same pattern, reusing
274    /// the pivot order and fill pattern without any symbolic work or pivot search.
275    /// The factorization is invalid if an error is returned.
276    pub fn refactor(&mut self, matrix: &CscMatrix) -> Result<(), SparseError> {
277        let n = self.pinv.len();
278        assert_eq!(n, matrix.height());
279        let mut work = vec![0.0; n * CHUNK];
280        let mut temp = vec![0.0; CHUNK * self.max_below()];
281        let mut tile = vec![
282            0.0;
283            CHUNK
284                * (0..self.sn_start.len() - 1)
285                    .map(|s| self.sn_start[s + 1] - self.sn_start[s])
286                    .max()
287                    .unwrap_or(0)
288        ];
289        let mut pointers = [0; CHUNK];
290        for s in 0..self.sn_start.len() - 1 {
291            let s1 = self.sn_start[s];
292            let s2 = self.sn_start[s + 1];
293            let s_width = s2 - s1;
294            let s_rows_start = self.sn_rows_ptr[s];
295            let s_m = self.sn_rows_ptr[s + 1] - s_rows_start;
296            let mut c1 = s1;
297            while c1 < s2 {
298                let c2 = s2.min(c1 + CHUNK);
299                let chunk = c2 - c1;
300                (c1..c2).zip(pointers.iter_mut()).for_each(|(j, pointer)| {
301                    let pinv = &self.pinv;
302                    let column = &mut work[(j - c1) * n..(j - c1 + 1) * n];
303                    matrix
304                        .column(self.q[j])
305                        .for_each(|(i, value)| column[pinv[i]] = *value);
306                    *pointer = self.u_col_ptr[j];
307                });
308                loop {
309                    let mut t = NONE;
310                    (c1..c2).zip(pointers.iter()).for_each(|(j, &pointer)| {
311                        if pointer < self.u_col_ptr[j + 1] - 1 {
312                            let row = self.u_row_idx[pointer];
313                            if row < c1 && (t == NONE || self.sn_of[row] < t) {
314                                t = self.sn_of[row];
315                            }
316                        }
317                    });
318                    if t == NONE {
319                        break;
320                    }
321                    let t1 = self.sn_start[t];
322                    let t2 = self.sn_start[t + 1];
323                    let width = t2 - t1;
324                    let consumed = width.min(c1 - t1);
325                    let rows = &self.sn_rows[self.sn_rows_ptr[t]..self.sn_rows_ptr[t + 1]];
326                    let m = rows.len();
327                    let below = m - width;
328                    let panel = &self.sn_values[self.sn_panel_ptr[t]..self.sn_panel_ptr[t + 1]];
329                    if consumed >= 4 {
330                        if chunk < CHUNK {
331                            tile[..width * CHUNK].fill(0.0);
332                        }
333                        (0..chunk).for_each(|b| {
334                            work[b * n + t1..b * n + t2]
335                                .iter()
336                                .zip(tile.chunks_exact_mut(CHUNK))
337                                .for_each(|(&value, row)| row[b] = value);
338                        });
339                        trisolve(&mut tile[..width * CHUNK], panel, m, consumed, width);
340                        (0..chunk).for_each(|b| {
341                            work[b * n + t1..b * n + t2]
342                                .iter_mut()
343                                .zip(tile.chunks_exact(CHUNK))
344                                .for_each(|(value, row)| *value = row[b]);
345                        });
346                        (c1..c2).zip(pointers.iter_mut()).for_each(|(j, pointer)| {
347                            let p_end = self.u_col_ptr[j + 1] - 1;
348                            let column = &work[(j - c1) * n..(j - c1 + 1) * n];
349                            while *pointer < p_end {
350                                let k = self.u_row_idx[*pointer];
351                                if k >= t1 + consumed {
352                                    break;
353                                }
354                                self.u_values[*pointer] = column[k];
355                                *pointer += 1;
356                            }
357                        });
358                    } else {
359                        (c1..c2).zip(pointers.iter_mut()).for_each(|(j, pointer)| {
360                            let p_end = self.u_col_ptr[j + 1] - 1;
361                            let column = &mut work[(j - c1) * n..(j - c1 + 1) * n];
362                            if consumed == width
363                                && width <= 3
364                                && *pointer + width <= p_end
365                                && self.u_row_idx[*pointer] == t1
366                                && self.u_row_idx[*pointer + width - 1] == t1 + width - 1
367                            {
368                                match width {
369                                    1 => self.u_values[*pointer] = column[t1],
370                                    2 => {
371                                        let u_0 = column[t1];
372                                        let u_1 = column[t1 + 1] - panel[1] * u_0;
373                                        column[t1 + 1] = u_1;
374                                        self.u_values[*pointer] = u_0;
375                                        self.u_values[*pointer + 1] = u_1;
376                                    }
377                                    _ => {
378                                        let u_0 = column[t1];
379                                        let u_1 = column[t1 + 1] - panel[1] * u_0;
380                                        let u_2 =
381                                            column[t1 + 2] - panel[2] * u_0 - panel[m + 2] * u_1;
382                                        column[t1 + 1] = u_1;
383                                        column[t1 + 2] = u_2;
384                                        self.u_values[*pointer] = u_0;
385                                        self.u_values[*pointer + 1] = u_1;
386                                        self.u_values[*pointer + 2] = u_2;
387                                    }
388                                }
389                                *pointer += width;
390                            } else {
391                                while *pointer < p_end {
392                                    let k = self.u_row_idx[*pointer];
393                                    if k >= c1 || k >= t2 {
394                                        break;
395                                    }
396                                    let c = k - t1;
397                                    let u = column[k];
398                                    self.u_values[*pointer] = u;
399                                    if u != 0.0 {
400                                        column[k + 1..t2]
401                                            .iter_mut()
402                                            .zip(panel[c * m + c + 1..c * m + width].iter())
403                                            .for_each(|(work_r, value)| *work_r -= value * u);
404                                    }
405                                    *pointer += 1;
406                                }
407                            }
408                        });
409                    }
410                    if below > 0 {
411                        temp[..CHUNK * below].fill(0.0);
412                        gemm_wide(
413                            &mut temp[..CHUNK * below],
414                            &work,
415                            n,
416                            panel,
417                            m,
418                            width,
419                            t1,
420                            consumed,
421                            below,
422                            chunk,
423                        );
424                        (0..chunk).for_each(|b| {
425                            let column = &mut work[b * n..(b + 1) * n];
426                            rows[width..]
427                                .iter()
428                                .zip(temp[b * below..(b + 1) * below].iter())
429                                .for_each(|(&row, value)| column[row] -= value);
430                        });
431                    }
432                    (0..chunk).for_each(|c| work[c * n + t1..c * n + t1 + consumed].fill(0.0));
433                }
434                for j in c1..c2 {
435                    let p_end = self.u_col_ptr[j + 1] - 1;
436                    let offset = (j - c1) * n;
437                    let mut pointer = pointers[j - c1];
438                    while pointer < p_end {
439                        let k = self.u_row_idx[pointer];
440                        let c = k - s1;
441                        let u = work[offset + k];
442                        self.u_values[pointer] = u;
443                        work[offset + k] = 0.0;
444                        if u != 0.0 {
445                            let start = self.sn_panel_ptr[s] + c * s_m;
446                            axpy(
447                                &mut work[offset + k + 1..offset + s2],
448                                &self.sn_values[start + c + 1..start + s_width],
449                                u,
450                            );
451                            self.sn_rows[s_rows_start + s_width..s_rows_start + s_m]
452                                .iter()
453                                .zip(self.sn_values[start + s_width..start + s_m].iter())
454                                .for_each(|(&row, value)| work[offset + row] -= value * u);
455                        }
456                        pointer += 1;
457                    }
458                    let pivot = work[offset + j];
459                    work[offset + j] = 0.0;
460                    if pivot.abs() < ABS_TOL {
461                        return Err(SparseError::Singular);
462                    }
463                    self.u_values[p_end] = pivot;
464                    let c = j - s1;
465                    let start = self.sn_panel_ptr[s] + c * s_m;
466                    self.sn_values[start + c] = 1.0;
467                    (c + 1..s_width).for_each(|local| {
468                        self.sn_values[start + local] = work[offset + s1 + local] / pivot;
469                        work[offset + s1 + local] = 0.0;
470                    });
471                    (s_width..s_m).for_each(|local| {
472                        let row = self.sn_rows[s_rows_start + local];
473                        self.sn_values[start + local] = work[offset + row] / pivot;
474                        work[offset + row] = 0.0;
475                    });
476                }
477                c1 = c2;
478            }
479        }
480        Ok(())
481    }
482    pub(super) fn max_below(&self) -> usize {
483        max_below(&self.sn_start, &self.sn_rows_ptr)
484    }
485}
486
487/// Eliminates the first `consumed` pivot columns of a panel from a transposed
488/// tile of CHUNK target columns, as a dense unit-lower triangular solve
489/// vectorized across the targets.
490fn trisolve(tile: &mut [Scalar], panel: &[Scalar], m: usize, consumed: usize, width: usize) {
491    #[cfg(target_arch = "x86_64")]
492    if super::simd() {
493        return unsafe { avx::trisolve(tile, panel, m, consumed, width) };
494    }
495    (0..consumed).for_each(|c| {
496        let (row, rest) = tile[c * CHUNK..].split_at_mut(CHUNK);
497        if row.iter().any(|&u| u != 0.0) {
498            rest[..(width - c - 1) * CHUNK]
499                .chunks_exact_mut(CHUNK)
500                .zip(panel[c * m + c + 1..c * m + width].iter())
501                .for_each(|(target, &value)| {
502                    target
503                        .iter_mut()
504                        .zip(row.iter())
505                        .for_each(|(entry, &u)| *entry -= value * u)
506                });
507        }
508    });
509}
510
511/// Nonzero pattern of the solution to Lx = b, as the topologically ordered reach
512/// of the pattern of b in the graph of L, placed in order[top..] with top returned.
513#[allow(clippy::too_many_arguments)]
514fn reach(
515    starts: impl Iterator<Item = usize>,
516    l_cols: &[Vec<(usize, Scalar)>],
517    pinv: &[usize],
518    tag: usize,
519    mark: &mut [usize],
520    order: &mut [usize],
521    stack: &mut [usize],
522    pstack: &mut [usize],
523) -> usize {
524    let mut top = order.len();
525    starts.for_each(|start| {
526        if mark[start] != tag {
527            mark[start] = tag;
528            stack[0] = start;
529            pstack[0] = 0;
530            let mut head = 0;
531            loop {
532                let i = stack[head];
533                let mut descended = false;
534                if pinv[i] != NONE {
535                    let column = &l_cols[pinv[i]];
536                    while pstack[head] < column.len() {
537                        let child = column[pstack[head]].0;
538                        pstack[head] += 1;
539                        if mark[child] != tag {
540                            mark[child] = tag;
541                            head += 1;
542                            stack[head] = child;
543                            pstack[head] = 0;
544                            descended = true;
545                            break;
546                        }
547                    }
548                }
549                if !descended {
550                    top -= 1;
551                    order[top] = i;
552                    if head == 0 {
553                        break;
554                    }
555                    head -= 1;
556                }
557            }
558        }
559    });
560    top
561}
562
563fn compress(
564    cols: Vec<Vec<(usize, Scalar)>>,
565    row_map: &[usize],
566) -> (Vec<usize>, Vec<usize>, Vec<Scalar>) {
567    let nnz = cols.iter().map(|col| col.len()).sum();
568    let mut col_ptr = Vec::with_capacity(cols.len() + 1);
569    let mut row_idx = Vec::with_capacity(nnz);
570    let mut values = Vec::with_capacity(nnz);
571    col_ptr.push(0);
572    cols.into_iter().for_each(|col| {
573        let mut mapped: Vec<(usize, Scalar)> = col
574            .into_iter()
575            .map(|(i, value)| (row_map[i], value))
576            .collect();
577        mapped.sort_unstable_by_key(|&(row, _)| row);
578        mapped.into_iter().for_each(|(row, value)| {
579            row_idx.push(row);
580            values.push(value);
581        });
582        col_ptr.push(row_idx.len());
583    });
584    (col_ptr, row_idx, values)
585}