@@ -18,31 +18,49 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1818 amplitude * ( -damping * x) . exp ( ) * ( omega * x + phi) . sin ( ) + offset
1919}
2020
21+ #[ inline]
22+ pub ( super ) fn value_grad_at ( param : & [ f64 ] , x : f64 , grad : & mut [ f64 ] ) -> f64 {
23+ debug_assert_eq ! ( grad. len( ) , 5 ) ;
24+
25+ let amplitude = param[ 0 ] ;
26+ let damping = param[ 1 ] ;
27+ let omega = param[ 2 ] ;
28+ let phi = param[ 3 ] ;
29+ let offset = param[ 4 ] ;
30+ let exp_part = ( -damping * x) . exp ( ) ;
31+ let angle = omega * x + phi;
32+ let sin_part = angle. sin ( ) ;
33+ let cos_part = angle. cos ( ) ;
34+
35+ grad[ 0 ] = exp_part * sin_part;
36+ grad[ 1 ] = -amplitude * x * exp_part * sin_part;
37+ grad[ 2 ] = amplitude * exp_part * cos_part * x;
38+ grad[ 3 ] = amplitude * exp_part * cos_part;
39+ grad[ 4 ] = 1.0 ;
40+
41+ amplitude * exp_part * sin_part + offset
42+ }
43+
2144pub ( super ) fn add_value_grad (
2245 x_values : & [ f64 ] ,
2346 param : & [ f64 ] ,
2447 value_first : & [ f64 ] ,
2548 gradient : & mut [ f64 ] ,
2649) {
27- let amplitude = param[ 0 ] ;
28- let damping = param[ 1 ] ;
29- let omega = param[ 2 ] ;
30- let phi = param[ 3 ] ;
50+ debug_assert_eq ! ( x_values. len( ) , value_first. len( ) ) ;
51+ debug_assert_eq ! ( gradient. len( ) , param. len( ) ) ;
3152
53+ let mut point_grad = [ 0.0 ; 5 ] ;
3254 let mut index = 0 ;
3355 while index < x_values. len ( ) {
34- let x = x_values[ index] ;
35- let exp_part = ( -damping * x) . exp ( ) ;
36- let angle = omega * x + phi;
37- let sin_part = angle. sin ( ) ;
38- let cos_part = angle. cos ( ) ;
39- let residual = value_first[ index] ;
40-
41- gradient[ 0 ] += residual * exp_part * sin_part;
42- gradient[ 1 ] += residual * ( -amplitude * x * exp_part * sin_part) ;
43- gradient[ 2 ] += residual * ( amplitude * exp_part * cos_part * x) ;
44- gradient[ 3 ] += residual * ( amplitude * exp_part * cos_part) ;
45- gradient[ 4 ] += residual;
56+ let upstream = value_first[ index] ;
57+ value_grad_at ( param, x_values[ index] , & mut point_grad) ;
58+
59+ gradient[ 0 ] += upstream * point_grad[ 0 ] ;
60+ gradient[ 1 ] += upstream * point_grad[ 1 ] ;
61+ gradient[ 2 ] += upstream * point_grad[ 2 ] ;
62+ gradient[ 3 ] += upstream * point_grad[ 3 ] ;
63+ gradient[ 4 ] += upstream * point_grad[ 4 ] ;
4664 index += 1 ;
4765 }
4866}
0 commit comments