Trainer로 학습 자동화하기
Trainer로 학습 자동화하기
PyTorch 코드를 LightningModule로 정리하고 나면, Trainer가 나머지를 전부 자동화해요. 연구·엔지니어링 보일러플레이트를 크게 줄이면서, 필요한 부분은 여전히 직접 제어할 수 있게 해주는 것이 Trainer의 역할이에요.
Trainer가 내부에서 하는 일
겉보기엔 "그냥 학습만 돌려주는 도구"처럼 보이지만, 내부에서는 루프의 세부를 알아서 처리해요. 대표적인 예가 이렇습니다.
- gradient 활성화/비활성화 자동 처리
- 학습·검증·테스트 데이터로더 실행
- 콜백을 적절한 시점에 호출
- 배치와 계산을 올바른 디바이스에 배치
학습 루프만 떼어 본 유사 코드는 이런 모양이에요.
# enable grads
torch.set_grad_enabled(True)
losses = []
for batch in train_dataloader:
on_train_batch_start() # 훅 호출
loss = training_step(batch) # 학습 스텝
optimizer.zero_grad() # 그래디언트 초기화
loss.backward() # 역전파
optimizer.step() # 파라미터 갱신
losses.append(loss)
기본 사용법
Trainer는 아주 간단하게 쓸 수 있어요.
model = MyLightningModule()
trainer = Trainer()
trainer.fit(model, train_dataloader, val_dataloader)
Python 스크립트에서 사용하기
스크립트에서는 main 함수 안에서 Trainer를 호출하는 걸 권장해요. accelerator와 devices를 명령줄 인자로 받을 수 있죠.
from argparse import ArgumentParser
def main(hparams):
model = LightningModule()
trainer = Trainer(accelerator=hparams.accelerator, devices=hparams.devices)
trainer.fit(model)
if __name__ == "__main__":
parser = ArgumentParser()
parser.add_argument("--accelerator", default=None)
parser.add_argument("--devices", default=None)
args = parser.parse_args()
main(args)
실행 시 이렇게 인자를 넘기면 됩니다.
python main.py --accelerator 'gpu' --devices 2
검증과 테스트
학습 루프 밖에서 별도의 검증 에폭을 수행하고 싶다면 validate를 써요. 초기화 직후나 학습 완료 후에 새 메트릭을 수집할 때 유용해요.
trainer.validate(model=model, dataloaders=val_dataloaders)
학습이 끝나고 논문을 내거나 프로덕션에 올리기 직전엔 테스트셋으로 확인해요.
trainer.test(dataloaders=test_dataloaders)
재현성 확보하기
완전히 재현 가능한 학습을 원한다면 시드를 설정하고 deterministic 플래그를 켜요.
from lightning.pytorch import Trainer, seed_everything
seed_everything(42, workers=True) # numpy, torch, python.random 시드 설정
model = Model()
trainer = Trainer(deterministic=True)
workers=True로 설정하면 데이터로더 워커와 프로세스마다 torch, numpy, 표준 random에 고유한 시드를 부여해요. 데이터 증강이 워커마다 반복되지 않게 보장해 주죠.
주요 accelerator 설정
accelerator 인자는 하드웨어 유형을 고를 수 있게 해줘요. "cpu", "gpu", "tpu", "hpu", "auto"를 지원하고, 커스텀 accelerator 인스턴스도 넘길 수 있어요.
# CPU accelerator
trainer = Trainer(accelerator="cpu")
# GPU accelerator, GPU 2개
trainer = Trainer(devices=2, accelerator="gpu")
# TPU accelerator, 8개 코어
trainer = Trainer(devices=8, accelerator="tpu")
# DistributedDataParallel 전략
trainer = Trainer(devices=4, accelerator="gpu", strategy="ddp")
"auto"는 현재 머신을 확인하고 적절한 Accelerator를 자동으로 골라 줘요.
# 머신에 GPU가 있으면 GPU accelerator를 자동 사용
trainer = Trainer(devices=2, accelerator="auto")