RRLforge
实验预计 50 分钟

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 版本文档

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_SIZEnum_generations一个 prompt 采几条
PROMPTS_PER_STEPper_device_train_batch_size⚠️ 语义不同,见下
MICRO_BATCH由前两者和累积步数推出TRL 自己算
gradient_accumulation_steps梯度累积
BETAbetaKL 系数
eps_low / eps_highepsilon / epsilon_highclip 阈值
max_tokensmax_completion_length生成上限
temperaturetemperature采样温度
util=0.30vllm_gpu_memory_utilization给 vLLM 的显存
sync_weights_to_vllm()自动TRL 每步替你做
engine.sleep()vllm_enable_sleep_mode部分版本有此项
logp_old 那次前向自动内部处理
序列/token 平权loss_type默认值随版本变,查一下
×per_device_train_batch_size 的语义陷阱

在 TRL 的 GRPO 里,这个参数指的是生成序列的条数,而不是 prompt 的个数。而且它必须能被 num_generations 整除。

也就是说 per_device_train_batch_size=8num_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]:
    ...

三条规则:

  1. completions 是本批所有生成的文本列表(不是单条)
  2. 返回列表长度必须与 completions 一致
  3. 数据集里的其它列会按名字通过 kwargs 传进来 —— 上面 answer=... 就是这么拿到的

多个 reward 函数用 reward_weights 加权求和,等价于你 L2 里的 total_reward

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 out

L2 的手写脚本里也该这么写。

跑起来,和 L2 的曲线放一起比

python -m scripts.trl_grpo.train 2>&1 | tee logs/trl-01.log
rl@forge
目标 0/3
  1. 1.确认你这个版本真实接受哪些参数
  2. 2.冒烟测试:确认能跑、显存够、指标合理
  3. 3.跑完 200 步,与 L2 的结果对比
TRL 实跑演练。重点是把它的指标和你 L2 的曲线对上。
goals 看目标,hint 要提示。
[rl@forge ~/rlforge]$
help 查看用法 · 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 分布会有细微不同。

查法:把它关掉再跑一次对比。

i其它常见对不齐来源
  • 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 数,最容易算错 batch
  • loss_type 默认是 token 平权,与手写的序列平权不等价,长度曲线会分岔
  • inspect.signature 问你装的那个版本,不要照抄博客
  • 分数对不上先查三处:batch 语义、loss_type、importance sampling 修正

延伸资料