Skip to content

Commit 9a989a7

Browse files
committed
added Newton-CG
1 parent 8253337 commit 9a989a7

8 files changed

Lines changed: 956 additions & 56 deletions

File tree

src/app.rs

Lines changed: 22 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -42,10 +42,11 @@ use self::i18n::{
4242
};
4343
use self::normalization::ParametricNormalization;
4444
use self::optimizer::{
45-
AdamInputState, LbfgsInputState, NelderMeadInputState, OptimizerPreset, OptimizerUiMode,
46-
SgdInputState, SteepestDescentInputState, adam_config_from_preset, infer_adam_preset,
47-
infer_lbfgs_preset, infer_nelder_mead_preset, infer_sgd_preset, infer_steepest_descent_preset,
48-
lbfgs_config_from_preset, nelder_mead_config_from_preset, optimizer_method_label,
45+
AdamInputState, LbfgsInputState, NelderMeadInputState, NewtonCgInputState, OptimizerPreset,
46+
OptimizerUiMode, SgdInputState, SteepestDescentInputState, adam_config_from_preset,
47+
infer_adam_preset, infer_lbfgs_preset, infer_nelder_mead_preset, infer_newton_cg_preset,
48+
infer_sgd_preset, infer_steepest_descent_preset, lbfgs_config_from_preset,
49+
nelder_mead_config_from_preset, newton_cg_config_from_preset, optimizer_method_label,
4950
optimizer_preset_label, sgd_config_from_preset, steepest_descent_config_from_preset,
5051
};
5152
use self::param_init::{
@@ -58,7 +59,7 @@ use self::replay::ReplayState;
5859
#[cfg(test)]
5960
use self::replay::{ReplayFrame, ReplayFramePayload};
6061
use crate::domain::{
61-
AdamConfig, CurveFamily, CurveParams, FitResult, LbfgsConfig, NelderMeadConfig,
62+
AdamConfig, CurveFamily, CurveParams, FitResult, LbfgsConfig, NelderMeadConfig, NewtonCgConfig,
6263
OptimizerConfig, OptimizerMethod, Point, Points, SgdConfig, SteepestDescentConfig,
6364
};
6465
use crate::fit::IterationMetricSnapshot;
@@ -542,6 +543,8 @@ pub struct CurveFitApp {
542543
nelder_mead_preset: OptimizerPreset,
543544
steepest_descent_inputs: SteepestDescentInputState,
544545
steepest_descent_preset: OptimizerPreset,
546+
newton_cg_inputs: NewtonCgInputState,
547+
newton_cg_preset: OptimizerPreset,
545548
sgd_inputs: SgdInputState,
546549
sgd_preset: OptimizerPreset,
547550
adam_inputs: AdamInputState,
@@ -611,6 +614,7 @@ impl CurveFitApp {
611614
OptimizerMethod::Lbfgs => self.lbfgs_preset,
612615
OptimizerMethod::NelderMead => self.nelder_mead_preset,
613616
OptimizerMethod::SteepestDescent => self.steepest_descent_preset,
617+
OptimizerMethod::NewtonCg => self.newton_cg_preset,
614618
OptimizerMethod::Sgd => self.sgd_preset,
615619
OptimizerMethod::Adam => self.adam_preset,
616620
}
@@ -621,6 +625,7 @@ impl CurveFitApp {
621625
OptimizerMethod::Lbfgs => self.lbfgs_preset = preset,
622626
OptimizerMethod::NelderMead => self.nelder_mead_preset = preset,
623627
OptimizerMethod::SteepestDescent => self.steepest_descent_preset = preset,
628+
OptimizerMethod::NewtonCg => self.newton_cg_preset = preset,
624629
OptimizerMethod::Sgd => self.sgd_preset = preset,
625630
OptimizerMethod::Adam => self.adam_preset = preset,
626631
}
@@ -643,6 +648,11 @@ impl CurveFitApp {
643648
);
644649
self.steepest_descent_preset = preset;
645650
}
651+
OptimizerMethod::NewtonCg => {
652+
self.newton_cg_inputs =
653+
NewtonCgInputState::from_config(&newton_cg_config_from_preset(preset));
654+
self.newton_cg_preset = preset;
655+
}
646656
OptimizerMethod::Sgd => {
647657
self.sgd_inputs = SgdInputState::from_config(&sgd_config_from_preset(preset));
648658
self.sgd_preset = preset;
@@ -665,6 +675,10 @@ impl CurveFitApp {
665675
.steepest_descent_inputs
666676
.to_config()
667677
.map(OptimizerConfig::SteepestDescent),
678+
OptimizerMethod::NewtonCg => self
679+
.newton_cg_inputs
680+
.to_config()
681+
.map(OptimizerConfig::NewtonCg),
668682
OptimizerMethod::Sgd => self.sgd_inputs.to_config().map(OptimizerConfig::Sgd),
669683
OptimizerMethod::Adam => self.adam_inputs.to_config().map(OptimizerConfig::Adam),
670684
}
@@ -1084,6 +1098,7 @@ impl Default for CurveFitApp {
10841098
let default_lbfgs = LbfgsConfig::default();
10851099
let default_nelder_mead = NelderMeadConfig::default();
10861100
let default_steepest_descent = SteepestDescentConfig::default();
1101+
let default_newton_cg = NewtonCgConfig::default();
10871102
let default_sgd = SgdConfig::default();
10881103
let default_adam = AdamConfig::default();
10891104

@@ -1104,6 +1119,8 @@ impl Default for CurveFitApp {
11041119
&default_steepest_descent,
11051120
),
11061121
steepest_descent_preset: infer_steepest_descent_preset(&default_steepest_descent),
1122+
newton_cg_inputs: NewtonCgInputState::from_config(&default_newton_cg),
1123+
newton_cg_preset: infer_newton_cg_preset(&default_newton_cg),
11071124
sgd_inputs: SgdInputState::from_config(&default_sgd),
11081125
sgd_preset: infer_sgd_preset(&default_sgd),
11091126
adam_inputs: AdamInputState::from_config(&default_adam),

src/app/optimizer.rs

Lines changed: 114 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
use crate::domain::{
2-
AdamConfig, LbfgsConfig, NelderMeadConfig, OptimizerMethod, SgdConfig, SteepestDescentConfig,
2+
AdamConfig, LbfgsConfig, NelderMeadConfig, NewtonCgConfig, OptimizerMethod, SgdConfig,
3+
SteepestDescentConfig,
34
};
45

56
use super::{C1_MIN, C2_MAX, STEP_MAX_MAX, STEP_MIN_MIN, UiLanguage};
@@ -12,11 +13,13 @@ pub(super) fn optimizer_method_label(
1213
(UiLanguage::English, OptimizerMethod::Lbfgs) => "LBFGS",
1314
(UiLanguage::English, OptimizerMethod::NelderMead) => "Nelder-Mead",
1415
(UiLanguage::English, OptimizerMethod::SteepestDescent) => "Steepest Descent",
16+
(UiLanguage::English, OptimizerMethod::NewtonCg) => "Newton-CG",
1517
(UiLanguage::English, OptimizerMethod::Sgd) => "SGD",
1618
(UiLanguage::English, OptimizerMethod::Adam) => "Adam",
1719
(UiLanguage::Russian, OptimizerMethod::Lbfgs) => "LBFGS",
1820
(UiLanguage::Russian, OptimizerMethod::NelderMead) => "Нелдер-Мид",
1921
(UiLanguage::Russian, OptimizerMethod::SteepestDescent) => "Наискорейший спуск",
22+
(UiLanguage::Russian, OptimizerMethod::NewtonCg) => "Ньютон-CG",
2023
(UiLanguage::Russian, OptimizerMethod::Sgd) => "SGD",
2124
(UiLanguage::Russian, OptimizerMethod::Adam) => "Adam",
2225
}
@@ -58,6 +61,29 @@ impl OptimizerPreset {
5861
pub(super) const ALL: [Self; 3] = [Self::Fast, Self::Balanced, Self::Precise];
5962
}
6063

64+
fn infer_preset_by<T, F>(config: &T, config_from_preset: F) -> OptimizerPreset
65+
where
66+
T: PartialEq,
67+
F: Fn(OptimizerPreset) -> T,
68+
{
69+
OptimizerPreset::ALL
70+
.into_iter()
71+
.find(|preset| config_from_preset(*preset).eq(config))
72+
.unwrap_or(OptimizerPreset::Custom)
73+
}
74+
75+
fn normalize_wolfe_line_search_inputs(
76+
c1: &mut f64,
77+
c2: &mut f64,
78+
step_min: &mut f64,
79+
step_max: &mut f64,
80+
) {
81+
*c1 = (*c1).clamp(C1_MIN, C2_MAX - 1e-4);
82+
*c2 = (*c2).clamp(*c1 + 1e-4, C2_MAX);
83+
*step_min = (*step_min).clamp(STEP_MIN_MIN, STEP_MAX_MAX - 1e-6);
84+
*step_max = (*step_max).clamp(*step_min + 1e-6, STEP_MAX_MAX);
85+
}
86+
6187
pub(super) fn lbfgs_config_from_preset(preset: OptimizerPreset) -> LbfgsConfig {
6288
match preset {
6389
OptimizerPreset::Fast => {
@@ -74,12 +100,7 @@ pub(super) fn lbfgs_config_from_preset(preset: OptimizerPreset) -> LbfgsConfig {
74100
}
75101

76102
pub(super) fn infer_lbfgs_preset(config: &LbfgsConfig) -> OptimizerPreset {
77-
for preset in OptimizerPreset::ALL {
78-
if &lbfgs_config_from_preset(preset) == config {
79-
return preset;
80-
}
81-
}
82-
OptimizerPreset::Custom
103+
infer_preset_by(config, lbfgs_config_from_preset)
83104
}
84105

85106
pub(super) fn nelder_mead_config_from_preset(preset: OptimizerPreset) -> NelderMeadConfig {
@@ -96,12 +117,7 @@ pub(super) fn nelder_mead_config_from_preset(preset: OptimizerPreset) -> NelderM
96117
}
97118

98119
pub(super) fn infer_nelder_mead_preset(config: &NelderMeadConfig) -> OptimizerPreset {
99-
for preset in OptimizerPreset::ALL {
100-
if &nelder_mead_config_from_preset(preset) == config {
101-
return preset;
102-
}
103-
}
104-
OptimizerPreset::Custom
120+
infer_preset_by(config, nelder_mead_config_from_preset)
105121
}
106122

107123
pub(super) fn steepest_descent_config_from_preset(
@@ -120,12 +136,26 @@ pub(super) fn steepest_descent_config_from_preset(
120136
}
121137

122138
pub(super) fn infer_steepest_descent_preset(config: &SteepestDescentConfig) -> OptimizerPreset {
123-
for preset in OptimizerPreset::ALL {
124-
if &steepest_descent_config_from_preset(preset) == config {
125-
return preset;
139+
infer_preset_by(config, steepest_descent_config_from_preset)
140+
}
141+
142+
pub(super) fn newton_cg_config_from_preset(preset: OptimizerPreset) -> NewtonCgConfig {
143+
match preset {
144+
OptimizerPreset::Fast => {
145+
NewtonCgConfig::try_new(80, 1e-8, 1e-8, 1e-4, 0.9, 1e-10, 1.0, 1e-8)
146+
.expect("fast Newton-CG preset must be valid")
147+
}
148+
OptimizerPreset::Balanced => NewtonCgConfig::default(),
149+
OptimizerPreset::Precise => {
150+
NewtonCgConfig::try_new(600, 1e-12, 0.0, 1e-4, 0.95, 1e-12, 10.0, 1e-12)
151+
.expect("precise Newton-CG preset must be valid")
126152
}
153+
OptimizerPreset::Custom => NewtonCgConfig::default(),
127154
}
128-
OptimizerPreset::Custom
155+
}
156+
157+
pub(super) fn infer_newton_cg_preset(config: &NewtonCgConfig) -> OptimizerPreset {
158+
infer_preset_by(config, newton_cg_config_from_preset)
129159
}
130160

131161
pub(super) fn sgd_config_from_preset(preset: OptimizerPreset) -> SgdConfig {
@@ -142,12 +172,7 @@ pub(super) fn sgd_config_from_preset(preset: OptimizerPreset) -> SgdConfig {
142172
}
143173

144174
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
175+
infer_preset_by(config, sgd_config_from_preset)
151176
}
152177

153178
pub(super) fn adam_config_from_preset(preset: OptimizerPreset) -> AdamConfig {
@@ -164,12 +189,7 @@ pub(super) fn adam_config_from_preset(preset: OptimizerPreset) -> AdamConfig {
164189
}
165190

166191
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
192+
infer_preset_by(config, adam_config_from_preset)
173193
}
174194

175195
#[derive(Debug, Clone, PartialEq)]
@@ -201,11 +221,12 @@ impl LbfgsInputState {
201221
}
202222

203223
pub(super) fn normalize_after_ui(&mut self) {
204-
self.c1 = self.c1.clamp(C1_MIN, C2_MAX - 1e-4);
205-
self.c2 = self.c2.clamp(self.c1 + 1e-4, C2_MAX);
206-
207-
self.step_min = self.step_min.clamp(STEP_MIN_MIN, STEP_MAX_MAX - 1e-6);
208-
self.step_max = self.step_max.clamp(self.step_min + 1e-6, STEP_MAX_MAX);
224+
normalize_wolfe_line_search_inputs(
225+
&mut self.c1,
226+
&mut self.c2,
227+
&mut self.step_min,
228+
&mut self.step_max,
229+
);
209230
}
210231

211232
pub(super) fn to_config(&self) -> Result<LbfgsConfig, String> {
@@ -294,10 +315,12 @@ impl SteepestDescentInputState {
294315
}
295316

296317
pub(super) fn normalize_after_ui(&mut self) {
297-
self.c1 = self.c1.clamp(C1_MIN, C2_MAX - 1e-4);
298-
self.c2 = self.c2.clamp(self.c1 + 1e-4, C2_MAX);
299-
self.step_min = self.step_min.clamp(STEP_MIN_MIN, STEP_MAX_MAX - 1e-6);
300-
self.step_max = self.step_max.clamp(self.step_min + 1e-6, STEP_MAX_MAX);
318+
normalize_wolfe_line_search_inputs(
319+
&mut self.c1,
320+
&mut self.c2,
321+
&mut self.step_min,
322+
&mut self.step_max,
323+
);
301324
}
302325

303326
pub(super) fn to_config(&self) -> Result<SteepestDescentConfig, String> {
@@ -313,6 +336,59 @@ impl SteepestDescentInputState {
313336
}
314337
}
315338

339+
#[derive(Debug, Clone, PartialEq)]
340+
pub(super) struct NewtonCgInputState {
341+
pub(super) max_iters: u64,
342+
pub(super) tol: f64,
343+
pub(super) curvature_threshold: f64,
344+
pub(super) c1: f64,
345+
pub(super) c2: f64,
346+
pub(super) step_min: f64,
347+
pub(super) step_max: f64,
348+
pub(super) width_tolerance: f64,
349+
}
350+
351+
impl NewtonCgInputState {
352+
pub(super) fn from_config(config: &NewtonCgConfig) -> Self {
353+
Self {
354+
max_iters: config.max_iters,
355+
tol: config.tol,
356+
curvature_threshold: config.curvature_threshold,
357+
c1: config.c1,
358+
c2: config.c2,
359+
step_min: config.step_min,
360+
step_max: config.step_max,
361+
width_tolerance: config.width_tolerance,
362+
}
363+
}
364+
365+
pub(super) fn normalize_after_ui(&mut self) {
366+
self.tol = self.tol.clamp(1e-14, 1e-2);
367+
self.curvature_threshold = self.curvature_threshold.clamp(0.0, 1e-2);
368+
normalize_wolfe_line_search_inputs(
369+
&mut self.c1,
370+
&mut self.c2,
371+
&mut self.step_min,
372+
&mut self.step_max,
373+
);
374+
self.width_tolerance = self.width_tolerance.clamp(0.0, 1e-3);
375+
}
376+
377+
pub(super) fn to_config(&self) -> Result<NewtonCgConfig, String> {
378+
NewtonCgConfig::try_new(
379+
self.max_iters,
380+
self.tol,
381+
self.curvature_threshold,
382+
self.c1,
383+
self.c2,
384+
self.step_min,
385+
self.step_max,
386+
self.width_tolerance,
387+
)
388+
.map_err(|error| error.to_string())
389+
}
390+
}
391+
316392
#[derive(Debug, Clone, PartialEq)]
317393
pub(super) struct SgdInputState {
318394
pub(super) max_iters: u64,

src/app/tests.rs

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,12 @@ fn optimizer_config_matches_selected_method() {
105105
Ok(OptimizerConfig::SteepestDescent(_))
106106
));
107107

108+
app.optimizer_method = OptimizerMethod::NewtonCg;
109+
assert!(matches!(
110+
app.optimizer_config(),
111+
Ok(OptimizerConfig::NewtonCg(_))
112+
));
113+
108114
app.optimizer_method = OptimizerMethod::Sgd;
109115
assert!(matches!(
110116
app.optimizer_config(),
@@ -186,6 +192,9 @@ fn optimizer_presets_are_stored_per_method() {
186192
app.optimizer_method = OptimizerMethod::SteepestDescent;
187193
app.apply_selected_optimizer_preset(OptimizerPreset::Fast);
188194

195+
app.optimizer_method = OptimizerMethod::NewtonCg;
196+
app.apply_selected_optimizer_preset(OptimizerPreset::Precise);
197+
189198
app.optimizer_method = OptimizerMethod::Sgd;
190199
app.apply_selected_optimizer_preset(OptimizerPreset::Balanced);
191200

@@ -201,6 +210,9 @@ fn optimizer_presets_are_stored_per_method() {
201210
app.optimizer_method = OptimizerMethod::SteepestDescent;
202211
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Fast);
203212

213+
app.optimizer_method = OptimizerMethod::NewtonCg;
214+
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Precise);
215+
204216
app.optimizer_method = OptimizerMethod::Sgd;
205217
assert_eq!(app.selected_optimizer_preset(), OptimizerPreset::Balanced);
206218

@@ -263,6 +275,33 @@ fn sgd_preset_changes_active_config_values() {
263275
assert!(precise_lr < fast_lr);
264276
}
265277

278+
#[test]
279+
fn newton_cg_preset_changes_active_config_values() {
280+
let mut app = CurveFitApp {
281+
optimizer_method: OptimizerMethod::NewtonCg,
282+
..Default::default()
283+
};
284+
app.apply_selected_optimizer_preset(OptimizerPreset::Fast);
285+
let fast_config = app
286+
.optimizer_config()
287+
.expect("optimizer config must be valid");
288+
289+
app.apply_selected_optimizer_preset(OptimizerPreset::Precise);
290+
let precise_config = app
291+
.optimizer_config()
292+
.expect("optimizer config must be valid");
293+
294+
let (fast_max_iters, precise_max_iters, fast_tol, precise_tol) =
295+
match (fast_config, precise_config) {
296+
(OptimizerConfig::NewtonCg(fast), OptimizerConfig::NewtonCg(precise)) => {
297+
(fast.max_iters, precise.max_iters, fast.tol, precise.tol)
298+
}
299+
_ => panic!("Newton-CG must remain active"),
300+
};
301+
assert!(precise_max_iters > fast_max_iters);
302+
assert!(precise_tol < fast_tol);
303+
}
304+
266305
#[test]
267306
fn adam_preset_changes_active_config_values() {
268307
let mut app = CurveFitApp {

0 commit comments

Comments
 (0)