@@ -871,35 +871,34 @@ def configure_optimizers(self):
871871 optimizer = self ._optimizer (self .parameters (), lr = self .learning_rate , ** self .config_train .optimizer_args )
872872
873873 # Scheduler
874+ self ._scheduler = self .config_train .scheduler
875+
874876 if self .continue_training :
875877 optimizer .load_state_dict (self .config_train .optimizer_state )
876878
877879 # Update initial learning rate to the last learning rate for continued training
878880 last_lr = float (optimizer .param_groups [0 ]["lr" ]) # Ensure it's a float
879881
880- batches_per_epoch = len (self .train_dataloader ())
881- total_batches_processed = self .start_epoch * batches_per_epoch
882-
883882 for param_group in optimizer .param_groups :
884883 param_group ["initial_lr" ] = (last_lr ,)
885884
885+ if self ._scheduler == torch .optim .lr_scheduler .OneCycleLR :
886+ log .warning ("OneCycleLR scheduler is not supported for continued training. Switching to ExponentialLR" )
887+ self ._scheduler = torch .optim .lr_scheduler .ExponentialLR
888+ self .config_train .scheduler_args = {"gamma" : 0.95 }
889+
890+ if self ._scheduler == torch .optim .lr_scheduler .OneCycleLR :
886891 lr_scheduler = self ._scheduler (
887892 optimizer ,
893+ max_lr = self .learning_rate ,
894+ total_steps = self .trainer .estimated_stepping_batches ,
888895 ** self .config_train .scheduler_args ,
889896 )
890897 else :
891- if self ._scheduler == torch .optim .lr_scheduler .OneCycleLR :
892- lr_scheduler = self ._scheduler (
893- optimizer ,
894- max_lr = self .learning_rate ,
895- total_steps = self .trainer .estimated_stepping_batches ,
896- ** self .config_train .scheduler_args ,
897- )
898- else :
899- lr_scheduler = self ._scheduler (
900- optimizer ,
901- ** self .config_train .scheduler_args ,
902- )
898+ lr_scheduler = self ._scheduler (
899+ optimizer ,
900+ ** self .config_train .scheduler_args ,
901+ )
903902
904903 return {"optimizer" : optimizer , "lr_scheduler" : lr_scheduler }
905904
0 commit comments