Trainer–inference mismatch

Trainer–inference mismatch is the difference between the token probabilities an inference engine used when sampling a rollout and the probabilities the trainer computes for the same tokens, in the same context, with the same weights. Because the policy-gradient update divides one by the other, the mismatch enters every update of an RL run. In exact arithmetic the two would be identical. In practice the sampler and trainer are different programs with different kernels, precisions and batch shapes, so their numbers disagree slightly on most tokens and substantially on a few.

Why it matters

A policy-gradient step uses the importance ratio

ρt=πθ(yt∣x,y<t)μ(yt∣x,y<t),\rho_t = \frac{\pi_\theta(y_t \mid x, y_{<t})}{\mu(y_t \mid x, y_{<t})},

where πθ\pi_\theta is the trainer's probability and μ\mu is the behavior probability recorded by the sampler (importance sampling). Nominally on-policy training assumes ρt=1\rho_t = 1 at the start of each step. Mismatch moves ρt\rho_t away from 1 even with zero policy lag, so training that looks on-policy is slightly off-policy in a way no staleness setting controls.

Most tokens have ratios very close to 1. The damage comes from the tail. A token that the sampler drew at probability 0.02 but the trainer scores at 0.10 has ρt=5\rho_t = 5, and its gradient is amplified accordingly. Such tokens concentrate where the distribution is flat, in rare tokens, and in long sequences where small errors in the context compound. Unchecked, they can drive the gradient norm up and training into collapse.

Causes

Precision. Inference commonly runs in bf16, or in FP8 with quantized weights, while the trainer keeps fp32 master weights and may use different mixed-precision rules. bf16 keeps only 7 explicit mantissa bits, so each matmul rounds differently depending on how it is computed. The final projection to vocabulary logits is especially sensitive: MiniMax-M1 traced its training–inference probability gap to high-magnitude activations in the LM head and raised the correlation between training-mode and inference-mode probabilities from about 0.9 to about 0.99 by computing that head in fp32.

Kernels. The trainer and the server implement attention, matmuls, normalization and activations with different kernels, fusions and reduction orders. Floating-point addition is not associative, so summing the same numbers in a different order changes the result.

Batch invariance. An inference kernel's output for one sequence can depend on what else is in the batch, because batch size changes how work is split and reduced. Server load varies, so the same prompt can get different logits on different calls. Horace He's Defeating Nondeterminism in LLM Inference identifies batch dependence as a main source of inference nondeterminism. Batch-invariant kernels make repeated inference deterministic. To get bitwise agreement between sampling and training, the authors also changed the training stack to match, at a throughput cost.

Mixture-of-experts routing. A top-kk router picks experts by comparing scores. A tiny numeric difference can flip which expert a token goes to, and a different expert produces a different output, not a slightly perturbed one. Ma et al. show that routing disagrees between training and inference and even between repeated inference passes, which makes MoE models especially prone to mismatch.

Sampling transforms. Temperature, top-p and top-k change the distribution a token was drawn from. If the sampler reports probabilities renormalized over the surviving top-p set and the trainer computes a full-vocabulary softmax, the ratio sits below 1 on every sampled token for reasons unrelated to numerics. The trainer has to apply the same temperature and the same truncation to compare like with like (policies).

Context. If the trainer's token sequence differs from the one the sampler saw, the two compute probabilities for different contexts. That is a rendering bug (token-level rendering). A cache reused across a weight update adds a related effect (KV cache).

Measuring it

Record the sampler's log-probability for every sampled token and compare it with the trainer's on the same tokens. Useful views:

  • Per-token log-ratio log⁡ρt\log \rho_t and its distribution, especially the tails.
  • A KL estimate between sampler and trainer. The estimator k3=ρ−log⁡ρ−1k_3 = \rho - \log \rho - 1 is non-negative per token. Over tokens sampled from μ\mu it averages to DKL(μ ∥ πθ)D_{\mathrm{KL}}(\mu\,\|\,\pi_\theta) when the two distributions have the same support, which a full-vocabulary trainer against a truncated sampler violates.
  • A probability scatter plot of trainer probability against sampler probability, which should lie on the diagonal.
  • Slices by position, token entropy, environment, and, for MoE models, how often expert selection differs.

Worked example. For the tail token above (μ=0.02\mu = 0.02, πθ=0.10\pi_\theta = 0.10), ρ=5\rho = 5 and k3=5−ln⁡5−1≈2.39k_3 = 5 - \ln 5 - 1 \approx 2.39. For a typical token with μ=0.50\mu = 0.50 and πθ=0.49\pi_\theta = 0.49, ρ=0.98\rho = 0.98 and k3=0.98−ln⁡0.98−1≈0.0002k_3 = 0.98 - \ln 0.98 - 1 \approx 0.0002. The mean over a batch can look small while a handful of tokens dominate the gradient, so tail statistics say more than the mean.

Mismatch and policy lag both move the ratio away from 1. To isolate numerics, measure on samples generated entirely by the weights the trainer holds at that step, or score a fixed set of rollouts with both engines at one checkpoint.

Corrections

The fixes fall into three kinds, which can be combined.

Align the numerics. Compute the LM head and MoE router logits in fp32. Match precision on both sides: Qi et al. find that switching trainer and sampler from bf16 to FP16, which has 10 mantissa bits, shrinks the mismatch and gives more stable optimization and faster convergence, with only a few lines of code changed. Use deterministic kernels where their cost is acceptable. Batch invariance removes dependence on batch composition, and Zhang et al. make matmul and reduction kernels invariant to tensor-parallel size, so a trainer on one GPU and a sampler sharded across several produce bit-identical results. Quantized inference widens the gap, so it needs the statistical corrections below.

Replay decisions from inference. Instead of recomputing discrete choices in the trainer, record them at sampling time and reuse them. Router replay (R3) sends the inference engine's expert selections to the trainer, and DeepSeek-V3.2 does the same under the name Keep Routing. Sampling-mask replay sends the set of token IDs that survived top-p or top-k, so the trainer renormalizes over the same support, which DeepSeek-V3.2 calls Keep Sampling Mask.

Correct statistically. Always use the sampler's recorded probability as μ\mu, never a trainer recomputation, so the ratio accounts for whatever mismatch remains. Then bound the tail, by capping ρt\rho_t (truncated importance sampling, which Yao et al. applied to this mismatch), masking tokens whose ratio leaves a band, or masking whole rollouts whose average divergence from μ\mu is too large. Each trades bias for variance by an amount that depends on the data and the thresholds (importance sampling).