Skip to content

Commit be1a355

Browse files
committed
added SGD and Adam
1 parent 37c5769 commit be1a355

10 files changed

Lines changed: 636 additions & 16 deletions

File tree

Cargo.lock

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

Cargo.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ egui_extras = { version = "0.34.0", default-features = false, features = [
2424
"svg_text",
2525
] }
2626
egui_plot = "0.35.0"
27+
stochastic_optimizers = "0.3.0"
2728

2829
[target.'cfg(not(target_arch = "wasm32"))'.dependencies]
2930
eframe = { version = "0.34.0", default-features = false, features = [

src/app.rs

Lines changed: 31 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -40,10 +40,11 @@ use self::i18n::{
4040
stop_icon_image, tool_icon_image, tool_label, tr, undo_icon_image, view_icon_image,
4141
};
4242
use self::optimizer::{
43-
LbfgsInputState, NelderMeadInputState, OptimizerPreset, OptimizerUiMode,
44-
SteepestDescentInputState, infer_lbfgs_preset, infer_nelder_mead_preset,
45-
infer_steepest_descent_preset, lbfgs_config_from_preset, nelder_mead_config_from_preset,
46-
optimizer_method_label, optimizer_preset_label, steepest_descent_config_from_preset,
43+
AdamInputState, LbfgsInputState, NelderMeadInputState, OptimizerPreset, OptimizerUiMode,
44+
SgdInputState, SteepestDescentInputState, adam_config_from_preset, infer_adam_preset,
45+
infer_lbfgs_preset, infer_nelder_mead_preset, infer_sgd_preset, infer_steepest_descent_preset,
46+
lbfgs_config_from_preset, nelder_mead_config_from_preset, optimizer_method_label,
47+
optimizer_preset_label, sgd_config_from_preset, steepest_descent_config_from_preset,
4748
};
4849
use self::param_init::{
4950
data_based_params_for_family, is_advanced_param_init_supported, polynomial_family,
@@ -55,8 +56,8 @@ use self::replay::ReplayState;
5556
#[cfg(test)]
5657
use self::replay::{ReplayFrame, ReplayFramePayload};
5758
use crate::domain::{
58-
CurveFamily, CurveParams, FitResult, LbfgsConfig, NelderMeadConfig, OptimizerConfig,
59-
OptimizerMethod, Point, Points, SteepestDescentConfig,
59+
AdamConfig, CurveFamily, CurveParams, FitResult, LbfgsConfig, NelderMeadConfig,
60+
OptimizerConfig, OptimizerMethod, Point, Points, SgdConfig, SteepestDescentConfig,
6061
};
6162
use crate::fit::IterationMetricSnapshot;
6263
use crate::fit::OptimizationLossMetric;
@@ -525,6 +526,10 @@ pub struct CurveFitApp {
525526
nelder_mead_preset: OptimizerPreset,
526527
steepest_descent_inputs: SteepestDescentInputState,
527528
steepest_descent_preset: OptimizerPreset,
529+
sgd_inputs: SgdInputState,
530+
sgd_preset: OptimizerPreset,
531+
adam_inputs: AdamInputState,
532+
adam_preset: OptimizerPreset,
528533
ui_language: UiLanguage,
529534
plot_tool: PlotTool,
530535
spray_points_per_second: usize,
@@ -590,6 +595,8 @@ impl CurveFitApp {
590595
OptimizerMethod::Lbfgs => self.lbfgs_preset,
591596
OptimizerMethod::NelderMead => self.nelder_mead_preset,
592597
OptimizerMethod::SteepestDescent => self.steepest_descent_preset,
598+
OptimizerMethod::Sgd => self.sgd_preset,
599+
OptimizerMethod::Adam => self.adam_preset,
593600
}
594601
}
595602

@@ -598,6 +605,8 @@ impl CurveFitApp {
598605
OptimizerMethod::Lbfgs => self.lbfgs_preset = preset,
599606
OptimizerMethod::NelderMead => self.nelder_mead_preset = preset,
600607
OptimizerMethod::SteepestDescent => self.steepest_descent_preset = preset,
608+
OptimizerMethod::Sgd => self.sgd_preset = preset,
609+
OptimizerMethod::Adam => self.adam_preset = preset,
601610
}
602611
}
603612

@@ -618,6 +627,14 @@ impl CurveFitApp {
618627
);
619628
self.steepest_descent_preset = preset;
620629
}
630+
OptimizerMethod::Sgd => {
631+
self.sgd_inputs = SgdInputState::from_config(&sgd_config_from_preset(preset));
632+
self.sgd_preset = preset;
633+
}
634+
OptimizerMethod::Adam => {
635+
self.adam_inputs = AdamInputState::from_config(&adam_config_from_preset(preset));
636+
self.adam_preset = preset;
637+
}
621638
}
622639
}
623640

