【Bug已解决】Does PPOTrainer use mini_batch_size to update parameters 解决方案

发布时间:2026/7/21 22:01:53

【Bug已解决】Does PPOTrainer use mini_batch_size to update parameters 解决方案 【Bug已解决】Does PPOTrainer use mini_batch_size to update parameters 解决方案一、现象长什么样在用PPOTrainer做 RLHF 时我们按文档设了mini_batch_size比如 rollout batch 是 64mini_batch_size 是 16期望每个 rollout 内做 4 次小批量梯度更新。但观察显存和步数后发现有问题期望每个 rollout(64) 拆成 4 个 mini_batch(16)各做一次 .step() 实际好像只做了 1 次大更新显存峰值等于整批 64具体现象设了mini_batch_size但显存峰值和不设这个值、用整批几乎一样 → 怀疑没生效训练步数统计显示每个 rollout 只更新 1 次而不是batch_size / mini_batch_size次当batch_size很大、mini_batch_size很小时直接 OOM说明更新时根本没按 mini_batch 切。于是核心疑问也是 issue 的标题PPOTrainer 到底有没有用mini_batch_size来切分更新答案是——当时的实现里step()直接对整个 rollout batch 做了一次前向反向mini_batch_size只被用来决定收集多少样本却没被用来分几次更新导致显存和多次小步更新的预期全部落空。二、背景PPO 的标准训练循环是收集 rollout用旧策略跑batch_size条样本prompt response 优势多轮小批量更新把这batch_size条样本打乱后切成mini_batch_size的小块每个小块做一次策略梯度更新可选多个 epoch重复。mini_batch_size的意义在于梯度更新时的实际 batch 大小它决定了显存峰值和优化稳定性。如果实现正确地切了那么即使batch_size64只要mini_batch_size16显存峰值就只对应 16 条且每个 rollout 做 4 次更新更好地利用样本、更稳。但PPOTrainer当时把mini_batch_size和收集批次大小混为一谈step()接收的就是已经收集好的整批内部直接loss.backward()一次没有任何按 mini_batch_size 再切的逻辑。于是mini_batch_size形同虚设对更新无影响显存峰值由batch_size决定而非mini_batch_size每 rollout 多步更新的期望落空。三、根因根因一句话PPOTrainer.step()把传入的整个 batch 当作一次更新的单位没有内部按mini_batch_size切分做多次小批量梯度更新mini_batch_size只控制了样本收集量没有控制更新粒度导致显存和更新次数都背离预期。具体无切分逻辑step(batch)内直接model(batch); loss.backward(); optimizer.step()batch 多大就一次更新多大。参数语义混淆batch_size收集与mini_batch_size更新被当成一回事或后者被忽略。epoch 循环缺失即使想做多个优化 epoch也因为没切 mini_batch 而无从下手。显存随 batch 线性增长大 batch 直接 OOMmini_batch_size 本应兜住却没兜住。本质是优化循环的 mini-batch 切分这一步被整体跳过了。四、最小可运行复现下面用纯 Python 模拟有没有按 mini_batch_size 切分对更新次数/显存的影响def ppo_step_no_split(batch, mini_batch_size): 旧实现忽略 mini_batch_size整批一次更新。 updates 1 peak len(batch) return updates, peak def ppo_step_with_split(batch, mini_batch_size): 正确实现按 mini_batch_size 切分多次更新。 updates 0 peak 0 for i in range(0, len(batch), mini_batch_size): mb batch[i:i mini_batch_size] updates 1 peak max(peak, len(mb)) return updates, peak def demo(): batch list(range(64)) u1, p1 ppo_step_no_split(batch, 16) u2, p2 ppo_step_with_split(batch, 16) print(f旧实现更新 {u1} 次, 峰值 batch{p1}) print(f正确实现更新 {u2} 次, 峰值 batch{p2}) if __name__ __main__: demo()输出旧实现更新 1 次, 峰值 batch64 正确实现更新 4 次, 峰值 batch16第一行就是 bug设了 mini_batch_size16 却只更新 1 次、峰值 64第二行才是预期——4 次更新、峰值 16。复现了mini_batch_size 没生效的核心差异。五、解决方案第一层在 step 内按 mini_batch_size 切分更新第一层给step()加上切分逻辑让mini_batch_size真正控制更新粒度import torch from typing import List, Dict def split_minibatches(batch: Dict[str, torch.Tensor], mini_batch_size: int): n batch[input_ids].shape[0] for i in range(0, n, mini_batch_size): yield {k: v[i:i mini_batch_size] for k, v in batch.items()} def ppo_step_fixed(trainer, batch, mini_batch_size, epochs1): 按 mini_batch_size 切分做 epochs 轮小批量更新。 total_updates 0 for _ in range(epochs): for mb in split_minibatches(batch, mini_batch_size): loss trainer.forward(mb) loss.backward() trainer.optimizer.step() trainer.optimizer.zero_grad() total_updates 1 return total_updates def demo(): batch {input_ids: torch.zeros(64, 4)} # 伪 trainer class T: def forward(self, mb): return mb[input_ids].sum() optimizer type(O, (), {step: lambda s: None, zero_grad: lambda s: None})() updates ppo_step_fixed(T(), batch, mini_batch_size16, epochs1) print(实际更新次数, updates, (应为 4)) if __name__ __main__: demo()核心是split_minibatchesstep()不再是整批一次而是按mini_batch_size切成若干小批各做一次反向更新。显存峰值降到 mini_batch 大小更新次数变成batch_size / mini_batch_size。六、解决方案第二层明确 batch_size 与 mini_batch_size 的语义分离第一层加了切分但要防止参数语义再次混淆。第二层在配置和文档层把两者厘清并加校验from dataclasses import dataclass from typing import Optional dataclass class PPOConfig: batch_size: int 64 # 每次 rollout 收集的样本数 mini_batch_size: int 16 # 每次梯度更新的样本数 ppo_epochs: int 1 # 每个 rollout 内重复优化的轮数 def __post_init__(self): if self.mini_batch_size 0: raise ValueError(mini_batch_size 必须 0) if self.batch_size % self.mini_batch_size ! 0: # 不允许不能整除避免最后一块大小不一导致形状问题 raise ValueError( fbatch_size({self.batch_size}) 必须能被 fmini_batch_size({self.mini_batch_size}) 整除 ) def expected_updates_per_rollout(cfg: PPOConfig) - int: return (cfg.batch_size // cfg.mini_batch_size) * cfg.ppo_epochs def demo(): cfg PPOConfig(batch_size64, mini_batch_size16, ppo_epochs2) print(每 rollout 期望更新次数, expected_updates_per_rollout(cfg), (64/16*28)) if __name__ __main__: demo()两者语义分离batch_size收集量mini_batch_size更新量ppo_epochs重复轮数校验整除避免最后一块形状不一致expected_updates_per_rollout给出明确预期便于监控是否真的做了这么多次更新。七、解决方案第三层监控更新次数 不变量测试第三层加监控与测试确保mini_batch_size 真的生效这件事可被观测、可回归from typing import List class UpdateCounter: def __init__(self): self.count 0 def step(self, mb): # 真实场景里这里做 backwardstep self.count 1 def train_with_monitor(batch_size, mini_batch_size, epochs, counter: UpdateCounter): for _ in range(epochs): for i in range(0, batch_size, mini_batch_size): counter.step(i) def test_minibatch_effective(): cfg PPOConfig(batch_size64, mini_batch_size16, ppo_epochs1) c UpdateCounter() train_with_monitor(cfg.batch_size, cfg.mini_batch_size, cfg.ppo_epochs, c) expected expected_updates_per_rollout(cfg) assert c.count expected, f更新次数 {c.count} ! 期望 {expected}mini_batch_size 未生效 print(fOK: 实际更新 {c.count} 次 期望 {expected} 次) if __name__ __main__: test_minibatch_effective()UpdateCounter记录真实更新次数测试断言它等于batch_size/mini_batch_size*epochs。任何把切分逻辑改回整批一次的改动都会让断言失败CI 直接拦下——把mini_batch_size 是否生效从靠猜变成可观测、可回归。八、落地建议如果你在 PPOTrainer 上确认 mini_batch_size 没生效建议改 step内部按mini_batch_size切分多次backwardstep。分离语义batch_size(收集) 与mini_batch_size(更新) 在 config 里明确分开并校验整除。支持 ppo_epochs每个 rollout 可重复多轮小批量优化。加监控打印每 rollout 实际更新次数对照batch_size/mini_batch_size*epochs。加测试锁住更新次数 期望防回归。显存验证设小 mini_batch_size 后显存峰值应下降作为生效证据。九、排查清单如果你怀疑 PPOTrainer 没用 mini_batch_size按顺序查看更新次数每 rollout 实际.step()几次应等于batch_size/mini_batch_size*epochs。看显存峰值设小 mini_batch_size 后峰值是否下降不降则说明没切分。搜 step 内部是否直接对整批loss.backward()没有split_minibatches。确认参数语义batch_size与mini_batch_size是否被混淆或后者被忽略。加 ppo_epochs是否需要每个 rollout 多轮优化切分后才能做。加更新次数监控/测试锁住更新次数期望。校验整除避免最后一块形状不一致导致形状错误。十、小结PPOTrainer设了mini_batch_size却不生效根因是**step()把整个 rollout batch 当作一次更新的单位没有内部按mini_batch_size切分做多次小批量梯度更新**——mini_batch_size只控制了样本收集量没控制更新粒度。结果是显存峰值由batch_size决定大 batch 直接 OOM、每个 rollout 只更新 1 次与多次小步的预期相悖小批量更新形同虚设。修复分三层第一层在step()内加split_minibatches按mini_batch_size切分、各做一次反向更新显存峰值降到 mini_batch 大小第二层在 config 里把batch_size(收集) 与mini_batch_size(更新) 语义分离并校验整除支持ppo_epochs多轮优化第三层加更新次数监控与更新次数期望不变量测试让 mini_batch_size 是否生效变得可观测、可回归。核心心法是mini_batch_size控制的是梯度更新的实际 batch 大小不是收集多少样本——任何 PPO 实现都必须把它落实到优化循环里的切分逻辑否则它只是一个没有作用的配置项。

相关新闻