Skip to main content

gammalooprs/utils/
hyperdual_utils.rs

1use std::fmt::{Display, LowerExp};
2use std::ops::AddAssign;
3
4use crate::cff::CutCFFIndex;
5use crate::utils::{F, FloatLike, PrecisionUpgradable};
6use itertools::{Itertools, iproduct};
7use spenso::algebra::{algebraic_traits::RefZero, complex::Complex};
8use symbolica::domains::dual::{DualNumberStructure, HyperDual};
9
10pub(crate) fn new_constant<T: Clone + RefZero>(shape: &HyperDual<T>, value: &T) -> HyperDual<T> {
11    let mut new = shape.clone();
12    new.values[0] = value.clone();
13    let new_values_iter = new.values.iter_mut().skip(1);
14    for v in new_values_iter {
15        *v = value.ref_zero();
16    }
17    new
18}
19
20pub(crate) fn new_from_values<T: Clone>(shape: &HyperDual<T>, values: &[T]) -> HyperDual<T> {
21    let mut new = shape.clone();
22    new.values = values.to_vec();
23    new
24}
25
26pub(crate) fn simple_n_deriv_shape(num_derivatives: usize) -> Vec<Vec<usize>> {
27    (0..=num_derivatives).map(|order| vec![order]).collect()
28}
29
30#[derive(Clone, Copy, Debug, PartialEq, Eq)]
31pub(crate) struct CutCFFVariableIndices {
32    pub lu_cut: Option<usize>,
33    pub left_threshold: Option<usize>,
34    pub right_threshold: Option<usize>,
35}
36
37pub(crate) fn variable_indices_from_cut_cff_index(
38    cut_cff_index: &CutCFFIndex,
39) -> CutCFFVariableIndices {
40    let mut next_index = 0;
41
42    let lu_cut = cut_cff_index
43        .lu_cut_order
44        .filter(|order| *order > 1)
45        .map(|_| {
46            let index = next_index;
47            next_index += 1;
48            index
49        });
50    let left_threshold = cut_cff_index
51        .left_threshold_order
52        .filter(|order| *order > 1)
53        .map(|_| {
54            let index = next_index;
55            next_index += 1;
56            index
57        });
58    let right_threshold = cut_cff_index
59        .right_threshold_order
60        .filter(|order| *order > 1)
61        .map(|_| {
62            let index = next_index;
63            next_index += 1;
64            index
65        });
66
67    CutCFFVariableIndices {
68        lu_cut,
69        left_threshold,
70        right_threshold,
71    }
72}
73
74pub(crate) fn shape_from_cut_cff_index(cut_cff_index: &CutCFFIndex) -> Option<Vec<Vec<usize>>> {
75    let max_derivative_shape = {
76        let mut max_derivative_shape = Vec::new();
77
78        if let Some(lu_cut_order) = cut_cff_index.lu_cut_order
79            && lu_cut_order > 1
80        {
81            max_derivative_shape.push(lu_cut_order - 1);
82        }
83
84        if let Some(left_th_order) = cut_cff_index.left_threshold_order
85            && left_th_order > 1
86        {
87            max_derivative_shape.push(left_th_order - 1);
88        }
89        if let Some(right_th_order) = cut_cff_index.right_threshold_order
90            && right_th_order > 1
91        {
92            max_derivative_shape.push(right_th_order - 1);
93        }
94
95        max_derivative_shape
96    };
97
98    if max_derivative_shape.is_empty() {
99        None
100    } else if max_derivative_shape.len() == 1 {
101        Some(
102            (0..=max_derivative_shape[0])
103                .map(|order| vec![order])
104                .collect(),
105        )
106    } else if max_derivative_shape.len() == 2 {
107        let mut result = iproduct!(0..=max_derivative_shape[0], 0..=max_derivative_shape[1])
108            .map(|(order1, order2)| vec![order1, order2])
109            .sorted()
110            .collect_vec();
111
112        result.sort_by(|a, b| {
113            let sum_a: usize = a.iter().sum();
114            let sum_b: usize = b.iter().sum();
115            sum_a.cmp(&sum_b)
116        });
117
118        result[1..3].sort_by(|a, b| b[0].cmp(&a[0]).then(b[1].cmp(&a[1])));
119
120        Some(result)
121    } else if max_derivative_shape.len() == 3 {
122        let mut result = iproduct!(
123            0..=max_derivative_shape[0],
124            0..=max_derivative_shape[1],
125            0..=max_derivative_shape[2]
126        )
127        .map(|(order1, order2, order3)| vec![order1, order2, order3])
128        .sorted()
129        .collect_vec();
130
131        result.sort_by(|a, b| {
132            let sum_a: usize = a.iter().sum();
133            let sum_b: usize = b.iter().sum();
134            sum_a.cmp(&sum_b)
135        });
136
137        result[1..4].sort_by(|a, b| b[0].cmp(&a[0]).then(b[1].cmp(&a[1])).then(b[2].cmp(&a[2])));
138
139        Some(result)
140    } else {
141        unreachable!("shape_from_cut_cff_index only supports up to 3 derivative orders")
142    }
143}
144
145impl<T> PrecisionUpgradable for HyperDual<T>
146where
147    T: PrecisionUpgradable,
148    T::Higher: Default,
149    T::Lower: Default,
150{
151    type Higher = HyperDual<T::Higher>;
152    type Lower = HyperDual<T::Lower>;
153
154    fn higher(&self) -> Self::Higher {
155        let shape = self.get_shape();
156        let owned_shape = shape
157            .into_iter()
158            .map(|slice| slice.to_vec())
159            .collect::<Vec<_>>();
160
161        let mut new_hyperdual = HyperDual::<T::Higher>::new(owned_shape);
162
163        new_hyperdual
164            .values
165            .iter_mut()
166            .zip(self.values.iter())
167            .for_each(|(new_v, old_v)| {
168                *new_v = old_v.higher();
169            });
170
171        new_hyperdual
172    }
173
174    fn lower(&self) -> Self::Lower {
175        let shape = self.get_shape();
176        let owned_shape = shape
177            .into_iter()
178            .map(|slice| slice.to_vec())
179            .collect::<Vec<_>>();
180
181        let mut new_hyperdual = HyperDual::<T::Lower>::new(owned_shape);
182
183        new_hyperdual
184            .values
185            .iter_mut()
186            .zip(self.values.iter())
187            .for_each(|(new_v, old_v)| {
188                *new_v = old_v.lower();
189            });
190
191        new_hyperdual
192    }
193}
194
195#[derive(Clone, Debug)]
196pub enum DualOrNot<T> {
197    Dual(HyperDual<T>),
198    NonDual(T),
199}
200
201impl<T> DualOrNot<T> {
202    pub fn unwrap_real(self) -> T {
203        match self {
204            DualOrNot::Dual(_dual) => panic!("Cannot unwrap real value from Dual variant"),
205            DualOrNot::NonDual(value) => value,
206        }
207    }
208}
209
210impl<T: LowerExp> Display for DualOrNot<T> {
211    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
212        match self {
213            DualOrNot::Dual(dual) => write!(
214                f,
215                "[{}]",
216                dual.values
217                    .iter()
218                    .map(|v| format!("{:+16e}", v))
219                    .collect::<Vec<_>>()
220                    .join(", ")
221            ),
222            DualOrNot::NonDual(value) => write!(f, "{:+16e}", value),
223        }
224    }
225}
226
227impl<T> AddAssign for DualOrNot<T>
228where
229    T: AddAssign<T>,
230{
231    fn add_assign(&mut self, rhs: Self) {
232        match (self, rhs) {
233            (DualOrNot::Dual(x), DualOrNot::Dual(y)) => {
234                assert_eq!(x.values.len(), y.values.len());
235                for (lhs, rhs) in x.values.iter_mut().zip(y.values) {
236                    *lhs += rhs;
237                }
238            }
239            (DualOrNot::NonDual(x), DualOrNot::NonDual(y)) => {
240                *x += y;
241            }
242            _ => panic!("Cannot add DualOrNot of different types"),
243        }
244    }
245}
246
247impl<T: Clone> DualOrNot<T> {
248    pub fn new_from_slice(shape: &Option<HyperDual<T>>, values: &[T]) -> Self {
249        match shape {
250            Some(dual_shape) => {
251                let hyperdual = new_from_values(dual_shape, values);
252                DualOrNot::Dual(hyperdual)
253            }
254            None => {
255                assert_eq!(values.len(), 1);
256                DualOrNot::NonDual(values[0].clone())
257            }
258        }
259    }
260}
261
262// this function assumes that the HyperDual has the correct shape for t-derivatives
263pub(crate) fn extract_t_derivatives_complex<T: FloatLike>(
264    dual: HyperDual<Complex<F<T>>>,
265) -> Vec<Complex<F<T>>> {
266    let mut result = Vec::with_capacity(dual.values.len());
267    let mut n_factorial = Complex::new_re(dual.values[0].re.one());
268
269    for (order, value) in dual.values.iter().enumerate() {
270        if order > 0 {
271            n_factorial = &n_factorial * Complex::new_re(value.re.from_usize(order));
272        }
273        result.push(value.clone() * &n_factorial);
274    }
275
276    result
277}
278
279// this function assumes that the HyperDual has the correct shape for t-derivatives
280pub(crate) fn extract_t_derivatives<T: FloatLike>(dual: HyperDual<F<T>>) -> Vec<F<T>> {
281    let mut result = Vec::with_capacity(dual.values.len());
282    let mut n_factorial = dual.values[0].one();
283
284    for (order, value) in dual.values.iter().enumerate() {
285        if order > 0 {
286            n_factorial = &n_factorial * value.from_usize(order);
287        }
288        result.push(value.clone() * &n_factorial);
289    }
290
291    result
292}
293
294pub(crate) fn dualize_dual_t_to_dual_r_t<T: FloatLike>(
295    t_dual: HyperDual<F<T>>,
296    target_shape: HyperDual<F<T>>,
297    variable: usize,
298) -> HyperDual<F<T>> {
299    let mut new_dual = new_constant(&target_shape, &t_dual.values[0]);
300    let n_variables_of_target_shape = target_shape.get_shape()[0].len();
301
302    debug_assert!(variable < n_variables_of_target_shape);
303
304    for (i, value) in t_dual.values.iter().enumerate().skip(1) {
305        let dual_shape_to_find = {
306            let mut shape = vec![0; n_variables_of_target_shape];
307            shape[variable] = i;
308            shape
309        };
310
311        #[allow(clippy::expect_fun_call)]
312        let index_of_derivative = target_shape
313            .get_shape()
314            .iter()
315            .position(|shape| shape == &dual_shape_to_find)
316            .expect(&format!(
317                "Could not find derivative shape: {:?} in {:?}",
318                dual_shape_to_find,
319                target_shape.get_shape()
320            ));
321
322        new_dual.values[index_of_derivative] = value.clone();
323    }
324
325    new_dual
326}
327
328pub(crate) fn extract_coefficient_t_duals<T: Clone + RefZero + Default>(
329    dual: &HyperDual<T>,
330    t_variable: usize,
331) -> (Vec<Vec<usize>>, Vec<HyperDual<T>>) {
332    let mixed_shape = dual
333        .get_shape()
334        .into_iter()
335        .map(|orders| orders.to_vec())
336        .collect_vec();
337    let n_variables = mixed_shape.first().map_or(0, Vec::len);
338
339    debug_assert!(t_variable < n_variables);
340
341    let max_t_order = mixed_shape
342        .iter()
343        .map(|orders| orders[t_variable])
344        .max()
345        .unwrap_or(0);
346    let t_dual_shape = HyperDual::new(simple_n_deriv_shape(max_t_order));
347
348    let mut coefficient_orders = Vec::<Vec<usize>>::new();
349    let mut coefficient_duals = Vec::<HyperDual<T>>::new();
350
351    for (orders, value) in mixed_shape.iter().zip(dual.values.iter()) {
352        let t_order = orders[t_variable];
353        let projected_orders = orders
354            .iter()
355            .enumerate()
356            .filter_map(|(index, order)| (index != t_variable).then_some(*order))
357            .collect_vec();
358
359        let coefficient_index = if let Some(existing_index) = coefficient_orders
360            .iter()
361            .position(|existing_orders| existing_orders == &projected_orders)
362        {
363            existing_index
364        } else {
365            coefficient_orders.push(projected_orders);
366            coefficient_duals.push(new_constant(&t_dual_shape, &value.ref_zero()));
367            coefficient_duals.len() - 1
368        };
369
370        coefficient_duals[coefficient_index].values[t_order] = value.clone();
371    }
372
373    (coefficient_orders, coefficient_duals)
374}
375
376#[cfg(test)]
377mod tests {
378    use super::*;
379
380    #[test]
381    fn dualize_dual_t_to_dual_r_t_uses_requested_target_variable() {
382        let t_dual = HyperDual::new(simple_n_deriv_shape(1));
383        let t_dual = new_from_values(&t_dual, &[F(3.0_f64), F(5.0_f64)]);
384
385        let target_shape = HyperDual::new(
386            shape_from_cut_cff_index(&CutCFFIndex {
387                left_threshold_order: Some(2),
388                right_threshold_order: None,
389                lu_cut_order: Some(2),
390            })
391            .unwrap(),
392        );
393
394        let dualized = dualize_dual_t_to_dual_r_t(t_dual, target_shape, 1);
395
396        assert_eq!(
397            dualized.values,
398            vec![F(3.0_f64), F(0.0_f64), F(5.0_f64), F(0.0_f64)]
399        );
400    }
401
402    #[test]
403    fn dualize_dual_t_to_dual_r_t_keeps_constant_inputs_constant() {
404        let t_dual = HyperDual::new(simple_n_deriv_shape(0));
405        let t_dual = new_from_values(&t_dual, &[F(7.0_f64)]);
406
407        let target_shape = HyperDual::new(
408            shape_from_cut_cff_index(&CutCFFIndex {
409                left_threshold_order: Some(2),
410                right_threshold_order: None,
411                lu_cut_order: Some(2),
412            })
413            .unwrap(),
414        );
415
416        let dualized = dualize_dual_t_to_dual_r_t(t_dual, target_shape, 1);
417
418        assert_eq!(
419            dualized.values,
420            vec![F(7.0_f64), F(0.0_f64), F(0.0_f64), F(0.0_f64)]
421        );
422    }
423
424    #[test]
425    fn extract_coefficient_t_duals_factorizes_single_threshold_mixed_dual() {
426        let mixed_shape = HyperDual::new(
427            shape_from_cut_cff_index(&CutCFFIndex {
428                left_threshold_order: Some(2),
429                right_threshold_order: None,
430                lu_cut_order: Some(2),
431            })
432            .unwrap(),
433        );
434        let mixed_dual = new_from_values(
435            &mixed_shape,
436            &[F(2.0_f64), F(3.0_f64), F(5.0_f64), F(7.0_f64)],
437        );
438
439        let (coefficient_orders, coefficient_duals) = extract_coefficient_t_duals(&mixed_dual, 0);
440
441        assert_eq!(coefficient_orders, vec![vec![0], vec![1]]);
442        assert_eq!(coefficient_duals[0].values, vec![F(2.0_f64), F(3.0_f64)]);
443        assert_eq!(coefficient_duals[1].values, vec![F(5.0_f64), F(7.0_f64)]);
444    }
445
446    #[test]
447    fn extract_coefficient_t_duals_factorizes_iterated_mixed_dual() {
448        let mixed_shape = HyperDual::new(
449            shape_from_cut_cff_index(&CutCFFIndex {
450                left_threshold_order: Some(2),
451                right_threshold_order: Some(2),
452                lu_cut_order: Some(2),
453            })
454            .unwrap(),
455        );
456        let mixed_dual = new_from_values(
457            &mixed_shape,
458            &[
459                F(1.0_f64),
460                F(2.0_f64),
461                F(3.0_f64),
462                F(4.0_f64),
463                F(5.0_f64),
464                F(6.0_f64),
465                F(7.0_f64),
466                F(8.0_f64),
467            ],
468        );
469
470        let (coefficient_orders, coefficient_duals) = extract_coefficient_t_duals(&mixed_dual, 0);
471
472        assert_eq!(
473            coefficient_orders,
474            vec![vec![0, 0], vec![1, 0], vec![0, 1], vec![1, 1]]
475        );
476        assert_eq!(coefficient_duals[0].values, vec![F(1.0_f64), F(2.0_f64)]);
477        assert_eq!(coefficient_duals[1].values, vec![F(3.0_f64), F(7.0_f64)]);
478        assert_eq!(coefficient_duals[2].values, vec![F(4.0_f64), F(6.0_f64)]);
479        assert_eq!(coefficient_duals[3].values, vec![F(5.0_f64), F(8.0_f64)]);
480    }
481
482    #[test]
483    fn shape_from_cut_cff_index_preserves_canonical_t_left_right_order() {
484        let three_axis_shape = shape_from_cut_cff_index(&CutCFFIndex {
485            left_threshold_order: Some(3),
486            right_threshold_order: Some(2),
487            lu_cut_order: Some(2),
488        })
489        .unwrap();
490
491        assert_eq!(three_axis_shape[0], vec![0, 0, 0]);
492        assert_eq!(three_axis_shape[1], vec![1, 0, 0]);
493        assert_eq!(three_axis_shape[2], vec![0, 1, 0]);
494        assert_eq!(three_axis_shape[3], vec![0, 0, 1]);
495        assert!(three_axis_shape.contains(&vec![0, 2, 1]));
496
497        assert_eq!(
498            shape_from_cut_cff_index(&CutCFFIndex {
499                left_threshold_order: Some(1),
500                right_threshold_order: Some(2),
501                lu_cut_order: Some(2),
502            }),
503            Some(vec![vec![0, 0], vec![1, 0], vec![0, 1], vec![1, 1]])
504        );
505
506        assert_eq!(
507            shape_from_cut_cff_index(&CutCFFIndex {
508                left_threshold_order: Some(2),
509                right_threshold_order: Some(1),
510                lu_cut_order: Some(1),
511            }),
512            Some(vec![vec![0], vec![1]])
513        );
514    }
515
516    #[test]
517    fn variable_indices_from_cut_cff_index_tracks_active_axes_without_reordering() {
518        assert_eq!(
519            variable_indices_from_cut_cff_index(&CutCFFIndex {
520                left_threshold_order: Some(2),
521                right_threshold_order: Some(2),
522                lu_cut_order: Some(2),
523            }),
524            CutCFFVariableIndices {
525                lu_cut: Some(0),
526                left_threshold: Some(1),
527                right_threshold: Some(2),
528            }
529        );
530
531        assert_eq!(
532            variable_indices_from_cut_cff_index(&CutCFFIndex {
533                left_threshold_order: Some(2),
534                right_threshold_order: Some(2),
535                lu_cut_order: Some(1),
536            }),
537            CutCFFVariableIndices {
538                lu_cut: None,
539                left_threshold: Some(0),
540                right_threshold: Some(1),
541            }
542        );
543
544        assert_eq!(
545            variable_indices_from_cut_cff_index(&CutCFFIndex {
546                left_threshold_order: Some(1),
547                right_threshold_order: Some(2),
548                lu_cut_order: Some(2),
549            }),
550            CutCFFVariableIndices {
551                lu_cut: Some(0),
552                left_threshold: None,
553                right_threshold: Some(1),
554            }
555        );
556    }
557}