Bắt đầu ngayBắt đầu miễn phí

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

Xem khóa học

Hướng dẫn bài tập

  • Import DataLoader.
  • Trả về một DataLoader tải self.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=____) 
Chỉnh sửa và Chạy Mã