use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use crate::qr::{Qr, QrError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LstSqError {
EmptyMatrix,
DimensionMismatch,
Underdetermined,
QrFailed,
}
impl core::fmt::Display for LstSqError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::DimensionMismatch => write!(f, "Matrix and vector dimensions do not match"),
Self::Underdetermined => {
write!(f, "System is underdetermined (more columns than rows)")
}
Self::QrFailed => write!(f, "QR decomposition failed"),
}
}
}
impl std::error::Error for LstSqError {}
impl From<QrError> for LstSqError {
fn from(_: QrError) -> Self {
Self::QrFailed
}
}
#[derive(Debug, Clone)]
pub struct LeastSquaresResult<T: Scalar> {
pub solution: Mat<T>,
pub residual: Mat<T>,
pub residual_norm_sq: T,
pub rank: usize,
}
pub fn lstsq<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
) -> Result<LeastSquaresResult<T>, LstSqError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(LstSqError::EmptyMatrix);
}
if b.nrows() != m {
return Err(LstSqError::DimensionMismatch);
}
if m < n {
return Err(LstSqError::Underdetermined);
}
let eps = <T as Scalar>::epsilon();
let qr = Qr::compute(a)?;
let q = qr.q();
let r = qr.r();
let mut qt_b = Mat::zeros(m, b.ncols());
for i in 0..m {
for j in 0..b.ncols() {
let mut sum = T::zero();
for k in 0..m {
sum = sum + q[(k, i)] * b[(k, j)];
}
qt_b[(i, j)] = sum;
}
}
let mut rank = n;
let r_max = if n > 0 {
Scalar::abs(r[(0, 0)])
} else {
T::one()
};
let tol = eps * T::from_f64(m.max(n) as f64).unwrap_or_else(T::zero) * r_max;
for i in 0..n {
if Scalar::abs(r[(i, i)]) < tol {
rank = i;
break;
}
}
let mut solution = Mat::zeros(n, b.ncols());
for col in 0..b.ncols() {
for i in (0..rank).rev() {
let mut sum = qt_b[(i, col)];
for j in (i + 1)..rank {
sum = sum - r[(i, j)] * solution[(j, col)];
}
if Scalar::abs(r[(i, i)]) > eps {
solution[(i, col)] = sum / r[(i, i)];
}
}
}
let mut residual = Mat::zeros(m, b.ncols());
let mut residual_norm_sq = T::zero();
for col in 0..b.ncols() {
for i in 0..m {
let mut ax_i = T::zero();
for j in 0..n {
ax_i = ax_i + a[(i, j)] * solution[(j, col)];
}
residual[(i, col)] = b[(i, col)] - ax_i;
residual_norm_sq = residual_norm_sq + residual[(i, col)] * residual[(i, col)];
}
}
Ok(LeastSquaresResult {
solution,
residual,
residual_norm_sq,
rank,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_lstsq_exact() {
let a = Mat::from_rows(&[&[1.0f64, 1.0], &[1.0, 2.0]]);
let b = Mat::from_rows(&[&[3.0], &[5.0]]);
let result = lstsq(a.as_ref(), b.as_ref()).unwrap();
let x = &result.solution;
assert!(approx_eq(x[(0, 0)], 1.0, 1e-10));
assert!(approx_eq(x[(1, 0)], 2.0, 1e-10));
assert!(result.residual_norm_sq < 1e-10);
}
#[test]
fn test_lstsq_overdetermined() {
let a = Mat::from_rows(&[&[1.0f64, 1.0], &[1.0, 2.0], &[1.0, 3.0]]);
let b = Mat::from_rows(&[&[6.0], &[8.0], &[10.0]]);
let result = lstsq(a.as_ref(), b.as_ref()).unwrap();
let x = &result.solution;
assert!(approx_eq(x[(0, 0)], 4.0, 1e-10));
assert!(approx_eq(x[(1, 0)], 2.0, 1e-10));
assert!(result.residual_norm_sq < 1e-10);
}
#[test]
fn test_lstsq_with_residual() {
let a = Mat::from_rows(&[&[1.0f64, 1.0], &[1.0, 2.0], &[1.0, 3.0]]);
let b = Mat::from_rows(&[&[6.0], &[8.0], &[11.0]]);
let result = lstsq(a.as_ref(), b.as_ref()).unwrap();
assert!(result.residual_norm_sq > 0.0);
let mut manual_residual_sq = 0.0;
for i in 0..3 {
let ax_i = a[(i, 0)] * result.solution[(0, 0)] + a[(i, 1)] * result.solution[(1, 0)];
let r_i = b[(i, 0)] - ax_i;
manual_residual_sq += r_i * r_i;
}
assert!(approx_eq(
result.residual_norm_sq,
manual_residual_sq,
1e-10
));
}
#[test]
fn test_lstsq_identity() {
let a = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0], &[0.0, 0.0]]);
let b = Mat::from_rows(&[&[3.0], &[4.0], &[0.0]]);
let result = lstsq(a.as_ref(), b.as_ref()).unwrap();
assert!(approx_eq(result.solution[(0, 0)], 3.0, 1e-10));
assert!(approx_eq(result.solution[(1, 0)], 4.0, 1e-10));
}
#[test]
fn test_lstsq_underdetermined() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0]]);
let b = Mat::from_rows(&[&[6.0]]);
let result = lstsq(a.as_ref(), b.as_ref());
assert!(matches!(result, Err(LstSqError::Underdetermined)));
}
#[test]
fn test_lstsq_dimension_mismatch() {
let a = Mat::from_rows(&[&[1.0f64, 1.0], &[1.0, 2.0]]);
let b = Mat::from_rows(&[&[1.0], &[2.0], &[3.0]]);
let result = lstsq(a.as_ref(), b.as_ref());
assert!(matches!(result, Err(LstSqError::DimensionMismatch)));
}
#[test]
fn test_lstsq_f32() {
let a = Mat::from_rows(&[&[1.0f32, 1.0], &[1.0, 2.0], &[1.0, 3.0]]);
let b = Mat::from_rows(&[&[6.0f32], &[8.0], &[10.0]]);
let result = lstsq(a.as_ref(), b.as_ref()).unwrap();
assert!((result.solution[(0, 0)] - 4.0).abs() < 1e-4);
assert!((result.solution[(1, 0)] - 2.0).abs() < 1e-4);
}
#[test]
fn test_lstsq_multiple_rhs() {
let a = Mat::from_rows(&[&[1.0f64, 1.0], &[1.0, 2.0], &[1.0, 3.0]]);
let b = Mat::from_rows(&[&[6.0, 3.0], &[8.0, 5.0], &[10.0, 7.0]]);
let result = lstsq(a.as_ref(), b.as_ref()).unwrap();
assert!(approx_eq(result.solution[(0, 0)], 4.0, 1e-10));
assert!(approx_eq(result.solution[(1, 0)], 2.0, 1e-10));
assert!(approx_eq(result.solution[(0, 1)], 1.0, 1e-10));
assert!(approx_eq(result.solution[(1, 1)], 2.0, 1e-10));
}
}