diff --git a/CHANGELOG.md b/CHANGELOG.md index 1c35a2ee..6e12875f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 3f6d9709..826bd4ed 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -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}) diff --git a/tests/models/__init__.py b/tests/models/__init__.py index 4992e70a..67206a63 100644 --- a/tests/models/__init__.py +++ b/tests/models/__init__.py @@ -24,7 +24,8 @@ from .mixins import ( LightInfTestDataloader, LightTestOptimizerWithSchedulingMixin, LightTestMultipleOptimizersWithSchedulingMixin, - LightTestOptimizersWithMixedSchedulingMixin + LightTestOptimizersWithMixedSchedulingMixin, + LightTestReduceLROnPlateauMixin ) diff --git a/tests/models/mixins.py b/tests/models/mixins.py index fd3f0dde..0be69172 100644 --- a/tests/models/mixins.py +++ b/tests/models/mixins.py @@ -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] diff --git a/tests/trainer/test_optimizers.py b/tests/trainer/test_optimizers.py index bc5dde5f..3ea0e3ff 100644 --- a/tests/trainer/test_optimizers.py +++ b/tests/trainer/test_optimizers.py @@ -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'