Val idx optional in validation_step (#108)

* made dataset_i only available with multiple datasets

* updated interface signature

* updated tests
This commit is contained in:
William Falcon
2019-08-13 11:37:37 -04:00
committed by GitHub
parent 905a2e5a12
commit 7f53e7bfb3
6 changed files with 48 additions and 34 deletions
@@ -105,7 +105,7 @@ class LightningTemplateModel(LightningModule):
# can also return just a scalar instead of a dict (return loss_val)
return output
def validation_step(self, data_batch, batch_i, dataloader_i):
def validation_step(self, data_batch, batch_i):
"""
Lightning calls this inside the validation loop
:param data_batch:
@@ -218,7 +218,7 @@ class LightningTemplateModel(LightningModule):
@pl.data_loader
def val_dataloader(self):
print('val data loader called')
return [self.__dataloader(train=False) for i in range(2)]
return self.__dataloader(train=False)
@pl.data_loader
def test_dataloader(self):
+28 -16
View File
@@ -377,6 +377,32 @@ class Trainer(TrainerIO):
self.tqdm_metrics[k] = v
def __validation_forward(self, model, data_batch, batch_i, dataloader_i):
# make dataloader_i arg in validation_step optional
args = [data_batch, batch_i]
if len(self.val_dataloader) > 1:
args.append(dataloader_i)
if self.use_ddp:
output = model(*args)
elif self.use_dp:
output = model(*args)
elif self.single_gpu:
# put inputs on gpu manually
gpu_id = self.data_parallel_device_ids[0]
for i, x in enumerate(data_batch):
if isinstance(x, torch.Tensor):
data_batch[i] = x.cuda(gpu_id)
# do non dp, ddp step
output = model.validation_step(*args)
else:
# CPU
output = model.validation_step(*args)
return output
def validate(self, model, dataloader, max_batches, dataloader_i):
"""
Run validation code
@@ -409,23 +435,9 @@ class Trainer(TrainerIO):
# -----------------
# RUN VALIDATION STEP
# -----------------
if self.use_ddp:
output = model(data_batch, batch_i, dataloader_i)
elif self.use_dp:
output = model(data_batch, batch_i, dataloader_i)
elif self.single_gpu:
# put inputs on gpu manually
gpu_id = self.data_parallel_device_ids[0]
for i, x in enumerate(data_batch):
if isinstance(x, torch.Tensor):
data_batch[i] = x.cuda(gpu_id)
# do non dp, ddp step
output = model.validation_step(data_batch, batch_i, dataloader_i)
else:
output = model.validation_step(data_batch, batch_i, dataloader_i)
output = self.__validation_forward(model, data_batch, batch_i, dataloader_i)
# track outputs for collation
outputs.append(output)
# batch done
+12 -10
View File
@@ -33,11 +33,21 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
"""
raise NotImplementedError
def validation_step(self, data_batch, batch_nb):
def training_step(self, *args, **kwargs):
"""
return loss, dict with metrics for tqdm
:param called with batch, batch_nb
additional: optimizer_i if multiple optimizers used
:return:
"""
raise NotImplementedError
def validation_step(self, *args, **kwargs):
"""
return whatever outputs will need to be aggregated in validation_end
OPTIONAL
:param data_batch:
:param called with batch, batch_nb
additional: dataset_i if multiple val datasets used
:return:
"""
pass
@@ -51,14 +61,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
"""
pass
def training_step(self, data_batch, batch_nb):
"""
return loss, dict with metrics for tqdm
:param data_batch:
:return:
"""
raise NotImplementedError
def configure_optimizers(self):
"""
Return a list of optimizers and a list of schedulers (could be empty)
+1 -1
View File
@@ -231,7 +231,7 @@ class LightningTestModel(LightningModule):
@data_loader
def val_dataloader(self):
return self.__dataloader(train=False)
return [self.__dataloader(train=False), self.__dataloader(train=False)]
@data_loader
def test_dataloader(self):
@@ -109,7 +109,7 @@ class NoValEndTestModel(LightningModule):
if self.trainer.batch_nb % 2 == 0:
return loss_val
def validation_step(self, data_batch, batch_i, dataloader_i):
def validation_step(self, data_batch, batch_nb):
"""
Lightning calls this inside the validation loop
:param data_batch:
@@ -135,16 +135,16 @@ class NoValEndTestModel(LightningModule):
val_acc = val_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_i % 1 == 0:
if batch_nb % 1 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
})
return output
if batch_i % 2 == 0:
if batch_nb % 2 == 0:
return val_acc
if batch_i % 3 == 0:
if batch_nb % 3 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
+1 -1
View File
@@ -767,7 +767,7 @@ def test_multiple_val_dataloader():
:return:
"""
hparams = get_hparams()
model = LightningTemplateModel(hparams)
model = LightningTestModel(hparams)
save_dir = init_save_dir()