Trainer 메서드 서브클래싱하기

Trainer 메서드 서브클래싱하기 (Subclassing Trainer methods)

Trainer 메서드를 서브클래싱하면 전체 루프를 다시 작성하지 않고도 학습 동작을 바꿀 수 있어요. 서브클래싱은 forward pass나 손실 계산 같은 학습 루프를 수정해요.

출처: 문서

본문

서브클래싱 전에, Trainer가 무엇을 계산하는지 바꿀 필요가 있는지, 아니면 언제 동작하고 동작할지 여부를 바꿔야 하는지 먼저 생각해 보세요. 시점과 조건 로직에는 Callback을 대신 사용해요. 콜백은 일이 언제 일어나는지(로깅, 평가, 조기 중지)를 제어하고, 서브클래싱은 무엇이 일어나는지(손실 계산, 데이터 로딩, 최적화)를 바꿔요.

[!NOTE] 서브클래싱할 수 있는 메서드의 전체 목록은 Trainer API 문서를 참고하세요. _save_checkpoint나 _evaluate처럼 _로 시작하는 프라이빗 메서드도 재정의할 수 있지만, 이런 것들은 예고 없이 바뀔 수 있어요.

get_train_dataloader

표준 get_train_dataloader() 메서드는 배치 하나를 로드해 학습하고, 버린 뒤 다음 배치를 로드해요.

def get_train_dataloader(self):
    return self._get_dataloader(
        batch_size=self._train_batch_size,
        ...
)

GRPO는 온라인 강화 학습 알고리즘으로, 학습 전에 완성(completion)을 생성해요. 매 스텝 완성을 생성하는 것은 자동회귀(autoregressive) 방식이라 비용이 커요. 512-토큰 완성은 학습 스텝의 forward pass 한 번에 비해 약 512번의 순차 forward pass가 필요해요. GRPOTrainer는 여러 스텝에 걸쳐 생성을 배치로 묶기 위해 get_train_dataloader()를 서브클래싱해요.

trl.GRPOTrainer.get_train_dataloader는 배치 크기에 steps_per_generation 인자를 곱해서 여러 학습 스텝의 생성 프롬프트 배치를 한 번에 로드해요. train_batch_size=4이고 steps_per_generation=8이면 dataloader는 크기 32의 배치를 생성해서 생성 비용을 8배 줄여요.

def get_train_dataloader(self):
    dataloader_params = {
        "batch_size": self._train_batch_size * self.args.steps_per_generation, # this is the only change
        ...
    }

compute_loss

compute_loss()는 모델이 계산한 교차 엔트로피(cross-entropy) 손실을 반환해요.

def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
    ...
    outputs = model(**inputs)
    ...
    loss = outputs["loss"] # get loss from model

    return (loss, outputs) if return_outputs else loss

DPO는 참조 모델(reference model)과 비교해 정책 모델이 선택된 응답을 거부된 응답보다 얼마나 선호하는지 측정해요. DPOTrainer는 손실 계산이 표준 교차 엔트로피와 여러 면에서 다르기 때문에 compute_loss()를 서브클래싱해요:

  • 모델은 라벨을 절대 보지 않아요. DPO가 log-prob을 계산할 수 있도록 logits만 반환해요.
  • 선택된 응답과 거부된 응답이 연결(concatenate)돼요.
  • 참조 모델이 자신의 log-prob을 계산해요.
  • 손실은 π_chosen, π_rejected, π_ref_chosen, π_ref_rejected의 함수예요.

위 어느 것도 표준 Trainer.compute_loss() 메서드에 들어맞지 않아요.

def compute_loss(
    self,
    model: PreTrainedModel | nn.Module,
    inputs: dict[str, torch.Tensor | Any],
    return_outputs=False,
    num_items_in_batch=None,
) -> torch.Tensor | tuple[torch.Tensor, dict[str, float]]:
    ...
    outputs = model(**inputs)
    logits = outputs.logits
    logps = get_logps(logits, inputs)
    chosen_logps, rejected_logps = logps.chunk(2, dim=0)  # batch is [chosen, rejected]
    ref_logits = self.ref_model(**inputs).logits
    ref_logps = get_logps(ref_logits, inputs)
    ref_chosen_logps, ref_rejected_logps = ref_logps.chunk(2, dim=0)  # batch is [chosen, rejected]
    chosen_scores = chosen_logps - ref_chosen_logps
    rejected_scores = rejected_logps - ref_rejected_logps
    per_sequence_loss = -F.logsigmoid(self.beta * chosen_scores - rejected_scores)
    loss = per_sequence_loss.mean()
    return (loss, outputs) if return_outputs else loss

다음 단계

  • 더 많은 실제 예시는 TRL에서 GRPOTrainer와 DPOTrainer가 Trainer를 어떻게 확장하는지, 또는 Axolotl이 그 위에 커스텀 트레이너를 어떻게 만드는지 확인해 보세요.
  • 학습 스텝 끝에 메트릭 로깅처럼 학습 이벤트 중에 일어나는 것을 커스터마이즈하기만 하면 된다면 Callbacks 가이드를 확인해 보세요.