use alloc::vec;
use crate::callback_runtime::{
self, CallbackRuntimeError, LeastSquaresCallbackFailure, LeastSquaresF32Callback,
LeastSquaresF64Callback,
};
use slatec_core::to_fortran_integer;
use slatec_sys::FortranInteger;
use super::{LeastSquaresError, LeastSquaresStatus};
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct LeastSquaresOptions<T = f64> {
pub tolerance: T,
}
impl Default for LeastSquaresOptions<f64> {
fn default() -> Self {
Self { tolerance: 1.0e-10 }
}
}
impl LeastSquaresOptions<f32> {
pub const fn single_precision() -> Self {
Self { tolerance: 1.0e-5 }
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct LeastSquaresResult<T = f64> {
pub parameters: alloc::vec::Vec<T>,
pub residuals: alloc::vec::Vec<T>,
pub cost: T,
pub residual_norm: T,
pub status: LeastSquaresStatus,
pub function_evaluations: usize,
}
fn workspace_len(
parameter_count: usize,
residual_count: usize,
) -> Result<usize, LeastSquaresError> {
let residual_plus_five = residual_count
.checked_add(5)
.ok_or(LeastSquaresError::WorkspaceOverflow)?;
parameter_count
.checked_mul(residual_plus_five)
.and_then(|value| value.checked_add(residual_count))
.ok_or(LeastSquaresError::WorkspaceOverflow)
}
fn native_integer(
value: usize,
argument: &'static str,
) -> Result<FortranInteger, LeastSquaresError> {
to_fortran_integer(value).map_err(|_| LeastSquaresError::IntegerOverflow { argument })
}
fn validate_f64(
initial: &[f64],
residual_count: usize,
options: LeastSquaresOptions<f64>,
) -> Result<(), LeastSquaresError> {
if initial.is_empty() {
return Err(LeastSquaresError::EmptyParameters);
}
if residual_count == 0 {
return Err(LeastSquaresError::EmptyResiduals);
}
if residual_count < initial.len() {
return Err(LeastSquaresError::Underdetermined {
residuals: residual_count,
parameters: initial.len(),
});
}
if let Some((index, _)) = initial
.iter()
.enumerate()
.find(|(_, value)| !value.is_finite())
{
return Err(LeastSquaresError::NonFiniteInitialValue { index });
}
if !options.tolerance.is_finite() || options.tolerance < 0.0 {
return Err(LeastSquaresError::InvalidTolerance);
}
Ok(())
}
fn validate_f32(
initial: &[f32],
residual_count: usize,
options: LeastSquaresOptions<f32>,
) -> Result<(), LeastSquaresError> {
if initial.is_empty() {
return Err(LeastSquaresError::EmptyParameters);
}
if residual_count == 0 {
return Err(LeastSquaresError::EmptyResiduals);
}
if residual_count < initial.len() {
return Err(LeastSquaresError::Underdetermined {
residuals: residual_count,
parameters: initial.len(),
});
}
if let Some((index, _)) = initial
.iter()
.enumerate()
.find(|(_, value)| !value.is_finite())
{
return Err(LeastSquaresError::NonFiniteInitialValue { index });
}
if !options.tolerance.is_finite() || options.tolerance < 0.0 {
return Err(LeastSquaresError::InvalidTolerance);
}
Ok(())
}
fn callback_failure(failure: LeastSquaresCallbackFailure) -> LeastSquaresError {
match failure {
LeastSquaresCallbackFailure::Panicked => LeastSquaresError::CallbackPanicked,
LeastSquaresCallbackFailure::NonFinite { index } => {
LeastSquaresError::CallbackReturnedNonFinite { index }
}
LeastSquaresCallbackFailure::InvalidPointer => LeastSquaresError::NativeContractViolation {
detail: "native least-squares callback pointer was null or overlapped",
},
LeastSquaresCallbackFailure::DimensionMismatch => {
LeastSquaresError::NativeContractViolation {
detail: "native M or N did not match the registered least-squares callback",
}
}
LeastSquaresCallbackFailure::UnexpectedFlag => LeastSquaresError::NativeContractViolation {
detail: "DNLS1E/SNLS1E IOPT=1 callback received an unexpected IFLAG",
},
}
}
fn callback_runtime_error(error: CallbackRuntimeError) -> LeastSquaresError {
match error {
CallbackRuntimeError::NestedCallback => LeastSquaresError::NestedNativeCallback,
}
}
fn native_status(status: FortranInteger) -> Result<LeastSquaresStatus, LeastSquaresError> {
match status {
1 => Ok(LeastSquaresStatus::ConvergedResidual),
2 => Ok(LeastSquaresStatus::ConvergedParameters),
3 => Ok(LeastSquaresStatus::ConvergedResidualAndParameters),
4 => Ok(LeastSquaresStatus::ConvergedOrthogonality),
5 => Ok(LeastSquaresStatus::MaximumEvaluations),
6 => Ok(LeastSquaresStatus::ResidualToleranceTooSmall),
7 => Ok(LeastSquaresStatus::ParameterToleranceTooSmall),
value => Err(LeastSquaresError::NativeStatus { status: value }),
}
}
fn norm_f64(values: &[f64]) -> f64 {
let mut scale = 0.0_f64;
let mut sum = 1.0_f64;
for value in values {
let magnitude = value.abs();
if magnitude != 0.0 {
if scale < magnitude {
sum = 1.0 + sum * (scale / magnitude) * (scale / magnitude);
scale = magnitude;
} else {
sum += (magnitude / scale) * (magnitude / scale);
}
}
}
if scale == 0.0 {
0.0
} else {
scale * sum.sqrt()
}
}
fn norm_f32(values: &[f32]) -> f32 {
let mut scale = 0.0_f32;
let mut sum = 1.0_f32;
for value in values {
let magnitude = value.abs();
if magnitude != 0.0 {
if scale < magnitude {
sum = 1.0 + sum * (scale / magnitude) * (scale / magnitude);
scale = magnitude;
} else {
sum += (magnitude / scale) * (magnitude / scale);
}
}
}
if scale == 0.0 {
0.0
} else {
scale * sum.sqrt()
}
}
fn run_f64<F>(
initial: &[f64],
residual_count: usize,
function: F,
options: LeastSquaresOptions<f64>,
) -> Result<LeastSquaresResult<f64>, LeastSquaresError>
where
F: FnMut(&[f64], &mut [f64]),
{
validate_f64(initial, residual_count, options)?;
let parameter_count = initial.len();
let workspace_length = workspace_len(parameter_count, residual_count)?;
let mut parameters = initial.to_vec();
let mut residuals = vec![0.0; residual_count];
let mut integer_workspace = vec![0; parameter_count];
let mut workspace = vec![0.0; workspace_length];
let mut m = native_integer(residual_count, "residual count")?;
let mut n = native_integer(parameter_count, "parameter count")?;
let mut lwa = native_integer(workspace_length, "workspace length")?;
let mut iopt = 1;
let mut tolerance = options.tolerance;
let mut nprint = 0;
let mut info = 0;
let invocation = callback_runtime::with_least_squares_f64(
parameter_count,
residual_count,
function,
|callback: LeastSquaresF64Callback| {
let _error_scope = crate::runtime::permit_recoverable_native_statuses();
unsafe {
slatec_sys::least_squares::dnls1e(
callback.ffi(),
&mut iopt,
&mut m,
&mut n,
parameters.as_mut_ptr(),
residuals.as_mut_ptr(),
&mut tolerance,
&mut nprint,
&mut info,
integer_workspace.as_mut_ptr(),
workspace.as_mut_ptr(),
&mut lwa,
);
}
},
)
.map_err(callback_runtime_error)?;
if let Some(failure) = invocation.failure {
return Err(callback_failure(failure));
}
let residual_norm = norm_f64(&residuals);
Ok(LeastSquaresResult {
parameters,
residuals,
cost: 0.5 * residual_norm * residual_norm,
residual_norm,
status: native_status(info)?,
function_evaluations: invocation.evaluations,
})
}
fn run_f32<F>(
initial: &[f32],
residual_count: usize,
function: F,
options: LeastSquaresOptions<f32>,
) -> Result<LeastSquaresResult<f32>, LeastSquaresError>
where
F: FnMut(&[f32], &mut [f32]),
{
validate_f32(initial, residual_count, options)?;
let parameter_count = initial.len();
let workspace_length = workspace_len(parameter_count, residual_count)?;
let mut parameters = initial.to_vec();
let mut residuals = vec![0.0; residual_count];
let mut integer_workspace = vec![0; parameter_count];
let mut workspace = vec![0.0; workspace_length];
let mut m = native_integer(residual_count, "residual count")?;
let mut n = native_integer(parameter_count, "parameter count")?;
let mut lwa = native_integer(workspace_length, "workspace length")?;
let mut iopt = 1;
let mut tolerance = options.tolerance;
let mut nprint = 0;
let mut info = 0;
let invocation = callback_runtime::with_least_squares_f32(
parameter_count,
residual_count,
function,
|callback: LeastSquaresF32Callback| {
let _error_scope = crate::runtime::permit_recoverable_native_statuses();
unsafe {
slatec_sys::least_squares::snls1e(
callback.ffi(),
&mut iopt,
&mut m,
&mut n,
parameters.as_mut_ptr(),
residuals.as_mut_ptr(),
&mut tolerance,
&mut nprint,
&mut info,
integer_workspace.as_mut_ptr(),
workspace.as_mut_ptr(),
&mut lwa,
);
}
},
)
.map_err(callback_runtime_error)?;
if let Some(failure) = invocation.failure {
return Err(callback_failure(failure));
}
let residual_norm = norm_f32(&residuals);
Ok(LeastSquaresResult {
parameters,
residuals,
cost: 0.5 * residual_norm * residual_norm,
residual_norm,
status: native_status(info)?,
function_evaluations: invocation.evaluations,
})
}
pub fn least_squares<F>(
initial: &[f64],
residual_count: usize,
residuals: F,
options: LeastSquaresOptions<f64>,
) -> Result<LeastSquaresResult<f64>, LeastSquaresError>
where
F: FnMut(&[f64], &mut [f64]),
{
run_f64(initial, residual_count, residuals, options)
}
pub fn least_squares_f32<F>(
initial: &[f32],
residual_count: usize,
residuals: F,
options: LeastSquaresOptions<f32>,
) -> Result<LeastSquaresResult<f32>, LeastSquaresError>
where
F: FnMut(&[f32], &mut [f32]),
{
run_f32(initial, residual_count, residuals, options)
}
#[cfg(test)]
mod tests {
use super::{LeastSquaresError, LeastSquaresOptions, native_status, norm_f64, workspace_len};
#[test]
fn workspace_formula_is_exact_and_checked() {
assert_eq!(workspace_len(2, 4), Ok(22));
assert_eq!(
workspace_len(usize::MAX, 1),
Err(LeastSquaresError::WorkspaceOverflow)
);
}
#[test]
fn validation_rejects_rectangular_contract_violations() {
let options = LeastSquaresOptions::default();
assert!(matches!(
super::validate_f64(&[], 1, options),
Err(LeastSquaresError::EmptyParameters)
));
assert!(matches!(
super::validate_f64(&[0.0, 1.0], 1, options),
Err(LeastSquaresError::Underdetermined { .. })
));
}
#[test]
fn statuses_and_norm_are_preserved() {
assert!(matches!(
native_status(5),
Ok(super::LeastSquaresStatus::MaximumEvaluations)
));
assert_eq!(norm_f64(&[3.0, 4.0]), 5.0);
}
}