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
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, [____, ____])