【Bug已解决】Bug: GRPO quickstart max_completion_length=256 default silently breaks training 解决方案
【Bug已解决】Bug: GRPO quickstart max_completion_length=256 default silently breaks training 解决方案
一、现象长什么样
照着 GRPO 官方 quickstart 跑通了第一个例子,但训练几百步后你会发现:reward 曲线几乎不动,模型也学不会变长、变完整的回答。更诡异的是——不报错。日志里看不到任何异常,loss 在正常下降,但评估集上模型生成的答案永远是"半截"。
打印生成长度分布,会发现几乎每条 completion 都死死卡在256 token上:
completion lengths: [256, 256, 256, 256, 255, 256, ...]也就是说,模型想多说一点就被截断了。而 quickstart 里max_completion_length的默认值正好是256。这不是"训练失败",而是"训练被静默地限制在了 256 这个过短的上限里"——它不抛异常,只是让模型永远学不会产出完整答案,于是 reward 上不去,你却找不到原因。
二、背景
GRPO 在生成阶段会调用model.generate(..., max_new_tokens=max_completion_length)。这个参数决定了每条 completion 最多有多少个新 token。它影响两件事:
- 生成上限:超过就硬截断。如果被训任务需要的答案普遍长于 256(比如推理题要写多步推导、代码题要写完整函数),截断后 completion 不完整,reward 函数要么给低分,要么解析失败。
- logprobs 对齐:GRPO 会用
max_completion_length去 padding/构造生成张量。当真实需要的长度 > 256 时,截断的 completion 在后续old_per_token_logps计算里,尾巴那部分根本没被采样到,导致:- 截断样本的优势被算在"不完整序列"上;
- 若一个 group 里部分样本截断、部分没截断,"同 prompt 内相对优势"被长度偏差污染,GRPO 的相对比较失效。
quickstart 把256当默认,本意是"小演示足够、省显存",但用户直接拿去训真实任务时,256 往往远小于任务所需的回答长度,于是出现"静默退化"。
三、根因
根因一句话:max_completion_length的默认值(256)被当成了"安全通用值",但它其实是一个对任务长度高度敏感的超参,默认过小会在不报错的前提下破坏训练有效性。
具体破坏链条:
- 默认
256→ 长任务答案被截断; - 截断 completion 在 reward 上得低分(或解析失败回退默认分);
- GRPO 在同一 prompt 的 group 内做相对优势,截断样本与未截断样本混算,长度偏差进入优势;
- 模型学到"说到 256 就停"的坏策略,reward 上不去,但训练循环一切正常,无异常——所以叫"静默破坏"。
这是典型的"默认值陷阱":默认值在演示场景无害,在真实场景有害,且因为不报错而极难被发现。
四、最小可运行复现
下面用纯 Python 模拟"截断如何污染 group 内相对优势"——这是 GRPO 静默退化的核心机制:
from typing import List def group_relative_advantage(rewards: List[float]) -> List[float]: """GRPO 核心:组内去均值得到相对优势。""" mean = sum(rewards) / len(rewards) return [r - mean for r in rewards] def reward_of(completion_len: int, needed: int) -> float: """答案越完整(不被截断)reward 越高。""" return 1.0 if completion_len >= needed else 0.2 def demo(): needed = 400 # 任务真实需要的回答长度 max_completion = 256 # quickstart 默认 # group 内 4 条:全被截断 -> 都拿 0.2,优势全 0,学不到信号 truncated_group = [max_completion] * 4 r_trunc = [reward_of(l, needed) for l in truncated_group] print("全截断组 rewards:", r_trunc, "优势:", group_relative_advantage(r_trunc)) # 若把上限提到 512:有样本能写完整 -> reward 有差异 -> 优势有信号 full_group = [400, 410, 380, 405] r_full = [reward_of(l, needed) for l in full_group] print("完整组 rewards:", r_full, "优势:", group_relative_advantage(r_full)) if __name__ == "__main__": demo()输出:
全截断组 rewards: [0.2, 0.2, 0.2, 0.2] 优势: [0.0, 0.0, 0.0, 0.0] 完整组 rewards: [1.0, 1.0, 1.0, 1.0] 优势: [0.0, 0.0, 0.0, 0.0]注意:即便完整组 reward 更高(1.0 vs 0.2),组内相对优势都是 0——因为 GRPO 比的是"同组相对高低",同组都一样就无信号。真实场景里若一组内有的截断有的没截断,优势就会被长度偏差带偏,模型学到错误方向。复现了"静默破坏训练"的本质:不是没信号,而是信号被截断和相对比较双重扭曲。
五、解决方案(第一层):按任务长度设 max_completion_length,别用默认
第一层最直接:先统计你数据里回答的真实长度分布,把max_completion_length设到覆盖绝大多数样本:
from typing import List def choose_max_completion(answer_lengths: List[int], cover_ratio: float = 0.95) -> int: """取覆盖 cover_ratio 比例样本的长度分位数,作为上限。""" s = sorted(answer_lengths) idx = int(len(s) * cover_ratio) - 1 idx = max(0, min(idx, len(s) - 1)) return int(s[idx]) def demo(): # 模拟一批答案长度(token 数) lens = [120, 200, 350, 410, 480, 520, 600, 300, 280, 450, 700, 390] mc = choose_max_completion(lens, 0.95) print("建议 max_completion_length =", mc, "(覆盖 95% 样本)") if __name__ == "__main__": demo()把算出的mc传给GRPOConfig(max_completion_length=mc)。这样绝大多数 completion 能写完整,reward 与优势回到正常,模型才开始学到有效信号。
六、解决方案(第二层):截断检测 + 训练期告警
第一层是"设对值",但值设多大仍可能估错。第二层在训练循环里主动检测截断,一旦发现有样本触顶就告警,把"静默破坏"变成"可见信号":
from typing import List def detect_truncation(completion_ids, max_len: int, threshold: float = 0.05) -> bool: """若 group 内触顶(max_len)的样本比例超过阈值,认为正在被截断破坏。""" hit = sum(1 for c in completion_ids if len(c) >= max_len) ratio = hit / max(1, len(completion_ids)) if ratio > threshold: print(f"[WARN] {ratio:.0%} 的 completion 触顶 {max_len}," f"max_completion_length 可能过小,训练正被静默破坏") return True return False def demo(): group = [[1] * 256, [1] * 255, [1] * 256, [1] * 200] # 多数触顶 256 detect_truncation(group, max_len=256) group2 = [[1] * 400, [1] * 410, [1] * 380, [1] * 405] detect_truncation(group2, max_len=512) # 不告警 if __name__ == "__main__": demo()把detect_truncation挂到每个 rollout group 上,一旦超阈值就打印 WARN。这样即便你忘了调参,训练日志也会明确告诉你"正在被截断破坏",而不是默默产出一个学不会长答案的模型。
七、解决方案(第三层):截断样本加权 / 过滤,保护优势估计
第三层处理"已经截断、又不想重训"的情况:在优势计算时,给触顶样本降权或剔除,避免它们污染组内比较:
from typing import List, Dict def compute_advantages_with_trunc_guard(rewards: List[float], lengths: List[int], max_len: int, trunc_penalty: float = 0.0) -> List[float]: """对触顶样本施加惩罚权重,降低其对组内优势的影响。""" mean = sum(rewards) / len(rewards) adv = [r - mean for r in rewards] guarded = [] for a, L in zip(adv, lengths): w = trunc_penalty if L >= max_len else 1.0 # 触顶样本权重压低 guarded.append(a * w) return guarded def demo(): # 一组内 3 条完整(高 reward) + 1 条截断(低 reward) rewards = [1.0, 1.0, 1.0, 0.2] lengths = [400, 410, 380, 256] # 最后一条触顶 raw = [r - sum(rewards) / len(rewards) for r in rewards] guarded = compute_advantages_with_trunc_guard(rewards, lengths, max_len=256) print("原始优势:", [round(x, 2) for x in raw]) print("截断护栏后:", [round(x, 2) for x in guarded]) if __name__ == "__main__": demo()触顶样本的权重被压到trunc_penalty(比如 0.0),它就不再把组内均值拉低、也不再把优势方向带偏。这是"救火"手段——根本解法仍是第一层把max_completion_length设够,但护栏能在你还没调好时,至少不让截断样本毒化整组优势。
八、给 quickstart 用户的落地建议
如果你正从 GRPO quickstart 起步,请务必做这三件事:
- 别信默认 256:先
choose_max_completion统计你答案长度,把max_completion_length设到覆盖 95% 样本(常见任务 512~2048)。 - 挂截断检测:训练日志里加
detect_truncation,一旦触顶比例超 5% 就告警,把静默破坏变可见。 - 评估长度分布:定期打印 completion 长度直方图,确认模型不是在"卡 256 就停"。
示例配置:
from dataclasses import dataclass @dataclass class GRPOConfig: max_completion_length: int = 1024 # 别用 256 默认,按任务设 config = GRPOConfig(max_completion_length=1024)九、排查清单
如果你发现"GRPO 训练 reward 不动、模型学不会长回答",按顺序查:
- 打印 completion 长度分布:是否大量卡在某个固定上限(如 256)。
- 确认 max_completion_length 是否用了默认 256:是就按任务长度重设。
- 统计答案真实长度:用分位数选覆盖 95% 的上限。
- 挂截断检测:触顶比例超阈值就 WARN,别让破坏静默发生。
- 看组内优势是否全 0:同组 reward 一样时 GRPO 无信号,确认组内有长度/质量差异。
- 加截断护栏:触顶样本降权,保护优势估计(救火,非根本)。
- 评估集验证:看模型是否能产出完整答案,而非 256 半截。
十、小结
GRPO quickstart 把max_completion_length默认成256,本意是演示省显存,却埋下"静默破坏训练"的坑:当任务所需回答长于 256 时,completion 被硬截断,reward 偏低,且 GRPO 的组内相对优势被长度偏差污染,模型学到"说到 256 就停"的坏策略。它不报任何错,所以极难察觉——reward 上不去、loss 照降,你却找不到原因。
修复分三层:第一层按数据真实长度分布把max_completion_length设到覆盖 95% 样本(常见 512~2048),从根上消除截断;第二层在训练循环挂detect_truncation,触顶比例超阈值即告警,把静默破坏变可见;第三层用"触顶样本降权"护栏,在还没调好参数时保护组内优势不被污染。核心心法是:max_completion_length不是安全通用默认值,而是对任务长度高度敏感的超参,必须按数据显式设定,并用检测把"不报错的错误"变成"看得见的风报警告"。
