From 0bcc858cef6c0570545bc3e0f3d7c22192c47af7 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 8 Jul 2019 19:11:16 -0400 Subject: [PATCH] moved sampler --- pytorch_lightning/models/trainer.py | 6 +++--- pytorch_lightning/root_module/root_module.py | 3 --- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 3daa759d..39d0b761 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -175,17 +175,17 @@ class Trainer(TrainerIO): self.tqdm_metrics = {} # determine number of training batches - self.nb_tng_batches = model.nb_batches(self.tng_dataloader) + self.nb_tng_batches = len(self.tng_dataloader) self.nb_tng_batches = int(self.nb_tng_batches * self.train_percent_check) # determine number of validation batches - self.nb_val_batches = model.nb_batches(self.val_dataloader) + self.nb_val_batches = len(self.val_dataloader) self.nb_val_batches = int(self.nb_val_batches * self.val_percent_check) self.nb_val_batches = max(1, self.nb_val_batches) self.nb_val_batches = self.nb_val_batches # determine number of test batches - self.nb_test_batches = model.nb_batches(self.test_dataloader) + self.nb_test_batches = len(self.test_dataloader) self.nb_test_batches = int(self.nb_test_batches * self.test_percent_check) # determine when to check validation diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index cd36b19f..78ef215e 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -93,9 +93,6 @@ class LightningModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks): model_summary = ModelSummary(self) print(model_summary) - def nb_batches(self, dataloader): - a = math.ceil(float(len(dataloader.dataset) / self.batch_size)) - return int(a) def freeze(self): for param in self.parameters():