use core::cmp::Ordering;
use core::ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Sub, SubAssign};
use crate::scalar::Numeric;
#[derive(Debug, Clone, Copy)]
pub struct HyperDual<T: Numeric> {
pub real: T,
pub eps1: T,
pub eps2: T,
pub eps1eps2: T,
}
impl<T: Numeric> HyperDual<T> {
#[inline]
pub fn new(real: T, eps1: T, eps2: T, eps1eps2: T) -> Self {
HyperDual {
real,
eps1,
eps2,
eps1eps2,
}
}
#[inline]
pub fn constant(real: T) -> Self {
HyperDual {
real,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
}
}
#[inline]
pub fn variable(real: T) -> Self {
HyperDual {
real,
eps1: T::ONE,
eps2: T::ONE,
eps1eps2: T::ZERO,
}
}
#[inline]
fn chain(self, val: T, d1: T, d2: T) -> Self {
HyperDual {
real: val,
eps1: d1 * self.eps1,
eps2: d1 * self.eps2,
eps1eps2: d1 * self.eps1eps2 + d2 * self.eps1 * self.eps2,
}
}
#[inline]
fn recip(self) -> Self {
let t = T::ONE / self.real;
self.chain(t, -(t * t), T::TWO * t * t * t)
}
}
impl<T: Numeric> Add for HyperDual<T> {
type Output = Self;
#[inline]
fn add(self, rhs: Self) -> Self {
HyperDual {
real: self.real + rhs.real,
eps1: self.eps1 + rhs.eps1,
eps2: self.eps2 + rhs.eps2,
eps1eps2: self.eps1eps2 + rhs.eps1eps2,
}
}
}
impl<T: Numeric> Sub for HyperDual<T> {
type Output = Self;
#[inline]
fn sub(self, rhs: Self) -> Self {
HyperDual {
real: self.real - rhs.real,
eps1: self.eps1 - rhs.eps1,
eps2: self.eps2 - rhs.eps2,
eps1eps2: self.eps1eps2 - rhs.eps1eps2,
}
}
}
impl<T: Numeric> Mul for HyperDual<T> {
type Output = Self;
#[inline]
fn mul(self, rhs: Self) -> Self {
HyperDual {
real: self.real * rhs.real,
eps1: self.real * rhs.eps1 + self.eps1 * rhs.real,
eps2: self.real * rhs.eps2 + self.eps2 * rhs.real,
eps1eps2: self.real * rhs.eps1eps2
+ self.eps1 * rhs.eps2
+ self.eps2 * rhs.eps1
+ self.eps1eps2 * rhs.real,
}
}
}
impl<T: Numeric> Div for HyperDual<T> {
type Output = Self;
#[inline]
#[allow(clippy::suspicious_arithmetic_impl)] fn div(self, rhs: Self) -> Self {
self * rhs.recip()
}
}
impl<T: Numeric> Neg for HyperDual<T> {
type Output = Self;
#[inline]
fn neg(self) -> Self {
HyperDual {
real: -self.real,
eps1: -self.eps1,
eps2: -self.eps2,
eps1eps2: -self.eps1eps2,
}
}
}
impl<T: Numeric> AddAssign for HyperDual<T> {
#[inline]
fn add_assign(&mut self, rhs: Self) {
*self = *self + rhs;
}
}
impl<T: Numeric> SubAssign for HyperDual<T> {
#[inline]
fn sub_assign(&mut self, rhs: Self) {
*self = *self - rhs;
}
}
impl<T: Numeric> MulAssign for HyperDual<T> {
#[inline]
fn mul_assign(&mut self, rhs: Self) {
*self = *self * rhs;
}
}
impl<T: Numeric> DivAssign for HyperDual<T> {
#[inline]
fn div_assign(&mut self, rhs: Self) {
*self = *self / rhs;
}
}
impl<T: Numeric> PartialEq for HyperDual<T> {
#[inline]
fn eq(&self, other: &Self) -> bool {
self.real == other.real
}
}
impl<T: Numeric> PartialOrd for HyperDual<T> {
#[inline]
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
self.real.partial_cmp(&other.real)
}
}
impl<T: Numeric> Numeric for HyperDual<T> {
const ZERO: Self = HyperDual {
real: T::ZERO,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const ONE: Self = HyperDual {
real: T::ONE,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const TWO: Self = HyperDual {
real: T::TWO,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const HALF: Self = HyperDual {
real: T::HALF,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const PI: Self = HyperDual {
real: T::PI,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const EPSILON: Self = HyperDual {
real: T::EPSILON,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const NAN: Self = HyperDual {
real: T::NAN,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const INFINITY: Self = HyperDual {
real: T::INFINITY,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const NEG_INFINITY: Self = HyperDual {
real: T::NEG_INFINITY,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const MAX: Self = HyperDual {
real: T::MAX,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
const MIN_POSITIVE: Self = HyperDual {
real: T::MIN_POSITIVE,
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
};
#[inline]
fn from_f64(value: f64) -> Self {
HyperDual::constant(T::from_f64(value))
}
#[inline]
fn from_u64(value: u64) -> Self {
HyperDual::constant(T::from_u64(value))
}
#[inline]
fn from_usize(value: usize) -> Self {
HyperDual::constant(T::from_usize(value))
}
#[inline]
fn abs(self) -> Self {
let sign = if self.real < T::ZERO { -T::ONE } else { T::ONE };
self.chain(self.real.abs(), sign, T::ZERO)
}
#[inline]
fn sqrt(self) -> Self {
let root = self.real.sqrt();
let d1 = T::ONE / (T::TWO * root);
let d2 = -(d1 / (T::TWO * self.real));
self.chain(root, d1, d2)
}
#[inline]
fn sin(self) -> Self {
self.chain(self.real.sin(), self.real.cos(), -(self.real.sin()))
}
#[inline]
fn cos(self) -> Self {
self.chain(self.real.cos(), -(self.real.sin()), -(self.real.cos()))
}
#[inline]
fn tan(self) -> Self {
let t = self.real.tan();
let sec2 = T::ONE + t * t;
self.chain(t, sec2, T::TWO * t * sec2)
}
#[inline]
fn exp(self) -> Self {
let e = self.real.exp();
self.chain(e, e, e)
}
#[inline]
fn ln(self) -> Self {
self.chain(
self.real.ln(),
T::ONE / self.real,
-(T::ONE / (self.real * self.real)),
)
}
#[inline]
fn atan2(self, other: Self) -> Self {
let (y, x) = (self, other);
let (yr, xr) = (y.real, x.real);
let r2 = xr * xr + yr * yr;
let r4 = r2 * r2;
let f_y = xr / r2;
let f_x = -yr / r2;
let f_yy = -(T::TWO * xr * yr) / r4;
let f_xx = (T::TWO * xr * yr) / r4;
let f_xy = (yr * yr - xr * xr) / r4;
HyperDual {
real: yr.atan2(xr),
eps1: f_y * y.eps1 + f_x * x.eps1,
eps2: f_y * y.eps2 + f_x * x.eps2,
eps1eps2: f_y * y.eps1eps2
+ f_x * x.eps1eps2
+ f_yy * y.eps1 * y.eps2
+ f_xx * x.eps1 * x.eps2
+ f_xy * (y.eps1 * x.eps2 + y.eps2 * x.eps1),
}
}
#[inline]
fn copysign(self, sign: Self) -> Self {
let s = if (self.real < T::ZERO) == (sign.real < T::ZERO) {
T::ONE
} else {
-T::ONE
};
self.chain(self.real.copysign(sign.real), s, T::ZERO)
}
#[inline]
fn floor(self) -> Self {
HyperDual {
real: self.real.floor(),
eps1: T::ZERO,
eps2: T::ZERO,
eps1eps2: T::ZERO,
}
}
#[inline]
fn is_nan(self) -> bool {
self.real.is_nan()
}
#[inline]
fn is_finite(self) -> bool {
self.real.is_finite()
}
}