@@ -632,6 +649,8 @@ impl CurveFitApp {
632649
.steepest_descent_inputs
633650
.to_config()
634651
.map(OptimizerConfig::SteepestDescent),
652+
OptimizerMethod::Sgd => self.sgd_inputs.to_config().map(OptimizerConfig::Sgd),
653+
OptimizerMethod::Adam => self.adam_inputs.to_config().map(OptimizerConfig::Adam),
635654
}
636655
}
637656

@@ -1055,6 +1074,8 @@ impl Default for CurveFitApp {
10551074
let default_lbfgs = LbfgsConfig::default();
10561075
let default_nelder_mead = NelderMeadConfig::default();
10571076
let default_steepest_descent = SteepestDescentConfig::default();
1077+
let default_sgd = SgdConfig::default();
1078+
let default_adam = AdamConfig::default();
10581079

10591080
Self {
10601081
points: PointsEditorState::default(),
@@ -1072,6 +1093,10 @@ impl Default for CurveFitApp {
10721093
&default_steepest_descent,
10731094
),
10741095
steepest_descent_preset: infer_steepest_descent_preset(&default_steepest_descent),
1096+
sgd_inputs: SgdInputState::from_config(&default_sgd),
1097+
sgd_preset: infer_sgd_preset(&default_sgd),
1098+
adam_inputs: AdamInputState::from_config(&default_adam),
1099+
adam_preset: infer_adam_preset(&default_adam),
10751100
ui_language: UiLanguage::English,
10761101
plot_tool: PlotTool::SinglePoint,
10771102
spray_points_per_second: 140,

src/app/optimizer.rs

Lines changed: 97 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,6 @@
1-
use crate::domain::{LbfgsConfig, NelderMeadConfig, OptimizerMethod, SteepestDescentConfig};
1+
use crate::domain::{
2+
AdamConfig, LbfgsConfig, NelderMeadConfig, OptimizerMethod, SgdConfig, SteepestDescentConfig,
3+
};
24

35
use super::{C1_MIN, C2_MAX, STEP_MAX_MAX, STEP_MIN_MIN, UiLanguage};
46

@@ -10,9 +12,13 @@ pub(super) fn optimizer_method_label(
1012
(UiLanguage::English, OptimizerMethod::Lbfgs) => "LBFGS",
1113
(UiLanguage::English, OptimizerMethod::NelderMead) => "Nelder-Mead",
1214
(UiLanguage::English, OptimizerMethod::SteepestDescent) => "Steepest Descent",
15+
(UiLanguage::English, OptimizerMethod::Sgd) => "SGD",
16+
(UiLanguage::English, OptimizerMethod::Adam) => "Adam",
1317
(UiLanguage::Russian, OptimizerMethod::Lbfgs) => "LBFGS",
1418
(UiLanguage::Russian, OptimizerMethod::NelderMead) => "Нелдер-Мид",
1519
(UiLanguage::Russian, OptimizerMethod::SteepestDescent) => "Наискорейший спуск",
20+
(UiLanguage::Russian, OptimizerMethod::Sgd) => "SGD",
21+
(UiLanguage::Russian, OptimizerMethod::Adam) => "Adam",
1622
}
1723
}
1824

@@ -122,6 +128,50 @@ pub(super) fn infer_steepest_descent_preset(config: &SteepestDescentConfig) -> O
122128
OptimizerPreset::Custom
123129
}
124130

