TRL GRPOTrainer:5090 单卡实跑
把手写脚本换成 TRL,跑同一个任务、比同一个分数。这是本站唯一在 5090 上被验证过的框架路径。
学完这节你能做到
- 用 GRPOTrainer 复现 L2 的实验,并对上分数
- 把自己写的 reward 函数接进 TRL 的 reward_funcs
- 把 TRL 的配置项逐个映射回你手写代码里的变量
为什么单卡框架层选 TRL
这条路径参考的两个工业框架都跑不了单卡(slime 最小 8×H100,Miles 的硬件列表是 A100/H100/B200/MI300 系列)。所以框架层实跑用 TRL:
| 候选 | 单卡 32GB 可行性 |
|---|---|
| TRL | ✅ 装得上、支持 colocate、文档跟得上 |
| Unsloth | ✅ 也可以,显存更省,但封装更厚 |
| verl | ⚠️ 官方入门是单卡 24GB 的 notebook,但整体为多卡设计,容易和 OOM 搏斗 |
| slime / Miles | ❌ Megatron + 多卡并行,单卡无从下手 |
选 TRL 的关键理由是它的抽象层次刚好:GRPOTrainer 帮你管住训练循环和 vLLM 协作,但 reward 函数还是你自己写的普通 Python。你 L2 写的 rewards.py 可以原样接进来。
最小可跑脚本
# scripts/trl_grpo/train.py
from datasets import load_dataset
from trl import GRPOConfig, GRPOTrainer
from scripts.grpo_scratch.rewards import correctness_reward, format_reward
MODEL = "./models/Qwen3-0.6B"
# ---------- reward:签名是固定的 ----------
def reward_correct(completions, answer, **kwargs):
"""
completions: list[str],本批所有生成
answer: list[str],数据集里同名列会自动按批传进来
返回: list[float],长度与 completions 一致
"""
return [correctness_reward(c, a) for c, a in zip(completions, answer)]
def reward_format(completions, **kwargs):
return [format_reward(c) for c in completions]
def main():
ds = load_dataset("openai/gsm8k", "main", split="train")
def prep(row):
# prompt 列会被自动套 chat template
return {
"prompt": [{"role": "user", "content": row["question"]}],
"answer": row["answer"].split("####")[-1].strip(),
}
ds = ds.map(prep, remove_columns=ds.column_names)
cfg = GRPOConfig(
output_dir="out/trl-grpo-01",
# ---- 对应你手写脚本里的变量 ----
num_generations=8, # ← GROUP_SIZE
per_device_train_batch_size=8, # ← PROMPTS_PER_STEP(见下方说明)
gradient_accumulation_steps=4,
max_completion_length=512, # ← max_tokens
max_prompt_length=256,
temperature=1.0,
top_p=1.0,
beta=0.0, # ← KL 系数
epsilon=0.2, # ← clip 下界
epsilon_high=0.28, # ← clip 上界(非对称)
learning_rate=1e-6,
adam_beta2=0.98,
max_grad_norm=1.0,
max_steps=200,
# ---- 单卡 colocate ----
use_vllm=True,
vllm_mode="colocate",
vllm_gpu_memory_utilization=0.30,
# ---- 显存 ----
bf16=True,
gradient_checkpointing=True,
# ---- 日志 ----
logging_steps=1,
save_steps=20,
report_to="none", # 想接 wandb 就改这里
)
trainer = GRPOTrainer(
model=MODEL,
args=cfg,
train_dataset=ds,
reward_funcs=[reward_correct, reward_format],
reward_weights=[1.0, 1.0], # 权重原则同 L2:便宜的分给得少
)
trainer.train()
if __name__ == "__main__":
main()
TRL 迭代很快,vllm_mode 的默认值、参数名、reward 函数签名在不同版本间都变过。以你装的那个版本的文档为准,不要照抄任何博客(包括这一页)。
python -c "import trl; print(trl.__version__)"
python -c "from trl import GRPOConfig; import inspect; print(inspect.signature(GRPOConfig.__init__))"第二条命令直接打出你这个版本真实接受的参数,比查文档快。
参数对照表
这张表是这一节的核心。左边是你 L2 亲手写的,右边是 TRL 的名字:
| 你的变量 | TRL 参数 | 说明 |
|---|---|---|
GROUP_SIZE | num_generations | 一个 prompt 采几条 |
PROMPTS_PER_STEP | per_device_train_batch_size | ⚠️ 语义不同,见下 |
MICRO_BATCH | 由前两者和累积步数推出 | TRL 自己算 |
| — | gradient_accumulation_steps | 梯度累积 |
BETA | beta | KL 系数 |
eps_low / eps_high | epsilon / epsilon_high | clip 阈值 |
max_tokens | max_completion_length | 生成上限 |
temperature | temperature | 采样温度 |
util=0.30 | vllm_gpu_memory_utilization | 给 vLLM 的显存 |
sync_weights_to_vllm() | 自动 | TRL 每步替你做 |
engine.sleep() | vllm_enable_sleep_mode | 部分版本有此项 |
logp_old 那次前向 | 自动 | 内部处理 |
| 序列/token 平权 | loss_type | 默认值随版本变,查一下 |
在 TRL 的 GRPO 里,这个参数指的是生成序列的条数,而不是 prompt 的个数。而且它必须能被 num_generations 整除。
也就是说 per_device_train_batch_size=8 配 num_generations=8,一步只处理 1 个 prompt 的 8 条生成。
想要「8 个 prompt × 8 条」,需要让 per_device_train_batch_size × gradient_accumulation_steps = 64。
这个语义差异是接框架时最容易算错 batch 的地方。跑起来后核对日志里的实际样本数,别凭参数名猜。
接自己的 reward
TRL 的 reward 函数约定:
def my_reward(completions, **kwargs) -> list[float]:
...
三条规则:
completions是本批所有生成的文本列表(不是单条)- 返回列表长度必须与
completions一致 - 数据集里的其它列会按名字通过
kwargs传进来 —— 上面answer=...就是这么拿到的
多个 reward 函数用 reward_weights 加权求和,等价于你 L2 里的 total_reward。
一条生成里出现意外格式,如果你的 reward 抛了异常,整个训练就停了。
在函数最外层套一个 try/except,出错返回 0.0 并打日志:
def reward_correct(completions, answer, **kwargs):
out = []
for c, a in zip(completions, answer):
try:
out.append(correctness_reward(c, a))
except Exception as e:
print(f"[reward] 跳过一条:{e}")
out.append(0.0)
return outL2 的手写脚本里也该这么写。
跑起来,和 L2 的曲线放一起比
python -m scripts.trl_grpo.train 2>&1 | tee logs/trl-01.log
- 1.确认你这个版本真实接受哪些参数
- 2.冒烟测试:确认能跑、显存够、指标合理
- 3.跑完 200 步,与 L2 的结果对比
TRL 实跑演练。重点是把它的指标和你 L2 的曲线对上。 goals 看目标,hint 要提示。
分数对不上时先看哪三个地方
如果 TRL 的结果和你手写版差很多,按这个顺序查:
一、batch 语义
per_device_train_batch_size 是序列条数不是 prompt 数。算错会导致实际 batch 差好几倍,而 batch 大小直接影响梯度噪声。
查法:看日志里 epoch 的推进速度,或直接打印每步的实际样本数。
二、loss_type(序列平权 vs token 平权)
这是上面演示里出现的差异来源。两者不等价,长度曲线会明显不同。
查法:print(cfg.loss_type),和你手写实现里的归一化方式对齐。
三、importance sampling 修正
TRL 默认开 vllm_importance_sampling_correction,它在修正 vLLM 与 transformers 的概率差异。你手写版(如果用的是「多跑一次前向拿 logp_old」的方案 A)不需要这个修正,两边的 ratio 分布会有细微不同。
查法:把它关掉再跑一次对比。
- chat template:TRL 自动套模板,你手写时是手动套的。打印出来逐字符比较。
- 答案抽取:GSM8K 的
answer列包含完整解题过程,####之后才是答案。两边都要正确切分。 - 随机种子:
seed不同时单次结果会有几个百分点波动。比较趋势不要比较单点。
TRL GRPO 里 `per_device_train_batch_size=8` 配 `num_generations=8`,一步处理几个 prompt?
TRL 跑出来的平均回答长度明显长于你手写版,最可能的原因是?
为什么这一节强调「不要照抄任何博客的 TRL 参数」?
这节课的落点
- 单卡框架层选 TRL:抽象层次刚好,reward 还是你自己的 Python
- 你 L2 写的
rewards.py能原样接进reward_funcs - reward 函数三条约定:收批量、返回等长列表、数据集其它列走
kwargs per_device_train_batch_size是序列条数不是 prompt 数,最容易算错 batchloss_type默认是 token 平权,与手写的序列平权不等价,长度曲线会分岔- 用
inspect.signature问你装的那个版本,不要照抄博客 - 分数对不上先查三处:batch 语义、loss_type、importance sampling 修正