LightningModule — 훈련 코드를 6개 섹션으로 조직화
LightningModule — 훈련 코드를 6개 섹션으로 조직화
LightningModule 은 PyTorch 코드를 6개 섹션으로 조직화해줘요. 코드가 추상화되는 게 아니라 정리(organized) 되는 것이라는 점이 핵심이에요. LightningModule 에 없는 나머지 코드는 Trainer 가 자동으로 처리합니다.
출처: https://lightning.ai/docs/pytorch/stable/common/lightning_module.html
6개 섹션
- Initialization (
__init__andsetup) - 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 |
LightningModule 은 torch.nn.Module 이지만 기능이 덧붙여진 것이므로 그렇게 사용하면 돼요.
net = Net.load_from_checkpoint(PATH)
net.freeze()
out = net(x)
더 알아보기
- 15분 입문은 Introduction 참고
- Trainer는 Trainer 참고
- 데이터 관리는 DataModule 참고