num_dual/datatypes/
hyperdual.rs1use 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#[derive(Copy, Clone, Debug)]
13#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
14pub struct HyperDual<T> {
15 pub re: T,
17 pub eps1: T,
19 pub eps2: T,
21 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 #[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 #[inline]
47 pub fn derivative1(mut self) -> Self {
48 self.eps1 = T::one();
49 self
50 }
51
52 #[inline]
54 pub fn derivative2(mut self) -> Self {
55 self.eps2 = T::one();
56 self
57 }
58}
59
60impl<T: DualNum> HyperDual<T> {
61 #[inline]
63 pub fn from_re(re: T) -> Self {
64 Self::new(re, T::zero(), T::zero(), T::zero())
65 }
66}
67
68impl<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
81impl<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
98impl<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
123impl<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]);