Papers
Topics
Authors
Recent
Search
2000 character limit reached

Diagnosing Training Inference Mismatch in LLM Reinforcement Learning

Published 14 May 2026 in cs.LG, cs.AI, and cs.CL | (2605.14220v1)

Abstract: Modern LLM RL systems separate rollout generation from policy optimization. These two stages are expected to produce token probabilities that match exactly. However, implementation differences can make them assign different values to the same sequence under the same model weights, inducing Training-Inference Mismatch (TIM). TIM is difficult to inspect because it is entangled with off-policy drift and common stabilization mechanisms. In this work, we isolate TIM in a zero-mismatch diagnostic setting (VeXact), and show that small token-level numerical disagreements can independently cause training collapse. We further show that TIM changes the effective optimization problem, and identify a set of remedies that could mitigate TIM. Our results suggest that TIM is not benign numerical noise, but a systems-level perturbation that should be treated as a first-order factor in analyzing LLM RL stability.

Summary

  • The paper demonstrates that training-inference mismatch alone can trigger full RL collapse, with MoE validation reward falling to 0.067 under vLLM while VeXact reaches 0.534.
  • The paper introduces VeXact, a deterministic rollout engine that achieves bit-wise alignment with FSDP and isolates numerical discrepancies from off-policy drift and other causes.
  • The paper finds that correction-ratio rejection combined with token-level importance sampling best approximates zero-mismatch training, while KL metrics may fail to detect degradation early.

Overview

This paper investigates Training-Inference Mismatch (TIM) as a first-order cause of instability in LLM reinforcement learning. TIM arises because training engines (FSDP, Megatron) and inference engines (vLLM, SGLang) implement the same model with different kernels and reduction orders, so that under identical weights and inputs they assign different token probabilities. The authors formalize this as a token-level discrepancy δt=log⁡πoldtrain(at∣st)−log⁡πoldrollout(at∣st)\delta_t = \log \pi^{\text{train}}_{\text{old}}(a_t|s_t) - \log \pi^{\text{rollout}}_{\text{old}}(a_t|s_t), which is objective-agnostic: it exists prior to any choice of REINFORCE, PPO, or GRPO.

The central methodological contribution is VeXact, a lightweight rollout engine built on VeRL that achieves bit-wise alignment of token log-probabilities with the FSDP trainer. VeXact eliminates both sources of mismatch — divergent model/kernel implementations and non-deterministic or batch-dependent kernel behavior — by reusing the HuggingFace model implementation inside FSDP and employing deterministic, batch-invariant kernels (RMSNorm, batched matmul, fused MoE, attention without KV splitting). Performance is retained through chunked prefill, CUDA graphs, pipeline parallelism, and optimistic KV allocation. This zero-mismatch baseline enables causal attribution of RL collapse to TIM alone, which prior work could not isolate from off-policy drift and stabilization mechanisms.

Isolating TIM's impact

The diagnostic design pairs VeXact against vLLM under REINFORCE with batch-whitened advantages, chosen deliberately because its single-update-per-batch structure avoids PPO ratio clipping that could mask TIM effects. Experiments cover Qwen3-1.7B (dense) and Qwen3-30B-A3B (MoE), trained on Sanity-Test-R1D-1.5B and DAPO respectively, evaluated on AIME 2024.

The results are stark. In the MoE setting, the vLLM run improves initially but degrades after step 280, with training reward falling from 0.574 to 0.255 and validation reward from 0.293 to 0.067; the VeXact run continues improving to 0.753 training and 0.534 validation reward. Since TIM is the only difference between the two configurations, the paper concludes that TIM by itself can trigger full training collapse — it is not benign numerical noise but an infrastructure-level perturbation that changes the effective optimization problem. Token-level inspection shows why small aggregates mislead: mean ∣δt∣|\delta_t| per batch is tiny, yet maximum values approach 1.0, including cases where the argmax token flips entirely (e.g., the training side preferring ":\n\n" over "that" at one position).

Failure modes in GRPO: recomputation versus bypass

