LightningDataModule — 데이터 파이프라인 캡슐화
LightningDataModule — 데이터 파이프라인 캡슐화
데이터모듈(datamodule)은 데이터를 처리하는 데 필요한 모든 단계를 담는 공유·재사용 가능한 클래스예요. PyTorch의 데이터 처리 다섯 단계를 하나로 캡슐화합니다.
출처: https://lightning.ai/docs/pytorch/stable/data/datamodule.html
- Download / tokenize / process.
- Clean and (maybe) save to disk.
- Load inside
Dataset. - Apply transforms (rotate, tokenize, etc…).
- 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_data와 setup을 직접 호출할 수도 있어요. Hyperparameter는 LightningModule처럼 self.save_hyperparameters()로 저장하고, 상태 저장이 필요하면 state_dict/load_state_dict를 정의해 체크포인트가 DataModule 상태를 추적·복원하게 할 수 있습니다.
더 알아보기
- LightningModule는 LightningModule 참고
- Trainer는 Trainer 참고
- 15분 입문은 Introduction 참고