Paper: 2609.20807 Authors: Martin Marek, Max Ryabinin Categories: cs.LG
The Gap
Reinforcement learning for large language models (PPO, GRPO, REINFORCE) is bottlenecked by rollout throughput. Generating rollouts from the full-precision training model in native PyTorch is painfully slow. Modern RL pipelines decouple training from inference: high-throughput serving systems (such as vLLM or TensorRT-LLM) generate rollouts using aggressive optimizations—FP8 or INT4 quantization, custom fused kernels, and asynchronous weight updates.
This creates Training-Inference Mismatch (TIM): the rollout distribution deviates subtly but persistently from the training policy .
Conventionally, practitioner folklore attributes RL collapse under TIM to high variance or out-of-distribution token probabilities, attempting to patch it with importance sampling (IS) ratios . But IS ratios explode under quantization discrepancies and require storing per-token logits from the inference engine, adding memory and network bottlenecks.
THE TRAINING-INFERENCE MISMATCH (TIM) TRAP
Rollout Engine (vLLM / FP8 Quant / Async)
| Generates trajectories
v
Divergence / Mismatch
|
v
Training Engine (BF16 PyTorch / Gradient Accumulation)
|
+-------------------------------------------------+
| Conventional View: Mismatch = High Variance |
| Common Fix: Importance Sampling Ratios |
| Consequence: Ratio explosion under quantization |
+-------------------------------------------------+
|
v
PAPER'S CORE INSIGHT: Mismatch = Persistent Systematic DRIFT
Discrepancies don't average out; they accumulate step-over-step
|
v
SOLUTION: Additive Score Centering Correction
Cancels the mean drift directly without unstable ratios
The Increment
One sentence: By identifying persistent gradient drift—rather than random variance—as the root cause of RL failure under training-inference mismatch, the authors derive an additive score-centering term that stabilizes training across 0.6B to 30B parameter scales while bypassing the ratio-explosion problems of importance sampling.
Core Mechanism
Under standard policy gradients, the gradient estimator is:
When , taking the expectation under the rollout distribution introduces a persistent non-zero bias:
Under on-policy conditions (), the expected score function holds identically. But under TIM, this expectation is non-zero and points consistently in the direction of the inference engine’s systematic distortions (e.g. rounding bias in FP8 quantization or stale parameter differences). As training progresses, this error acts like a constant non-zero force pushing parameters off the optimization cliff.
Score Centering subtracts this expected score vector directly:
where tracks the empirical mean of the score function under rollout samples.
SCORE CENTERING VS IMPORTANCE SAMPLING
Raw Policy Gradient under TIM:
Gradient = Signal + [Persistent Systematic DRIFT] -> Training Diverges
Importance Sampling Fix:
Gradient = Signal * [Ratio w] -> w explodes when rollout != train
Score Centering (This Paper):
Gradient = Signal + Drift - [Estimated Drift Center] -> Zero-Mean Robust
The structural metaphor is a rowing boat steering across a river with a constant cross-current.
- The target destination is policy improvement.
- Random turbulence is gradient variance (normal SGD noise).
- The cross-current is TIM drift caused by quantized rollouts: every stroke pushes the boat subtly toward the bank. If you ignore the current, even if your strokes are powerful, you end up miles downstream on the rocks.
- Importance sampling tries to look through high-powered magnifying glasses at every tiny wave, multiplying each stroke’s angle by an erratic compensation factor that frequently violently jerks the oars.
- Score centering simply measures the steady current once and offsets the rudder by that exact angle: a calm, additive counter-torque that cancels the drift while keeping every stroke smooth.
Key Concepts
- Training-Inference Mismatch (TIM): The systematic difference between the execution graph and arithmetic precision used to generate trajectory tokens versus the graph used to compute backward passes.
- Persistent Drift vs. Noise: Zero-mean noise in stochastic gradients averages out over batches. Systematic bias from numerical precision or parameter staleness has a non-zero mean and compounds monotonically across optimization steps.
- Additive vs. Multiplicative Correction: Importance sampling uses multiplicative weights () which suffer from unbounded variance when support differs. Score centering is additive (), guaranteeing stable magnitude regardless of distribution divergence.
Framework Shift
Before (Multiplicative Ratio Compensation):
Quantized Rollout -> Compute per-token logits -> Form ratio p_train / p_rollout
-> Ratios explode under FP8/INT4 -> Optimization collapses unless clipped hard
After (Additive Score Centering):
Quantized Rollout -> Compute standard score vector -> Subtract estimated drift center
-> Additive correction preserves gradient scale -> Stable training from 0.6B to 30B
From “correcting rollout probabilities by multiplying per-token likelihood ratios,” the core shift is centering the gradient score vector additively to cancel systematic engine drift.
Expert Assessment
Problem choice: Critical for real-world scaling. Modern post-training clusters burn hundreds of thousands of GPU hours running vLLM rollouts paired with Megatron/FSDP training. TIM is one of the dirtiest open secrets in LLM post-training, frequently causing silent training divergence or forcing teams onto expensive full-precision rollout setups.
Method maturity: Grounded and mathematically clean. Moving from multiplicative ratios to additive centering of the score function is intuitive yet theoretically sound. The property that score centering composes gracefully with importance sampling when handling stale checkpoints makes it directly applicable to asynchronous distributed RL.
Experimental integrity: Strong breadth of scale. Experiments span small models (0.6B) up to 30B parameters. The evaluation tests both quantization-induced TIM (FP8/INT8) and staleness-induced TIM (delayed actor updates), clearly charting where score centering alone suffices and where hybrid composition with IS provides the best stability.
Writing quality: Transparent and precise. The distinction between variance-induced failure and drift-induced failure is articulated crisply with both theoretical derivations and empirical trajectories.
Verdict: strong accept — A principled, low-overhead fix for a pervasive systems bottleneck in large-scale LLM reinforcement learning.
Takeaways
- If your RL run becomes unstable when switching to quantized rollout engines, do not assume you need smaller learning rates or tighter PPO clipping; suspect systematic TIM drift.
- Implement score centering by tracking running score means during batch evaluation.
- For asynchronous RL with stale rollout actors, compose score centering with moderate importance sampling rather than relying on importance sampling alone.
论文: 2609.20807 作者: Martin Marek, Max Ryabinin 分类: cs.LG
缺口
大语言模型的强化学习阶段(如 PPO、GRPO 或策略梯度方法)最大的吞吐瓶颈在于 Rollout(采样生成)。 如果在 PyTorch 原生训练框架下用全精度模型生成 Rollout,速度会慢到无法忍受。 因此,工业界通行的做法是将采样与训练解耦:利用 vLLM 或 TensorRT-LLM 等推理引擎,配合 FP8/INT4 量化、算子融合以及异步参数同步等极致优化来狂飙采样吞吐。
这必然引发训推不一致(Training-Inference Mismatch,简称 TIM):用于采样的推理引擎 与实际进行反向传播的训练模型 之间存在细微但持续的数值偏差。
长期以来,工程经验普遍认为 TIM 导致的训练崩盘是因为「方差过大」或「偏离分布」,因而常规做法是使用重要性采样(Importance Sampling, IS)权重 进行再平衡。 但在量化精度不一致时,IS 权重极易发散爆炸,且它强制要求推理端保存并回传每个 token 的 logit,带来了巨大的显存与网络带宽开销。
训推不一致(TIM)的工程困局
推理引擎 (vLLM / FP8量化 / 异步更新)
| 产生采样轨迹
v
数值与精度系统偏差
|
v
训练引擎 (BF16 PyTorch / 梯度反向传播)
|
+-------------------------------------------------+
| 传统认知:TIM = 采样带来了高方差噪声 |
| 常见解法:使用重要性采样(IS)比值乘法修正 |
| 致命代价:量化偏差下 IS 比值爆炸,强行裁剪又失真 |
+-------------------------------------------------+
|
v
本文核心洞见:TIM = 持续累积的系统性「漂移(Drift)」
推理差异在单步微不足道,但每一步方向一致,不断累加
|
v
全新解法:加性「得分中心化(Score Centering)」修正
直接对齐并抵消均值漂移,彻底避开乘法比值的数值发散
增量
一句话: 本文揭示了训推不一致导致强化学习崩坏的元凶是长期累积的系统性梯度漂移而非随机方差,并推导出一种加性「得分中心化」修正项,在 0.6B 到 30B 参数规模下全面稳定了训练,且摆脱了重要性采样比值爆炸的宿疾。
核心机制
在标准的策略梯度更新中,梯度估计量为:
在理想的同策略(On-policy)情形下,期望得分函数恒为零:。 但当存在 TIM 时,在采样分布下的期望值不再为零:
这一非零分量并不具备零均值的随机抵消特性,而是严格沿着推理引擎量化舍入或参数迟滞的系统性偏差方向持续存在。 随着训练步数推移,这个持续存在的系统性漂移力就像一股看不见的恒力,硬生生把模型参数推下性能崩塌的悬崖。
得分中心化(Score Centering) 的做法极其干净优雅:通过估计该得分向量在采样分布下的均值 ,直接将其从梯度中加性扣除:
得分中心化与重要性采样的机制对比
原始 TIM 策略梯度:
梯度 = 真实信号 + [持续累积的系统性漂移] -> 训练发散崩溃
传统重要性采样修正:
梯度 = 真实信号 * [概率比值 w] -> 量化下比值 w 剧烈发散
本文得分中心化修正:
梯度 = 真实信号 + 漂移 - [估计的漂移中心] -> 零均值稳态
这里的核喻是在存在固定侧风的河道上划皮划艇。
- 终点是策略的最优状态。
- 水面的微小水花是正常的 SGD 梯度随机噪声。
- TIM 漂移则是侧向的恒定水流:量化带来的细微舍入偏差让皮划艇每划一桨都往右舷偏转一厘米。如果不抵消这股暗流,任凭你划桨力道再大,最终也会被冲上几公里外的浅滩。
- 重要性采样就像是一个神经过敏的舵手,每划一桨都用高倍放大镜去测量每个细小水花,随后猛打方向舵(乘法调节),导致小艇左右剧烈晃荡,甚至直接翻船。
- 得分中心化则是稳重的老舵手:测算出侧流的稳定推力后,直接在船舵上加装一块固定配重的反向偏压板(加性扣除中心均值),水流直接被抵消,小艇始终平稳向前。
关键概念
- 训推不一致(TIM):在大模型后训练过程中,为了提升吞吐,前向生成采用的算子、量化级别或权重版本,与反向传播更新采用的计算图存在的不一致现象。
- 系统性漂移 vs. 随机噪声:零均值的随机噪声能通过增大 Batch Size 自然平均平滑;但由数值精度引发的偏差在参数空间具有确定方向,会在多步更新后像复利一样线性或指数级累积。
- 加性修正的鲁棒优势:重要性采样属于乘法修正(),只要支撑集存在微小缝隙,比值就会趋于无穷大;得分中心化属于加性修正(),具有天生的数值上界稳定性。
框架转变
之前(基于乘法概率比值的修补方案):
量化推理采样 -> 保存并回传全序列 logits -> 计算 p_train / p_rollout 比值
-> 比值因量化截断频繁爆炸 -> 必须进行粗暴截断,导致梯度估计失效
之后(基于加性得分中心化的均值校正):
量化推理采样 -> 计算标准得分函数 -> 加性减去批次经验中心漂移量
-> 无需保存推理端 logits,无比值发散风险 -> 从 0.6B 到 30B 均平稳收敛
从「死磕每个 token 的概率比值乘法平衡」,核心转变在于:直接通过在梯度层面加性平移中心,一举消解训推引擎之间的系统性几何漂移。
专家评审
选题眼光: 极具现实工业价值。 当前几乎所有一线大模型实验室都在为 RL 采样阶段昂贵的算力开销发愁,TIM 是工业级后训练集群中最常见却又最棘手的「暗礁」。 把研究对准 TIM 的机理与解法,抓住了工程落地的咽喉。
方法成熟度: 理论推导清晰优美。 抛弃了学术界惯性套用重要性采样的路径依赖,从得分函数零期望的核心性质切入,洞察到「非零漂移」这一本质。 方法为纯加性形式,计算开销可忽略不计,且能无缝与异步分布式系统结合。
实验诚意: 模型跨度扎实(0.6B 至 30B)。 全面覆盖了量化精度不一致(FP8/INT8)与异步参数陈旧(Staleness)两类最核心的 TIM 场景,量化深入、对比公平。
写作功力: 逻辑链条非常紧凑,对误差为何累积的数学直觉解释得非常透彻。
判决: 强接收 (strong accept) — 强化学习系统工程领域的代表性硬核好文,直击核心痛点且解法极简高效。
要点总结
- 当大模型在开启低比特量化引擎做 RL 采样出现训练崩塌时,不要盲目调小学习率或加重裁剪,应当首先排查系统性 TIM 梯度漂移。
- 在训练步中轻量维护得分函数的滑动均值 ,即可零成本实施得分中心化。
- 面对异步分布式 RL 中的高延迟参数陈旧问题,将得分中心化与适度截断的重要性采样结合使用,能取得最佳的稳定收敛效果。