mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-08-21 11:20:03 +08:00
* add doctest to circleci * Revert "add doctest to circleci" This reverts commit c45b34ea911a81f87989f6c3a832b1e8d8c471c6. * Revert "Revert "add doctest to circleci"" This reverts commit 41fca97fdcfe1cf4f6bdb3bbba75d25fa3b11f70. * doctest docs rst files * Revert "doctest docs rst files" This reverts commit b4a2e83e3da5ed1909de500ec14b6b614527c07f. * doctest only rst * doctest debugging.rst * doctest apex * doctest callbacks * doctest early stopping * doctest for child modules * doctest experiment reporting * indentation * doctest fast training * doctest for hyperparams * doctests for lr_finder * doctests multi-gpu * more doctest * make doctest drone * fix label build error * update fast training * update invalid imports * fix problem with int device count * rebase stuff * wip * wip * wip * intro guide * add missing code block * circleci * logger import for doctest * test if doctest runs on drone * fix mnist download * also run install deps for building docs * install cmake * try sudo * hide output * try pip stuff * try to mock horovod * Tranfer -> Transfer * add torchvision to extras * revert pip stuff * mlflow file location * do not mock torch * torchvision * drone extra req. * try higher sphinx version * Revert "try higher sphinx version" This reverts commit 490ac28e46d6fd52352640dfdf0d765befa56988. * try coverage command * try coverage command * try undoc flag * newline * undo drone * report coverage * review Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com> * remove torchvision from extras * skip tests only if torchvision not available * fix testoutput torchvision Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com>
74 lines
2.1 KiB
ReStructuredText
74 lines
2.1 KiB
ReStructuredText
.. testsetup:: *
|
|
|
|
from pytorch_lightning.core.lightning import LightningModule
|
|
|
|
Multiple Datasets
|
|
=================
|
|
Lightning supports multiple dataloaders in a few ways.
|
|
|
|
1. Create a dataloader that iterates both datasets under the hood.
|
|
2. In the validation and test loop you also have the option to return multiple dataloaders
|
|
which lightning will call sequentially.
|
|
|
|
Multiple training dataloaders
|
|
-----------------------------
|
|
For training, the best way to use multiple-dataloaders is to create a Dataloader class
|
|
which wraps both your dataloaders. (This of course also works for testing and validation
|
|
dataloaders).
|
|
|
|
(`reference <https://discuss.pytorch.org/t/train-simultaneously-on-two-datasets/649/2>`_)
|
|
|
|
.. testcode::
|
|
|
|
class ConcatDataset(torch.utils.data.Dataset):
|
|
def __init__(self, *datasets):
|
|
self.datasets = datasets
|
|
|
|
def __getitem__(self, i):
|
|
return tuple(d[i] for d in self.datasets)
|
|
|
|
def __len__(self):
|
|
return min(len(d) for d in self.datasets)
|
|
|
|
class LitModel(LightningModule):
|
|
|
|
def train_dataloader(self):
|
|
concat_dataset = ConcatDataset(
|
|
datasets.ImageFolder(traindir_A),
|
|
datasets.ImageFolder(traindir_B)
|
|
)
|
|
|
|
loader = torch.utils.data.DataLoader(
|
|
concat_dataset,
|
|
batch_size=args.batch_size,
|
|
shuffle=True,
|
|
num_workers=args.workers,
|
|
pin_memory=True
|
|
)
|
|
return loader
|
|
|
|
def val_dataloader(self):
|
|
# SAME
|
|
...
|
|
|
|
def test_dataloader(self):
|
|
# SAME
|
|
...
|
|
|
|
Test/Val dataloaders
|
|
--------------------
|
|
For validation, test dataloaders lightning also gives you the additional
|
|
option of passing in multiple dataloaders back from each call.
|
|
|
|
See the following for more details:
|
|
|
|
- :meth:`~pytorch_lightning.core.LightningModule.val_dataloader`
|
|
- :meth:`~pytorch_lightning.core.LightningModule.test_dataloader`
|
|
|
|
.. testcode::
|
|
|
|
def val_dataloader(self):
|
|
loader_1 = Dataloader()
|
|
loader_2 = Dataloader()
|
|
return [loader_1, loader_2]
|