行业资讯

【Bug已解决】Optimize old_per_token_logps recomputation in GRPOTrainer: per-rollout window check instead

发布时间:2026/7/22 7:42:58
【Bug已解决】Optimize old_per_token_logps recomputation in GRPOTrainer: per-rollout window check instead 【Bug已解决】Optimize old_per_token_logps recomputation in GRPOTrainer: per-rollout window check instead of static modulo 解决方案.md原始报错Optimize old_per_token_logps recomputation in GRPOTrainer: per-rollout window check instead of static modulo 场景GRPOTrainer 需要周期性地重新计算旧策略的逐 token 对数概率old_per_token_logps用作 PPO/GRPO 优势计算里的参考。现在的实现用静态取模每训练 N 步step % N 0就全体重算一次。但很多 rollout 在这 N 步里根本没被更新过/没必要重算全体重算浪费算力而且静态取模不考虑哪些 rollout 还有效、哪些已过期。优化方向是改成按 rollout 的窗口检查只在某个 rollout 超出有效窗口或确实变旧时才重算它而非无差别全体重算。 关键词old_per_token_logps、重计算优化、静态取模、per-rollout 窗口、参考策略、惰性重算、步数调度、算力浪费。一、现象长什么样重算太勤、且不分对象每隔固定 N 步step % N 0训练器把所有rollout 的 old logps 重新算一遍但很多 rollout 在这 N 步里早就被新采样覆盖了旧值本就该丢重算它们没意义另一些 rollout 还在有效窗口内、没被新数据取代本不用重算却被算了全体重算让每 N 步出现一次算力尖峰训练被周期性拖慢静态取模不考虑每个 rollout 自己的新旧程度粒度太粗表现训练整体变慢且重算的算力大比例是浪费在已经无效/无需更新的 rollout 上。核心问题重算策略用全局静态步数取模代替按 rollout 个体有效性判断粒度粗、浪费大。二、背景为什么全体每 N 步重算是浪费old_per_token_logps 是参考策略旧策略对当前 batch 的对数概率用于计算重要性比π_new / π_old。它的有效性取决于这批 rollout 是否还是当前参考策略产生的。合理的重算时机应该是当参考策略本身被更新了训练步推进让 π_old 变了旧的 logps 才失效需要重算或者某个 rollout 超出了有效窗口太老不再代表当前策略需要重算或丢弃。而全局step % N 0全体重算是一种粗粒度近似它不管每个 rollout 的实际新旧一律到点全算。这有两个浪费无效 rollout 也被算已经被新采样覆盖的 rollout算它的 old logps 毫无意义有效 rollout 被迫跟着算还在窗口内的 rollout 其实不用算但被全体带了进去。优化就是把这个全局定时全体重算改成逐 rollout 判断是否需要重算——只在某 rollout 确实超出窗口/参考策略已变时才算它。这就是per-rollout window check。三、根因重算用全局静态取模不区分 rollout 个体根因拆解静态取模if step % N 0: recompute_all()粒度全局全体重算到点把所有 rollout 一起算含无效/无需更新的无个体状态没记录每个 rollout 的产生步/最后有效步无法判断个体新旧算力尖峰每 N 步一次全体重算周期性拖慢浪费大重算算力多数花在已无效 rollout 上窗口概念缺没有有效窗口概念无法按窗口淘汰/重算。下面用最小模型复现全体每 N 步重算含无效 rollout再给逐 rollout 窗口检查的修复。四、最小可运行复现class Rollout: def __init__(self, rid, born_step): self.rid rid self.born_step born_step # 产生时的步数 self.old_logps None def recompute_all_static(rollouts, step, N, window): 错误step % N 0 时全体重算不管个体是否有效。 if step % N 0: for r in rollouts: r.old_logps fcomputed{step} # 含已超出窗口的无效 rollout return sum(1 for r in rollouts if r.old_logps) if __name__ __main__: rollouts [Rollout(r1, 0), Rollout(r2, 0), Rollout(r3, 95)] # 在 step100, N50 时全体重算但 r3 在 step95 产生step100 仍在 window10 内本不需算 n recompute_all_static(rollouts, 100, N50, window10) print(全体重算数量(含不需算的):, n) # 3浪费运行可见 step100 把 3 个 rollout 全算但 r3 还在窗口内本不必算——浪费现场。五、方案逐 rollout 窗口检查只重算超窗/失效的第一层每个 rollout 记录产生步重算时逐个判断是否超出有效窗口只重算超窗的def recompute_per_rollout(rollouts, step, window): 正确逐 rollout 判断是否超窗只重算失效的。 recomputed 0 for r in rollouts: age step - r.born_step if age window: # 超出有效窗口 - 需重算 r.old_logps fcomputed{step} r.born_step step # 刷新产生步 recomputed 1 return recomputed if __name__ __main__: rollouts [Rollout(r1, 0), Rollout(r2, 0), Rollout(r3, 95)] n recompute_per_rollout(rollouts, 100, window10) print(逐 rollout 重算数量(仅超窗的):, n) # 2r1,r2 超窗r3 在窗内逐 rollout 窗口检查只重算真正失效的浪费大减。六、方案参考策略变更时标记失效而非定时全体第二层除了窗口还在参考策略被更新时标记相关 rollout 失效下次用到才惰性重算lazyclass RolloutStore: def __init__(self, window): self.window window self.rollouts {} self.ref_version 0 def mark_stale_on_ref_update(self): self.ref_version 1 # 参考策略变了 for r in self.rollouts.values(): r.stale True # 全部标记失效但暂不算 def get_logps(self, rid, step, compute_fn): r self.rollouts.setdefault(rid, Rollout(rid, step)) age step - r.born_step if getattr(r, stale, False) or age self.window: r.old_logps compute_fn(r) # 惰性重算仅失效/超窗时 r.born_step step r.stale False return r.old_logps if __name__ __main__: store RolloutStore(window10) store.rollouts[r1] Rollout(r1, 0) store.mark_stale_on_ref_update() # 参考策略更新r1 失效 # 用到 r1 时才重算惰性未用到的不浪费 print(惰性重算:, store.get_logps(r1, 100, lambda r: logps))参考策略变更标记失效 惰性重算没被用到的 rollout 绝不浪费算力。七、方案重算计数与预算避免尖峰第三层把全体定时重算的算力尖峰改成分散的按需重算并用预算限制每步重算数量平滑开销def recompute_budgeted(store, step, budget): 每步最多重算 budget 个失效 rollout平滑算力。 stale [r for r in store.rollouts.values() if getattr(r, stale, False) or (step - r.born_step) store.window] done 0 for r in stale: if done budget: break r.old_logps fcomputed{step} r.born_step step r.stale False done 1 return done if __name__ __main__: store RolloutStore(window10) for i in range(5): store.rollouts[fr{i}] Rollout(fr{i}, 0) store.mark_stale_on_ref_update() # 每步最多重算 2 个分多步平滑不再一次性尖峰 print(分步重算(预算2):, [recompute_budgeted(store, 100, 2) for _ in range(3)])预算限制把每 N 步全体尖峰摊成每步少量训练更平稳。八、验证把只重算失效 rollout锁进测试def test_only_stale_recomputed(): rollouts [Rollout(r1, 0), Rollout(r2, 0), Rollout(r3, 95)] n recompute_per_rollout(rollouts, 100, window10) assert n 2 # 只 r1,r2 超窗 def test_window_keeps_fresh(): r3 Rollout(r3, 95) rollouts [r3] recompute_per_rollout(rollouts, 100, window10) assert r3.old_logps is None # 在窗内不被重算 def test_budget_smooths(): store RolloutStore(window10) for i in range(5): store.rollouts[fr{i}] Rollout(fr{i}, 0) store.mark_stale_on_ref_update() total sum(recompute_budgeted(store, 100, 2) for _ in range(3)) assert total 5 if __name__ __main__: test_only_stale_recomputed() test_window_keeps_fresh() test_budget_smooths() print(old logps 按需重算测试通过。)九、排查清单old logps 重算浪费按顺序查静态取模是否step % N 0全体重算是则粒度太粗。全体重算到点是否所有 rollout 一起算含无效/无需更新的。个体状态是否记录每个 rollout 的产生步/有效性没记录则无法判断个体。窗口概念是否有有效窗口淘汰机制没有则无法按需。参考策略变更参考策略更新时是否标记失效、惰性重算还是定时全体算力尖峰是否每 N 步一次全体重算拖慢训练是则改分散。预算平滑是否用预算限制每步重算数有则开销平稳。十、小结old_per_token_logps 重算用静态取模太浪费是重算策略用全局step % N全体重算不区分 rollout 个体有效性很多 rollout 已无效或仍在窗口内无需重算却被无差别一起算造成周期性算力尖峰与大量浪费**。修复三层逐 rollout 窗口检查记录每个 rollout 产生步只重算超出有效窗口的失效标记 惰性重算参考策略更新时标记失效用到时才重算未用的不浪费预算平滑每步限制重算数量把全体尖峰摊成分步按需训练更稳。核心原则周期性重算旧策略 logps不应是全局静态取模的全体重算而应是基于每个 rollout 的新旧/窗口判断的按需重算。凡是每 N 步step % N全体 recompute的写法都应改为 per-rollout 窗口检查 惰性重算 预算平滑——只算真正失效的算力不被浪费。