Skip to main content

ad_trait/
differentiable_function.rs

1use crate::forward_ad::adfn::adfn;
2use crate::forward_ad::ForwardADTrait;
3#[cfg(feature = "std")]
4use crate::reverse_ad::adr::{adr, GlobalComputationGraph};
5use crate::AD;
6use alloc::rc::Rc;
7use alloc::sync::Arc;
8use alloc::vec::Vec;
9use alloc::{format, vec};
10use core::marker::PhantomData;
11use nalgebra::{DMatrix, DVector};
12#[cfg(feature = "std")]
13use rand::distributions::Distribution;
14#[cfg(feature = "std")]
15use rand::distributions::Uniform;
16#[cfg(feature = "std")]
17use rand::thread_rng;
18#[cfg(feature = "std")]
19use rand::Rng;
20#[cfg(feature = "std")]
21use std::sync::{Mutex, RwLock};
22
23#[cfg(feature = "nightly")]
24use crate::simd::f64xn::f64xn;
25
26/// A trait for types that can be "reparameterized" with a different `AD` type.
27///
28/// This is a critical feature for automatic differentiation, as it allows a function
29/// originally defined for `f64` to be converted into a version that uses `adr` or `adfn`
30/// for derivative tracking.
31pub trait Reparameterize {
32    /// The type of the function after reparameterization with type `T2`.
33    type SelfType<T2: AD>: DifferentiableFunctionTrait<T2>;
34}
35
36impl<R: Reparameterize> Reparameterize for Rc<R> {
37    type SelfType<T2: AD> = R::SelfType<T2>;
38}
39impl<R: Reparameterize> Reparameterize for Arc<R> {
40    type SelfType<T2: AD> = R::SelfType<T2>;
41}
42#[cfg(feature = "std")]
43impl<R: Reparameterize> Reparameterize for Mutex<R> {
44    type SelfType<T2: AD> = R::SelfType<T2>;
45}
46#[cfg(feature = "std")]
47impl<R: Reparameterize> Reparameterize for RwLock<R> {
48    type SelfType<T2: AD> = R::SelfType<T2>;
49}
50
51/*
52pub trait DifferentiableFunctionClass {
53    type FunctionType<T: AD> : DifferentiableFunctionTrait<T>;
54}
55impl DifferentiableFunctionClass for () {
56    type FunctionType<T: AD> = ();
57}
58*/
59
60/// Defines the interface for a function that can be differentiated.
61///
62/// Implementors must provide the `call` method to evaluate the function for a given
63/// `AD` type `T`, and specify the number of inputs and outputs.
64pub trait DifferentiableFunctionTrait<T: AD> {
65    // type FunctionClass: DifferentiableFunctionClass;
66    /// A human-readable name for the function.
67    const NAME: &'static str;
68
69    /// Evaluates the function.
70    ///
71    /// # Arguments
72    /// * `inputs` - A slice of input values of type `T`.
73    /// * `freeze` - If true, certain caches or state updates might be skipped (used in optimizations).
74    fn call(&self, inputs: &[T], freeze: bool) -> Vec<T>;
75
76    /// The number of input variables the function expects.
77    fn num_inputs(&self) -> usize;
78
79    /// The number of output variables the function returns.
80    fn num_outputs(&self) -> usize;
81}
82
83pub trait ToOtherADType: Reparameterize {
84    fn to_other_ad_type<T2: AD>(&self) -> <Self as Reparameterize>::SelfType<T2>;
85}
86
87impl<T: AD, F: DifferentiableFunctionTrait<T>> DifferentiableFunctionTrait<T> for Rc<F> {
88    const NAME: &'static str = F::NAME;
89
90    fn call(&self, inputs: &[T], freeze: bool) -> Vec<T> {
91        (**self).call(inputs, freeze)
92    }
93
94    fn num_inputs(&self) -> usize {
95        (**self).num_inputs()
96    }
97
98    fn num_outputs(&self) -> usize {
99        (**self).num_outputs()
100    }
101}
102impl<T: AD, F: DifferentiableFunctionTrait<T>> DifferentiableFunctionTrait<T> for Arc<F> {
103    const NAME: &'static str = F::NAME;
104
105    fn call(&self, inputs: &[T], freeze: bool) -> Vec<T> {
106        (**self).call(inputs, freeze)
107    }
108
109    fn num_inputs(&self) -> usize {
110        (**self).num_inputs()
111    }
112
113    fn num_outputs(&self) -> usize {
114        (**self).num_outputs()
115    }
116}
117#[cfg(feature = "std")]
118impl<T: AD, F: DifferentiableFunctionTrait<T>> DifferentiableFunctionTrait<T> for Mutex<F> {
119    const NAME: &'static str = F::NAME;
120
121    fn call(&self, inputs: &[T], freeze: bool) -> Vec<T> {
122        self.lock().unwrap().call(inputs, freeze)
123    }
124
125    fn num_inputs(&self) -> usize {
126        self.lock().unwrap().num_inputs()
127    }
128
129    fn num_outputs(&self) -> usize {
130        self.lock().unwrap().num_outputs()
131    }
132}
133#[cfg(feature = "std")]
134impl<T: AD, F: DifferentiableFunctionTrait<T>> DifferentiableFunctionTrait<T> for RwLock<F> {
135    const NAME: &'static str = F::NAME;
136
137    fn call(&self, inputs: &[T], freeze: bool) -> Vec<T> {
138        self.read().unwrap().call(inputs, freeze)
139    }
140
141    fn num_inputs(&self) -> usize {
142        self.read().unwrap().num_inputs()
143    }
144
145    fn num_outputs(&self) -> usize {
146        self.read().unwrap().num_outputs()
147    }
148}
149
150impl<T: AD> DifferentiableFunctionTrait<T> for () {
151    const NAME: &'static str = "()";
152
153    fn call(&self, _inputs: &[T], _freeze: bool) -> Vec<T> {
154        vec![]
155    }
156
157    fn num_inputs(&self) -> usize {
158        0
159    }
160
161    fn num_outputs(&self) -> usize {
162        0
163    }
164}
165impl Reparameterize for () {
166    type SelfType<T2: AD> = ();
167}
168
169/*
170pub struct DifferentiableFunctionClassZero;
171impl DifferentiableFunctionClass for DifferentiableFunctionClassZero {
172    type FunctionType<T: AD> = DifferentiableFunctionZero;
173}
174*/
175
176#[derive(Clone)]
177pub struct DifferentiableFunctionZero {
178    num_inputs: usize,
179    num_outputs: usize,
180}
181impl DifferentiableFunctionZero {
182    pub fn new(num_inputs: usize, num_outputs: usize) -> Self {
183        Self {
184            num_inputs,
185            num_outputs,
186        }
187    }
188}
189impl<T: AD> DifferentiableFunctionTrait<T> for DifferentiableFunctionZero {
190    const NAME: &'static str = "DifferentiableFunctionZero";
191
192    fn call(&self, _inputs: &[T], _frozen_freeze: bool) -> Vec<T> {
193        vec![T::zero(); self.num_outputs]
194    }
195
196    fn num_inputs(&self) -> usize {
197        self.num_inputs
198    }
199
200    fn num_outputs(&self) -> usize {
201        self.num_outputs
202    }
203}
204
205impl Reparameterize for DifferentiableFunctionZero {
206    type SelfType<T2: AD> = DifferentiableFunctionZero;
207}
208
209pub trait DerivativeMethodClass {
210    type DerivativeMethod: DerivativeMethodTrait;
211}
212impl DerivativeMethodClass for () {
213    type DerivativeMethod = ();
214}
215
216/// Defines a method for computing the derivative of a `DifferentiableFunctionTrait`.
217pub trait DerivativeMethodTrait: Clone {
218    /// The `AD` type used by this method (e.g., `f64`, `adr`, `adfn`).
219    type T: AD;
220
221    /// Computes the function's value and its Jacobian matrix at the given input point.
222    ///
223    /// # Arguments
224    /// * `inputs` - The input values as `f64`.
225    /// * `function` - The function to differentiate, which must be reparameterizable.
226    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
227        &self,
228        inputs: &[f64],
229        function: &D,
230    ) -> (Vec<f64>, DMatrix<f64>);
231}
232impl DerivativeMethodTrait for () {
233    type T = f64;
234
235    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
236        &self,
237        _inputs: &[f64],
238        _function: &D,
239    ) -> (Vec<f64>, DMatrix<f64>) {
240        panic!("derivative should not actually be called on ()");
241    }
242}
243
244////////////////////////////////////////////////////////////////////////////////////////////////////
245
246pub struct DerivativeMethodClassFiniteDifferencing;
247impl DerivativeMethodClass for DerivativeMethodClassFiniteDifferencing {
248    type DerivativeMethod = FiniteDifferencing;
249}
250
251/// Computes derivatives using the Finite Differencing method.
252///
253/// This method approximates the Jacobian by evaluating the function at slightly
254/// perturbed points. It's safe and works on any function, but can be numerically
255/// unstable and slow for a large number of inputs.
256#[derive(Clone)]
257pub struct FiniteDifferencing {}
258impl FiniteDifferencing {
259    pub fn new() -> Self {
260        Self {}
261    }
262}
263impl DerivativeMethodTrait for FiniteDifferencing {
264    type T = f64;
265
266    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
267        &self,
268        inputs: &[f64],
269        function: &D,
270    ) -> (Vec<f64>, DMatrix<f64>) {
271        let num_inputs = inputs.len();
272        let num_outputs = function.num_outputs();
273        let mut out_derivative = DMatrix::zeros(num_outputs, num_inputs);
274
275        let h = 0.0000001;
276
277        let x0 = inputs.to_vec();
278        // let f0 = D::call(&x0, args);
279        let f0 = function.call(&x0, false);
280
281        for col_idx in 0..num_inputs {
282            let mut xh = x0.clone();
283            xh[col_idx] += h;
284            // let fh = D::call(&xh, args);
285            let fh = function.call(&xh, true);
286            for row_idx in 0..num_outputs {
287                out_derivative[(row_idx, col_idx)] = (fh[row_idx] - f0[row_idx]) / h;
288            }
289        }
290
291        (f0, out_derivative)
292    }
293}
294
295#[cfg(feature = "std")]
296pub struct DerivativeMethodClassReverseAD;
297#[cfg(feature = "std")]
298impl DerivativeMethodClass for DerivativeMethodClassReverseAD {
299    type DerivativeMethod = ReverseAD;
300}
301
302#[cfg(feature = "std")]
303/// Computes derivatives using Reverse-mode Automatic Differentiation.
304///
305/// This method uses a global computation graph to track operations and then
306/// performs a backward pass to compute gradients efficiently. It is ideal
307/// for functions with few outputs and many inputs.
308#[derive(Clone)]
309pub struct ReverseAD {}
310#[cfg(feature = "std")]
311impl ReverseAD {
312    pub fn new() -> Self {
313        Self {}
314    }
315}
316#[cfg(feature = "std")]
317impl DerivativeMethodTrait for ReverseAD {
318    type T = adr;
319
320    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
321        &self,
322        inputs: &[f64],
323        function: &D,
324    ) -> (Vec<f64>, DMatrix<f64>) {
325        let num_inputs = inputs.len();
326        let num_outputs = function.num_outputs();
327        let mut out_derivative = DMatrix::zeros(num_outputs, num_inputs);
328
329        GlobalComputationGraph::get().reset();
330
331        let mut inputs_ad = vec![];
332        for input in inputs.iter() {
333            inputs_ad.push(adr::new_variable(*input, false));
334        }
335
336        let f = function.call(&inputs_ad, false);
337        assert_eq!(f.len(), num_outputs);
338        let out_value = f.iter().map(|x| x.value()).collect();
339
340        for row_idx in 0..num_outputs {
341            if f[row_idx].is_constant() {
342                for col_idx in 0..num_inputs {
343                    out_derivative[(row_idx, col_idx)] = 0.0;
344                }
345            } else {
346                let grad_output = f[row_idx].get_backwards_mode_grad();
347                for col_idx in 0..num_inputs {
348                    let d = grad_output.wrt(&inputs_ad[col_idx]);
349                    out_derivative[(row_idx, col_idx)] = d;
350                }
351            }
352        }
353
354        (out_value, out_derivative)
355    }
356}
357
358pub struct DerivativeMethodClassForwardAD;
359impl DerivativeMethodClass for DerivativeMethodClassForwardAD {
360    type DerivativeMethod = ForwardAD;
361}
362
363/// Computes derivatives using Forward-mode Automatic Differentiation (Single Tangent).
364///
365/// This method propagates a single tangent value alongside each computation.
366/// To compute a full Jacobian, it evaluates the function once per input dimension.
367#[derive(Clone)]
368pub struct ForwardAD {}
369impl ForwardAD {
370    pub fn new() -> Self {
371        Self {}
372    }
373}
374impl DerivativeMethodTrait for ForwardAD {
375    type T = adfn<1>;
376
377    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
378        &self,
379        inputs: &[f64],
380        function: &D,
381    ) -> (Vec<f64>, DMatrix<f64>) {
382        let num_inputs = inputs.len();
383        let num_outputs = function.num_outputs();
384        let mut out_derivative = DMatrix::zeros(num_outputs, num_inputs);
385        let mut out_value = vec![];
386
387        for col_idx in 0..num_inputs {
388            let mut inputs_ad = vec![];
389            for (i, input) in inputs.iter().enumerate() {
390                if i == col_idx {
391                    inputs_ad.push(adfn::new(*input, [1.0]))
392                } else {
393                    inputs_ad.push(adfn::new(*input, [0.0]))
394                }
395            }
396
397            // let f = D::call(&inputs_ad, args);
398            let freeze = if col_idx == 0 { false } else { true };
399            let f = function.call(&inputs_ad, freeze);
400            assert_eq!(
401                f.len(),
402                num_outputs,
403                "{}",
404                format!("does not match {}, {}", f.len(), num_outputs)
405            );
406            for (row_idx, res) in f.iter().enumerate() {
407                if out_value.len() < num_outputs {
408                    out_value.push(res.value);
409                }
410                if res.tangent[0].is_nan() {
411                    out_derivative[(row_idx, col_idx)] = res.tangent[0];
412                } else {
413                    out_derivative[(row_idx, col_idx)] = res.tangent[0];
414                }
415            }
416        }
417
418        (out_value, out_derivative)
419    }
420}
421
422/// Defines a method for computing the Hessian of a `DifferentiableFunctionTrait`.
423///
424/// Implementations of this trait are used by `FunctionEngine::hessian` to extract
425/// second-order information from recursive AD types.
426///
427/// # Implementors
428/// * `HessianAD<N>`: For Forward-over-Forward Hessians.
429/// * `HessianAD_FOR<N>`: For Forward-over-Reverse Hessians.
430#[diagnostic::on_unimplemented(
431    message = "the derivative method `{Self}` does not support Hessian computation",
432    label = "this method does not implement `HessianMethodTrait`",
433    note = "Hessian computation requires recursive AD types. Use `HessianAD<N>` (Forward-over-Forward) or `HessianAD_FOR<N>` (Forward-over-Reverse) instead."
434)]
435pub trait HessianMethodTrait: DerivativeMethodTrait {
436    /// Computes the function's value, Jacobian, and Hessian matrices at the given input point.
437    fn hessian<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
438        &self,
439        inputs: &[f64],
440        function: &D,
441    ) -> (Vec<f64>, DMatrix<f64>, Vec<DMatrix<f64>>);
442}
443
444#[cfg(feature = "hessian")]
445use crate::hyper_ad::hyper::HyperAD_ADFN;
446
447#[cfg(feature = "hessian")]
448#[derive(Clone)]
449pub struct HessianAD<const N: usize> {}
450#[cfg(feature = "hessian")]
451impl<const N: usize> HessianAD<N> {
452    pub fn new() -> Self {
453        Self {}
454    }
455}
456
457#[cfg(feature = "hessian")]
458impl<const N: usize> DerivativeMethodTrait for HessianAD<N> {
459    type T = HyperAD_ADFN<N>;
460
461    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
462        &self,
463        inputs: &[f64],
464        function: &D,
465    ) -> (Vec<f64>, DMatrix<f64>) {
466        // HessianAD can still be used for just derivatives
467        let num_inputs = inputs.len();
468        let num_outputs = function.num_outputs();
469        let mut out_derivative = DMatrix::zeros(num_outputs, num_inputs);
470        let mut out_value = vec![];
471
472        let mut inputs_ad = vec![];
473        for (i, input) in inputs.iter().enumerate() {
474            let mut inner = adfn::<N>::constant(*input);
475            if i < N {
476                inner.set_tangent_value(i, 1.0);
477            }
478            let mut outer = HyperAD_ADFN::<N>::new_inner_constant(inner);
479            if i < N {
480                outer.set_tangent_value(i, 1.0);
481            }
482            inputs_ad.push(outer);
483        }
484
485        let f = function.call(&inputs_ad, false);
486        for (row_idx, res) in f.iter().enumerate() {
487            out_value.push(res.value());
488            let grad = res.tangent_as_vec();
489            for (col_idx, g) in grad.iter().enumerate() {
490                if col_idx < num_inputs {
491                    out_derivative[(row_idx, col_idx)] = *g;
492                }
493            }
494        }
495
496        (out_value, out_derivative)
497    }
498}
499
500#[cfg(feature = "hessian")]
501impl<const N: usize> HessianMethodTrait for HessianAD<N> {
502    fn hessian<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
503        &self,
504        inputs: &[f64],
505        function: &D,
506    ) -> (Vec<f64>, DMatrix<f64>, Vec<DMatrix<f64>>) {
507        let num_inputs = inputs.len();
508        let num_outputs = function.num_outputs();
509        let mut out_value = vec![];
510        let mut out_jacobian = DMatrix::zeros(num_outputs, num_inputs);
511        let mut out_hessians = vec![DMatrix::zeros(num_inputs, num_inputs); num_outputs];
512
513        // Loop over rows and columns in batches of N
514        for row_batch_start in (0..num_inputs).step_by(N) {
515            for col_batch_start in (0..num_inputs).step_by(N) {
516                let mut inputs_ad = vec![];
517                for (i, input) in inputs.iter().enumerate() {
518                    let mut inner = adfn::<N>::constant(*input);
519                    if i >= col_batch_start && i < col_batch_start + N {
520                        inner.set_tangent_value(i - col_batch_start, 1.0);
521                    }
522                    let mut outer = HyperAD_ADFN::<N>::new_inner_constant(inner);
523                    if i >= row_batch_start && i < row_batch_start + N {
524                        outer.set_tangent_value(i - row_batch_start, 1.0);
525                    }
526                    inputs_ad.push(outer);
527                }
528
529                let f = function.call(&inputs_ad, row_batch_start > 0 || col_batch_start > 0);
530                for (row_idx, res) in f.iter().enumerate() {
531                    if row_batch_start == 0 && col_batch_start == 0 {
532                        out_value.push(res.value());
533                    }
534                    
535                    // Extract Jacobian (only need to do this for one set of column batches per row batch)
536                    if row_batch_start == 0 {
537                        let grad = res.inner_value().tangent_as_vec();
538                        for i in 0..N {
539                            if col_batch_start + i < num_inputs {
540                                out_jacobian[(row_idx, col_batch_start + i)] = grad[i];
541                            }
542                        }
543                    }
544
545                    // Extract Hessian block
546                    for i in 0..N {
547                        let r_idx = row_batch_start + i;
548                        if r_idx >= num_inputs { break; }
549                        
550                        let hess_row_chunk = res.tangent[i].tangent_as_vec();
551                        for j in 0..N {
552                            let c_idx = col_batch_start + j;
553                            if c_idx >= num_inputs { break; }
554                            out_hessians[row_idx][(r_idx, c_idx)] = hess_row_chunk[j];
555                        }
556                    }
557                }
558            }
559        }
560
561        (out_value, out_jacobian, out_hessians)
562    }
563}
564
565#[cfg(all(feature = "hessian", feature = "std"))]
566use crate::hyper_ad::hyper_adr::HyperAD_ADR;
567
568#[cfg(all(feature = "hessian", feature = "std"))]
569#[derive(Clone)]
570#[allow(non_camel_case_types)]
571pub struct HessianAD_FOR<const N: usize> {}
572#[cfg(all(feature = "hessian", feature = "std"))]
573impl<const N: usize> HessianAD_FOR<N> {
574    pub fn new() -> Self {
575        Self {}
576    }
577}
578
579#[cfg(all(feature = "hessian", feature = "std"))]
580impl<const N: usize> DerivativeMethodTrait for HessianAD_FOR<N> {
581    type T = HyperAD_ADR<N>;
582
583    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
584        &self,
585        inputs: &[f64],
586        function: &D,
587    ) -> (Vec<f64>, DMatrix<f64>) {
588        let res = self.hessian(inputs, function);
589        (res.0, res.1)
590    }
591}
592
593#[cfg(all(feature = "hessian", feature = "std"))]
594impl<const N: usize> HessianMethodTrait for HessianAD_FOR<N> {
595    fn hessian<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
596        &self,
597        inputs: &[f64],
598        function: &D,
599    ) -> (Vec<f64>, DMatrix<f64>, Vec<DMatrix<f64>>) {
600        let num_inputs = inputs.len();
601        let num_outputs = function.num_outputs();
602        let mut out_value = vec![];
603        let mut out_jacobian = DMatrix::zeros(num_outputs, num_inputs);
604        let mut out_hessians = vec![DMatrix::zeros(num_inputs, num_inputs); num_outputs];
605
606        // Loop over rows in batches of N. Each pass recovers full columns via backprop.
607        for row_batch_start in (0..num_inputs).step_by(N) {
608            let mut inputs_ad = vec![];
609            let mut inputs_adr = vec![];
610            for (i, input) in inputs.iter().enumerate() {
611                // Reset computation graph only on the very first batch
612                let adr_var = crate::reverse_ad::adr::adr::new_variable(*input, i == 0 && row_batch_start == 0);
613                inputs_adr.push(adr_var);
614                let mut outer = HyperAD_ADR::<N>::new_inner_constant(adr_var);
615                if i >= row_batch_start && i < row_batch_start + N {
616                    outer.set_tangent_value(i - row_batch_start, 1.0);
617                }
618                inputs_ad.push(outer);
619            }
620
621            let f = function.call(&inputs_ad, row_batch_start > 0);
622            for (row_idx, res) in f.iter().enumerate() {
623                if row_batch_start == 0 {
624                    out_value.push(res.value());
625                    
626                    // Extract full Jacobian from primal ADR
627                    let grad = res.value.get_backwards_mode_grad();
628                    for (col_idx, adr_var) in inputs_adr.iter().enumerate() {
629                        out_jacobian[(row_idx, col_idx)] = grad.wrt(adr_var);
630                    }
631                }
632
633                // Extract Hessian rows for this batch
634                for i in 0..N {
635                    let r_idx = row_batch_start + i;
636                    if r_idx >= num_inputs { break; }
637                    
638                    let grad_hess = res.tangent[i].get_backwards_mode_grad();
639                    for (c_idx, adr_var) in inputs_adr.iter().enumerate() {
640                        out_hessians[row_idx][(r_idx, c_idx)] = grad_hess.wrt(adr_var);
641                    }
642                }
643            }
644        }
645
646        (out_value, out_jacobian, out_hessians)
647    }
648}
649
650pub struct DerivativeMethodClassForwardADMulti<A: AD + ForwardADTrait>(PhantomData<A>);
651impl<A: AD + ForwardADTrait> DerivativeMethodClass for DerivativeMethodClassForwardADMulti<A> {
652    type DerivativeMethod = ForwardADMulti<A>;
653}
654
655/// Computes derivatives using Forward-mode Automatic Differentiation (Multi-Tangent).
656///
657/// This method allows for computing multiple columns of the Jacobian in a single pass
658/// by propagating a vector of tangents. This can significantly speed up computation
659/// by taking advantage of SIMD and reducing function overhead.
660#[derive(Clone)]
661pub struct ForwardADMulti<A: AD + ForwardADTrait> {
662    phantom_data: PhantomData<A>,
663}
664impl<A: AD + ForwardADTrait> ForwardADMulti<A> {
665    pub fn new() -> Self {
666        Self {
667            phantom_data: PhantomData::default(),
668        }
669    }
670}
671impl<A: AD + ForwardADTrait> DerivativeMethodTrait for ForwardADMulti<A> {
672    type T = A;
673
674    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
675        &self,
676        inputs: &[f64],
677        function: &D,
678    ) -> (Vec<f64>, DMatrix<f64>) {
679        let num_inputs = inputs.len();
680        let num_outputs = function.num_outputs();
681        let mut out_derivative = DMatrix::zeros(num_outputs, num_inputs);
682        let mut out_value = vec![];
683
684        let mut curr_idx = 0;
685
686        let mut freeze = false;
687        let k = Self::T::tangent_size();
688        'l1: loop {
689            let mut inputs_ad = vec![];
690            for input in inputs.iter() {
691                // inputs_ad.push(adf::new(*input, [0.0; K]))
692                inputs_ad.push(Self::T::constant(*input));
693            }
694
695            'l2: for i in 0..k {
696                if curr_idx + i >= num_inputs {
697                    break 'l2;
698                }
699                // inputs_ad[curr_idx+i].tangent[i] = 1.0;
700                inputs_ad[curr_idx + i].set_tangent_value(i, 1.0);
701            }
702
703            let f = function.call(&inputs_ad, freeze);
704            freeze = true;
705            assert_eq!(f.len(), num_outputs);
706
707            for (row_idx, res) in f.iter().enumerate() {
708                if out_value.len() < num_outputs {
709                    out_value.push(res.value());
710                }
711                let curr_tangent = res.tangent_as_vec();
712                'l3: for i in 0..k {
713                    if curr_idx + i >= num_inputs {
714                        break 'l3;
715                    }
716                    // out_derivative[(row_idx, curr_idx+i)] = res.tangent[i];
717                    if curr_tangent[i].is_nan() {
718                        out_derivative[(row_idx, curr_idx + i)] = curr_tangent[i];
719                    } else {
720                        out_derivative[(row_idx, curr_idx + i)] = curr_tangent[i];
721                    }
722                }
723            }
724
725            curr_idx += k;
726            if curr_idx >= num_inputs {
727                break 'l1;
728            }
729        }
730
731        return (out_value, out_derivative);
732    }
733}
734
735#[cfg(feature = "nightly")]
736pub struct DerivativeMethodClassFiniteDifferencingMulti<const K: usize>;
737#[cfg(feature = "nightly")]
738impl<const K: usize> DerivativeMethodClass for DerivativeMethodClassFiniteDifferencingMulti<K> {
739    type DerivativeMethod = FiniteDifferencingMulti2<K>;
740}
741
742#[cfg(feature = "nightly")]
743#[derive(Clone)]
744pub struct FiniteDifferencingMulti2<const K: usize>;
745#[cfg(feature = "nightly")]
746impl<const K: usize> FiniteDifferencingMulti2<K> {
747    pub fn new() -> Self {
748        Self {}
749    }
750}
751#[cfg(feature = "nightly")]
752impl<const K: usize> DerivativeMethodTrait for FiniteDifferencingMulti2<K> {
753    type T = f64xn<K>;
754
755    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
756        &self,
757        inputs: &[f64],
758        function: &D,
759    ) -> (Vec<f64>, DMatrix<f64>) {
760        let num_inputs = inputs.len();
761        let num_outputs = function.num_outputs();
762        let mut out_derivative = DMatrix::zeros(num_outputs, num_inputs);
763        let mut out_value = vec![];
764
765        let h = 0.0000001;
766
767        let mut curr_idx = 0;
768        let mut first_loop = true;
769
770        'l1: loop {
771            let mut inputs_ad = vec![];
772            for input in inputs.iter() {
773                inputs_ad.push(f64xn::<K>::splat(*input));
774            }
775
776            if first_loop {
777                'l2: for i in 0..K {
778                    if curr_idx + i >= num_inputs {
779                        break 'l2;
780                    }
781                    if i + 1 >= K {
782                        break 'l2;
783                    }
784                    inputs_ad[curr_idx + i].value[i + 1] += h;
785                }
786            } else {
787                'l2: for i in 0..K {
788                    if curr_idx + i >= num_inputs {
789                        break 'l2;
790                    }
791                    if i >= K {
792                        break 'l2;
793                    }
794                    inputs_ad[curr_idx + i].value[i] += h;
795                }
796            }
797
798            // let f = D::call(&inputs_ad, args);
799            let f = function.call(&inputs_ad, false);
800            assert_eq!(f.len(), num_outputs);
801
802            if first_loop {
803                for res in f.iter() {
804                    out_value.push(res.value[0]);
805                }
806            }
807
808            for (row_idx, res) in f.iter().enumerate() {
809                if first_loop {
810                    'l3: for i in 0..K {
811                        if curr_idx + i >= num_inputs {
812                            break 'l3;
813                        }
814                        if i + 1 >= K {
815                            break 'l3;
816                        }
817                        out_derivative[(row_idx, curr_idx + i)] =
818                            (res.value[i + 1] - out_value[row_idx]) / h;
819                    }
820                } else {
821                    'l3: for i in 0..K {
822                        if curr_idx + i >= num_inputs {
823                            break 'l3;
824                        }
825                        if i >= K {
826                            break 'l3;
827                        }
828                        out_derivative[(row_idx, curr_idx + i)] =
829                            (res.value[i] - out_value[row_idx]) / h;
830                    }
831                }
832            }
833
834            if first_loop {
835                first_loop = false;
836                curr_idx += K - 1;
837            } else {
838                curr_idx += K;
839            }
840
841            if curr_idx >= num_inputs {
842                break 'l1;
843            }
844        }
845
846        return (out_value, out_derivative);
847    }
848}
849
850#[cfg(feature = "std")]
851#[derive(Clone)]
852pub struct WASP {
853    cache: Arc<RwLock<WASPCache>>,
854    num_f_calls: Arc<RwLock<usize>>,
855    d_theta: f64,
856    d_ell: f64,
857}
858#[cfg(feature = "std")]
859impl WASP {
860    pub fn new(n: usize, m: usize, orthonormal_delta_x: bool, d_theta: f64, d_ell: f64) -> Self {
861        Self {
862            cache: Arc::new(RwLock::new(WASPCache::new(n, m, orthonormal_delta_x))),
863            num_f_calls: Arc::new(RwLock::new(0)),
864            d_theta,
865            d_ell,
866        }
867    }
868    pub fn reset_cache(&self) {
869        self.cache.write().unwrap().reset();
870    }
871    pub fn new_default(n: usize, m: usize) -> Self {
872        Self::new(n, m, true, 0.3, 0.3)
873    }
874    pub fn num_f_calls(&self) -> usize {
875        return self.num_f_calls.read().unwrap().clone();
876    }
877}
878#[cfg(feature = "std")]
879impl DerivativeMethodTrait for WASP {
880    type T = f64;
881
882    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
883        &self,
884        inputs: &[f64],
885        function: &D,
886    ) -> (Vec<f64>, DMatrix<f64>) {
887        let mut num_f_calls = 0;
888        let f_k = function.call(inputs, false);
889        let f_k_dv = DVector::from_column_slice(&f_k);
890        num_f_calls += 1;
891        let epsilon = 0.000001;
892
893        let mut cache = self.cache.write().unwrap();
894        let n = inputs.len();
895
896        let x = DVector::<f64>::from_column_slice(inputs);
897
898        loop {
899            let i = cache.i.clone();
900
901            let delta_x_i = cache.delta_x.column(i);
902
903            let x_k_plus_delta_x_i: DVector<f64> = &x + epsilon * &delta_x_i;
904            let f_k_plus_delta_x_i = DVector::<f64>::from_column_slice(
905                &function.call(x_k_plus_delta_x_i.as_slice(), true),
906            );
907            num_f_calls += 1;
908            let delta_f_i = (&f_k_plus_delta_x_i - &f_k_dv) / epsilon;
909            let delta_f_i_hat = cache.delta_f_t.row(i);
910            let delta_f_i_hat = DVector::from_column_slice(delta_f_i_hat.transpose().as_slice());
911            let return_result = close_enough(&delta_f_i, &delta_f_i_hat, self.d_theta, self.d_ell);
912
913            cache.delta_f_t.set_row(i, &delta_f_i.transpose());
914            let c_1_mat = &cache.c_1[i];
915            let c_2_mat = &cache.c_2[i];
916            let delta_f_t = &cache.delta_f_t;
917
918            let d_t_star = c_1_mat * delta_f_t + c_2_mat * delta_f_i.transpose();
919            let d_star = d_t_star.transpose();
920
921            let tmp = &d_star * &cache.delta_x;
922            cache.delta_f_t = tmp.transpose();
923
924            let mut new_i = i + 1;
925            if new_i >= n {
926                new_i = 0;
927            }
928            cache.i = new_i;
929
930            if return_result {
931                *self.num_f_calls.write().unwrap() = num_f_calls;
932                return (f_k, d_star);
933            }
934        }
935    }
936}
937
938#[cfg(feature = "std")]
939#[derive(Clone, Debug)]
940pub struct WASPCache {
941    pub n: usize,
942    pub m: usize,
943    pub i: usize,
944    pub delta_f_t: DMatrix<f64>,
945    pub delta_x: DMatrix<f64>,
946    pub c_1: Vec<DMatrix<f64>>,
947    pub c_2: Vec<DVector<f64>>,
948}
949#[cfg(feature = "std")]
950impl WASPCache {
951    pub fn new(n: usize, m: usize, orthonormal_delta_x: bool) -> Self {
952        let delta_f_t = DMatrix::<f64>::identity(n, m);
953        let delta_x = get_tangent_matrix(n, orthonormal_delta_x);
954        let mut c_1 = vec![];
955        let mut c_2 = vec![];
956
957        let a_mat: DMatrix<f64> = 2.0 * &delta_x * &delta_x.transpose();
958        let a_inv_mat = a_mat.try_inverse().unwrap();
959
960        for i in 0..n {
961            let delta_x_i = DVector::<f64>::from_column_slice(delta_x.column(i).as_slice());
962            let s_i = (delta_x_i.transpose() * &a_inv_mat * &delta_x_i)[(0, 0)];
963            let s_i_inv = 1.0 / s_i;
964            let c_1_mat = &a_inv_mat
965                * (DMatrix::<f64>::identity(n, n)
966                    - s_i_inv * &delta_x_i * delta_x_i.transpose() * &a_inv_mat)
967                * 2.0
968                * &delta_x;
969            let c_2_mat = s_i_inv * &a_inv_mat * delta_x_i;
970            c_1.push(c_1_mat);
971            c_2.push(c_2_mat);
972        }
973
974        return Self {
975            n,
976            m,
977            i: 0,
978            delta_f_t,
979            delta_x,
980            c_1,
981            c_2,
982        };
983    }
984    pub fn reset(&mut self) {
985        self.delta_f_t = DMatrix::<f64>::identity(self.n, self.m);
986        self.i = 0;
987    }
988}
989
990/*
991#[derive(Clone)]
992pub struct WASP2 {
993    cache: Arc<RwLock<WASPCache2>>,
994    num_f_calls: Arc<RwLock<usize>>,
995    d_theta: f64,
996    d_ell: f64
997}
998impl WASP2 {
999    pub fn new(n: usize, m: usize, alpha:f64, orthonormal_delta_x: bool, d_theta: f64, d_ell: f64) -> Self {
1000        Self {
1001            cache: Arc::new(RwLock::new(WASPCache2::new(n, m, alpha, orthonormal_delta_x))),
1002            num_f_calls: Arc::new(RwLock::new(0)),
1003            d_theta,
1004            d_ell,
1005        }
1006    }
1007    pub fn new_default(n: usize, m: usize) -> Self {
1008        Self::new(n, m, 0.98, true, 0.3, 0.3)
1009    }
1010    pub fn num_f_calls(&self) -> usize {
1011        return self.num_f_calls.read().unwrap().clone()
1012    }
1013}
1014impl DerivativeMethodTrait for WASP2 {
1015    type T = f64;
1016
1017    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(&self, inputs: &[f64], function: &D) -> (Vec<f64>, DMatrix<f64>) {
1018        let mut num_f_calls = 0;
1019        let f_k = function.call(inputs, false);
1020        let f_k_dv = DVector::from_column_slice(&f_k);
1021        num_f_calls += 1;
1022        let epsilon = 0.000001;
1023
1024        let mut cache = self.cache.write().unwrap();
1025        let n = inputs.len();
1026
1027        let x = DVector::<f64>::from_column_slice(inputs);
1028
1029        loop {
1030            let i = cache.i.clone();
1031
1032            let delta_x_i = cache.delta_x.column(i);
1033
1034            let x_k_plus_delta_x_i: DVector<f64> = &x + epsilon*&delta_x_i;
1035            let f_k_plus_delta_x_i = DVector::<f64>::from_column_slice(&function.call(x_k_plus_delta_x_i.as_slice(), true));
1036            num_f_calls += 1;
1037            let delta_f_i = (&f_k_plus_delta_x_i - &f_k_dv) / epsilon;
1038            let delta_f_i_hat = &cache.curr_d * &delta_x_i;
1039            // let delta_f_i_hat = cache.delta_f_t.row(i);
1040            // let delta_f_i_hat = DVector::from_column_slice(delta_f_i_hat.transpose().as_slice());
1041            let return_result = close_enough(&delta_f_i, &delta_f_i_hat, self.d_theta, self.d_ell);
1042
1043            cache.delta_f_t.set_row(i, &delta_f_i.transpose());
1044            let c_1_mat = &cache.c_1[i];
1045            let c_2_mat = &cache.c_2[i];
1046            let delta_f_t = &cache.delta_f_t;
1047
1048            let d_t_star = c_1_mat*delta_f_t + c_2_mat*delta_f_i.transpose();
1049            let d_star = d_t_star.transpose();
1050            cache.curr_d = d_star.clone();
1051
1052            let mut new_i = i + 1;
1053            if new_i >= n { new_i = 0; }
1054            cache.i = new_i;
1055
1056            if return_result {
1057                *self.num_f_calls.write().unwrap() = num_f_calls;
1058                return (f_k, d_star);
1059            }
1060        }
1061    }
1062}
1063
1064pub struct WASPCache2 {
1065    pub i: usize,
1066    pub curr_d: DMatrix<f64>,
1067    pub delta_f_t: DMatrix<f64>,
1068    pub delta_x: DMatrix<f64>,
1069    pub c_1: Vec<DMatrix<f64>>,
1070    pub c_2: Vec<DVector<f64>>
1071}
1072impl WASPCache2 {
1073    pub fn new(n: usize, m: usize, alpha: f64, orthonormal_delta_x: bool) -> Self {
1074        assert!(alpha > 0.0 && alpha < 1.0);
1075
1076        let curr_d = DMatrix::<f64>::identity(m, n);
1077        let delta_f_t = DMatrix::<f64>::identity(n, m);
1078        let delta_x = get_tangent_matrix(n, orthonormal_delta_x);
1079        let mut c_1 = vec![];
1080        let mut c_2 = vec![];
1081
1082        for i in 0..n {
1083            let delta_x_i = DVector::<f64>::from_column_slice(delta_x.column(i).as_slice());
1084            let mut w_i = DMatrix::<f64>::zeros(n, n);
1085            for j in 0..n {
1086                let exponent = math_mod(i as i32 - j as i32, n as i32) as f64 / (n as i32 - 1) as f64;
1087                w_i[(j, j)] = alpha * (1.0 - alpha).powf(exponent);
1088            }
1089            let w_i_2 = &w_i * &w_i;
1090
1091            let a_i = 2.0 * &delta_x * &w_i_2 * &delta_x.transpose();
1092            let a_i_inv = a_i.clone().try_inverse().unwrap();
1093
1094            let s_i = (delta_x_i.transpose() * &a_i_inv * &delta_x_i)[(0,0)];
1095            let s_i_inv = 1.0 / s_i;
1096            let c_1_mat = &a_i_inv * (DMatrix::<f64>::identity(n, n) - s_i_inv * &delta_x_i * delta_x_i.transpose() * &a_i_inv) * 2.0 * &delta_x * &w_i_2;
1097            let c_2_mat = s_i_inv * &a_i_inv * delta_x_i;
1098            c_1.push(c_1_mat);
1099            c_2.push(c_2_mat);
1100        }
1101
1102        return Self {
1103            i: 0,
1104            curr_d,
1105            delta_f_t,
1106            delta_x,
1107            c_1,
1108            c_2,
1109        }
1110    }
1111}
1112
1113pub fn math_mod(a: i32, b: i32) -> i32 {
1114    return ((a % b) + b) % b;
1115}
1116*/
1117
1118#[cfg(feature = "std")]
1119pub(crate) fn get_tangent_matrix(n: usize, orthogonal: bool) -> DMatrix<f64> {
1120    let mut rng = thread_rng();
1121    let uniform = Uniform::new(-1.0, 1.0);
1122
1123    let t = DMatrix::<f64>::from_fn(n, n, |_, _| uniform.sample(&mut rng));
1124
1125    return if orthogonal {
1126        let svd = t.svd(true, true);
1127        let delta_x = svd.u.as_ref().unwrap() * svd.v_t.as_ref().unwrap();
1128        delta_x
1129    } else {
1130        t
1131    };
1132}
1133
1134pub(crate) fn close_enough(a: &DVector<f64>, b: &DVector<f64>, d_theta: f64, d_ell: f64) -> bool {
1135    let a_n = a.norm();
1136    let b_n = b.norm();
1137
1138    let tmp = ((a.dot(&b) / (a_n * b_n)) - 1.0).abs();
1139    if tmp > d_theta {
1140        return false;
1141    }
1142
1143    let tmp1 = if b_n != 0.0 {
1144        ((a_n / b_n) - 1.0).abs()
1145    } else {
1146        f64::MAX
1147    };
1148    let tmp2 = if a_n != 0.0 {
1149        ((b_n / a_n) - 1.0).abs()
1150    } else {
1151        f64::MAX
1152    };
1153
1154    if f64::min(tmp1, tmp2) > d_ell {
1155        return false;
1156    }
1157
1158    return true;
1159}
1160
1161/*
1162
1163pub fn math_modulus(a: i64, b: i64) -> usize {
1164    (((a % b) + b) % b) as usize
1165}
1166
1167pub fn get_tangent_matrix(n: usize, orthonormalize: bool) -> DMatrix<f64> {
1168    let mut out = DMatrix::zeros(n, n);
1169    let mut rng = rand::rng();
1170
1171    for i in 0..n {
1172        for j in 0..n {
1173            out[(i, j)] = rng.random_range(-1.0..=1.0);
1174        }
1175    }
1176
1177    if orthonormalize {
1178        let svd = out.svd(true, true);
1179        out = svd.u.as_ref().unwrap()*svd.v_t.as_ref().unwrap();
1180    }
1181
1182    return out;
1183}
1184
1185pub fn wasp_projection<D: DifferentiableFunctionTrait<f64> + ?Sized>(f: &D, f_x_k: &DVector<f64>, x_k: &[f64], cache: &WASPCache) -> DMatrix<f64> {
1186    let epsilon = 0.00001;
1187    let x_k = DVector::from_column_slice(x_k);
1188    let i = cache.i.lock().unwrap();
1189    let c_1_mat = &cache.c_1_mats[*i];
1190    let c_2_mat = &cache.c_2_mats[*i];
1191    let delta_x_i = DVector::from_column_slice(cache.delta_x_mat.column(*i).as_slice());
1192    let f_x_k_delta = DVector::from_column_slice(&f.call((&x_k + epsilon*&delta_x_i).as_slice(), true));
1193    let delta_f_i = (f_x_k_delta - f_x_k) / epsilon;
1194    let mut delta_f_hat_t = cache.delta_f_mat_t.lock().unwrap();
1195    delta_f_hat_t.set_row(*i, &delta_f_i.transpose());
1196    return c_1_mat*&*delta_f_hat_t + c_2_mat*&delta_f_i.transpose();
1197}
1198
1199pub fn wasp_projection2<D: DifferentiableFunctionTrait<f64> + ?Sized>(f: &D, f_x_k: &DVector<f64>, x_k: &[f64], cache: &WASPCache2) -> DMatrix<f64> {
1200    let epsilon = 0.00001;
1201    let x_k = DVector::from_column_slice(x_k);
1202    let i = cache.i.lock().unwrap();
1203    let c_1_mat = &cache.c_1_mats[*i];
1204    let c_2_mat = &cache.c_2_mats[*i];
1205    let delta_x_i = DVector::from_column_slice(cache.delta_x_mat.column(*i).as_slice());
1206    let f_x_k_delta = DVector::from_column_slice(&f.call((&x_k + epsilon*&delta_x_i).as_slice(), true));
1207    let delta_f_i = (f_x_k_delta - f_x_k) / epsilon;
1208    let mut delta_f_hat_t = cache.delta_f_mat_t.lock().unwrap();
1209    delta_f_hat_t.set_row(*i, &delta_f_i.transpose());
1210    return c_1_mat*&*delta_f_hat_t + c_2_mat*&delta_f_i.transpose();
1211}
1212
1213pub fn close_enough(d_a_t_mat: &DMatrix<f64>, d_b_t_mat: &DMatrix<f64>, l: usize, m: usize, d_theta: f64) -> bool {
1214    let mut numbers: Vec<usize> = (0..m).collect();
1215
1216    let mut rng = rng();
1217    numbers.shuffle(&mut rng);
1218
1219    let js: Vec<usize> = numbers.into_iter().take(l).collect();
1220
1221    // println!("{}", d_a_t_mat);
1222    // println!("{}", d_b_t_mat);
1223    // println!("---");
1224
1225    for j in js {
1226        let d_a = DVector::from_column_slice(d_a_t_mat.column(j).as_slice());
1227        let d_b = DVector::from_column_slice(d_b_t_mat.column(j).as_slice());
1228
1229        let d_a_n = d_a.norm();
1230        let d_b_n = d_b.norm();
1231
1232        let dot = d_a.dot(&d_b);
1233        let angle = (dot / (d_a_n * d_b_n)).acos();
1234        // println!("{:?}", angle);
1235
1236        if angle > d_theta { return false; }
1237    }
1238
1239    return true;
1240}
1241
1242#[inline(always)]
1243pub fn close_enough2(a: &DVector<f64>, b: &DVector<f64>, d_theta: f64, d_l: f64) -> bool {
1244    let an = a.norm();
1245    let bn = b.norm();
1246    let d = a.dot(b);
1247
1248    if (d / (an * bn) - 1.0).abs() > d_theta { return false; }
1249    if (an / bn - 1.0).abs() > d_l { return false; }
1250
1251    return true;
1252}
1253
1254pub fn derivative_angular_distance(d_a_t_mat: &DMatrix<f64>, d_b_t_mat: &DMatrix<f64>) -> f64 {
1255    let m = d_a_t_mat.ncols();
1256
1257    let mut max_angle = f64::MIN;
1258
1259    for j in 0..m {
1260        let d_a = DVector::from_column_slice(d_a_t_mat.column(j).as_slice());
1261        let d_b = DVector::from_column_slice(d_b_t_mat.column(j).as_slice());
1262
1263        let d_a_n = d_a.norm();
1264        let d_b_n = d_b.norm();
1265
1266        let dot = d_a.dot(&d_b);
1267        let angle = (dot / (d_a_n * d_b_n)).acos();
1268        if angle > max_angle { max_angle = angle; }
1269    }
1270
1271    return max_angle;
1272}
1273
1274#[derive(Clone)]
1275pub struct WASPCache {
1276    pub delta_f_mat_t: Arc<Mutex<DMatrix<f64>>>,
1277    pub delta_x_mat: DMatrix<f64>,
1278    pub c_1_mats: Vec<DMatrix<f64>>,
1279    pub c_2_mats: Vec<DVector<f64>>,
1280    pub i: Arc<Mutex<usize>>
1281}
1282impl WASPCache {
1283    pub fn new(n: usize, m: usize, alpha: f64, orthonormalize: bool) -> Self {
1284        let delta_x_mat = get_tangent_matrix(n, orthonormalize);
1285        let mut c_1_mats = vec![];
1286        let mut c_2_mats = vec![];
1287
1288        for i in 0..n {
1289            let delta_x_i = DVector::from_column_slice(delta_x_mat.column(i).as_slice());
1290            let mut w_i_mat = DMatrix::zeros(n, n);
1291            for j in 0..n {
1292                let exp = math_modulus(i as i64 - j as i64, n as i64) as f64 / ((n - 1) as f64);
1293                w_i_mat[(j,j)] = alpha*(1.0 - alpha).pow(  exp );
1294            }
1295            let w_i_mat_2 = &w_i_mat * & w_i_mat;
1296            let a_i_mat = 2.0*(&delta_x_mat * &w_i_mat_2 * &delta_x_mat.transpose());
1297            let a_i_mat_inv = a_i_mat.clone().try_inverse().unwrap();
1298            let s_i = (&delta_x_i.transpose() * &a_i_mat_inv * &delta_x_i)[0];
1299            let s_i_inv = 1.0 / s_i;
1300            let c_1_mat = &a_i_mat_inv*(DMatrix::identity(n, n) - s_i_inv*&delta_x_i*&delta_x_i.transpose()*&a_i_mat_inv)*2.0*&delta_x_mat*&w_i_mat_2;
1301            let c_2_mat = s_i_inv*&a_i_mat_inv*&delta_x_i;
1302            c_1_mats.push(c_1_mat);
1303            c_2_mats.push(c_2_mat);
1304        }
1305
1306        Self {
1307            delta_f_mat_t: Arc::new(Mutex::new(DMatrix::zeros(n, m))),
1308            delta_x_mat,
1309            c_1_mats,
1310            c_2_mats,
1311            i: Arc::new(Mutex::new(0)),
1312        }
1313    }
1314}
1315
1316#[derive(Clone)]
1317pub struct WASPCache2 {
1318    pub delta_f_mat_t: Arc<Mutex<DMatrix<f64>>>,
1319    pub delta_x_mat: DMatrix<f64>,
1320    pub c_1_mats: Vec<DMatrix<f64>>,
1321    pub c_2_mats: Vec<DVector<f64>>,
1322    pub i: Arc<Mutex<usize>>
1323}
1324impl WASPCache2 {
1325    pub fn new(n: usize, m: usize, orthonormalize: bool) -> Self {
1326        let delta_x_mat = get_tangent_matrix(n, orthonormalize);
1327        let mut c_1_mats = vec![];
1328        let mut c_2_mats = vec![];
1329
1330        for i in 0..n {
1331            let delta_x_i = DVector::from_column_slice(delta_x_mat.column(i).as_slice());
1332            let a_i_mat = 2.0*(&delta_x_mat * &delta_x_mat.transpose());
1333            let a_i_mat_inv = a_i_mat.clone().try_inverse().unwrap();
1334            let s_i = (&delta_x_i.transpose() * &a_i_mat_inv * &delta_x_i)[0];
1335            let s_i_inv = 1.0 / s_i;
1336            let c_1_mat = &a_i_mat_inv*(DMatrix::identity(n, n) - s_i_inv*&delta_x_i*&delta_x_i.transpose()*&a_i_mat_inv)*2.0*&delta_x_mat;
1337            let c_2_mat = s_i_inv*&a_i_mat_inv*&delta_x_i;
1338            c_1_mats.push(c_1_mat);
1339            c_2_mats.push(c_2_mat);
1340        }
1341
1342        Self {
1343            delta_f_mat_t: Arc::new(Mutex::new(DMatrix::zeros(n, m))),
1344            delta_x_mat,
1345            c_1_mats,
1346            c_2_mats,
1347            i: Arc::new(Mutex::new(0)),
1348        }
1349    }
1350}
1351
1352pub struct DerivativeMethodClassWASP;
1353impl DerivativeMethodClass for DerivativeMethodClassWASP {
1354    type DerivativeMethod = WASP;
1355}
1356
1357#[derive(Clone)]
1358pub struct WASP {
1359    pub cache: WASPCache2,
1360    pub d_theta: f64,
1361    pub d_l: f64,
1362    pub num_f_calls: Arc<Mutex<usize>>
1363}
1364impl WASP {
1365    pub fn new(n: usize, m: usize, d_theta: f64, d_l: f64, orthonormalize: bool) -> Self {
1366        Self {
1367            cache: WASPCache2::new(n, m, orthonormalize),
1368            d_theta,
1369            d_l,
1370            num_f_calls: Arc::new(Mutex::new(0)),
1371        }
1372    }
1373
1374    pub fn get_num_f_calls(&self) -> usize {
1375        self.num_f_calls.lock().unwrap().clone()
1376    }
1377}
1378impl DerivativeMethodTrait for WASP {
1379    type T = f64;
1380
1381    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(&self, inputs: &[f64], function: &D) -> (Vec<f64>, DMatrix<f64>) {
1382        let mut num_f_calls = self.num_f_calls.lock().unwrap();
1383        *num_f_calls = 0;
1384        let f_x_k_vec = function.call(inputs, false);
1385        let f_x_k = DVector::from_column_slice(&f_x_k_vec);
1386        let x_k = DVector::from_column_slice(inputs);
1387        *num_f_calls += 1;
1388
1389        let mut return_result;
1390
1391        loop {
1392            let mut i = self.cache.i.lock().unwrap();
1393
1394            let epsilon = 0.00001;
1395            let c_1_mat = &self.cache.c_1_mats[*i];
1396            let c_2_mat = &self.cache.c_2_mats[*i];
1397            let delta_x_i = DVector::from_column_slice(self.cache.delta_x_mat.column(*i).as_slice());
1398            let f_x_k_delta = DVector::from_column_slice(&function.call((&x_k + epsilon * &delta_x_i).as_slice(), true));
1399            *num_f_calls += 1;
1400            let delta_f_i = (f_x_k_delta - &f_x_k) / epsilon;
1401            let mut delta_f_hat_t = self.cache.delta_f_mat_t.lock().unwrap();
1402            let delta_f_i_hat = DVector::from_column_slice(delta_f_hat_t.row(*i).transpose().as_slice());
1403
1404            return_result = close_enough2(&delta_f_i, &delta_f_i_hat, self.d_theta, self.d_l);
1405
1406            delta_f_hat_t.set_row(*i, &delta_f_i.transpose());
1407            let d_t = c_1_mat * &*delta_f_hat_t + c_2_mat * &delta_f_i.transpose();
1408            *delta_f_hat_t = &self.cache.delta_x_mat.transpose() * &d_t;
1409
1410            *i = (*i + 1) % inputs.len();
1411
1412            if return_result {
1413                return (f_x_k_vec, d_t.transpose());
1414            }
1415        }
1416    }
1417}
1418
1419pub struct DerivativeMethodClassWASPNec;
1420impl DerivativeMethodClass for DerivativeMethodClassWASPNec {
1421    type DerivativeMethod = WASPNec;
1422}
1423#[derive(Clone)]
1424pub struct WASPNec {
1425    pub cache: WASPCache,
1426    pub first_call: Arc<Mutex<bool>>,
1427    pub num_f_calls: Arc<Mutex<usize>>
1428}
1429impl WASPNec {
1430    pub fn new(n: usize, m: usize, alpha: f64, orthonormalize: bool) -> Self {
1431        Self {
1432            cache: WASPCache::new(n, m, alpha, orthonormalize),
1433            first_call: Arc::new(Mutex::new(true)),
1434            num_f_calls: Arc::new(Mutex::new(0)),
1435        }
1436    }
1437
1438    pub fn get_num_f_calls(&self) -> usize {
1439        self.num_f_calls.lock().unwrap().clone()
1440    }
1441}
1442impl DerivativeMethodTrait for WASPNec {
1443    type T = f64;
1444
1445    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(&self, inputs: &[f64], function: &D) -> (Vec<f64>, DMatrix<f64>) {
1446        let mut num_f_calls = self.num_f_calls.lock().unwrap();
1447        *num_f_calls = 0;
1448        let f_x_k_vec = function.call(inputs, false);
1449        let f_x_k = DVector::from_column_slice(&f_x_k_vec);
1450        *num_f_calls += 1;
1451
1452        let mut first_call = self.first_call.lock().unwrap();
1453        if *first_call {
1454            let epsilon = 0.00001;
1455            let x_k = DVector::from_column_slice(inputs);
1456            let n = inputs.len();
1457            for i in 0..n {
1458                let delta_x_i = DVector::from_column_slice(self.cache.delta_x_mat.column(i).as_slice());
1459                let f_x_k_delta = DVector::from_column_slice(&function.call((&x_k + epsilon*&delta_x_i).as_slice(), true));
1460                let delta_f_i = (f_x_k_delta - &f_x_k) / epsilon;
1461                let mut delta_f_hat_t = self.cache.delta_f_mat_t.lock().unwrap();
1462                delta_f_hat_t.set_row(i, &delta_f_i.transpose());
1463            }
1464            *first_call = false;
1465        }
1466
1467        let d_t = wasp_projection(function, &f_x_k, inputs, &self.cache);
1468        *num_f_calls += 1;
1469        let mut i = self.cache.i.lock().unwrap();
1470        *i = (*i + 1) % inputs.len();
1471
1472        return (f_x_k_vec, d_t.transpose());
1473    }
1474}
1475
1476pub struct DerivativeMethodClassWASPEc;
1477impl DerivativeMethodClass for DerivativeMethodClassWASPEc {
1478    type DerivativeMethod = WASPEc;
1479}
1480
1481#[derive(Clone)]
1482pub struct WASPEc {
1483    pub cache_a: WASPCache,
1484    pub cache_b: WASPCache,
1485    pub l: usize,
1486    pub d_theta: f64,
1487    pub num_f_calls: Arc<Mutex<usize>>
1488}
1489impl WASPEc {
1490    pub fn new(n: usize, m: usize, alpha: f64, orthonormalize: bool, l: usize, d_theta: f64) -> Self {
1491        assert!(l <= m);
1492
1493        Self {
1494            cache_a: WASPCache::new(n, m, alpha, orthonormalize),
1495            cache_b: WASPCache::new(n, m, alpha, orthonormalize),
1496            l,
1497            d_theta,
1498            num_f_calls: Arc::new(Mutex::new(0)),
1499        }
1500    }
1501
1502    pub fn get_num_f_calls(&self) -> usize {
1503        self.num_f_calls.lock().unwrap().clone()
1504    }
1505}
1506impl DerivativeMethodTrait for WASPEc {
1507    type T = f64;
1508
1509    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(&self, inputs: &[f64], function: &D) -> (Vec<f64>, DMatrix<f64>) {
1510        let mut num_f_calls = self.num_f_calls.lock().unwrap();
1511        *num_f_calls = 0;
1512        let f_x_k_vec = function.call(inputs, false);
1513        let f_x_k = DVector::from_column_slice(&f_x_k_vec);
1514        *num_f_calls += 1;
1515
1516        loop {
1517            let d_a_t = wasp_projection(function, &f_x_k, inputs, &self.cache_a);
1518            let d_b_t = wasp_projection(function, &f_x_k, inputs, &self.cache_b);
1519            *num_f_calls += 2;
1520            let mut i_a = self.cache_a.i.lock().unwrap();
1521            let mut i_b = self.cache_b.i.lock().unwrap();
1522            *i_a = (*i_a + 1) % inputs.len();
1523            *i_b = (*i_b + 1) % inputs.len();
1524
1525            if close_enough(&d_a_t, &d_b_t, self.l, function.num_outputs(), self.d_theta) {
1526                return (f_x_k_vec, (d_a_t.transpose() + d_b_t.transpose()) * 0.5);
1527            }
1528        }
1529    }
1530}
1531
1532*/
1533
1534#[cfg(feature = "std")]
1535#[derive(Clone)]
1536pub struct SPSA;
1537#[cfg(feature = "std")]
1538impl SPSA {
1539    pub fn new() -> Self {
1540        Self {}
1541    }
1542}
1543#[cfg(feature = "std")]
1544impl DerivativeMethodTrait for SPSA {
1545    type T = f64;
1546
1547    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
1548        &self,
1549        inputs: &[f64],
1550        function: &D,
1551    ) -> (Vec<f64>, DMatrix<f64>) {
1552        let f0 = function.call(inputs, false);
1553
1554        let mut rng = rand::thread_rng();
1555
1556        let epsilon = 0.00000001;
1557
1558        let r: Vec<f64> = (0..inputs.len())
1559            .into_iter()
1560            .map(|_x| rng.gen_range(-1.0..=1.0))
1561            .collect();
1562        let x = DVector::from_column_slice(inputs);
1563        let delta_k = DVector::from_column_slice(&r);
1564        let xpos: DVector<f64> = &x + epsilon * &delta_k;
1565        let xneg: DVector<f64> = &x - epsilon * &delta_k;
1566        let fpos = DVector::from_column_slice(&function.call(xpos.as_slice(), false));
1567        let fneg = DVector::from_column_slice(&function.call(xneg.as_slice(), false));
1568        let v = (&fpos - &fneg) / (2.0 * epsilon);
1569        let delta_k_inverse =
1570            DVector::from_column_slice(&delta_k.iter().map(|x| 1.0 / *x).collect::<Vec<f64>>());
1571        let out = &v * &delta_k_inverse.transpose();
1572
1573        (f0, out)
1574    }
1575}
1576
1577pub struct DerivativeMethodClassAlwaysZero;
1578impl DerivativeMethodClass for DerivativeMethodClassAlwaysZero {
1579    type DerivativeMethod = DerivativeAlwaysZero;
1580}
1581
1582#[derive(Clone)]
1583pub struct DerivativeAlwaysZero;
1584impl DerivativeAlwaysZero {
1585    pub fn new() -> Self {
1586        Self {}
1587    }
1588}
1589impl DerivativeMethodTrait for DerivativeAlwaysZero {
1590    type T = f64;
1591
1592    fn derivative<D: DifferentiableFunctionTrait<Self::T> + ?Sized>(
1593        &self,
1594        _inputs: &[f64],
1595        function: &D,
1596    ) -> (Vec<f64>, DMatrix<f64>) {
1597        let num_outputs = function.num_outputs();
1598        let num_inputs = function.num_inputs();
1599        (
1600            vec![0.0; num_outputs],
1601            DMatrix::from_vec(num_outputs, num_inputs, vec![0.0; num_outputs * num_inputs]),
1602        )
1603    }
1604}