Skip to main content

num_dual/datatypes/
dual.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 dual number for the calculations of first derivatives.
12#[derive(Copy, Clone, Debug)]
13#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
14pub struct Dual<T> {
15    /// Real part of the dual number
16    pub re: T,
17    /// Derivative part of the dual number
18    pub eps: T,
19}
20
21#[cfg(feature = "ndarray")]
22impl<T: DualNum> ndarray::ScalarOperand for Dual<T> {}
23
24pub type Dual32 = Dual<f32>;
25pub type Dual64 = Dual<f64>;
26
27impl<T> Dual<T> {
28    /// Create a new dual number from its fields.
29    #[inline]
30    pub fn new(re: T, eps: T) -> Self {
31        Self { re, eps }
32    }
33}
34
35impl<T: Zero> Dual<T> {
36    /// Create a new dual number from the real part.
37    #[inline]
38    pub fn from_re(re: T) -> Self {
39        Self::new(re, T::zero())
40    }
41}
42
43impl<T: One> Dual<T> {
44    /// Set the derivative part to 1.
45    /// ```
46    /// # use num_dual::{Dual64, DualNum};
47    /// let x = Dual64::from_re(5.0).derivative().powi(2);
48    /// assert_eq!(x.re, 25.0);
49    /// assert_eq!(x.eps, 10.0);
50    /// ```
51    #[inline]
52    pub fn derivative(mut self) -> Self {
53        self.eps = T::one();
54        self
55    }
56}
57
58/* chain rule */
59impl<T: DualNum> Dual<T> {
60    #[inline]
61    fn chain_rule(&self, f0: T, f1: T) -> Self {
62        Self::new(f0, self.eps.clone() * f1)
63    }
64}
65
66/* product rule */
67impl<T: DualNum> Mul<&Dual<T>> for &Dual<T> {
68    type Output = Dual<T>;
69    #[inline]
70    fn mul(self, other: &Dual<T>) -> Self::Output {
71        Dual::new(
72            self.re.clone() * other.re.clone(),
73            self.eps.clone() * other.re.clone() + other.eps.clone() * self.re.clone(),
74        )
75    }
76}
77
78/* quotient rule */
79impl<T: DualNum> Div<&Dual<T>> for &Dual<T> {
80    type Output = Dual<T>;
81    #[inline]
82    fn div(self, other: &Dual<T>) -> Dual<T> {
83        let inv = other.re.recip();
84        Dual::new(
85            self.re.clone() * inv.clone(),
86            (self.eps.clone() * other.re.clone() - other.eps.clone() * self.re.clone())
87                * inv.clone()
88                * inv,
89        )
90    }
91}
92
93/* string conversions */
94impl<T: DualNum> fmt::Display for Dual<T> {
95    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
96        write!(f, "{} + {}ε", self.re, self.eps)
97    }
98}
99
100impl_first_derivatives!(Dual, [eps]);
101impl_dual!(Dual, [eps]);
102#[cfg(feature = "nalgebra")]
103impl_nalgebra!(Dual, [eps]);