mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
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:
committed by
William Falcon
parent
89ececb32b
commit
a6d64ac013
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user