When Do Larger Batches Help Scale LLM Reinforcement Learning?
Ziniu Li, Jinbo Wang, Guanhua Huang, Feiyuan Zhang, Pengbo Li, Alex Chen
cs.LG, cs.AI
2026-08-29
腾讯混元在固定硬件上拆开LLM RL的batch效应。Adam按平方根重调学习率后,GRPO把prompt batch从128提到1024,到达77%目标的墙钟从11.90小时降到8.42小时,约省29%;学习率钉死再翻倍batch,墙钟反而多42%。
LLM 做强化学习后训练时,把 batch 做大几乎是默认动作:一步吃更多 rollout,梯度更稳,机器也更容易喂饱。监督训练里这条路通常要加卡,数据已经在手里。LLM RL 不一样。样本得先自回归生成出来,再拿去更新。batch 变大,每一步消耗的新样本也变多,生成和训练都可能更慢。墙钟上更快还是更慢,并不清楚。
腾讯混元把问题钉在固定硬件上:不加加速器,只把 batch 调大,什么时候能缩短到达目标分数的时间。他们用 GRPO 训 Qwen3-30B-A3B-Instruct-2507,用 PPO 训内部混元 MoE(3B 激活参数),把算法效应和系统效应拆开比。
墙钟进度拆成两项相乘:每消耗一个样本,验证分涨多少;单位时间能消耗多少样本。固定目标分数后,大 batch 更快的条件很硬:吞吐增益必须盖过到达目标所需样本的惩罚比 rN。吞吐涨 1.2 倍、样本却要多花 1.5 倍,墙钟照样变差。
要让 rN 接近 1,他们测的是 batch-size invariance:按 batch 重调超参之后,用累计样本而不是更新步数画学习曲线,不同 batch 应该大致叠在一起。实验里只动 Adam 学习率,按 Malladi 等人的平方根规则初始化,η(B)=η0√(B/B0),其余超参不动。系统侧吃的是生成和训练的不对称。解码在低并发时算术强度低,常常卡在内存带宽上,权重要按步从显存流一遍;并发序列变多,这次搬运摊到更多 token 上,生成时间对 batch 经常亚线性。训练是大矩阵乘,算量大致跟 token 数成正比。固定硬件上把 rollout 并发做大,有机会提高端到端吞吐,而不必加卡。
具体设定:
平方根重调学习率之后,GRPO 在 P=64 到 1024 之间,按累计训练回复画的曲线大致对齐;P=2048 和 4096 开始偏离。PPO 在 B=256 到 2048 对齐,B=4096 同样偏离。对齐区间里 actor 梯度范数接近 B^{-1/2} 的噪声主导缩放,GRPO 拟合斜率 -0.468,PPO -0.471。再翻倍就扁了:GRPO 从 P=2048 到 4096 中位梯度范数只降到 0.823 倍,PPO 从 2048 到 4096 只降到 0.905 倍,理论值是 0.707 倍。
学习率钉死会直接打破对齐。P 从 128 加倍到 256、η 仍是 10^{-6},到达目标所需样本多 67%。理想不变性预测更新步数减半;缩放学习率那组落在 0.50–0.67 倍,固定学习率那组只有 0.75–0.83 倍。
生成侧,固定硬件确实吃得到亚线性。PPO 把训练 batch 从 256 提到 1024,每步回复数乘 4,边界收集时间从 39 秒涨到 68 秒(1.74 倍),生成吞吐 2.29 倍。GRPO 从 P=128 到 512,回复乘 4,收集时间 2.93 倍,吞吐只涨到 1.36 倍;actor 更新时间几乎线性,P=128 到 256 从 101.3 秒到 208.6 秒(2.06 倍)。
墙钟会计钉在 J=77%:
| P | 学习率 | 保留样本(K) | 墙钟(h) | rN | 相对吞吐 |
| 128 | 1×10^{-6} | 122.88 | 11.90 | 1.00 | 1.00 |
| 256 | √2×10^{-6} | 122.88 | 9.54 | 1.00 | 1.25 |
| 1024 | 2√2×10^{-6} | 122.88 | 8.42 | 1.00 | 1.41 |
| 2048 | 4×10^{-6} | 196.61 | 14.68 | 1.60 | 1.30 |
| 256 固定 LR | 1×10^{-6} | 204.80 | 16.93 | 1.67 | 1.17 |
P=1024 是测到的最快点,相对 P=128 约 0.71 倍墙钟,省 29%。P=2048 之后 rN 跳到 1.60,吞吐还在 1.3 倍附近,墙钟变成 1.23 倍。固定学习率的 P=256 吞吐 1.17 倍,样本惩罚 1.67 倍,墙钟 1.42 倍。
GRPO 里真正 batch 单位是 P×G,不是 prompt 数。G=8、P=128 和 G=16、P=64 对齐在同一条累计回复曲线上。PPO 的 critic 梯度范数几乎不随 batch 下降,actor 和 value model 可能各有一套临界 batch。
可落地的判断是:先按平方根把 Adam 学习率跟着 batch 改掉,画出按累计样本对齐的曲线;再加大生成并发,直到吞吐增益盖不过样本惩罚。混元写成 Align, then accelerate。
这不是 batch 越大越好。固定硬件上 LLM RL 比监督训练更有机会,因为解码经常没喂饱;过了训练侧临界 batch,样本效率塌掉,吞吐再涨也救不回来。跟加卡保步时的大 batch 不是同一件事:这里卡数锁死,快慢全看吞吐和 rN 谁赢。
比较不同 P 或 G 时不要按 step 数对齐,按保留下来的训练回复数对齐。G 和 P 在测过的范围内可以互换,关键是乘积。学习率钉死再把 batch 翻倍,就是更慢的那条 16.93 小时曲线。
超参只动了学习率。clip、KL、温度、过滤阈值都没按 batch 重调,不变性失败时分不清是统计饱和还是调参不够。Hilton 等人给 PPO 拆过 EWMA 近端项才做出更强不变性;这篇没引入那套构造,只在单次全局更新加异步 partial rollout 的设定里做迁移测试。
不变区间是局部的。换模型、数据、奖励或训练阶段都要重测。G 只试了 8 和 16,PPO 的 critic 临界 batch 没做完。生成临界 batch 绑在模型、解码引擎、回复长度和硬件上,2.29 倍和 29% 不能当可迁移常数。
会计口径也有选择。目标钉在 77%,K 取某个验证 checkpoint;P=256 缩放学习率那组验证分 76.72,略低于目标,仍被算进 rN=1.00。硬件配置和卡数正文没有展开,吞吐数字换集群无法直接复现。两边任务都偏数学推理,更长回复或工具调用会改写生成侧的亚线性区间。