use crate::{
core::{
Bounds, Error, LinearEqualities, LinearInequalities, LinearResidual, Loss, ProblemClass,
ProblemSummary, TikhonovRegularization,
},
solve::{DefaultSolver, SolveResult, Solver},
};
#[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,
}
}
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());
}
}