use core::marker::PhantomData;
use crate::numerical_derivative::derivator::{DerivatorMultiVariable, DerivatorSingleVariable};
use crate::scalar::{Dual, HyperDual, Jet, Numeric, ScalarFn, ScalarFnN};
use crate::utils::error_codes::CalcError;
const MAX_ORDER: usize = 6;
#[derive(Debug, Clone, Copy)]
pub struct AutoDiffSingle<T = f64> {
_marker: PhantomData<T>,
}
impl<T> Default for AutoDiffSingle<T> {
fn default() -> Self {
AutoDiffSingle {
_marker: PhantomData,
}
}
}
impl<T: Numeric> DerivatorSingleVariable for AutoDiffSingle<T> {
type Scalar = T;
fn get<F: ScalarFn>(&self, order: usize, func: &F, point: T) -> Result<T, CalcError> {
match order {
0 => Err(CalcError::DerivativeOrderZero),
1 => Ok(func.eval(Dual::variable(point)).deriv),
2 => Ok(func.eval(HyperDual::variable(point)).eps1eps2),
o if o <= MAX_ORDER => Ok(func
.eval(Jet::<T, { MAX_ORDER + 1 }>::variable(point))
.derivative(o)),
_ => Err(CalcError::DerivativeOrderUnsupported),
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct AutoDiffMulti<T = f64> {
_marker: PhantomData<T>,
}
impl<T> Default for AutoDiffMulti<T> {
fn default() -> Self {
AutoDiffMulti {
_marker: PhantomData,
}
}
}
impl<T: Numeric> DerivatorMultiVariable for AutoDiffMulti<T> {
type Scalar = T;
fn get<F: ScalarFnN<NUM_VARS>, const NUM_VARS: usize, const NUM_ORDER: usize>(
&self,
func: &F,
idx_to_differentiate: &[usize; NUM_ORDER],
point: &[T; NUM_VARS],
) -> Result<T, CalcError> {
if NUM_ORDER == 0 {
return Err(CalcError::DerivativeOrderZero);
}
for &idx in idx_to_differentiate {
if idx >= NUM_VARS {
return Err(CalcError::IndexOutOfRange);
}
}
match NUM_ORDER {
1 => {
let i = idx_to_differentiate[0];
let mut seed: [Dual<T>; NUM_VARS] =
core::array::from_fn(|k| Dual::constant(point[k]));
seed[i] = Dual::variable(point[i]);
Ok(func.eval(&seed).deriv)
}
2 => {
let i = idx_to_differentiate[0];
let j = idx_to_differentiate[1];
let mut seed: [HyperDual<T>; NUM_VARS] =
core::array::from_fn(|k| HyperDual::constant(point[k]));
seed[i].eps1 = T::ONE;
seed[j].eps2 = T::ONE;
Ok(func.eval(&seed).eps1eps2)
}
3 => {
let i = idx_to_differentiate[0];
let j = idx_to_differentiate[1];
let k = idx_to_differentiate[2];
let seed: [Dual<HyperDual<T>>; NUM_VARS] = core::array::from_fn(|m| {
let a = if m == i { T::ONE } else { T::ZERO };
let b = if m == j { T::ONE } else { T::ZERO };
let c = if m == k { T::ONE } else { T::ZERO };
Dual::new(
HyperDual::new(point[m], a, b, T::ZERO),
HyperDual::new(c, T::ZERO, T::ZERO, T::ZERO),
)
});
Ok(func.eval(&seed).deriv.eps1eps2)
}
_ => Err(CalcError::DerivativeOrderUnsupported),
}
}
}