[PyTorch] 모델 가져오기

Jeonghyun·2022년 9월 29일

PyTorch

목록 보기
4/6

model.save() : 학습의 결과를 저장하기 위한 함수

  • 모델 형태(architecture)와 parameter를 저장
  • 모델 학습 중간 과정의 저장을 통해 최선의 결과 모델을 선택
  • 만들어진 모델을 다른 연구자와 공유하여 학습 재연성 향상
print("Model's state_dict:")
for param_tensorin model.state_dict():  # state_dict : 모델의 파라미터 표시
	print(param_tensor,"\t",model.state_dict()[param_tensor].size())

torch.save(model.state_dict(), os.path.join(MODEL_PATH, "model.pt"))  # 모델의 파리미터 저장

new_model = TheModelClass()
# 같은 모델의 형태에서 파라미터만 load
new_model.load_state_dict(torch.load(os.path.join(MODEL_PATH, "model.pt")))

# 모델의 architecture와 함께 저장
torch.save(model, os.path.join(MODEL_PATH, "model.pt"))
model = torch.load(os.path.join(MODEL_PATH,"model.pt"))

checkpoints

  • 학습의 중간 결과를 저장하여 최선의 결과 선택
  • earlystopping 기법 사용 시 이전 학습의 겨로가물 저장
  • 일반적으로 epoch, loss, metric을 함께 저장1
# 모델의 정보를 epoch과 함께 저장
torch.save({
			'epoch': e,
			'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(), 
            'loss': epoch_loss,
            },
f"saved/checkpoint_model_{e}_{epoch_loss/len(dataloader)}_{epoch_acc/len(dataloader)}.pt")

checkpoint = torch.load(PATH)
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
epoch = checkpoint['epoch']
loss = checkpoint['loss']
import warnings
warnings.filterwarnings("ignore") # user warning이 많을 때 무시할 수 있음

Transfer learning

  • 다른 데이터셋으로 만든 모델을 현재 데이터에 적용
  • 일반적으로 대용량 데이터셋으로 만들어진 모델의 성능 향상
  • backbone architecture가 잘 학습된 모델에서 일부분만 변경하여 학습 수행

Freezing : pretrained model을 활용시 모델의 일부분을 frozen시킴(파라미터 변경 x)




출처 - 부스트캠프 AI tech 교육자료


[부스트캠프 AI Tech] Week 2 - Day 4

0개의 댓글