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)
- TRL 문서
- Fine Tuning with TRL 강연
더 알아보기 (Learn more)
- Unsloth 문서에서 TRL 트레이너를 확장하는 방법을 살펴보세요.
- Transformers 커스텀 모델 문서