Skip to main content

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

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