Zacznij terazZacznij za darmo

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

Zobacz kurs

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, [____, ____])
Edytuj i uruchom kod