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];