Skip to content

Commit 00f2e25

Browse files
committed
update for onecyclelr
1 parent 9fe3401 commit 00f2e25

2 files changed

Lines changed: 15 additions & 16 deletions

File tree

neuralprophet/forecaster.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -436,6 +436,7 @@ def __init__(
436436
batch_size: Optional[int] = None,
437437
loss_func: Union[str, torch.nn.modules.loss._Loss, Callable] = "SmoothL1Loss",
438438
optimizer: Union[str, Type[torch.optim.Optimizer]] = "AdamW",
439+
scheduler: Optional[str] = "onecyclelr",
439440
newer_samples_weight: float = 2,
440441
newer_samples_start: float = 0.0,
441442
quantiles: List[float] = [],
@@ -451,7 +452,6 @@ def __init__(
451452
accelerator: Optional[str] = None,
452453
trainer_config: dict = {},
453454
prediction_frequency: Optional[dict] = None,
454-
scheduler: Optional[str] = "onecyclelr",
455455
):
456456
self.config = locals()
457457
self.config.pop("self")

neuralprophet/time_net.py

Lines changed: 14 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)