use crate::core::{Error, Matrix, Vector};
#[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());
}
}