建立訓練用 DataLoader
現在我們已經完成資料集的切分,接下來需要定義一個資料載入器,在訓練期間提供批次資料。DataLoader 能有效率地將資料載入記憶體,並支援隨機打亂以提升泛化能力。在這個練習中,你將完成 train_dataloader 方法。
本練習屬於課程
Scalable AI Models with PyTorch Lightning
練習說明
- 匯入
DataLoader。 - 回傳一個能載入
self.train_data的DataLoader,並啟用隨機打亂以提升泛化能力。
動手互動練習
試著完成這個範例程式碼,體驗一下這個練習。
# 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=____)