Začněte nyníZačněte zdarma

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

Zobrazit kurz

Pokyny k cvičení

  • Naimportuj DataLoader.
  • Vrať DataLoader, který načítá self.train_data s 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=____) 
Upravit a spustit kód