SFTTrainer — 지시 데이터로 모델 훈련하기

SFTTrainer — 지시 데이터로 모델 훈련하기

SFTTrainer 는 TRL이 제공하는 지도 파인튜닝 트레이너예요. 여러 줄이면 언어 모델을 지시 데이터로 훈련할 수 있어요.

기본 사용법

from trl import SFTTrainer
from datasets import load_dataset

trainer = SFTTrainer(
    model="Qwen/Qwen3-0.6B",
    train_dataset=load_dataset("trl-lib/Capybara", split="train"),
)
trainer.train()

데이터 형식

  • language modeling: {"text": ...} 형태.
  • prompt-completion: {"prompt": ..., "completion": ...} 형태.
  • conversational: {"messages": [...]} 형태. 채팅 템플릿을 자동 적용해요.

손실 계산

SFT는 토큰 단위 교차 엔트로피로, 모델이 이전 토큰을 보고 다음 토큰을 예측하도록 훈련해요. 패딩 토큰은 손실에서 마스킹하고, 라벨은 한 토큰만큼 오른쪽으로 밀어서(one-token shift) 맞춰요.

유용한 설정

  • packing=True: 여러 예제를 한 시퀀스에 채워 효율 향상.
  • assistant_only_loss=True: 어시스턴트 응답에만 손실.
  • completion_only_loss: prompt-completion에서 완성부에만 손실.
  • 메모리 절약을 위해 chunked_nll 손실(기본값) 사용.

더 알아보기