131+
pub(super) fn sgd_config_from_preset(preset: OptimizerPreset) -> SgdConfig {
132+
match preset {
133+
OptimizerPreset::Fast => {
134+
SgdConfig::try_new(250, 3e-2).expect("fast SGD preset must be valid")
135+
}
136+
OptimizerPreset::Balanced => SgdConfig::default(),
137+
OptimizerPreset::Precise => {
138+
SgdConfig::try_new(4_000, 3e-3).expect("precise SGD preset must be valid")
139+
}
140+
OptimizerPreset::Custom => SgdConfig::default(),
141+
}
142+
}
143+
144+
pub(super) fn infer_sgd_preset(config: &SgdConfig) -> OptimizerPreset {
145+
for preset in OptimizerPreset::ALL {
146+
if &sgd_config_from_preset(preset) == config {
147+
return preset;
148+
}
149+
}
150+
OptimizerPreset::Custom
151+
}
152+
153+
pub(super) fn adam_config_from_preset(preset: OptimizerPreset) -> AdamConfig {
154+
match preset {
155+
OptimizerPreset::Fast => {
156+
AdamConfig::try_new(200, 2e-2).expect("fast Adam preset must be valid")
157+
}
158+
OptimizerPreset::Balanced => AdamConfig::default(),
159+
OptimizerPreset::Precise => {
160+
AdamConfig::try_new(3_000, 1e-3).expect("precise Adam preset must be valid")
161+
}
162+
OptimizerPreset::Custom => AdamConfig::default(),
163+
}
164+
}
165+
166+
pub(super) fn infer_adam_preset(config: &AdamConfig) -> OptimizerPreset {
167+
for preset in OptimizerPreset::ALL {
168+
if &adam_config_from_preset(preset) == config {
169+
return preset;
170+
}
171+
}
172+
OptimizerPreset::Custom
173+
}
174+
125175
#[derive(Debug, Clone, PartialEq)]
126176
pub(super) struct LbfgsInputState {
127177
pub(super) history_size: usize,
@@ -262,3 +312,49 @@ impl SteepestDescentInputState {
262312
.map_err(|error| error.to_string())
263313
}
264314
}
315+
316+
#[derive(Debug, Clone, PartialEq)]
317+
pub(super) struct SgdInputState {
318+
pub(super) max_iters: u64,
319+
pub(super) learning_rate: f64,
320+
}
321+
322+
impl SgdInputState {
323+
pub(super) fn from_config(config: &SgdConfig) -> Self {
324+
Self {
325+
max_iters: config.max_iters,
326+
learning_rate: config.learning_rate,
327+
}
328+
}
329+
330+
pub(super) fn normalize_after_ui(&mut self) {
331+
self.learning_rate = self.learning_rate.clamp(1e-6, 1.0);
332+
}
333+
334+
pub(super) fn to_config(&self) -> Result<SgdConfig, String> {
335+
SgdConfig::try_new(self.max_iters, self.learning_rate).map_err(|error| error.to_string())
336+
}
337+
}
338+
339+
#[derive(Debug, Clone, PartialEq)]
340+
pub(super) struct AdamInputState {
341+
pub(super) max_iters: u64,
342+
pub(super) learning_rate: f64,
343+
}
344+
345+
impl AdamInputState {
346+
pub(super) fn from_config(config: &AdamConfig) -> Self {
347+
Self {
348+
max_iters: config.max_iters,
349+
learning_rate: config.learning_rate,
350+
}
351+
}
352+
353+
pub(super) fn normalize_after_ui(&mut self) {
354+
self.learning_rate = self.learning_rate.clamp(1e-6, 1.0);
355+
}
356+
357+
pub(super) fn to_config(&self) -> Result<AdamConfig, String> {
358+
AdamConfig::try_new(self.max_iters, self.learning_rate).map_err(|error| error.to_string())
359+
}
360+
}

src/app/tests.rs

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -104,6 +104,18 @@ fn optimizer_config_matches_selected_method() {
104104
app.optimizer_config(),
105105
Ok(OptimizerConfig::SteepestDescent(_))
106106
));
107+
108+
app.optimizer_method = OptimizerMethod::Sgd;
109+
assert!(matches!(
110+
app.optimizer_config(),
111+
Ok(OptimizerConfig::Sgd(_))
112+
));
113+
114+
app.optimizer_method = OptimizerMethod::Adam;
115+
assert!(matches!(
116+
app.optimizer_config(),
117+
Ok(OptimizerConfig::Adam(_))
118+
));
107119
}
108120

