resopt 0.3.0

Declarative constrained residual optimization in Rust
Documentation
use crate::{
    core::{
        Bounds, Error, LinearEqualities, LinearInequalities, LinearResidual, Loss, ProblemClass,
        ProblemSummary, TikhonovRegularization,
    },
    solve::{DefaultSolver, SolveResult, Solver},
};

/// Declarative constrained residual minimization problem.
///
/// This type always represents a minimization problem of the form
///
/// minimize_x   loss(Ax - b)
///            + lambda / 2 * ||Lx - x_ref||_2^2   (optional)
/// subject to   linear equalities
///              linear inequalities
///              variable bounds
#[derive(Debug, Clone, PartialEq)]
pub struct ConstrainedResidualProblem {
    residual: LinearResidual,
    loss: Loss,
    equalities: Vec<LinearEqualities>,
    inequalities: Vec<LinearInequalities>,
    bounds: Option<Bounds>,
    regularization: Option<TikhonovRegularization>,
}

impl ConstrainedResidualProblem {
    pub fn new(residual: LinearResidual, loss: Loss) -> Result<Self, Error> {
        loss.validate()?;

        let problem = Self {
            residual,
            loss,
            equalities: Vec::new(),
            inequalities: Vec::new(),
            bounds: None,
            regularization: None,
        };

        problem.validate()?;
        Ok(problem)
    }

    pub fn residual(&self) -> &LinearResidual {
        &self.residual
    }

    pub fn loss(&self) -> &Loss {
        &self.loss
    }

    pub fn equalities(&self) -> &[LinearEqualities] {
        &self.equalities
    }

    pub fn inequalities(&self) -> &[LinearInequalities] {
        &self.inequalities
    }

    pub fn bounds(&self) -> Option<&Bounds> {
        self.bounds.as_ref()
    }

    pub fn regularization(&self) -> Option<&TikhonovRegularization> {
        self.regularization.as_ref()
    }

    pub fn x_dim(&self) -> usize {
        self.residual.x_dim()
    }

    pub fn residual_dim(&self) -> usize {
        self.residual.residual_dim()
    }

    pub fn add_equalities(mut self, block: LinearEqualities) -> Result<Self, Error> {
        self.check_constraint_cols(block.matrix().ncols(), "equality")?;
        self.equalities.push(block);
        Ok(self)
    }

    pub fn add_inequalities(mut self, block: LinearInequalities) -> Result<Self, Error> {
        self.check_constraint_cols(block.matrix().ncols(), "inequality")?;
        self.inequalities.push(block);
        Ok(self)
    }

    pub fn with_bounds(mut self, bounds: Bounds) -> Result<Self, Error> {
        if bounds.len() != self.x_dim() {
            return Err(Error::DimensionMismatch {
                message: format!(
                    "bounds dimension ({}) must match x dimension ({})",
                    bounds.len(),
                    self.x_dim()
                ),
            });
        }

        self.bounds = Some(bounds);
        Ok(self)
    }

    pub fn with_regularization(
        mut self,
        regularization: TikhonovRegularization,
    ) -> Result<Self, Error> {
        self.check_constraint_cols(regularization.matrix().ncols(), "regularization")?;
        self.regularization = Some(regularization);
        Ok(self)
    }

    pub fn validate(&self) -> Result<(), Error> {
        self.loss.validate()?;
        let x_dim = self.x_dim();

        for eq in &self.equalities {
            if eq.matrix().ncols() != x_dim {
                return Err(Error::DimensionMismatch {
                    message: format!(
                        "equality matrix column count ({}) must match x dimension ({})",
                        eq.matrix().ncols(),
                        x_dim
                    ),
                });
            }
        }

        for ineq in &self.inequalities {
            if ineq.matrix().ncols() != x_dim {
                return Err(Error::DimensionMismatch {
                    message: format!(
                        "inequality matrix column count ({}) must match x dimension ({})",
                        ineq.matrix().ncols(),
                        x_dim
                    ),
                });
            }
        }

        if let Some(bounds) = &self.bounds {
            if bounds.len() != x_dim {
                return Err(Error::DimensionMismatch {
                    message: format!(
                        "bounds dimension ({}) must match x dimension ({})",
                        bounds.len(),
                        x_dim
                    ),
                });
            }
        }

        if let Some(regularization) = &self.regularization {
            if regularization.matrix().ncols() != x_dim {
                return Err(Error::DimensionMismatch {
                    message: format!(
                        "regularization matrix column count ({}) must match x dimension ({})",
                        regularization.matrix().ncols(),
                        x_dim
                    ),
                });
            }
        }

        Ok(())
    }

    pub fn summary(&self) -> ProblemSummary {
        let equality_rows = self.equalities.iter().map(|b| b.rows()).sum();
        let inequality_rows = self.inequalities.iter().map(|b| b.rows()).sum();

        ProblemSummary {
            x_dim: self.x_dim(),
            residual_dim: self.residual_dim(),
            regularization_dim: self.regularization.as_ref().map_or(0, |reg| reg.rows()),
            equality_blocks: self.equalities.len(),
            equality_rows,
            inequality_blocks: self.inequalities.len(),
            inequality_rows,
            has_bounds: self.bounds.is_some(),
            has_regularization: self.regularization.is_some(),
            loss: self.loss.clone(),
            class: self.class(),
        }
    }

