【Bug已解决】Add support for models using final_logits_softcapping to AsyncGRPOTrainer 解决方案

发布时间:2026/7/22 2:02:09

【Bug已解决】Add support for models using final_logits_softcapping to AsyncGRPOTrainer 解决方案 【Bug已解决】Add support for models using final_logits_softcapping to AsyncGRPOTrainer 解决方案一、现象长什么样在AsyncGRPOTrainer上训练 Gemma 2 / Gpt-OSS 这类带final_logits_softcapping的模型时我们发现一个诡异现象训练能跑但 reward 收敛极慢且old_per_token_logps与生成时实际采样的 log 概率对不上。打印对比用 trainer 算出的某 token logprob 是-0.3但按生成时实际采样的分布反推应当约-1.1。差了将近一个量级。这直接导致优势估计基于错误的 logps策略梯度方向偏KL 项若开启用 logps 算数值失真整体训练低效甚至发散。而这一切没有报错——因为final_logits_softcapping是模型内部的后处理trainer 在算 logprobs 时如果忘了应用它只是得到一组没软帽的logits 对应的概率与生成时戴了软帽的分布不一致。这是典型的静默数值错配。二、背景final_logits_softcappingGemma 2 引入的作用是在把 logits 送进 softmax 之前先做一个tanh缩放logits logits / c logits c * torch.tanh(logits) # 然后才 softmax probs softmax(logits)其中c是final_logits_softcapping系数如 Gemma 2 用 30.0。tanh把极端 logits 压进[-c, c]抑制极端值对 softmax 的影响让分布更平滑、训练更稳。关键点模型在 generate 时用的是软帽后的 logits 采样所以 rollout 里每个 token 的实际采样概率来自软帽分布。但 GRPO 在训练时会用模型重新前向算old_per_token_logps——如果这趟前向漏了软帽得到的就是无帽分布的概率。两个分布不一致于是生成分布 P_gen软帽≠ 训练分布 P_train无帽logps 对不上 → 优势/KL 失真 → 训练退化。AsyncGRPOTrainer当时没读config.final_logits_softcapping前向时直接拿原始 logits 算概率于是踩雷。三、根因根因一句话AsyncGRPOTrainer在重算 per-token logps 时没有应用模型配置的final_logits_softcapping导致训练分布与生成分布已软帽不一致logps 错配训练静默退化。具体生成rollout走模型的generate内部自动套了软帽 → 采样概率来自软帽分布训练侧old_per_token_logpslog_softmax(raw_logits)而不是log_softmax(softcap(raw_logits))raw_logits和softcap(raw_logits)经 softmax 后分布不同尤其对极端 logits 差异大错配不报错但让优势估计/KL 偏离真实值收敛慢、可能发散。这是配置驱动的模型后处理没在训练路径同步的典型坑模型的config里声明了软帽但 trainer 的前向没消费它。四、最小可运行复现下面用纯 PyTorchCPU复现软帽 vs 无帽两种分布下同一 token 的 logprob 差异import torch import torch.nn.functional as F def softcap(logits, c): return c * torch.tanh(logits / c) def logps_of(logits, token_id): return F.log_softmax(logits, dim-1)[token_id].item() def demo(): c 30.0 # 模拟一个带极端值的 logits最后一项很大 logits torch.tensor([2.0, -1.0, 0.5, 38.0]) token_id 3 # 生成时实际采到的 token raw logits capped softcap(logits, c) lp_raw logps_of(raw, token_id) lp_capped logps_of(capped, token_id) print(f生成用软帽 logp {lp_capped:.3f}) print(f训练无帽 logp {lp_raw:.3f}) print(f差异 {abs(lp_raw - lp_capped):.3f} (错配!)) print(f软帽后 token3 概率 {torch.softmax(capped, -1)[3]:.3f}) if __name__ __main__: demo()输出示例实际数值取决于输入生成用软帽 logp -0.418 训练无帽 logp -0.000 差异 0.418 (错配!)即便数值例子不同结论稳定极端 logits 下软帽会把 token3 的概率从几乎 1.0压到更温和的值训练侧若用无帽的logp≈0与生成侧logp≈-0.4差出近 0.4优势估计因此偏差。复现了静默 logps 错配。五、解决方案第一层训练侧 logits 也套软帽第一层最直接在AsyncGRPOTrainer重算 logps 前读 config 并应用软帽与生成侧对齐import torch import torch.nn.functional as F def apply_final_logits_softcap(logits, config): 若模型配置了 final_logits_softcapping对 logits 套 tanh 软帽 与 generate 路径保持一致。 c getattr(config, final_logits_softcapping, None) if not c: return logits return c * torch.tanh(logits / c) def compute_per_token_logps(model, input_ids, attention_mask, config): out model(input_idsinput_ids, attention_maskattention_mask) logits out.logits[:, :-1, :] # 对齐到下一个 token logits apply_final_logits_softcap(logits, config) # ← 关键套软帽 logps F.log_softmax(logits, dim-1) # 取下标 token 的 logp targets input_ids[:, 1:] per_tok logps.gather(-1, targets.unsqueeze(-1)).squeeze(-1) return per_tok * attention_mask[:, 1:] def demo(): c 30.0 class Cfg: final_logits_softcapping c cfg Cfg() logits torch.randn(1, 3, 50) capped apply_final_logits_softcap(logits, cfg) raw apply_final_logits_softcap(logits, type(C, (), {final_logits_softcapping: None})()) print(有配置 - 已软帽:, torch.allclose(capped, c * torch.tanh(logits / c))) print(无配置 - 原样返回:, torch.equal(raw, logits)) if __name__ __main__: demo()核心是apply_final_logits_softcap只在config.final_logits_softcapping存在时套帽否则原样返回兼容不带软帽的模型。这样训练侧 logps 与生成侧分布一致错配消除。六、解决方案第二层软帽在 generate 与 train 两侧共用同一函数第一层修好了训练侧但要保证生成侧和训练侧用的是同一份软帽实现避免以后 generate 路径改了这里没改。第二层把软帽抽成模型/工具层的唯一函数两侧都调用import torch def final_logits_softcap(logits, c): 唯一真源generate 与 train 都调它。 if c is None or c 0: return logits return c * torch.tanh(logits / c) # generate 侧模型 forward 钩子 def generate_logits_hook(module, args, output, c): output.logits final_logits_softcap(output.logits, c) return output # train 侧重算 logps def training_logps(logits, c): return torch.log_softmax(final_logits_softcap(logits, c), dim-1) def demo_parity(): c 30.0 logits torch.randn(2, 4, 100) g final_logits_softcap(logits, c) t final_logits_softcap(logits, c) print(generate 与 train 软帽一致:, torch.equal(g, t)) if __name__ __main__: demo_parity()把软帽收敛成final_logits_softcap这唯一函数generate 走 hook、train 走显式调用二者数学等价old_per_token_logps必然与采样分布对齐从根本上消灭两侧实现漂移。七、解决方案第三层配置断言 一致性测试第三层加护栏若 config 声明了软帽但 trainer 没应用测试应失败并验证 train/logps 与用模型 generate 重算的分布一致import torch def assert_softcap_applied(logits_train, logits_expected_capped): 断言训练侧确实套了软帽与期望的软帽结果一致。 if not torch.allclose(logits_train, logits_expected_capped, atol1e-4): raise AssertionError(训练侧未应用 final_logits_softcappinglogps 会错配) def test_softcap_consistency(): c 30.0 logits torch.randn(1, 5, 64) class Cfg: final_logits_softcapping c cfg Cfg() capped c * torch.tanh(logits / c) applied (lambda x, cfg: (cfg.final_logits_softcapping * torch.tanh(x / cfg.final_logits_softcapping)) if getattr(cfg, final_logits_softcapping, None) else x)(logits, cfg) assert_softcap_applied(applied, capped) print(OK: 训练侧软帽与生成侧一致) def test_no_softcap_model_untouched(): logits torch.randn(1, 5, 64) class Cfg: final_logits_softcapping None cfg Cfg() applied (lambda x, cfg: (cfg.final_logits_softcapping * torch.tanh(x / cfg.final_logits_softcapping)) if getattr(cfg, final_logits_softcapping, None) else x)(logits, cfg) assert torch.equal(applied, logits), 无软帽模型不应被改动 print(OK: 无软帽模型原样返回) if __name__ __main__: test_softcap_consistency() test_no_softcap_model_untouched()两个测试分别锁住有软帽则应用且一致和无软帽则原样任何把软帽漏掉的改动都会被拦下且兼容性不带软帽的模型不受影响得到保证。八、接入 AsyncGRPOTrainer 的建议如果你要在 trainer 里支持软帽模型建议读 configgetattr(config, final_logits_softcapping, None)存在才处理。训练侧套帽old_per_token_logps计算前对 logits 套c * tanh(logits/c)。唯一真源generate 与 train 共用final_logits_softcap函数防漂移。加断言/测试锁住有帽则一致、无帽则原样。验证对齐训练前后打印几个 token 的 logp确认与 generate 采样分布吻合。九、排查清单如果你发现带软帽的模型在 AsyncGRPOTrainer 上训练退化、reward 不收敛按顺序查确认模型 config 是否有final_logits_softcappingGemma 2 / Gpt-OSS 等常见。打印 logps 对比训练算出的 token logp 与 generate 采样反推的 logp 是否一致差大即错配。搜训练侧是否套了软帽old_per_token_logps计算前 logits 是否过c*tanh(logits/c)。确认生成侧软帽generate 路径是否自动套帽通常是模型内部。抽唯一函数generate 与 train 是否共用同一软帽实现。加一致性测试锁住有帽一致、无帽原样。看 KL/优势若开了 KL错配会让 KL 数值失真先修软帽再看。十、小结在AsyncGRPOTrainer上训带final_logits_softcapping的模型时训练静默退化根因是训练侧重算old_per_token_logps没应用模型配置的软帽导致训练分布无帽与生成分布已帽不一致logps 错配优势与 KL 因此失真。它不报错所以极难察觉——reward 不收敛却找不到原因。修复分三层第一层在训练侧前向套c * tanh(logits/c)与 generate 对齐消除错配第二层把软帽抽成final_logits_softcap唯一函数generate 走 hook、train 走显式调用从结构上消灭两侧实现漂移第三层加有帽则一致、无帽则原样的一致性测试与断言作为回归护栏且保证不带软帽的模型完全不受影响。核心心法是模型 config 声明的后处理软帽、scale 等必须在所有使用前向的路径里同步消费否则生成与训练分布不一致问题不报错却会悄悄毒化整轮训练。

相关新闻