Skip to content

Commit d58f135

Browse files
committed
added fit results export to json
1 parent c859d13 commit d58f135

8 files changed

Lines changed: 1057 additions & 17 deletions

File tree

Cargo.lock

Lines changed: 286 additions & 10 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,11 +17,14 @@ argmin-math = { version = "0.5.1", default-features = false, features = [
1717
"vec",
1818
"ndarray_v0_16-nolinalg",
1919
] }
20+
chrono = "0.4.42"
2021
egui_extras = { version = "0.34.1", default-features = false, features = [
2122
"svg_text",
2223
] }
2324
egui_plot = "0.35.0"
2425
ndarray = "0.16.1"
26+
serde = { version = "1.0.228", features = ["derive"] }
27+
serde_json = "1.0.149"
2528
stochastic_optimizers = "0.3.0"
2629

2730
[target.'cfg(not(target_arch = "wasm32"))'.dependencies]
@@ -31,6 +34,7 @@ eframe = { version = "0.34.1", default-features = false, features = [
3134
"x11",
3235
"wayland",
3336
] }
37+
egui-file-dialog = "0.13.0"
3438
sys-locale = "0.3.2"
3539

3640
[target.'cfg(target_arch = "wasm32")'.dependencies]

src/app.rs

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@ use std::time::Instant;
77
use web_time::Instant;
88

