From 619143a734a641e2f22d74ad639488a2809b408c Mon Sep 17 00:00:00 2001 From: Jeffrey Ling Date: Tue, 19 Nov 2019 15:38:54 -0800 Subject: [PATCH] Fix incorrect handling of on_batch_end edge cases in run_training_batch (#509) * Fix returning only 2 values on an early exit. This fixes a bug `ValueError: not enough values to unpack (expected 3, got 2)` * Update train_loop_mixin.py * Change to return dict The return value was actually a dict even though that variable is initialized as a list. --- pytorch_lightning/trainer/train_loop_mixin.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pytorch_lightning/trainer/train_loop_mixin.py b/pytorch_lightning/trainer/train_loop_mixin.py index 68c34cf4..7e9fde01 100644 --- a/pytorch_lightning/trainer/train_loop_mixin.py +++ b/pytorch_lightning/trainer/train_loop_mixin.py @@ -155,7 +155,7 @@ class TrainerTrainLoopMixin(object): all_log_metrics = [] if batch is None: - return 0, grad_norm_dic + return 0, grad_norm_dic, {} # hook if self.is_function_implemented('on_batch_start'): @@ -163,7 +163,7 @@ class TrainerTrainLoopMixin(object): response = model_ref.on_batch_start(batch) if response == -1: - return -1, grad_norm_dic + return -1, grad_norm_dic, {} splits = [batch] if self.truncated_bptt_steps is not None: