開始使用免費開始

建立訓練用 DataLoader

現在我們已經完成資料集的切分,接下來需要定義一個資料載入器,在訓練期間提供批次資料。DataLoader 能有效率地將資料載入記憶體,並支援隨機打亂以提升泛化能力。在這個練習中,你將完成 train_dataloader 方法。

本練習屬於課程

Scalable AI Models with PyTorch Lightning

檢視課程

練習說明

  • 匯入 DataLoader
  • 回傳一個能載入 self.train_dataDataLoader,並啟用隨機打亂以提升泛化能力。

動手互動練習

試著完成這個範例程式碼,體驗一下這個練習。

# 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=____) 
編輯並執行程式碼