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
Pokyny k cvičení
- Importuj
random_splitpro 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, [____, ____])