use super::{Method, bounded};
use crate::core::math::{MatrixIndex, Scalar, VectorIndex, VectorLen};
use crate::core::parallel::{MaybeSend, MaybeSync};
use crate::core::problem::{Gradient, Jacobian};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum DerivativeSource {
Analytic,
FiniteDifference,
}
#[derive(Debug)]
#[non_exhaustive]
pub enum DerivativeCheckError<E, F: Scalar = f64> {
Evaluation(E),
InvalidInput(&'static str),
OutputSize {
expected: usize,
actual: usize,
},
NonFiniteEvaluation {
point: Vec<F>,
output: usize,
value: F,
},
NonFiniteDerivative {
source: DerivativeSource,
output: usize,
coordinate: Option<usize>,
value: F,
},
NoFeasibleDirection,
}
impl<E: std::fmt::Display, F: Scalar> std::fmt::Display
for DerivativeCheckError<E, F>
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Evaluation(e) => {
write!(f, "derivative-check callback failed: {e}")
}
Self::InvalidInput(message) => f.write_str(message),
Self::OutputSize { expected, actual } => {
write!(f, "expected {expected} residual entries, got {actual}")
}
Self::NonFiniteEvaluation {
point,
output,
value,
} => write!(
f,
"non-finite function value {value:?} at output {output}, point {point:?}"
),
Self::NonFiniteDerivative {
source,
output,
coordinate,
value,
} => write!(
f,
"non-finite {source:?} derivative {value:?} at output {output}, coordinate {coordinate:?}"
),
Self::NoFeasibleDirection => f.write_str(
"no feasible representable probe along the direction",
),
}
}
}
impl<E: std::error::Error + 'static, F: Scalar> std::error::Error
for DerivativeCheckError<E, F>
{
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Evaluation(e) => Some(e),
_ => None,
}
}
}
#[derive(Clone, Debug)]
pub struct DerivativeComparison<F: Scalar = f64> {
pub output: usize,
pub coordinate: Option<usize>,
pub analytic: F,
pub finite_difference: F,
pub absolute_error: F,
pub relative_error: F,
pub scaled_error: F,
}
impl<F: Scalar> DerivativeComparison<F> {
pub fn passed(&self) -> bool {
self.scaled_error <= F::one()
}
}
#[derive(Clone, Debug)]
pub struct DerivativeCheckReport<F: Scalar = f64> {
pub comparisons: Vec<DerivativeComparison<F>>,
pub skipped_coordinates: Vec<usize>,
pub direction: Option<Vec<F>>,
}
impl<F: Scalar> DerivativeCheckReport<F> {
pub fn passed(&self) -> bool {
self.comparisons.iter().all(DerivativeComparison::passed)
}
}
#[derive(Clone, Debug)]
pub struct DerivativeChecker<F: Scalar = f64> {
method: Method,
precision: F,
step: Option<F>,
absolute_tolerance: F,
relative_tolerance: F,
bounds: Option<(Vec<F>, Vec<F>)>,
}
impl<F: Scalar> Default for DerivativeChecker<F> {
fn default() -> Self {
Self::new()
}
}
impl<F: Scalar> DerivativeChecker<F> {
pub fn new() -> Self {
Self {
method: Method::Central,
precision: F::epsilon(),
step: None,
absolute_tolerance: F::epsilon().sqrt(),
relative_tolerance: F::epsilon().sqrt(),
bounds: None,
}
}
pub fn method(mut self, method: Method) -> Self {
self.method = method;
self
}
pub fn with_absolute_tolerance(mut self, tolerance: F) -> Self {
assert!(
tolerance.is_finite() && tolerance >= F::zero(),
"invalid absolute tolerance"
);
self.absolute_tolerance = tolerance;
self
}
pub fn with_relative_tolerance(mut self, tolerance: F) -> Self {
assert!(
tolerance.is_finite() && tolerance >= F::zero(),
"invalid relative tolerance"
);
self.relative_tolerance = tolerance;
self
}
pub fn function_precision(mut self, precision: F) -> Self {
assert!(
precision.is_finite() && precision >= F::zero(),
"invalid function precision"
);
self.precision = precision.max(F::epsilon());
self
}
pub fn with_step(mut self, step: F) -> Self {
assert!(
step.is_finite() && step > F::zero(),
"invalid finite-difference step"
);
self.step = Some(step);
self
}
pub fn with_bounds<V: VectorLen + VectorIndex<F>>(
mut self,
lower: V,
upper: V,
) -> Self {
assert_eq!(lower.vec_len(), upper.vec_len(), "bound length mismatch");
let lower: Vec<_> = (0..lower.vec_len())
.map(|i| {
let v = lower.get_scalar(i);
if v.is_finite() { v } else { F::neg_infinity() }
})
.collect();
let upper: Vec<_> = (0..upper.vec_len())
.map(|i| {
let v = upper.get_scalar(i);
if v.is_finite() { v } else { F::infinity() }
})
.collect();
assert!(
lower.iter().zip(&upper).all(|(l, u)| l <= u),
"lower bound exceeds upper bound"
);
self.bounds = Some((lower, upper));
self
}
fn options(&self) -> bounded::Options<F> {
bounded::Options {
method: self.method,
precision: self.precision,
step: self.step,
minpack: false,
}
}
fn bounds(&self) -> Option<(&[F], &[F])> {
self.bounds
.as_ref()
.map(|(l, u)| (l.as_slice(), u.as_slice()))
}
fn fixed(&self, j: usize) -> bool {
self.bounds.as_ref().is_some_and(|(l, u)| l[j] == u[j])
}
fn validate_point<V: VectorLen + VectorIndex<F>, E>(
&self,
x: &V,
) -> Result<(), DerivativeCheckError<E, F>> {
if let Some((l, _)) = &self.bounds {
if l.len() != x.vec_len() {
return Err(DerivativeCheckError::InvalidInput(
"point and bounds length mismatch",
));
}
}
for j in 0..x.vec_len() {
let v = x.get_scalar(j);
if !v.is_finite() {
return Err(DerivativeCheckError::InvalidInput(
"point must be finite",
));
}
if self
.bounds
.as_ref()
.is_some_and(|(l, u)| v < l[j] || v > u[j])
{
return Err(DerivativeCheckError::InvalidInput(
"point is outside probe bounds",
));
}
}
Ok(())
}
}
impl<F: Scalar + MaybeSend + MaybeSync> DerivativeChecker<F> {
pub fn check_gradient<P>(
&self,
problem: &P,
x: &P::Param,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<P::Error, F>>
where
P: Gradient<Output = F> + MaybeSync,
P::Param: Clone + VectorLen + VectorIndex<F> + MaybeSync,
P::Gradient: VectorLen + VectorIndex<F>,
P::Error: MaybeSend,
{
self.gradient(problem, x, None)
}
pub fn check_gradient_direction<P>(
&self,
problem: &P,
x: &P::Param,
direction: &P::Param,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<P::Error, F>>
where
P: Gradient<Output = F> + MaybeSync,
P::Param: Clone + VectorLen + VectorIndex<F> + MaybeSync,
P::Gradient: VectorLen + VectorIndex<F>,
P::Error: MaybeSend,
{
self.gradient(problem, x, Some(direction))
}
fn gradient<P>(
&self,
problem: &P,
x: &P::Param,
direction: Option<&P::Param>,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<P::Error, F>>
where
P: Gradient<Output = F> + MaybeSync,
P::Param: Clone + VectorLen + VectorIndex<F> + MaybeSync,
P::Gradient: VectorLen + VectorIndex<F>,
P::Error: MaybeSend,
{
self.validate_point(x)?;
let g = problem
.gradient(x)
.map_err(DerivativeCheckError::Evaluation)?;
if g.vec_len() != x.vec_len() {
return Err(DerivativeCheckError::InvalidInput(
"gradient and point length mismatch",
));
}
let analytic =
(0..g.vec_len()).map(|j| vec![g.get_scalar(j)]).collect();
self.run(x, 1, analytic, direction, |x| {
problem.cost(x).map(|f| vec![f])
})
}
pub fn check_jacobian<P>(
&self,
problem: &P,
x: &P::Param,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<P::Error, F>>
where
P: Jacobian + MaybeSync,
P::Param: Clone + VectorLen + VectorIndex<F> + MaybeSync,
P::Output: VectorLen + VectorIndex<F>,
P::Jacobian: MatrixIndex<F>,
P::Error: MaybeSend,
{
self.jacobian(problem, x, None)
}
pub fn check_jacobian_direction<P>(
&self,
problem: &P,
x: &P::Param,
direction: &P::Param,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<P::Error, F>>
where
P: Jacobian + MaybeSync,
P::Param: Clone + VectorLen + VectorIndex<F> + MaybeSync,
P::Output: VectorLen + VectorIndex<F>,
P::Jacobian: MatrixIndex<F>,
P::Error: MaybeSend,
{
self.jacobian(problem, x, Some(direction))
}
fn jacobian<P>(
&self,
problem: &P,
x: &P::Param,
direction: Option<&P::Param>,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<P::Error, F>>
where
P: Jacobian + MaybeSync,
P::Param: Clone + VectorLen + VectorIndex<F> + MaybeSync,
P::Output: VectorLen + VectorIndex<F>,
P::Jacobian: MatrixIndex<F>,
P::Error: MaybeSend,
{
self.validate_point(x)?;
let a = problem
.jacobian(x)
.map_err(DerivativeCheckError::Evaluation)?;
if a.matrix_cols() != x.vec_len() {
return Err(DerivativeCheckError::InvalidInput(
"Jacobian column count and point length mismatch",
));
}
let analytic = (0..a.matrix_cols())
.map(|j| {
(0..a.matrix_rows()).map(|i| a.matrix_entry(i, j)).collect()
})
.collect();
self.run(x, a.matrix_rows(), analytic, direction, |x| {
problem
.residual(x)
.map(|r| (0..r.vec_len()).map(|i| r.get_scalar(i)).collect())
})
}
fn run<V, E, C>(
&self,
x: &V,
rows: usize,
analytic: Vec<Vec<F>>,
direction: Option<&V>,
evaluate: C,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<E, F>>
where
V: Clone + VectorLen + VectorIndex<F> + MaybeSync,
E: MaybeSend,
C: Fn(&V) -> Result<Vec<F>, E> + MaybeSync,
{
let evaluate = |probe: &V| {
self.validate_point(probe)?;
let values =
evaluate(probe).map_err(DerivativeCheckError::Evaluation)?;
if values.len() != rows {
return Err(DerivativeCheckError::OutputSize {
expected: rows,
actual: values.len(),
});
}
for (output, &value) in values.iter().enumerate() {
if !value.is_finite() {
return Err(DerivativeCheckError::NonFiniteEvaluation {
point: (0..probe.vec_len())
.map(|j| probe.get_scalar(j))
.collect(),
output,
value,
});
}
}
Ok(values)
};
if let Some(d) = direction {
return self.directional(x, rows, analytic, d, evaluate);
}
let mut report = DerivativeCheckReport {
comparisons: Vec::new(),
skipped_coordinates: Vec::new(),
direction: None,
};
let (_, numeric) =
bounded::columns(x, self.options(), self.bounds(), evaluate)?;
for (j, (a, n)) in analytic.iter().zip(&numeric).enumerate() {
if self.fixed(j) {
report.skipped_coordinates.push(j);
continue;
}
for i in 0..rows {
report.comparisons.push(self.compare(
a[i],
n[i],
i,
Some(j),
)?);
}
}
Ok(report)
}
fn directional<V, E, C>(
&self,
x: &V,
rows: usize,
analytic: Vec<Vec<F>>,
direction: &V,
evaluate: C,
) -> Result<DerivativeCheckReport<F>, DerivativeCheckError<E, F>>
where
V: Clone + VectorLen + VectorIndex<F> + MaybeSync,
E: MaybeSend,
C: Fn(&V) -> Result<Vec<F>, DerivativeCheckError<E, F>> + MaybeSync,
{
if direction.vec_len() != x.vec_len() {
return Err(DerivativeCheckError::InvalidInput(
"direction and point length mismatch",
));
}
let mut d: Vec<_> = (0..direction.vec_len())
.map(|j| direction.get_scalar(j))
.collect();
let mut magnitude = F::zero();
let mut pivot = 0;
for (j, &v) in d.iter().enumerate() {
if !v.is_finite() {
return Err(DerivativeCheckError::InvalidInput(
"direction must be finite",
));
}
if v != F::zero() && self.fixed(j) {
return Err(DerivativeCheckError::NoFeasibleDirection);
}
if v.abs() > magnitude {
magnitude = v.abs();
pivot = j;
}
}
if magnitude == F::zero() {
return Err(DerivativeCheckError::InvalidInput(
"direction must be nonzero",
));
}
for v in &mut d {
*v = *v / magnitude;
}
let sign = d[pivot];
let anchor = x.get_scalar(pivot);
let mut lower = F::neg_infinity();
let mut upper = F::infinity();
let scale = if matches!(self.method, Method::Forward) {
self.precision.sqrt()
} else {
self.precision.cbrt()
};
let mut step = F::infinity();
for (j, &v) in d.iter().enumerate() {
if v == F::zero() {
continue;
}
step =
step.min(scale * x.get_scalar(j).abs().max(F::one()) / v.abs());
if let Some((lo, hi)) = &self.bounds {
let slope = v / sign;
let a = anchor + (lo[j] - x.get_scalar(j)) / slope;
let b = anchor + (hi[j] - x.get_scalar(j)) / slope;
lower = lower.max(a.min(b));
upper = upper.min(a.max(b));
}
}
if lower >= upper || lower > anchor || upper < anchor {
return Err(DerivativeCheckError::NoFeasibleDirection);
}
let mut options = self.options();
options.step = Some(self.step.unwrap_or(step));
let (_, numeric) = bounded::columns(
&vec![anchor],
options,
Some((&[lower], &[upper])),
|s| {
let t = (s[0] - anchor) / sign;
let mut probe = x.clone();
for (j, &v) in d.iter().enumerate() {
let value = if j == pivot {
s[0]
} else if v == F::zero() {
x.get_scalar(j)
} else {
t.mul_add(v, x.get_scalar(j))
};
if let Some((lo, hi)) = &self.bounds {
if value < lo[j] || value > hi[j] {
return Err(
DerivativeCheckError::NoFeasibleDirection,
);
}
}
probe.set_scalar(j, value);
}
evaluate(&probe)
},
)?;
let mut report = DerivativeCheckReport {
comparisons: Vec::new(),
skipped_coordinates: Vec::new(),
direction: Some(d.clone()),
};
for i in 0..rows {
let mut product = F::zero();
for (j, &v) in d.iter().enumerate() {
if v != F::zero() {
finite(
analytic[j][i],
DerivativeSource::Analytic,
i,
Some(j),
)?;
product = product + analytic[j][i] * v;
}
}
report.comparisons.push(self.compare(
product,
numeric[0][i] * sign,
i,
None,
)?);
}
Ok(report)
}
fn compare<E>(
&self,
analytic: F,
numeric: F,
output: usize,
coordinate: Option<usize>,
) -> Result<DerivativeComparison<F>, DerivativeCheckError<E, F>> {
finite(analytic, DerivativeSource::Analytic, output, coordinate)?;
finite(
numeric,
DerivativeSource::FiniteDifference,
output,
coordinate,
)?;
let absolute_error = (analytic - numeric).abs();
let magnitude = analytic.abs().max(numeric.abs());
let relative_error = if magnitude == F::zero() {
F::zero()
} else if absolute_error.is_finite() {
absolute_error / magnitude
} else {
(analytic / magnitude - numeric / magnitude).abs()
};
let scaled_error = if absolute_error == F::zero() {
F::zero()
} else {
relative_error
/ (self.absolute_tolerance / magnitude
+ self.relative_tolerance)
};
Ok(DerivativeComparison {
output,
coordinate,
analytic,
finite_difference: numeric,
absolute_error,
relative_error,
scaled_error,
})
}
}
fn finite<E, F: Scalar>(
value: F,
source: DerivativeSource,
output: usize,
coordinate: Option<usize>,
) -> Result<(), DerivativeCheckError<E, F>> {
if value.is_finite() {
Ok(())
} else {
Err(DerivativeCheckError::NonFiniteDerivative {
source,
output,
coordinate,
value,
})
}
}