Skip to content

Commit 1247c85

Browse files
committed
added gradient diagnostic
1 parent 7f5e17f commit 1247c85

11 files changed

Lines changed: 278 additions & 15 deletions

src/app.rs

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -43,7 +43,8 @@ use crate::fit::{
4343
default_spline_initial_knot_y, sample_curve,
4444
};
4545
use crate::fit::{
46-
IncrementalFitRunner, IncrementalFitStep, IncrementalSplineFitRunner, IncrementalSplineFitStep,
46+
GradientIterationDiagnostics, IncrementalFitRunner, IncrementalFitStep,
47+
IncrementalSplineFitRunner, IncrementalSplineFitStep,
4748
};
4849

4950
use self::diagnostics::{IterationDiagnostics, diagnostics_plot_y_axis_width};

src/app/diagnostics.rs

Lines changed: 34 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -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

src/app/fit_worker.rs

Lines changed: 22 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -74,7 +74,12 @@ impl CurveFitApp {
7474
replay_frames: &mut Vec<ReplayFrame>,
7575
) {
7676
for entry in trace {
77-
diagnostics.append(entry.iteration, entry.metrics, &entry.params);
77+
diagnostics.append(
78+
entry.iteration,
79+
entry.metrics,
80+
entry.gradient_diagnostics,
81+
&entry.params,
82+
);
7883
Self::upsert_buffered_parametric_replay_frame(
7984
replay_frames,
8085
entry.iteration,
@@ -89,7 +94,12 @@ impl CurveFitApp {
8994
replay_frames: &mut Vec<ReplayFrame>,
9095
) {
9196
for entry in trace {
92-
diagnostics.append_spline(entry.iteration, entry.metrics, &entry.knot_y);
97+
diagnostics.append_spline(
98+
entry.iteration,
99+
entry.metrics,
100+
entry.gradient_diagnostics,
101+
&entry.knot_y,
102+
);
93103
Self::upsert_buffered_spline_replay_frame(
94104
replay_frames,
95105
entry.iteration,
@@ -218,7 +228,7 @@ impl CurveFitApp {
218228
// Финальные метрики пересчитываем по тем точкам и параметрам, которые видит пользователь.
219229
let (metrics, snapshot_metrics, snapshot_residuals) =
220230
self.parametric_metrics_and_residuals(points, &result.params);
221-
diagnostics.append(result.iterations, metrics, &result.params);
231+
diagnostics.append(result.iterations, metrics, None, &result.params);
222232
result_metrics = Some(snapshot_metrics);
223233
residual_plot_points = snapshot_residuals;
224234
}
@@ -257,7 +267,7 @@ impl CurveFitApp {
257267
Self::apply_spline_trace_to_buffers(trace, &mut diagnostics, &mut replay_frames);
258268
let knot_y = result.knots.iter().map(|knot| knot[1]).collect::<Vec<_>>();
259269
let spline_plot_curve = Self::plot_points_from_pairs(result.curve.iter().copied());
260-
diagnostics.append_spline(result.iterations, metrics, &knot_y);
270+
diagnostics.append_spline(result.iterations, metrics, None, &knot_y);
261271
Self::upsert_buffered_spline_replay_frame(
262272
&mut replay_frames,
263273
result.iterations,
@@ -466,6 +476,7 @@ impl CurveFitApp {
466476
iteration,
467477
mse: _,
468478
metrics,
479+
gradient_diagnostics,
469480
params,
470481
}) => {
471482
let params = if let Some(normalization) = normalization {
@@ -498,6 +509,7 @@ impl CurveFitApp {
498509
iteration_trace.push(ParametricIterationTraceEntry {
499510
iteration,
500511
metrics,
512+
gradient_diagnostics,
501513
params,
502514
});
503515
}
@@ -546,12 +558,14 @@ impl CurveFitApp {
546558
iteration,
547559
mse: _,
548560
metrics,
561+
gradient_diagnostics,
549562
knot_y,
550563
curve,
551564
}) => {
552565
iteration_trace.push(SplineIterationTraceEntry {
553566
iteration,
554567
metrics,
568+
gradient_diagnostics,
555569
knot_y,
556570
curve,
557571
});
@@ -620,6 +634,7 @@ impl CurveFitApp {
620634
iteration,
621635
mse: _,
622636
metrics,
637+
gradient_diagnostics,
623638
params,
624639
}) => {
625640
let params = if let Some(normalization) = normalization {
@@ -643,6 +658,7 @@ impl CurveFitApp {
643658
iteration_trace.push(ParametricIterationTraceEntry {
644659
iteration,
645660
metrics,
661+
gradient_diagnostics,
646662
params,
647663
});
648664
}
@@ -735,12 +751,14 @@ impl CurveFitApp {
735751
iteration,
736752
mse: _,
737753
metrics,
754+
gradient_diagnostics,
738755
knot_y,
739756
curve,
740757
}) => {
741758
iteration_trace.push(SplineIterationTraceEntry {
742759
iteration,
743760
metrics,
761+
gradient_diagnostics,
744762
knot_y,
745763
curve,
746764
});

src/app/panel_state.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
pub(super) enum DiagnosticsTab {
66
#[default]
77
Loss,
8+
Gradient,
89
Residuals,
910
}
1011

src/app/state.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -190,6 +190,7 @@ pub(super) enum WasmFitJob {
190190
pub(super) struct ParametricIterationTraceEntry {
191191
pub(super) iteration: u64,
192192
pub(super) metrics: IterationMetricSnapshot,
193+
pub(super) gradient_diagnostics: Option<GradientIterationDiagnostics>,
193194
pub(super) params: CurveParams,
194195
}
195196

@@ -198,6 +199,7 @@ pub(super) struct ParametricIterationTraceEntry {
198199
pub(super) struct SplineIterationTraceEntry {
199200
pub(super) iteration: u64,
200201
pub(super) metrics: IterationMetricSnapshot,
202+
pub(super) gradient_diagnostics: Option<GradientIterationDiagnostics>,
201203
pub(super) knot_y: Vec<f64>,
202204
pub(super) curve: Vec<[f64; 2]>,
203205
}

src/app/tests.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,8 @@ use super::{
2424
};
2525
use crate::domain::{CurveFamily, CurveParams, FitResult, OptimizerConfig, Point, Points};
2626
use crate::fit::{
27-
DEFAULT_METRIC_QUANTIZATION_DECIMAL_PLACES, IterationMetricSnapshot, MetricQuantization,
28-
OptimizationLossMetric, SplineResult,
27+
DEFAULT_METRIC_QUANTIZATION_DECIMAL_PLACES, GradientIterationDiagnostics,
28+
IterationMetricSnapshot, MetricQuantization, OptimizationLossMetric, SplineResult,
2929
};
3030

3131
// Все import-ы специально поднимаются в этот модуль,

src/app/tests/fit_lifecycle.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -245,11 +245,13 @@ fn finished_message_applies_buffered_parametric_trace_in_single_poll() {
245245
super::ParametricIterationTraceEntry {
246246
iteration: 1,
247247
metrics: metrics_snapshot(2.0, 2.0, 1.4, 1.2, 1.1, 0.1, 2.0),
248+
gradient_diagnostics: None,
248249
params: CurveParams::Linear { a: 0.4, b: 0.8 },
249250
},
250251
super::ParametricIterationTraceEntry {
251252
iteration: 2,
252253
metrics: metrics_snapshot(0.5, 0.5, 0.7, 0.6, 0.55, 0.7, 0.9),
254+
gradient_diagnostics: None,
253255
params: CurveParams::Linear { a: 1.2, b: 0.5 },
254256
},
255257
];

src/app/tests/replay_diagnostics.rs

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,11 +43,13 @@ fn diagnostics_append_replaces_duplicate_iteration() {
4343
diagnostics.append(
4444
2,
4545
metrics_snapshot(5.0, 5.0, 5.0_f64.sqrt(), 2.0, 1.75, -1.5, 3.0),
46+
None,
4647
&CurveParams::Linear { a: 1.0, b: 0.0 },
4748
);
4849
diagnostics.append(
4950
2,
5051
metrics_snapshot(3.0, 3.0, 3.0_f64.sqrt(), 1.5, 1.1, -2.0, 2.0),
52+
None,
5153
&CurveParams::Linear { a: -1.5, b: 0.5 },
5254
);
5355

@@ -65,6 +67,25 @@ fn diagnostics_append_replaces_duplicate_iteration() {
6567
assert_eq!(diagnostics.parameter_series[1][1], [2.0, 0.5]);
6668
}
6769

70+
#[test]
71+
fn diagnostics_append_tracks_gradient_norm_and_direction() {
72+
let mut diagnostics = IterationDiagnostics::default();
73+
let params = CurveParams::Linear { a: 1.0, b: 0.0 };
74+
75+
diagnostics.append(
76+
3,
77+
metrics_snapshot(1.0, 1.0, 1.0, 1.0, 1.0, 0.0, 1.0),
78+
Some(GradientIterationDiagnostics {
79+
gradient_l2_norm: 1e-3,
80+
gradient_cosine: Some(-0.25),
81+
}),
82+
&params,
83+
);
84+
85+
assert_eq!(diagnostics.gradient_log_l2_norm_points, vec![[3.0, -3.0]]);
86+
assert_eq!(diagnostics.gradient_cosine_points, vec![[3.0, -0.25]]);
87+
}
88+
6889
#[test]
6990
fn diagnostics_append_resets_when_family_changes() {
7091
let points = line_points();
@@ -78,6 +99,7 @@ fn diagnostics_append_resets_when_family_changes() {
7899
diagnostics.append(
79100
4,
80101
metrics_snapshot(1.0, 1.0, 1.0, 1.0, 0.8, 0.2, 1.0),
102+
None,
81103
&CurveParams::Quadratic {
82104
a: 1.0,
83105
b: -2.0,
@@ -103,11 +125,13 @@ fn diagnostics_append_spline_tracks_knot_parameters() {
103125
diagnostics.append_spline(
104126
1,
105127
metrics_snapshot(2.5, 2.5, 2.5_f64.sqrt(), 1.2, 1.6, -0.25, 2.0),
128+
None,
106129
&[0.5, -1.0],
107130
);
108131
diagnostics.append_spline(
109132
2,
110133
metrics_snapshot(1.5, 1.5, 1.5_f64.sqrt(), 0.9, 1.0, 0.25, 1.4),
134+
None,
111135
&[0.75, -0.25],
112136
);
113137

0 commit comments

Comments
 (0)