@@ -94,7 +94,7 @@ class Train:
9494 optimizer : Union [str , Type [torch .optim .Optimizer ]]
9595 quantiles : List [float ] = field (default_factory = list )
9696 optimizer_args : dict = field (default_factory = dict )
97- scheduler : Optional [Type [torch .optim .lr_scheduler .OneCycleLR ]] = None
97+ scheduler : Optional [Type [torch .optim .lr_scheduler ._LRScheduler ]] = None
9898 scheduler_args : dict = field (default_factory = dict )
9999 newer_samples_weight : float = 1.0
100100 newer_samples_start : float = 0.0
@@ -193,16 +193,59 @@ def set_scheduler(self):
193193 Set the scheduler and scheduler args.
194194 The scheduler is not initialized yet as this is done in configure_optimizers in TimeNet.
195195 """
196- self .scheduler = torch .optim .lr_scheduler .OneCycleLR
197- self .scheduler_args .update (
198- {
199- "pct_start" : 0.3 ,
200- "anneal_strategy" : "cos" ,
201- "div_factor" : 10.0 ,
202- "final_div_factor" : 10.0 ,
203- "three_phase" : True ,
204- }
205- )
196+ self .scheduler_args .clear ()
197+ if isinstance (self .scheduler , str ):
198+ if self .scheduler .lower () == "onecyclelr" :
199+ self .scheduler = torch .optim .lr_scheduler .OneCycleLR
200+ self .scheduler_args .update (
201+ {
202+ "pct_start" : 0.3 ,
203+ "anneal_strategy" : "cos" ,
204+ "div_factor" : 10.0 ,
205+ "final_div_factor" : 10.0 ,
206+ "three_phase" : True ,
207+ }
208+ )
209+ elif self .scheduler .lower () == "steplr" :
210+ self .scheduler = torch .optim .lr_scheduler .StepLR
211+ self .scheduler_args .update (
212+ {
213+ "step_size" : 10 ,
214+ "gamma" : 0.1 ,
215+ }
216+ )
217+ elif self .scheduler .lower () == "exponentiallr" :
218+ self .scheduler = torch .optim .lr_scheduler .ExponentialLR
219+ self .scheduler_args .update (
220+ {
221+ "gamma" : 0.95 ,
222+ }
223+ )
224+ elif self .scheduler .lower () == "reducelronplateau" :
225+ self .scheduler = torch .optim .lr_scheduler .ReduceLROnPlateau
226+ self .scheduler_args .update (
227+ {
228+ "mode" : "min" ,
229+ "factor" : 0.1 ,
230+ "patience" : 10 ,
231+ }
232+ )
233+ elif self .scheduler .lower () == "cosineannealinglr" :
234+ self .scheduler = torch .optim .lr_scheduler .CosineAnnealingLR
235+ self .scheduler_args .update (
236+ {
237+ "T_max" : 50 ,
238+ }
239+ )
240+ else :
241+ raise NotImplementedError (f"Scheduler { self .scheduler } is not supported." )
242+ elif self .scheduler is None :
243+ self .scheduler = torch .optim .lr_scheduler .ExponentialLR
244+ self .scheduler_args .update (
245+ {
246+ "gamma" : 0.95 ,
247+ }
248+ )
206249
207250 def set_lr_finder_args (self , dataset_size , num_batches ):
208251 """
0 commit comments