Determined PyTorch API — PyTorchTrial로 모델 훈련하기

Determined PyTorch API — PyTorchTrial로 모델 훈련하기

Determined에서 PyTorch 모델을 학습하려면 PyTorchTrial을 상속하는 trial 클래스를 구현하고, 그것을 experiment 설정의 entrypoint로 지정해야 해요. trial 클래스 안에서 학습 절차의 각 구성 요소를 담당하는 함수들을 오버라이드하면 됩니다.

출처: https://docs.determined.ai/latest/model-dev-guide/api-guides/apis-howto/api-pytorch-ug.html

trial 클래스 구현

determined.pytorch.PyTorchTrial을 상속하고, 학습에 쓰일 구성 요소를 오버라이드해요. 기본적으로 Distributed Backend는 Horovod를 쓰고, torch.distributed와 DistributedDataParallel을 선택할 수도 있어요.

데이터 로딩

데이터를 로딩하는 건 build_training_data_loader()build_validation_data_loader()라는 두 함수로 정의해요. 각 함수는 determined.pytorch.DataLoader 인스턴스를 반환해야 해요.

객체 초기화

학습에 쓸 객체(model, optimizer, learning rate scheduler, 커스텀 loss·metric 함수)는 PyTorchTrial의 생성자 __init__에서 제공된 context를 사용해 초기화해야 해요.

class MyTrial(pytorch.PyTorchTrial):
    def __init__(self, context):
        self.context = context
        # model, optimizer, scheduler, loss 등을 context로 초기화

주의: 객체를 반드시 wrap하세요

PyTorchTrial이 다루는 model·optimizer·scheduler는 Determined가 제공하는 방식으로 **반드시 감싸(wrap)**야 해요. 감싸지 않으면, 일시정지 후 재개된 trial의 메트릭이 그렇지 않은 trial과 크게 달라질 수 있어요. 이유는 체크포인트에서 모델 상태가 정확히 복원되지 않을 수 있기 때문이에요. PyTorch API를 올바르게 쓰지 않으면 이런 상태 복원 문제가 생겨요.

체크포인팅

체크포인트에는 모델 정의, 실험 설정, 네트워크 아키텍처, 가중치·하이퍼파라미터 값이 포함돼요. stateful optimizer 사용 시 optimizer 상태(학습률)도 함께 저장돼요.

더 알아보기

handling 데이터 유형·디버깅 팁은 공식 PyTorch API 문서에서 확인할 수 있어요.