LightningDataModule — 데이터 파이프라인 캡슐화

LightningDataModule — 데이터 파이프라인 캡슐화

데이터모듈(datamodule)은 데이터를 처리하는 데 필요한 모든 단계를 담는 공유·재사용 가능한 클래스예요. PyTorch의 데이터 처리 다섯 단계를 하나로 캡슐화합니다.

출처: https://lightning.ai/docs/pytorch/stable/data/datamodule.html

  1. Download / tokenize / process.
  2. Clean and (maybe) save to disk.
  3. Load inside Dataset.
  4. Apply transforms (rotate, tokenize, etc…).
  5. Wrap inside a DataLoader.

이 클래스는 어디서든 공유해 쓸 수 있어요.

model = LitClassifier()
trainer = Trainer()
imagenet = ImagenetDataModule()
trainer.fit(model, datamodule=imagenet)
cifar10 = CIFAR10DataModule()
trainer.fit(model, datamodule=cifar10)

왜 DataModule이 필요할까

일반 PyTorch 코드에서는 데이터 정리·준비가 여러 파일에 흩어져 있어서 split과 transform을 프로젝트 간에 공유·재사용하기가 불가능해요. "무슨 split을 썼지?", "어떤 transform을 썼지?", "어떻게 토크나이즈했지?" 같은 질문을 한 적이 있다면 데이터모듈이 해답이에요.

구현 예시

import lightning as L
from torch.utils.data import random_split, DataLoader
# Note - you must have torchvision installed for this example
from torchvision.datasets import MNIST
from torchvision import transforms

class MNISTDataModule(L.LightningDataModule):
    def __init__(self, data_dir: str = "./"):
        super().__init__()
        self.data_dir = data_dir
        self.transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.1307,), (0.3081,))
        ])

    def prepare_data(self):
        # download
        MNIST(self.data_dir, train=True, download=True)
        MNIST(self.data_dir, train=False, download=True)

    def setup(self, stage: str):
        # Assign train/val datasets for use in dataloaders
        if stage == "fit":
            mnist_full = MNIST(self.data_dir, train=True, transform=self.transform)
            self.mnist_train, self.mnist_val = random_split(
                mnist_full, [55000, 5000],
                generator=torch.Generator().manual_seed(42)
            )
        # Assign test dataset for use in dataloader(s)
        if stage == "test":
            self.mnist_test = MNIST(self.data_dir, train=False, transform=self.transform)
        if stage == "predict":
            self.mnist_predict = MNIST(self.data_dir, train=False, transform=self.transform)

    def train_dataloader(self):
        return DataLoader(self.mnist_train, batch_size=32)
    def val_dataloader(self):
        return DataLoader(self.mnist_val, batch_size=32)
    def test_dataloader(self):
        return DataLoader(self.mnist_test, batch_size=32)
    def predict_dataloader(self):
        return DataLoader(self.mnist_predict, batch_size=32)

핵심 훅

  • prepare_data — 다운로드·토크나이즈처럼 한 번 처리하는 단계. Lightning은 이를 CPU의 단일 프로세스에서만 호출해 다중 프로세스(분산) 환경에서 데이터가 손상되는 것을 막아줘요. prepare_data에 상태 할당(self.x=y)은 권장하지 않아요(단일 프로세스에서 호출돼 다른 프로세스에 전달되지 않기 때문).
  • setup — 클래스 수, 어휘 구축, train/val/test split, 데이터셋 생성, transform 적용처럼 모든 GPU에서 수행해야 하는 작업용. stage 인자를 받아 trainer.{fit,validate,test,predict}별로 로직을 분리합니다. 모든 노드의 모든 프로세스에서 호출되므로 여기서 상태를 설정하는 걸 권장해요.
  • train_dataloader / val_dataloader / test_dataloader / predict_dataloader — 각각의 데이터로더 생성.
  • teardown — 상태 정리용. 모든 노드의 모든 프로세스에서 호출.

사용법

추천하는 사용 방식은:

dm = MNISTDataModule()
model = Model()
trainer.fit(model, datamodule=dm)
trainer.test(datamodule=dm)
trainer.validate(datamodule=dm)
trainer.predict(datamodule=dm)

모델을 만들기 위해 데이터셋 정보가 필요하면 prepare_datasetup을 직접 호출할 수도 있어요. Hyperparameter는 LightningModule처럼 self.save_hyperparameters()로 저장하고, 상태 저장이 필요하면 state_dict/load_state_dict를 정의해 체크포인트가 DataModule 상태를 추적·복원하게 할 수 있습니다.

더 알아보기