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}