Views
No views yet

SFTTrainer, GRPOTrainer, DPOTrainer, RewardTrainer and more.pip:pip install trlpip install git+https://github.com/huggingface/trl.gitgit clone https://github.com/huggingface/trl.gitSFTTrainerSFTTrainer:1from trl import SFTTrainer
2from datasets import load_dataset
3
4dataset = load_dataset("trl-lib/Capybara", split="train")
5
6trainer = SFTTrainer(
7 model="Qwen/Qwen2.5-0.5B",
8 train_dataset=dataset,
9)
10trainer.train()GRPOTrainerGRPOTrainer implements the Group Relative Policy Optimization (GRPO) algorithm that is more memory-efficient than PPO and was used to train Deepseek AI's R1.1from datasets import load_dataset
2from trl import GRPOTrainer
3from trl.rewards import accuracy_reward
4
5dataset = load_dataset("trl-lib/DeepMath-103K", split="train")
6
7trainer = GRPOTrainer(
8 model="Qwen/Qwen2.5-0.5B-Instruct",
9 reward_funcs=accuracy_reward,
10 train_dataset=dataset,
11)
12trainer.train()[!NOTE] For reasoning models, use thereasoning_accuracy_reward()function for better results.
DPOTrainerDPOTrainer implements the popular Direct Preference Optimization (DPO) algorithm that was used to post-train Llama 3 and many other models. Here is a basic example of how to use the DPOTrainer:1from datasets import load_dataset
2from trl import DPOTrainer
3
4dataset = load_dataset("trl-lib/ultrafeedback_binarized", split="train")
5
6trainer = DPOTrainer(
7 model="Qwen3/Qwen-0.6B",
8 train_dataset=dataset,
9)
10trainer.train()RewardTrainerRewardTrainer:1from trl import RewardTrainer
2from datasets import load_dataset
3
4dataset = load_dataset("trl-lib/ultrafeedback_binarized", split="train")
5
6trainer = RewardTrainer(
7 model="Qwen/Qwen2.5-0.5B-Instruct",
8 train_dataset=dataset,
9)
10trainer.train()1trl sft --model_name_or_path Qwen/Qwen2.5-0.5B \
2 --dataset_name trl-lib/Capybara \
3 --output_dir Qwen2.5-0.5B-SFT1trl dpo --model_name_or_path Qwen/Qwen2.5-0.5B-Instruct \
2 --dataset_name argilla/Capybara-Preferences \
3 --output_dir Qwen2.5-0.5B-DPO --help for more details.trl or customize it to your needs make sure to read the contribution guide and make sure you make a dev install:1git clone https://github.com/huggingface/trl.git
2cd trl/
3pip install -e .[dev]trl.experimental for unstable / fast-evolving features. Anything there may change or be removed in any release without notice.from trl.experimental.new_trainer import NewTrainer1@software{vonwerra2020trl,
2 title = {{TRL: Transformers Reinforcement Learning}},
3 author = {von Werra, Leandro and Belkada, Younes and Tunstall, Lewis and Beeching, Edward and Thrush, Tristan and Lambert, Nathan and Huang, Shengyi and Rasul, Kashif and Gallouédec, Quentin},
4 license = {Apache-2.0},
5 url = {https://github.com/huggingface/trl},
6 year = {2020}
7}