Support torch.optim.lr_scheduler.ReduceLROnPlateau (#320)

* feat: add reducelronplateau callback

* feat: use reducelronplateau callback in trainer

* feat: only on unsupported lr schedulers

* feat: last but not the least merge of master

* feat: merge master

* feat: support only on scheduler in reduceLrOnPlateauScheduler

* refactor: code style

* Update pt_callbacks.py

* Update trainer.py

* Update train_loop_mixin.py

* Update trainer.py

* Update train_loop_mixin.py
This commit is contained in:
Mary Trofimova
2019-12-03 07:59:41 -05:00
committed by William Falcon
parent 89ececb32b
commit a6d64ac013
2 changed files with 21 additions and 2 deletions
+11 -2
View File
@@ -150,7 +150,8 @@ When this flag is enabled each batch is split into sequences of size truncated_b
"""
import numpy as np
import tqdm
from pytorch_lightning.utilities.debugging import MisconfigurationException
try:
from apex import amp
@@ -213,7 +214,15 @@ class TrainerTrainLoopMixin(object):
# update LR schedulers
if self.lr_schedulers is not None:
for lr_scheduler in self.lr_schedulers:
lr_scheduler.step(self.current_epoch)
lr_scheduler.step(epoch=self.current_epoch)
if self.reduce_lr_on_plateau_scheduler is not None:
val_loss = self.callback_metrics.get('val_loss')
if val_loss is None:
avail_metrics = ','.join(list(self.callback_metrics.keys()))
m = f'ReduceLROnPlateau conditioned on metric val_loss ' \
f'which is not available. Available metrics are: {avail_metrics}'
raise MisconfigurationException(m)
self.reduce_lr_on_plateau_scheduler.step(val_loss, epoch=self.current_epoch)
# early stopping
met_min_epochs = epoch_nb > self.min_nb_epochs
+10
View File
@@ -188,6 +188,8 @@ class Trainer(TrainerIOMixin,
self.early_stop_callback = None
self.configure_early_stopping(early_stop_callback, logger)
self.reduce_lr_on_plateau_scheduler = None
# configure checkpoint callback
self.checkpoint_callback = checkpoint_callback
self.weights_save_path = weights_save_path
@@ -378,12 +380,20 @@ class Trainer(TrainerIOMixin,
# two lists
elif len(optimizers) == 2 and isinstance(optimizers[0], list):
optimizers, lr_schedulers = optimizers
lr_schedulers, self.reduce_lr_on_plateau_scheduler = self.configure_schedulers(lr_schedulers)
return optimizers, lr_schedulers
# single list or tuple
elif isinstance(optimizers, list) or isinstance(optimizers, tuple):
return optimizers, []
def configure_schedulers(self, schedulers):
for i, scheduler in enumerate(schedulers):
if isinstance(scheduler, torch.optim.lr_scheduler.ReduceLROnPlateau):
reduce_lr_on_plateau_scheduler = schedulers.pop(i)
return schedulers, reduce_lr_on_plateau_scheduler
return schedulers, None
def run_pretrain_routine(self, model):
"""
Sanity check a few things before starting actual training