Skip to main content

rusolver/
cholesky.rs

1// SPDX-License-Identifier: Apache-2.0
2use 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    /// Full input symmetry is checked. Accepted roundoff differences are not
11    /// averaged: the lower triangle defines the effective symmetric matrix.
12    pub symmetry_tolerance: Tolerance,
13    /// Reject pivots <= max(abs, rel*max_abs(A)). Default only rejects <= zero.
14    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    /// `A = L L^T`; non-positive-definite inputs return an error, never a hidden jitter.
29    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}