torch.save — 객체 직렬화 저장 API

torch.save — 객체 직렬화 저장 API

torch.save(obj, f) 는 객체를 직렬화해 디스크에 저장하는 함수예요. Python의 pickle을 이용해 모델, 텐서, 딕셔너리 등 대부분 객체를 저장할 수 있어요.

특징

  • 모델의 경우 통상 model.state_dict() 를 저장해, 나중에 같은 구조로 재현한 모델에 로드한다.
  • 여러 객체(모델 A/B, 옵티마이저들, 에폭 등)는 딕셔너리로 묶어 한 번에 저장한다.
  • 확장자는 관례이며 .pt/.pth/.tar 를 쓴다.
  • 파일 경로(문자열) 또는 파일 객체를 받는다. 경로로 주면 torch가 압축·포맷을 처리한다.

권장

  • 저장 대상은 가중치 상태(state_dict) 중심으로 하되, 재개를 위해 옵티마이저 상태와 메타 정보도 함께 담는다.
  • 분산 학습에서는 각 랭크의 상태를 잘 맞춰 저장하고, 로드 시 디바이스 매핑(map_location)을 지정한다.

더 알아보기