From 019b4d16d02446caeb60412b60da2301e6fd62f8 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Sun, 4 Aug 2019 13:08:14 -0500 Subject: [PATCH] formatting --- pytorch_lightning/root_module/hooks.py | 1 + pytorch_lightning/root_module/model_saving.py | 5 ++--- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/pytorch_lightning/root_module/hooks.py b/pytorch_lightning/root_module/hooks.py index 88abe80d..849826a8 100644 --- a/pytorch_lightning/root_module/hooks.py +++ b/pytorch_lightning/root_module/hooks.py @@ -1,5 +1,6 @@ import torch + class ModelHooks(torch.nn.Module): def on_batch_start(self, data_batch): pass diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/root_module/model_saving.py index 39c5ae7b..34e4485b 100644 --- a/pytorch_lightning/root_module/model_saving.py +++ b/pytorch_lightning/root_module/model_saving.py @@ -1,7 +1,6 @@ import torch import os import re -import pdb from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel, LightningDataParallel @@ -77,7 +76,7 @@ class TrainerIO(object): optimizer_states.append(optimizer.state_dict()) checkpoint['optimizer_states'] = optimizer_states - + # save lr schedulers lr_schedulers = [] for i, scheduler in enumerate(self.lr_schedulers): @@ -141,7 +140,7 @@ class TrainerIO(object): optimizer_states = checkpoint['optimizer_states'] for optimizer, opt_state in zip(self.optimizers, optimizer_states): optimizer.load_state_dict(opt_state) - + # restore the lr schedulers lr_schedulers = checkpoint['lr_schedulers'] for scheduler, lrs_state in zip(self.lr_schedulers, lr_schedulers):