diff --git a/README.md b/README.md index f2c51e0f..f0e2cfea 100644 --- a/README.md +++ b/README.md @@ -109,6 +109,7 @@ class CoolSystem(pl.LightningModule): def configure_optimizers(self): # REQUIRED # can return multiple optimizers and learning_rate schedulers + # (LBFGS it is automatically supported, no need for closure function) return torch.optim.Adam(self.parameters(), lr=0.02) @pl.data_loader diff --git a/docs/LightningModule/RequiredTrainerInterface.md b/docs/LightningModule/RequiredTrainerInterface.md index 345c5ac6..3cf630db 100644 --- a/docs/LightningModule/RequiredTrainerInterface.md +++ b/docs/LightningModule/RequiredTrainerInterface.md @@ -200,8 +200,7 @@ Set up as many optimizers and (optionally) learning rate schedulers as you need. Lightning will call .backward() and .step() on each one in every epoch. If you use 16 bit precision it will also handle that. **Note:** If you use multiple optimizers, training_step will have an additional ```optimizer_idx``` parameter. - - +**Note 2:** If you use LBFGS lightning handles the closure function automatically for you. ##### Return Return any of these 3 options: diff --git a/pytorch_lightning/testing/lm_test_module_base.py b/pytorch_lightning/testing/lm_test_module_base.py index 36c6a8ab..bd7d2f87 100644 --- a/pytorch_lightning/testing/lm_test_module_base.py +++ b/pytorch_lightning/testing/lm_test_module_base.py @@ -115,11 +115,11 @@ class LightningTestModelBase(LightningModule): # --------------------- def configure_optimizers(self): """ - return whatever optimizers we want here + return whatever optimizers we want here. :return: list of optimizers """ # try no scheduler for this model (testing purposes) - if self.hparams.optimizer == 'lbfgs': + 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) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index a6765318..c004a4cf 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -1264,7 +1264,8 @@ class Trainer(TrainerIO): # wrap the forward step in a closure so second order methods work def optimizer_closure(): # forward pass - closure_loss, model_specific_tqdm_metrics = self.__training_forward(batch, batch_nb, opt_idx) + output = self.__training_forward(batch, batch_nb, opt_idx) + closure_loss, model_specific_tqdm_metrics = output # track metrics self.__add_tqdm_metrics(model_specific_tqdm_metrics) @@ -1312,7 +1313,8 @@ class Trainer(TrainerIO): # calls .step(), .zero_grad() # override function to modify this behavior model = self.__get_model() - model.optimizer_step(self.current_epoch, batch_nb, optimizer, opt_idx, optimizer_closure) + model.optimizer_step(self.current_epoch, batch_nb, + optimizer, opt_idx, optimizer_closure) # calculate running loss for display self.running_loss.append(self.batch_loss_value) diff --git a/tests/test_models.py b/tests/test_models.py index b2aa5391..c70663d0 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -42,7 +42,6 @@ def test_default_logger_callbacks_cpu_model(): Test each of the trainer options :return: """ - reset_seed() trainer_options = dict( @@ -68,14 +67,16 @@ def test_lbfgs_cpu_model(): Test each of the trainer options :return: """ + reset_seed() + trainer_options = dict( max_nb_epochs=1, gradient_clip_val=1.0, overfit_pct=0.20, print_nan_grads=True, show_progress_bar=False, - train_percent_check=0.01, - val_percent_check=0.01 + train_percent_check=0.1, + val_percent_check=0.1 ) model, hparams = get_model(use_test_model=True, lbfgs=True) @@ -85,6 +86,7 @@ def test_lbfgs_cpu_model(): model.freeze() model.unfreeze() + def test_multi_gpu_model_ddp2(): """ Make sure DDP2 works @@ -1537,7 +1539,7 @@ def get_model(use_test_model=False, lbfgs=False): # set up model with these hyperparams hparams = get_hparams() if lbfgs: - hparams.optimizer = 'lbfgs' + setattr(hparams, 'optimizer_name', 'lbfgs') if use_test_model: model = LightningTestModel(hparams)