use num_traits::{FromPrimitive, One};
use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LdltError {
Singular {
index: usize,
},
NotSquare {
nrows: usize,
ncols: usize,
},
DimensionMismatch {
expected: usize,
actual: usize,
},
}
impl core::fmt::Display for LdltError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
LdltError::Singular { index } => {
write!(f, "Matrix is singular at index {index}")
}
LdltError::NotSquare { nrows, ncols } => {
write!(f, "Matrix is not square: {nrows}×{ncols}")
}
LdltError::DimensionMismatch { expected, actual } => {
write!(f, "Dimension mismatch: expected {expected}, got {actual}")
}
}
}
}
impl std::error::Error for LdltError {}
#[derive(Clone, Debug)]
pub struct Ldlt<T: Scalar> {
ld: Mat<T>,
}
impl<T: Field + Real + bytemuck::Zeroable> Ldlt<T> {
pub fn compute(a: MatRef<'_, T>) -> Result<Self, LdltError> {
Self::compute_with_tol(a, None)
}
pub fn compute_with_tol(a: MatRef<'_, T>, tol: Option<T>) -> Result<Self, LdltError> {
let n = a.nrows();
if n != a.ncols() {
return Err(LdltError::NotSquare {
nrows: n,
ncols: a.ncols(),
});
}
if n == 0 {
return Ok(Ldlt {
ld: Mat::zeros(0, 0),
});
}
let mut ld = Mat::zeros(n, n);
for i in 0..n {
for j in 0..=i {
ld[(i, j)] = a[(i, j)];
}
}
let default_tol = <T as Scalar>::epsilon()
* <T as FromPrimitive>::from_usize(n).unwrap_or(<T as One>::one());
let tolerance = tol.unwrap_or(default_tol);
for j in 0..n {
let mut d_j = ld[(j, j)];
for k in 0..j {
let l_jk = ld[(j, k)];
let d_k = ld[(k, k)];
d_j = d_j - l_jk * l_jk * d_k;
}
if Scalar::abs(d_j) <= tolerance {
return Err(LdltError::Singular { index: j });
}
ld[(j, j)] = d_j;
for i in (j + 1)..n {
let mut l_ij = ld[(i, j)];
for k in 0..j {
l_ij = l_ij - ld[(i, k)] * ld[(k, k)] * ld[(j, k)];
}
ld[(i, j)] = l_ij / d_j;
}
}
Ok(Ldlt { ld })
}
#[inline]
pub fn size(&self) -> usize {
self.ld.nrows()
}
pub fn l_factor(&self) -> Mat<T> {
let n = self.size();
let mut l = Mat::zeros(n, n);
for i in 0..n {
l[(i, i)] = T::one(); for j in 0..i {
l[(i, j)] = self.ld[(i, j)];
}
}
l
}
pub fn d_diagonal(&self) -> Vec<T> {
let n = self.size();
(0..n).map(|i| self.ld[(i, i)]).collect()
}
pub fn d_factor(&self) -> Mat<T> {
let n = self.size();
let mut d = Mat::zeros(n, n);
for i in 0..n {
d[(i, i)] = self.ld[(i, i)];
}
d
}
pub fn determinant(&self) -> T {
let n = self.size();
if n == 0 {
return T::one();
}
let mut det = T::one();
for i in 0..n {
det = det * self.ld[(i, i)];
}
det
}
pub fn inertia(&self) -> (usize, usize, usize) {
let n = self.size();
let mut pos = 0;
let mut neg = 0;
let zero_tol = <T as Scalar>::epsilon();
for i in 0..n {
let d = self.ld[(i, i)];
if d > zero_tol {
pos += 1;
} else if d < -zero_tol {
neg += 1;
}
}
(pos, neg, n - pos - neg)
}
pub fn is_positive_definite(&self) -> bool {
let (pos, _, zero) = self.inertia();
pos == self.size() && zero == 0
}
pub fn is_negative_definite(&self) -> bool {
let (_, neg, zero) = self.inertia();
neg == self.size() && zero == 0
}
pub fn solve(&self, b: MatRef<'_, T>) -> Result<Mat<T>, LdltError> {
let n = self.size();
if b.nrows() != n {
return Err(LdltError::DimensionMismatch {
expected: n,
actual: b.nrows(),
});
}
let m = b.ncols();
let mut x = Mat::zeros(n, m);
for j in 0..m {
for i in 0..n {
x[(i, j)] = b[(i, j)];
}
}
for k in 0..n {
for i in (k + 1)..n {
let mult = self.ld[(i, k)];
for j in 0..m {
let val = x[(i, j)] - mult * x[(k, j)];
x[(i, j)] = val;
}
}
}
for k in 0..n {
let d_inv = T::one() / self.ld[(k, k)];
for j in 0..m {
x[(k, j)] = x[(k, j)] * d_inv;
}
}
for k in (0..n).rev() {
for i in 0..k {
let mult = self.ld[(k, i)];
for j in 0..m {
let val = x[(i, j)] - mult * x[(k, j)];
x[(i, j)] = val;
}
}
}
Ok(x)
}
pub fn inverse(&self) -> Result<Mat<T>, LdltError> {
let n = self.size();
let identity = Mat::<T>::eye(n);
self.solve(identity.as_ref())
}
pub fn log_abs_determinant(&self) -> (T, i32) {
let n = self.size();
if n == 0 {
return (T::zero(), 1);
}
let mut log_det = T::zero();
let mut sign = 1i32;
for i in 0..n {
let d = self.ld[(i, i)];
log_det = log_det + Real::ln(Scalar::abs(d));
if d < T::zero() {
sign *= -1;
}
}
(log_det, sign)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_ldlt_positive_definite() {
let a: Mat<f64> = Mat::from_rows(&[&[4.0, 2.0], &[2.0, 5.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let l = ldlt.l_factor();
let d = ldlt.d_diagonal();
assert!((l[(0, 0)] - 1.0).abs() < 1e-10);
assert!((l[(1, 1)] - 1.0).abs() < 1e-10);
assert!(l[(0, 1)].abs() < 1e-10);
assert!(d[0] > 0.0);
assert!(d[1] > 0.0);
let d_mat = ldlt.d_factor();
let n = a.nrows();
let mut ldlt_prod: Mat<f64> = Mat::zeros(n, n);
let mut ld: Mat<f64> = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
for k in 0..n {
ld[(i, j)] = ld[(i, j)] + l[(i, k)] * d_mat[(k, j)];
}
}
}
for i in 0..n {
for j in 0..n {
for k in 0..n {
ldlt_prod[(i, j)] = ldlt_prod[(i, j)] + ld[(i, k)] * l[(j, k)];
}
}
}
for i in 0..n {
for j in 0..n {
let diff = ldlt_prod[(i, j)] - a[(i, j)];
assert!(
diff.abs() < 1e-10,
"LDL^T[{},{}] = {}, A[{},{}] = {}",
i,
j,
ldlt_prod[(i, j)],
i,
j,
a[(i, j)]
);
}
}
}
#[test]
fn test_ldlt_indefinite() {
let a: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[2.0, 1.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let (pos, neg, zero) = ldlt.inertia();
assert_eq!(pos, 1);
assert_eq!(neg, 1);
assert_eq!(zero, 0);
assert!(!ldlt.is_positive_definite());
assert!(!ldlt.is_negative_definite());
}
#[test]
fn test_ldlt_solve() {
let a: Mat<f64> = Mat::from_rows(&[&[4.0, 2.0], &[2.0, 5.0]]);
let b: Mat<f64> = Mat::from_rows(&[&[8.0], &[11.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let x = ldlt.solve(b.as_ref()).expect("Should solve");
let ax0 = 4.0 * x[(0, 0)] + 2.0 * x[(1, 0)];
let ax1 = 2.0 * x[(0, 0)] + 5.0 * x[(1, 0)];
assert!((ax0 - 8.0).abs() < 1e-10, "Ax[0] = {}", ax0);
assert!((ax1 - 11.0).abs() < 1e-10, "Ax[1] = {}", ax1);
}
#[test]
fn test_ldlt_determinant() {
let a: Mat<f64> = Mat::from_rows(&[&[4.0, 2.0], &[2.0, 5.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let det = ldlt.determinant();
assert!((det - 16.0).abs() < 1e-10, "det = {}", det);
}
#[test]
fn test_ldlt_inverse() {
let a: Mat<f64> = Mat::from_rows(&[&[4.0, 2.0], &[2.0, 5.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let a_inv = ldlt.inverse().expect("Should invert");
assert!((a_inv[(0, 0)] - 0.3125).abs() < 1e-10);
assert!((a_inv[(0, 1)] + 0.125).abs() < 1e-10);
assert!((a_inv[(1, 0)] + 0.125).abs() < 1e-10);
assert!((a_inv[(1, 1)] - 0.25).abs() < 1e-10);
}
#[test]
fn test_ldlt_3x3() {
let a: Mat<f64> = Mat::from_rows(&[
&[4.0, 12.0, -16.0],
&[12.0, 37.0, -43.0],
&[-16.0, -43.0, 98.0],
]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let b: Mat<f64> = Mat::from_rows(&[&[1.0], &[2.0], &[3.0]]);
let x = ldlt.solve(b.as_ref()).expect("Should solve");
for i in 0..3 {
let mut ax_i = 0.0;
for j in 0..3 {
ax_i += a[(i, j)] * x[(j, 0)];
}
assert!((ax_i - b[(i, 0)]).abs() < 1e-10, "Ax[{}] = {}", i, ax_i);
}
}
#[test]
fn test_ldlt_identity() {
let eye: Mat<f64> = Mat::eye(3);
let ldlt = Ldlt::compute(eye.as_ref()).expect("Should decompose");
let d = ldlt.d_diagonal();
for di in d {
assert!((di - 1.0).abs() < 1e-10);
}
let det = ldlt.determinant();
assert!((det - 1.0).abs() < 1e-10);
}
#[test]
fn test_ldlt_singular() {
let a: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[2.0, 4.0]]);
let result = Ldlt::compute(a.as_ref());
assert!(result.is_err());
}
#[test]
fn test_ldlt_empty() {
let a: Mat<f64> = Mat::zeros(0, 0);
let ldlt = Ldlt::compute(a.as_ref()).expect("Empty should succeed");
assert_eq!(ldlt.size(), 0);
assert_eq!(ldlt.determinant(), 1.0);
}
#[test]
fn test_ldlt_not_square() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let result = Ldlt::compute(a.as_ref());
assert!(matches!(result, Err(LdltError::NotSquare { .. })));
}
#[test]
fn test_ldlt_f32() {
let a: Mat<f32> = Mat::from_rows(&[&[4.0f32, 2.0], &[2.0, 5.0]]);
let b: Mat<f32> = Mat::from_rows(&[&[8.0f32], &[11.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let x = ldlt.solve(b.as_ref()).expect("Should solve");
let ax0 = 4.0 * x[(0, 0)] + 2.0 * x[(1, 0)];
let ax1 = 2.0 * x[(0, 0)] + 5.0 * x[(1, 0)];
assert!((ax0 - 8.0).abs() < 1e-5, "Ax[0] = {}", ax0);
assert!((ax1 - 11.0).abs() < 1e-5, "Ax[1] = {}", ax1);
}
#[test]
fn test_ldlt_log_abs_determinant() {
let a: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[2.0, 1.0]]);
let ldlt = Ldlt::compute(a.as_ref()).expect("Should decompose");
let (log_det, sign) = ldlt.log_abs_determinant();
assert!((log_det - 3.0f64.ln()).abs() < 1e-10);
assert_eq!(sign, -1);
}
}