Lightning 체크포인트 — 훈련 상태 저장과 복원
Lightning 체크포인트 — 훈련 상태 저장과 복원
PyTorch Lightning은 훈련 중 모델의 전체 상태를 자동으로 저장하는 체크포인트 기능을 제공해요. 하이퍼파라미터, 옵티마이저 상태, 에포크, 배치 인덱스 등을 한 번에 기록해서 중단된 훈련을 이어가거나 배포할 수 있습니다.
출처: https://lightning.ai/docs/pytorch/stable/common/checkpointing.html
기본 동작
기본 설정에서는 매 훈련 에포크가 끝날 때마다 epoch=0-step=...ckpt 형태의 체크포인트가 현재 작업 디렉터리에 저장됩니다.
trainer = Trainer()
trainer.fit(model, datamodule=dm)
# 저장 예: epoch=0-step=1161.ckpt
체크포인트에 담기는 내용
기본 체크포인트에는 다음이 저장됩니다.
- 16비트 정밀도(AMP) 스케일러 상태
- 현재 에포크
- 글로벌 스텝
LightningModule의 state_dict- 모든 옵티마이저의 state_dict
- 모든 러너 및 옵티마이저 스케줄러의 state_dict
save_hyperparameters()로 저장된 하이퍼파라미터
모델 로드
먼저 모델을 생성하고 load_from_checkpoint로 불러옵니다. 하이퍼파라미터는 체크포인트 안에 있으므로 코드에서 save_hyperparameters()를 함께 쓰는 게 권장돼요.
model = MyLightningModule.load_from_checkpoint("/path/to/checkpoint.ckpt")
trainer.fit(model)
하이퍼파라미터를 명시적으로 덮어쓰려면 인자로 넘기면 됩니다.
model = MyLightningModule.load_from_checkpoint(ckpt_path, learning_rate=1e-3)
더 알아보기
- 모델 정의는 LightningModule 참고
- 훈련 루프 제어는 Trainer 참고
- 데이터는 DataModule 참고