Crearea unui DataLoader pentru antrenament
Acum că am împărțit setul de date, trebuie să definim un data loader care să furnizeze loturi de date în timpul antrenamentului. DataLoader încarcă datele eficient în memorie și permite amestecarea lor pentru o generalizare mai bună. În acest exercițiu, vei completa metoda train_dataloader.
Acest exercițiu face parte din cursul
Modele AI scalabile cu PyTorch Lightning
Instrucțiuni pentru exercițiu
- Importă
DataLoader. - Returnează un
DataLoadercare încarcăself.train_data, activând amestecarea datelor pentru o generalizare mai bună.
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
# 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=____)