ÎncepețiÎncepe gratuit

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

Vezi cursul

Instrucțiuni pentru exercițiu

  • Importă DataLoader.
  • Returnează un DataLoader care î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=____) 
Editează și rulează codul