Podział danych za pomocą LightningDataModule
Uzupełnij metodę setup w klasie LightningDataModule. Odpowiedni podział zbioru danych sprawia, że model jest trenowany na jednym podzbiorze, a walidowany na innym – co pomaga unikać przetrenowania.
dataset został już wcześniej zaimportowany.
To ćwiczenie jest częścią kursu
Skalowalne modele AI z PyTorch Lightning
Instrukcje do ćwiczenia
- Zaimportuj
random_split, aby podzielić zbiór danych na podzbiory treningowy i walidacyjny. - Podziel zbiór danych na podzbiór treningowy (80%) i walidacyjny (20%) za pomocą
random_split.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# 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, [____, ____])