11use crate :: domain:: {
2- AdamConfig , LbfgsConfig , NelderMeadConfig , OptimizerMethod , SgdConfig , SteepestDescentConfig ,
2+ AdamConfig , LbfgsConfig , NelderMeadConfig , NewtonCgConfig , OptimizerMethod , SgdConfig ,
3+ SteepestDescentConfig ,
34} ;
45
56use 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+
6187pub ( 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
76102pub ( 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
85106pub ( 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
98119pub ( 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
107123pub ( super ) fn steepest_descent_config_from_preset (
@@ -120,12 +136,26 @@ pub(super) fn steepest_descent_config_from_preset(
120136}
121137
122138pub ( 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
131161pub ( 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
144174pub ( 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
153178pub ( 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
166191pub ( 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 ) ]
317393pub ( super ) struct SgdInputState {
318394 pub ( super ) max_iters : u64 ,
0 commit comments