LightningModule — 훈련 코드를 6개 섹션으로 조직화

LightningModule — 훈련 코드를 6개 섹션으로 조직화

LightningModule 은 PyTorch 코드를 6개 섹션으로 조직화해줘요. 코드가 추상화되는 게 아니라 정리(organized) 되는 것이라는 점이 핵심이에요. LightningModule 에 없는 나머지 코드는 Trainer 가 자동으로 처리합니다.

출처: https://lightning.ai/docs/pytorch/stable/common/lightning_module.html

6개 섹션

  • Initialization (__init__ and setup)
  • Train Loop (training_step)
  • Validation Loop (validation_step)
  • Test Loop (test_step)
  • Prediction Loop (predict_step)
  • Optimizers and LR Schedulers (configure_optimizers)

devicesss를 신경 쓰지 않아요

Lightning을 쓰면 .cuda().to(device) 호출이 필요 없어요. Lightning이 알아서 해줍니다.

# don't do in Lightning
x = torch.Tensor(2, 3)
x = x.cuda()
x = x.to(device)

# do this instead
x = x  # leave it alone!

분산 전략에서 돌릴 때도 Lightning이 기본으로 분산 샘플러를 처리해줘요.

# Don't do in Lightning...
data = MNIST(...)
sampler = DistributedSampler(data)
DataLoader(data, sampler=sampler)

# do this instead
data = MNIST(...)
DataLoader(data)

스타터 예시

필수 메서드는 이 정도예요.

import lightning as L
import torch

from lightning.pytorch.demos import Transformer

class LightningTransformer(L.LightningModule):
    def __init__(self, vocab_size):
        super().__init__()
        self.model = Transformer(vocab_size=vocab_size)

    def forward(self, inputs, target):
        return self.model(inputs, target)

    def training_step(self, batch, batch_idx):
        inputs, target = batch
        output = self(inputs, target)
        loss = torch.nn.functional.nll_loss(output, target.view(-1))
        return loss

    def configure_optimizers(self):
        return torch.optim.SGD(self.model.parameters(), lr=0.1)

훈련은 이렇게 해요.

from lightning.pytorch.demos import WikiText2
from torch.utils.data import DataLoader

dataset = WikiText2()
dataloader = DataLoader(dataset)
model = LightningTransformer(vocab_size=dataset.vocab_size)

trainer = L.Trainer(fast_dev_run=100)
trainer.fit(model=model, train_dataloaders=dataloader)

핵심 메서드 요약

Name Description
__init__ and setup Define initialization here
forward To run data through your model only (separate from training_step)
training_step the complete training step
validation_step the complete validation step
test_step the complete test step
predict_step the complete prediction step

LightningModuletorch.nn.Module 이지만 기능이 덧붙여진 것이므로 그렇게 사용하면 돼요.

net = Net.load_from_checkpoint(PATH)
net.freeze()
out = net(x)

더 알아보기