use ndarray::{Array1, Array2};
pub use crate::numeric_derivative::{FdDerivative, FdVerdict, RiddersConfig, ridders_derivative};
pub fn ridders_partial_derivative<F>(
mut objective: F,
x: &Array1<f64>,
coord: usize,
config: RiddersConfig,
) -> FdDerivative
where
F: FnMut(&Array1<f64>) -> f64,
{
assert!(
coord < x.len(),
"ridders_partial_derivative: coordinate {coord} out of range for length {}",
x.len()
);
ridders_derivative(
|t| {
let mut probe = x.clone();
probe[coord] += t;
objective(&probe)
},
config,
)
}
pub fn numerical_gradient_central_diff<F>(mut f: F, x: &Array1<f64>, eps: f64) -> Array1<f64>
where
F: FnMut(&Array1<f64>) -> f64,
{
let mut grad = Array1::zeros(x.len());
for i in 0..x.len() {
let mut xp = x.clone();
let mut xm = x.clone();
xp[i] += eps;
xm[i] -= eps;
grad[i] = (f(&xp) - f(&xm)) / (2.0 * eps);
}
grad
}
pub fn directional_central_diff<F>(
mut f: F,
x: &Array1<f64>,
direction: &Array1<f64>,
eps: f64,
) -> Array1<f64>
where
F: FnMut(&Array1<f64>) -> Array1<f64>,
{
assert_eq!(
x.len(),
direction.len(),
"directional_central_diff: x and direction must have equal length"
);
let xp = x + &(direction * eps);
let xm = x - &(direction * eps);
(f(&xp) - f(&xm)) / (2.0 * eps)
}
pub fn numerical_hessian_central_diff<F>(mut f: F, x: &Array1<f64>, eps: f64) -> Array2<f64>
where
F: FnMut(&Array1<f64>) -> f64,
{
let n = x.len();
let mut hess = Array2::zeros((n, n));
for i in 0..n {
for j in 0..n {
let mut pp = x.clone();
let mut pm = x.clone();
let mut mp = x.clone();
let mut mm = x.clone();
pp[i] += eps;
pp[j] += eps;
pm[i] += eps;
pm[j] -= eps;
mp[i] -= eps;
mp[j] += eps;
mm[i] -= eps;
mm[j] -= eps;
hess[[i, j]] = (f(&pp) - f(&pm) - f(&mp) + f(&mm)) / (4.0 * eps * eps);
}
}
hess
}
pub fn verify_gradient_vs_fd<F>(
objective: F,
analytic_grad: &Array1<f64>,
x: &Array1<f64>,
eps: f64,
tol: f64,
) -> Result<(), String>
where
F: FnMut(&Array1<f64>) -> f64,
{
if analytic_grad.len() != x.len() {
return Err(format!(
"verify_gradient_vs_fd: analytic gradient length {} != x length {}",
analytic_grad.len(),
x.len()
));
}
let fd = numerical_gradient_central_diff(objective, x, eps);
for i in 0..x.len() {
let bound = tol * (1.0 + fd[i].abs());
let gap = (analytic_grad[i] - fd[i]).abs();
if gap > bound {
return Err(format!(
"verify_gradient_vs_fd: coordinate {i} disagrees: analytic={:.6e}, fd={:.6e}, gap={:.3e}, tol={:.3e} (bound {:.3e})",
analytic_grad[i], fd[i], gap, tol, bound
));
}
}
Ok(())
}
pub fn assert_matrix_derivativefd(fd: &Array2<f64>, analytic: &Array2<f64>, tol: f64, label: &str) {
assert_eq!(analytic.dim(), fd.dim(), "{} dimensions must match", label);
for i in 0..analytic.nrows() {
for j in 0..analytic.ncols() {
let analytic_ij = analytic[[i, j]];
let fd_ij = fd[[i, j]];
let diff = (analytic_ij - fd_ij).abs();
if analytic_ij.abs() > tol && fd_ij.abs() > tol {
assert_eq!(
analytic_ij.signum(),
fd_ij.signum(),
"{} sign mismatch at ({}, {}): analytic={}, fd={}",
label,
i,
j,
analytic_ij,
fd_ij
);
}
assert!(
diff <= tol,
"{} value mismatch at ({}, {}): analytic={}, fd={}, abs_diff={}, tol={}",
label,
i,
j,
analytic_ij,
fd_ij,
diff,
tol
);
}
}
}
pub fn assert_matrix_derivativefd_rel(
fd: &Array2<f64>,
analytic: &Array2<f64>,
rel_tol: f64,
label: &str,
) {
assert_eq!(analytic.dim(), fd.dim(), "{} dimensions must match", label);
for i in 0..analytic.nrows() {
for j in 0..analytic.ncols() {
let analytic_ij = analytic[[i, j]];
let fd_ij = fd[[i, j]];
let tol = rel_tol * (1.0 + analytic_ij.abs());
if analytic_ij.abs() > tol && fd_ij.abs() > tol {
assert_eq!(
analytic_ij.signum(),
fd_ij.signum(),
"{} sign mismatch at ({}, {}): analytic={}, fd={}",
label,
i,
j,
analytic_ij,
fd_ij
);
}
let diff = (analytic_ij - fd_ij).abs();
assert!(
diff <= tol,
"{} value mismatch at ({}, {}): analytic={}, fd={}, abs_diff={}, rel_tol={}, tol={}",
label,
i,
j,
analytic_ij,
fd_ij,
diff,
rel_tol,
tol
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_abs_diff_eq;
use ndarray::array;
#[test]
fn quadratic_gradient_and_directional_match_closed_form() {
let a = array![[3.0, 0.5, -0.2], [0.5, 2.0, 0.4], [-0.2, 0.4, 1.5]];
let b = array![0.3, -1.1, 0.7];
let x = array![0.9, -0.4, 1.3];
let objective = |v: &Array1<f64>| 0.5 * v.dot(&a.dot(v)) + b.dot(v);
let analytic_grad = a.dot(&x) + &b;
let eps = 1e-6;
let fd = numerical_gradient_central_diff(objective, &x, eps);
for i in 0..x.len() {
assert_abs_diff_eq!(fd[i], analytic_grad[i], epsilon = 1e-6);
}
verify_gradient_vs_fd(objective, &analytic_grad, &x, eps, 1e-5)
.expect("analytic gradient matches FD of the quadratic");
let direction = array![0.6, -0.8, 0.2];
let grad_map = |v: &Array1<f64>| a.dot(v) + &b;
let hvp_fd = directional_central_diff(grad_map, &x, &direction, eps);
let hvp_exact = a.dot(&direction);
for i in 0..direction.len() {
assert_abs_diff_eq!(hvp_fd[i], hvp_exact[i], epsilon = 1e-6);
}
let hess_fd = numerical_hessian_central_diff(objective, &x, 1e-4);
for i in 0..x.len() {
for j in 0..x.len() {
assert_abs_diff_eq!(hess_fd[[i, j]], a[[i, j]], epsilon = 1e-6);
}
}
assert_matrix_derivativefd(&hess_fd, &a, 1e-5, "quadratic curvature");
assert_matrix_derivativefd_rel(&hess_fd, &a, 1e-5, "quadratic curvature (relative)");
}
#[test]
fn verify_rejects_wrong_gradient() {
let x = array![1.0, 2.0];
let objective = |v: &Array1<f64>| v[0] * v[0] + v[1] * v[1];
let exact = array![2.0, 4.0];
verify_gradient_vs_fd(objective, &exact, &x, 1e-6, 1e-5).expect("exact gradient passes");
let wrong = array![2.0, 4.5];
let err = verify_gradient_vs_fd(objective, &wrong, &x, 1e-6, 1e-5)
.expect_err("perturbed gradient must be rejected");
assert!(
err.contains("coordinate 1"),
"error should name coord 1: {err}"
);
}
#[test]
fn ridders_partial_picks_the_requested_coordinate() {
let x = array![0.7, -1.3, 2.1];
let objective = |v: &Array1<f64>| v[0] * v[1] + (v[2] * v[2]).sin();
for (coord, exact) in [
(0usize, x[1]),
(1, x[0]),
(2, 2.0 * x[2] * f64::cos(x[2] * x[2])),
] {
let measured =
ridders_partial_derivative(objective, &x, coord, RiddersConfig::default());
assert!(
(measured.value - exact).abs() <= 1e-9,
"coord {coord}: {:.10e} vs {exact:.10e}",
measured.value
);
}
}
#[test]
#[should_panic(expected = "identity curvature value mismatch at (1, 1)")]
fn matrix_assert_rejects_a_disagreeing_entry() {
let analytic = array![[1.0, 0.0], [0.0, 1.0]];
let fd = array![[1.0, 0.0], [0.0, 1.5]];
assert_matrix_derivativefd(&fd, &analytic, 1e-6, "identity curvature");
}
}