Î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
Instrucțiuni pentru exercițiu
- Importă
random_splitpentru 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, [____, ____])