ReduceLROnPlateau bug fix (#1126)

* bug fix and test

* update CHANGELOG.md

Co-authored-by: Nicki Skafte <nugginea@gmail.com>
This commit is contained in:
Nicki Skafte
2020-03-16 14:35:10 -04:00
committed by GitHub
co-authored by Nicki Skafte
parent 774d9be357
commit 384e124490
5 changed files with 51 additions and 5 deletions
+1 -1
View File
@@ -28,7 +28,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/).
### Fixed
-
- Fixed bug related to type cheking of `ReduceLROnPlateau` lr schedulers([#1114](https://github.com/PyTorchLightning/pytorch-lightning/issues/1114))
## [0.7.1] - 2020-03-07
+2 -2
View File
@@ -707,8 +707,8 @@ class Trainer(
if 'scheduler' not in scheduler:
raise ValueError(f'Lr scheduler should have key `scheduler`',
' with item being a lr scheduler')
scheduler['reduce_on_plateau'] = \
isinstance(scheduler, optim.lr_scheduler.ReduceLROnPlateau)
scheduler['reduce_on_plateau'] = isinstance(
scheduler['scheduler'], optim.lr_scheduler.ReduceLROnPlateau)
lr_schedulers.append({**default_config, **scheduler})
+2 -1
View File
@@ -24,7 +24,8 @@ from .mixins import (
LightInfTestDataloader,
LightTestOptimizerWithSchedulingMixin,
LightTestMultipleOptimizersWithSchedulingMixin,
LightTestOptimizersWithMixedSchedulingMixin
LightTestOptimizersWithMixedSchedulingMixin,
LightTestReduceLROnPlateauMixin
)
+10
View File
@@ -678,6 +678,16 @@ class LightTestOptimizersWithMixedSchedulingMixin:
[{'scheduler': lr_scheduler1, 'interval': 'step'}, lr_scheduler2]
class LightTestReduceLROnPlateauMixin:
def configure_optimizers(self):
if self.hparams.optimizer_name == 'lbfgs':
optimizer = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
lr_scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer)
return [optimizer], [lr_scheduler]
def _get_output_metric(output, name):
if isinstance(output, dict):
val = output[name]
+36 -1
View File
@@ -10,9 +10,12 @@ from pytorch_lightning import Trainer
from tests.models import (
TestModelBase,
LightTrainDataloader,
LightValidationStepMixin,
LightValidationMixin,
LightTestOptimizerWithSchedulingMixin,
LightTestMultipleOptimizersWithSchedulingMixin,
LightTestOptimizersWithMixedSchedulingMixin
LightTestOptimizersWithMixedSchedulingMixin,
LightTestReduceLROnPlateauMixin
)
@@ -144,3 +147,35 @@ def test_multi_optimizer_with_scheduling_stepping(tmpdir):
# Called every 3 steps, meaning for 1 epoch of 11 batches, it is called 3 times
assert init_lr * 0.1 == adjusted_lr2, \
'lr for optimizer 2 not adjusted correctly'
def test_reduce_lr_on_plateau_scheduling(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightTestReduceLROnPlateauMixin,
LightTrainDataloader,
LightValidationMixin,
LightValidationStepMixin,
TestModelBase):
pass
hparams = tutils.get_hparams()
model = CurrentTestModel(hparams)
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=1,
val_percent_check=0.1,
train_percent_check=0.2
)
# fit model
trainer = Trainer(**trainer_options)
results = trainer.fit(model)
assert trainer.lr_schedulers[0] == \
dict(scheduler=trainer.lr_schedulers[0]['scheduler'], monitor='val_loss',
interval='epoch', frequency=1, reduce_on_plateau=True), \
'lr schduler was not correctly converted to dict'