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
262pub(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
279pub(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}