99
use eframe::egui;
10+
#[cfg(not(target_arch = "wasm32"))]
11+
use egui_file_dialog::FileDialog;
1012
use egui_plot::{
1113
Legend, Line, LineStyle, Plot, PlotBounds, PlotPoint, PlotPoints, PlotResponse,
1214
Points as PlotPointsItem, VLine,
@@ -25,6 +27,7 @@ mod plot_utils;
2527
mod points_state;
2628
mod points_text;
2729
mod replay;
30+
mod result_export;
2831
mod ui;
2932

3033
use self::diagnostics::{IterationDiagnostics, diagnostics_plot_y_axis_width};
@@ -62,6 +65,7 @@ use self::points_text::{
6265
use self::replay::ReplayState;
6366
#[cfg(test)]
6467
use self::replay::{ReplayFrame, ReplayFramePayload};
68+
use self::result_export::FitExportRecord;
6569
use crate::domain::{
6670
AdamConfig, CurveFamily, CurveParams, FitResult, LbfgsConfig, NelderMeadConfig, NewtonCgConfig,
6771
OptimizerConfig, OptimizerMethod, Point, Points, SgdConfig, SteepestDescentConfig,
@@ -792,6 +796,12 @@ pub struct CurveFitApp {
792796
fit_in_progress: bool,
793797
fit_loss_metric: OptimizationLossMetric,
794798
fit_metric_quantization: MetricQuantization,
799+
fit_optimizer_method: OptimizerMethod,
800+
fit_export_record: Option<FitExportRecord>,
801+
#[cfg(not(target_arch = "wasm32"))]
802+
fit_export_file_dialog: FileDialog,
803+
#[cfg(not(target_arch = "wasm32"))]
804+
fit_export_pending_json: Option<String>,
795805
fit_preview_params: Option<CurveParams>,
796806
fit_preview_iteration: Option<u64>,
797807
fit_started_at: Option<Instant>,
@@ -1125,6 +1135,7 @@ impl CurveFitApp {
11251135
self.reset_fit_timer();
11261136
self.fit_result = None;
11271137
self.spline_result = None;
1138+
self.clear_fit_export_state();
11281139
self.active_fit_points = None;
11291140
self.result_metrics = None;
11301141
self.residual_plot_points.clear();
@@ -1493,6 +1504,16 @@ impl Default for CurveFitApp {
14931504
fit_in_progress: false,
14941505
fit_loss_metric: OptimizationLossMetric::default(),
14951506
fit_metric_quantization: MetricQuantization::Disabled,
1507+
fit_optimizer_method: OptimizerMethod::default(),
1508+
fit_export_record: None,
1509+
#[cfg(not(target_arch = "wasm32"))]
1510+
fit_export_file_dialog: FileDialog::new()
1511+
.title("Save fit result JSON")
1512+
.add_save_extension("JSON files", "json")
1513+
.default_save_extension("JSON files")
1514+
.default_file_name("fit-result.json"),
1515+
#[cfg(not(target_arch = "wasm32"))]
1516+
fit_export_pending_json: None,
14961517
fit_preview_params: None,
14971518
fit_preview_iteration: None,
14981519
fit_started_at: None,
@@ -1527,6 +1548,8 @@ impl eframe::App for CurveFitApp {
15271548
self.maybe_run_pending_auto_refit();
15281549
self.tick_replay(ctx);
15291550
self.poll_points_clipboard_import(ctx);
1551+
#[cfg(not(target_arch = "wasm32"))]
1552+
self.poll_fit_export_save_dialog(ctx);
15301553
self.maybe_refresh_points_cache_after_debounce();
15311554

15321555
if !self.fit_in_progress {

src/app/fit_worker.rs

Lines changed: 30 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -137,6 +137,7 @@ impl CurveFitApp {
137137
Ok(FitWorkerMessage::Stopped) => {
138138
self.fit_in_progress = false;
139139
self.reset_fit_timer();
140+
self.clear_fit_export_state();
140141
self.active_fit_points = None;
141142
self.finalize_replay_after_fit_stopped();
142143
if !self.discard_fit_worker_updates {
@@ -151,6 +152,10 @@ impl CurveFitApp {
151152
self.fit_in_progress = false;
152153
let fit_points = self.active_fit_points.take();
153154
if !self.discard_fit_worker_updates {
155+
let point_count = fit_points
156+
.as_ref()
157+
.map(Points::len)
158+
.unwrap_or(self.residual_plot_points.len());
154159
if let Some(points) = fit_points.as_ref() {
155160
self.update_parametric_result_metrics(points, &result.params);
156161
let metrics = calculate_iteration_metrics_with_quantization(
@@ -170,11 +175,13 @@ impl CurveFitApp {
170175
result.params.clone(),
171176
);
172177
self.finalize_replay_after_fit_completion();
173-
self.fit_result = Some(result);
174178
self.complete_fit_timer_successfully();
179+
self.store_parametric_fit_export_record(&result, point_count);
180+
self.fit_result = Some(result);
175181
self.status = Some(StatusMessage::FitCompleted);
176182
} else {
177183
self.reset_fit_timer();
184+
self.clear_fit_export_state();
178185
self.set_fit_stopped_status_if_fitting();
179186
}
180187
keep_receiver = false;
@@ -183,6 +190,7 @@ impl CurveFitApp {
183190
Ok(FitWorkerMessage::SplineFinished { result, metrics }) => {
184191
self.fit_in_progress = false;
185192
if !self.discard_fit_worker_updates {
193+
let point_count = result.residuals.len();
186194
let knot_y = result.knots.iter().map(|knot| knot[1]).collect::<Vec<_>>();
187195
let spline_plot_curve =
188196
Self::plot_points_from_pairs(result.curve.iter().copied());
@@ -194,11 +202,13 @@ impl CurveFitApp {
194202
);
195203
self.upsert_spline_replay_frame(result.iterations, spline_plot_curve);
196204
self.finalize_replay_after_fit_completion();
197-
self.spline_result = Some(result);
198205
self.complete_fit_timer_successfully();
206+
self.store_spline_fit_export_record(&result, point_count);
207+
self.spline_result = Some(result);
199208
self.status = Some(StatusMessage::FitCompleted);
200209
} else {
201210
self.reset_fit_timer();
211+
self.clear_fit_export_state();
202212
self.set_fit_stopped_status_if_fitting();
203213
}
204214
self.active_fit_points = None;
@@ -208,6 +218,7 @@ impl CurveFitApp {
208218
Ok(FitWorkerMessage::Failed(error)) => {
209219
self.fit_in_progress = false;
210220
self.reset_fit_timer();
221+
self.clear_fit_export_state();
211222
self.active_fit_points = None;
212223
if !self.discard_fit_worker_updates {
213224
self.status = Some(StatusMessage::Error(error));
@@ -221,6 +232,7 @@ impl CurveFitApp {
221232
Err(TryRecvError::Disconnected) => {
222233
self.fit_in_progress = false;
223234
self.reset_fit_timer();
235+
self.clear_fit_export_state();
224236
self.active_fit_points = None;
225237
if !self.discard_fit_worker_updates {
226238
self.status = Some(StatusMessage::Error(
@@ -323,6 +335,7 @@ impl CurveFitApp {
323335
Err(error) => {
324336
self.fit_in_progress = false;
325337
self.reset_fit_timer();
338+
self.clear_fit_export_state();
326339
self.status = Some(StatusMessage::Error(error));
327340
self.active_fit_points = None;
328341
break;
@@ -354,13 +367,18 @@ impl CurveFitApp {
354367
Ok(params) => params,
355368
Err(error) => {
356369
self.reset_fit_timer();
370+
self.clear_fit_export_state();
357371
self.status = Some(StatusMessage::Error(error));
358372
self.active_fit_points = None;
359373
break;
360374
}
361375
};
362376
}
363377
let fit_points = self.active_fit_points.take();
378+
let point_count = fit_points
379+
.as_ref()
380+
.map(Points::len)
381+
.unwrap_or(self.residual_plot_points.len());
364382
if let Some(points) = fit_points.as_ref() {
365383
let (mse, rmse) = calculate_metrics_with_quantization(
366384
points,
@@ -384,14 +402,16 @@ impl CurveFitApp {
384402
}
385403
self.upsert_parametric_replay_frame(result.iterations, result.params.clone());
386404
self.finalize_replay_after_fit_completion();
387-
self.fit_result = Some(result);
388405
self.complete_fit_timer_successfully();
406+
self.store_parametric_fit_export_record(&result, point_count);
407+
self.fit_result = Some(result);
389408
self.status = Some(StatusMessage::FitCompleted);
390409
break;
391410
}
392411
Ok(IncrementalFitStep::Cancelled) => {
393412
self.fit_in_progress = false;
394413
self.reset_fit_timer();
414+
self.clear_fit_export_state();
395415
self.finalize_replay_after_fit_stopped();
396416
self.status = Some(StatusMessage::FitStopped);
397417
self.active_fit_points = None;
@@ -400,6 +420,7 @@ impl CurveFitApp {
400420
Err(error) => {
401421
self.fit_in_progress = false;
402422
self.reset_fit_timer();
423+
self.clear_fit_export_state();
403424
self.status = Some(StatusMessage::Error(error.to_string()));
404425
self.active_fit_points = None;
405426
break;
@@ -426,6 +447,7 @@ impl CurveFitApp {
426447
}
427448
Ok(IncrementalSplineFitStep::Finished { result, metrics }) => {
428449
self.fit_in_progress = false;
450+
let point_count = result.residuals.len();
429451
let knot_y = result.knots.iter().map(|knot| knot[1]).collect::<Vec<_>>();
430452
let spline_plot_curve =
431453
Self::plot_points_from_pairs(result.curve.iter().copied());
@@ -434,15 +456,17 @@ impl CurveFitApp {
434456
.append_spline(result.iterations, metrics, &knot_y);
435457
self.upsert_spline_replay_frame(result.iterations, spline_plot_curve);
436458
self.finalize_replay_after_fit_completion();
437-
self.spline_result = Some(result);
438459
self.complete_fit_timer_successfully();
460+
self.store_spline_fit_export_record(&result, point_count);
461+
self.spline_result = Some(result);
439462
self.status = Some(StatusMessage::FitCompleted);
440463
self.active_fit_points = None;
441464
break;
442465
}
443466
Ok(IncrementalSplineFitStep::Cancelled) => {
444467
self.fit_in_progress = false;
445468
self.reset_fit_timer();
469+
self.clear_fit_export_state();
446470
self.finalize_replay_after_fit_stopped();
447471
self.status = Some(StatusMessage::FitStopped);
448472
self.active_fit_points = None;
@@ -451,6 +475,7 @@ impl CurveFitApp {
451475
Err(error) => {
452476
self.fit_in_progress = false;
453477
self.reset_fit_timer();
478+
self.clear_fit_export_state();
454479
self.status = Some(StatusMessage::Error(error.to_string()));
455480
self.active_fit_points = None;
456481
break;
@@ -634,6 +659,7 @@ impl CurveFitApp {
634659
return;
635660
}
636661

662+
self.fit_optimizer_method = self.optimizer_method;
637663
self.last_right_panel_fit_snapshot = Some(self.capture_right_panel_fit_snapshot());
638664
self.auto_refit_pending_rerun = false;
639665

0 commit comments

Comments
 (0)