1use crate::{
3 Matrix, MatrixView, SolverError, Tolerance
4};
5use crate::numerics::{
6 checked, sum_iter, symmetric
7};
8#[derive(Clone, Copy, Debug)]
9pub struct CholeskyOptions {
10 pub symmetry_tolerance: Tolerance,
13 pub pivot_tolerance: Tolerance,
15}
16impl Default for CholeskyOptions {
17 fn default()->Self {
18 Self {
19 symmetry_tolerance: Tolerance::default(), pivot_tolerance: Tolerance::EXACT
20 }
21 }
22}
23#[derive(Clone, Debug)]
24pub struct Cholesky {
25 lower: Matrix
26}
27impl Cholesky {
28 pub fn factor(a: MatrixView<'_>, options: CholeskyOptions)->Result<Self, SolverError> {
30 let n=a.square()?;
31 symmetric(a, options.symmetry_tolerance)?;
32 let cutoff=options.pivot_tolerance.threshold(a.max_abs())?;
33 let mut l=Matrix::zeros(n, n)?;
34 for j in 0..n {
35 let sum=sum_iter((0..j).map(|k| l.data[j*n+k]*l.data[j*n+k]))?;
36 let pivot=checked(a.at(j, j)-sum, "Cholesky diagonal")?;
37 if pivot<=cutoff {
38 return Err(SolverError::NotPositiveDefinite{
39 index: j, pivot
40 });
41 }
42 let diagonal=pivot.sqrt();
43 l.data[j*n+j]=diagonal;
44 for i in j+1..n {
45 let sum=sum_iter((0..j).map(|k| l.data[i*n+k]*l.data[j*n+k]))?;
46 l.data[i*n+j]=checked((a.at(i, j)-sum)/diagonal, "Cholesky column")?;
47 }
48 }
49 Ok(Self{
50 lower: l
51 })
52 }
53 pub fn lower(&self)->&Matrix {
54 &self.lower
55 }
56 pub fn order(&self)->usize {
57 self.lower.rows
58 }
59 pub fn solve(&self, b: MatrixView<'_>)->Result<Matrix, SolverError> {
60 let n=self.order();
61 let m=b.columns();
62 if b.rows()!=n {
63 return Err(SolverError::Shape("Cholesky RHS row count"));
64 }
65 let mut x=b.to_owned()?;
66 for i in 0..n {
67 for c in 0..m {
68 let sum=sum_iter((0..i).map(|k|self.lower.data[i*n+k]*x.data[k*m+c]))?;
69 x.data[i*m+c]=checked((x.data[i*m+c]-sum)/self.lower.data[i*n+i], "Cholesky forward solve")?;
70 }
71 }
72 for i in (0..n).rev() {
73 for c in 0..m {
74 let sum=sum_iter((i+1..n).map(|k|self.lower.data[k*n+i]*x.data[k*m+c]))?;
75 x.data[i*m+c]=checked((x.data[i*m+c]-sum)/self.lower.data[i*n+i], "Cholesky back solve")?;
76 }
77 }
78 Ok(x)
79 }
80 pub fn log_determinant(&self)->Result<f64, SolverError> {
81 checked(2.0*sum_iter((0..self.order()).map(|i|self.lower.data[i*self.order()+i].ln()))?, "Cholesky log determinant")
82 }
83}