PyTorch 튜토리얼 — 모델 저장·불러오기와 체크포인트

PyTorch 튜토리얼 — 모델 저장·불러오기와 체크포인트

PyTorch는 모델·텐서·딕셔너리 저장을 세 핵심 함수로 다뤄요: torch.save, torch.load, load_state_dict. 이 튜토리얼은 다양한 사용 사례별 저장/로드를 안내해요.

state_dict 란

nn.Module의 학습 가능한 파라미터(가중치·편향)와 버퍼를 레이어 → 텐서로 매핑한 딕셔너리야. 옵티마이저도 state_dict를 가져 상태·하이퍼파라미터를 담아요.

일반 체크포인트 저장/로드

torch.save({'epoch': epoch, 'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(), 'loss': loss}, PATH)
...
checkpoint = torch.load(PATH, weights_only=True)
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
  • 재개(resume)를 위해 모델 외에 옵티마이저 상태·에폭·손실 도 함께 저장해요(모델만 보다 2~3배 큼, 보통 .tar 확장자).
  • 불러온 뒤 추론이면 model.eval(), 재학습이면 model.train() 을 호출해야 dropout/BN 동작이 일관돼요.
  • 일부 파라미터만 옮기는 전이 학습엔 load_state_dict(..., strict=False) 로 미일치 키를 무시해요.

더 알아보기