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

Chia dữ liệu với LightningDataModule

Bạn sẽ hoàn thiện phương thức setup trong một LightningDataModule. Phân chia tập dữ liệu đúng cách giúp mô hình được huấn luyện trên một phần và được validation trên phần khác, từ đó tránh overfitting.

dataset đã được nhập sẵn.

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 random_split để chia tập dữ liệu thành huấn luyện và validation.
  • Chia tập dữ liệu thành huấn luyện (80%) và validation (20%) bằng random_split.

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 
import lightning.pytorch as pl
from torch.utils.data import ____

class SplitDataModule(pl.LightningDataModule):
    def __init__(self):
        super().__init__()
        self.train_data = None
        self.val_data = None
    def setup(self, stage=None):
        # Split the dataset into training (80%) and validation (20%)
        self.____, self.____ = random_split(dataset, [____, ____])
Chỉnh sửa và Chạy Mã