Vytvoření train DataLoaderu
Teď, když máme dataset rozdělený, je potřeba definovat data loader, který bude během trénování poskytovat dávky dat. DataLoader efektivně načítá data do paměti a umožňuje jejich promíchání pro lepší generalizaci. V tomto cvičení dokončíš metodu train_dataloader.
Toto cvičení je součástí kurzu
Škálovatelné modely AI s PyTorch Lightning
Pokyny k cvičení
- Naimportuj
DataLoader. - Vrať
DataLoader, který načítáself.train_datas povoleným promícháváním dat pro lepší generalizaci.
Interaktivní cvičení na vyzkoušení si v praxi
Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.
# Import libraries
from torch.utils.data import ____
import lightning.pytorch as pl
class LoaderDataModule(pl.LightningDataModule):
def __init__(self):
super().__init__()
self.train_data = None
self.val_data = None
def setup(self, stage=None):
self.train_data, self.val_data = random_split(dataset, [80, 20])
def train_dataloader(self):
# Complete DataLoader
return ____(____, batch_size=16, shuffle=____)