【Bug已解决】Add warning when all completions are truncated 解决方案
原始报错:Add warning when all completions are truncated 场景:在 GRPO / RL 训练的生成阶段,一个 prompt 会采样出一组(group)completion。如果这组里所有 completion 都因为达到
max_completion_length而被截断,说明这个 prompt 在现有长度限制下根本无法完整生成。此时继续拿这些残缺 completion 算奖励,信号全错。现状是没有任何告警,训练静默地在这些"全截断"样本上浪费。提议:当一组 completion 全部截断时,发出明确警告。 关键词:生成截断、全截断检测、GRPO 组、训练健康告警、奖励信号保护。
一、现象长什么样
训练日志里:
- 某个 prompt 采样了 8 条 completion,全部撞到
max_completion_length上限被截断; - 这 8 条残缺文本进入奖励计算,分数要么是截断处的残值,要么全 0;
- 整个组的相对优势估计基于"全是残缺文本",毫无意义;
- 没有任何日志提示"这组全截断了",训练照常推进,指标悄悄失真;
- 大量这类 prompt 累积,训练效果莫名变差,排查时才知道问题。
这个 issue 是"至少给个警告"——让全截断从隐形变成可见,研究者才能决定是调大长度、跳过该 prompt、还是换数据。
二、背景:为什么"全截断"特别危险
单条截断已经不好(见相关长度默认值问题),但整组全截断更糟:
- 单条截断:组内还有完整 completion,相对优势仍可估计;
- 整组全截断:组内没有一条完整,相对优势在"一堆残缺"里算,等于用噪声当信号;
- 更严重的是,全截断往往意味着
max_completion_length对该 prompt系统性不够(任务本身长),不是偶然。
所以全截断是一个高价值告警信号:它不只是"这条样本有问题",而是"这个长度配置对这类 prompt 不适用"。
三、根因:无全截断检测与告警
根因拆解:
- 无组级检测:代码只逐条生成,没汇总"这组几条截了";
- 全截断静默:即使检测到,也没告警,训练继续;
- 无聚合统计:不统计"全截断 prompt 占比",问题规模化后看不见;
- 无处置:告警后没有标准动作(跳过/调大/记录),只是一句日志;
- 与长度默认脱节:全截断频发本应推动调大默认,但没人看见。
下面用最小模型复现"一组全截断但无告警",再给修复。
四、最小可运行复现
def generate_one(true_len, max_len): gen = min(true_len, max_len) return gen, gen < true_len # (生成长度, 是否截断) def process_group(prompt_true_len, group_size, max_len): completions = [generate_one(prompt_true_len, max_len) for _ in range(group_size)] truncated = [c for c in completions if c[1]] # 错误:只处理,不检测全截断,也不告警 return len(truncated) if __name__ == "__main__": # 任务需要 600,max=256,8 条全截 n = process_group(600, group_size=8, max_len=256) print(f"截断条数: {n}/8") # 8/8 全截断,但无任何警告运行输出 8/8,但程序一声不吭——全截断被完全忽略。
五、方案:组内全截断检测
第一层:生成一组后,汇总截断情况,当整组全截断时明确标记:
def analyze_group(completions, max_len): # completions: 每条 (gen_len, is_truncated) total = len(completions) trunc = sum(1 for _, t in completions if t) all_truncated = (trunc == total) return { "total": total, "truncated": trunc, "all_truncated": all_truncated, "ratio": trunc / total, } def process_group_fixed(prompt_true_len, group_size, max_len): comps = [generate_one(prompt_true_len, max_len) for _ in range(group_size)] info = analyze_group(comps, max_len) if info["all_truncated"]: print(f"[WARN] 整组 {info['total']} 条全部截断 " f"(任务需 {prompt_true_len} > max {max_len}),奖励信号无效") return info if __name__ == "__main__": process_group_fixed(600, group_size=8, max_len=256) # [WARN] 整组 8 条全部截断...全截断从隐形变成显式 WARN,研究者立刻知道这个 prompt 在当前长度下无解。
六、方案:分级告警(单条 / 整组 / 整批)
第二层:告警分三级,分别处理,避免要么静默要么刷屏:
class TruncationWatcher: def __init__(self, warn_group_ratio=1.0, warn_batch_ratio=0.1): self.group_all = 0 self.group_total = 0 self.warn_group_ratio = warn_group_ratio self.warn_batch_ratio = warn_batch_ratio def on_group(self, info): self.group_total += 1 if info["all_truncated"]: self.group_all += 1 # 单组级:直接 WARN(每条全截断都值得看) print(f"[WARN] prompt#{self.group_total} 整组全截断") # 批级:攒够一定量再汇总 if self.group_total % 50 == 0: ratio = self.group_all / self.group_total if ratio > self.warn_batch_ratio: print(f"[BATCH WARN] 全截断 prompt 占比 {ratio:.1%}," f"建议调大 max_completion_length") if __name__ == "__main__": w = TruncationWatcher() for _ in range(50): info = process_group_fixed(600, 8, 256) # 全截断 w.on_group(info) # 每 50 组输出一次批级汇总分级让单组全截断即时可见、批量趋势用汇总呈现,既不漏也不刷屏。
七、方案:告警后处置(跳过 / 调大 / 记录)
第三层:告警不能只是日志,要有标准处置——全截断的组应从奖励计算里排除(或标记无效),并记录到报告供后续决策:
def decide_group(info, max_len, task_len): if info["all_truncated"]: return { "action": "skip", # 无效奖励,跳过该组 "reason": "all_truncated", "suggest_max_len": max(task_len, max_len * 2), } if info["ratio"] >= 0.5: return {"action": "keep_warn", "reason": "high_truncation"} return {"action": "keep", "reason": "ok"} if __name__ == "__main__": info = process_group_fixed(600, 8, 256) decision = decide_group(info, 256, 600) print("处置:", decision) # {'action': 'skip', 'reason': 'all_truncated', 'suggest_max_len': 512}全截断的组被 skip(不污染奖励),同时给出"建议调大到 512"的可操作建议,把告警变成决策依据。
八、验证:把"全截断检测"锁进测试
def test_all_truncated_detected(): comps = [(256, True)] * 8 # 8 条全截断 info = analyze_group(comps, 256) assert info["all_truncated"] is True assert info["truncated"] == 8 def test_not_all_truncated(): comps = [(256, True), (100, False)] * 4 # 半截 info = analyze_group(comps, 256) assert info["all_truncated"] is False def test_skip_decision_on_all_truncated(): info = analyze_group([(256, True)] * 8, 256) d = decide_group(info, 256, 600) assert d["action"] == "skip" if __name__ == "__main__": test_all_truncated_detected() test_not_all_truncated() test_skip_decision_on_all_truncated() print("全截断检测与处置测试通过。")九、排查清单("全截断无告警"按顺序查)
- 组级检测:生成一组后是否汇总"几条截了"?还是只逐条处理?
- 全截断告警:整组全截时是否有明确 WARN?静默则危险。
- 分级:是否区分单组告警与批量趋势?避免刷屏或遗漏。
- 聚合统计:全截断 prompt 占比是否被统计?规模化问题才看得见。
- 处置:告警后是否 skip/调大/记录?只日志不够。
- 奖励保护:全截断的组是否从奖励计算排除?否则污染信号。
- 与长度联动:全截断频发是否推动调大 max_completion_length?
十、小结
"所有 completion 都截断时应告警"是生成阶段缺少组级截断检测,全截断被静默消化导致奖励信号系统性失真。修复三层:
- 组级检测:生成一组后汇总截断数,整组全截断即标记;
- 分级告警:单组全截断即时 WARN、批量趋势汇总呈现,不漏不刷;
- 告警即处置:全截断组从奖励计算跳过,并给出"调大长度"的可操作建议。
核心原则:全截断不是普通截断——它意味着当前长度配置对该 prompt 系统性失效,组内无任何完整样本,奖励信号完全不可信。把它从隐形变成显式告警并自动 skip,训练才不会在"一堆残缺文本"上悄悄学歪。