109121
#[test]
@@ -174,6 +186,12 @@ fn optimizer_presets_are_stored_per_method() {
174186
app.optimizer_method = OptimizerMethod::SteepestDescent;
175187
app.apply_selected_optimizer_preset(OptimizerPreset::Fast);
176188

189+
app.optimizer_method = OptimizerMethod::Sgd;
190+
app.apply_selected_optimizer_preset(OptimizerPreset::Balanced);
191+
192+
app.optimizer_method = OptimizerMethod::Adam;
193+
app.apply_selected_optimizer_preset(OptimizerPreset::Precise);
194+
177195
app.optimizer_method = OptimizerMethod::Lbfgs;
178196
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Fast);
179197

@@ -182,6 +200,12 @@ fn optimizer_presets_are_stored_per_method() {
182200

183201
app.optimizer_method = OptimizerMethod::SteepestDescent;
184202
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Fast);
203+
204+
app.optimizer_method = OptimizerMethod::Sgd;
205+
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Balanced);
206+
207+
app.optimizer_method = OptimizerMethod::Adam;
208+
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Precise);
185209
}
186210

187211
#[test]
@@ -209,6 +233,66 @@ fn optimizer_preset_changes_active_config_values() {
209233
assert!(precise_max_iters > fast_max_iters);
210234
}
211235

236+
#[test]
237+
fn sgd_preset_changes_active_config_values() {
238+
let mut app = CurveFitApp {
239+
optimizer_method: OptimizerMethod::Sgd,
240+
..Default::default()
241+
};
242+
app.apply_selected_optimizer_preset(OptimizerPreset::Fast);
243+
let fast_config = app
244+
.optimizer_config()
245+
.expect("optimizer config must be valid");
246+
247+
app.apply_selected_optimizer_preset(OptimizerPreset::Precise);
248+
let precise_config = app
249+
.optimizer_config()
250+
.expect("optimizer config must be valid");
251+
252+
let (fast_max_iters, precise_max_iters, fast_lr, precise_lr) =
253+
match (fast_config, precise_config) {
254+
(OptimizerConfig::Sgd(fast), OptimizerConfig::Sgd(precise)) => (
255+
fast.max_iters,
256+
precise.max_iters,
257+
fast.learning_rate,
258+
precise.learning_rate,
259+
),
260+
_ => panic!("SGD must remain active"),
261+
};
262+
assert!(precise_max_iters > fast_max_iters);
263+
assert!(precise_lr < fast_lr);
264+
}
265+
266+
#[test]
267+
fn adam_preset_changes_active_config_values() {
268+
let mut app = CurveFitApp {
269+
optimizer_method: OptimizerMethod::Adam,
270+
..Default::default()
271+
};
272+
app.apply_selected_optimizer_preset(OptimizerPreset::Fast);
273+
let fast_config = app
274+
.optimizer_config()
275+
.expect("optimizer config must be valid");
276+
277+
app.apply_selected_optimizer_preset(OptimizerPreset::Precise);
278+
let precise_config = app
279+
.optimizer_config()
280+
.expect("optimizer config must be valid");
281+
282+
let (fast_max_iters, precise_max_iters, fast_lr, precise_lr) =
283+
match (fast_config, precise_config) {
284+
(OptimizerConfig::Adam(fast), OptimizerConfig::Adam(precise)) => (
285+
fast.max_iters,
286+
precise.max_iters,
287+
fast.learning_rate,
288+
precise.learning_rate,
289+
),
290+
_ => panic!("Adam must remain active"),
291+
};
292+
assert!(precise_max_iters > fast_max_iters);
293+
assert!(precise_lr < fast_lr);
294+
}
295+
212296
#[test]
213297
fn diagnostics_initialize_stores_iteration_zero_state() {
214298
let points = line_points();

0 commit comments

Comments
 (0)