Skip to main content

conspire/math/tensor/rank_2/exponential/
mod.rs

1#[cfg(test)]
2mod test;
3
4use crate::math::Quantity;
5use crate::math::assert::Assert;
6use crate::units::Dimensionless;
7
8use super::{
9    super::{
10        Rank2, Tensor, TensorArray, TensorError, rank_0::list::TensorRank0List, rank_4::TensorRank4,
11    },
12    TensorRank2,
13    eigen::{find_orthonormal_eigenvectors, reconstruct_symmetric, solve_cubic_symmetric},
14};
15
16/// Whether the eigenvalues from the cubic are far enough apart for the spectral
17/// decomposition to be accurate: the roots of a nearly repeated pair are ill-conditioned,
18/// and the error they carry (about a part in 10^9 for a gap of 10^-11) passes straight
19/// into the exponential.
20fn well_separated(eigenvalues: &TensorRank0List<3>, norm: f64) -> bool {
21    let gap = |i: usize, j: usize| (eigenvalues[i] - eigenvalues[j]).abs();
22    gap(0, 1).min(gap(0, 2)).min(gap(1, 2)) >= 1e-2 * (1.0 + norm)
23}
24
25impl<I> TensorRank2<3, I, I, Dimensionless> {
26    /// Returns the matrix exponential of the 3x3 tensor.
27    ///
28    /// Diagonal tensors go entrywise; symmetric tensors (exactly or up to
29    /// round-off) with well-separated eigenvalues through the spectral decomposition;
30    /// anything with a small enough norm through a truncated Taylor series; and
31    /// everything else, including symmetric tensors with a nearly repeated
32    /// eigenvalue, through scaling and squaring of that series.
33    pub fn expm(&self) -> Result<Self, TensorError> {
34        if self.is_diagonal() {
35            let mut expm = TensorRank2::zero();
36            expm.iter_mut()
37                .enumerate()
38                .zip(self.iter())
39                .for_each(|((i, expm_i), self_i)| expm_i[i] = self_i[i].exp());
40            return Ok(expm);
41        }
42        let norm = self.norm().value();
43        if norm < 1e-2 {
44            return Ok(self.expm_series());
45        }
46        let transpose = self.transpose();
47        if self.is_symmetric() || (self - &transpose).norm().value() < 1e-9 * (1.0 + norm) {
48            let symmetric = (self + transpose) * 0.5;
49            let mut eigenvalues = solve_cubic_symmetric(symmetric.invariants())?;
50            if well_separated(&eigenvalues, norm) {
51                let eigenvectors = find_orthonormal_eigenvectors(&eigenvalues, &symmetric);
52                eigenvalues
53                    .iter_mut()
54                    .for_each(|eigenvalue| *eigenvalue = eigenvalue.exp());
55                return Ok(reconstruct_symmetric(eigenvalues, eigenvectors));
56            }
57        }
58        let squarings = (norm / 5e-3).log2().ceil().max(1.0) as u32;
59        let mut expm = (self / 2.0_f64.powi(squarings as i32)).expm_series();
60        (0..squarings).for_each(|_| expm = &expm * &expm);
61        Ok(expm)
62    }
63    /// The truncated Taylor series `Σ Aᵏ/k!`; accurate only for a small norm.
64    fn expm_series(&self) -> Self {
65        let num_terms = match self.norm().value() {
66            norm if norm < 1e-4 => 3,
67            norm if norm < 1e-3 => 5,
68            _ => 8,
69        };
70        let mut expm = self + TensorRank2::identity();
71        let mut power = self.clone();
72        let mut factorial = 1.0;
73        (2..=num_terms).for_each(|k| {
74            power *= self;
75            factorial *= k as f64;
76            expm += &power / factorial;
77        });
78        expm
79    }
80    /// Returns the derivative of the matrix exponential of the 3x3 tensor.
81    ///
82    /// The Frechet derivative $`\mathrm{d}\exp(\mathbf{A})/\mathrm{d}\mathbf{A}`$, formed
83    /// diagonally entrywise, from a truncated series near zero, from scaling and
84    /// squaring of that series for a general (non-symmetric, larger-norm) tensor or a
85    /// symmetric one with a nearly repeated eigenvalue, and otherwise, for a symmetric
86    /// tensor with well-separated eigenvalues, from the spectral decomposition with the
87    /// divided differences
88    /// ```math
89    /// \frac{e^{\lambda_i} - e^{\lambda_j}}{\lambda_i - \lambda_j},
90    /// \qquad e^{\lambda_j} \text{ for } \lambda_i = \lambda_j.
91    /// ```
92    pub fn dexpm(&self) -> Result<TensorRank4<3, I, I, I, I, Dimensionless>, TensorError> {
93        if self.is_diagonal() {
94            let mut dexpm = TensorRank4::zero();
95            dexpm.iter_mut().enumerate().for_each(|(i, dexpm_i)| {
96                dexpm_i.iter_mut().enumerate().for_each(|(j, dexpm_ij)| {
97                    dexpm_ij.iter_mut().enumerate().for_each(|(k, dexpm_ijk)| {
98                        dexpm_ijk
99                            .iter_mut()
100                            .enumerate()
101                            .filter(|(l, _)| i == k && &j == l)
102                            .for_each(|(_, dexpm_ijkl)| {
103                                *dexpm_ijkl = if Assert::default()
104                                    .eq_within_tols(self[i][i], &self[j][j])
105                                    .is_ok()
106                                {
107                                    self[j][j].exp()
108                                } else {
109                                    (self[i][i].exp() - self[j][j].exp())
110                                        / (self[i][i] - self[j][j])
111                                }
112                            })
113                    })
114                })
115            });
116            Ok(dexpm)
117        } else {
118            let norm = self.norm().value();
119            if norm < 1e-2 {
120                //
121                // d(A^n)[H] = sum_{p=0}^{n-1} A^p . H . A^{n-1-p}, so the truncated series
122                // gives dexpm_{ijkl} = sum_n (1/n!) sum_p (A^p)_{ik} (A^{n-1-p})_{lj}.
123                //
124                let num_terms = if norm < 1e-4 {
125                    3
126                } else if norm < 1e-3 {
127                    5
128                } else {
129                    8
130                };
131                let mut power = Self::identity();
132                let mut powers = vec![power.clone()];
133                (1..num_terms).for_each(|_| {
134                    power *= self;
135                    powers.push(power.clone())
136                });
137                let mut dexpm = TensorRank4::zero();
138                let mut factorial = 1.0;
139                for n in 1..=num_terms {
140                    factorial *= n as f64;
141                    for p in 0..n {
142                        let (left, right) = (&powers[p], &powers[n - 1 - p]);
143                        for i in 0..3 {
144                            for j in 0..3 {
145                                for k in 0..3 {
146                                    for l in 0..3 {
147                                        dexpm[i][j][k][l] += Quantity::new(
148                                            left[i][k].value() * right[l][j].value() / factorial,
149                                        )
150                                    }
151                                }
152                            }
153                        }
154                    }
155                }
156                Ok(dexpm)
157            } else {
158                let transpose = self.transpose();
159                let nearly_symmetric =
160                    self.is_symmetric() || (self - &transpose).norm().value() < 1e-9 * (1.0 + norm);
161                let spectral = if nearly_symmetric {
162                    let symmetric = (self + transpose.clone()) * 0.5;
163                    let eigenvalues = solve_cubic_symmetric(symmetric.invariants())?;
164                    well_separated(&eigenvalues, norm).then_some((symmetric, eigenvalues))
165                } else {
166                    None
167                };
168                let Some((symmetric, eigenvalues)) = spectral else {
169                    //
170                    // Non-symmetric, or symmetric with a nearly repeated eigenvalue (whose
171                    // divided differences would lose accuracy): scaling and squaring of the
172                    // Fréchet derivative.
173                    // With E = exp(B), L = dexp(B), the squaring B → 2B gives
174                    // E → E² and L → L·E + E·L (contracting the middle index);
175                    // one final 1/scale converts d/dB back to d/dA.
176                    //
177                    let squarings = (norm / 5e-3).log2().ceil().max(1.0) as u32;
178                    let scale = 2.0_f64.powi(squarings as i32);
179                    let mut expm = (self / scale).expm_series();
180                    let mut dexpm = (self / scale).dexpm()?;
181                    for _ in 0..squarings {
182                        let mut next = TensorRank4::zero();
183                        for i in 0..3 {
184                            for j in 0..3 {
185                                for k in 0..3 {
186                                    for l in 0..3 {
187                                        let mut value = 0.0;
188                                        for p in 0..3 {
189                                            value += dexpm[i][p][k][l].value() * expm[p][j].value()
190                                                + expm[i][p].value() * dexpm[p][j][k][l].value();
191                                        }
192                                        next[i][j][k][l] = Quantity::new(value);
193                                    }
194                                }
195                            }
196                        }
197                        dexpm = next;
198                        expm = &expm * &expm;
199                    }
200                    dexpm.iter_mut().for_each(|dexpm_i| {
201                        dexpm_i.iter_mut().for_each(|dexpm_ij| {
202                            dexpm_ij.iter_mut().for_each(|dexpm_ijk| {
203                                dexpm_ijk
204                                    .iter_mut()
205                                    .for_each(|dexpm_ijkl| *dexpm_ijkl /= scale)
206                            })
207                        })
208                    });
209                    return Ok(dexpm);
210                };
211                let divided_difference: Self = eigenvalues
212                    .iter()
213                    .map(|eigenvalue_i| {
214                        eigenvalues
215                            .iter()
216                            .map(|eigenvalue_j| {
217                                if Assert::default()
218                                    .eq_within_tols(eigenvalue_i, eigenvalue_j)
219                                    .is_ok()
220                                {
221                                    eigenvalue_j.exp()
222                                } else {
223                                    (eigenvalue_i.exp() - eigenvalue_j.exp())
224                                        / (eigenvalue_i - eigenvalue_j)
225                                }
226                            })
227                            .collect()
228                    })
229                    .collect();
230                let eigenvectors =
231                    find_orthonormal_eigenvectors(&eigenvalues, &symmetric).transpose();
232                Ok(eigenvectors.iter().map(|eigenvector_i|
233                    eigenvectors.iter().map(|eigenvector_j|
234                        eigenvectors.iter().map(|eigenvector_k|
235                            eigenvectors.iter().map(|eigenvector_l|
236                                eigenvector_i.iter().zip(eigenvector_k.iter().zip(divided_difference.iter())).map(|(eigenvector_ip, (eigenvector_kp, divided_difference_p))|
237                                    eigenvector_j.iter().zip(eigenvector_l.iter().zip(divided_difference_p.iter())).map(|(eigenvector_jq, (eigenvector_lq, divided_difference_pq))|
238                                        eigenvector_ip * eigenvector_kp * divided_difference_pq * eigenvector_jq * eigenvector_lq
239                                    ).sum::<Quantity>()
240                                ).sum::<Quantity>()
241                            ).collect()
242                        ).collect()
243                    ).collect()
244                ).collect())
245            }
246        }
247    }
248    /// Applies the inverse matrix-exponential Fréchet derivative at `self` (the
249    /// algebra element `σ`) to `rate`.
250    ///
251    /// The Bernoulli commutator series, truncated at four terms (exact to fifth
252    /// order); with `\mathrm{ad}_\sigma(A) = \sigma A - A\sigma`,
253    /// ```math
254    /// \mathrm{dexpinv}_\sigma(A) = A - \tfrac{1}{2}[\sigma, A]
255    ///     + \tfrac{1}{12}[\sigma, [\sigma, A]]
256    ///     - \tfrac{1}{720}[\sigma, [\sigma, [\sigma, [\sigma, A]]]] .
257    /// ```
258    /// Pure matrix products — total on any input, no symmetry needed.
259    pub fn dexpinv(&self, rate: &Self) -> Self {
260        let mut term = rate.clone();
261        let mut result = term.clone() * DEXPINV_COEFFICIENTS[0];
262        for &coefficient in DEXPINV_COEFFICIENTS.iter().skip(1) {
263            term = self * &term - &term * self;
264            if coefficient != 0.0 {
265                result += term.clone() * coefficient;
266            }
267        }
268        result
269    }
270    /// The directional derivative of [`Self::dexpinv`] in the direction
271    /// `(d_sigma, d_rate)`.
272    ///
273    /// Forward-mode through the same truncated series: with `T_0 = A` and
274    /// `T_k = [\sigma, T_{k-1}]`,
275    /// ```math
276    /// \mathrm{d}T_0 = \mathrm{d}A, \qquad
277    /// \mathrm{d}T_k = [\mathrm{d}\sigma, T_{k-1}] + [\sigma, \mathrm{d}T_{k-1}] .
278    /// ```
279    pub fn dexpinv_tangent(&self, rate: &Self, d_sigma: &Self, d_rate: &Self) -> Self {
280        let mut term = rate.clone();
281        let mut d_term = d_rate.clone();
282        let mut result = d_term.clone() * DEXPINV_COEFFICIENTS[0];
283        for &coefficient in DEXPINV_COEFFICIENTS.iter().skip(1) {
284            d_term = d_sigma * &term - &term * d_sigma + (self * &d_term - &d_term * self);
285            term = self * &term - &term * self;
286            if coefficient != 0.0 {
287                result += d_term.clone() * coefficient;
288            }
289        }
290        result
291    }
292}
293
294/// The Bernoulli numbers `Bₖ/k!` of the `dexpinv` series, truncated at four terms.
295const DEXPINV_COEFFICIENTS: [f64; 5] = [1.0, -0.5, 1.0 / 12.0, 0.0, -1.0 / 720.0];