From 81df2259ef0adb7f510ffb998d661ed289981548 Mon Sep 17 00:00:00 2001 From: Alok Singh <8325708+alok@users.noreply.github.com> Date: Fri, 6 Sep 2019 22:08:09 -0700 Subject: [PATCH] Make print_nan_grads print grad (#208) This seems more useful for debugging. --- pytorch_lightning/trainer/trainer.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index aa673abd..70a41368 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -1091,10 +1091,10 @@ class Trainer(TrainerIO): torch.nn.utils.clip_grad_norm_(model.parameters(), self.gradient_clip) def __print_nan_grads(self): - if self.print_nan_grads: - model = self.__get_model() - for param in model.parameters(): - print(param.grad.float().sum()) + model = self.__get_model() + for param in model.parameters(): + if torch.isnan(param.grad.float()).any(): + print(param, param.grad) def __run_tng_batch(self, data_batch, batch_nb): if data_batch is None: @@ -1137,7 +1137,8 @@ class Trainer(TrainerIO): model_ref.on_after_backward() # nan grads - self.__print_nan_grads() + if self.print_nan_grads: + self.__print_nan_grads() # track total loss for logging (avoid mem leaks) self.batch_loss_value += loss.item()