From 0ae3dd9ed46440f54b1014370261a96bcac8f9b7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ayberk=20Ayd=C4=B1n?= Date: Tue, 14 Jan 2020 06:12:04 +0300 Subject: [PATCH] Fix GAN training. (#603) * fix dangling gradients make sure only the gradients of the current optimizer's paramaters are calculated in the training step. * add note about multiple optimizer gradient update * Update training_loop.py --- pytorch_lightning/core/lightning.py | 2 ++ pytorch_lightning/trainer/training_loop.py | 7 +++++++ 2 files changed, 9 insertions(+) diff --git a/pytorch_lightning/core/lightning.py b/pytorch_lightning/core/lightning.py index d2ba6c8a..2e62e03b 100644 --- a/pytorch_lightning/core/lightning.py +++ b/pytorch_lightning/core/lightning.py @@ -664,6 +664,8 @@ class LightningModule(GradInformation, ModelIO, ModelHooks): .. note:: If you use multiple optimizers, training_step will have an additional `optimizer_idx` parameter. .. note:: If you use LBFGS lightning handles the closure function automatically for you. + + .. note:: If you use multiple optimizers, gradients will be calculated only for the parameters of current optimizer at each training step. Example ------- diff --git a/pytorch_lightning/trainer/training_loop.py b/pytorch_lightning/trainer/training_loop.py index 5056fdd7..cc9e6e21 100644 --- a/pytorch_lightning/trainer/training_loop.py +++ b/pytorch_lightning/trainer/training_loop.py @@ -454,6 +454,13 @@ class TrainerTrainLoopMixin(ABC): # call training_step once per optimizer for opt_idx, optimizer in enumerate(self.optimizers): + # make sure only the gradients of the current optimizer's paramaters are calculated + # in the training step to prevent dangling gradients in multiple-optimizer setup. + for param in self.get_model().parameters(): + param.requires_grad = False + for group in optimizer.param_groups: + for param in group['params']: + param.requires_grad = True # wrap the forward step in a closure so second order methods work def optimizer_closure():