From 3f761524707ddd1814603ef02b61ed53f83a8366 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Sun, 21 Jul 2019 18:23:48 -0400 Subject: [PATCH] added on_after_backward --- pytorch_lightning/models/trainer.py | 5 +++++ pytorch_lightning/root_module/hooks.py | 7 +++++++ 2 files changed, 12 insertions(+) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 55f084af..8108e01d 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -775,6 +775,11 @@ class Trainer(TrainerIO): else: loss.backward() + # insert after step hook + if self.__is_function_implemented('on_after_backward'): + model_ref = self.__get_model() + response = model_ref.on_after_backward() + if self.print_nan_grads: model = self.__get_model() for param in model.parameters(): diff --git a/pytorch_lightning/root_module/hooks.py b/pytorch_lightning/root_module/hooks.py index d0e2bbaa..88abe80d 100644 --- a/pytorch_lightning/root_module/hooks.py +++ b/pytorch_lightning/root_module/hooks.py @@ -36,3 +36,10 @@ class ModelHooks(torch.nn.Module): """ pass + def on_after_backward(self): + """ + Called after loss.backward() and before optimizers do anything + :return: + """ + pass +