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)을 지정한다.