added on_after_backward

This commit is contained in:
William Falcon
2019-07-21 18:23:48 -04:00
parent d98b9f2f93
commit 3f76152470
2 changed files with 12 additions and 0 deletions
+5
View File
@@ -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():
+7
View File
@@ -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