Mind Lab Toolkit (MinT)
使用 MinT

RL (GRPO)

RL (GRPO) 适用于可用 reward / verifier / 环境反馈打分的场景,例如数学推理(答案可程序化校验)、代码(沙箱执行)、对话质量(judge model 评分)等。MinT 把 GRPO、PPO 等 RL 算法封装成与 SFT 同形态的 API,减少自研成本。

RL 训练循环

GRPO 循环分四个阶段:

  1. Sampling:当前 policy 为每个 prompt 生成多条 response。
  2. Scoring:reward 函数(verifier / judge / 环境)为每条 response 打分。
  3. Advantage 计算:在同一 prompt 的一组样本内做 reward 中心化:advantage[i] = reward[i] - mean_r
  4. Training:用 loss_fn="importance_sampling" 训练;reward 高于 mean_r 的样本拿到正梯度,低于的拿到负梯度。

构造 Datum

Prompt 位置的 weight 为 0,response 位置填入 advantage 值;同时需要传入每个 token 的 logprobs(从 sampling 步骤取得)和 advantages(zero-padding 覆盖 prompt prefix)。

mean_r = sum(rewards) / num_samples
for i, (seq, reward) in enumerate(zip(sequences, rewards)):
    advantage = reward - mean_r
    response_len = len(seq.tokens) - prompt_len
    advantages   = [0.0] * (prompt_len - 1) + [advantage] * response_len
    datums.append(types.Datum(
        model_input=types.ModelInput.from_ints(tokens=seq.tokens[:-1]),
        loss_fn_inputs={
            "target_tokens": seq.tokens[1:],
            "weights":       [0.0] * (prompt_len - 1) + [1.0] * response_len,
            "logprobs":      seq.logprobs[1:],
            "advantages":   advantages,
        },
    ))

参数

参数类型默认值含义
loss_fnstr"importance_sampling"GRPO 默认;服务端还支持 "ppo""cispo""dro"
group_sizeint4每个 prompt 采样条数;越大方差越低但显存更多,典型 4–16
groups_per_batchint8每步训练的 prompt 数
max_tokensint16最大生成长度;数学 ≈ 16,对话 ≈ 128,代码 ≈ 256
learning_ratefloat2e-5RL 通常低于 SFT,典型 1e-5 到 4e-5
temperaturefloat0.8采样温度,RL 典型 0.7–1.0
kl_penalty_coeffloat0.0> 0 时对 policy 偏离 reference 施加 KL 惩罚,防止 collapse
base_modelstr"Qwen/Qwen3-0.6B"底座模型 ID
rankint16LoRA 秩

SFT + RL 两阶段

对于同时拥有监督数据和奖励信号的场景,推荐先用 SFT 做暖启动、再切到 GRPO。参考路径见社区版 quickstart.py

下一步

本页目录