The paper then dissects two standard ways of obtaining πold\pi_{\text{old}} in PPO/GRPO pipelines. Under recomputation, the trainer re-evaluates sampled tokens, making the PPO denominator πoldtrain\pi^{\text{train}}_{\text{old}} rather than the distribution that actually generated the data. Under bypass, the rollout engine transmits its own log-probabilities directly. Both instantiate the same clipped surrogate but with different denominators.

Empirically, VeXact holds training reward near 0.93 while vLLM recomputation degrades from ~0.87 to ~0.40 within 650 steps, partially recovers, then collapses to near-zero after step ~1665; bypass shows single-stage degradation to ~0.4 without collapsing. Notably, the failure signals decouple across metrics: bypass degrades reward without synchronized loss spikes, and recomputation enters degradation before gradient-norm anomalies appear.

Two analytical findings follow:

  • KL estimators are insufficient early-warning signals. In recomputation mode, K1K_1 and K3K_3 estimators remain flat for the first 700 steps while reward is already failing. The paper attributes the stealthiness to the zero-centered loss contribution C(r)=−(r−1)AC(r) = -(r-1)A: TIM skews advantage-weighted contributions asymmetrically across positive and negative advantages, converting symmetric numerical noise into a sign-imbalanced, non-zero-mean distortion of gradients before aggregate probability-space divergence becomes visible.
  • Bypass fails through optimization-space misalignment. Even when the correct behavioral probabilities are used as denominators, the numerator and score-function gradients are computed on the trainer's numerical path. Updates exploit artifacts of the trainer's forward pass that do not translate to behavioral improvement under the rollout engine, producing silent policy degradation.

Ablating algorithmic compensations

Using VeXact as ground truth, the paper evaluates post-hoc corrections along two axes: the masking signal (rcorr=πoldtrain/πoldrolloutr_{\text{corr}} = \pi^{\text{train}}_{\text{old}}/\pi^{\text{rollout}}_{\text{old}} versus the rollout-side PPO ratio) and granularity (token-level truncated importance sampling plus sequence-level rejection). Three findings emerge:

  1. Sequence-level rejection keyed on rcorrr_{\text{corr}} outperforms rejection keyed on the PPO ratio, because the correction ratio constrains where the update's base distribution sits relative to the sampling space, whereas the PPO ratio only bounds how far the policy moves.
  2. Adding TIS on top of rcorrr_{\text{corr}}-based rejection most closely tracks VeXact, consistent with the earlier analysis: TIS effectively repairs the loss-function ratio from ∣δt∣|\delta_t|0 toward ∣δt∣|\delta_t|1. The choice between ∣δt∣|\delta_t|2 and ∣δt∣|\delta_t|3 for sequence scoring has minor effect (thresholds used: ∣δt∣|\delta_t|4, ∣δt∣|\delta_t|5).
  3. These corrections remain post-hoc: they suppress already-generated samples, may discard useful learning signal, and their thresholds cannot be principledly calibrated without a zero-mismatch reference like VeXact.

Limitations

The authors concede that evaluation covers a finite set of models, tasks, and system configurations; generalization of the mitigations to broader RL settings, and whether they introduce side effects invisible in these experiments, remain open. The claim that symmetric numerical noise becomes skewed through interaction with clipping is stated as a hypothesis rather than proven analytically. Whether zero-mismatch execution remains necessary or merely beneficial in heavily asynchronous agentic RL regimes, where off-policyness dominates by design, is not resolved here.

Conclusion

The paper establishes TIM as an independently sufficient cause of LLM RL collapse, provides a mechanistic account of why recomputation and bypass fail differently, and demonstrates that existing corrections can approximate but not replace zero-mismatch execution. Its practical legacy is twofold: VeXact as a calibration instrument for tuning correction thresholds against a noise-free reference, and a case for treating numerical determinism across the training-inference boundary as a systems requirement rather than an implementation detail.

Paper to Video (Beta)

No one has generated a video about this paper yet.

Whiteboard

No one has generated a whiteboard explanation for this paper yet.

Open Problems

We haven't generated a list of open problems mentioned in this paper yet.