Skip to main content

num_dual/datatypes/
hyperdual.rs

1use crate::{DualNum, DualNumFloat, DualStruct};
2use num_traits::{FloatConst, FromPrimitive, Inv, Num, One, Signed, Zero};
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Serialize};
5use std::fmt;
6use std::iter::{Product, Sum};
7use std::ops::{
8    Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Rem, RemAssign, Sub, SubAssign,
9};
10
11/// A scalar hyper-dual number for the calculation of second partial derivatives.
12#[derive(Copy, Clone, Debug)]
13#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
14pub struct HyperDual<T> {
15    /// Real part of the hyper-dual number
16    pub re: T,
17    /// Partial derivative part of the hyper-dual number
18    pub eps1: T,
19    /// Partial derivative part of the hyper-dual number
20    pub eps2: T,
21    /// Second partial derivative part of the hyper-dual number
22    pub eps1eps2: T,
23}
24
25#[cfg(feature = "ndarray")]
26impl<T: DualNum> ndarray::ScalarOperand for HyperDual<T> {}
27
28pub type HyperDual32 = HyperDual<f32>;
29pub type HyperDual64 = HyperDual<f64>;
30
31impl<T> HyperDual<T> {
32    /// Create a new hyper-dual number from its fields.
33    #[inline]
34    pub fn new(re: T, eps1: T, eps2: T, eps1eps2: T) -> Self {
35        Self {
36            re,
37            eps1,
38            eps2,
39            eps1eps2,
40        }
41    }
42}
43
44impl<T: DualNum> HyperDual<T> {
45    /// Set the partial derivative part w.r.t. the 1st variable to 1.
46    #[inline]
47    pub fn derivative1(mut self) -> Self {
48        self.eps1 = T::one();
49        self
50    }
51
52    /// Set the partial derivative part w.r.t. the 2nd variable to 1.
53    #[inline]
54    pub fn derivative2(mut self) -> Self {
55        self.eps2 = T::one();
56        self
57    }
58}
59
60impl<T: DualNum> HyperDual<T> {
61    /// Create a new hyper-dual number from the real part.
62    #[inline]
63    pub fn from_re(re: T) -> Self {
64        Self::new(re, T::zero(), T::zero(), T::zero())
65    }
66}
67
68/* chain rule */
69impl<T: DualNum> HyperDual<T> {
70    #[inline]
71    fn chain_rule(&self, f0: T, f1: T, f2: T) -> Self {
72        Self::new(
73            f0,
74            self.eps1.clone() * f1.clone(),
75            self.eps2.clone() * f1.clone(),
76            self.eps1eps2.clone() * f1 + self.eps1.clone() * self.eps2.clone() * f2,
77        )
78    }
79}
80
81/* product rule */
82impl<T: DualNum> Mul<&HyperDual<T>> for &HyperDual<T> {
83    type Output = HyperDual<T>;
84    #[inline]
85    fn mul(self, other: &HyperDual<T>) -> HyperDual<T> {
86        HyperDual::new(
87            self.re.clone() * other.re.clone(),
88            other.eps1.clone() * self.re.clone() + self.eps1.clone() * other.re.clone(),
89            other.eps2.clone() * self.re.clone() + self.eps2.clone() * other.re.clone(),
90            other.eps1eps2.clone() * self.re.clone()
91                + self.eps1.clone() * other.eps2.clone()
92                + other.eps1.clone() * self.eps2.clone()
93                + self.eps1eps2.clone() * other.re.clone(),
94        )
95    }
96}
97
98/* quotient rule */
99impl<T: DualNum> Div<&HyperDual<T>> for &HyperDual<T> {
100    type Output = HyperDual<T>;
101    #[inline]
102    fn div(self, other: &HyperDual<T>) -> HyperDual<T> {
103        let inv = other.re.recip();
104        let inv2 = inv.clone() * &inv;
105        HyperDual::new(
106            self.re.clone() * &inv,
107            (self.eps1.clone() * other.re.clone() - other.eps1.clone() * self.re.clone())
108                * inv2.clone(),
109            (self.eps2.clone() * other.re.clone() - other.eps2.clone() * self.re.clone())
110                * inv2.clone(),
111            self.eps1eps2.clone() * inv.clone()
112                - (other.eps1eps2.clone() * self.re.clone()
113                    + self.eps1.clone() * other.eps2.clone()
114                    + other.eps1.clone() * self.eps2.clone())
115                    * inv2.clone()
116                + other.eps1.clone()
117                    * other.eps2.clone()
118                    * ((T::one() + T::one()) * self.re.clone() * inv2 * inv),
119        )
120    }
121}
122
123/* string conversions */
124impl<T: DualNum> fmt::Display for HyperDual<T> {
125    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
126        fmt::Display::fmt(&self.re, f)?;
127        write!(f, " + ")?;
128        fmt::Display::fmt(&self.eps1, f)?;
129        write!(f, "ε1 + ")?;
130        fmt::Display::fmt(&self.eps2, f)?;
131        write!(f, "ε2 + ")?;
132        fmt::Display::fmt(&self.eps1eps2, f)?;
133        write!(f, "ε1ε2")
134    }
135}
136
137impl_second_derivatives!(HyperDual, [eps1, eps2, eps1eps2]);
138impl_dual!(HyperDual, [eps1, eps2, eps1eps2]);
139#[cfg(feature = "nalgebra")]
140impl_nalgebra!(HyperDual, [eps1, eps2, eps1eps2]);