Tạo DataLoader cho train
Sau khi đã tách dataset, bạn cần định nghĩa một data loader để cung cấp các batch dữ liệu trong quá trình huấn luyện. DataLoader tải dữ liệu vào bộ nhớ hiệu quả và cho phép xáo trộn để mô hình tổng quát hóa tốt hơn. Trong bài này, bạn sẽ hoàn thiện phương thức train_dataloader.
Bài tập này là một phần của khóa học
Mô hình AI có khả năng mở rộng với PyTorch Lightning
Hướng dẫn bài tập
- Import
DataLoader. - Trả về một
DataLoadertảiself.train_data, bật xáo trộn để giúp tổng quát hóa tốt hơn.
Bài tập tương tác thực hành trực tiếp
Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.
# 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=____)