num_dual/datatypes/
dual3.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
9#[derive(Copy, Clone, Debug)]
11#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
12pub struct Dual3<T> {
13 pub re: T,
15 pub v1: T,
17 pub v2: T,
19 pub v3: T,
21}
22
23#[cfg(feature = "ndarray")]
24impl<T: DualNum> ndarray::ScalarOperand for Dual3<T> {}
25
26pub type Dual3_32 = Dual3<f32>;
27pub type Dual3_64 = Dual3<f64>;
28
29impl<T> Dual3<T> {
30 #[inline]
32 pub fn new(re: T, v1: T, v2: T, v3: T) -> Self {
33 Self { re, v1, v2, v3 }
34 }
35}
36
37impl<T: One + Zero> Dual3<T> {
38 #[inline]
40 pub fn from_re(re: T) -> Self {
41 Self::new(re, T::zero(), T::zero(), T::zero())
42 }
43
44 #[inline]
54 pub fn derivative(mut self) -> Self {
55 self.v1 = T::one();
56 self
57 }
58}
59
60impl<T: DualNum> Dual3<T> {
61 #[inline]
62 fn chain_rule(&self, f0: T, f1: T, f2: T, f3: T) -> Self {
63 let three = T::one() + T::one() + T::one();
64 Self::new(
65 f0,
66 f1.clone() * &self.v1,
67 f2.clone() * &self.v1 * &self.v1 + f1.clone() * &self.v2,
68 f3 * &self.v1 * &self.v1 * &self.v1 + three * f2 * &self.v1 * &self.v2 + f1 * &self.v3,
69 )
70 }
71}
72
73impl<T: DualNum> Mul<&Dual3<T>> for &Dual3<T> {
74 type Output = Dual3<T>;
75 #[inline]
76 fn mul(self, rhs: &Dual3<T>) -> Dual3<T> {
77 let two = T::one() + T::one();
78 let three = T::one() + &two;
79 Dual3::new(
80 self.re.clone() * &rhs.re,
81 self.v1.clone() * &rhs.re + self.re.clone() * &rhs.v1,
82 self.v2.clone() * &rhs.re + two * &self.v1 * &rhs.v1 + self.re.clone() * &rhs.v2,
83 self.v3.clone() * &rhs.re
84 + three * (self.v2.clone() * &rhs.v1 + self.v1.clone() * &rhs.v2)
85 + self.re.clone() * &rhs.v3,
86 )
87 }
88}
89
90impl<T: DualNum> Div<&Dual3<T>> for &Dual3<T> {
91 type Output = Dual3<T>;
92 #[inline]
93 fn div(self, rhs: &Dual3<T>) -> Dual3<T> {
94 let rec = T::one() / &rhs.re;
95 let f0 = rec.clone();
96 let f1 = -f0.clone() * &rec;
97 let f2 = -f1.clone() * &rec * T::Primitive::TWO;
98 let f3 = -f2.clone() * rec * T::Primitive::THREE;
99 self * rhs.chain_rule(f0, f1, f2, f3)
100 }
101}
102
103impl<T: fmt::Display> fmt::Display for Dual3<T> {
105 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
106 write!(
107 f,
108 "{} + {}v1 + {}v2 + {}v3",
109 self.re, self.v1, self.v2, self.v3
110 )
111 }
112}
113
114impl_third_derivatives!(Dual3, [v1, v2, v3]);
115impl_dual!(Dual3, [v1, v2, v3]);
116#[cfg(feature = "nalgebra")]
117impl_nalgebra!(Dual3, [v1, v2, v3]);