Создание DataLoader для обучения
Теперь, когда набор данных разбит на части, нам нужно определить загрузчик данных, который будет подавать пакеты данных в процессе обучения. DataLoader эффективно загружает данные в память и поддерживает перемешивание для улучшения обобщающей способности модели. В этом упражнении вы завершите реализацию метода train_dataloader.
Это упражнение является частью курса
Масштабируемые модели ИИ с PyTorch Lightning
Инструкции к упражнению
- Импортируйте
DataLoader. - Верните
DataLoader, который загружаетself.train_dataс включённым перемешиванием для улучшения обобщающей способности модели.
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# 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=____)