    pub fn class(&self) -> ProblemClass {
        let has_eq = !self.equalities.is_empty();
        let has_ineq = !self.inequalities.is_empty();
        let has_bounds = self.bounds.is_some();

        match (has_eq, has_ineq, has_bounds) {
            (false, false, false) => ProblemClass::Unconstrained,
            (true, false, false) => ProblemClass::EqualityConstrained,
            (false, true, false) => ProblemClass::InequalityConstrained,
            (false, false, true) => ProblemClass::Bounded,
            _ => ProblemClass::MixedConstrained,
        }
    }

    /// Simple path: delegates to the default solver.
    pub fn solve(&self) -> Result<SolveResult, Error> {
        DefaultSolver::new().solve(self)
    }

    fn check_constraint_cols(&self, cols: usize, kind: &str) -> Result<(), Error> {
        if cols != self.x_dim() {
            return Err(Error::DimensionMismatch {
                message: format!(
                    "{} matrix column count ({}) must match x dimension ({})",
                    kind,
                    cols,
                    self.x_dim()
                ),
            });
        }

        Ok(())
    }
}

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

    fn make_problem_2x2() -> ConstrainedResidualProblem {
        let m = Matrix::from_row_major(2, 2, vec![1.0, 0.0, 0.0, 1.0]).unwrap();
        let r = LinearResidual::new(m, vec![1.0, 2.0]).unwrap();
        ConstrainedResidualProblem::new(r, Loss::L2Squared).unwrap()
    }

    #[test]
    fn new_unconstrained() {
        let p = make_problem_2x2();
        assert_eq!(p.x_dim(), 2);
        assert_eq!(p.residual_dim(), 2);
        assert_eq!(p.class(), ProblemClass::Unconstrained);
        assert!(p.equalities().is_empty());
        assert!(p.inequalities().is_empty());
        assert!(p.bounds().is_none());
        assert!(p.regularization().is_none());
    }

    #[test]
    fn add_equalities_ok() {
        let eq_m = Matrix::from_row_major(1, 2, vec![1.0, 1.0]).unwrap();
        let eq = LinearEqualities::new(eq_m, vec![1.0]).unwrap();
        let p = make_problem_2x2().add_equalities(eq).unwrap();
        assert_eq!(p.equalities().len(), 1);
        assert_eq!(p.class(), ProblemClass::EqualityConstrained);
    }

    #[test]
    fn add_equalities_wrong_dim() {
        let eq_m = Matrix::from_row_major(1, 5, vec![1.0; 5]).unwrap();
        let eq = LinearEqualities::new(eq_m, vec![1.0]).unwrap();
        assert!(make_problem_2x2().add_equalities(eq).is_err());
    }

    #[test]
    fn add_inequalities_ok() {
        let ineq_m = Matrix::from_row_major(1, 2, vec![1.0, -1.0]).unwrap();
        let ineq = LinearInequalities::new(ineq_m, vec![0.0]).unwrap();
        let p = make_problem_2x2().add_inequalities(ineq).unwrap();
        assert_eq!(p.class(), ProblemClass::InequalityConstrained);
    }

    #[test]
    fn with_bounds_ok() {
        let p = make_problem_2x2().with_bounds(Bounds::free(2)).unwrap();
        assert_eq!(p.class(), ProblemClass::Bounded);
    }

    #[test]
    fn with_bounds_wrong_dim() {
        assert!(make_problem_2x2().with_bounds(Bounds::free(10)).is_err());
    }

    #[test]
    fn with_regularization_ok() {
        let regularization = TikhonovRegularization::ridge(2, 1.0).unwrap();
        let problem = make_problem_2x2()
            .with_regularization(regularization)
            .unwrap();
        let summary = problem.summary();

        assert!(problem.regularization().is_some());
        assert!(summary.has_regularization);
        assert_eq!(summary.regularization_dim, 2);
    }

    #[test]
    fn with_regularization_wrong_dim() {
        let regularization = TikhonovRegularization::ridge(10, 1.0).unwrap();
        assert!(make_problem_2x2()
            .with_regularization(regularization)
            .is_err());
    }

    #[test]
    fn summary_unconstrained() {
        let s = make_problem_2x2().summary();
        assert_eq!(s.x_dim, 2);
        assert_eq!(s.residual_dim, 2);
        assert_eq!(s.regularization_dim, 0);
        assert_eq!(s.equality_blocks, 0);
        assert_eq!(s.inequality_blocks, 0);
        assert!(!s.has_bounds);
        assert!(!s.has_regularization);
        assert_eq!(s.class, ProblemClass::Unconstrained);
    }

    #[test]
    fn validate_ok() {
        assert!(make_problem_2x2().validate().is_ok());
    }

    #[test]
    fn clone_and_eq() {
        let p = make_problem_2x2();
        assert_eq!(p, p.clone());
    }
}