diff --git a/docs/source/introduction_guide.rst b/docs/source/introduction_guide.rst index c3dbb57e..765ba8fd 100644 --- a/docs/source/introduction_guide.rst +++ b/docs/source/introduction_guide.rst @@ -494,13 +494,13 @@ In this method we do all the preparation we need to do once (instead of on every self.test_dataset = mnist_test def train_dataloader(self): - return DataLoader(train_dataset, batch_size=64) + return DataLoader(self.train_dataset, batch_size=64) def val_dataloader(self): - return DataLoader(mnist_val, batch_size=64) + return DataLoader(self.mnist_val, batch_size=64) def test_dataloader(self): - return DataLoader(mnist_test, batch_size=64) + return DataLoader(self.mnist_test, batch_size=64) The `prepare_data` method is also a good place to do any data processing that needs to be done only once (ie: download or tokenize, etc...). diff --git a/pytorch_lightning/core/__init__.py b/pytorch_lightning/core/__init__.py index 39479cf2..fac76ffd 100644 --- a/pytorch_lightning/core/__init__.py +++ b/pytorch_lightning/core/__init__.py @@ -236,11 +236,11 @@ Data preparation ---------------- Data preparation in PyTorch follows 5 steps: -1. Download -2. Clean and (maybe) save to disk -3. Load inside dataset -4. Apply transforms (rotate, tokenize, etc...) -5. Wrap inside a dataloader + 1. Download + 2. Clean and (maybe) save to disk + 3. Load inside dataset + 4. Apply transforms (rotate, tokenize, etc...) + 5. Wrap inside a dataloader When working in distributed settings, steps 1 and 2 have to be done from a single GPU, otherwise you will overwrite these files from @@ -251,8 +251,10 @@ allow for this def prepare_data(self): # download - mnist_train = MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()) - mnist_test = MNIST(os.getcwd(), train=False, download=True, transform=transforms.ToTensor()) + mnist_train = MNIST(os.getcwd(), train=True, download=True, + transform=transforms.ToTensor()) + mnist_test = MNIST(os.getcwd(), train=False, download=True, + transform=transforms.ToTensor()) # train/val split mnist_train, mnist_val = random_split(mnist_train, [55000, 5000]) @@ -263,13 +265,13 @@ allow for this self.test_dataset = mnist_test def train_dataloader(self): - return DataLoader(train_dataset, batch_size=64) + return DataLoader(self.train_dataset, batch_size=64) def val_dataloader(self): - return DataLoader(mnist_val, batch_size=64) + return DataLoader(self.mnist_val, batch_size=64) def test_dataloader(self): - return DataLoader(mnist_test, batch_size=64) + return DataLoader(self.mnist_test, batch_size=64) .. note:: ``prepare_data`` is called once. @@ -305,6 +307,9 @@ Check out this `COLAB `_ for a live demo. +LightningModule Class +--------------------- + """ from .decorators import data_loader