11use super :: common:: { is_finite_non_negative, scale_and_mirror_upper_hessian, stabilize_hessian} ;
22use ndarray:: Array2 ;
33
4+ /// Вычисляет арктангенс-ступень:
5+ /// `f(x) = amplitude * atan(slope * (x - x0)) + offset`,
6+ /// где:
7+ /// - `amplitude` — амплитуда перехода,
8+ /// - `slope` — крутизна перехода,
9+ /// - `x0` — центр перехода по оси `x`,
10+ /// - `offset` — вертикальный сдвиг.
411#[ inline]
512pub ( super ) fn eval ( param : & [ f64 ] , x : f64 ) -> f64 {
6- param[ 0 ] * ( param[ 1 ] * ( x - param[ 2 ] ) ) . atan ( ) + param[ 3 ]
13+ let amplitude = param[ 0 ] ;
14+ let slope = param[ 1 ] ;
15+ let x0 = param[ 2 ] ;
16+ let offset = param[ 3 ] ;
17+ amplitude * ( slope * ( x - x0) ) . atan ( ) + offset
718}
819
920pub ( super ) fn accumulate_gradient < L > (
@@ -16,20 +27,24 @@ pub(super) fn accumulate_gradient<L>(
1627 L : FnMut ( f64 , f64 ) -> f64 ,
1728{
1829 debug_assert_eq ! ( x_values. len( ) , y_values. len( ) ) ;
30+ let amplitude = param[ 0 ] ;
31+ let slope = param[ 1 ] ;
32+ let x0 = param[ 2 ] ;
33+ let offset = param[ 3 ] ;
1934
2035 let mut index = 0 ;
2136 while index < x_values. len ( ) {
2237 let x = x_values[ index] ;
2338 let y = y_values[ index] ;
24- let z = param [ 1 ] * ( x - param [ 2 ] ) ;
39+ let z = slope * ( x - x0 ) ;
2540 let atan_z = z. atan ( ) ;
2641 let inv_den = 1.0 / ( 1.0 + z * z) ;
27- let model = param [ 0 ] * atan_z + param [ 3 ] ;
42+ let model = amplitude * atan_z + offset ;
2843 let residual = loss_derivative_from_prediction ( model, y) ;
2944
3045 gradient[ 0 ] += residual * atan_z;
31- gradient[ 1 ] += residual * ( param [ 0 ] * ( x - param [ 2 ] ) * inv_den) ;
32- gradient[ 2 ] += residual * ( -param [ 0 ] * param [ 1 ] * inv_den) ;
46+ gradient[ 1 ] += residual * ( amplitude * ( x - x0 ) * inv_den) ;
47+ gradient[ 2 ] += residual * ( -amplitude * slope * inv_den) ;
3348 gradient[ 3 ] += residual;
3449 index += 1 ;
3550 }
@@ -53,20 +68,21 @@ where
5368 let sample_count = x_values. len ( ) ;
5469 let sample_scale = 1.0 / sample_count as f64 ;
5570 let mut hessian = Array2 :: zeros ( ( 4 , 4 ) ) ;
71+ let amplitude = param[ 0 ] ;
72+ let slope = param[ 1 ] ;
73+ let x0 = param[ 2 ] ;
74+ let offset = param[ 3 ] ;
5675
5776 let mut index = 0 ;
5877 while index < sample_count {
5978 let x = x_values[ index] ;
6079 let y = y_values[ index] ;
61- let a = param[ 0 ] ;
62- let b = param[ 1 ] ;
63- let c = param[ 2 ] ;
64- let u = x - c;
65- let z = b * u;
80+ let u = x - x0;
81+ let z = slope * u;
6682 let atan_z = z. atan ( ) ;
6783 let inv_den = 1.0 / ( 1.0 + z * z) ;
6884 let d2_shape_dz2 = -2.0 * z * inv_den * inv_den;
69- let model = a * atan_z + param [ 3 ] ;
85+ let model = amplitude * atan_z + offset ;
7086 if !model. is_finite ( ) {
7187 return None ;
7288 }
@@ -78,15 +94,15 @@ where
7894 }
7995
8096 let jac_a = atan_z;
81- let jac_b = a * inv_den * u;
82- let jac_c = -a * inv_den * b ;
97+ let jac_b = amplitude * inv_den * u;
98+ let jac_c = -amplitude * inv_den * slope ;
8399 let jac_d = 1.0 ;
84100
85101 let d2_model_dadb = inv_den * u;
86- let d2_model_dadc = -inv_den * b ;
87- let d2_model_dbdb = a * d2_shape_dz2 * u * u;
88- let d2_model_dbdc = a * ( d2_shape_dz2 * ( -b ) * u - inv_den) ;
89- let d2_model_dcdc = a * d2_shape_dz2 * b * b ;
102+ let d2_model_dadc = -inv_den * slope ;
103+ let d2_model_dbdb = amplitude * d2_shape_dz2 * u * u;
104+ let d2_model_dbdc = amplitude * ( d2_shape_dz2 * ( -slope ) * u - inv_den) ;
105+ let d2_model_dcdc = amplitude * d2_shape_dz2 * slope * slope ;
90106
91107 hessian[ [ 0 , 0 ] ] += loss_second * jac_a * jac_a;
92108 hessian[ [ 0 , 1 ] ] += loss_second * jac_a * jac_b + loss_first * d2_model_dadb;
0 commit comments