@@ -7,18 +7,49 @@ pub enum OptimizationLossMetric {
77 Mse ,
88 Mae ,
99 SoftL1 ,
10+ Chebyshev ,
11+ Msle ,
1012}
1113
1214impl OptimizationLossMetric {
1315 /// Полный список вариантов для UI и переборов.
14- pub const ALL : [ Self ; 3 ] = [ Self :: Mse , Self :: Mae , Self :: SoftL1 ] ;
16+ pub const ALL : [ Self ; 5 ] = [
17+ Self :: Mse ,
18+ Self :: Mae ,
19+ Self :: SoftL1 ,
20+ Self :: Chebyshev ,
21+ Self :: Msle ,
22+ ] ;
1523
1624 /// Короткое имя метрики для подписи в легенде.
1725 pub fn id ( self ) -> & ' static str {
1826 match self {
1927 Self :: Mse => "mse" ,
2028 Self :: Mae => "mae" ,
2129 Self :: SoftL1 => "soft_l1" ,
30+ Self :: Chebyshev => "chebyshev" ,
31+ Self :: Msle => "msle" ,
32+ }
33+ }
34+
35+ #[ inline]
36+ pub ( super ) fn simd_fast_path_supported ( self ) -> bool {
37+ matches ! ( self , Self :: Mse | Self :: Mae | Self :: SoftL1 )
38+ }
39+
40+ #[ inline]
41+ pub ( super ) fn requires_numerical_hessian ( self ) -> bool {
42+ matches ! ( self , Self :: Mae | Self :: Chebyshev )
43+ }
44+
45+ #[ inline]
46+ fn signum_or_zero ( value : f64 ) -> f64 {
47+ if value > 0.0 {
48+ 1.0
49+ } else if value < 0.0 {
50+ -1.0
51+ } else {
52+ 0.0
2253 }
2354 }
2455
@@ -27,32 +58,52 @@ impl OptimizationLossMetric {
2758 Self :: Mse => residual * residual,
2859 Self :: Mae => residual. abs ( ) ,
2960 Self :: SoftL1 => 2.0 * ( ( 1.0 + residual * residual) . sqrt ( ) - 1.0 ) ,
61+ Self :: Chebyshev => residual. abs ( ) ,
62+ Self :: Msle => {
63+ let log_term = ( 1.0 + residual. abs ( ) ) . ln ( ) ;
64+ log_term * log_term
65+ }
3066 }
3167 }
3268
3369 pub ( super ) fn residual_derivative ( self , residual : f64 ) -> f64 {
3470 match self {
3571 Self :: Mse => 2.0 * residual,
36- Self :: Mae => {
37- if residual > 0.0 {
38- 1.0
39- } else if residual < 0.0 {
40- -1.0
41- } else {
42- 0.0
43- }
44- }
72+ Self :: Mae | Self :: Chebyshev => Self :: signum_or_zero ( residual) ,
4573 Self :: SoftL1 => 2.0 * residual / ( 1.0 + residual * residual) . sqrt ( ) ,
74+ Self :: Msle => {
75+ let abs_residual = residual. abs ( ) ;
76+ let log_term = ( 1.0 + abs_residual) . ln ( ) ;
77+ Self :: signum_or_zero ( residual) * ( 2.0 * log_term / ( 1.0 + abs_residual) )
78+ }
4679 }
4780 }
4881
4982 pub ( super ) fn residual_second_derivative ( self , residual : f64 ) -> f64 {
5083 match self {
5184 Self :: Mse => 2.0 ,
52- Self :: Mae => 0.0 ,
85+ Self :: Mae | Self :: Chebyshev => 0.0 ,
5386 Self :: SoftL1 => 2.0 / ( 1.0 + residual * residual) . powf ( 1.5 ) ,
87+ Self :: Msle => {
88+ let abs_residual = residual. abs ( ) ;
89+ let one_plus_abs = 1.0 + abs_residual;
90+ let log_term = one_plus_abs. ln ( ) ;
91+ 2.0 * ( 1.0 - log_term) / ( one_plus_abs * one_plus_abs)
92+ }
5493 }
5594 }
95+
96+ pub ( super ) fn value_from_prediction ( self , prediction : f64 , target : f64 ) -> f64 {
97+ self . value_from_residual ( prediction - target)
98+ }
99+
100+ pub ( super ) fn prediction_derivative ( self , prediction : f64 , target : f64 ) -> f64 {
101+ self . residual_derivative ( prediction - target)
102+ }
103+
104+ pub ( super ) fn prediction_second_derivative ( self , prediction : f64 , target : f64 ) -> f64 {
105+ self . residual_second_derivative ( prediction - target)
106+ }
56107}
57108
58109/// Значение по умолчанию для числа знаков после запятой в режиме квантизации метрик.
@@ -125,11 +176,6 @@ impl ResidualQuantizer {
125176 Self :: Enabled { scale } => ( value * scale) . round ( ) / scale,
126177 }
127178 }
128-
129- #[ inline]
130- pub ( super ) fn residual ( self , predicted : f64 , observed : f64 ) -> f64 {
131- self . quantize_value ( predicted) - self . quantize_value ( observed)
132- }
133179}
134180
135181#[ derive( Debug , Clone , Copy , PartialEq ) ]
@@ -211,6 +257,7 @@ pub(super) struct ScalarMetrics {
211257 pub ( super ) rmse : f64 ,
212258 pub ( super ) mae : f64 ,
213259 pub ( super ) soft_l1 : f64 ,
260+ pub ( super ) msle : f64 ,
214261 pub ( super ) r2 : f64 ,
215262 pub ( super ) max_abs_error : f64 ,
216263}
@@ -223,6 +270,8 @@ pub(super) fn scalar_loss_value(
223270 OptimizationLossMetric :: Mse => metrics. mse ,
224271 OptimizationLossMetric :: Mae => metrics. mae ,
225272 OptimizationLossMetric :: SoftL1 => metrics. soft_l1 ,
273+ OptimizationLossMetric :: Chebyshev => metrics. max_abs_error ,
274+ OptimizationLossMetric :: Msle => metrics. msle ,
226275 }
227276}
228277
@@ -246,13 +295,17 @@ where
246295 let mut sse = 0.0 ;
247296 let mut sae = 0.0 ;
248297 let mut soft_l1_sum = 0.0 ;
298+ let mut msle_sum = 0.0 ;
249299 let mut max_abs_error = 0.0_f64 ;
250300 for point in points. as_slice ( ) {
251- let residual = quantizer. residual ( evaluate ( point. x ( ) ) , point. y ( ) ) ;
301+ let predicted = quantizer. quantize_value ( evaluate ( point. x ( ) ) ) ;
302+ let observed = quantizer. quantize_value ( point. y ( ) ) ;
303+ let residual = predicted - observed;
252304 let abs_residual = residual. abs ( ) ;
253305 sse += residual * residual;
254306 sae += abs_residual;
255307 soft_l1_sum += OptimizationLossMetric :: SoftL1 . value_from_residual ( residual) ;
308+ msle_sum += OptimizationLossMetric :: Msle . value_from_prediction ( predicted, observed) ;
256309 max_abs_error = max_abs_error. max ( abs_residual) ;
257310 }
258311
@@ -268,6 +321,7 @@ where
268321 let rmse = mse. sqrt ( ) ;
269322 let mae = sae / sample_count;
270323 let soft_l1 = soft_l1_sum / sample_count;
324+ let msle = msle_sum / sample_count;
271325 let r2 = if sst <= 1e-15 {
272326 if sse <= 1e-15 { 1.0 } else { 0.0 }
273327 } else {
@@ -279,6 +333,7 @@ where
279333 rmse,
280334 mae,
281335 soft_l1,
336+ msle,
282337 r2,
283338 max_abs_error,
284339 }
@@ -289,6 +344,7 @@ pub(super) struct EvaluatorMetrics {
289344 pub ( super ) rmse : f64 ,
290345 pub ( super ) mae : f64 ,
291346 pub ( super ) soft_l1 : f64 ,
347+ pub ( super ) msle : f64 ,
292348 pub ( super ) r2 : f64 ,
293349 pub ( super ) max_abs_error : f64 ,
294350 pub ( super ) residuals : Vec < [ f64 ; 2 ] > ,
@@ -316,6 +372,7 @@ where
316372 rmse : scalar. rmse ,
317373 mae : scalar. mae ,
318374 soft_l1 : scalar. soft_l1 ,
375+ msle : scalar. msle ,
319376 r2 : scalar. r2 ,
320377 max_abs_error : scalar. max_abs_error ,
321378 residuals,
0 commit comments