Tworzenie DataLoadera dla danych treningowych
Skoro podzieliliśmy już zbiór danych, musimy zdefiniować data loader, który będzie dostarczał partie danych podczas treningu. DataLoader sprawnie wczytuje dane do pamięci i umożliwia ich losowanie, co poprawia generalizację modelu. W tym ćwiczeniu uzupełnisz metodę train_dataloader.
To ćwiczenie jest częścią kursu
Skalowalne modele AI z PyTorch Lightning
Instrukcje do ćwiczenia
- Zaimportuj
DataLoader. - Zwróć
DataLoader, który wczytujeself.train_dataz włączonym losowaniem kolejności danych, aby poprawić generalizację.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# Import libraries
from torch.utils.data import ____
import lightning.pytorch as pl
class LoaderDataModule(pl.LightningDataModule):
def __init__(self):
super().__init__()
self.train_data = None
self.val_data = None
def setup(self, stage=None):
self.train_data, self.val_data = random_split(dataset, [80, 20])
def train_dataloader(self):
# Complete DataLoader
return ____(____, batch_size=16, shuffle=____)