TRL

TRL

TRL은 기초 모델(foundation model)용 사후 학습 프레임워크예요. SFT, GRPO, DPO 같은 방법을 지원하며, Transformers의 Trainer 클래스 위에 구축돼요.

출처: 문서

본문

TRL은 기초 모델용 사후 학습 프레임워크예요. SFT, GRPO, DPO 같은 방법을 포함해요. 각 방법은 Trainer 클래스 위에 구축된 전용 트레이너를 가지며, 단일 GPU에서 다중 노드 클러스터까지 확장돼요.

from datasets import load_dataset
from trl import GRPOTrainer
from trl.rewards import accuracy_reward

dataset = load_dataset("trl-lib/DeepMath-103K", split="train")

trainer = GRPOTrainer(
    model="Qwen/Qwen2-0.5B-Instruct",
    reward_funcs=accuracy_reward,
    train_dataset=dataset,
)
trainer.train()

Transformers 통합

TRL은 Transformers API를 확장하고 메서드별 설정을 추가해요.

  • TRL 트레이너는 Trainer 위에 구축돼요. GRPOTrainer 같은 메서드별 트레이너는 생성, 보상 스코어링, 손실 계산을 추가해요. 설정 클래스는 TrainingArguments를 확장해 메서드별 필드를 더해요.

  • 모델 로딩은 AutoConfig.from_pretrained()를 사용하고, 그다음 config에서 해당 클래스의 from_pretrained로 모델 클래스를 인스턴스화해요.

리소스 (Resources)

더 알아보기 (Learn more)