From f6416f737d73f5405e1a2f380022b80b7f1e5910 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Sun, 21 Jul 2019 18:15:58 -0400 Subject: [PATCH] added grad hook --- pytorch_lightning/models/trainer.py | 5 +++++ pytorch_lightning/root_module/hooks.py | 14 ++++++++++++++ 2 files changed, 19 insertions(+) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 79d2c2dc..55f084af 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -795,6 +795,11 @@ class Trainer(TrainerIO): for optimizer in self.optimizers: optimizer.step() + # insert after step hook + if self.__is_function_implemented('on_before_zero_grad'): + model_ref = self.__get_model() + response = model_ref.on_before_zero_grad(optimizer) + # clear gradients optimizer.zero_grad() diff --git a/pytorch_lightning/root_module/hooks.py b/pytorch_lightning/root_module/hooks.py index 6d5c5dcd..d0e2bbaa 100644 --- a/pytorch_lightning/root_module/hooks.py +++ b/pytorch_lightning/root_module/hooks.py @@ -22,3 +22,17 @@ class ModelHooks(torch.nn.Module): def on_tng_metrics(self, metrics): pass + def on_before_zero_grad(self, optimizer): + """ + Called after optimizer.step() and before optimizer.zero_grad() + + for optimizer in optimizers: + optimizer.step() + model.on_before_zero_grad(optimizer) # < ---- called here + optimizer.zero_grad + + :param optimizer: + :return: + """ + pass +