Compare commits

..
16 Commits
4 changed files with 30 additions and 19 deletions
+24 -16
View File
@@ -33,6 +33,8 @@ class Trainer(TrainerIO):
log_save_interval=1, add_log_row_interval=1, log_save_interval=1, add_log_row_interval=1,
lr_scheduler_milestones=None, lr_scheduler_milestones=None,
use_amp=False, use_amp=False,
check_grad_nans=False,
amp_level='O2',
nb_sanity_val_steps=5): nb_sanity_val_steps=5):
# Transfer params # Transfer params
@@ -58,6 +60,8 @@ class Trainer(TrainerIO):
self.nb_sanity_val_steps = nb_sanity_val_steps self.nb_sanity_val_steps = nb_sanity_val_steps
self.lr_scheduler_milestones = [] if lr_scheduler_milestones is None else [int(x.strip()) for x in lr_scheduler_milestones.split(',')] self.lr_scheduler_milestones = [] if lr_scheduler_milestones is None else [int(x.strip()) for x in lr_scheduler_milestones.split(',')]
self.lr_schedulers = [] self.lr_schedulers = []
self.amp_level = amp_level
self.check_grad_nans = check_grad_nans
# training state # training state
self.optimizers = None self.optimizers = None
@@ -122,21 +126,21 @@ class Trainer(TrainerIO):
self.tqdm_metrics = {} self.tqdm_metrics = {}
# determine number of training batches # determine number of training batches
nb_tng_batches = self.model.nb_batches(self.tng_dataloader) self.nb_tng_batches = self.model.nb_batches(self.tng_dataloader)
self.nb_tng_batches = int(nb_tng_batches * self.train_percent_check) self.nb_tng_batches = int(self.nb_tng_batches * self.train_percent_check)
# determine number of validation batches # determine number of validation batches
nb_val_batches = self.model.nb_batches(self.val_dataloader) self.nb_val_batches = self.model.nb_batches(self.val_dataloader)
nb_val_batches = int(nb_val_batches * self.val_percent_check) self.nb_val_batches = int(self.nb_val_batches * self.val_percent_check)
nb_val_batches = max(1, nb_val_batches) self.nb_val_batches = max(1, self.nb_val_batches)
self.nb_val_batches = nb_val_batches self.nb_val_batches = self.nb_val_batches
# determine number of test batches # determine number of test batches
nb_test_batches = self.model.nb_batches(self.test_dataloader) self.nb_test_batches = self.model.nb_batches(self.test_dataloader)
self.nb_test_batches = int(nb_test_batches * self.test_percent_check) self.nb_test_batches = int(self.nb_test_batches * self.test_percent_check)
# determine when to check validation # determine when to check validation
self.val_check_batch = int(nb_tng_batches * self.val_check_interval) self.val_check_batch = int(self.nb_tng_batches * self.val_check_interval)
def __add_tqdm_metrics(self, metrics): def __add_tqdm_metrics(self, metrics):
for k, v in metrics.items(): for k, v in metrics.items():
@@ -163,19 +167,19 @@ class Trainer(TrainerIO):
outputs = [] outputs = []
# run training # run training
for i, data_batch in enumerate(dataloader): for batch_i, data_batch in enumerate(dataloader):
if data_batch is None: if data_batch is None:
continue continue
# stop short when on fast dev run # stop short when on fast dev run
if max_batches is not None and i >= max_batches: if max_batches is not None and batch_i >= max_batches:
break break
# ----------------- # -----------------
# RUN VALIDATION STEP # RUN VALIDATION STEP
# ----------------- # -----------------
output = model.validation_step(data_batch) output = model.validation_step(data_batch, batch_i)
outputs.append(output) outputs.append(output)
# batch done # batch done
@@ -222,7 +226,7 @@ class Trainer(TrainerIO):
if self.use_amp: if self.use_amp:
# An example # An example
self.model, optimizer = amp.initialize( self.model, optimizer = amp.initialize(
self.model, self.optimizers[0], opt_level="O2", self.model, self.optimizers[0], opt_level=self.amp_level,
) )
self.optimizers[0] = optimizer self.optimizers[0] = optimizer
model.trainer = self model.trainer = self
@@ -290,7 +294,7 @@ class Trainer(TrainerIO):
# --------------- # ---------------
# RUN TRAIN STEP # RUN TRAIN STEP
# --------------- # ---------------
batch_result = self.__run_tng_batch(data_batch) batch_result = self.__run_tng_batch(data_batch, batch_nb)
early_stop_epoch = batch_result == -1 early_stop_epoch = batch_result == -1
# --------------- # ---------------
@@ -348,7 +352,7 @@ class Trainer(TrainerIO):
return return
def __run_tng_batch(self, data_batch): def __run_tng_batch(self, data_batch, batch_nb):
if data_batch is None: if data_batch is None:
return 0 return 0
@@ -363,7 +367,7 @@ class Trainer(TrainerIO):
# forward pass # forward pass
# return a scalar value and a dic with tqdm metrics # return a scalar value and a dic with tqdm metrics
loss, model_specific_tqdm_metrics_dic = self.model.training_step(data_batch) loss, model_specific_tqdm_metrics_dic = self.model.training_step(data_batch, batch_nb)
self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic) self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic)
# backward pass # backward pass
@@ -374,6 +378,10 @@ class Trainer(TrainerIO):
else: else:
loss.backward() loss.backward()
if self.check_grad_nans:
for param in self.model.parameters():
print(param.grad.float().sum())
self.batch_loss_value += loss.item() self.batch_loss_value += loss.item()
# gradient update with accumulated gradients # gradient update with accumulated gradients
+2 -2
View File
@@ -51,7 +51,7 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
""" """
raise NotImplementedError raise NotImplementedError
def validation_step(self, data_batch): def validation_step(self, data_batch, batch_nb):
""" """
return whatever outputs will need to be aggregated in validation_end return whatever outputs will need to be aggregated in validation_end
:param data_batch: :param data_batch:
@@ -67,7 +67,7 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
""" """
raise NotImplementedError raise NotImplementedError
def training_step(self, data_batch): def training_step(self, data_batch, batch_nb):
""" """
return loss, dict with metrics for tqdm return loss, dict with metrics for tqdm
:param data_batch: :param data_batch:
+3
View File
@@ -51,6 +51,9 @@ def add_default_args(parser, root_dir, rand_seed=None, possible_model_names=None
parser.add_argument('--disable_cuda', dest='disable_cuda', action='store_true') parser.add_argument('--disable_cuda', dest='disable_cuda', action='store_true')
parser.add_argument('--default_tensor_type', default='torch.cuda.FloatTensor', type=str) parser.add_argument('--default_tensor_type', default='torch.cuda.FloatTensor', type=str)
parser.add_argument('--use_amp', dest='use_amp', action='store_true') parser.add_argument('--use_amp', dest='use_amp', action='store_true')
parser.add_argument('--check_grad_nans', dest='check_grad_nans', action='store_true')
parser.add_argument('--amp_level', default='O2',type=str)
# run on hpc # run on hpc
parser.add_argument('--on_cluster', dest='on_cluster', action='store_true') parser.add_argument('--on_cluster', dest='on_cluster', action='store_true')
+1 -1
View File
@@ -7,7 +7,7 @@ from setuptools import setup, find_packages
# http://blog.ionelmc.ro/2014/05/25/python-packaging/ # http://blog.ionelmc.ro/2014/05/25/python-packaging/
setup( setup(
name="pytorch-lightning", name="pytorch-lightning",
version='0.1.dev182', version='0.1.dev1832',
description="The Keras for ML researchers using PyTorch", description="The Keras for ML researchers using PyTorch",
author="William Falcon", author="William Falcon",
author_email="waf2107@columbia.edu", author_email="waf2107@columbia.edu",