Skip to content

Commit 9fe3401

Browse files
committed
enable setting the scheduler
1 parent b14d20b commit 9fe3401

3 files changed

Lines changed: 82 additions & 18 deletions

File tree

neuralprophet/configure.py

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

neuralprophet/forecaster.py

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -451,6 +451,7 @@ def __init__(
451451
accelerator: Optional[str] = None,
452452
trainer_config: dict = {},
453453
prediction_frequency: Optional[dict] = None,
454+
scheduler: Optional[str] = "onecyclelr",
454455
):
455456
self.config = locals()
456457
self.config.pop("self")
@@ -509,6 +510,7 @@ def __init__(
509510
self.config_train = configure.Train(
510511
quantiles=quantiles,
511512
learning_rate=learning_rate,
513+
scheduler=scheduler,
512514
epochs=epochs,
513515
batch_size=batch_size,
514516
loss_func=loss_func,
@@ -921,6 +923,7 @@ def fit(
921923
continue_training: bool = False,
922924
num_workers: int = 0,
923925
deterministic: bool = False,
926+
scheduler: Optional[str] = None,
924927
):
925928
"""Train, and potentially evaluate model.
926929
@@ -986,6 +989,18 @@ def fit(
986989
if continue_training and epochs is None:
987990
raise ValueError("Continued training requires setting the number of epochs to train for.")
988991

992+
if continue_training:
993+
if scheduler is not None:
994+
self.config_train.scheduler = scheduler
995+
else:
996+
self.config_train.scheduler = None
997+
self.config_train.set_scheduler()
998+
999+
if scheduler is not None and not continue_training:
1000+
log.warning(
1001+
"Scheduler can only be set in fit when continuing training. Please set the scheduler when initializing the model."
1002+
)
1003+
9891004
# Configuration
9901005
if epochs is not None:
9911006
self.config_train.epochs = epochs
@@ -2681,7 +2696,6 @@ def _init_train_loader(self, df, num_workers=0):
26812696
config_seasonality=self.config_seasonality,
26822697
)
26832698

2684-
print("Changepoints:", self.config_trend.changepoints)
26852699
df = _normalize(df=df, config_normalization=self.config_normalization)
26862700
if not self.fitted:
26872701
if self.config_trend.changepoints is not None:

neuralprophet/time_net.py

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -883,16 +883,23 @@ def configure_optimizers(self):
883883
for param_group in optimizer.param_groups:
884884
param_group["initial_lr"] = (last_lr,)
885885

886-
lr_scheduler = lr_scheduler = torch.optim.lr_scheduler.ExponentialLR(
887-
optimizer, gamma=0.95, last_epoch=total_batches_processed - 1
888-
)
889-
else:
890886
lr_scheduler = self._scheduler(
891887
optimizer,
892-
max_lr=self.learning_rate,
893-
total_steps=self.trainer.estimated_stepping_batches,
894888
**self.config_train.scheduler_args,
895889
)
890+
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+
)
896903

897904
return {"optimizer": optimizer, "lr_scheduler": lr_scheduler}
898905

0 commit comments

Comments
 (0)