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
11pub 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 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 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 pub fn nonzeros(&self) -> usize {
278 self.fill
279 }
280 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#[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}