RRLforge
先搞懂原理 路线9 / 13 · 用代码验证你的理解退出路线
实验预计 60 分钟

拼成训练循环:第一次真正的 RL 训练

把 rollout、reward、loss 串起来,加上权重同步和显存管理,在 5090 上跑完第一次训练。

学完这节你能做到

  • 跑完一次完整训练,看到 reward 曲线上升
  • 实现训练权重到 vLLM 的同步,并解释不同步会发生什么
  • 在 32GB 内安排好训练与推理的显存,扛住不 OOM

主循环

前三节的零件都齐了,这一节把它们串起来。骨架就五步:

for step in range(N):
    ① 采样      rollout(engine, prompts, G)      ← 用 vLLM
    ② 打分      reward(text, gold)               ← 纯 Python
    ③ 算优势    (r - mean) / std,按组            ← 一行
    ④ 更新      forward + loss + backward        ← transformers
    ⑤ 同步      把新权重推给 vLLM                 ← 关键,别漏

第⑤步是新手最容易漏的一步,漏了会导致一个很隐蔽的问题——待会儿讲。

完整脚本

# scripts/grpo_scratch/train.py
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

from .rollout import build_engine, rollout
from .rewards import total_reward
from .grpo import sequence_logprobs, group_advantages, grpo_loss

MODEL = "./models/Qwen3-0.6B"
GROUP_SIZE = 8
PROMPTS_PER_STEP = 8
MICRO_BATCH = 2
LR = 1e-6
BETA = 0.0          # KL 系数,先设 0,只记录不惩罚
STEPS = 200


def main():
    tok = AutoTokenizer.from_pretrained(MODEL)

    # 训练侧模型
    model = AutoModelForCausalLM.from_pretrained(
        MODEL, dtype=torch.bfloat16, attn_implementation="sdpa"
    ).cuda()
    model.gradient_checkpointing_enable()
    model.config.use_cache = False        # 训练时关掉 KV cache,省显存

    opt = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=0.0, betas=(0.9, 0.98))

    # 推理侧引擎(同一张卡)
    engine = build_engine(MODEL, util=0.30, max_len=1024)

    dataset = load_gsm8k_train()          # 见 L0「数据集与评测基线」

    for step in range(STEPS):
        batch = dataset.sample(PROMPTS_PER_STEP)

        # ---------- ① 采样 ----------
        engine.wake_up()                              # 把显存要回来
        samples = rollout(engine, tok, [b["question"] for b in batch],
                          group_size=GROUP_SIZE, max_tokens=512, temperature=1.0)
        engine.sleep(level=1)                         # 把 KV cache 的显存让出去

        # ---------- ② 打分 ----------
        rewards = []
        for s in samples:
            gold = batch[s.group_id]["answer"]
            truncated = len(s.completion_ids) >= 512
            rewards.append(total_reward(s.text, gold, truncated)["total"])
        rewards_t = torch.tensor(rewards, dtype=torch.float32, device="cuda")

        # ---------- ③ 算优势 ----------
        advantages = group_advantages(rewards_t, GROUP_SIZE)

        # 整批优势全为 0(所有组都退化)时,这一步没有任何可学的东西
        if advantages.abs().max() < 1e-6:
            print(f"[{step}] 所有组都退化,跳过")
            continue

        # ---------- ④ 更新 ----------
        input_ids, attn_mask, comp_mask = collate(samples, tok.pad_token_id)

        # logp_old:用更新前的参数跑一遍,存下来
        with torch.no_grad():
            logp_old, _ = sequence_logprobs(model, input_ids, attn_mask, comp_mask)

        opt.zero_grad(set_to_none=True)
        n_micro = (len(samples) + MICRO_BATCH - 1) // MICRO_BATCH
        metrics_acc = []

        for i in range(0, len(samples), MICRO_BATCH):
            sl = slice(i, i + MICRO_BATCH)
            logp_new, mask = sequence_logprobs(
                model, input_ids[sl], attn_mask[sl], comp_mask[sl]
            )
            loss, m = grpo_loss(
                logp_new, logp_old[sl], None,     # 参考模型见下方讨论
                advantages[sl], mask, beta=BETA,
            )
            (loss / n_micro).backward()            # 梯度累积
            metrics_acc.append(m)

        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        opt.step()

        # ---------- ⑤ 同步权重 ----------
        sync_weights_to_vllm(model, engine)

        log(step, rewards, advantages, metrics_acc, grad_norm, samples)


