use crate::cost::CostFunction;
use crate::cost::CostFunctionType;
use crate::error::{NllsProblemError, ParameterBlockStorageError, ResidualBlockBuildingError};
use crate::loss::LossFunction;
use crate::parameter_block::{ParameterBlockOrIndex, ParameterBlockStorage};
use crate::residual_block::{ResidualBlock, ResidualBlockId};
use crate::solver::{SolverOptions, SolverSummary};
use ceres_solver_sys::cxx::UniquePtr;
use ceres_solver_sys::ffi;
use std::pin::Pin;
pub struct NllsProblem<'cost> {
inner: UniquePtr<ffi::Problem<'cost>>,
parameter_storage: ParameterBlockStorage,
residual_blocks: Vec<ResidualBlock>,
}
impl<'cost> NllsProblem<'cost> {
pub fn new() -> Self {
Self {
inner: ffi::new_problem(),
parameter_storage: ParameterBlockStorage::new(),
residual_blocks: Vec::new(),
}
}
pub fn residual_block_builder(self) -> ResidualBlockBuilder<'cost> {
ResidualBlockBuilder {
problem: self,
cost: None,
loss: None,
parameters: Vec::new(),
}
}
#[inline]
fn inner(&self) -> &ffi::Problem<'cost> {
self.inner
.as_ref()
.expect("Underlying C++ unique_ptr<Problem> must hold non-null pointer")
}
#[inline]
fn inner_mut(&mut self) -> Pin<&mut ffi::Problem<'cost>> {
self.inner
.as_mut()
.expect("Underlying C++ unique_ptr<Problem> must hold non-null pointer")
}
pub fn set_parameter_block_constant(
&mut self,
block_index: usize,
) -> Result<(), ParameterBlockStorageError> {
let block_pointer = self.parameter_storage.get_block(block_index)?.pointer_mut();
unsafe {
self.inner_mut().SetParameterBlockConstant(block_pointer);
}
Ok(())
}
pub fn set_parameter_block_variable(
&mut self,
block_index: usize,
) -> Result<(), ParameterBlockStorageError> {
let block_pointer = self.parameter_storage.get_block(block_index)?.pointer_mut();
unsafe {
self.inner_mut().SetParameterBlockVariable(block_pointer);
}
Ok(())
}
pub fn is_parameter_block_constant(
&self,
block_index: usize,
) -> Result<bool, ParameterBlockStorageError> {
let block_pointer = self.parameter_storage.get_block(block_index)?.pointer_mut();
unsafe { Ok(self.inner().IsParameterBlockConstant(block_pointer)) }
}
pub fn solve(
mut self,
options: &SolverOptions,
) -> Result<NllsProblemSolution, NllsProblemError> {
if self.residual_blocks.is_empty() {
return Err(NllsProblemError::NoResidualBlocks);
}
let mut summary = SolverSummary::new();
ffi::solve(
options
.0
.as_ref()
.expect("Underlying C++ SolverOptions must hold non-null pointer"),
self.inner_mut(),
summary
.0
.as_mut()
.expect("Underlying C++ unique_ptr<SolverSummary> must hold non-null pointer"),
);
Ok(NllsProblemSolution {
parameters: self.parameter_storage.to_values(),
summary,
})
}
}
impl Default for NllsProblem<'_> {
fn default() -> Self {
Self::new()
}
}
pub struct NllsProblemSolution {
pub parameters: Vec<Vec<f64>>,
pub summary: SolverSummary,
}
pub struct ResidualBlockBuilder<'cost> {
problem: NllsProblem<'cost>,
cost: Option<(CostFunctionType<'cost>, usize)>,
loss: Option<LossFunction>,
parameters: Vec<ParameterBlockOrIndex>,
}
impl<'cost> ResidualBlockBuilder<'cost> {
pub fn set_cost(
mut self,
func: impl Into<CostFunctionType<'cost>>,
num_residuals: usize,
) -> Self {
self.cost = Some((func.into(), num_residuals));
self
}
pub fn set_loss(mut self, loss: LossFunction) -> Self {
self.loss = Some(loss);
self
}
pub fn set_parameters<P>(mut self, parameters: impl IntoIterator<Item = P>) -> Self
where
P: Into<ParameterBlockOrIndex>,
{
self.parameters = parameters.into_iter().map(|p| p.into()).collect();
self
}
pub fn add_parameter<P>(mut self, parameter_block: P) -> Self
where
P: Into<ParameterBlockOrIndex>,
{
self.parameters.push(parameter_block.into());
self
}
pub fn build_into_problem(
self,
) -> Result<(NllsProblem<'cost>, ResidualBlockId), ResidualBlockBuildingError> {
let Self {
mut problem,
cost,
loss,
parameters,
} = self;
if parameters.is_empty() {
return Err(ResidualBlockBuildingError::MissingParameters);
}
let parameter_indices = problem.parameter_storage.extend(parameters)?;
let parameter_sizes: Vec<_> = parameter_indices
.iter()
.map(|&index| problem.parameter_storage.blocks()[index].len())
.collect();
let parameter_pointers: Pin<Vec<_>> = Pin::new(
parameter_indices
.iter()
.map(|&index| problem.parameter_storage.blocks()[index].pointer_mut())
.collect(),
);
let cost = if let Some((func, num_redisuals)) = cost {
CostFunction::new(func, parameter_sizes, num_redisuals)
} else {
return Err(ResidualBlockBuildingError::MissingCost);
};
let residual_block_id = unsafe {
ffi::add_residual_block(
problem
.inner
.as_mut()
.expect("Underlying C++ unique_ptr<Problem> must hold non-null pointer"),
cost.into_inner(),
loss.map(|loss| loss.into_inner())
.unwrap_or_else(UniquePtr::null),
parameter_pointers.as_ptr(),
parameter_indices.len() as i32,
)
};
problem.residual_blocks.push(ResidualBlock {
id: residual_block_id.clone(),
parameter_pointers,
});
for &index in parameter_indices.iter() {
let block = &problem.parameter_storage.blocks()[index];
if let Some(lower_bound) = block.lower_bounds() {
for (i, lower_bound) in lower_bound.iter().enumerate() {
if let Some(lower_bound) = lower_bound {
unsafe {
problem
.inner
.as_mut()
.expect(
"Underlying C++ unique_ptr<Problem> must hold non-null pointer",
)
.SetParameterLowerBound(block.pointer_mut(), i as i32, *lower_bound)
}
}
}
}
}
for &index in parameter_indices.iter() {
let block = &problem.parameter_storage.blocks()[index];
if let Some(upper_bound) = block.upper_bounds() {
for (i, upper_bound) in upper_bound.iter().enumerate() {
if let Some(upper_bound) = upper_bound {
unsafe {
problem
.inner
.as_mut()
.expect(
"Underlying C++ unique_ptr<Problem> must hold non-null pointer",
)
.SetParameterUpperBound(block.pointer_mut(), i as i32, *upper_bound)
}
}
}
}
}
Ok((problem, residual_block_id))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cost::CostFunctionType;
use crate::loss::{LossFunction, LossFunctionType};
use approx::assert_abs_diff_eq;
fn simple_end_to_end_test_with_loss(loss: LossFunction) {
const NUM_OBSERVATIONS: usize = 67;
const NDIM: usize = 2;
let data: [[f64; NDIM]; NUM_OBSERVATIONS] = [
0.000000e+00,
1.133898e+00,
7.500000e-02,
1.334902e+00,
1.500000e-01,
1.213546e+00,
2.250000e-01,
1.252016e+00,
3.000000e-01,
1.392265e+00,
3.750000e-01,
1.314458e+00,
4.500000e-01,
1.472541e+00,
5.250000e-01,
1.536218e+00,
6.000000e-01,
1.355679e+00,
6.750000e-01,
1.463566e+00,
7.500000e-01,
1.490201e+00,
8.250000e-01,
1.658699e+00,
9.000000e-01,
1.067574e+00,
9.750000e-01,
1.464629e+00,
1.050000e+00,
1.402653e+00,
1.125000e+00,
1.713141e+00,
1.200000e+00,
1.527021e+00,
1.275000e+00,
1.702632e+00,
1.350000e+00,
1.423899e+00,
1.425000e+00,
1.543078e+00,
1.500000e+00,
1.664015e+00,
1.575000e+00,
1.732484e+00,
1.650000e+00,
1.543296e+00,
1.725000e+00,
1.959523e+00,
1.800000e+00,
1.685132e+00,
1.875000e+00,
1.951791e+00,
1.950000e+00,
2.095346e+00,
2.025000e+00,
2.361460e+00,
2.100000e+00,
2.169119e+00,
2.175000e+00,
2.061745e+00,
2.250000e+00,
2.178641e+00,
2.325000e+00,
2.104346e+00,
2.400000e+00,
2.584470e+00,
2.475000e+00,
1.914158e+00,
2.550000e+00,
2.368375e+00,
2.625000e+00,
2.686125e+00,
2.700000e+00,
2.712395e+00,
2.775000e+00,
2.499511e+00,
2.850000e+00,
2.558897e+00,
2.925000e+00,
2.309154e+00,
3.000000e+00,
2.869503e+00,
3.075000e+00,
3.116645e+00,
3.150000e+00,
3.094907e+00,
3.225000e+00,
2.471759e+00,
3.300000e+00,
3.017131e+00,
3.375000e+00,
3.232381e+00,
3.450000e+00,
2.944596e+00,
3.525000e+00,
3.385343e+00,
3.600000e+00,
3.199826e+00,
3.675000e+00,
3.423039e+00,
3.750000e+00,
3.621552e+00,
3.825000e+00,
3.559255e+00,
3.900000e+00,
3.530713e+00,
3.975000e+00,
3.561766e+00,
4.050000e+00,
3.544574e+00,
4.125000e+00,
3.867945e+00,
4.200000e+00,
4.049776e+00,
4.275000e+00,
3.885601e+00,
4.350000e+00,
4.110505e+00,
4.425000e+00,
4.345320e+00,
4.500000e+00,
4.161241e+00,
4.575000e+00,
4.363407e+00,
4.650000e+00,
4.161576e+00,
4.725000e+00,
4.619728e+00,
4.800000e+00,
4.737410e+00,
4.875000e+00,
4.727863e+00,
4.950000e+00,
4.669206e+00,
]
.chunks_exact(NDIM)
.map(|chunk| chunk.try_into().unwrap())
.collect::<Vec<_>>()
.try_into()
.unwrap();
let cost: CostFunctionType = Box::new(move |parameters, residuals, mut jacobians| {
let m = parameters[0][0];
let c = parameters[1][0];
for ((i, row), residual) in data.into_iter().enumerate().zip(residuals.iter_mut()) {
let x = row[0];
let y = row[1];
*residual = y - f64::exp(m * x + c);
if let Some(jacobians) = jacobians.as_mut() {
if let Some(d_dm) = jacobians[0].as_mut() {
d_dm[i][0] = -x * f64::exp(m * x + c);
}
if let Some(d_dc) = jacobians[1].as_mut() {
d_dc[i][0] = -f64::exp(m * x + c);
}
}
}
true
});
let initial_guess = vec![vec![0.0], vec![0.0]];
let NllsProblemSolution {
parameters: solution,
summary,
} = NllsProblem::new()
.residual_block_builder()
.set_cost(cost, NUM_OBSERVATIONS)
.set_parameters(initial_guess)
.set_loss(loss)
.build_into_problem()
.unwrap()
.0
.solve(&SolverOptions::default())
.unwrap();
assert!(summary.is_solution_usable());
println!("{}", summary.full_report());
let m = solution[0][0];
let c = solution[1][0];
assert_abs_diff_eq!(0.3, m, epsilon = 0.02);
assert_abs_diff_eq!(0.1, c, epsilon = 0.04);
}
#[test]
fn simple_end_to_end_test_trivial_custom_loss() {
let loss: LossFunctionType = Box::new(|squared_norm: f64, out: &mut [f64; 3]| {
out[0] = squared_norm;
out[1] = 1.0;
out[2] = 0.0;
});
simple_end_to_end_test_with_loss(LossFunction::custom(loss));
}
#[test]
fn simple_end_to_end_test_arctan_stock_loss() {
simple_end_to_end_test_with_loss(LossFunction::arctan(1.0));
}
}