@@ -14,6 +14,8 @@ pub(super) struct IterationDiagnostics {
1414 pub ( super ) soft_l1_points : Vec < [ f64 ; 2 ] > ,
1515 pub ( super ) r2_abs_points : Vec < [ f64 ; 2 ] > ,
1616 pub ( super ) max_abs_error_points : Vec < [ f64 ; 2 ] > ,
17+ pub ( super ) gradient_log_l2_norm_points : Vec < [ f64 ; 2 ] > ,
18+ pub ( super ) gradient_cosine_points : Vec < [ f64 ; 2 ] > ,
1719 pub ( super ) parameter_names : Vec < String > ,
1820 pub ( super ) parameter_series : Vec < Vec < [ f64 ; 2 ] > > ,
1921}
@@ -42,13 +44,14 @@ impl IterationDiagnostics {
4244 loss_metric,
4345 metric_quantization,
4446 ) ;
45- self . append ( 0 , metrics, params) ;
47+ self . append ( 0 , metrics, None , params) ;
4648 }
4749
4850 pub ( super ) fn append (
4951 & mut self ,
5052 iteration : u64 ,
5153 metrics : IterationMetricSnapshot ,
54+ gradient_diagnostics : Option < GradientIterationDiagnostics > ,
5255 params : & CurveParams ,
5356 ) {
5457 let family = params. family ( ) ;
@@ -61,7 +64,7 @@ impl IterationDiagnostics {
6164 }
6265
6366 let iteration = iteration as f64 ;
64- self . upsert_metrics ( iteration, metrics) ;
67+ self . upsert_metrics ( iteration, metrics, gradient_diagnostics ) ;
6568 params. with_values ( |values| {
6669 for ( series, value) in self . parameter_series . iter_mut ( ) . zip ( values. iter ( ) . copied ( ) ) {
6770 upsert_iteration_point ( series, iteration, value) ;
@@ -73,6 +76,7 @@ impl IterationDiagnostics {
7376 & mut self ,
7477 iteration : u64 ,
7578 metrics : IterationMetricSnapshot ,
79+ gradient_diagnostics : Option < GradientIterationDiagnostics > ,
7680 knot_y : & [ f64 ] ,
7781 ) {
7882 let parameter_count = knot_y. len ( ) ;
@@ -81,13 +85,18 @@ impl IterationDiagnostics {
8185 }
8286
8387 let iteration = iteration as f64 ;
84- self . upsert_metrics ( iteration, metrics) ;
88+ self . upsert_metrics ( iteration, metrics, gradient_diagnostics ) ;
8589 for ( series, value) in self . parameter_series . iter_mut ( ) . zip ( knot_y. iter ( ) . copied ( ) ) {
8690 upsert_iteration_point ( series, iteration, value) ;
8791 }
8892 }
8993
90- fn upsert_metrics ( & mut self , iteration : f64 , metrics : IterationMetricSnapshot ) {
94+ fn upsert_metrics (
95+ & mut self ,
96+ iteration : f64 ,
97+ metrics : IterationMetricSnapshot ,
98+ gradient_diagnostics : Option < GradientIterationDiagnostics > ,
99+ ) {
91100 upsert_iteration_point ( & mut self . loss_points , iteration, metrics. loss ) ;
92101 upsert_iteration_point ( & mut self . mse_points , iteration, metrics. mse ) ;
93102 upsert_iteration_point ( & mut self . rmse_points , iteration, metrics. rmse ) ;
@@ -99,6 +108,25 @@ impl IterationDiagnostics {
99108 iteration,
100109 metrics. max_abs_error ,
101110 ) ;
111+ if let Some ( gradient_diagnostics) = gradient_diagnostics
112+ && gradient_diagnostics. gradient_l2_norm . is_finite ( )
113+ && gradient_diagnostics. gradient_l2_norm > 0.0
114+ {
115+ upsert_iteration_point (
116+ & mut self . gradient_log_l2_norm_points ,
117+ iteration,
118+ gradient_diagnostics. gradient_l2_norm . log10 ( ) ,
119+ ) ;
120+ if let Some ( gradient_cosine) = gradient_diagnostics. gradient_cosine
121+ && gradient_cosine. is_finite ( )
122+ {
123+ upsert_iteration_point (
124+ & mut self . gradient_cosine_points ,
125+ iteration,
126+ gradient_cosine,
127+ ) ;
128+ }
129+ }
102130 }
103131
104132 fn reset_for_family ( & mut self , family : CurveFamily ) {
@@ -136,6 +164,8 @@ impl IterationDiagnostics {
136164 self . soft_l1_points . clear ( ) ;
137165 self . r2_abs_points . clear ( ) ;
138166 self . max_abs_error_points . clear ( ) ;
167+ self . gradient_log_l2_norm_points . clear ( ) ;
168+ self . gradient_cosine_points . clear ( ) ;
139169 }
140170}
141171
0 commit comments