数据截至 (上游 commit 67dfbe211a07)
01 · Trainer 家族与公共骨架
这一章讲什么: TRL 十几个 Trainer 共同站在什么地基上。读完你会知道
_BaseTrainer/_BaseConfig到底给了什么(答案:很少,这是故意的)、稳定版和experimental/怎么分层、以及每个 Trainer 构造函数里那段几乎一模一样的「配方」长什么样。
1. 它要解决的小问题
后 训练方法有一整个动物园:SFT、DPO、KTO、GRPO、RLOO、蒸馏……每种方法要复用一大堆与方法无关的东西:
- 分布式训练(DDP/DeepSpeed/FSDP)、混合精度、梯度累积、checkpoint、日志上报;
- 模型加载(给个 Hub id 就能训)、量化、LoRA 包装;
- 命令行参数解析、模型卡生成。
TRL 的选择:自己不造这层。 训练循环和分布式全部直接用 transformers 的 Trainer;TRL 只在它上面加一层极薄的基座,再把「怎么加载模型、怎么处理数据集、怎么算 loss」按统一配方写进每个子类。
2. 思路:继承链上的三层
transformers.Trainer ← 训练循环、accelerate 集成、日志、checkpoint
▲
trl._BaseTrainer ← 只加两件事:匿名遥测、模型卡生成
▲
SFTTrainer / DPOTrainer / … ← 各自的数据准备 + compute_loss
配置侧是对称的三层:
transformers.TrainingArguments
▲
trl._BaseConfig ← 只覆盖几个 TRL 偏好的默认值
▲
SFTConfig / DPOConfig / … ← 各自方法特有的超参
关键认识:基座薄是设计,不是偷懒。 训练基础设施越薄地压在 transformers 上,TRL 就越能免费吃到上游的 bug 修复和新硬件支持;方法之间的可比性也越强——两个 Trainer 的差异只剩数据准备和 loss。
3. 真实实现:基座里到底有什么
3.1 _BaseTrainer:遥测 + 模型卡
_BaseTrainer(trl/trainer/base_trainer.py:69)继承 Trainer,构造时只多做一件事——发一条匿名使用遥测(_send_telemetry,trl/trainer/base_trainer.py:79)。这条遥测本身是个不错的「防御性编程」样本:
- 只在主进程发,避免按 world size 放大;
- world size 分桶上报(
"1"/"2-8"/"9-64"/"65+"),不给精确集群规模; - Trainer 类名有白名单
_TELEMETRY_TRAINERS(trl/trainer/base_trainer.py:35),用户的自定义子类一律上报为"other",不泄露私有名字; - 数据集只报来源类型(hub / files / memory / other),不报内容。
另一件事是 create_model_card(trl/trainer/base_trainer.py:141):训练结束时用 trl/templates/ 下的模板(如 lm_model_card.md)生成 README 模型卡,自动带上论文引用——每个 Trainer 用类属性声明自己的论文,例如 DPOTrainer 的 _name = "DPO"、_tag_names = ["trl", "dpo"](trl/trainer/dpo_trainer.py:500-501)和 _paper 字典。
3.2 _BaseConfig:只改默认值
_BaseConfig(trl/trainer/base_config.py:21)继承 TrainingArguments,做的事用一只手数得完:
| 覆盖项 | TRL 默认值 | 为什么 |
|---|---|---|
logging_steps | 10 | 后训练步数少,默认 500 太稀 |
gradient_checkpointing | True | 后训练显存紧,默认开 |
bf16 | 未设 fp16 时默认 True(__post_init__,trl/trainer/base_config.py:104-107) | 后训练几乎没人用 fp32/fp16 |
lr_scheduler_kwargs 等 | 修类型标注 | 绕过 transformers 旧版本的 argparse bug(注释见 trl/trainer/base_config.py:75-88) |
每个方法自己的 Config 再往下加专有参数——比如 SFTConfig 加 packing、assistant_only_loss,GRPOConfig 加 num_generations、beta。因为它们都是 TrainingArguments,所以 transformers.HfArgumentParser 能直接把 Config 类变成命令行参数,TRL 的 CLI(trl/cli/main.py:32)和 examples 里的脚本都靠这一点。
4. 稳定层 vs experimental 层
v1 的组织方式(MIGRATION.md:1-16):
| 层 | 位置 | 内容 |
|---|---|---|
| 稳定 API | trl/trainer/ | SFT、DPO、GRPO、KTO、Reward、RLOO、Distillation 七个 Trainer |
| 实验 API | trl/experimental/ | PPO、ORPO、CPO、BCO、Online DPO、XPO、NashMD、GKD、异步 GRPO 等 20+ |
从 trl.experimental 导入会立刻收 到一条「API 不稳定,随时可能改或删」的警告(trl/experimental/__init__.py:32)。
怎么理解这个分层: 稳定层是「论文已被广泛复现、接口承诺向后兼容」的方法;experimental 是「论文较新、或实现还在快速迭代」的方法。著名的 PPO(trl/experimental/ppo/ppo_trainer.py:297)在 v1 被挪进 experimental——社区重心已从经典 RLHF 转向 GRPO 系,PPO 实现进入维护模式。
5. 每个 Trainer 构造函数的通用配方
读任何一个 XxxTrainer.__init__,都会看到同一段流程。以 SFTTrainer.__init__(trl/trainer/sft_trainer.py:916)为例:
# 示意,非源码:每个 TRL Trainer 的构造配方
def __init__(self, model, args=None, train_dataset=None, processing_class=None, ...):
args = args or XxxConfig(f"{model_name}-XXX") # ① 没给配置就用默认
if isinstance(model, str): # ② 给了模型 id 就现场加载
model = create_model_from_path(model, **args.model_init_kwargs)
if processing_class is None: # ③ 没给 tokenizer 就自动加载
processing_class = AutoProcessor.from_pretrained(model.config)
# ④ 校验数据集类型与格式(IterableDataset 要关掉 dispatch_batches 等)
# ⑤ 准备数据集:tokenize、加掩码、(可选)packing —— 各方法不同
# ⑥ 选 collator —— 各方法不同
# ⑦ 若给了 peft_config,包一层 LoRA
if peft_config is not None:
model = get_peft_model(model, peft_config) # sft_trainer.py:1124
super().__init__(model, args, train_dataset, data_collator, ...) # ⑧ 交给 transformers
逐步对应真实代码:
- 补默认配置:没传
args时用模型名造一个(trl/trainer/sft_trainer.py:941-944);传了普通TrainingArguments也会转成对应 Config。 - 模型 id → 模型对象:
create_model_from_path(trl/trainer/utils.py:1145)先读 config 推断架构类(含 remote-code 模型的auto_map回退),再from_pretrained;分布式训练时强制device_map=None(trl/trainer/sft_trainer.py:980-981)。 - 自动加载 processing_class:
AutoProcessor.from_pretrained(trl/trainer/sft_trainer.py:998-1001)——用 AutoProcessor 而不是 AutoTokenizer,是为了同时覆盖纯文本模型和 VLM。 - 数据集校验:比如 SFT 拒绝
with_transform过的数据集,因为惰性 transform 会被map()烘焙进缓存(trl/trainer/sft_trainer.py:1434-1441)。 _prepare_dataset:方法特定,第 2、3 章细讲。- collator:方法特定(SFT 的语言建模 collator、DPO 的偏好 collator……)。
- PEFT 包装:
get_peft_model(trl/trainer/sft_trainer.py:1124);独立的ModelConfig(trl/trainer/model_config.py:19)还提供了从命令行描述 LoRA/量化的一整套字段,get_peft_config/get_quantization_config(trl/trainer/utils.py:255、:236)把它翻译成 peft/bitsandbytes 的对象。 - 交给父类:之后的
trainer.train()就是 transformers 的训练循环,TRX 的 Trainer 只覆写compute_loss等钩子。
重点看: ①②③⑦⑧ 在所有 Trainer 里逐字雷同;方法之间的全部差异集中在 ⑤⑥ 和 compute_loss。这就是为什么这个库适合对比阅读。
6. 关键细节与坑
model_init_kwargs只在传模型 id 时生效。 模型已经实例化再传model_init_kwargs,会被忽略并打 warning(trl/trainer/sft_trainer.py:984-989)——新手常见误会。- dtype 默认 float32。
create_model_from_path里dtype缺省是"float32"而不是 transformers v5 的 auto(trl/trainer/utils.py:1166),训大模型记得在model_init_kwargs里显式给"bfloat16"。 - 遥测可关。 设
HF_HUB_DISABLE_TELEMETRY=1或HF_HUB_OFFLINE=1即可(trl/trainer/base_trainer.py:81注释)。 - 想加自己的方法? 抄配方:继承
_BaseTrainer+ 写一个_BaseConfig子类 + 实现_prepare_dataset和compute_loss,就获得了全套分布式/日志/模型卡能力。这也是为什么 HF 生态的新论文实现大多先出现在trl/experimental/。