ÎncepețiÎncepe gratuit

Împărțirea datelor cu LightningDataModule

Vei completa metoda setup dintr-un LightningDataModule. Împărțirea corectă a setului de date asigură că modelul este antrenat pe un subset și validat pe altul, prevenind astfel supraadaptarea (overfitting).

dataset a fost deja importat în prealabil.

Acest exercițiu face parte din cursul

Modele AI scalabile cu PyTorch Lightning

Vezi cursul

Instrucțiuni pentru exercițiu

  • Importă random_split pentru a împărți setul de date în seturi de antrenament și validare.
  • Împarte setul de date în antrenament (80%) și validare (20%) folosind random_split.

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

# 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, [____, ____])
Editează și rulează codul