전체 모델 대신 언제나 state_dict를 저장해
모델을 저장하는 방법은 두 가지야:
- state_dict(권장): 매개변수와 버퍼만 Python OrderedDict로 저장해. 불러올 때는 모델 클래스를 먼저 만들고
load_state_dict를 호출해. 코드 변경, 프레임워크 버전 차이, 클래스 이름 변경에도 비교적 견고해. - torch.save(model): 전체 Python 객체를 피클로 저장해. 불러올 때 정확히 같은 클래스 정의가 필요해서 이름 변경이나 리팩터링에 쉽게 깨져.
항상 state_dict 방식을 사용해.
학습 체크포인트에 넣을 것
추론만 한다면 모델 state_dict로 충분해. 학습을 재개하려면 다음도 필요해:
- 모델 state_dict
- 옵티마이저 state_dict. Adam의 누적 모멘트는 간단히 다시 만들 수 없어.
- 학습률 조정기 state_dict. 현재 단계와 last_lr 같은 상태가 들어 있어.
- 현재 에포크와 단계 번호
- 지금까지 가장 좋은 검증 평가지표. 조기 종료를 이어 가는 데 필요해.
- 정확한 재현성이 중요하다면 난수 생성기 상태
weights_only=True
불러올 때 weights_only=True를 명시해. PyTorch 2.6부터는 pickle_module을 따로 넘기지 않으면 이 값이 기본이지만, 의도를 분명히 하고 이전 버전과 호환하려면 명시하는 편이 좋아. 임의의 Python 객체를 역직렬화하지 않고 텐서 데이터만 읽으므로, 악성 .pt 파일이 로드 과정에서 코드를 실행하는 실제 공격을 막는 데 도움이 돼.
조기 종료
검증 손실이 개선될 때마다 '지금까지 가장 좋은' 체크포인트를 저장해. 개선이 없는 에포크 수를 세다가 N에 도달하면 학습을 멈춰. 몇 줄의 상태만 관리하면 불필요한 계산을 막을 수 있어.