Lightning 15분 입문
Lightning 15분 입문
PyTorch Lightning의 전형적인 워크플로우는 일곱 개의 핵심 단계로 이루어져요. 이미지 오토인코더를 예로 들면서, 모델 정의부터 학습, 추론, 시각화까지 한 흐름으로 보여드릴게요. 순수 PyTorch로 짰다면 많았을 보일러플레이트가 Lightning에서는 어떻게 줄어드는지 바로 확인할 수 있어요.
1. LightningModule 정의하기
LightningModule은 nn.Module이 training_step 안에서 더 복잡한 방식으로 동작하게 해줘요. 학습 디바이스(무엇이든 하나)를 지정하면 자동으로 옮겨줘요.
import lightning as L
class LitAutoEncoder(L.LightningModule):
def __init__(self, encoder, decoder):
super().__init__()
self.encoder = encoder
self.decoder = decoder
def training_step(self, batch, batch_idx):
x, _ = batch
x = x.view(x.size(0), -1)
z = self.encoder(x)
x_hat = self.decoder(z)
loss = nn.functional.mse_loss(x_hat, x)
self.log("train_loss", loss) # 기본적으로 TensorBoard에 로깅
return loss
def configure_optimizers(self):
return optim.Adam(self.parameters(), lr=1e-3)
2. 데이터셋 정의하기
train/val/test/predict 분할에 어떤 iterable이든 사용할 수 있어요. DataLoader, numpy 등 모두 가능해요.
dataset = MNIST(os.getcwd(), download=True, transform=ToTensor())
train_loader = utils.data.DataLoader(dataset)
3. 학습하기
Trainer가 LightningModule과 데이터셋을 섞어서, 스케일을 위한 엔지니어링 복잡도를 추상화해요.
trainer = L.Trainer(limit_train_batches=100, max_epochs=1)
trainer.fit(model=autoencoder, train_dataloaders=train_loader)
4. 모델 사용하기
학습이 끝나면 체크포인트를 로드해 추론을 하거나, ONNX·TorchScript로 내보내 프로덕션에 올릴 수 있어요.
checkpoint = "./lightning_logs/version_0/checkpoints/epoch=0-step=100.ckpt"
autoencoder = LitAutoEncoder.load_from_checkpoint(checkpoint, encoder=encoder, decoder=decoder)
encoder = autoencoder.encoder
encoder.eval()
fake_image_batch = torch.rand(4, 28 * 28, device=autoencoder.device)
embeddings = encoder(fake_image_batch)
5. 학습 가속화하기
고급 학습 기능은 Trainer 인자로 켜요. 최신 기법들이 학습 루프에 자동으로 통합되기 때문에 코드 변경이 필요 없어요.
# 4 GPU 학습
trainer = L.Trainer(devices=4, accelerator="gpu")
# DeepSpeed/FSDP로 1TB+ 파라미터 모델 학습
trainer = L.Trainer(devices=4, accelerator="gpu",
strategy="deepspeed_stage_2", precision=16)
# 반복 실험에 유용한 플래그
trainer = L.Trainer(max_epochs=10, min_epochs=5, overfit_batches=1)
유연성 극대화하기
Lightning의 핵심 원칙은 PyTorch를 결코 숨기지 않으면서 최대한의 유연성을 주는 거예요. 학습 루프 어디에든 훅을 통해 커스텀 코드를 주입할 수 있고, 반복되는 코드는 콜백으로 묶어 켜고 끌 수 있어요.
class LitAutoEncoder(L.LightningModule):
def backward(self, loss):
loss.backward()
trainer = Trainer(callbacks=[AWSCheckpoints()])
최첨단 연구처럼 최적화나 학습 루프를 완전히 직접 제어하고 싶다면, 원시 PyTorch 루프를 쓰는 것도 허용돼요.