Move logger initialization (#750)

This commit is contained in:
Vadim Bereznyuk
2020-01-26 09:42:57 -05:00
committed by William Falcon
parent cc12ff36a9
commit 7deec2c14e
3 changed files with 21 additions and 19 deletions
+1 -17
View File
@@ -2,7 +2,6 @@ import os
from abc import ABC
from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping
from pytorch_lightning.logging import TensorBoardLogger
class TrainerCallbackConfigMixin(ABC):
@@ -50,7 +49,7 @@ class TrainerCallbackConfigMixin(ABC):
if self.weights_save_path is None:
self.weights_save_path = self.default_save_path
def configure_early_stopping(self, early_stop_callback, logger):
def configure_early_stopping(self, early_stop_callback):
if early_stop_callback is True:
self.early_stop_callback = EarlyStopping(
monitor='val_loss',
@@ -75,18 +74,3 @@ class TrainerCallbackConfigMixin(ABC):
else:
self.early_stop_callback = early_stop_callback
self.enable_early_stop = True
# configure logger
if logger is True:
# default logger
self.logger = TensorBoardLogger(
save_dir=self.default_save_path,
version=self.slurm_job_id,
name='lightning_logs'
)
self.logger.rank = 0
elif logger is False:
self.logger = None
else:
self.logger = logger
self.logger.rank = 0
+16
View File
@@ -3,6 +3,7 @@ from abc import ABC
import torch
from pytorch_lightning.core import memory
from pytorch_lightning.logging import TensorBoardLogger
class TrainerLoggingMixin(ABC):
@@ -21,6 +22,21 @@ class TrainerLoggingMixin(ABC):
self.use_ddp2 = None
self.num_gpus = None
def configure_logger(self, logger):
if logger is True:
# default logger
self.logger = TensorBoardLogger(
save_dir=self.default_save_path,
version=self.slurm_job_id,
name='lightning_logs'
)
self.logger.rank = 0
elif logger is False:
self.logger = None
else:
self.logger = logger
self.logger.rank = 0
def log_metrics(self, metrics, grad_norm_dic, step=None):
"""Logs the metric dict passed in.
+4 -2
View File
@@ -558,10 +558,12 @@ class Trainer(TrainerIOMixin,
self.current_epoch = 0
self.total_batches = 0
# configure logger
self.configure_logger(logger)
# configure early stop callback
# creates a default one if none passed in
self.early_stop_callback = None
self.configure_early_stopping(early_stop_callback, logger)
self.configure_early_stopping(early_stop_callback)
self.reduce_lr_on_plateau_scheduler = None