diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index 47cad131..c4e836bc 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -92,16 +92,20 @@ class LightningModule(GradInformation, ModelIO, ModelHooks): """ raise NotImplementedError - def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i): + def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i, second_order_closure=None): """ Do something instead of the standard optimizer behavior :param epoch_nb: :param batch_nb: :param optimizer: :param optimizer_i: + :param second_order_closure: closure for second order methods :return: """ - optimizer.step() + if isinstance(optimizer, torch.optim.LBFGS): + optimizer.step(second_order_closure) + else: + optimizer.step() # clear gradients optimizer.zero_grad() diff --git a/pytorch_lightning/testing/lm_test_module_base.py b/pytorch_lightning/testing/lm_test_module_base.py index ab768241..36c6a8ab 100644 --- a/pytorch_lightning/testing/lm_test_module_base.py +++ b/pytorch_lightning/testing/lm_test_module_base.py @@ -119,7 +119,10 @@ class LightningTestModelBase(LightningModule): :return: list of optimizers """ # try no scheduler for this model (testing purposes) - optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate) + if self.hparams.optimizer == 'lbfgs': + optimizer = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate) + else: + optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate) # test returning only 1 list instead of 2 return optimizer diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 6eae341b..b54329f7 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -275,6 +275,8 @@ class Trainer(TrainerIO): raise ModuleNotFoundError(msg) def __configure_accumulated_gradients(self, accumulate_grad_batches): + self.accumulate_grad_batches = None + if isinstance(accumulate_grad_batches, dict): self.accumulation_scheduler = GradientAccumulationScheduler(accumulate_grad_batches) elif isinstance(accumulate_grad_batches, int): @@ -1267,27 +1269,34 @@ class Trainer(TrainerIO): # call training_step once per optimizer for opt_idx, optimizer in enumerate(self.optimizers): - # forward pass - loss, model_specific_tqdm_metrics = self.__training_forward(batch, batch_nb, opt_idx) + # 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) - # track metrics - self.__add_tqdm_metrics(model_specific_tqdm_metrics) + # track metrics + self.__add_tqdm_metrics(model_specific_tqdm_metrics) - # accumulate loss - # (if accumulate_grad_batches = 1 no effect) - loss = loss / self.accumulate_grad_batches + # accumulate loss + # (if accumulate_grad_batches = 1 no effect) + closure_loss = closure_loss / self.accumulate_grad_batches - # backward pass - if self.use_amp: - with amp.scale_loss(loss, optimizer) as scaled_loss: - scaled_loss.backward() - else: - loss.backward() + # backward pass + if self.use_amp: + with amp.scale_loss(closure_loss, optimizer) as scaled_loss: + scaled_loss.backward() + else: + closure_loss.backward() - # insert after step hook - if self.__is_function_implemented('on_after_backward'): - model_ref = self.__get_model() - model_ref.on_after_backward() + # insert after step hook + if self.__is_function_implemented('on_after_backward'): + model_ref = self.__get_model() + model_ref.on_after_backward() + + return closure_loss + + # calculate loss + loss = optimizer_closure() # nan grads if self.print_nan_grads: @@ -1311,7 +1320,7 @@ 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) + 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 d6075474..5cc290c2 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -44,7 +44,6 @@ def test_default_logger_callbacks_cpu_model(): Test each of the trainer options :return: """ - trainer_options = dict( max_nb_epochs=1, gradient_clip_val=1.0, @@ -63,6 +62,28 @@ def test_default_logger_callbacks_cpu_model(): model.unfreeze() +def test_lbfgs_cpu_model(): + """ + Test each of the trainer options + :return: + """ + 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 + ) + + model, hparams = get_model(use_test_model=True, lbfgs=True) + run_model_test_no_loggers(trainer_options, model, hparams, on_gpu=False) + + # test freeze on cpu + model.freeze() + model.unfreeze() + def test_multi_gpu_model_ddp2(): """ Make sure DDP2 works @@ -1447,9 +1468,11 @@ def get_hparams(continue_training=False, hpc_exp_number=0): return hparams -def get_model(use_test_model=False): +def get_model(use_test_model=False, lbfgs=False): # set up model with these hyperparams hparams = get_hparams() + if lbfgs: + hparams.optimizer = 'lbfgs' if use_test_model: model = LightningTestModel(hparams)