resopt 0.3.0

Declarative constrained residual optimization in Rust
Documentation
use crate::core::{Error, Matrix, Vector};

/// Tikhonov regularization term
/// `lambda / 2 * ||Lx - x_ref||_2^2`.
#[derive(Debug, Clone, PartialEq)]
pub struct TikhonovRegularization {
    lambda: f64,
    matrix: Matrix,
    target: Vector,
}

impl TikhonovRegularization {
    pub fn new(lambda: f64, matrix: Matrix, target: Vector) -> Result<Self, Error> {
        if lambda <= 0.0 {
            return Err(Error::InvalidParameter {
                message: "Tikhonov lambda must be strictly positive".to_string(),
            });
        }

        if matrix.nrows() != target.len() {
            return Err(Error::DimensionMismatch {
                message: format!(
                    "regularization matrix row count ({}) must match target length ({})",
                    matrix.nrows(),
                    target.len()
                ),
            });
        }

        Ok(Self {
            lambda,
            matrix,
            target,
        })
    }

    pub fn ridge(x_dim: usize, lambda: f64) -> Result<Self, Error> {
        let mut data = vec![0.0; x_dim * x_dim];

        for i in 0..x_dim {
            data[i * x_dim + i] = 1.0;
        }

        Self::new(
            lambda,
            Matrix::from_row_major(x_dim, x_dim, data)?,
            vec![0.0; x_dim],
        )
    }

    pub fn lambda(&self) -> f64 {
        self.lambda
    }

    pub fn matrix(&self) -> &Matrix {
        &self.matrix
    }

    pub fn target(&self) -> &[f64] {
        &self.target
    }

    pub fn rows(&self) -> usize {
        self.matrix.nrows()
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn valid_general_regularization() {
        let reg = TikhonovRegularization::new(
            0.5,
            Matrix::from_row_major(2, 3, vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0]).unwrap(),
            vec![1.0, -1.0],
        )
        .unwrap();

        assert!((reg.lambda() - 0.5).abs() < f64::EPSILON);
        assert_eq!(reg.rows(), 2);
        assert_eq!(reg.target(), &[1.0, -1.0]);
    }

    #[test]
    fn ridge_regularization_is_identity() {
        let reg = TikhonovRegularization::ridge(3, 2.0).unwrap();

        assert_eq!(reg.rows(), 3);
        assert_eq!(reg.matrix().ncols(), 3);
        assert_eq!(reg.target(), &[0.0, 0.0, 0.0]);
        assert_eq!(
            reg.matrix().data(),
            &[1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]
        );
    }

    #[test]
    fn lambda_must_be_positive() {
        assert!(TikhonovRegularization::ridge(2, 0.0).is_err());
        assert!(TikhonovRegularization::ridge(2, -1.0).is_err());
    }

    #[test]
    fn target_dimension_must_match_matrix_rows() {
        let reg = TikhonovRegularization::new(
            1.0,
            Matrix::from_row_major(2, 2, vec![1.0, 0.0, 0.0, 1.0]).unwrap(),
            vec![0.0],
        );

        assert!(reg.is_err());
    }
}