if __name__ == "__main__":
    main()

第⑤步:为什么每步都要同步权重

GRPO 是 on-policy 算法:优势估计只对「产生这批数据的那个策略」有效。

如果不同步,会发生这样的事:

step 0   vLLM 用参数 θ₀ 采样  →  训练侧更新到 θ₁      ✅ 数据来自 θ₀,训 θ₀→θ₁,正确
step 1   vLLM 还是用 θ₀ 采样  →  训练侧从 θ₁ 更新到 θ₂  ❌ 数据来自 θ₀,却在训 θ₁
step 50  vLLM 还在用 θ₀      →  训练侧已经是 θ₅₀       ❌❌ 差得离谱

到 step 50 时,你实际在做的是「用一个 50 步前的旧模型采的数据,去更新当前模型」——这是严重的 off-policy,ratio 会离 1 很远,clip 大面积触发,训练基本停滞。

症状:reward 前几步涨一点然后就平了,clip_frac 一路走高。

×这个 bug 不报错,只是「训不动」

新手最容易在这里卡很久,因为一切看起来都在运行:loss 有值、梯度在更新、日志在刷。只是效果不涨。

排查方法:打印 ratio_mean。如果它随 step 单调偏离 1(比如跑到 0.6 或 1.5),基本就是权重同步没做或做错了。

同步的实现(简化版):

def sync_weights_to_vllm(model, engine):
    """把训练侧的权重刷进 vLLM 的模型实例。"""
    state = {k: v for k, v in model.state_dict().items()}
    llm_model = engine.llm_engine.model_executor.driver_worker.model_runner.model
    llm_model.load_weights(state.items())
i工业框架在这一步花了很大力气

单卡上同步就是一次显存内的拷贝,1~3 秒。跨节点就麻烦了:训练侧的权重是切分在多张卡上的(TP/PP),推理侧的切分方式还不一样,要先聚合再重新分发。

所以你会在框架文档里看到:Miles 的 P2P weight transfer(走 RDMA 直传,跳过 CPU)、delta weight sync(只传变化的部分);slime 把权重同步列为「不做多推理后端抽象」的理由之一。

这些设计解决的就是你刚刚用三行代码搞定的那件事——在 8×H100 上它要花几十秒,占单步时间的一大块。

参考模型:单卡上的取舍

算 KL 需要参考模型(训练开始前的那份权重)。单卡上有三个选择:

方案显存KL 可用适用
常驻一份 bf16 副本+1.1 GiB(0.6B)显存宽裕时的默认选择
不要参考模型,β=00❌ 连监控都没有显存紧张、短训练
用 LoRA:禁用适配器即得基座0LoRA 训练时的免费午餐

上面的脚本里我传了 NoneBETA=0,也就是第二种。这是个有意的简化:先让主循环跑通,第一次训练不看 KL。

跑通之后建议加上参考模型(0.6B 只要 1.1 GiB),因为 KL 曲线是判断「模型有没有跑飞」最直接的指标,L4 闯关里会大量用到。

LoRA 用户的免费午餐

如果你用 LoRA,参考模型就是「禁用适配器后的模型」:

with model.disable_adapter():
    logp_ref, _ = sequence_logprobs(model, input_ids, attn_mask, comp_mask)

一分显存都不用多花。这是 LoRA 在 RL 场景下一个被低估的好处。

sleep / wake:单卡的显存腾挪

采样和训练抢同一张卡。vLLM 的 sleep mode 让这件事变得可行:

采样阶段  engine.wake_up()   vLLM 申请 KV cache 显存,开始生成
              ↓
训练阶段  engine.sleep(1)    KV cache 显存归还,训练侧拿去用
              ↓
下一轮    engine.wake_up()   再申请回来

level=1 只释放 KV cache,保留权重;level=2 连权重也卸载(省更多,但唤醒更慢)。单卡 colocate 用 level=1 就够。

代价是每轮切换有几秒开销——L4 的时间账计算器里有一项「colocate 切换」就是它。

显存安排:0.6B 全参 + colocate 的一组实测参考
vLLM (util=0.30,含 KV cache)      约  9.5 GiB
训练侧 权重 + 梯度 + AdamW           约  9.0 GiB
激活(micro_batch=2,开检查点)       约  0.2 GiB
logits(float32,1024 长度)         约  1.2 GiB
CUDA context + 碎片                  约  1.5 GiB
─────────────────────────────────────────────
合计                                 约 21.4 GiB   ← 32 GiB 上有余量

