Skip to main content

num_dual/
impl_derivatives.rs

1#[macro_export]
2macro_rules! impl_derivatives {
3    ($deriv:ident, $nderiv:expr, $struct:ident, [$($im:ident),*]$(, [$($dim:tt),*]$(, [$($ddim:tt),*])*)?) => {
4        impl<T: DualNum<Primitive = F>, F: DualNumFloat$($(, $dim: Dim)*)?> DualNum for $struct<T$($(, $dim)*)?>
5        where
6        $($(DefaultAllocator: Allocator<$($ddim,)*>),*)?
7        {
8            type Primitive = F;
9
10            const NDERIV: usize = T::NDERIV + $nderiv;
11
12            type InnerDual = T;
13            fn from_re(inner: Self::InnerDual) -> Self {
14                Self::from_re(inner)
15            }
16
17            #[inline]
18            fn recip(&self) -> Self {
19                let rec = self.re.recip();
20                let f0 = rec.clone();
21                first!($deriv, let f1 = -f0.clone() * &rec;);
22                second!($deriv, let f2 = -f1.clone() * &rec * T::Primitive::TWO;);
23                third!($deriv, let f3 = -f2.clone() * rec * T::Primitive::THREE;);
24                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
25            }
26
27            #[inline]
28            fn powi(&self, exp: i32) -> Self {
29                match exp {
30                    0 => Self::one(),
31                    1 => self.clone(),
32                    2 => self * self,
33                    _ => {
34                        let pow3 = self.re.powi(exp - 3);
35                        let f0 = pow3.clone() * &self.re * &self.re * &self.re;
36                        first!($deriv, let f1 = pow3.clone() * &self.re * &self.re * T::Primitive::from_i32(exp).unwrap(););
37                        second!($deriv, let f2 = pow3.clone() * &self.re * T::Primitive::from_i32(exp * (exp - 1)).unwrap(););
38                        third!($deriv, let f3 = pow3 * T::Primitive::from_i32(exp * (exp - 1) * (exp - 2)).unwrap(););
39                        chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
40                    }
41                }
42            }
43
44            #[inline]
45            fn powf(&self, n: F) -> Self {
46                if n.is_zero() {
47                    Self::one()
48                } else if n.is_one() {
49                    self.clone()
50                } else if (n - F::one() - F::one()).abs() < F::epsilon() {
51                    self * self
52                } else {
53                    let n1 = n - F::one();
54                    let n2 = n1 - F::one();
55                    let n3 = n2 - F::one();
56                    let pow3 = self.re.powf(n3);
57                    let f0 = pow3.clone() * &self.re * &self.re * &self.re;
58                    first!($deriv, let f1 = pow3.clone() * &self.re * &self.re * n;);
59                    second!($deriv, let f2 = pow3.clone() * &self.re * n * n1;);
60                    third!($deriv, let f3 = pow3 * n * n1 * n2;);
61                    chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
62                }
63            }
64
65            #[inline]
66            fn sqrt(&self) -> Self {
67                first!($deriv, let rec = self.re.recip(););
68                first!($deriv, let half = T::Primitive::HALF;);
69                let f0 = self.re.sqrt();
70                first!($deriv, let f1 = f0.clone() * &rec * half;);
71                second!($deriv, let f2 = -f1.clone() * &rec * half;);
72                third!($deriv, let f3 = f2.clone() * rec * (-F::one() - half););
73                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
74            }
75
76            #[inline]
77            fn cbrt(&self) -> Self {
78                first!($deriv, let rec = self.re.recip(););
79                first!($deriv, let third = T::Primitive::THIRD;);
80                let f0 = self.re.cbrt();
81                first!($deriv, let f1 = f0.clone() * &rec * third;);
82                second!($deriv, let f2 = f1.clone() * &rec * (third - F::one()););
83                third!($deriv, let f3 = f2.clone() * rec * (third - F::one() - F::one()););
84                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
85            }
86
87
88            #[inline]
89            fn exp(&self) -> Self {
90                let f = self.re.exp();
91                chain_rule!($deriv, Self::chain_rule(self, f.clone(), f.clone(), f.clone(), f))
92            }
93
94            #[inline]
95            fn exp2(&self) -> Self {
96                first!($deriv, let ln2 = T::Primitive::TWO.ln(););
97                let f0 = self.re.exp2();
98                first!($deriv, let f1 = f0.clone() * ln2;);
99                second!($deriv, let f2 = f1.clone() * ln2;);
100                third!($deriv, let f3 = f2.clone() * ln2;);
101                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
102            }
103
104            #[inline]
105            fn exp_m1(&self) -> Self {
106                let f0 = self.re.exp_m1();
107                first!($deriv, let f1 = self.re.exp(););
108                chain_rule!($deriv, Self::chain_rule(self, f0, f1.clone(), f1.clone(), f1))
109            }
110
111            #[inline]
112            fn ln(&self) -> Self {
113                first!($deriv, let rec = self.re.recip(););
114                let f0 = self.re.ln();
115                first!($deriv, let f1 = rec.clone(););
116                second!($deriv, let f2 = -f1.clone() * &rec;);
117                third!($deriv, let f3 = -f2.clone() * rec * T::Primitive::TWO;);
118                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
119            }
120
121            #[inline]
122            fn log(&self, base: F) -> Self {
123                first!($deriv, let rec = self.re.recip(););
124                let f0 = self.re.log(base);
125                first!($deriv, let f1 = rec.clone() / base.ln(););
126                second!($deriv, let f2 = -f1.clone() * &rec;);
127                third!($deriv, let f3 = -f2.clone() * rec * T::Primitive::TWO;);
128                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
129            }
130
131            #[inline]
132            fn log2(&self) -> Self {
133                first!($deriv, let rec = self.re.recip(););
134                let f0 = self.re.log2();
135                first!($deriv, let f1 = rec.clone() / (F::one() + F::one()).ln(););
136                second!($deriv, let f2 = -f1.clone() * &rec;);
137                third!($deriv, let f3 = -f2.clone() * rec * T::Primitive::TWO;);
138                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
139            }
140
141            #[inline]
142            fn log10(&self) -> Self {
143                first!($deriv, let rec = self.re.recip(););
144                let f0 = self.re.log10();
145                first!($deriv, let f1 = rec.clone() / T::Primitive::LN_10(););
146                second!($deriv, let f2 = -f1.clone() * &rec;);
147                third!($deriv, let f3 = -f2.clone() * rec * T::Primitive::TWO;);
148                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
149            }
150
151            #[inline]
152            fn ln_1p(&self) -> Self {
153                first!($deriv, let rec = (self.re.clone() + F::one()).recip(););
154                let f0 = self.re.ln_1p();
155                first!($deriv, let f1 = rec.clone(););
156                second!($deriv, let f2 = -f1.clone() * &rec;);
157                third!($deriv, let f3 = -f2.clone() * rec * T::Primitive::TWO;);
158                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
159            }
160
161            #[inline]
162            fn sin(&self) -> Self {
163                zeroth!($deriv, let s = self.re.sin(););
164                first!($deriv, let (s, c) = self.re.sin_cos(););
165                chain_rule!($deriv, Self::chain_rule(self, s.clone(), c.clone(), -s, -c))
166            }
167
168            #[inline]
169            fn cos(&self) -> Self {
170                zeroth!($deriv, let c = self.re.cos(););
171                first!($deriv, let (s, c) = self.re.sin_cos(););
172                chain_rule!($deriv, Self::chain_rule(self, c.clone(), -s.clone(), -c, s))
173            }
174
175            #[inline]
176            fn sin_cos(&self) -> (Self, Self) {
177                let (s, c) = self.re.sin_cos();
178                (
179                    chain_rule!($deriv, Self::chain_rule(self, s.clone(), c.clone(), -s.clone(), -c.clone())),
180                    chain_rule!($deriv, Self::chain_rule(self, c.clone(), -s.clone(), -c, s)))
181            }
182
183            #[inline]
184            fn tan(&self) -> Self {
185                let (sin, cos) = self.sin_cos();
186                sin / cos
187            }
188
189            #[inline]
190            fn asin(&self) -> Self {
191                first!($deriv, let rec = (T::one() - self.re.clone() * &self.re).recip(););
192                let f0 = self.re.asin();
193                first!($deriv, let f1 = rec.sqrt(););
194                second!($deriv, let f2 = self.re.clone() * &f1 * &rec;);
195                third!($deriv, let f3 = (self.re.clone() * &self.re * (F::one() + F::one()) + F::one()) * &f1 * &rec * rec;);
196                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
197            }
198
199            #[inline]
200            fn acos(&self) -> Self {
201                first!($deriv, let rec = (T::one() - self.re.clone() * &self.re).recip(););
202                let f0 = self.re.acos();
203                first!($deriv, let f1 = -rec.sqrt(););
204                second!($deriv, let f2 = self.re.clone() * &f1 * &rec;);
205                third!($deriv, let f3 = (self.re.clone() * &self.re * (F::one() + F::one()) + F::one()) * &f1 * &rec * rec;);
206                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
207            }
208
209            #[inline]
210            fn atan(&self) -> Self {
211                first!($deriv, let rec = (T::one() + self.re.clone() * &self.re).recip(););
212                let f0 = self.re.atan();
213                first!($deriv, let f1 = rec.clone(););
214                second!($deriv, let two = F::one() + F::one(););
215                second!($deriv, let f2 = -self.re.clone() * &f1 * &rec * two;);
216                third!($deriv, let f3 = (self.re.clone() * &self.re * T::Primitive::SIX - two) * &f1 * &rec * rec;);
217                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
218            }
219
220            #[inline]
221            fn atan2(&self, other: Self) -> Self {
222                let mut res = (self / other.clone()).atan();
223                res.re = self.re.atan2(other.re);
224                res
225            }
226
227            #[inline]
228            fn sinh(&self) -> Self {
229                let s = self.re.sinh();
230                first!($deriv, let c = self.re.cosh(););
231                chain_rule!($deriv, Self::chain_rule(self, s.clone(), c.clone(), s, c))
232            }
233
234            #[inline]
235            fn cosh(&self) -> Self {
236                first!($deriv, let s = self.re.sinh(););
237                let c = self.re.cosh();
238                chain_rule!($deriv, Self::chain_rule(self, c.clone(), s.clone(), c, s))
239            }
240
241            #[inline]
242            fn tanh(&self) -> Self {
243                self.sinh() / self.cosh()
244            }
245
246            #[inline]
247            fn asinh(&self) -> Self {
248                first!($deriv, let rec = (T::one() + self.re.clone() * &self.re).recip(););
249                let f0 = self.re.asinh();
250                first!($deriv, let f1 = rec.sqrt(););
251                second!($deriv, let f2 = -self.re.clone() * &f1 * &rec;);
252                third!($deriv, let f3 = (self.re.clone() * &self.re * (F::one() + F::one()) - F::one()) * &f1 * &rec * rec;);
253                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
254            }
255
256            #[inline]
257            fn acosh(&self) -> Self {
258                first!($deriv, let rec = (self.re.clone() * &self.re - F::one()).recip(););
259                let f0 = self.re.acosh();
260                first!($deriv, let f1 = rec.sqrt(););
261                second!($deriv, let f2 = -self.re.clone() * &f1 * &rec;);
262                third!($deriv, let f3 = (self.re.clone() * &self.re * (F::one() + F::one()) + F::one()) * &f1 * &rec * rec;);
263                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
264            }
265
266            #[inline]
267            fn atanh(&self) -> Self {
268                first!($deriv, let rec = (T::one() - self.re.clone() * &self.re).recip(););
269                let f0 = self.re.atanh();
270                first!($deriv, let f1 = rec.clone(););
271                second!($deriv, let two = F::one() + F::one(););
272                second!($deriv, let f2 = self.re.clone() * &f1 * &rec * two;);
273                third!($deriv, let f3 = (self.re.clone() * &self.re * T::Primitive::SIX + two) * &f1 * &rec * rec;);
274                chain_rule!($deriv, Self::chain_rule(self, f0, f1, f2, f3))
275            }
276
277            #[inline]
278            fn sph_j0(&self) -> Self {
279                if self.re().abs() < F::epsilon() {
280                    Self::one() - self * self / T::Primitive::SIX
281                } else {
282                    self.sin() / self
283                }
284            }
285
286            #[inline]
287            fn sph_j1(&self) -> Self {
288                if self.re().abs() < F::epsilon() {
289                    self.clone() / T::Primitive::THREE
290                } else {
291                    let (s, c) = self.sin_cos();
292                    (s - self * c) / (self * self)
293                }
294            }
295
296            #[inline]
297            fn sph_j2(&self) -> Self {
298                if self.re().abs() < F::epsilon() {
299                    self * self / T::Primitive::FIFTEEN
300                } else {
301                    let (s, c) = self.sin_cos();
302                    let s2 = self * self;
303                    ((&s - self * c) * T::Primitive::THREE - &s2 * s) / (s2 * self)
304                }
305            }
306        }
307    };
308}
309
310#[macro_export]
311macro_rules! zeroth {
312    (zeroth, $($code:tt)*) => {
313        $($code)*
314    };
315    (first, $($code:tt)*) => {};
316    (second, $($code:tt)*) => {};
317    (third, $($code:tt)*) => {};
318}
319
320#[macro_export]
321macro_rules! first {
322    (zeroth, $($code:tt)*) => {};
323    (first, $($code:tt)*) => {
324         $($code)*
325    };
326    (second, $($code:tt)*) => {
327        $($code)*
328    };
329    (third, $($code:tt)*) => {
330        $($code)*
331    };
332}
333
334#[macro_export]
335macro_rules! second {
336    (zeroth, $($code:tt)*) => {};
337    (first, $($code:tt)*) => {};
338    (second, $($code:tt)*) => {
339        $($code)*
340    };
341    (third, $($code:tt)*) => {
342        $($code)*
343    };
344}
345
346#[macro_export]
347macro_rules! third {
348    (zeroth, $($code:tt)*) => {};
349    (first, $($code:tt)*) => {};
350    (second, $($code:tt)*) => {};
351    (third, $($code:tt)*) => {
352        $($code)*
353    };
354}
355
356#[macro_export]
357macro_rules! chain_rule {
358    (zeroth, Self::chain_rule($self:ident, $f0:expr, $f1:expr, $f2:expr, $f3:expr)) => {
359        Self::chain_rule($self, $f0)
360    };
361    (first, Self::chain_rule($self:ident, $f0:expr, $f1:expr, $f2:expr, $f3:expr)) => {
362        Self::chain_rule($self, $f0, $f1)
363    };
364    (second, Self::chain_rule($self:ident, $f0:expr, $f1:expr, $f2:expr, $f3:expr)) => {
365        Self::chain_rule($self, $f0, $f1, $f2)
366    };
367    (third, Self::chain_rule($self:ident, $f0:expr, $f1:expr, $f2:expr, $f3:expr)) => {
368        Self::chain_rule($self, $f0, $f1, $f2, $f3)
369    };
370}
371
372#[macro_export]
373macro_rules! impl_zeroth_derivatives {
374    ($struct:ident, [$($im:ident),*]$(, [$($dim:tt),*]$(, [$($ddim:tt),*])*)?) => {
375        impl_derivatives!(zeroth, 0, $struct, [$($im),*]$(, [$($dim),*]$(, [$($ddim),*])*)?);
376    };
377}
378
379#[macro_export]
380macro_rules! impl_first_derivatives {
381    ($struct:ident, [$($im:ident),*]$(, [$($dim:tt),*]$(, [$($ddim:tt),*])*)?) => {
382        impl_derivatives!(first, 1, $struct, [$($im),*]$(, [$($dim),*]$(, [$($ddim),*])*)?);
383    };
384}
385
386#[macro_export]
387macro_rules! impl_second_derivatives {
388    ($struct:ident, [$($im:ident),*]$(, [$($dim:tt),*]$(, [$($ddim:tt),*])*)?) => {
389        impl_derivatives!(second, 2, $struct, [$($im),*]$(, [$($dim),*]$(, [$($ddim),*])*)?);
390    };
391}
392
393#[macro_export]
394macro_rules! impl_third_derivatives {
395    ($struct:ident, [$($im:ident),*]$(, [$($dim:tt),*]$(, [$($ddim:tt),*])*)?) => {
396        impl_derivatives!(third, 3, $struct, [$($im),*]$(, [$($dim),*]$(, [$($ddim),*])*)?);
397    };
398}