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

Rozdělení dat pomocí LightningDataModule

Doplníš metodu setup ve třídě LightningDataModule. Správné rozdělení datasetu zajistí, že model bude trénovaný na jedné části dat a validovaný na druhé – a pomůže tak předejít přetrénování.

dataset je již předem naimportovaný.

Toto cvičení je součástí kurzu

Škálovatelné modely AI s PyTorch Lightning

Zobrazit kurz

Pokyny k cvičení

  • Importuj random_split pro rozdělení datasetu na trénovací a validační část.
  • Rozděl dataset na trénovací (80 %) a validační část (20 %) pomocí random_split.

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 
import lightning.pytorch as pl
from torch.utils.data import ____

class SplitDataModule(pl.LightningDataModule):
    def __init__(self):
        super().__init__()
        self.train_data = None
        self.val_data = None
    def setup(self, stage=None):
        # Split the dataset into training (80%) and validation (20%)
        self.____, self.____ = random_split(dataset, [____, ____])
Upravit a spustit kód