余量可以拿去做两件事之一:加参考模型(+1.1 GiB),或者把 vLLM 的 util 提到 0.40 让采样更快。优先选后者,因为采样是瓶颈。

该记什么日志

RL 的调试全靠这几条曲线。一开始就把它们记全,出问题时不用重跑:

def log(step, rewards, advantages, metrics, grad_norm, samples):
    import statistics as st
    lengths = [len(s.completion_ids) for s in samples]
    print(
        f"[{step:4d}] "
        f"reward={st.mean(rewards):.3f} "          # 最重要:在涨吗
        f"acc={sum(r > 0.9 for r in rewards)/len(rewards):.2%} "  # 只看正确性
        f"adv_std={advantages.std().item():.3f} "  # 接近 0 = 白跑
        f"kl={st.mean(m['kl'] for m in metrics):.4f} "
        f"clip={st.mean(m['clip_frac'] for m in metrics):.3f} "
        f"ratio={st.mean(m['ratio_mean'] for m in metrics):.4f} "  # 该接近 1
        f"len={st.mean(lengths):.0f} "             # 长度爆炸的预警
        f"trunc={sum(l >= 512 for l in lengths)/len(lengths):.1%} "
        f"gnorm={grad_norm:.2f}"
    )

八个数,各自的作用:

指标健康表现异常含义
reward缓慢上升平了 = 没在学;暴涨 = 可能在刷分
acc和 reward 同步上升reward 涨但 acc 不涨 = reward hacking
adv_std稳定在 0.5~1.2趋近 0 = 大量组退化,白跑
kl缓慢上行陡增 = 跑飞
clip_frac几个百分点> 20% = 单步太猛,降学习率
ratio_mean接近 1.0偏离 = 权重没同步
len稳定单调暴涨 = 长度爆炸
trunc< 10%高 = max_tokens 不够
!reward 和 acc 一定要分开记

这是最重要的一条监控设计。

reward 是总分(含格式分),acc 只看答案对不对。两者背离就是 reward hacking 的确诊信号:分数在涨,能力没涨。

只记 reward 的话,你会以为训练很成功,直到评测时发现准确率没变。

启动,然后等

python -m scripts.grpo_scratch.train 2>&1 | tee logs/run-01.log

先跑 10 步的冒烟测试,确认:

  • 不 OOM
  • ratio_mean ≈ 1.0
  • adv_std 不是 0
  • 显存占用符合预期

确认之后再放长到 200 步。0.6B 上一步约 2040 秒(取决于你 L0 测出的吞吐),200 步大约 1.52 小时。

rl@forge
目标 0/3
  1. 1.跑 10 步冒烟测试,确认不 OOM 且各项指标正常
  2. 2.完整跑一次训练
  3. 3.故意关掉权重同步,认出它的症状
第一次训练。先冒烟测试,再看曲线,出问题时按 ratio 定位。
goals 看目标,hint 要提示。
[rl@forge ~/rlforge]$
help 查看用法 · goals 看目标 · hint 要提示 · ↑↓ 翻历史
检查点单选

训练跑着,loss 有值、梯度在更新,但 reward 前几步涨一点之后就平了,同时 ratio_mean 从 1.0 单调跌到 0.6。最可能是什么问题?

检查点单选

为什么日志里 reward 和 acc 一定要分开记?

检查点单选

训练到后期 adv_std 从 1.0 缓慢降到 0.6,说明什么?

这节课的落点

  • 五步主循环:采样 → 打分 → 优势 → 更新 → 同步权重
  • 不同步权重会变成严重 off-policy,特征指纹是 ratio_mean 单调偏离 1
  • 单卡靠 engine.sleep(1) / wake_up() 腾挪显存,每轮几秒开销
  • 参考模型单卡上可选:常驻 1.1 GiB / 完全不要(β=0)/ LoRA 禁用适配器免费拿
  • 八个必记指标,其中 reward 与 acc 必须分开,ratio_mean 必须记
  • 先跑 10 步冒烟测试再放长;0.6B 上 200 步约 2 小时
  • 长度上涨本身不是问题,acc 停了而 len 还在涨才是

延伸资料